This commit is contained in:
@@ -0,0 +1,40 @@
|
||||
#!/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()
|
||||
Reference in New Issue
Block a user