Resource Counting
Total_compute = 6 * model_weight * token_num Forward pass: 2(# data points)(# parameters) FLOPs Backward pass: 4(# data points)(# parameters) FLOPs
bf16 用于参数、激活值、梯度 fp32 用于优化器状态 optimizer states
num_parameters: D * D * L flops: 6 * B * num_parameters
混合精度训练: Pytorch 会自动优化数据精度,如果是做矩阵乘法,fp16是安全的,但如果是指数运算,则会保留为fp32
缩小维度的矩阵运算并不能提高提升运算速度,要看 计算密度 是否满足了硬件的计算密度规格, 硬件计算密度是规格上的 单位时间浮点数计算量 / 单位时间通信字节量, 计算密度是算法的 浮点数计算量 / 总搬运字节量, 可以由此来判断是 memory bound 还是 compute bound
Total memory: parameter_memory = 2 * (D * D * L) #(2 bytes for bf16) activation_memory = 2 * B * D * L #(2 bytes for bf16) gradient_memory = 2 * num_parameters #(2 bytes for bf16) optimizer_state_memory = 4 * num_parameters #(4 bytes for fp32)s Memory 主要有两个作用:
- 需要在 HBM 中存储这些数据
- 这些数据必须传输到加速器上
如何减少训练时的显存占用? rematerialization Key idea: 对于前向传播,只保留一部分层的激活值,反向传播在遇到最近的checkpoint的时候把缺失的层激活值再算出来。