Pull Request已成功合入, 合并人@ascend-robot
(感谢 yuht9 的贡献)变更摘要
此 PR 新增了 ExternalStream 类,作为 torch_npu.npu.Stream 的子类,用于包装外部库创建的 aclrtStream 句柄,使 torch_npu 能够在该外部流上执行算子以实现多库数据交换。同时扩展了 C++ 层的流构造逻辑以支持通过 stream_ptr 创建外部流,并添加了对应的单元测试覆盖创建、切换、运算、同步限制等场景。
主要改动
-
新增
ExternalStream类:在torch_npu/npu/streams.py中定义ExternalStream(Stream)子类,通过__new__接收stream_ptr整数值和可选的device参数,并在设备上下文中调用父类构造;类文档详细列出了synchronize()、query()、record_event()等不支持的使用场景及其根因说明。 -
C++ 层支持
stream_ptr构造外部流:在torch_npu/csrc/npu/Stream.cpp的THNPStream_pynew函数中增加对stream_ptr参数的判断分支,当传入stream_ptr时校验priority和is_sync_launch必须为默认值,并调用c10_npu::getStreamFromExternal创建流对象。 -
模块导出更新:在
torch_npu/npu/__init__.py中将ExternalStream加入__all__列表并从streams模块导入,确保用户可通过torch_npu.npu.ExternalStream访问。 -
新增测试用例:在
test/npu/test_stream.py中新增TestExternalStream测试类,覆盖外部流创建、设为当前流、执行张量运算、相同指针返回同一对象、isinstance检查以及synchronize()和query()的限制行为验证。


代码审查
All four changed files are reviewed. Let me now submit the closing summary.
审查总结
按优先级统计
| 优先级 | 数量 |
|---|---|
| P0 | 0 |
| P1 | 0 |
| P2 | 1 |
| P3 | 0 |
各文件审查结果
-
test/npu/test_stream.py— 新增TestExternalStream测试类,7 个测试用例覆盖创建、设为当前流、执行算子、相同指针、isinstance、synchronize/query 限制。_get_external_stream上下文管理器正确创建和销毁原始流。无问题。 -
torch_npu/csrc/npu/Stream.cpp— 新增stream_ptr参数解析和两条 TORCH_CHECK 校验,修改三元表达式链以支持外部流路径。存在 1 个 P2 问题:三元链中stream_id || device_index || device_type优先于stream_ptr,当两者同时传入时外部流创建被绕过。 -
torch_npu/npu/__init__.py— 在__all__和 import 行中新增ExternalStream。无问题。 -
torch_npu/npu/streams.py— 新增ExternalStream(Stream)类,通过__new__将device参数转为设备上下文并透传stream_ptr到 C++ 层。设计合理,参数约束由 C++ 层兜底。无问题。
整体风险评估
低风险。唯一的 P2 问题需要调用者同时传入 stream_ptr 和 stream_id 才会触发,正常通过 ExternalStream Python API 使用不会命中该路径。建议修复三元表达式顺序以消除隐患。
| 类型 | 数量 |
|---|---|
| 🔴 阻塞 | 0 |
| 🟡 建议 | 2 |
💬 仅评论


Pull Request 已合并或已关闭。
If you want to solve this problem, you can click here to do it in the FAQs.




【合入来源】
【修改方案】
ExternalStream是torch_npu.npu.Stream的子类,用于包装外部库创建的aclrtStream句柄,使 torch_npu 能在该流上 launch 算子,实现多库数据交换。类定义
class ExternalStream(Stream): def __new__(cls, stream_ptr, device=None, **kwargs): with torch_npu.npu.device(device): return super().__new__(cls, stream_ptr=stream_ptr, **kwargs)构造参数
stream_ptraclrtStream的整数值devicepriorityis_sync_launch【资料变更】
不涉及
【接口变更】
不涉及
【功能验证】
DT验证通过
【CheckList】