""" 模型表示提取器接口 功能: - 从预存的表示文件中加载 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)