DeepSeek-R1 WebGPU(六):浏览器本地跑推理的完整实现

一、项目整体架构

这是一个典型的双线程架构

  • 主线程(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]);
  }
}

逐点解析:

  1. static 静态属性tokenizermodel 挂载在类本身,而非实例上,全局唯一。
  2. ??= 空值合并赋值 :只有变量为 null / undefined 时才执行右侧赋值。第一次调用触发下载加载,后续直接复用内存中已就绪的对象。
  3. dtype: "q4f16" :4bit 量化权重,把 1.5B 参数模型压缩到约 1GB 显存以内,浏览器才能跑得动。
  4. device: "webgpu" :推理运算交给浏览器 GPU,纯 CPU 跑 1.5B 模型几乎不可用。
  5. 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;

执行流程:

  1. 传入 past_key_values_cache 作为历史缓存
  2. 模型只对新输入的 token 做注意力计算
  3. 输出更新后的 past_key_values,包含全部历史
  4. 赋值回全局缓存,下一轮对话继续复用

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;

经典流式渲染套路:

  1. 生成开始时:往消息数组 push 一条空的 assistant 消息占位
  2. 每次流式更新:找到最后一条 assistant 消息,把新文本拼接到 content 后面
  3. 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 输出的完整链路:

  1. 用户回车发送onEnter 执行,setMessages 追加 user 消息
  2. useEffect 检测 → messages 变化,最后一条是 user,发送 type: "generate"
  3. Worker 收到消息stopping_criteria.reset() 清除旧标记,调用 generate(messages)
  4. 模型单例加载 → 第一次下载权重,后续直接复用
  5. 模板拼接apply_chat_template 把对话数组转成模型认识的 prompt 格式
  6. 流式生成model.generate 逐 token 输出,TextStreamer 双回调处理
  7. token 回调 → 统计 TPS、检测 `` 切换状态
  8. 文本回调postMessage 把文本 + 状态 + TPS 发回主线程
  9. 前端渲染 → 追加到最后一条 AI 消息,打字机效果
  10. 生成结束 → 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 的流式聊天,核心设计思想都是相通的。

希望这篇源码拆解能帮你彻底搞懂浏览器端大模型推理的实现细节。如果觉得有帮助,欢迎点赞收藏。

相关推荐
新时代牛马1 小时前
Vue.js 响应式原理详解
前端·javascript·vue.js
BreezeJiang1 小时前
为什么 Next.js 要把组件切成两半?从一个 SEO 死穴说起
javascript·react.js
Molecular_Chat1 小时前
Biotin SNA 生物素偶联凝集素结合特异性、生物素标记效率定量表征研究
javascript·python
进击的蛋蛋1 小时前
JS对象的遍历
javascript·面试
mONESY1 小时前
React 前端如何不傻等后端接口?
前端·javascript·后端
xiaominlaopodaren1 小时前
three.js地图数学基础(七):地图相机
javascript·gis·three.js
vx-程序开发2 小时前
django汽车租赁系统---附源码25360
java·javascript·spring boot·python·eclipse·django·php
满栀5852 小时前
vue3动态路由详细效果
前端·javascript·vue.js·typescript·前端框架
山荷枝3 小时前
05-Vue
前端·javascript·vue.js