#!/usr/bin/env python3 """Add the original-to-compact token lookup to an already exported student ONNX.""" from __future__ import annotations import argparse import sys from pathlib import Path import numpy as np import onnx from onnx import helper, numpy_helper def patch(args: argparse.Namespace) -> None: if args.output.exists(): raise FileExistsError(f"Refusing to overwrite {args.output}") model = onnx.load(str(args.model)) if not any(value.name == "input_ids" for value in model.graph.input): raise ValueError("Model has no input_ids graph input") if any(item.name == "original_to_compact_token_id" for item in model.graph.initializer): raise ValueError("Model already contains original_to_compact_token_id") consumers = 0 for node in model.graph.node: for index, name in enumerate(node.input): if name == "input_ids": node.input[index] = "mapped_input_ids" consumers += 1 if consumers == 0: raise ValueError("No input_ids consumers found") token_map = np.load(args.token_map).astype(np.int64, copy=False) model.graph.initializer.append(numpy_helper.from_array(token_map, "original_to_compact_token_id")) model.graph.node.insert( 0, helper.make_node( "Gather", ["original_to_compact_token_id", "input_ids"], ["mapped_input_ids"], axis=0, name="MapOriginalTokenIds", ), ) onnx.checker.check_model(model) onnx.save(model, str(args.output)) print(f"token_map_entries={token_map.size} consumers={consumers} output_bytes={args.output.stat().st_size}") def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--model", type=Path, required=True) parser.add_argument("--token-map", type=Path, required=True) parser.add_argument("--output", type=Path, required=True) return parser.parse_args() if __name__ == "__main__": sys.exit(patch(parse_args()) or 0)