""" 向量索引模块 支持 numpy 暴力检索(默认,无需 faiss)或 faiss 加速。 特征矩阵 (N, 768),L2 归一化,用内积(=余弦相似度)检索。 """ import json import numpy as np class VectorIndex: """简单的向量检索索引""" def __init__(self, features, card_ids, metas=None, use_faiss=False): """ Args: features: np.ndarray (N, 768) L2归一化 card_ids: list[str] 长度 N metas: list[dict] 长度 N(可选,每条卡牌元数据) use_faiss: 是否用 faiss 加速 """ self.features = features.astype(np.float32) self.card_ids = card_ids self.metas = metas or [None] * len(card_ids) self.use_faiss = use_faiss if use_faiss: import faiss self.index = faiss.IndexFlatIP(features.shape[1]) self.index.add(self.features) else: self.index = None def search(self, query_feats, top_k=5, filter_lang=None): """ 检索。 Args: query_feats: np.ndarray (M, 768) 或 (768,) top_k: 返回前 K 个 filter_lang: 只在某语言的卡里检索(简中/繁中/tcg us/tcg jp),None=全库 Returns: list of list of (card_id, score, meta),长度 M """ q = np.atleast_2d(query_feats.astype(np.float32)) valid = ~np.isnan(q).any(axis=1) M = q.shape[0] # 按语言过滤候选索引(只算符合语言的卡) if filter_lang and self.metas: lang_mask = np.array( [1 if (m and m.get("language") == filter_lang) else 0 for m in self.metas], dtype=bool) sub_idx = np.where(lang_mask)[0] else: sub_idx = None results = [] for i in range(M): if not valid[i]: results.append([]) continue if self.use_faiss and filter_lang is None: scores, indices = self.index.search(q[i:i+1], top_k) idxs, scs = indices[0], scores[0] else: feats = self.features[sub_idx] if sub_idx is not None else self.features sims = q[i] @ feats.T # (n_sub,) k = min(top_k, sims.shape[0]) if k <= 0: results.append([]) continue local = np.argpartition(-sims, k - 1)[:k] local = local[np.argsort(-sims[local])] scs = sims[local] idxs = (sub_idx[local] if sub_idx is not None else local) items = [(self.card_ids[int(j)], float(scs[r]), self.metas[int(j)]) for r, j in enumerate(idxs)] results.append(items) return results # ---------- 持久化 ---------- def save(self, features_path, meta_path): np.save(features_path, self.features) with open(meta_path, "w", encoding="utf-8") as f: json.dump({ "card_ids": self.card_ids, "metas": self.metas, }, f, ensure_ascii=False) @classmethod def load(cls, features_path, meta_path, use_faiss=False): features = np.load(features_path) with open(meta_path, "r", encoding="utf-8") as f: data = json.load(f) return cls(features, data["card_ids"], data.get("metas"), use_faiss=use_faiss)