1. 为什么选择在浏览器端跑 AI 推理:一张图背后的架构变化
先说一个我自己的真实场景。去年我做过一个内部图片检索工具,最初的方案非常“标准”:后端挂一个 Python 服务,用 ONNX Runtime 加载 MobileNet 提取特征,再把特征向量存入数据库中做相似度检索。这套东西跑得很顺,直到有一次我在客户现场做演示,会议室网络出奇地差,上传图片到服务器等待推理结果,那张图转了十几秒都没出来,非常尴尬。后来我就在想:如果特征提取这一步能直接放到浏览器里做,数据不出端,服务端只负责保存和检索特征向量,整个体验会不会完全不同?
事实证明,这个方向完全可行。现在的浏览器性能和几年前已经不是一个量级,尤其是 WebGPU 接口的普及,让浏览器端终于能调用 GPU 做通用计算,跑起卷积神经网络来不再只是“玩具级别”。这篇文章我不会去讲大而全的理论,就用“以图搜图”这个具体需求,完整带大家走一遍:从如何把 MobileNet 转成浏览器能用的格式,到怎么在浏览器里用 WebGPU 加速推理,再到最终如何把图片变成 1024 维向量,最后实现一套还算能用的纯前端以图搜图方案。
这个项目适合谁看?如果你已经了解一些深度学习基础,想做图形检索、商品识别、图片分类这类端侧 AI 应用,或者你只是好奇“浏览器里到底能不能流畅跑模型”,那这篇文章应该能给你一个完整的参考。如果你是完全零基础的前端工程师,也没关系,涉及模型原理的部分我会尽量用大白话解释,你跟着操作也能跑通。
1.1 传统以图搜图链路里的三次网络开销
传统以图搜图的服务端架构,至少存在三段明显的网络成本:第一段,用户把原始图片上传到服务器,图片体积从几十 KB 到几 MB 不等,移动端弱网环境下非常痛苦;第二段,服务端拿到图片后要调用推理服务,内部可能还有一次内部网络转发,即使在同一台机器上也有序列化和进程通信的开销;第三段,检索结果返回时,除非只返回名字,否则还要回传缩略图或基础信息。
我当时的痛点就在第一段。一张 2MB 的图片,在普通 4G 网络下上传已经需要好几秒了,再加上服务端排队、推理、返回,用户感知的总延迟往往超过 10 秒。优化手段无非是压缩图片、加 CDN、加缓存,但这些都没有从根上解决“数据非得跑一趟服务端”这个问题。
WebGPU 方案则完全不同。图片从用户本地相册选完,直接用 canvas 缩放到模型需要的尺寸,这时候图片在内存里已经变成一堆数字了。特征提取的整个过程完全发生在本地 GPU 上,没有任何网络请求。等拿到 1024 维向量之后,客户端只需要把这个几 KB 的向量发给服务端做检索,或者更极端一点,直接把特征库都放到本地,那连这一步都省了。这种架构不是小修小补,是把整个链路里的主要瓶颈直接砍掉。
1.2 WebGPU 出现之前,端侧推理为什么“能用但难受”
很多人听到浏览器里跑模型,第一反应是“用 TensorFlow.js 啊”。没错,TensorFlow.js 确实很早就支持了 WebGL 后端,但如果你真的在浏览器里跑过 MobileNet 甚至更大的模型,应该能体会到那种憋屈感。
WebGL 本身是为图形渲染设计的 API,它是按“顶点着色器 + 片元着色器”这套图形管线来工作的。要在 WebGL 里做通用计算,得把张量数据编码成纹理,把卷积运算写成着色器代码,这中间有大量数据排列的额外开销。矩阵乘法这种操作,在 WebGL 里跑起来可能勉强还行,但遇到注意力机制、动态 shape、复杂的分支逻辑,编写和调试的成本非常高。TensorFlow.js 已经把很多脏活封装掉了,但底层执行效率依然受限。
WebGPU 是图形 API 的一次正面革新,它从设计之初就考虑了通用计算场景。通过 Compute Shader,你可以直接把数据塞进 GPU 缓冲区,定义 workgroup 和线程布局,用更接近现代 GPU 硬件模型的方式执行并行计算。不需要再像 WebGL 那样把所有运算伪装成“画一张图”,这带来的性能收益是实打实的。我实测同一台 MacBook Pro 上,同一个 MobileNet ONNX 模型,TensorFlow.js WebGL 后端的单张推理耗时大约是 70 毫秒,换成 WebGPU 后端后降到 20 毫秒左右,这还只是 MobileNet 这种轻量模型,模型越大差距越明显。
1.3 选型结论:WebGPU + MobileNet + 1024 维向量
回到项目本身。以图搜图的核心诉求是:给一张参考图,从库里找回最相似的若干张。这里最常用也最稳妥的做法不是端到端训练一个检索模型,而是先用一个预训练的图像分类模型作为特征提取器,把图片映射成固定长度的向量,再把向量之间的相似度作为图片相似度。
MobileNet 是这个场景里非常合适的选择。它的显存占用小、推理快,同时因为它在 ImageNet 上做过大规模预训练,中间层输出的特征对通用图像内容有很强的表征能力。MobileNetV2 的最后一个卷积层输出经过全局平均池化后,得到的是一个 1280 维的向量;而 MobileNetV1 在同样的操作下得到的是 1024 维向量。1024 维是一个很不错的设计:维度足够表达图像内容,又没有高到让普通设备的余弦相似度计算和存储开销变得离谱。标题里写 1024 维,用的就是 MobileNetV1 这条线。
至于推理框架,我选的是 onnxruntime-web。它的优势在于:第一,模型从 PyTorch 导出一路走 ONNX 格式,整个工具链非常成熟;第二,onnxruntime-web 从 1.17 版本左右开始支持 WebGPU 后端,虽然当时标注为实验性,但已经能用;第三,它提供统一的 Session API,后续想切回 WebGL 或者 CPU 后端做降级处理很方便,代码改动很小。接下来我按完整流程详细说。
2. 工程准备:推理框架、模型格式转换与 WebGPU 运行环境
这一节是纯工程准备环节。看似琐碎,但很多人项目做不下去,就是在准备阶段踩了坑:模型转出来有问题、浏览器版本不支持、WebGPU 设备请求失败,哪一个都致命。我按照自己的操作顺序,把每一步的关键点都过一遍。
2.1 onnxruntime-web 为什么比纯 WebGL 方案更适合
先说清楚选型理由,这样你后续遇到问题才知道往哪个方向排查。
onnxruntime-web 本质上是一个在浏览器里运行的推理引擎,它能把 ONNX 格式的模型编译成可在多种后端执行的代码。后端可以是 CPU 的 WASM、GPU 的 WebGL,也可以是 WebGPU。相比 TensorFlow.js,onnxruntime-web 有几个很实际的优势。
其一,ONNX 生态的兼容性。PyTorch 官方提供了torch.onnx.export,HuggingFace 的模型也大量提供 ONNX 权重,TensorFlow 这边也有tf2onnx。你几乎可以用任何主流框架训练模型,然后统一导出成 ONNX 给 onnxruntime-web 用。TensorFlow.js 只能吃 TF 系的模型,转换链路相对封闭,遇到 PyTorch 模型还得先绕一圈。
其二,WebGPU 后端的执行效率更可控。onnxruntime-web 的 WebGPU EP(Execution Provider)实现了算子级别的 GPU 内核,MobileNet 这种标准 CNN 结构里的 Conv、BatchNorm、Relu、GlobalAveragePool 都有完整支持。我用下来,同一个模型 WebGPU 后端比 WebGL 后端有很明显的速度提升,这是我选它的核心原因。
其三,API 设计比较直接。核心就三件事:创建一个 InferenceSession、用session.run传入输入 Tensor、拿到输出 Tensor。配置executionProviders时可以指定['webgpu', 'wasm'],onnxruntime 会优先尝试 WebGPU,如果环境不支持会自动回退到 WASM,这个平缓降级机制在实际项目里非常有用。
顺便说一句,如果你完全不想碰 WebGPU 的底层细节,onnxruntime-web 确实是目前最省心的浏览器端推理方案。我后面踩到的坑,大多不是框架本身的问题,而是模型转换和图像预处理这些“周边工作”没做到位。
2.2 把 MobileNet 导出成 ONNX:具体命令与参数
在浏览器里跑模型的第一步,是先把模型转换成 ONNX 格式。这里强烈建议:如果你只是想做以图搜图,不要自己从头训练 MobileNet,直接用预训练权重做特征提取器就足够了。ImageNet 预训练模型学到的视觉特征非常通用,用于检索场景效果足够好,自己从头训练既费时间又很难超过它。
我使用的是 PyTorch 官方的torchvision.models.mobilenet_v1对应的预训练权重。不过这里有个小坑:torchvision 原生提供的 MobileNetV1 权重实际上是从 TensorFlow 移植过来的,并且输出头的最后一个分类层是 1000 类全连接层。做特征提取时,我们不需要最后的分类层,只需要到全局平均池化那一步的输出。所以在导出之前,需要先把模型结构改一下。
import torch from torchvision import models # 加载预训练权重 weights = models.MobileNetV1Weights.DEFAULT model = models.mobilenet_v1(weights=weights) # 去掉最后的分类层,保留 feature 部分和全局平均池化 class FeatureExtractor(torch.nn.Module): def __init__(self, backbone): super().__init__() self.features = backbone.features self.pool = torch.nn.AdaptiveAvgPool2d((1, 1)) def forward(self, x): x = self.features(x) x = self.pool(x) x = torch.flatten(x, 1) return x model = FeatureExtractor(model) model.eval() # 导出 ONNX dummy_input = torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy_input, "mobilenet_v1_feature.onnx", input_names=["input_image"], output_names=["feature_vector"], dynamic_axes={ "input_image": {0: "batch_size"}, "feature_vector": {0: "batch_size"}, }, opset_version=17, )这里有三个关键点需要说明。
第一,dynamic_axes。虽然检索场景里我们一般一次只处理一张图,但保留 batch 维度动态化可以让同一个模型在后续批量建库时直接用,不用再重新导出。opset_version 建议用 17 或更高,onnxruntime-web 对较新的 opset 支持度较好。
第二,输入尺寸。MobileNetV1 的原始输入尺寸是 224x224,这个尺寸下特征图是 7x7,全局平均池化后得到特征向量。如果你把输入尺寸改成 192 或 128,模型的浮点运算量会进一步降低,但特征质量也会受影响。我实测 224 是检索准确率和速度都比较平衡的点,不推荐轻易改小。
第三,输出名称。output_names里的feature_vector会在后面前端代码中用到,保持名称一致能减少很多麻烦。
导出完成后,可以用onnxruntime在本地先跑一次推理验证输出维度:
import onnxruntime as ort import numpy as np sess = ort.InferenceSession("mobilenet_v1_feature.onnx") test_input = np.random.randn(1, 3, 224, 224).astype(np.float32) output = sess.run(["feature_vector"], {"input_image": test_input})[0] print(output.shape) # 预期输出 (1, 1024)看到(1, 1024)就说明导出成功。这个 1024 维向量就是后续所有检索操作的基础。
2.3 浏览器端 WebGPU 适配检查与降级策略
跑到浏览器这一步,第一个要确认的是浏览器是否支持 WebGPU。目前 Chrome、Edge 从较新版本开始默认启用 WebGPU,Safari 和 Firefox 的支持状态则比较暧昧。更稳妥的做法是运行时检测,而不是默认用户一定支持。
async function isWebGPUSupported() { if (!navigator.gpu) { return false; } try { const adapter = await navigator.gpu.requestAdapter(); return !!adapter; } catch (e) { return false; } }注意navigator.gpu存在不代表 WebGPU 一定能用,因为还需要 GPU 适配器能够成功创建。我遇到过一种情况:浏览器支持 WebGPU,但设备是虚拟机环境没有 GPU,requestAdapter()返回null,这种情况下只能走降级方案。
onnxruntime-web 的降级处理比较方便,构造 session 时指定多个 execution providers,按优先级排列即可:
const session = await ort.InferenceSession.create("./mobilenet_v1_feature.onnx", { executionProviders: ["webgpu", "wasm"], });如果当前设备支持 WebGPU,onnxruntime 会优先使用;如果不支持,会自动落到 WASM 后端。WASM 后端走的是 CPU 计算,MobileNet 在普通 PC 上大概 100 到 200 毫秒,虽然比 GPU 慢,但至少功能可用,不会白屏。移动端如果 WebGPU 不可用,WASM 的表现也可以接受,毕竟 MobileNet 这个模型本身不算大。
还有一点,WebGPU 后端有时会因为浏览器版本太旧或驱动问题导致推理失败。这种异常往往不是同步抛错,而是 Promise reject。建议在调用session.run时做一层 try/catch,一旦发现 WebGPU 执行失败,可以动态重建一个只走 WASM 的 session 继续跑。把这个逻辑封装成一个函数,用户在弱设备上至少不会完全无法使用。
3. 特征提取链路:从图片到 1024 维向量的完整环节
模型和推理框架都准备好了,接下来进入核心环节:怎么把用户选中的一张图片,变成可供检索的 1024 维向量。看着简单,但图片预处理如果做得不对,特征表达会很差,检索精度直接崩。很多人模型部署跑通了,发现搜出来结果不准,毛病大多出在这一步。
3.1 预处理:resize、归一化、NCHW 与 Float32
MobileNet 的输入要求是三通道 RGB 图像,尺寸 224x224,数值范围经过归一化。在前端拿到图片后,标准处理流程是:
第一步,用 canvas 把图片缩放到 224x224。这一步有两个细节:一是保持宽高比还是直接拉伸?MobileNet 在训练时用的是“等比缩放后再中心裁剪”,但实际检索场景中直接拉伸也可以,因为模型见过各种变形;二是 canvas 在绘制前最好对 EXIF 方向做一次修正,否则手机拍摄的竖图可能被旋转。我之前的做法是用createImageBitmap配合imageOrientation: "from-image"选项,浏览器会自动处理方向,之后再绘制到 canvas。
async function loadImageToCanvas(file) { const bitmap = await createImageBitmap(file, { imageOrientation: "from-image" }); const canvas = document.createElement("canvas"); canvas.width = 224; canvas.height = 224; const ctx = canvas.getContext("2d"); ctx.drawImage(bitmap, 0, 0, 224, 224); return canvas; }第二步,把 canvas 像素数据取出,转成 Float32 数组,并按 ImageNet 的均值标准差归一化。ImageNet 数据集的归一化参数是 mean = [0.485, 0.456, 0.406],std = [0.229, 0.224, 0.225],注意这三个值分别对应 RGB 三个通道。
function preprocessCanvas(canvas) { const ctx = canvas.getContext("2d"); const imageData = ctx.getImageData(0, 0, 224, 224); const data = imageData.data; // RGBA 顺序 const float32Data = new Float32Array(3 * 224 * 224); const mean = [0.485, 0.456, 0.406]; const std = [0.229, 0.224, 0.225]; // 这里直接输出 NCHW 格式 for (let i = 0; i < 224 * 224; i++) { const r = data[i * 4] / 255.0; const g = data[i * 4 + 1] / 255.0; const b = data[i * 4 + 2] / 255.0; float32Data[i] = (r - mean[0]) / std[0]; // R 通道 float32Data[224 * 224 + i] = (g - mean[1]) / std[1]; // G 通道 float32Data[2 * 224 * 224 + i] = (b - mean[2]) / std[2]; // B 通道 } return float32Data; }这里有个特别容易踩的坑:ctx.getImageData返回的是 RGBA 顺序,而 MobileNet 期望的输入是 RGB 顺序,且最常见的 ONNX 模型默认输入布局是 NCHW,也就是通道在 HW 之前。很多从前端转过来的同学习惯写 HWC 布局,直接把像素数组 flatten 后丢给模型,出来的结果就是完全错的。上述代码省略了 A 通道,并直接排成 NCHW,需要特别注意。
第三步,构建ort.Tensor。注意类型必须指定为"float32",维度是[1, 3, 224, 224]。
const tensor = new ort.Tensor("float32", float32Data, [1, 3, 224, 224]);以上准备无误后,session.run就可以拿到特征向量了。
3.2 推理输出:为什么不再需要输出层
你可能会好奇,为什么做特征提取时要删掉 MobileNet 最后的全连接分类层,只保留到全局平均池化?
MobileNet 的结构大致是这样的:卷积层 + 深度可分离卷积层堆叠,最后接一个全局平均池化和全连接分类层。全连接层的作用是拉伸成 1000 个数值,每个值对应 ImageNet 中一个类别的概率。这些概率是高度特化的,它告诉你“这张图片像猫还是像狗”,但对“这张图整体长什么样”的表达并不充分。
全局平均池化输出的 1024 维特征向量则不同。它保留了图片在语义空间中的分布信息,是一种稠密向量表示。两张语义相近的图片,它们的 1024 维向量在欧氏空间中的距离就比较近;语义差异大的图片,向量方向差异就大。这个性质使得余弦相似度或点积成为衡量图片相似度的有效方式。
还有一点,去掉分类层后模型变得更小更快。MobileNetV1 完整模型约 4.2MB,去掉最后的全连接层后大约 3.7MB,差别不算大,但省掉的少量计算对移动端还是有意义的。更重要的是,删掉分类层后模型不再依赖 ImageNet 的 1000 类标签,它输出的是纯特征,后续你可以在特征之上连接自己的分类头,也可以直接用于检索。
3.3 一条真实查询链路的前端代码走读
把前面的单个步骤串起来,一张图片从选中到拿到 1024 维向量的完整流程是这样的:
async function extractFeature(file, session) { // 1. 加载图片到 canvas,完成 resize 和方向修正 const canvas = await loadImageToCanvas(file); // 2. 预处理:归一化,转 NCHW Float32 const float32Data = preprocessCanvas(canvas); // 3. 构造 Tensor 并推理 const tensor = new ort.Tensor("float32", float32Data, [1, 3, 224, 224]); const feeds = { input_image: tensor }; const results = await session.run(feeds); // 4. 取出输出特征向量 const featureVector = results.feature_vector.data; return Array.from(featureVector); // 1024 个 float32 数值 }这里的results.feature_vector是 onnxruntime 返回的输出张量,.data属性是一个 Float32Array,长度 1024。如果你导出模型时用的是我上面的output_names=["feature_vector"],这里字段名直接对应上即可。
推理完之后,我习惯马上对向量做一次 L2 归一化。这样后续检索时可以直接用点积代替余弦相似度,省掉每次计算范数的时间。
function normalizeVector(vec) { const norm = Math.sqrt(vec.reduce((sum, x) => sum + x * x, 0)); if (norm === 0) return vec; return vec.map((x) => x / norm); }从用户角度,从点击图片到拿到特征向量,整个过程在本机完成,没有任何网络请求,这是浏览器端推理最大的体验优势。实测下来,在带有 WebGPU 的 PC 上,这一步耗时稳定在 20 毫秒左右,用户几乎感知不到延迟。
4. 向量数据构建与检索:没有数据库也能做以图搜图
特征提取做好之后,剩下的问题就是检索了。以图搜图本质上分两段:第一段是“建库”,把候选图库里的每张图片都跑一遍特征提取,存成特征向量;第二段是“查询”,给定一张查询图,提取特征向量后和库里的所有向量算相似度,取 TopK。
4.1 特征库怎么建:内存数组与本地持久化
如果图片库不太大(几百到几千张),直接在浏览器内存里维护一个数组就够用了。每个条目可以这样存:
const featureDB = [ { id: "image_001", path: "./images/001.jpg", vector: Float32Array(1024) }, { id: "image_002", path: "./images/002.jpg", vector: Float32Array(1024) }, // ... ];建库的过程就是对每张候选图片调用一次extractFeature,得到的向量归一化后存入数组。这里有个体感建议:建库是一次性成本,可以在用户空闲时后台逐张处理,或者做成一个“导入图库”的上传页面,把图片路径和向量都持久化下来。
浏览器端持久化最方便的是 IndexedDB。直接把整个 featureDB 数组存进去,下次打开页面直接读取,不需要重新提取一遍特征。IndexedDB 存取二进制数组的性能不错,1024 维向量存 1000 张图大约 4MB 数据,完全在可接受范围内。
如果你有后端支撑,也可以把特征向量上报到服务端保存。查询时有两种玩法:一种是查询图片的特征向量在浏览器本地算好,只发送 1024 维向量给后端;另一种是连特征提取都在本地做,后端只做一个向量检索服务。前一种方案数据量小、网络开销低,是 WebGPU 端侧推理最常见的架构形态。
4.2 余弦相似度、归一化与 TopK 查询实现
查询的核心就是相似度计算。由于建库和查询时都有做 L2 归一化,余弦相似度可以直接退化成点积。对每个库中的向量和查询向量做内积,值越大表示越相似。
function cosineSimilarity(a, b) { let dot = 0; for (let i = 0; i < a.length; i++) { dot += a[i] * b[i]; } return dot; } function searchTopK(queryVec, featureDB, k = 10) { const results = []; for (let i = 0; i < featureDB.length; i++) { const score = cosineSimilarity(queryVec, featureDB[i].vector); results.push({ id: featureDB[i].id, path: featureDB[i].path, score }); } results.sort((a, b) => b.score - a.score); return results.slice(0, k); }这个实现非常简单,但在特征库低于 1 万条时性能完全够用。因为 1 次查询要做 N 次 1024 维向量点积,1 万条就是 1 千万次浮点乘法,现代浏览器里也就几十毫秒的事。
关于相似度阈值,我建议根据实际数据分布去定,不建议拍脑袋。先拿一批同类别图片跑一遍查询,记录最低相似度,拿不同类别的图片跑一遍,观察区分度。通常来说,相似度 0.8 以上可以认为近似同图,0.6 到 0.8 是语义相近但不完全一致,低于 0.5 基本就没什么相关性了。这组数值会随数据集变化,最好做一次小规模验证再定。
4.3 数据量大一些怎么办:分桶近似与后端兜底
如果图片库涨到几万张以上,每次都线性扫全部向量会感觉到明显卡顿。此时可以从两个维度优化。
第一个维度是降维。1024 维向量可以直接用 PCA、随机投影之类的算法压到 128 维或 64 维,检索速度可以提升数倍,同时精度损失通常可控。完全用 JavaScript 实现 PCA 不是不行,但数据量大了容易把主线程卡住,建议放到 Web Worker 里做。
第二个维度是建索引。前端可以做一个粗粒度的哈希分桶:对向量做符号随机投影(也就是 SimHash),得到 64 位二进制指纹,然后按指纹的前 8 位分桶。查询时只和候选桶内的向量算精细相似度。这样做的原理是:如果两个向量在原始空间里比较接近,它们的 SimHash 指纹也大概率相同。分桶是一种召回策略,能保证快速排除大部分无关向量,然后由精细相似度负责精排。
这里提一句,如果项目本身有后端,那么向量检索建议直接交给专业工具做。比如用 FAISS 或者专门的向量数据库,建好 HNSW 索引,几百万条向量也能在毫秒级返回。浏览器端方案适合数据量可控、离线优先的场景,它是架构选型的一个选项,不是银弹。我在生产环境里的做法是:前端优先走本地检索,命中结果直接显示;如果本地没有足够多的候选,再把查询向量发给后端,由后端的向量库补充结果。
5. 实测数据、调优方向与踩坑记录
最后这部分,我把项目运行过程中记录的真实数据和踩坑经验分享出来。这些东西在官方文档里很难找到,但实际工程里价值很高。
5.1 我这组测试数据:不同设备上的推理耗时
我分别在三类设备上做了同一模型的推理测试,统一输入一张 224x224 的猫图片,运行环境是 onnxruntime-web + WebGPU 后端(不支持时降级 WASM)。
| 设备 | 后端 | 单张推理耗时 | 内存增量 | 备注 |
|---|---|---|---|---|
| MacBook Pro M1 Pro | WebGPU | 18-25 ms | 约 120MB | 稳定,温控良好 |
| 中端 Windows PC (GTX 1660) | WebGPU | 22-30 ms | 约 150MB | 需要用 Chrome 系浏览器 |
| Android 中端手机 (骁龙 778G) | WASM | 90-130 ms | 约 80MB | WebGPU 可用性不稳定 |
| 低端 Android 手机 | WASM | 180-250 ms | 约 80MB | 可接受但不丝滑 |
MobileNet 本身是一个非常轻量的网络,在 WebGPU 加持下,PC 端基本能做到“随点随算”,用户无感知。移动端即使走 WASM,200 毫秒左右的延迟对于一次离线图片检索也是能接受的。如果你想追求移动端 GPU 加速,可以再等等 WebGPU 在移动浏览器上的普及,或者考虑用 WebGL 后端做降级中间层。
建库耗时也顺便记录一下:1000 张图片在 MacBook Pro 上跑完全部特征提取,大约需要 20 秒左右。这个耗时分布在图片解码、canvas 绘制和 GPU 推理上。图片解码是最大的瓶颈,尤其是 JPEG 大图。后续可以考虑把图片缩略图先做出来,推理时直接用压缩后的图。
5.2 三个影响很大的调优细节
第一个细节是输入张量的内存复用。如果用session.run时每次都重新new ort.Tensor,底层会有频繁的内存分配和释放。对于单张查询影响不大,但如果做建库,图片是批量处理的,建议预分配一块输入缓冲区,重复使用同一个 Tensor,只替换数据内容。不过要注意,onnxruntime 的 Tensor 对象底层数据是共享的,直接修改原始 Float32Array 的内容是可行的,但要确保 shape 和 dtype 一致。
第二个细节是 canvas 的willReadFrequently选项。getImageData在某些浏览器上会触发 canvas 从 GPU 回读操作,速度较慢。创建 2D context 时加上{ willReadFrequently: true },是在告诉浏览器“我经常读取像素数据”,浏览器会据此优化存储方式,实测可以明显减少getImageData的耗时。
第三个细节是 Web Worker 的使用。特征提取的预处理、vector 归一化、TopK 排序这些计算量都可能阻塞主线程,尤其是前两者涉及循环。把整个推理和检索逻辑放进 Web Worker,从主线程只接收图片文件或 ImageBitmap 和查询结果,UI 的流畅度会好很多。移动端页面对主线程阻塞非常敏感,哪怕只卡 100 毫秒,滚动动画都会掉帧。建议从一开始就走 Worker 方案,后面省得重构。
5.3 踩过的坑与解决记录
坑一:模型输入 channel order 写反。我第一次跑通的时候,检索结果惨不忍睹,同类的图片排到很后面。排查了半天,最后发现是getImageData返回 RGBA,我把它当成 RGB 直接用了,导致 R 通道实际存的是绿通道的数据。这个问题极难用肉眼发现,因为模型不会报错,只是结果不对。后来我在预处理函数里加了一个自检:对一张纯红色图片做推理,打印特征向量的前几个值,和从 Python 侧得到的结果进行对比,一旦不一致立即检查通道顺序。
坑二:WebGPU 后端偶发推理失败。在 Windows 某些浏览器版本上,onnxruntime-web 的 WebGPU 后端跑一段时间后会报A requested resource is already in use,这种错误没有固定复现规律。我的处理方案是给推理服务加一个降级开关:一旦捕获到 WebGPU 执行异常,就自动销毁当前 session,创建一个新的 WASM session 继续跑。同时把错误上报到一个本地日志模块,方便后续定位。
坑三:ONNX 模型的动态 shape 支持。我最初导出模型时没有加dynamic_axes,导致每次推理只能输入固定的 batch size 1。后来想批量建库,用 batch size 8 的 Tensor 去跑,直接报 shape 不匹配。重新导出模型加上动态轴后问题解决。建议从一开始就加上动态轴支持,不然后面还得反复导出模型。
坑四:IndexedDB 存取 Float32Array 的坑。直接存 Array、再转 Float32Array 的话,大数据量下速度很慢。IndexedDB 原生支持ArrayBuffer,直接把 Float32Array 的buffer存进去,读取时再new Float32Array(result.buffer)包裹一层,性能好很多。存的等待时间从原来的几百毫秒降到几十毫秒。
说实话,这整套流程做下来,最大的感触不是 WebGPU 有多快,而是“端侧推理”这件事终于不再是一个工程噱头了。它真正把传统以图搜图架构里的网络瓶颈砍掉了,让应用可以在弱网、局域网甚至完全离线的环境里都有着不错的体验。我自己在实际项目中已经把这个方案跑成了主链路,本地 5000 多张图片的检索库,从选图到出结果基本在一秒以内,其中绝大部分时间还是花在数据读取和界面渲染上。
最后分享一个小技巧:在调试阶段,可以在页面里做一个简单的相似度分布可视化,把当前的查询向量和库中所有向量的相似度画成直方图。你会非常直观地看到阈值该定多少、模型在哪些图库上表现差,这比闷头调代码高效得多。希望这篇文章能帮你少走点弯路。