Load a state dict from a ControlNet trained with [MONAI Generative](https://github.com/Project-MONAI/GenerativeModels). Args: old_state_dict: state dict from the old ControlNet model.
(self, old_state_dict: dict, verbose=False)
| 423 | return down_block_res_samples, mid_block_res_sample |
| 424 | |
| 425 | def load_old_state_dict(self, old_state_dict: dict, verbose=False) -> None: |
| 426 | """ |
| 427 | Load a state dict from a ControlNet trained with |
| 428 | [MONAI Generative](https://github.com/Project-MONAI/GenerativeModels). |
| 429 | |
| 430 | Args: |
| 431 | old_state_dict: state dict from the old ControlNet model. |
| 432 | """ |
| 433 | |
| 434 | new_state_dict = self.state_dict() |
| 435 | # if all keys match, just load the state dict |
| 436 | if all(k in new_state_dict for k in old_state_dict): |
| 437 | print("All keys match, loading state dict.") |
| 438 | self.load_state_dict(old_state_dict) |
| 439 | return |
| 440 | |
| 441 | if verbose: |
| 442 | # print all new_state_dict keys that are not in old_state_dict |
| 443 | for k in new_state_dict: |
| 444 | if k not in old_state_dict: |
| 445 | print(f"key {k} not found in old state dict") |
| 446 | # and vice versa |
| 447 | print("----------------------------------------------") |
| 448 | for k in old_state_dict: |
| 449 | if k not in new_state_dict: |
| 450 | print(f"key {k} not found in new state dict") |
| 451 | |
| 452 | # copy over all matching keys |
| 453 | for k in new_state_dict: |
| 454 | if k in old_state_dict: |
| 455 | new_state_dict[k] = old_state_dict.pop(k) |
| 456 | |
| 457 | # fix the attention blocks |
| 458 | attention_blocks = [k.replace(".out_proj.weight", "") for k in new_state_dict if "out_proj.weight" in k] |
| 459 | for block in attention_blocks: |
| 460 | # projection |
| 461 | new_state_dict[f"{block}.out_proj.weight"] = old_state_dict.pop(f"{block}.to_out.0.weight") |
| 462 | new_state_dict[f"{block}.out_proj.bias"] = old_state_dict.pop(f"{block}.to_out.0.bias") |
| 463 | |
| 464 | if verbose: |
| 465 | # print all remaining keys in old_state_dict |
| 466 | print("remaining keys in old_state_dict:", old_state_dict.keys()) |
| 467 | self.load_state_dict(new_state_dict) |