确保图片压缩都采用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
+1 -1
View File
@@ -468,7 +468,7 @@ class FAISSSearcher:
def main():
"""主函数"""
# 配置参数
MODEL_PATH = "model/embedding_20250917_145342/best_embedding_model.pth"
MODEL_PATH = "model/embedding_20250930_102826/best_embedding_model.pth"
TRAIN_DIR = "dataset/train"
OUTPUT_DIR = "faiss_vector_db/faiss_index"
BATCH_SIZE = 16
@@ -1124,7 +1124,7 @@ class EmbeddingFoodClassifierApp:
query_tensor = torch.from_numpy(query_embedding.reshape(1, -1)) # [1,512]
# 打印query_tensor的前十个数字
print(query_embedding[:10])
# print(query_embedding[:50])
if self.db_matrix is None:
raise RuntimeError("向量库未加载")
@@ -1136,8 +1136,8 @@ class EmbeddingFoodClassifierApp:
topk_scores, topk_indices = torch.topk(sims, k=min(k, sims.shape[1]), dim=1)
indices = topk_indices.cpu().numpy()
scores = topk_scores.cpu().numpy()
print("相似度索引:", indices)
print("相似度分数:", scores)
# print("相似度索引:", indices)
# print("相似度分数:", scores)
# 收集相似图片的类别
similar_classes = []
+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:
+2 -2
View File
@@ -17,7 +17,7 @@ def main():
# 1. 加载训练好的embedding模型权重
# base_model = create_resnet50_embedding(embedding_dim=512, pretrained=True)
base_model = create_mobile_resnet50_embedding(embedding_dim=512, pretrained=True)
model_path = "../model/embedding_20250917_145342/best_embedding_model.pth"
model_path = "../model/embedding_20250930_102826/best_embedding_model.pth"
if not os.path.exists(model_path):
print(f"错误:模型文件不存在 {model_path}")
@@ -85,7 +85,7 @@ def main():
traced_model = torch.jit.trace(mobile_wrapper, single_input)
# 保存模型
output_path = "../model/embedding_20250917_145342/best_embedding_model_mobile.pt"
output_path = "../model/embedding_20250930_102826/best_embedding_model_mobile.pt"
traced_model.save(output_path)
print(f"✓ TorchScript模型保存成功: {output_path}")