JevCodeLocator-0.6B / inference_example.py
RowletQwQ's picture
JevCodeLocator-0.6B
cfb6957 verified
Raw History Blame Contribute Delete
2.17 kB
"""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()