41 lines
1.5 KiB
Python
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()
|