CoolFace
Apppublic

NeuralInternet/Text-Generation_Playground

sourceHugging Facemitupdated 4y agoView on Hugging Face
2likes
convert-to-flexgen.py61 linesDownload Raw Back to root
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