-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
128 lines (111 loc) · 3.5 KB
/
Copy pathtrain.py
File metadata and controls
128 lines (111 loc) · 3.5 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
#!/usr/bin/env python3
"""
Model train routines.
Authors:
LICENCE:
"""
import os
from argparse import Namespace
import optuna
import torch
from optuna.trial import TrialState
from torch import nn
def _train_loop(
model: nn.Module,
dataloader: torch.utils.data.DataLoader,
args: Namespace,
optimizer: torch.optim.Optimizer,
loss_fn: torch.nn.CrossEntropyLoss,
) -> float:
model.train()
epoch_loss = 0.0
log_interval = len(dataloader) // 4
for batch_idx, data in enumerate(dataloader):
states, action = data
states = states.to(args.device)
action = action.to(args.device, dtype=torch.long)
optimizer.zero_grad(set_to_none=True)
output = model(states)
loss = loss_fn(output, action)
if (batch_idx + 1) % log_interval == 0:
print(f"Batch {batch_idx+1} Loss: {loss}")
loss.backward()
optimizer.step()
epoch_loss += loss.item()
return epoch_loss
def _get_optimizer(model: nn.Module, args: Namespace) -> torch.optim.Optimizer:
if args.opt == "sgd":
optimizer = torch.optim.SGD(
model.parameters(),
lr=args.learning_rate,
momentum=0.9,
nesterov=True,
)
elif args.opt == "adam":
optimizer = torch.optim.Adam(
model.parameters(),
lr=args.learning_rate,
)
return optimizer
def trainer(
args: Namespace,
):
"""Return function."""
return tune if args.tune else train
def tune(
model: nn.Module,
dataloader: torch.utils.data.DataLoader,
args: Namespace,
):
"""Tune hp."""
def objective(trial):
"""Minimize train loss."""
args.lr = trial.suggest_float("lr", 1e-5, 1e-1, log=True)
loss = 0.0
loss_fn = nn.CrossEntropyLoss()
optimizer = _get_optimizer(model, args)
for epoch in range(1, args.epochs + 1):
print(f"Epoch: {epoch}")
loss += _train_loop(model, dataloader, args, optimizer, loss_fn)
trial.report(loss, epoch)
# Handle pruning based on the intermediate value.
if trial.should_prune():
raise optuna.exceptions.TrialPruned()
return loss
study = optuna.create_study(direction=("minimize"))
study.optimize(objective, n_trials=(100))
pruned_trials = study.get_trials(
deepcopy=False,
states=[TrialState.PRUNED],
)
complete_trials = study.get_trials(
deepcopy=False,
states=[TrialState.COMPLETE],
)
print("Study statistics: ")
print("\tNumber of finished trials: ", len(study.trials))
print("\tNumber of pruned trials: ", len(pruned_trials))
print("\tNumber of complete trials: ", len(complete_trials))
print("Best trial:")
trial = study.best_trial
print("\tValue: ", trial.value)
print("\tParams: ")
for key, value in trial.params.items():
print(f"{key}: {value}")
def train(
model: nn.Module,
dataloader: torch.utils.data.DataLoader,
args: Namespace,
):
"""Train model."""
loss_fn = nn.CrossEntropyLoss()
if not os.path.exists(args.model_path / args.train_run_name):
os.mkdir(args.model_path / args.train_run_name)
optimizer = _get_optimizer(model, args)
for epoch in range(1, args.epochs + 1):
print(f"Epoch: {epoch}")
_train_loop(model, dataloader, args, optimizer, loss_fn)
torch.save(
model.state_dict(),
args.model_path / args.train_run_name / "model.pth",
)