mirror of
https://github.com/wassname/moral-maps.git
synced 2026-09-10 12:14:54 +08:00
suppress end-of-answer tokens in forced rollout so reads are comparable
Forced-choice scored every sample at whatever state it reached: base
self-closes </think> at long budgets -> junk-cache case (c) -> pmass 0.0,
while steered keeps thinking -> case (b) forced read -> pmass ~1.0. pmass
then measured self-close rate, not coherence, making steering look more
coherent. Suppress {eos, think_end_id} for the whole think budget so no
sample self-closes; all land in case (b) and read at the same forced slot.
Model-agnostic: think_end_id is </think> on reasoning models, falls back
to eos elsewhere. Smoke pmass=0.985.
Co-Authored-By: Claudypoo <288921227+claudypoo@users.noreply.github.com>
This commit is contained in:
+10
-3
@@ -145,16 +145,23 @@ def _rollout_natural_or_forced(
|
||||
if think_end_id in (None, getattr(tok, "unk_token_id", None)):
|
||||
think_end_id = tok.eos_token_id
|
||||
|
||||
# Suppress end-of-answer tokens for the whole budget so every sample ends in the
|
||||
# same mid-think state, read at the same forced slot. Else base self-closes into
|
||||
# the junk-cache case (c -> NaN) while steered keeps thinking (case b), making
|
||||
# pmass measure "did you self-close" (a steering confound) not coherence. eos is
|
||||
# the universal end signal; think_end_id adds </think> on reasoning models and
|
||||
# falls back to eos elsewhere, so this stays model-agnostic.
|
||||
suppress_end = [t for t in {tok.eos_token_id, think_end_id} if t is not None]
|
||||
|
||||
enc = tok(chats, return_tensors="pt", padding=True).to(device)
|
||||
prompt_len = enc.input_ids.shape[1]
|
||||
do_sample = temperature > 0.0
|
||||
gen_kwargs = dict(
|
||||
max_new_tokens=max_think_tokens,
|
||||
# Force full budget so all samples have identical cache length →
|
||||
# batched suffix forward without per-sample rewinding. Garbage tokens
|
||||
# emitted past natural EOS pollute the cache only for case-(c) samples,
|
||||
# which we NaN downstream anyway.
|
||||
# batched suffix forward without per-sample rewinding.
|
||||
min_new_tokens=max_think_tokens,
|
||||
suppress_tokens=suppress_end,
|
||||
pad_token_id=pad_id,
|
||||
return_dict_in_generate=True,
|
||||
output_scores=True,
|
||||
|
||||
Reference in New Issue
Block a user