From ada854c499ae5d0db50455dbedc84c5dced876c1 Mon Sep 17 00:00:00 2001
From: wassname <1103714+wassname@users.noreply.github.com>
Date: Tue, 5 May 2026 22:19:11 +0800
Subject: [PATCH] guided_rollout_multibool: switch to 12 single-slot KV-forks +
assert full-attention
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit
The chained-fill design (one suffix with all foundations, two passes for true/false)
hit a non-recoverable conv-state issue on hybrid linear-attention layers (Qwen3.5):
splitting prefix and suffix forwards via past_key_values silently produced
wrong logits (pmass dropping to 0.04, top token leaking to ' "' = 0.72).
Switched to 12 independent single-slot completions per prompt:
for (frame, foundation) in {is_violation, is_ok} × foundations:
cache scoring_prefix once, fork suffix `\n{"": {"":`,
read logits at the last token (predicting `true|false`).
final[f] = 0.5 * (lr_violation[f] - lr_ok[f])
Framing flip cancels per-key prior bias the same way true/false fill did,
without the chained-slot causality that interacts badly with split forwards.
Added _assert_full_attention(): checks model.config.layer_types and fails
loudly on hybrid models. Verified parity vs flat forward on Qwen3-0.6B
(Δ ≤ 0.13 nats; signal of interest is ≫1 nat) and assert fires on Qwen3.5-0.8B.
Co-Authored-By: Claude Opus 4.7
---
src/tinymfv/guided.py | 226 +++++++++++++++++++++++++++++++++++++++++-
1 file changed, 225 insertions(+), 1 deletion(-)
diff --git a/src/tinymfv/guided.py b/src/tinymfv/guided.py
index 3f646cb..9fe1b60 100644
--- a/src/tinymfv/guided.py
+++ b/src/tinymfv/guided.py
@@ -1,7 +1,8 @@
from __future__ import annotations
import contextlib
-from dataclasses import dataclass
+import copy
+from dataclasses import dataclass, field
import torch
import torch.nn.functional as F
@@ -305,6 +306,229 @@ def guided_rollout_batch(
return results
+@dataclass
+class MultiBoolResult:
+ user_prompt: str
+ think_text: str
+ logratios: dict[str, float] # foundation -> 0.5*(lr_violation - lr_ok); positive = violates
+ pmass_format: dict[str, float] # foundation -> avg(pm_violation, pm_ok); ≈1 if formatting holds
+ lr_violation: dict[str, float] = field(default_factory=dict) # diagnostic: per-frame raw logratios
+ lr_ok: dict[str, float] = field(default_factory=dict)
+ think_tokens: int = 0
+ emitted_close: bool = False
+
+
+_DEFAULT_MULTIBOOL_HINT: str = (
+ 'After thinking, output JSON: {"is_violation": {: true|false}} '
+ "marking which Moral Foundations Theory foundation(s) the scenario violates."
+)
+
+_DEFAULT_FOUNDATIONS: tuple[str, ...] = (
+ "care", "fairness", "loyalty", "authority", "sanctity", "liberty",
+)
+
+
+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 (causal-conv1d, flash-linear-attention). 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"guided_rollout_multibool requires full-attention layers; "
+ f"model.config.layer_types contains {non_full}. Hybrid models lose "
+ f"conv/SSM state through past_key_values, breaking KV-fork. "
+ f"Use a full-attention model (e.g. Qwen/Qwen3-0.6B) instead."
+ )
+
+
+@torch.no_grad()
+def guided_rollout_multibool(
+ model, tok,
+ user_prompts: list[str],
+ foundations: list[str] | None = None,
+ *,
+ max_think_tokens: int = 256,
+ schema_hint: str = _DEFAULT_MULTIBOOL_HINT,
+) -> list[MultiBoolResult]:
+ """Score per-foundation violation logratio via 12 KV-forked single-slot completions.
+
+ For each prompt:
+ 1. Batched think generate (≤max_think_tokens).
+ 2. Cache scoring_prefix = chat + think + (one batched fwd).
+ 3. For each (frame, foundation) ∈ {is_violation, is_ok} × foundations,
+ fork from the cache and complete with `\\n{"": {"":`. Read
+ logits at the LAST suffix token (predicting ` true|false`).
+ 4. final[f] = 0.5 * (lr_violation[f] - lr_ok[f]). Flipping the framing
+ flips which token means "violates", so averaging cancels per-key
+ priors (e.g. JSON-true bias) without the chained-slot causality of
+ the all-true / all-false fill design.
+
+ Cost: 1 prefix forward + 12 small suffix forwards (each ≤14 tokens),
+ all batched over B prompts.
+
+ Requires a full-attention model -- asserts at entry."""
+ _assert_full_attention(model)
+ if tok.padding_side != "left":
+ raise ValueError("tok.padding_side must be 'left' for batched rollout")
+ if foundations is None:
+ foundations = list(_DEFAULT_FOUNDATIONS)
+ device = next(model.parameters()).device
+ pad_id = tok.pad_token_id if tok.pad_token_id is not None else tok.eos_token_id
+
+ # === Phase 1: think generation (batched) ===
+ chats = []
+ for up in user_prompts:
+ full = f"{up}\n\n{schema_hint}" if schema_hint else up
+ msgs = [{"role": "user", "content": full}]
+ try:
+ p = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True)
+ except TypeError:
+ p = tok.apply_chat_template(msgs, tokenize=False)
+ chats.append(p + "\n")
+
+ think_end_id = tok.convert_tokens_to_ids("")
+ if think_end_id in (None, getattr(tok, "unk_token_id", None)):
+ think_end_id = tok.eos_token_id
+
+ enc = tok(chats, return_tensors="pt", padding=True).to(device)
+ prompt_len = enc.input_ids.shape[1]
+ phase1 = model.generate(
+ **enc,
+ max_new_tokens=max_think_tokens,
+ do_sample=False,
+ eos_token_id=think_end_id,
+ pad_token_id=pad_id,
+ )
+
+ # === Build per-row scoring_prefix ===
+ sp_per_row: list[str] = []
+ sp_ids_per_row: list[list[int]] = []
+ think_per_row: list[tuple[str, int, bool]] = [] # (think_text, n_think, emitted_close)
+ for i, p in enumerate(chats):
+ gen_ids = phase1[i, prompt_len:]
+ keep = gen_ids != pad_id
+ gen_ids = gen_ids[keep] if keep.any() else gen_ids[: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
+ if emitted_close:
+ think_text, _ = gen_text.split(_CLOSE_MARKER, 1)
+ sp = p + think_text + _CLOSE_MARKER
+ else:
+ think_text = gen_text
+ sp = p + gen_text + "\nI should answer now." + _CLOSE_MARKER
+ sp_per_row.append(sp)
+ sp_ids_per_row.append(tok(sp, add_special_tokens=False)["input_ids"])
+ think_per_row.append((think_text, n_think, emitted_close))
+
+ B = len(sp_per_row)
+ P_max = max(len(s) for s in sp_ids_per_row)
+
+ # === Phase 2a: batched prefix forward (left-padded) ===
+ pref_input = torch.full((B, P_max), pad_id, dtype=torch.long, device=device)
+ pref_attn = torch.zeros((B, P_max), dtype=torch.long, device=device)
+ pref_real = torch.zeros(B, dtype=torch.long, device=device)
+ for i, sp_ids in enumerate(sp_ids_per_row):
+ L = len(sp_ids)
+ pref_input[i, P_max - L:] = torch.tensor(sp_ids, device=device)
+ pref_attn[i, P_max - L:] = 1
+ pref_real[i] = L
+ pref_out = model(input_ids=pref_input, attention_mask=pref_attn, use_cache=True)
+ cache = pref_out.past_key_values
+
+ # === Phase 2b: 12 forked suffix forwards. Cache is left-padded, so
+ # position_ids must use the per-row REAL prefix length (not P_max).
+ a_ids, b_ids = choice_token_ids_tf(tok)
+ a_t = torch.tensor(a_ids, device=device, dtype=torch.long)
+ b_t = torch.tensor(b_ids, device=device, dtype=torch.long)
+ all_ids = torch.tensor(a_ids + b_ids, device=device, dtype=torch.long)
+
+ def fork(suffixes: list[list[int]]) -> torch.Tensor:
+ """Run one suffix forward with deepcopied cache. Returns [B, V] logp at
+ the LAST real suffix token per row."""
+ 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) # right-pad
+ suf_mask[i, :L] = 1
+ last_pos[i] = L - 1
+ full_attn = torch.cat([pref_attn, suf_mask], dim=1)
+ position_ids = pref_real[:, None] + torch.arange(J_max, device=device)[None, :]
+ 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)
+ return logp[torch.arange(B, device=device), last_pos] # [B, V]
+
+ def suf_ids_for(frame: str, foundation: str) -> list[list[int]]:
+ suf_text = f'\n{{"{frame}": {{"{foundation}":'
+ return [tok(sp + suf_text, add_special_tokens=False)["input_ids"][len(sp_ids):]
+ for sp, sp_ids in zip(sp_per_row, sp_ids_per_row)]
+
+ lr_per: dict[tuple[str, str], torch.Tensor] = {}
+ pm_per: dict[tuple[str, str], torch.Tensor] = {}
+ for frame in ("is_violation", "is_ok"):
+ for f in foundations:
+ lp = fork(suf_ids_for(frame, f)) # [B, V]
+ la = torch.logsumexp(lp[:, a_t], dim=-1)
+ lb = torch.logsumexp(lp[:, b_t], dim=-1)
+ lr_per[(frame, f)] = (la - lb).cpu()
+ pm_per[(frame, f)] = lp[:, all_ids].exp().sum(-1).cpu()
+
+ # === Aggregate: final = 0.5*(lr_violation - lr_ok), pmass = avg ===
+ results: list[MultiBoolResult] = []
+ for i, (think_text, n_think, emitted_close) in enumerate(think_per_row):
+ per_logr: dict[str, float] = {}
+ per_pm: dict[str, float] = {}
+ per_lr_v: dict[str, float] = {}
+ per_lr_o: dict[str, float] = {}
+ for f in foundations:
+ lr_v = float(lr_per[("is_violation", f)][i].item())
+ lr_o = float(lr_per[("is_ok", f)][i].item())
+ per_lr_v[f] = lr_v
+ per_lr_o[f] = lr_o
+ per_logr[f] = 0.5 * (lr_v - lr_o)
+ pm_v = float(pm_per[("is_violation", f)][i].item())
+ pm_o = float(pm_per[("is_ok", f)][i].item())
+ per_pm[f] = 0.5 * (pm_v + pm_o)
+
+ results.append(MultiBoolResult(
+ user_prompt=user_prompts[i],
+ think_text=think_text,
+ logratios=per_logr,
+ pmass_format=per_pm,
+ lr_violation=per_lr_v,
+ lr_ok=per_lr_o,
+ think_tokens=n_think,
+ emitted_close=emitted_close,
+ ))
+
+ # SHOULD: pmass≈1 at every (foundation, frame). Aggregate-once warning.
+ worst = min(
+ ((i, f, r.pmass_format[f]) for i, r in enumerate(results) for f in foundations),
+ key=lambda x: x[2], default=None,
+ )
+ if worst is not None and worst[2] < 0.5:
+ logger.warning(
+ f"multibool pmass<0.5 at {worst[1]} on row {worst[0]} (pm={worst[2]:.3f}); "
+ "schema may be drifting under steering"
+ )
+ return results
+
+
def choice_token_ids_tf(tok) -> list[list[int]]:
def _variants(words):
seen = []