PyTorch与TensorFlow.js提取ResNet152特征嵌入结果差异排查
问题:PyTorch与TensorFlow.js提取ResNet152特征嵌入结果差异过大
我之前用PyTorch结合ResNet152提取图像特征嵌入效果良好,现在尝试用TensorFlow.js实现相同逻辑,但两者输出差异极大,甚至TensorFlow.js的结果出现大量0值,请问操作是否有误?
PyTorch实现代码及输出
import torch import torchvision.models as models from torchvision import transforms from PIL import Image # Load the model resnet152_torch = models.resnet152(pretrained=True) # 去掉最后一层全连接层,保留到平均池化层 resnet152 = torch.nn.Sequential(*(list(resnet152_torch.children())[:-1])) # 设置为评估模式 resnet152_torch.eval() # 加载并预处理图像(已为224x224) image_path = "test.png" img = Image.open(image_path).convert("RGB") preprocess = transforms.Compose([ transforms.ToTensor(), # 将0-255像素值转为0-1 transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) img_tensor = preprocess(img).unsqueeze(0) with torch.no_grad(): img_features = resnet152(img_tensor) print(img_features.squeeze())
输出:
tensor([0.2098, 0.4687, 0.0914, ..., 0.0309, 0.0919, 0.0480])
TensorFlow.js实现步骤及输出
1. 导出Keras格式ResNet152
from tensorflow.keras.applications import ResNet152 from tensorflow.keras.models import save_model # 加载预训练ResNet152模型 resnet152 = ResNet152(weights='imagenet') resnet152.trainable = False save_model(resnet152, "resnet152.h5")
2. 转换为TensorFlow.js格式
tensorflowjs_converter --input_format keras resnet152.h5 resnet152
3. JavaScript提取逻辑
import * as tf from '@tensorflow/tfjs-node'; import fs from 'fs'; async function main() { const model = await tf.loadLayersModel('file://resnet152/model.json'); const modelWithoutFinalLayer = tf.model({ inputs: model.input, outputs: model.getLayer('avg_pool').output }); const image = fs.readFileSync('example_images/test.png'); const imageTensor = tf.node.decodeImage(image, 3); const preprocessedInput = tf.div(tf.sub(imageTensor, [123.68, 116.779, 103.939]), [58.393, 57.12, 57.375]); const batchedInput = preprocessedInput.expandDims(0); const embeddings = modelWithoutFinalLayer.predict(batchedInput).squeeze(); embeddings.print(); } await main();
输出:
Tensor [0, 0, 0, ..., 0, 0, 0.029606]
问题原因及解决方案
1. 模型结构不匹配(核心问题)
PyTorch中你去掉了ResNet152的最后一层全连接(fc)层,只保留到平均池化层;但Keras导出时默认加载了带顶层fc层的完整模型,虽然你在JS中取了avg_pool的输出,但PyTorch与Keras的ResNet152在池化层后的结构细节、权重初始化存在差异,且带顶层fc层的模型可能间接影响池化层输出的特征分布。
修正方案:导出Keras模型时,明确指定include_top=False,去掉顶层fc层,直接保留到平均池化层的结构,与PyTorch对齐:
from tensorflow.keras.applications import ResNet152 from tensorflow.keras.models import save_model # 加载不带顶层全连接层的ResNet152,pooling='avg'直接输出2048维特征向量 resnet152 = ResNet152(weights='imagenet', include_top=False, pooling='avg') resnet152.trainable = False save_model(resnet152, "resnet152_no_top.h5")
2. 图像预处理细节缺失
TensorFlow.js中tf.node.decodeImage返回的是uint8类型的张量(0-255像素值),直接进行减法和除法运算会导致精度损失,需先转换为float32类型。
修正后的JavaScript代码:
import * as tf from '@tensorflow/tfjs-node'; import fs from 'fs'; async function main() { // 加载修正后的模型 const model = await tf.loadLayersModel('file://resnet152_no_top/model.json'); const imageBuffer = fs.readFileSync('example_images/test.png'); // 解码图像并转换为float32类型 const imageTensor = tf.node.decodeImage(imageBuffer, 3).cast('float32'); // Keras ResNet标准预处理:基于0-255像素值减均值、除以标准差 const mean = [123.68, 116.779, 103.939]; const std = [58.393, 57.12, 57.375]; const preprocessedInput = tf.div(tf.sub(imageTensor, mean), std); // 添加batch维度 const batchedInput = preprocessedInput.expandDims(0); const embeddings = model.predict(batchedInput).squeeze(); embeddings.print(); } await main();
3. 模型转换验证
重新执行转换命令,确保模型转换完整:
tensorflowjs_converter --input_format keras resnet152_no_top.h5 resnet152_no_top
补充说明
即使修正后,PyTorch与TensorFlow.js的输出数值仍会存在细微差异,这是因为两者的ResNet152预训练权重来自不同框架的训练实现(初始化、优化器细节等),但特征向量的语义一致性是一致的,不会影响后续的特征匹配、检索等任务。
内容的提问来源于stack exchange,提问作者Tombatron
相关产品推荐
相关产品推荐

