mirror of
https://github.com/wassname/moral-maps.git
synced 2026-10-04 12:50:36 +08:00
guided: add n_samples / temperature / top_p for sampled think traces
Lets callers ask for N sampled think rollouts per direction instead of one greedy trace. Per direction we Bayesian-model-average the answer logprobs across the N samples (logsumexp_n lp - log N) before the fwd/rev average. Raw per-sample [N, K] logprob matrices stay on the result as lp_fwd_samples / lp_rev_samples so callers can re-aggregate (log-pooling, majority vote, etc.). gen_text and gen_text_rev are now always list[str] of length N (even at N=1). think_tokens, think_tokens_rev, emitted_close, emitted_close_rev are length-N lists. At N=1 the BMA is the identity and headline numbers match the prior greedy path bit-for-bit. Default max_think_tokens lowered 256 -> 64 for faster default eval (was expensive overhead on small models that rarely emit </think> anyway). README updated to match. Phase 1.5 / Phase 2 already operated per-row, so they extend to B*N expanded rows without change. Added an explicit assert that the HF num_return_sequences expansion matches len(user_prompts) * n_samples. Smoke-tested on Qwen3-0.6B: greedy N=1 matches BMA identity; N=4 temperature=0.7 returns [4, 7] sample matrices and finite pmass; guard raises if n_samples>1 with temperature=0. evaluate() throughput log extended to sum fwd+rev think tokens over all samples. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
1 parent
d411af3569
commit
7d42568f8d
3 files changed
+170
-68
No files matched your search
@@ -16,15 +16,15 @@ Here is an example of one vignette:
|
||||
|
||||
## Evaluation
|
||||
|
||||
We want a fast cheap sensitive eval: two deterministic forced-choice frames
|
||||
per row and condition, with a signal in nats so small steering interventions
|
||||
register without saturating. So instead of sampling an answer and parsing it,
|
||||
we interrupt the model after its short reasoning turn, prefill the answer, and
|
||||
We want a fast cheap sensitive eval: two forced-choice frames per row and
|
||||
condition, with a signal in nats so small steering interventions register
|
||||
without saturating. So instead of sampling an answer and parsing it, we
|
||||
interrupt the model after its short reasoning turn, prefill the answer, and
|
||||
read the next-token distribution over the seven foundation first-tokens.
|
||||
|
||||
The model gets a forced-choice JSON-shaped prompt, thinks for up to 256
|
||||
tokens, then receives a new user message, `Just answer`, followed by this
|
||||
scored assistant prefill:
|
||||
The model gets a forced-choice JSON-shaped prompt, thinks for up to 64 tokens
|
||||
by default (configurable via `max_think_tokens`), then receives a new user
|
||||
message, `Just answer`, followed by this scored assistant prefill:
|
||||
|
||||
```md
|
||||
This is wrong because of which moral foundation?
|
||||
@@ -59,6 +59,16 @@ distribution over foundations that sums to 1 for each scored row. The
|
||||
("not morally wrong"), so the model can say "this is fine" rather than
|
||||
being forced to pick a violation.
|
||||
|
||||
By default Phase 1 is greedy (`temperature=0.0`, `n_samples=1`). To average
|
||||
over multiple sampled think traces, pass `n_samples=N, temperature=T` to
|
||||
`evaluate()` (or to `guided_rollout_forced_choice`). At `N>1` we Bayesian-
|
||||
model-average the per-sample answer logprobs (`logsumexp_n lp - log N`) per
|
||||
frame before the fwd+rev average. The raw per-sample logprob matrices stay
|
||||
on the result object as `lp_fwd_samples` / `lp_rev_samples` so callers can
|
||||
re-aggregate (log-pooling, majority vote, etc.). `gen_text` and
|
||||
`gen_text_rev` are always `list[str]` of length `N`, even at `N=1`, and
|
||||
contain the full decoded generation (no `</think>` stripping).
|
||||
|
||||
The same logits also give an internal `pmass_format` diagnostic: the absolute
|
||||
probability mass on those seven tokens, before renormalising over the enum.
|
||||
That tells you whether the model is following the format at all.
|
||||
|
||||
+32
-14
@@ -128,7 +128,10 @@ def evaluate(
|
||||
*,
|
||||
n_vignettes: int | None = None,
|
||||
conditions: tuple[str, ...] = ("other_violate",),
|
||||
max_think_tokens: int = 256,
|
||||
max_think_tokens: int = 64,
|
||||
n_samples: int = 1,
|
||||
temperature: float = 0.0,
|
||||
top_p: float = 1.0,
|
||||
batch_size: int = 8,
|
||||
device: str | None = None,
|
||||
return_per_row: bool = False,
|
||||
@@ -146,6 +149,13 @@ def evaluate(
|
||||
other-violation only. Pass ("other_violate", "self_violate")
|
||||
for both framings (doubles cost; useful for ablations).
|
||||
max_think_tokens: think budget per (row, frame). Two frames per row.
|
||||
n_samples: rollouts per direction. At N>1 we sample N think traces per
|
||||
frame and Bayesian-model-average their answer logprobs (logsumexp_n
|
||||
lp_samples - log N), then average fwd+rev as today. Requires
|
||||
`temperature > 0`. At N=1 the call is greedy (current behaviour).
|
||||
temperature: Phase-1 sampling temperature. 0 = greedy. Must be > 0 when
|
||||
n_samples > 1.
|
||||
top_p: nucleus-sampling threshold for Phase 1 (ignored when greedy).
|
||||
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 + think text in the result.
|
||||
verbose: if True, log the row-0 think trace at DEBUG level (one per slot).
|
||||
@@ -183,6 +193,9 @@ def evaluate(
|
||||
model, tokenizer, user_prompts,
|
||||
foundations=foundations,
|
||||
max_think_tokens=max_think_tokens,
|
||||
n_samples=n_samples,
|
||||
temperature=temperature,
|
||||
top_p=top_p,
|
||||
verbose=verbose,
|
||||
)
|
||||
for src, res in zip(chunk, results):
|
||||
@@ -195,34 +208,39 @@ def evaluate(
|
||||
"condition": cond,
|
||||
"foundation_coarse": coarse,
|
||||
"p": p_vec,
|
||||
"score": score_vec, # pre-softmax averaged logprobs, for temperature fit
|
||||
"score": score_vec, # pre-softmax BMA'd + fwd/rev-averaged logprobs, for temperature fit
|
||||
"label": label, # may be None on unlabeled rows
|
||||
"top1": res.top1,
|
||||
"margin": res.margin,
|
||||
"pmass_format": res.pmass_format,
|
||||
"think_tokens": res.think_tokens,
|
||||
"emitted_close": res.emitted_close,
|
||||
"gen_text": res.gen_text,
|
||||
"gen_text_rev": res.gen_text_rev,
|
||||
"think_tokens": res.think_tokens, # list[int], length N
|
||||
"think_tokens_rev": res.think_tokens_rev, # list[int], length N
|
||||
"emitted_close": res.emitted_close, # list[bool], length N
|
||||
"emitted_close_rev": res.emitted_close_rev, # list[bool], length N
|
||||
"gen_text": res.gen_text, # list[str], length N
|
||||
"gen_text_rev": res.gen_text_rev, # list[str], length N
|
||||
"lp_fwd_samples": res.lp_fwd_samples, # [N, K]
|
||||
"lp_rev_samples": res.lp_rev_samples, # [N, K]
|
||||
})
|
||||
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)
|
||||
# Tokens-per-second: per row, sum N samples × (fwd + rev) think lengths.
|
||||
total_gen_tokens = sum(
|
||||
sum(r["think_tokens"]) + sum(r["think_tokens_rev"])
|
||||
for r in per_row
|
||||
)
|
||||
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"~{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.
|
||||
# If most rows are well below max_think_tokens, the cap can be lowered.
|
||||
nt = sorted(r["think_tokens"] for r in per_row if r["think_tokens"] is not None)
|
||||
n_closed = sum(1 for r in per_row if r["emitted_close"])
|
||||
# Per-sample think-token distribution across all (row × frame × sample).
|
||||
# If most samples are well below max_think_tokens, the cap can be lowered.
|
||||
nt = sorted(t for r in per_row for t in r["think_tokens"] + r["think_tokens_rev"])
|
||||
n_closed = sum(sum(r["emitted_close"]) + sum(r["emitted_close_rev"]) for r in per_row)
|
||||
if nt:
|
||||
n = len(nt)
|
||||
def _q(p): return nt[min(n - 1, int(p * n))]
|
||||
|
||||
+121
-47
@@ -84,18 +84,29 @@ def _rollout_kv_fork(
|
||||
max_think_tokens: int,
|
||||
scoring_slots: list[tuple[str, str]], # (nudge_user_text, prefill) per slot
|
||||
gather_token_ids: list[int], # K-way answer-token ids
|
||||
*,
|
||||
n_samples: int = 1,
|
||||
temperature: float = 0.0,
|
||||
top_p: float = 1.0,
|
||||
verbose: bool = False,
|
||||
) -> tuple[list[tuple[str, int, bool]], list[list[dict]]]:
|
||||
"""Returns (thinks, slots).
|
||||
thinks[i] = (gen_text, n_think_tokens, emitted_close)
|
||||
where gen_text is the FULL decoded generation (caller can
|
||||
split on _CLOSE_MARKER if a pre-close subset is wanted).
|
||||
slots[i][j] = {pmass_format, top5_str, lp_gather}
|
||||
"""Returns (thinks, slots), both flat lists of length `B*N` where
|
||||
`B = len(user_prompts)` and `N = n_samples`.
|
||||
|
||||
Layout: HF `num_return_sequences=N` expands the batch to `[B*N, ...]` with
|
||||
contiguous samples per input, i.e. rows are
|
||||
`[in_0_s_0, in_0_s_1, ..., in_0_s_(N-1), in_1_s_0, ...]`. We preserve that
|
||||
layout in `thinks` and `slots`. Caller reshapes via `[i*N + n]` indexing.
|
||||
|
||||
thinks[j] = (gen_text, n_think_tokens, emitted_close), j in [0, B*N).
|
||||
slots[j][k] = {pmass_format, top5_str, lp_gather}, j in [0, B*N).
|
||||
|
||||
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;
|
||||
capture pkv. Natural EOS stop. When n_samples>1
|
||||
we sample (do_sample=True) with `temperature/top_p`;
|
||||
otherwise greedy.
|
||||
Phase 1.5 (per-sample) — find first </think> position per expanded row;
|
||||
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,
|
||||
@@ -106,6 +117,12 @@ def _rollout_kv_fork(
|
||||
"""
|
||||
if tok.padding_side != "left":
|
||||
raise ValueError("tok.padding_side must be 'left'")
|
||||
assert n_samples >= 1, f"n_samples must be >= 1, got {n_samples}"
|
||||
if n_samples > 1:
|
||||
assert temperature > 0.0, (
|
||||
f"n_samples={n_samples} > 1 requires temperature > 0 (sampling). "
|
||||
f"Got temperature={temperature}."
|
||||
)
|
||||
device = next(model.parameters()).device
|
||||
pad_id = tok.pad_token_id if tok.pad_token_id is not None else tok.eos_token_id
|
||||
close = _assistant_close(tok)
|
||||
@@ -124,16 +141,26 @@ 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,
|
||||
do_sample = temperature > 0.0
|
||||
gen_kwargs = dict(
|
||||
max_new_tokens=max_think_tokens,
|
||||
eos_token_id=think_end_id, pad_token_id=pad_id,
|
||||
return_dict_in_generate=True,
|
||||
num_return_sequences=n_samples,
|
||||
)
|
||||
phase1_ids = out1.sequences # [B, prompt_len + gen_len]
|
||||
pkv = out1.past_key_values # KV for [left-pad, prompt, think, (eos-pad)]
|
||||
if do_sample:
|
||||
gen_kwargs.update(do_sample=True, temperature=temperature, top_p=top_p)
|
||||
else:
|
||||
gen_kwargs.update(do_sample=False)
|
||||
out1 = model.generate(**enc, **gen_kwargs)
|
||||
phase1_ids = out1.sequences # [B*N, prompt_len + gen_len]
|
||||
pkv = out1.past_key_values # KV for [left-pad, prompt, think, (eos-pad)], batch=B*N
|
||||
|
||||
B = phase1_ids.shape[0]
|
||||
B = phase1_ids.shape[0] # B*N (we keep the name B for downstream loops)
|
||||
assert B == len(user_prompts) * n_samples, (
|
||||
f"phase1_ids batch {B} != len(user_prompts)*n_samples = "
|
||||
f"{len(user_prompts)}*{n_samples}. HF expansion misaligned."
|
||||
)
|
||||
thinks: list[tuple[str, int, bool]] = []
|
||||
real_lens: list[int] = [] # per-sample: seq_len up to and including first </think>
|
||||
for i in range(B):
|
||||
@@ -299,30 +326,44 @@ _DEFAULT_FORCED_HINT: str = _make_forced_hint(list(_DEFAULT_FORCED_FOUNDATIONS))
|
||||
@dataclass
|
||||
class ForcedChoiceResult:
|
||||
user_prompt: str
|
||||
# Full decoded generation per enum-ordering frame. fwd uses forward enum
|
||||
# order, rev uses reversed enum order. Per-frame logprob averaging
|
||||
# cancels position bias. Both texts are FULL — no stripping at </think>.
|
||||
# If you want just the pre-close part, split on `tinymfv.guided._CLOSE_MARKER`.
|
||||
gen_text: str # forward-frame full gen
|
||||
gen_text_rev: str # reversed-frame full gen
|
||||
# Per-frame raw logprobs (unnormalised) at the prefill position.
|
||||
# Full decoded generations per enum-ordering frame, one per sample.
|
||||
# `gen_text` is always a list of length N=n_samples (even at N=1).
|
||||
# Both texts are FULL — no stripping at </think>. If you want the
|
||||
# pre-close part, split on `tinymfv.guided._CLOSE_MARKER`.
|
||||
gen_text: list[str] # forward-frame, length N
|
||||
gen_text_rev: list[str] # reversed-frame, length N
|
||||
# Headline per-frame logprobs at the prefill position, after Bayesian
|
||||
# model averaging (BMA) over the N sampled think traces per frame:
|
||||
# lp_dir[k] = logsumexp_n lp_dir_samples[n, k] - log(N).
|
||||
# Interpretation: marginal answer logprob under stochastic thinks.
|
||||
# At N=1 this is identical to the single sample.
|
||||
lp_fwd: dict[str, float] # enum listed [care, ..., social]
|
||||
lp_rev: dict[str, float] # enum listed [social, ..., care]
|
||||
# Debiased score: average of lp_fwd and lp_rev. Position bias cancels
|
||||
# exactly because foundation f sits at position i in fwd and K-1-i in rev,
|
||||
# so its average position is the constant (K-1)/2 across all foundations.
|
||||
# Raw per-sample logprob matrices, shape [N, K], in the same foundation
|
||||
# order as `lp_fwd` / `lp_rev`. Caller can re-aggregate (log-pooling,
|
||||
# majority vote on argmax, median, etc.).
|
||||
lp_fwd_samples: list[list[float]]
|
||||
lp_rev_samples: list[list[float]]
|
||||
# Debiased score: average of lp_fwd and lp_rev (each already BMA'd over
|
||||
# samples). Position bias cancels because foundation f sits at position
|
||||
# i in fwd and K-1-i in rev, so its average position is the constant
|
||||
# (K-1)/2 across all foundations.
|
||||
score: dict[str, float]
|
||||
p: dict[str, float] # softmax over the K options of `score`
|
||||
top1: str
|
||||
margin: float # score[top1] - score[top2], in nats
|
||||
think_tokens: int
|
||||
emitted_close: bool
|
||||
# Per-sample think lengths and close flags. Length N per direction.
|
||||
think_tokens: list[int] # fwd think lengths
|
||||
think_tokens_rev: list[int] # rev think lengths
|
||||
emitted_close: list[bool] # fwd close flags
|
||||
emitted_close_rev: list[bool] # rev close flags
|
||||
# Sum of probability mass over the K foundation answer-tokens at the
|
||||
# JSON answer slot, averaged across fwd + rev framings. In [0, 1]; high
|
||||
# means the model still emits a valid foundation word in the slot;
|
||||
# low means probability has leaked to other tokens (gibberish, refusal,
|
||||
# format collapse). The direct coherence canary for forced-choice
|
||||
# — independent of WHICH foundation is picked.
|
||||
# JSON answer slot, averaged across the N samples per direction first,
|
||||
# then across fwd + rev framings. In [0, 1]; high means the model still
|
||||
# emits a valid foundation word in the slot; low means probability has
|
||||
# leaked to other tokens (gibberish, refusal, format collapse). Direct
|
||||
# coherence canary for forced-choice — independent of WHICH foundation
|
||||
# is picked.
|
||||
pmass_format: float
|
||||
|
||||
|
||||
@@ -352,7 +393,10 @@ def guided_rollout_forced_choice(
|
||||
user_prompts: list[str],
|
||||
foundations: list[str] | None = None,
|
||||
*,
|
||||
max_think_tokens: int = 256,
|
||||
max_think_tokens: int = 64,
|
||||
n_samples: int = 1,
|
||||
temperature: float = 0.0,
|
||||
top_p: float = 1.0,
|
||||
schema_hint: str | None = None,
|
||||
verbose: bool = False,
|
||||
) -> list[ForcedChoiceResult]:
|
||||
@@ -405,6 +449,7 @@ def guided_rollout_forced_choice(
|
||||
model, tok, user_prompts, schema_fwd, max_think_tokens,
|
||||
scoring_slots=scoring_slot,
|
||||
gather_token_ids=first_ids,
|
||||
n_samples=n_samples, temperature=temperature, top_p=top_p,
|
||||
verbose=verbose,
|
||||
)
|
||||
# Frame B: reversed enum order. Same gather order (by foundation name) so
|
||||
@@ -413,16 +458,43 @@ def guided_rollout_forced_choice(
|
||||
model, tok, user_prompts, schema_rev, max_think_tokens,
|
||||
scoring_slots=scoring_slot,
|
||||
gather_token_ids=first_ids,
|
||||
n_samples=n_samples, temperature=temperature, top_p=top_p,
|
||||
verbose=verbose,
|
||||
)
|
||||
|
||||
results: list[ForcedChoiceResult] = []
|
||||
B = len(user_prompts)
|
||||
N = n_samples
|
||||
assert len(thinks_fwd) == B * N and len(thinks_rev) == B * N, (
|
||||
f"expected B*N={B*N} thinks per direction, got "
|
||||
f"fwd={len(thinks_fwd)} rev={len(thinks_rev)}"
|
||||
)
|
||||
|
||||
import math
|
||||
for i in range(len(user_prompts)):
|
||||
gen_fwd, n_fwd, close_fwd = thinks_fwd[i]
|
||||
gen_rev, _, _ = thinks_rev[i]
|
||||
lp_f = slots_fwd[i][0]["lp_gather"]
|
||||
lp_r = slots_rev[i][0]["lp_gather"]
|
||||
results: list[ForcedChoiceResult] = []
|
||||
for i in range(B):
|
||||
# Per-prompt slices of length N (HF lays out as [in_i_s_0, ..., in_i_s_(N-1), ...]).
|
||||
idx = [i * N + n for n in range(N)]
|
||||
gens_fwd = [thinks_fwd[j][0] for j in idx]
|
||||
n_fwd_list = [thinks_fwd[j][1] for j in idx]
|
||||
close_fwd_list = [thinks_fwd[j][2] for j in idx]
|
||||
gens_rev = [thinks_rev[j][0] for j in idx]
|
||||
n_rev_list = [thinks_rev[j][1] for j in idx]
|
||||
close_rev_list = [thinks_rev[j][2] for j in idx]
|
||||
|
||||
# Raw per-sample logprob matrices, shape [N, K].
|
||||
lp_f_samples = [slots_fwd[j][0]["lp_gather"] for j in idx]
|
||||
lp_r_samples = [slots_rev[j][0]["lp_gather"] for j in idx]
|
||||
log_N = math.log(N)
|
||||
# BMA per direction: logsumexp_n lp_samples[n, k] - log(N).
|
||||
def _bma(samples: list[list[float]]) -> list[float]:
|
||||
out = []
|
||||
for k in range(K):
|
||||
vals = [samples[n][k] for n in range(N)]
|
||||
m = max(vals)
|
||||
out.append(m + math.log(sum(math.exp(v - m) for v in vals)) - log_N)
|
||||
return out
|
||||
lp_f = _bma(lp_f_samples)
|
||||
lp_r = _bma(lp_r_samples)
|
||||
score = [(lp_f[k] + lp_r[k]) / 2.0 for k in range(K)]
|
||||
|
||||
m = max(score)
|
||||
@@ -432,25 +504,27 @@ def guided_rollout_forced_choice(
|
||||
order_sorted = sorted(range(K), key=lambda k: -score[k])
|
||||
top1 = foundations[order_sorted[0]]
|
||||
margin = score[order_sorted[0]] - score[order_sorted[1]]
|
||||
# Average pmass_format across framings: coherence canary independent
|
||||
# of WHICH foundation is picked. Sum prob mass over the K answer
|
||||
# tokens at the JSON slot; drops when model emits non-foundation
|
||||
# tokens (gibberish, refusal, format collapse).
|
||||
pm_f = slots_fwd[i][0]["pmass_format"]
|
||||
pm_r = slots_rev[i][0]["pmass_format"]
|
||||
# Average pmass_format across N samples per direction, then across
|
||||
# fwd + rev framings.
|
||||
pm_f = sum(slots_fwd[j][0]["pmass_format"] for j in idx) / N
|
||||
pm_r = sum(slots_rev[j][0]["pmass_format"] for j in idx) / N
|
||||
pm = 0.5 * (pm_f + pm_r)
|
||||
results.append(ForcedChoiceResult(
|
||||
user_prompt=user_prompts[i],
|
||||
gen_text=gen_fwd,
|
||||
gen_text_rev=gen_rev,
|
||||
gen_text=gens_fwd,
|
||||
gen_text_rev=gens_rev,
|
||||
lp_fwd={foundations[k]: lp_f[k] for k in range(K)},
|
||||
lp_rev={foundations[k]: lp_r[k] for k in range(K)},
|
||||
lp_fwd_samples=lp_f_samples,
|
||||
lp_rev_samples=lp_r_samples,
|
||||
score={foundations[k]: score[k] for k in range(K)},
|
||||
p=p,
|
||||
top1=top1,
|
||||
margin=float(margin),
|
||||
think_tokens=n_fwd,
|
||||
emitted_close=close_fwd,
|
||||
think_tokens=n_fwd_list,
|
||||
think_tokens_rev=n_rev_list,
|
||||
emitted_close=close_fwd_list,
|
||||
emitted_close_rev=close_rev_list,
|
||||
pmass_format=float(pm),
|
||||
))
|
||||
|
||||
|
||||
Reference in new issue
Block a user