Files
Raptor/validation/tools/split_onnx_prefixes.py
ilgeco f3a4e19f7c
Validate Operations / validate-operations (push) Has been cancelled
Python script for compare
2026-07-24 12:47:37 +02:00

41 lines
1.5 KiB
Python

#!/usr/bin/env python3
import argparse
from pathlib import Path
import onnx
def split_prefixes(model_path: Path, output_dir: Path, name: str) -> None:
model = onnx.shape_inference.infer_shapes(onnx.load(model_path))
initializer_names = {initializer.name for initializer in model.graph.initializer}
input_names = [value.name for value in model.graph.input if value.name not in initializer_names]
extractor = onnx.utils.Extractor(model)
output_dir.mkdir(parents=True, exist_ok=True)
for depth, node in enumerate(model.graph.node):
output_name = next(output for output in node.output if output)
prefix = extractor.extract_model(input_names, [output_name])
prefix.ir_version = max(prefix.ir_version, 4)
onnx.checker.check_model(prefix)
depth_name = f"depth_{depth:02d}"
depth_dir = output_dir / depth_name
depth_dir.mkdir(parents=True, exist_ok=True)
output_path = depth_dir / f"{name}_{depth_name}.onnx"
onnx.save(prefix, output_path)
print(f"{depth_name}: {node.op_type} -> {output_name} ({len(prefix.graph.node)} nodes)")
def main() -> None:
parser = argparse.ArgumentParser(description="Split an ONNX graph into one ancestor prefix per node.")
parser.add_argument("model", type=Path)
parser.add_argument("output_dir", type=Path)
parser.add_argument("--name", required=True)
args = parser.parse_args()
split_prefixes(args.model, args.output_dir, args.name)
if __name__ == "__main__":
main()