(self)
| 176 | change_forward(self.model.diffusion_model) |
| 177 | |
| 178 | def configure_optimizers(self): |
| 179 | lr = self.learning_rate |
| 180 | params = [] |
| 181 | if self.freeze_model == 'crossattn-kv': |
| 182 | for x in self.model.diffusion_model.named_parameters(): |
| 183 | if 'transformer_blocks' in x[0]: |
| 184 | if 'attn2.to_k' in x[0] or 'attn2.to_v' in x[0]: |
| 185 | params += [x[1]] |
| 186 | print(x[0]) |
| 187 | elif self.freeze_model == 'crossattn': |
| 188 | for x in self.model.diffusion_model.named_parameters(): |
| 189 | if 'transformer_blocks' in x[0]: |
| 190 | if 'attn2' in x[0]: |
| 191 | params += [x[1]] |
| 192 | print(x[0]) |
| 193 | else: |
| 194 | params = list(self.model.parameters()) |
| 195 | |
| 196 | if self.cond_stage_trainable: |
| 197 | print(f"{self.__class__.__name__}: Also optimizing conditioner params!") |
| 198 | if self.add_token: |
| 199 | params = params + list(self.cond_stage_model.transformer.text_model.embeddings.token_embedding.parameters()) |
| 200 | else: |
| 201 | params = params + list(self.cond_stage_model.parameters()) |
| 202 | |
| 203 | if self.learn_logvar: |
| 204 | print('Diffusion model optimizing logvar') |
| 205 | params.append(self.logvar) |
| 206 | opt = torch.optim.AdamW(params, lr=lr) |
| 207 | if self.use_scheduler: |
| 208 | assert 'target' in self.scheduler_config |
| 209 | scheduler = instantiate_from_config(self.scheduler_config) |
| 210 | |
| 211 | print("Setting up LambdaLR scheduler...") |
| 212 | scheduler = [ |
| 213 | { |
| 214 | 'scheduler': LambdaLR(opt, lr_lambda=scheduler.schedule), |
| 215 | 'interval': 'step', |
| 216 | 'frequency': 1 |
| 217 | }] |
| 218 | return [opt], scheduler |
| 219 | return opt |
| 220 | |
| 221 | def p_losses(self, x_start, cond, t, mask=None, noise=None): |
| 222 | noise = default(noise, lambda: torch.randn_like(x_start)) |
nothing calls this directly
no outgoing calls
no test coverage detected