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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 04:59:54