确保图片压缩都采用F.interpolate(),这样和安卓端得到的结果基本差不多。
This commit is contained in:
@@ -468,7 +468,7 @@ class FAISSSearcher:
|
|||||||
def main():
|
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"
|
TRAIN_DIR = "dataset/train"
|
||||||
OUTPUT_DIR = "faiss_vector_db/faiss_index"
|
OUTPUT_DIR = "faiss_vector_db/faiss_index"
|
||||||
BATCH_SIZE = 16
|
BATCH_SIZE = 16
|
||||||
|
|||||||
@@ -1124,7 +1124,7 @@ class EmbeddingFoodClassifierApp:
|
|||||||
query_tensor = torch.from_numpy(query_embedding.reshape(1, -1)) # [1,512]
|
query_tensor = torch.from_numpy(query_embedding.reshape(1, -1)) # [1,512]
|
||||||
# 打印query_tensor的前十个数字
|
# 打印query_tensor的前十个数字
|
||||||
|
|
||||||
print(query_embedding[:10])
|
# print(query_embedding[:50])
|
||||||
|
|
||||||
if self.db_matrix is None:
|
if self.db_matrix is None:
|
||||||
raise RuntimeError("向量库未加载")
|
raise RuntimeError("向量库未加载")
|
||||||
@@ -1136,8 +1136,8 @@ class EmbeddingFoodClassifierApp:
|
|||||||
topk_scores, topk_indices = torch.topk(sims, k=min(k, sims.shape[1]), dim=1)
|
topk_scores, topk_indices = torch.topk(sims, k=min(k, sims.shape[1]), dim=1)
|
||||||
indices = topk_indices.cpu().numpy()
|
indices = topk_indices.cpu().numpy()
|
||||||
scores = topk_scores.cpu().numpy()
|
scores = topk_scores.cpu().numpy()
|
||||||
print("相似度索引:", indices)
|
# print("相似度索引:", indices)
|
||||||
print("相似度分数:", scores)
|
# print("相似度分数:", scores)
|
||||||
|
|
||||||
# 收集相似图片的类别
|
# 收集相似图片的类别
|
||||||
similar_classes = []
|
similar_classes = []
|
||||||
|
|||||||
+11
-7
@@ -56,11 +56,9 @@ class ResNet50EmbeddingNet(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 图片预处理变换(用于推理)
|
# 图片预处理变换(用于推理)
|
||||||
# TODO 这里应该要修改归一化的逻辑,不然不一样。
|
# 使用与移动端一致的路径:仅转换为Tensor,缩放与归一化在 preprocess_image 中用 F.interpolate 完成
|
||||||
self.preprocess = transforms.Compose([
|
self.preprocess = transforms.Compose([
|
||||||
transforms.Resize((224, 224), interpolation=transforms.InterpolationMode.BILINEAR),
|
transforms.ToTensor()
|
||||||
transforms.ToTensor(),
|
|
||||||
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
|
|
||||||
])
|
])
|
||||||
|
|
||||||
# 内部预处理变换(用于已经是tensor但未归一化的数据)
|
# 内部预处理变换(用于已经是tensor但未归一化的数据)
|
||||||
@@ -88,11 +86,17 @@ class ResNet50EmbeddingNet(nn.Module):
|
|||||||
else:
|
else:
|
||||||
raise ValueError("输入必须是PIL Image或numpy array")
|
raise ValueError("输入必须是PIL Image或numpy array")
|
||||||
|
|
||||||
# 应用预处理变换
|
# 应用预处理变换(先转为tensor,值范围0-1)
|
||||||
tensor = self.preprocess(image)
|
tensor = self.preprocess(image)
|
||||||
# 添加batch维度
|
# 添加batch维度 [1, C, H, W]
|
||||||
tensor = tensor.unsqueeze(0)
|
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
|
return tensor
|
||||||
|
|
||||||
def mobile_preprocess(self, x: torch.Tensor) -> torch.Tensor:
|
def mobile_preprocess(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ def main():
|
|||||||
# 1. 加载训练好的embedding模型权重
|
# 1. 加载训练好的embedding模型权重
|
||||||
# base_model = create_resnet50_embedding(embedding_dim=512, pretrained=True)
|
# base_model = create_resnet50_embedding(embedding_dim=512, pretrained=True)
|
||||||
base_model = create_mobile_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):
|
if not os.path.exists(model_path):
|
||||||
print(f"错误:模型文件不存在 {model_path}")
|
print(f"错误:模型文件不存在 {model_path}")
|
||||||
@@ -85,7 +85,7 @@ def main():
|
|||||||
traced_model = torch.jit.trace(mobile_wrapper, single_input)
|
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)
|
traced_model.save(output_path)
|
||||||
print(f"✓ TorchScript模型保存成功: {output_path}")
|
print(f"✓ TorchScript模型保存成功: {output_path}")
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user