forked from Beomi/KoAlpaca
-
Notifications
You must be signed in to change notification settings - Fork 3
Expand file tree
/
Copy pathconvert_to_onnx.py
More file actions
33 lines (26 loc) · 785 Bytes
/
Copy pathconvert_to_onnx.py
File metadata and controls
33 lines (26 loc) · 785 Bytes
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
import torch
from transformers import GPT2LMHeadModel, GPT2Tokenizer
from pathlib import Path
model_name = "KoAlpaca.cpp"
output_path = "output_model.onnx"
model = GPT2LMHeadModel.from_pretrained(model_name)
tokenizer = GPT2Tokenizer.from_pretrained(model_name)
model.eval()
# Get the input text dynamically
input_text = input("Enter input text: ")
input_ids = tokenizer.encode(input_text, return_tensors="pt")
input_names = ["input_ids"]
output_names = ["output_0"]
dynamic_axes = {
"input_ids": {0: "batch_size", 1: "sequence_length"},
"output_0": {0: "batch_size", 1: "sequence_length"},
}
torch.onnx.export(
model,
input_ids,
output_path,
input_names=input_names,
output_names=output_names,
dynamic_axes=dynamic_axes,
opset_version=12,
)