WebAssembly AI 插件:浏览器端推理引擎的设计与 Rust 实践

cover

一、AI 推理不一定要在服务器,浏览器也能跑

我第一次在浏览器里跑通一个 MNIST 手写数字识别模型时,激动了好久——不需要服务器、不需要 API Key、打开网页就能推理。虽然浏览器端的算力远不如 GPU 服务器,但对于小模型(< 50MB)和低延迟场景(< 100ms),浏览器端推理有独特优势:零服务器成本、数据不出浏览器、离线可用。

WebAssembly 是浏览器端 AI 推理的关键技术:它让 Rust/C++ 编写的推理引擎可以编译为 .wasm 文件,在浏览器中以接近原生的速度运行。这篇文章记录我用 Rust + WASM 构建浏览器端推理插件的完整过程。

二、WASM AI 推理插件的架构

flowchart TB
    A[Rust 推理引擎源码] --> B[wasm-pack 编译]
    B --> C[.wasm 文件 + JS 胶水代码]

    C --> D[浏览器加载]
    D --> D1[WebAssembly.instantiate<br/>初始化 WASM 模块]
    D1 --> D2[分配 WASM 线性内存<br/>加载模型权重]

    D2 --> E[推理接口]
    E --> E1[infer(input_ptr, len)<br/>执行前向传播]
    E --> E2[get_output_ptr()<br/>读取推理结果]

    E1 --> F[WASM 线性内存]
    E2 --> F

    F --> G[JavaScript 侧<br/>TypedArray 交互]

    subgraph 性能优化
        H[SIMD 指令<br/>wasm-simd128]
        I[Web Workers<br/>多线程推理]
        J[模型量化<br/>INT8 权重]
    end

    H --> E1
    I --> E1
    J --> D2

    style B fill:#e3f2fd
    style F fill:#fff3e0
    style H fill:#e8f5e9

WASM AI 推理插件的核心是 Rust 推理引擎编译为 .wasm 文件,通过 wasm-pack 生成 JS 胶水代码。数据交互通过 WASM 线性内存完成——JS 侧用 TypedArray 写入输入数据,Rust 侧读取并执行推理,结果写回线性内存供 JS 读取。性能优化依赖 SIMD 指令、Web Workers 多线程和模型量化。

三、代码实现与分析

3.1 Rust 推理引擎核心

// src/lib.rs
use wasm_bindgen::prelude::*;

/// 简单的全连接层推理引擎
#[wasm_bindgen]
pub struct NeuralNetwork {
    weights: Vec<f32>,
    biases: Vec<f32>,
    input_size: usize,
    output_size: usize,
}

#[wasm_bindgen]
impl NeuralNetwork {
    /// 创建新的网络实例
    #[wasm_bindgen(constructor)]
    pub fn new(input_size: usize, output_size: usize) -> Self {
        // 初始化随机权重(实际场景从文件加载)
        let weight_count = input_size * output_size;
        let mut weights = Vec::with_capacity(weight_count);
        for i in 0..weight_count {
            // Xavier 初始化
            let scale = (2.0 / (input_size + output_size) as f32).sqrt();
            weights.push(pseudo_random(i) * scale);
        }

        let biases = vec![0.0f32; output_size];

        Self {
            weights,
            biases,
            input_size,
            output_size,
        }
    }

    /// 从字节数组加载模型权重
    pub fn load_weights(&mut self, data: &[u8]) -> Result<(), JsValue> {
        if data.len() != self.weights.len() * 4 + self.biases.len() * 4 {
            return Err(JsValue::from_str("权重数据长度不匹配"));
        }

        let float_data: Vec<f32> = data
            .chunks_exact(4)
            .map(|chunk| {
                f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]])
            })
            .collect();

        let weight_count = self.weights.len();
        self.weights = float_data[..weight_count].to_vec();
        self.biases = float_data[weight_count..].to_vec();

        Ok(())
    }

    /// 执行前向推理
    pub fn infer(&self, input: &[f32]) -> Vec<f32> {
        if input.len() != self.input_size {
            return vec![];
        }

        let mut output = vec![0.0f32; self.output_size];

        // 矩阵乘法:output = weights * input + biases
        for j in 0..self.output_size {
            let mut sum = self.biases[j];
            for i in 0..self.input_size {
                sum += self.weights[j * self.input_size + i] * input[i];
            }
            // ReLU 激活
            output[j] = sum.max(0.0);
        }

        output
    }

    /// 获取模型信息
    pub fn model_info(&self) -> String {
        format!(
            "输入维度: {}, 输出维度: {}, 参数量: {}",
            self.input_size,
            self.output_size,
            self.weights.len() + self.biases.len(),
        )
    }
}

/// 简单伪随机数(确定性,用于初始化)
fn pseudo_random(seed: usize) -> f32 {
    let x = (seed as u64).wrapping_mul(6364136223846793005)
        .wrapping_add(1442695040888963407);
    (x >> 33) as f32 / u32::MAX as f32 * 2.0 - 1.0
}

/// Softmax 函数(用于分类模型的输出层)
#[wasm_bindgen]
pub fn softmax(input: &[f32]) -> Vec<f32> {
    if input.is_empty() {
        return vec![];
    }

    let max_val = input.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
    let exps: Vec<f32> = input.iter().map(|&x| (x - max_val).exp()).collect();
    let sum: f32 = exps.iter().sum();

    exps.iter().map(|&x| x / sum).collect()
}

3.2 JavaScript 侧交互

// pkg/ai_inference.js(由 wasm-pack 生成,此处展示使用方式)
import init, { NeuralNetwork, softmax } from './pkg/ai_inference.js';

async function runInference() {
    // 初始化 WASM 模块
    await init();

    // 创建网络实例
    const net = new NeuralNetwork(784, 10);  // MNIST: 28x28 输入, 10 分类

    // 加载预训练权重
    const weightResponse = await fetch('./model_weights.bin');
    const weightData = new Uint8Array(await weightResponse.arrayBuffer());
    net.load_weights(weightData);

    console.log(net.model_info());

    // 准备输入数据(MNIST 图像归一化到 0-1)
    const imageData = new Float32Array(784);
    // ... 填充图像数据 ...

    // 执行推理
    const t0 = performance.now();
    const logits = net.infer(imageData);
    const inferenceTime = performance.now() - t0;

    // Softmax 得到概率
    const probabilities = softmax(logits);

    // 获取预测类别
    const predictedClass = probabilities.indexOf(Math.max(...probabilities));
    console.log(`预测类别: ${predictedClass}, 置信度: ${probabilities[predictedClass].toFixed(4)}`);
    console.log(`推理耗时: ${inferenceTime.toFixed(1)}ms`);
}

// Web Worker 中运行推理(避免阻塞 UI)
// worker.js
self.onmessage = async (e) => {
    const { input } = e.data;
    const net = new NeuralNetwork(784, 10);
    // ... 加载权重和推理 ...
    self.postMessage({ result: probabilities });
};

3.3 构建配置与性能优化

# Cargo.toml
[package]
name = "ai-inference"
version = "0.1.0"
edition = "2021"

[lib]
crate-type = ["cdylib", "rlib"]

[dependencies]
wasm-bindgen = "0.2"
js-sys = "0.3"
web-sys = { version = "0.3", features = ["Window", "Performance"] }

[profile.release]
opt-level = 3
lto = true               # 链接时优化,减小 .wasm 体积
codegen-units = 1         # 单编译单元,更好的优化

[features]
default = ["simd"]
simd = []                 # 启用 WASM SIMD
# 构建命令
# 1. 安装 wasm-pack
cargo install wasm-pack

# 2. 编译为 WASM(启用 SIMD)
wasm-pack build --target web -- --features simd

# 3. 检查 .wasm 文件大小
ls -lh pkg/ai_inference_bg.wasm
# 目标:< 500KB(未量化模型权重需单独加载)

四、WASM AI 推理的边界与权衡

模型大小限制:WASM 线性内存默认上限 4GB,但浏览器实际可用内存远小于此。模型权重 + 推理中间结果 + WASM 模块本身,总内存占用应控制在 500MB 以内。超过这个限制,移动端浏览器可能崩溃。建议对模型做 INT8 量化,将权重体积压缩 4 倍。

SIMD 的浏览器兼容性:WASM SIMD (wasm-simd128) 在 Chrome 91+、Firefox 89+、Safari 16.4+ 支持。旧版浏览器需要回退到标量实现。建议用 wasm-feature-detect 库检测 SIMD 支持,动态加载对应版本。

多线程推理的复杂性:Web Workers + SharedArrayBuffer 可以实现多线程推理,但 SharedArrayBuffer 要求服务器设置特定的 CORS 头(Cross-Origin-Opener-Policy: same-origin 和 Cross-Origin-Embedder-Policy: require-corp)。很多 CDN 和静态托管服务不支持这些头,导致多线程推理无法使用。

推理精度与量化的取舍:INT8 量化可以显著减小模型体积和加速推理,但会损失精度。对于分类任务,INT8 的精度损失通常可接受(< 1%);对于回归任务或需要精确数值的场景,建议保持 FP32 或使用混合精度。

五、总结

WebAssembly 让 Rust 编写的 AI 推理引擎可以在浏览器中以接近原生的速度运行,实现零服务器成本的端侧推理。本文的关键实践为:用 wasm-bindgen 暴露 Rust 推理接口给 JavaScript、通过 WASM 线性内存传递输入输出数据、用 LTO 和 codegen-units 优化 .wasm 体积、用 Web Workers 避免推理阻塞 UI。WASM AI 推理适合小模型和低延迟场景,大模型和训练场景仍需服务器端 GPU。浏览器兼容性和内存限制是当前的主要约束。

Logo

openEuler 是由开放原子开源基金会孵化的全场景开源操作系统项目,面向数字基础设施四大核心场景(服务器、云计算、边缘计算、嵌入式),全面支持 ARM、x86、RISC-V、loongArch、PowerPC、SW-64 等多样性计算架构

更多推荐