Files
Aletheia/tools/tts/patch_booktts_student_token_map.py
T

55 lines
2.0 KiB
Python

#!/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)