Skip to main content

Fast Euclid equivariant operations for JAX

Project description

⚙️ e3j-ops

This package contains the CUDA/C++ source for e3j. The Python bindings to XLA handlers are bundled into the e3j_ops shared object, which the main e3j package wraps in custom JAX primitives via the ffi_call API.

Note: The e3j_ops ABI should not be considered stable for now, but considered private within the e3j Python API.

Building from source

The CMake build recipes are defined in CMakeLists.txt. The Makefile of e3j also defines fine-graind recipes to build and test individual CUDA/C++ objects.

Project structure

The source is organized as follows:

  • cuda : CUDA kernel implementations
  • ffi : XLA-FFI handlers declarations and pybind11 module definition
  • xla : vendored XLA FFI headers
  • tests : C++ kernel tests

CUDA/C++ source

Each operation is defined in its own namespace following a common interface (see e.g. tensor_product.cuh for detailed signatures):

  • a namespace e3j::op_name within the enclosing e3j namespace,
  • a struct e3j::op_name::Params for passing hyperparameters,
  • a __global__ CUDA function e3j::op_name::kernel(),
  • a __host__ launcher e3j::op_name::launch()

Although object-oriented patterns are hard to mix with __device__ code, this name-based ABI is subject to change.

namespace e3j {
namespace op_name {

    struct Params;

    template <typename Idx, typename Val>
    __global__ void kernel (Params p, ...);

    template <typename Idx, typename Val>
    e3j::Error launch (..., Params p, cudaStream_t stream);

} // namespace op_name
} // namespace e3j

Most kernels are templated across a range of index and/or value data types. Some e3j primitives automatically dispatch to narrow index dtypes (uint8 / uint16) when feature spaces dimensions are small enough. Half-precision float16 arithmetic for values is not yet supported, but planned for a soon upcoming release.

XLA-FFI handlers

The FFI layer uses the XLA FFI API to bind kernel launchers as XLA custom calls. Each handler is registered with XLA_FFI_DEFINE_HANDLER, receiving typed buffers and attributes directly (no opaque descriptor packing):

xla::Error OpNameHandler(
    cudaStream_t stream,
    int32_t num_out,
    xla::AnyBuffer x,
    xla::Result<xla::AnyBuffer> out
) {
    // ... dtype dispatch ...
    return e3j::op_name::launch<Idx, Val>(
        x.typed_data<Idx>(),
        out->typed_data<Val>(),
        params, stream
    ).to_xla();
}

XLA_FFI_DEFINE_HANDLER(
    xla_op_name,
    OpNameHandler,
    xla::Ffi::Bind()
        .Ctx<xla::PlatformStream<cudaStream_t>>()
        .Attr<int32_t>("num_out")
        .Arg<xla::AnyBuffer>()
        .Ret<xla::AnyBuffer>()
);

Handlers are exposed to Python as PyCapsules via pyEncapsulateFunction in e3j_ops.h, and registered in the pybind11 module defined in e3j_ops.cpp.

Contributing

Although it is too early for e3j_ops to accept significant external contributions, bug reports or questions are very welcome via GitHub issues and discussions.

Citing

If you use e3j within your work, we kindly ask you to cite the following preprint:

@article{Peltre26-e3j,
    title   = {{E3J}: an Efficient and Open-Source Euclidean Equivariance Backend},
    author  = {Peltre, Olivier and Picard, Armand and Pichard, Adrien and Giacomoni, Luca and Braganca, Miguel and Heyraud, Valentin and Brunken, Christoph and Tilly, Jules},
    journal = {preprint},
    year    = {2026},
    url     = {(preprint)}
  }
}

Project details


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distributions

No source distribution files available for this release.See tutorial on generating distribution archives.

Built Distributions

If you're not sure about the file name format, learn more about wheel file names.

e3j_ops-0.1.0b4-cp314-cp314-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl (3.9 MB view details)

Uploaded CPython 3.14manylinux: glibc 2.24+ x86-64manylinux: glibc 2.28+ x86-64

e3j_ops-0.1.0b4-cp313-cp313-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl (3.9 MB view details)

Uploaded CPython 3.13manylinux: glibc 2.24+ x86-64manylinux: glibc 2.28+ x86-64

e3j_ops-0.1.0b4-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl (3.9 MB view details)

Uploaded CPython 3.12manylinux: glibc 2.24+ x86-64manylinux: glibc 2.28+ x86-64

e3j_ops-0.1.0b4-cp311-cp311-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl (3.9 MB view details)

Uploaded CPython 3.11manylinux: glibc 2.24+ x86-64manylinux: glibc 2.28+ x86-64

File details

Details for the file e3j_ops-0.1.0b4-cp314-cp314-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl.

File metadata

File hashes

Hashes for e3j_ops-0.1.0b4-cp314-cp314-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl
Algorithm Hash digest
SHA256 d8823b57113afd74b093606845c0d7b2e57189c90d9f895fd70e886cdf9fc9c3
MD5 a8d78bfa645def904370c855b8e28193
BLAKE2b-256 81c6a042c2845f1d64be8c35993337245adb331bfcfc28b0715e5afcaa06ad89

See more details on using hashes here.

File details

Details for the file e3j_ops-0.1.0b4-cp313-cp313-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl.

File metadata

File hashes

Hashes for e3j_ops-0.1.0b4-cp313-cp313-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl
Algorithm Hash digest
SHA256 025f2a1ce71a4c4722023ee2973e273d4bc482aff8297a9f974c7fc1b484f7c2
MD5 8a8069ad0d1c270992fc48d9d9f25bfd
BLAKE2b-256 ed1ced9c405595dcca286a8fc7fcc40650308d7ea01a6c0e98cacaef36d15753

See more details on using hashes here.

File details

Details for the file e3j_ops-0.1.0b4-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl.

File metadata

File hashes

Hashes for e3j_ops-0.1.0b4-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl
Algorithm Hash digest
SHA256 51c94145f7faf8e4167c8e6e11d49d3b97883870fc3054b7dbc838aa57e82fa0
MD5 8bd2fde733395d3e1d2bce12523007b4
BLAKE2b-256 eb690c5e289bcf94a2ea14d266f9e75ef25973e8ce92f5b12e65d006adb93e46

See more details on using hashes here.

File details

Details for the file e3j_ops-0.1.0b4-cp311-cp311-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl.

File metadata

File hashes

Hashes for e3j_ops-0.1.0b4-cp311-cp311-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl
Algorithm Hash digest
SHA256 34a34cf52e418e1e6123078d341b9b8ebb5dfd8313ad5ceef82bc22a6503b40e
MD5 a6b6c77baa10b8fa8886e6574ddbdd2c
BLAKE2b-256 982e8359333330f7a4547e188a672486e0debc722d8b605295e2fedff8427896

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page