MCPcopy Create free account
hub / github.com/UX-Decoder/Semantic-SAM / align_and_update_state_dicts

Function align_and_update_state_dicts

utils/model.py:29–58  ·  view source on GitHub ↗
(model_state_dict, ckpt_state_dict)

Source from the content-addressed store, hash-verified

27 return cls
28
29def align_and_update_state_dicts(model_state_dict, ckpt_state_dict):
30 model_keys = sorted(model_state_dict.keys())
31 ckpt_keys = sorted(ckpt_state_dict.keys())
32 result_dicts = {}
33 matched_log = []
34 unmatched_log = []
35 unloaded_log = []
36 for model_key in model_keys:
37 model_weight = model_state_dict[model_key]
38 if model_key in ckpt_keys:
39 ckpt_weight = ckpt_state_dict[model_key]
40 if model_weight.shape == ckpt_weight.shape:
41 result_dicts[model_key] = ckpt_weight
42 ckpt_keys.pop(ckpt_keys.index(model_key))
43 matched_log.append("Loaded {}, Model Shape: {} <-> Ckpt Shape: {}".format(model_key, model_weight.shape, ckpt_weight.shape))
44 else:
45 unmatched_log.append("*UNMATCHED* {}, Model Shape: {} <-> Ckpt Shape: {}".format(model_key, model_weight.shape, ckpt_weight.shape))
46 else:
47 unloaded_log.append("*UNLOADED* {}, Model Shape: {}".format(model_key, model_weight.shape))
48
49 if is_main_process():
50 for info in matched_log:
51 logger.info(info)
52 for info in unloaded_log:
53 logger.warning(info)
54 for key in ckpt_keys:
55 logger.warning("$UNUSED$ {}, Ckpt Shape: {}".format(key, ckpt_state_dict[key].shape))
56 for info in unmatched_log:
57 logger.warning(info)
58 return result_dicts

Callers 1

from_pretrainedMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected