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