(self, text_encoder_lr, unet_lr, default_lr)
| 297 | return info |
| 298 | |
| 299 | def prepare_optimizer_params(self, text_encoder_lr, unet_lr, default_lr): |
| 300 | self.requires_grad_(True) |
| 301 | all_params = [] |
| 302 | |
| 303 | def enumerate_params(loras): |
| 304 | params = [] |
| 305 | for lora in loras: |
| 306 | params.extend(lora.parameters()) |
| 307 | return params |
| 308 | |
| 309 | if self.text_encoder_loras: |
| 310 | param_data = {"params": enumerate_params(self.text_encoder_loras)} |
| 311 | if text_encoder_lr is not None: |
| 312 | param_data["lr"] = text_encoder_lr |
| 313 | all_params.append(param_data) |
| 314 | |
| 315 | if self.unet_loras: |
| 316 | param_data = {"params": enumerate_params(self.unet_loras)} |
| 317 | if unet_lr is not None: |
| 318 | param_data["lr"] = unet_lr |
| 319 | all_params.append(param_data) |
| 320 | |
| 321 | return all_params |
| 322 | |
| 323 | def enable_gradient_checkpointing(self): |
| 324 | pass |
nothing calls this directly
no outgoing calls
no test coverage detected