MCPcopy Create free account
hub / github.com/GuyTevet/motion-diffusion-model / __init__

Method __init__

train/training_loop.py:38–128  ·  view source on GitHub ↗
(self, args, train_platform, model, diffusion, data)

Source from the content-addressed store, hash-verified

36
37class TrainLoop:
38 def __init__(self, args, train_platform, model, diffusion, data):
39 self.args = args
40 self.dataset = args.dataset
41 self.train_platform = train_platform
42 self.model = model
43 self.model_avg = None
44 if self.args.use_ema:
45 self.model_avg = copy.deepcopy(self.model)
46 self.model_for_eval = self.model_avg if self.args.use_ema else self.model
47 if args.gen_guidance_param != 1:
48 self.model_for_eval = ClassifierFreeSampleModel(self.model_for_eval) # wrapping model with the classifier-free sampler
49 self.diffusion = diffusion
50 self.cond_mode = model.cond_mode
51 self.data = data
52 self.batch_size = args.batch_size
53 self.microbatch = args.batch_size # deprecating this option
54 self.lr = args.lr
55 self.log_interval = args.log_interval
56 self.save_interval = args.save_interval
57 self.resume_checkpoint = args.resume_checkpoint
58 self.use_fp16 = False # deprecating this option
59 self.fp16_scale_growth = 1e-3 # deprecating this option
60 self.weight_decay = args.weight_decay
61 self.lr_anneal_steps = args.lr_anneal_steps
62
63 self.step = 0
64 self.resume_step = 0
65 self.global_batch = self.batch_size # * dist.get_world_size()
66 self.num_steps = args.num_steps
67 self.num_epochs = self.num_steps // len(self.data) + 1
68
69 self.sync_cuda = torch.cuda.is_available()
70
71 self._load_and_sync_parameters()
72 self.mp_trainer = MixedPrecisionTrainer(
73 model=self.model,
74 use_fp16=self.use_fp16,
75 fp16_scale_growth=self.fp16_scale_growth,
76 )
77
78 self.save_dir = args.save_dir
79 self.overwrite = args.overwrite
80
81 if self.args.use_ema:
82 self.opt = AdamW(
83 # with amp, we don't need to use the mp_trainer's master_params
84 (self.model.parameters()
85 if self.use_fp16 else self.mp_trainer.master_params),
86 lr=self.lr,
87 weight_decay=self.weight_decay,
88 betas=(0.9, self.args.adam_beta2),
89 )
90 else:
91 self.opt = AdamW(
92 self.mp_trainer.master_params, lr=self.lr, weight_decay=self.weight_decay
93 )
94
95 if self.resume_step:

Callers

nothing calls this directly

Calls 7

_load_optimizer_stateMethod · 0.95
get_dataset_loaderFunction · 0.90
EvaluatorMDMWrapperClass · 0.90

Tested by

no test coverage detected