- 深度学习
- 机器学习
- 人工智能
【免费下载链接】mxnet
Lightweight, Portable, Flexible Distributed/Mobile Deep Learning with Dynamic, Mutation-aware Dataflow Dep Scheduler; for Python, R, Julia, Scala, Go, Javascript and more
本文是 Apache MXNet 官方 API 参考文档中 mxnet.visualization 模块索引页 的深度展开。该索引页通过 Sphinx
automodule指令将mxnet.visualization模块的全部成员自动生成为 API 文档,其背后对应仓库中 python/mxnet/visualization.py 这一实际实现模块。读完本文,你将掌握如何使用mxnet.viz.plot_network将 Symbol 计算图渲染为 Graphviz 有向图、如何使用mxnet.viz.print_summary输出类 Keras 风格的逐层文本摘要,并理解其底层的形状推断、参数统计与权重隐藏等实现机制。
一、模块概览:mxnet.visualization是什么
mxnet.visualization是 MXNet Python 前端中专门负责神经网络结构可视化的工具模块。它对外提供两个核心函数:
| 函数 | 作用 | 输出形态 |
|---|---|---|
print_summary(symbol, ...) | 以文本表格形式打印网络的逐层结构 | 终端/日志文本 |
plot_network(symbol, ...) | 将计算图渲染为 Graphviz 有向图(digraph)对象 | GraphvizDigraph对象,可保存或展示 |
在 python/mxnet/init.py 中,模块被同时以全名和短名两种方式导出:
from . import visualization from . import visualization as viz因此,import mxnet之后既可以写mxnet.visualization.plot_network(...),也可以使用短形式mxnet.viz.plot_network(...)。本文以下统一使用mxnet.viz短形式,与官方文档、教程及仓库内示例代码保持一致。
需要说明的是,plot_network的依赖graphviz(Python 库)与系统级 Graphviz 程序均需预先安装,否则函数会在导入阶段直接抛出ImportError(详见 python/mxnet/visualization.py)。
二、print_summary:文本化的网络结构摘要
2.1 函数签名与参数
print_summary(symbol, shape=None, line_length=120, positions=[.44, .64, .74, 1.])各参数说明(依据 python/mxnet/visualization.py 中的 docstring):
symbol(必填):待可视化的Symbol对象。若非Symbol类型,会抛出TypeError("symbol must be Symbol")。shape(可选,dict):输入形状字典,键为输入符号名(str),值为形状元组(tuple)。提供后,摘要中的 Output Shape 列会被填充,同时 BatchNorm 等层的参数统计也会依赖形状信息;若形状推断不完整,会抛出ValueError("Input shape is incomplete")。line_length(可选,int,默认 120):打印行的总长度(字符数)。positions(可选,list,默认[.44, .64, .74, 1.]):表格各列在行内的相对(小于等于 1)或绝对位置。源码在positions[-1] <= 1时会将其换算为int(line_length * p),即按比例扩展为绝对列宽(python/mxnet/visualization.py)。
2.2 输出格式与列含义
print_summary输出一个以分隔线框起的表格,表头固定为四列(python/mxnet/visualization.py):
Layer (type) Output Shape Param # Previous Layer- Layer (type):层名 + 括号内的算子类型,例如
conv1(Convolution); - Output Shape:该层输出的张量形状,多个维度用
x连接(如32x64x64); - Param #:该层可学习参数量;
- Previous Layer:直接上游层名;若一层有多个输入,会以追加行列出其余上游层。
函数返回None,所有信息直接print到标准输出,最后一行输出Total params: {params}汇总整个网络的总参数量(python/mxnet/visualization.py)。
2.3 参数统计的源码实现
各算子的参数量并非运行时从权重张量统计,而是基于算子属性推算(python/mxnet/visualization.py):
- Convolution:
pre_filter * num_filter // num_group * prod(kernel),若no_bias不为'True'则再加num_filter(偏置项)。 - FullyConnected:
(pre_filter + 1) * num_hidden(含偏置);若no_bias=True则为pre_filter * num_hidden。 - BatchNorm:
num_filter * 2,其中num_filter来自输出形状的第二维(即 gamma、beta 两个可学习向量)。 - Embedding:
input_dim * output_dim。
其中pre_filter是上一层输出的通道数(取输出形状的第一维),_str2tuple内部通过正则re.findall(r"\d+", string)从"3,3"这类属性字符串中提取各维度数值(python/mxnet/visualization.py)。
2.4 完整示例(与官方单测一致)
仓库的 tests/python/unittest/test_viz.py 给出了同时覆盖"无形状"与"带形状"两种调用的示例,直接可运行:
import mxnet as mx data = mx.sym.Variable('data') bias = mx.sym.Variable('fc1_bias', lr_mult=1.0) emb1 = mx.symbol.Embedding(data=data, name='emb1', input_dim=100, output_dim=28) conv1 = mx.symbol.Convolution(data=emb1, name='conv1', num_filter=32, kernel=(3,3), stride=(2,2)) bn1 = mx.symbol.BatchNorm(data=conv1, name="bn1") act1 = mx.symbol.Activation(data=bn1, name='relu1', act_type="relu") mp1 = mx.symbol.Pooling(data=act1, name='mp1', kernel=(2,2), stride=(2,2), pool_type='max') fc1 = mx.sym.FullyConnected(data=mp1, bias=bias, name='fc1', num_hidden=10, lr_mult=0) fc2 = mx.sym.FullyConnected(data=fc1, name='fc2', num_hidden=10, wd_mult=0.5) sc1 = mx.symbol.SliceChannel(data=fc2, num_outputs=10, name="slice_1", squeeze_axis=0) # 不带形状:Output Shape 列留空 mx.viz.print_summary(sc1) # 带输入形状:输出每一层的形状与参数 shape = {"data": (1, 3, 28)} mx.viz.print_summary(sc1, shape)运行后将看到类似下面的摘要(分隔线数量由line_length控制):
________________________________________________________________________________________________________________________ Layer (type) Output Shape Param # Previous Layer ======================================================================================================================== emb1(Embedding) 1x28x100x28 2800 conv1(Convolution) 1x32x50x14 8064 emb1 bn1(BatchNorm) 1x32x50x14 64 conv1 relu1(Activation) 1x32x50x14 0 bn1 mp1(Pooling) 1x32x25x7 0 relu1 fc1(FullyConnected) 1x10 5600 mp1 fc2(FullyConnected) 1x10 110 fc1 slice_1(SliceChannel) 1x10 0 fc2 ======================================================================================================================== Total params: 14638 ________________________________________________________________________________________________________________________说明:
SliceChannel未命中源码中的参数统计分支,其Param #计为 0;上表输出形状为按shape={'data': (1,3,28)}推断的示意结果,实际打印以你运行环境输出的数值为准。
三、plot_network:Graphviz 计算图可视化
3.1 函数签名与参数
plot_network(symbol, title="plot", save_format='pdf', shape=None, dtype=None, node_attrs={}, hide_weights=True)各参数说明(依据 python/mxnet/visualization.py 中的 docstring):
symbol(必填):Symbol对象,可视化其计算所需的子图部分。title(可选,str,默认"plot"):生成的可视化标题,也是Digraph(name=title, ...)的图名。save_format(可选,str,默认'pdf'):保存格式,直接透传给 GraphvizDigraph(format=...),支持pdf、png、svg等 Graphviz 支持的格式。shape(可选,dict):输入张量形状字典,键为输入符号名(str),值为形状元组。指定后,节点之间的边上会标注张量形状(如100x200)。dtype(可选,dict):输入张量类型字典,键为输入符号名,值为类型(如numpy.float32)。指定后,边上会追加(类型名)标注,例如100x200(float32)。node_attrs(可选,dict):Graphviz 节点属性覆盖字典。默认节点属性为{"shape": "box", "fixedsize": "true", "width": "1.3", "height": "0.8034", "style": "filled"}(python/mxnet/visualization.py),传入的键值会覆盖默认值。例如node_attrs={"shape":"oval","fixedsize":"false"}会将节点改为椭圆并允许节点随内容自适应大小。hide_weights(可选,bool,默认True):为True时,隐藏名称以_weight、_bias、_beta、_gamma、_moving_var、_moving_mean、_running_var、_running_mean结尾的节点,使图更清爽(python/mxnet/visualization.py)。
返回值:一个 GraphvizDigraph对象。可调用.view()直接打开预览,或通过.render(...)/.save(...)输出到文件。
3.2 节点的样式与配色规则
从源码的节点渲染逻辑(python/mxnet/visualization.py)可以看到完整的样式约定:
- 输入节点(op 为
null):形状为椭圆(oval),使用配色#8dd3c7(浅青),标签为变量名; - Convolution:标签为
Convolution\n{kernel}/{stride}, {num_filter},配色#fb8072(浅红); - FullyConnected:标签为
FullyConnected\n{num_hidden},配色#fb8072; - Activation:标签为
Activation\n{act_type},配色#ffffb3(浅黄); - LeakyReLU:标签为
LeakyReLU\n{act_type}(未显式指定时默认Leaky),配色#ffffb3; - BatchNorm:配色
#bebada(浅紫); - Pooling:标签为
Pooling\n{pool_type}, {kernel}/{stride},配色#80b1d3(浅蓝); - Concat / Flatten / Reshape:配色
#fdb462(浅橙); - Softmax:配色
#fccde5(浅粉); - 其他算子:统一使用
#b3de69(浅绿);Custom算子则以attrs["op_type"]作为标签。
这一 8 色调色板定义于 python/mxnet/visualization.py,使不同算子类别在图中可快速区分。
3.3 官方示例:构建并可视化一个三层 MLP
plot_networkdocstring 中的示例(python/mxnet/visualization.py)构建了一个data → fc1 → relu1 → fc2 → out的 MLP 并可视化:
import mxnet as mx net = mx.sym.Variable('data') net = mx.sym.FullyConnected(data=net, name='fc1', num_hidden=128) net = mx.sym.Activation(data=net, name='relu1', act_type="relu") net = mx.sym.FullyConnected(data=net, name='fc2', num_hidden=10) net = mx.sym.SoftmaxOutput(data=net, name='out') digraph = mx.viz.plot_network(net, shape={'data': (100, 200)}, node_attrs={"fixedsize": "false"}) digraph.view()digraph.view()会调用系统默认 PDF 查看器打开图形;若要保存为文件,可使用digraph.render('net')(生成的文件扩展名由save_format决定)。
3.4 带形状与类型标注
在 tests/python/unittest/test_viz.py 的测试中,同时传入了shape与dtype,边上的标签会形如100x200(float32):
import numpy as np digraph = mx.viz.plot_network(net, shape={'data': (100, 200)}, dtype={'data': np.float32}, node_attrs={"fixedsize": "false"})形状与类型标注的生成逻辑位于 python/mxnet/visualization.py:对每个非输入节点的上游边,从infer_shape/infer_type得到的结果字典中取出对应输出条目(多输出算子还会拼接num_outputs下标),将形状(x连接)与dtype.__name__写入边的label。
四、源码级实现原理
4.1 形状与类型推断
两个函数都依赖 Symbol 的静态推断能力:
internals = symbol.get_internals() # 展开为完整计算图内部节点 _, out_shapes, _ = internals.infer_shape(**shape) _, out_types, _ = internals.infer_type(**dtype)- 推断失败(返回
None)时抛出ValueError("Input shape is incomplete")或ValueError("Input type is incomplete")(python/mxnet/visualization.py); - 随后通过
dict(zip(internals.list_outputs(), out_shapes))建立"输出名 → 形状"的查找表,print_summary与边标注都从该表中取值; - 展示形状时统一去掉批次维度:
out_shape = shape_dict[key][1:],因此图中标注的是去掉 batch 后的张量形状。
4.2 重复节点名检测与环警告
构建图之前,源码会检查图中是否存在重名节点(python/mxnet/visualization.py):
if len(nodes) != len(set([node["name"] for node in nodes])): ... warning_message = "There are multiple variables with the same name in your graph, " \ "this may result in cyclic graph. Repeated names: " + ','.join(repeated) warnings.warn(warning_message, RuntimeWarning)重名通常意味着变量被复用,可能导致渲染出环。仓库测试 tests/python/unittest/test_viz.py 专门构造了两个都叫fc的全连接层,断言会收到一条包含重名提示的RuntimeWarning。
4.3 权重隐藏与特殊算子处理
- 权重隐藏:
hide_weights=True时,looks_like_weight命中的节点被加入hidden_nodes,连边阶段会跳过指向这些节点的边(python/mxnet/visualization.py);hide_weights=False时这些节点仍会以空椭圆渲染。 - 多输出算子:边标注会依据
num_outputs属性定位正确的输出下标(python/mxnet/visualization.py)。 _contrib_BilinearResize2D特例:只保留第一个输入参与画边,规避其特殊的多输入语义(python/mxnet/visualization.py)。- 边方向:边统一使用
dir="back"与arrowtail="open",即以反向箭头从下游节点指向上游节点,配合节点布局形成清晰的计算流向(python/mxnet/visualization.py)。
五、仓库内的实战案例
5.1 SSD 检测网络可视化工具
example/ssd/tools/visualize_net.py 是一个独立可运行的工具脚本,用plot_network渲染 SSD 目标检测网络:
a = mx.viz.plot_network(net, shape={"data": (1, 3, args.data_shape, args.data_shape)}, ...)该脚本通过命令行参数指定网络配置、数据形状与输出文件,是"将plot_network封装为工程工具"的参考模板。
5.2 推荐系统矩阵分解网络
example/recommenders/demo1-MF.ipynb 在 Notebook 中直接对 user/item 双塔 Embedding 网络调用:
mx.viz.plot_network(net1(mx.sym.var('user'), mx.sym.var('item')), node_attrs={"fixedsize": "false"})这与 docs/static_site/src/pages/api/faq/visualize_graph.md 中演示的矩阵分解可视化流程一致:user、item两个输入分别经 Embedding 查表,内积求和后接LinearRegressionOutput,plot_network能清晰地呈现出"输入节点(椭圆)→ 计算节点(矩形)→ 输出节点"的完整数据流。
5.3 模型转换与量化场景
plot_network还常用于检查模型结构:
- ONNX 模型导入后可视化:见 docs/python_docs/python/tutorials/packages/onnx/inference_on_onnx_model.md;
- MKLDNN 量化前后对比:见 docs/python_docs/python/tutorials/performance/backend/mkldnn/mkldnn_quantization.md,其中对原符号图、量化符号图、反量化符号图分别调用
mx.viz.plot_network(sym)生成对比图。
六、常见问题与排查
| 现象 | 原因 | 处理方式 |
|---|---|---|
ImportError("Draw network requires graphviz library") | 未安装 Python 包graphviz或系统 Graphviz | pip install graphviz,并确保系统已安装 Graphviz 程序(如apt-get install graphviz) |
TypeError("symbol must be Symbol") | 传入的不是mx.sym.Symbol实例 | 确认传入对象由mx.sym.*算子构建,而非 NDArray 或 Gluon Block |
ValueError("Input shape is incomplete") | shape字典缺少某个输入的形状,或形状信息不足以推断所有节点 | 为每个输入变量补全{"输入名": (维度元组)} |
RuntimeWarning: There are multiple variables with the same name... | 图中存在重名节点,可能形成环 | 检查网络定义,为变量/层使用唯一名称 |
| 图中看不到权重/偏置节点 | hide_weights=True(默认)隐藏了*_weight、*_bias等节点 | 设置hide_weights=False显示全部节点 |
七、延伸阅读
- 模块完整实现:python/mxnet/visualization.py
- 官方单元测试:tests/python/unittest/test_viz.py
- 模块在
mxnet顶层 API 中的位置:docs/python_docs/python/api/mxnet/index.rst - 可视化官方教程入口:docs/python_docs/python/tutorials/packages/viz/index.rst
- FAQ 指南:docs/static_site/src/pages/api/faq/visualize_graph.md
mxnet.visualization以极简的 API 提供了从"终端文本摘要"到"Graphviz 计算图"的完整网络结构观测手段:print_summary适合在训练脚本里快速核对层序与参数量,plot_network适合在论文、报告或模型调试阶段直观呈现数据流与算子分布。两者均基于 Symbol 的静态形状推断,无需真正执行网络即可获得结构信息,是 MXNet 使用者理解、调试与展示模型的重要工具。
- 深度学习
- 机器学习
- 人工智能
【免费下载链接】mxnet
Lightweight, Portable, Flexible Distributed/Mobile Deep Learning with Dynamic, Mutation-aware Dataflow Dep Scheduler; for Python, R, Julia, Scala, Go, Javascript and more
相关推荐
MXNet 计算图可视化完全指南:mxnet.visualization 的 print_summary 与 plot_network 深度解析
MXNet 计算图可视化完全指南:mxnet.visualization 的 print_summary 与 plot_network 深度解析 mxnet.v
深度学习人工智能机器学习分布式训练终极指南:如何高效批量下载网页资源 - Chrome扩展完全解析
终极指南:如何高效批量下载网页资源 Chrome扩展完全解析 你是否曾经为了下载一个网页上的所有图片、CSS和JavaScript文件而头疼不已?手动一个个保存
开发工具深入解析 PyG 的图神经网络可解释性模块 torch_geometric.explain
深入解析 PyG 的图神经网络可解释性模块 torch_geometric.explain 本指南以 docs/source/modules/explain.r
人工智能机器学习深度学习图计算
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考