"""Self-contained inference code for JevCodeLocator-0.6B. The model is a Qwen3-0.6B backbone followed by a small set-attention head that scores a *set* of candidate code locations (``path:start-end``) against a query: backbone(last hidden state at EOS of each candidate path) -> LayerNorm -> scalar head (per-candidate logit) -> set-attention over the K candidates (context-dependent correction) One forward pass scores every candidate at once; no embeddings or vector index are involved. Candidate text is the compact summary produced by the data pipeline (``path | signature + first lines``). Example ------- >>> from modeling_jev import load_jev_model, score_candidates >>> model, tok = load_jev_model("JevCodeLocator-0.6B", device="cuda") >>> cands = { ... "src/auth.py:10-42": "def login(user, pwd) | validates credentials and returns a session token", ... "src/db.py:5-30": "def connect(dsn) | opens a database connection pool", ... } >>> for key, prob in score_candidates(model, tok, "where are credentials validated?", cands): ... print(f"{prob:.4f} {key}") """ from __future__ import annotations import json from pathlib import Path from typing import Dict, List, Sequence, Tuple import torch import torch.nn.functional as F from torch import nn INSTRUCTIONS = ("Select the single code location (file path and line range) that best answers the query. " "Each option is a candidate location; its text includes a compact signature/code summary. " "Only use the state and the option texts.") class JevDecisionModel(nn.Module): """Qwen3 backbone + LayerNorm/scalar head + set-attention over candidates.""" def __init__(self, backbone: nn.Module, set_head: str = "attention") -> None: super().__init__() self.backbone = backbone hidden = backbone.config.hidden_size self.norm = nn.LayerNorm(hidden) self.scalar = nn.Linear(hidden, 1) nn.init.normal_(self.scalar.weight, std=0.02) nn.init.zeros_(self.scalar.bias) self.set_head = set_head if set_head == "attention": self.set_project = nn.Linear(hidden + 1, 128) self.set_attention = nn.MultiheadAttention(128, 4, dropout=0.0, batch_first=True) self.set_output = nn.Linear(128, 1) nn.init.zeros_(self.set_output.weight) nn.init.zeros_(self.set_output.bias) def forward(self, examples: Sequence[dict], pad_token: int): """examples: [{leaf_tokens: [[int, ...], ...], candidate_ids: [...], type: "choice"}].""" paths = [ids for ex in examples for ids in ex["leaf_tokens"]] device = self.scalar.weight.device lengths = torch.tensor([len(ids) for ids in paths], device=device) width = int(lengths.max()) tokens = torch.full((len(paths), width), pad_token, dtype=torch.long, device=device) for i, ids in enumerate(paths): tokens[i, :len(ids)] = torch.tensor(ids, device=device) attention = torch.arange(width, device=device)[None, :] < lengths[:, None] hidden = self.backbone(input_ids=tokens, attention_mask=attention, use_cache=False).last_hidden_state leaves = hidden[torch.arange(len(paths), device=device), lengths - 1] kmax = max(len(ex["candidate_ids"]) for ex in examples) h = leaves.new_zeros((len(examples), kmax, leaves.shape[-1])) valid = torch.zeros((len(examples), kmax), dtype=torch.bool, device=device) offset = 0 for i, ex in enumerate(examples): n = len(ex["leaf_tokens"]) h[i, :n] = leaves[offset:offset + n] valid[i, :len(ex["candidate_ids"])] = True offset += n h = self.norm(h) z = self.scalar(h).squeeze(-1).float() choice = torch.tensor([i for i, ex in enumerate(examples) if ex["type"] == "choice"], device=device) if self.set_head == "attention" and len(choice): log_k = valid[choice].sum(-1).float().log()[:, None, None].expand(-1, kmax, 1) u = self.set_project(torch.cat([h[choice], log_k.to(h.dtype)], dim=-1)) mixed, _ = self.set_attention(u, u, u, key_padding_mask=~valid[choice], need_weights=False) delta = self.set_output(torch.tanh(u + mixed)).squeeze(-1).float() z = z.index_add(0, choice, delta) out = [] for i, ex in enumerate(examples): if ex["type"] == "boolean": out.append(F.pad(torch.stack([z[i, 0] * 0, z[i, 0]]), (0, kmax - 2))) else: out.append(z[i]) return torch.stack(out).masked_fill(~valid, -1e9), valid def load_jev_model(checkpoint_dir: str | Path, device: str = "cuda", dtype=torch.float32): """Load the backbone, head weights and tokenizer from a checkpoint directory.""" from safetensors.torch import load_file from transformers import AutoConfig, AutoModel, AutoTokenizer root = Path(checkpoint_dir).expanduser().resolve() cfg = json.loads((root / "config.json").read_text(encoding="utf-8")) body_config = AutoConfig.from_pretrained(str(root / "backbone_config"), local_files_only=True) backbone = AutoModel.from_config(body_config, attn_implementation="sdpa") model = JevDecisionModel(backbone, set_head=cfg.get("set_head", "attention")) weights = root / "model.safetensors" if not weights.is_file(): weights = root / "best.safetensors" missing, unexpected = model.load_state_dict(load_file(str(weights)), strict=False) if any(k.startswith("backbone.") or k.startswith("norm.") or k.startswith("scalar.") for k in missing): raise RuntimeError(f"checkpoint is missing head/backbone weights: {list(missing)[:5]}") model = model.to(device=device, dtype=dtype).eval() tokenizer = AutoTokenizer.from_pretrained(str(root / "tokenizer"), local_files_only=True) return model, tokenizer def prepare_examples(payload: dict, tokenizer, max_length: int = 2048) -> List[dict]: """Encode state + candidate texts into one leaf-token path per candidate.""" if tokenizer.eos_token_id is None: raise ValueError("tokenizer has no eos_token_id") examples = [] for row in payload["states"]: for qid, q in row["questions"].items(): ids = list(q["criteria"]) texts = [f"{key}: {q['criteria'][key]}" for key in ids] segments = [f"State:\n{row['state']}\n", f"Question type: choice\nQuestion:\n{q.get('instructions', INSTRUCTIONS)}\n"] prefix = sum([tokenizer.encode(t, add_special_tokens=False) for t in segments], []) leaves = [prefix + tokenizer.encode(f"Candidate:\n{t}\nDecision:", add_special_tokens=False) + [tokenizer.eos_token_id] for t in texts] largest = max(map(len, leaves)) if largest > max_length: raise ValueError(f"{row['id']}:{qid} candidate path is {largest} tokens > max_length={max_length}") examples.append({"id": f"{row['id']}:{qid}", "state_id": row["id"], "qid": qid, "type": "choice", "candidate_ids": ids, "candidate_texts": texts, "leaf_tokens": leaves}) return examples @torch.inference_mode() def score_candidates(model, tokenizer, query: str, candidates: Dict[str, str], repo: str = "repo", max_length: int = 2048) -> List[Tuple[str, float]]: """Return [(candidate_key, probability)] sorted by probability, best first. ``candidates`` maps ``path:start-end`` -> compact summary text; order is preserved in the input but the returned list is ranked by the model. """ state = ("task: fast_context_code_location\n" f"repo: {repo}\n" f"candidate_count: {len(candidates)}\n" "query:\n" f"{query}\n") row = {"id": "query", "state": state, "questions": {"q_loc": {"type": "choice", "instructions": INSTRUCTIONS, "criteria": dict(candidates)}}} ex = prepare_examples({"states": [row]}, tokenizer, max_length)[0] logits, valid = model([ex], tokenizer.pad_token_id) probs = logits[0, :len(ex["candidate_ids"])].float().softmax(-1).tolist() return sorted(zip(ex["candidate_ids"], probs), key=lambda kv: -kv[1])