File size: 6,541 Bytes
4afe981 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 | #!/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
|