diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-07 11:13:47 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-07 11:13:47 +0200 |
| commit | 96b4edc5e291e253747956083933339040ad07c3 (patch) | |
| tree | ae6fdc64fda6b87b092b810d632715d13b663468 | |
| parent | 58ef69d71be7cdaa005a4ee21ec5ba43511e7196 (diff) | |
Add model conversion script
| -rw-r--r-- | gpt2/scripts/base.py | 11 | ||||
| -rw-r--r-- | gpt2/scripts/conv.py | 105 |
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() |
