Download code/decoding.py from Duoia/duovlm-40m-v1: direct link, hf CLI and curl.
- Browser
- Download file 6.54 kB
-
https://huggingface.co/Duoia/duovlm-40m-v1/resolve/main/code/decoding.py
- Command line
-
hf download hf://Duoia/duovlm-40m-v1/code/decoding.py
-
curl -L -o decoding.py https://huggingface.co/Duoia/duovlm-40m-v1/resolve/main/code/decoding.py
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 | |
| 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 | |