opensysone / source /decision_model.py
andyshu's picture
Back up verified OpenSysOne training snapshot and pinned source
2d5c26a verified
Raw History Blame
9.26 kB
"""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))