add augment noise and num layer=6
This commit is contained in:
+46
-26
@@ -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']:
|
||||
|
||||
Reference in New Issue
Block a user