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

TensorFlow.js:食品分割模型掩码与原图融合全黑问题求助

问题:食品分割模型融合结果全黑,如何修复?

尝试运行谷歌移动端食品分割模型,想要将分割掩码与原始图像融合显示,结果输出全黑图像。期望实现灰度分割掩码叠加在图像上,高亮识别出的食品区域。

提供的测试代码:

<!DOCTYPE html>
<html>
  <head>
    <title>Image Transform with TensorFlow.js</title>
    <!-- Load TensorFlow.js library -->
    <script src="https://cdn.jsdelivr.net/npm/@tensorflow/tfjs@4.1.0/dist/tf.min.js"></script>
    <script src="https://cdn.jsdelivr.net/npm/@tensorflow/tfjs-tflite@0.0.1-alpha.6/dist/tf-tflite.min.js"></script>
  </head>
  <body>
    <h1>Image Transform with TensorFlow.js</h1>
    <!-- Load the image to be transformed -->
    <img id="input-image" src="https://upload.wikimedia.org/wikipedia/en/7/7d/Lenna_%28test_image%29.png" crossorigin="anonymous" alt="Original image">
    <br>
    <!-- Button to trigger image transform -->
    <button id="transform-button">Transform</button>
    <!-- Display the transformed image -->
    <h1>Transformed Image</h1>
    <canvas id="output-image"></canvas>
    <!-- Script to perform the transformation -->
    <script>
      const transformButton = document.getElementById('transform-button');
      const inputImage = document.getElementById('input-image');
      const canvas = document.getElementById('output-image');
      async function transform(inputImageData) {
        const model = await tflite.loadTFLiteModel('https://storage.googleapis.com/tfhub-lite-models/google/lite-model/seefood/segmenter/mobile_food_segmenter_V1/1.tflite');
        // Loop through the model's input tensors
        for (const input of model.inputs) {
          console.log(`Input Name: ${input.name}`);
          console.log(`Input Shape: ${input.shape}`);
          console.log(`Input Type: ${input.dtype}`);
          console.log();
        }
        // Loop through the model's output tensors
        for (const output of model.outputs) {
          console.log(`Output Name: ${output.name}`);
          console.log(`Output Shape: ${output.shape}`);
          console.log(`Output Type: ${output.dtype}`);
          console.log();
        }
        // Transform
        let inputTensor = tf.browser.fromPixels(inputImageData);
        inputTensor = tf.image.resizeBilinear(inputTensor, [
          513,
          513
        ]).expandDims(0);
        inputTensor = tf.cast (inputTensor, 'int32')
        
        // transform
        let outputTensor = model.predict(inputTensor);
        inputTensor = inputTensor.squeeze(0);
        outputTensor = outputTensor.squeeze(0);
        
        //merge output channels
        // TODO
        let scalarWeight = 0.5;
        let averageOutput = outputTensor.mul(scalarWeight).max(2);
        averageOutput = averageOutput.expandDims(-1);
        
        // merge input and output
        const inputNorm = inputTensor.div(255);
        const outputNorm = averageOutput.div(255);
        console.log(inputNorm.shape);
        console.log(outputNorm.shape);
        let combined = await inputNorm.mul(outputNorm);
        // Create the imageData object
        let dataArray = await tf.browser.toPixels(combined);
        const outputImageData = new ImageData(dataArray, combined.shape[1], combined.shape[0]);
        return outputImageData;
      }
      // When the button is clicked, update the image
      transformButton.addEventListener('click', async () => {
        // Get the canvas element
        const canvas = document.querySelector("canvas");
        const ctx = canvas.getContext('2d');
        // Set the canvas size to the size of the original image
        canvas.width = inputImage.width;
        canvas.height = inputImage.height;
        // Draw the original image on the canvas
        ctx.drawImage(inputImage, 0, 0);
        let imageData = ctx.getImageData(0, 0, canvas.width, canvas.height);
        const result = await transform(imageData);
        // Update the image data on the canvas
        canvas.width = result.width;
        canvas.height = result.height;
        ctx.putImageData(result, 0, 0);
      });
    </script>
  </body>
</html>

问题分析与修复步骤

1. 输入数据类型与归一化错误

该模型要求输入为float32类型,且像素值需归一化到[0,1]区间。原代码将输入转为int32且未做归一化,导致模型输出异常。
修复:

inputTensor = tf.cast(inputTensor, 'float32').div(255);

2. 输出掩码提取错误

模型输出是二分类概率图:通道0对应背景,通道1对应食品区域。原代码直接对所有通道取max,无法正确提取食品掩码。
修复:

// 提取食品区域的概率掩码
let foodMask = outputTensor.slice([0,0,1], [-1,-1,1]);
// 转为3通道,方便和原图融合
foodMask = foodMask.repeat(3, 2);

3. 图像融合逻辑错误

原代码用inputNorm.mul(outputNorm)会让非掩码区域变为0(全黑),正确的叠加方式应保留原图并高亮掩码区域。
修复(示例:给食品区域添加半透明高亮):

const maskOverlay = foodMask.mul(0.3); // 掩码透明度设为0.3
let combined = inputTensor.add(maskOverlay).clipByValue(0, 1); // 确保像素值在0-1之间

4. 输出尺寸匹配问题

模型输出为513x513,需将结果resize回原始图像尺寸,避免显示比例异常。
修复:

combined = tf.image.resizeBilinear(combined, [originalHeight, originalWidth]);

5. 模型重复加载优化

每次点击按钮都重新加载模型会浪费性能,应改为页面加载时预加载一次。


修复后的完整代码

<!DOCTYPE html>
<html>
  <head>
    <title>食品分割与图像融合</title>
    <script src="https://cdn.jsdelivr.net/npm/@tensorflow/tfjs@4.1.0/dist/tf.min.js"></script>
    <script src="https://cdn.jsdelivr.net/npm/@tensorflow/tfjs-tflite@0.0.1-alpha.6/dist/tf-tflite.min.js"></script>
  </head>
  <body>
    <h1>食品分割演示</h1>
    <img id="input-image" src="https://upload.wikimedia.org/wikipedia/en/7/7d/Lenna_%28test_image%29.png" crossorigin="anonymous" alt="测试图像">
    <br>
    <button id="transform-button">开始分割</button>
    <h1>融合结果</h1>
    <canvas id="output-image"></canvas>
    <script>
      const transformButton = document.getElementById('transform-button');
      const inputImage = document.getElementById('input-image');
      const canvas = document.getElementById('output-image');
      let model = null;

      // 预加载模型
      async function loadModel() {
        model = await tflite.loadTFLiteModel('https://storage.googleapis.com/tfhub-lite-models/google/lite-model/seefood/segmenter/mobile_food_segmenter_V1/1.tflite');
        console.log('模型加载完成');
      }

      async function transform(inputImageData, originalWidth, originalHeight) {
        if (!model) await loadModel();

        // 处理输入:resize到模型要求的513x513,转float32并归一化到[0,1]
        let inputTensor = tf.browser.fromPixels(inputImageData);
        inputTensor = tf.image.resizeBilinear(inputTensor, [513, 513]).expandDims(0);
        inputTensor = tf.cast(inputTensor, 'float32').div(255);

        // 推理得到输出
        let outputTensor = model.predict(inputTensor);
        inputTensor = inputTensor.squeeze(0);
        outputTensor = outputTensor.squeeze(0);

        // 提取食品区域掩码(通道1是食品,通道0是背景)
        let foodMask = outputTensor.slice([0,0,1], [-1,-1,1]);
        // 将掩码转成3通道,方便和原图融合
        foodMask = foodMask.repeat(3, 2);

        // 融合逻辑:给食品区域添加半透明高亮
        const maskOverlay = foodMask.mul(0.3);
        let combined = inputTensor.add(maskOverlay).clipByValue(0, 1);

        // 将结果resize回原始图像尺寸
        combined = tf.image.resizeBilinear(combined, [originalHeight, originalWidth]);

        // 转成ImageData返回
        let dataArray = await tf.browser.toPixels(combined);
        const outputImageData = new ImageData(dataArray, originalWidth, originalHeight);
        
        // 清理张量避免内存泄漏
        tf.dispose([inputTensor, outputTensor, foodMask, maskOverlay, combined]);
        return outputImageData;
      }

      transformButton.addEventListener('click', async () => {
        const ctx = canvas.getContext('2d');
        canvas.width = inputImage.width;
        canvas.height = inputImage.height;
        ctx.drawImage(inputImage, 0, 0);
        const imageData = ctx.getImageData(0, 0, canvas.width, canvas.height);
        
        const result = await transform(imageData, inputImage.width, inputImage.height);
        ctx.putImageData(result, 0, 0);
      });

      // 页面加载时预加载模型
      window.onload = loadModel;
    </script>
  </body>
</html>

额外说明

参考实现中的segmentationMap是tfjs-models中deeplab封装库返回的结构化结果,而你直接使用的是原始TFLite模型,因此需要手动处理输出张量,提取对应通道的掩码。

内容的提问来源于stack exchange,提问作者Hugh Pearse

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 19:10:31