NeuralInternet/Text-Generation_Playground
2
1'''2 3Converts a transformers model to a format compatible with flexgen.4 5'''6 7import argparse8import os9from pathlib import Path10 11import numpy as np12import torch13from tqdm import tqdm14from transformers import AutoModelForCausalLM, AutoTokenizer15 16parser = argparse.ArgumentParser(formatter_class=lambda prog: argparse.HelpFormatter(prog,max_help_position=54))17parser.add_argument('MODEL', type=str, default=None, nargs='?', help="Path to the input model.")18args = parser.parse_args()19 20def disable_torch_init():21 """22 Disable the redundant torch default initialization to accelerate model creation.23 """24 import torch25 global torch_linear_init_backup26 global torch_layer_norm_init_backup27 28 torch_linear_init_backup = torch.nn.Linear.reset_parameters29 setattr(torch.nn.Linear, "reset_parameters", lambda self: None)30 31 torch_layer_norm_init_backup = torch.nn.LayerNorm.reset_parameters32 setattr(torch.nn.LayerNorm, "reset_parameters", lambda self: None)33 34def restore_torch_init():35 """Rollback the change made by disable_torch_init."""36 import torch37 setattr(torch.nn.Linear, "reset_parameters", torch_linear_init_backup)38 setattr(torch.nn.LayerNorm, "reset_parameters", torch_layer_norm_init_backup)39 40if __name__ == '__main__':41 path = Path(args.MODEL)42 model_name = path.name43 44 print(f"Loading {model_name}...")45 #disable_torch_init()46 model = AutoModelForCausalLM.from_pretrained(path, torch_dtype=torch.float16, low_cpu_mem_usage=True)47 #restore_torch_init()48 49 tokenizer = AutoTokenizer.from_pretrained(path)50 51 out_folder = Path(f"models/{model_name}-np")52 if not Path(out_folder).exists():53 os.mkdir(out_folder)54 55 print(f"Saving the converted model to {out_folder}...")56 for name, param in tqdm(list(model.model.named_parameters())):57 name = name.replace("decoder.final_layer_norm", "decoder.layer_norm")58 param_path = os.path.join(out_folder, name)59 with open(param_path, "wb") as f:60 np.save(f, param.cpu().detach().numpy())61 