File size: 10,716 Bytes
5e6d9f5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
"""推理:压缩 KV cache 增量解码 + 基于 MTP 头的自投机解码。

    python generate.py --ckpt checkpoints/best.pt --prompt "在深夜," --spec

自投机解码(self-speculative decoding)的思路:
MTP 头本来就是训来预测「再下一个」token 的,那它天然就是一个 draft model,
而且和主干共享 KV cache,几乎不额外花钱。

每一轮:
  1. 主干一次前向同时吃 [上一个真 token, 上一轮的草稿 token]
  2. 第 1 个位置的输出告诉我们「真正的下一个 token 应该是什么」→ 用来验草稿
  3. 草稿被接受 → 这一轮白赚一个 token(第 2 个位置的输出直接就是再下一个)
     草稿被拒绝 → 回滚一格,损失仅一次多余的 KV 写入
验收用的是标准 speculative sampling 的接受-重采样规则,
输出分布和不开投机时**严格一致**,不是近似。
"""

import argparse
import sys
import time
from typing import Optional, Tuple

import torch
import torch.nn.functional as F

from config import ModelConfig
from dataio.tokenizer import load_tokenizer
from model import DeepSeekFlashMini


# --------------------------------------------------------------- 采样工具
def filter_logits(logits: torch.Tensor, temperature: float, top_k: int, top_p: float):
    """返回过滤+归一化后的概率分布 (1, V)。"""
    if temperature <= 0:                      # 贪心:退化成 one-hot
        probs = torch.zeros_like(logits)
        probs.scatter_(-1, logits.argmax(-1, keepdim=True), 1.0)
        return probs
    logits = logits / temperature
    if top_k and top_k > 0:
        k = min(top_k, logits.size(-1))
        thresh = logits.topk(k, dim=-1)[0][..., -1:]
        logits = logits.masked_fill(logits < thresh, float("-inf"))
    if top_p and 0 < top_p < 1.0:
        sorted_logits, sorted_idx = logits.sort(dim=-1, descending=True)
        cum = sorted_logits.softmax(-1).cumsum(-1)
        remove = cum - sorted_logits.softmax(-1) > top_p
        sorted_logits = sorted_logits.masked_fill(remove, float("-inf"))
        logits = torch.empty_like(logits).scatter_(-1, sorted_idx, sorted_logits)
    return logits.softmax(-1)


def sample_from(probs: torch.Tensor) -> int:
    return int(torch.multinomial(probs, num_samples=1).item())


# --------------------------------------------------------------- 生成器
class Generator:
    def __init__(self, model: DeepSeekFlashMini, tokenizer, device, max_seq_len=None,
                 attn_impl: str = "naive"):
        self.model = model.eval()
        self.tok = tokenizer
        self.device = device
        self.max_seq_len = max_seq_len or model.cfg.max_seq_len
        model.set_attn_impl(attn_impl)
        model.setup_cache(1, self.max_seq_len, device, torch.float32)

    def _t(self, ids):
        return torch.tensor([ids], dtype=torch.long, device=self.device)

    @torch.no_grad()
    def generate(self, prompt: str, max_new_tokens: int = 200, temperature: float = 0.8,
                 top_k: int = 50, top_p: float = 0.95, speculative: bool = False,
                 stream: bool = True, stop_on_eos: bool = True) -> Tuple[str, dict]:
        ids = self.tok.encode(prompt, bos=True)
        keep = max(1, self.max_seq_len - max_new_tokens - 2)
        ids = ids[-keep:] or [self.tok.bos_id]
        self.eos = self.tok.eos_id if stop_on_eos else -1
        gen = self._spec_loop if speculative else self._plain_loop
        t0 = time.time()
        out_ids, meta = gen(ids, max_new_tokens, temperature, top_k, top_p, stream)
        out_ids = out_ids[:max_new_tokens]      # 投机解码一轮可能吐 2 个,这里截齐
        meta["seconds"] = time.time() - t0
        meta["tokens"] = len(out_ids)
        meta["tok_per_s"] = len(out_ids) / max(meta["seconds"], 1e-9)
        meta["ids"] = out_ids
        meta["prompt_ids"] = ids
        text = self.tok.decode(ids + out_ids)
        return text, meta

    # -------------------------------------------------- 普通增量解码
    def _plain_loop(self, ids, max_new, temperature, top_k, top_p, stream):
        m = self.model
        pos = 0
        pending = list(ids)
        out = []
        printed = 0
        while len(out) < max_new:
            _, logits, _ = m.forward_trunk(self._t(pending), start_pos=pos)
            pos += len(pending)
            probs = filter_logits(logits[:, -1].float(), temperature, top_k, top_p)
            t = sample_from(probs)
            if t == self.eos:
                break
            out.append(t)
            pending = [t]
            printed = self._stream(out, printed, stream)
            if pos + 2 >= self.max_seq_len:
                break
        if stream:
            print()
        return out, {"mode": "plain", "accepted": 0, "rounds": len(out)}

    # -------------------------------------------------- MTP 自投机解码
    def _spec_loop(self, ids, max_new, temperature, top_k, top_p, stream):
        m = self.model
        assert m.cfg.n_mtp > 0, "该模型没有 MTP 头,无法投机解码"
        pos = 0
        pending = list(ids)
        out = []
        printed = 0
        draft: Optional[int] = None
        draft_q: Optional[torch.Tensor] = None
        rounds = accepted = 0

        while len(out) < max_new:
            T = len(pending)
            h, logits, _ = m.forward_trunk(self._t(pending), start_pos=pos)
            rounds += 1
            eos_hit = False

            if draft is None:
                # 首轮(prefill):只出一个真 token
                p = filter_logits(logits[:, -1].float(), temperature, top_k, top_p)
                t_new = sample_from(p)
                pos += T
                if t_new == self.eos:
                    break
                out.append(t_new)
                next_tokens = pending[1:] + [t_new]
                h_chunk, chunk_start = h, pos - T
            else:
                # pending = [上一个真 token, 草稿],用位置 T-2 的输出来验草稿
                p = filter_logits(logits[:, T - 2].float(), temperature, top_k, top_p)
                q = draft_q
                ratio = (p[0, draft] / q[0, draft].clamp_min(1e-10)).clamp(max=1.0)
                if torch.rand(1, device=ratio.device) < ratio:
                    # ---- 接受草稿:这一轮吐 2 个 token ----
                    accepted += 1
                    out.append(draft)
                    p2 = filter_logits(logits[:, T - 1].float(), temperature, top_k, top_p)
                    t_new = sample_from(p2)
                    pos += T
                    if draft == self.eos or t_new == self.eos:
                        eos_hit = True
                    else:
                        out.append(t_new)
                    next_tokens = pending[1:] + [t_new]
                    h_chunk, chunk_start = h, pos - T
                else:
                    # ---- 拒绝:从残差分布 max(0, p-q) 重采样,保证分布无偏 ----
                    resid = (p - q).clamp_min(0)
                    s = resid.sum()
                    resid = p if s < 1e-9 else resid / s     # p==q 的退化情形
                    t_new = sample_from(resid)
                    pos += T - 1                     # 回滚草稿占的那一格
                    if t_new == self.eos:
                        eos_hit = True
                    else:
                        out.append(t_new)
                    next_tokens = pending[1:T - 1] + [t_new]
                    h_chunk, chunk_start = h[:, :T - 1], pos - (T - 1)

            printed = self._stream(out, printed, stream)
            if eos_hit or pos + 3 >= self.max_seq_len or len(out) >= max_new:
                break

            # ---- 用 MTP 头造下一轮的草稿(顺带把 MTP 的 cache 补齐)----
            _, mtp_logits, _ = m.mtp_forward(
                h_chunk, self._t(next_tokens), start_pos=chunk_start, depth=0)
            draft_q = filter_logits(mtp_logits[:, -1].float(), temperature, top_k, top_p)
            draft = sample_from(draft_q)
            pending = [out[-1], draft]

        if stream:
            print()
        rate = accepted / max(rounds - 1, 1)
        return out, {"mode": "speculative", "accepted": accepted, "rounds": rounds,
                     "accept_rate": rate}

    def _stream(self, out, printed, stream):
        if not stream:
            return printed
        text = self.tok.decode(out)
        if len(text) > printed:
            sys.stdout.write(text[printed:])
            sys.stdout.flush()
            printed = len(text)
        return printed


# --------------------------------------------------------------- CLI
def load_model(ckpt_path, device):
    ck = torch.load(ckpt_path, map_location=device, weights_only=False)
    cfg = ModelConfig.from_dict(ck["config"])
    model = DeepSeekFlashMini(cfg).to(device)
    model.load_state_dict(ck["model"])
    return model, cfg, ck


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--ckpt", default="checkpoints/best.pt")
    ap.add_argument("--tokenizer", default="")
    ap.add_argument("--prompt", default="在深夜,")
    ap.add_argument("--max-new-tokens", type=int, default=160)
    ap.add_argument("--temperature", type=float, default=0.8)
    ap.add_argument("--top-k", type=int, default=50)
    ap.add_argument("--top-p", type=float, default=0.95)
    ap.add_argument("--spec", action="store_true", help="开启 MTP 自投机解码")
    ap.add_argument("--attn", default="naive", choices=["naive", "absorb"])
    ap.add_argument("--device", default="auto")
    ap.add_argument("--seed", type=int, default=0)
    args = ap.parse_args()

    torch.manual_seed(args.seed)
    device = (torch.device("cuda") if torch.cuda.is_available() else
              torch.device("mps") if torch.backends.mps.is_available() else
              torch.device("cpu")) if args.device == "auto" else torch.device(args.device)

    model, cfg, ck = load_model(args.ckpt, device)
    tok_path = args.tokenizer or f"{ck.get('data_dir', 'data')}/tokenizer.json"
    tok = load_tokenizer(tok_path)

    gen = Generator(model, tok, device, attn_impl=args.attn)
    print(f"--- {'MTP 投机解码' if args.spec else '普通解码'} | {args.attn} 注意力 ---")
    print(args.prompt, end="")
    _, meta = gen.generate(args.prompt, args.max_new_tokens, args.temperature,
                           args.top_k, args.top_p, speculative=args.spec)
    line = (f"[{meta['tokens']} tokens / {meta['seconds']:.2f}s = "
            f"{meta['tok_per_s']:.1f} tok/s")
    if args.spec:
        line += f" | 草稿接受率 {meta['accept_rate']*100:.0f}%"
    print(line + "]")


if __name__ == "__main__":
    main()