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