-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathtrain.py
More file actions
58 lines (51 loc) · 1.73 KB
/
Copy pathtrain.py
File metadata and controls
58 lines (51 loc) · 1.73 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
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader
import argparse
from copy import deepcopy
import wandb
from src.model.flow_matching import ConditionalFlowMatching
from src.model.base_models.unet_mlp import TDFlowUnet
from src.training import TDFlowTrainer
from src.datasets import PointMassMazeDataset
def main(task = 'reach_top_left', num_epochs=100, loss_type='td2_cfm'):
dataset = PointMassMazeDataset(task=task)
train_loader = DataLoader(dataset=dataset, batch_size=1024, shuffle=True)
gamma = 0.99
ema = 1e-3
optimizer_config = {
'lr':1e-4,
'weight_decay': 1e-3
}
velocity = TDFlowUnet()
fm = ConditionalFlowMatching(velocity, obs_dim=(4, ))
fm_target = deepcopy(fm)
trainer = TDFlowTrainer(
fm=fm,
fm_target=fm_target,
train_loader=train_loader,
optimizer_config=optimizer_config,
gamma=gamma,
ema=ema,
device='auto',
task=task,
loss_type=loss_type
)
try:
trainer.fit(num_epochs)
finally:
wandb.finish()
torch.save(trainer.fm.model.state_dict(), f'checkpoints/{loss_type}_model_{task}.pth')
torch.save(trainer.fm_target.model.state_dict(), f'checkpoints/{loss_type}_target_model_{task}.pth')
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument('--task', type=str, default='reach_top_left')
parser.add_argument('--num_epochs', type=int, default=100)
parser.add_argument('--loss_type', type=str, default='td2_cfm')
args = parser.parse_args()
print(f"Starting training with args: {args}")
main(
task=args.task,
num_epochs=args.num_epochs,
loss_type=args.loss_type
)