提交提案之前,请先检索仓库内是否已有相同的提案,如已有请在同一提案中进行讨论。
1、需求背景:部分网络使用torch_npu的moe_token_permute类接口,包括npu_moe_token_permute与对应的反向接口。过程中发现存在内存问题,为反向存储了不必要的数据,存在优化空间。
2、需求价值: 训练内存优化
3、需求详细: 描述: 开发过程中发现在正向的过程中会保留多余的数据用于反向。 例如npu_moe_token_permute及其反向接口配置如下,在反向aclnn算子中事实上并不需要用到tokens和indices,但为了获得grad_tokens和num_topk入参用到了tokens和indices,这影响内存的使用效率。如tokens的shape为[SBH],当序列长度和模型层数变大之后,影响随之加大。
解决方案: 1、框架层正向完成相应元数据计算,保存在context中,从context获取对应数值 2、为了保证兼容性,增加grad_v2接口,反向绑定到v2接口上
4、测试方案:UT测试验证,精度性能无劣化;性能在真实网络训练可优化
欢迎加入社区,感谢您对社区的贡献 🎉!
提交提案之前,请先检索仓库内是否已有相同的提案,如已有请在同一提案中进行讨论。
💻 需求背景、当前现状、期望实现的功能内容、具体的设计方案、以及测试方案
1、需求背景:部分网络使用torch_npu的moe_token_permute类接口,包括npu_moe_token_permute与对应的反向接口。过程中发现存在内存问题,为反向存储了不必要的数据,存在优化空间。
2、需求价值:
训练内存优化
3、需求详细:
描述:
开发过程中发现在正向的过程中会保留多余的数据用于反向。
例如npu_moe_token_permute及其反向接口配置如下,在反向aclnn算子中事实上并不需要用到tokens和indices,但为了获得grad_tokens和num_topk入参用到了tokens和indices,这影响内存的使用效率。如tokens的shape为[SBH],当序列长度和模型层数变大之后,影响随之加大。
解决方案:
1、框架层正向完成相应元数据计算,保存在context中,从context获取对应数值
2、为了保证兼容性,增加grad_v2接口,反向绑定到v2接口上
4、测试方案:UT测试验证,精度性能无劣化;性能在真实网络训练可优化
替代方案
补充说明
欢迎加入社区,感谢您对社区的贡献 🎉!