Files
FoodClassifier/exp_multimodal/build_dish_name_index.py

173 lines
5.5 KiB
Python
Raw Permalink 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.
"""
构建菜品名向量索引
将labels.py中的菜品名编码为向量并构建FAISS索引(一次性任务)
"""
import argparse
import json
import os
import sys
import time
from typing import List
import faiss
import numpy as np
# 兼容:支持直接运行脚本或用 -m 模块方式运行
if __name__ == "__main__" and __package__ is None:
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from exp_multimodal.labels import build_labels
from exp_multimodal.text_embedder import OllamaEmbedder
def build_index(
dish_names: List[str],
embedder: OllamaEmbedder,
output_dir: str,
batch_size: int = 100,
) -> None:
"""
构建FAISS索引
参数:
dish_names: 菜品名列表
embedder: OllamaEmbedder实例
output_dir: 输出目录
batch_size: 批量编码大小
"""
os.makedirs(output_dir, exist_ok=True)
print(f"[BuildIndex] Total dishes={len(dish_names)} batch_size={batch_size}")
# 批量编码(带重试机制)
all_embeddings = []
for i in range(0, len(dish_names), batch_size):
batch = dish_names[i:i+batch_size]
batch_num = i//batch_size + 1
total_batches = (len(dish_names)-1)//batch_size + 1
print(f"[BuildIndex] Encoding batch {batch_num}/{total_batches} (size={len(batch)})...")
# 重试机制:最多3次,指数退避
max_retries = 3
for retry in range(max_retries):
try:
batch_embs = embedder.encode(batch)
all_embeddings.append(batch_embs)
break # 成功则跳出重试循环
except Exception as e:
if retry < max_retries - 1:
wait_time = 5 * (retry + 1) # 5s, 10s, 15s
print(f"[BuildIndex] Batch {batch_num} failed (attempt {retry+1}/{max_retries}): {e}")
print(f"[BuildIndex] Retrying in {wait_time} seconds...")
time.sleep(wait_time)
else:
print(f"[BuildIndex] Batch {batch_num} failed after {max_retries} attempts: {e}")
raise
# 合并所有向量
embeddings = np.vstack(all_embeddings)
print(f"[BuildIndex] Concatenated embeddings shape={embeddings.shape}")
# 归一化向量(用于余弦相似度)
norms = np.linalg.norm(embeddings, axis=1, keepdims=True)
embeddings = embeddings / (norms + 1e-8)
print(f"[BuildIndex] Normalized embeddings")
# 构建FAISS索引(IndexFlatIP = 内积索引,适合归一化后的向量)
dim = embeddings.shape[1]
index = faiss.IndexFlatIP(dim)
index.add(embeddings.astype(np.float32))
print(f"[BuildIndex] Built FAISS index dim={dim} ntotal={index.ntotal}")
# 保存文件
names_path = os.path.join(output_dir, "dish_names.json")
embeddings_path = os.path.join(output_dir, "dish_embeddings.npy")
index_path = os.path.join(output_dir, "faiss_index.bin")
with open(names_path, "w", encoding="utf-8") as f:
json.dump(dish_names, f, ensure_ascii=False, indent=2)
print(f"[BuildIndex] Saved dish names to {names_path}")
np.save(embeddings_path, embeddings)
print(f"[BuildIndex] Saved embeddings to {embeddings_path}")
faiss.write_index(index, index_path)
print(f"[BuildIndex] Saved FAISS index to {index_path}")
print(f"[BuildIndex] ✅ Index build complete! Output dir: {output_dir}")
def main():
ap = argparse.ArgumentParser(description="构建菜品名向量索引")
ap.add_argument(
"--mode",
choices=["dish", "whole", "processed"],
default="dish",
help="数据集模式(默认: dish"
)
ap.add_argument(
"--output",
default=None,
help="输出目录(默认: faiss_vector_db/<mode>_names"
)
ap.add_argument(
"--embedder_url",
default="http://192.168.1.250:11434",
help="Ollama服务地址"
)
ap.add_argument(
"--embedder_model",
default="quentinz/bge-large-zh-v1.5",
help="Embedding模型名称"
)
ap.add_argument(
"--batch_size",
type=int,
default=100,
help="批量编码大小"
)
args = ap.parse_args()
# 确定输出目录
if args.output:
output_dir = args.output
else:
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
output_dir = os.path.join(project_root, "faiss_vector_db", f"{args.mode}_names")
print(f"[Main] mode={args.mode} output_dir={output_dir}")
print(f"[Main] embedder_url={args.embedder_url} model={args.embedder_model}")
# 构建菜品名列表
print(f"[Main] Building labels from mode={args.mode}...")
dish_names = build_labels(args.mode, alias_map_path=None)
print(f"[Main] Built {len(dish_names)} dish names")
if not dish_names:
print("[Main] ❌ No dish names found, abort")
return
# 初始化Embedder
print(f"[Main] Initializing OllamaEmbedder...")
embedder = OllamaEmbedder(
base_url=args.embedder_url,
model=args.embedder_model,
timeout=180, # 3分钟超时
)
# 构建索引
print(f"[Main] Building FAISS index...")
build_index(
dish_names=dish_names,
embedder=embedder,
output_dir=output_dir,
batch_size=args.batch_size,
)
print(f"[Main] ✅ All done!")
if __name__ == "__main__":
main()