Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions src/array/array.cc
Original file line number Diff line number Diff line change
Expand Up @@ -665,6 +665,14 @@ std::vector<NDArray> CSRGetDataAndIndices(

CSRMatrix CSRTranspose(CSRMatrix csr) {
CSRMatrix ret;
#ifdef DGL_USE_ASCEND
if (csr.indptr->ctx.device_type == kDGLAscend) {
ATEN_ID_TYPE_SWITCH(csr.indptr->dtype, IdType, {
ret = impl::CSRTranspose<kDGLAscend, IdType>(csr);
});
return ret;
}
#endif
ATEN_XPU_SWITCH_CUDA(csr.indptr->ctx.device_type, XPU, "CSRTranspose", {
ATEN_ID_TYPE_SWITCH(csr.indptr->dtype, IdType, {
ret = impl::CSRTranspose<XPU, IdType>(csr);
Expand Down
35 changes: 35 additions & 0 deletions src/array/ascend/csr_transpose.cc
Original file line number Diff line number Diff line change
@@ -0,0 +1,35 @@
// ============================================================================
// CSRTranspose Ascend 实现 — CSR 转置(等价于 CSR→CSC)
// ============================================================================
//
// 实现策略:复用已有 Ascend 算子链路
// CSRTranspose(csr) = COOToCSR(COOTranspose(CSRToCOO(csr, false)))
//
// 依赖的 Ascend 已适配算子:
// - CSRToCOO<kDGLAscend, IdType> (src/array/ascend/csr_to_coo.cc)
// - COOTranspose (纯元数据交换,无数据拷贝)
// - COOToCSR<kDGLAscend, IdType> (src/array/ascend/coo2csr.cc)
//
// 与 CUDA int64 路径一致(cuda/csr_transpose.cc:86-88)。
// ============================================================================

#include <dgl/array.h>
#include "../array_op.h"

namespace dgl {
namespace aten {
namespace impl {

template <>
CSRMatrix CSRTranspose<kDGLAscend, int32_t>(CSRMatrix csr) {
return COOToCSR(COOTranspose(CSRToCOO(csr, false)));
}

template <>
CSRMatrix CSRTranspose<kDGLAscend, int64_t>(CSRMatrix csr) {
return COOToCSR(COOTranspose(CSRToCOO(csr, false)));
}

} // namespace impl
} // namespace aten
} // namespace dgl
Loading