Download source/decision_model.py from andyshu/opensysone: direct link, hf CLI and curl.
- Browser
- Download file 9.26 kB
-
https://huggingface.co/andyshu/opensysone/resolve/7364912c6bdfe66b6c3a24b308376a18cc482bee/source/decision_model.py
- Command line
-
hf download hf://andyshu/opensysone@7364912c6bdfe66b6c3a24b308376a18cc482bee/source/decision_model.py
-
curl -L -o decision_model.py https://huggingface.co/andyshu/opensysone/resolve/7364912c6bdfe66b6c3a24b308376a18cc482bee/source/decision_model.py
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 = {} | |
| def device(self): | |
| return self.head.weight.device | |
| 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 | |
| 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)) | |
| 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) | |
| 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)) | |