-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
120 lines (101 loc) · 4 KB
/
Copy pathtrain.py
File metadata and controls
120 lines (101 loc) · 4 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
# coding=utf-8
import os
from omegaconf import DictConfig, OmegaConf
import hydra
from hydra.core.hydra_config import HydraConfig
import pytorch_lightning as pl
from pytorch_lightning.callbacks import ModelCheckpoint, LearningRateMonitor, Callback
from pytorch_lightning.loggers import TensorBoardLogger
import k2
import optuna
from optuna.integration import PyTorchLightningPruningCallback
class MetricsCallback(Callback):
"""Callback to track metrics for Optuna pruning."""
def __init__(self):
super().__init__()
self.metrics = []
def on_validation_end(self, trainer, pl_module):
current_per = trainer.callback_metrics.get("val/per", None)
if current_per is not None:
self.metrics.append(float(current_per))
@hydra.main(version_base=None, config_path="configs", config_name="cfg")
def train(cfg: DictConfig):
assert cfg.symbols.num_graphemes == len(
cfg.symbols.graphemes
), f"{cfg.symbols.num_graphemes} != {len(cfg.symbols.graphemes)}"
assert cfg.symbols.num_phonemes == len(
cfg.symbols.phonemes
), f"{cfg.symbols.num_phonemes} != {len(cfg.symbols.phonemes)}"
dataset = hydra.utils.instantiate(cfg.dataset)
from g2p_module import G2pModule
from data_module import G2pDataModule
datamodule = G2pDataModule(
dataset=dataset,
batch_size=cfg.train.batch_size,
eval_split=cfg.train.eval_split,
seed=cfg.train.get("seed", 42),
)
training_module = G2pModule(
model=cfg.model,
muon_lr=cfg.train.muon_lr,
adam_lr=cfg.train.adam_lr,
weight_decay=cfg.train.weight_decay,
lr_scheduler=cfg.train.lr_scheduler,
ilm_loss_weight=cfg.train.ilm_loss_weight,
ilm_scheduled_sampling_ratio=cfg.train.ilm_scheduled_sampling_ratio,
grad_clip=cfg.train.grad_clip,
loss_device=(
"cuda" if k2.with_cuda and cfg.train.accelerator == "gpu" else "cpu"
),
graphemes=OmegaConf.to_object(dataset.graphemes),
phonemes=OmegaConf.to_object(dataset.phonemes),
print_debug_preview=cfg.train.print_debug_preview,
logs_stressless_error_rates=cfg.train.get("logs_stressless_error_rates", False),
)
output_dir = HydraConfig.get().runtime.output_dir
checkpoint_callback = ModelCheckpoint(
dirpath=os.path.join(output_dir, "checkpoints"),
monitor="val/score",
mode="min",
save_top_k=1,
filename="g2p-best-score",
)
periodic_checkpoint = ModelCheckpoint(
dirpath=os.path.join(output_dir, "checkpoints"),
every_n_epochs=cfg.train.save_interval,
filename="g2p-{epoch:03d}",
save_top_k=-1,
)
periodic_checkpoint.CHECKPOINT_EQUALS_CHAR = "-"
lr_monitor = LearningRateMonitor(logging_interval="epoch")
callbacks_list = [checkpoint_callback, periodic_checkpoint, lr_monitor]
metrics_callback = MetricsCallback()
callbacks_list.append(metrics_callback)
trial = None
try:
trial = optuna.integration.get_current_trial()
if trial is not None:
print(f"Running Optuna trial {trial.number}")
pruning_callback = PyTorchLightningPruningCallback(trial, monitor="val/per")
callbacks_list.append(pruning_callback)
optuna_tb_callback = optuna.integration.TensorBoardCallback(
log_dir=output_dir, metric_name="val/per"
)
callbacks_list.append(optuna_tb_callback)
except:
pass
tensorboard_logger = TensorBoardLogger(save_dir=output_dir, name="", version="")
trainer = pl.Trainer(
max_epochs=cfg.train.epochs,
accelerator=cfg.train.accelerator,
devices=1,
callbacks=callbacks_list,
logger=tensorboard_logger,
gradient_clip_val=cfg.train.grad_clip,
log_every_n_steps=10,
enable_progress_bar=True,
)
trainer.fit(training_module, datamodule=datamodule)
return float(training_module.score_min) # optuna objective
if __name__ == "__main__":
train()