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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions .github/workflows/requirements_sycl_ci.txt
Original file line number Diff line number Diff line change
@@ -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
41 changes: 40 additions & 1 deletion .github/workflows/verify_extension_build.yml
Original file line number Diff line number Diff line change
Expand Up @@ -41,4 +41,43 @@ jobs:

- name: Test JAX extension build
run: |
XLA_DIRECT_DOWNLOAD=1 pip install -e "./openequivariance_extjax" --no-build-isolation
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 \
Comment thread
vbharadwaj-bk marked this conversation as resolved.
| 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
34 changes: 29 additions & 5 deletions docs/installation.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -143,4 +152,19 @@ on a major cluster, send us a pull request to add your configuration!

conda activate <your-conda-env>
export CC=cc
export CXX=CC
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
38 changes: 36 additions & 2 deletions openequivariance/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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()
3 changes: 2 additions & 1 deletion openequivariance/openequivariance/_torch/E3NNConv.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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)

Expand Down
17 changes: 9 additions & 8 deletions openequivariance/openequivariance/_torch/E3NNTensorProduct.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand Down Expand Up @@ -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
Expand All @@ -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)

Expand All @@ -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)

Expand Down
3 changes: 2 additions & 1 deletion openequivariance/openequivariance/_torch/FlashTPConv.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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),
)

Expand Down
113 changes: 71 additions & 42 deletions openequivariance/openequivariance/_torch/NPDoubleBackwardMixin.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import torch
from openequivariance._torch.extlib import DEVICE_TYPE


def _none_to_zeros(values, refs):
Expand All @@ -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)

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
Loading
Loading