MCPcopy Create free account
hub / github.com/adobe-research/custom-diffusion / configure_optimizers

Method configure_optimizers

src/model.py:178–219  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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))

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected