如何使用TensorFlow JS运行UNet分割模型完成网页图像分割任务
图像输入UNet模型的完整处理流程
你需要按照「图像读取→预处理匹配模型输入→模型推理→结果后处理」的流程实现,具体如下:
核心处理逻辑
- 第一步:从
img标签读取图像并转为张量
用tf.browser.fromPixels()直接读取img元素,得到的是形状为[高度, 宽度, 3]、值范围0-255的uint8张量 - 第二步:预处理对齐模型输入要求
你的模型输入要求为[batch, 1024, 1024, 3],需要做三个处理:- 调整图像尺寸到1024*1024
- 数值归一化,和你Python训练时的预处理逻辑保持一致(通常是除以255将值缩放到0~1区间)
- 扩展出batch维度,匹配模型输入的四维要求
- 第三步:执行推理
调用model.predict()得到输出张量,注意用tf.tidy()自动清理中间张量,避免浏览器显存泄漏 - 第四步:后处理分割结果
将模型输出的张量转为可可视化的格式,比如二分类分割可以设定阈值输出掩码图,多分类可以取类别对应通道后映射颜色
完整示例代码
// 提前加载好的全局可用模型 let model; const MODEL_INPUT_SIZE = 1024; async function getSegmentationResult(imgElement) { // 读取并预处理图像 const imgTensor = tf.browser.fromPixels(imgElement); const preprocessed = tf.tidy(() => { // 调整尺寸到模型要求的1024*1024 const resized = tf.image.resizeBilinear(imgTensor, [MODEL_INPUT_SIZE, MODEL_INPUT_SIZE]); // 归一化到0~1,和Python训练时预处理逻辑保持一致即可 const normalized = resized.cast('float32').div(tf.scalar(255)); // 扩展batch维度,变成[1, 1024, 1024, 3] return normalized.expandDims(0); }); // 模型推理 const predictions = model.predict(preprocessed); // 后处理(以下为二分类分割示例,可根据自身需求调整) const segmentationMask = tf.tidy(() => { // 去掉batch维度,取输出第一通道(二分类输出为单通道的场景) const mask = predictions.squeeze().slice([0,0,0], [-1,-1,1]); // 阈值0.5,大于0.5判定为前景,否则为背景 return mask.greater(tf.scalar(0.5)).cast('float32'); }); // 张量结果同步到CPU const maskData = await segmentationMask.data(); // 手动清理冗余张量,避免显存泄漏 segmentationMask.dispose(); predictions.dispose(); preprocessed.dispose(); imgTensor.dispose(); return { maskData: maskData, maskSize: MODEL_INPUT_SIZE }; } // 分割结果渲染到canvas示例 async function renderMaskToCanvas(maskData, canvasElement) { const ctx = canvasElement.getContext('2d'); canvasElement.width = MODEL_INPUT_SIZE; canvasElement.height = MODEL_INPUT_SIZE; const imageData = ctx.createImageData(MODEL_INPUT_SIZE, MODEL_INPUT_SIZE); // 示例将前景设为半透明红色,背景完全透明,可自定义配色 for (let i = 0; i < maskData.length; i++) { const idx = i * 4; if (maskData[i] === 1) { imageData.data[idx] = 255; imageData.data[idx + 1] = 0; imageData.data[idx + 2] = 0; imageData.data[idx + 3] = 128; } else { imageData.data[idx + 3] = 0; } } ctx.putImageData(imageData, 0, 0); } // 调用示例 // 假设你的img标签id为inputImg,canvas标签id为maskCanvas const inputImg = document.getElementById('inputImg'); const maskCanvas = document.getElementById('maskCanvas'); inputImg.onload = async () => { const segResult = await getSegmentationResult(inputImg); await renderMaskToCanvas(segResult.maskData, maskCanvas); }
注意事项
- 预处理逻辑必须和Python训练时的逻辑完全对齐,如果训练时用的是归一化到
[-1, 1],把归一化代码改成normalized = resized.cast('float32').div(tf.scalar(127.5)).sub(tf.scalar(1))即可 - 如果训练数据是经过等比例缩放+pad填充到1024*1024的,前端预处理也要采用相同逻辑,不要直接拉伸图像,避免分割精度下降
- 推理完成后不需要的张量要及时调用
dispose()释放,避免浏览器内存占用过高导致卡顿
内容的提问来源于stack exchange,提问作者Frank
相关产品推荐
相关产品推荐

