MCPcopy Create free account
hub / github.com/Project-MONAI/MONAI / load_old_state_dict

Method load_old_state_dict

monai/networks/nets/controlnet.py:425–467  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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)

Calls 3

popMethod · 0.80
state_dictMethod · 0.45
load_state_dictMethod · 0.45