#!/usr/bin/env python3 """安全解码:n-gram 禁复读 + 重复惩罚 + 循环截断 + 信息量守卫(服务端统一走这里) 实测背景(STAGE1_ROADMAP §4.5):38M 容量的模型在 greedy 下 **eos 几乎不可能是 argmax** (512 个位置里只有最后一位该吐 eos),所以输出一旦超过 1~2 句就不会自停,会一直复读到 max_new。 数据侧两条假设(语料重复 / 文档边界太密)已实测证伪,所以这个问题**只能在解码侧解决**。 四层防线(每层都对应一种实测到的失败形态): 1. no_repeat_ngram_size = 3 禁任何逐字 3-gram 重复(实测退化 5/5→0/5、uniq 0.088→0.40) 2. repetition_penalty = 1.1~1.2 压低已出现 token 的分数(叠温度 0.7 时 uniq 到 0.75) 3. 循环截断 尾部出现周期 ≤12 的循环 → 停并剪掉重复段(不顶上限) 4. 信息量守卫 尾部窗口 2-gram 去重率过低 → 停(挡"改词绕过 ngram 禁复读"那种 "small, and very small area. The ocean is a small, large area" 式打转) **服务侧不要依赖模型自停**:按任务硬性设 max_new(VQA 12 / 描述 30~40), 上面的防线只是保证在长输出时不会退化成噪声。 实测(2026-09-27,S3 成品模型):greedy / 采样 / 对抗输入下都一句话就自停 —— 因为 Stage 3 的短答案 SFT 教会了 eos。防线是给 OOD 输入和长文本场景兜底的。 """ from __future__ import annotations import torch EOS = 1 def loop_period(hist: list[int], max_period: int = 12) -> int: """尾部是否已进入周期循环。返回周期 p(0 = 没循环,可以继续)。 p<=2 要求连续 3 次("a a a" / "ab ab ab" 才算),p>=3 要求 2 次(整句重复一遍就停)。 """ n = len(hist) for p in range(1, max_period + 1): reps = 3 if p <= 2 else 2 need = p * reps if n < need: continue tail = hist[-need:] if all(tail[i] == tail[i % p] for i in range(need)): return p return 0 def truncate_loop(hist: list[int], p: int) -> list[int]: """把尾部重复的部分剪掉,保留第一次出现。""" if p <= 0: return hist reps = 3 if p <= 2 else 2 need = p * reps return hist[: len(hist) - need + p] def degeneracy_index(hist: list[int], window: int = 24, min_distinct2: float = 0.6) -> int: """返回第一次判为"原地打转"的位置(0 = 没打转)。 用尾部 window 个 token 的 **2-gram 去重率**衡量。正常英文句子这个值很高(>0.85), 而"small, small, and very small area"这种改词打转会掉到 0.6 以下。 """ for i in range(window, len(hist) + 1): tail = hist[i - window:i] bg = list(zip(tail, tail[1:])) if len(set(bg)) / len(bg) < min_distinct2: return i return 0 @torch.no_grad() def decode_safe(model, ids: torch.Tensor, feats, max_new: int = 12, ngram: int = 3, rep_penalty: float = 1.0, temperature: float = 0.0, top_k: int = 0, loop_break: bool = True, deg_min_distinct2: float = 0.6, deg_window: int = 24, eos_id: int = EOS, seed: int | None = None): """返回 (每行 token 列表, 统计字典)。ids: [B,T];feats: 特征或 None。""" B = ids.size(0) dev = ids.device lens = torch.full((B,), ids.size(1) - 1, dtype=torch.long, device=dev) out: list[list[int]] = [[] for _ in range(B)] done = [False] * B stats = {"steps": 0, "loop_rows": 0, "deg_rows": 0, "eos_rows": 0, "min_period": 0, "cut_tokens": 0} gen = torch.Generator(device="cpu") if seed is not None: gen.manual_seed(seed) cur = ids for _ in range(max_new): logits = model(cur, feats)[0][torch.arange(B, device=dev), lens].float() for b in range(B): if done[b]: continue hist = out[b] row = logits[b] # 1) n-gram 禁复读 if ngram > 0 and len(hist) >= ngram: prefix = tuple(hist[-(ngram - 1):]) banned = {hist[i + ngram - 1] for i in range(len(hist) - ngram + 1) if tuple(hist[i:i + ngram - 1]) == prefix} if banned: row[list(banned)] = -float("inf") # 2) 重复惩罚 if rep_penalty > 1.0 and hist: idx = torch.tensor(sorted(set(hist)), device=dev, dtype=torch.long) v = row[idx] row[idx] = torch.where(v > 0, v / rep_penalty, v * rep_penalty) # 3) 采样 or 贪心 if temperature > 0: lg = row / temperature if top_k > 0: kth = torch.topk(lg, min(top_k, lg.numel())).values[-1] lg = lg.masked_fill(lg < kth, -float("inf")) nxt = int(torch.multinomial(torch.softmax(lg.cpu(), -1), 1, generator=gen).to(dev).item()) else: nxt = int(torch.argmax(row).item()) hist.append(nxt) if nxt == eos_id: hist.pop() # eos 不进正文 done[b] = True stats["eos_rows"] += 1 continue if loop_break: # 4a) 周期循环截断 p = loop_period(hist) if p: kept = truncate_loop(hist, p) stats["loop_rows"] += 1 stats["cut_tokens"] += len(hist) - len(kept) stats["min_period"] = p out[b] = kept done[b] = True continue # 4b) 信息量守卫(挡改词打转) if deg_min_distinct2 > 0: dj = degeneracy_index(hist, deg_window, deg_min_distinct2) if dj: stats["deg_rows"] += 1 stats["cut_tokens"] += len(hist) - dj out[b] = hist[:dj] done[b] = True continue stats["steps"] += 1 cur = torch.cat( [cur, torch.tensor([[out[b][-1] if out[b] else eos_id] for b in range(B)], device=dev, dtype=cur.dtype)], dim=1) lens = lens + 1 if all(done): break return out, stats