| 242 | # ----------------------------------------------------------------------------- |
| 243 | |
| 244 | def parse_args() -> argparse.Namespace: |
| 245 | parser = argparse.ArgumentParser(description="Train Conformer on LibriSpeech") |
| 246 | # === dataset === |
| 247 | parser.add_argument("--root", type=str, default="/shared/LIBRISPEECH", help="LibriSpeech root directory") |
| 248 | parser.add_argument( |
| 249 | "--train-sets", |
| 250 | type=str, |
| 251 | default="train-clean-100,train-clean-360,train-other-500", |
| 252 | help="Comma‑separated subset names for training", |
| 253 | ) |
| 254 | parser.add_argument("--valid-set", type=str, default="dev-clean", help="Validation subset name") |
| 255 | # === training hyper‑params === |
| 256 | parser.add_argument("--epochs", type=int, default=50, help="Number of training epochs") |
| 257 | parser.add_argument("--batch-size", type=int, default=64, |
| 258 | help="Mini-batch per gradient update") |
| 259 | parser.add_argument("--lr", type=float, default=4e-4, |
| 260 | help="Peak LR after warm-up") |
| 261 | parser.add_argument("--save-dir", type=str, default="/shared/conformer/checkpoints", |
| 262 | help="Where to write epoch checkpoints") |
| 263 | parser.add_argument( |
| 264 | "--betas", |
| 265 | type=str, |
| 266 | default="0.9,0.999", |
| 267 | help="Adam betas as a comma‑separated pair", |
| 268 | ) |
| 269 | parser.add_argument("--weight-decay", type=float, default=1e-4, help="Adam weight decay") |
| 270 | parser.add_argument("--warmup-epochs", type=float, default=1.0, help="Noam LR warm‑up steps") |
| 271 | parser.add_argument("--grad-clip", type=float, default=5.0, help="Gradient clipping threshold (L2 norm)") |
| 272 | # --- gradient-accumulation --- |
| 273 | parser.add_argument( |
| 274 | "--accum-steps", |
| 275 | type=int, |
| 276 | default=1, |
| 277 | help="How many mini-batches to accumulate gradients over before " |
| 278 | "performing an optimizer update (≃ effective batch-size multiplier)", |
| 279 | ) |
| 280 | parser.add_argument("--num-workers", type=int, default=8, help="DataLoader workers for training") |
| 281 | parser.add_argument("--val-num-workers", type=int, default=4, help="DataLoader workers for validation") |
| 282 | # === data augmentation / front‑end === |
| 283 | parser.add_argument("--sample-rate", type=int, default=16_000) |
| 284 | parser.add_argument("--n-mels", type=int, default=80) |
| 285 | parser.add_argument("--n-fft", type=int, default=512) |
| 286 | parser.add_argument("--hop-length", type=int, default=160) |
| 287 | parser.add_argument("--time-mask-ratio", type=float, default=0.05) |
| 288 | parser.add_argument( |
| 289 | "--speeds", |
| 290 | type=str, |
| 291 | default="0.9,1.0,1.1", |
| 292 | help="Comma‑separated speed perturb factors", |
| 293 | ) |
| 294 | parser.add_argument( |
| 295 | "--no-augment", |
| 296 | action="store_true", |
| 297 | help="If set, disables SpecAugment & speed perturb during training" |
| 298 | ) |
| 299 | parser.add_argument( |
| 300 | "--sp-model", |
| 301 | type=str, |