🌱 Stage 1

On the Fly Batching

本文概念: 其他,-先在这里放着

对于一个 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化计算。

分成三步

  1. graph definition
  2. operation batching
  3. 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,为了进行矩阵运算, 还需要内存拷贝。

内存拷贝上能否有什么优化?

其他, 先在这里放着

💬 评论加载中...