| 12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394 |
- """
- 向量索引模块
- 支持 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)
|