nur-dev's picture
Add files using upload-large-folder tool
e69b72a verified
Raw
History Blame Contribute Delete
4.67 kB
"""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)
# Guarantee at least one scored token and a non-empty context.
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) # [B, S-1]
# Target at shifted index j predicts absolute position (j+1). Score it
# when context_len <= j+1 < real_len.
positions = torch.arange(1, max_len, device=device).unsqueeze(0) # [1, S-1]
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