first commit asr model
This commit is contained in:
+357
@@ -0,0 +1,357 @@
|
||||
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/train.tsv',
|
||||
audio_dir=workspace_dir / '.data/clips',
|
||||
tokenizer=tokenizer,
|
||||
batch_size=CONFIG['batch_size'],
|
||||
shuffle=True,
|
||||
)
|
||||
|
||||
val_loader = create_dataloader(
|
||||
tsv_path=workspace_dir / '.data/dev.tsv',
|
||||
audio_dir=workspace_dir / '.data/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()
|
||||
Reference in New Issue
Block a user