一份真正能跑在浏览器里的 DeepSeek-R1 推理代码,核心逻辑全在一个 Web Worker 文件里。这篇文章不聊虚的,直接按代码块逐段拆解,每一行都讲清楚"它为什么这么写"。
一、导入与模型下载地址
javascript
import {
AutoTokenizer, // 分词器:把文本切成 token id
AutoModelForCausalLM, // 大模型:负责真正的推理
TextStreamer, // 流式输出:边生成边往外吐字
InterruptableStoppingCriteria, // 可中断的停止条件
env, // 环境配置:改下载地址等
} from "@huggingface/transformers";
// 国内访问 HuggingFace 被墙,走 hf-mirror 镜像下载模型
env.remoteHost = "https://hf-mirror.com";
逐行看:
AutoTokenizer:分词器。大模型不吃字符串,只认 token id(数字)。它负责「文本 → id」和「id → 文本」的双向转换。AutoModelForCausalLM:因果语言模型(CausalLM),也就是"根据前面的 token 预测下一个 token"的模型。这是推理的核心。TextStreamer:流式输出器。模型是一个 token 一个 token 往外蹦的,它负责每蹦一个就触发一次回调,实现打字机效果。InterruptableStoppingCriteria:停止条件。生成是循环的,每轮预测一个 token,这个对象用来"注入一个可被外部置位的开关",用户点了停止,循环就断。env.remoteHost:因为 HuggingFace 在国内访问慢,改成镜像站hf-mirror.com。注意这一步必须在任何下载之前执行,否则模型还是从原站拉。
二、核心中的核心:TextGenerationPipeline 单例类
kotlin
class TextGenerationPipeline {
static model_id = "onnx-community/DeepSeek-R1-Distill-Qwen-1.5B-ONNX";
static async getInstance(progress_callback = null) {
this.tokenizer ??= AutoTokenizer.from_pretrained(this.model_id, {
progress_callback,
});
this.model ??= AutoModelForCausalLM.from_pretrained(this.model_id, {
dtype: "q4f16", // 4-bit 量化
device: "webgpu", // 用 GPU 推理
progress_callback,
});
return Promise.all([this.tokenizer, this.model]);
}
}
1. 为什么要用单例?
from_pretrained 的开销非常大:要下载几百 MB 的权重文件 ,还要做 WebGPU shader 编译。如果每次对话都重新 new 一次、重新加载一次,用户等下载就等疯了。
所以这个类用单例模式:模型只初始化一次,之后一直复用。
2. static 方法里的 this 是谁?
这是新手最容易懵的点。getInstance 是 static(静态)方法,静态方法里的 this 指向类本身 ,也就是 TextGenerationPipeline。
所以:
kotlin
this.tokenizer ??= ... // 等价于 TextGenerationPipeline.tokenizer ??= ...
缓存是挂在类 上的,而不是某个实例上,天然全局唯一。这就是单例的实现方式------不靠 new,靠静态属性存一份。
3. ??= 到底做了什么(重点)
kotlin
this.tokenizer ??= AutoTokenizer.from_pretrained(...);
??= 是空值合并赋值运算符,语义是:
当左侧变量是
null或undefined时,才执行右侧并赋值;其他情况(包括false、0、'')一律跳过。
等价展开:
kotlin
if (this.tokenizer === null || this.tokenizer === undefined) {
this.tokenizer = AutoTokenizer.from_pretrained(...);
}
执行流程:
- 第一次调用 :
this.tokenizer是undefined→ 满足条件 → 执行from_pretrained,把返回的 Promise 存进this.tokenizer。 - 第二次及以后 :
this.tokenizer已经是一个 Promise 对象(不是 null/undefined)→ 跳过,直接复用,不会重复下载。
4. 缓存的是 Promise,不是结果(最精妙的一层)
注意:AutoTokenizer.from_pretrained() 是同步返回一个 Promise 对象 的,真正下载是异步进行。所以 ??= 存进去的是那个 Promise 本身,而不是"已经下载完的结果"。
这带来两个好处:
好处一:并发去重。 如果两个请求几乎同时进来,第一个请求已经把 Promise 存进 this.tokenizer 了;第二个进来时 ??= 发现已经有值,就直接复用同一个还没完成的 Promise。结果:两个请求共享同一次下载,不会下载两遍。
好处二:后续调用零等待。 Promise 一旦 resolve,之后再 await 它,几乎瞬间返回。
5. dtype 和 device 两个参数
arduino
{
dtype: "q4f16", // 量化到 4-bit,权重体积缩小约 1/4,显存/内存占用大降
device: "webgpu", // 让推理跑在 GPU 上,而不是 CPU
}
q4f16:4-bit 量化,牺牲一点精度换速度,是端侧/浏览器跑模型的标配。webgpu:WebGPU 是浏览器新一代 GPU 接口,比 WebGL 更适合做 AI 计算。
6. 最后 Promise.all
kotlin
return Promise.all([this.tokenizer, this.model]);
分词器和模型互不依赖 ,所以用 Promise.all 并行等两个下载任务都完成,而不是先等一个再等另一个(串行会慢一倍)。
三、中断机制与 KV 缓存
csharp
// 可被外部中断的停止条件实例
const stopping_criteria = new InterruptableStoppingCriteria();
// 缓存上一次的注意力计算结果
let past_key_values_cache = null;
1. InterruptableStoppingCriteria
生成是一个循环:每轮预测一个 token,然后判断"要不要停"。这个对象提供了一个开关,用户点"停止"时,外部调用 .interrupt() 把开关置位,模型下一轮检测到这个标志就停下。
arduino
case "interrupt":
stopping_criteria.interrupt(); // 置位 → 模型检测到后停止
break;
2. past_key_values_cache(KV 缓存)
每次对话,模型都要做大量 KV 注意力计算 ,非常耗算力。如果上一轮已经算过一部分,下一轮其实可以跳过重复计算,直接复用之前的缓存。
csharp
case "reset":
past_key_values_cache = null; // 清空缓存,重新开始
stopping_criteria.reset();
break;
多轮对话时缓存能大幅提速;"重置"对话时则要清空,否则会带上旧上下文。
四、generate():生成函数(最长的部分)
1. 拿到模型 + 套模板
javascript
async function generate(messages) {
const [tokenizer, model] = await TextGenerationPipeline.getInstance();
const inputs = tokenizer.apply_chat_template(messages, {
add_generation_prompt: true, // 自动追加 <|im_start|>assistant\n
return_dict: true,
});
apply_chat_template:把对话数组套上 DeepSeek 训练时用的特定模板 ,拼成模型认识的字符串。格式类似<|im_start|>user 你好 <|im_end|>。add_generation_prompt: true:在末尾自动追加<|im_start|>assistant\n,表示"接下来该助手续写了",引导模型开始生成回答。
2. 区分「思考」和「回答」
DeepSeek-R1 的生成分两部分:先思考(<think>...</think>),再正式回答。
csharp
const [START_THINKING_TOKEN_ID, END_THINKING_TOKEN_ID] = tokenizer.encode(
"<think></think>",
{ add_special_tokens: false },
);
let state = "thinking"; // 当前是思考阶段还是回答阶段
let startTime; // 计时起点
let numTokens = 0; // 已生成 token 总数
let tps; // 每秒生成 token 数
3. 统计生成速度 TPS
ini
const token_callback_function = (tokens) => {
startTime ??= performance.now(); // 第一次才记录开始时间
if (numTokens++ > 0) {
tps = (numTokens / (performance.now() - startTime)) * 1000;
}
if (tokens[0] == END_THINKING_TOKEN_ID) {
state = "answering"; // 遇到 </think> 结束标记,切换到回答阶段
}
};
逐行拆:
startTime ??= performance.now():又是一个??=。只有第一次(startTime是undefined)才记录时间戳,之后保持不变------避免反复重置计时起点,保证 TPS 算得准。tps = (numTokens / (performance.now() - startTime)) * 1000:token 数除以毫秒数,再乘 1000,得到「每秒生成多少个 token」。这是衡量推理速度的核心指标。tokens[0] == END_THINKING_TOKEN_ID:检测到</think>结束标记,说明思考结束,进入正式回答。
4. 把状态回传主线程
javascript
const callback_function = (output) => {
self.postMessage({
status: "update",
output, // 当前生成的文本
tps, // 实时速度
numTokens, // 已生成 token 数
state, // thinking / answering
});
};
Worker 里不能直接操作 DOM,所以通过 self.postMessage 把数据发回主线程,由主线程去更新界面。
5. 配置流式输出器
arduino
const streamer = new TextStreamer(tokenizer, {
skip_prompt: true, // 跳过输入的 prompt,只输出新生成的
skip_special_tokens: true, // 跳过 <|im_end|> 这类特殊标记
callback_function, // 每生成一段触发,回传主线程
token_callback_function, // 每个 token 触发,用于统计
});
6. 真正的生成调用
rust
self.postMessage({ status: "start" });
const { past_key_values, sequences } = await model.generate({
...inputs, // 上面套好模板的 token id
do_sample: false, // 贪心解码:取概率最大的 token,结果稳定
max_new_tokens: 2048, // 最多生成 2048 个 token
streamer, // 边生成边回调
stopping_criteria, // 可被外部中断
return_dict_in_generate: true, // 返回完整字典(含 past_key_values)
});
past_key_values_cache = past_key_values; // 保存 KV 缓存供下轮复用
几个关键点:
do_sample: false:贪心策略,每次取概率最高的 token,输出确定性强。适合需要稳定回答的场景。return_dict_in_generate: true:让generate除了返回sequences(token 序列)外,还返回past_key_values(KV 缓存),这样多轮对话能复用。past_key_values_cache = past_key_values:把这次算出的 KV 缓存存下来,供下一轮跳过重复计算。
7. 解码最终结果
php
const decoded = tokenizer.batch_decode(sequences, {
skip_special_tokens: true,
});
self.postMessage({
status: "complete",
output: decoded,
});
batch_decode:把 token id 批量转回文本。skip_special_tokens 过滤掉特殊标记,得到干净的回答。
五、check():检测 WebGPU 是否可用
javascript
async function check() {
try {
const adapter = await navigator.gpu.requestAdapter();
if (!adapter) {
throw new Error("WebGPU is not supported (no adapter found)");
}
} catch (e) {
self.postMessage({ status: "error", data: e.toString() });
}
}
navigator.gpu.requestAdapter():请求 GPU 适配器(adapter 是 GPU 的抽象,后续所有 WebGPU 计算都通过它执行)。- 如果返回
null(拿不到 adapter),说明浏览器/设备不支持 WebGPU,抛出错误并回传主线程。
六、load():加载 + 预热
php
async function load() {
self.postMessage({ status: "loading", data: "Loading model..." });
const [tokenizer, model] = await TextGenerationPipeline.getInstance((x) => {
self.postMessage(x); // 把下载进度回传主线程
});
self.postMessage({
status: "loading",
data: "Compiling shaders and warming up model...",
});
// 用 "a" 跑一次,触发 shader 编译和模型预热
const inputs = tokenizer("a");
await model.generate({ ...inputs, max_new_tokens: 1 });
self.postMessage({ status: "ready" });
}
关键一步是"预热" :WebGPU 的 shader 编译很耗时,如果等用户第一次提问时才编译,会明显卡顿。所以这里先用一个极小的输入(字符 "a")跑一遍,提前支付编译成本,用户真正使用时就是秒出。
七、消息分发:Worker 的入口
typescript
self.addEventListener("message", async (e) => {
const { type, data } = e.data;
switch (type) {
case "check": // 检查 WebGPU 支持
check();
break;
case "load": // 加载模型(懒加载 + 进度 + 预热)
load();
break;
case "generate": // 生成文本
stopping_criteria.reset(); // 每次生成前先复位中断开关
generate(data);
break;
case "interrupt": // 中断生成
stopping_criteria.interrupt();
break;
case "reset": // 重置对话
past_key_values_cache = null; // 清空 KV 缓存
stopping_criteria.reset();
break;
}
});
主线程通过 worker.postMessage({ type: "generate", data: messages }) 发指令,Worker 根据 type 分发到对应函数。整个通信就是「主线程发指令 → Worker 干活 → postMessage 回报结果」的闭环。
八、整份代码的"一句话"总结
| 代码块 | 解决了什么问题 |
|---|---|
env.remoteHost |
国内下载模型走镜像 |
??= + 静态属性 |
单例懒加载,模型只下载一次 |
| 缓存 Promise | 并发去重 + 后续零开销 |
Promise.all |
分词器和模型并行下载 |
past_key_values_cache |
多轮对话复用注意力计算 |
InterruptableStoppingCriteria |
让生成能被用户中途打断 |
TextStreamer + 回调 |
流式输出 + 实时统计 TPS |
tokenizer("a") 预热 |
提前编译 shader,避免首次卡顿 |
整份代码最值得记住的,还是那一行:
kotlin
this.tokenizer ??= AutoTokenizer.from_pretrained(this.model_id, { progress_callback });
一行 ??=,同时干了「懒加载 + 单例缓存 + 并发去重」三件事。下次你在任何"初始化开销大、只需要做一次"的场景,直接照搬这个套路就行。