add augment noise and num layer=6

This commit is contained in:
2026-05-09 14:11:58 +06:00
parent d31233a79a
commit 96cd0a20cb
6 changed files with 155 additions and 211 deletions
+46 -26
View File
@@ -31,11 +31,11 @@ CONFIG = {
'input_dim': 256,
'num_heads': 8,
'ffn_dim': 2048,
'num_layers': 8,
'dropout': 0.1,
'num_layers': 6,
'dropout': 0.15,
# 保存和评估
'early_stopping_patience': 12,
'early_stopping_patience': 10,
}
@@ -168,17 +168,13 @@ def validate(model: ASRModel, dataloader: DataLoader, criterion: CTCLoss, device
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):
def save_checkpoint(model: ASRModel, optimizer: Optimizer, scheduler: OneCycleLR, global_step: int, epoch: int, train_loss: float, val_loss: float, cer: float, wer: float, current_prob: float, save_path: Path):
checkpoint = {
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
@@ -189,6 +185,7 @@ def save_checkpoint(model: ASRModel, optimizer: Optimizer, scheduler: OneCycleLR
'epoch': epoch,
'cer': cer,
'wer': wer,
'current_prob': current_prob,
}
torch.save(checkpoint, save_path)
@@ -206,14 +203,15 @@ def load_checkpoint(file_path: Path, model: ASRModel, optimizer: AdamW, schedule
global_step = checkpoint['global_step']
best_cer = checkpoint['cer']
best_wer = checkpoint['wer']
return epoch, global_step, best_cer, best_wer
current_prob = checkpoint['current_prob']
return epoch, global_step, best_cer, best_wer, current_prob
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
final_prob = 0.8
warmup_epochs = 12
current_prob = 0.0
# ============ 创建数据加载器 ============
@@ -221,6 +219,7 @@ def main():
tsv_path=workspace_dir / '.data/ug/train_new.tsv',
audio_dir=workspace_dir / '.data/ug/clips',
noise_dir='/mnt/dataset/dataset/audio/noise',
corridor_noise_dir= workspace_dir / 'data/corridor',
tokenizer=tokenizer,
batch_size=CONFIG['batch_size'],
shuffle=True,
@@ -228,15 +227,28 @@ def main():
augment_prob=current_prob
)
val_loader = create_dataloader(
val_loader_clean = 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',
corridor_noise_dir= workspace_dir / 'data/corridor',
tokenizer=tokenizer,
batch_size=CONFIG['batch_size'],
shuffle=False,
augment=False
)
val_loader_noisy = 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',
corridor_noise_dir= workspace_dir / 'data/corridor',
tokenizer=tokenizer,
batch_size=CONFIG['batch_size'],
shuffle=False,
augment=True,
augment_prob=0.5
)
# ============ 初始化模型 ============
model = ASRModel(
@@ -259,9 +271,9 @@ def main():
max_lr=CONFIG['learning_rate'],
epochs=CONFIG['num_epochs'],
steps_per_epoch=len(train_loader),
pct_start=0.15,
pct_start=0.2,
anneal_strategy='cos',
div_factor=10.0, # 初始 lr = max_lr / 10
div_factor=25.0, # 初始 lr = max_lr / 10
final_div_factor=1e4, # 最终 lr = max_lr / 10000
)
scaler = GradScaler()
@@ -286,7 +298,7 @@ def main():
latest = find_latest_checkpoint(checkpoint_dir)
if latest:
start_epoch, global_step, best_cer, best_wer = load_checkpoint(latest, model, optimizer, scheduler)
start_epoch, global_step, best_cer, best_wer, current_prob = load_checkpoint(latest, model, optimizer, scheduler)
start_epoch += 1 # 从下一个 epoch 开始
print(f"✅ 恢复训练: 从 epoch={start_epoch} 开始, global_step={global_step}")
else:
@@ -308,25 +320,26 @@ def main():
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}")
val_loss_c, val_cer_c, val_wer_c = validate(model=model, dataloader=val_loader_clean, criterion=criterion, device=device, tokenizer=tokenizer, writer=writer, global_step=global_step)
val_loss_n, val_cer_n, val_wer_n = validate(model=model, dataloader=val_loader_noisy, criterion=criterion, device=device, tokenizer=tokenizer, writer=writer, global_step=global_step)
print(f"\n📊 Step {global_step} | Clean Val Loss: {val_loss_c:.4f} | Clean Val CER: {val_cer_c:.4f} | Clean Val WER: {val_wer_c:.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)
save_checkpoint(model=model, optimizer=optimizer, scheduler=scheduler, epoch=epoch, global_step=global_step, train_loss=train_loss, val_loss=val_loss_c, cer=val_cer_c, wer=val_wer_c, current_prob=current_prob, save_path=checkpoint_path)
# 分别保存最佳 CER 和 WER 模型
improved = False
if val_cer < best_cer:
best_cer = val_cer
if val_cer_c < best_cer:
best_cer = val_cer_c
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)
save_checkpoint(model=model, optimizer=optimizer, scheduler=scheduler, epoch=epoch, global_step=global_step, train_loss=train_loss, val_loss=val_loss_c, cer=val_cer_c, wer=val_wer_c, current_prob=current_prob, save_path=best_cer_path)
improved = True
if val_wer < best_wer:
best_wer = val_wer
if val_wer_c < best_wer:
best_wer = val_wer_c
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)
save_checkpoint(model=model, optimizer=optimizer, scheduler=scheduler, epoch=epoch, global_step=global_step, train_loss=train_loss, val_loss=val_loss_c, cer=val_cer_c, wer=val_wer_c, current_prob=current_prob, save_path=best_wer_path)
improved = True
# Early Stopping 逻辑
@@ -343,9 +356,16 @@ def main():
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")
writer.add_scalar('Val/EpochLoss', val_loss_c, epoch)
writer.add_scalar('Val/Loss', val_loss_c, global_step)
writer.add_scalar('Val/CER', val_cer_c, global_step)
writer.add_scalar('Val/WER', val_wer_c, global_step)
writer.add_scalar('Val_Noisy/EpochLoss', val_loss_n, epoch)
writer.add_scalar('Val_Noisy/Loss', val_loss_n, global_step)
writer.add_scalar('Val_Noisy/CER', val_cer_n, global_step)
writer.add_scalar('Val_Noisy/WER', val_wer_n, global_step)
print(f"✅ Epoch {epoch} 完成 | Train Avg Loss: {train_loss:.4f} | Val Avg Loss: {val_loss_c:.4f} | Best CER: {best_cer:.4f} | Best WER: {best_wer:.4f}\n")
# Early Stopping 检查
if patience_counter >= CONFIG['early_stopping_patience']: