diff --git a/CHANGELOG.md b/CHANGELOG.md index 1c9f9847..f3eabca6 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,7 @@ ## Latest Changes +- Removed OEQ's direct cuBLAS and rocBLAS dependencies. Use torch's cuda c shim instead. + ### v0.7.0 (2026-09-10) **Added**: - Public XLA FFI registration provider diff --git a/openequivariance/CMakeLists.txt b/openequivariance/CMakeLists.txt index 7bfbd499..44caad4a 100644 --- a/openequivariance/CMakeLists.txt +++ b/openequivariance/CMakeLists.txt @@ -117,7 +117,6 @@ endfunction() find_package(CUDAToolkit QUIET) find_package(hip QUIET) -find_package(rocblas QUIET) if(CUDAToolkit_FOUND) message(STATUS "Building stable extension with CUDA backend.") @@ -138,7 +137,6 @@ if(CUDAToolkit_FOUND) CUDA::cudart CUDA::cuda_driver CUDA::nvrtc - CUDA::cublas cuda_stub_lib ) add_stable_extension(oeq_stable_cuda CUDA_BACKEND "${CUDA_LINK_LIBS}") @@ -146,6 +144,7 @@ endif() if(hip_FOUND) message(STATUS "Building stable extension with HIP backend.") + find_package(hiprtc REQUIRED) add_library(hip_stub_lib SHARED ${EXT_DIR}/stubs/stream.cpp) @@ -159,16 +158,10 @@ if(hip_FOUND) CXX_STANDARD 17 ) - if(TARGET roc::rocblas) - set(HIP_BLAS_LIB roc::rocblas) - else() - set(HIP_BLAS_LIB rocblas) - endif() - set(HIP_LINK_LIBS hip_stub_lib hip::host - ${HIP_BLAS_LIB} + hiprtc::hiprtc ) add_stable_extension(torch_stable_hip HIP_BACKEND "${HIP_LINK_LIBS}") endif() diff --git a/openequivariance/openequivariance/_torch/extlib/__init__.py b/openequivariance/openequivariance/_torch/extlib/__init__.py index 2a114b3e..b38080ac 100644 --- a/openequivariance/openequivariance/_torch/extlib/__init__.py +++ b/openequivariance/openequivariance/_torch/extlib/__init__.py @@ -112,7 +112,7 @@ def load_jit_extension(): ], ) if torch.version.cuda: - extra_link_args.extend(["-lcuda", "-lcudart", "-lnvrtc", "-lcublas"]) + extra_link_args.extend(["-lcuda", "-lcudart", "-lnvrtc", "-ltorch_cuda"]) try: torch_libs, cuda_libs = library_paths("cuda") @@ -127,6 +127,8 @@ def load_jit_extension(): elif torch.version.hip: torch_libs = library_paths("cuda")[0] extra_link_args.append("-Wl,-rpath," + torch_libs) + extra_link_args.extend("-L" + path for path in library_paths("cuda")) + extra_link_args.extend(["-ltorch_hip", "-lhiprtc"]) extra_cflags.append("-DHIP_BACKEND") torch_sources = [oeq_root + "/extension/" + src for src in torch_sources] diff --git a/openequivariance/openequivariance/extension/group_mm.hpp b/openequivariance/openequivariance/extension/group_mm.hpp index 19249c39..247b6b3d 100644 --- a/openequivariance/openequivariance/extension/group_mm.hpp +++ b/openequivariance/openequivariance/extension/group_mm.hpp @@ -1,141 +1,73 @@ #pragma once +#include #include #include #include +#include -#ifdef CUDA_BACKEND - #include "cublas_v2.h" - #include +#include - struct BlasHandle { - cublasHandle_t handle; - BlasHandle() { - if (cublasCreate(&handle) != CUBLAS_STATUS_SUCCESS) - throw std::logic_error("CUBLAS initialization failed"); - } - ~BlasHandle() { cublasDestroy(handle); } - }; -#elif defined(HIP_BACKEND) - #include "rocblas/rocblas.h" - #include - - struct BlasHandle { - rocblas_handle handle; - BlasHandle() { - if (rocblas_create_handle(&handle) != rocblas_status_success) - throw std::logic_error("rocBLAS initialization failed"); - } - ~BlasHandle() { rocblas_destroy_handle(handle); } - }; -#endif +namespace oeq { -inline BlasHandle& get_blas_handle() { - static BlasHandle handle; - return handle; +inline void check_group_mm_shim(AOTITorchError status) { + if (status != AOTI_TORCH_SUCCESS) + throw std::runtime_error("group_gemm: PyTorch C shim failed"); } -template -void group_gemm_blas(void* A_raw, void* B_raw, void* C_raw, - int64_t* ragged_counts, int num_W, int batch_size, int m, int k, int ragged_inner) { +using GroupMMTensor = std::unique_ptr< + std::remove_pointer_t, + decltype(&aoti_torch_delete_tensor_object)>; + +inline GroupMMTensor group_mm_view( + AtenTensorHandle tensor, std::array sizes, + std::array strides, int64_t offset) { + AtenTensorHandle view = nullptr; + check_group_mm_shim(aoti_torch__reinterpret_tensor( + tensor, 3, sizes.data(), strides.data(), offset, &view)); + return GroupMMTensor(view, aoti_torch_delete_tensor_object); +} - auto& blas = get_blas_handle(); - T alpha = 1.0, beta = 0.0; - T* A_base = reinterpret_cast(A_raw); - T* B_base = reinterpret_cast(B_raw); - T* C_base = reinterpret_cast(C_raw); +inline void group_gemm_torch( + AtenTensorHandle A, AtenTensorHandle B, AtenTensorHandle C, + const int64_t* ragged_counts, int64_t num_W, int64_t batch_size, + int64_t m, int64_t k, int64_t ragged_inner) { + if (batch_size == 0 || m == 0 || k == 0) + return; - int64_t ragged_offset = 0; - for (int i = 0; i < num_W; i++) { - int M, K, N, lda, ldb, ldc, strideA, strideB, strideC; - T *A, *B, *C; -#ifdef CUDA_BACKEND - cublasOperation_t transa, transb; -#elif defined(HIP_BACKEND) - rocblas_operation transa, transb; -#endif + int64_t offset = 0; + for (int64_t i = 0; i < num_W; ++i) { + const int64_t n = ragged_counts[i]; + if (n == 0) + continue; if (ragged_inner == 0) { - M = m; K = k; N = static_cast(ragged_counts[i]); - A = A_base + (m * k * batch_size * i); - lda = k; strideA = M * K; - B = B_base + (k * batch_size * ragged_offset); - ldb = K * batch_size; strideB = K; - C = C_base + (m * batch_size * ragged_offset); - ldc = M * batch_size; strideC = M; -#ifdef CUDA_BACKEND - transa = CUBLAS_OP_T; transb = CUBLAS_OP_N; -#elif defined(HIP_BACKEND) - transa = rocblas_operation_transpose; transb = rocblas_operation_none; -#endif + auto input = group_mm_view(B, + {batch_size, n, k}, {k, batch_size * k, 1}, + offset * batch_size * k); + auto weight = group_mm_view(A, + {batch_size, k, m}, {m * k, 1, k}, + i * batch_size * m * k); + auto output = group_mm_view(C, + {batch_size, n, m}, {m, batch_size * m, 1}, + offset * batch_size * m); + check_group_mm_shim(aoti_torch_cuda_bmm_out( + output.get(), input.get(), weight.get())); } else { - M = k; K = static_cast(ragged_counts[i]); N = m; - A = B_base + (k * batch_size * ragged_offset); - lda = k * batch_size; strideA = M; - B = A_base + (m * batch_size * ragged_offset); - ldb = m * batch_size; strideB = N; - C = C_base + (m * k * batch_size * i); - ldc = k; strideC = M * N; -#ifdef CUDA_BACKEND - transa = CUBLAS_OP_N; transb = CUBLAS_OP_T; -#elif defined(HIP_BACKEND) - transa = rocblas_operation_none; transb = rocblas_operation_transpose; -#endif - } - ragged_offset += ragged_counts[i]; - - if (ragged_counts[i] > 0) { -#ifdef CUDA_BACKEND - cublasStatus_t stat; - if (std::is_same::value) { - stat = cublasSgemmStridedBatched(blas.handle, - transa, transb, M, N, K, - reinterpret_cast(&alpha), - reinterpret_cast(A), lda, strideA, - reinterpret_cast(B), ldb, strideB, - reinterpret_cast(&beta), - reinterpret_cast(C), ldc, strideC, - batch_size); - } else if (std::is_same::value) { - stat = cublasDgemmStridedBatched(blas.handle, - transa, transb, M, N, K, - reinterpret_cast(&alpha), - reinterpret_cast(A), lda, strideA, - reinterpret_cast(B), ldb, strideB, - reinterpret_cast(&beta), - reinterpret_cast(C), ldc, strideC, - batch_size); - } else { - throw std::logic_error("Unsupported datatype for grouped GEMM!"); - } - if (stat != CUBLAS_STATUS_SUCCESS) - throw std::logic_error("Grouped GEMM failed!"); -#elif defined(HIP_BACKEND) - rocblas_status stat; - if (std::is_same::value) { - stat = rocblas_sgemm_strided_batched(blas.handle, - transa, transb, M, N, K, - reinterpret_cast(&alpha), - reinterpret_cast(A), lda, strideA, - reinterpret_cast(B), ldb, strideB, - reinterpret_cast(&beta), - reinterpret_cast(C), ldc, strideC, - batch_size); - } else if (std::is_same::value) { - stat = rocblas_dgemm_strided_batched(blas.handle, - transa, transb, M, N, K, - reinterpret_cast(&alpha), - reinterpret_cast(A), lda, strideA, - reinterpret_cast(B), ldb, strideB, - reinterpret_cast(&beta), - reinterpret_cast(C), ldc, strideC, - batch_size); - } else { - throw std::logic_error("Unsupported datatype for grouped GEMM!"); - } - if (stat != rocblas_status_success) - throw std::logic_error("Grouped GEMM failed!"); -#endif + auto left = group_mm_view(A, + {batch_size, m, n}, {m, 1, batch_size * m}, + offset * batch_size * m); + auto right = group_mm_view(B, + {batch_size, n, k}, {k, batch_size * k, 1}, + offset * batch_size * k); + auto output = group_mm_view(C, + {batch_size, m, k}, {m * k, k, 1}, + i * batch_size * m * k); + check_group_mm_shim(aoti_torch_cuda_bmm_out( + output.get(), left.get(), right.get())); } + offset += n; } } + +} diff --git a/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp b/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp index ddabd0bb..cf458306 100644 --- a/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp +++ b/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp @@ -2,7 +2,7 @@ #include #ifdef CUDA_BACKEND - #include + #include #endif #ifdef HIP_BACKEND @@ -10,9 +10,11 @@ #endif #include +#include #include #include #include +#include #include using Tensor = torch::Tensor; @@ -29,6 +31,20 @@ constexpr Dtype kByte = torch::kByte; #define REGISTER_LIBRARY_IMPL TORCH_LIBRARY_IMPL #define REGISTER_LIBRARY TORCH_LIBRARY +namespace { + +class TensorDeviceGuard { + c10::DeviceGuard guard; +public: + explicit TensorDeviceGuard(const Tensor& tensor) : guard(tensor.device()) {} +}; + +AtenTensorHandle tensor_handle(Tensor& tensor) { + return torch::aot_inductor::tensor_pointer_to_tensor_handle(&tensor); +} + +} + #include "torch_core.hpp" Tensor tensor_to_cpu_contiguous(const Tensor &tensor) { diff --git a/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp b/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp index 6bf3d51f..642ab2b0 100644 --- a/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp +++ b/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp @@ -27,6 +27,20 @@ constexpr Dtype kByte = torch::headeronly::ScalarType::Byte; #define REGISTER_LIBRARY_IMPL STABLE_TORCH_LIBRARY_IMPL #define REGISTER_LIBRARY STABLE_TORCH_LIBRARY +namespace { + +class TensorDeviceGuard { + torch::stable::accelerator::DeviceGuard guard; +public: + explicit TensorDeviceGuard(const Tensor& tensor) : guard(tensor.get_device_index()) {} +}; + +AtenTensorHandle tensor_handle(Tensor& tensor) { + return tensor.get(); +} + +} + #include "torch_core.hpp" Tensor tensor_to_cpu_contiguous(const Tensor &tensor) { diff --git a/openequivariance/openequivariance/extension/stubs/stream.cpp b/openequivariance/openequivariance/extension/stubs/stream.cpp index fd011c35..346493c4 100644 --- a/openequivariance/openequivariance/extension/stubs/stream.cpp +++ b/openequivariance/openequivariance/extension/stubs/stream.cpp @@ -1,8 +1,15 @@ +#define USE_CUDA + #include -#include +#include extern "C" { AOTITorchError aoti_torch_get_current_cuda_stream(int32_t device_index, void** ret_stream) { - return 0; + return AOTI_TORCH_FAILURE; + } + + AOTITorchError aoti_torch_cuda_bmm_out( + AtenTensorHandle out, AtenTensorHandle self, AtenTensorHandle mat2) { + return AOTI_TORCH_FAILURE; } -} \ No newline at end of file +} diff --git a/openequivariance/openequivariance/extension/torch_core.hpp b/openequivariance/openequivariance/extension/torch_core.hpp index ab78d96a..52dd3331 100644 --- a/openequivariance/openequivariance/extension/torch_core.hpp +++ b/openequivariance/openequivariance/extension/torch_core.hpp @@ -621,13 +621,40 @@ inline tuple jit_conv_double_backward( inline Tensor group_gemm( Tensor A, Tensor B, Tensor ragged_counts, int64_t num_W, int64_t batch_size, int64_t m, int64_t k, int64_t ragged_inner) { + TCHECK(A.is_cuda() && B.is_cuda(), "group_gemm: A and B must be GPU tensors"); + TCHECK(A.get_device() == B.get_device(), "group_gemm: A and B must be on the same device"); TCHECK(A.scalar_type() == B.scalar_type(), "group_gemm: A and B must have the same dtype"); + TCHECK(A.scalar_type() == kFloat || A.scalar_type() == kDouble, + "group_gemm: unsupported dtype, expected float32 or float64"); + TCHECK(ragged_counts.is_cpu(), "group_gemm: ragged_counts must be on the CPU"); TCHECK(ragged_counts.scalar_type() == kLong, "group_gemm: ragged_counts must be int64"); + TCHECK(num_W >= 0 && batch_size >= 0 && m >= 0 && k >= 0, + "group_gemm: dimensions must be nonnegative"); + TCHECK(ragged_inner == 0 || ragged_inner == 1, "group_gemm: ragged_inner must be 0 or 1"); + TCHECK(ragged_counts.dim() == 1 && ragged_counts.size(0) == num_W, + "group_gemm: ragged_counts must contain num_W entries"); + TCHECK(B.dim() == 3, "group_gemm: B must have shape [rows, batch_size, k]"); + const int64_t rows = B.size(0); + check_tensor(B, {rows, batch_size, k}, A.scalar_type(), "group_gemm B"); + if (ragged_inner == 0) + check_tensor(A, {num_W, batch_size, m, k}, A.scalar_type(), "group_gemm A"); + else + check_tensor(A, {rows, batch_size, m}, A.scalar_type(), "group_gemm A"); + + Tensor rc_c = tensor_contiguous(ragged_counts); + const auto* rc_ptr = static_cast(data_ptr(rc_c)); + int64_t row_count = 0; + for (int64_t i = 0; i < num_W; ++i) { + TCHECK(rc_ptr[i] >= 0 && rc_ptr[i] <= rows - row_count, + "group_gemm: ragged_counts must be nonnegative and sum to the number of rows"); + row_count += rc_ptr[i]; + } + TCHECK(row_count == rows, + "group_gemm: ragged_counts must sum to the number of rows"); + TensorDeviceGuard device_guard(A); Tensor A_c = tensor_contiguous(A); Tensor B_c = tensor_contiguous(B); - Tensor rc_c = tensor_contiguous(ragged_counts); - int64_t* rc_ptr = reinterpret_cast(data_ptr(rc_c)); Tensor C; if (ragged_inner == 0) { @@ -637,15 +664,8 @@ inline Tensor group_gemm( C = tensor_zeros_like(A, make_sizes({num_W, batch_size, m, k})); } - if (A.scalar_type() == kFloat) { - group_gemm_blas(data_ptr(A_c), data_ptr(B_c), data_ptr(C), rc_ptr, - (int)num_W, (int)batch_size, (int)m, (int)k, (int)ragged_inner); - } else if (A.scalar_type() == kDouble) { - group_gemm_blas(data_ptr(A_c), data_ptr(B_c), data_ptr(C), rc_ptr, - (int)num_W, (int)batch_size, (int)m, (int)k, (int)ragged_inner); - } else { - throw std::logic_error("group_gemm: unsupported dtype, expected float32 or float64"); - } + oeq::group_gemm_torch(tensor_handle(A_c), tensor_handle(B_c), tensor_handle(C), + rc_ptr, num_W, batch_size, m, k, ragged_inner); return C; }