diff --git a/.github/workflows/requirements_sycl_ci.txt b/.github/workflows/requirements_sycl_ci.txt new file mode 100644 index 00000000..d2d92ac3 --- /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.13.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..358d9340 100644 --- a/.github/workflows/verify_extension_build.yml +++ b/.github/workflows/verify_extension_build.yml @@ -41,4 +41,43 @@ 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' + + - 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 + + - 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/docs/installation.rst b/docs/installation.rst index 5ade5c0c..17feab14 100644 --- a/docs/installation.rst +++ b/docs/installation.rst @@ -7,13 +7,22 @@ Installation You need the following to install OpenEquivariance: -- A Linux system equipped with an NVIDIA / AMD 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 +- A Linux system equipped with an NVIDIA / AMD / Intel graphics card. +- 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 ``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`` + (PyTorch locates the SYCL toolchain from it). Tensors passed to the + kernels live on the ``xpu`` device rather than ``cuda``. The precompiled + extension and the JAX frontend both remain CUDA/HIP only. + .. tab:: PyTorch Installation is one easy command, followed by import verification: @@ -143,4 +152,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 restore + 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..3f3ad3dc 100644 --- a/openequivariance/CMakeLists.txt +++ b/openequivariance/CMakeLists.txt @@ -173,6 +173,40 @@ 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.") +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) + + set(SYCL_LINK_LIBS + sycl_stub_lib + sycl + ) + add_stable_extension(oeq_stable_sycl SYCL_BACKEND "${SYCL_LINK_LIBS}") + + 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..c8599df0 100644 --- a/openequivariance/openequivariance/_torch/extlib/__init__.py +++ b/openequivariance/openequivariance/_torch/extlib/__init__.py @@ -22,11 +22,34 @@ extension_module = None -assert torch.version.cuda or torch.version.hip, ( - "Only CUDA and HIP backends are supported" + +def _detect_backend(): + """HIP builds report a ``torch.version.cuda`` of None, so HIP is tested first.""" + if torch.version.hip: + return "hip" + if torch.version.cuda: + return "cuda" + if getattr(torch.version, "xpu", None): + 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" + +if IS_SYCL and Version(torch.__version__) < Version("2.10"): + raise RuntimeError( + f"The SYCL backend requires PyTorch >= 2.10, found {torch.__version__}." + ) + +DEVICE_TYPE = "xpu" if IS_SYCL else "cuda" @contextlib.contextmanager @@ -111,7 +134,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 +147,35 @@ 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": + 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" + + extra_cflags.extend(["-fsycl", "-DSYCL_BACKEND"]) + 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) + extra_link_args.append("-L" + lib_dir) 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 +207,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 +233,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 +261,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/_torch/symmetric_contraction/SymmetricContraction.py b/openequivariance/openequivariance/_torch/symmetric_contraction/SymmetricContraction.py index 83c56933..7f3c38fc 100644 --- a/openequivariance/openequivariance/_torch/symmetric_contraction/SymmetricContraction.py +++ b/openequivariance/openequivariance/_torch/symmetric_contraction/SymmetricContraction.py @@ -123,6 +123,10 @@ def U_matrix_real( class GroupMM: def __init__(self, dtype, num_elements, batch_size): + 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 +410,6 @@ def backward(ctx, grad_output): ) -if extlib.BUILT_EXTENSION: +if extlib.BUILT_EXTENSION and extlib.DEVICE_TYPE != "xpu": register_torch_fakes() register_autograd() 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..5aabb502 100644 --- a/openequivariance/openequivariance/core/ComputationSchedule.py +++ b/openequivariance/openequivariance/core/ComputationSchedule.py @@ -62,7 +62,9 @@ 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 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) 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..7399ea8b 100644 --- a/openequivariance/openequivariance/core/utils.py +++ b/openequivariance/openequivariance/core/utils.py @@ -173,13 +173,17 @@ def benchmark(func, num_warmup, num_iter, mode="gpu_time", kernel_names=[]): else: from torch.profiler import ProfilerActivity, profile, record_function + 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 +262,10 @@ def transpose_irrep_layout( ) return out + + +def accelerator_device_type(): + """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 4ecef72a..9dd224c8 100644 --- a/openequivariance/openequivariance/extension/backend/backend_cuda.hpp +++ b/openequivariance/openequivariance/extension/backend/backend_cuda.hpp @@ -323,7 +323,8 @@ class __attribute__((visibility("default"))) CUJITKernel { } } - void execute(int kernel_id, void* args[], KernelLaunchConfig config) { + 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 2eb4ed85..a68e5433 100644 --- a/openequivariance/openequivariance/extension/backend/backend_hip.hpp +++ b/openequivariance/openequivariance/extension/backend/backend_hip.hpp @@ -292,7 +292,8 @@ 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[], [[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 new file mode 100644 index 00000000..9a8c225c --- /dev/null +++ b/openequivariance/openequivariance/extension/backend/backend_sycl.hpp @@ -0,0 +1,302 @@ +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include + +using namespace std; +namespace syclex = sycl::ext::oneapi::experimental; + +using Stream = sycl::queue *; + +Stream get_current_stream(); + +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(); + } +}; + +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()); + + 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; + + 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)) + { } +}; + +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]; + + 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); + } + } + + 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) { + 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..55a6fdd2 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 = 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) { - void *fixup_args[] = {&workspace, &L3_out}; + auto fixup_args = KernelArgs(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 = 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) { - void *fixup_args[] = {&workspace, &L1_grad}; + auto fixup_args = KernelArgs(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 = 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, 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 = KernelArgs(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 = KernelArgs(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/kernel_args.hpp b/openequivariance/openequivariance/extension/kernel_args.hpp new file mode 100644 index 00000000..aa4c1ecb --- /dev/null +++ b/openequivariance/openequivariance/extension/kernel_args.hpp @@ -0,0 +1,31 @@ +#pragma once + +#include +#include +#include + +template +class KernelArgs { + std::array ptrs_; + std::array sizes_; + + template + 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)...} {} + + 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 +KernelArgs(Ts &...) -> KernelArgs; diff --git a/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp b/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp index ddabd0bb..c7f9bdb6 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,17 @@ Stream get_current_stream() { #ifdef HIP_BACKEND return c10::hip::getCurrentHIPStream(); #endif +#ifdef SYCL_BACKEND + 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..dc8ec85b 100644 --- a/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp +++ b/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp @@ -11,6 +11,11 @@ #include #include #include +#ifdef SYCL_BACKEND + // 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 using Tensor = torch::stable::Tensor; @@ -67,14 +72,27 @@ 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 + 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(). + 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 +101,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..16280c13 100644 --- a/openequivariance/openequivariance/extension/stubs/stream.cpp +++ b/openequivariance/openequivariance/extension/stubs/stream.cpp @@ -2,7 +2,14 @@ #include 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..de248509 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 = 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)); } 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 = 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)); } 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 = 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; - 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..ac1daa54 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,7 +26,15 @@ using GPU_Allocator = HIP_Allocator; #endif -#include "group_mm.hpp" +#ifdef SYCL_BACKEND + #include "backend_sycl.hpp" + using JITKernel = SYCLJITKernel; + using GPU_Allocator = SYCL_Allocator; +#endif + +#ifndef SYCL_BACKEND + #include "group_mm.hpp" +#endif #include "tensorproducts.hpp" #include "convolution.hpp" @@ -43,6 +52,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 +131,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 +196,21 @@ inline std::unordered_map lock(mut); + tp_cache.clear(); + conv_cache.clear(); +} + +inline void register_kernel_cache_cleanup() { + struct RegisterOnce { + RegisterOnce() { std::atexit(release_kernel_caches); } + }; + static RegisterOnce registered; +} +#endif + inline std::pair*, KernelProp> compile_tp_with_caching(const Tensor &json_bytes, int64_t hash) { @@ -220,6 +245,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 +287,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}; } @@ -618,6 +649,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) { @@ -649,10 +681,19 @@ inline Tensor group_gemm( return C; } +#endif // =========================================================== -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)); @@ -661,7 +702,9 @@ REGISTER_LIBRARY_IMPL(libtorch_tp_jit, CUDA, 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) { @@ -673,5 +716,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 f880544f..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.IS_HIP, 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 9234158f..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.IS_HIP, + "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 23a6f63a..74e127ef 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() +DEVICE_TYPE = "cuda" + platform = "CUDA" if IS_HIP: platform = "ROCM" @@ -16,4 +18,5 @@ __all__ = [ "GPUTimer", "DeviceProp", + "DEVICE_TYPE", ] 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..8b3e7082 100644 --- a/openequivariance/openequivariance/templates/jinja_utils.py +++ b/openequivariance/openequivariance/templates/jinja_utils.py @@ -18,8 +18,13 @@ 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): + """: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}'") + env = Environment( loader=PackageLoader("openequivariance"), extensions=["jinja2.ext.do"] ) @@ -28,18 +33,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["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: + 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");' + "__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..547c0256 100644 --- a/openequivariance/openequivariance/templates/macros.jinja +++ b/openequivariance/openequivariance/templates/macros.jinja @@ -236,3 +236,11 @@ 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 %} + +{%- macro declare_smem(schedule) %} + {%- if is_sycl %} + SYCL_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..6cd91a6e --- /dev/null +++ b/openequivariance/openequivariance/templates/sycl_compat.cuh @@ -0,0 +1,80 @@ +#include + +#include +#include + +namespace syclex = sycl::ext::oneapi::experimental; +namespace twi = sycl::ext::oneapi::this_work_item; + +#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)) + +#define __device__ +#define __host__ +#define __forceinline__ inline +#define __restrict__ __restrict + +#define __launch_bounds__(...) + +struct SyclIndex1D { + size_t x; + operator size_t() const { return x; } +}; + +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 _sycl_thread_idx() +#define blockIdx _sycl_block_idx() +#define blockDim _sycl_block_dim() +#define gridDim _sycl_grid_dim() + +static inline void _sycl_syncwarp() { + sycl::group_barrier(twi::get_sub_group()); +} + +static inline void _sycl_syncthreads() { + sycl::group_barrier(twi::get_nd_item<1>().get_group()); +} + +#define __syncthreads() _sycl_syncthreads() + +template +static inline T _sycl_shfl_down(T val, int offset) { + return sycl::shift_group_left(twi::get_sub_group(), val, offset); +} + +template +static inline T _sycl_atomic_add(T* address, T val) { + sycl::atomic_ref ref(*address); + return ref.fetch_add(val); +} + +template +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 _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 _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 77715028..287835b6 100644 --- a/tests/batch_test.py +++ b/tests/batch_test.py @@ -21,6 +21,8 @@ import openequivariance as oeq +from conftest import device_type + @pytest.fixture(params=[np.float32, np.float64], ids=["F32", "F64"], scope="module") def dtype(request): @@ -426,7 +428,7 @@ def test_submodule_dtype_conversion(self, parent_module_and_problem): parent, problem = parent_module_and_problem batch_size = 10 - device = "cuda" + 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 4a515664..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" @@ -18,3 +20,21 @@ def pytest_addoption(parser): @pytest.fixture(scope="session") def with_jax(request): return request.config.getoption("--jax") + + +def 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(): + import torch + + return getattr(torch, device_type()) diff --git a/tests/conv_test.py b/tests/conv_test.py index 1adb5553..3eb2456d 100644 --- a/tests/conv_test.py +++ b/tests/conv_test.py @@ -22,6 +22,8 @@ nequip_oam_problems, ) +from conftest import device_type + @pytest.fixture(params=[np.float32, np.float64], ids=["F32", "F64"], scope="module") def dtype(request): @@ -418,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 = "cuda" + 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 bf51d3ed..f51be3f5 100644 --- a/tests/example_test.py +++ b/tests/example_test.py @@ -1,6 +1,8 @@ import pytest import os +from conftest import device_type + def test_tutorial_torch(with_jax): if with_jax: @@ -9,19 +11,21 @@ def test_tutorial_torch(with_jax): import torch import e3nn.o3 as o3 - gen = torch.Generator(device="cuda") + 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="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_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("cuda") - W = torch.rand(batch_size, tp_e3nn.weight_numel, device="cuda", 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)) @@ -51,13 +55,15 @@ def test_tutorial_torch(with_jax): [0, 1, 1, 2], # Receiver [1, 0, 2, 1], ], # Sender - device="cuda", + device=device_type(), 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_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 6aba9690..1b77df7a 100644 --- a/tests/export_test.py +++ b/tests/export_test.py @@ -13,6 +13,8 @@ from openequivariance._torch.E3NNTensorProduct import E3NNTensorProduct +from conftest import device_type + @pytest.fixture(scope="session") def problem_and_irreps(): @@ -28,7 +30,7 @@ def problem_and_irreps(): weight_dtype=np.float32, ) - gen = torch.Generator(device="cuda") + gen = torch.Generator(device=device_type()) gen.manual_seed(0) return ( @@ -42,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="cuda") + 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="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_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="cuda", 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="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_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), ( @@ -143,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="cuda") + gen = torch.Generator(device=device_type()) 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_type()) + e3nn_tp = E3NNTensorProduct(problem).e3nn_tp.to(device_type()) # 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_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 9b38d55e..5a621e49 100644 --- a/tests/input_validation_test.py +++ b/tests/input_validation_test.py @@ -5,6 +5,8 @@ from openequivariance import TPProblem, TensorProduct, TensorProductConv +from conftest import device_type + @pytest.fixture def tpp(): @@ -26,7 +28,7 @@ def edge_index(): ], sort_order="row", sparse_size=(3, 4), - device="cuda", + device=device_type(), dtype=torch.long, ) ei.fill_cache_() @@ -35,28 +37,30 @@ def edge_index(): @pytest.fixture def tp_buffers(tpp): - gen = torch.Generator(device="cuda") + gen = torch.Generator(device=device_type()) 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_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="cuda") + gen = torch.Generator(device=device_type()) 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_type(), 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_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="cuda", 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..7f973793 100644 --- a/tests/stream_test.py +++ b/tests/stream_test.py @@ -15,6 +15,8 @@ from openequivariance import TensorProduct, TensorProductConv, TPProblem +from conftest import device_type, torch_accelerator + class KernelExpectation(NamedTuple): kernel_name: str @@ -34,12 +36,9 @@ def __call__(self) -> Any: return self.func(*self.buffers) -cuda = torch.device("cuda") - - @pytest.fixture def gen(): - return torch.Generator(device="cuda") + return torch.Generator(device=device_type()) @pytest.fixture @@ -55,7 +54,7 @@ def edge_index(): [1, 0, 2, 1], # Sender ], sparse_size=(3, 4), - device="cuda", + device=device_type(), dtype=torch.long, ) @@ -73,21 +72,23 @@ 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_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="cuda", 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="cuda", 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="cuda", generator=gen) return (X, Y, W, edge_index[0], edge_index[1]) @@ -139,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="cuda") + dummy_grad = torch.tensor(1.0, device=device_type()) dummy.backward( dummy_grad, retain_graph=True, @@ -217,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="cuda") + dummy_grad = torch.tensor(1.0, device=device_type()) dummy.backward( dummy_grad, retain_graph=True, @@ -297,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="cuda") + dummy_grad = torch.tensor(1.0, device=device_type()) dummy.backward( dummy_grad, retain_graph=True, @@ -345,8 +346,9 @@ 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): + 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 bbd105ac..c94728e0 100644 --- a/tests/symmetric_contraction_test.py +++ b/tests/symmetric_contraction_test.py @@ -9,6 +9,14 @@ from openequivariance._torch.symmetric_contraction import SymmetricContraction +from conftest import 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 @@ -25,7 +33,6 @@ ], ) -DEVICE = torch.device("cuda") SC_CONFIGS = [ SCConfig( @@ -34,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"), @@ -42,7 +49,7 @@ 3, 3, [0, 1, 2, 0, 1, 2, 0, 1], - DEVICE, + torch.device(device_type()), ), SCConfig( o3.Irreps("4x0e + 4x1o"), @@ -50,7 +57,7 @@ 2, 5, [0, 1, 2, 3, 4, 0, 1, 2, 3, 4], - DEVICE, + torch.device(device_type()), ), ]