import torch from torch import Tensor, device, no_grad, cuda from torch.optim import AdamW, Optimizer from torch.optim.lr_scheduler import OneCycleLR from torch.nn import CTCLoss, utils from torch.utils.data import DataLoader from torch.utils.tensorboard import SummaryWriter from pathlib import Path from tqdm import tqdm from datetime import datetime from torch.amp.grad_scaler import GradScaler from tokenizer import ASRTokenizer from dataset import Batch, create_dataloader from model import ASRModel # ============ 全局配置 ============ CONFIG = { # 数据配置 'batch_size': 96, # 训练配置 'num_epochs': 50, 'learning_rate': 1e-4, 'weight_decay': 1e-4, 'grad_clip_norm': 1.0, # 模型配置 'input_dim': 640, 'num_heads': 8, 'ffn_dim': 2048, 'num_layers': 8, 'dropout': 0.1, # 保存和评估 'early_stopping_patience': 12, } def calculate_cer(pred: str, target: str) -> float: """ 计算字符错误率 (Character Error Rate) 使用编辑距离算法""" if len(target) == 0: return 0.0 if len(pred) == 0 else 1.0 d = [[0] * (len(target) + 1) for _ in range(len(pred) + 1)] for i in range(len(pred) + 1): d[i][0] = i for j in range(len(target) + 1): d[0][j] = j for i in range(1, len(pred) + 1): for j in range(1, len(target) + 1): if pred[i-1] == target[j-1]: d[i][j] = d[i-1][j-1] else: d[i][j] = min(d[i-1][j], d[i][j-1], d[i-1][j-1]) + 1 return d[len(pred)][len(target)] / len(target) def calculate_wer(pred: str, target: str) -> float: """ 计算词错误率 (Word Error Rate) """ pred_words = pred.split() target_words = target.split() if len(target_words) == 0: return 0.0 if len(pred_words) == 0 else 1.0 d = [[0] * (len(target_words) + 1) for _ in range(len(pred_words) + 1)] for i in range(len(pred_words) + 1): d[i][0] = i for j in range(len(target_words) + 1): d[0][j] = j for i in range(1, len(pred_words) + 1): for j in range(1, len(target_words) + 1): if pred_words[i-1] == target_words[j-1]: d[i][j] = d[i-1][j-1] else: d[i][j] = min(d[i-1][j], d[i][j-1], d[i-1][j-1]) + 1 return d[len(pred_words)][len(target_words)] / len(target_words) def train_one_epoch(model: ASRModel, dataloader: DataLoader, criterion: CTCLoss, optimizer: Optimizer, scheduler: OneCycleLR, device: device, epoch: int, writer: SummaryWriter, global_step: int, scaler: GradScaler) -> tuple[float, int]: model.train() num_batches = 0 total_train_loss = 0 progress_bar = tqdm(dataloader, desc=f"Epoch {epoch}") for batch_index, batch in enumerate(progress_bar): batch: Batch mel_specs = batch['mel_specs'].to(device) targets = batch['targets'].to(device) mel_lengths = batch['mel_lengths'].to(device) target_lengths = batch['target_lengths'].to(device) optimizer.zero_grad() with torch.autocast(device_type='cuda', dtype=torch.bfloat16): log_probs, lengths = model(mel_specs=mel_specs, mel_lengths=mel_lengths) log_probs: Tensor log_probs_ctc = log_probs.permute(1, 0, 2) loss: Tensor = criterion(log_probs=log_probs_ctc, targets=targets, input_lengths=lengths, target_lengths=target_lengths) scaler.scale(loss).backward() scaler.unscale_(optimizer) utils.clip_grad_norm_(model.parameters(), max_norm=CONFIG['grad_clip_norm']) scaler.step(optimizer) scaler.update() scheduler.step() total_train_loss += loss.item() num_batches += 1 global_step += 1 writer.add_scalar('Train/Loss', loss.item(), global_step) writer.add_scalar('Train/LearningRate', optimizer.param_groups[0]['lr'], global_step) progress_bar.set_postfix({'epoch': f"{epoch}/{CONFIG['num_epochs']}",'loss': f'{loss.item():.4f}', 'step': global_step}) train_avg_loss = total_train_loss / num_batches return train_avg_loss, global_step def validate(model: ASRModel, dataloader: DataLoader, criterion: CTCLoss, device: device, tokenizer: ASRTokenizer, writer: SummaryWriter, global_step: int) -> tuple[float, float, float]: model.eval() total_loss = 0 total_cer = 0 total_wer = 0 num_samples = 0 num_batches = 0 examples = [] with no_grad(): progress_bar = tqdm(dataloader, desc="Validate", leave=False) for batch in progress_bar: batch: Batch mel_specs = batch['mel_specs'].to(device) targets = batch['targets'].to(device) mel_lengths = batch['mel_lengths'].to(device) target_lengths = batch['target_lengths'].to(device) with torch.autocast(device_type='cuda', dtype=torch.bfloat16): log_probs, lengths = model(mel_specs=mel_specs, mel_lengths=mel_lengths) log_probs: Tensor log_probs_ctc = log_probs.permute(1, 0, 2) loss: Tensor = criterion(log_probs=log_probs_ctc, targets=targets, input_lengths=lengths, target_lengths=target_lengths) total_loss += loss.item() for i in range(log_probs.shape[0]): pred_text = tokenizer.ctc_greedy_decode(log_probs=log_probs[i]) true_text = batch['target_texts'][i] cer = calculate_cer(pred=pred_text, target=true_text) wer = calculate_wer(pred=pred_text, target=true_text) total_cer += cer total_wer += wer num_samples += 1 if len(examples) < 3: examples.append((true_text, pred_text, cer, wer)) num_batches += 1 avg_loss = total_loss / num_batches avg_cer = total_cer / num_samples avg_wer = total_wer / num_samples writer.add_scalar('Val/Loss', avg_loss, global_step) writer.add_scalar('Val/CER', avg_cer, global_step) writer.add_scalar('Val/WER', avg_wer, global_step) for index, (true_text, pred_text, cer, wer) in enumerate(examples): writer.add_text(f'Val/Example_{index}', f'True: {true_text}\nPred: {pred_text}\nCER: {cer:.4f} | WER: {wer:.4f}', global_step) model.train() return avg_loss, avg_cer, avg_wer def save_checkpoint(model: ASRModel, optimizer: Optimizer, scheduler: OneCycleLR, global_step: int, epoch: int, train_loss: float, val_loss: float, cer: float, wer: float, save_path: Path): checkpoint = { 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'scheduler_state_dict': scheduler.state_dict(), 'global_step': global_step, 'train_loss': train_loss, 'val_loss': val_loss, 'epoch': epoch, 'cer': cer, 'wer': wer, } torch.save(checkpoint, save_path) def find_latest_checkpoint(checkpoint_dir: Path) -> Path | None: checkpoints = sorted(checkpoint_dir.glob('checkpoint_epoch_*.pt'), key=lambda p: int(p.stem.split('_')[-1])) return checkpoints[-1] if checkpoints else None def load_checkpoint(file_path: Path, model: ASRModel, optimizer: AdamW, scheduler: OneCycleLR): checkpoint = torch.load(file_path, weights_only=False) model.load_state_dict(checkpoint['model_state_dict']) optimizer.load_state_dict(checkpoint['optimizer_state_dict']) scheduler.load_state_dict(checkpoint['scheduler_state_dict']) epoch = checkpoint['epoch'] global_step = checkpoint['global_step'] best_cer = checkpoint['cer'] best_wer = checkpoint['wer'] return epoch, global_step, best_cer, best_wer def main(): workspace_dir = Path(__file__).parent.parent device = torch.device('cuda:0') tokenizer = ASRTokenizer(workspace_dir / 'config/asr_vocab.json') # ============ 创建数据加载器 ============ train_loader = create_dataloader( tsv_path=workspace_dir / '.data/ug/train.tsv', audio_dir=workspace_dir / '.data/ug/clips', tokenizer=tokenizer, batch_size=CONFIG['batch_size'], shuffle=True, augment=True ) val_loader = create_dataloader( tsv_path=workspace_dir / '.data/ug/dev.tsv', audio_dir=workspace_dir / '.data/ug/clips', tokenizer=tokenizer, batch_size=CONFIG['batch_size'], shuffle=False, augment=False ) # ============ 初始化模型 ============ model = ASRModel( vocab_size=tokenizer.vocab_size(), input_dim=CONFIG['input_dim'], num_heads=CONFIG['num_heads'], ffn_dim=CONFIG['ffn_dim'], num_layers=CONFIG['num_layers'], dropout=CONFIG['dropout'], ).to(device) print(f"🤖 模型参数量: {model.get_num_params() / 1e6:.2f}M") criterion = CTCLoss(blank=tokenizer.get_special_token_id(''), zero_infinity=True) optimizer = AdamW(model.parameters(), lr=CONFIG['learning_rate'], weight_decay=CONFIG['weight_decay'], foreach=True) # 使用 OneCycleLR:自动处理 warmup 和衰减 scheduler = OneCycleLR( optimizer, max_lr=CONFIG['learning_rate'], epochs=CONFIG['num_epochs'], steps_per_epoch=len(train_loader), pct_start=0.1, # 前 10% 步数用于 warmup anneal_strategy='cos', div_factor=25.0, # 初始 lr = max_lr / 25 final_div_factor=1e4, # 最终 lr = max_lr / 10000 ) scaler = GradScaler() # ============ TensorBoard ============ log_dir = workspace_dir / 'runs' / datetime.now().strftime('%Y%m%d_%H%M%S') writer = SummaryWriter(log_dir) config_text = '\n'.join([f'{k}: {v}' for k, v in CONFIG.items()]) writer.add_text('Config', config_text, 0) print(f"📊 TensorBoard 日志: {log_dir}") checkpoint_dir = workspace_dir / '.checkpoints' checkpoint_dir.mkdir(exist_ok=True) # ============ 训练循环 ============ best_cer = float('inf') best_wer = float('inf') patience_counter = 0 global_step = 0 start_epoch = 0 latest = find_latest_checkpoint(checkpoint_dir) if latest: start_epoch, global_step, best_cer, best_wer = load_checkpoint(latest, model, optimizer, scheduler) start_epoch += 1 # 从下一个 epoch 开始 print(f"✅ 恢复训练: 从 epoch={start_epoch} 开始, global_step={global_step}") else: start_epoch = 0 print("🆕 从头开始训练\n") for epoch in range(start_epoch, CONFIG['num_epochs']): train_loss, global_step = train_one_epoch( model=model, dataloader=train_loader, criterion=criterion, optimizer=optimizer, scheduler=scheduler, device=device, epoch=epoch, writer=writer, global_step=global_step, scaler=scaler, ) val_loss, val_cer, val_wer = validate(model=model, dataloader=val_loader, criterion=criterion, device=device, tokenizer=tokenizer, writer=writer, global_step=global_step) print(f"\n📊 Step {global_step} | Val Loss: {val_loss:.4f} | Val CER: {val_cer:.4f} | Val WER: {val_wer:.4f}") # 保存常规 checkpoint checkpoint_path = checkpoint_dir / f'checkpoint_epoch_{epoch}.pt' save_checkpoint(model=model, optimizer=optimizer, scheduler=scheduler, epoch=epoch, global_step=global_step, train_loss=train_loss, val_loss=val_loss, cer=val_cer, wer=val_wer, save_path=checkpoint_path) # 分别保存最佳 CER 和 WER 模型 improved = False if val_cer < best_cer: best_cer = val_cer best_cer_path = checkpoint_dir / 'best_cer_model.pt' save_checkpoint(model=model, optimizer=optimizer, scheduler=scheduler, epoch=epoch, global_step=global_step, train_loss=train_loss, val_loss=val_loss, cer=val_cer, wer=val_wer, save_path=best_cer_path) improved = True if val_wer < best_wer: best_wer = val_wer best_wer_path = checkpoint_dir / 'best_wer_model.pt' save_checkpoint(model=model, optimizer=optimizer, scheduler=scheduler, epoch=epoch, global_step=global_step, train_loss=train_loss, val_loss=val_loss, cer=val_cer, wer=val_wer, save_path=best_wer_path) improved = True # Early Stopping 逻辑 if improved: patience_counter = 0 else: patience_counter += 1 print(f"⚠️ 验证指标未改善,patience: {patience_counter}/{CONFIG['early_stopping_patience']}") # 删除旧的 checkpoint(保留最近3个) old_checkpoints = sorted(checkpoint_dir.glob('checkpoint_epoch_*.pt'), key=lambda p: int(p.stem.split('_')[-1])) for old in old_checkpoints[:-3]: old.unlink() cuda.empty_cache() writer.add_scalar('Val/EpochLoss', val_loss, epoch) writer.add_scalar('Train/EpochLoss', train_loss, epoch) print(f"✅ Epoch {epoch} 完成 | Train Avg Loss: {train_loss:.4f} | Val Avg Loss: {val_loss:.4f} | Best CER: {best_cer:.4f} | Best WER: {best_wer:.4f}\n") # Early Stopping 检查 if patience_counter >= CONFIG['early_stopping_patience']: print(f"\n🛑 Early Stopping: 验证指标连续 {CONFIG['early_stopping_patience']} 次没有改善") print(f"🏆 最佳 CER: {best_cer:.4f}") print(f"🏆 最佳 WER: {best_wer:.4f}") break writer.close() print(f"\n{'='*60}") print(f"🎉 训练完成!") print(f"🏆 最佳 CER: {best_cer:.4f}") print(f"🏆 最佳 WER: {best_wer:.4f}") print(f"📊 TensorBoard: tensorboard --logdir={workspace_dir / 'runs'}") print(f"{'='*60}\n") if __name__ == "__main__": main()