-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
29 lines (20 loc) · 1.19 KB
/
Copy pathtrain.py
File metadata and controls
29 lines (20 loc) · 1.19 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
from argparse import ArgumentParser
from torch.utils.data import DataLoader
from dataset import VideoDataset, DatasetRepeater
from augmentation import get_transform
from train_class import FirstOrderMotionModel
if __name__ == "__main__":
parser = ArgumentParser()
parser.add_argument("--data_path", default="./moving-gif/train", help="path to training data")
parser.add_argument("--config_path", default="./configs/mgif.yaml", help="path to config file")
parser.add_argument("--log_path", default='logs', help="path to log")
parser.add_argument("--checkpoint_path", default=None, help="path to save the checkpoint")
args = parser.parse_args()
dataset = VideoDataset(data_path=args.data_path, id_sampling=False, transform=get_transform("mgif"))
dataset = DatasetRepeater(dataset, num_repeats=25)
dataloader = DataLoader(dataset, batch_size=16, shuffle=True, num_workers=2, pin_memory=True)
model = FirstOrderMotionModel(config_path=args.config_path,
log_path=args.log_path,
checkpoint_path=args.checkpoint_path)
print("Training...")
model.train(dataloader)