summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-07 11:13:47 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-07 11:13:47 +0200
commit96b4edc5e291e253747956083933339040ad07c3 (patch)
treeae6fdc64fda6b87b092b810d632715d13b663468
parent58ef69d71be7cdaa005a4ee21ec5ba43511e7196 (diff)
Add model conversion script
-rw-r--r--gpt2/scripts/base.py11
-rw-r--r--gpt2/scripts/conv.py105
2 files changed, 103 insertions, 13 deletions
diff --git a/gpt2/scripts/base.py b/gpt2/scripts/base.py
new file mode 100644
index 0000000..0553ffd
--- /dev/null
+++ b/gpt2/scripts/base.py
@@ -0,0 +1,11 @@
+#!/usr/bin/env python3
+
+from transformers import AutoTokenizer, AutoModelForCausalLM
+
+model_id = "gpt2"
+
+tokenizer = AutoTokenizer.from_pretrained(model_id)
+tokenizer.save_pretrained("../models/base")
+
+model = AutoModelForCausalLM.from_pretrained(model_id)
+model.save_pretrained("../models/base")
diff --git a/gpt2/scripts/conv.py b/gpt2/scripts/conv.py
index e3773e2..e820ec2 100644
--- a/gpt2/scripts/conv.py
+++ b/gpt2/scripts/conv.py
@@ -1,15 +1,94 @@
-# /// script
-# dependencies = [
-# "torch",
-# "transformers",
-# "optimum",
-# "optimum[onnxruntime]",
-# ]
-# ///
-
-from transformers import AutoTokenizer, AutoModelForCausalLM
+#!/usr/bin/env python3
+
+import argparse
+import os
+import shutil
+import tempfile
+
+import numpy as np
+import onnx
+from onnx import helper, TensorProto, numpy_helper
from optimum.onnxruntime import ORTModelForCausalLM
-model_id = "gpt2"
-model = ORTModelForCausalLM.from_pretrained(model_id, export=True, use_cache=True)
-model.save_pretrained("../models/base") \ No newline at end of file
+
+def parse_args():
+ parser = argparse.ArgumentParser()
+
+ parser.add_argument('--model', required=True, help='directory containing model.safetensors')
+
+ return parser.parse_args()
+
+
+def add_log_probs(model_path: str, output_path: str):
+ m = onnx.load(model_path)
+
+ g = m.graph
+
+ INT64_MAX = 9223372036854775807
+
+ def init(name, values, dtype=np.int64):
+ g.initializer.append(numpy_helper.from_array(np.array(values, dtype=dtype), name=name))
+
+ # 1D scalar-like constants (shape [1]) for use in Slice starts/ends/axes
+ init("_lp_zero", [0])
+ init("_lp_one", [1])
+ init("_lp_int_max", [INT64_MAX])
+ init("_lp_axis1", [1])
+ init("_lp_axis2", [2])
+
+ def node(op, inputs, outputs, **attrs):
+ g.node.append(helper.make_node(op, inputs=inputs, outputs=outputs, **attrs))
+
+ # log_softmax(logits) over vocab dim → [1, seq, vocab]
+ node("LogSoftmax", ["logits"], ["_lp_lsm"], axis=2)
+
+ # seq_len as shape-[1] tensor
+ node("Shape", ["logits"], ["_lp_shape"])
+ node("Gather", ["_lp_shape", "_lp_one"], ["_lp_seq_len"]) # shape [1]
+
+ # seq_len - 1 → shape [1]
+ node("Sub", ["_lp_seq_len", "_lp_one"], ["_lp_seq_len_m1"])
+
+ # log_softmax[:, 0:seq_len-1, :] → [1, seq-1, vocab]
+ node("Slice", ["_lp_lsm", "_lp_zero", "_lp_seq_len_m1", "_lp_axis1"], ["_lp_lsm_s"])
+
+ # input_ids[:, 1:] → [1, seq-1]
+ node("Slice", ["input_ids", "_lp_one", "_lp_int_max", "_lp_axis1"], ["_lp_ids_s"])
+
+ # unsqueeze → [1, seq-1, 1] (axes as input for opset 13+)
+ node("Unsqueeze", ["_lp_ids_s", "_lp_axis2"], ["_lp_ids_3d"])
+
+ # gather the log prob for each target token → [1, seq-1, 1]
+ node("GatherElements", ["_lp_lsm_s", "_lp_ids_3d"], ["_lp_gathered"], axis=2)
+
+ # squeeze trailing dim → [1, seq-1]
+ node("Squeeze", ["_lp_gathered", "_lp_axis2"], ["token_logprobs"])
+
+ g.output.append(
+ helper.make_tensor_value_info("token_logprobs", TensorProto.FLOAT, [None, None])
+ )
+
+ onnx.save(m, output_path)
+
+
+def main():
+ args = parse_args()
+
+ model_dir = os.path.abspath(args.model)
+
+ with tempfile.TemporaryDirectory() as tmp:
+ no_cache_dir = os.path.join(tmp, "base")
+
+ ORTModelForCausalLM.from_pretrained(model_dir, export=True, use_cache=False).save_pretrained(no_cache_dir)
+ shutil.copy2(os.path.join(no_cache_dir, "model.onnx"), os.path.join(model_dir, "model.onnx"))
+
+ cache_dir = os.path.join(tmp, "cache")
+
+ ORTModelForCausalLM.from_pretrained(model_dir, export=True, use_cache=True).save_pretrained(cache_dir)
+ shutil.copy2(os.path.join(cache_dir, "model.onnx"), os.path.join(model_dir, "model_cache.onnx"))
+
+ add_log_probs(os.path.join(model_dir, "model.onnx"), os.path.join(model_dir, "model_eval.onnx"))
+
+
+if __name__ == '__main__':
+ main()