-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain_dl.py
More file actions
136 lines (111 loc) · 4.74 KB
/
Copy pathtrain_dl.py
File metadata and controls
136 lines (111 loc) · 4.74 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
#!/usr/bin/env python3
"""Deep learning training entry point.
Usage:
python train_dl.py --data-dir data --epochs 20 --batch-size 4
"""
from __future__ import annotations
import argparse
from pathlib import Path
import torchaudio
import torch
from torch.utils.data import DataLoader
import pytorch_lightning as pl
from pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping, LearningRateMonitor
import os
from src.ml.deep_model import Wav2VecClassifier
from pytorch_lightning.loggers import TensorBoardLogger
logger = TensorBoardLogger("lightning_logs", name="wav2vec_classifier")
# For prepare_dataset
def prepare_dataset(root: Path, sample_rate: int = 16000):
"""
Scans dataset directory and returns list of (file_path, label_idx) tuples.
"""
label_dirs = [d for d in root.iterdir() if d.is_dir() and not d.name.startswith('.')]
label_to_idx = {d.name: i for i, d in enumerate(sorted(label_dirs))}
items = []
for label_dir in label_dirs:
for audio_file in label_dir.glob("**/*"):
if audio_file.suffix.lower() not in {".wav", ".mp3", ".m4a"}:
continue
items.append((audio_file, label_to_idx[label_dir.name]))
return items, label_to_idx
# For AudioDataset
class AudioDataset(torch.utils.data.Dataset):
"""
Dataset for loading audio files with optional fixed-length cropping/padding.
"""
def __init__(self, items, sample_rate: int = 16000, max_len_sec: float | None = None):
self.items = items
self.sample_rate = sample_rate
self.max_len_samples = None
if max_len_sec is not None:
self.max_len_samples = int(max_len_sec * sample_rate)
def __len__(self):
return len(self.items)
# For __getitem__
def __getitem__(self, idx):
"""
Loads audio file, resamples, converts to mono, and applies length constraints.
"""
path, label = self.items[idx]
waveform, sr = torchaudio.load(str(path))
if sr != 16000:
waveform = torchaudio.functional.resample(waveform, sr, 16000)
# mono
if waveform.shape[0] > 1:
waveform = waveform.mean(dim=0, keepdim=True)
# Crop or pad to max length if set
if self.max_len_samples is not None:
if waveform.shape[-1] > self.max_len_samples:
waveform = waveform[: self.max_len_samples]
elif waveform.shape[-1] < self.max_len_samples:
pad_len = self.max_len_samples - waveform.shape[-1]
waveform = torch.nn.functional.pad(waveform, (0, pad_len))
return waveform, label
# For collate_fn
def collate_fn(batch):
"""
Stacks waveforms and labels, removing channel dimension.
"""
waveforms, labels = zip(*batch)
waveforms = [w.squeeze(0) for w in waveforms] # remove channel dim
return torch.stack(waveforms), torch.tensor(labels)
# For main
def main():
"""
Trains Wav2VecClassifier using Lightning Trainer with checkpoints and early stopping.
"""
parser = argparse.ArgumentParser()
parser.add_argument("--data-dir", type=str, default="data")
parser.add_argument("--epochs", type=int, default=20)
parser.add_argument("--batch-size", type=int, default=1)
parser.add_argument("--lr", type=float, default=1e-4)
parser.add_argument("--device", type=str, default="cpu", choices=["cpu", "mps", "cuda", "auto"], help="Accelerator to use")
parser.add_argument("--max-sec", type=float, default=5.0, help="Crop/pad each clip to this many seconds")
args = parser.parse_args()
data_dir = Path(args.data_dir)
items, label_to_idx = prepare_dataset(data_dir)
num_classes = len(label_to_idx)
# simple split 80/20
split = int(0.8 * len(items))
train_items = items[:split]
val_items = items[split:]
train_ds = AudioDataset(train_items, max_len_sec=args.max_sec)
val_ds = AudioDataset(val_items, max_len_sec=args.max_sec)
train_loader = DataLoader(train_ds, batch_size=args.batch_size, shuffle=True, collate_fn=collate_fn, num_workers=0)
val_loader = DataLoader(val_ds, batch_size=args.batch_size, shuffle=False, collate_fn=collate_fn, num_workers=0)
model = Wav2VecClassifier(num_classes=num_classes, lr=args.lr)
ckpt_cb = ModelCheckpoint(monitor="val_acc", mode="max", save_top_k=1)
early_stop = EarlyStopping(monitor="val_acc", mode="max", patience=5)
lr_monitor = LearningRateMonitor(logging_interval="epoch")
trainer = pl.Trainer(
max_epochs=args.epochs,
accelerator=args.device,
callbacks=[ckpt_cb, early_stop, lr_monitor],
precision=32,
logger=logger
)
trainer.fit(model, train_loader, val_loader)
print("Best model saved at", ckpt_cb.best_model_path)
if __name__ == "__main__":
main()