(self, parameters)
| 358 | # Get param groups that will be passed into the optimizer. Models can override this, e.g. SDXL |
| 359 | # supports separate learning rates for unet and text encoders. |
| 360 | def get_param_groups(self, parameters): |
| 361 | return [{'params': parameters}] |
| 362 | |
| 363 | # Default loss_fn. MSE between output and target, with mask support. |
| 364 | def get_loss_fn(self): |