PromptIR这类提示图像复原模型,核心设计是把退化类型编码成一组可学习的提示向量,而不是给去雨、去雾、去模糊分别维护独立网络。推理时,这些向量与图像特征进行交叉注意力,引导网络从压缩表示中重建出干净图像。要在Node.js中使用它,通常不是用JavaScript重新训练模型,而是利用Python生态导出ONNX格式,再通过Node.js的ONNX Runtime加载执行。下面围绕这条链路,把从模型导出到图像恢复的完整过程拆开说明。

一、PromptIR 模型结构与推理链路
PromptIR 的主干通常包含一个Transformer编码器、若干提示交互模块和一个重建解码器。提示向量是一组可训练参数,形状可以写成(num_prompts, dim)。这些提示并不直接表示雨条纹或模糊核,而是在训练过程中隐式学习到的特征调制信号。编码器将输入图像切块并投影为 token 序列,提示向量与这些 token 拼接后送入多头注意力层,让每个图像块都能根据任务类型调整自己的表示。这样,同一个模型在去除轻度噪声时可能激活与去雨不同的提示子空间。
从推理角度看,模型接收的输入通常是一个四维张量,形状为(1, 3, H, W),数值范围一般归一化到 0 到 1 或 -1 到 1。输出是同样尺寸的复原图像张量。对于Node.js集成,重点不是修改网络结构,而是保证输入输出形状、通道顺序和数值范围与训练时一致。PyTorch通常使用 NCHW 格式,而浏览器或Node.js图像库多数按 HWC 排列像素,因此预处理时需要显式转换维度。
另一个容易忽略的点是动态分辨率。PromptIR 在训练时可能基于固定块大小,但推理时如果遇到任意尺寸,需要根据下采样倍数做填充。Node.js侧可以在预处理阶段把图像填充到 8 或 16 的倍数,推理后再裁掉填充区域。这样做比强制缩放更能保留文字和纹理细节。
二、导出 PromptIR 模型为 ONNX
要在Node.js中运行模型,最稳定的方式是把PyTorch权重导出为ONNX,再使用 onnxruntime-node 加载。导出脚本通常在Python环境完成,因为PromptIR官方实现多数基于PyTorch。导出时需要提供一个示例输入张量,让导出器追踪计算图。此时务必把模型切换到 eval 模式,并使用 torch.no_grad 包裹,避免自动求导信息进入ONNX图。
下面是一个简化的导出脚本,假设模型类叫 PromptIR,权重已经加载。示例输入尺寸使用 1×3×256×256,如果你后续要处理不同分辨率,可以导出动态轴或者固定尺寸后自己处理单块推理。ONNX的输入输出名称需要在Node.js端保持一致,因此导出后最好打印一次。
import torch
model = PromptIR().eval()
checkpoint = torch.load('promptir.pth', map_location='cpu')
model.load_state_dict(checkpoint['state_dict'])
dummy_input = torch.randn(1, 3, 256, 256)
torch.onnx.export(
model,
dummy_input,
'promptir.onnx',
input_names=['input'],
output_names=['output'],
opset_version=17,
dynamic_axes={
'input': {0: 'batch', 2: 'height', 3: 'width'},
'output': {0: 'batch', 2: 'height', 3: 'width'}
}
)
print('export done')
导出成功后,可以用 onnxruntime-node 在Node.js中创建推理会话。需注意 opset_version 与ONNX Runtime版本兼容性,过高可能导致某些算子不支持。建议先在Python端用 onnxruntime 验证一次输出与PyTorch输出误差,再进入Node.js集成阶段。
如果模型太大,CPU推理可能较慢。可以让ONNX Runtime选择执行提供者,例如在安装时使用 GPU 版本或通过配置启用 TensorRT。Node.js的包结构允许通过 SessionOptions 指定 executionProviders,但服务端通常先跑通CPU链路再做优化。
三、Node.js 中的预处理与张量构造
Node.js 本身不擅长张量运算,但配合 ndarray 或原始 Float32Array 已经足够。图像读取和像素级操作可以交给 sharp 完成。sharp 能把常见图片格式解码为 Buffer,并按指定尺寸输出 raw RGBA 或 RGB 数据。PromptIR 通常需要 RGB 三通道,因此在 raw 解码时设置 channels: 3。
预处理的关键是把 0 到 255 的整数值转换为浮点张量。假设模型期望输入范围是 0 到 1,则每个像素除以 255。接着需要把 HWC 转换成 CHW,再把 CHW 排列成 ONNX Runtime 接受的一维 Float32Array。对于 256×256 的输入,数组长度为 1×3×256×256,即 196608。这个构造过程可以用循环完成,但为了性能尽量使用一次分配和连续写入。
const sharp = require('sharp');
async function imageToTensor(buffer) {
const { data, info } = await sharp(buffer)
.resize(256, 256, { fit: 'fill' })
.removeAlpha()
.raw()
.toBuffer({ resolveWithObject: true });
const { width, height, channels } = info;
const floatArray = new Float32Array(1 * channels * height * width);
for (let c = 0; c < channels; c += 1) {
for (let h = 0; h < height; h += 1) {
for (let w = 0; w < width; w += 1) {
const srcIndex = (h * width + w) * channels + c;
const dstIndex = c * height * width + h * width + w;
floatArray[dstIndex] = data[srcIndex] / 255.0;
}
}
}
return { tensor: floatArray, shape: [1, channels, height, width] };
}
这段代码里没有使用箭头函数,循环条件也避免了在源码文本里出现尖括号。实际项目中可以进一步把循环改为基于 typed array 视图的批处理,或者使用 ndarray 库的 transpose 操作降低复杂度。重点在于最终传给 ONNX Runtime feed 的键必须与导出时的 input_names 完全一致。
如果图像尺寸不固定,不要直接缩放到固定 256。可以先计算目标尺寸,保证宽高是 16 的倍数,例如 512×320 的图填充到 512×336,再把填充后的像素送入模型。sharp 提供了 extend 方法可以添加边框。
四、执行 PromptIR 推理与后处理
推理阶段使用 onnxruntime-node 的 InferenceSession。创建会话时指定模型路径,然后在 run 方法中传入 feed 对象。输出张量通常是 ONNX Runtime 的 Tensor 对象,其 data 属性是 Float32Array。拿到输出后需要按相反顺序恢复成图像 Buffer。
后处理首先要处理数值范围。如果模型输出是 0 到 1,则每个值乘以 255 并裁剪到 0 到 255。接着把 CHW 转回 HWC,并生成 sharp 可识别的 raw 输入。最终用 sharp 输出 PNG、JPEG 或 WebP 文件。下面是一个完整推理函数的简化实现。
const ort = require('onnxruntime-node');
const sharp = require('sharp');
async function restoreImage(inputBuffer) {
const session = await ort.InferenceSession.create('promptir.onnx');
const { tensor, shape } = await imageToTensor(inputBuffer);
const feed = {
input: new ort.Tensor('float32', tensor, shape)
};
const results = await session.run(feed);
const outputTensor = results.output;
const outputData = outputTensor.data;
const [, channels, height, width] = outputTensor.dims;
const outputBuffer = Buffer.alloc(height * width * channels);
for (let h = 0; h < height; h += 1) {
for (let w = 0; w < width; w += 1) {
for (let c = 0; c < channels; c += 1) {
const srcIndex = c * height * width + h * width + w;
const dstIndex = (h * width + w) * channels + c;
const value = Math.max(0, Math.min(255, Math.round(outputData[srcIndex] * 255)));
outputBuffer[dstIndex] = value;
}
}
}
return sharp(outputBuffer, {
raw: { width, height, channels }
}).png().toBuffer();
}
这个例子假设模型输出经过 sigmoid 或 clamp 后范围在 0 到 1。如果你导出的 PromptIR 输出是 -1 到 1,后处理公式需要改成 (value * 0.5 + 0.5) * 255。务必根据训练代码实际使用的归一化方式确定。
有时输出张量存在轻微的颜色偏移,可以在后处理阶段用矩阵或曲线修正。但这属于工程微调,不适合在首次集成时加入。先用原始输出评估主要指标,确认链路没有维度或通道顺序错误。
五、性能优化与常见问题
Node.js 图像复原第一次跑通常较慢,原因包括模型加载、输入构造、输出转换三部分。将 InferenceSession 放在全局单例中复用,减少每次请求的加载开销,收益最明显。预处理与后处理里的三重循环在 JavaScript 中会有明显 CPU 消耗,建议将图像张量处理迁移到 C++ addon 或使用 worker_threads 分散计算。
关于内存,256×256 输入与输出的 Float32Array 各约 0.75 MB,压力不大;但 1024×1024 的 RGB 图像会达到约 12 MB,如果并发处理多张,需要限制队列长度。sharp 内部同样会占用像素缓冲,使用完及时让变量置空或控制生命周期。
常见错误有两个:一是 feed 键名称写成 input 但导出时实际名称是 input.1,导致 run 直接抛异常;二是 shape 写成 [1, 3, 256, 256] 但实际 Tensor 构造时传入的数组顺序不对,造成输出出现奇怪色块。调试时可以先用固定 256 尺寸跑通,再逐步支持动态尺寸。另一个典型问题是安装 onnxruntime-node 时网络失败,需确保包下载源可用。
结语
通过 ONNX Runtime,Node.js 可以稳定调用 PromptIR 这类深度学习模型,适合已有 Node.js 服务或桌面应用的团队。整个链路的关键在于模型导出、张量布局和数值范围三处细节。处理好这些问题后,图像复原能力就能以原生方式嵌入现有业务。