增加了导出安卓的一层网络,可以避免安卓使用向量数据库。

This commit is contained in:
2025-10-09 13:29:53 +08:00
parent a14aa7f8e1
commit 065192fff8
2 changed files with 225 additions and 0 deletions
@@ -1100,6 +1100,8 @@ class EmbeddingFoodClassifierApp:
# 在FAISS索引中搜索最相似的k张图片
scores, indices = self.faiss_index.search(query_embedding, k)
print("最相似的图片索引:", indices)
print("最相似的图片分数:", scores)
# 收集相似图片的类别
similar_classes = []
+223
View File
@@ -0,0 +1,223 @@
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/faiss_index/embeddings.json"
output_path = "../faiss_vector_db/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()