Files
FoodClassifier/toAndroid/toAndroidSearch.py
T
2025-10-17 18:24:04 +08:00

223 lines
8.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import os
import sys
import json
import math
from typing import List, Tuple, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
# 添加项目根目录到路径(与 toAndroidEmbedding.py 保持一致)
# current_dir = os.path.dirname(os.path.abspath(__file__))
# parent_dir = os.path.dirname(current_dir)
# sys.path.insert(0, parent_dir)
print(sys.path)
def _load_embeddings_json(json_path: str) -> torch.Tensor:
"""
读取 ../faiss_vector_db/embeddings.json,尽量兼容常见格式,并返回 [N, 512] 的 float32 Tensor。
兼容格式示例:
1) [ [float]*512, [float]*512, ... ]
2) [ {"embedding": [float]*512, ...}, {"embedding": [...]}, ... ]
3) {"embeddings": [ [float]*512, ... ] }
4) {"items": [ {"embedding": [...]}, ... ] }
如维度不为512或为空,抛出异常。
"""
if not os.path.exists(json_path):
raise FileNotFoundError(f"未找到 embeddings.json: {json_path}")
with open(json_path, "r", encoding="utf-8") as f:
data = json.load(f)
def extract_vectors(obj: Union[list, dict]) -> List[List[float]]:
if isinstance(obj, list):
# list of vectors or list of dicts
if len(obj) == 0:
return []
if isinstance(obj[0], list):
return obj
if isinstance(obj[0], dict):
vecs = []
for item in obj:
if "embedding" in item and isinstance(item["embedding"], list):
vecs.append(item["embedding"])
return vecs
return []
elif isinstance(obj, dict):
# possible keys: "embeddings", "items", "data"
for key in ("embeddings", "items", "data", "vectors"):
if key in obj:
inner = obj[key]
return extract_vectors(inner)
# maybe dict of id->vector
# ensure values are lists
values = list(obj.values())
if values and isinstance(values[0], list):
return values
if values and isinstance(values[0], dict) and "embedding" in values[0]:
return [v["embedding"] for v in values if "embedding" in v]
return []
else:
return []
vectors = extract_vectors(data)
if not vectors:
raise ValueError("embeddings.json 未解析到任何向量,请确认文件结构。")
# 验证维度
dim = len(vectors[0])
if dim != 512:
raise ValueError(f"向量维度为 {dim},与期望的512不符。请检查 embeddings.json。")
# 转为 Tensor
emb_tensor = torch.tensor(vectors, dtype=torch.float32) # [N, 512]
if emb_tensor.ndim != 2 or emb_tensor.shape[1] != 512:
raise ValueError("向量张量形状异常,期望 [N, 512]。")
return emb_tensor
class TwoLayerSimilarityNet(nn.Module):
"""
两层网络:
- layer1: Linear(512->512) 无偏置,初始化为恒等(可视作占位层,满足“两层”要求)
- layer2: Linear(512->N) 无偏置,权重为数据库向量(每行一个向量),用于点积相似度
为近似余弦相似度:
- 构造时将数据库向量 L2 归一化后作为 layer2.weight
- 前向时对输入向量做 L2 归一化,再依次通过 layer1、layer2
"""
def __init__(self, db_vectors: torch.Tensor):
"""
db_vectors: [N, 512] float32
"""
super().__init__()
if db_vectors.ndim != 2 or db_vectors.shape[1] != 512:
raise ValueError("db_vectors 形状必须为 [N, 512]")
N = db_vectors.shape[0]
# (已移除恒等层,直接使用单层相似度投影)
# 单层:权重为数据库向量(行表示样本),无偏置
self.layer2 = nn.Linear(512, N, bias=False)
# 归一化数据库向量 -> 近似余弦相似度
# db_norm = F.normalize(db_vectors, p=2, dim=1) # [N, 512], 每行单位范数
with torch.no_grad():
# nn.Linear 的 weight 形状为 [out_features, in_features] = [N, 512]
# self.layer2.weight.copy_(db_norm)
self.layer2.weight.copy_(db_vectors)
# 冻结所有参数(移动端推理仅用作检索头)
for p in self.parameters():
p.requires_grad = False
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
x: [B, 512] float32
返回: [B, N] 相似度(点积,因归一化近似余弦)
"""
if x.ndim != 2 or x.shape[1] != 512:
raise ValueError("输入张量形状必须为 [batch, 512]")
# 输入归一化
x_norm = F.normalize(x, p=2, dim=1) # [B, 512]
# 单层线性变换
sims = self.layer2(x_norm) # [B, N]
# sims = self.layer2(x) # [B, N]
return sims
def build_and_export_similarity_head(
embeddings_json_path: str = "faiss_vector_db/embeddings.json",
output_path: str = "faiss_vector_db/similarity_head.pt",
) -> Tuple[str, int]:
"""
构建两层相似度网络并导出 TorchScript。
返回 (输出路径, 样本数)
"""
print("开始构建相似度检索头(TwoLayerSimilarityNet...")
# 1) 加载原始向量
db_vectors = _load_embeddings_json(embeddings_json_path) # [N, 512]
N = db_vectors.shape[0]
print(f"✓ 已加载数据库向量,样本数: {N},维度: 512")
# 2) 构建网络
model = TwoLayerSimilarityNet(db_vectors)
model.eval()
print("✓ 单层检索头构建完成(layer2=向量库权重,无偏置)")
# 3) 进行一次简单测试
with torch.no_grad():
test_input = torch.rand(2, 512, dtype=torch.float32)
test_output = model(test_input)
print(f"✓ 前向测试成功,输入形状: {tuple(test_input.shape)},输出形状: {tuple(test_output.shape)}")
print(f" 输出值范围: [{test_output.min().item():.4f}, {test_output.max().item():.4f}]")
# 4) 导出为 TorchScript(使用 trace,输入为 [1, 512]
try:
example = torch.rand(1, 512, dtype=torch.float32)
traced = torch.jit.trace(model, example)
traced.save(output_path)
print(f"✓ TorchScript 模型已保存到: {output_path}")
# 验证加载
loaded = torch.jit.load(output_path)
with torch.no_grad():
out = loaded(example)
if out.shape != (1, N):
raise RuntimeError("加载后的 TorchScript 输出形状异常")
print("✓ 导出模型验证通过")
except Exception as e:
print(f"✗ TorchScript 导出失败: {e}")
raise
# 5) 信息统计
model_size_mb = os.path.getsize(output_path) / (1024 * 1024.0)
total_params = sum(p.numel() for p in model.parameters())
print("\n" + "=" * 60)
print("检索头模型信息:")
print("=" * 60)
print(f"样本数 (输出维度): {N}")
print(f"模型文件大小: {model_size_mb:.2f} MB")
print(f"参数总数: {total_params:,} (全部冻结)")
print("=" * 60)
return output_path, N
def main():
print("开始转换检索头为移动端格式(TorchScript...")
embeddings_path = "../faiss_vector_db/DishClassification/faiss_index/embeddings.json"
output_path = "../faiss_vector_db/DishClassification/faiss_index/similarity_head.pt"
try:
path, n_items = build_and_export_similarity_head(embeddings_path, output_path)
print("\nAndroid端使用说明:")
print("1) 输入为 [batch, 512] 的 float32 张量(由移动端Embedding模型输出)")
print("2) 输出为 [batch, N] 的相似度,其中 N 为向量库样本数")
print("3) 打包到APK后,使用 PyTorch Mobile 加载 similarity_head.pt 进行推理")
print("4) 如需做 Top-K 检索,取每行的最大/Top-K即可")
print("\n示例(Kotlin/Java 伪代码):")
print(" val module = LiteModuleLoader.load(assetFilePath(\"similarity_head.pt\"))")
print(" val inputTensor = Tensor.fromBlob(inputArray, longArrayOf(batch.toLong(), 512L))")
print(" val outputTensor = module.forward(IValue.from(inputTensor)).toTensor()")
print(" val sims = outputTensor.dataAsFloatArray // 长度 = batch * N")
print("\n完成。")
except Exception as e:
print(f"✗ 转换失败: {e}")
import traceback
traceback.print_exc()
main()