feat: change waveform

This commit is contained in:
2026-05-07 11:29:21 +06:00
commit d31233a79a
21 changed files with 5330 additions and 0 deletions
+367
View File
@@ -0,0 +1,367 @@
import math
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': 32,
# 训练配置
'num_epochs': 50,
'learning_rate': 2e-4,
'weight_decay': 1e-4,
'grad_clip_norm': 1.0,
# 模型配置
'input_dim': 256,
'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
waveforms = batch['waveforms'].to(device)
targets = batch['targets'].to(device)
waveform_lengths = batch['waveform_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(waveforms=waveforms, waveform_lengths=waveform_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
waveforms = batch['waveforms'].to(device)
targets = batch['targets'].to(device)
waveform_lengths = batch['waveform_lengths'].to(device)
target_lengths = batch['target_lengths'].to(device)
with torch.autocast(device_type='cuda', dtype=torch.bfloat16):
log_probs, lengths = model(waveforms=waveforms, waveform_lengths=waveform_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')
final_prob = 0.5
warmup_epochs = 8
current_prob = 0.0
# ============ 创建数据加载器 ============
train_loader = create_dataloader(
tsv_path=workspace_dir / '.data/ug/train_new.tsv',
audio_dir=workspace_dir / '.data/ug/clips',
noise_dir='/mnt/dataset/dataset/audio/noise',
tokenizer=tokenizer,
batch_size=CONFIG['batch_size'],
shuffle=True,
augment=True,
augment_prob=current_prob
)
val_loader = create_dataloader(
tsv_path=workspace_dir / '.data/ug/val_new.tsv',
audio_dir=workspace_dir / '.data/ug/clips',
noise_dir='/mnt/dataset/dataset/audio/noise',
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.15,
anneal_strategy='cos',
div_factor=10.0, # 初始 lr = max_lr / 10
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']):
current_prob = final_prob * (1 - math.cos(math.pi * epoch / warmup_epochs)) / 2
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()