[TVM代码剖析] NNVM计算图抽象机制

Hint: 建议配合NNVM以及TVM源码享用,风味最佳。

设计概要

图的抽象

简而言之就是用Symbol同时用来抽象计算图中的OperatorOperand

  • Variable Symbol
  • Functor Symbol(AtomicSymbol), Callable语义

准确地讲,Symbol这个抽象概念表示一种具有多个输入和输出的过程Symbol对象仅仅包含一个vector<NodeEntry> outputs成员,用来记录在图中它的所有输出节点;而每一个Node记录的则是自己的属性和所有的输入inputs。所以说真正负责建图的对象是Node,它能够看到自己的所有输入;然后在计算图之上所定义的抽象概念,例如OperandOperator使用Symbol来表示。

在图的表示中,可以用Operator来将一些Operandcompose为新的一个Operand,也就是图中的节点。最后我们拿到一个Symbol,从它就可以回溯出整个图,转换为Graph对象来完成建图操作。

不同的Operator会有自己的属性,常用的属性定义在op_attr_types.h里面,基本上就是不同的类型声明。

核心数据结构

  • Node: 表示IndexedGraph中的一个节点,节点基本上分为OperatorVariable
  • NodeEntry: 表示图中某个节点的输入。因此对于这样一个输入,我们需要在它的数据结构内部记录它属于哪个节点的输出,以及是第几个输出。

Python接口

首先描述一下MXNet对计算图基本要素的定义。对于单个Symbol以及组合出来的Symbol也就是Network来说,arguments指的是能够前向运行网络所必须的数据,例如输入数据datalabel、各个层的权重,例如Convolutionkernels,以及Batch Norm\(\gamma\)\(\beta\)auxiliary_states指的是一些运行时需要的,并且有必要被保存下来的特殊状态,例如Batch Normmoving_meanmoving_var之类。另外,inputs指的是argumentsauxiliary_states的并集。最后outputs当然就是指网络的最终输出,它多数情况下都是某种Loss,不过也有可能只是计算出的数据,例如GANGenerator

NNVM这里的接口实现非常漂亮,首先C API定义的导出接口就很少,Python代码中仅仅定义最核心的Symbol类、运算符重载以及必要的ctypes胶水代码。至于如何用Operator来建图,则是在C++代码中通过静态定义注册进去,然后在import nnvm的时候动态注册进Python。具体实现里面有更多的细节。

实现细节

自动求导

基本原理

自动求导的原理比较简单,注意所有变量都是Tensor。考虑操作符Operator的基本性质与函数稍有区别,是从多个输入映射到多个输出。所以设当前节点的操作符接受m个输入,返回n个输出,则形式化表示为:

$$Op(x_1, x_2, ..., x_m) = [f_1(x_1, ..., x_m), ..., f_n(x_1, ..., x_m)]$$

其中\(f_1, ..., f_n\)为不同输出所对应的运算过程(函数)。

为每个操作符定义Gradient运算过程,输入为误差\(y\)对当前节点所有输出的导数(一个list),输出为误差对当前节点所有输入的导数(同样是一个list)。输入形式化表示为\([\frac{\partial y}{\partial f_1(...)}, \frac{\partial y}{\partial f_2(...)}, ..., \frac{\partial y}{\partial f_n(...)}]\),那么输出就应该是:

$$ [\sum_{i=1}^{n}\frac{\partial y}{\partial f_i}\frac{\partial f_i}{\partial x_1}, \sum_{i=1}^{n}\frac{\partial y}{\partial f_i}\frac{\partial f_i}{\partial x_2}, ..., \sum_{i=1}^{n}\frac{\partial y}{\partial f_i}\frac{\partial f_i}{\partial x_m}] $$

另外,\(\frac{\partial y}{\partial f_i(...)}\)是把从\(y\)到操作符Op的某个输出\(f_i(...)\)的所有偏导路径都求和之后,累加得出的结果。这也就意味着所有操作符都需要定义累加运算的属性,如果没有的话,在MXNet中默认会尝试调用__ewise_sum__进行累加。注意以上这些运算过程都是符号运算而非数值运算,也就是运算会产生图上新的节点。

接下来我们就直接按照逆拓扑序遍历网络,把当前节点的所有输出的导数的聚合结果(一个节点)\([\frac{\partial y}{\partial f_i(...)}]\)求出来,然后用操作符的Gradient方法求出对所有输入的导数\([\frac{\partial y}{\partial x_i}]\),再不断反向传播即可。逆拓扑序可以保证遍历到当前节点的时候,它的所有输出节点,也就是反向传播时的所有输入节点,都已经求导完成了。

在神经网络中,我们做梯度下降需要误差\(y\)对权重矩阵\(W\)的导数,所以在自动求导的过程中,我们只需要将关心的这些权重节点的导数节点输出即可。事实上,如果我们知道哪些节点是权重,我们就可以让自动求导器产生他们的导数节点。然后在训练的时候,为这些梯度项指定storage,这样遍历一遍计算图之后,就能够直接取出我们想要的梯度进行更新了。

MXNet实现细节

下图为一个简单的两层MLP的Forward过程计算图:

MLP Forward

那么在反向求导的过程就是这个样子:

  1. 求出softmaxfc3softmax_label的导数softmax_backward.
  2. 此时已知\(y\)fc3节点的所有输出的导数是[softmax_backward],那么就可以计算\(y\)fc3的所有输入[relu2, fc3_weight, fc3_bias]的导数。
  3. 以此类推,逆拓扑序完成整个求导过程……

MLP Gradient

这里我们注意到,softmax具有两个shape明显不同的输入节点,但却只有一个导数节点softmax_backward;同样地fc3有三个不同的输入,但却也只有一个导数节点fc3_backward,其他节点的情况都与此类似。这里实际上是MXNet重构计算图设计的历史遗留问题。因为绝大多数旧操作符并不是按照上述设计来实现的——他们并不包含一个Gradient方法,能够根据节点的输入直接按照符号运算求出导数节点。所以为了兼容,MXNet注册操作符的时候在NNVM中为这些旧版本的操作符注册了加前缀_backward_的反向操作符。在对计算图进行求导的时候,会直接生成一个Operator_backward_版本的节点,并且让操作符声明自己在反向求导过程中所依赖的数据。然后这个_backward_节点在实际执行计算图的时候,会按照反向的行为进行计算,得到正确的导数结果。

Type/Shape Inference

通过自动求导构建计算图以后,我们就能以输入Tensorshape,推导出图中所有节点的shape,从而用于后续的内存分配等过程。这一步要求各个Operator实现一个InferShape方法,其输入是当前节点指针以及该操作符的所有输入数据的shape,返回值是该操作符所有输出数据的shape。对于每个操作符来说,实现这样的功能是很简单的:比如对于element-wise加法,那么输出数据的shape就等于输入数据的shape;对于矩阵乘法则需要考虑是否转置之类的情况。

由于上文提到的MXNet的历史遗留问题,对于反向过程要更麻烦一些,需要针对性的特殊处理。这里面用到一个“控制流依赖”(control flow dependencies)的关系来进行反向过程的类型推导。至于具体怎么实现的……这代码我是不想继续纠结了。

内存分配

首先贴上架构设计文档:Optimizing Memory Consumption in Deep Learning

为计算图的所有节点分配内存的问题可以抽象为:给定一个内存块Request/Free的操作序列,要求满足所有分配的需求,并且使得总的分配内存数量最少。所以这是个NP问题么?

内存分配器会有一个参数match_range_,用来表示在[size/match_range_, size*match_range_]的范围内来寻找内存块。这里面的trick在于先试图分配大的内存块,然后找不到的话再试图分配小的内存块。当然这里并不是真的在分配内存,而只是预先规划我要分配怎样的内存,如果找到小的内存块,肯定不满足我们的需求,我们此时把它放大到我们想要的size即可,我们现在是在记录需求,反正最终运行的时候才会真的分配内存。

接下来分析实现细节。整体上分为初始化阶段和按照拓扑序遍历计算图的阶段。

初始化阶段

  1. 内存分配阶段要依赖于ShapeTypeInference,这是显然的,不然分配个毛啊。在注册这个Pass的时候,会指定这种依赖关系。
  2. 然后要计算所有非Variable节点的出度,作为refcount;有些操作符具有FIgnoreInputs属性,并不需要输入数据(只要shape),比如zeroslike这样的操作符,所以遍历的时候不要算这部分的引用计数。
  3. 输出节点要额外加一个引用计数(出度+1),保证在计算图执行到结束的时候也不会回收这些内存。这一点很重要,我就踩过坑。

拓扑序遍历阶段

这一阶段直接是一个for循环,以拓扑序遍历整个计算图,循环体内所做的事情如下:

  1. 首先是检查是否能做in-place运算优化。Operator可以设置自己支持inplace操作来显式优化内存分配,所以内存分配的时候是先处理能够inplace的情况,然后再操作正常的内存分配。另外inplace优化实际上可能是一对多的关系,就是说运算符可以指定一个输入节点的内存可能被复用给多个输出节点,因为可能有的输出节点只需要shape信息,不需要数据本身,根本不用给他分配空间。最后,inplace优化需要满足一个挺复杂的条件:
    • 输入节点只对应一个输出(出度为1)
    • 输出节点有被其他节点引用(否则就不需要为它分配内存,因为根本不用算它)
    • 输出节点尚未分配内存
    • 输入节点已分配内存(拓扑序遍历的话,这个条件应该是默认满足的)
    • 数据类型、大小匹配
  2. 接下来就开始遍历当前节点的所有输出了,把所有还没分配内存的节点记录下来排个序,从小到大依次向内存分配器请求内存即可。
  3. 然后我们就可以更新引用计数了:把所有输入节点(排除FIgnoreInputs节点)的refcount - 1,如果refcount == 0,就可以释放这个节点的内存。另外这时会遇到有些节点出度本来就是零,这是因为inplace优化导致的,跳过就行了。
  4. 最后我们还需要遍历一遍输出节点,把那些出度为零的节点的内存释放掉,标记为不需要分配内存,因为他们根本不被用到,对用户来说处于“不可见状态”。

设备分配

这个pass的功能就是在计算的时候设置不同的节点/子图在哪个设备上进行计算。如果出现了跨设备的数据依赖,就增加一个设备之间数据拷贝的节点。显然,这样的设备分配策略会改变计算图的结构,所以这里采用经典的持久化数据结构,即只增加不修改,从而基于原计算图得到新的计算图。它的策略看起来相对简单:

  1. 初始状态下,为所有节点赋予对应的设备编号device_id,默认为-1invalid);
  2. 然后开始拓扑序遍历计算图,每个节点可以包含一个属性(string),指明这个节点在计算的时候属于哪个分组(group);同时计算图自己也会有一个设备分配的映射关系,是从group到计算设备(device_id)的映射。这样如果一个节点有分组信息,就可以直接根据计算图的映射来找到它所对应的设备了。如果节点木有设备分组属性的话,就直接把当前节点的输入节点的设备作为当前节点的运算设备,这是显然的。
  3. 我们还需要逆拓扑序遍历一遍。猜测这是因为正拓扑序遍历能够保证所有前向运算分配到合适的设备,但是逆拓扑序才能保证给反向运算分配合适的设备。此时所有的节点就都具有自己的device_id了。
  4. 接下来开始为所有跨设备的运算操作之间插入数据复制的Operator。这是一个相对复杂的过程:
    • 首先要进行一个合法性检查,因为NNVM的设计允许Operator去就地mutate它的inputs,但如果当前节点与它的输入节点并不在一个设备上面,就不能这样做,因为这样的就地mutate不能通过插入一个copy节点来实现。
    • 然后检查对于当前节点是否需要变更图结构。假设我们知道经过所有操作以后,我们生成了一个新图,那么显然在新图中的每个节点和原图存在一一对应的关系。这里采用了一个new_node_map的哈希表,来表示某个节点在新图中所对应的节点。
    • 如果当前节点的输入inputs有节点需要被映射到新的节点,那么我们就也需要为当前节点创建一个新的映射节点,然后将那个对应的输入节点指向它所对应的新节点。
    • 并且,如果当前节点的某个输入所对应的device_id与其自身的device_id不同,那么我们就插入一个新的copy节点,用于在执行的时候将数据正确地从不同设备之间进行拷贝。然后这个新的节点就可以被加入new_node_map中。
  5. 最后我们返回一个新图,方法就是把旧图中的outputs节点替换成new_node_map中的对应节点。

一个简单的例子,考虑如下代码:

:::Python
import mxnet as mx
a = mx.sym.Variable('a')
b = mx.sym.Variable('b')
c = mx.sym.Variable('c')
with mx.AttrScope(ctx_group='dev1'):
    net = a * b
with mx.AttrScope(ctx_group='dev2'):
    net = net + c
e = net.simple_bind(mx.cpu(), a=(10, 10), grad_req='write',
                    group2ctx={'dev1': mx.cpu(), 'dev2': mx.gpu()})
mx.viz.plot_network(net)

也就是说,默认整个net运行在CPU上面,指定net = a + b运行在CPU上;但是指定其中一个步骤net = net + c运行在GPU上面。

则生成的计算图如下:

Gradients

经过了Plan Device Pass之后,该计算图被修改为如下所示:

Gradients

可以看到,为了满足中间计算过程的跨设备要求,c以及a * b的运算结果被显式拷贝操作符复制到GPU设备上面,进行加法运算以及反向加法求导之后,又被拷贝回CPU端来继续进行乘法求导。

算子融合

算子融合(Operator Fusion)的核心目的是为了降低片外访存,从而提高硬件带宽利用率。并不是所有算子都能够相互融合,NNVM将其由简单到复杂,分为几种主要的融合模式(Fuse Pattern):

  1. ElemWise:两个相同shape的Tensor进行元素与元素之间的算数操作,例如elementwise_add。这是最简单的场景,data locality也最好。
  2. BroadCast:输出Tensor的每一个元素都可以被唯一映射到输入Tensor的相应元素,但是要求axis是保序的。似乎只有broadcast系列算子满足这个条件——两个shape不同的Tensor之间的算数操作,在运算时,较小的那个Tensor会被broadcast到较大的Tensor的shape,然后再执行计算,其映射规则例如\(out_{i, j}=\sum_{m,n}in_{i,m,j,n}\)。反例如transpose就不符合这个条件,它的映射是\(out_{i,j}=in_{j,i}\)
  3. Injective:满足上面条件的前半句但却不满足后半句的算子属于这个类型。与上一个类型的本质区别就是由于axis不保序,导致代码执行的时候访存的data locality降低了。
  4. CommReduce:
  5. OutEWiseFusable:复杂的算子例如convolution,最多只能在输出的时候fuse上一个简单的elementwise算子,但是在其代码内部并不能做更加复杂的融合操作。
  6. Opaque:完全无法相互融合的算子,例如topk

所有算子在定义的时候,需要为自己注册一个TOpPattern属性,表示自己的固有融合模式,其值属于上述几种类型之一。

算子融合的实际执行分为两个主要步骤:

  1. 将整个计算图切分成一系列子图,每个子图内部的算子将被融合——相当于将子图里的所有OpNode合并为一个多输入、多输出的复杂OpNode
  2. 按照上个阶段的子图划分,真正建立新的一系列子图。
  3. 为各个子图调用编译后端,生成相应的kernel代码。

子图切分

首先以正拓扑序遍历计算图,根据每个计算节点(OpNode)的预定义信息,为当前和前驱节点其求出几个重要的属性。因为这个pass并没有修改计算图本身的结构,所以实际上是用这些属性来描述了子图会怎样被切分。在代码实现中均以vector的形式存放这些属性,最后会注册给Graph对象:

  • fuse_rule:标识这个节点是否应该被fuse到其他节点。这里只会有两种可选项:kFuseToMaster以及kRealize。前者表示将当前节点合并到相应子图的master节点;后者表示直接为其生成代码无需融合。
  • master_node:表示如果要融合的话,这块子图以哪个节点为主来生成代码,把所有其他节点都合并过来。由于这个pass不修改计算图本身,只追加新的属性,因此这个属性同样相当于为图进行了着色,从而能够划分出各个子图。
  • fuse_pattern:表示这块子图如果作为整体的OpNode来看待的话,它的融合模式是什么。虽然单个算子OpNode都有自己固有的融合模式,但是并不能代表融合子图之后也具有同样的融合模式,要以其中最复杂的fusion模式为准。例如,如果整个子图都属于elementwise算子之间的融合,那么这块子图的融合模式就是ElemWise。我们可以在生成代码的时候直接把数据进行flatten操作;但如果其中有一个是convolution跟其他的elementwise融合了,那么这个子图的融合模式就属于OutEWiseFusable,我们就不能去flatten数据了。

正向遍历计算图时,具体的执行逻辑为:

  • 如果当前节点是一个输入参数(argument,也就是TF里的placeholder),那么无需做额外处理。
  • 读取当前节点自身的融合模式,如果属于最简单的ElemWise和Broadcast,则遍历所有输入(前驱)节点并检查其所对应子图的fuse_pattern是什么
    • 如果是ElemWise、Broadcast、Injective以及OutEWiseFusable的话,就能直接与当前节点融合。
    • PS:OutEWiseFusable的判断条件比较复杂,当前节点只能与一个属性为OutEWiseFusable的输入节点进行融合,并且输入的shape必须与当前节点的输出shape能够match才行。如果找到了能够融合的输入节点,当前节点的master_node也要设置为那个输入节点的master_node。简而言之,原则就是以复杂的OpNode作为master。
    • 否则的话,无法与当前节点融合,选择为其直接生成代码。
    • 这里由于融合了其他节点,因此整块子图的融合模式可能会发生变化,所以需要根据实际情况进行相应的更新。
  • 如果当前节点的融合模式属于较复杂的Injective或者CommReduce,那么只能跟输入中复杂度小于等于Injective的子图进行融合,否则的话就只能选择直接生成代码。但是为什么两个相连的CommReduce无法融合?个人理解是这种情况并不会发生,因为两个reduction完全可以合并成一个,除非要用到中间结果——此时融合也没有意义了。
  • 对于其他情况,选择直接为相应算子生成代码。
  • 在循环的末尾,我们需要根据实际融合的情况,来更新当前节点所对应子图的fuse_pattern——可能比节点自身的融合模式更加复杂。

接下来我们需要逆拓扑序遍历计算图,从而得到一个新的属性:group_root。这里所谓root与之前提到的group_master有一个细微的区别,root表示节点所属的整块子图最终往外输出的节点是哪个,用它来标志(染色)整个子图;而master表示代码生成的时候以谁为主,将其他节点fuse过来。所以这里我们需要逆序遍历计算图,这样才能够优先处理拓扑序靠后的节点——也就是最终的输出节点;之后再逐渐向前递推,从而得出整个图中所有节点的group_root属性。循环体内部主要的逻辑分为三部分:

  • 首先检查当前节点是否具有group_root属性,如果尚未设置,则将当前节点设置为root。由于是逆拓扑序遍历,所以我们保证了所有子图内对外输出的节点会被优先设置为root。
  • 然后检查当前节点的那些标记为需要做融合的输入节点是否同时有OutEWiseFusable类型以及Injective类型,如果是的话,就相当于如下图所示的模式(op表示一个不可融合的算子),此时我们其实已经把addsigmoid进行了融合,这样生成的代码在写回内存时访存模式为\(mem[add]_{i,j}=mem[conv]_{i,j}+mem[op]_{i,j}+\sigma_{i,j}\),可以省去一次内存读取。但是当前节点的master_node还是指向输入的OutEWiseFusable节点(因为上一次遍历时是这样处理的),需要把它改成指向自己;并且同一子图内的所有节点的master_node也都要改为当前节点。个人理解这一步有些画蛇添足,应该是算法上的补丁。

    | | | conv2d op sigmoid \ | / \ | / add |

  • 最后我们再次遍历输入节点,向其传播group_root属性——也就是设置所有需要融合的输入节点的group_root指向自己的group_root即可。此外,记得在第二步的时候,如果我们更新了同一子图内所有节点的master_node,现在我们也需要把这个属性传播给所有的输入节点——如果其fuse_pattern为Injective的话。

经过上述算法的处理之后,仍然存在如下所示的子图是无法被融合的,因为conv2d的输出被多个不同的算子使用。

    conv2d
    /  |  \
   /   |   \
 op    op   op
  |    |    |
  |    |    |

但是考虑如下一种特殊的场景,如果几个不同分支很快又交汇到一起,并且他们又都是elementwise算子,那么就可以都融合到一起。这种场景在ResNet里面非常多,所以专门有一段逻辑进行优化。

    conv2d
    /  |  \
   /   |   \
 op    op   op
  \    |    /
   \   |   /
  elemwise add
       |

事实上,如果遇到这样的模式,已执行的算法会把图中下半部分,即中间的三个OpNode以及下面的汇节点给融合成一个算子,相当于如下所示的这个样子:

    conv2d
   /   |   \
  /    |    \
  \    |    /
   \   |   /
   fused op
       |

接下来所需要做的只是将上面的conv2d与下面的融合算子再次进行融合。算法主要分为两步:

  • 首先我们逆拓扑序遍历一次计算图,求出有哪些节点的输出被不止一个节点使用,以及具体是被哪些节点使用——记录下他们的group_root用于后续处理。
  • 然后再次逆拓扑序遍历一遍计算图:
    • 这次要检查那些输出会被多个子节点所使用的节点,其所有子节点是否都属于同一个子图,并且都属于Broadcast/ElemWise模式,也就是最简单的融合模式。
    • 如果是的话。说明这个节点的各个输出节点都已经被融合到一起去了,我们可以放心地把这个节点与它的输出节点再次进行融合。注意这次融合相当于删除了一个子图,因此我们需要把删掉的子图的master_nodefuse_pattern以及group_root属性都进行更新。

新图构造

经过上述处理,我们将计算图划分为一系列可以融合成独立算子的子图,并为每个节点附加了master_nodegroup_root以及fuse_pattern这三个属性,从而能够标识出这些子图的划分。与上一个pass不同的是,这个pass中我们会基于现有的子图切分,为每块子图填充一个称之为FuseEntry的数据结构:

struct FuseEntry {
  // 用来表示一块需要进行融合的子图
  Graph subgraph;
  // 将这块子图的输入NodeEntry映射到新图对应的NodeEntry
  std::unordered_map<IndexedGraph::NodeEntry, nnvm::NodeEntry, ...> imap;
  // 将新图中的Node映射到旧图中的NodeEntry
  std::unordered_map<const Node *, IndexedGraph::NodeEntry> reverse_imap;
  // TVM Placeholder for inputs
  std::unordered_map<const Node *, Tensor> input_info;
  // Whether we can flatten data
  bool flatten_data;
  // The corresponding function.
  GraphFunc compiled_func;
};

这个过程本质上非常类似于图的镜像复制(请先脑补相应算法),建立出的数据结构实际上相当于原图的镜像,但在结构上发生了一些变化,从而用于下一步的图编译。

在第一遍正拓扑序遍历中,主要是找到每个子图中来自于其他子图或者是placeholder的输入NodeEntry,然后为这些输入创建对应的新Variable节点,最后填充FuseEntry中的几个哈希表,从而维护好新旧两个图中NodeEntry的映射关系——通过FuseEntry.imap可以直接得到原图中IndexedGraph::NodeEntry所对应的新nnvm::NodeEntry

在第二次遍历中,我们以正拓扑序遍历图中每一个节点,并为其创建一个新的镜像节点Node::Create())。首先我们检查各个节点的所有输入节点:

  • 如果是来自其他子图的输出,就找到上次遍历中新建的Variable的对应输出,连接到这个镜像Node的输入(即追加到Node::inputs)。
  • 否则的话,就用原本那个输入节点所对应的镜像节点的nnvm::NodeEntry来作为输入。这里由于是正拓扑序遍历,因此可以保证输入的映射都已经被正确设置好了。

其次我们检查当前节点是否是group_root,如果不是,我们需要填充上面用到的映射关系;如果是的话,那么它的输出不会被子图内的节点所用到,所以我们需要将其加入FuseEntry.subgraph.outputs里面,作为整块子图的输出。

由此我们可以看到,这个过程与图的镜像复制的主要区别,就在于我们额外新建了一些Variable节点来作为每块子图的输入,而不改变子图内部的连接关系,从而使得每个FuseEntry.subgraph真正成为了一个新的独立的子图。

经过上述的过程,我们实际上相当于得到了一系列的独立的子图(std::vector<FuseEntry>),例如下面这样:

在上图中,conv2d2group_master,而relu2则是group_root,注意其中的区别;来自于其他子图的输入则被替换成了input?

代码生成

最终我们根据一系列子图信息(FuseEntry),调用TVM的codegen后端来生成实际的算子代码。这里主要分为如下几个步骤:

  • 首先遍历所有的子图,调用GraphLower()函数将其编译到最终的low level IR(也就是HalideIR),主要是应用相应的loop transformation,得到变换过后的代码。注意,每一个OP会对应一系列的LoweredFunc对象。
  • 然后我们需要生成新的计算图。注意这次建图与上一步的区别——先前我们只是建立了各个子图的结构,从而针对各个融合后的子图生成了代码,但计算图本身并没有被调整为融合过后的结构,所以还不能真正用于执行。所以这一步我们要将计算图重新调整为经过融合的,并且具有正确的依赖关系,可以执行的结构。算法的实现思路其实很简单,我们只需要正拓扑序遍历原始计算图,为每块子图建立一个对应的新Node,设置其属性,然后正确设置子图之间的连接关系即可,也就是填充好Node.inputs这个vector
  • 最终,我们调用编译后端,将所有的LoweredFunc编译成可执行的二进制代码,并设置到Graph"module"属性里面去。至此,图编译的流程就完成了;接下来还需要再做一些运行时相关的优化,例如内存分配等等,这就与图编译无关了。

一个融合过后的计算图示意如下,删去了输入label使得节点名称更加清晰(Resnet-18):

可以看到,经过融合后的计算图虽然还是有飞线,但已经完全变成了串行的形式。

代码风格&槽点

  • 一句话概括TVM的整体风格就是有种用Python思想写C++的感觉。
  • Operator的注册机制(Registry)对全局状态用的有点多,而且有些优化trick显得意义不大,宏接口设计的倒是比较漂亮。虽然将算子分散到各个源文件来注册不同属性是个挺不错的主意,但其实并不能实现任意覆盖已注册的属性:这就一脚踩进了C++在编译单元之间静态变量的初始化顺序未定义的坑里——事实上是按照链接顺序的逆序执行初始化的。
  • 尽管template到处飞,但是并没有很好地用静态类型来约束代码,到处出现的Op::GetAttr<FGradient>()里面那个参数字符串看得人难受。
  • 计算图优化的部分采用一个pass一个编译单元的结构,代码可读性一般,建图的数据结构目测是借鉴了LLVM,没有具体benchmark过,姑且认为IndexedGraph这样的设计能起到一定的优化效果吧。
  • Python接口与ctypes部分写得很赞,与NNVM_REGISTER_OP的协作非常漂亮,使得C++实现的算子能够在import的时候自动注册到modules里面。就是……读起来有点绕😂
comments powered by Disqus
Published:
2018-12-01
分类:
Tag: