一、项目整体架构
这是一个典型的双线程架构:
- 主线程(React) :负责 UI 渲染、用户交互、状态管理,不碰任何模型计算
- Worker 子线程:负责模型加载、WebGPU 推理、token 生成,不访问 DOM
两者通过 postMessage 双向通信,消息类型包括 check / load / generate / interrupt / reset。
技术栈
- 前端:React + Vite + TailwindCSS
- 推理引擎:
@huggingface/transformers.js - 模型:
DeepSeek-R1-Distill-Qwen-1.5B-ONNX(4bit 量化) - 加速:WebGPU 浏览器原生 GPU 计算
二、Worker 端:模型单例加载设计
2.1 为什么要用单例模式
大模型加载成本极高:下载权重、反序列化、上传 GPU 显存、编译着色器,每一步都耗时。如果每次对话都重新加载,用户体验会极差。
所以项目用了静态属性 + 空值合并赋值 ??= 实现单例:
js
运行
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,
});
// 模型:4bit 量化 + WebGPU 设备
this.model ??= AutoModelForCausalLM.from_pretrained(this.model_id, {
dtype: "q4f16",
device: "webgpu",
progress_callback,
});
return Promise.all([this.tokenizer, this.model]);
}
}
逐点解析:
static静态属性 :tokenizer和model挂载在类本身,而非实例上,全局唯一。??=空值合并赋值 :只有变量为null/undefined时才执行右侧赋值。第一次调用触发下载加载,后续直接复用内存中已就绪的对象。dtype: "q4f16":4bit 量化权重,把 1.5B 参数模型压缩到约 1GB 显存以内,浏览器才能跑得动。device: "webgpu":推理运算交给浏览器 GPU,纯 CPU 跑 1.5B 模型几乎不可用。Promise.all:分词器和模型并行下载加载,减少总等待时间。
💡 设计亮点:单例 + 懒加载。页面打开不加载模型,用户点击 "Load model" 才开始下载;一旦加载完成,常驻 Worker 内存,后续对话零延迟启动。
2.2 国内镜像代理配置
HuggingFace 国内访问受限,项目通过 Vite 中间件做同源代理,完美绕过跨域和网络问题:
js
运行
ini
import { env } from "@huggingface/transformers";
// 所有模型下载请求走 /hf-mirror 代理
env.remoteHost = "/hf-mirror";
原理很简单:请求从 huggingface.co 重写到本地 /hf-mirror 路径,Vite 开发服务器转发到 hf-mirror.com 镜像站。同源请求彻底解决 CORS 问题。
三、核心机制:可中断停止条件
3.1 原生停止条件的局限
transformers.js 原生的 StoppingCriteria 只能基于 token 内容判断停止:最大长度、命中 stop 词等。它不支持外部中途打断正在运行的生成循环。
如果用户点了 "停止" 按钮,原生实现要等下一个 stop 条件命中才会停,体验非常差。
3.2 InterruptableStoppingCriteria 解决方案
项目引入了 InterruptableStoppingCriteria,在原生停止条件基础上增加了外部中断能力:
js
运行
arduino
// Worker 全局唯一实例
const stopping_criteria = new InterruptableStoppingCriteria();
对外暴露三个核心方法:
表格
| 方法 | 作用 |
|---|---|
interrupt() |
外部调用,打上中断标记 |
reset() |
清除中断标记,准备新一轮生成 |
_check() |
生成循环每步调用,检测是否停止 |
消息分发中的使用:
js
运行
typescript
switch (type) {
case "generate":
stopping_criteria.reset(); // 生成前先清除旧标记
generate(data);
break;
case "interrupt":
stopping_criteria.interrupt(); // 用户点停止,打标记
break;
case "reset":
past_key_values_cache = null;
stopping_criteria.reset();
break;
}
工作原理:
model.generate 内部每生成一个 token,都会调用停止条件的 _check() 方法。一旦检测到 interrupted 标志为 true,立即终止生成循环,返回已产出的序列。
💡 关键点:
interrupt()不会强行杀死函数,只是设置 flag。真正停止靠生成循环内部每一步主动检测。这是协作式中断,而非抢占式终止。
四、KV Cache:对话提速的核心
4.1 为什么需要 KV 缓存
大语言模型每生成一个 token,都要计算整个上下文的注意力。如果每一轮对话都从头计算所有历史 token,速度会随对话长度线性下降。
KV Cache 的思路很简单:把已经算过的 key /value 张量存起来,下一轮直接复用,只计算新增的 token。
js
运行
csharp
// 全局缓存,Worker 内存常驻
let past_key_values_cache = null;
4.2 生成流程中的使用
js
运行
php
const { past_key_values, sequences } = await model.generate({
...inputs,
// past_key_values: past_key_values_cache, // 注释状态,可开启复用
do_sample: false,
max_new_tokens: 2048,
streamer,
stopping_criteria,
return_dict_in_generate: true,
});
// 保存本轮输出的 KV,供下一轮使用
past_key_values_cache = past_key_values;
执行流程:
- 传入
past_key_values_cache作为历史缓存 - 模型只对新输入的 token 做注意力计算
- 输出更新后的
past_key_values,包含全部历史 - 赋值回全局缓存,下一轮对话继续复用
4.3 Reset 清空缓存
js
运行
ini
case "reset":
past_key_values_cache = null;
stopping_criteria.reset();
break;
用户点击 "Reset" 开启新对话时,把缓存置为 null,丢弃全部历史上下文,相当于全新会话。
⚠️ 注意:KV Cache 是 GPU 张量对象,不能通过
postMessage传给主线程,只能留在 Worker 内部。这也是为什么推理必须放 Worker 的原因之一。
五、流式输出:TextStreamer 深度解析
5.1 为什么需要 TextStreamer
如果没有流式工具,只能等所有 token 生成完毕一次性返回,用户看到的是 "等半天然后整段文字蹦出来",体验很差。
TextStreamer 负责把模型逐 token 输出的数字 ID,实时解码成可读文本,实现打字机效果。
js
运行
arduino
const streamer = new TextStreamer(tokenizer, {
skip_prompt: true,
skip_special_tokens: true,
callback_function,
token_callback_function,
});
5.2 四个关键参数详解
1. skip_prompt: true
跳过输入的原始 prompt 文本。如果设为 false,会把用户的提问也回调输出,前端就会重复渲染用户消息。对话场景必须开 true。
2. skip_special_tokens: true
过滤掉 <|im_start|>、<|im_end|>、、 这类模型控制标记。用户看到的是干净的正文,不会夹杂特殊标签。
3. token_callback_function(原始 token 回调)
在解码成文本之前触发,拿到原始 token ID,用来做底层逻辑判断:
js
运行
ini
const token_callback_function = (tokens) => {
// 第一个 token 到来时记录开始时间
startTime ??= performance.now();
// 计算 TPS(跳过第一个预热 token)
if (numTokens++ > 0) {
tps = (numTokens / (performance.now() - startTime)) * 1000;
}
// 检测思考结束标记,切换状态
if (tokens[0] == END_THINKING_TOKEN_ID) {
state = "answering";
}
};
- TPS 统计 :
tokens per second,衡量 WebGPU 推理速度的核心指标 - 思考状态切换 :DeepSeek-R1 输出格式是
推理过程正式回答,检测到 `` 的 token ID 就切换状态
4. callback_function(文本回调)
解码完成后触发,拿到干净的文本字符串,直接推给前端渲染:
js
运行
lua
const callback_function = (output) => {
self.postMessage({
status: "update",
output,
tps,
numTokens,
state,
});
};
5.3 双回调的设计巧妙之处
这里有个很精妙的设计:
skip_special_tokens: true让文本输出不包含标签,用户看不到 ``- 但
token_callback_function依然能收到原始 token ID,代码可以识别思考结束的位置
页面干净和逻辑判断两者兼得。
六、思考 / 回答双状态机制
DeepSeek-R1 是推理模型,输出分为两阶段:先内部思考打草稿,再给出正式回答。前端需要区分渲染。
6.1 获取思考标签的 token ID
js
运行
php
const [START_THINKING_TOKEN_ID, END_THINKING_TOKEN_ID] = tokenizer.encode(
"",
{ add_special_tokens: false },
);
add_special_tokens: false:不要自动追加首尾特殊标记,只单纯把标签本身转成 ID- 解构得到两个数字:开始标记 ID 和结束标记 ID
6.2 状态切换逻辑
js
运行
ini
let state = "thinking"; // 初始状态:思考中
// ... 每个 token 回调时检测
if (tokens[0] == END_THINKING_TOKEN_ID) {
state = "answering";
}
初始状态是 thinking,当模型生成出 `` 这个 token 时,切换为 answering。
前端收到 state 字段后,可以做差异化渲染:
- 思考阶段:折叠框、灰色字体、"正在思考..." 提示
- 回答阶段:正常正文样式
七、React 主线程:通信与状态管理
7.1 Worker 初始化与消息监听
js
运行
javascript
useEffect(() => {
if (!worker.current) {
worker.current = new Worker(new URL("./worker.js", import.meta.url), {
type: "module",
});
worker.current.postMessage({ type: "check" });
}
const onMessageReceived = (e) => {
switch (e.data.status) {
case "loading": /* 更新加载文字 */ break;
case "progress": /* 更新下载进度条 */ break;
case "ready": /* 模型就绪,进入聊天页 */ break;
case "start": /* 开始生成,插入空 AI 消息 */ break;
case "update": /* 流式追加文本 */ break;
case "complete": /* 生成结束,统计耗时 */ break;
case "error": /* 显示错误 */ break;
}
};
worker.current.addEventListener("message", onMessageReceived);
return () => {
worker.current.removeEventListener("message", onMessageReceived);
};
}, []);
只执行一次的 useEffect:组件挂载时创建 Worker,注册消息监听,卸载时清理。
7.2 流式渲染的实现
js
运行
ini
case "start":
setIsRunning(true);
setNumTokens(0);
setTps(0);
timeRef.current = performance.now();
// 先插入一条空的助手消息
setMessages((prev) => [...prev, { role: "assistant", content: "" }]);
break;
case "update":
setNumTokens(e.data.numTokens);
setTps(e.data.tps);
// 追加到最后一条助手消息
setMessages((prev) => {
const next = [...prev];
const last = next[next.length - 1];
if (last && last.role === "assistant") {
next[next.length - 1] = {
...last,
content: last.content + e.data.output
};
}
return next;
});
break;
经典流式渲染套路:
- 生成开始时:往消息数组 push 一条空的 assistant 消息占位
- 每次流式更新:找到最后一条 assistant 消息,把新文本拼接到 content 后面
- React 自动重渲染,视觉上就是打字机逐字出现的效果
7.3 useEffect 自动触发生成
js
运行
ini
useEffect(() => {
// 没有用户消息:不做事
if (messages.filter((x) => x.role === "user").length === 0) {
return;
}
// 最后一条是 AI 回复:不重复请求
if (messages.at(-1).role === "assistant") {
return;
}
// 最后一条是用户消息:触发生成
worker.current.postMessage({ type: "generate", data: messages });
}, [messages]);
为什么不直接在 onEnter 里发消息?
因为 setMessages 是 React 异步更新。onEnter 执行完的那一刻,messages 还没变成最新值。
靠 useEffect 监听 messages 变化来触发,保证拿到的永远是更新完成的最新数组。
7.4 中断与重置
js
运行
lua
function onInterrupt() {
worker.current?.postMessage({ type: "interrupt" });
setIsRunning(false);
}
用户点击停止按钮,发送 interrupt 消息。Worker 端设置中断标记,生成循环检测到后立刻终止。
重置按钮则发送 reset 消息,清空 KV 缓存和消息列表,开启全新对话。
八、完整端到端数据流
梳理一下从用户输入到 AI 输出的完整链路:
- 用户回车发送 →
onEnter执行,setMessages追加 user 消息 - useEffect 检测 → messages 变化,最后一条是 user,发送
type: "generate" - Worker 收到消息 →
stopping_criteria.reset()清除旧标记,调用generate(messages) - 模型单例加载 → 第一次下载权重,后续直接复用
- 模板拼接 →
apply_chat_template把对话数组转成模型认识的 prompt 格式 - 流式生成 →
model.generate逐 token 输出,TextStreamer 双回调处理 - token 回调 → 统计 TPS、检测 `` 切换状态
- 文本回调 →
postMessage把文本 + 状态 + TPS 发回主线程 - 前端渲染 → 追加到最后一条 AI 消息,打字机效果
- 生成结束 → Worker 保存 KV 缓存,发送
complete;前端关闭 loading,展示统计
九、设计亮点与思考
9.1 架构层面
- 职责分离:UI 和推理彻底分离,Worker 负责重计算,主线程只负责渲染
- 单例懒加载:模型只加载一次,对话复用,性能最优
- 消息驱动:所有交互通过 type 消息分发,结构清晰易扩展
9.2 细节处理
- 协作式中断:不暴力杀线程,靠 flag 标记优雅停止
- 双回调流式:原始 token 做逻辑判断,解码文本做渲染,各司其职
- KV 缓存复用:多轮对话增量推理,速度不随对话长度线性下降
- 同源代理方案:巧妙解决 HuggingFace 国内访问 + CORS 双重问题
9.3 可以优化的点
- 源码中
past_key_values_cache传入被注释掉了,实际跑的时候每轮都会全量计算,速度会慢。开启 KV 缓存复用是重要的性能优化点。 - 目前思考和回答的文本是拼接在一起的,可以根据
state做分段存储,前端渲染折叠效果会更好。 - 缺少错误重试和降级方案,WebGPU 编译失败时的体验可以更友好。
十、完整核心代码
Worker 端(worker.js)
js
运行
ini
import {
AutoTokenizer,
AutoModelForCausalLM,
TextStreamer,
InterruptableStoppingCriteria,
env,
} from "@huggingface/transformers";
env.remoteHost = "/hf-mirror";
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",
device: "webgpu",
progress_callback,
});
return Promise.all([this.tokenizer, this.model]);
}
}
const stopping_criteria = new InterruptableStoppingCriteria();
let past_key_values_cache = null;
async function generate(messages) {
const [tokenizer, model] = await TextGenerationPipeline.getInstance();
const inputs = tokenizer.apply_chat_template(messages, {
add_generation_prompt: true,
return_dict: true,
});
const [START_THINKING_TOKEN_ID, END_THINKING_TOKEN_ID] = tokenizer.encode(
"",
{ add_special_tokens: false },
);
let state = "thinking";
let startTime;
let numTokens = 0;
let tps;
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";
}
};
const callback_function = (output) => {
self.postMessage({
status: "update",
output,
tps,
numTokens,
state,
});
};
const streamer = new TextStreamer(tokenizer, {
skip_prompt: true,
skip_special_tokens: true,
callback_function,
token_callback_function,
});
self.postMessage({ status: "start" });
const { past_key_values, sequences } = await model.generate({
...inputs,
do_sample: false,
max_new_tokens: 2048,
streamer,
stopping_criteria,
return_dict_in_generate: true,
});
past_key_values_cache = past_key_values;
const decoded = tokenizer.batch_decode(sequences, {
skip_special_tokens: true,
});
self.postMessage({
status: "complete",
output: decoded,
});
}
self.addEventListener("message", async (e) => {
const { type, data } = e.data;
switch (type) {
case "generate":
stopping_criteria.reset();
generate(data);
break;
case "interrupt":
stopping_criteria.interrupt();
break;
case "reset":
past_key_values_cache = null;
stopping_criteria.reset();
break;
}
});
写在最后
浏览器本地跑大模型正在从 "玩具" 走向 "可用"。WebGPU 的普及、模型量化技术的进步,加上 transformers.js 这样的工程化封装,让普通人打开浏览器就能跑一个 1.5B 的推理模型。
这套架构的思路非常经典:重计算放 Worker、主线程只做 UI、消息驱动、流式渲染、缓存复用。不管是 WebGPU 本地模型,还是调用后端 API 的流式聊天,核心设计思想都是相通的。
希望这篇源码拆解能帮你彻底搞懂浏览器端大模型推理的实现细节。如果觉得有帮助,欢迎点赞收藏。