MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / main

Function main

codegeex/megatron/convert_ckpt_parallel.py:48–120  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

46
47
48def main():
49 parser = argparse.ArgumentParser()
50 parser = get_change_ckpt_args(parser)
51 args, _ = parser.parse_known_args()
52
53 print(f"Load ckpt from {args.load_ckpt_path}...")
54 state_dict = torch.load(args.load_ckpt_path, map_location="cpu")
55
56 print(f"Spliting ckpt into {args.target_tensor_model_parallel_size} parts...")
57 output_state_dict = []
58 for i in range(args.target_tensor_model_parallel_size):
59 output_state_dict.append({})
60
61 print("Converting Embedding layers...")
62 word_embeddings = state_dict['module']['language_model']['embedding']['word_embeddings']['weight']
63 position_embeddings = state_dict['module']['language_model']['embedding']['position_embeddings']['weight']
64 out_word_embeddings = torch.chunk(word_embeddings, args.target_tensor_model_parallel_size, dim=0)
65
66 for i in range(args.target_tensor_model_parallel_size):
67 pos_emb_dict = get_element_from_dict_by_path(
68 output_state_dict[i], "module.language_model.embedding.position_embeddings"
69 )
70 pos_emb_dict["weight"] = position_embeddings
71
72 word_emb_dict = get_element_from_dict_by_path(
73 output_state_dict[i], "module.language_model.embedding.word_embeddings"
74 )
75 word_emb_dict["weight"] = out_word_embeddings[i].clone()
76
77 print("Converting QueryEmbedding layers...")
78 query_embeddings = state_dict['module']['language_model']['topQueryEmbedding']['top_query_embeddings']['weight']
79 out_query_embeddings = torch.chunk(query_embeddings, args.target_tensor_model_parallel_size, dim=0)
80
81 for i in range(args.target_tensor_model_parallel_size):
82 query_emb_dict = get_element_from_dict_by_path(
83 output_state_dict[i], "module.language_model.topQueryEmbedding.top_query_embeddings"
84 )
85 query_emb_dict["weight"] = out_query_embeddings[i].clone()
86
87 print("Converting Transformer layers...")
88 for layer_name in state_dict['module']['language_model']['transformer'].keys():
89 params = state_dict['module']['language_model']['transformer'][layer_name]
90 if "layernorm" in layer_name:
91 pass
92 elif "attention" in layer_name and "weight" in layer_name:
93 if "dense" in layer_name:
94 params = torch.chunk(params, args.target_tensor_model_parallel_size, dim=1)
95 else:
96 params = torch.chunk(params, args.target_tensor_model_parallel_size, dim=0)
97 elif "weight" in layer_name and "dense" in layer_name:
98 if "h_to_4h" in layer_name:
99 params = torch.chunk(params, args.target_tensor_model_parallel_size, dim=0)
100 else:
101 params = torch.chunk(params, args.target_tensor_model_parallel_size, dim=1)
102 elif "bias" in layer_name:
103 if "dense" not in layer_name or "mlp" in layer_name:
104 if "4h_to_h" in layer_name:
105 pass

Callers 1

Calls 2

get_change_ckpt_argsFunction · 0.70

Tested by

no test coverage detected