"""Small decoder decision scorer; all full and cached paths use identical tokens.""" from __future__ import annotations import time import torch from torch import nn from torch.nn import functional as F from transformers import AutoModelForCausalLM, AutoTokenizer from transformers.cache_utils import DynamicCache class DecisionScorer(nn.Module): """Score arbitrary candidate text with a scalar readout of its final token.""" max_tokens = 2048 def __init__(self, model_path, train_layers=2, device="cuda", dtype=torch.float32): super().__init__() self.tokenizer = AutoTokenizer.from_pretrained(model_path, local_files_only=True) self.lm = AutoModelForCausalLM.from_pretrained( model_path, dtype=dtype, attn_implementation="sdpa", local_files_only=True, ).to(device) layers = self.lm.model.layers if not isinstance(train_layers, int) or not 0 <= train_layers <= len(layers): raise ValueError(f"train_layers must be an integer in [0, {len(layers)}]") self.lm.requires_grad_(False) if train_layers: for layer in layers[-train_layers:]: layer.requires_grad_(True) self.head = nn.Linear(self.lm.config.hidden_size, 1, device=device, dtype=torch.float32) nn.init.zeros_(self.head.weight) nn.init.zeros_(self.head.bias) self.pad_id = self.tokenizer.pad_token_id if self.pad_id is None: self.pad_id = self.tokenizer.eos_token_id if self.pad_id is None: raise ValueError("Tokenizer must define a padding or EOS token") self.last_shared_timings = {} @property def device(self): return self.head.weight.device @staticmethod def _text(value, name): if not isinstance(value, str) or not value.strip(): raise ValueError(f"{name} must be a nonempty string") return value def prefix_ids(self, state): state = self._text(state, "state") return self.tokenizer.encode( "Does the candidate correctly answer the question about the state? " "Answer yes or no.\n\nSTATE:\n" + state, add_special_tokens=False, ) def branch_ids(self, question, candidate): question = self._text(question, "question") candidate = self._text(candidate, "candidate") return self.tokenizer.encode( f"\n\nQUESTION:\n{question}\n\nCANDIDATE:\n{candidate}\n\nDECISION:", add_special_tokens=False, ) def _choices(self, example): self._text(example["question"], "question") choices = example["choices"] if not isinstance(choices, (list, tuple)) or len(choices) < 2: raise ValueError("choices must contain at least two candidates") for choice in choices: self._text(choice, "candidate") if len({choice.strip() for choice in choices}) != len(choices): raise ValueError("choices must be unique") return choices def _check_length(self, prefix, branch): if len(prefix) + len(branch) > self.max_tokens: raise ValueError( f"Input has {len(prefix) + len(branch)} tokens; limit is " f"{self.max_tokens}. Inputs are never silently truncated." ) def _pad(self, sequences): lengths = torch.tensor([len(ids) for ids in sequences], device=self.device) width = max(map(len, sequences)) ids = torch.full((len(sequences), width), self.pad_id, device=self.device, dtype=torch.long) for row, sequence in enumerate(sequences): ids[row, :len(sequence)] = torch.tensor(sequence, device=self.device) mask = torch.arange(width, device=self.device).unsqueeze(0) < lengths.unsqueeze(1) return ids, mask.long(), lengths @staticmethod def _last(hidden, lengths): rows = torch.arange(hidden.shape[0], device=hidden.device) return hidden[rows, lengths - 1] def _full_hidden(self, examples): if not examples: raise ValueError("examples must not be empty") sequences, counts = [], [] for example in examples: prefix = self.prefix_ids(example["state"]) choices = self._choices(example) counts.append(len(choices)) for candidate in choices: branch = self.branch_ids(example["question"], candidate) self._check_length(prefix, branch) sequences.append(prefix + branch) ids, mask, lengths = self._pad(sequences) output = self.lm.model(input_ids=ids, attention_mask=mask, use_cache=False) return self._last(output.last_hidden_state, lengths), counts def score_examples(self, examples): """Differentiable full-sequence scores, returned as one [choices] tensor each.""" hidden, counts = self._full_hidden(examples) return list(self.head(hidden.float()).squeeze(-1).split(counts)) @torch.inference_mode() def scores_token_baseline(self, examples): """Next-token logit(yes) - logit(no), without allocating vocabulary logits.""" hidden, counts = self._full_hidden(examples) token_ids = [self.tokenizer.encode(word, add_special_tokens=False) for word in (" yes", " no")] if any(len(ids) != 1 for ids in token_ids): raise ValueError("The yes/no baseline requires single-token ' yes' and ' no'") index = torch.tensor([ids[0] for ids in token_ids], device=self.device) weight = self.lm.lm_head.weight.index_select(0, index) bias = self.lm.lm_head.bias if bias is not None: bias = bias.index_select(0, index) logits = F.linear(hidden, weight, bias).float() return list((logits[:, 0] - logits[:, 1]).split(counts)) def _sync(self): if self.device.type == "cuda": torch.cuda.synchronize(self.device) @torch.inference_mode() def scores_shared(self, state, questions, branch_batch_size=16): """Prefill state once, then score independent candidate branches in chunks. Call eval() before comparing inference paths. Timing includes device transfers and cache replication, and excludes tokenization. cache_bytes measures the single stored prefix, not its temporary batched copies. """ if not questions: raise ValueError("questions must not be empty") if not isinstance(branch_batch_size, int) or branch_batch_size < 1: raise ValueError("branch_batch_size must be a positive integer") prefix = self.prefix_ids(state) branches, counts = [], [] for question in questions: choices = self._choices(question) counts.append(len(choices)) for candidate in choices: branch = self.branch_ids(question["question"], candidate) self._check_length(prefix, branch) branches.append(branch) self._sync() started = time.perf_counter() prefix_tensor = torch.tensor([prefix], device=self.device) prefill = self.lm.model( input_ids=prefix_tensor, attention_mask=torch.ones_like(prefix_tensor), use_cache=True, ) prefix_cache = prefill.past_key_values self._sync() prefilled = time.perf_counter() scores = [] for start in range(0, len(branches), branch_batch_size): chunk = branches[start:start + branch_batch_size] # DynamicCache mutates in place. Each chunk owns new K/V tensors; # repeat_interleave allocates even when this chunk has one branch. cache = DynamicCache( ddp_cache_data=[ (key.repeat_interleave(len(chunk), dim=0), value.repeat_interleave(len(chunk), dim=0)) for key, value, *_ in prefix_cache ], config=self.lm.config, ) ids, branch_mask, lengths = self._pad(chunk) prefix_mask = torch.ones((len(chunk), len(prefix)), dtype=torch.long, device=self.device) mask = torch.cat((prefix_mask, branch_mask), dim=1) positions = (torch.arange(ids.shape[1], device=self.device) + len(prefix)).unsqueeze(0) output = self.lm.model( input_ids=ids, attention_mask=mask, position_ids=positions, past_key_values=cache, use_cache=True, ) hidden = self._last(output.last_hidden_state, lengths) scores.append(self.head(hidden.float()).squeeze(-1)) self._sync() finished = time.perf_counter() cache_bytes = sum( key.numel() * key.element_size() + value.numel() * value.element_size() for key, value, *_ in prefix_cache ) self.last_shared_timings = { "prefill_ms": (prefilled - started) * 1000, "branch_ms": (finished - prefilled) * 1000, "prefix_tokens": len(prefix), "branches": len(branches), "branch_tokens": sum(map(len, branches)), "cache_bytes": cache_bytes, } return list(torch.cat(scores).split(counts))