| """Causal-LM sequence scoring used by the zero-shot evaluation tasks. |
| |
| The core primitive returns, for each input sequence, the model's summed (or |
| length-normalised) log-probability over a scored token region. Grammaticality |
| (minimal pairs) scores the whole sentence; multiple-choice scores only the |
| continuation tokens after a shared context. |
| |
| Batched scoring is exact: sequences are right-padded, and only real, |
| in-region target positions contribute. Causal attention means real tokens never |
| attend to right-padding, so per-token scores are identical to unpadded scoring |
| (verified in tests). |
| """ |
|
|
| from __future__ import annotations |
|
|
| import torch |
|
|
|
|
| def common_prefix_len(a: list[int], b: list[int]) -> int: |
| n = 0 |
| for x, y in zip(a, b): |
| if x != y: |
| break |
| n += 1 |
| return n |
|
|
|
|
| def encode_with_continuation(tokenizer, context_text: str, continuation_text: str) -> tuple[list[int], int]: |
| """Tokenise ``context_text + continuation_text`` and return ``(ids, context_len)``. |
| |
| ``context_len`` is the number of leading tokens that belong to the context |
| (and are therefore *not* scored). Computed as the shared token prefix between |
| the context alone and the full string, which is robust to subword-merge |
| effects at the boundary. |
| """ |
|
|
| context_ids = list(tokenizer.encode(context_text, add_bos=True).input_ids) |
| full_ids = list(tokenizer.encode(context_text + continuation_text, add_bos=True).input_ids) |
| context_len = common_prefix_len(context_ids, full_ids) |
| |
| context_len = max(1, min(context_len, len(full_ids) - 1)) |
| return full_ids, context_len |
|
|
|
|
| def _autocast(device_type: str, precision: str): |
| if precision == "fp32" or device_type == "cpu": |
| return torch.autocast(device_type=device_type, enabled=False) |
| dtype = torch.bfloat16 if precision == "bf16" else torch.float16 |
| return torch.autocast(device_type=device_type, dtype=dtype) |
|
|
|
|
| @torch.no_grad() |
| def score_sequences( |
| model, |
| sequences: list[list[int]], |
| context_lens: list[int], |
| *, |
| device: torch.device, |
| pad_id: int = 0, |
| batch_size: int = 16, |
| precision: str = "bf16", |
| length_normalize: bool = False, |
| predicate_memory_intervention: str = "none", |
| predicate_memory_residual_scale: float | None = None, |
| graph_object_residual_scale: float | None = None, |
| ) -> list[float]: |
| """Return a score per sequence: summed log-prob over positions ``[context_len, L)``. |
| |
| With ``length_normalize=True`` the sum is divided by the number of scored |
| tokens (recommended for comparing continuations of different lengths). |
| """ |
|
|
| if len(sequences) != len(context_lens): |
| raise ValueError("sequences and context_lens must have equal length") |
| model.eval() |
| device_type = device.type |
| scores: list[float] = [] |
|
|
| for start in range(0, len(sequences), batch_size): |
| chunk = sequences[start : start + batch_size] |
| chunk_ctx = context_lens[start : start + batch_size] |
| max_len = max(len(seq) for seq in chunk) |
| batch = torch.full((len(chunk), max_len), pad_id, dtype=torch.long) |
| real_len = torch.empty(len(chunk), dtype=torch.long) |
| ctx = torch.tensor(chunk_ctx, dtype=torch.long) |
| for i, seq in enumerate(chunk): |
| batch[i, : len(seq)] = torch.tensor(seq, dtype=torch.long) |
| real_len[i] = len(seq) |
| batch = batch.to(device) |
| attention_mask = (batch != pad_id).long() |
|
|
| with _autocast(device_type, precision): |
| logits = model( |
| batch, |
| attention_mask=attention_mask, |
| predicate_memory_intervention=predicate_memory_intervention, |
| predicate_memory_residual_scale=predicate_memory_residual_scale, |
| graph_object_residual_scale=graph_object_residual_scale, |
| ).logits |
| logprobs = torch.log_softmax(logits[:, :-1, :].float(), dim=-1) |
| targets = batch[:, 1:] |
| token_lp = logprobs.gather(-1, targets.unsqueeze(-1)).squeeze(-1) |
|
|
| |
| |
| positions = torch.arange(1, max_len, device=device).unsqueeze(0) |
| ctx = ctx.to(device).unsqueeze(1) |
| rl = real_len.to(device).unsqueeze(1) |
| mask = (positions >= ctx) & (positions < rl) |
| summed = (token_lp * mask).sum(dim=1) |
| counts = mask.sum(dim=1).clamp(min=1) |
| result = summed / counts if length_normalize else summed |
| scores.extend(result.tolist()) |
|
|
| return scores |
|
|