如何将TensorBoard中的神经网络计算图导出为可读文件(如JSON)
当然可以,下面是几种实用的方法,帮你把ResNet这类神经网络的计算图(以操作算子为节点)提取为JSON格式的可读文件:
方法1:利用PyTorch JIT追踪解析计算图
通过PyTorch的JIT追踪功能,我们可以直接获取计算图的节点信息,再将其序列化为JSON文件:
import torch import torchvision.models as models import json # 初始化模型与输入张量 rn18 = models.resnet18().cuda() x_rn18 = torch.rand(1, 3, 224, 224).cuda() # 追踪模型得到带计算图的ScriptModule traced_model = torch.jit.trace(rn18, x_rn18) # 定义节点解析函数,提取关键信息 def parse_node(node): return { "node_name": node.name(), "operator_kind": node.kind(), "input_tensors": [inp.debugName() for inp in node.inputs()], "output_tensors": [out.debugName() for out in node.outputs()] } # 遍历计算图所有节点并解析 graph_nodes = [parse_node(node) for node in traced_model.graph().nodes()] # 保存为JSON文件 with open("resnet18_graph.json", "w", encoding="utf-8") as f: json.dump({"computational_graph": graph_nodes}, f, indent=2)
方法2:解析TensorBoard生成的事件文件
你已经通过add_graph生成了TensorBoard的可视化数据,可直接解析对应的事件文件提取计算图:
from tensorboard.backend.event_processing.event_accumulator import EventAccumulator import tensorflow as tf from google.protobuf.json_format import MessageToJson import json # 替换为你的SummaryWriter保存路径(默认是./runs) event_acc = EventAccumulator("./runs") event_acc.Reload() # 提取计算图的protobuf数据 graph_proto_str = event_acc.Tensors("graph")[-1].tensor_proto.string_val[0] # 解析protobuf并转为JSON graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(graph_proto_str) graph_json = MessageToJson(graph_def) # 保存文件 with open("resnet18_graph_from_tb.json", "w", encoding="utf-8") as f: f.write(graph_json)
注意:这个方法需要安装tensorflow和protobuf库,可通过pip install tensorflow protobuf安装。
方法3:使用torchviz提取结构化图数据
torchviz不仅能生成可视化图,还能让我们解析计算图的节点与边信息:
import torch import torchvision.models as models from torchviz import make_dot import json # 初始化模型并前向传播得到输出 rn18 = models.resnet18().cuda() x_rn18 = torch.rand(1, 3, 224, 224).cuda() output = rn18(x_rn18) # 生成计算图的Dot对象 dot_graph = make_dot(output, params=dict(rn18.named_parameters())) # 解析节点与边信息 structured_graph = { "nodes": [ {"name": node.name, "operator_label": node.attr.get("label", "")} for node in dot_graph.body.nodes() ], "edges": [ {"from_node": edge[0], "to_node": edge[1]} for edge in dot_graph.body.edges() ] } # 保存为JSON文件 with open("resnet18_graph_torchviz.json", "w", encoding="utf-8") as f: json.dump(structured_graph, f, indent=2)
内容的提问来源于stack exchange,提问作者Maryam
相关产品推荐
相关产品推荐

