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
相关产品推荐
相关产品推荐

