对于一个 RNN, naive 实现是每次只计算一个句子。
func rnn( xs:[vec{8}] ) // 假设词向量长度8
n = length(xs)
h_0 = 0
for x in xs
h_t = tanh(W * [h_{t-1}; x_t] + b)
\hat{y} = U * h_n + c
Loss = dist(\hat{y}, y)
在一个batch 中, 计算多个句子的 loss,求导,更新参数。
如果要更高效地计算,需要 padding 和 mask, padding 把一个 batch 中的句子长度补齐, mask 记录每个句子长度。
func rnn_batch(xss:[vec{8}]*b, length:map{int,int})
xss:[mat{8 * b}]
M = 0
for i = i:b
M[i,length(i)] = 1
H_0 = 0
for t = 1:max(length(:))
H_t = tanh(W * [H_{t-1};X_t] + b
\hat{Y}_t = U * H_t + c
L_t = dist( \hat{Y_t} - Y)*(M[:,t] * 1^T))
Loss = sum L_t
L_t 会在每一步都一直往下传,只有在一个句子的结束时,Mask 的值为1,这个计算才有效。
如何自动去batch化计算。
分成三步
- graph definition
- operation batching
- computation
1、3步是框架自带的。
第一步中,一般动态网络可能把计算图构建与 forward 计算同时进行, 这里需要改为 lazy evaluation。
只有在用户调用 forward 时才执行 forward 计算。这样允许把计算任务积攒起来。
计算 compatibility groups
下面 operation 和 node 指相同的东西。
然后在一个计算图中, 把能够一起计算的 operation batch 起来。 先把节点分成 compatibility groups, 每个 group 中的节点能够 batching。通过给每个 node 附一个 signature, signature 记录一个 operation 满足 batch 的必要信息, 如果两个 operation 的 signature 相同则在同一个 group。
如 tanh, log 这种操作, signature 就是函数名, 如果一个操作依赖于输入的大小, 则 signature 还需记录其大小。 如果一个操作依赖于某个特定输入 node,则 signature 记录 node 名字, 比如乘法操作固定左操作数, 将右操作数 batch。
决定执行顺序
寻找一个执行顺序, 满足(1)数据依赖,(2)相同的 signature 的操作 batch 化。 找到一个最优执行顺序是 NP hard的,(这个居然有证明的,后面再说。
两种启发式顺序:
depth-based, 计算节点的深度,定义为叶子到这个节点的最长距离。 把有相同 depth 和 signature 的 node batch计算。 即用深度保证计算顺序满足数据依赖。 有些不同深度的其实也能并行算, 这种方法不能利用。
agenda-based, 本文贡献, 用一个 agenda 放 operation 数据依赖满足了的operation, 对每个 operation 维护一个计数器, 初始为其 input 个数。agenda 初始为没有数据依赖的叶子节点,每次计算从 agenda 中选一批 signature 相同的操作一起算,并从 agenda 中删除,他们的后继的依赖计数器减一,等于0的加入到 agenda 中。 重复到 agenda 里面没有操作。
问题:从 agenda 中选择时,能一次选多个吗, 如果一次只能选一个,是贪心地选最大的 compatibility group 吗? 这个 agenda 机制在框架里应该也有类似的机制,两者怎么结合。
从 agenda 中选择 group 时, 又用了一个启发算法, 计算 signature 的平均 depth,优先选择平均 depth 小的, 直觉是,有些类似 operation 可能在相近的高度,可以避免过早的把一些积累的并不多的 operation 执行了。
评价: 我觉得这方法几乎完美了。
也许可以从 batch 数据选择上补充一下吧。。
在 forward 计算时,需要吧 single nodes 转化为 batched node,为了进行矩阵运算, 还需要内存拷贝。
内存拷贝上能否有什么优化?