| 27 | return cls |
| 28 | |
| 29 | def 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 |