| import os |
| import torch |
| import torch.nn as nn |
| import torch.optim as optim |
| from torch.utils.data import Dataset, DataLoader |
| import numpy as np |
| import random |
| import matplotlib.pyplot as plt |
| import json |
| from datasets import load_dataset |
| from tqdm import tqdm |
| import re |
| import datetime |
| from tokenizers import Tokenizer |
| from tokenizers.models import BPE |
| from tokenizers.trainers import BpeTrainer |
| from tokenizers.pre_tokenizers import Whitespace |
| import torch.backends.cudnn as cudnn |
| import time |
| import subprocess |
| import psutil |
|
|
| |
| |
| |
|
|
| |
| if torch.backends.mps.is_available(): |
| device = torch.device("mps") |
| print("使用设备: Apple Silicon (MPS)") |
| else: |
| device = torch.device("cpu") |
| print("使用设备: CPU") |
|
|
| |
| torch.manual_seed(42) |
| np.random.seed(42) |
| random.seed(42) |
|
|
| |
| |
| |
|
|
| class HardwareGuard: |
| def __init__(self, device_type="mps", max_temp=90, cooldown_time=900): |
| """ |
| MPS设备硬件保护系统 |
| :param device_type: 设备类型 (mps/cpu) |
| :param max_temp: 最大允许温度 (°C) |
| :param cooldown_time: 冷却时间 (秒) |
| """ |
| self.device_type = device_type |
| self.max_temp = max_temp |
| self.cooldown_time = cooldown_time |
| self.last_check_time = time.time() |
| self.check_interval = 300 |
| self.low_power_mode = False |
| |
| |
| self.last_temp = self.get_current_temperature() |
| print(f"硬件保护系统启动 - 最大温度: {max_temp}°C, 冷却时间: {cooldown_time}秒") |
| |
| def get_current_temperature(self): |
| """获取当前设备温度""" |
| if self.device_type == "mps": |
| try: |
| |
| result = subprocess.run(['istats', 'cpu', 'temp'], capture_output=True, text=True) |
| if result.returncode == 0: |
| |
| temp_str = result.stdout.split(":")[1].split("°")[0].strip() |
| return float(temp_str) |
| except: |
| pass |
| |
| |
| try: |
| temps = psutil.sensors_temperatures() |
| if 'coretemp' in temps: |
| return max([entry.current for entry in temps['coretemp']]) |
| elif 'acpitz' in temps: |
| return max([entry.current for entry in temps['acpitz']]) |
| except: |
| pass |
| |
| return 60 |
| |
| def get_power_mode(self): |
| """获取当前电源模式""" |
| try: |
| result = subprocess.run(['pmset', '-g'], capture_output=True, text=True) |
| if 'lowpowermode' in result.stdout.lower(): |
| return "low_power" |
| return "normal" |
| except: |
| return "normal" |
| |
| def enforce_low_power_mode(self): |
| """启用低功耗模式""" |
| if not self.low_power_mode: |
| print("⚠️ 温度过高,启用低功耗模式") |
| try: |
| subprocess.run(['sudo', 'pmset', '-a', 'lowpowermode', '1']) |
| self.low_power_mode = True |
| except: |
| print("无法启用低功耗模式") |
| |
| def restore_normal_power_mode(self): |
| """恢复正常功耗模式""" |
| if self.low_power_mode: |
| print("✅ 温度正常,恢复标准功耗模式") |
| try: |
| subprocess.run(['sudo', 'pmset', '-a', 'lowpowermode', '0']) |
| self.low_power_mode = False |
| except: |
| print("无法恢复标准功耗模式") |
| |
| def check_temperature(self): |
| """检查温度并采取保护措施""" |
| current_time = time.time() |
| if current_time - self.last_check_time > self.check_interval: |
| temp = self.get_current_temperature() |
| self.last_check_time = current_time |
| |
| print(f"🌡️ 当前温度: {temp}°C | 电源模式: {self.get_power_mode()}") |
| |
| if temp > self.max_temp: |
| self.enforce_low_power_mode() |
| return True |
| else: |
| self.restore_normal_power_mode() |
| |
| return False |
| |
| def cooldown_protocol(self): |
| """执行冷却协议""" |
| print(f"🔥 温度过高,开始冷却协议 ({self.cooldown_time}秒)...") |
| |
| |
| torch.mps.empty_cache() |
| |
| |
| try: |
| subprocess.run(['pmset', 'displaysleepnow']) |
| except: |
| pass |
| |
| |
| cooldown_start = time.time() |
| while time.time() - cooldown_start < self.cooldown_time: |
| time.sleep(60) |
| temp = self.get_current_temperature() |
| print(f"冷却中... 当前温度: {temp}°C") |
| |
| |
| if temp < self.max_temp - 5: |
| break |
| |
| print("✅ 冷却完成,恢复训练") |
| try: |
| subprocess.run(['caffeinate', '-u', '-t', '1']) |
| except: |
| pass |
|
|
| |
| |
| |
|
|
| |
| if device.type == "mps": |
| torch.mps.set_per_process_memory_fraction(0.8) |
| torch.mps.empty_cache() |
|
|
| |
| |
| |
|
|
| def train_bpe_tokenizer(text_list, vocab_size=50000, save_path="bpe_tokenizer.json"): |
| """训练BPE分词器""" |
| print(f"训练BPE分词器 (词汇表大小: {vocab_size})") |
| |
| |
| tokenizer = Tokenizer(BPE(unk_token="<UNK>")) |
| trainer = BpeTrainer( |
| vocab_size=vocab_size, |
| special_tokens=["<PAD>", "<UNK>", "<SEP>", "<CLS>"], |
| min_frequency=2 |
| ) |
| tokenizer.pre_tokenizer = Whitespace() |
| |
| |
| tokenizer.train_from_iterator(text_list, trainer=trainer) |
| |
| |
| tokenizer.save(save_path) |
| print(f"BPE分词器已保存至 {save_path}") |
| |
| return tokenizer |
|
|
| def load_bpe_tokenizer(tokenizer_path="bpe_tokenizer.json"): |
| """加载BPE分词器""" |
| if os.path.exists(tokenizer_path): |
| tokenizer = Tokenizer.from_file(tokenizer_path) |
| print(f"已加载BPE分词器,词汇表大小: {tokenizer.get_vocab_size()}") |
| return tokenizer |
| else: |
| print(f"未找到分词器文件 {tokenizer_path}") |
| return None |
|
|
| def build_vocab_from_tokenizer(tokenizer): |
| """从BPE分词器构建词汇表""" |
| vocab = tokenizer.get_vocab() |
| reverse_vocab = {idx: token for token, idx in vocab.items()} |
| print(f"BPE词汇表大小: {len(vocab)}") |
| return vocab, reverse_vocab |
|
|
| |
| |
| |
|
|
| class TextDataset(Dataset): |
| def __init__(self, text_list, tokenizer, max_sequence_length=128): |
| self.tokenizer = tokenizer |
| self.text_indices = [] |
| self.max_sequence_length = max_sequence_length |
| |
| for text in tqdm(text_list, desc="预处理文本"): |
| |
| token_ids = tokenizer.encode(text).ids |
| |
| |
| for i in range(0, len(token_ids), max_sequence_length): |
| segment = token_ids[i:i + max_sequence_length] |
| if len(segment) >= 2: |
| self.text_indices.append(segment) |
|
|
| def __len__(self): |
| return len(self.text_indices) |
|
|
| def __getitem__(self, idx): |
| token_ids = self.text_indices[idx] |
| input_seq = token_ids[:-1] |
| target_seq = token_ids[1:] |
| return torch.tensor(input_seq, dtype=torch.long), torch.tensor(target_seq, dtype=torch.long) |
|
|
| class DynamicPadder: |
| def __init__(self, pad_token_id, max_length=128): |
| self.pad_token_id = pad_token_id |
| self.max_length = max_length |
| |
| def __call__(self, batch): |
| """ |
| 动态填充批次中的序列并限制最大长度 |
| """ |
| batch_size = len(batch) |
| |
| |
| max_len = min(max(len(seq[0]) for seq in batch), self.max_length) |
| |
| padded_inputs = torch.full((batch_size, max_len), self.pad_token_id, dtype=torch.long) |
| padded_targets = torch.full((batch_size, max_len), self.pad_token_id, dtype=torch.long) |
|
|
| for i, (input_seq, target_seq) in enumerate(batch): |
| seq_len = min(len(input_seq), max_len) |
| padded_inputs[i, :seq_len] = input_seq[:seq_len] |
| padded_targets[i, :seq_len] = target_seq[:seq_len] |
|
|
| return padded_inputs, padded_targets |
|
|
| def clean_text(text): |
| """改进的文本清洗函数""" |
| |
| text = text.lower() |
| |
| text = re.sub(r'[^a-z0-9\s.,!?;:\'"-]', '', text) |
| |
| text = re.sub(r'\s+', ' ', text).strip() |
| return text |
|
|
| def load_dataset_text(dataset_name="BEE-spoke-data/fineweb-1M_en-med", split="train", max_samples=5000): |
| """ |
| 加载数据集并限制样本数量和长度 |
| """ |
| print(f"加载数据集: {dataset_name}...") |
| text_list = [] |
| total_samples = max_samples |
| |
| try: |
| |
| dataset = load_dataset(dataset_name, split=split, streaming=True) |
| dataset = dataset.take(max_samples) |
| print(f"使用流式数据加载,取 {max_samples} 个样本") |
| except Exception as e: |
| print(f"流式加载失败: {str(e)},尝试完整加载...") |
| try: |
| dataset = load_dataset(dataset_name, split=split) |
| if max_samples is not None and max_samples < len(dataset): |
| dataset = dataset.select(range(max_samples)) |
| total_samples = max_samples |
| else: |
| total_samples = len(dataset) |
| print(f"使用完整数据集加载,共 {total_samples} 个样本") |
| except Exception as e2: |
| print(f"完整加载也失败: {str(e2)}") |
| return text_list |
| |
| |
| count = 0 |
| for sample in tqdm(dataset, desc="处理样本", total=total_samples): |
| if 'text' in sample: |
| text = sample['text'] |
| elif 'content' in sample: |
| text = sample['content'] |
| else: |
| continue |
| |
| |
| text = clean_text(text) |
| |
| |
| words = text.split()[:500] |
| if len(words) > 5: |
| text_list.append(" ".join(words)) |
| count += 1 |
| |
| print(f"成功提取 {count} 个文本样本") |
| return text_list |
|
|
| |
| |
| |
|
|
| class StabilizedDenoisingModel(nn.Module): |
| def __init__(self, vocab_size, embed_dim, hidden_dim, num_layers): |
| super(StabilizedDenoisingModel, self).__init__() |
|
|
| |
| self.embedding = nn.Embedding(vocab_size, embed_dim) |
| |
| |
| self.row_transform = nn.Linear(embed_dim, hidden_dim) |
| self.dim_transform = nn.Linear(hidden_dim, hidden_dim) |
| |
| |
| self.norm = nn.LayerNorm(hidden_dim) |
|
|
| |
| self.denoise_layers = nn.ModuleList([ |
| nn.Sequential( |
| nn.Linear(hidden_dim, hidden_dim), |
| nn.ReLU(), |
| nn.Linear(hidden_dim, hidden_dim) |
| ) |
| for _ in range(num_layers) |
| ]) |
|
|
| |
| self.output_layer = nn.Linear(hidden_dim, vocab_size) |
| |
| |
| self._init_weights() |
| |
| def _init_weights(self): |
| """改进的权重初始化,增强训练稳定性""" |
| |
| nn.init.normal_(self.embedding.weight, mean=0.0, std=0.02) |
| |
| |
| for layer in [self.row_transform, self.dim_transform, self.output_layer]: |
| nn.init.xavier_uniform_(layer.weight) |
| if layer.bias is not None: |
| nn.init.zeros_(layer.bias) |
| |
| |
| for layer_seq in self.denoise_layers: |
| for i, layer in enumerate(layer_seq): |
| if isinstance(layer, nn.Linear): |
| if i % 2 == 0: |
| nn.init.kaiming_normal_(layer.weight, nonlinearity='relu') |
| else: |
| nn.init.xavier_uniform_(layer.weight) |
| if layer.bias is not None: |
| nn.init.zeros_(layer.bias) |
|
|
| def forward(self, input_seq): |
| |
| embedded_seq = self.embedding(input_seq) |
| |
| |
| hidden_space = self.row_transform(embedded_seq) |
| hidden_space = self.dim_transform(hidden_space) |
| hidden_space = self.norm(hidden_space) |
|
|
| |
| for denoise_layer in self.denoise_layers: |
| signal = denoise_layer(hidden_space) |
| |
| |
| gate = torch.sigmoid(signal) |
| denoised = hidden_space - gate * signal + (1 - gate) * torch.relu(signal) |
| |
| |
| hidden_space = self.norm(hidden_space + denoised) |
|
|
| |
| logits = self.output_layer(hidden_space) |
| return logits |
|
|
| |
| |
| |
|
|
| def train_model(text_list, embed_dim, hidden_dim, num_layers, batch_size, |
| num_epochs, lr=1e-3, model_path="model.pth", tokenizer_path="bpe_tokenizer.json"): |
| """ |
| 训练生成式去噪模型 (MPS优化版) |
| """ |
| |
| hardware_guard = HardwareGuard( |
| device_type=device.type, |
| max_temp=85, |
| cooldown_time=900 |
| ) |
| |
| |
| tokenizer = load_bpe_tokenizer(tokenizer_path) |
| if tokenizer is None: |
| print("分词器无效或不存在,将重新训练...") |
| tokenizer = train_bpe_tokenizer(text_list, vocab_size=50000) |
| |
| |
| vocab, reverse_vocab = build_vocab_from_tokenizer(tokenizer) |
| vocab_size = tokenizer.get_vocab_size() |
| pad_token_id = vocab["<PAD>"] |
| print(f"词汇表大小: {vocab_size}") |
|
|
| |
| print("构建文本数据集...") |
| dataset = TextDataset(text_list, tokenizer, max_sequence_length=128) |
| print(f"数据集大小: {len(dataset)}") |
| |
| |
| num_workers = min(4, os.cpu_count() // 2) |
| |
| |
| padder = DynamicPadder(pad_token_id, max_length=128) |
| |
| data_loader = DataLoader( |
| dataset, |
| batch_size=batch_size, |
| shuffle=True, |
| num_workers=num_workers, |
| pin_memory=(device.type != "mps"), |
| collate_fn=padder |
| ) |
| |
| print(f"数据加载器使用 {num_workers} workers") |
|
|
| |
| print(f"初始化模型 (词汇表大小: {vocab_size}, 嵌入维度: {embed_dim}, 隐藏维度: {hidden_dim}, 层数: {num_layers})") |
| model = StabilizedDenoisingModel(vocab_size, embed_dim, hidden_dim, num_layers).to(device) |
| |
| |
| optimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-5) |
| criterion = nn.CrossEntropyLoss(ignore_index=pad_token_id) |
|
|
| |
| scheduler = optim.lr_scheduler.ReduceLROnPlateau( |
| optimizer, |
| mode='min', |
| factor=0.5, |
| patience=3 |
| ) |
|
|
| |
| start_epoch = 0 |
| if os.path.exists(model_path): |
| checkpoint = torch.load(model_path, map_location=device) |
| |
| if checkpoint["model_state_dict"]["embedding.weight"].size(0) == vocab_size: |
| model.load_state_dict(checkpoint["model_state_dict"]) |
| if "optimizer_state_dict" in checkpoint: |
| optimizer.load_state_dict(checkpoint["optimizer_state_dict"]) |
| start_epoch = checkpoint.get("epoch", 0) + 1 |
| print(f"已加载模型参数,继续从第 {start_epoch} 轮训练") |
| else: |
| print("词汇表大小不匹配,将从头开始训练") |
| else: |
| print(f"未找到模型文件 {model_path},从头开始训练") |
| |
| |
| print(f"开始训练,共 {num_epochs} 轮,批次大小: {batch_size}") |
| best_loss = float('inf') |
| patience = 5 |
| patience_counter = 0 |
| train_losses = [] |
| |
| |
| log_file = f"training_log_{datetime.datetime.now().strftime('%Y%m%d_%H%M%S')}.txt" |
| with open(log_file, "w") as f: |
| f.write(f"训练开始时间: {datetime.datetime.now()}\n") |
| f.write(f"设备: {device}\n") |
| f.write(f"模型参数: embed_dim={embed_dim}, hidden_dim={hidden_dim}, num_layers={num_layers}, batch_size={batch_size}, lr={lr}\n") |
| f.write(f"词汇表大小: {vocab_size}\n") |
| f.write(f"数据集大小: {len(dataset)}\n") |
| f.write(f"硬件保护: 最大温度=85°C, 冷却时间=900秒\n") |
| f.write("="*80 + "\n") |
| |
| for epoch in range(start_epoch, start_epoch + num_epochs): |
| model.train() |
| total_loss = 0 |
| total_batches = len(data_loader) |
| batch_counter = 0 |
| last_cooldown = time.time() |
| |
| progress_bar = tqdm(data_loader, desc=f"Epoch {epoch+1}/{start_epoch + num_epochs}") |
| |
| for input_seq, target_seq in progress_bar: |
| try: |
| |
| if hardware_guard.check_temperature(): |
| hardware_guard.cooldown_protocol() |
| last_cooldown = time.time() |
| |
| |
| if time.time() - last_cooldown > 1800: |
| print("🕒 定期冷却激活") |
| hardware_guard.cooldown_protocol() |
| last_cooldown = time.time() |
| |
| |
| input_seq, target_seq = input_seq.to(device), target_seq.to(device) |
| |
| |
| if input_seq.size(1) == 0: |
| continue |
| |
| |
| logits = model(input_seq) |
| |
| |
| logits_flat = logits.view(-1, logits.size(-1)) |
| targets_flat = target_seq.view(-1) |
| loss = criterion(logits_flat, targets_flat) |
| |
| |
| loss.backward() |
| |
| |
| torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) |
| |
| |
| optimizer.step() |
| optimizer.zero_grad() |
| |
| total_loss += loss.item() |
| batch_counter += 1 |
| avg_loss = total_loss / batch_counter |
| progress_bar.set_postfix(loss=avg_loss, lr=optimizer.param_groups[0]['lr']) |
| |
| |
| if device.type == "mps" and batch_counter % 50 == 0: |
| torch.mps.empty_cache() |
| |
| except RuntimeError as e: |
| if 'out of memory' in str(e).lower(): |
| print(f"显存不足,尝试减小批次大小...") |
| torch.mps.empty_cache() if device.type == "mps" else torch.cuda.empty_cache() |
| |
| |
| batch_size = max(4, batch_size // 2) |
| print(f"新批次大小: {batch_size}") |
| |
| |
| data_loader = DataLoader( |
| dataset, |
| batch_size=batch_size, |
| shuffle=True, |
| num_workers=num_workers, |
| pin_memory=(device.type != "mps"), |
| collate_fn=padder |
| ) |
| progress_bar = tqdm(data_loader, desc=f"Epoch {epoch+1}/{start_epoch + num_epochs}") |
| else: |
| print(f"训练出错: {str(e)},跳过该批次") |
| torch.mps.empty_cache() if device.type == "mps" else torch.cuda.empty_cache() |
| |
| |
| avg_epoch_loss = total_loss / batch_counter if batch_counter > 0 else 0 |
| train_losses.append(avg_epoch_loss) |
| print(f"Epoch {epoch + 1}/{start_epoch + num_epochs}, 平均损失: {avg_epoch_loss:.4f}") |
| |
| |
| scheduler.step(avg_epoch_loss) |
| |
| |
| current_lr = optimizer.param_groups[0]['lr'] |
| print(f"当前学习率: {current_lr:.2e}") |
| |
| |
| torch.save({ |
| "model_state_dict": model.state_dict(), |
| "optimizer_state_dict": optimizer.state_dict(), |
| "epoch": epoch, |
| "batch_size": batch_size |
| }, model_path) |
| print(f"模型已保存至 {model_path}") |
| |
| |
| print("\n本轮训练后生成样本:") |
| start_text = "The future of" |
| generated_text = generate_text( |
| model, |
| start_text, |
| max_len=50, |
| tokenizer=tokenizer, |
| temperature=0.8 |
| ) |
| print(f"输入: {start_text}") |
| print(f"生成: {generated_text}") |
| |
| |
| with open(log_file, "a") as f: |
| f.write(f"\nEpoch {epoch + 1}/{start_epoch + num_epochs}\n") |
| f.write(f"训练损失: {avg_epoch_loss:.4f}\n") |
| f.write(f"学习率: {current_lr:.2e}\n") |
| f.write(f"批次大小: {batch_size}\n") |
| f.write(f"输入: {start_text}\n") |
| f.write(f"生成: {generated_text}\n") |
| f.write("="*80 + "\n") |
| |
| |
| if device.type == "mps": |
| torch.mps.empty_cache() |
| print(f"当前内存使用: {psutil.virtual_memory().percent}%") |
| |
| |
| if avg_epoch_loss < best_loss: |
| best_loss = avg_epoch_loss |
| patience_counter = 0 |
| |
| torch.save(model.state_dict(), "best_model.pth") |
| print(f"新最佳模型保存至 best_model.pth (损失: {best_loss:.4f})") |
| else: |
| patience_counter += 1 |
| if patience_counter >= patience: |
| print(f"验证损失连续 {patience} 轮未提升,提前停止训练") |
| break |
| |
| |
| plt.figure(figsize=(10, 6)) |
| plt.plot(range(start_epoch, start_epoch + len(train_losses)), train_losses, label="training loss") |
| plt.xlabel("epochs") |
| plt.ylabel("loss") |
| plt.title("training loss") |
| plt.legend() |
| plt.grid(True) |
| plt.savefig("training_loss.png") |
| print("训练损失曲线已保存至 training_loss.png") |
| |
| |
| model.load_state_dict(torch.load("best_model.pth", map_location=device)) |
| print("已加载最佳模型") |
| |
| return model, tokenizer |
|
|
| |
| |
| |
|
|
| def generate_text(model, start_text, max_len, tokenizer, temperature=0.8): |
| model.eval() |
| |
| |
| start_text = clean_text(start_text) |
| |
| |
| input_ids = tokenizer.encode(start_text).ids |
| input_tensor = torch.tensor([input_ids], dtype=torch.long).to(device) |
| |
| generated_ids = input_ids.copy() |
| |
| for _ in range(max_len): |
| try: |
| with torch.no_grad(): |
| |
| if input_tensor.size(1) > 100: |
| input_tensor = input_tensor[:, -100:] |
| |
| |
| logits = model(input_tensor) |
| next_token_logits = logits[:, -1, :] / temperature |
| |
| |
| probs = torch.softmax(next_token_logits, dim=-1) |
| |
| |
| probs[probs < 0.01] = 0 |
| probs = probs / probs.sum() |
| |
| |
| next_token = torch.multinomial(probs, num_samples=1).item() |
| |
| |
| if next_token == tokenizer.token_to_id("<SEP>"): |
| break |
| |
| |
| generated_ids.append(next_token) |
| next_token_tensor = torch.tensor([[next_token]], device=device, dtype=torch.long) |
| input_tensor = torch.cat([input_tensor, next_token_tensor], dim=1) |
| |
| except Exception as e: |
| print(f"生成时出错: {str(e)}") |
| break |
| |
| |
| return tokenizer.decode(generated_ids) |
|
|
| |
| |
| |
|
|
| if __name__ == "__main__": |
| |
| text_list = load_dataset_text( |
| dataset_name="BEE-spoke-data/fineweb-1M_en-med", |
| max_samples=50000 |
| ) |
| |
| |
| model, tokenizer = train_model( |
| text_list, |
| embed_dim=256, |
| hidden_dim=512, |
| num_layers=16, |
| batch_size=64, |
| num_epochs=159, |
| lr=1e-4, |
| model_path="model.pth", |
| tokenizer_path="bpe_tokenizer.json" |
| ) |
|
|
| |
| print("\n最终文本生成示例:") |
| start_text = "The future of artificial intelligence" |
| generated_text = generate_text( |
| model, |
| start_text, |
| max_len=100, |
| tokenizer=tokenizer, |
| temperature=0.8 |
| ) |
| print(f"输入: {start_text}") |
| print(f"生成: {generated_text}") |
| |
| |
| with open("final_generation.txt", "w") as f: |
| f.write(f"输入: {start_text}\n") |
| f.write(f"生成: {generated_text}\n") |
| f.write(f"生成时间: {datetime.datetime.now()}\n") |
| if device.type == "mps": |
| f.write(f"最终内存使用: {psutil.virtual_memory().percent}%\n") |
| |