guided: rewind KV cache to natural EOS per-sample before forcing answer

Drop force_min_new_tokens — banning EOS to force a 2048-token think
generates ~250 tokens of real reasoning + </think>, then ~1800 tokens of
post-EOS sycophancy spew. Measuring pmass at the forced-answer slot with
that spew in the KV cache corrupted the coherence signal.

Replace with per-sample Phase 1.5: find each sample's first </think> in
phase1_ids, slice the batched DynamicCache (B, n_heads_kv, seq, d_head)
down to one sample × end_pos seq via _slice_pkv_one. The Phase 2 suffix
forward then runs per-sample over the rewound cache so the answer slot
sees only the coherent thinking trace.

GQA-safe (slices batch + seq, not heads). Phase 1 stays batched, Phase
1.5/2 loop adds ~5-10% wall-clock for the bs=1 forward.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
wassname
2026-05-20 00:27:02 +00:00
co-authored by Claude Opus 4.7
parent f58586408c
commit 5eabe37f8e
2 changed files with 139 additions and 95 deletions
+22 -8
View File
@@ -33,13 +33,13 @@ import torch
from loguru import logger
from tqdm.auto import tqdm
from .data import load_vignettes, ConfigName
from .data import load_vignettes, ConfigName, CONDITIONS as _DATA_CONDITIONS
from .guided import (
guided_rollout_forced_choice,
_DEFAULT_FORCED_FOUNDATIONS,
)
CONDITIONS = ("other_violate", "self_violate")
CONDITIONS = tuple(_DATA_CONDITIONS)
# Probe word -> dataset coarse label.
_PROBE_TO_COARSE: dict[str, str] = {
@@ -47,7 +47,6 @@ _PROBE_TO_COARSE: dict[str, str] = {
"authority": "Authority", "sanctity": "Sanctity", "liberty": "Liberty",
"social": "SocialNorms",
}
_COARSE_TO_PROBE: dict[str, str] = {v: k for k, v in _PROBE_TO_COARSE.items()}
# Some Clifford rows use "Social Norms" with a space; normalise.
_COARSE_NORM = {"Social Norms": "SocialNorms"}
@@ -127,11 +126,13 @@ def evaluate(
name: ConfigName = "classic",
vignettes: list[dict] | None = None,
*,
n_vignettes: int | None = None,
conditions: tuple[str, ...] = CONDITIONS,
max_think_tokens: int = 256,
batch_size: int = 8,
device: str | None = None,
return_per_row: bool = False,
verbose: bool = False,
) -> dict[str, Any]:
"""Run forced-choice 7-way probe per (vignette, condition).
@@ -139,18 +140,24 @@ def evaluate(
model, tokenizer: HuggingFace causal LM + matching tokenizer with chat template.
name: dataset config (`classic` / `scifi` / `ai-actor`).
vignettes: optional pre-loaded list (overrides `name`).
n_vignettes: optional slice — keep only the first N (after loading).
conditions: which condition strings to score. Default = both.
max_think_tokens: think budget per (row, frame). Two frames per row.
batch_size: rows per forced-choice call (KV cache = batch * 2 * max_think_tokens).
return_per_row: if True, include the per-row 7-vec p in the result.
return_per_row: if True, include the per-row 7-vec p + think text in the result.
verbose: if True, log the row-0 think trace at DEBUG level (one per slot).
Returns:
Dict with `table`, `profile`, `mean_js`, `mean_nll`, `mean_nll_T`,
`median_nll_T`, `T`, `top1_acc`, and `info`. If `return_per_row=True`,
also includes `per_row` with the row-level distributions and scores.
`median_nll_T`, `T`, `top1_acc`, `mean_pmass_format`, and `info`.
With `return_per_row=True`, also includes `per_row` with per-row
`p`, `score` (debiased logp per foundation), `pmass_format`,
`think_text` / `think_text_rev`, and `top1` / `margin`.
"""
if vignettes is None:
vignettes = load_vignettes(name)
if n_vignettes is not None:
vignettes = vignettes[:n_vignettes]
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
@@ -172,6 +179,7 @@ def evaluate(
model, tokenizer, user_prompts,
foundations=foundations,
max_think_tokens=max_think_tokens,
verbose=verbose,
)
for src, res in zip(chunk, results):
p_vec = np.array([res.p[f] for f in foundations], dtype=float)
@@ -190,15 +198,21 @@ def evaluate(
"pmass_format": res.pmass_format,
"think_tokens": res.think_tokens,
"emitted_close": res.emitted_close,
"think_text": res.think_text,
"gen_text_full": res.gen_text_full,
})
pbar.update(len(chunk))
elapsed = time.time() - t0
n_rows = len(per_row)
n_labeled = sum(1 for r in per_row if r["label"] is not None)
# Tokens-per-second: 2 frames per row (fwd + rev), each generates think_tokens.
# think_tokens on the result is the fwd count; rev cost is the same order.
total_gen_tokens = 2 * sum(r["think_tokens"] for r in per_row if r["think_tokens"] is not None)
tps = total_gen_tokens / elapsed if elapsed > 0 else 0.0
logger.info(
f"{name}: {n_rows} rows in {elapsed:.1f}s ({n_rows/elapsed:.1f} rows/s); "
f"{n_labeled}/{n_rows} have label dist"
f"{name}: {n_rows} rows in {elapsed:.1f}s ({n_rows/elapsed:.1f} rows/s, "
f"~{tps:.0f} tok/s); {n_labeled}/{n_rows} have label dist"
)
# Per-row think-token distribution — main eval-cost driver. Rows are
# 2 frames × n_vignettes; we average across frames before reporting.
+117 -87
View File
@@ -4,14 +4,20 @@ Public API: `guided_rollout_forced_choice` (K-way moral-foundation probe with
two-pass enum-reversal position-bias debias).
Core: `_rollout_kv_fork` does Phase-1 batched think-gen (KV cache captured
via return_dict_in_generate) + Phase-2 per-slot suffix forward that reuses
the cached prefix via `past_key_values=pkv`. Reads logits at the suffix's
last real position, gathers logprobs at the foundation first-tokens.
via return_dict_in_generate) + Phase-1.5 per-sample rewind to first </think>
+ Phase-2 per-sample suffix forward over the rewound pkv. Reads logits at the
suffix's last position, gathers logprobs at the foundation first-tokens.
Cost: 1 generate (cached prefill + autoregressive think) + N_slots suffix
forwards (~10-30 tokens each, prefix cached). Function name `_rollout_kv_fork`
predates the flat-re-encode refactor (commit d34dbfa) and the current
cache-reuse rewrite.
Why per-sample rewind: HF generate() with a batch stops each sample at its
own EOS but keeps the cache full-length (pad-filled after stop). If we just
appended a batched suffix at J_max, the suffix's position embeddings would
land far past the model's actual stopping point, polluting the pmass
measurement with post-EOS context. Per-sample slicing puts the suffix
immediately after each sample's real content.
Cost: 1 generate (batched) + B suffix forwards (one per sample, ~10-30
tokens each, prefix cached via past_key_values). Function name predates
the cache-reuse rewrite.
Why turn-boundary close+nudge: matches what a chat UI emits when a human
interrupts a partial assistant turn. On-policy in chat-tuned data, where the
@@ -44,12 +50,30 @@ def _assistant_close(tok) -> str:
return closed.split(_ASSISTANT_SENTINEL, 1)[1]
def _split_choice_ids(choice_token_ids: list) -> tuple[list[int], list[int]]:
if len(choice_token_ids) == 2 and all(isinstance(x, (list, tuple)) for x in choice_token_ids):
return list(choice_token_ids[0]), list(choice_token_ids[1])
return list(choice_token_ids), []
def _slice_pkv_one(pkv, i: int, end_pos: int):
"""Slice the batched KV cache to sample i, keeping only the first `end_pos`
seq positions. Returns a per-sample DynamicCache usable as
`past_key_values=` in a subsequent forward.
GQA-safe: slices only batch and seq dims; n_heads_kv (which may be <
n_heads_q) is preserved. NOTE: sliding-window-attention layers in models
like Gemma-2 cap the cached seq_len at window_size; for budgets >
window_size, end_pos may exceed cache length — we clamp to the actual
cached length per layer.
transformers 5.x DynamicCache exposes per-layer `.layers[l].keys` /
`.values` ([B, n_heads_kv, seq, d_head]). We slice each and rebuild a
fresh DynamicCache via .update().
"""
from transformers.cache_utils import DynamicCache
out = DynamicCache()
for layer_idx, layer in enumerate(pkv.layers):
k = layer.keys
v = layer.values
kk = k[i:i+1, :, :min(end_pos, k.shape[2]), :]
vv = v[i:i+1, :, :min(end_pos, v.shape[2]), :]
out.update(kk, vv, layer_idx)
return out
@torch.no_grad()
@@ -59,22 +83,24 @@ def _rollout_kv_fork(
schema_hint: str,
max_think_tokens: int,
scoring_slots: list[tuple[str, str]], # (nudge_user_text, prefill) per slot
choice_token_ids: list, # [a_ids, b_ids]
gather_token_ids: list[int], # K-way answer-token ids
verbose: bool = False,
gather_token_ids: list[int] | None = None,
) -> tuple[list[tuple[str, int, bool]], list[list[dict]]]:
) -> tuple[list[tuple[str, int, bool, str]], list[list[dict]]]:
"""Returns (thinks, slots).
thinks[i] = (think_text, n_think_tokens, emitted_close)
slots[i][j] = {pmass_format, logratio, p_true, top5_str, [lp_gather]}
thinks[i] = (think_text, n_think_tokens, emitted_close, gen_text_full)
slots[i][j] = {pmass_format, top5_str, lp_gather}
Two-phase rollout:
Phase 1 — generate up to max_think_tokens with cache=True, capture pkv.
Phase 2 — for each scoring slot, forward only the suffix
(close + interrupt + nudge + prefill) with past_key_values=pkv,
read logits at the suffix's last real token.
Three-phase rollout:
Phase 1 (batched) — generate up to max_think_tokens with cache=True,
capture pkv. Natural EOS stop (no min_new_tokens).
Phase 1.5 (per-sample) — find first </think> position per sample;
rewind pkv to that position so post-EOS spew
does not pollute the answer-slot measurement.
Phase 2 (per-sample) — forward the scoring suffix with rewound pkv,
read logits at the suffix's last position.
If `gather_token_ids` is provided, slot dict also has `lp_gather`:
log-probs at last suffix position for those token ids.
`pmass_format` is Σ exp(logp) over `gather_token_ids` at the slot.
`lp_gather` is the per-id logp vector at the slot.
"""
if tok.padding_side != "left":
raise ValueError("tok.padding_side must be 'left'")
@@ -97,7 +123,8 @@ def _rollout_kv_fork(
enc = tok(chats, return_tensors="pt", padding=True).to(device)
prompt_len = enc.input_ids.shape[1]
out1 = model.generate(
**enc, max_new_tokens=max_think_tokens, do_sample=False,
**enc,
max_new_tokens=max_think_tokens, do_sample=False,
eos_token_id=think_end_id, pad_token_id=pad_id,
return_dict_in_generate=True,
)
@@ -105,27 +132,33 @@ def _rollout_kv_fork(
pkv = out1.past_key_values # KV for [left-pad, prompt, think, (eos-pad)]
B = phase1_ids.shape[0]
thinks: list[tuple[str, int, bool]] = []
thinks: list[tuple[str, int, bool, str]] = []
real_lens: list[int] = [] # per-sample: seq_len up to and including first </think>
for i in range(B):
gen_ids = phase1_ids[i, prompt_len:]
keep = gen_ids != pad_id
gen_ids = gen_ids[keep] if keep.any() else gen_ids[:0]
gen_ids_full = phase1_ids[i, prompt_len:]
keep = gen_ids_full != pad_id
gen_ids = gen_ids_full[keep] if keep.any() else gen_ids_full[:0]
gen_text = tok.decode(gen_ids, skip_special_tokens=True)
n_think = int(gen_ids.shape[0])
emitted_close = _CLOSE_MARKER in gen_text
think_text = gen_text.split(_CLOSE_MARKER, 1)[0] if emitted_close else gen_text
thinks.append((think_text, n_think, emitted_close))
thinks.append((think_text, n_think, emitted_close, gen_text))
# Attention mask for the cached prefix. Real tokens = left-padded prompt
# tokens + generated tokens up to eos; pad_id positions on either end are
# masked out so suffix attention doesn't see them.
# Phase 1.5: rewind position = first think_end_id in gen (inclusive),
# so the answer slot's KV context ends at the natural stopping point —
# not at the post-EOS spew (which would corrupt pmass).
eos_mask = (gen_ids_full == think_end_id)
if eos_mask.any():
first_eos = int(eos_mask.nonzero(as_tuple=True)[0][0].item())
real_lens.append(prompt_len + first_eos + 1)
else:
real_lens.append(phase1_ids.shape[1]) # no EOS → keep full budget
# Attention mask for the full cached prefix (per-sample slices reuse this).
pref_attn = (phase1_ids != pad_id).long()
# === Phase 2: per-slot suffix forward, reusing Phase 1's KV cache ===
a_ids, b_ids = _split_choice_ids(choice_token_ids)
a_t = torch.tensor(a_ids, device=device, dtype=torch.long) if a_ids else None
b_t = torch.tensor(b_ids, device=device, dtype=torch.long) if b_ids else None
all_ids = torch.tensor(a_ids + b_ids, device=device, dtype=torch.long)
# === Phase 2: per-sample suffix forward over rewound pkv ===
gid_t = torch.tensor(gather_token_ids, device=device, dtype=torch.long)
def suf_ids_for(nudge: str, prefill: str) -> list[list[int]]:
"""Per-row suffix: optional </think> close + assistant-turn close +
@@ -136,42 +169,47 @@ def _rollout_kv_fork(
tokenize=False, continue_final_message=True,
)
suffixes = []
for _, _, emitted_close in thinks:
for _, _, emitted_close, _ in thinks:
head = "" if emitted_close else _CLOSE_MARKER
suf_text = head + close + interrupt
suffixes.append(tok(suf_text, add_special_tokens=False)["input_ids"])
return suffixes
def fork(suffixes: list[list[int]]) -> torch.Tensor:
"""Forward only suffix tokens with pkv from Phase 1.
Returns [B, V] logp at suffix's last real token."""
J_max = max(len(s) for s in suffixes)
suf_input = torch.full((B, J_max), pad_id, dtype=torch.long, device=device)
suf_mask = torch.zeros((B, J_max), dtype=torch.long, device=device)
last_pos = torch.zeros(B, dtype=torch.long, device=device)
for i, s in enumerate(suffixes):
L = len(s)
suf_input[i, :L] = torch.tensor(s, device=device)
suf_mask[i, :L] = 1
last_pos[i] = L - 1
# attention_mask must span both cached and new tokens.
full_attn = torch.cat([pref_attn, suf_mask], dim=1)
out = model(
input_ids=suf_input,
attention_mask=full_attn,
past_key_values=pkv,
use_cache=False, # don't grow / mutate the cache between slots
)
# out.logits is [B, J_max, V] — only suffix positions.
logp = F.log_softmax(out.logits.float(), dim=-1)
return logp[torch.arange(B, device=device), last_pos]
def fork_per_sample(suffixes: list[list[int]]) -> torch.Tensor:
"""Per-sample forward: rewind pkv to first-EOS for each sample,
forward only that sample's suffix, return [B, V] logp at the suffix's
last position.
Per-sample (bs=1) because each sample's rewind position differs;
batching would require padding pkv along seq_len with attention-mask
gymnastics on a heterogeneous-length cache. Heavy lifting (Phase 1)
was already batched, so this loop is a thin extra cost.
"""
V = model.config.vocab_size
lp_last = torch.zeros((B, V), device=device, dtype=torch.float32)
for i in range(B):
end_pos = real_lens[i]
pkv_i = _slice_pkv_one(pkv, i, end_pos)
pref_attn_i = pref_attn[i:i+1, :end_pos]
suf_i = torch.tensor([suffixes[i]], device=device, dtype=torch.long)
L = suf_i.shape[1]
suf_mask_i = torch.ones((1, L), dtype=torch.long, device=device)
full_attn_i = torch.cat([pref_attn_i, suf_mask_i], dim=1)
out = model(
input_ids=suf_i,
attention_mask=full_attn_i,
past_key_values=pkv_i,
use_cache=False,
)
lp_last[i] = F.log_softmax(out.logits[0, -1].float(), dim=-1)
return lp_last
slots: list[list[dict]] = [[] for _ in range(B)]
for j, (nudge, prefill) in enumerate(scoring_slots):
suf_ids = suf_ids_for(nudge, prefill)
if verbose:
# DEBUG: shows row 0 only. Keeps trace in the user's verbose
# sidecar but out of any downstream INFO sink.
# DEBUG: shows row 0 only. Independent generate from raw ids
# (does not use the cache) so it still works after the rewind.
real0 = phase1_ids[0][phase1_ids[0] != pad_id]
prefix_text = tok.decode(real0, skip_special_tokens=False)
suf_text_0 = tok.decode(suf_ids[0], skip_special_tokens=False)
@@ -184,32 +222,19 @@ def _rollout_kv_fork(
f"--- slot {j} (nudge={nudge!r}, prefill={prefill!r}) ---\n"
f"{prefix_text}{suf_text_0}<<<MODEL CONTINUES>>>{free}\n--- end slot {j} ---"
)
lp_last = fork(suf_ids)
pmass = lp_last[:, all_ids].exp().sum(-1)
if a_t is not None and b_t is not None:
la = torch.logsumexp(lp_last[:, a_t], dim=-1)
lb = torch.logsumexp(lp_last[:, b_t], dim=-1)
logratio = la - lb
p_true = torch.softmax(torch.stack([la, lb], dim=-1), dim=-1)[:, 0]
else:
logratio = torch.full((B,), float("nan"), device=device)
p_true = torch.full((B,), float("nan"), device=device)
lp_last = fork_per_sample(suf_ids)
pmass = lp_last[:, gid_t].exp().sum(-1)
for i in range(B):
top5 = lp_last[i].topk(5)
top5_str = " ".join(
f"{tok.decode([int(idx)])!r}:{float(prob.exp()):.3f}"
for idx, prob in zip(top5.indices, top5.values)
)
d = {
slots[i].append({
"pmass_format": float(pmass[i].item()),
"logratio": float(logratio[i].item()),
"p_true": float(p_true[i].item()),
"top5_str": top5_str,
}
if gather_token_ids is not None:
gid_t = torch.tensor(gather_token_ids, device=device, dtype=torch.long)
d["lp_gather"] = lp_last[i, gid_t].cpu().tolist()
slots[i].append(d)
"lp_gather": lp_last[i, gid_t].cpu().tolist(),
})
return thinks, slots
@@ -293,6 +318,12 @@ class ForcedChoiceResult:
# format collapse). The direct coherence canary for forced-choice
# — independent of WHICH foundation is picked.
pmass_format: float
# Forward-frame full decoded gen including anything past </think>. With the
# natural-EOS rewind in `_rollout_kv_fork`, this is normally identical to
# `think_text` (model stopped at </think>); preserved as a separate field
# for sidecar inspection when generation hits max_think_tokens without
# emitting close.
gen_text_full: str = ""
def _resolve_first_token_ids(tok, words: list[str]) -> tuple[list[int], dict[str, int]]:
@@ -373,25 +404,23 @@ def guided_rollout_forced_choice(
thinks_fwd, slots_fwd = _rollout_kv_fork(
model, tok, user_prompts, schema_fwd, max_think_tokens,
scoring_slots=scoring_slot,
choice_token_ids=[[first_ids[0]]], # unused; satisfies API
verbose=verbose,
gather_token_ids=first_ids,
verbose=verbose,
)
# Frame B: reversed enum order. Same gather order (by foundation name) so
# lp_rev[f] is comparable to lp_fwd[f].
thinks_rev, slots_rev = _rollout_kv_fork(
model, tok, user_prompts, schema_rev, max_think_tokens,
scoring_slots=scoring_slot,
choice_token_ids=[[first_ids[0]]],
verbose=verbose,
gather_token_ids=first_ids,
verbose=verbose,
)
results: list[ForcedChoiceResult] = []
import math
for i in range(len(user_prompts)):
think_fwd, n_fwd, close_fwd = thinks_fwd[i]
think_rev, _, _ = thinks_rev[i]
think_fwd, n_fwd, close_fwd, gen_text_full_fwd = thinks_fwd[i]
think_rev, _, _, _ = thinks_rev[i]
lp_f = slots_fwd[i][0]["lp_gather"]
lp_r = slots_rev[i][0]["lp_gather"]
score = [(lp_f[k] + lp_r[k]) / 2.0 for k in range(K)]
@@ -422,6 +451,7 @@ def guided_rollout_forced_choice(
margin=float(margin),
think_tokens=n_fwd,
emitted_close=close_fwd,
gen_text_full=gen_text_full_fwd,
pmass_format=float(pm),
))