| import onnx | |
| import argparse | |
| def print_node_types_sorted(model_path): | |
| # 加载ONNX模型 | |
| model = onnx.load(model_path) | |
| # 获取模型中的所有节点 | |
| nodes = model.graph.node | |
| # 创建一个字典来存储每个类型的计数 | |
| type_counts = {} | |
| # 遍历所有节点并统计类型 | |
| for node in nodes: | |
| node_type = node.op_type | |
| if node_type in type_counts: | |
| type_counts[node_type] += 1 | |
| else: | |
| type_counts[node_type] = 1 | |
| # 将类型名称按照字母顺序排序 | |
| sorted_types = sorted(type_counts.items()) | |
| # 打印排序后的类型名称和它们的数量 | |
| for type_name, count in sorted_types: | |
| print(f"{type_name}: {count}") | |
| if __name__ == "__main__": | |
| # 设置命令行参数解析 | |
| parser = argparse.ArgumentParser(description='Process an ONNX model.') | |
| parser.add_argument('model_path', type=str, help='Path to the ONNX model file') | |
| # 解析命令行参数 | |
| args = parser.parse_args() | |
| # 使用从命令行获取的模型路径调用函数 | |
| print_node_types_sorted(args.model_path) | |