MCPcopy Create free account
hub / github.com/FreedomGu/Diffportrait360 / load_state_dict

Function load_state_dict

diffportrait360_release/code/train.py:32–37  ·  view source on GitHub ↗
(model, ckpt_path, strict=True, map_location="cpu")

Source from the content-addressed store, hash-verified

30print(f"TORCH_VERSION={TORCH_VERSION} FP16_DTYPE={FP16_DTYPE}")
31torch.set_default_tensor_type(torch.FloatTensor)
32def load_state_dict(model, ckpt_path, strict=True, map_location="cpu"):
33 print(f"Loading model state dict from {ckpt_path} ...")
34 state_dict = load_from_pretrain(ckpt_path, map_location=map_location)
35 state_dict = state_dict.get('state_dict', state_dict)
36 model.load_state_dict(state_dict, strict=strict)
37 del state_dict
38
39def get_cond_control(args, batch_data, control_type, device, model=None,batch_size=None, train=True, seg_model=None):
40 # Single-control

Callers 1

mainFunction · 0.70

Calls 1

load_from_pretrainFunction · 0.90

Tested by

no test coverage detected