"""Dynamic candidate readout over a causal Qwen3.5 text backbone. Candidate endpoints retain their contextual vectors. A final global-query vector can incorporate all options before a shared bilinear + MLP scorer scores every candidate. This is a research architecture, not a Jev claim. """ import hashlib import json import math from pathlib import Path import torch from torch import nn import torch.nn.functional as F from safetensors.torch import load_file, save_file from transformers import AutoTokenizer, Qwen3_5ForConditionalGeneration from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5TextModel PROMPT_VERSION = "structured-segmented-candidate-endpoints-global-query-v2" MAX_OPTIONS = 255 def canonical(value): return json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":")) def payload(value): return value if isinstance(value, str) else canonical(value) def segments(row): opts = row["options"] if not 2 <= len(opts) <= MAX_OPTIONS: raise ValueError(f"{row['id']}: expected 2..255 options") if not all(isinstance(o["key"], str) for o in opts): raise ValueError("Option keys must be strings") if len({o["key"] for o in opts}) != len(opts): raise ValueError("Duplicate option keys") prefix = f"Context:\n{payload(row['state'])}\n\nTask type: {row.get('task_type', 'choice')}\nQuestion:\n{payload(row['instructions'])}\nOptions:" # Tokenize each part separately. This deliberately fixes boundaries and # avoids guessing endpoint indices from merged BPE character offsets. options = ["\n" for o in opts] suffix = "\n\nSelect the single option best supported by the context and instructions.\nDecision:" return prefix, options, suffix def render(row): prefix, opts, suffix = segments(row) return prefix + "".join(opts) + suffix def encode(row, tokenizer, max_length=16384): prefix, opts, suffix = segments(row) ids = tokenizer.encode(prefix, add_special_tokens=False) candidate_positions = [] for option in opts: part = tokenizer.encode(option, add_special_tokens=False) if not part: raise ValueError("Empty tokenized candidate") ids.extend(part) candidate_positions.append(len(ids) - 1) ids.extend(tokenizer.encode(suffix, add_special_tokens=False)) if len(ids) > max_length: raise ValueError(f"{row['id']}: {len(ids)} tokens exceeds max_length={max_length}; no truncation allowed") label = row.get("label", -1) if label != -1 and not 0 <= label < len(opts): raise ValueError("Invalid label") prompt = prefix + "".join(opts) + suffix return {"id": row["id"], "ids": ids, "label": label, "nopts": len(opts), "family": row.get("family", "unspecified"), "candidate_positions": candidate_positions, "query_position": len(ids) - 1, "target_probs": row.get("target_probs"), "prompt_sha256": hashlib.sha256(prompt.encode()).hexdigest(), "token_ids_sha256": hashlib.sha256(canonical(ids).encode()).hexdigest(), "segmented_tokenization": True} def collate(items, pad_id): length = ((max(len(x["ids"]) for x in items) + 31) // 32) * 32 nopts = max(x["nopts"] for x in items) ids = torch.full((len(items), length), pad_id, dtype=torch.long) mask = torch.zeros_like(ids) positions = torch.zeros((len(items), nopts), dtype=torch.long) candidate_mask = torch.zeros((len(items), nopts), dtype=torch.bool) for i, item in enumerate(items): if len(item["candidate_positions"]) != item["nopts"]: raise ValueError("Candidate count does not match endpoint count") if not all(0 <= p < item["query_position"] < len(item["ids"]) for p in item["candidate_positions"]): raise ValueError("Candidate endpoints must precede global query") if len(set(item["candidate_positions"])) != item["nopts"]: raise ValueError("Duplicate candidate endpoint") ids[i, :len(item["ids"])] = torch.tensor(item["ids"]) mask[i, :len(item["ids"])] = 1 positions[i, :item["nopts"]] = torch.tensor(item["candidate_positions"]) candidate_mask[i, :item["nopts"]] = True return {"input_ids": ids, "attention_mask": mask, "candidate_positions": positions, "candidate_mask": candidate_mask, "query_positions": torch.tensor([x["query_position"] for x in items]), "labels": torch.tensor([x["label"] for x in items]), "nopts": torch.tensor([x["nopts"] for x in items]), "ids": [x["id"] for x in items], "families": [x["family"] for x in items]} class CandidateHead(nn.Module): def __init__(self, hidden_size, head_dim=256): super().__init__() self.head_dim = head_dim self.candidate_norm = nn.LayerNorm(hidden_size) self.query_norm = nn.LayerNorm(hidden_size) self.key = nn.Linear(hidden_size, head_dim, bias=False) self.query = nn.Linear(hidden_size, head_dim, bias=False) self.candidate_mlp = nn.Linear(hidden_size, head_dim, bias=True) self.query_mlp = nn.Linear(hidden_size, head_dim, bias=False) self.scalar = nn.Linear(head_dim, 1, bias=False) nn.init.normal_(self.scalar.weight, mean=0., std=0.01) def forward(self, candidates, query): # Keep the small shared head in FP32 even when the backbone uses BF16. # The v1 letter head exhibited BF16 ties sensitive to batch padding. with torch.autocast(device_type=candidates.device.type, enabled=False): c = self.candidate_norm(candidates.float()) q = self.query_norm(query.float()) bilinear = (self.key(c) * self.query(q)[:, None, :]).sum(-1) / math.sqrt(self.head_dim) interaction = self.scalar(F.gelu(self.candidate_mlp(c) + self.query_mlp(q)[:, None, :])).squeeze(-1) return bilinear + interaction class DecisionModel(nn.Module): def __init__(self, backbone, head, metadata): super().__init__() self.backbone, self.head, self.metadata = backbone, head, metadata @classmethod def from_base(cls, path, revision="local", dtype=torch.bfloat16, attention="sdpa", head_dim=256): tokenizer = AutoTokenizer.from_pretrained(path, local_files_only=True) full, info = Qwen3_5ForConditionalGeneration.from_pretrained(path, dtype=dtype, local_files_only=True, attn_implementation=attention, output_loading_info=True) if any(info.get(k) for k in ("missing_keys", "mismatched_keys", "error_msgs")): raise RuntimeError(f"Incomplete base loading: {info}") backbone = full.model.language_model backbone.config.use_cache = False head = CandidateHead(backbone.config.hidden_size, head_dim) metadata = {"base_revision": revision, "text_parameter_count": sum(p.numel() for p in backbone.parameters()), "prompt_version": PROMPT_VERSION, "attention": attention, "head_dim": head_dim, "max_options": MAX_OPTIONS, "architecture": "contextual-candidate-endpoint-plus-global-query-shared-bilinear-mlp", "head_initialization": "random-shared-content-scorer", "head_precision": "float32-outside-autocast"} return cls(backbone, head, metadata), tokenizer @classmethod def from_decision_checkpoint(cls, path, dtype=torch.bfloat16, attention="sdpa", head_dim=256): """Warm-start the backbone of a trained v1 model; initialize a new head.""" path = Path(path) metadata = json.loads((path / "decision_config.json").read_text()) backbone = Qwen3_5TextModel.from_pretrained(path / "backbone", dtype=dtype, local_files_only=True, attn_implementation=attention) backbone.config.use_cache = False head = CandidateHead(backbone.config.hidden_size, head_dim) metadata.update({"prompt_version": PROMPT_VERSION, "head_dim": head_dim, "max_options": MAX_OPTIONS, "architecture": "contextual-candidate-endpoint-plus-global-query-shared-bilinear-mlp", "head_initialization": "random-shared-content-scorer", "warm_start": "trained-v1-text-backbone", "head_precision": "float32-outside-autocast"}) return cls(backbone, head, metadata), AutoTokenizer.from_pretrained(path, local_files_only=True) @classmethod def from_checkpoint(cls, path, dtype=torch.bfloat16, attention="sdpa"): path = Path(path) metadata = json.loads((path / "decision_config.json").read_text()) if metadata["prompt_version"] != PROMPT_VERSION: raise ValueError("Not a pointer-v2 checkpoint; use from_decision_checkpoint for warm start") backbone = Qwen3_5TextModel.from_pretrained(path / "backbone", dtype=dtype, local_files_only=True, attn_implementation=attention) head = CandidateHead(backbone.config.hidden_size, metadata["head_dim"]) head.load_state_dict(load_file(path / "decision_head.safetensors")) return cls(backbone, head, metadata), AutoTokenizer.from_pretrained(path, local_files_only=True) def forward(self, input_ids, attention_mask, candidate_positions, candidate_mask, query_positions, **unused): hidden = self.backbone(input_ids=input_ids, attention_mask=attention_mask, use_cache=False).last_hidden_state batches = torch.arange(hidden.shape[0], device=hidden.device) candidates = hidden[batches[:, None], candidate_positions] query = hidden[batches, query_positions] scores = self.head(candidates, query).float() return scores.masked_fill(~candidate_mask, -float("inf")) def save(self, path, tokenizer): path = Path(path); path.mkdir(parents=True, exist_ok=True) self.backbone.save_pretrained(path / "backbone", safe_serialization=True, max_shard_size="4GB") save_file({n: v.detach().cpu().contiguous() for n, v in self.head.state_dict().items()}, str(path / "decision_head.safetensors")) tokenizer.save_pretrained(path) (path / "decision_config.json").write_text(json.dumps(self.metadata, indent=2) + "\n") def classification_loss(logits, labels): return F.cross_entropy(logits, labels)