From 43ec31ac2ac3c2c887c41f5b194268c8354082b1 Mon Sep 17 00:00:00 2001 From: Abhishek Bagusetty Date: Wed, 26 Aug 2026 16:16:51 -0500 Subject: [PATCH 01/12] feat: add SYCL backend for Intel GPU (XPU) support Kernels are generated as SYCL free functions and compiled at runtime through the oneAPI kernel compiler. The Jinja templates are shared with CUDA and HIP via templates/sycl_compat.cuh; the backend is now a string rather than an is_hip boolean, and JITKernel::execute also takes argument sizes. --- CHANGELOG.md | 16 + README.md | 11 +- docs/index.rst | 6 +- docs/installation.rst | 31 +- openequivariance/CMakeLists.txt | 48 ++- .../openequivariance/_torch/E3NNConv.py | 3 +- .../_torch/E3NNTensorProduct.py | 17 +- .../openequivariance/_torch/FlashTPConv.py | 3 +- .../_torch/NPDoubleBackwardMixin.py | 113 +++--- .../openequivariance/_torch/TensorProduct.py | 24 +- .../_torch/TensorProductConv.py | 49 +-- .../_torch/extlib/__init__.py | 91 ++++- .../openequivariance/benchmark/correctness.py | 3 +- .../core/ComputationSchedule.py | 6 +- .../openequivariance/core/ConvolutionBase.py | 36 +- .../openequivariance/core/LoopUnrollConv.py | 4 +- .../openequivariance/core/LoopUnrollTP.py | 4 +- .../core/TensorProductBase.py | 55 ++- .../openequivariance/core/utils.py | 25 +- .../extension/backend/backend_cuda.hpp | 5 +- .../extension/backend/backend_hip.hpp | 5 +- .../extension/backend/backend_sycl.hpp | 337 ++++++++++++++++++ .../extension/convolution.hpp | 47 ++- .../openequivariance/extension/group_mm.hpp | 75 ++++ .../extension/kernel_args.hpp | 32 ++ .../extension/libtorch_tp_jit.cpp | 16 + .../extension/libtorch_tp_jit_stable.cpp | 38 +- .../extension/stubs/stream.cpp | 14 +- .../extension/tensorproducts.hpp | 26 +- .../openequivariance/extension/torch_core.hpp | 53 ++- .../openequivariance/jax/TensorProduct.py | 2 +- .../openequivariance/jax/TensorProductConv.py | 2 +- .../openequivariance/jax/extlib/__init__.py | 4 + .../openequivariance/templates/common.cuh | 4 + .../openequivariance/templates/jinja_utils.py | 43 ++- .../templates/loop_unroll_batch.cuh | 11 +- .../templates/loop_unroll_conv_atomic.cuh | 11 +- .../templates/loop_unroll_conv_det.cuh | 11 +- .../openequivariance/templates/macros.jinja | 11 + .../templates/sycl_compat.cuh | 116 ++++++ tests/batch_test.py | 6 +- tests/conftest.py | 22 ++ tests/conv_test.py | 6 +- tests/example_test.py | 22 +- tests/export_test.py | 34 +- tests/input_validation_test.py | 22 +- tests/multidevice_test.py | 11 +- tests/stream_test.py | 33 +- tests/symmetric_contraction_test.py | 6 +- 49 files changed, 1302 insertions(+), 268 deletions(-) create mode 100644 openequivariance/openequivariance/extension/backend/backend_sycl.hpp create mode 100644 openequivariance/openequivariance/extension/kernel_args.hpp create mode 100644 openequivariance/openequivariance/templates/sycl_compat.cuh diff --git a/CHANGELOG.md b/CHANGELOG.md index 1c9f9847..9d2f9d38 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,21 @@ ## Latest Changes +**Added**: +- SYCL backend, bringing support for Intel GPUs through PyTorch's + `xpu` device. Kernels are generated as SYCL free functions and compiled + at runtime through the oneAPI kernel compiler extension. The backend is + selected automatically from the active PyTorch build. + +**Changed**: +- The kernel backend is now identified by a string (`"cuda"` / `"hip"` / + `"sycl"`) rather than an `is_hip` boolean, in the Jinja environment and + the `LoopUnrollTP` / `LoopUnrollConv` constructors. +- `JITKernel::execute` also takes the kernel argument sizes, which SYCL + requires to launch with raw arguments. CUDA and HIP ignore them. +- float64 Clebsch-Gordon coefficients are emitted without the `L` + (long double) literal suffix. The values are unchanged — a hex float + literal is already exactly a double — but SPIR-V targets reject the suffix. + ### v0.7.0 (2026-09-10) **Added**: - Public XLA FFI registration provider diff --git a/README.md b/README.md index 3efe59a7..a750e6a4 100644 --- a/README.md +++ b/README.md @@ -6,8 +6,9 @@ [[JAX Examples]](#jax-examples) [[Citation and Acknowledgements]](#citation-and-acknowledgements) -OpenEquivariance is a CUDA and HIP kernel generator for the Clebsch-Gordon tensor product, +OpenEquivariance is a CUDA, HIP, and SYCL kernel generator for the Clebsch-Gordon tensor product, a key kernel in rotation-equivariant deep neural networks. +It targets NVIDIA, AMD, and Intel GPUs. It implements some of the tensor products that [e3nn](https://e3nn.org/) supports commonly found in graph neural networks @@ -20,6 +21,11 @@ and GCC 9+ available before installing our package via pip install openequivariance ``` +On Intel GPUs, install a PyTorch build with XPU support and make the +oneAPI DPC++ compiler (`icpx`) available on your `PATH`; the kernels are +compiled at runtime through SYCL. Tensors live on the `xpu` device there +instead of `cuda`. + We provide up to an order of magnitude acceleration over e3nn perform on par with the latest version of [NVIDIA cuEquivariance](https://github.com/NVIDIA/cuEquivariance), which has a closed-source kernel package. @@ -74,7 +80,8 @@ print(torch.norm(Z)) ``` And here's the same tensor product using openequivariance. We require that your -tensors are stored on a CUDA device for this to work: +tensors are stored on a GPU device for this to work +(``cuda`` for NVIDIA and AMD GPUs, ``xpu`` for Intel GPUs): ```python import openequivariance as oeq diff --git a/docs/index.rst b/docs/index.rst index 3d5f055b..f58c14bc 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -6,15 +6,15 @@ OpenEquivariance ============================== -`OpenEquivariance `_ is a CUDA and -HIP kernel generator for the Clebsch-Gordon +`OpenEquivariance `_ is a CUDA, +HIP, and SYCL kernel generator for the Clebsch-Gordon tensor product, a key kernel in equivariant graph neural networks. We offer an identical interface to e3nn and produce the same results (up to numerical roundoff). Our package exhibits up to an order of magnitude speedup over e3nn and competitive performance with NVIDIA's cuEquivariance. Here, you can find our API reference, installation instructions, -and troubleshooting guide. We support for both NVIDIA and AMD GPUs through +and troubleshooting guide. We support NVIDIA, AMD, and Intel GPUs through our PyTorch interface, including support for JITScript compilation accessible from C++. diff --git a/docs/installation.rst b/docs/installation.rst index 5ade5c0c..4c0b737c 100644 --- a/docs/installation.rst +++ b/docs/installation.rst @@ -7,13 +7,21 @@ Installation You need the following to install OpenEquivariance: -- A Linux system equipped with an NVIDIA / AMD graphics card. +- A Linux system equipped with an NVIDIA / AMD / Intel graphics card. - Either PyTorch >= 2.4 (>= 2.8 for AOTI and export), or JAX>0.5.0 - with CUDA or RocM support. -- GCC 9+ and the CUDA / HIP toolkit. The command + with CUDA, RocM, or XPU support. +- GCC 9+ and the CUDA / HIP toolkit, or the oneAPI DPC++ compiler + (``icpx``) for Intel GPUs. The command ``c++ --version`` should return >= 9.0; see below for details on setting an alternate compiler. +.. note:: + + On Intel GPUs the kernels are generated as SYCL and compiled at runtime + through the oneAPI kernel compiler, so ``icpx`` must be on your ``PATH``. + Tensors passed to the kernels live on the ``xpu`` device rather than + ``cuda``. The JAX frontend currently supports CUDA and HIP only. + .. tab:: PyTorch Installation is one easy command, followed by import verification: @@ -143,4 +151,19 @@ on a major cluster, send us a pull request to add your configuration! conda activate export CC=cc - export CXX=CC \ No newline at end of file + export CXX=CC + + +.. tab:: ALCF Aurora (Intel Data Center GPU Max) + + Aurora provides a PyTorch build with XPU support in its ``frameworks`` + module; no separate PyTorch install is needed. + + .. code-block:: bash + :caption: env.sh (last updated August 2026) + + module load oneapi/release + module load frameworks + + export CC=icx + export CXX=icpx \ No newline at end of file diff --git a/openequivariance/CMakeLists.txt b/openequivariance/CMakeLists.txt index 7bfbd499..df05b1ee 100644 --- a/openequivariance/CMakeLists.txt +++ b/openequivariance/CMakeLists.txt @@ -173,6 +173,50 @@ if(hip_FOUND) add_stable_extension(torch_stable_hip HIP_BACKEND "${HIP_LINK_LIBS}") endif() -if(NOT CUDAToolkit_FOUND AND NOT hip_FOUND) - message(WARNING "Neither CUDAToolkit nor HIP was found. The stable extension will not be built.") +# SYCL: detected via the compiler accepting -fsycl (IntelLLVM / icpx). +set(OEQ_SYCL_FOUND FALSE) +if(CMAKE_CXX_COMPILER_ID MATCHES "IntelLLVM") + set(OEQ_SYCL_FOUND TRUE) +endif() + +if(OEQ_SYCL_FOUND) + message(STATUS "Building stable extension with SYCL backend.") + + add_library(sycl_stub_lib SHARED ${EXT_DIR}/stubs/stream.cpp) + + target_include_directories(sycl_stub_lib PRIVATE + ${LIBTORCH_INCLUDE_DIR} + ) + + set_target_properties(sycl_stub_lib PROPERTIES + OUTPUT_NAME "torch_xpu" + POSITION_INDEPENDENT_CODE ON + CXX_STANDARD 17 + ) + + target_compile_definitions(sycl_stub_lib PRIVATE SYCL_BACKEND=1) + + find_package(MKL QUIET COMPONENTS SYCL) + if(TARGET MKL::MKL_SYCL::BLAS) + set(SYCL_BLAS_LIB MKL::MKL_SYCL::BLAS) + else() + set(SYCL_BLAS_LIB mkl_sycl_blas) + endif() + + set(SYCL_LINK_LIBS + sycl_stub_lib + sycl + ${SYCL_BLAS_LIB} + ) + add_stable_extension(oeq_stable_sycl SYCL_BACKEND "${SYCL_LINK_LIBS}") + + # -fsycl is required on both the compile and the link line. + foreach(tgt oeq_stable_sycl oeq_stable_sycl_aoti) + target_compile_options(${tgt} PRIVATE -fsycl) + target_link_options(${tgt} PRIVATE -fsycl) + endforeach() +endif() + +if(NOT CUDAToolkit_FOUND AND NOT hip_FOUND AND NOT OEQ_SYCL_FOUND) + message(WARNING "None of CUDAToolkit, HIP or SYCL was found. The stable extension will not be built.") endif() diff --git a/openequivariance/openequivariance/_torch/E3NNConv.py b/openequivariance/openequivariance/_torch/E3NNConv.py index 4cc20662..b4c3e61a 100644 --- a/openequivariance/openequivariance/_torch/E3NNConv.py +++ b/openequivariance/openequivariance/_torch/E3NNConv.py @@ -6,6 +6,7 @@ ) from openequivariance._torch.E3NNTensorProduct import E3NNTensorProduct from openequivariance._torch.NPDoubleBackwardMixin import NumpyDoubleBackwardMixinConv +from openequivariance._torch.extlib import DEVICE_TYPE class E3NNConv(ConvolutionBase, NumpyDoubleBackwardMixinConv): @@ -31,7 +32,7 @@ def __init__(self, config, *, idx_dtype=np.int64, torch_op=True): path_normalization=config.path_normalization, internal_weights=config.internal_weights, shared_weights=config.shared_weights, - ).to(device="cuda") + ).to(device=DEVICE_TYPE) self.reference_tp = E3NNTensorProduct(config) diff --git a/openequivariance/openequivariance/_torch/E3NNTensorProduct.py b/openequivariance/openequivariance/_torch/E3NNTensorProduct.py index df696ad6..1fe21eca 100644 --- a/openequivariance/openequivariance/_torch/E3NNTensorProduct.py +++ b/openequivariance/openequivariance/_torch/E3NNTensorProduct.py @@ -13,6 +13,7 @@ from openequivariance.core.e3nn_lite import TPProblem from openequivariance.core.logging import getLogger from openequivariance._torch.NPDoubleBackwardMixin import NumpyDoubleBackwardMixin +from openequivariance._torch.extlib import DEVICE_TYPE TORCH_COMPILE_AUTOTUNING_DIR = pathlib.Path("triton_autotuning") @@ -48,7 +49,7 @@ def __init__(self, config: TPProblem, torch_op=True): path_normalization=config.path_normalization, internal_weights=config.internal_weights, shared_weights=config.shared_weights, - ).to(device="cuda") + ).to(device=DEVICE_TYPE) if config.irrep_dtype == np.float64: torch.set_default_dtype(torch.float32) # Reset to default @@ -62,9 +63,9 @@ def forward_cpu( L3_out: np.ndarray, weights: np.ndarray, ) -> None: - torch_L1_in = torch.tensor(L1_in, device="cuda") - torch_L2_in = torch.tensor(L2_in, device="cuda") - torch_weights = torch.tensor(weights, device="cuda") + torch_L1_in = torch.tensor(L1_in, device=DEVICE_TYPE) + torch_L2_in = torch.tensor(L2_in, device=DEVICE_TYPE) + torch_weights = torch.tensor(weights, device=DEVICE_TYPE) torch_L3_out = self.e3nn_tp(torch_L1_in, torch_L2_in, torch_weights) @@ -80,13 +81,13 @@ def backward_cpu( weights: np.ndarray, weights_grad: np.ndarray, ) -> None: - torch_L1_in = torch.tensor(L1_in, requires_grad=True, device="cuda") - torch_L2_in = torch.tensor(L2_in, requires_grad=True, device="cuda") - torch_weights = torch.tensor(weights, requires_grad=True, device="cuda") + torch_L1_in = torch.tensor(L1_in, requires_grad=True, device=DEVICE_TYPE) + torch_L2_in = torch.tensor(L2_in, requires_grad=True, device=DEVICE_TYPE) + torch_weights = torch.tensor(weights, requires_grad=True, device=DEVICE_TYPE) torch_out = self.e3nn_tp(torch_L1_in, torch_L2_in, torch_weights) - torch_L3_grad_in = torch.tensor(L3_grad, device="cuda") + torch_L3_grad_in = torch.tensor(L3_grad, device=DEVICE_TYPE) torch_out.backward(gradient=torch_L3_grad_in) diff --git a/openequivariance/openequivariance/_torch/FlashTPConv.py b/openequivariance/openequivariance/_torch/FlashTPConv.py index 9ec5c409..f29cda71 100644 --- a/openequivariance/openequivariance/_torch/FlashTPConv.py +++ b/openequivariance/openequivariance/_torch/FlashTPConv.py @@ -6,6 +6,7 @@ import numpy as np from openequivariance.core.ConvolutionBase import ConvolutionBase from openequivariance.core.utils import oeq_to_torch_dtype +from openequivariance._torch.extlib import DEVICE_TYPE class FlashTPConv(ConvolutionBase): @@ -30,7 +31,7 @@ def __init__(self, config, *, idx_dtype=np.int64, torch_op=True): config.irreps_in2, config.irreps_out, instructions, - device="cuda", + device=DEVICE_TYPE, dtype=oeq_to_torch_dtype(config.irrep_dtype), ) diff --git a/openequivariance/openequivariance/_torch/NPDoubleBackwardMixin.py b/openequivariance/openequivariance/_torch/NPDoubleBackwardMixin.py index 5411377d..9dd77247 100644 --- a/openequivariance/openequivariance/_torch/NPDoubleBackwardMixin.py +++ b/openequivariance/openequivariance/_torch/NPDoubleBackwardMixin.py @@ -1,4 +1,5 @@ import torch +from openequivariance._torch.extlib import DEVICE_TYPE def _none_to_zeros(values, refs): @@ -19,14 +20,18 @@ def double_backward_cpu( ): assert self.torch_op - in1_torch = torch.tensor(in1, device="cuda", requires_grad=True) - in2_torch = torch.tensor(in2, device="cuda", requires_grad=True) - weights_torch = torch.tensor(weights, device="cuda", requires_grad=True) - out_grad_torch = torch.tensor(out_grad, device="cuda", requires_grad=True) - in1_dgrad_torch = torch.tensor(in1_dgrad, device="cuda", requires_grad=False) - in2_dgrad_torch = torch.tensor(in2_dgrad, device="cuda", requires_grad=False) + in1_torch = torch.tensor(in1, device=DEVICE_TYPE, requires_grad=True) + in2_torch = torch.tensor(in2, device=DEVICE_TYPE, requires_grad=True) + weights_torch = torch.tensor(weights, device=DEVICE_TYPE, requires_grad=True) + out_grad_torch = torch.tensor(out_grad, device=DEVICE_TYPE, requires_grad=True) + in1_dgrad_torch = torch.tensor( + in1_dgrad, device=DEVICE_TYPE, requires_grad=False + ) + in2_dgrad_torch = torch.tensor( + in2_dgrad, device=DEVICE_TYPE, requires_grad=False + ) weights_dgrad_torch = torch.tensor( - weights_dgrad, device="cuda", requires_grad=False + weights_dgrad, device=DEVICE_TYPE, requires_grad=False ) out_torch = self.forward(in1_torch, in2_torch, weights_torch) @@ -67,20 +72,30 @@ def triple_backward_cpu( ): assert self.torch_op - in1_torch = torch.tensor(in1, device="cuda", requires_grad=True) - in2_torch = torch.tensor(in2, device="cuda", requires_grad=True) - weights_torch = torch.tensor(weights, device="cuda", requires_grad=True) - out_grad_torch = torch.tensor(out_grad, device="cuda", requires_grad=True) - in1_dgrad_torch = torch.tensor(in1_dgrad, device="cuda", requires_grad=True) - in2_dgrad_torch = torch.tensor(in2_dgrad, device="cuda", requires_grad=True) + in1_torch = torch.tensor(in1, device=DEVICE_TYPE, requires_grad=True) + in2_torch = torch.tensor(in2, device=DEVICE_TYPE, requires_grad=True) + weights_torch = torch.tensor(weights, device=DEVICE_TYPE, requires_grad=True) + out_grad_torch = torch.tensor(out_grad, device=DEVICE_TYPE, requires_grad=True) + in1_dgrad_torch = torch.tensor( + in1_dgrad, device=DEVICE_TYPE, requires_grad=True + ) + in2_dgrad_torch = torch.tensor( + in2_dgrad, device=DEVICE_TYPE, requires_grad=True + ) weights_dgrad_torch = torch.tensor( - weights_dgrad, device="cuda", requires_grad=True + weights_dgrad, device=DEVICE_TYPE, requires_grad=True + ) + out_tgrad_torch = torch.tensor( + out_tgrad, device=DEVICE_TYPE, requires_grad=False + ) + in1_tgrad_torch = torch.tensor( + in1_tgrad, device=DEVICE_TYPE, requires_grad=False + ) + in2_tgrad_torch = torch.tensor( + in2_tgrad, device=DEVICE_TYPE, requires_grad=False ) - out_tgrad_torch = torch.tensor(out_tgrad, device="cuda", requires_grad=False) - in1_tgrad_torch = torch.tensor(in1_tgrad, device="cuda", requires_grad=False) - in2_tgrad_torch = torch.tensor(in2_tgrad, device="cuda", requires_grad=False) weights_tgrad_torch = torch.tensor( - weights_tgrad, device="cuda", requires_grad=False + weights_tgrad, device=DEVICE_TYPE, requires_grad=False ) out_torch = self.forward(in1_torch, in2_torch, weights_torch) @@ -147,19 +162,23 @@ def double_backward_cpu( ): assert self.torch_op - in1_torch = torch.tensor(in1, device="cuda", requires_grad=True) - in2_torch = torch.tensor(in2, device="cuda", requires_grad=True) - weights_torch = torch.tensor(weights, device="cuda", requires_grad=True) - out_grad_torch = torch.tensor(out_grad, device="cuda", requires_grad=True) - in1_dgrad_torch = torch.tensor(in1_dgrad, device="cuda", requires_grad=False) - in2_dgrad_torch = torch.tensor(in2_dgrad, device="cuda", requires_grad=False) + in1_torch = torch.tensor(in1, device=DEVICE_TYPE, requires_grad=True) + in2_torch = torch.tensor(in2, device=DEVICE_TYPE, requires_grad=True) + weights_torch = torch.tensor(weights, device=DEVICE_TYPE, requires_grad=True) + out_grad_torch = torch.tensor(out_grad, device=DEVICE_TYPE, requires_grad=True) + in1_dgrad_torch = torch.tensor( + in1_dgrad, device=DEVICE_TYPE, requires_grad=False + ) + in2_dgrad_torch = torch.tensor( + in2_dgrad, device=DEVICE_TYPE, requires_grad=False + ) weights_dgrad_torch = torch.tensor( - weights_dgrad, device="cuda", requires_grad=False + weights_dgrad, device=DEVICE_TYPE, requires_grad=False ) - torch_rows = torch.tensor(graph.rows, device="cuda") - torch_cols = torch.tensor(graph.cols, device="cuda") - torch_transpose_perm = torch.tensor(graph.transpose_perm, device="cuda") + torch_rows = torch.tensor(graph.rows, device=DEVICE_TYPE) + torch_cols = torch.tensor(graph.cols, device=DEVICE_TYPE) + torch_transpose_perm = torch.tensor(graph.transpose_perm, device=DEVICE_TYPE) out_torch = self.forward( in1_torch, @@ -208,25 +227,35 @@ def triple_backward_cpu( ): assert self.torch_op - in1_torch = torch.tensor(in1, device="cuda", requires_grad=True) - in2_torch = torch.tensor(in2, device="cuda", requires_grad=True) - weights_torch = torch.tensor(weights, device="cuda", requires_grad=True) - out_grad_torch = torch.tensor(out_grad, device="cuda", requires_grad=True) - in1_dgrad_torch = torch.tensor(in1_dgrad, device="cuda", requires_grad=True) - in2_dgrad_torch = torch.tensor(in2_dgrad, device="cuda", requires_grad=True) + in1_torch = torch.tensor(in1, device=DEVICE_TYPE, requires_grad=True) + in2_torch = torch.tensor(in2, device=DEVICE_TYPE, requires_grad=True) + weights_torch = torch.tensor(weights, device=DEVICE_TYPE, requires_grad=True) + out_grad_torch = torch.tensor(out_grad, device=DEVICE_TYPE, requires_grad=True) + in1_dgrad_torch = torch.tensor( + in1_dgrad, device=DEVICE_TYPE, requires_grad=True + ) + in2_dgrad_torch = torch.tensor( + in2_dgrad, device=DEVICE_TYPE, requires_grad=True + ) weights_dgrad_torch = torch.tensor( - weights_dgrad, device="cuda", requires_grad=True + weights_dgrad, device=DEVICE_TYPE, requires_grad=True + ) + out_tgrad_torch = torch.tensor( + out_tgrad, device=DEVICE_TYPE, requires_grad=False + ) + in1_tgrad_torch = torch.tensor( + in1_tgrad, device=DEVICE_TYPE, requires_grad=False + ) + in2_tgrad_torch = torch.tensor( + in2_tgrad, device=DEVICE_TYPE, requires_grad=False ) - out_tgrad_torch = torch.tensor(out_tgrad, device="cuda", requires_grad=False) - in1_tgrad_torch = torch.tensor(in1_tgrad, device="cuda", requires_grad=False) - in2_tgrad_torch = torch.tensor(in2_tgrad, device="cuda", requires_grad=False) weights_tgrad_torch = torch.tensor( - weights_tgrad, device="cuda", requires_grad=False + weights_tgrad, device=DEVICE_TYPE, requires_grad=False ) - torch_rows = torch.tensor(graph.rows, device="cuda") - torch_cols = torch.tensor(graph.cols, device="cuda") - torch_transpose_perm = torch.tensor(graph.transpose_perm, device="cuda") + torch_rows = torch.tensor(graph.rows, device=DEVICE_TYPE) + torch_cols = torch.tensor(graph.cols, device=DEVICE_TYPE) + torch_transpose_perm = torch.tensor(graph.transpose_perm, device=DEVICE_TYPE) out_torch = self.forward( in1_torch, diff --git a/openequivariance/openequivariance/_torch/TensorProduct.py b/openequivariance/openequivariance/_torch/TensorProduct.py index 00fe8c07..3a17a898 100644 --- a/openequivariance/openequivariance/_torch/TensorProduct.py +++ b/openequivariance/openequivariance/_torch/TensorProduct.py @@ -44,7 +44,7 @@ def _init_class(self): self, self.input_args["problem"], dp, - extlib.IS_HIP, + extlib.BACKEND, self.input_args["torch_op"], ) @@ -146,9 +146,9 @@ def forward_cpu( weights, not self.config.shared_weights ) - torch_L1_in = torch.tensor(L1_in, device="cuda") - torch_L2_in = torch.tensor(L2_in, device="cuda") - torch_weights = torch.tensor(weights_chunked, device="cuda") + torch_L1_in = torch.tensor(L1_in, device=extlib.DEVICE_TYPE) + torch_L2_in = torch.tensor(L2_in, device=extlib.DEVICE_TYPE) + torch_weights = torch.tensor(weights_chunked, device=extlib.DEVICE_TYPE) torch_L3_out = self.forward(torch_L1_in, torch_L2_in, torch_weights) L3_out[:] = torch_L3_out.numpy(force=True) @@ -160,13 +160,15 @@ def backward_cpu( weights, not self.config.shared_weights ) - torch_L1_in = torch.tensor(L1_in, requires_grad=True, device="cuda") - torch_L2_in = torch.tensor(L2_in, requires_grad=True, device="cuda") - torch_weights = torch.tensor(weights_chunked, requires_grad=True, device="cuda") + torch_L1_in = torch.tensor(L1_in, requires_grad=True, device=extlib.DEVICE_TYPE) + torch_L2_in = torch.tensor(L2_in, requires_grad=True, device=extlib.DEVICE_TYPE) + torch_weights = torch.tensor( + weights_chunked, requires_grad=True, device=extlib.DEVICE_TYPE + ) torch_out = self.forward(torch_L1_in, torch_L2_in, torch_weights) - torch_L3_grad_in = torch.tensor(L3_grad, device="cuda") + torch_L3_grad_in = torch.tensor(L3_grad, device=extlib.DEVICE_TYPE) torch_out.backward(gradient=torch_L3_grad_in) @@ -364,13 +366,13 @@ def register_autocast(): import torch torch.library.register_autocast( - "libtorch_tp_jit::jit_tp_forward", "cuda", torch.float32 + "libtorch_tp_jit::jit_tp_forward", extlib.DEVICE_TYPE, torch.float32 ) torch.library.register_autocast( - "libtorch_tp_jit::jit_tp_backward", "cuda", torch.float32 + "libtorch_tp_jit::jit_tp_backward", extlib.DEVICE_TYPE, torch.float32 ) torch.library.register_autocast( - "libtorch_tp_jit::jit_tp_double_backward", "cuda", torch.float32 + "libtorch_tp_jit::jit_tp_double_backward", extlib.DEVICE_TYPE, torch.float32 ) diff --git a/openequivariance/openequivariance/_torch/TensorProductConv.py b/openequivariance/openequivariance/_torch/TensorProductConv.py index d052f909..f6015fce 100644 --- a/openequivariance/openequivariance/_torch/TensorProductConv.py +++ b/openequivariance/openequivariance/_torch/TensorProductConv.py @@ -4,7 +4,8 @@ import torch from openequivariance._torch.extlib import ( - IS_HIP, + BACKEND, + DEVICE_TYPE, DeviceProp, BUILT_EXTENSION, ) @@ -78,7 +79,7 @@ def _init_class(self): self, self.input_args["problem"], dp, - IS_HIP, + BACKEND, idx_dtype=np.int64, torch_op=self.input_args["torch_op"], deterministic=self.input_args["deterministic"], @@ -87,7 +88,9 @@ def _init_class(self): self.allocate_workspace(self.workspace_size) - self.dummy_transpose_perm = torch.zeros(1, dtype=torch.int64, device="cuda") + self.dummy_transpose_perm = torch.zeros( + 1, dtype=torch.int64, device=DEVICE_TYPE + ) self.weight_numel = self.config.weight_numel self.kernel = string_to_tensor(self.kernel_string) self.L3_dim = self.kernel_prop["L3_dim"] @@ -187,7 +190,7 @@ def forward( def allocate_workspace(self, size_bytes): self.workspace_size = size_bytes self.workspace_buffer = torch.zeros( - size_bytes, dtype=torch.uint8, device="cuda" + size_bytes, dtype=torch.uint8, device=DEVICE_TYPE ) self.workspace_ptr = self.workspace_buffer.data_ptr() logger.info(f"Convolution requires {size_bytes // 1000000}MB of workspace.") @@ -214,14 +217,14 @@ def forward_cpu(self, L1_in, L2_in, weights, L3_out, graph): weights, not self.config.shared_weights ) - torch_L1_in = torch.tensor(L1_in, device="cuda") - torch_L2_in = torch.tensor(L2_in, device="cuda") - torch_weights = torch.tensor(weights_chunked, device="cuda") - torch_rows = torch.tensor(graph.rows, device="cuda") - torch_cols = torch.tensor(graph.cols, device="cuda") + torch_L1_in = torch.tensor(L1_in, device=DEVICE_TYPE) + torch_L2_in = torch.tensor(L2_in, device=DEVICE_TYPE) + torch_weights = torch.tensor(weights_chunked, device=DEVICE_TYPE) + torch_rows = torch.tensor(graph.rows, device=DEVICE_TYPE) + torch_cols = torch.tensor(graph.cols, device=DEVICE_TYPE) if self.deterministic: - torch_sender_perm = torch.tensor(graph.transpose_perm, device="cuda") + torch_sender_perm = torch.tensor(graph.transpose_perm, device=DEVICE_TYPE) else: torch_sender_perm = None @@ -245,15 +248,17 @@ def backward_cpu( weights, not self.config.shared_weights ) - torch_L1_in = torch.tensor(L1_in, requires_grad=True, device="cuda") - torch_L2_in = torch.tensor(L2_in, requires_grad=True, device="cuda") - torch_weights = torch.tensor(weights_chunked, requires_grad=True, device="cuda") - torch_L3_grad = torch.tensor(L3_grad, device="cuda") - torch_rows = torch.tensor(graph.rows, device="cuda") - torch_cols = torch.tensor(graph.cols, device="cuda") + torch_L1_in = torch.tensor(L1_in, requires_grad=True, device=DEVICE_TYPE) + torch_L2_in = torch.tensor(L2_in, requires_grad=True, device=DEVICE_TYPE) + torch_weights = torch.tensor( + weights_chunked, requires_grad=True, device=DEVICE_TYPE + ) + torch_L3_grad = torch.tensor(L3_grad, device=DEVICE_TYPE) + torch_rows = torch.tensor(graph.rows, device=DEVICE_TYPE) + torch_cols = torch.tensor(graph.cols, device=DEVICE_TYPE) if self.deterministic: - torch_sender_perm = torch.tensor(graph.transpose_perm, device="cuda") + torch_sender_perm = torch.tensor(graph.transpose_perm, device=DEVICE_TYPE) else: torch_sender_perm = None @@ -285,7 +290,9 @@ def register_torch_fakes(): def fake_forward( kernel, hash, L1_in, L2_in, W, L3_dim, rows, cols, workspace_buffer, sender_perm ): - return torch.empty(L1_in.shape[0], L3_dim, device="cuda", dtype=L1_in.dtype) + return torch.empty( + L1_in.shape[0], L3_dim, device=DEVICE_TYPE, dtype=L1_in.dtype + ) @torch.library.register_fake("libtorch_tp_jit::jit_conv_backward") def fake_backward( @@ -558,13 +565,13 @@ def triple_backward(ctx, t_L1_grad, t_L2_grad, t_W_grad, t_L3_dgrad): def register_autocast(): torch.library.register_autocast( - "libtorch_tp_jit::jit_conv_forward", "cuda", torch.float32 + "libtorch_tp_jit::jit_conv_forward", DEVICE_TYPE, torch.float32 ) torch.library.register_autocast( - "libtorch_tp_jit::jit_conv_backward", "cuda", torch.float32 + "libtorch_tp_jit::jit_conv_backward", DEVICE_TYPE, torch.float32 ) torch.library.register_autocast( - "libtorch_tp_jit::jit_conv_double_backward", "cuda", torch.float32 + "libtorch_tp_jit::jit_conv_double_backward", DEVICE_TYPE, torch.float32 ) diff --git a/openequivariance/openequivariance/_torch/extlib/__init__.py b/openequivariance/openequivariance/_torch/extlib/__init__.py index 2a114b3e..2fb244fb 100644 --- a/openequivariance/openequivariance/_torch/extlib/__init__.py +++ b/openequivariance/openequivariance/_torch/extlib/__init__.py @@ -22,11 +22,35 @@ extension_module = None -assert torch.version.cuda or torch.version.hip, ( - "Only CUDA and HIP backends are supported" + +def _detect_backend(): + """ + Determines which GPU backend this PyTorch build targets. + + Returns one of ``"cuda"``, ``"hip"`` or ``"sycl"``. HIP builds report a + ``torch.version.cuda`` of ``None``, so HIP must be tested first. + """ + if torch.version.hip: + return "hip" + if torch.version.cuda: + return "cuda" + if hasattr(torch, "xpu") and torch.xpu.is_available(): + return "sycl" + return None + + +BACKEND = _detect_backend() + +assert BACKEND is not None, ( + "Only the CUDA, HIP and XPU (SYCL) backends are supported. " + "No supported accelerator was detected in this PyTorch build." ) -IS_HIP = bool(torch.version.hip) +IS_HIP = BACKEND == "hip" +IS_SYCL = BACKEND == "sycl" + +# The torch device type that tensors passed to the kernels must live on. +DEVICE_TYPE = "xpu" if IS_SYCL else "cuda" @contextlib.contextmanager @@ -111,7 +135,7 @@ def load_jit_extension(): f"-l{python_lib_name}", ], ) - if torch.version.cuda: + if BACKEND == "cuda": extra_link_args.extend(["-lcuda", "-lcudart", "-lnvrtc", "-lcublas"]) try: @@ -124,16 +148,48 @@ def load_jit_extension(): getLogger().info(str(e)) extra_cflags.append("-DCUDA_BACKEND") - elif torch.version.hip: + elif BACKEND == "hip": torch_libs = library_paths("cuda")[0] extra_link_args.append("-Wl,-rpath," + torch_libs) extra_cflags.append("-DHIP_BACKEND") + elif BACKEND == "sycl": + # torch.utils.cpp_extension compiles with $CXX (default c++), + # which must be the oneAPI DPC++ driver for -fsycl to work. + import shutil + + cxx = os.environ.get("CXX", "") + if "icpx" not in os.path.basename(cxx): + if shutil.which("icpx") is None: + BUILT_EXTENSION_ERROR = ( + "The SYCL backend requires the oneAPI DPC++ compiler. " + "Put 'icpx' on your PATH or set CXX to it." + ) + return + os.environ["CXX"] = "icpx" + + # SYCL sources must be compiled and linked by the SYCL compiler + # driver; -fsycl is required on both the compile and link lines. + extra_cflags.extend(["-fsycl", "-DSYCL_BACKEND"]) + extra_link_args.extend( + ["-fsycl", "-ltorch_xpu", "-lc10_xpu", "-lmkl_sycl_blas"] + ) + + for lib_dir in library_paths("xpu"): + extra_link_args.append("-Wl,-rpath," + lib_dir) + extra_link_args.append("-L" + lib_dir) + + mkl_root = os.environ.get("MKLROOT") + if mkl_root: + mkl_lib = os.path.join(mkl_root, "lib") + extra_link_args.append("-L" + mkl_lib) + extra_link_args.append("-Wl,-rpath," + mkl_lib) + extra_include_dirs.append(os.path.join(mkl_root, "include")) torch_sources = [oeq_root + "/extension/" + src for src in torch_sources] include_dirs = ( [oeq_root + "/extension/" + d for d in include_dirs] + extra_include_dirs - + include_paths("cuda") + + include_paths("xpu" if BACKEND == "sycl" else "cuda") ) with warnings.catch_warnings(): @@ -165,10 +221,12 @@ def load_precompiled_extension(): True # Doesn't actually use libpython, just set this as true anyway ) try: - if torch.version.cuda: + if BACKEND == "cuda": import openequivariance._torch.extlib.oeq_stable_cuda as extension_module - elif torch.version.hip: + elif BACKEND == "hip": import openequivariance._torch.extlib.oeq_stable_hip as extension_module + elif BACKEND == "sycl": + import openequivariance._torch.extlib.oeq_stable_sycl as extension_module torch.ops.load_library(extension_module.__file__) BUILT_EXTENSION = True @@ -189,13 +247,17 @@ def load_precompiled_extension(): WARNING_MESSAGE += f"PyTorch version {torch.__version__} is < 2.10, minimum required for precompiled extension. Please upgrade to 2.10.\n" USE_PRECOMPILED_EXTENSION = False -if torch.version.hip: +if BACKEND == "hip": WARNING_MESSAGE += "HIP does not support precompiled extension yet.\n" USE_PRECOMPILED_EXTENSION = False -if not os.path.exists( - os.path.join(os.path.dirname(__file__), "liboeq_stable_cuda_aoti.so") -): +if BACKEND == "sycl": + WARNING_MESSAGE += "SYCL does not support precompiled extension yet.\n" + USE_PRECOMPILED_EXTENSION = False + +AOTI_SO_NAME = f"liboeq_stable_{BACKEND}_aoti.so" + +if not os.path.exists(os.path.join(os.path.dirname(__file__), AOTI_SO_NAME)): WARNING_MESSAGE += "Precompiled extension shared object not found.\n" USE_PRECOMPILED_EXTENSION = False @@ -213,10 +275,7 @@ def torch_ext_so_path(): return extension_module.__file__ else: dirname = os.path.dirname(extension_module.__file__) - if torch.version.cuda: - return os.path.join(dirname, "liboeq_stable_cuda_aoti.so") - elif torch.version.hip: - return os.path.join(dirname, "liboeq_stable_hip_aoti.so") + return os.path.join(dirname, AOTI_SO_NAME) sys.modules["oeq_utilities"] = extension_module diff --git a/openequivariance/openequivariance/benchmark/correctness.py b/openequivariance/openequivariance/benchmark/correctness.py index fa95f826..8c355df7 100644 --- a/openequivariance/openequivariance/benchmark/correctness.py +++ b/openequivariance/openequivariance/benchmark/correctness.py @@ -19,6 +19,7 @@ from openequivariance.core.e3nn_lite import TPProblem from openequivariance.core.TensorProductBase import TensorProductBase from openequivariance.core.utils import transpose_irrep_layout +from openequivariance._torch.extlib import DEVICE_TYPE logger = getLogger() @@ -562,7 +563,7 @@ def correctness_forward_conv( args["transpose_perm"] = graph.transpose_perm for key in args: - args[key] = torch.tensor(args[key], device="cuda") + args[key] = torch.tensor(args[key], device=DEVICE_TYPE) run_out[:] = tp.forward(**args).cpu().numpy() diff --git a/openequivariance/openequivariance/core/ComputationSchedule.py b/openequivariance/openequivariance/core/ComputationSchedule.py index f9f10013..304db645 100644 --- a/openequivariance/openequivariance/core/ComputationSchedule.py +++ b/openequivariance/openequivariance/core/ComputationSchedule.py @@ -62,7 +62,11 @@ def __init__(self, src_irreps, src_views, idxs): class CGTensor: def __init__(self, l1, l2, l3, normalization_factor, dtype): - suffix_map = {np.float32: "f", np.float64: "L"} + # A hex float literal is a double by default and represents the value + # exactly, so float64 needs no suffix. An "L" (long double) suffix + # would not change the value but is rejected by SPIR-V targets, which + # have no 128-bit float type. + suffix_map = {np.float32: "f", np.float64: ""} tensor = wigner_3j(l1, l2, l3) coord1, coord2, coord3 = [ diff --git a/openequivariance/openequivariance/core/ConvolutionBase.py b/openequivariance/openequivariance/core/ConvolutionBase.py index 116a21b3..4f5ab97d 100644 --- a/openequivariance/openequivariance/core/ConvolutionBase.py +++ b/openequivariance/openequivariance/core/ConvolutionBase.py @@ -6,7 +6,7 @@ get_random_buffers_forward_conv, ) from openequivariance.core.e3nn_lite import wigner_3j -from openequivariance.core.utils import benchmark +from openequivariance.core.utils import accelerator_device_type, benchmark logger = getLogger() @@ -142,14 +142,14 @@ def benchmark_forward( assert graph.rows.dtype == self.idx_dtype assert graph.cols.dtype == self.idx_dtype - torch_L1_in = torch.tensor(L1_in, device="cuda") - torch_L2_in = torch.tensor(L2_in, device="cuda") - torch_weights = torch.tensor(weights, device="cuda") + torch_L1_in = torch.tensor(L1_in, device=accelerator_device_type()) + torch_L2_in = torch.tensor(L2_in, device=accelerator_device_type()) + torch_weights = torch.tensor(weights, device=accelerator_device_type()) - torch_rows = torch.tensor(graph.rows, device="cuda") - torch_cols = torch.tensor(graph.cols, device="cuda") + torch_rows = torch.tensor(graph.rows, device=accelerator_device_type()) + torch_cols = torch.tensor(graph.cols, device=accelerator_device_type()) torch_transpose_perm = ( - torch.tensor(graph.transpose_perm, device="cuda") + torch.tensor(graph.transpose_perm, device=accelerator_device_type()) if self.deterministic else None ) @@ -200,19 +200,27 @@ def benchmark_backward( assert graph.rows.dtype == self.idx_dtype assert graph.cols.dtype == self.idx_dtype - torch_L1_in = torch.tensor(in1, device="cuda", requires_grad=True) - torch_L2_in = torch.tensor(in2, device="cuda", requires_grad=True) - torch_weights = torch.tensor(weights, device="cuda", requires_grad=True) + torch_L1_in = torch.tensor( + in1, device=accelerator_device_type(), requires_grad=True + ) + torch_L2_in = torch.tensor( + in2, device=accelerator_device_type(), requires_grad=True + ) + torch_weights = torch.tensor( + weights, device=accelerator_device_type(), requires_grad=True + ) - torch_rows = torch.tensor(graph.rows, device="cuda").detach() - torch_cols = torch.tensor(graph.cols, device="cuda").detach() - torch_transpose_perm = torch.tensor(graph.transpose_perm, device="cuda") + torch_rows = torch.tensor(graph.rows, device=accelerator_device_type()).detach() + torch_cols = torch.tensor(graph.cols, device=accelerator_device_type()).detach() + torch_transpose_perm = torch.tensor( + graph.transpose_perm, device=accelerator_device_type() + ) fwd_args = [torch_L1_in, torch_L2_in, torch_weights, torch_rows, torch_cols] if self.deterministic: fwd_args.append(torch_transpose_perm) torch_out = self.forward(*fwd_args) - torch_L3_grad = torch.tensor(out_grad, device="cuda") + torch_L3_grad = torch.tensor(out_grad, device=accelerator_device_type()) mode = "gpu_time" if self.torch_op else "torch_kernel_time" diff --git a/openequivariance/openequivariance/core/LoopUnrollConv.py b/openequivariance/openequivariance/core/LoopUnrollConv.py index 17869760..c2e86948 100644 --- a/openequivariance/openequivariance/core/LoopUnrollConv.py +++ b/openequivariance/openequivariance/core/LoopUnrollConv.py @@ -20,7 +20,7 @@ def __init__( self, config, dp, - is_hip, + backend, *, idx_dtype: type[np.generic] = np.int64, torch_op: bool = False, @@ -34,7 +34,7 @@ def __init__( if kahan: assert deterministic - env = get_jinja_environment(is_hip=is_hip) + env = get_jinja_environment(backend=backend, warp_size=dp.warpsize) template = env.get_template("loop_unroll_conv_atomic.cuh") analysis = filter_and_analyze_problem(config) diff --git a/openequivariance/openequivariance/core/LoopUnrollTP.py b/openequivariance/openequivariance/core/LoopUnrollTP.py index e2969041..1de89f1b 100644 --- a/openequivariance/openequivariance/core/LoopUnrollTP.py +++ b/openequivariance/openequivariance/core/LoopUnrollTP.py @@ -17,10 +17,10 @@ class LoopUnrollTP(TensorProductBase): - def __init__(self, config, dp, is_hip, torch_op): + def __init__(self, config, dp, backend, torch_op): super().__init__(config, torch_op=torch_op) - env = get_jinja_environment(is_hip=is_hip) + env = get_jinja_environment(backend=backend, warp_size=dp.warpsize) template = env.get_template("loop_unroll_batch.cuh") analysis = filter_and_analyze_problem(config) diff --git a/openequivariance/openequivariance/core/TensorProductBase.py b/openequivariance/openequivariance/core/TensorProductBase.py index c6fc83f8..38277f5f 100644 --- a/openequivariance/openequivariance/core/TensorProductBase.py +++ b/openequivariance/openequivariance/core/TensorProductBase.py @@ -2,7 +2,7 @@ from openequivariance.core.e3nn_lite import TPProblem from openequivariance.core.logging import getLogger -from openequivariance.core.utils import benchmark +from openequivariance.core.utils import accelerator_device_type, benchmark logger = getLogger() @@ -77,9 +77,11 @@ def benchmark_forward( with_torch_overhead: bool = True, kernel_names=["forward"], ) -> np.ndarray: - torch_L1_in = torch.tensor(L1_in).to(device="cuda").detach() - torch_L2_in = torch.tensor(L2_in).to(device="cuda").detach() - torch_weights = torch.tensor(weights).to(device="cuda").detach() + torch_L1_in = torch.tensor(L1_in).to(device=accelerator_device_type()).detach() + torch_L2_in = torch.tensor(L2_in).to(device=accelerator_device_type()).detach() + torch_weights = ( + torch.tensor(weights).to(device=accelerator_device_type()).detach() + ) mode = "gpu_time" if with_torch_overhead else "torch_kernel_time" return benchmark( @@ -101,11 +103,17 @@ def benchmark_backward( with_torch_overhead: bool = True, kernel_names=["backward"], ) -> np.ndarray: - torch_L1_in = torch.tensor(L1_in, requires_grad=True, device="cuda") - torch_L2_in = torch.tensor(L2_in, requires_grad=True, device="cuda") - torch_weights = torch.tensor(weights, requires_grad=True, device="cuda") + torch_L1_in = torch.tensor( + L1_in, requires_grad=True, device=accelerator_device_type() + ) + torch_L2_in = torch.tensor( + L2_in, requires_grad=True, device=accelerator_device_type() + ) + torch_weights = torch.tensor( + weights, requires_grad=True, device=accelerator_device_type() + ) torch_out = self.forward(torch_L1_in, torch_L2_in, torch_weights) - torch_L3_grad_in = torch.tensor(L3_buffer, device="cuda") + torch_L3_grad_in = torch.tensor(L3_buffer, device=accelerator_device_type()) mode = "gpu_time" if with_torch_overhead else "torch_kernel_time" @@ -134,13 +142,22 @@ def benchmark_double_backward( with_torch_overhead: bool = True, kernel_names=["double_backward_A", "double_backward_B"], ) -> np.ndarray: - torch_L1_in = torch.tensor(L1_in, requires_grad=True, device="cuda") - torch_L2_in = torch.tensor(L2_in, requires_grad=True, device="cuda") - torch_weights = torch.tensor(weights, requires_grad=True, device="cuda") + torch_L1_in = torch.tensor( + L1_in, requires_grad=True, device=accelerator_device_type() + ) + torch_L2_in = torch.tensor( + L2_in, requires_grad=True, device=accelerator_device_type() + ) + torch_weights = torch.tensor( + weights, requires_grad=True, device=accelerator_device_type() + ) torch_out = self(torch_L1_in, torch_L2_in, torch_weights) torch_out_grad = ( - torch_out.clone().detach().to(device="cuda").requires_grad_(True) + torch_out.clone() + .detach() + .to(device=accelerator_device_type()) + .requires_grad_(True) ) (torch_L1_grad, torch_L2_grad, torch_weights_grad) = torch.autograd.grad( @@ -156,12 +173,18 @@ def benchmark_double_backward( + torch.norm(torch_L2_grad) + torch.norm(torch_weights_grad) ) - dummy_grad = torch.tensor(float(dummy), device="cuda", requires_grad=True) + dummy_grad = torch.tensor( + float(dummy), device=accelerator_device_type(), requires_grad=True + ) - torch_L1_grad = torch.tensor(L1_in, requires_grad=True, device="cuda") - torch_L2_grad = torch.tensor(L2_in, requires_grad=True, device="cuda") + torch_L1_grad = torch.tensor( + L1_in, requires_grad=True, device=accelerator_device_type() + ) + torch_L2_grad = torch.tensor( + L2_in, requires_grad=True, device=accelerator_device_type() + ) torch_weights_grad = torch.tensor( - weights_grad, requires_grad=True, device="cuda" + weights_grad, requires_grad=True, device=accelerator_device_type() ) mode = "gpu_time" if with_torch_overhead else "torch_kernel_time" diff --git a/openequivariance/openequivariance/core/utils.py b/openequivariance/openequivariance/core/utils.py index 53638422..3afb5b31 100644 --- a/openequivariance/openequivariance/core/utils.py +++ b/openequivariance/openequivariance/core/utils.py @@ -173,13 +173,19 @@ def benchmark(func, num_warmup, num_iter, mode="gpu_time", kernel_names=[]): else: from torch.profiler import ProfilerActivity, profile, record_function + # The profiler activity is per-accelerator: XPU kernels are not + # recorded under the CUDA activity. + accelerator_activity = ( + ProfilerActivity.XPU + if accelerator_device_type() == "xpu" + else ProfilerActivity.CUDA + ) + trace_file = tempfile.NamedTemporaryFile().name for i in range(num_iter): timer.clear_L2_cache() - with profile( - activities=[ProfilerActivity.CUDA], record_shapes=True - ) as prof: + with profile(activities=[accelerator_activity], record_shapes=True) as prof: with record_function("profile"): func() @@ -258,3 +264,16 @@ def transpose_irrep_layout( ) return out + + +def accelerator_device_type(): + """ + Returns the ``torch`` device type the kernels run on: ``"xpu"`` for the + SYCL backend, ``"cuda"`` for CUDA and HIP (PyTorch exposes HIP tensors + under the ``cuda`` device type). + + Imported lazily so that the backend-agnostic core does not pull in torch. + """ + from openequivariance._torch.extlib import DEVICE_TYPE + + return DEVICE_TYPE diff --git a/openequivariance/openequivariance/extension/backend/backend_cuda.hpp b/openequivariance/openequivariance/extension/backend/backend_cuda.hpp index 4ecef72a..9b427412 100644 --- a/openequivariance/openequivariance/extension/backend/backend_cuda.hpp +++ b/openequivariance/openequivariance/extension/backend/backend_cuda.hpp @@ -323,7 +323,10 @@ class __attribute__((visibility("default"))) CUJITKernel { } } - void execute(int kernel_id, void* args[], KernelLaunchConfig config) { + void execute(int kernel_id, void* args[], const size_t arg_sizes[], + size_t num_args, KernelLaunchConfig config) { + (void) arg_sizes; // The CUDA driver infers argument sizes from the kernel signature. + (void) num_args; if(kernel_id >= kernels.size()) throw std::logic_error("Kernel index out of range!"); diff --git a/openequivariance/openequivariance/extension/backend/backend_hip.hpp b/openequivariance/openequivariance/extension/backend/backend_hip.hpp index 2eb4ed85..4066daa0 100644 --- a/openequivariance/openequivariance/extension/backend/backend_hip.hpp +++ b/openequivariance/openequivariance/extension/backend/backend_hip.hpp @@ -292,7 +292,10 @@ class __attribute__((visibility("default"))) HIPJITKernel { // Ignore for AMD GPUs } - void execute(int kernel_id, void* args[], KernelLaunchConfig config) { + void execute(int kernel_id, void* args[], const size_t arg_sizes[], + size_t num_args, KernelLaunchConfig config) { + (void) arg_sizes; // The HIP driver infers argument sizes from the kernel signature. + (void) num_args; int device_id; HIP_ERRCHK(hipGetDevice(&device_id)); if(device_id != kernels->device) { kernels.reset(); diff --git a/openequivariance/openequivariance/extension/backend/backend_sycl.hpp b/openequivariance/openequivariance/extension/backend/backend_sycl.hpp new file mode 100644 index 00000000..b1dd2fc8 --- /dev/null +++ b/openequivariance/openequivariance/extension/backend/backend_sycl.hpp @@ -0,0 +1,337 @@ +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include + +using namespace std; +namespace syclex = sycl::ext::oneapi::experimental; + +/* +* SYCL streams are queues. Unlike CUDA / HIP, a sycl::queue is a +* reference-counted handle rather than an opaque pointer, so the "stream" type +* is a pointer to the queue owned by the caller (PyTorch). +*/ +using Stream = sycl::queue *; + +// Defined by the translation unit that binds this header to the framework +// (PyTorch / JAX). Returns the queue of the framework's current stream. +Stream get_current_stream(); + +// Returns the queue the kernels should be submitted to. A null stream means +// "no queue was supplied", in which case we fall back to the framework's +// current stream, and only then to a process-wide default queue. +inline sycl::queue &resolve_queue(Stream stream) { + if (stream != nullptr) { + return *stream; + } + if (Stream current = get_current_stream()) { + return *current; + } + static sycl::queue default_queue{sycl::gpu_selector_v}; + return default_queue; +} + +class SYCL_Allocator { +public: + static void* gpu_alloc (size_t size) { + sycl::queue &q = resolve_queue(nullptr); + void* ptr = sycl::malloc_device(size, q); + if (ptr == nullptr) { + throw std::runtime_error("SYCL device allocation failed!"); + } + return ptr; + } + + static void gpu_free (void* ptr) { + sycl::queue &q = resolve_queue(nullptr); + sycl::free(ptr, q); + } + + static void copy_host_to_device (void* host, void* device, size_t size) { + sycl::queue &q = resolve_queue(nullptr); + q.memcpy(device, host, size).wait(); + } + + static void copy_device_to_host (void* host, void* device, size_t size) { + sycl::queue &q = resolve_queue(nullptr); + q.memcpy(host, device, size).wait(); + } +}; + +/* +* SYCL has no direct equivalent of cudaEvent elapsed time that works without +* enabling profiling on the queue, so the timer brackets a wall-clock interval +* around a queue synchronization. +*/ +class GPUTimer { + std::chrono::time_point start_time; + +public: + GPUTimer() = default; + + void start() { + sycl::queue &q = resolve_queue(nullptr); + q.wait(); + start_time = std::chrono::steady_clock::now(); + } + + float stop_clock_get_elapsed() { + sycl::queue &q = resolve_queue(nullptr); + q.wait(); + auto stop_time = std::chrono::steady_clock::now(); + std::chrono::duration elapsed = stop_time - start_time; + return elapsed.count(); + } + + void clear_L2_cache() { + size_t element_count = 25000000; + sycl::queue &q = resolve_queue(nullptr); + + int* ptr = (int*) sycl::malloc_device(element_count * sizeof(int), q); + q.memset(ptr, 42, element_count * sizeof(int)).wait(); + sycl::free(ptr, q); + q.wait(); + } + + ~GPUTimer() = default; +}; + +class __attribute__((visibility("default"))) DeviceProp { +public: + std::string name; + int warpsize; + int major, minor; + int multiprocessorCount; + int maxSharedMemPerBlock; + int maxSharedMemoryPerMultiprocessor; + + DeviceProp(int device_id) { + auto devices = sycl::device::get_devices(sycl::info::device_type::gpu); + if (devices.empty()) { + throw std::runtime_error("No SYCL GPU devices found!"); + } + if (device_id < 0 || static_cast(device_id) >= devices.size()) { + device_id = 0; + } + sycl::device dev = devices[device_id]; + + name = dev.get_info(); + multiprocessorCount = + static_cast(dev.get_info()); + + // A SYCL sub-group is the analogue of a CUDA warp / HIP wavefront. Pick + // the largest supported size that the kernel generator can target. + auto sg_sizes = dev.get_info(); + warpsize = 32; + if (!sg_sizes.empty()) { + if (std::find(sg_sizes.begin(), sg_sizes.end(), size_t(32)) != sg_sizes.end()) { + warpsize = 32; + } else { + warpsize = static_cast( + *std::max_element(sg_sizes.begin(), sg_sizes.end())); + } + } + + maxSharedMemPerBlock = + static_cast(dev.get_info()); + maxSharedMemoryPerMultiprocessor = maxSharedMemPerBlock; + + // SYCL exposes no compute-capability equivalent. These fields exist + // only for parity with the CUDA backend and are unused on SYCL. + major = 0; + minor = 0; + } +}; + +class __attribute__((visibility("default"))) KernelLaunchConfig { +public: + uint32_t num_blocks = 0; + uint32_t num_threads = 0; + uint32_t warp_size = 32; + uint32_t smem = 0; + Stream hStream = nullptr; + + KernelLaunchConfig() = default; + ~KernelLaunchConfig() = default; + + KernelLaunchConfig(uint32_t num_blocks, uint32_t num_threads_per_block, uint32_t smem) : + num_blocks(num_blocks), + num_threads(num_threads_per_block), + smem(smem) + { } + + KernelLaunchConfig(int64_t num_blocks_i, int64_t num_threads_i, int64_t smem_i) : + KernelLaunchConfig( static_cast(num_blocks_i), + static_cast(num_threads_i), + static_cast(smem_i)) + { } +}; + +/* +* Runtime compilation uses the SYCL kernel_compiler extension with +* source_language::sycl, documented at +* https://github.com/intel/llvm/blob/sycl/sycl/doc/extensions/experimental/sycl_ext_oneapi_kernel_compiler_sycl.asciidoc +* +* The generated kernels are free functions marked with nd_range_kernel, so they +* are launched with raw (untyped) arguments exactly like cuLaunchKernel takes a +* void* array. +*/ +class __attribute__((visibility("default"))) SYCLJITKernel { +private: + bool compiled = false; + + vector kernel_names; + vector kernels; + std::unique_ptr> bundle; + +public: + string kernel_plaintext; + + SYCLJITKernel(string plaintext) : + kernel_plaintext(plaintext) { } + + void compile(string kernel_name, const vector template_params, int opt_level=3) { + vector kernel_names_i = {kernel_name}; + vector> template_param_list = {template_params}; + compile(kernel_names_i, template_param_list, opt_level); + } + + void compile(vector kernel_names_i, vector> template_param_list, int opt_level=3) { + if(compiled) { + throw std::logic_error("JIT object has already been compiled!"); + } + + if(kernel_names_i.size() != template_param_list.size()) { + throw std::logic_error("Kernel names and template parameters must have the same size!"); + } + + for(unsigned int kernel = 0; kernel < kernel_names_i.size(); kernel++) { + string kernel_name = kernel_names_i[kernel]; + vector &template_params = template_param_list[kernel]; + + // Step 1: Generate kernel names from the template parameters + if(template_params.size() == 0) { + kernel_names.push_back(kernel_name); + } + else { + std::string result = kernel_name + "<"; + for(unsigned int i = 0; i < template_params.size(); i++) { + result += std::to_string(template_params[i]); + if(i != template_params.size() - 1) { + result += ","; + } + } + result += ">"; + kernel_names.push_back(result); + } + } + + // Build against the context the kernels will actually run in, so the + // resulting bundle is valid for every device that context spans. + sycl::queue &q = resolve_queue(nullptr); + sycl::context build_context = q.get_context(); + + if(!q.get_device().ext_oneapi_can_build(syclex::source_language::sycl)) { + throw std::runtime_error( + "The SYCL device does not support runtime compilation of SYCL source."); + } + + std::string opt_arg = "-O" + std::to_string(opt_level); + std::vector build_opts = {opt_arg, "-ffast-math"}; + + std::string log; + try { + auto source_bundle = syclex::create_kernel_bundle_from_source( + build_context, + syclex::source_language::sycl, + kernel_plaintext); + + auto exe_bundle = syclex::build( + source_bundle, + syclex::properties{ + syclex::build_options{build_opts}, + syclex::save_log{&log}}); + + bundle = std::make_unique< + sycl::kernel_bundle>( + std::move(exe_bundle)); + } catch (const sycl::exception &e) { + throw std::logic_error("SYCL runtime compilation failed: " + + std::string(e.what()) + "\nlog: " + log); + } + + compiled = true; + + for (size_t i = 0; i < kernel_names.size(); i++) { + kernels.push_back(bundle->ext_oneapi_get_kernel(kernel_names[i])); + } + } + + void set_max_smem(int kernel_id, uint32_t max_smem_bytes) { + // Shared (local) memory is declared statically inside the generated + // kernel via work_group_static, so there is no opt-in to perform here. + // Validate the request against the device limit so an oversubscription + // fails with a clear message instead of at launch. + if(!compiled) + throw std::logic_error("JIT object has not been compiled!"); + if(static_cast(kernel_id) >= kernels.size()) + throw std::logic_error("Kernel index out of range!"); + + sycl::queue &q = resolve_queue(nullptr); + size_t local_mem = q.get_device().get_info(); + if(static_cast(max_smem_bytes) > local_mem) { + throw std::runtime_error("Requested shared memory (" + + std::to_string(max_smem_bytes) + + " bytes) exceeds the device local memory size (" + + std::to_string(local_mem) + " bytes)."); + } + } + + void execute(int kernel_id, void* args[], const size_t arg_sizes[], + size_t num_args, KernelLaunchConfig config) { + if(!compiled) + throw std::logic_error("JIT object has not been compiled!"); + if(static_cast(kernel_id) >= kernels.size()) + throw std::logic_error("Kernel index out of range!"); + + sycl::queue &q = resolve_queue(config.hStream); + + std::vector raw_args; + raw_args.reserve(num_args); + for (size_t i = 0; i < num_args; i++) { + raw_args.emplace_back(args[i], arg_sizes[i]); + } + + sycl::nd_range<1> range{ + sycl::range<1>(static_cast(config.num_blocks) * + static_cast(config.num_threads)), + sycl::range<1>(static_cast(config.num_threads))}; + + sycl::kernel &k = kernels[kernel_id]; + + syclex::submit(q, [&](sycl::handler &cgh) { + for (size_t i = 0; i < raw_args.size(); i++) { + cgh.set_arg(static_cast(i), raw_args[i]); + } + cgh.parallel_for(range, k); + }); + } + + ~SYCLJITKernel() = default; +}; + +inline KernelLaunchConfig with_stream(const KernelLaunchConfig& config, Stream stream) { + KernelLaunchConfig new_config = config; + new_config.hStream = stream; + return new_config; +} diff --git a/openequivariance/openequivariance/extension/convolution.hpp b/openequivariance/openequivariance/extension/convolution.hpp index 83ad58b4..57211dec 100644 --- a/openequivariance/openequivariance/extension/convolution.hpp +++ b/openequivariance/openequivariance/extension/convolution.hpp @@ -4,6 +4,8 @@ #include #include +#include "kernel_args.hpp" + struct ConvData { void* rows; void* cols; @@ -89,11 +91,12 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { ConvData conv_data = {rows, cols, nnz, node_count}; - void *args[] = {&L1_in, &L2_in, &weights, &L3_out, &conv_data, &workspace}; - jit.execute(0, args, with_stream(forward_config_ref, stream)); + auto args = make_kernel_args(L1_in, L2_in, weights, L3_out, conv_data, workspace); + jit.execute(0, args.data(), args.arg_sizes(), args.count(), + with_stream(forward_config_ref, stream)); if(reinterpret_cast(workspace) != 0) { - void *fixup_args[] = {&workspace, &L3_out}; + auto fixup_args = make_kernel_args(workspace, L3_out); KernelLaunchConfig fixup_config( forward_config_ref.num_blocks, @@ -102,7 +105,8 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { ); fixup_config.hStream = stream; - jit.execute(2, fixup_args, fixup_config); + jit.execute(2, fixup_args.data(), fixup_args.arg_sizes(), + fixup_args.count(), fixup_config); } } @@ -118,11 +122,14 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { Stream stream) { ConvData conv_data = {rows, cols, nnz, node_count}; - void *args[] = {&L1_in, &L1_grad, &L2_in, &L2_grad, &weight, &weight_grad, &L3_grad, &conv_data, &workspace, &transpose_perm}; - jit.execute(1, args, with_stream(backward_config_ref, stream)); + auto args = make_kernel_args(L1_in, L1_grad, L2_in, L2_grad, weight, + weight_grad, L3_grad, conv_data, workspace, + transpose_perm); + jit.execute(1, args.data(), args.arg_sizes(), args.count(), + with_stream(backward_config_ref, stream)); if(reinterpret_cast(workspace) != 0) { - void *fixup_args[] = {&workspace, &L1_grad}; + auto fixup_args = make_kernel_args(workspace, L1_grad); KernelLaunchConfig fixup_config( backward_config_ref.num_blocks, @@ -131,7 +138,8 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { ); fixup_config.hStream = stream; - jit.execute(3, fixup_args, fixup_config); + jit.execute(3, fixup_args.data(), fixup_args.arg_sizes(), + fixup_args.count(), fixup_config); } } @@ -145,33 +153,36 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { Stream stream) { ConvData conv_data = {rows, cols, nnz, node_count}; - void* args[] = { - &L1_in, &L2_in, &W, &L3_grad, &L1_dgrad, &L2_dgrad, &w_dgrad, - &L1_grad, &L2_grad, &W_grad, &L3_dgrad, &conv_data, &wspace, &transpose_perm - }; + auto args = make_kernel_args( + L1_in, L2_in, W, L3_grad, L1_dgrad, L2_dgrad, w_dgrad, + L1_grad, L2_grad, W_grad, L3_dgrad, conv_data, wspace, transpose_perm); - jit.execute(4, args, with_stream(forward_config_ref, stream)); + jit.execute(4, args.data(), args.arg_sizes(), args.count(), + with_stream(forward_config_ref, stream)); if(reinterpret_cast(wspace) != 0) { - void *fixup_args[] = {&wspace, &L3_dgrad}; + auto fixup_args = make_kernel_args(wspace, L3_dgrad); KernelLaunchConfig fixup_config( forward_config_ref.num_blocks, forward_config_ref.num_threads, 0 ); fixup_config.hStream = stream; - jit.execute(2, fixup_args, fixup_config); + jit.execute(2, fixup_args.data(), fixup_args.arg_sizes(), + fixup_args.count(), fixup_config); } - jit.execute(5, args, with_stream(double_backward_config_ref, stream)); + jit.execute(5, args.data(), args.arg_sizes(), args.count(), + with_stream(double_backward_config_ref, stream)); if(reinterpret_cast(wspace) != 0) { - void *fixup_args[] = {&wspace, &L1_grad}; + auto fixup_args = make_kernel_args(wspace, L1_grad); KernelLaunchConfig fixup_config( double_backward_config_ref.num_blocks, double_backward_config_ref.num_threads, 0 ); fixup_config.hStream = stream; - jit.execute(6, fixup_args, fixup_config); + jit.execute(6, fixup_args.data(), fixup_args.arg_sizes(), + fixup_args.count(), fixup_config); } } diff --git a/openequivariance/openequivariance/extension/group_mm.hpp b/openequivariance/openequivariance/extension/group_mm.hpp index 19249c39..e2d0196a 100644 --- a/openequivariance/openequivariance/extension/group_mm.hpp +++ b/openequivariance/openequivariance/extension/group_mm.hpp @@ -3,6 +3,7 @@ #include #include #include +#include #ifdef CUDA_BACKEND #include "cublas_v2.h" @@ -28,6 +29,35 @@ } ~BlasHandle() { rocblas_destroy_handle(handle); } }; +#elif defined(SYCL_BACKEND) + #include + #include + + // oneMKL takes the queue per call rather than a persistent handle, so this + // exists only to keep a single shape across the three backends. + struct BlasHandle { + BlasHandle() = default; + ~BlasHandle() = default; + }; + + // A small shared-USM array. oneMKL's pointer-array GEMM reads the pointer + // lists from the device, so they cannot live on the host stack. + template + class UsmArray { + sycl::queue &q_; + T *ptr_; + public: + UsmArray(sycl::queue &q, size_t n) : q_(q), + ptr_(sycl::malloc_shared(n, q)) { + if (ptr_ == nullptr) + throw std::runtime_error("Shared USM allocation failed!"); + } + ~UsmArray() { sycl::free(ptr_, q_); } + UsmArray(const UsmArray &) = delete; + UsmArray &operator=(const UsmArray &) = delete; + T &operator[](size_t i) { return ptr_[i]; } + T *get() { return ptr_; } + }; #endif inline BlasHandle& get_blas_handle() { @@ -53,6 +83,8 @@ void group_gemm_blas(void* A_raw, void* B_raw, void* C_raw, cublasOperation_t transa, transb; #elif defined(HIP_BACKEND) rocblas_operation transa, transb; +#elif defined(SYCL_BACKEND) + oneapi::mkl::transpose transa, transb; #endif if (ragged_inner == 0) { @@ -67,6 +99,9 @@ void group_gemm_blas(void* A_raw, void* B_raw, void* C_raw, transa = CUBLAS_OP_T; transb = CUBLAS_OP_N; #elif defined(HIP_BACKEND) transa = rocblas_operation_transpose; transb = rocblas_operation_none; +#elif defined(SYCL_BACKEND) + transa = oneapi::mkl::transpose::trans; + transb = oneapi::mkl::transpose::nontrans; #endif } else { M = k; K = static_cast(ragged_counts[i]); N = m; @@ -80,6 +115,9 @@ void group_gemm_blas(void* A_raw, void* B_raw, void* C_raw, transa = CUBLAS_OP_N; transb = CUBLAS_OP_T; #elif defined(HIP_BACKEND) transa = rocblas_operation_none; transb = rocblas_operation_transpose; +#elif defined(SYCL_BACKEND) + transa = oneapi::mkl::transpose::nontrans; + transb = oneapi::mkl::transpose::trans; #endif } ragged_offset += ragged_counts[i]; @@ -135,6 +173,43 @@ void group_gemm_blas(void* A_raw, void* B_raw, void* C_raw, } if (stat != rocblas_status_success) throw std::logic_error("Grouped GEMM failed!"); +#elif defined(SYCL_BACKEND) + (void) blas; + // Submit onto the same queue the kernels use so the GEMM is + // ordered against them. + sycl::queue &q = resolve_queue(get_current_stream()); + + // oneMKL's strided batch API requires stride_c >= ldc * n, which + // this interleaved layout deliberately violates (consecutive + // matrices overlap in the batch dimension). The pointer-array form + // imposes no such constraint, so build explicit pointer lists. + // oneMKL reads these arrays on the device, so they live in shared + // USM rather than on the host stack. + UsmArray a_ptrs(q, batch_size), b_ptrs(q, batch_size); + UsmArray c_ptrs(q, batch_size); + for (int j = 0; j < batch_size; j++) { + a_ptrs[j] = A + static_cast(strideA) * j; + b_ptrs[j] = B + static_cast(strideB) * j; + c_ptrs[j] = C + static_cast(strideC) * j; + } + + int64_t m64 = M, n64 = N, k64 = K; + int64_t lda64 = lda, ldb64 = ldb, ldc64 = ldc; + int64_t group_size = batch_size; + + try { + oneapi::mkl::blas::column_major::gemm_batch( + q, &transa, &transb, &m64, &n64, &k64, &alpha, + a_ptrs.get(), &lda64, + b_ptrs.get(), &ldb64, + &beta, c_ptrs.get(), &ldc64, + 1, &group_size) + // The pointer arrays are freed when this scope exits, so + // the submission has to complete before then. + .wait(); + } catch (const sycl::exception &e) { + throw std::logic_error("Grouped GEMM failed: " + std::string(e.what())); + } #endif } } diff --git a/openequivariance/openequivariance/extension/kernel_args.hpp b/openequivariance/openequivariance/extension/kernel_args.hpp new file mode 100644 index 00000000..d1ad64af --- /dev/null +++ b/openequivariance/openequivariance/extension/kernel_args.hpp @@ -0,0 +1,32 @@ +#pragma once + +#include +#include + +/* +* Packs kernel arguments as a list of (pointer, size) pairs. +* +* The CUDA and HIP driver APIs only need the pointer array, but SYCL launches +* free-function kernels with raw (untyped) arguments and therefore also needs +* the size of each argument. Collecting both here keeps a single call shape in +* the backend-independent code. +* +* As with the raw `void*[]` form this replaces, the caller must keep the +* referenced objects alive until the launch has been enqueued. +*/ +template +struct KernelArgs { + std::array ptrs; + std::array sizes; + + void **data() { return ptrs.data(); } + const size_t *arg_sizes() const { return sizes.data(); } + static constexpr size_t count() { return N; } +}; + +template +inline KernelArgs make_kernel_args(Ts &...args) { + return KernelArgs{ + {static_cast(&args)...}, + {sizeof(Ts)...}}; +} diff --git a/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp b/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp index ddabd0bb..560068ab 100644 --- a/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp +++ b/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp @@ -9,6 +9,10 @@ #include #endif +#ifdef SYCL_BACKEND + #include +#endif + #include #include #include @@ -81,6 +85,18 @@ Stream get_current_stream() { #ifdef HIP_BACKEND return c10::hip::getCurrentHIPStream(); #endif +#ifdef SYCL_BACKEND + // The queue is owned by PyTorch and outlives the kernel launch. + return &c10::xpu::getCurrentXPUStream().queue(); +#endif +} + +bool tensor_is_on_gpu(const Tensor &tensor) { +#ifdef SYCL_BACKEND + return tensor.is_xpu(); +#else + return tensor.is_cuda(); +#endif } namespace py=pybind11; diff --git a/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp b/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp index 6bf3d51f..6372710c 100644 --- a/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp +++ b/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp @@ -1,4 +1,8 @@ -#define USE_CUDA +#ifdef SYCL_BACKEND + #define USE_XPU +#else + #define USE_CUDA +#endif #include #include @@ -11,6 +15,9 @@ #include #include #include +#ifdef SYCL_BACKEND + #include +#endif using Tensor = torch::stable::Tensor; @@ -67,14 +74,28 @@ void *data_ptr(const Tensor &tensor) { } Stream get_current_stream() { - auto device_idx = torch::stable::accelerator::getCurrentDeviceIndex(); void* stream_ptr = nullptr; - TORCH_ERROR_CODE_CHECK(aoti_torch_get_current_cuda_stream(device_idx, &stream_ptr)); - #ifdef CUDA_BACKEND - return static_cast(stream_ptr); - #elif defined(HIP_BACKEND) - return static_cast(stream_ptr); + #ifdef SYCL_BACKEND + // Returns the sycl::queue* backing the current XPU stream. + TORCH_ERROR_CODE_CHECK(aoti_torch_get_current_sycl_queue(&stream_ptr)); + #else + auto device_idx = torch::stable::accelerator::getCurrentDeviceIndex(); + TORCH_ERROR_CODE_CHECK(aoti_torch_get_current_cuda_stream(device_idx, &stream_ptr)); + #endif + + return static_cast(stream_ptr); +} + +bool tensor_is_on_gpu(const Tensor &tensor) { + #ifdef SYCL_BACKEND + // The stable Tensor has no is_xpu(), so compare the device type directly. + int32_t device_type; + TORCH_ERROR_CODE_CHECK( + aoti_torch_get_device_type(tensor.get(), &device_type)); + return device_type == aoti_torch_device_type_xpu(); + #else + return tensor.is_cuda(); #endif } @@ -83,6 +104,9 @@ Stream get_current_stream() { #endif #ifdef HIP_BACKEND #define EXTENSION_NAME oeq_stable_hip +#endif +#ifdef SYCL_BACKEND + #define EXTENSION_NAME oeq_stable_sycl #endif #ifdef INCLUDE_NB_EXTENSION diff --git a/openequivariance/openequivariance/extension/stubs/stream.cpp b/openequivariance/openequivariance/extension/stubs/stream.cpp index fd011c35..310944c7 100644 --- a/openequivariance/openequivariance/extension/stubs/stream.cpp +++ b/openequivariance/openequivariance/extension/stubs/stream.cpp @@ -1,8 +1,20 @@ #include #include +/* +* Weak stand-ins for the stream accessors the AOTI shim declares. The real +* symbols come from libtorch_cuda / libtorch_xpu at load time; these exist so +* the extension links without a hard dependency on the accelerator runtime. +*/ extern "C" { +#ifdef SYCL_BACKEND + AOTITorchError aoti_torch_get_current_sycl_queue(void** ret_queue) { + *ret_queue = nullptr; + return 0; + } +#else AOTITorchError aoti_torch_get_current_cuda_stream(int32_t device_index, void** ret_stream) { return 0; } -} \ No newline at end of file +#endif +} diff --git a/openequivariance/openequivariance/extension/tensorproducts.hpp b/openequivariance/openequivariance/extension/tensorproducts.hpp index b4b4d84b..ff85722a 100644 --- a/openequivariance/openequivariance/extension/tensorproducts.hpp +++ b/openequivariance/openequivariance/extension/tensorproducts.hpp @@ -6,6 +6,8 @@ #include #include +#include "kernel_args.hpp" + template class __attribute__ ((visibility ("default"))) JITTPImpl { public: @@ -80,8 +82,9 @@ class __attribute__ ((visibility ("default"))) JITTPImpl { void* weights, Stream stream) { - void *args[] = { &num_products, &L1_in, &L2_in, &L3_out, &weights}; - jit.execute(0, args, with_stream(forward_config_ref, stream)); + auto args = make_kernel_args(num_products, L1_in, L2_in, L3_out, weights); + jit.execute(0, args.data(), args.arg_sizes(), args.count(), + with_stream(forward_config_ref, stream)); } void backward( @@ -90,8 +93,10 @@ class __attribute__ ((visibility ("default"))) JITTPImpl { void* L2_in, void* L2_grad, void* weight, void* weight_grad, void* L3_grad, Stream stream) { - void *args[] = { &num_products, &L1_in, &L1_grad, &L2_in, &L2_grad, &weight, &weight_grad, &L3_grad}; - jit.execute(1, args, with_stream(backward_config_ref, stream)); + auto args = make_kernel_args(num_products, L1_in, L1_grad, L2_in, L2_grad, + weight, weight_grad, L3_grad); + jit.execute(1, args.data(), args.arg_sizes(), args.count(), + with_stream(backward_config_ref, stream)); } void double_backward( @@ -100,13 +105,14 @@ class __attribute__ ((visibility ("default"))) JITTPImpl { void* L1_dgrad, void* L2_dgrad, void* w_dgrad, // Gradients w.r.t outputs of backward op void* L1_grad, void* L2_grad, void* W_grad, void* L3_dgrad, Stream stream) { - void* args[] = { - &num_products, &L1_in, &L2_in, &W, &L3_grad, &L1_dgrad, &L2_dgrad, &w_dgrad, - &L1_grad, &L2_grad, &W_grad, &L3_dgrad - }; + auto args = make_kernel_args( + num_products, L1_in, L2_in, W, L3_grad, L1_dgrad, L2_dgrad, w_dgrad, + L1_grad, L2_grad, W_grad, L3_dgrad); double_backward_config_ref.hStream = stream; - jit.execute(2, args, with_stream(forward_config_ref, stream)); - jit.execute(3, args, with_stream(double_backward_config_ref, stream)); + jit.execute(2, args.data(), args.arg_sizes(), args.count(), + with_stream(forward_config_ref, stream)); + jit.execute(3, args.data(), args.arg_sizes(), args.count(), + with_stream(double_backward_config_ref, stream)); } ~JITTPImpl() = default; diff --git a/openequivariance/openequivariance/extension/torch_core.hpp b/openequivariance/openequivariance/extension/torch_core.hpp index ab78d96a..d4f0faee 100644 --- a/openequivariance/openequivariance/extension/torch_core.hpp +++ b/openequivariance/openequivariance/extension/torch_core.hpp @@ -1,6 +1,7 @@ #pragma once #include +#include #include #include #include @@ -25,6 +26,12 @@ using GPU_Allocator = HIP_Allocator; #endif +#ifdef SYCL_BACKEND + #include "backend_sycl.hpp" + using JITKernel = SYCLJITKernel; + using GPU_Allocator = SYCL_Allocator; +#endif + #include "group_mm.hpp" #include "tensorproducts.hpp" @@ -43,6 +50,7 @@ void tensor_zero_(Tensor &tensor); void alert_not_deterministic(const char *name); Stream get_current_stream(); +bool tensor_is_on_gpu(const Tensor &tensor); const uint8_t *tensor_data_ptr_u8(const Tensor &tensor); void *data_ptr(const Tensor &tensor); @@ -121,7 +129,7 @@ inline void check_tensor(const Tensor &tensor, "Shape mismatch for tensor '", tensor_name, "'. Expected: ", shape_to_string(expected_shape), ". Got: ", tensor_sizes_str(tensor)); - TCHECK(tensor.is_cuda(), "Tensor '", tensor_name, "' is not on the GPU."); + TCHECK(tensor_is_on_gpu(tensor), "Tensor '", tensor_name, "' is not on the GPU."); TCHECK(tensor.scalar_type() == expected_dtype, "Dtype mismatch for tensor '", tensor_name, "'. Expected: ", static_cast(expected_dtype), @@ -186,6 +194,33 @@ inline std::unordered_map lock(mut); + tp_cache.clear(); + conv_cache.clear(); +} + +inline void register_kernel_cache_cleanup() { + static const bool registered = [] { + std::atexit(release_kernel_caches); + return true; + }(); + (void) registered; +} +#endif + inline std::pair*, KernelProp> compile_tp_with_caching(const Tensor &json_bytes, int64_t hash) { @@ -220,6 +255,9 @@ inline std::pair*, KernelProp> std::make_pair(std::move(jit_tp_impl), KernelProp(kernel_prop_map, false))}); it = tp_cache.find(hash); +#ifdef SYCL_BACKEND + register_kernel_cache_cleanup(); +#endif } return {it->second.first.get(), it->second.second}; } @@ -259,6 +297,9 @@ inline std::pair*, KernelProp> std::make_pair(std::move(jit_conv_impl), KernelProp(kernel_prop_map, true))}); it = conv_cache.find(hash); +#ifdef SYCL_BACKEND + register_kernel_cache_cleanup(); +#endif } return {it->second.first.get(), it->second.second}; } @@ -652,7 +693,15 @@ inline Tensor group_gemm( // =========================================================== -REGISTER_LIBRARY_IMPL(libtorch_tp_jit, CUDA, m) { +// The dispatch key must match the device the tensors live on: XPU for SYCL, +// CUDA for both CUDA and HIP (PyTorch maps HIP tensors onto the CUDA key). +#ifdef SYCL_BACKEND + #define OEQ_DISPATCH_KEY XPU +#else + #define OEQ_DISPATCH_KEY CUDA +#endif + +REGISTER_LIBRARY_IMPL(libtorch_tp_jit, OEQ_DISPATCH_KEY, m) { m.impl("jit_tp_forward", BOX(&jit_tp_forward)); m.impl("jit_tp_backward", BOX(&jit_tp_backward)); m.impl("jit_tp_double_backward", BOX(&jit_tp_double_backward)); diff --git a/openequivariance/openequivariance/jax/TensorProduct.py b/openequivariance/openequivariance/jax/TensorProduct.py index f880544f..419f3b33 100644 --- a/openequivariance/openequivariance/jax/TensorProduct.py +++ b/openequivariance/openequivariance/jax/TensorProduct.py @@ -16,7 +16,7 @@ class TensorProduct(LoopUnrollTP): def __init__(self, problem: TPProblem): dp = extlib.DeviceProp(0) - super().__init__(problem, dp, extlib.IS_HIP, torch_op=False) + super().__init__(problem, dp, extlib.BACKEND, torch_op=False) self.kernel = self.kernel_string self.weight_numel = problem.weight_numel diff --git a/openequivariance/openequivariance/jax/TensorProductConv.py b/openequivariance/openequivariance/jax/TensorProductConv.py index 9234158f..daf6ff7b 100644 --- a/openequivariance/openequivariance/jax/TensorProductConv.py +++ b/openequivariance/openequivariance/jax/TensorProductConv.py @@ -42,7 +42,7 @@ def __init__( super().__init__( config, dp, - extlib.IS_HIP, + extlib.BACKEND, idx_dtype=np.int32, torch_op=False, deterministic=deterministic, diff --git a/openequivariance/openequivariance/jax/extlib/__init__.py b/openequivariance/openequivariance/jax/extlib/__init__.py index 23a6f63a..fdd7553d 100644 --- a/openequivariance/openequivariance/jax/extlib/__init__.py +++ b/openequivariance/openequivariance/jax/extlib/__init__.py @@ -3,6 +3,8 @@ IS_HIP = oeq_extjax.is_hip() +BACKEND = "hip" if IS_HIP else "cuda" + platform = "CUDA" if IS_HIP: platform = "ROCM" @@ -16,4 +18,6 @@ __all__ = [ "GPUTimer", "DeviceProp", + "BACKEND", + "IS_HIP", ] diff --git a/openequivariance/openequivariance/templates/common.cuh b/openequivariance/openequivariance/templates/common.cuh index dff311aa..e8dcb231 100644 --- a/openequivariance/openequivariance/templates/common.cuh +++ b/openequivariance/openequivariance/templates/common.cuh @@ -1,3 +1,7 @@ +{%- if is_sycl %} +{% include 'sycl_compat.cuh' %} +{%- endif %} + #define ROW_OPERATION(ROW_LEN, LOOP_VAR, ...) \ _Pragma ("unroll") \ for(int LOOP_VAR = 0; LOOP_VAR < ROW_LEN; LOOP_VAR += THREADS_PER_WARP) { \ diff --git a/openequivariance/openequivariance/templates/jinja_utils.py b/openequivariance/openequivariance/templates/jinja_utils.py index 076fd198..024841da 100644 --- a/openequivariance/openequivariance/templates/jinja_utils.py +++ b/openequivariance/openequivariance/templates/jinja_utils.py @@ -18,8 +18,19 @@ def sizeof(dtype): raise Exception("Provided undefined datatype to sizeof!") -@lru_cache(maxsize=2) -def get_jinja_environment(is_hip=False): +@lru_cache(maxsize=8) +def get_jinja_environment(backend="cuda", warp_size=32): + """ + Builds the Jinja environment used to render the kernel templates. + + :param backend: one of ``"cuda"``, ``"hip"`` or ``"sycl"``. + :param warp_size: size of a warp / wavefront / sub-group. Only consulted by + the SYCL backend, which must bake the sub-group size into + the generated kernel as a compile-time property. + """ + if backend not in ("cuda", "hip", "sycl"): + raise ValueError(f"Unknown kernel backend '{backend}'") + env = Environment( loader=PackageLoader("openequivariance"), extensions=["jinja2.ext.do"] ) @@ -28,18 +39,32 @@ def get_jinja_environment(is_hip=False): env.globals["sizeof"] = sizeof env.globals["enumerate"] = enumerate + is_hip = backend == "hip" + is_sycl = backend == "sycl" + + env.globals["backend"] = backend env.globals["is_hip"] = is_hip - env.globals["syncwarp"] = ( - '__builtin_amdgcn_fence(__ATOMIC_RELEASE, "wavefront");__builtin_amdgcn_wave_barrier();__builtin_amdgcn_fence(__ATOMIC_ACQUIRE, "wavefront");' - if is_hip - else "__syncwarp()" - ) - env.globals["atomic_add"] = "unsafeAtomicAdd" if is_hip else "atomicAdd" + env.globals["is_sycl"] = is_sycl + env.globals["warp_size"] = warp_size - if is_hip: + if is_sycl: + # Provided by templates/sycl_compat.cuh. + env.globals["syncwarp"] = "oeq_syncwarp()" + env.globals["atomic_add"] = "oeq_atomic_add" + env.globals["shfl_down"] = lambda val, offset: f"oeq_shfl_down({val}, {offset})" + elif is_hip: + env.globals["syncwarp"] = ( + '__builtin_amdgcn_fence(__ATOMIC_RELEASE, "wavefront");' + "__builtin_amdgcn_wave_barrier();" + '__builtin_amdgcn_fence(__ATOMIC_ACQUIRE, "wavefront");' + ) + env.globals["atomic_add"] = "unsafeAtomicAdd" env.globals["shfl_down"] = lambda val, offset: f"__shfl_down( {val}, {offset})" else: + env.globals["syncwarp"] = "__syncwarp()" + env.globals["atomic_add"] = "atomicAdd" env.globals["shfl_down"] = ( lambda val, offset: f"__shfl_down_sync(FULL_MASK, {val}, {offset})" ) + return env diff --git a/openequivariance/openequivariance/templates/loop_unroll_batch.cuh b/openequivariance/openequivariance/templates/loop_unroll_batch.cuh index 83e0e0d2..0e60561d 100644 --- a/openequivariance/openequivariance/templates/loop_unroll_batch.cuh +++ b/openequivariance/openequivariance/templates/loop_unroll_batch.cuh @@ -5,7 +5,8 @@ transpose_load, transpose_store, load_ir_segments, load_ir_segments_force, store_ir_segments, declare_smem_variables, - set_launch_bound_variables, launch_bounds + set_launch_bound_variables, launch_bounds, + declare_smem with context%} {%- from 'loop_unroll_tp.cuh' import @@ -26,7 +27,7 @@ __global__ void {{ launch_bounds(forward_schedule) }} forward(size_t num_products, IRREP_T* L1_in, IRREP_T* L2_in, IRREP_T* L3_out, WEIGHT_T* weights) { - extern __shared__ char s[]; + {{ declare_smem(forward_schedule) }} {{ set_launch_bound_variables(forward_schedule.launch_config) }} {%- set tpp = forward_schedule.updated_config %} char* smem = s + {{forward_schedule.memory_per_warp}} * warp_loc; @@ -73,7 +74,7 @@ backward(size_t num_products, IRREP_T* L2_in, IRREP_T* L2_grad, WEIGHT_T* weights, WEIGHT_T* weights_grad, IRREP_T* L3_grad) { - extern __shared__ char s[]; + {{ declare_smem(backward_schedule) }} {{ set_launch_bound_variables(backward_schedule.launch_config) }} char* smem = s + {{backward_schedule.memory_per_warp}} * warp_loc; @@ -153,7 +154,7 @@ double_backward_A( IRREP_T* L1_dgrad, IRREP_T* L2_dgrad, IRREP_T* W_dgrad, // Gradients w.r.t outputs of backward op IRREP_T* L1_grad, IRREP_T* L2_grad, WEIGHT_T* W_grad, IRREP_T* L3_dgrad) { - extern __shared__ char s[]; + {{ declare_smem(forward_schedule) }} {{ set_launch_bound_variables(forward_schedule.launch_config) }} {%- set tpp = forward_schedule.updated_config %} char* smem = s + {{forward_schedule.memory_per_warp}} * warp_loc; @@ -223,7 +224,7 @@ double_backward_B( IRREP_T* L1_dgrad, IRREP_T* L2_dgrad, IRREP_T* W_dgrad, // Gradients w.r.t outputs of backward op IRREP_T* L1_grad, IRREP_T* L2_grad, WEIGHT_T* W_grad, IRREP_T* L3_dgrad) { - extern __shared__ char s[]; + {{ declare_smem(schedule) }} {{ set_launch_bound_variables(schedule.launch_config) }} char* smem = s + {{schedule.memory_per_warp}} * warp_loc; diff --git a/openequivariance/openequivariance/templates/loop_unroll_conv_atomic.cuh b/openequivariance/openequivariance/templates/loop_unroll_conv_atomic.cuh index 3d461dbc..6d50a9a4 100644 --- a/openequivariance/openequivariance/templates/loop_unroll_conv_atomic.cuh +++ b/openequivariance/openequivariance/templates/loop_unroll_conv_atomic.cuh @@ -7,7 +7,8 @@ load_ir_segments, load_ir_segments_force, store_ir_segments, declare_smem_variables, - set_launch_bound_variables, launch_bounds + set_launch_bound_variables, launch_bounds, + declare_smem with context %} #define THREADS_PER_WARP {{ forward_schedule.launch_config.warp_size }} // Warp size should be the same for forward and backward @@ -58,7 +59,7 @@ forward(IRREP_T* L1_in, ConvData c, void* workspace) { - extern __shared__ char s[]; + {{ declare_smem(forward_schedule) }} size_t num_products = c.nnz; unsigned {{idx_type}}* rows = (unsigned {{idx_type}}*) c.rows; unsigned {{idx_type}}* cols = (unsigned {{idx_type}}*) c.cols; @@ -111,7 +112,7 @@ backward(IRREP_T* L1_in, IRREP_T* L1_grad, WEIGHT_T* weights, WEIGHT_T* weights_grad, IRREP_T* L3_grad, ConvData c, void* workspace, unsigned {{idx_type}}* transpose_perm) { - extern __shared__ char s[]; + {{ declare_smem(backward_schedule) }} size_t num_products = c.nnz; unsigned {{idx_type}}* rows = (unsigned {{idx_type}}*) c.rows; unsigned {{idx_type}}* cols = (unsigned {{idx_type}}*) c.cols; @@ -186,7 +187,7 @@ double_backward_A(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, IRREP_T* L1_grad, IRREP_T* L2_grad, WEIGHT_T* W_grad, IRREP_T* L3_dgrad, ConvData c, void* workspace, unsigned {{idx_type}}* transpose_perm) { - extern __shared__ char s[]; + {{ declare_smem(forward_schedule) }} size_t num_products = c.nnz; unsigned {{idx_type}}* rows = (unsigned {{idx_type}}*) c.rows; unsigned {{idx_type}}* cols = (unsigned {{idx_type}}*) c.cols; @@ -268,7 +269,7 @@ double_backward_B(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, unsigned {{idx_type}}* rows = (unsigned {{idx_type}}*) c.rows; unsigned {{idx_type}}* cols = (unsigned {{idx_type}}*) c.cols; - extern __shared__ char s[]; + {{ declare_smem(double_backward_schedule) }} {{ set_launch_bound_variables(schedule.launch_config) }} char* smem = s + {{schedule.memory_per_warp}} * warp_loc; diff --git a/openequivariance/openequivariance/templates/loop_unroll_conv_det.cuh b/openequivariance/openequivariance/templates/loop_unroll_conv_det.cuh index f5bb56a4..5f363bf8 100644 --- a/openequivariance/openequivariance/templates/loop_unroll_conv_det.cuh +++ b/openequivariance/openequivariance/templates/loop_unroll_conv_det.cuh @@ -6,7 +6,8 @@ transpose_load, transpose_store, load_ir_segments, store_ir_segments, declare_smem_variables, - set_launch_bound_variables, launch_bounds + set_launch_bound_variables, launch_bounds, + declare_smem with context %} #define THREADS_PER_WARP {{ forward_schedule.launch_config.warp_size }} // Warp size should be the same for forward and backward @@ -103,7 +104,7 @@ forward( ConvData c, void* workspace_raw) { - extern __shared__ char s[]; + {{ declare_smem(forward_schedule) }} size_t num_products = c.nnz; {{idx_type}}* rows = ({{idx_type}}*) c.rows; {{idx_type}}* cols = ({{idx_type}}*) c.cols; @@ -192,7 +193,7 @@ backward(IRREP_T* L1_in, IRREP_T* L1_grad, IRREP_T* L3_grad, ConvData c, void* workspace_raw, {{idx_type}}* transpose_perm) { - extern __shared__ char s[]; + {{ declare_smem(backward_schedule) }} size_t num_products = c.nnz; // Note the transpose below (cols -> rows, rows -> cols) @@ -300,7 +301,7 @@ double_backward_A(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, IRREP_T* L1_grad, IRREP_T* L2_grad, WEIGHT_T* W_grad, IRREP_T* L3_dgrad, ConvData c, void* workspace_raw, unsigned {{idx_type}}* transpose_perm) { - extern __shared__ char s[]; + {{ declare_smem(forward_schedule) }} size_t num_products = c.nnz; unsigned {{idx_type}}* rows = (unsigned {{idx_type}}*) c.rows; unsigned {{idx_type}}* cols = (unsigned {{idx_type}}*) c.cols; @@ -421,7 +422,7 @@ double_backward_B(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, {{idx_type}}* cols = ({{idx_type}}*) c.rows; {{idx_type}}* tperm = ({{idx_type}}*) transpose_perm; - extern __shared__ char s[]; + {{ declare_smem(double_backward_schedule) }} {{ set_launch_bound_variables(schedule.launch_config) }} char* smem = s + {{schedule.memory_per_warp}} * warp_loc; diff --git a/openequivariance/openequivariance/templates/macros.jinja b/openequivariance/openequivariance/templates/macros.jinja index ce2a840b..dd10b1df 100644 --- a/openequivariance/openequivariance/templates/macros.jinja +++ b/openequivariance/openequivariance/templates/macros.jinja @@ -236,3 +236,14 @@ Keys map to lists of tuples with (name, dtype, num_elements) of each subarray. {%- macro launch_bounds(schedule) %} __launch_bounds__({{schedule.launch_config.num_threads}}) {%- endmacro %} + +{# Declares the per-block shared memory buffer `s`. CUDA and HIP size the + allocation at launch; SYCL runtime compilation has no dynamic local memory + for free-function kernels, so the size is baked in from the schedule. #} +{%- macro declare_smem(schedule) %} + {%- if is_sycl %} + OEQ_DECLARE_SMEM({{ schedule.launch_config.smem }}) + {%- else %} + extern __shared__ char s[]; + {%- endif %} +{%- endmacro %} diff --git a/openequivariance/openequivariance/templates/sycl_compat.cuh b/openequivariance/openequivariance/templates/sycl_compat.cuh new file mode 100644 index 00000000..34468ee5 --- /dev/null +++ b/openequivariance/openequivariance/templates/sycl_compat.cuh @@ -0,0 +1,116 @@ +{# +Compatibility shim that lets the CUDA/HIP-flavored kernel templates compile as +SYCL free-function kernels under runtime compilation. Included only when +targeting the SYCL backend; CUDA and HIP see none of this. + +The generated source is compiled by the SYCL kernel_compiler extension, so it +must be self-contained: everything the kernel body relies on is declared here. +#} +#include + +#include +#include + +namespace syclex = sycl::ext::oneapi::experimental; +namespace twi = sycl::ext::oneapi::this_work_item; + +// A CUDA __global__ kernel becomes a SYCL free-function nd_range kernel. The +// sub-group size is fixed to the warp size the schedule was generated against, +// which is what makes the warp-level code below well-defined. +#define OEQ_SUBGROUP_SIZE {{ warp_size }} +#define __global__ extern "C" SYCL_EXTERNAL \ + SYCL_EXT_ONEAPI_FUNCTION_PROPERTY((syclex::nd_range_kernel<1>)) \ + SYCL_EXT_ONEAPI_FUNCTION_PROPERTY((syclex::sub_group_size)) + +#define __device__ +#define __host__ +#define __forceinline__ inline +#define __restrict__ __restrict + +// Occupancy hints have no runtime-compilation equivalent; the work-group size +// is supplied at launch instead. +#define __launch_bounds__(...) + +// --------------------------------------------------------------------------- +// Thread / block indexing +// --------------------------------------------------------------------------- +// All generated kernels are launched as 1D nd_ranges, so only .x is meaningful. +struct OeqIndex1D { + size_t x; + operator size_t() const { return x; } +}; + +static inline OeqIndex1D oeq_thread_idx() { return {twi::get_nd_item<1>().get_local_id(0)}; } +static inline OeqIndex1D oeq_block_idx() { return {twi::get_nd_item<1>().get_group(0)}; } +static inline OeqIndex1D oeq_block_dim() { return {twi::get_nd_item<1>().get_local_range(0)}; } +static inline OeqIndex1D oeq_grid_dim() { return {twi::get_nd_item<1>().get_group_range(0)}; } + +#define threadIdx oeq_thread_idx() +#define blockIdx oeq_block_idx() +#define blockDim oeq_block_dim() +#define gridDim oeq_grid_dim() + +// --------------------------------------------------------------------------- +// Synchronization +// --------------------------------------------------------------------------- +static inline void oeq_syncwarp() { + sycl::group_barrier(twi::get_sub_group()); +} + +static inline void oeq_syncthreads() { + sycl::group_barrier(twi::get_nd_item<1>().get_group()); +} + +#define __syncthreads() oeq_syncthreads() +#define __threadfence_block() oeq_syncwarp() + +// --------------------------------------------------------------------------- +// Warp-level primitives +// --------------------------------------------------------------------------- +template +static inline T oeq_shfl_down(T val, int offset) { + return sycl::shift_group_left(twi::get_sub_group(), val, offset); +} + +// --------------------------------------------------------------------------- +// Atomics +// --------------------------------------------------------------------------- +template +static inline T oeq_atomic_add(T* address, T val) { + sycl::atomic_ref ref(*address); + return ref.fetch_add(val); +} + +// --------------------------------------------------------------------------- +// min / max +// --------------------------------------------------------------------------- +// CUDA provides these as device builtins over mixed integer types. Templating +// on both operands keeps the mixed-width call sites in the templates working. +template +static inline auto oeq_min(A a, B b) -> typename std::common_type::type { + using C = typename std::common_type::type; + return static_cast(a) < static_cast(b) ? static_cast(a) : static_cast(b); +} + +template +static inline auto oeq_max(A a, B b) -> typename std::common_type::type { + using C = typename std::common_type::type; + return static_cast(a) > static_cast(b) ? static_cast(a) : static_cast(b); +} + +#define min oeq_min +#define max oeq_max + +// --------------------------------------------------------------------------- +// Shared memory +// --------------------------------------------------------------------------- +// CUDA's `extern __shared__ char s[]` sizes the allocation at launch. SYCL +// runtime compilation has no dynamic-local-memory equivalent for free-function +// kernels, so each kernel declares a function-scope work_group_static buffer +// sized to the shared memory its own schedule requires. +#define OEQ_DECLARE_SMEM(BYTES) \ + static syclex::work_group_static oeq_smem_buf; \ + char* s = &oeq_smem_buf[0]; diff --git a/tests/batch_test.py b/tests/batch_test.py index 77715028..eeff2b4a 100644 --- a/tests/batch_test.py +++ b/tests/batch_test.py @@ -21,6 +21,10 @@ import openequivariance as oeq +from conftest import device_type + +DEVICE = device_type() + @pytest.fixture(params=[np.float32, np.float64], ids=["F32", "F64"], scope="module") def dtype(request): @@ -426,7 +430,7 @@ def test_submodule_dtype_conversion(self, parent_module_and_problem): parent, problem = parent_module_and_problem batch_size = 10 - device = "cuda" + device = DEVICE input_dtype = self._problem_dtype(problem) in1, in2, weights = self._make_inputs(problem, batch_size, input_dtype, device) diff --git a/tests/conftest.py b/tests/conftest.py index 4a515664..af026dc5 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -18,3 +18,25 @@ def pytest_addoption(parser): @pytest.fixture(scope="session") def with_jax(request): return request.config.getoption("--jax") + + +def device_type(): + """ + The torch device type the kernels run on for the detected backend: + ``"xpu"`` for SYCL, ``"cuda"`` for CUDA and HIP. + """ + from openequivariance._torch.extlib import DEVICE_TYPE + + return DEVICE_TYPE + + +def torch_accelerator(): + """The ``torch.cuda`` / ``torch.xpu`` module matching the active backend.""" + import torch + + return getattr(torch, device_type()) + + +@pytest.fixture(scope="session") +def device(): + return device_type() diff --git a/tests/conv_test.py b/tests/conv_test.py index 1adb5553..7dfa2dce 100644 --- a/tests/conv_test.py +++ b/tests/conv_test.py @@ -22,6 +22,10 @@ nequip_oam_problems, ) +from conftest import device_type + +DEVICE = device_type() + @pytest.fixture(params=[np.float32, np.float64], ids=["F32", "F64"], scope="module") def dtype(request): @@ -418,7 +422,7 @@ def _make_inputs(self, problem, graph, rng, dtype, device, deterministic): def test_submodule_dtype_conversion(self, parent_module_and_problem, graph): parent, problem = parent_module_and_problem - device = "cuda" + device = DEVICE rng = np.random.default_rng(12345) input_dtype = self._problem_dtype(problem) diff --git a/tests/example_test.py b/tests/example_test.py index bf51d3ed..2b363864 100644 --- a/tests/example_test.py +++ b/tests/example_test.py @@ -1,6 +1,10 @@ import pytest import os +from conftest import device_type + +DEVICE = device_type() + def test_tutorial_torch(with_jax): if with_jax: @@ -9,19 +13,19 @@ def test_tutorial_torch(with_jax): import torch import e3nn.o3 as o3 - gen = torch.Generator(device="cuda") + gen = torch.Generator(device=DEVICE) batch_size = 1000 X_ir, Y_ir, Z_ir = o3.Irreps("1x2e"), o3.Irreps("1x3e"), o3.Irreps("1x2e") - X = torch.rand(batch_size, X_ir.dim, device="cuda", generator=gen) - Y = torch.rand(batch_size, Y_ir.dim, device="cuda", generator=gen) + X = torch.rand(batch_size, X_ir.dim, device=DEVICE, generator=gen) + Y = torch.rand(batch_size, Y_ir.dim, device=DEVICE, generator=gen) instructions = [(0, 0, 0, "uvu", True)] tp_e3nn = o3.TensorProduct( X_ir, Y_ir, Z_ir, instructions, shared_weights=False, internal_weights=False - ).to("cuda") - W = torch.rand(batch_size, tp_e3nn.weight_numel, device="cuda", generator=gen) + ).to(DEVICE) + W = torch.rand(batch_size, tp_e3nn.weight_numel, device=DEVICE, generator=gen) Z = tp_e3nn(X, Y, W) print(torch.norm(Z)) @@ -51,13 +55,13 @@ def test_tutorial_torch(with_jax): [0, 1, 1, 2], # Receiver [1, 0, 2, 1], ], # Sender - device="cuda", + device=DEVICE, dtype=torch.long, ) - X = torch.rand(node_ct, X_ir.dim, device="cuda", generator=gen) - Y = torch.rand(nonzero_ct, Y_ir.dim, device="cuda", generator=gen) - W = torch.rand(nonzero_ct, problem.weight_numel, device="cuda", generator=gen) + X = torch.rand(node_ct, X_ir.dim, device=DEVICE, generator=gen) + Y = torch.rand(nonzero_ct, Y_ir.dim, device=DEVICE, generator=gen) + W = torch.rand(nonzero_ct, problem.weight_numel, device=DEVICE, generator=gen) tp_conv = oeq.TensorProductConv( problem, deterministic=False diff --git a/tests/export_test.py b/tests/export_test.py index 6aba9690..d7bd169e 100644 --- a/tests/export_test.py +++ b/tests/export_test.py @@ -13,6 +13,10 @@ from openequivariance._torch.E3NNTensorProduct import E3NNTensorProduct +from conftest import device_type + +DEVICE = device_type() + @pytest.fixture(scope="session") def problem_and_irreps(): @@ -28,7 +32,7 @@ def problem_and_irreps(): weight_dtype=np.float32, ) - gen = torch.Generator(device="cuda") + gen = torch.Generator(device=DEVICE) gen.manual_seed(0) return ( @@ -42,30 +46,30 @@ def problem_and_irreps(): @pytest.fixture(params=["batch", "conv_det", "conv_atomic"], scope="session") def tp_and_inputs(request, problem_and_irreps): problem, X_ir, Y_ir, _ = problem_and_irreps - gen = torch.Generator(device="cuda") + gen = torch.Generator(device=DEVICE) gen.manual_seed(0) if request.param == "batch": batch_size = 1000 - X = torch.rand(batch_size, X_ir.dim, device="cuda", generator=gen) - Y = torch.rand(batch_size, Y_ir.dim, device="cuda", generator=gen) - W = torch.rand(batch_size, problem.weight_numel, device="cuda", generator=gen) + X = torch.rand(batch_size, X_ir.dim, device=DEVICE, generator=gen) + Y = torch.rand(batch_size, Y_ir.dim, device=DEVICE, generator=gen) + W = torch.rand(batch_size, problem.weight_numel, device=DEVICE, generator=gen) return oeq.TensorProduct(problem), (X, Y, W) else: node_ct, nonzero_ct = 3, 4 # Receiver, sender indices for message passing GNN edge_index = EdgeIndex( - [[0, 1, 1, 2], [1, 0, 2, 1]], device="cuda", dtype=torch.long + [[0, 1, 1, 2], [1, 0, 2, 1]], device=DEVICE, dtype=torch.long ) _, sender_perm = edge_index.sort_by("col") edge_index, _ = edge_index.sort_by("row") edge_index = [edge_index[0].detach(), edge_index[1].detach()] - X = torch.rand(node_ct, X_ir.dim, device="cuda", generator=gen) - Y = torch.rand(nonzero_ct, Y_ir.dim, device="cuda", generator=gen) - W = torch.rand(nonzero_ct, problem.weight_numel, device="cuda", generator=gen) + X = torch.rand(node_ct, X_ir.dim, device=DEVICE, generator=gen) + Y = torch.rand(nonzero_ct, Y_ir.dim, device=DEVICE, generator=gen) + W = torch.rand(nonzero_ct, problem.weight_numel, device=DEVICE, generator=gen) if request.param == "conv_atomic": return oeq.TensorProductConv(problem, torch_op=True, deterministic=False), ( @@ -143,18 +147,18 @@ def test_aoti_cpp_inference(problem_and_irreps): cmake_prefix_path = torch.utils.cmake_prefix_path torch_ext_so_path = oeq.torch_ext_so_path() - gen = torch.Generator(device="cuda") + gen = torch.Generator(device=DEVICE) gen.manual_seed(0) batch_size = 1000 # Create models - oeq_tp = oeq.TensorProduct(problem).to("cuda") - e3nn_tp = E3NNTensorProduct(problem).e3nn_tp.to("cuda") + oeq_tp = oeq.TensorProduct(problem).to(DEVICE) + e3nn_tp = E3NNTensorProduct(problem).e3nn_tp.to(DEVICE) # Prepare inputs for export - X = torch.rand(batch_size, X_ir.dim, device="cuda", generator=gen) - Y = torch.rand(batch_size, Y_ir.dim, device="cuda", generator=gen) - W = torch.rand(batch_size, problem.weight_numel, device="cuda", generator=gen) + X = torch.rand(batch_size, X_ir.dim, device=DEVICE, generator=gen) + Y = torch.rand(batch_size, Y_ir.dim, device=DEVICE, generator=gen) + W = torch.rand(batch_size, problem.weight_numel, device=DEVICE, generator=gen) inputs = (X, Y, W) with ( diff --git a/tests/input_validation_test.py b/tests/input_validation_test.py index 9b38d55e..683db47d 100644 --- a/tests/input_validation_test.py +++ b/tests/input_validation_test.py @@ -5,6 +5,10 @@ from openequivariance import TPProblem, TensorProduct, TensorProductConv +from conftest import device_type + +DEVICE = device_type() + @pytest.fixture def tpp(): @@ -26,7 +30,7 @@ def edge_index(): ], sort_order="row", sparse_size=(3, 4), - device="cuda", + device=DEVICE, dtype=torch.long, ) ei.fill_cache_() @@ -35,28 +39,28 @@ def edge_index(): @pytest.fixture def tp_buffers(tpp): - gen = torch.Generator(device="cuda") + gen = torch.Generator(device=DEVICE) gen.manual_seed(42) N = 1000 - X = torch.rand(N, tpp.irreps_in1.dim, device="cuda", generator=gen) - Y = torch.rand(N, tpp.irreps_in2.dim, device="cuda", generator=gen) - W = torch.rand(N, tpp.weight_numel, device="cuda", generator=gen) + X = torch.rand(N, tpp.irreps_in1.dim, device=DEVICE, generator=gen) + Y = torch.rand(N, tpp.irreps_in2.dim, device=DEVICE, generator=gen) + W = torch.rand(N, tpp.weight_numel, device=DEVICE, generator=gen) return [X, Y, W] @pytest.fixture def conv_buffers(edge_index, tpp): - gen = torch.Generator(device="cuda") + gen = torch.Generator(device=DEVICE) gen.manual_seed(42) X = torch.rand( - edge_index.num_rows, tpp.irreps_in1.dim, device="cuda", generator=gen + edge_index.num_rows, tpp.irreps_in1.dim, device=DEVICE, generator=gen ) Y = torch.rand( - edge_index.num_cols, tpp.irreps_in2.dim, device="cuda", generator=gen + edge_index.num_cols, tpp.irreps_in2.dim, device=DEVICE, generator=gen ) - W = torch.rand(edge_index.num_cols, tpp.weight_numel, device="cuda", generator=gen) + W = torch.rand(edge_index.num_cols, tpp.weight_numel, device=DEVICE, generator=gen) _, inv_perm = edge_index.get_csc() return [X, Y, W, edge_index[0], edge_index[1], inv_perm] diff --git a/tests/multidevice_test.py b/tests/multidevice_test.py index 7b7b48c7..be4b5b7b 100644 --- a/tests/multidevice_test.py +++ b/tests/multidevice_test.py @@ -3,6 +3,13 @@ import subprocess import os +# Resolved directly rather than through conftest: this file is also executed +# as a standalone script by torch.distributed.run below, where conftest is not +# importable. +from openequivariance._torch.extlib import DEVICE_TYPE as DEVICE + +ACCEL = getattr(torch, DEVICE) + def test_multidevice(): result = subprocess.run( @@ -40,7 +47,7 @@ def test_multidevice(): problem = mace_problems()[0] local_rank = int(os.environ["LOCAL_RANK"]) - device = f"cuda:{local_rank}" + device = f"{DEVICE}:{local_rank}" torch.set_default_device(device) X_ir, Y_ir, Z_ir = problem.irreps_in1, problem.irreps_in2, problem.irreps_out @@ -53,5 +60,5 @@ def test_multidevice(): Y = torch.rand(batch_size, Y_ir.dim, device=device, generator=gen) W = torch.rand(batch_size, problem.weight_numel, device=device, generator=gen) - with torch.cuda.device(device): + with ACCEL.device(device): result = tp.forward(X, Y, W) diff --git a/tests/stream_test.py b/tests/stream_test.py index 42ac4dd2..582e0055 100644 --- a/tests/stream_test.py +++ b/tests/stream_test.py @@ -15,6 +15,10 @@ from openequivariance import TensorProduct, TensorProductConv, TPProblem +from conftest import device_type, torch_accelerator + +DEVICE = device_type() + class KernelExpectation(NamedTuple): kernel_name: str @@ -34,12 +38,13 @@ def __call__(self) -> Any: return self.func(*self.buffers) -cuda = torch.device("cuda") +accel_device = torch.device(DEVICE) +ACCEL = torch_accelerator() @pytest.fixture def gen(): - return torch.Generator(device="cuda") + return torch.Generator(device=DEVICE) @pytest.fixture @@ -55,7 +60,7 @@ def edge_index(): [1, 0, 2, 1], # Sender ], sparse_size=(3, 4), - device="cuda", + device=DEVICE, dtype=torch.long, ) @@ -73,21 +78,21 @@ def tpp(): @pytest.fixture def tp_buffers(N, tpp, gen): - X = torch.rand(N, tpp.irreps_in1.dim, device="cuda", generator=gen) - Y = torch.rand(N, tpp.irreps_in2.dim, device="cuda", generator=gen) - W = torch.rand(N, tpp.weight_numel, device="cuda", generator=gen) + X = torch.rand(N, tpp.irreps_in1.dim, device=DEVICE, generator=gen) + Y = torch.rand(N, tpp.irreps_in2.dim, device=DEVICE, generator=gen) + W = torch.rand(N, tpp.weight_numel, device=DEVICE, generator=gen) return (X, Y, W) @pytest.fixture def conv_buffers(edge_index, tpp, gen): X = torch.rand( - edge_index.num_rows, tpp.irreps_in1.dim, device="cuda", generator=gen + edge_index.num_rows, tpp.irreps_in1.dim, device=DEVICE, generator=gen ) Y = torch.rand( - edge_index.num_cols, tpp.irreps_in2.dim, device="cuda", generator=gen + edge_index.num_cols, tpp.irreps_in2.dim, device=DEVICE, generator=gen ) - W = torch.rand(edge_index.num_cols, tpp.weight_numel, device="cuda", generator=gen) + W = torch.rand(edge_index.num_cols, tpp.weight_numel, device=DEVICE, generator=gen) return (X, Y, W, edge_index[0], edge_index[1]) @@ -139,7 +144,7 @@ def double_backward_fn(X, Y, W): dummy = torch.norm(in1_grad) + torch.norm(in2_grad) + torch.norm(w_grad) # Second backward - dummy_grad = torch.tensor(1.0, device="cuda") + dummy_grad = torch.tensor(1.0, device=DEVICE) dummy.backward( dummy_grad, retain_graph=True, @@ -217,7 +222,7 @@ def double_backward_fn(X, Y, W, receivers, senders): dummy = torch.norm(in1_grad) + torch.norm(in2_grad) + torch.norm(w_grad) # Second backward - dummy_grad = torch.tensor(1.0, device="cuda") + dummy_grad = torch.tensor(1.0, device=DEVICE) dummy.backward( dummy_grad, retain_graph=True, @@ -297,7 +302,7 @@ def double_backward_fn(X, Y, W, receivers, senders): dummy = torch.norm(in1_grad) + torch.norm(in2_grad) + torch.norm(w_grad) # Second backward - dummy_grad = torch.tensor(1.0, device="cuda") + dummy_grad = torch.tensor(1.0, device=DEVICE) dummy.backward( dummy_grad, retain_graph=True, @@ -345,8 +350,8 @@ def test_separate_streams(request, tmp_path, executable: Executable): ) as prof: streams = [-1, -2] for priority in streams: - s = torch.cuda.Stream(device=cuda, priority=priority) - with torch.cuda.stream(s): + s = ACCEL.Stream(device=accel_device, priority=priority) + with ACCEL.stream(s): with record_function(f"executable_{priority}"): for _ in range(COUNT): executable() diff --git a/tests/symmetric_contraction_test.py b/tests/symmetric_contraction_test.py index bbd105ac..bc3459a6 100644 --- a/tests/symmetric_contraction_test.py +++ b/tests/symmetric_contraction_test.py @@ -9,6 +9,10 @@ from openequivariance._torch.symmetric_contraction import SymmetricContraction +from conftest import device_type + +DEVICE = device_type() + mace_symmetric_contraction = pytest.importorskip("mace.modules.symmetric_contraction") MaceSymmetricContraction = mace_symmetric_contraction.SymmetricContraction @@ -25,7 +29,7 @@ ], ) -DEVICE = torch.device("cuda") +DEVICE = torch.device(DEVICE) SC_CONFIGS = [ SCConfig( From a5525497f8930b23a68d6fe55f3e3a83b382f752 Mon Sep 17 00:00:00 2001 From: Abhishek Bagusetty Date: Thu, 3 Sep 2026 10:49:30 -0500 Subject: [PATCH 02/12] ci: verify the SYCL extension build on a GPU-less runner Mirrors the CUDA job. GitHub offers no Intel GPU runner, so only the build is exercised. --- .github/workflows/requirements_sycl_ci.txt | 7 +++ .github/workflows/verify_extension_build.yml | 44 ++++++++++++++++++- CHANGELOG.md | 3 ++ .../_torch/extlib/__init__.py | 6 ++- 4 files changed, 58 insertions(+), 2 deletions(-) create mode 100644 .github/workflows/requirements_sycl_ci.txt diff --git a/.github/workflows/requirements_sycl_ci.txt b/.github/workflows/requirements_sycl_ci.txt new file mode 100644 index 00000000..2ad949fa --- /dev/null +++ b/.github/workflows/requirements_sycl_ci.txt @@ -0,0 +1,7 @@ +--extra-index-url https://download.pytorch.org/whl/xpu +numpy==2.2.5 +torch==2.10.0+xpu +pytest==9.0.3 +ninja==1.11.1.4 +nanobind==2.10.2 +scikit-build-core==0.11.6 diff --git a/.github/workflows/verify_extension_build.yml b/.github/workflows/verify_extension_build.yml index 6c903205..234405ca 100644 --- a/.github/workflows/verify_extension_build.yml +++ b/.github/workflows/verify_extension_build.yml @@ -41,4 +41,46 @@ jobs: - name: Test JAX extension build run: | - XLA_DIRECT_DOWNLOAD=1 pip install -e "./openequivariance_extjax" --no-build-isolation \ No newline at end of file + XLA_DIRECT_DOWNLOAD=1 pip install -e "./openequivariance_extjax" --no-build-isolation + + verify_sycl_extension: + if: ${{ github.event.label.name == 'ci-ready' || github.event_name != 'pull_request' }} + runs-on: ubuntu-latest + + steps: + - uses: actions/checkout@v4 + - name: Set up Python + uses: actions/setup-python@v5 + with: + python-version: "3.10" + cache: 'pip' + cache-dependency-path: '**/requirements_sycl_ci.txt' + + # The SYCL backend compiles the extension with the oneAPI DPC++ driver + # (icpx), so the compiler and the oneMKL headers both have to be present. + # No Intel GPU is attached to this runner; only the build is exercised. + - name: Install oneAPI DPC++ and oneMKL + run: | + wget -qO- https://apt.repos.intel.com/intel-gpg-keys/GPG-PUB-KEY-INTEL-SW-PRODUCTS.PUB \ + | gpg --dearmor | sudo tee /usr/share/keyrings/oneapi-archive-keyring.gpg > /dev/null + echo "deb [signed-by=/usr/share/keyrings/oneapi-archive-keyring.gpg] https://apt.repos.intel.com/oneapi all main" \ + | sudo tee /etc/apt/sources.list.d/oneAPI.list + sudo apt-get update + sudo apt-get install -y intel-oneapi-compiler-dpcpp-cpp intel-oneapi-mkl-devel + + - name: Install dependencies + run: | + python -m pip install --upgrade pip + pip install -r .github/workflows/requirements_sycl_ci.txt + pip install -e "./openequivariance" + + - name: Test SYCL extension build via import + run: | + source /opt/intel/oneapi/setvars.sh + export CXX=icpx + + pytest tests/import_test.py + + export OEQ_JIT_EXTENSION=1 + + pytest tests/import_test.py \ No newline at end of file diff --git a/CHANGELOG.md b/CHANGELOG.md index 9d2f9d38..4f7099e9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,9 @@ `xpu` device. Kernels are generated as SYCL free functions and compiled at runtime through the oneAPI kernel compiler extension. The backend is selected automatically from the active PyTorch build. +- A CI job that verifies the SYCL extension builds. Like the CUDA job it + runs on a GPU-less runner and only exercises the build, since GitHub + offers no Intel GPU runner. **Changed**: - The kernel backend is now identified by a string (`"cuda"` / `"hip"` / diff --git a/openequivariance/openequivariance/_torch/extlib/__init__.py b/openequivariance/openequivariance/_torch/extlib/__init__.py index 2fb244fb..fb571b15 100644 --- a/openequivariance/openequivariance/_torch/extlib/__init__.py +++ b/openequivariance/openequivariance/_torch/extlib/__init__.py @@ -29,12 +29,16 @@ def _detect_backend(): Returns one of ``"cuda"``, ``"hip"`` or ``"sycl"``. HIP builds report a ``torch.version.cuda`` of ``None``, so HIP must be tested first. + + All three checks are build-time properties of the PyTorch install, not + runtime device queries, so importing works on a machine with no + accelerator attached (a CI builder, for instance). """ if torch.version.hip: return "hip" if torch.version.cuda: return "cuda" - if hasattr(torch, "xpu") and torch.xpu.is_available(): + if getattr(torch.version, "xpu", None): return "sycl" return None From ec8b0af78129d05c973f07e7e3b229d663c86b80 Mon Sep 17 00:00:00 2001 From: Abhishek Bagusetty Date: Thu, 3 Sep 2026 11:52:37 -0500 Subject: [PATCH 03/12] ci: pin the SYCL job to a verified PyTorch --- .github/workflows/requirements_sycl_ci.txt | 6 +++++- CHANGELOG.md | 2 ++ docs/installation.rst | 14 +++++++++----- .../openequivariance/_torch/extlib/__init__.py | 7 +++++++ 4 files changed, 23 insertions(+), 6 deletions(-) diff --git a/.github/workflows/requirements_sycl_ci.txt b/.github/workflows/requirements_sycl_ci.txt index 2ad949fa..2d3182e4 100644 --- a/.github/workflows/requirements_sycl_ci.txt +++ b/.github/workflows/requirements_sycl_ci.txt @@ -1,6 +1,10 @@ +# PyTorch >= 2.7 is the minimum for the SYCL backend: 2.6 adds the XPU +# device, shim_xpu.h and cpp_extension's SYCL support, and 2.7 adds +# torch.library.register_autocast, which the operators register with. +# Pinned to the newest XPU wheel this was verified against. --extra-index-url https://download.pytorch.org/whl/xpu numpy==2.2.5 -torch==2.10.0+xpu +torch==2.12.1+xpu pytest==9.0.3 ninja==1.11.1.4 nanobind==2.10.2 diff --git a/CHANGELOG.md b/CHANGELOG.md index 4f7099e9..d8802793 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,8 @@ - A CI job that verifies the SYCL extension builds. Like the CUDA job it runs on a GPU-less runner and only exercises the build, since GitHub offers no Intel GPU runner. +- The SYCL backend requires PyTorch >= 2.7 and raises at import if an + older version is installed. **Changed**: - The kernel backend is now identified by a string (`"cuda"` / `"hip"` / diff --git a/docs/installation.rst b/docs/installation.rst index 4c0b737c..f840780e 100644 --- a/docs/installation.rst +++ b/docs/installation.rst @@ -8,8 +8,8 @@ Installation You need the following to install OpenEquivariance: - A Linux system equipped with an NVIDIA / AMD / Intel graphics card. -- Either PyTorch >= 2.4 (>= 2.8 for AOTI and export), or JAX>0.5.0 - with CUDA, RocM, or XPU support. +- Either PyTorch >= 2.4 (>= 2.8 for AOTI and export, >= 2.7 for Intel + GPUs), or JAX>0.5.0 with CUDA, RocM, or XPU support. - GCC 9+ and the CUDA / HIP toolkit, or the oneAPI DPC++ compiler (``icpx``) for Intel GPUs. The command ``c++ --version`` should return >= 9.0; see below for details on @@ -18,9 +18,13 @@ You need the following to install OpenEquivariance: .. note:: On Intel GPUs the kernels are generated as SYCL and compiled at runtime - through the oneAPI kernel compiler, so ``icpx`` must be on your ``PATH``. - Tensors passed to the kernels live on the ``xpu`` device rather than - ``cuda``. The JAX frontend currently supports CUDA and HIP only. + through the oneAPI kernel compiler, so ``icpx`` must be on your ``PATH`` + (PyTorch locates the SYCL toolchain from it). Tensors passed to the + kernels live on the ``xpu`` device rather than ``cuda``. PyTorch 2.7 is + the minimum: 2.6 introduces the XPU device and the SYCL support in + ``torch.utils.cpp_extension``, and 2.7 adds + ``torch.library.register_autocast``. The precompiled extension and the + JAX frontend both remain CUDA/HIP only. .. tab:: PyTorch diff --git a/openequivariance/openequivariance/_torch/extlib/__init__.py b/openequivariance/openequivariance/_torch/extlib/__init__.py index fb571b15..065f8b7f 100644 --- a/openequivariance/openequivariance/_torch/extlib/__init__.py +++ b/openequivariance/openequivariance/_torch/extlib/__init__.py @@ -53,6 +53,13 @@ def _detect_backend(): IS_HIP = BACKEND == "hip" IS_SYCL = BACKEND == "sycl" +# The SYCL backend needs torch.library.register_autocast, added in 2.7. The +# XPU device itself and cpp_extension's SYCL support arrive in 2.6. +if IS_SYCL and Version(torch.__version__) < Version("2.7"): + raise RuntimeError( + f"The SYCL backend requires PyTorch >= 2.7, found {torch.__version__}." + ) + # The torch device type that tensors passed to the kernels must live on. DEVICE_TYPE = "xpu" if IS_SYCL else "cuda" From 7db6cc1b34dbd21f24b1e9cf56fdb25ba9dd39db Mon Sep 17 00:00:00 2001 From: Abhishek Bagusetty Date: Thu, 3 Sep 2026 12:29:05 -0500 Subject: [PATCH 04/12] docs: align the SYCL PyTorch floor with the project's existing requirement --- .github/workflows/requirements_sycl_ci.txt | 4 +--- CHANGELOG.md | 4 ++-- docs/installation.rst | 13 +++++-------- .../openequivariance/_torch/extlib/__init__.py | 10 ++++++---- 4 files changed, 14 insertions(+), 17 deletions(-) diff --git a/.github/workflows/requirements_sycl_ci.txt b/.github/workflows/requirements_sycl_ci.txt index 2d3182e4..a8b0a39c 100644 --- a/.github/workflows/requirements_sycl_ci.txt +++ b/.github/workflows/requirements_sycl_ci.txt @@ -1,6 +1,4 @@ -# PyTorch >= 2.7 is the minimum for the SYCL backend: 2.6 adds the XPU -# device, shim_xpu.h and cpp_extension's SYCL support, and 2.7 adds -# torch.library.register_autocast, which the operators register with. +# The SYCL backend follows the project's existing PyTorch >= 2.8 floor. # Pinned to the newest XPU wheel this was verified against. --extra-index-url https://download.pytorch.org/whl/xpu numpy==2.2.5 diff --git a/CHANGELOG.md b/CHANGELOG.md index d8802793..47870c4a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,8 +8,8 @@ - A CI job that verifies the SYCL extension builds. Like the CUDA job it runs on a GPU-less runner and only exercises the build, since GitHub offers no Intel GPU runner. -- The SYCL backend requires PyTorch >= 2.7 and raises at import if an - older version is installed. +- The SYCL backend requires PyTorch >= 2.8, the floor the project already + sets for AOTI and export, and raises at import on an older version. **Changed**: - The kernel backend is now identified by a string (`"cuda"` / `"hip"` / diff --git a/docs/installation.rst b/docs/installation.rst index f840780e..1a401ab2 100644 --- a/docs/installation.rst +++ b/docs/installation.rst @@ -8,8 +8,8 @@ Installation You need the following to install OpenEquivariance: - A Linux system equipped with an NVIDIA / AMD / Intel graphics card. -- Either PyTorch >= 2.4 (>= 2.8 for AOTI and export, >= 2.7 for Intel - GPUs), or JAX>0.5.0 with CUDA, RocM, or XPU support. +- Either PyTorch >= 2.4 (>= 2.8 for AOTI, export, and Intel GPUs), or + JAX>0.5.0 with CUDA, RocM, or XPU support. - GCC 9+ and the CUDA / HIP toolkit, or the oneAPI DPC++ compiler (``icpx``) for Intel GPUs. The command ``c++ --version`` should return >= 9.0; see below for details on @@ -20,11 +20,8 @@ You need the following to install OpenEquivariance: On Intel GPUs the kernels are generated as SYCL and compiled at runtime through the oneAPI kernel compiler, so ``icpx`` must be on your ``PATH`` (PyTorch locates the SYCL toolchain from it). Tensors passed to the - kernels live on the ``xpu`` device rather than ``cuda``. PyTorch 2.7 is - the minimum: 2.6 introduces the XPU device and the SYCL support in - ``torch.utils.cpp_extension``, and 2.7 adds - ``torch.library.register_autocast``. The precompiled extension and the - JAX frontend both remain CUDA/HIP only. + kernels live on the ``xpu`` device rather than ``cuda``. The precompiled + extension and the JAX frontend both remain CUDA/HIP only. .. tab:: PyTorch @@ -166,7 +163,7 @@ on a major cluster, send us a pull request to add your configuration! .. code-block:: bash :caption: env.sh (last updated August 2026) - module load oneapi/release + module restore module load frameworks export CC=icx diff --git a/openequivariance/openequivariance/_torch/extlib/__init__.py b/openequivariance/openequivariance/_torch/extlib/__init__.py index 065f8b7f..84d81fe6 100644 --- a/openequivariance/openequivariance/_torch/extlib/__init__.py +++ b/openequivariance/openequivariance/_torch/extlib/__init__.py @@ -53,11 +53,13 @@ def _detect_backend(): IS_HIP = BACKEND == "hip" IS_SYCL = BACKEND == "sycl" -# The SYCL backend needs torch.library.register_autocast, added in 2.7. The -# XPU device itself and cpp_extension's SYCL support arrive in 2.6. -if IS_SYCL and Version(torch.__version__) < Version("2.7"): +# The SYCL backend needs at least the 2.7 APIs (torch.library.register_autocast; +# 2.6 for the XPU device and cpp_extension's SYCL support), but it is only +# tested against the 2.8 floor the rest of the project already requires for +# AOTI and export, so that is what is enforced. +if IS_SYCL and Version(torch.__version__) < Version("2.8"): raise RuntimeError( - f"The SYCL backend requires PyTorch >= 2.7, found {torch.__version__}." + f"The SYCL backend requires PyTorch >= 2.8, found {torch.__version__}." ) # The torch device type that tensors passed to the kernels must live on. From cb16f22860d9dcaa01feaf7265ad995ac0fd710b Mon Sep 17 00:00:00 2001 From: Abhishek Bagusetty Date: Thu, 3 Sep 2026 12:44:41 -0500 Subject: [PATCH 05/12] docs: correct the AOTI PyTorch floor to 2.10 The stable-ABI headers the extension includes are not all present before 2.10, which the precompiled-extension check in extlib has required all along. The SYCL guard now enforces the same floor. --- .github/workflows/requirements_sycl_ci.txt | 2 +- CHANGELOG.md | 9 ++++++++- docs/installation.rst | 2 +- .../openequivariance/_torch/extlib/__init__.py | 10 +++++----- 4 files changed, 15 insertions(+), 8 deletions(-) diff --git a/.github/workflows/requirements_sycl_ci.txt b/.github/workflows/requirements_sycl_ci.txt index a8b0a39c..5d01c496 100644 --- a/.github/workflows/requirements_sycl_ci.txt +++ b/.github/workflows/requirements_sycl_ci.txt @@ -1,4 +1,4 @@ -# The SYCL backend follows the project's existing PyTorch >= 2.8 floor. +# The SYCL backend follows the project's existing PyTorch >= 2.10 floor. # Pinned to the newest XPU wheel this was verified against. --extra-index-url https://download.pytorch.org/whl/xpu numpy==2.2.5 diff --git a/CHANGELOG.md b/CHANGELOG.md index 47870c4a..204b6dda 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,7 +8,7 @@ - A CI job that verifies the SYCL extension builds. Like the CUDA job it runs on a GPU-less runner and only exercises the build, since GitHub offers no Intel GPU runner. -- The SYCL backend requires PyTorch >= 2.8, the floor the project already +- The SYCL backend requires PyTorch >= 2.10, the floor the project already sets for AOTI and export, and raises at import on an older version. **Changed**: @@ -21,6 +21,13 @@ (long double) literal suffix. The values are unchanged — a hex float literal is already exactly a double — but SPIR-V targets reject the suffix. +**Fixed**: +- The documented PyTorch floor for AOTI and export is 2.10, not 2.8. The + stable-ABI headers the extension includes (`torch/csrc/stable/tensor_struct.h`, + `torch/csrc/stable/accelerator.h`, `torch/headeronly/core/DeviceType.h`) are + not all present before 2.10, which is the version the precompiled-extension + check in `extlib` has required all along. + ### v0.7.0 (2026-09-10) **Added**: - Public XLA FFI registration provider diff --git a/docs/installation.rst b/docs/installation.rst index 1a401ab2..17feab14 100644 --- a/docs/installation.rst +++ b/docs/installation.rst @@ -8,7 +8,7 @@ Installation You need the following to install OpenEquivariance: - A Linux system equipped with an NVIDIA / AMD / Intel graphics card. -- Either PyTorch >= 2.4 (>= 2.8 for AOTI, export, and Intel GPUs), or +- Either PyTorch >= 2.4 (>= 2.10 for AOTI, export, and Intel GPUs), or JAX>0.5.0 with CUDA, RocM, or XPU support. - GCC 9+ and the CUDA / HIP toolkit, or the oneAPI DPC++ compiler (``icpx``) for Intel GPUs. The command diff --git a/openequivariance/openequivariance/_torch/extlib/__init__.py b/openequivariance/openequivariance/_torch/extlib/__init__.py index 84d81fe6..213d2688 100644 --- a/openequivariance/openequivariance/_torch/extlib/__init__.py +++ b/openequivariance/openequivariance/_torch/extlib/__init__.py @@ -53,13 +53,13 @@ def _detect_backend(): IS_HIP = BACKEND == "hip" IS_SYCL = BACKEND == "sycl" -# The SYCL backend needs at least the 2.7 APIs (torch.library.register_autocast; -# 2.6 for the XPU device and cpp_extension's SYCL support), but it is only -# tested against the 2.8 floor the rest of the project already requires for +# The SYCL backend's own APIs arrive earlier (2.6 for the XPU device and +# cpp_extension's SYCL support, 2.7 for torch.library.register_autocast), but +# it is only tested against the 2.10 floor the rest of the project requires for # AOTI and export, so that is what is enforced. -if IS_SYCL and Version(torch.__version__) < Version("2.8"): +if IS_SYCL and Version(torch.__version__) < Version("2.10"): raise RuntimeError( - f"The SYCL backend requires PyTorch >= 2.8, found {torch.__version__}." + f"The SYCL backend requires PyTorch >= 2.10, found {torch.__version__}." ) # The torch device type that tensors passed to the kernels must live on. From 4f25be5420ed663f9d2a84d616b08b006f03a67d Mon Sep 17 00:00:00 2001 From: Abhishek Bagusetty Date: Thu, 3 Sep 2026 12:50:16 -0500 Subject: [PATCH 06/12] docs: shorten the CHANGELOG entry for the AOTI version floor --- CHANGELOG.md | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 204b6dda..c50608df 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -22,11 +22,7 @@ literal is already exactly a double — but SPIR-V targets reject the suffix. **Fixed**: -- The documented PyTorch floor for AOTI and export is 2.10, not 2.8. The - stable-ABI headers the extension includes (`torch/csrc/stable/tensor_struct.h`, - `torch/csrc/stable/accelerator.h`, `torch/headeronly/core/DeviceType.h`) are - not all present before 2.10, which is the version the precompiled-extension - check in `extlib` has required all along. +- The documented PyTorch floor for AOTI and export is 2.10, not 2.8. ### v0.7.0 (2026-09-10) **Added**: From 353e326c0356d584c1cee54367f9d3a5d02588ed Mon Sep 17 00:00:00 2001 From: Abhishek Bagusetty Date: Thu, 3 Sep 2026 12:51:32 -0500 Subject: [PATCH 07/12] docs: simplify the CHANGELOG entry for the float64 literal suffix --- CHANGELOG.md | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index c50608df..7af7d76e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -18,8 +18,7 @@ - `JITKernel::execute` also takes the kernel argument sizes, which SYCL requires to launch with raw arguments. CUDA and HIP ignore them. - float64 Clebsch-Gordon coefficients are emitted without the `L` - (long double) literal suffix. The values are unchanged — a hex float - literal is already exactly a double — but SPIR-V targets reject the suffix. + (long double) literal suffix, which SPIR-V rejects. The values are unchanged. **Fixed**: - The documented PyTorch floor for AOTI and export is 2.10, not 2.8. From b7381ce8385b1d025d12b7a1f619a1a5c7cfde37 Mon Sep 17 00:00:00 2001 From: Abhishek Bagusetty Date: Wed, 16 Sep 2026 13:29:04 -0500 Subject: [PATCH 08/12] fix: drop the dead __threadfence_block shim from sycl_compat.cuh Upstream #217 replaced HIP's __threadfence_block() with the amdgcn wave barrier, so no template emits the token any more. --- openequivariance/openequivariance/templates/sycl_compat.cuh | 1 - 1 file changed, 1 deletion(-) diff --git a/openequivariance/openequivariance/templates/sycl_compat.cuh b/openequivariance/openequivariance/templates/sycl_compat.cuh index 34468ee5..4c3cbb16 100644 --- a/openequivariance/openequivariance/templates/sycl_compat.cuh +++ b/openequivariance/openequivariance/templates/sycl_compat.cuh @@ -62,7 +62,6 @@ static inline void oeq_syncthreads() { } #define __syncthreads() oeq_syncthreads() -#define __threadfence_block() oeq_syncwarp() // --------------------------------------------------------------------------- // Warp-level primitives From 5fe3c7764f7f46dd59a1cb14a32794d38cc19893 Mon Sep 17 00:00:00 2001 From: Abhishek Bagusetty Date: Mon, 21 Sep 2026 10:38:31 -0500 Subject: [PATCH 09/12] address PR comments --- .github/workflows/requirements_sycl_ci.txt | 2 - .github/workflows/verify_extension_build.yml | 7 +- CHANGELOG.md | 23 ----- README.md | 11 +-- docs/index.rst | 6 +- openequivariance/CMakeLists.txt | 10 -- .../_torch/extlib/__init__.py | 31 +------ .../SymmetricContraction.py | 8 +- .../core/ComputationSchedule.py | 6 +- .../openequivariance/core/utils.py | 11 +-- .../extension/backend/backend_cuda.hpp | 6 +- .../extension/backend/backend_hip.hpp | 6 +- .../extension/backend/backend_sycl.hpp | 46 +++------- .../extension/convolution.hpp | 14 +-- .../openequivariance/extension/group_mm.hpp | 75 --------------- .../extension/kernel_args.hpp | 34 +++---- .../extension/libtorch_tp_jit.cpp | 1 - .../extension/libtorch_tp_jit_stable.cpp | 13 +-- .../extension/stubs/stream.cpp | 5 - .../extension/tensorproducts.hpp | 6 +- .../openequivariance/extension/torch_core.hpp | 36 ++++---- .../openequivariance/jax/TensorProduct.py | 4 +- .../openequivariance/jax/TensorProductConv.py | 2 +- .../openequivariance/jax/extlib/__init__.py | 7 +- .../openequivariance/templates/jinja_utils.py | 20 ++-- .../openequivariance/templates/macros.jinja | 5 +- .../templates/sycl_compat.cuh | 91 ++++++------------- tests/batch_test.py | 4 +- tests/conftest.py | 20 ++-- tests/conv_test.py | 4 +- tests/example_test.py | 24 ++--- tests/export_test.py | 38 ++++---- tests/input_validation_test.py | 22 ++--- tests/stream_test.py | 35 ++++--- tests/symmetric_contraction_test.py | 13 ++- 35 files changed, 208 insertions(+), 438 deletions(-) diff --git a/.github/workflows/requirements_sycl_ci.txt b/.github/workflows/requirements_sycl_ci.txt index 5d01c496..08e7a70b 100644 --- a/.github/workflows/requirements_sycl_ci.txt +++ b/.github/workflows/requirements_sycl_ci.txt @@ -1,5 +1,3 @@ -# The SYCL backend follows the project's existing PyTorch >= 2.10 floor. -# Pinned to the newest XPU wheel this was verified against. --extra-index-url https://download.pytorch.org/whl/xpu numpy==2.2.5 torch==2.12.1+xpu diff --git a/.github/workflows/verify_extension_build.yml b/.github/workflows/verify_extension_build.yml index 234405ca..358d9340 100644 --- a/.github/workflows/verify_extension_build.yml +++ b/.github/workflows/verify_extension_build.yml @@ -56,17 +56,14 @@ jobs: cache: 'pip' cache-dependency-path: '**/requirements_sycl_ci.txt' - # The SYCL backend compiles the extension with the oneAPI DPC++ driver - # (icpx), so the compiler and the oneMKL headers both have to be present. - # No Intel GPU is attached to this runner; only the build is exercised. - - name: Install oneAPI DPC++ and oneMKL + - name: Install oneAPI DPC++ run: | wget -qO- https://apt.repos.intel.com/intel-gpg-keys/GPG-PUB-KEY-INTEL-SW-PRODUCTS.PUB \ | gpg --dearmor | sudo tee /usr/share/keyrings/oneapi-archive-keyring.gpg > /dev/null echo "deb [signed-by=/usr/share/keyrings/oneapi-archive-keyring.gpg] https://apt.repos.intel.com/oneapi all main" \ | sudo tee /etc/apt/sources.list.d/oneAPI.list sudo apt-get update - sudo apt-get install -y intel-oneapi-compiler-dpcpp-cpp intel-oneapi-mkl-devel + sudo apt-get install -y intel-oneapi-compiler-dpcpp-cpp - name: Install dependencies run: | diff --git a/CHANGELOG.md b/CHANGELOG.md index 7af7d76e..1c9f9847 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,28 +1,5 @@ ## Latest Changes -**Added**: -- SYCL backend, bringing support for Intel GPUs through PyTorch's - `xpu` device. Kernels are generated as SYCL free functions and compiled - at runtime through the oneAPI kernel compiler extension. The backend is - selected automatically from the active PyTorch build. -- A CI job that verifies the SYCL extension builds. Like the CUDA job it - runs on a GPU-less runner and only exercises the build, since GitHub - offers no Intel GPU runner. -- The SYCL backend requires PyTorch >= 2.10, the floor the project already - sets for AOTI and export, and raises at import on an older version. - -**Changed**: -- The kernel backend is now identified by a string (`"cuda"` / `"hip"` / - `"sycl"`) rather than an `is_hip` boolean, in the Jinja environment and - the `LoopUnrollTP` / `LoopUnrollConv` constructors. -- `JITKernel::execute` also takes the kernel argument sizes, which SYCL - requires to launch with raw arguments. CUDA and HIP ignore them. -- float64 Clebsch-Gordon coefficients are emitted without the `L` - (long double) literal suffix, which SPIR-V rejects. The values are unchanged. - -**Fixed**: -- The documented PyTorch floor for AOTI and export is 2.10, not 2.8. - ### v0.7.0 (2026-09-10) **Added**: - Public XLA FFI registration provider diff --git a/README.md b/README.md index a750e6a4..3efe59a7 100644 --- a/README.md +++ b/README.md @@ -6,9 +6,8 @@ [[JAX Examples]](#jax-examples) [[Citation and Acknowledgements]](#citation-and-acknowledgements) -OpenEquivariance is a CUDA, HIP, and SYCL kernel generator for the Clebsch-Gordon tensor product, +OpenEquivariance is a CUDA and HIP kernel generator for the Clebsch-Gordon tensor product, a key kernel in rotation-equivariant deep neural networks. -It targets NVIDIA, AMD, and Intel GPUs. It implements some of the tensor products that [e3nn](https://e3nn.org/) supports commonly found in graph neural networks @@ -21,11 +20,6 @@ and GCC 9+ available before installing our package via pip install openequivariance ``` -On Intel GPUs, install a PyTorch build with XPU support and make the -oneAPI DPC++ compiler (`icpx`) available on your `PATH`; the kernels are -compiled at runtime through SYCL. Tensors live on the `xpu` device there -instead of `cuda`. - We provide up to an order of magnitude acceleration over e3nn perform on par with the latest version of [NVIDIA cuEquivariance](https://github.com/NVIDIA/cuEquivariance), which has a closed-source kernel package. @@ -80,8 +74,7 @@ print(torch.norm(Z)) ``` And here's the same tensor product using openequivariance. We require that your -tensors are stored on a GPU device for this to work -(``cuda`` for NVIDIA and AMD GPUs, ``xpu`` for Intel GPUs): +tensors are stored on a CUDA device for this to work: ```python import openequivariance as oeq diff --git a/docs/index.rst b/docs/index.rst index f58c14bc..3d5f055b 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -6,15 +6,15 @@ OpenEquivariance ============================== -`OpenEquivariance `_ is a CUDA, -HIP, and SYCL kernel generator for the Clebsch-Gordon +`OpenEquivariance `_ is a CUDA and +HIP kernel generator for the Clebsch-Gordon tensor product, a key kernel in equivariant graph neural networks. We offer an identical interface to e3nn and produce the same results (up to numerical roundoff). Our package exhibits up to an order of magnitude speedup over e3nn and competitive performance with NVIDIA's cuEquivariance. Here, you can find our API reference, installation instructions, -and troubleshooting guide. We support NVIDIA, AMD, and Intel GPUs through +and troubleshooting guide. We support for both NVIDIA and AMD GPUs through our PyTorch interface, including support for JITScript compilation accessible from C++. diff --git a/openequivariance/CMakeLists.txt b/openequivariance/CMakeLists.txt index df05b1ee..3f3ad3dc 100644 --- a/openequivariance/CMakeLists.txt +++ b/openequivariance/CMakeLists.txt @@ -173,7 +173,6 @@ if(hip_FOUND) add_stable_extension(torch_stable_hip HIP_BACKEND "${HIP_LINK_LIBS}") endif() -# SYCL: detected via the compiler accepting -fsycl (IntelLLVM / icpx). set(OEQ_SYCL_FOUND FALSE) if(CMAKE_CXX_COMPILER_ID MATCHES "IntelLLVM") set(OEQ_SYCL_FOUND TRUE) @@ -196,21 +195,12 @@ if(OEQ_SYCL_FOUND) target_compile_definitions(sycl_stub_lib PRIVATE SYCL_BACKEND=1) - find_package(MKL QUIET COMPONENTS SYCL) - if(TARGET MKL::MKL_SYCL::BLAS) - set(SYCL_BLAS_LIB MKL::MKL_SYCL::BLAS) - else() - set(SYCL_BLAS_LIB mkl_sycl_blas) - endif() - set(SYCL_LINK_LIBS sycl_stub_lib sycl - ${SYCL_BLAS_LIB} ) add_stable_extension(oeq_stable_sycl SYCL_BACKEND "${SYCL_LINK_LIBS}") - # -fsycl is required on both the compile and the link line. foreach(tgt oeq_stable_sycl oeq_stable_sycl_aoti) target_compile_options(${tgt} PRIVATE -fsycl) target_link_options(${tgt} PRIVATE -fsycl) diff --git a/openequivariance/openequivariance/_torch/extlib/__init__.py b/openequivariance/openequivariance/_torch/extlib/__init__.py index 213d2688..ef89048a 100644 --- a/openequivariance/openequivariance/_torch/extlib/__init__.py +++ b/openequivariance/openequivariance/_torch/extlib/__init__.py @@ -24,16 +24,7 @@ def _detect_backend(): - """ - Determines which GPU backend this PyTorch build targets. - - Returns one of ``"cuda"``, ``"hip"`` or ``"sycl"``. HIP builds report a - ``torch.version.cuda`` of ``None``, so HIP must be tested first. - - All three checks are build-time properties of the PyTorch install, not - runtime device queries, so importing works on a machine with no - accelerator attached (a CI builder, for instance). - """ + """HIP builds report a ``torch.version.cuda`` of None, so HIP is tested first.""" if torch.version.hip: return "hip" if torch.version.cuda: @@ -53,16 +44,11 @@ def _detect_backend(): IS_HIP = BACKEND == "hip" IS_SYCL = BACKEND == "sycl" -# The SYCL backend's own APIs arrive earlier (2.6 for the XPU device and -# cpp_extension's SYCL support, 2.7 for torch.library.register_autocast), but -# it is only tested against the 2.10 floor the rest of the project requires for -# AOTI and export, so that is what is enforced. if IS_SYCL and Version(torch.__version__) < Version("2.10"): raise RuntimeError( f"The SYCL backend requires PyTorch >= 2.10, found {torch.__version__}." ) -# The torch device type that tensors passed to the kernels must live on. DEVICE_TYPE = "xpu" if IS_SYCL else "cuda" @@ -166,8 +152,6 @@ def load_jit_extension(): extra_link_args.append("-Wl,-rpath," + torch_libs) extra_cflags.append("-DHIP_BACKEND") elif BACKEND == "sycl": - # torch.utils.cpp_extension compiles with $CXX (default c++), - # which must be the oneAPI DPC++ driver for -fsycl to work. import shutil cxx = os.environ.get("CXX", "") @@ -180,24 +164,13 @@ def load_jit_extension(): return os.environ["CXX"] = "icpx" - # SYCL sources must be compiled and linked by the SYCL compiler - # driver; -fsycl is required on both the compile and link lines. extra_cflags.extend(["-fsycl", "-DSYCL_BACKEND"]) - extra_link_args.extend( - ["-fsycl", "-ltorch_xpu", "-lc10_xpu", "-lmkl_sycl_blas"] - ) + extra_link_args.extend(["-fsycl", "-ltorch_xpu", "-lc10_xpu"]) for lib_dir in library_paths("xpu"): extra_link_args.append("-Wl,-rpath," + lib_dir) extra_link_args.append("-L" + lib_dir) - mkl_root = os.environ.get("MKLROOT") - if mkl_root: - mkl_lib = os.path.join(mkl_root, "lib") - extra_link_args.append("-L" + mkl_lib) - extra_link_args.append("-Wl,-rpath," + mkl_lib) - extra_include_dirs.append(os.path.join(mkl_root, "include")) - torch_sources = [oeq_root + "/extension/" + src for src in torch_sources] include_dirs = ( [oeq_root + "/extension/" + d for d in include_dirs] diff --git a/openequivariance/openequivariance/_torch/symmetric_contraction/SymmetricContraction.py b/openequivariance/openequivariance/_torch/symmetric_contraction/SymmetricContraction.py index 83c56933..833d1bc8 100644 --- a/openequivariance/openequivariance/_torch/symmetric_contraction/SymmetricContraction.py +++ b/openequivariance/openequivariance/_torch/symmetric_contraction/SymmetricContraction.py @@ -123,6 +123,11 @@ def U_matrix_real( class GroupMM: def __init__(self, dtype, num_elements, batch_size): + # group_gemm is not built for SYCL, so the operator does not exist there. + if extlib.DEVICE_TYPE == "xpu": + raise NotImplementedError( + "Symmetric contraction is not supported on the SYCL backend." + ) self.num_elements = num_elements self.batch_size = batch_size @@ -406,6 +411,7 @@ def backward(ctx, grad_output): ) -if extlib.BUILT_EXTENSION: +# group_gemm is not built for SYCL, so the operator does not exist there. +if extlib.BUILT_EXTENSION and extlib.DEVICE_TYPE != "xpu": register_torch_fakes() register_autograd() diff --git a/openequivariance/openequivariance/core/ComputationSchedule.py b/openequivariance/openequivariance/core/ComputationSchedule.py index 304db645..5aabb502 100644 --- a/openequivariance/openequivariance/core/ComputationSchedule.py +++ b/openequivariance/openequivariance/core/ComputationSchedule.py @@ -62,10 +62,8 @@ def __init__(self, src_irreps, src_views, idxs): class CGTensor: def __init__(self, l1, l2, l3, normalization_factor, dtype): - # A hex float literal is a double by default and represents the value - # exactly, so float64 needs no suffix. An "L" (long double) suffix - # would not change the value but is rejected by SPIR-V targets, which - # have no 128-bit float type. + # A hex float literal is already exactly a double, so float64 needs no + # suffix; an "L" suffix would be rejected by SPIR-V targets. suffix_map = {np.float32: "f", np.float64: ""} tensor = wigner_3j(l1, l2, l3) diff --git a/openequivariance/openequivariance/core/utils.py b/openequivariance/openequivariance/core/utils.py index 3afb5b31..11b660aa 100644 --- a/openequivariance/openequivariance/core/utils.py +++ b/openequivariance/openequivariance/core/utils.py @@ -173,8 +173,7 @@ def benchmark(func, num_warmup, num_iter, mode="gpu_time", kernel_names=[]): else: from torch.profiler import ProfilerActivity, profile, record_function - # The profiler activity is per-accelerator: XPU kernels are not - # recorded under the CUDA activity. + # Profiler activity is per-accelerator. accelerator_activity = ( ProfilerActivity.XPU if accelerator_device_type() == "xpu" @@ -267,13 +266,7 @@ def transpose_irrep_layout( def accelerator_device_type(): - """ - Returns the ``torch`` device type the kernels run on: ``"xpu"`` for the - SYCL backend, ``"cuda"`` for CUDA and HIP (PyTorch exposes HIP tensors - under the ``cuda`` device type). - - Imported lazily so that the backend-agnostic core does not pull in torch. - """ + """Imported lazily so the backend-agnostic core does not pull in torch.""" from openequivariance._torch.extlib import DEVICE_TYPE return DEVICE_TYPE diff --git a/openequivariance/openequivariance/extension/backend/backend_cuda.hpp b/openequivariance/openequivariance/extension/backend/backend_cuda.hpp index 9b427412..9dd224c8 100644 --- a/openequivariance/openequivariance/extension/backend/backend_cuda.hpp +++ b/openequivariance/openequivariance/extension/backend/backend_cuda.hpp @@ -323,10 +323,8 @@ class __attribute__((visibility("default"))) CUJITKernel { } } - void execute(int kernel_id, void* args[], const size_t arg_sizes[], - size_t num_args, KernelLaunchConfig config) { - (void) arg_sizes; // The CUDA driver infers argument sizes from the kernel signature. - (void) num_args; + void execute(int kernel_id, void* args[], [[maybe_unused]] const size_t arg_sizes[], + [[maybe_unused]] size_t num_args, KernelLaunchConfig config) { if(kernel_id >= kernels.size()) throw std::logic_error("Kernel index out of range!"); diff --git a/openequivariance/openequivariance/extension/backend/backend_hip.hpp b/openequivariance/openequivariance/extension/backend/backend_hip.hpp index 4066daa0..a68e5433 100644 --- a/openequivariance/openequivariance/extension/backend/backend_hip.hpp +++ b/openequivariance/openequivariance/extension/backend/backend_hip.hpp @@ -292,10 +292,8 @@ class __attribute__((visibility("default"))) HIPJITKernel { // Ignore for AMD GPUs } - void execute(int kernel_id, void* args[], const size_t arg_sizes[], - size_t num_args, KernelLaunchConfig config) { - (void) arg_sizes; // The HIP driver infers argument sizes from the kernel signature. - (void) num_args; + void execute(int kernel_id, void* args[], [[maybe_unused]] const size_t arg_sizes[], + [[maybe_unused]] size_t num_args, KernelLaunchConfig config) { int device_id; HIP_ERRCHK(hipGetDevice(&device_id)); if(device_id != kernels->device) { kernels.reset(); diff --git a/openequivariance/openequivariance/extension/backend/backend_sycl.hpp b/openequivariance/openequivariance/extension/backend/backend_sycl.hpp index b1dd2fc8..f3112a27 100644 --- a/openequivariance/openequivariance/extension/backend/backend_sycl.hpp +++ b/openequivariance/openequivariance/extension/backend/backend_sycl.hpp @@ -16,20 +16,12 @@ using namespace std; namespace syclex = sycl::ext::oneapi::experimental; -/* -* SYCL streams are queues. Unlike CUDA / HIP, a sycl::queue is a -* reference-counted handle rather than an opaque pointer, so the "stream" type -* is a pointer to the queue owned by the caller (PyTorch). -*/ using Stream = sycl::queue *; -// Defined by the translation unit that binds this header to the framework -// (PyTorch / JAX). Returns the queue of the framework's current stream. Stream get_current_stream(); -// Returns the queue the kernels should be submitted to. A null stream means -// "no queue was supplied", in which case we fall back to the framework's -// current stream, and only then to a process-wide default queue. +// A null stream falls back to the framework's current stream, then to a +// process-wide default queue. inline sycl::queue &resolve_queue(Stream stream) { if (stream != nullptr) { return *stream; @@ -68,11 +60,8 @@ class SYCL_Allocator { } }; -/* -* SYCL has no direct equivalent of cudaEvent elapsed time that works without -* enabling profiling on the queue, so the timer brackets a wall-clock interval -* around a queue synchronization. -*/ +// No cudaEvent equivalent works without enabling profiling on the queue, so +// this brackets a wall-clock interval around a queue synchronization. class GPUTimer { std::chrono::time_point start_time; @@ -129,8 +118,6 @@ class __attribute__((visibility("default"))) DeviceProp { multiprocessorCount = static_cast(dev.get_info()); - // A SYCL sub-group is the analogue of a CUDA warp / HIP wavefront. Pick - // the largest supported size that the kernel generator can target. auto sg_sizes = dev.get_info(); warpsize = 32; if (!sg_sizes.empty()) { @@ -146,8 +133,7 @@ class __attribute__((visibility("default"))) DeviceProp { static_cast(dev.get_info()); maxSharedMemoryPerMultiprocessor = maxSharedMemPerBlock; - // SYCL exposes no compute-capability equivalent. These fields exist - // only for parity with the CUDA backend and are unused on SYCL. + // Unused on SYCL; present for parity with the CUDA backend. major = 0; minor = 0; } @@ -177,15 +163,8 @@ class __attribute__((visibility("default"))) KernelLaunchConfig { { } }; -/* -* Runtime compilation uses the SYCL kernel_compiler extension with -* source_language::sycl, documented at -* https://github.com/intel/llvm/blob/sycl/sycl/doc/extensions/experimental/sycl_ext_oneapi_kernel_compiler_sycl.asciidoc -* -* The generated kernels are free functions marked with nd_range_kernel, so they -* are launched with raw (untyped) arguments exactly like cuLaunchKernel takes a -* void* array. -*/ +// Uses the SYCL kernel_compiler extension with source_language::sycl: +// https://github.com/intel/llvm/blob/sycl/sycl/doc/extensions/experimental/sycl_ext_oneapi_kernel_compiler_sycl.asciidoc class __attribute__((visibility("default"))) SYCLJITKernel { private: bool compiled = false; @@ -219,7 +198,6 @@ class __attribute__((visibility("default"))) SYCLJITKernel { string kernel_name = kernel_names_i[kernel]; vector &template_params = template_param_list[kernel]; - // Step 1: Generate kernel names from the template parameters if(template_params.size() == 0) { kernel_names.push_back(kernel_name); } @@ -236,8 +214,8 @@ class __attribute__((visibility("default"))) SYCLJITKernel { } } - // Build against the context the kernels will actually run in, so the - // resulting bundle is valid for every device that context spans. + // Build against the context the kernels run in, so the bundle is valid + // for every device that context spans. sycl::queue &q = resolve_queue(nullptr); sycl::context build_context = q.get_context(); @@ -278,10 +256,8 @@ class __attribute__((visibility("default"))) SYCLJITKernel { } void set_max_smem(int kernel_id, uint32_t max_smem_bytes) { - // Shared (local) memory is declared statically inside the generated - // kernel via work_group_static, so there is no opt-in to perform here. - // Validate the request against the device limit so an oversubscription - // fails with a clear message instead of at launch. + // Local memory is declared statically in the generated kernel, so there + // is nothing to opt into; just validate against the device limit. if(!compiled) throw std::logic_error("JIT object has not been compiled!"); if(static_cast(kernel_id) >= kernels.size()) diff --git a/openequivariance/openequivariance/extension/convolution.hpp b/openequivariance/openequivariance/extension/convolution.hpp index 57211dec..55a6fdd2 100644 --- a/openequivariance/openequivariance/extension/convolution.hpp +++ b/openequivariance/openequivariance/extension/convolution.hpp @@ -91,12 +91,12 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { ConvData conv_data = {rows, cols, nnz, node_count}; - auto args = make_kernel_args(L1_in, L2_in, weights, L3_out, conv_data, workspace); + auto args = KernelArgs(L1_in, L2_in, weights, L3_out, conv_data, workspace); jit.execute(0, args.data(), args.arg_sizes(), args.count(), with_stream(forward_config_ref, stream)); if(reinterpret_cast(workspace) != 0) { - auto fixup_args = make_kernel_args(workspace, L3_out); + auto fixup_args = KernelArgs(workspace, L3_out); KernelLaunchConfig fixup_config( forward_config_ref.num_blocks, @@ -122,14 +122,14 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { Stream stream) { ConvData conv_data = {rows, cols, nnz, node_count}; - auto args = make_kernel_args(L1_in, L1_grad, L2_in, L2_grad, weight, + auto args = KernelArgs(L1_in, L1_grad, L2_in, L2_grad, weight, weight_grad, L3_grad, conv_data, workspace, transpose_perm); jit.execute(1, args.data(), args.arg_sizes(), args.count(), with_stream(backward_config_ref, stream)); if(reinterpret_cast(workspace) != 0) { - auto fixup_args = make_kernel_args(workspace, L1_grad); + auto fixup_args = KernelArgs(workspace, L1_grad); KernelLaunchConfig fixup_config( backward_config_ref.num_blocks, @@ -153,14 +153,14 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { Stream stream) { ConvData conv_data = {rows, cols, nnz, node_count}; - auto args = make_kernel_args( + auto args = KernelArgs( L1_in, L2_in, W, L3_grad, L1_dgrad, L2_dgrad, w_dgrad, L1_grad, L2_grad, W_grad, L3_dgrad, conv_data, wspace, transpose_perm); jit.execute(4, args.data(), args.arg_sizes(), args.count(), with_stream(forward_config_ref, stream)); if(reinterpret_cast(wspace) != 0) { - auto fixup_args = make_kernel_args(wspace, L3_dgrad); + auto fixup_args = KernelArgs(wspace, L3_dgrad); KernelLaunchConfig fixup_config( forward_config_ref.num_blocks, forward_config_ref.num_threads, @@ -174,7 +174,7 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { jit.execute(5, args.data(), args.arg_sizes(), args.count(), with_stream(double_backward_config_ref, stream)); if(reinterpret_cast(wspace) != 0) { - auto fixup_args = make_kernel_args(wspace, L1_grad); + auto fixup_args = KernelArgs(wspace, L1_grad); KernelLaunchConfig fixup_config( double_backward_config_ref.num_blocks, double_backward_config_ref.num_threads, diff --git a/openequivariance/openequivariance/extension/group_mm.hpp b/openequivariance/openequivariance/extension/group_mm.hpp index e2d0196a..19249c39 100644 --- a/openequivariance/openequivariance/extension/group_mm.hpp +++ b/openequivariance/openequivariance/extension/group_mm.hpp @@ -3,7 +3,6 @@ #include #include #include -#include #ifdef CUDA_BACKEND #include "cublas_v2.h" @@ -29,35 +28,6 @@ } ~BlasHandle() { rocblas_destroy_handle(handle); } }; -#elif defined(SYCL_BACKEND) - #include - #include - - // oneMKL takes the queue per call rather than a persistent handle, so this - // exists only to keep a single shape across the three backends. - struct BlasHandle { - BlasHandle() = default; - ~BlasHandle() = default; - }; - - // A small shared-USM array. oneMKL's pointer-array GEMM reads the pointer - // lists from the device, so they cannot live on the host stack. - template - class UsmArray { - sycl::queue &q_; - T *ptr_; - public: - UsmArray(sycl::queue &q, size_t n) : q_(q), - ptr_(sycl::malloc_shared(n, q)) { - if (ptr_ == nullptr) - throw std::runtime_error("Shared USM allocation failed!"); - } - ~UsmArray() { sycl::free(ptr_, q_); } - UsmArray(const UsmArray &) = delete; - UsmArray &operator=(const UsmArray &) = delete; - T &operator[](size_t i) { return ptr_[i]; } - T *get() { return ptr_; } - }; #endif inline BlasHandle& get_blas_handle() { @@ -83,8 +53,6 @@ void group_gemm_blas(void* A_raw, void* B_raw, void* C_raw, cublasOperation_t transa, transb; #elif defined(HIP_BACKEND) rocblas_operation transa, transb; -#elif defined(SYCL_BACKEND) - oneapi::mkl::transpose transa, transb; #endif if (ragged_inner == 0) { @@ -99,9 +67,6 @@ void group_gemm_blas(void* A_raw, void* B_raw, void* C_raw, transa = CUBLAS_OP_T; transb = CUBLAS_OP_N; #elif defined(HIP_BACKEND) transa = rocblas_operation_transpose; transb = rocblas_operation_none; -#elif defined(SYCL_BACKEND) - transa = oneapi::mkl::transpose::trans; - transb = oneapi::mkl::transpose::nontrans; #endif } else { M = k; K = static_cast(ragged_counts[i]); N = m; @@ -115,9 +80,6 @@ void group_gemm_blas(void* A_raw, void* B_raw, void* C_raw, transa = CUBLAS_OP_N; transb = CUBLAS_OP_T; #elif defined(HIP_BACKEND) transa = rocblas_operation_none; transb = rocblas_operation_transpose; -#elif defined(SYCL_BACKEND) - transa = oneapi::mkl::transpose::nontrans; - transb = oneapi::mkl::transpose::trans; #endif } ragged_offset += ragged_counts[i]; @@ -173,43 +135,6 @@ void group_gemm_blas(void* A_raw, void* B_raw, void* C_raw, } if (stat != rocblas_status_success) throw std::logic_error("Grouped GEMM failed!"); -#elif defined(SYCL_BACKEND) - (void) blas; - // Submit onto the same queue the kernels use so the GEMM is - // ordered against them. - sycl::queue &q = resolve_queue(get_current_stream()); - - // oneMKL's strided batch API requires stride_c >= ldc * n, which - // this interleaved layout deliberately violates (consecutive - // matrices overlap in the batch dimension). The pointer-array form - // imposes no such constraint, so build explicit pointer lists. - // oneMKL reads these arrays on the device, so they live in shared - // USM rather than on the host stack. - UsmArray a_ptrs(q, batch_size), b_ptrs(q, batch_size); - UsmArray c_ptrs(q, batch_size); - for (int j = 0; j < batch_size; j++) { - a_ptrs[j] = A + static_cast(strideA) * j; - b_ptrs[j] = B + static_cast(strideB) * j; - c_ptrs[j] = C + static_cast(strideC) * j; - } - - int64_t m64 = M, n64 = N, k64 = K; - int64_t lda64 = lda, ldb64 = ldb, ldc64 = ldc; - int64_t group_size = batch_size; - - try { - oneapi::mkl::blas::column_major::gemm_batch( - q, &transa, &transb, &m64, &n64, &k64, &alpha, - a_ptrs.get(), &lda64, - b_ptrs.get(), &ldb64, - &beta, c_ptrs.get(), &ldc64, - 1, &group_size) - // The pointer arrays are freed when this scope exits, so - // the submission has to complete before then. - .wait(); - } catch (const sycl::exception &e) { - throw std::logic_error("Grouped GEMM failed: " + std::string(e.what())); - } #endif } } diff --git a/openequivariance/openequivariance/extension/kernel_args.hpp b/openequivariance/openequivariance/extension/kernel_args.hpp index d1ad64af..2c7dd517 100644 --- a/openequivariance/openequivariance/extension/kernel_args.hpp +++ b/openequivariance/openequivariance/extension/kernel_args.hpp @@ -3,30 +3,22 @@ #include #include -/* -* Packs kernel arguments as a list of (pointer, size) pairs. -* -* The CUDA and HIP driver APIs only need the pointer array, but SYCL launches -* free-function kernels with raw (untyped) arguments and therefore also needs -* the size of each argument. Collecting both here keeps a single call shape in -* the backend-independent code. -* -* As with the raw `void*[]` form this replaces, the caller must keep the -* referenced objects alive until the launch has been enqueued. -*/ template -struct KernelArgs { - std::array ptrs; - std::array sizes; +class KernelArgs { + std::array ptrs_; + std::array sizes_; - void **data() { return ptrs.data(); } - const size_t *arg_sizes() const { return sizes.data(); } +public: + template + explicit KernelArgs(Ts &...args) + : ptrs_{static_cast(&args)...}, sizes_{sizeof(Ts)...} { + static_assert(sizeof...(Ts) == N, "argument count mismatch"); + } + + void **data() { return ptrs_.data(); } + const size_t *arg_sizes() const { return sizes_.data(); } static constexpr size_t count() { return N; } }; template -inline KernelArgs make_kernel_args(Ts &...args) { - return KernelArgs{ - {static_cast(&args)...}, - {sizeof(Ts)...}}; -} +KernelArgs(Ts &...) -> KernelArgs; diff --git a/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp b/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp index 560068ab..c7f9bdb6 100644 --- a/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp +++ b/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp @@ -86,7 +86,6 @@ Stream get_current_stream() { return c10::hip::getCurrentHIPStream(); #endif #ifdef SYCL_BACKEND - // The queue is owned by PyTorch and outlives the kernel launch. return &c10::xpu::getCurrentXPUStream().queue(); #endif } diff --git a/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp b/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp index 6372710c..dc8ec85b 100644 --- a/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp +++ b/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp @@ -1,8 +1,4 @@ -#ifdef SYCL_BACKEND - #define USE_XPU -#else - #define USE_CUDA -#endif +#define USE_CUDA #include #include @@ -16,7 +12,9 @@ #include #include #ifdef SYCL_BACKEND - #include + // Declared in shim_xpu.h behind USE_XPU; declared here so the SYCL build + // does not have to define that macro as well as SYCL_BACKEND. + extern "C" AOTITorchError aoti_torch_get_current_sycl_queue(void** ret_queue); #endif @@ -77,7 +75,6 @@ Stream get_current_stream() { void* stream_ptr = nullptr; #ifdef SYCL_BACKEND - // Returns the sycl::queue* backing the current XPU stream. TORCH_ERROR_CODE_CHECK(aoti_torch_get_current_sycl_queue(&stream_ptr)); #else auto device_idx = torch::stable::accelerator::getCurrentDeviceIndex(); @@ -89,7 +86,7 @@ Stream get_current_stream() { bool tensor_is_on_gpu(const Tensor &tensor) { #ifdef SYCL_BACKEND - // The stable Tensor has no is_xpu(), so compare the device type directly. + // The stable Tensor has no is_xpu(). int32_t device_type; TORCH_ERROR_CODE_CHECK( aoti_torch_get_device_type(tensor.get(), &device_type)); diff --git a/openequivariance/openequivariance/extension/stubs/stream.cpp b/openequivariance/openequivariance/extension/stubs/stream.cpp index 310944c7..16280c13 100644 --- a/openequivariance/openequivariance/extension/stubs/stream.cpp +++ b/openequivariance/openequivariance/extension/stubs/stream.cpp @@ -1,11 +1,6 @@ #include #include -/* -* Weak stand-ins for the stream accessors the AOTI shim declares. The real -* symbols come from libtorch_cuda / libtorch_xpu at load time; these exist so -* the extension links without a hard dependency on the accelerator runtime. -*/ extern "C" { #ifdef SYCL_BACKEND AOTITorchError aoti_torch_get_current_sycl_queue(void** ret_queue) { diff --git a/openequivariance/openequivariance/extension/tensorproducts.hpp b/openequivariance/openequivariance/extension/tensorproducts.hpp index ff85722a..de248509 100644 --- a/openequivariance/openequivariance/extension/tensorproducts.hpp +++ b/openequivariance/openequivariance/extension/tensorproducts.hpp @@ -82,7 +82,7 @@ class __attribute__ ((visibility ("default"))) JITTPImpl { void* weights, Stream stream) { - auto args = make_kernel_args(num_products, L1_in, L2_in, L3_out, weights); + auto args = KernelArgs(num_products, L1_in, L2_in, L3_out, weights); jit.execute(0, args.data(), args.arg_sizes(), args.count(), with_stream(forward_config_ref, stream)); } @@ -93,7 +93,7 @@ class __attribute__ ((visibility ("default"))) JITTPImpl { void* L2_in, void* L2_grad, void* weight, void* weight_grad, void* L3_grad, Stream stream) { - auto args = make_kernel_args(num_products, L1_in, L1_grad, L2_in, L2_grad, + auto args = KernelArgs(num_products, L1_in, L1_grad, L2_in, L2_grad, weight, weight_grad, L3_grad); jit.execute(1, args.data(), args.arg_sizes(), args.count(), with_stream(backward_config_ref, stream)); @@ -105,7 +105,7 @@ class __attribute__ ((visibility ("default"))) JITTPImpl { void* L1_dgrad, void* L2_dgrad, void* w_dgrad, // Gradients w.r.t outputs of backward op void* L1_grad, void* L2_grad, void* W_grad, void* L3_dgrad, Stream stream) { - auto args = make_kernel_args( + auto args = KernelArgs( num_products, L1_in, L2_in, W, L3_grad, L1_dgrad, L2_dgrad, w_dgrad, L1_grad, L2_grad, W_grad, L3_dgrad); double_backward_config_ref.hStream = stream; diff --git a/openequivariance/openequivariance/extension/torch_core.hpp b/openequivariance/openequivariance/extension/torch_core.hpp index d4f0faee..ddc41599 100644 --- a/openequivariance/openequivariance/extension/torch_core.hpp +++ b/openequivariance/openequivariance/extension/torch_core.hpp @@ -32,7 +32,9 @@ using GPU_Allocator = SYCL_Allocator; #endif -#include "group_mm.hpp" +#ifndef SYCL_BACKEND + #include "group_mm.hpp" +#endif #include "tensorproducts.hpp" #include "convolution.hpp" @@ -195,29 +197,23 @@ inline std::unordered_map lock(mut); tp_cache.clear(); conv_cache.clear(); } +/* +* Registered after the first compile, not at static init: libsycl-jit is +* dlopened on the first runtime compilation and registers its own teardown +* then. atexit runs handlers in reverse order of registration, so registering +* later guarantees the caches are cleared before the JIT library unloads. +*/ inline void register_kernel_cache_cleanup() { - static const bool registered = [] { - std::atexit(release_kernel_caches); - return true; - }(); - (void) registered; + struct RegisterOnce { + RegisterOnce() { std::atexit(release_kernel_caches); } + }; + static RegisterOnce registered; } #endif @@ -659,6 +655,7 @@ inline tuple jit_conv_double_backward( // =========================================================== +#ifndef SYCL_BACKEND 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) { @@ -690,6 +687,7 @@ inline Tensor group_gemm( return C; } +#endif // =========================================================== @@ -710,7 +708,9 @@ REGISTER_LIBRARY_IMPL(libtorch_tp_jit, OEQ_DISPATCH_KEY, m) { m.impl("jit_conv_backward", BOX(&jit_conv_backward)); m.impl("jit_conv_double_backward", BOX(&jit_conv_double_backward)); +#ifndef SYCL_BACKEND m.impl("group_gemm", BOX(&group_gemm)); +#endif }; REGISTER_LIBRARY(libtorch_tp_jit, m) { @@ -722,5 +722,7 @@ REGISTER_LIBRARY(libtorch_tp_jit, m) { m.def("jit_conv_backward(Tensor json_bytes, int hash, Tensor L1_in, Tensor L2_in, Tensor W, Tensor L3_grad, Tensor rows, Tensor cols, Tensor workspace, Tensor transpose_perm) -> (Tensor, Tensor, Tensor)"); m.def("jit_conv_double_backward(Tensor json_bytes, int hash, Tensor L1_in, Tensor L2_in, Tensor W, Tensor L3_grad, Tensor L1_dgrad, Tensor L2_dgrad, Tensor W_dgrad, Tensor rows, Tensor cols, Tensor workspace, Tensor transpose_perm) -> (Tensor, Tensor, Tensor, Tensor)"); +#ifndef SYCL_BACKEND m.def("group_gemm(Tensor A, Tensor B, Tensor ragged_counts, int num_W, int batch_size, int m, int k, int ragged_inner) -> Tensor"); +#endif }; diff --git a/openequivariance/openequivariance/jax/TensorProduct.py b/openequivariance/openequivariance/jax/TensorProduct.py index 419f3b33..086b7e19 100644 --- a/openequivariance/openequivariance/jax/TensorProduct.py +++ b/openequivariance/openequivariance/jax/TensorProduct.py @@ -16,7 +16,9 @@ class TensorProduct(LoopUnrollTP): def __init__(self, problem: TPProblem): dp = extlib.DeviceProp(0) - super().__init__(problem, dp, extlib.BACKEND, torch_op=False) + super().__init__( + problem, dp, "hip" if extlib.IS_HIP else "cuda", torch_op=False + ) self.kernel = self.kernel_string self.weight_numel = problem.weight_numel diff --git a/openequivariance/openequivariance/jax/TensorProductConv.py b/openequivariance/openequivariance/jax/TensorProductConv.py index daf6ff7b..ced7625d 100644 --- a/openequivariance/openequivariance/jax/TensorProductConv.py +++ b/openequivariance/openequivariance/jax/TensorProductConv.py @@ -42,7 +42,7 @@ def __init__( super().__init__( config, dp, - extlib.BACKEND, + "hip" if extlib.IS_HIP else "cuda", idx_dtype=np.int32, torch_op=False, deterministic=deterministic, diff --git a/openequivariance/openequivariance/jax/extlib/__init__.py b/openequivariance/openequivariance/jax/extlib/__init__.py index fdd7553d..cf7219b1 100644 --- a/openequivariance/openequivariance/jax/extlib/__init__.py +++ b/openequivariance/openequivariance/jax/extlib/__init__.py @@ -3,7 +3,9 @@ IS_HIP = oeq_extjax.is_hip() -BACKEND = "hip" if IS_HIP else "cuda" +# The JAX frontend supports CUDA and ROCm only, both of which use torch's +# "cuda" device naming. +DEVICE_TYPE = "cuda" platform = "CUDA" if IS_HIP: @@ -18,6 +20,5 @@ __all__ = [ "GPUTimer", "DeviceProp", - "BACKEND", - "IS_HIP", + "DEVICE_TYPE", ] diff --git a/openequivariance/openequivariance/templates/jinja_utils.py b/openequivariance/openequivariance/templates/jinja_utils.py index 024841da..8b3e7082 100644 --- a/openequivariance/openequivariance/templates/jinja_utils.py +++ b/openequivariance/openequivariance/templates/jinja_utils.py @@ -20,14 +20,8 @@ def sizeof(dtype): @lru_cache(maxsize=8) def get_jinja_environment(backend="cuda", warp_size=32): - """ - Builds the Jinja environment used to render the kernel templates. - - :param backend: one of ``"cuda"``, ``"hip"`` or ``"sycl"``. - :param warp_size: size of a warp / wavefront / sub-group. Only consulted by - the SYCL backend, which must bake the sub-group size into - the generated kernel as a compile-time property. - """ + """:param warp_size: only consulted by SYCL, which bakes the sub-group size + into the generated kernel as a compile-time property.""" if backend not in ("cuda", "hip", "sycl"): raise ValueError(f"Unknown kernel backend '{backend}'") @@ -42,16 +36,16 @@ def get_jinja_environment(backend="cuda", warp_size=32): is_hip = backend == "hip" is_sycl = backend == "sycl" - env.globals["backend"] = backend env.globals["is_hip"] = is_hip env.globals["is_sycl"] = is_sycl env.globals["warp_size"] = warp_size if is_sycl: - # Provided by templates/sycl_compat.cuh. - env.globals["syncwarp"] = "oeq_syncwarp()" - env.globals["atomic_add"] = "oeq_atomic_add" - env.globals["shfl_down"] = lambda val, offset: f"oeq_shfl_down({val}, {offset})" + env.globals["syncwarp"] = "_sycl_syncwarp()" + env.globals["atomic_add"] = "_sycl_atomic_add" + env.globals["shfl_down"] = ( + lambda val, offset: f"_sycl_shfl_down({val}, {offset})" + ) elif is_hip: env.globals["syncwarp"] = ( '__builtin_amdgcn_fence(__ATOMIC_RELEASE, "wavefront");' diff --git a/openequivariance/openequivariance/templates/macros.jinja b/openequivariance/openequivariance/templates/macros.jinja index dd10b1df..547c0256 100644 --- a/openequivariance/openequivariance/templates/macros.jinja +++ b/openequivariance/openequivariance/templates/macros.jinja @@ -237,12 +237,9 @@ Keys map to lists of tuples with (name, dtype, num_elements) of each subarray. __launch_bounds__({{schedule.launch_config.num_threads}}) {%- endmacro %} -{# Declares the per-block shared memory buffer `s`. CUDA and HIP size the - allocation at launch; SYCL runtime compilation has no dynamic local memory - for free-function kernels, so the size is baked in from the schedule. #} {%- macro declare_smem(schedule) %} {%- if is_sycl %} - OEQ_DECLARE_SMEM({{ schedule.launch_config.smem }}) + SYCL_DECLARE_SMEM({{ schedule.launch_config.smem }}) {%- else %} extern __shared__ char s[]; {%- endif %} diff --git a/openequivariance/openequivariance/templates/sycl_compat.cuh b/openequivariance/openequivariance/templates/sycl_compat.cuh index 4c3cbb16..e41fe633 100644 --- a/openequivariance/openequivariance/templates/sycl_compat.cuh +++ b/openequivariance/openequivariance/templates/sycl_compat.cuh @@ -1,11 +1,3 @@ -{# -Compatibility shim that lets the CUDA/HIP-flavored kernel templates compile as -SYCL free-function kernels under runtime compilation. Included only when -targeting the SYCL backend; CUDA and HIP see none of this. - -The generated source is compiled by the SYCL kernel_compiler extension, so it -must be self-contained: everything the kernel body relies on is declared here. -#} #include #include @@ -14,68 +6,53 @@ must be self-contained: everything the kernel body relies on is declared here. namespace syclex = sycl::ext::oneapi::experimental; namespace twi = sycl::ext::oneapi::this_work_item; -// A CUDA __global__ kernel becomes a SYCL free-function nd_range kernel. The -// sub-group size is fixed to the warp size the schedule was generated against, -// which is what makes the warp-level code below well-defined. -#define OEQ_SUBGROUP_SIZE {{ warp_size }} +// The sub-group size is fixed to the warp size the schedule was generated +// against, which is what makes the warp-level code below well-defined. +#define SYCL_SUBGROUP_SIZE {{ warp_size }} #define __global__ extern "C" SYCL_EXTERNAL \ SYCL_EXT_ONEAPI_FUNCTION_PROPERTY((syclex::nd_range_kernel<1>)) \ - SYCL_EXT_ONEAPI_FUNCTION_PROPERTY((syclex::sub_group_size)) + SYCL_EXT_ONEAPI_FUNCTION_PROPERTY((syclex::sub_group_size)) #define __device__ #define __host__ #define __forceinline__ inline #define __restrict__ __restrict -// Occupancy hints have no runtime-compilation equivalent; the work-group size -// is supplied at launch instead. #define __launch_bounds__(...) -// --------------------------------------------------------------------------- -// Thread / block indexing -// --------------------------------------------------------------------------- // All generated kernels are launched as 1D nd_ranges, so only .x is meaningful. -struct OeqIndex1D { +struct SyclIndex1D { size_t x; operator size_t() const { return x; } }; -static inline OeqIndex1D oeq_thread_idx() { return {twi::get_nd_item<1>().get_local_id(0)}; } -static inline OeqIndex1D oeq_block_idx() { return {twi::get_nd_item<1>().get_group(0)}; } -static inline OeqIndex1D oeq_block_dim() { return {twi::get_nd_item<1>().get_local_range(0)}; } -static inline OeqIndex1D oeq_grid_dim() { return {twi::get_nd_item<1>().get_group_range(0)}; } +static inline SyclIndex1D _sycl_thread_idx() { return {twi::get_nd_item<1>().get_local_id(0)}; } +static inline SyclIndex1D _sycl_block_idx() { return {twi::get_nd_item<1>().get_group(0)}; } +static inline SyclIndex1D _sycl_block_dim() { return {twi::get_nd_item<1>().get_local_range(0)}; } +static inline SyclIndex1D _sycl_grid_dim() { return {twi::get_nd_item<1>().get_group_range(0)}; } -#define threadIdx oeq_thread_idx() -#define blockIdx oeq_block_idx() -#define blockDim oeq_block_dim() -#define gridDim oeq_grid_dim() +#define threadIdx _sycl_thread_idx() +#define blockIdx _sycl_block_idx() +#define blockDim _sycl_block_dim() +#define gridDim _sycl_grid_dim() -// --------------------------------------------------------------------------- -// Synchronization -// --------------------------------------------------------------------------- -static inline void oeq_syncwarp() { +static inline void _sycl_syncwarp() { sycl::group_barrier(twi::get_sub_group()); } -static inline void oeq_syncthreads() { +static inline void _sycl_syncthreads() { sycl::group_barrier(twi::get_nd_item<1>().get_group()); } -#define __syncthreads() oeq_syncthreads() +#define __syncthreads() _sycl_syncthreads() -// --------------------------------------------------------------------------- -// Warp-level primitives -// --------------------------------------------------------------------------- template -static inline T oeq_shfl_down(T val, int offset) { +static inline T _sycl_shfl_down(T val, int offset) { return sycl::shift_group_left(twi::get_sub_group(), val, offset); } -// --------------------------------------------------------------------------- -// Atomics -// --------------------------------------------------------------------------- template -static inline T oeq_atomic_add(T* address, T val) { +static inline T _sycl_atomic_add(T* address, T val) { sycl::atomic_ref -static inline auto oeq_min(A a, B b) -> typename std::common_type::type { +static inline auto _sycl_min(A a, B b) -> typename std::common_type::type { using C = typename std::common_type::type; return static_cast(a) < static_cast(b) ? static_cast(a) : static_cast(b); } template -static inline auto oeq_max(A a, B b) -> typename std::common_type::type { +static inline auto _sycl_max(A a, B b) -> typename std::common_type::type { using C = typename std::common_type::type; return static_cast(a) > static_cast(b) ? static_cast(a) : static_cast(b); } -#define min oeq_min -#define max oeq_max - -// --------------------------------------------------------------------------- -// Shared memory -// --------------------------------------------------------------------------- -// CUDA's `extern __shared__ char s[]` sizes the allocation at launch. SYCL -// runtime compilation has no dynamic-local-memory equivalent for free-function -// kernels, so each kernel declares a function-scope work_group_static buffer -// sized to the shared memory its own schedule requires. -#define OEQ_DECLARE_SMEM(BYTES) \ - static syclex::work_group_static oeq_smem_buf; \ - char* s = &oeq_smem_buf[0]; +#define min _sycl_min +#define max _sycl_max + +// SYCL runtime compilation has no dynamic-local-memory equivalent for +// free-function kernels, so each kernel declares a function-scope buffer sized +// to the shared memory its own schedule requires. +#define SYCL_DECLARE_SMEM(BYTES) \ + static syclex::work_group_static _sycl_smem_buf; \ + char* s = &_sycl_smem_buf[0]; diff --git a/tests/batch_test.py b/tests/batch_test.py index eeff2b4a..287835b6 100644 --- a/tests/batch_test.py +++ b/tests/batch_test.py @@ -23,8 +23,6 @@ from conftest import device_type -DEVICE = device_type() - @pytest.fixture(params=[np.float32, np.float64], ids=["F32", "F64"], scope="module") def dtype(request): @@ -430,7 +428,7 @@ def test_submodule_dtype_conversion(self, parent_module_and_problem): parent, problem = parent_module_and_problem batch_size = 10 - device = DEVICE + device = device_type() input_dtype = self._problem_dtype(problem) in1, in2, weights = self._make_inputs(problem, batch_size, input_dtype, device) diff --git a/tests/conftest.py b/tests/conftest.py index af026dc5..5bc998a1 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,4 +1,6 @@ import os +import sys + import pytest os.environ["JAX_ENABLE_X64"] = "True" @@ -21,22 +23,18 @@ def with_jax(request): def device_type(): - """ - The torch device type the kernels run on for the detected backend: - ``"xpu"`` for SYCL, ``"cuda"`` for CUDA and HIP. - """ - from openequivariance._torch.extlib import DEVICE_TYPE + # Called at module scope, before fixtures exist, so --jax is read from the + # command line rather than the with_jax fixture. This keeps JAX-only runs + # from importing the torch extension module. + if "--jax" in sys.argv: + from openequivariance.jax.extlib import DEVICE_TYPE + else: + from openequivariance._torch.extlib import DEVICE_TYPE return DEVICE_TYPE def torch_accelerator(): - """The ``torch.cuda`` / ``torch.xpu`` module matching the active backend.""" import torch return getattr(torch, device_type()) - - -@pytest.fixture(scope="session") -def device(): - return device_type() diff --git a/tests/conv_test.py b/tests/conv_test.py index 7dfa2dce..3eb2456d 100644 --- a/tests/conv_test.py +++ b/tests/conv_test.py @@ -24,8 +24,6 @@ from conftest import device_type -DEVICE = device_type() - @pytest.fixture(params=[np.float32, np.float64], ids=["F32", "F64"], scope="module") def dtype(request): @@ -422,7 +420,7 @@ def _make_inputs(self, problem, graph, rng, dtype, device, deterministic): def test_submodule_dtype_conversion(self, parent_module_and_problem, graph): parent, problem = parent_module_and_problem - device = DEVICE + device = device_type() rng = np.random.default_rng(12345) input_dtype = self._problem_dtype(problem) diff --git a/tests/example_test.py b/tests/example_test.py index 2b363864..f51be3f5 100644 --- a/tests/example_test.py +++ b/tests/example_test.py @@ -3,8 +3,6 @@ from conftest import device_type -DEVICE = device_type() - def test_tutorial_torch(with_jax): if with_jax: @@ -13,19 +11,21 @@ def test_tutorial_torch(with_jax): import torch import e3nn.o3 as o3 - gen = torch.Generator(device=DEVICE) + gen = torch.Generator(device=device_type()) batch_size = 1000 X_ir, Y_ir, Z_ir = o3.Irreps("1x2e"), o3.Irreps("1x3e"), o3.Irreps("1x2e") - X = torch.rand(batch_size, X_ir.dim, device=DEVICE, generator=gen) - Y = torch.rand(batch_size, Y_ir.dim, device=DEVICE, generator=gen) + X = torch.rand(batch_size, X_ir.dim, device=device_type(), generator=gen) + Y = torch.rand(batch_size, Y_ir.dim, device=device_type(), generator=gen) instructions = [(0, 0, 0, "uvu", True)] tp_e3nn = o3.TensorProduct( X_ir, Y_ir, Z_ir, instructions, shared_weights=False, internal_weights=False - ).to(DEVICE) - W = torch.rand(batch_size, tp_e3nn.weight_numel, device=DEVICE, generator=gen) + ).to(device_type()) + W = torch.rand( + batch_size, tp_e3nn.weight_numel, device=device_type(), generator=gen + ) Z = tp_e3nn(X, Y, W) print(torch.norm(Z)) @@ -55,13 +55,15 @@ def test_tutorial_torch(with_jax): [0, 1, 1, 2], # Receiver [1, 0, 2, 1], ], # Sender - device=DEVICE, + device=device_type(), dtype=torch.long, ) - X = torch.rand(node_ct, X_ir.dim, device=DEVICE, generator=gen) - Y = torch.rand(nonzero_ct, Y_ir.dim, device=DEVICE, generator=gen) - W = torch.rand(nonzero_ct, problem.weight_numel, device=DEVICE, generator=gen) + X = torch.rand(node_ct, X_ir.dim, device=device_type(), generator=gen) + Y = torch.rand(nonzero_ct, Y_ir.dim, device=device_type(), generator=gen) + W = torch.rand( + nonzero_ct, problem.weight_numel, device=device_type(), generator=gen + ) tp_conv = oeq.TensorProductConv( problem, deterministic=False diff --git a/tests/export_test.py b/tests/export_test.py index d7bd169e..1b77df7a 100644 --- a/tests/export_test.py +++ b/tests/export_test.py @@ -15,8 +15,6 @@ from conftest import device_type -DEVICE = device_type() - @pytest.fixture(scope="session") def problem_and_irreps(): @@ -32,7 +30,7 @@ def problem_and_irreps(): weight_dtype=np.float32, ) - gen = torch.Generator(device=DEVICE) + gen = torch.Generator(device=device_type()) gen.manual_seed(0) return ( @@ -46,30 +44,34 @@ def problem_and_irreps(): @pytest.fixture(params=["batch", "conv_det", "conv_atomic"], scope="session") def tp_and_inputs(request, problem_and_irreps): problem, X_ir, Y_ir, _ = problem_and_irreps - gen = torch.Generator(device=DEVICE) + gen = torch.Generator(device=device_type()) gen.manual_seed(0) if request.param == "batch": batch_size = 1000 - X = torch.rand(batch_size, X_ir.dim, device=DEVICE, generator=gen) - Y = torch.rand(batch_size, Y_ir.dim, device=DEVICE, generator=gen) - W = torch.rand(batch_size, problem.weight_numel, device=DEVICE, generator=gen) + X = torch.rand(batch_size, X_ir.dim, device=device_type(), generator=gen) + Y = torch.rand(batch_size, Y_ir.dim, device=device_type(), generator=gen) + W = torch.rand( + batch_size, problem.weight_numel, device=device_type(), generator=gen + ) return oeq.TensorProduct(problem), (X, Y, W) else: node_ct, nonzero_ct = 3, 4 # Receiver, sender indices for message passing GNN edge_index = EdgeIndex( - [[0, 1, 1, 2], [1, 0, 2, 1]], device=DEVICE, dtype=torch.long + [[0, 1, 1, 2], [1, 0, 2, 1]], device=device_type(), dtype=torch.long ) _, sender_perm = edge_index.sort_by("col") edge_index, _ = edge_index.sort_by("row") edge_index = [edge_index[0].detach(), edge_index[1].detach()] - X = torch.rand(node_ct, X_ir.dim, device=DEVICE, generator=gen) - Y = torch.rand(nonzero_ct, Y_ir.dim, device=DEVICE, generator=gen) - W = torch.rand(nonzero_ct, problem.weight_numel, device=DEVICE, generator=gen) + X = torch.rand(node_ct, X_ir.dim, device=device_type(), generator=gen) + Y = torch.rand(nonzero_ct, Y_ir.dim, device=device_type(), generator=gen) + W = torch.rand( + nonzero_ct, problem.weight_numel, device=device_type(), generator=gen + ) if request.param == "conv_atomic": return oeq.TensorProductConv(problem, torch_op=True, deterministic=False), ( @@ -147,18 +149,20 @@ def test_aoti_cpp_inference(problem_and_irreps): cmake_prefix_path = torch.utils.cmake_prefix_path torch_ext_so_path = oeq.torch_ext_so_path() - gen = torch.Generator(device=DEVICE) + gen = torch.Generator(device=device_type()) gen.manual_seed(0) batch_size = 1000 # Create models - oeq_tp = oeq.TensorProduct(problem).to(DEVICE) - e3nn_tp = E3NNTensorProduct(problem).e3nn_tp.to(DEVICE) + oeq_tp = oeq.TensorProduct(problem).to(device_type()) + e3nn_tp = E3NNTensorProduct(problem).e3nn_tp.to(device_type()) # Prepare inputs for export - X = torch.rand(batch_size, X_ir.dim, device=DEVICE, generator=gen) - Y = torch.rand(batch_size, Y_ir.dim, device=DEVICE, generator=gen) - W = torch.rand(batch_size, problem.weight_numel, device=DEVICE, generator=gen) + X = torch.rand(batch_size, X_ir.dim, device=device_type(), generator=gen) + Y = torch.rand(batch_size, Y_ir.dim, device=device_type(), generator=gen) + W = torch.rand( + batch_size, problem.weight_numel, device=device_type(), generator=gen + ) inputs = (X, Y, W) with ( diff --git a/tests/input_validation_test.py b/tests/input_validation_test.py index 683db47d..5a621e49 100644 --- a/tests/input_validation_test.py +++ b/tests/input_validation_test.py @@ -7,8 +7,6 @@ from conftest import device_type -DEVICE = device_type() - @pytest.fixture def tpp(): @@ -30,7 +28,7 @@ def edge_index(): ], sort_order="row", sparse_size=(3, 4), - device=DEVICE, + device=device_type(), dtype=torch.long, ) ei.fill_cache_() @@ -39,28 +37,30 @@ def edge_index(): @pytest.fixture def tp_buffers(tpp): - gen = torch.Generator(device=DEVICE) + gen = torch.Generator(device=device_type()) gen.manual_seed(42) N = 1000 - X = torch.rand(N, tpp.irreps_in1.dim, device=DEVICE, generator=gen) - Y = torch.rand(N, tpp.irreps_in2.dim, device=DEVICE, generator=gen) - W = torch.rand(N, tpp.weight_numel, device=DEVICE, generator=gen) + X = torch.rand(N, tpp.irreps_in1.dim, device=device_type(), generator=gen) + Y = torch.rand(N, tpp.irreps_in2.dim, device=device_type(), generator=gen) + W = torch.rand(N, tpp.weight_numel, device=device_type(), generator=gen) return [X, Y, W] @pytest.fixture def conv_buffers(edge_index, tpp): - gen = torch.Generator(device=DEVICE) + gen = torch.Generator(device=device_type()) gen.manual_seed(42) X = torch.rand( - edge_index.num_rows, tpp.irreps_in1.dim, device=DEVICE, generator=gen + edge_index.num_rows, tpp.irreps_in1.dim, device=device_type(), generator=gen ) Y = torch.rand( - edge_index.num_cols, tpp.irreps_in2.dim, device=DEVICE, generator=gen + edge_index.num_cols, tpp.irreps_in2.dim, device=device_type(), generator=gen + ) + W = torch.rand( + edge_index.num_cols, tpp.weight_numel, device=device_type(), generator=gen ) - W = torch.rand(edge_index.num_cols, tpp.weight_numel, device=DEVICE, generator=gen) _, inv_perm = edge_index.get_csc() return [X, Y, W, edge_index[0], edge_index[1], inv_perm] diff --git a/tests/stream_test.py b/tests/stream_test.py index 582e0055..7f973793 100644 --- a/tests/stream_test.py +++ b/tests/stream_test.py @@ -17,8 +17,6 @@ from conftest import device_type, torch_accelerator -DEVICE = device_type() - class KernelExpectation(NamedTuple): kernel_name: str @@ -38,13 +36,9 @@ def __call__(self) -> Any: return self.func(*self.buffers) -accel_device = torch.device(DEVICE) -ACCEL = torch_accelerator() - - @pytest.fixture def gen(): - return torch.Generator(device=DEVICE) + return torch.Generator(device=device_type()) @pytest.fixture @@ -60,7 +54,7 @@ def edge_index(): [1, 0, 2, 1], # Sender ], sparse_size=(3, 4), - device=DEVICE, + device=device_type(), dtype=torch.long, ) @@ -78,21 +72,23 @@ def tpp(): @pytest.fixture def tp_buffers(N, tpp, gen): - X = torch.rand(N, tpp.irreps_in1.dim, device=DEVICE, generator=gen) - Y = torch.rand(N, tpp.irreps_in2.dim, device=DEVICE, generator=gen) - W = torch.rand(N, tpp.weight_numel, device=DEVICE, generator=gen) + X = torch.rand(N, tpp.irreps_in1.dim, device=device_type(), generator=gen) + Y = torch.rand(N, tpp.irreps_in2.dim, device=device_type(), generator=gen) + W = torch.rand(N, tpp.weight_numel, device=device_type(), generator=gen) return (X, Y, W) @pytest.fixture def conv_buffers(edge_index, tpp, gen): X = torch.rand( - edge_index.num_rows, tpp.irreps_in1.dim, device=DEVICE, generator=gen + edge_index.num_rows, tpp.irreps_in1.dim, device=device_type(), generator=gen ) Y = torch.rand( - edge_index.num_cols, tpp.irreps_in2.dim, device=DEVICE, generator=gen + edge_index.num_cols, tpp.irreps_in2.dim, device=device_type(), generator=gen + ) + W = torch.rand( + edge_index.num_cols, tpp.weight_numel, device=device_type(), generator=gen ) - W = torch.rand(edge_index.num_cols, tpp.weight_numel, device=DEVICE, generator=gen) return (X, Y, W, edge_index[0], edge_index[1]) @@ -144,7 +140,7 @@ def double_backward_fn(X, Y, W): dummy = torch.norm(in1_grad) + torch.norm(in2_grad) + torch.norm(w_grad) # Second backward - dummy_grad = torch.tensor(1.0, device=DEVICE) + dummy_grad = torch.tensor(1.0, device=device_type()) dummy.backward( dummy_grad, retain_graph=True, @@ -222,7 +218,7 @@ def double_backward_fn(X, Y, W, receivers, senders): dummy = torch.norm(in1_grad) + torch.norm(in2_grad) + torch.norm(w_grad) # Second backward - dummy_grad = torch.tensor(1.0, device=DEVICE) + dummy_grad = torch.tensor(1.0, device=device_type()) dummy.backward( dummy_grad, retain_graph=True, @@ -302,7 +298,7 @@ def double_backward_fn(X, Y, W, receivers, senders): dummy = torch.norm(in1_grad) + torch.norm(in2_grad) + torch.norm(w_grad) # Second backward - dummy_grad = torch.tensor(1.0, device=DEVICE) + dummy_grad = torch.tensor(1.0, device=device_type()) dummy.backward( dummy_grad, retain_graph=True, @@ -350,8 +346,9 @@ def test_separate_streams(request, tmp_path, executable: Executable): ) as prof: streams = [-1, -2] for priority in streams: - s = ACCEL.Stream(device=accel_device, priority=priority) - with ACCEL.stream(s): + accel = torch_accelerator() + s = accel.Stream(device=torch.device(device_type()), priority=priority) + with accel.stream(s): with record_function(f"executable_{priority}"): for _ in range(COUNT): executable() diff --git a/tests/symmetric_contraction_test.py b/tests/symmetric_contraction_test.py index bc3459a6..c94728e0 100644 --- a/tests/symmetric_contraction_test.py +++ b/tests/symmetric_contraction_test.py @@ -11,7 +11,11 @@ from conftest import device_type -DEVICE = device_type() +if device_type() == "xpu": + pytest.skip( + "Symmetric contraction is not supported on the SYCL backend.", + allow_module_level=True, + ) mace_symmetric_contraction = pytest.importorskip("mace.modules.symmetric_contraction") MaceSymmetricContraction = mace_symmetric_contraction.SymmetricContraction @@ -29,7 +33,6 @@ ], ) -DEVICE = torch.device(DEVICE) SC_CONFIGS = [ SCConfig( @@ -38,7 +41,7 @@ 2, 4, [0, 2, 3, 2, 0, 0, 2, 3, 2, 2], - DEVICE, + torch.device(device_type()), ), SCConfig( o3.Irreps("1x0e + 1x1o + 1x2e"), @@ -46,7 +49,7 @@ 3, 3, [0, 1, 2, 0, 1, 2, 0, 1], - DEVICE, + torch.device(device_type()), ), SCConfig( o3.Irreps("4x0e + 4x1o"), @@ -54,7 +57,7 @@ 2, 5, [0, 1, 2, 3, 4, 0, 1, 2, 3, 4], - DEVICE, + torch.device(device_type()), ), ] From ff37c65d20283827ef211b281b97325e288010a9 Mon Sep 17 00:00:00 2001 From: Abhishek Bagusetty Date: Mon, 21 Sep 2026 11:15:20 -0500 Subject: [PATCH 10/12] fix SYCL CI: link libsycl explicitly and bump the torch pin to 2.13 --- .github/workflows/requirements_sycl_ci.txt | 2 +- openequivariance/openequivariance/_torch/extlib/__init__.py | 5 ++++- 2 files changed, 5 insertions(+), 2 deletions(-) diff --git a/.github/workflows/requirements_sycl_ci.txt b/.github/workflows/requirements_sycl_ci.txt index 08e7a70b..d2d92ac3 100644 --- a/.github/workflows/requirements_sycl_ci.txt +++ b/.github/workflows/requirements_sycl_ci.txt @@ -1,6 +1,6 @@ --extra-index-url https://download.pytorch.org/whl/xpu numpy==2.2.5 -torch==2.12.1+xpu +torch==2.13.0+xpu pytest==9.0.3 ninja==1.11.1.4 nanobind==2.10.2 diff --git a/openequivariance/openequivariance/_torch/extlib/__init__.py b/openequivariance/openequivariance/_torch/extlib/__init__.py index ef89048a..e2fd7fe4 100644 --- a/openequivariance/openequivariance/_torch/extlib/__init__.py +++ b/openequivariance/openequivariance/_torch/extlib/__init__.py @@ -164,8 +164,11 @@ def load_jit_extension(): return os.environ["CXX"] = "icpx" + # -lsycl is explicit: with -fsycl alone the driver may leave libsycl + # out of DT_NEEDED, which only shows up as an undefined symbol at + # import time. extra_cflags.extend(["-fsycl", "-DSYCL_BACKEND"]) - extra_link_args.extend(["-fsycl", "-ltorch_xpu", "-lc10_xpu"]) + extra_link_args.extend(["-fsycl", "-lsycl", "-ltorch_xpu", "-lc10_xpu"]) for lib_dir in library_paths("xpu"): extra_link_args.append("-Wl,-rpath," + lib_dir) From a9f39123a068b7369eebdef60f542c53f58fdfc1 Mon Sep 17 00:00:00 2001 From: Abhishek Bagusetty Date: Tue, 22 Sep 2026 06:43:55 -0500 Subject: [PATCH 11/12] address a few more PR comments --- .../SymmetricContraction.py | 2 -- .../openequivariance/core/utils.py | 1 - .../extension/backend/backend_sycl.hpp | 2 -- .../extension/kernel_args.hpp | 23 ++++++++++++------- .../openequivariance/extension/torch_core.hpp | 6 ----- .../templates/sycl_compat.cuh | 4 ---- 6 files changed, 15 insertions(+), 23 deletions(-) diff --git a/openequivariance/openequivariance/_torch/symmetric_contraction/SymmetricContraction.py b/openequivariance/openequivariance/_torch/symmetric_contraction/SymmetricContraction.py index 833d1bc8..7f3c38fc 100644 --- a/openequivariance/openequivariance/_torch/symmetric_contraction/SymmetricContraction.py +++ b/openequivariance/openequivariance/_torch/symmetric_contraction/SymmetricContraction.py @@ -123,7 +123,6 @@ def U_matrix_real( class GroupMM: def __init__(self, dtype, num_elements, batch_size): - # group_gemm is not built for SYCL, so the operator does not exist there. if extlib.DEVICE_TYPE == "xpu": raise NotImplementedError( "Symmetric contraction is not supported on the SYCL backend." @@ -411,7 +410,6 @@ def backward(ctx, grad_output): ) -# group_gemm is not built for SYCL, so the operator does not exist there. if extlib.BUILT_EXTENSION and extlib.DEVICE_TYPE != "xpu": register_torch_fakes() register_autograd() diff --git a/openequivariance/openequivariance/core/utils.py b/openequivariance/openequivariance/core/utils.py index 11b660aa..7399ea8b 100644 --- a/openequivariance/openequivariance/core/utils.py +++ b/openequivariance/openequivariance/core/utils.py @@ -173,7 +173,6 @@ def benchmark(func, num_warmup, num_iter, mode="gpu_time", kernel_names=[]): else: from torch.profiler import ProfilerActivity, profile, record_function - # Profiler activity is per-accelerator. accelerator_activity = ( ProfilerActivity.XPU if accelerator_device_type() == "xpu" diff --git a/openequivariance/openequivariance/extension/backend/backend_sycl.hpp b/openequivariance/openequivariance/extension/backend/backend_sycl.hpp index f3112a27..52473928 100644 --- a/openequivariance/openequivariance/extension/backend/backend_sycl.hpp +++ b/openequivariance/openequivariance/extension/backend/backend_sycl.hpp @@ -214,8 +214,6 @@ class __attribute__((visibility("default"))) SYCLJITKernel { } } - // Build against the context the kernels run in, so the bundle is valid - // for every device that context spans. sycl::queue &q = resolve_queue(nullptr); sycl::context build_context = q.get_context(); diff --git a/openequivariance/openequivariance/extension/kernel_args.hpp b/openequivariance/openequivariance/extension/kernel_args.hpp index 2c7dd517..aa4c1ecb 100644 --- a/openequivariance/openequivariance/extension/kernel_args.hpp +++ b/openequivariance/openequivariance/extension/kernel_args.hpp @@ -2,22 +2,29 @@ #include #include +#include template class KernelArgs { std::array ptrs_; std::array sizes_; -public: template - explicit KernelArgs(Ts &...args) - : ptrs_{static_cast(&args)...}, sizes_{sizeof(Ts)...} { - static_assert(sizeof...(Ts) == N, "argument count mismatch"); - } + static constexpr bool is_arg_pack = + sizeof...(Ts) == N && + !(sizeof...(Ts) == 1 && + (std::is_same_v, KernelArgs> && ...)); + +public: + template >> + explicit constexpr KernelArgs(Ts &...args) noexcept + : ptrs_{const_cast( + static_cast(&args))...}, + sizes_{sizeof(Ts)...} {} - void **data() { return ptrs_.data(); } - const size_t *arg_sizes() const { return sizes_.data(); } - static constexpr size_t count() { return N; } + constexpr void **data() noexcept { return ptrs_.data(); } + constexpr const size_t *arg_sizes() const noexcept { return sizes_.data(); } + static constexpr size_t count() noexcept { return N; } }; template diff --git a/openequivariance/openequivariance/extension/torch_core.hpp b/openequivariance/openequivariance/extension/torch_core.hpp index ddc41599..ac1daa54 100644 --- a/openequivariance/openequivariance/extension/torch_core.hpp +++ b/openequivariance/openequivariance/extension/torch_core.hpp @@ -203,12 +203,6 @@ inline void release_kernel_caches() { conv_cache.clear(); } -/* -* Registered after the first compile, not at static init: libsycl-jit is -* dlopened on the first runtime compilation and registers its own teardown -* then. atexit runs handlers in reverse order of registration, so registering -* later guarantees the caches are cleared before the JIT library unloads. -*/ inline void register_kernel_cache_cleanup() { struct RegisterOnce { RegisterOnce() { std::atexit(release_kernel_caches); } diff --git a/openequivariance/openequivariance/templates/sycl_compat.cuh b/openequivariance/openequivariance/templates/sycl_compat.cuh index e41fe633..6cd91a6e 100644 --- a/openequivariance/openequivariance/templates/sycl_compat.cuh +++ b/openequivariance/openequivariance/templates/sycl_compat.cuh @@ -6,8 +6,6 @@ namespace syclex = sycl::ext::oneapi::experimental; namespace twi = sycl::ext::oneapi::this_work_item; -// The sub-group size is fixed to the warp size the schedule was generated -// against, which is what makes the warp-level code below well-defined. #define SYCL_SUBGROUP_SIZE {{ warp_size }} #define __global__ extern "C" SYCL_EXTERNAL \ SYCL_EXT_ONEAPI_FUNCTION_PROPERTY((syclex::nd_range_kernel<1>)) \ @@ -20,7 +18,6 @@ namespace twi = sycl::ext::oneapi::this_work_item; #define __launch_bounds__(...) -// All generated kernels are launched as 1D nd_ranges, so only .x is meaningful. struct SyclIndex1D { size_t x; operator size_t() const { return x; } @@ -60,7 +57,6 @@ static inline T _sycl_atomic_add(T* address, T val) { return ref.fetch_add(val); } -// Templated on both operands to keep the mixed-width call sites working. template static inline auto _sycl_min(A a, B b) -> typename std::common_type::type { using C = typename std::common_type::type; From f8e3afdb215a2bff74040afb4143a4fb2043f0bc Mon Sep 17 00:00:00 2001 From: Abhishek Bagusetty Date: Tue, 22 Sep 2026 07:50:48 -0500 Subject: [PATCH 12/12] cleanup a few more comments --- .../openequivariance/_torch/extlib/__init__.py | 3 --- .../openequivariance/extension/backend/backend_sycl.hpp | 9 --------- openequivariance/openequivariance/jax/extlib/__init__.py | 2 -- 3 files changed, 14 deletions(-) diff --git a/openequivariance/openequivariance/_torch/extlib/__init__.py b/openequivariance/openequivariance/_torch/extlib/__init__.py index e2fd7fe4..c8599df0 100644 --- a/openequivariance/openequivariance/_torch/extlib/__init__.py +++ b/openequivariance/openequivariance/_torch/extlib/__init__.py @@ -164,9 +164,6 @@ def load_jit_extension(): return os.environ["CXX"] = "icpx" - # -lsycl is explicit: with -fsycl alone the driver may leave libsycl - # out of DT_NEEDED, which only shows up as an undefined symbol at - # import time. extra_cflags.extend(["-fsycl", "-DSYCL_BACKEND"]) extra_link_args.extend(["-fsycl", "-lsycl", "-ltorch_xpu", "-lc10_xpu"]) diff --git a/openequivariance/openequivariance/extension/backend/backend_sycl.hpp b/openequivariance/openequivariance/extension/backend/backend_sycl.hpp index 52473928..9a8c225c 100644 --- a/openequivariance/openequivariance/extension/backend/backend_sycl.hpp +++ b/openequivariance/openequivariance/extension/backend/backend_sycl.hpp @@ -20,8 +20,6 @@ using Stream = sycl::queue *; Stream get_current_stream(); -// A null stream falls back to the framework's current stream, then to a -// process-wide default queue. inline sycl::queue &resolve_queue(Stream stream) { if (stream != nullptr) { return *stream; @@ -60,8 +58,6 @@ class SYCL_Allocator { } }; -// No cudaEvent equivalent works without enabling profiling on the queue, so -// this brackets a wall-clock interval around a queue synchronization. class GPUTimer { std::chrono::time_point start_time; @@ -133,7 +129,6 @@ class __attribute__((visibility("default"))) DeviceProp { static_cast(dev.get_info()); maxSharedMemoryPerMultiprocessor = maxSharedMemPerBlock; - // Unused on SYCL; present for parity with the CUDA backend. major = 0; minor = 0; } @@ -163,8 +158,6 @@ class __attribute__((visibility("default"))) KernelLaunchConfig { { } }; -// Uses the SYCL kernel_compiler extension with source_language::sycl: -// https://github.com/intel/llvm/blob/sycl/sycl/doc/extensions/experimental/sycl_ext_oneapi_kernel_compiler_sycl.asciidoc class __attribute__((visibility("default"))) SYCLJITKernel { private: bool compiled = false; @@ -254,8 +247,6 @@ class __attribute__((visibility("default"))) SYCLJITKernel { } void set_max_smem(int kernel_id, uint32_t max_smem_bytes) { - // Local memory is declared statically in the generated kernel, so there - // is nothing to opt into; just validate against the device limit. if(!compiled) throw std::logic_error("JIT object has not been compiled!"); if(static_cast(kernel_id) >= kernels.size()) diff --git a/openequivariance/openequivariance/jax/extlib/__init__.py b/openequivariance/openequivariance/jax/extlib/__init__.py index cf7219b1..74e127ef 100644 --- a/openequivariance/openequivariance/jax/extlib/__init__.py +++ b/openequivariance/openequivariance/jax/extlib/__init__.py @@ -3,8 +3,6 @@ IS_HIP = oeq_extjax.is_hip() -# The JAX frontend supports CUDA and ROCm only, both of which use torch's -# "cuda" device naming. DEVICE_TYPE = "cuda" platform = "CUDA"