确保图片压缩都采用F.interpolate(),这样和安卓端得到的结果基本差不多。

This commit is contained in:
2025-10-10 10:58:59 +08:00
parent 86a4f1ff67
commit 9eed8acb71
4 changed files with 17 additions and 13 deletions
+11 -7
View File
@@ -56,11 +56,9 @@ class ResNet50EmbeddingNet(nn.Module):
)
# 图片预处理变换(用于推理)
# TODO 这里应该要修改归一化的逻辑,不然不一样。
# 使用与移动端一致的路径:仅转换为Tensor,缩放与归一化在 preprocess_image 中用 F.interpolate 完成
self.preprocess = transforms.Compose([
transforms.Resize((224, 224), interpolation=transforms.InterpolationMode.BILINEAR),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
transforms.ToTensor()
])
# 内部预处理变换(用于已经是tensor但未归一化的数据)
@@ -88,11 +86,17 @@ class ResNet50EmbeddingNet(nn.Module):
else:
raise ValueError("输入必须是PIL Image或numpy array")
# 应用预处理变换
# 应用预处理变换(先转为tensor,值范围0-1
tensor = self.preprocess(image)
# 添加batch维度
# 添加batch维度 [1, C, H, W]
tensor = tensor.unsqueeze(0)
# 使用与移动端一致的双线性插值缩放到224x224
tensor = F.interpolate(tensor, size=(224, 224), mode='bilinear', align_corners=False)
# ImageNet标准化(与移动端保持一致)
mean = torch.tensor([0.485, 0.456, 0.406], device=tensor.device).view(1, 3, 1, 1)
std = torch.tensor([0.229, 0.224, 0.225], device=tensor.device).view(1, 3, 1, 1)
tensor = (tensor - mean) / std
return tensor
def mobile_preprocess(self, x: torch.Tensor) -> torch.Tensor: