#!/usr/bin/env python """ 从预训练模型提取文本表示 使用现有的模型权重,为每个领域的文本提取最后一层隐藏状态, 计算 mean pooling 后保存为 H ∈ R^{d×N} 矩阵。 """ import os # 设置环境变量以减少显存碎片化 os.environ["PYTORCH_ALLOC_CONF"] = "expandable_segments:True" import gc import torch import json from pathlib import Path from typing import List, Dict from transformers import AutoTokenizer, AutoModel import warnings warnings.filterwarnings("ignore") def load_texts(filepath: str) -> List[str]: """从 JSONL 文件加载文本""" texts = [] with open(filepath, "r", encoding="utf-8") as f: for line in f: record = json.loads(line.strip()) texts.append(record.get("text", "")) return texts def extract_representations( model_path: str, texts: List[str], batch_size: int = 4, # 减小 batch_size 以避免 OOM max_length: int = 512 ) -> torch.Tensor: """ 从文本列表提取模型表示 Args: model_path: 模型路径(本地目录) texts: 文本列表 batch_size: batch 大小 max_length: 最大序列长度 Returns: H: 表示矩阵,shape = (d, N) """ print(f" 加载模型:{model_path}") tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True) # 修复:某些模型(如 Mistral)缺少 pad_token,需要手动设置 if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token print(f" 设置 pad_token = eos_token: {repr(tokenizer.pad_token)}") model = AutoModel.from_pretrained( model_path, torch_dtype=torch.float32, device_map="auto", trust_remote_code=True, output_hidden_states=True ) model.eval() print(f" 模型设备:{model.device}") print(f" 文本数量:{len(texts)}") all_embeddings = [] with torch.no_grad(): for i in range(0, len(texts), batch_size): batch_texts = texts[i:i + batch_size] # Tokenize inputs = tokenizer( batch_texts, return_tensors="pt", padding=True, truncation=True, max_length=max_length ) inputs = inputs.to(model.device) # Forward pass outputs = model(**inputs) # Mean pooling over sequence length # hidden_states: (batch, seq_len, hidden_dim) hidden = outputs.last_hidden_state attention_mask = inputs["attention_mask"].unsqueeze(-1).float() # 加权平均(考虑 padding) sum_embeddings = (hidden * attention_mask).sum(dim=1) valid_tokens = attention_mask.sum(dim=1) mean_embeddings = sum_embeddings / (valid_tokens + 1e-9) all_embeddings.append(mean_embeddings.cpu()) print(f" 进度:{min(i + batch_size, len(texts))}/{len(texts)}") # 拼接所有 embedding: (N, d) embeddings = torch.cat(all_embeddings, dim=0) # 转置为 (d, N) 格式 H = embeddings.T # 清理 GPU 内存 del model, tokenizer if torch.cuda.is_available(): torch.cuda.empty_cache() return H def save_representation( H: torch.Tensor, model_name: str, domain: str, output_dir: str = "model/representations" ) -> str: """保存表示矩阵""" output_path = Path(output_dir) / f"{model_name}_{domain}.pt" output_path.parent.mkdir(parents=True, exist_ok=True) torch.save(H, output_path) # 保存元数据 meta_path = output_path.with_suffix(".json") meta = { "model_name": model_name, "domain": domain, "d": int(H.shape[0]), "N": int(H.shape[1]), "dtype": str(H.dtype) } with open(meta_path, "w") as f: json.dump(meta, f, indent=2) return str(output_path) def get_model_name(model_path: str) -> str: """从路径提取模型简称""" path = Path(model_path) name = path.name.lower() # 简化命名 if "qwen" in name: return "qwen_7b" elif "mistral" in name: return "mistral_7b" elif "llama3" in name or "llama-3" in name: return "llama3_8b" elif "gemma" in name: return "gemma_9b" elif "bert" in name: return "bert" elif "gpt2" in name: return "gpt2" else: return path.name.replace("-", "_") def main(): print("=" * 60) print("模型表示提取") print("=" * 60) print() # 可用模型 models = [ "model/weights/Qwen2.5-7B-Instruct", "model/weights/Mistral-7B-v0.3", "model/weights/llama3-8b", ] # 可用领域(6 个原始领域 + 5 个 GLUE 领域) domains = [ # 原始领域 ("news_en", "database/corpus/news_en/texts.jsonl"), ("news_zh", "database/corpus/news_zh/texts.jsonl"), ("academic", "database/corpus/academic/texts.jsonl"), ("code", "database/corpus/code/texts.jsonl"), ("dialogue", "database/corpus/dialogue/texts.jsonl"), ("literature", "database/corpus/literature/texts.jsonl"), # GLUE 领域 ("glue_mnli", "database/corpus/glue_mnli/texts.jsonl"), ("glue_qnli", "database/corpus/glue_qnli/texts.jsonl"), ("glue_sst2", "database/corpus/glue_sst2/texts.jsonl"), ("glue_stsb", "database/corpus/glue_stsb/texts.jsonl"), ("glue_cola", "database/corpus/glue_cola/texts.jsonl"), ] # 检查哪些文本文件存在 available_domains = [] for domain, path in domains: if Path(path).exists(): available_domains.append((domain, path)) else: print(f"跳过(文件不存在):{path}") if not available_domains: print("错误:没有找到任何文本语料文件!") return print(f"可用领域:{[d[0] for d in available_domains]}") print(f"模型数量:{len(models)}") print() # 对每个模型 × 领域组合提取表示 for model_path in models: if not Path(model_path).exists(): print(f"跳过(模型不存在):{model_path}") continue model_name = get_model_name(model_path) print(f"\n[模型:{model_name}]") for domain, texts_path in available_domains: print(f"\n [{domain}]") # 加载文本 texts = load_texts(texts_path) print(f" 加载了 {len(texts)} 条文本") if len(texts) == 0: print(" 跳过:没有文本") continue # 提取表示 H = extract_representations(model_path, texts) print(f" 表示矩阵形状:{H.shape}") print(f" dtype: {H.dtype}") # 保存 output_path = save_representation(H, model_name, domain) print(f" 已保存:{output_path}") # 每个领域结束后清理 GPU 缓存 del H gc.collect() if torch.cuda.is_available(): torch.cuda.empty_cache() print() print("=" * 60) print("提取完成!") print("=" * 60) print() print("下一步:运行 MVP 实验") print(" python -m experiments.verify --mode mvp") if __name__ == "__main__": main()