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