File size: 2,173 Bytes
cfb6957
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
"""Minimal end-to-end example for JevCodeLocator-0.6B.

    python inference_example.py

Shows the intended serving pattern: a cheap retriever (BM25 / symbol search)
produces a shortlist of `path:start-end` candidates, and this model reranks them
with one forward pass.
"""
from __future__ import annotations

import sys
from pathlib import Path

import torch

from modeling_jev import load_jev_model, score_candidates

# Directory containing model.safetensors + config.json (default: this folder).
CHECKPOINT = sys.argv[1] if len(sys.argv) > 1 else str(Path(__file__).resolve().parent)

# A shortlist as a retriever would produce it: key = path:start-end, value = compact summary.
CANDIDATES = {
    "src/auth/session.py:41-88": "def create_session(user, password) | validates credentials against the user store and returns a signed session token",
    "src/db/pool.py:12-60": "def connect(dsn, max_size) | opens and configures the database connection pool",
    "src/api/routes/login.py:18-35": "def login_handler(request) | parses the JSON body and calls create_session, mapping failures to HTTP 401",
    "src/auth/tokens.py:7-29": "def sign(payload, key) | HMAC-signs a payload with the application secret",
    "docs/auth.md:1-40": "Authentication overview and deployment notes",
}

QUERY = "where are user credentials validated before a session token is created?"


def main() -> None:
    device = "cuda" if torch.cuda.is_available() else "cpu"
    model, tokenizer = load_jev_model(CHECKPOINT, device=device, dtype=torch.float32)
    ranked = score_candidates(model, tokenizer, QUERY, CANDIDATES, repo="example", max_length=2048)
    print(f"query: {QUERY}\n")
    for rank, (key, prob) in enumerate(ranked, 1):
        print(f"{rank}. {prob:.4f}  {key}")

    # File-level aggregation: sum the probabilities of all candidates in one file.
    files: dict[str, float] = {}
    for key, prob in ranked:
        files[key.rsplit(":", 1)[0]] = files.get(key.rsplit(":", 1)[0], 0.0) + prob
    print("\nfile ranking:")
    for path, score in sorted(files.items(), key=lambda kv: -kv[1]):
        print(f"  {score:.4f}  {path}")


if __name__ == "__main__":
    main()