Decision-1.0-Sol-2B / code /decision_model.py
Xunzhuo's picture
Release measured Decision 1.0 decoder
b6a4ce8 verified
Raw History Blame
10.1 kB
"""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<option>\n" + canonical({"key": o["key"], "description": o.get("description")}) + "\n</option>" 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)