You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.11 12:05:21