TensorFlow.js与Python TensorFlow预测结果不一致问题求助
Chrome扩展中TensorFlow.js与Python预测结果不一致的排查与修复
核心问题分析
预测结果差异本质是Python与TensorFlow.js的图像预处理流程不一致,同时存在模型加载路径错误和图像拉伸变形的问题:
- 图像缩放方式错误:直接设置
img.width/height会导致图像拉伸,破坏原始比例,和Python训练时的中心裁剪/等比例缩放逻辑不符 - 模型加载路径错误:Service Worker无法直接读取相对路径,需用
chrome.runtime.getURL获取模型资源 - 缺少归一化处理:Python训练时通常会将像素值归一化到
[0,1]或[-1,1],TF.js中未执行该步骤 - 数据传递效率问题:将
ImageData转为普通数组传递会增加内存开销,可优化为Transferable Objects
分步修复代码
1. 修复contentScript.js:正确处理图像缩放与数据传递
const IMAGE_SIZE = 224; const MIN_IMG_SIZE = 128; function loadImageAndSendDataBack(src, sendResponse) { const img = new Image(); img.crossOrigin = 'anonymous'; img.onerror = function(e) { console.warn(`无法加载外部图片: ${src}`); sendResponse({rawImageData: undefined}); return; }; img.onload = function(e) { // 检查最小尺寸 if (img.height <= MIN_IMG_SIZE || img.width <= MIN_IMG_SIZE) { console.warn(`图片尺寸过小: [${img.height} x ${img.width}],要求最小[${MIN_IMG_SIZE} x ${MIN_IMG_SIZE}]`); sendResponse({rawImageData: undefined}); return; } // 计算等比例缩放尺寸,避免拉伸,和Python训练逻辑对齐 const scale = Math.max(IMAGE_SIZE / img.width, IMAGE_SIZE / img.height); const scaledWidth = Math.floor(img.width * scale); const scaledHeight = Math.floor(img.height * scale); const offsetX = Math.floor((scaledWidth - IMAGE_SIZE) / 2); const offsetY = Math.floor((scaledHeight - IMAGE_SIZE) / 2); const canvas = new OffscreenCanvas(IMAGE_SIZE, IMAGE_SIZE); const ctx = canvas.getContext('2d'); // 先缩放到足够大,再裁剪中心区域 ctx.drawImage(img, offsetX, offsetY, IMAGE_SIZE, IMAGE_SIZE, 0, 0, IMAGE_SIZE, IMAGE_SIZE); const imageData = ctx.getImageData(0, 0, IMAGE_SIZE, IMAGE_SIZE); // 使用Transferable Objects传递数据,减少内存开销 sendResponse({ rawImageData: imageData.data, width: IMAGE_SIZE, height: IMAGE_SIZE, }, [imageData.data.buffer]); }; img.src = src; } let changeImg = function() { let images = document.getElementsByTagName("img"); for (let i = 0; i < images.length; i++) { let src = images[i].src; loadImageAndSendDataBack(src, function(data){ chrome.runtime.sendMessage({ msg: "image", index: i, data: data}, () => {}); }); } } changeImg();
2. 修复background.js:正确加载模型与图像预处理
import * as tf from '@tensorflow/tfjs'; let model; const loadModel = async () => { console.log("加载模型中..."); const startTime = performance.now(); try { // 正确获取模型的绝对路径 const modelUrl = chrome.runtime.getURL('src/model/model.json'); model = await tf.loadLayersModel(modelUrl); const totalTime = Math.floor(performance.now() - startTime); console.log(`模型加载完成,耗时 ${totalTime} ms`); } catch (e) { console.error('模型加载失败', e); } }; // 单独注册消息监听,避免重复绑定 chrome.runtime.onMessage.addListener((message, sender, sendResponse) => { if (!model) { console.warn('模型尚未加载完成'); sendResponse({success: false}); return true; // 保持消息通道开放 } if (message.msg === "image" && message.data?.rawImageData) { tf.tidy(() => { const imageData = new ImageData( new Uint8ClampedArray(message.data.rawImageData), message.data.width, message.data.height ); // 执行与Python一致的预处理流程 let tensor = tf.browser.fromPixels(imageData); tensor = tf.expandDims(tensor); // 归一化到[0,1],匹配Keras训练时的rescale=1/255配置 tensor = tensor.div(255.0); const prediction = model.predict(tensor); prediction.print(); sendResponse({success: true, prediction: prediction.dataSync()}); }); return true; // 异步响应需返回true } }); loadModel();
3. 验证manifest.json配置
确保模型资源路径与实际文件位置匹配,当前配置已满足需求:
"web_accessible_resources": [ { "resources": ["src/model/model.json", "src/model/group1-shard1of3.bin", "src/model/group1-shard2of3.bin", "src/model/group1-shard3of3.bin"], "matches": ["<all_urls>"] } ]
额外验证步骤
- 对比输入张量:在Python中打印预处理后的图像张量,在TF.js中用
tensor.print()输出,确保数值范围、形状完全一致 - 检查模型转换:使用
tfjs-converter转换Keras模型时,确保指定--input_shape参数匹配输入尺寸 - 跨域图片测试:确认目标网站允许跨域访问图片,或在manifest中添加
"cross-origin-isolation": "use-site-per-process"(如需要)
内容的提问来源于stack exchange,提问作者John Jam
相关产品推荐
相关产品推荐

