diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-05-04 15:21:55 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-05-04 15:28:10 +0200 |
| commit | 3f4748105edabbf5cefb761d1d61993d0b4d9028 (patch) | |
| tree | 3e7b782233444d805d03d3c319d37f17e0173c96 | |
| parent | 70602ac424113eb3f3d08ecdf1e018871b7c3997 (diff) | |
Handle external data
| -rw-r--r-- | gpt2/scripts/conv.py | 12 |
1 files changed, 8 insertions, 4 deletions
diff --git a/gpt2/scripts/conv.py b/gpt2/scripts/conv.py index e820ec2..ec72d75 100644 --- a/gpt2/scripts/conv.py +++ b/gpt2/scripts/conv.py @@ -8,19 +8,20 @@ import tempfile import numpy as np import onnx from onnx import helper, TensorProto, numpy_helper -from optimum.onnxruntime import ORTModelForCausalLM +from optimum.exporters.onnx import main_export def parse_args(): parser = argparse.ArgumentParser() parser.add_argument('--model', required=True, help='directory containing model.safetensors') + parser.add_argument('--no-post-process', action='store_true', help='skip ONNX post-processing') return parser.parse_args() def add_log_probs(model_path: str, output_path: str): - m = onnx.load(model_path) + m = onnx.load(model_path, load_external_data=False) g = m.graph @@ -79,12 +80,15 @@ def main(): 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) + main_export(model_dir, output=no_cache_dir, task="text-generation", do_validation=False, no_post_process=args.no_post_process) shutil.copy2(os.path.join(no_cache_dir, "model.onnx"), os.path.join(model_dir, "model.onnx")) + if os.path.exists(os.path.join(no_cache_dir, "model.onnx_data")): + shutil.copy2(os.path.join(no_cache_dir, "model.onnx_data"), os.path.join(model_dir, "model.onnx_data")) + cache_dir = os.path.join(tmp, "cache") - ORTModelForCausalLM.from_pretrained(model_dir, export=True, use_cache=True).save_pretrained(cache_dir) + main_export(model_dir, output=cache_dir, task="text-generation-with-past", do_validation=False, no_post_process=args.no_post_process) 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")) |
