You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使用TensorFlow JS运行UNet分割模型完成网页图像分割任务

图像输入UNet模型的完整处理流程

你需要按照「图像读取→预处理匹配模型输入→模型推理→结果后处理」的流程实现,具体如下:

核心处理逻辑

  • 第一步:从img标签读取图像并转为张量
    用tf.browser.fromPixels()直接读取img元素,得到的是形状为[高度, 宽度, 3]、值范围0-255的uint8张量
  • 第二步:预处理对齐模型输入要求
    你的模型输入要求为[batch, 1024, 1024, 3],需要做三个处理:
    1. 调整图像尺寸到1024*1024
    2. 数值归一化,和你Python训练时的预处理逻辑保持一致(通常是除以255将值缩放到0~1区间)
    3. 扩展出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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.05 02:18:01