| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209 |
- """
- 模型表示提取器接口
- 功能:
- - 从预存的表示文件中加载 H 矩阵
- - 批量加载多个模型的表示
- - 将表示保存到文件
- 数据格式约定:
- - 表示文件:.pt 或 .npy 格式,存储 H ∈ R^{d×N}
- - 元数据文件:.json,存储模型名、数据集名、提取层等信息
- """
- import torch
- import json
- import os
- from pathlib import Path
- from typing import Dict, List, Optional, Tuple, Any
- from dataclasses import dataclass
- @dataclass
- class RepresentationMeta:
- """表示矩阵的元数据"""
- model_name: str
- dataset_name: str
- layer_name: str
- d: int # 特征维度
- N: int # 样本数
- dtype: str
- extra: Dict[str, Any] = None
- class RepresentationLoader:
- """
- 从预存文件加载模型表示
- """
- def __init__(self, representations_dir: str = "model/representations"):
- self.base_dir = Path(representations_dir)
- self.base_dir.mkdir(parents=True, exist_ok=True)
- def get_representation_path(
- self,
- model_name: str,
- dataset_name: str,
- suffix: str = ".pt"
- ) -> Path:
- """获取表示文件路径"""
- filename = f"{model_name}_{dataset_name}{suffix}"
- return self.base_dir / filename
- def load_representation(
- self,
- model_name: str,
- dataset_name: str,
- suffix: str = ".pt"
- ) -> torch.Tensor:
- """
- 加载表示矩阵 H ∈ R^{d×N}
- Args:
- model_name: 模型标识
- dataset_name: 数据集标识
- suffix: 文件后缀 (.pt 或 .npy)
- Returns:
- H: 表示矩阵,shape = (d, N)
- """
- path = self.get_representation_path(model_name, dataset_name, suffix)
- if suffix == ".pt":
- # 兼容旧版本 PyTorch(不支持 weights_only 参数)
- try:
- H = torch.load(path, weights_only=True)
- except TypeError:
- H = torch.load(path)
- elif suffix == ".npy":
- try:
- H = torch.from_numpy(torch.load(path, weights_only=True))
- except TypeError:
- H = torch.from_numpy(torch.load(path))
- else:
- raise ValueError(f"Unsupported suffix: {suffix}")
- # 确保 shape = (d, N)
- if H.ndim != 2:
- raise ValueError(f"H 必须是 2D 矩阵,当前 shape: {H.shape}")
- return H
- def load_multiple_representations(
- self,
- model_dataset_pairs: List[Tuple[str, str]],
- suffix: str = ".pt"
- ) -> Dict[str, torch.Tensor]:
- """
- 批量加载多个表示
- Args:
- model_dataset_pairs: [(model_name, dataset_name), ...]
- Returns:
- {(model, dataset): H, ...}
- """
- result = {}
- for model_name, dataset_name in model_dataset_pairs:
- key = f"{model_name}__{dataset_name}"
- H = self.load_representation(model_name, dataset_name, suffix)
- result[key] = H
- return result
- def save_representation(
- self,
- H: torch.Tensor,
- model_name: str,
- dataset_name: str,
- suffix: str = ".pt",
- meta: Optional[RepresentationMeta] = None
- ) -> None:
- """
- 保存表示矩阵到文件
- Args:
- H: 表示矩阵
- model_name: 模型标识
- dataset_name: 数据集标识
- suffix: 文件后缀
- meta: 可选元数据
- """
- path = self.get_representation_path(model_name, dataset_name, suffix)
- torch.save(H, path)
- # 保存元数据
- if meta:
- meta_path = path.with_suffix(".json")
- with open(meta_path, "w") as f:
- json.dump({
- "model_name": meta.model_name,
- "dataset_name": meta.dataset_name,
- "layer_name": meta.layer_name,
- "d": meta.d,
- "N": meta.N,
- "dtype": meta.dtype,
- "extra": meta.extra or {}
- }, f, indent=2)
- def list_available_representations(self) -> List[Dict[str, str]]:
- """列出所有可用的表示文件"""
- available = []
- for path in self.base_dir.glob("*.pt"):
- parts = path.stem.split("_", 1)
- if len(parts) == 2:
- available.append({
- "model_name": parts[0],
- "dataset_name": parts[1],
- "path": str(path)
- })
- return available
- # ============= 便捷函数 =============
- def load_H(
- model_name: str,
- dataset_name: str,
- representations_dir: str = "model/representations",
- suffix: str = ".pt"
- ) -> torch.Tensor:
- """便捷函数:加载单个表示"""
- loader = RepresentationLoader(representations_dir)
- return loader.load_representation(model_name, dataset_name, suffix)
- def load_all_H(
- model_names: List[str],
- dataset_name: str,
- representations_dir: str = "model/representations",
- suffix: str = ".pt"
- ) -> Dict[str, torch.Tensor]:
- """便捷函数:加载多个模型在同一数据集上的表示"""
- loader = RepresentationLoader(representations_dir)
- pairs = [(m, dataset_name) for m in model_names]
- return loader.load_multiple_representations(pairs, suffix)
- def save_H(
- H: torch.Tensor,
- model_name: str,
- dataset_name: str,
- representations_dir: str = "model/representations",
- suffix: str = ".pt",
- meta: Optional[Dict] = None
- ) -> None:
- """便捷函数:保存表示"""
- loader = RepresentationLoader(representations_dir)
- if meta:
- meta_obj = RepresentationMeta(
- model_name=model_name,
- dataset_name=dataset_name,
- layer_name=meta.get("layer_name", "last"),
- d=H.shape[0],
- N=H.shape[1],
- dtype=str(H.dtype),
- extra=meta.get("extra")
- )
- else:
- meta_obj = None
- loader.save_representation(H, model_name, dataset_name, suffix, meta_obj)
|