duovlm-40m-v1 / code /decoding.py
Duoia's picture
DuoVLM-40M v1: from-scratch 40M vision-language model (frozen CLIP + MiniPile-pretrained LM)
4afe981 verified
Raw History Blame Contribute Delete
6.54 kB
#!/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