alignment_spectrum.py 20 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600
  1. #!/usr/bin/env python
  2. """
  3. E4 实验:对齐即谱工程
  4. 目的:验证对齐(SFT/Instruct)会使模型在指令集上的谱更集中(r_eff 下降)
  5. 假设:
  6. 1. 在指令集上:Instruct 版的 r_eff < Base 版(谱压缩)
  7. 2. 在通用集上:两者 r_eff 相近(基础能力不变)
  8. 3. 域专一化 (STG) = |r_eff_instruction - r_eff_general|,Instruct 版应更大
  9. 用法:
  10. python experiments/alignment_spectrum.py
  11. """
  12. import os
  13. import json
  14. import torch
  15. import numpy as np
  16. from pathlib import Path
  17. from typing import Dict, List, Tuple, Optional, Any
  18. from dataclasses import dataclass, asdict
  19. from datetime import datetime
  20. # 设置环境变量
  21. os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
  22. from transformers import AutoTokenizer, AutoModelForCausalLM
  23. from tqdm import tqdm
  24. # 导入谱计算模块
  25. import sys
  26. sys.path.insert(0, str(Path(__file__).parent.parent))
  27. from model.spectrum import compute_gram_spectrum, compute_spectrum_array
  28. @dataclass
  29. class E4Config:
  30. """E4 实验配置"""
  31. # 模型对:Base vs Instruct
  32. base_model_name: str = "llama3-8b"
  33. instruct_model_name: str = "llama3-8b-instruct"
  34. base_model_path: str = "model/weights/llama3-8b"
  35. instruct_model_path: str = "model/weights/llama3-8b-instruct"
  36. # 数据集
  37. instruction_dataset: str = "alpaca" # 指令集
  38. general_dataset: str = "flores200" # 通用集(对照)
  39. # 实验参数
  40. batch_size: int = 4
  41. max_length: int = 512
  42. pooling: str = "mean"
  43. sample_size: int = 200 # 每个数据集采样 200 条
  44. output_dir: str = "experiments/output/e4"
  45. device: str = "cuda"
  46. @dataclass
  47. class E4Result:
  48. """E4 实验结果"""
  49. timestamp: str
  50. config: dict
  51. # 指令集结果
  52. instruction_r_eff: Dict[str, float] # {model_name: r_eff}
  53. instruction_rs_cross: float # Base vs Instruct 的差异
  54. # 通用集结果(对照)
  55. general_r_eff: Dict[str, float]
  56. general_rs_cross: float
  57. # 域专一化指标
  58. stg: Dict[str, float] # {model_name: |r_eff_instr - r_eff_general|}
  59. # 谱距离
  60. spd_instruction: float # Base vs Instruct 在指令集上的 SPD
  61. spd_general: float # Base vs Instruct 在通用集上的 SPD
  62. # 假设验证
  63. hypothesis_support: Dict[str, Any]
  64. class AlignmentSpectrumAnalyzer:
  65. """对齐谱分析器"""
  66. def __init__(self, config: E4Config):
  67. self.config = config
  68. self.output_dir = Path(config.output_dir)
  69. self.output_dir.mkdir(parents=True, exist_ok=True)
  70. # 模型缓存
  71. self.loaded_models: Dict[str, Tuple[AutoTokenizer, AutoModelForCausalLM]] = {}
  72. def load_model(self, model_name: str) -> Tuple[AutoTokenizer, AutoModelForCausalLM]:
  73. """加载模型"""
  74. if model_name in self.loaded_models:
  75. return self.loaded_models[model_name]
  76. model_path = (self.config.base_model_path if model_name == self.config.base_model_name
  77. else self.config.instruct_model_path)
  78. print(f"加载模型:{model_name} (from {model_path})")
  79. tokenizer = AutoTokenizer.from_pretrained(
  80. model_path,
  81. trust_remote_code=True
  82. )
  83. if tokenizer.pad_token is None:
  84. tokenizer.pad_token = tokenizer.eos_token
  85. if tokenizer.eos_token_id is None:
  86. tokenizer.eos_token_id = 128009 # LLaMA3 default
  87. model = AutoModelForCausalLM.from_pretrained(
  88. model_path,
  89. torch_dtype=torch.float16,
  90. device_map="auto",
  91. output_hidden_states=True,
  92. trust_remote_code=True
  93. )
  94. model.eval()
  95. self.loaded_models[model_name] = (tokenizer, model)
  96. return tokenizer, model
  97. def unload_model(self, model_name: str):
  98. """卸载模型释放显存"""
  99. if model_name in self.loaded_models:
  100. del self.loaded_models[model_name]
  101. import gc
  102. gc.collect()
  103. if torch.cuda.is_available():
  104. torch.cuda.empty_cache()
  105. print(f"已卸载:{model_name}")
  106. def load_texts(self, dataset_name: str, limit: int = 200) -> List[str]:
  107. """加载数据集文本"""
  108. from database.corpus import CorpusLoader
  109. loader = CorpusLoader("database/corpus")
  110. try:
  111. texts = loader.load_texts(dataset_name, limit=limit)
  112. print(f"加载 {dataset_name}: {len(texts)} 条文本")
  113. return texts
  114. except FileNotFoundError:
  115. print(f"警告:数据集 {dataset_name} 不存在,使用备用数据")
  116. # 备用:生成简单指令
  117. if dataset_name == "alpaca":
  118. return [f"Instruction {i}: Please explain the concept of machine learning."
  119. for i in range(limit)]
  120. return [f"Text {i}" for i in range(limit)]
  121. @torch.no_grad()
  122. def extract_hidden_states(
  123. self,
  124. model_name: str,
  125. texts: List[str]
  126. ) -> torch.Tensor:
  127. """
  128. 提取隐层表示
  129. Returns:
  130. H: 表示矩阵 (d, N)
  131. """
  132. tokenizer, model = self.load_model(model_name)
  133. all_features = []
  134. for i in range(0, len(texts), self.config.batch_size):
  135. batch = texts[i:i + self.config.batch_size]
  136. inputs = tokenizer(
  137. batch,
  138. return_tensors="pt",
  139. padding=True,
  140. truncation=True,
  141. max_length=self.config.max_length
  142. ).to(model.device)
  143. outputs = model(**inputs)
  144. hidden = outputs.hidden_states[-1] # 最后一层
  145. attention_mask = inputs["attention_mask"]
  146. if self.config.pooling == "mean":
  147. mask_expanded = attention_mask.unsqueeze(-1).float()
  148. feat = (hidden * mask_expanded).sum(1) / (mask_expanded.sum(1) + 1e-10)
  149. else:
  150. feat = hidden[:, 0, :] # CLS token
  151. all_features.append(feat.cpu().float())
  152. # 转置为 (d, N)
  153. H = torch.cat(all_features, dim=0).T
  154. return H
  155. def compute_r_eff(self, H: torch.Tensor) -> float:
  156. """计算有效秩"""
  157. return compute_gram_spectrum(H)
  158. def compute_spectra(self, H: torch.Tensor) -> np.ndarray:
  159. """计算谱分布"""
  160. _, prob = compute_spectrum_array(H)
  161. return prob
  162. def compute_spd(self, prob1: np.ndarray, prob2: np.ndarray) -> float:
  163. """计算谱 Platonic 距离 (Wasserstein-1 距离)"""
  164. from model.spectrum import wasserstein1_distance
  165. # 对齐长度
  166. max_len = max(len(prob1), len(prob2))
  167. p1_pad = np.pad(prob1, (0, max_len - len(prob1)))
  168. p2_pad = np.pad(prob2, (0, max_len - len(prob2)))
  169. return wasserstein1_distance(p1_pad, p2_pad)
  170. def analyze_pair(
  171. self,
  172. dataset_name: str,
  173. texts: List[str]
  174. ) -> Dict[str, Any]:
  175. """
  176. 分析一对模型在特定数据集上的谱
  177. Returns:
  178. {
  179. "r_effs": {model: r_eff},
  180. "spectra": {model: spectrum},
  181. "rs_cross": float, # 方差
  182. "spd": float
  183. }
  184. """
  185. result = {"r_effs": {}, "spectra": {}}
  186. # 提取 Base 模型
  187. print(f" 提取 {self.config.base_model_name}...")
  188. H_base = self.extract_hidden_states(self.config.base_model_name, texts)
  189. r_eff_base = self.compute_r_eff(H_base)
  190. spec_base = self.compute_spectra(H_base)
  191. result["r_effs"][self.config.base_model_name] = r_eff_base
  192. result["spectra"][self.config.base_model_name] = spec_base
  193. # 卸载 Base 模型,加载 Instruct 模型
  194. self.unload_model(self.config.base_model_name)
  195. # 提取 Instruct 模型
  196. print(f" 提取 {self.config.instruct_model_name}...")
  197. H_instruct = self.extract_hidden_states(self.config.instruct_model_name, texts)
  198. r_eff_instruct = self.compute_r_eff(H_instruct)
  199. spec_instruct = self.compute_spectra(H_instruct)
  200. result["r_effs"][self.config.instruct_model_name] = r_eff_instruct
  201. result["spectra"][self.config.instruct_model_name] = spec_instruct
  202. # 计算 RS_cross (方差) 和 SPD
  203. r_eff_values = list(result["r_effs"].values())
  204. result["rs_cross"] = float(np.var(r_eff_values))
  205. result["spd"] = self.compute_spd(spec_base, spec_instruct)
  206. return result
  207. def run_full_analysis(self) -> E4Result:
  208. """运行完整的 E4 分析"""
  209. print("=" * 60)
  210. print("E4 实验:对齐即谱工程")
  211. print("=" * 60)
  212. print(f"\n模型对:{self.config.base_model_name} vs {self.config.instruct_model_name}")
  213. print(f"指令集:{self.config.instruction_dataset}")
  214. print(f"通用集:{self.config.general_dataset}")
  215. print()
  216. # 加载数据
  217. print("[1/4] 加载指令集...")
  218. instruct_texts = self.load_texts(self.config.instruction_dataset,
  219. limit=self.config.sample_size)
  220. print("[2/4] 加载通用集...")
  221. general_texts = self.load_texts(self.config.general_dataset,
  222. limit=self.config.sample_size)
  223. # 分析指令集
  224. print("\n[3/4] 分析指令集上的谱...")
  225. instruct_result = self.analyze_pair(self.config.instruction_dataset, instruct_texts)
  226. # 卸载所有模型
  227. self.unload_model(self.config.instruct_model_name)
  228. # 分析通用集
  229. print("\n[4/4] 分析通用集上的谱...")
  230. general_result = self.analyze_pair(self.config.general_dataset, general_texts)
  231. # 计算域专一化 (STG)
  232. stg = {}
  233. for model_name in [self.config.base_model_name, self.config.instruct_model_name]:
  234. r_instr = instruct_result["r_effs"][model_name]
  235. r_gen = general_result["r_effs"][model_name]
  236. stg[model_name] = abs(r_instr - r_gen)
  237. # 假设验证
  238. hypothesis = self._evaluate_hypothesis(
  239. instruct_result, general_result, stg
  240. )
  241. # 打包结果
  242. result = E4Result(
  243. timestamp=datetime.now().isoformat(),
  244. config=asdict(self.config),
  245. instruction_r_eff=instruct_result["r_effs"],
  246. instruction_rs_cross=instruct_result["rs_cross"],
  247. general_r_eff=general_result["r_effs"],
  248. general_rs_cross=general_result["rs_cross"],
  249. stg=stg,
  250. spd_instruction=instruct_result["spd"],
  251. spd_general=general_result["spd"],
  252. hypothesis_support=hypothesis
  253. )
  254. self._save_result(result)
  255. return result
  256. def _evaluate_hypothesis(
  257. self,
  258. instruct: Dict,
  259. general: Dict,
  260. stg: Dict
  261. ) -> Dict[str, Any]:
  262. """评估实验假设"""
  263. base_name = self.config.base_model_name
  264. instruct_name = self.config.instruct_model_name
  265. # H1: Instruct 版在指令集上 r_eff 更低
  266. h1_supported = (instruct["r_effs"][instruct_name] <
  267. instruct["r_effs"][base_name])
  268. # H2: Instruct 版的域专一化 (STG) 更高
  269. h2_supported = stg[instruct_name] > stg[base_name]
  270. # H3: 指令集上的谱距离 > 通用集上的谱距离
  271. # (这意味着对齐主要影响指令理解)
  272. h3_data_needed = True # 需要更多模型对才能验证
  273. return {
  274. "h1_spectrum_compression": {
  275. "supported": h1_supported,
  276. "interpretation": "Instruct 版在指令集上谱更集中" if h1_supported else "假设未获支持"
  277. },
  278. "h2_domain_specialization": {
  279. "supported": h2_supported,
  280. "interpretation": "Instruct 版域专一化程度更高" if h2_supported else "假设未获支持"
  281. },
  282. "overall": "假设获支持" if (h1_supported and h2_supported) else "部分支持"
  283. }
  284. def _save_result(self, result: E4Result):
  285. """保存结果"""
  286. result_dict = {
  287. "timestamp": result.timestamp,
  288. "config": result.config,
  289. "instruction_r_eff": result.instruction_r_eff,
  290. "instruction_rs_cross": result.instruction_rs_cross,
  291. "general_r_eff": result.general_r_eff,
  292. "general_rs_cross": result.general_rs_cross,
  293. "stg": result.stg,
  294. "spd_instruction": result.spd_instruction,
  295. "spd_general": result.spd_general,
  296. "hypothesis_support": result.hypothesis_support
  297. }
  298. # JSON 结果
  299. result_file = self.output_dir / "e4_result.json"
  300. with open(result_file, "w", encoding="utf-8") as f:
  301. json.dump(result_dict, f, indent=2, ensure_ascii=False)
  302. # Markdown 报告
  303. self._generate_report(result)
  304. print(f"\n结果已保存到:{result_file}")
  305. def _generate_report(self, result: E4Result):
  306. """生成 Markdown 报告"""
  307. r = result
  308. report = f"""# E4 实验报告:对齐即谱工程
  309. **实验时间**: {r.timestamp}
  310. **模型对**: {r.config['base_model_name']} (Base) vs {r.config['instruct_model_name']} (Instruct)
  311. ---
  312. ## 一、实验目标
  313. 验证对齐(SFT/Instruct 微调)对模型谱结构的影响:
  314. - **假设 H1**: Instruct 版在指令集上的 r_eff 更低(谱压缩)
  315. - **假设 H2**: Instruct 版的域专一化程度更高(STG 更大)
  316. ---
  317. ## 二、实验配置
  318. | 配置项 | 值 |
  319. |--------|-----|
  320. | Base 模型 | {r.config['base_model_name']} |
  321. | Instruct 模型 | {r.config['instruct_model_name']} |
  322. | 指令数据集 | {r.config['instruction_dataset']} |
  323. | 通用数据集 | {r.config['general_dataset']} |
  324. | 样本数量 | {r.config['sample_size']} 条 |
  325. ---
  326. ## 三、核心结果
  327. ### 3.1 指令集上的 r_eff 对比
  328. | 模型 | r_eff |
  329. |------|-------|
  330. | {r.config['base_model_name']} | {r.instruction_r_eff[r.config['base_model_name']]:.4f} |
  331. | {r.config['instruct_model_name']} | {r.instruction_r_eff[r.config['instruct_model_name']]:.4f} |
  332. | **RS_cross** | {r.instruction_rs_cross:.4f} |
  333. ### 3.2 通用集上的 r_eff 对比(对照)
  334. | 模型 | r_eff |
  335. |------|-------|
  336. | {r.config['base_model_name']} | {r.general_r_eff[r.config['base_model_name']]:.4f} |
  337. | {r.config['instruct_model_name']} | {r.general_r_eff[r.config['instruct_model_name']]:.4f} |
  338. | **RS_cross** | {r.general_rs_cross:.4f} |
  339. ### 3.3 域专一化 (STG)
  340. STG = |r_eff_instruction - r_eff_general|
  341. | 模型 | STG |
  342. |------|-----|
  343. | {r.config['base_model_name']} | {r.stg[r.config['base_model_name']]:.4f} |
  344. | {r.config['instruct_model_name']} | {r.stg[r.config['instruct_model_name']]:.4f} |
  345. ### 3.4 谱距离 (SPD)
  346. | 数据集 | SPD (Base vs Instruct) |
  347. |--------|----------------------|
  348. | 指令集 | {r.spd_instruction:.4f} |
  349. | 通用集 | {r.spd_general:.4f} |
  350. ---
  351. ## 四、假设验证
  352. ### H1: 谱压缩假设
  353. **预测**: Instruct 版在指令集上的 r_eff < Base 版
  354. **结果**: {'✅ 支持' if r.hypothesis_support['h1_spectrum_compression']['supported'] else '❌ 不支持'}
  355. {r.hypothesis_support['h1_spectrum_compression']['interpretation']}
  356. ### H2: 域专一化假设
  357. **预测**: Instruct 版的 STG > Base 版
  358. **结果**: {'✅ 支持' if r.hypothesis_support['h2_domain_specialization']['supported'] else '❌ 不支持'}
  359. {r.hypothesis_support['h2_domain_specialization']['interpretation']}
  360. ---
  361. ## 五、总体结论
  362. **实验结论**: {r.hypothesis_support['overall']}
  363. ### 解释
  364. """
  365. # 添加具体解释
  366. delta_r = (r.instruction_r_eff[r.config['base_model_name']] -
  367. r.instruction_r_eff[r.config['instruct_model_name']])
  368. if delta_r > 0:
  369. report += f"""1. **谱压缩效应**: Instruct 微调使 r_eff 降低了 {delta_r:.4f} ({delta_r/r.instruction_r_eff[r.config['base_model_name']]*100:.1f}%)
  370. - 这表明对齐过程使模型在指令理解上使用更少的表示维度
  371. - 与"对齐即谱工程"的假设一致
  372. 2. **域专一化**:
  373. """
  374. else:
  375. report += f"""1. **谱压缩效应不显著**: Instruct 版 r_eff 反而高 {abs(delta_r):.4f}
  376. - 可能的原因:指令微调增加了表示的多样性
  377. - 需要更多模型对验证
  378. 2. **域专一化**:
  379. """
  380. stg_diff = r.stg[r.config['instruct_model_name']] - r.stg[r.config['base_model_name']]
  381. if stg_diff > 0:
  382. report += f""" - Instruct 版的 STG 比 Base 版高 {stg_diff:.4f}
  383. - 表明对齐增强了模型对指令域的专一化
  384. """
  385. else:
  386. report += f""" - Instruct 版的 STG 比 Base 版低 {abs(stg_diff):.4f}
  387. - 表明对齐没有显著增强域专一化
  388. """
  389. report += """
  390. ---
  391. ## 六、局限性
  392. 1. **单一模型对**: 仅使用 LLaMA-3-8B Base/Instruct 一对,结论普适性有限
  393. 2. **数据集有限**: 仅使用 alpaca 作为指令集,可能需要更多指令数据验证
  394. 3. **缺少 RLHF/DPO**: 仅对比 SFT,没有包含更强的对齐方式
  395. ## 七、未来工作
  396. - 增加 Qwen2.5-7B Base/Instruct 对验证
  397. - 添加 DPO/RLHF 版本对比
  398. - 分析逐层谱变化(LCP)
  399. ---
  400. *报告生成时间:{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}*
  401. """
  402. report_file = self.output_dir / "e4_report.md"
  403. with open(report_file, "w", encoding="utf-8") as f:
  404. f.write(report)
  405. print(f"报告已保存到:{report_file}")
  406. def main():
  407. """主函数"""
  408. import argparse
  409. parser = argparse.ArgumentParser(description="E4 实验:对齐即谱工程")
  410. parser.add_argument("--model-pair", type=str, default="qwen",
  411. choices=["qwen", "llama3", "mistral"],
  412. help="选择模型对:qwen (Qwen2.5-7B), llama3 (LLaMA-3-8B), 或 mistral (Mistral-7B)")
  413. parser.add_argument("--sample-size", type=int, default=200,
  414. help="采样文本数量")
  415. parser.add_argument("--output-dir", type=str, default=None,
  416. help="输出目录")
  417. args = parser.parse_args()
  418. # 根据选择配置模型对
  419. if args.model_pair == "qwen":
  420. # Qwen2.5-7B 对(真正的 Base/Instruct)
  421. config = E4Config(
  422. base_model_name="Qwen2.5-7B-Base",
  423. instruct_model_name="Qwen2.5-7B-Instruct",
  424. base_model_path="model/weights/Qwen2.5-7B-Base",
  425. instruct_model_path="model/weights/Qwen2.5-7B-Instruct",
  426. instruction_dataset="alpaca",
  427. general_dataset="flores200",
  428. sample_size=args.sample_size,
  429. batch_size=4,
  430. output_dir=args.output_dir or "experiments/output/e4_qwen"
  431. )
  432. elif args.model_pair == "llama3":
  433. # LLaMA-3-8B 对(注意:之前发现两个权重相同,仅用于测试)
  434. config = E4Config(
  435. base_model_name="llama3-8b",
  436. instruct_model_name="llama3-8b-instruct",
  437. base_model_path="model/weights/llama3-8b",
  438. instruct_model_path="model/weights/llama3-8b-instruct",
  439. instruction_dataset="alpaca",
  440. general_dataset="flores200",
  441. sample_size=args.sample_size,
  442. batch_size=4,
  443. output_dir=args.output_dir or "experiments/output/e4_llama3"
  444. )
  445. else: # mistral
  446. # Mistral-7B 对
  447. config = E4Config(
  448. base_model_name="Mistral-7B-v0.3",
  449. instruct_model_name="Mistral-7B-Instruct-v0.3",
  450. base_model_path="model/weights/Mistral-7B-v0.3",
  451. instruct_model_path="model/weights/Mistral-7B-Instruct-v0.3",
  452. instruction_dataset="alpaca",
  453. general_dataset="flores200",
  454. sample_size=args.sample_size,
  455. batch_size=4,
  456. output_dir=args.output_dir or "experiments/output/e4_mistral"
  457. )
  458. analyzer = AlignmentSpectrumAnalyzer(config)
  459. result = analyzer.run_full_analysis()
  460. print("\n" + "=" * 60)
  461. print("E4 实验完成!")
  462. print("=" * 60)
  463. print(f"\n核心发现:")
  464. print(f" 指令集 r_eff: {result.instruction_r_eff[result.config['base_model_name']]:.4f} (Base) -> "
  465. f"{result.instruction_r_eff[result.config['instruct_model_name']]:.4f} (Instruct)")
  466. print(f" 域专一化 STG: {result.stg[result.config['base_model_name']]:.4f} (Base) -> "
  467. f"{result.stg[result.config['instruct_model_name']]:.4f} (Instruct)")
  468. print(f"\n结论:{result.hypothesis_support['overall']}")
  469. if __name__ == "__main__":
  470. main()