| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402 |
- #!/usr/bin/env python
- """
- 能力涌现谱突变检测 - 方向 9 实验 E3
- 功能:
- - 检测 r_eff 曲线中的突变点
- - 使用 Ruptures 库进行变化点检测
- - 关联突变点与模型能力涌现
- 用法:
- python experiments/emergence_detection.py --output-dir output/e3
- """
- import os
- import json
- import torch
- import numpy as np
- from pathlib import Path
- from typing import Dict, List, Tuple, Optional, Any
- from dataclasses import dataclass, asdict
- from datetime import datetime
- from transformers import AutoTokenizer, AutoModelForCausalLM
- # 导入谱计算模块
- import sys
- sys.path.insert(0, str(Path(__file__).parent.parent))
- from model.spectrum import compute_gram_spectrum
- from model.extractor import save_H, load_H
- from experiments.cross_model_convergence import CrossModelConvergence, ExperimentConfig
- try:
- import ruptures as rpt
- HAS_RUPTURES = True
- except ImportError:
- HAS_RUPTURES = False
- print("警告:ruptures 未安装,变化点检测功能不可用")
- @dataclass
- class EmergenceResult:
- """能力涌现检测结果"""
- timestamp: str
- ability_name: str
- r_eff_curve: List[float]
- model_order: List[str]
- model_mmlu: List[float]
- max_jump: dict
- changepoints: List[int]
- jump_significance: float
- class EmergenceDetector:
- """能力涌现检测器"""
- def __init__(
- self,
- model_sequence: List[str],
- model_paths: Dict[str, str],
- model_mmlu: Dict[str, float],
- output_dir: str = "experiments/output/emergence"
- ):
- self.model_sequence = model_sequence
- self.model_paths = model_paths
- self.model_mmlu = model_mmlu
- self.output_dir = Path(output_dir)
- self.output_dir.mkdir(parents=True, exist_ok=True)
- # 按 MMLU 排序模型
- self.sorted_models = sorted(
- model_sequence,
- key=lambda m: model_mmlu.get(m, 0)
- )
- def load_texts_for_ability(
- self,
- ability_name: str
- ) -> List[str]:
- """根据能力类型加载探针文本"""
- from database.corpus import CorpusLoader
- loader = CorpusLoader("database/corpus")
- # 能力 - 数据集映射
- ability_mapping = {
- "cot_reasoning": ["bigbench"],
- "arithmetic": ["gsm8k", "math"],
- "code_generation": ["humaneval", "code"],
- "instruction_following": ["alpaca"],
- }
- datasets = ability_mapping.get(ability_name, ["gsm8k"])
- texts = []
- for ds in datasets:
- try:
- texts.extend(loader.load_texts(ds, limit=200))
- except FileNotFoundError:
- continue
- return texts[:500]
- def compute_r_eff_curve(
- self,
- texts: List[str],
- ability_name: str
- ) -> Dict[str, float]:
- """计算每个模型的 r_eff"""
- r_effs = {}
- for model_name in self.sorted_models:
- print(f" 处理 {model_name}...")
- # 构建唯一的数据集名称
- dataset_name = f"emergence_{ability_name}"
- # 提取表示
- H = self._extract_hidden_states(model_name, texts)
- # 计算 r_eff
- H_tensor = torch.tensor(H.numpy() if hasattr(H, "cpu") else H) if not isinstance(H, torch.Tensor) else H
- r_eff = compute_gram_spectrum(H_tensor)
- r_effs[model_name] = r_eff
- # 保存表示
- rep_dir = self.output_dir.parent / "representations"
- rep_dir.mkdir(parents=True, exist_ok=True)
- save_H(H, model_name, dataset_name,
- representations_dir=str(rep_dir))
- return r_effs
- def _extract_hidden_states(
- self,
- model_name: str,
- texts: List[str]
- ) -> torch.Tensor:
- """提取隐层表示"""
- model_path = self.model_paths.get(
- model_name,
- f"model/weights/{model_name}"
- )
- tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
- # 设置 pad_token
- if tokenizer.pad_token is None:
- tokenizer.pad_token = tokenizer.eos_token
- if tokenizer.eos_token_id is None:
- tokenizer.eos_token_id = 50256
- model = AutoModelForCausalLM.from_pretrained(
- model_path,
- torch_dtype=torch.float16,
- device_map="auto",
- output_hidden_states=True,
- trust_remote_code=True
- )
- model.eval()
- all_features = []
- batch_size = 8
- with torch.no_grad():
- for i in range(0, len(texts), batch_size):
- batch = texts[i:i + batch_size]
- inputs = tokenizer(
- batch,
- return_tensors="pt",
- padding=True,
- truncation=True,
- max_length=512
- ).to(model.device)
- outputs = model(**inputs)
- hidden = outputs.hidden_states[-1]
- attention_mask = inputs["attention_mask"]
- # Mean pooling
- mask_expanded = attention_mask.unsqueeze(-1).float()
- feat = (hidden * mask_expanded).sum(1) / (mask_expanded.sum(1) + 1e-10)
- all_features.append(feat.cpu().float())
- H = torch.cat(all_features, dim=0).T
- return H
- def detect_changepoints(
- self,
- r_effs: Dict[str, float],
- ability_name: str
- ) -> EmergenceResult:
- """检测变化点"""
- # 按能力排序的 r_eff 曲线
- r_eff_curve = [r_effs[m] for m in self.sorted_models]
- mmlu_scores = [self.model_mmlu.get(m, 0) for m in self.sorted_models]
- # 方法 1:一阶差分寻找最大跃变
- if len(r_eff_curve) > 1:
- diffs = np.abs(np.diff(r_eff_curve))
- jump_idx = int(np.argmax(diffs))
- jump_magnitude = float(diffs[jump_idx])
- else:
- jump_idx = 0
- jump_magnitude = 0
- # 方法 2:Ruptures 变化点检测
- changepoints = []
- if HAS_RUPTURES and len(r_eff_curve) >= 3:
- try:
- # PELT 算法
- algo = rpt.Pelt(model="rbf").fit(np.array(r_eff_curve).reshape(-1, 1))
- changepoints = algo.predict(pen=2)
- except Exception as e:
- print(f" Ruptures 检测失败:{e}")
- # 降级:使用二阶差分
- if len(r_eff_curve) >= 3:
- second_diff = np.diff(np.diff(r_eff_curve))
- changepoints = [int(np.argmax(np.abs(second_diff)) + 1) + 1]
- # 计算跃变显著性
- if len(r_eff_curve) >= 2:
- mean_r_eff = np.mean(r_eff_curve)
- std_r_eff = np.std(r_eff_curve) if len(r_eff_curve) > 1 else 1
- jump_significance = jump_magnitude / (std_r_eff + 1e-10)
- else:
- jump_significance = 0
- result = EmergenceResult(
- timestamp=datetime.now().isoformat(),
- ability_name=ability_name,
- r_eff_curve=r_eff_curve,
- model_order=self.sorted_models.copy(),
- model_mmlu=mmlu_scores,
- max_jump={
- "model_index": jump_idx,
- "model_name": self.sorted_models[jump_idx] if jump_idx < len(self.sorted_models) else "N/A",
- "mmlu_at_jump": mmlu_scores[jump_idx] if jump_idx < len(mmlu_scores) else 0,
- "jump_magnitude": jump_magnitude
- },
- changepoints=changepoints,
- jump_significance=jump_significance
- )
- return result
- def run_full_analysis(
- self,
- abilities: Optional[List[str]] = None
- ) -> Dict[str, EmergenceResult]:
- """运行完整的能力涌现分析"""
- if abilities is None:
- abilities = [
- "cot_reasoning",
- "arithmetic",
- "code_generation",
- "instruction_following"
- ]
- results = {}
- for ability in abilities:
- print(f"\n{'='*60}")
- print(f"分析能力:{ability}")
- print(f"{'='*60}")
- # 加载探针文本
- texts = self.load_texts_for_ability(ability)
- if not texts:
- print(f" 跳过:无可用文本")
- continue
- print(f" 文本数:{len(texts)}")
- # 计算 r_eff 曲线
- r_effs = self.compute_r_eff_curve(texts, ability)
- # 检测变化点
- result = self.detect_changepoints(r_effs, ability)
- results[ability] = result
- # 打印摘要
- print(f"\n 结果摘要:")
- print(f" 最大跃变模型:{result.max_jump['model_name']}")
- print(f" MMLU@跃变点:{result.max_jump['mmlu_at_jump']:.1f}%")
- print(f" 跃变幅度:{result.max_jump['jump_magnitude']:.3f}")
- print(f" 跃变显著性:{result.jump_significance:.2f}σ")
- if result.changepoints:
- print(f" 变化点索引:{result.changepoints}")
- # 保存结果
- self._save_results(results)
- return results
- def _save_results(self, results: Dict[str, EmergenceResult]):
- """保存结果"""
- results_dict = {}
- for name, result in results.items():
- results_dict[name] = asdict(result)
- result_file = self.output_dir / "emergence_results.json"
- with open(result_file, "w") as f:
- json.dump(results_dict, f, indent=2)
- print(f"\n结果已保存到:{result_file}")
- # 生成摘要报告
- self._generate_summary(results)
- def _generate_summary(self, results: Dict[str, EmergenceResult]):
- """生成摘要报告"""
- # 获取第一个结果来访问 model_order 和 model_mmlu
- first_result = list(results.values())[0]
- model_order = first_result.model_order
- model_mmlu = first_result.model_mmlu
- summary = [
- "# E3 能力涌现谱突变检测摘要",
- "",
- f"分析时间:{datetime.now().strftime('%Y-%m-%d %H:%M')}",
- "",
- "## 检测到的涌现点",
- "",
- "| 能力 | 涌现模型 | MMLU@涌现 | 跃变幅度 | 显著性 |",
- "|------|---------|-----------|----------|--------|"
- ]
- for name, result in results.items():
- summary.append(
- f"| {name} | {result.max_jump['model_name']} | "
- f"{result.max_jump['mmlu_at_jump']:.1f}% | "
- f"{result.max_jump['jump_magnitude']:.3f} | "
- f"{result.jump_significance:.2f}σ |"
- )
- summary.extend([
- "",
- "## 模型序列 (按 MMLU 排序)",
- ""
- ])
- for i, model in enumerate(model_order):
- mmlu = model_mmlu[i]
- summary.append(f"{i+1}. {model} (MMLU: {mmlu:.1f}%)")
- summary_file = self.output_dir / "emergence_summary.md"
- with open(summary_file, "w") as f:
- f.write("\n".join(summary))
- print(f"摘要报告已保存到:{summary_file}")
- def main():
- import argparse
- parser = argparse.ArgumentParser(description="能力涌现谱突变检测")
- parser.add_argument("--output-dir", type=str,
- default="experiments/output/emergence",
- help="输出目录")
- parser.add_argument("--abilities", type=str, nargs="+",
- default=["cot_reasoning", "arithmetic", "code_generation", "instruction_following"],
- help="要分析的能力列表")
- args = parser.parse_args()
- # 模型序列和 MMLU 分数(文献值)
- MODEL_SEQUENCE = [
- "llama3.2-3b-instruct",
- "Mistral-7B-v0.3",
- "llama3-8b",
- "Qwen2.5-7B-Instruct",
- "gemma-2-9b-it"
- ]
- MODEL_PATHS = {
- "llama3.2-3b-instruct": "model/weights/llama3.2-3b-instruct",
- "Mistral-7B-v0.3": "model/weights/Mistral-7B-v0.3",
- "llama3-8b": "model/weights/llama3-8b",
- "Qwen2.5-7B-Instruct": "model/weights/Qwen2.5-7B-Instruct",
- "gemma-2-9b-it": "model/weights/gemma-2-9b-it"
- }
- # 近似 MMLU 分数(用于排序)
- MODEL_MMLU = {
- "llama3.2-3b-instruct": 58.0,
- "Mistral-7B-v0.3": 62.5,
- "llama3-8b": 68.4,
- "Qwen2.5-7B-Instruct": 86.3,
- "gemma-2-9b-it": 82.0
- }
- detector = EmergenceDetector(
- model_sequence=MODEL_SEQUENCE,
- model_paths=MODEL_PATHS,
- model_mmlu=MODEL_MMLU,
- output_dir=args.output_dir
- )
- results = detector.run_full_analysis(args.abilities)
- if __name__ == "__main__":
- main()
|