extractor.py 6.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209
  1. """
  2. 模型表示提取器接口
  3. 功能:
  4. - 从预存的表示文件中加载 H 矩阵
  5. - 批量加载多个模型的表示
  6. - 将表示保存到文件
  7. 数据格式约定:
  8. - 表示文件:.pt 或 .npy 格式,存储 H ∈ R^{d×N}
  9. - 元数据文件:.json,存储模型名、数据集名、提取层等信息
  10. """
  11. import torch
  12. import json
  13. import os
  14. from pathlib import Path
  15. from typing import Dict, List, Optional, Tuple, Any
  16. from dataclasses import dataclass
  17. @dataclass
  18. class RepresentationMeta:
  19. """表示矩阵的元数据"""
  20. model_name: str
  21. dataset_name: str
  22. layer_name: str
  23. d: int # 特征维度
  24. N: int # 样本数
  25. dtype: str
  26. extra: Dict[str, Any] = None
  27. class RepresentationLoader:
  28. """
  29. 从预存文件加载模型表示
  30. """
  31. def __init__(self, representations_dir: str = "model/representations"):
  32. self.base_dir = Path(representations_dir)
  33. self.base_dir.mkdir(parents=True, exist_ok=True)
  34. def get_representation_path(
  35. self,
  36. model_name: str,
  37. dataset_name: str,
  38. suffix: str = ".pt"
  39. ) -> Path:
  40. """获取表示文件路径"""
  41. filename = f"{model_name}_{dataset_name}{suffix}"
  42. return self.base_dir / filename
  43. def load_representation(
  44. self,
  45. model_name: str,
  46. dataset_name: str,
  47. suffix: str = ".pt"
  48. ) -> torch.Tensor:
  49. """
  50. 加载表示矩阵 H ∈ R^{d×N}
  51. Args:
  52. model_name: 模型标识
  53. dataset_name: 数据集标识
  54. suffix: 文件后缀 (.pt 或 .npy)
  55. Returns:
  56. H: 表示矩阵,shape = (d, N)
  57. """
  58. path = self.get_representation_path(model_name, dataset_name, suffix)
  59. if suffix == ".pt":
  60. # 兼容旧版本 PyTorch(不支持 weights_only 参数)
  61. try:
  62. H = torch.load(path, weights_only=True)
  63. except TypeError:
  64. H = torch.load(path)
  65. elif suffix == ".npy":
  66. try:
  67. H = torch.from_numpy(torch.load(path, weights_only=True))
  68. except TypeError:
  69. H = torch.from_numpy(torch.load(path))
  70. else:
  71. raise ValueError(f"Unsupported suffix: {suffix}")
  72. # 确保 shape = (d, N)
  73. if H.ndim != 2:
  74. raise ValueError(f"H 必须是 2D 矩阵,当前 shape: {H.shape}")
  75. return H
  76. def load_multiple_representations(
  77. self,
  78. model_dataset_pairs: List[Tuple[str, str]],
  79. suffix: str = ".pt"
  80. ) -> Dict[str, torch.Tensor]:
  81. """
  82. 批量加载多个表示
  83. Args:
  84. model_dataset_pairs: [(model_name, dataset_name), ...]
  85. Returns:
  86. {(model, dataset): H, ...}
  87. """
  88. result = {}
  89. for model_name, dataset_name in model_dataset_pairs:
  90. key = f"{model_name}__{dataset_name}"
  91. H = self.load_representation(model_name, dataset_name, suffix)
  92. result[key] = H
  93. return result
  94. def save_representation(
  95. self,
  96. H: torch.Tensor,
  97. model_name: str,
  98. dataset_name: str,
  99. suffix: str = ".pt",
  100. meta: Optional[RepresentationMeta] = None
  101. ) -> None:
  102. """
  103. 保存表示矩阵到文件
  104. Args:
  105. H: 表示矩阵
  106. model_name: 模型标识
  107. dataset_name: 数据集标识
  108. suffix: 文件后缀
  109. meta: 可选元数据
  110. """
  111. path = self.get_representation_path(model_name, dataset_name, suffix)
  112. torch.save(H, path)
  113. # 保存元数据
  114. if meta:
  115. meta_path = path.with_suffix(".json")
  116. with open(meta_path, "w") as f:
  117. json.dump({
  118. "model_name": meta.model_name,
  119. "dataset_name": meta.dataset_name,
  120. "layer_name": meta.layer_name,
  121. "d": meta.d,
  122. "N": meta.N,
  123. "dtype": meta.dtype,
  124. "extra": meta.extra or {}
  125. }, f, indent=2)
  126. def list_available_representations(self) -> List[Dict[str, str]]:
  127. """列出所有可用的表示文件"""
  128. available = []
  129. for path in self.base_dir.glob("*.pt"):
  130. parts = path.stem.split("_", 1)
  131. if len(parts) == 2:
  132. available.append({
  133. "model_name": parts[0],
  134. "dataset_name": parts[1],
  135. "path": str(path)
  136. })
  137. return available
  138. # ============= 便捷函数 =============
  139. def load_H(
  140. model_name: str,
  141. dataset_name: str,
  142. representations_dir: str = "model/representations",
  143. suffix: str = ".pt"
  144. ) -> torch.Tensor:
  145. """便捷函数:加载单个表示"""
  146. loader = RepresentationLoader(representations_dir)
  147. return loader.load_representation(model_name, dataset_name, suffix)
  148. def load_all_H(
  149. model_names: List[str],
  150. dataset_name: str,
  151. representations_dir: str = "model/representations",
  152. suffix: str = ".pt"
  153. ) -> Dict[str, torch.Tensor]:
  154. """便捷函数:加载多个模型在同一数据集上的表示"""
  155. loader = RepresentationLoader(representations_dir)
  156. pairs = [(m, dataset_name) for m in model_names]
  157. return loader.load_multiple_representations(pairs, suffix)
  158. def save_H(
  159. H: torch.Tensor,
  160. model_name: str,
  161. dataset_name: str,
  162. representations_dir: str = "model/representations",
  163. suffix: str = ".pt",
  164. meta: Optional[Dict] = None
  165. ) -> None:
  166. """便捷函数:保存表示"""
  167. loader = RepresentationLoader(representations_dir)
  168. if meta:
  169. meta_obj = RepresentationMeta(
  170. model_name=model_name,
  171. dataset_name=dataset_name,
  172. layer_name=meta.get("layer_name", "last"),
  173. d=H.shape[0],
  174. N=H.shape[1],
  175. dtype=str(H.dtype),
  176. extra=meta.get("extra")
  177. )
  178. else:
  179. meta_obj = None
  180. loader.save_representation(H, model_name, dataset_name, suffix, meta_obj)