File size: 9,258 Bytes
2d5c26a | 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 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 | """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))
|