确保图片压缩都采用F.interpolate(),这样和安卓端得到的结果基本差不多。
This commit is contained in:
+11
-7
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user