From 5704b00175bd8bea8cf266e480e02eb5c362354a Mon Sep 17 00:00:00 2001
From: wassname <1103714+wassname@users.noreply.github.com>
Date: Tue, 28 Apr 2026 21:47:33 +0800
Subject: [PATCH] gemma4: disable thinking mode via enable_thinking=False in
apply_chat_template
Gemma 4 (E2B/E4B) uses channel-based thinking tokens (<|think|>, <|channel>).
chat_template_extras() detects this via template string and passes
enable_thinking=False to all apply_chat_template calls in data gen,
dilemmas eval, and KL calib (via build_chat_text). Qwen3 and Gemma 3
return {} (existing thinking-mode handling unchanged).
---
src/ws/_tok_extras.py | 15 ++++
src/ws/data.py | 5 +-
src/ws/eval/_steer_common.py | 154 +++++++++++++++++++++++++++++++++++
src/ws/eval/dilemmas.py | 2 +
4 files changed, 175 insertions(+), 1 deletion(-)
create mode 100644 src/ws/_tok_extras.py
create mode 100644 src/ws/eval/_steer_common.py
diff --git a/src/ws/_tok_extras.py b/src/ws/_tok_extras.py
new file mode 100644
index 0000000..dfd0d2c
--- /dev/null
+++ b/src/ws/_tok_extras.py
@@ -0,0 +1,15 @@
+"""Tiny tokenizer utilities with no ws imports (avoids circular deps)."""
+
+
+def chat_template_extras(tok) -> dict:
+ """Extra kwargs for apply_chat_template that vary by model family.
+
+ Gemma 4 family is identified by <|think|>/<|channel> in the Jinja template.
+ Pass enable_thinking=False explicitly so outputs skip the thought channel
+ even if the model would otherwise default to thinking mode.
+ Qwen3 and Gemma 3 have no such kwarg and return {}.
+ """
+ template = tok.chat_template or ""
+ if "<|think|>" in template or "<|channel>" in template:
+ return {"enable_thinking": False}
+ return {}
diff --git a/src/ws/data.py b/src/ws/data.py
index 6fe8c6f..54ee171 100644
--- a/src/ws/data.py
+++ b/src/ws/data.py
@@ -27,6 +27,8 @@ from loguru import logger
from tqdm.auto import tqdm
from transformers import AutoModelForCausalLM, AutoTokenizer
+from ws._tok_extras import chat_template_extras
+
REPO_ROOT = Path(__file__).resolve().parents[2]
DATA_DIR = REPO_ROOT / "data"
@@ -213,7 +215,8 @@ def _build_specs(topics, n_personas: int, n_samples: int, behavior: str):
@torch.no_grad()
def _gen(model, tok, sys_prompt: str, user_prompt: str, max_new_tokens: int, temperature: float):
msgs = [{"role": "system", "content": sys_prompt}, {"role": "user", "content": user_prompt}]
- text = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True)
+ text = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True,
+ **chat_template_extras(tok))
inputs = tok(text, return_tensors="pt").to(model.device)
out = model.generate(
**inputs,
diff --git a/src/ws/eval/_steer_common.py b/src/ws/eval/_steer_common.py
new file mode 100644
index 0000000..d608c23
--- /dev/null
+++ b/src/ws/eval/_steer_common.py
@@ -0,0 +1,154 @@
+"""Shared steering primitives used by both KL calibration and dilemma eval.
+
+Why share this module: prompt formatting, special-token boundaries, and
+steering-context wiring are exactly the surface where bugs hide. If calib and
+eval don't share this code, you can fix calib without fixing eval (or vice
+versa) and never notice. Everything here is what both scripts call.
+
+Provides:
+ - chat-template builders (text + ids)
+ - unified steering_context: dW / repe / prompt / base under one with-block
+ - greedy_generate_under_steering: greedy-roll n_new_tokens with steering on
+ - teacher_force_logp: forward fixed ids, return log-probs at last n positions
+ - log_sample_prompt: dumps the full chat-templated string with special tokens
+ visible (\n's, <|im_start|>, etc.) so prompt-formatting bugs surface in logs
+"""
+
+from __future__ import annotations
+
+from contextlib import contextmanager
+
+import torch
+from baukit import TraceDict
+from loguru import logger
+from torch import Tensor
+
+from ws._tok_extras import chat_template_extras # noqa: F401 (re-export)
+from ws.eval.activation_baseline import _edit_all_tokens_per_layer
+from ws.steer import weight_steer
+
+
+THINK_OPEN = ""
+THINK_CLOSE = ""
+
+
+def build_chat_text(tok, system: str, user: str, assistant_prefix: str,
+ *, thinking: bool = False) -> str:
+ """Render [sys?, user, assistant=prefix] through the model's chat template.
+
+ `continue_final_message=True` means the assistant turn stays open, so the
+ next-token distribution is over the *continuation* of `assistant_prefix`,
+ not over a fresh assistant turn header.
+
+ If `thinking=True`, post-process the rendered text so the assistant turn
+ ends inside an *open* `` block — Qwen3's chat template auto-injects
+ `\\n\\n\\n\\n` when the prefix doesn't start with ``.
+ We snip everything after the last `` so the next-token distribution
+ is over reasoning tokens, matching the gist's "20 thinking tokens" budget.
+ """
+ msgs = []
+ if system:
+ msgs.append({"role": "system", "content": system})
+ msgs.append({"role": "user", "content": user})
+ msgs.append({"role": "assistant", "content": assistant_prefix})
+ text = tok.apply_chat_template(
+ msgs, tokenize=False,
+ continue_final_message=True, add_generation_prompt=False,
+ **chat_template_extras(tok),
+ )
+ if thinking:
+ idx = text.rfind(THINK_OPEN)
+ if idx >= 0:
+ text = text[: idx + len(THINK_OPEN)] + "\n"
+ return text
+
+
+def build_chat_ids(tok, system: str, user: str, assistant_prefix: str,
+ max_total: int = 512, *, thinking: bool = False) -> Tensor:
+ text = build_chat_text(tok, system, user, assistant_prefix, thinking=thinking)
+ enc = tok(text, return_tensors="pt", truncation=True, max_length=max_total)
+ return enc.input_ids.squeeze(0)
+
+
+@contextmanager
+def steering_context(method: str, alpha: float, *, model,
+ w=None, repe_dirs=None, repe_layers=None):
+ """Unified steering for dW: / repe / prompt: / base.
+
+ `prompt:` and `base` are nullcontext — their "steering" is the system
+ prompt baked into input_ids upstream, not a runtime hook.
+ """
+ if method.startswith("dW:"):
+ with weight_steer(model, w, alpha):
+ yield
+ elif method == "repe":
+ hooks = [f"model.layers.{L}" for L in repe_layers]
+ edit = _edit_all_tokens_per_layer(repe_dirs, list(repe_layers), alpha)
+ with TraceDict(model, hooks, edit_output=edit):
+ yield
+ elif method.startswith("prompt:") or method == "base":
+ yield
+ else:
+ raise ValueError(f"unknown method: {method}")
+
+
+@torch.no_grad()
+def greedy_generate_under_steering(
+ model, tok, input_ids: Tensor, *, method: str, alpha: float,
+ n_new_tokens: int, w=None, repe_dirs=None, repe_layers=None,
+) -> tuple[Tensor, Tensor]:
+ """Greedy-generate n_new_tokens under steering. Returns (gen_ids[T], logp_steered[T,V]).
+
+ `output_scores=True` with `do_sample=False` returns the raw next-token
+ logits at each generation step — these are the steered model's actual
+ distribution at each rolled position.
+ """
+ with steering_context(method, alpha, model=model, w=w,
+ repe_dirs=repe_dirs, repe_layers=repe_layers):
+ out = model.generate(
+ input_ids.unsqueeze(0).to(model.device),
+ max_new_tokens=n_new_tokens, do_sample=False, temperature=1.0,
+ return_dict_in_generate=True, output_scores=True,
+ pad_token_id=tok.pad_token_id, eos_token_id=tok.eos_token_id,
+ )
+ new_ids = out.sequences[0, input_ids.shape[0]:].cpu()
+ # output_scores: tuple of [B, V] tensors, one per generated step
+ logp_steered = torch.stack(
+ [s[0].float().log_softmax(-1) for s in out.scores], dim=0
+ ).cpu()
+ # If gen stopped early on EOS, scores has one extra step than new_ids; trim
+ logp_steered = logp_steered[: new_ids.shape[0]]
+ return new_ids, logp_steered
+
+
+@torch.no_grad()
+def teacher_force_logp(model, full_ids: Tensor, n_tokens: int) -> Tensor:
+ """Forward `full_ids` once, return log-probs at the last n_tokens positions.
+
+ Specifically: returns log-probs of distributions that *predict* the last
+ n_tokens of `full_ids` (i.e. positions [-n_tokens-1 : -1] of the logits).
+ """
+ out = model(input_ids=full_ids.unsqueeze(0).to(model.device))
+ logits = out.logits[0, -n_tokens - 1:-1]
+ return logits.float().log_softmax(-1).cpu()
+
+
+def log_sample_prompt(tok, text: str, *, generated_ids: Tensor | None = None,
+ label: str = "sample", max_chars: int = 1200) -> None:
+ """Log the full chat-templated prompt with special tokens visible.
+
+ Use this once per method/per script run. The point is: if the chat
+ template silently changes between calib and eval, you see it in the log
+ before debugging metrics.
+ """
+ snippet = text if len(text) <= max_chars else (text[:max_chars] + f"…[+{len(text) - max_chars} chars]")
+ logger.info(f"[{label}] full prompt (special tokens included):\n{snippet}")
+ ids = tok(text, return_tensors="pt").input_ids[0]
+ first = tok.convert_ids_to_tokens(ids[: min(8, len(ids))].tolist())
+ last = tok.convert_ids_to_tokens(ids[-min(8, len(ids)):].tolist())
+ logger.info(f"[{label}] n_input_tokens={ids.shape[0]} first8={first} last8={last}")
+ if generated_ids is not None and len(generated_ids) > 0:
+ gen_text = tok.decode(generated_ids, skip_special_tokens=False)
+ gen_toks = tok.convert_ids_to_tokens(generated_ids.tolist())
+ logger.info(f"[{label}] greedy gen ({len(generated_ids)} toks): {gen_text!r}")
+ logger.info(f"[{label}] greedy gen tokens: {gen_toks}")
diff --git a/src/ws/eval/dilemmas.py b/src/ws/eval/dilemmas.py
index d85fbfb..b8f118f 100644
--- a/src/ws/eval/dilemmas.py
+++ b/src/ws/eval/dilemmas.py
@@ -28,6 +28,7 @@ from torch import Tensor
from torch.utils.data import DataLoader
from transformers import AutoModelForCausalLM, AutoTokenizer, DataCollatorWithPadding
+from ws._tok_extras import chat_template_extras
from ws.eval.sycophancy import get_choice_ids
from ws.steer import weight_steer
@@ -84,6 +85,7 @@ def _format_row(row: dict, tok, max_tokens: int, system_prompt: str = "") -> dic
return_tensors="pt",
truncation=True,
max_length=max_tokens,
+ **chat_template_extras(tok),
)
input_ids = encoded.input_ids.squeeze(0) if hasattr(encoded, "input_ids") else encoded.squeeze(0)