vector_index.py 3.4 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394
  1. """
  2. 向量索引模块
  3. 支持 numpy 暴力检索(默认,无需 faiss)或 faiss 加速。
  4. 特征矩阵 (N, 768),L2 归一化,用内积(=余弦相似度)检索。
  5. """
  6. import json
  7. import numpy as np
  8. class VectorIndex:
  9. """简单的向量检索索引"""
  10. def __init__(self, features, card_ids, metas=None, use_faiss=False):
  11. """
  12. Args:
  13. features: np.ndarray (N, 768) L2归一化
  14. card_ids: list[str] 长度 N
  15. metas: list[dict] 长度 N(可选,每条卡牌元数据)
  16. use_faiss: 是否用 faiss 加速
  17. """
  18. self.features = features.astype(np.float32)
  19. self.card_ids = card_ids
  20. self.metas = metas or [None] * len(card_ids)
  21. self.use_faiss = use_faiss
  22. if use_faiss:
  23. import faiss
  24. self.index = faiss.IndexFlatIP(features.shape[1])
  25. self.index.add(self.features)
  26. else:
  27. self.index = None
  28. def search(self, query_feats, top_k=5, filter_lang=None):
  29. """
  30. 检索。
  31. Args:
  32. query_feats: np.ndarray (M, 768) 或 (768,)
  33. top_k: 返回前 K 个
  34. filter_lang: 只在某语言的卡里检索(简中/繁中/tcg us/tcg jp),None=全库
  35. Returns:
  36. list of list of (card_id, score, meta),长度 M
  37. """
  38. q = np.atleast_2d(query_feats.astype(np.float32))
  39. valid = ~np.isnan(q).any(axis=1)
  40. M = q.shape[0]
  41. # 按语言过滤候选索引(只算符合语言的卡)
  42. if filter_lang and self.metas:
  43. lang_mask = np.array(
  44. [1 if (m and m.get("language") == filter_lang) else 0 for m in self.metas],
  45. dtype=bool)
  46. sub_idx = np.where(lang_mask)[0]
  47. else:
  48. sub_idx = None
  49. results = []
  50. for i in range(M):
  51. if not valid[i]:
  52. results.append([])
  53. continue
  54. if self.use_faiss and filter_lang is None:
  55. scores, indices = self.index.search(q[i:i+1], top_k)
  56. idxs, scs = indices[0], scores[0]
  57. else:
  58. feats = self.features[sub_idx] if sub_idx is not None else self.features
  59. sims = q[i] @ feats.T # (n_sub,)
  60. k = min(top_k, sims.shape[0])
  61. if k <= 0:
  62. results.append([])
  63. continue
  64. local = np.argpartition(-sims, k - 1)[:k]
  65. local = local[np.argsort(-sims[local])]
  66. scs = sims[local]
  67. idxs = (sub_idx[local] if sub_idx is not None else local)
  68. items = [(self.card_ids[int(j)], float(scs[r]), self.metas[int(j)])
  69. for r, j in enumerate(idxs)]
  70. results.append(items)
  71. return results
  72. # ---------- 持久化 ----------
  73. def save(self, features_path, meta_path):
  74. np.save(features_path, self.features)
  75. with open(meta_path, "w", encoding="utf-8") as f:
  76. json.dump({
  77. "card_ids": self.card_ids,
  78. "metas": self.metas,
  79. }, f, ensure_ascii=False)
  80. @classmethod
  81. def load(cls, features_path, meta_path, use_faiss=False):
  82. features = np.load(features_path)
  83. with open(meta_path, "r", encoding="utf-8") as f:
  84. data = json.load(f)
  85. return cls(features, data["card_ids"], data.get("metas"), use_faiss=use_faiss)