浏览器里跑 DeepSeek-R1,这份 worker.js 到底写了什么

一份真正能跑在浏览器里的 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 是谁?

这是新手最容易懵的点。getInstancestatic(静态)方法,静态方法里的 this 指向类本身 ,也就是 TextGenerationPipeline

所以:

kotlin 复制代码
this.tokenizer ??= ...   // 等价于 TextGenerationPipeline.tokenizer ??= ...

缓存是挂在 上的,而不是某个实例上,天然全局唯一。这就是单例的实现方式------不靠 new,靠静态属性存一份。

3. ??= 到底做了什么(重点)

kotlin 复制代码
this.tokenizer ??= AutoTokenizer.from_pretrained(...);

??=空值合并赋值运算符,语义是:

当左侧变量是 nullundefined 时,才执行右侧并赋值;其他情况(包括 false0'')一律跳过。

等价展开:

kotlin 复制代码
if (this.tokenizer === null || this.tokenizer === undefined) {
  this.tokenizer = AutoTokenizer.from_pretrained(...);
}

执行流程:

  • 第一次调用this.tokenizerundefined → 满足条件 → 执行 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. dtypedevice 两个参数

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():又是一个 ??=。只有第一次(startTimeundefined)才记录时间戳,之后保持不变------避免反复重置计时起点,保证 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 });

一行 ??=,同时干了「懒加载 + 单例缓存 + 并发去重」三件事。下次你在任何"初始化开销大、只需要做一次"的场景,直接照搬这个套路就行。

相关推荐
用户6919026813391 小时前
用浏览器 WebGPU跑DPSK大模型(1) - 模型的下载和前端下载进度的显示
javascript·react.js·架构
夏幻灵2 小时前
Vue 2 与 Vue 3 在应用初始化上的设计差异与技术演进
前端·javascript·vue.js
悟空瞎说2 小时前
Three.js 架构实践:用 MVC 模式组织你的 3D 应用
前端·javascript
王林不想说话2 小时前
JavaScript 从入门到进阶:一篇覆盖核心知识、运行机制与工程实践
javascript
windliang2 小时前
Claude Code 源码分析(十一):一段会话怎样保存、恢复与续写
前端·javascript·人工智能
小小尚@6 小时前
AE脚本-AE Actions v1.1.8 操作动作记录器
开发语言·前端·javascript·jupyter·postman
大家的林语冰6 小时前
👍 超越 ESLint,Oxc 优先采用 TypeScript 7,Rust 和 Go 梦幻联动!
前端·javascript·typescript
胡萝卜术8 小时前
在浏览器中跑 DeepSeek-R1:WebGPU 推理全流程深度解析
前端·javascript·面试
xiaominlaopodaren9 小时前
three.js地图数学基础(六):齐次坐标与矩阵
javascript·gis·three.js