(self)
| 142 | self.opt.load_state_dict(state_dict) |
| 143 | |
| 144 | def run_loop(self): |
| 145 | while ( |
| 146 | not self.lr_anneal_steps |
| 147 | or self.step + self.resume_step < self.lr_anneal_steps |
| 148 | ): |
| 149 | batch, cond = next(self.data) |
| 150 | self.run_step(batch, cond) |
| 151 | if self.step % self.log_interval == 0: |
| 152 | logger.dumpkvs() |
| 153 | if self.step % self.save_interval == 0 and self.step > 0: |
| 154 | self.save() |
| 155 | # Run for a finite amount of time in integration tests. |
| 156 | if os.environ.get("DIFFUSION_TRAINING_TEST", "") and self.step > 0: |
| 157 | return |
| 158 | self.step += 1 |
| 159 | # Save the last checkpoint if it wasn't already saved. |
| 160 | if (self.step - 1) % self.save_interval != 0: |
| 161 | self.save() |
| 162 | |
| 163 | def run_step(self, batch, cond): |
| 164 | self.forward_backward(batch, cond) |
no test coverage detected