已合并
fix: support packed LSTM output shape transition #5761
ymdsall创建于 9 天前
fix: support packed LSTM output shape transition #5761
已合并
Pull Request已成功合入, 合并人@ascend-robot
(感谢 ymdsall 的贡献)ascend-robot
9 天前 评论:
9 天前 评论:
atomgit-bot
9 天前 评论:
9 天前 评论:
变更摘要
本 PR 修复 packed LSTM 在 NPU kernel 上的输出 shape 兼容问题:将 batch_sizes 的设备适配与输出 reshape 逻辑下沉到 op-plugin 的 LstmKernelNpu.cpp(ACLop)和 LSTMKernelNpuOpApi.cpp(ACLNN)中。核心规则是仅在 batch_sizes 位于 CPU 时将输出 y reshape 为 PackedSequence 所需的二维形状,若旧 PTA monkey-patch 传入 NPU batch_sizes 则保留原输出、由 Python patch 完成 reshape,从而同时兼容保留和删除旧 LSTM.forward patch 的 PTA 版本。同时补充了 ACLop 直调输出形状、BF16、混合 dtype 及参数数量 fallback 等测试。
主要改动
- ACLop 输出 reshape 下沉:在
op_plugin/ops/aclops/LstmKernelNpu.cpp的lstm中新增should_reshape_output = batch_sizes.device().is_cpu()判断,为真时对输出y执行y.reshape({-1, y.size(-1)}),内部仍继续使用 CPUbatch_sizes。 - ACLNN 设备适配与 reshape:在
op_plugin/ops/opapi/LSTMKernelNpuOpApi.cpp的lstm中新增batch_sizes与data设备不一致时经to(data.device())转换并传入_lstm_npu,同时按同一should_reshape_output规则对output_y做二维 reshape,保证与 ACLop 路径行为一致。 - 兼容旧 monkey-patch 场景:两处 kernel 均以
batch_sizes.device().is_cpu()为分界,NPUbatch_sizes时保持原输出不变,由旧 PTA 的 Python patch 负责 reshape,避免重复处理。 - 补充 packed LSTM 测试:在
test/test_base_ops/test_lstm.py中为_build_lstm/_make_packed_inputs增加num_layers参数,新增混合 dtype、BF16、参数数量 fallback 及 ACLop 直调输出形状测试,将断言收紧为y.shape == (packed.data.size(0), 6),并通过PackedSequence回填pad_packed_sequence验证输出可还原为(5, 3, 6)。


atomgit-bot
9 天前 评论:
9 天前 评论:
9 天前 添加了label:needs-issue
此处折叠了48条消息 查看更多
wanglijun55
6 天前 评论:
6 天前 评论:
/lgtm


chengpeng25
6 天前 评论:
6 天前 评论:
/approve


6 天前 添加了label:approvedlgtm
6 天前 合入了pull request
6 天前 关联了里程碑:TorchNPU-v26.2.0
【合入来源】
【修改方案】
batch_sizes时,在 C++ 层将输出 reshape 为 PackedSequence 所需的二维数据,用于直调和删除 Python patch 后的 PTA。batch_sizes时保留原输出形状,由旧 PTA 的LSTM.forwardmonkey-patch 完成 reshape,避免重复处理。batch_sizes设备;ACLop 继续在内部使用 CPUbatch_sizes,ACLNN 在调用算子前按需转到数据设备。本 PR 是 LSTM.forward monkey-patch 删除前的兼容修改。合入顺序为:本 PR → pytorch PR 43893 删除 monkey-patch → op-plugin MR 5675 rebase 后合入完整修改和测试。
【资料变更】
不涉及。
【接口变更】
不涉及。
【功能验证】
git diff --check:通过。batch_sizes路径。【CheckList】