2022-12-21 17:28:19 +01:00
|
|
|
'''
|
|
|
|
Converts a transformers model to .pt, which is faster to load.
|
|
|
|
|
2023-01-07 04:04:52 +01:00
|
|
|
Example:
|
2023-01-07 20:54:49 +01:00
|
|
|
python convert-to-torch.py models/opt-1.3b
|
2022-12-21 17:28:19 +01:00
|
|
|
|
2023-01-07 20:54:49 +01:00
|
|
|
The output will be written to torch-dumps/name-of-the-model.pt
|
2022-12-21 17:28:19 +01:00
|
|
|
'''
|
2023-02-10 19:57:55 +01:00
|
|
|
|
2023-01-07 20:33:43 +01:00
|
|
|
from pathlib import Path
|
2023-02-10 19:57:55 +01:00
|
|
|
from sys import argv
|
|
|
|
|
|
|
|
import torch
|
|
|
|
from transformers import AutoModelForCausalLM
|
2022-12-21 17:28:19 +01:00
|
|
|
|
2023-01-07 20:33:43 +01:00
|
|
|
path = Path(argv[1])
|
|
|
|
model_name = path.name
|
2023-01-07 04:04:52 +01:00
|
|
|
|
|
|
|
print(f"Loading {model_name}...")
|
2023-01-11 03:41:35 +01:00
|
|
|
model = AutoModelForCausalLM.from_pretrained(path, low_cpu_mem_usage=True, torch_dtype=torch.float16).cuda()
|
2023-01-16 20:35:45 +01:00
|
|
|
print(f"Model loaded.\nSaving to torch-dumps/{model_name}.pt")
|
2023-01-07 20:33:43 +01:00
|
|
|
torch.save(model, Path(f"torch-dumps/{model_name}.pt"))
|