358 lines
14 KiB
Python
358 lines
14 KiB
Python
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('<BLANK>'), 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() |