More frequent EMA update (#20)
* move ema update within pretraining epoch * update pretraining ema configs * add lift and can epoch 8000 checkpoint url * add note about EMA issue in pretraining instruction
This commit is contained in:
@@ -36,6 +36,10 @@ class TrainDiffusionAgent(PreTrainAgent):
|
||||
|
||||
self.optimizer.step()
|
||||
self.optimizer.zero_grad()
|
||||
|
||||
# update ema
|
||||
if self.epoch % self.update_ema_freq == 0:
|
||||
self.step_ema()
|
||||
loss_train = np.mean(loss_train_epoch)
|
||||
|
||||
# validate
|
||||
@@ -53,10 +57,6 @@ class TrainDiffusionAgent(PreTrainAgent):
|
||||
# update lr
|
||||
self.lr_scheduler.step()
|
||||
|
||||
# update ema
|
||||
if self.epoch % self.update_ema_freq == 0:
|
||||
self.step_ema()
|
||||
|
||||
# save model
|
||||
if self.epoch % self.save_model_freq == 0 or self.epoch == self.n_epochs:
|
||||
self.save_model()
|
||||
|
||||
@@ -44,6 +44,10 @@ class TrainGaussianAgent(PreTrainAgent):
|
||||
|
||||
self.optimizer.step()
|
||||
self.optimizer.zero_grad()
|
||||
|
||||
# update ema
|
||||
if self.epoch % self.update_ema_freq == 0:
|
||||
self.step_ema()
|
||||
loss_train = np.mean(loss_train_epoch)
|
||||
ent_train = np.mean(ent_train_epoch)
|
||||
|
||||
@@ -65,10 +69,6 @@ class TrainGaussianAgent(PreTrainAgent):
|
||||
# update lr
|
||||
self.lr_scheduler.step()
|
||||
|
||||
# update ema
|
||||
if self.epoch % self.update_ema_freq == 0:
|
||||
self.step_ema()
|
||||
|
||||
# save model
|
||||
if self.epoch % self.save_model_freq == 0 or self.epoch == self.n_epochs:
|
||||
self.save_model()
|
||||
|
||||
Reference in New Issue
Block a user