File size: 8,379 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 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 | """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])
|