mirror of
https://github.com/wassname/moral-maps.git
synced 2026-09-09 11:27:22 +08:00
Refactor guided rollout scoring to use flat prefix+suffix approach and remove full attention assertion
This commit is contained in:
+22
-32
@@ -1,14 +1,17 @@
|
|||||||
"""Guided rollout: think + KV-forked suffix scoring for forced-choice probes.
|
"""Guided rollout: think + flat prefix+suffix scoring for forced-choice probes.
|
||||||
|
|
||||||
Public API: `guided_rollout_forced_choice` (K-way moral-foundation probe with
|
Public API: `guided_rollout_forced_choice` (K-way moral-foundation probe with
|
||||||
two-pass enum-reversal position-bias debias).
|
two-pass enum-reversal position-bias debias).
|
||||||
|
|
||||||
Core: `_rollout_kv_fork` does Phase-1 batched think-gen + Phase-2a one prefix
|
Core: `_rollout_kv_fork` does Phase-1 batched think-gen + Phase-2a prefix
|
||||||
forward + Phase-2b deepcopy-cache + per-slot suffix forward. Reads logits at
|
forward (for free per-row prompt-NLL) + Phase-2b per-slot flat prefix+suffix
|
||||||
the prefill's last position, gathers logprobs at the foundation first-tokens.
|
forward. Reads logits at the prefill's last position, gathers logprobs at the
|
||||||
|
foundation first-tokens.
|
||||||
|
|
||||||
Cost: 1 prefix forward + N_slots small suffix forwards (each <=14 tokens),
|
Cost: 1 prefix forward + N_slots full prefix+suffix forwards (suffix <=14
|
||||||
batched over B prompts. Requires a full-attention model -- asserts at entry.
|
tokens). We re-encode the prefix per slot to stay correct on hybrid attention
|
||||||
|
(sliding-window / SSM) where KV-fork loses cached context past the window.
|
||||||
|
Function is still named `_rollout_kv_fork` for caller-API stability.
|
||||||
|
|
||||||
Why turn-boundary close+nudge: matches what a chat UI emits when a human
|
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
|
interrupts a partial assistant turn. On-policy in chat-tuned data, where the
|
||||||
@@ -16,7 +19,6 @@ prior `\\nI should answer now.</think>` mid-turn splice was OOD.
|
|||||||
"""
|
"""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import copy
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -48,19 +50,6 @@ def _split_choice_ids(choice_token_ids: list) -> tuple[list[int], list[int]]:
|
|||||||
return list(choice_token_ids), []
|
return list(choice_token_ids), []
|
||||||
|
|
||||||
|
|
||||||
def _assert_full_attention(model) -> None:
|
|
||||||
"""Hybrid models (Qwen3.5 linear attention, Mamba SSM, RWKV) carry conv/recurrent
|
|
||||||
state that doesn't round-trip through past_key_values without fast-path libs.
|
|
||||||
Result: KV-fork suffix logits diverge from a flat forward by 0.1-0.7 nats per
|
|
||||||
slot. Fail loud rather than silently produce wrong logratios."""
|
|
||||||
layer_types = getattr(model.config, "layer_types", None)
|
|
||||||
if layer_types is None:
|
|
||||||
return
|
|
||||||
non_full = sorted({t for t in layer_types if t != "full_attention"})
|
|
||||||
assert not non_full, (
|
|
||||||
f"requires full-attention layers; model.config.layer_types contains {non_full}. "
|
|
||||||
f"Hybrid models lose conv/SSM state through past_key_values, breaking KV-fork."
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
@@ -89,7 +78,6 @@ def _rollout_kv_fork(
|
|||||||
of log-probs at last suffix position for those token ids (in the same order).
|
of log-probs at last suffix position for those token ids (in the same order).
|
||||||
Used by forced-choice (K-way) probing where a 2-way logratio is insufficient.
|
Used by forced-choice (K-way) probing where a 2-way logratio is insufficient.
|
||||||
"""
|
"""
|
||||||
_assert_full_attention(model)
|
|
||||||
if tok.padding_side != "left":
|
if tok.padding_side != "left":
|
||||||
raise ValueError("tok.padding_side must be 'left'")
|
raise ValueError("tok.padding_side must be 'left'")
|
||||||
device = next(model.parameters()).device
|
device = next(model.parameters()).device
|
||||||
@@ -152,8 +140,7 @@ def _rollout_kv_fork(
|
|||||||
pref_input[i, P_max - L:] = torch.tensor(sp_ids, device=device)
|
pref_input[i, P_max - L:] = torch.tensor(sp_ids, device=device)
|
||||||
pref_attn[i, P_max - L:] = 1
|
pref_attn[i, P_max - L:] = 1
|
||||||
pref_real[i] = L
|
pref_real[i] = L
|
||||||
pref_out = model(input_ids=pref_input, attention_mask=pref_attn, use_cache=True)
|
pref_out = model(input_ids=pref_input, attention_mask=pref_attn, use_cache=False)
|
||||||
cache = pref_out.past_key_values
|
|
||||||
|
|
||||||
# === Per-row prompt NLL (free; reuses pref_out.logits) ===
|
# === Per-row prompt NLL (free; reuses pref_out.logits) ===
|
||||||
# Teacher-forcing on the chat tokens (= rendered user turn + generation
|
# Teacher-forcing on the chat tokens (= rendered user turn + generation
|
||||||
@@ -180,7 +167,12 @@ def _rollout_kv_fork(
|
|||||||
nll = -logp.gather(-1, targets[:, None]).squeeze(-1).mean()
|
nll = -logp.gather(-1, targets[:, None]).squeeze(-1).mean()
|
||||||
nll_prompts.append(float(nll.item()))
|
nll_prompts.append(float(nll.item()))
|
||||||
|
|
||||||
# === Phase 2b: per-slot KV-fork ===
|
# === Phase 2b: per-slot flat prefix+suffix forward ===
|
||||||
|
# We re-encode the prefix per slot rather than re-using a KV cache. Slower
|
||||||
|
# (no cache reuse, ~7× per-slot cost on full-attention models) but correct
|
||||||
|
# for both full-attention AND sliding-window/SSM models. KV-fork was
|
||||||
|
# incompatible with hybrid attention (Gemma-2/3) where the cache drops
|
||||||
|
# context past the window.
|
||||||
a_ids, b_ids = _split_choice_ids(choice_token_ids)
|
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
|
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
|
b_t = torch.tensor(b_ids, device=device, dtype=torch.long) if b_ids else None
|
||||||
@@ -199,7 +191,7 @@ def _rollout_kv_fork(
|
|||||||
for sp, sp_ids in zip(sp_per_row, sp_ids_per_row)]
|
for sp, sp_ids in zip(sp_per_row, sp_ids_per_row)]
|
||||||
|
|
||||||
def fork(suffixes: list[list[int]]) -> torch.Tensor:
|
def fork(suffixes: list[list[int]]) -> torch.Tensor:
|
||||||
"""Forward suffix with deepcopied cache. Returns [B, V] logp at last real suffix token."""
|
"""Flat prefix+suffix forward. Returns [B, V] logp at last real suffix token."""
|
||||||
J_max = max(len(s) for s in suffixes)
|
J_max = max(len(s) for s in suffixes)
|
||||||
suf_input = torch.full((B, J_max), pad_id, dtype=torch.long, device=device)
|
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)
|
suf_mask = torch.zeros((B, J_max), dtype=torch.long, device=device)
|
||||||
@@ -209,15 +201,13 @@ def _rollout_kv_fork(
|
|||||||
suf_input[i, :L] = torch.tensor(s, device=device)
|
suf_input[i, :L] = torch.tensor(s, device=device)
|
||||||
suf_mask[i, :L] = 1
|
suf_mask[i, :L] = 1
|
||||||
last_pos[i] = L - 1
|
last_pos[i] = L - 1
|
||||||
|
full_input = torch.cat([pref_input, suf_input], dim=1)
|
||||||
full_attn = torch.cat([pref_attn, suf_mask], dim=1)
|
full_attn = torch.cat([pref_attn, suf_mask], dim=1)
|
||||||
position_ids = pref_real[:, None] + torch.arange(J_max, device=device)[None, :]
|
out = model(input_ids=full_input, attention_mask=full_attn, use_cache=False)
|
||||||
forked = copy.deepcopy(cache)
|
|
||||||
out = model(
|
|
||||||
input_ids=suf_input, attention_mask=full_attn,
|
|
||||||
position_ids=position_ids, past_key_values=forked, use_cache=False,
|
|
||||||
)
|
|
||||||
logp = F.log_softmax(out.logits.float(), dim=-1)
|
logp = F.log_softmax(out.logits.float(), dim=-1)
|
||||||
return logp[torch.arange(B, device=device), last_pos]
|
# Suffix's last real token is at absolute position P_max + last_pos[i].
|
||||||
|
absolute_last = P_max + last_pos
|
||||||
|
return logp[torch.arange(B, device=device), absolute_last]
|
||||||
|
|
||||||
slots: list[list[dict]] = [[] for _ in range(B)]
|
slots: list[list[dict]] = [[] for _ in range(B)]
|
||||||
for j, (nudge, prefill) in enumerate(scoring_slots):
|
for j, (nudge, prefill) in enumerate(scoring_slots):
|
||||||
|
|||||||
Reference in New Issue
Block a user