()
| 46 | |
| 47 | |
| 48 | def 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 |
no test coverage detected