From b7a2a78cc355e7c59d3f316ad1fb8febe13fdd29 Mon Sep 17 00:00:00 2001 From: asglover <140220574+asglover@users.noreply.github.com> Date: Sun, 13 Sep 2026 00:52:50 -0700 Subject: [PATCH 1/5] framework managed workspace allocations --- CHANGELOG.md | 5 + .../_torch/TensorProductConv.py | 26 +--- .../openequivariance/core/ConvolutionBase.py | 1 - .../extension/convolution.hpp | 59 +++++--- .../extension/libtorch_tp_jit.cpp | 4 + .../extension/libtorch_tp_jit_stable.cpp | 6 + .../openequivariance/extension/torch_core.hpp | 44 ++++-- .../openequivariance/jax/TensorProductConv.py | 5 - .../openequivariance/jax/jvp/conv_prim.py | 143 ++++++++++-------- .../openequivariance/jax/utils.py | 14 ++ .../openequivariance/jax/vjp/conv_func.py | 81 ++++++---- .../templates/loop_unroll_conv_atomic.cuh | 18 --- openequivariance_extjax/src/ffi_handlers.cpp | 41 ++--- tests/cuda_graph_test.py | 123 +++++++++++++++ tests/multidevice_test.py | 18 +++ tests/stream_test.py | 7 +- 16 files changed, 397 insertions(+), 198 deletions(-) create mode 100644 tests/cuda_graph_test.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 1c9f9847..da747680 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,10 @@ ## Latest Changes +### Unreleased +**Changed**: +- The deterministic convolution workspace is now allocated by the framework (PyTorch / JAX) each call +- Atomic convolutions no longer compile or launch fixup kernels. + ### v0.7.0 (2026-09-10) **Added**: - Public XLA FFI registration provider diff --git a/openequivariance/openequivariance/_torch/TensorProductConv.py b/openequivariance/openequivariance/_torch/TensorProductConv.py index d052f909..94e42b4c 100644 --- a/openequivariance/openequivariance/_torch/TensorProductConv.py +++ b/openequivariance/openequivariance/_torch/TensorProductConv.py @@ -85,7 +85,6 @@ def _init_class(self): kahan=self.input_args["kahan"], ) - self.allocate_workspace(self.workspace_size) self.dummy_transpose_perm = torch.zeros(1, dtype=torch.int64, device="cuda") self.weight_numel = self.config.weight_numel @@ -180,18 +179,9 @@ def forward( self.L3_dim, rows, cols, - self.workspace_buffer, sender_perm, ) - def allocate_workspace(self, size_bytes): - self.workspace_size = size_bytes - self.workspace_buffer = torch.zeros( - size_bytes, dtype=torch.uint8, device="cuda" - ) - self.workspace_ptr = self.workspace_buffer.data_ptr() - logger.info(f"Convolution requires {size_bytes // 1000000}MB of workspace.") - def reorder_weights_from_e3nn(self, weights, has_batch_dim=True): return reorder_torch( self.forward_schedule, weights, "forward", not self.config.shared_weights @@ -282,9 +272,7 @@ def register_torch_fakes(): import torch @torch.library.register_fake("libtorch_tp_jit::jit_conv_forward") - def fake_forward( - kernel, hash, L1_in, L2_in, W, L3_dim, rows, cols, workspace_buffer, sender_perm - ): + def fake_forward(kernel, hash, L1_in, L2_in, W, L3_dim, rows, cols, sender_perm): return torch.empty(L1_in.shape[0], L3_dim, device="cuda", dtype=L1_in.dtype) @torch.library.register_fake("libtorch_tp_jit::jit_conv_backward") @@ -297,7 +285,6 @@ def fake_backward( L3_grad, rows, cols, - workspace_buffer, sender_perm, ): return torch.empty_like(L1_in), torch.empty_like(L2_in), torch.empty_like(W) @@ -315,7 +302,6 @@ def fake_double_backward( w_dgrad, rows, cols, - workspace_buffer, transpose_perm=None, ): return [ @@ -345,7 +331,6 @@ def setup_context(ctx, inputs, output): ctx.L3_dim, ctx.rows, ctx.cols, - ctx.workspace_buffer, ctx.sender_perm, ) = inputs @@ -359,10 +344,9 @@ def backward(ctx, grad_output): grad_output, ctx.rows, ctx.cols, - ctx.workspace_buffer, ctx.sender_perm, ) - return None, None, L1_grad, L2_grad, W_grad, None, None, None, None, None + return None, None, L1_grad, L2_grad, W_grad, None, None, None, None torch.library.register_autograd( "libtorch_tp_jit::jit_conv_forward", backward, setup_context=setup_context @@ -378,7 +362,6 @@ def setup_context_double_backward(ctx, inputs, output): ctx.grad_output, ctx.rows, ctx.cols, - ctx.workspace_buffer, ctx.sender_perm, ) = inputs ctx.inputs = inputs @@ -396,7 +379,6 @@ def double_backward(ctx, E, F, G): G, ctx.rows, ctx.cols, - ctx.workspace_buffer, ctx.sender_perm, ) return ( @@ -409,7 +391,6 @@ def double_backward(ctx, E, F, G): None, None, None, - None, ) torch.library.register_autograd( @@ -431,7 +412,6 @@ def setup_context_triple_backward(ctx, inputs, output): ctx.W_dgrad, ctx.rows, ctx.cols, - ctx.workspace_buffer, ctx.sender_perm, ) = inputs @@ -444,7 +424,6 @@ def triple_backward(ctx, t_L1_grad, t_L2_grad, t_W_grad, t_L3_dgrad): common_args = ( ctx.rows, ctx.cols, - ctx.workspace_buffer, ctx.sender_perm, ) @@ -546,7 +525,6 @@ def triple_backward(ctx, t_L1_grad, t_L2_grad, t_W_grad, t_L3_dgrad): None, None, None, - None, ) torch.library.register_autograd( diff --git a/openequivariance/openequivariance/core/ConvolutionBase.py b/openequivariance/openequivariance/core/ConvolutionBase.py index 116a21b3..200c8385 100644 --- a/openequivariance/openequivariance/core/ConvolutionBase.py +++ b/openequivariance/openequivariance/core/ConvolutionBase.py @@ -112,7 +112,6 @@ def __init__( global torch import torch - self.workspace_ptr = 0 self.workspace_size = 0 def reorder_weights_from_e3nn(self, weights, has_batch_dim=True): diff --git a/openequivariance/openequivariance/extension/convolution.hpp b/openequivariance/openequivariance/extension/convolution.hpp index 83ad58b4..d732dfae 100644 --- a/openequivariance/openequivariance/extension/convolution.hpp +++ b/openequivariance/openequivariance/extension/convolution.hpp @@ -20,33 +20,49 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { KernelLaunchConfig backward_config_ref; KernelLaunchConfig double_backward_config_ref; int opt_level; + bool deterministic; + + enum Kernel { + FORWARD = 0, + BACKWARD = 1, + DOUBLE_BACKWARD_A = 2, + DOUBLE_BACKWARD_B = 3, + FIXUP_FORWARD = 4, + FIXUP_BACKWARD = 5, + FIXUP_DOUBLE_BACKWARD_B = 6 + }; JITConvImpl( std::string jit_kernel, KernelLaunchConfig forward_config_i, KernelLaunchConfig backward_config_i, KernelLaunchConfig double_backward_config_i, - int opt_level_i) : + int opt_level_i, + bool deterministic_i) : jit(jit_kernel), forward_config_ref(forward_config_i), backward_config_ref(backward_config_i), double_backward_config_ref(double_backward_config_i), - opt_level(opt_level_i) { + opt_level(opt_level_i), + deterministic(deterministic_i) { - vector kernels = {"forward", "backward", "fixup_forward", "fixup_backward", "double_backward_A", "double_backward_B", "fixup_double_backwardB"}; - jit.compile(kernels, {{}, {}, {}, {}, {}, {}, {}}, opt_level); + vector kernels = {"forward", "backward", "double_backward_A", "double_backward_B"}; + if(deterministic) { + kernels.insert(kernels.end(), {"fixup_forward", "fixup_backward", "fixup_double_backwardB"}); + } + jit.compile(kernels, vector>(kernels.size()), opt_level); if(forward_config_ref.smem > 0) { - jit.set_max_smem(0, forward_config_ref.smem); - jit.set_max_smem(4, forward_config_ref.smem); + jit.set_max_smem(FORWARD, forward_config_ref.smem); + jit.set_max_smem(DOUBLE_BACKWARD_A, forward_config_ref.smem); } if(backward_config_ref.smem > 0) { - jit.set_max_smem(1, backward_config_ref.smem); + jit.set_max_smem(BACKWARD, backward_config_ref.smem); } if(double_backward_config_ref.smem > 0) { - jit.set_max_smem(5, double_backward_config_ref.smem); + jit.set_max_smem(DOUBLE_BACKWARD_B, double_backward_config_ref.smem); } } @@ -73,7 +89,8 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { dbl_bwd_dict["num_threads"], dbl_bwd_dict["smem"] ), - static_cast(kernel_dims["opt_level"])) { } + static_cast(kernel_dims["opt_level"]), + kernel_dims["deterministic"] != 0) { } void exec_conv( void* L1_in, @@ -90,9 +107,9 @@ 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)); + jit.execute(FORWARD, args, with_stream(forward_config_ref, stream)); - if(reinterpret_cast(workspace) != 0) { + if(deterministic) { void *fixup_args[] = {&workspace, &L3_out}; KernelLaunchConfig fixup_config( @@ -102,7 +119,7 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { ); fixup_config.hStream = stream; - jit.execute(2, fixup_args, fixup_config); + jit.execute(FIXUP_FORWARD, fixup_args, fixup_config); } } @@ -119,9 +136,9 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { 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)); + jit.execute(BACKWARD, args, with_stream(backward_config_ref, stream)); - if(reinterpret_cast(workspace) != 0) { + if(deterministic) { void *fixup_args[] = {&workspace, &L1_grad}; KernelLaunchConfig fixup_config( @@ -131,7 +148,7 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { ); fixup_config.hStream = stream; - jit.execute(3, fixup_args, fixup_config); + jit.execute(FIXUP_BACKWARD, fixup_args, fixup_config); } } @@ -150,8 +167,8 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { &L1_grad, &L2_grad, &W_grad, &L3_dgrad, &conv_data, &wspace, &transpose_perm }; - jit.execute(4, args, with_stream(forward_config_ref, stream)); - if(reinterpret_cast(wspace) != 0) { + jit.execute(DOUBLE_BACKWARD_A, args, with_stream(forward_config_ref, stream)); + if(deterministic) { void *fixup_args[] = {&wspace, &L3_dgrad}; KernelLaunchConfig fixup_config( forward_config_ref.num_blocks, @@ -159,11 +176,11 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { 0 ); fixup_config.hStream = stream; - jit.execute(2, fixup_args, fixup_config); + jit.execute(FIXUP_FORWARD, fixup_args, fixup_config); } - jit.execute(5, args, with_stream(double_backward_config_ref, stream)); - if(reinterpret_cast(wspace) != 0) { + jit.execute(DOUBLE_BACKWARD_B, args, with_stream(double_backward_config_ref, stream)); + if(deterministic) { void *fixup_args[] = {&wspace, &L1_grad}; KernelLaunchConfig fixup_config( double_backward_config_ref.num_blocks, @@ -171,7 +188,7 @@ class __attribute__ ((visibility ("default"))) JITConvImpl { 0 ); fixup_config.hStream = stream; - jit.execute(6, fixup_args, fixup_config); + jit.execute(FIXUP_DOUBLE_BACKWARD_B, fixup_args, fixup_config); } } diff --git a/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp b/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp index ddabd0bb..997733bb 100644 --- a/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp +++ b/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp @@ -47,6 +47,10 @@ Tensor tensor_zeros_like(const Tensor &ref, const std::vector &sizes) { return torch::zeros(sizes, ref.options()); } +Tensor tensor_zeros_bytes(const Tensor &ref, int64_t nbytes) { + return torch::zeros({nbytes}, ref.options().dtype(torch::kByte)); +} + void tensor_zero_(Tensor &tensor) { tensor.zero_(); } diff --git a/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp b/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp index 6bf3d51f..f0fb85a9 100644 --- a/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp +++ b/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp @@ -50,6 +50,12 @@ Tensor tensor_zeros_like(const Tensor &ref, const std::vector &sizes) { return out; } +Tensor tensor_zeros_bytes(const Tensor &ref, int64_t nbytes) { + std::vector sizes = {nbytes}; + auto sizes_ref = torch::headeronly::IntHeaderOnlyArrayRef(sizes.data(), sizes.size()); + return torch::stable::new_zeros(ref, sizes_ref, kByte); +} + void tensor_zero_(Tensor &tensor) { torch::stable::zero_(tensor); } diff --git a/openequivariance/openequivariance/extension/torch_core.hpp b/openequivariance/openequivariance/extension/torch_core.hpp index ab78d96a..fc255be8 100644 --- a/openequivariance/openequivariance/extension/torch_core.hpp +++ b/openequivariance/openequivariance/extension/torch_core.hpp @@ -5,6 +5,7 @@ #include #include #include +#include #include #include #include @@ -39,6 +40,7 @@ Tensor tensor_to_cpu_contiguous(const Tensor &tensor); Tensor tensor_contiguous(const Tensor &tensor); Tensor tensor_empty_like(const Tensor &ref, const std::vector &sizes); Tensor tensor_zeros_like(const Tensor &ref, const std::vector &sizes); +Tensor tensor_zeros_bytes(const Tensor &ref, int64_t nbytes); void tensor_zero_(Tensor &tensor); void alert_not_deterministic(const char *name); @@ -422,7 +424,6 @@ inline Tensor jit_conv_forward( int64_t L3_dim, Tensor rows, Tensor cols, - Tensor workspace, Tensor transpose_perm) { auto [jit_kernel, k] = compile_conv_with_caching(json_bytes, hash); @@ -433,7 +434,6 @@ inline Tensor jit_conv_forward( check_tensor(L1_in, {node_count, k.L1_dim}, k.irrep_dtype, "L1_in"); check_tensor(L2_in, {nnz, k.L2_dim}, k.irrep_dtype, "L2_in"); - check_tensor(workspace, {k.workspace_size}, k.workspace_dtype, "workspace"); check_tensor(rows, {nnz}, k.idx_dtype, "rows"); check_tensor(cols, {nnz}, k.idx_dtype, "cols"); @@ -454,7 +454,13 @@ inline Tensor jit_conv_forward( Tensor W_contig = tensor_contiguous(W); Tensor rows_contig = tensor_contiguous(rows); Tensor cols_contig = tensor_contiguous(cols); - Tensor workspace_contig = tensor_contiguous(workspace); + + std::optional workspace; + void *workspace_ptr = nullptr; + if (k.deterministic) { + workspace.emplace(tensor_zeros_bytes(L1_in, k.workspace_size)); + workspace_ptr = data_ptr(*workspace); + } jit_kernel->exec_conv( data_ptr(L1_contig), @@ -464,7 +470,7 @@ inline Tensor jit_conv_forward( data_ptr(rows_contig), data_ptr(cols_contig), nnz, node_count, - data_ptr(workspace_contig), + workspace_ptr, stream); return L3_out; @@ -478,7 +484,6 @@ inline tuple jit_conv_backward( Tensor L3_grad, Tensor rows, Tensor cols, - Tensor workspace, Tensor transpose_perm) { auto [jit_kernel, k] = compile_conv_with_caching(json_bytes, hash); @@ -490,7 +495,6 @@ inline tuple jit_conv_backward( check_tensor(L1_in, {node_count, k.L1_dim}, k.irrep_dtype, "L1_in"); check_tensor(L2_in, {nnz, k.L2_dim}, k.irrep_dtype, "L2_in"); check_tensor(L3_grad, {node_count, k.L3_dim}, k.irrep_dtype, "L3_grad"); - check_tensor(workspace, {k.workspace_size}, k.workspace_dtype, "workspace"); check_tensor(rows, {nnz}, k.idx_dtype, "rows"); check_tensor(cols, {nnz}, k.idx_dtype, "cols"); @@ -516,12 +520,18 @@ inline tuple jit_conv_backward( Tensor rows_contig = tensor_contiguous(rows); Tensor cols_contig = tensor_contiguous(cols); - Tensor workspace_contig = tensor_contiguous(workspace); Tensor transpose_perm_contig = tensor_contiguous(transpose_perm); if (k.shared_weights) tensor_zero_(W_grad); + std::optional workspace; + void *workspace_ptr = nullptr; + if (k.deterministic) { + workspace.emplace(tensor_zeros_bytes(L1_in, k.workspace_size)); + workspace_ptr = data_ptr(*workspace); + } + jit_kernel->backward( data_ptr(L1_in_contig), data_ptr(L1_grad), data_ptr(L2_in_contig), data_ptr(L2_grad), @@ -529,7 +539,7 @@ inline tuple jit_conv_backward( data_ptr(L3_grad_contig), data_ptr(rows_contig), data_ptr(cols_contig), nnz, node_count, - data_ptr(workspace_contig), + workspace_ptr, data_ptr(transpose_perm_contig), stream); @@ -547,7 +557,6 @@ inline tuple jit_conv_double_backward( Tensor W_dgrad, Tensor rows, Tensor cols, - Tensor workspace, Tensor transpose_perm) { auto [jit_kernel, k] = compile_conv_with_caching(json_bytes, hash); @@ -561,7 +570,6 @@ inline tuple jit_conv_double_backward( check_tensor(L3_grad, {node_count, k.L3_dim}, k.irrep_dtype, "L3_grad"); check_tensor(L1_dgrad, {node_count, k.L1_dim}, k.irrep_dtype, "L1_dgrad"); check_tensor(L2_dgrad, {nnz, k.L2_dim}, k.irrep_dtype, "L2_dgrad"); - check_tensor(workspace, {k.workspace_size}, k.workspace_dtype, "workspace"); check_tensor(rows, {nnz}, k.idx_dtype, "rows"); check_tensor(cols, {nnz}, k.idx_dtype, "cols"); @@ -594,12 +602,18 @@ inline tuple jit_conv_double_backward( Tensor rows_contig = tensor_contiguous(rows); Tensor cols_contig = tensor_contiguous(cols); - Tensor workspace_contig = tensor_contiguous(workspace); Tensor transpose_perm_contig = tensor_contiguous(transpose_perm); if (k.shared_weights) tensor_zero_(W_grad); + std::optional workspace; + void *workspace_ptr = nullptr; + if (k.deterministic) { + workspace.emplace(tensor_zeros_bytes(L1_in, k.workspace_size)); + workspace_ptr = data_ptr(*workspace); + } + jit_kernel->double_backward( data_ptr(L1_in_contig), data_ptr(L2_in_contig), data_ptr(W_contig), data_ptr(L3_grad_contig), @@ -609,7 +623,7 @@ inline tuple jit_conv_double_backward( data_ptr(W_grad), data_ptr(L3_dgrad), data_ptr(rows_contig), data_ptr(cols_contig), nnz, node_count, - data_ptr(workspace_contig), data_ptr(transpose_perm_contig), + workspace_ptr, data_ptr(transpose_perm_contig), stream ); @@ -669,9 +683,9 @@ REGISTER_LIBRARY(libtorch_tp_jit, m) { m.def("jit_tp_backward(Tensor json_bytes, int hash, Tensor L1_in, Tensor L2_in, Tensor W, Tensor L3_grad) -> (Tensor, Tensor, Tensor)"); m.def("jit_tp_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, Tensor, Tensor, Tensor)"); - m.def("jit_conv_forward(Tensor json_bytes, int hash, Tensor L1_in, Tensor L2_in, Tensor W, int L3_dim, Tensor rows, Tensor cols, Tensor workspace, Tensor transpose_perm) -> Tensor"); - 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)"); + m.def("jit_conv_forward(Tensor json_bytes, int hash, Tensor L1_in, Tensor L2_in, Tensor W, int L3_dim, Tensor rows, Tensor cols, Tensor transpose_perm) -> Tensor"); + 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 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 transpose_perm) -> (Tensor, Tensor, Tensor, Tensor)"); 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"); }; diff --git a/openequivariance/openequivariance/jax/TensorProductConv.py b/openequivariance/openequivariance/jax/TensorProductConv.py index 9234158f..fbd2292d 100644 --- a/openequivariance/openequivariance/jax/TensorProductConv.py +++ b/openequivariance/openequivariance/jax/TensorProductConv.py @@ -53,10 +53,6 @@ def __init__( self.weight_numel = config.weight_numel self.L3_dim = self.config.irreps_out.dim - self.workspace = jnp.zeros((self.workspace_size,), dtype=jnp.uint8) - logger.info( - f"Convolution requires {self.workspace_size // (2**20)}MB of workspace." - ) self.dummy_transpose_perm = jnp.zeros((1,), dtype=jnp.int32) def forward( @@ -86,7 +82,6 @@ def forward( W, rows, cols, - self.workspace, sender_perm, L3_dim=self.L3_dim, kernel=self.kernel, diff --git a/openequivariance/openequivariance/jax/jvp/conv_prim.py b/openequivariance/openequivariance/jax/jvp/conv_prim.py index 8d386463..994a475e 100644 --- a/openequivariance/openequivariance/jax/jvp/conv_prim.py +++ b/openequivariance/openequivariance/jax/jvp/conv_prim.py @@ -2,7 +2,11 @@ import jax.numpy as jnp from jax.extend import core from jax.interpreters import mlir, ad, batching -from openequivariance.jax.utils import clean_tensors +from openequivariance.jax.utils import ( + clean_tensors, + conv_workspace_shape, + conv_workspace_zeros, +) # ============================================================================== # 1. Forward Primitive @@ -11,16 +15,29 @@ conv_fwd_p = core.Primitive("conv_fwd") -def conv_fwd_impl(X, Y, W, rows, cols, workspace, sender_perm, *, L3_dim, kernel, hash): +def conv_fwd_impl(X, Y, W, rows, cols, sender_perm, *, L3_dim, kernel, hash): irrep_dtype = X.dtype out_shape = jax.ShapeDtypeStruct((X.shape[0], L3_dim), irrep_dtype) - call = jax.ffi.ffi_call("conv_forward", out_shape) - return call(X, Y, W, rows, cols, workspace, sender_perm, kernel=kernel, hash=hash) + call = jax.ffi.ffi_call( + "conv_forward", + (out_shape, conv_workspace_shape(kernel)), + input_output_aliases={5: 1}, + ) + out, _workspace = call( + X, + Y, + W, + rows, + cols, + conv_workspace_zeros(kernel), + sender_perm, + kernel=kernel, + hash=hash, + ) + return out -def conv_fwd_abstract_eval( - X, Y, W, rows, cols, workspace, sender_perm, *, L3_dim, kernel, hash -): +def conv_fwd_abstract_eval(X, Y, W, rows, cols, sender_perm, *, L3_dim, kernel, hash): return jax.core.ShapedArray((X.shape[0], L3_dim), X.dtype) @@ -42,22 +59,31 @@ def conv_fwd_abstract_eval( conv_bwd_p.multiple_results = True -def conv_bwd_impl(X, Y, W, dZ, rows, cols, workspace, sender_perm, *, kernel, hash): +def conv_bwd_impl(X, Y, W, dZ, rows, cols, sender_perm, *, kernel, hash): irrep_dtype = X.dtype out_shapes = ( jax.ShapeDtypeStruct(X.shape, irrep_dtype), jax.ShapeDtypeStruct(Y.shape, irrep_dtype), jax.ShapeDtypeStruct(W.shape, irrep_dtype), + conv_workspace_shape(kernel), ) - call = jax.ffi.ffi_call("conv_backward", out_shapes) - return call( - X, Y, W, dZ, rows, cols, workspace, sender_perm, kernel=kernel, hash=hash + call = jax.ffi.ffi_call("conv_backward", out_shapes, input_output_aliases={6: 3}) + dX, dY, dW, _workspace = call( + X, + Y, + W, + dZ, + rows, + cols, + conv_workspace_zeros(kernel), + sender_perm, + kernel=kernel, + hash=hash, ) + return dX, dY, dW -def conv_bwd_abstract_eval( - X, Y, W, dZ, rows, cols, workspace, sender_perm, *, kernel, hash -): +def conv_bwd_abstract_eval(X, Y, W, dZ, rows, cols, sender_perm, *, kernel, hash): irrep_dtype = X.dtype return ( jax.core.ShapedArray(X.shape, irrep_dtype), @@ -85,7 +111,7 @@ def conv_bwd_abstract_eval( def conv_dbwd_impl( - X, Y, W, dZ, ddX, ddY, ddW, rows, cols, workspace, sender_perm, *, kernel, hash + X, Y, W, dZ, ddX, ddY, ddW, rows, cols, sender_perm, *, kernel, hash ): irrep_dtype = X.dtype out_shapes = ( @@ -93,9 +119,12 @@ def conv_dbwd_impl( jax.ShapeDtypeStruct(Y.shape, irrep_dtype), jax.ShapeDtypeStruct(W.shape, irrep_dtype), jax.ShapeDtypeStruct(dZ.shape, irrep_dtype), + conv_workspace_shape(kernel), + ) + call = jax.ffi.ffi_call( + "conv_double_backward", out_shapes, input_output_aliases={9: 4} ) - call = jax.ffi.ffi_call("conv_double_backward", out_shapes) - return call( + gX, gY, gW, gdZ, _workspace = call( X, Y, W, @@ -105,15 +134,16 @@ def conv_dbwd_impl( ddW, rows, cols, - workspace, + conv_workspace_zeros(kernel), sender_perm, kernel=kernel, hash=hash, ) + return gX, gY, gW, gdZ def conv_dbwd_abstract_eval( - X, Y, W, dZ, ddX, ddY, ddW, rows, cols, workspace, sender_perm, *, kernel, hash + X, Y, W, dZ, ddX, ddY, ddW, rows, cols, sender_perm, *, kernel, hash ): irrep_dtype = X.dtype return ( @@ -142,10 +172,10 @@ def conv_dbwd_abstract_eval( def conv_fwd_jvp_impl( - X, Y, W, dX, dY, dW, rows, cols, workspace, sender_perm, *, L3_dim, kernel, hash + X, Y, W, dX, dY, dW, rows, cols, sender_perm, *, L3_dim, kernel, hash ): kwargs = dict(L3_dim=L3_dim, kernel=kernel, hash=hash) - args_meta = (rows, cols, workspace, sender_perm) + args_meta = (rows, cols, sender_perm) term1 = conv_fwd_p.bind(dX, Y, W, *args_meta, **kwargs) term2 = conv_fwd_p.bind(X, dY, W, *args_meta, **kwargs) @@ -154,7 +184,7 @@ def conv_fwd_jvp_impl( def conv_fwd_jvp_abstract_eval( - X, Y, W, dX, dY, dW, rows, cols, workspace, sender_perm, *, L3_dim, kernel, hash + X, Y, W, dX, dY, dW, rows, cols, sender_perm, *, L3_dim, kernel, hash ): return jax.core.ShapedArray((X.shape[0], L3_dim), X.dtype) @@ -179,15 +209,15 @@ def conv_fwd_jvp_abstract_eval( def conv_fwd_jvp_transpose( - ct, X, Y, W, dX, dY, dW, rows, cols, workspace, sender_perm, *, L3_dim, kernel, hash + ct, X, Y, W, dX, dY, dW, rows, cols, sender_perm, *, L3_dim, kernel, hash ): X, Y, W = clean_tensors(X, Y, W) grad_X, grad_Y, grad_W = conv_bwd_p.bind( - X, Y, W, ct, rows, cols, workspace, sender_perm, kernel=kernel, hash=hash + X, Y, W, ct, rows, cols, sender_perm, kernel=kernel, hash=hash ) - return (None, None, None, grad_X, grad_Y, grad_W, None, None, None, None) + return (None, None, None, grad_X, grad_Y, grad_W, None, None, None) ad.primitive_transposes[conv_fwd_jvp_p] = conv_fwd_jvp_transpose @@ -199,8 +229,8 @@ def conv_fwd_jvp_transpose( def conv_fwd_jvp_rule(primals, tangents, *, L3_dim, kernel, hash): - X, Y, W, rows, cols, workspace, sender_perm = primals - dX, dY, dW, drows, dcols, dworkspace, dsender_perm = tangents + X, Y, W, rows, cols, sender_perm = primals + dX, dY, dW, drows, dcols, dsender_perm = tangents dX, dY, dW = clean_tensors(dX, dY, dW) out_primal = conv_fwd_p.bind( @@ -209,7 +239,6 @@ def conv_fwd_jvp_rule(primals, tangents, *, L3_dim, kernel, hash): W, rows, cols, - workspace, sender_perm, L3_dim=L3_dim, kernel=kernel, @@ -224,7 +253,6 @@ def conv_fwd_jvp_rule(primals, tangents, *, L3_dim, kernel, hash): dW, rows, cols, - workspace, sender_perm, L3_dim=L3_dim, kernel=kernel, @@ -245,9 +273,9 @@ def conv_fwd_jvp_rule(primals, tangents, *, L3_dim, kernel, hash): def conv_fwd_jvp_jvp_rule(primals, tangents, *, L3_dim, kernel, hash): tangents_clean = tuple(clean_tensors(*tangents)) - def func(x, y, w, dx, dy, dw, r, c, ws, sp): + def func(x, y, w, dx, dy, dw, r, c, sp): return conv_fwd_jvp_impl( - x, y, w, dx, dy, dw, r, c, ws, sp, L3_dim=L3_dim, kernel=kernel, hash=hash + x, y, w, dx, dy, dw, r, c, sp, L3_dim=L3_dim, kernel=kernel, hash=hash ) return jax.jvp(func, primals, tangents_clean) @@ -265,10 +293,10 @@ def func(x, y, w, dx, dy, dw, r, c, ws, sp): def conv_bwd_jvp_impl( - X, Y, W, dZ, tX, tY, tW, tdZ, rows, cols, workspace, sender_perm, *, kernel, hash + X, Y, W, dZ, tX, tY, tW, tdZ, rows, cols, sender_perm, *, kernel, hash ): kwargs = dict(kernel=kernel, hash=hash) - args_meta = (rows, cols, workspace, sender_perm) + args_meta = (rows, cols, sender_perm) term_dZ = conv_bwd_p.bind(X, Y, W, tdZ, *args_meta, **kwargs) term_X = conv_bwd_p.bind(tX, Y, W, dZ, *args_meta, **kwargs) @@ -283,7 +311,7 @@ def conv_bwd_jvp_impl( def conv_bwd_jvp_abstract_eval( - X, Y, W, dZ, tX, tY, tW, tdZ, rows, cols, workspace, sender_perm, *, kernel, hash + X, Y, W, dZ, tX, tY, tW, tdZ, rows, cols, sender_perm, *, kernel, hash ): irrep_dtype = X.dtype return ( @@ -324,7 +352,6 @@ def conv_bwd_jvp_transpose( tdZ, rows, cols, - workspace, sender_perm, *, kernel, @@ -344,10 +371,10 @@ def conv_bwd_jvp_transpose( tensors_clean = clean_tensors(X, Y, W, dZ, ddX, ddY, ddW) g_X, g_Y, g_W, g_dZ = conv_dbwd_p.bind( - *tensors_clean, rows, cols, workspace, sender_perm, kernel=kernel, hash=hash + *tensors_clean, rows, cols, sender_perm, kernel=kernel, hash=hash ) - return (None, None, None, None, g_X, g_Y, g_W, g_dZ, None, None, None, None) + return (None, None, None, None, g_X, g_Y, g_W, g_dZ, None, None, None) ad.primitive_transposes[conv_bwd_jvp_p] = conv_bwd_jvp_transpose @@ -361,9 +388,9 @@ def conv_bwd_jvp_transpose( def conv_bwd_jvp_jvp_rule(primals, tangents, *, kernel, hash): tangents_clean = tuple(clean_tensors(*tangents)) - def func(x, y, w, dz, tx, ty, tw, tdz, r, c, ws, sp): + def func(x, y, w, dz, tx, ty, tw, tdz, r, c, sp): return conv_bwd_jvp_impl( - x, y, w, dz, tx, ty, tw, tdz, r, c, ws, sp, kernel=kernel, hash=hash + x, y, w, dz, tx, ty, tw, tdz, r, c, sp, kernel=kernel, hash=hash ) return jax.jvp(func, primals, tangents_clean) @@ -378,13 +405,13 @@ def func(x, y, w, dz, tx, ty, tw, tdz, r, c, ws, sp): def conv_bwd_jvp_rule(primals, tangents, *, kernel, hash): - X, Y, W, dZ, rows, cols, workspace, sender_perm = primals - tX, tY, tW, tdZ, drows, dcols, dworkspace, dsender_perm = tangents + X, Y, W, dZ, rows, cols, sender_perm = primals + tX, tY, tW, tdZ, drows, dcols, dsender_perm = tangents tX, tY, tW, tdZ = clean_tensors(tX, tY, tW, tdZ) out_primal = conv_bwd_p.bind( - X, Y, W, dZ, rows, cols, workspace, sender_perm, kernel=kernel, hash=hash + X, Y, W, dZ, rows, cols, sender_perm, kernel=kernel, hash=hash ) out_tangent = conv_bwd_jvp_p.bind( X, @@ -397,7 +424,6 @@ def conv_bwd_jvp_rule(primals, tangents, *, kernel, hash): tdZ, rows, cols, - workspace, sender_perm, kernel=kernel, hash=hash, @@ -424,7 +450,6 @@ def conv_dbwd_slow( ddW, rows, cols, - workspace, sender_perm, *, L3_dim, @@ -432,7 +457,7 @@ def conv_dbwd_slow( hash, ): kwargs = dict(kernel=kernel, hash=hash) - args_meta = (rows, cols, workspace, sender_perm) + args_meta = (rows, cols, sender_perm) op1 = conv_bwd_p.bind(ddX, ddY, W, dZ, *args_meta, **kwargs) op2 = conv_bwd_p.bind(X, Y, ddW, dZ, *args_meta, **kwargs) @@ -461,7 +486,7 @@ def conv_dbwd_jvp_rule(primals, tangents, *, kernel, hash): dZ = primals[3] # Infer L3_dim from dZ (4th input) L3_dim = dZ.shape[1] - def func(x, y, w, dz, ddx, ddy, ddw, r, c, ws, sp): + def func(x, y, w, dz, ddx, ddy, ddw, r, c, sp): return conv_dbwd_slow( x, y, @@ -472,7 +497,6 @@ def func(x, y, w, dz, ddx, ddy, ddw, r, c, ws, sp): ddw, r, c, - ws, sp, L3_dim=L3_dim, kernel=kernel, @@ -492,13 +516,13 @@ def func(x, y, w, dz, ddx, ddy, ddw, r, c, ws, sp): def conv_dbwd_transpose( - ct, X, Y, W, dZ, ddX, ddY, ddW, rows, cols, workspace, sender_perm, *, kernel, hash + ct, X, Y, W, dZ, ddX, ddY, ddW, rows, cols, sender_perm, *, kernel, hash ): L3_dim = dZ.shape[1] X, Y, W, dZ, ddX, ddY, ddW = clean_tensors(X, Y, W, dZ, ddX, ddY, ddW) - def func(x, y, w, dz, ddx, ddy, ddw, r, c, ws, sp): + def func(x, y, w, dz, ddx, ddy, ddw, r, c, sp): return conv_dbwd_slow( x, y, @@ -509,16 +533,13 @@ def func(x, y, w, dz, ddx, ddy, ddw, r, c, ws, sp): ddw, r, c, - ws, sp, L3_dim=L3_dim, kernel=kernel, hash=hash, ) - _, vjp_fun = jax.vjp( - func, X, Y, W, dZ, ddX, ddY, ddW, rows, cols, workspace, sender_perm - ) + _, vjp_fun = jax.vjp(func, X, Y, W, dZ, ddX, ddY, ddW, rows, cols, sender_perm) input_grads = vjp_fun(ct) return input_grads @@ -555,17 +576,16 @@ def flatten_args(vector_arg_values, batch_axes): B = find_batch_size(vector_arg_values, batch_axes) new_args = [] - for i, (arg, batch_axis) in enumerate(zip(vector_arg_values, batch_axes)): - if i != len(vector_arg_values) - 2: - if batch_axis is None and arg is not None: - arg = jnp.broadcast_to(arg, (B,) + arg.shape) - elif batch_axis is not None and batch_axis != 0: - arg = jnp.moveaxis(arg, batch_axis, 0) + for arg, batch_axis in zip(vector_arg_values, batch_axes): + if batch_axis is None and arg is not None: + arg = jnp.broadcast_to(arg, (B,) + arg.shape) + elif batch_axis is not None and batch_axis != 0: + arg = jnp.moveaxis(arg, batch_axis, 0) new_args.append(arg) vector_arg_values = new_args - rows, cols, workspace, sender_perm = vector_arg_values[-4:] + rows, cols, sender_perm = vector_arg_values[-3:] rows_offset, cols_offset, sender_perm_offset = rows, cols, sender_perm if B > 1: batch_offsets = (jnp.arange(B) * num_nodes).astype(rows.dtype) @@ -575,10 +595,9 @@ def flatten_args(vector_arg_values, batch_axes): if sender_perm is not None: sender_perm_offset = sender_perm + batch_offsets[:, None] - new_args = [arg.reshape(-1, *arg.shape[2:]) for arg in vector_arg_values[:-4]] + [ + new_args = [arg.reshape(-1, *arg.shape[2:]) for arg in vector_arg_values[:-3]] + [ jnp.ravel(rows_offset), jnp.ravel(cols_offset), - workspace, jnp.ravel(sender_perm_offset), ] diff --git a/openequivariance/openequivariance/jax/utils.py b/openequivariance/openequivariance/jax/utils.py index 371b0ae5..c59a68ed 100644 --- a/openequivariance/openequivariance/jax/utils.py +++ b/openequivariance/openequivariance/jax/utils.py @@ -1,9 +1,23 @@ +import functools +import json + import jax import jax.numpy as jnp import numpy as np from jax.interpreters import ad +@functools.cache +def conv_workspace_shape(kernel: str) -> jax.ShapeDtypeStruct: + size = json.loads(kernel)["kernel_prop"]["workspace_size"] + return jax.ShapeDtypeStruct((int(size),), jnp.uint8) + + +def conv_workspace_zeros(kernel: str) -> jax.Array: + shape = conv_workspace_shape(kernel) + return jnp.zeros(shape.shape, shape.dtype) + + def reorder_jax_helper(schedule, weights_in, direction, has_batch_dim): assert direction in ["forward", "backward"] diff --git a/openequivariance/openequivariance/jax/vjp/conv_func.py b/openequivariance/openequivariance/jax/vjp/conv_func.py index 20d7ccc2..4f35e084 100644 --- a/openequivariance/openequivariance/jax/vjp/conv_func.py +++ b/openequivariance/openequivariance/jax/vjp/conv_func.py @@ -2,30 +2,46 @@ import jax.numpy as jnp from functools import partial +from openequivariance.jax.utils import conv_workspace_shape, conv_workspace_zeros + def zeros_like(x): return jnp.zeros_like(x) -@partial(jax.custom_vjp, nondiff_argnums=(5, 6, 7, 8, 9)) -def forward(X, Y, W, rows, cols, workspace, sender_perm, L3_dim, kernel, hash): +@partial(jax.custom_vjp, nondiff_argnums=(5, 6, 7, 8)) +def forward(X, Y, W, rows, cols, sender_perm, L3_dim, kernel, hash): forward_call = jax.ffi.ffi_call( - "conv_forward", jax.ShapeDtypeStruct((X.shape[0], L3_dim), X.dtype) + "conv_forward", + ( + jax.ShapeDtypeStruct((X.shape[0], L3_dim), X.dtype), + conv_workspace_shape(kernel), + ), + input_output_aliases={5: 1}, ) - return forward_call( - X, Y, W, rows, cols, workspace, sender_perm, kernel=kernel, hash=hash + out, _workspace = forward_call( + X, + Y, + W, + rows, + cols, + conv_workspace_zeros(kernel), + sender_perm, + kernel=kernel, + hash=hash, ) + return out -def forward_fwd(X, Y, W, rows, cols, workspace, sender_perm, L3_dim, kernel, hash): - out = forward(X, Y, W, rows, cols, workspace, sender_perm, L3_dim, kernel, hash) +def forward_fwd(X, Y, W, rows, cols, sender_perm, L3_dim, kernel, hash): + out = forward(X, Y, W, rows, cols, sender_perm, L3_dim, kernel, hash) return out, (X, Y, W, rows, cols) -def forward_bwd(workspace, sender_perm, L3_dim, kernel, hash, res, dZ): +def forward_bwd(sender_perm, L3_dim, kernel, hash, res, dZ): X, Y, W, rows, cols = res dX, dY, dW = backward( - X, Y, W, dZ, rows, cols, workspace, sender_perm, kernel=kernel, hash=hash + X, Y, W, dZ, rows, cols, sender_perm, kernel=kernel, hash=hash ) return dX, dY, dW, None, None @@ -33,27 +49,39 @@ def forward_bwd(workspace, sender_perm, L3_dim, kernel, hash, res, dZ): forward.defvjp(forward_fwd, forward_bwd) -@partial(jax.custom_vjp, nondiff_argnums=(6, 7, 8, 9)) -def backward(X, Y, W, dZ, rows, cols, workspace, sender_perm, kernel, hash): +@partial(jax.custom_vjp, nondiff_argnums=(6, 7, 8)) +def backward(X, Y, W, dZ, rows, cols, sender_perm, kernel, hash): backward_call = jax.ffi.ffi_call( "conv_backward", ( jax.ShapeDtypeStruct(X.shape, X.dtype), jax.ShapeDtypeStruct(Y.shape, Y.dtype), jax.ShapeDtypeStruct(W.shape, W.dtype), + conv_workspace_shape(kernel), ), + input_output_aliases={6: 3}, ) - return backward_call( - X, Y, W, dZ, rows, cols, workspace, sender_perm, kernel=kernel, hash=hash + dX, dY, dW, _workspace = backward_call( + X, + Y, + W, + dZ, + rows, + cols, + conv_workspace_zeros(kernel), + sender_perm, + kernel=kernel, + hash=hash, ) + return dX, dY, dW -def backward_fwd(X, Y, W, dZ, rows, cols, workspace, sender_perm, kernel, hash): - out = backward(X, Y, W, dZ, rows, cols, workspace, sender_perm, kernel, hash) +def backward_fwd(X, Y, W, dZ, rows, cols, sender_perm, kernel, hash): + out = backward(X, Y, W, dZ, rows, cols, sender_perm, kernel, hash) return out, (X, Y, W, dZ, rows, cols) -def backward_bwd(workspace, sender_perm, kernel, hash, res, derivatives): +def backward_bwd(sender_perm, kernel, hash, res, derivatives): X, Y, W, dZ, rows, cols = res ddX, ddY, ddW = derivatives @@ -67,7 +95,6 @@ def backward_bwd(workspace, sender_perm, kernel, hash, res, derivatives): ddW, rows, cols, - workspace, sender_perm, kernel, hash, @@ -79,10 +106,8 @@ def backward_bwd(workspace, sender_perm, kernel, hash, res, derivatives): backward.defvjp(backward_fwd, backward_bwd) -@partial(jax.custom_vjp, nondiff_argnums=(9, 10, 11, 12)) -def double_backward( - X, Y, W, dZ, ddX, ddY, ddW, rows, cols, workspace, sender_perm, kernel, hash -): +@partial(jax.custom_vjp, nondiff_argnums=(9, 10, 11)) +def double_backward(X, Y, W, dZ, ddX, ddY, ddW, rows, cols, sender_perm, kernel, hash): double_backward_call = jax.ffi.ffi_call( "conv_double_backward", ( @@ -90,9 +115,11 @@ def double_backward( jax.ShapeDtypeStruct(Y.shape, Y.dtype), jax.ShapeDtypeStruct(W.shape, W.dtype), jax.ShapeDtypeStruct(dZ.shape, dZ.dtype), + conv_workspace_shape(kernel), ), + input_output_aliases={9: 4}, ) - return double_backward_call( + gX, gY, gW, gdZ, _workspace = double_backward_call( X, Y, W, @@ -102,24 +129,24 @@ def double_backward( ddW, rows, cols, - workspace, + conv_workspace_zeros(kernel), sender_perm, kernel=kernel, hash=hash, ) + return gX, gY, gW, gdZ def double_backward_fwd( - X, Y, W, dZ, ddX, ddY, ddW, rows, cols, workspace, sender_perm, kernel, hash + X, Y, W, dZ, ddX, ddY, ddW, rows, cols, sender_perm, kernel, hash ): out = double_backward( - X, Y, W, dZ, ddX, ddY, ddW, rows, cols, workspace, sender_perm, kernel, hash + X, Y, W, dZ, ddX, ddY, ddW, rows, cols, sender_perm, kernel, hash ) return out, (X, Y, W, dZ, ddX, ddY, ddW, rows, cols) def triple_backward( - workspace, sender_perm, kernel, hash, @@ -129,7 +156,7 @@ def triple_backward( X, Y, W, dZ, ddX, ddY, ddW, rows, cols = residuals t_dX, t_dY, t_dW, t_ddZ = tangent_outputs - common_args = (rows, cols, workspace, sender_perm, kernel, hash) + common_args = (rows, cols, sender_perm, kernel, hash) op1_inputs = (ddX, ddY, W, dZ, t_dX, t_dY, zeros_like(W)) g1_ddX, g1_ddY, g1_W, g1_dZ = double_backward(*op1_inputs, *common_args) diff --git a/openequivariance/openequivariance/templates/loop_unroll_conv_atomic.cuh b/openequivariance/openequivariance/templates/loop_unroll_conv_atomic.cuh index 3d461dbc..9250d077 100644 --- a/openequivariance/openequivariance/templates/loop_unroll_conv_atomic.cuh +++ b/openequivariance/openequivariance/templates/loop_unroll_conv_atomic.cuh @@ -31,24 +31,6 @@ struct ConvData { unsigned long node_count; }; -__global__ void -{{ launch_bounds(forward_schedule) }} -fixup_forward(void* workspace, IRREP_T* dst_ptr) { - // Empty, no fixup -} - -__global__ void -{{ launch_bounds(backward_schedule) }} -fixup_backward(void* workspace, IRREP_T* dst_ptr) { - // Empty, no fixup -} - -__global__ void -{{ launch_bounds(double_backward_schedule) }} -fixup_double_backwardB(void* workspace, IRREP_T* dst_ptr) { - // Empty, no fixup -} - __global__ void {{ launch_bounds(forward_schedule) }} forward(IRREP_T* L1_in, diff --git a/openequivariance_extjax/src/ffi_handlers.cpp b/openequivariance_extjax/src/ffi_handlers.cpp index fb3a702e..6b2d2999 100644 --- a/openequivariance_extjax/src/ffi_handlers.cpp +++ b/openequivariance_extjax/src/ffi_handlers.cpp @@ -466,9 +466,10 @@ ffi::Error conv_forward_impl( ffi::AnyBuffer W, ffi::AnyBuffer rows, ffi::AnyBuffer cols, - ffi::AnyBuffer workspace, + ffi::AnyBuffer workspace_in, ffi::AnyBuffer transpose_perm, ffi::Result L3_out, + ffi::Result workspace, stream_t stream, std::string_view kernel_json, int64_t hash) { @@ -477,19 +478,18 @@ ffi::Error conv_forward_impl( kernel_json, hash, true); const int64_t nnz = rows.dimensions()[0]; const int64_t node_count = L1_in.dimensions()[0]; - void* workspace_ptr = data_ptr(workspace); + void* workspace_ptr = nullptr; check_tensor(L1_in, {node_count, k.L1_dim}, k.irrep_dtype, "L1_in"); check_tensor(L2_in, {nnz, k.L2_dim}, k.irrep_dtype, "L2_in"); - check_tensor(workspace, {k.workspace_size}, k.workspace_dtype, "workspace"); + check_tensor(workspace_in, {k.workspace_size}, k.workspace_dtype, "workspace"); + check_tensor(*workspace, {k.workspace_size}, k.workspace_dtype, "workspace"); check_tensor(rows, {nnz}, k.idx_dtype, "rows"); check_tensor(cols, {nnz}, k.idx_dtype, "cols"); - if (k.deterministic){ + if (k.deterministic) { check_tensor(transpose_perm, {nnz}, k.idx_dtype, "transpose perm"); - } - else { - workspace_ptr = nullptr; + workspace_ptr = data_ptr(workspace); } zero_buffer(*L3_out, stream); @@ -522,8 +522,9 @@ ffi::Error conv_backward_impl( ffi::Result W_grad, ffi::AnyBuffer rows, ffi::AnyBuffer cols, - ffi::AnyBuffer workspace, + ffi::AnyBuffer workspace_in, ffi::AnyBuffer transpose_perm, + ffi::Result workspace, stream_t stream, std::string_view kernel_json, int64_t hash) { @@ -532,20 +533,19 @@ ffi::Error conv_backward_impl( kernel_json, hash, true); const int64_t nnz = rows.dimensions()[0]; const int64_t node_count = L1_in.dimensions()[0]; - void* workspace_ptr = data_ptr(workspace); + void* workspace_ptr = nullptr; check_tensor(L1_in, {node_count, k.L1_dim}, k.irrep_dtype, "L1_in"); check_tensor(L2_in, {nnz, k.L2_dim}, k.irrep_dtype, "L2_in"); check_tensor(L3_grad, {node_count, k.L3_dim}, k.irrep_dtype, "L3_grad"); - check_tensor(workspace, {k.workspace_size}, k.workspace_dtype, "workspace"); + check_tensor(workspace_in, {k.workspace_size}, k.workspace_dtype, "workspace"); + check_tensor(*workspace, {k.workspace_size}, k.workspace_dtype, "workspace"); check_tensor(rows, {nnz}, k.idx_dtype, "rows"); check_tensor(cols, {nnz}, k.idx_dtype, "cols"); if (k.deterministic) { check_tensor(transpose_perm, {nnz}, k.idx_dtype, "transpose perm"); - } - else { - workspace_ptr = nullptr; + workspace_ptr = data_ptr(workspace); } zero_buffer(*L1_grad, stream); zero_buffer(*L2_grad, stream); @@ -591,8 +591,9 @@ ffi::Error conv_double_backward_impl( ffi::Result L3_dgrad, ffi::AnyBuffer rows, ffi::AnyBuffer cols, - ffi::AnyBuffer workspace, + ffi::AnyBuffer workspace_in, ffi::AnyBuffer transpose_perm, + ffi::Result workspace, stream_t stream, std::string_view kernel_json, int64_t hash) { @@ -601,22 +602,21 @@ ffi::Error conv_double_backward_impl( kernel_json, hash, true); const int64_t nnz = rows.dimensions()[0]; const int64_t node_count = L1_in.dimensions()[0]; - void* workspace_ptr = data_ptr(workspace); + void* workspace_ptr = nullptr; check_tensor(L1_in, {node_count, k.L1_dim}, k.irrep_dtype, "L1_in"); check_tensor(L2_in, {nnz, k.L2_dim}, k.irrep_dtype, "L2_in"); check_tensor(L3_grad, {node_count, k.L3_dim}, k.irrep_dtype, "L3_grad"); check_tensor(L1_dgrad, {node_count, k.L1_dim}, k.irrep_dtype, "L1_dgrad"); check_tensor(L2_dgrad, {nnz, k.L2_dim}, k.irrep_dtype, "L2_dgrad"); - check_tensor(workspace, {k.workspace_size}, k.workspace_dtype, "workspace"); + check_tensor(workspace_in, {k.workspace_size}, k.workspace_dtype, "workspace"); + check_tensor(*workspace, {k.workspace_size}, k.workspace_dtype, "workspace"); check_tensor(rows, {nnz}, k.idx_dtype, "rows"); check_tensor(cols, {nnz}, k.idx_dtype, "cols"); if (k.deterministic) { check_tensor(transpose_perm, {nnz}, k.idx_dtype, "transpose perm"); - } - else { - workspace_ptr = nullptr; + workspace_ptr = data_ptr(workspace); } zero_buffer(*L1_grad, stream); zero_buffer(*L2_grad, stream); @@ -694,6 +694,7 @@ XLA_FFI_DEFINE_HANDLER_SYMBOL( .Arg() .Arg() .Ret() + .Ret() .Ctx>() .Attr("kernel") .Attr("hash"), @@ -713,6 +714,7 @@ XLA_FFI_DEFINE_HANDLER_SYMBOL( .Arg() .Arg() .Arg() + .Ret() .Ctx>() .Attr("kernel") .Attr("hash"), @@ -736,6 +738,7 @@ XLA_FFI_DEFINE_HANDLER_SYMBOL( .Arg() .Arg() .Arg() + .Ret() .Ctx>() .Attr("kernel") .Attr("hash"), diff --git a/tests/cuda_graph_test.py b/tests/cuda_graph_test.py new file mode 100644 index 00000000..c0663827 --- /dev/null +++ b/tests/cuda_graph_test.py @@ -0,0 +1,123 @@ +import pytest +import torch + +import openequivariance as oeq +from openequivariance.benchmark.problems import mace_problems + + +def _sorted_graph(node_count, nnz, gen): + rows = torch.randint(0, node_count, (nnz,), device="cuda", generator=gen) + cols = torch.randint(0, node_count, (nnz,), device="cuda", generator=gen) + order = torch.argsort(rows * node_count + cols) + rows, cols = rows[order], cols[order] + sender_perm = torch.argsort(cols * node_count + rows) + return rows, cols, sender_perm + + +@pytest.fixture(scope="module") +def problem(): + return mace_problems()[0] + + +@pytest.fixture(params=[False, True], ids=["atomic", "deterministic"]) +def conv_and_inputs(request, problem): + deterministic = request.param + gen = torch.Generator(device="cuda") + gen.manual_seed(0) + + node_count, nnz = 2000, 40000 + conv = oeq.TensorProductConv(problem, deterministic=deterministic) + X = torch.randn(node_count, problem.irreps_in1.dim, device="cuda", generator=gen) + Y = torch.randn(nnz, problem.irreps_in2.dim, device="cuda", generator=gen) + W = torch.randn(nnz, problem.weight_numel, device="cuda", generator=gen) + rows, cols, sender_perm = _sorted_graph(node_count, nnz, gen) + if not deterministic: + sender_perm = None + G = torch.randn(node_count, problem.irreps_out.dim, device="cuda", generator=gen) + return conv, deterministic, (X, Y, W, rows, cols, sender_perm, G) + + +def _fwd_bwd(conv, X, Y, W, rows, cols, sender_perm, G): + out = conv(X, Y, W, rows, cols, sender_perm) + gX, gY, gW = torch.autograd.grad((out * G).sum(), (X, Y, W)) + return out, gX, gY, gW + + +def _assert_close(actual, expected, deterministic): + for a, e in zip(actual, expected): + if deterministic: + assert torch.equal(a, e) + else: + assert torch.allclose(a, e, atol=1e-4, rtol=1e-4) + + +def test_eager_cuda_graph_capture(conv_and_inputs): + conv, deterministic, (X, Y, W, rows, cols, sender_perm, G) = conv_and_inputs + X, Y, W = (t.clone().requires_grad_(True) for t in (X, Y, W)) + + reference = [ + t.detach().clone() for t in _fwd_bwd(conv, X, Y, W, rows, cols, sender_perm, G) + ] + + side = torch.cuda.Stream() + side.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(side): + for _ in range(3): + _fwd_bwd(conv, X, Y, W, rows, cols, sender_perm, G) + torch.cuda.current_stream().wait_stream(side) + + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + static_outputs = _fwd_bwd(conv, X, Y, W, rows, cols, sender_perm, G) + + for _ in range(3): + for t in static_outputs: + t.zero_() + graph.replay() + torch.cuda.synchronize() + _assert_close(static_outputs, reference, deterministic) + + +def test_compile_reduce_overhead(conv_and_inputs): + conv, deterministic, (X, Y, W, rows, cols, sender_perm, G) = conv_and_inputs + X, Y, W = (t.clone().requires_grad_(True) for t in (X, Y, W)) + + reference = [ + t.detach().clone() for t in _fwd_bwd(conv, X, Y, W, rows, cols, sender_perm, G) + ] + + compiled = torch.compile(conv, mode="reduce-overhead") + + def step(): + return _fwd_bwd(compiled, X, Y, W, rows, cols, sender_perm, G) + + for _ in range(4): + outputs = [t.detach().clone() for t in step()] + torch.cuda.synchronize() + _assert_close(outputs, reference, deterministic) + + +def test_concurrent_streams_share_no_state(conv_and_inputs): + conv, deterministic, (X, Y, W, rows, cols, sender_perm, G) = conv_and_inputs + X2, Y2, W2, G2 = (t.flip(0).contiguous() for t in (X, Y, W, G)) + + def run(X, Y, W, G): + X, Y, W = (t.clone().requires_grad_(True) for t in (X, Y, W)) + return [ + t.detach().clone() + for t in _fwd_bwd(conv, X, Y, W, rows, cols, sender_perm, G) + ] + + ref1 = run(X, Y, W, G) + ref2 = run(X2, Y2, W2, G2) + torch.cuda.synchronize() + + s1, s2 = torch.cuda.Stream(), torch.cuda.Stream() + for _ in range(5): + with torch.cuda.stream(s1): + out1 = run(X, Y, W, G) + with torch.cuda.stream(s2): + out2 = run(X2, Y2, W2, G2) + torch.cuda.synchronize() + _assert_close(out1, ref1, deterministic) + _assert_close(out2, ref2, deterministic) diff --git a/tests/multidevice_test.py b/tests/multidevice_test.py index 7b7b48c7..4e3a24a8 100644 --- a/tests/multidevice_test.py +++ b/tests/multidevice_test.py @@ -55,3 +55,21 @@ def test_multidevice(): with torch.cuda.device(device): result = tp.forward(X, Y, W) + + node_count, nnz = 500, 5000 + rows = torch.randint(0, node_count, (nnz,), device=device, generator=gen) + cols = torch.randint(0, node_count, (nnz,), device=device, generator=gen) + order = torch.argsort(rows * node_count + cols) + rows, cols = rows[order], cols[order] + sender_perm = torch.argsort(cols * node_count + rows) + + conv = oeq.TensorProductConv(problem, deterministic=True) + Xc = torch.rand(node_count, X_ir.dim, device=device, generator=gen) + Yc = torch.rand(nnz, Y_ir.dim, device=device, generator=gen) + Wc = torch.rand(nnz, problem.weight_numel, device=device, generator=gen) + + with torch.cuda.device(device): + conv_result = conv.forward(Xc, Yc, Wc, rows, cols, sender_perm) + torch.cuda.synchronize(device) + assert conv_result.device == Xc.device + assert torch.isfinite(conv_result).all() diff --git a/tests/stream_test.py b/tests/stream_test.py index 42ac4dd2..2d99c976 100644 --- a/tests/stream_test.py +++ b/tests/stream_test.py @@ -242,7 +242,7 @@ def double_backward_fn(X, Y, W, receivers, senders): def oeq_conv_det_fwd(tpp, conv_buffers): tp_conv = TensorProductConv(tpp, torch_op=True, deterministic=False) - return Executable(tp_conv, conv_buffers, [KE("forward", 1), KE("fixup_forward", 1)]) + return Executable(tp_conv, conv_buffers, [KE("forward", 1)]) @pytest.fixture @@ -265,9 +265,7 @@ def backward_fn(X, Y, W, receivers, senders): conv_buffers, [ KE("forward", 1), - KE("fixup_forward", 1), KE("backward", 1), - KE("fixup_backward", 1), ], ) @@ -311,12 +309,9 @@ def double_backward_fn(X, Y, W, receivers, senders): conv_buffers, [ KE("forward", 1), - KE("fixup_forward", 2), KE("backward", 1), - KE("fixup_backward", 1), KE("double_backward_A", 1), KE("double_backward_B", 1), - KE("fixup_double_backwardB", 1), ], ) From b0aee00c49dee5b115be902267fe09c2345c333c Mon Sep 17 00:00:00 2001 From: asglover <140220574+asglover@users.noreply.github.com> Date: Sun, 13 Sep 2026 19:28:50 -0700 Subject: [PATCH 2/5] make default zero so the allocation can be eliminated in atomic mode --- openequivariance/openequivariance/core/LoopUnrollConv.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/openequivariance/openequivariance/core/LoopUnrollConv.py b/openequivariance/openequivariance/core/LoopUnrollConv.py index 17869760..fbe84e06 100644 --- a/openequivariance/openequivariance/core/LoopUnrollConv.py +++ b/openequivariance/openequivariance/core/LoopUnrollConv.py @@ -139,7 +139,7 @@ def generate_double_backward_schedule(warps_per_block): self.backward_workspace_offset = None self.double_backwardB_offset = None - self.workspace_size = 1 + self.workspace_size = 0 if deterministic: destination_index_bytes = 32 # Add extra to account for padding self.workspace_size = max( From 9be39aa30e05c98089be86d6584efc5eee813dd2 Mon Sep 17 00:00:00 2001 From: asglover <140220574+asglover@users.noreply.github.com> Date: Sun, 13 Sep 2026 19:43:18 -0700 Subject: [PATCH 3/5] fmt --- openequivariance/openequivariance/_torch/TensorProductConv.py | 1 - 1 file changed, 1 deletion(-) diff --git a/openequivariance/openequivariance/_torch/TensorProductConv.py b/openequivariance/openequivariance/_torch/TensorProductConv.py index 94e42b4c..dad4808e 100644 --- a/openequivariance/openequivariance/_torch/TensorProductConv.py +++ b/openequivariance/openequivariance/_torch/TensorProductConv.py @@ -85,7 +85,6 @@ def _init_class(self): kahan=self.input_args["kahan"], ) - self.dummy_transpose_perm = torch.zeros(1, dtype=torch.int64, device="cuda") self.weight_numel = self.config.weight_numel self.kernel = string_to_tensor(self.kernel_string) From 3afa8a81f421dbd9ee2e6abd93720343fc886729 Mon Sep 17 00:00:00 2001 From: asglover <140220574+asglover@users.noreply.github.com> Date: Tue, 15 Sep 2026 23:23:55 -0700 Subject: [PATCH 4/5] change to empty for performance, bug fix comes separately --- .../openequivariance/extension/libtorch_tp_jit.cpp | 4 ++-- .../openequivariance/extension/libtorch_tp_jit_stable.cpp | 4 ++-- .../openequivariance/extension/torch_core.hpp | 8 ++++---- 3 files changed, 8 insertions(+), 8 deletions(-) diff --git a/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp b/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp index 997733bb..81e80fc8 100644 --- a/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp +++ b/openequivariance/openequivariance/extension/libtorch_tp_jit.cpp @@ -47,8 +47,8 @@ Tensor tensor_zeros_like(const Tensor &ref, const std::vector &sizes) { return torch::zeros(sizes, ref.options()); } -Tensor tensor_zeros_bytes(const Tensor &ref, int64_t nbytes) { - return torch::zeros({nbytes}, ref.options().dtype(torch::kByte)); +Tensor tensor_empty_bytes(const Tensor &ref, int64_t nbytes) { + return torch::empty({nbytes}, ref.options().dtype(torch::kByte)); } void tensor_zero_(Tensor &tensor) { diff --git a/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp b/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp index f0fb85a9..8e4a8db6 100644 --- a/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp +++ b/openequivariance/openequivariance/extension/libtorch_tp_jit_stable.cpp @@ -50,10 +50,10 @@ Tensor tensor_zeros_like(const Tensor &ref, const std::vector &sizes) { return out; } -Tensor tensor_zeros_bytes(const Tensor &ref, int64_t nbytes) { +Tensor tensor_empty_bytes(const Tensor &ref, int64_t nbytes) { std::vector sizes = {nbytes}; auto sizes_ref = torch::headeronly::IntHeaderOnlyArrayRef(sizes.data(), sizes.size()); - return torch::stable::new_zeros(ref, sizes_ref, kByte); + return torch::stable::new_empty(ref, sizes_ref, kByte); } void tensor_zero_(Tensor &tensor) { diff --git a/openequivariance/openequivariance/extension/torch_core.hpp b/openequivariance/openequivariance/extension/torch_core.hpp index fc255be8..ce7c78ad 100644 --- a/openequivariance/openequivariance/extension/torch_core.hpp +++ b/openequivariance/openequivariance/extension/torch_core.hpp @@ -40,7 +40,7 @@ Tensor tensor_to_cpu_contiguous(const Tensor &tensor); Tensor tensor_contiguous(const Tensor &tensor); Tensor tensor_empty_like(const Tensor &ref, const std::vector &sizes); Tensor tensor_zeros_like(const Tensor &ref, const std::vector &sizes); -Tensor tensor_zeros_bytes(const Tensor &ref, int64_t nbytes); +Tensor tensor_empty_bytes(const Tensor &ref, int64_t nbytes); void tensor_zero_(Tensor &tensor); void alert_not_deterministic(const char *name); @@ -458,7 +458,7 @@ inline Tensor jit_conv_forward( std::optional workspace; void *workspace_ptr = nullptr; if (k.deterministic) { - workspace.emplace(tensor_zeros_bytes(L1_in, k.workspace_size)); + workspace.emplace(tensor_empty_bytes(L1_in, k.workspace_size)); workspace_ptr = data_ptr(*workspace); } @@ -528,7 +528,7 @@ inline tuple jit_conv_backward( std::optional workspace; void *workspace_ptr = nullptr; if (k.deterministic) { - workspace.emplace(tensor_zeros_bytes(L1_in, k.workspace_size)); + workspace.emplace(tensor_empty_bytes(L1_in, k.workspace_size)); workspace_ptr = data_ptr(*workspace); } @@ -610,7 +610,7 @@ inline tuple jit_conv_double_backward( std::optional workspace; void *workspace_ptr = nullptr; if (k.deterministic) { - workspace.emplace(tensor_zeros_bytes(L1_in, k.workspace_size)); + workspace.emplace(tensor_empty_bytes(L1_in, k.workspace_size)); workspace_ptr = data_ptr(*workspace); } From 9d0c2c21b5b94c97319ffd3133fdcf8109e8e973 Mon Sep 17 00:00:00 2001 From: asglover <140220574+asglover@users.noreply.github.com> Date: Wed, 16 Sep 2026 00:07:40 -0700 Subject: [PATCH 5/5] remove zeroing for jax workspace as well --- openequivariance/openequivariance/jax/jvp/conv_prim.py | 8 ++++---- openequivariance/openequivariance/jax/utils.py | 4 ++-- openequivariance/openequivariance/jax/vjp/conv_func.py | 8 ++++---- 3 files changed, 10 insertions(+), 10 deletions(-) diff --git a/openequivariance/openequivariance/jax/jvp/conv_prim.py b/openequivariance/openequivariance/jax/jvp/conv_prim.py index 994a475e..9d609d64 100644 --- a/openequivariance/openequivariance/jax/jvp/conv_prim.py +++ b/openequivariance/openequivariance/jax/jvp/conv_prim.py @@ -5,7 +5,7 @@ from openequivariance.jax.utils import ( clean_tensors, conv_workspace_shape, - conv_workspace_zeros, + conv_workspace_empty, ) # ============================================================================== @@ -29,7 +29,7 @@ def conv_fwd_impl(X, Y, W, rows, cols, sender_perm, *, L3_dim, kernel, hash): W, rows, cols, - conv_workspace_zeros(kernel), + conv_workspace_empty(kernel), sender_perm, kernel=kernel, hash=hash, @@ -75,7 +75,7 @@ def conv_bwd_impl(X, Y, W, dZ, rows, cols, sender_perm, *, kernel, hash): dZ, rows, cols, - conv_workspace_zeros(kernel), + conv_workspace_empty(kernel), sender_perm, kernel=kernel, hash=hash, @@ -134,7 +134,7 @@ def conv_dbwd_impl( ddW, rows, cols, - conv_workspace_zeros(kernel), + conv_workspace_empty(kernel), sender_perm, kernel=kernel, hash=hash, diff --git a/openequivariance/openequivariance/jax/utils.py b/openequivariance/openequivariance/jax/utils.py index c59a68ed..fc8fc74e 100644 --- a/openequivariance/openequivariance/jax/utils.py +++ b/openequivariance/openequivariance/jax/utils.py @@ -13,9 +13,9 @@ def conv_workspace_shape(kernel: str) -> jax.ShapeDtypeStruct: return jax.ShapeDtypeStruct((int(size),), jnp.uint8) -def conv_workspace_zeros(kernel: str) -> jax.Array: +def conv_workspace_empty(kernel: str) -> jax.Array: shape = conv_workspace_shape(kernel) - return jnp.zeros(shape.shape, shape.dtype) + return jnp.empty(shape.shape, shape.dtype) def reorder_jax_helper(schedule, weights_in, direction, has_batch_dim): diff --git a/openequivariance/openequivariance/jax/vjp/conv_func.py b/openequivariance/openequivariance/jax/vjp/conv_func.py index 4f35e084..5e2e5eb6 100644 --- a/openequivariance/openequivariance/jax/vjp/conv_func.py +++ b/openequivariance/openequivariance/jax/vjp/conv_func.py @@ -2,7 +2,7 @@ import jax.numpy as jnp from functools import partial -from openequivariance.jax.utils import conv_workspace_shape, conv_workspace_zeros +from openequivariance.jax.utils import conv_workspace_shape, conv_workspace_empty def zeros_like(x): @@ -25,7 +25,7 @@ def forward(X, Y, W, rows, cols, sender_perm, L3_dim, kernel, hash): W, rows, cols, - conv_workspace_zeros(kernel), + conv_workspace_empty(kernel), sender_perm, kernel=kernel, hash=hash, @@ -68,7 +68,7 @@ def backward(X, Y, W, dZ, rows, cols, sender_perm, kernel, hash): dZ, rows, cols, - conv_workspace_zeros(kernel), + conv_workspace_empty(kernel), sender_perm, kernel=kernel, hash=hash, @@ -129,7 +129,7 @@ def double_backward(X, Y, W, dZ, ddX, ddY, ddW, rows, cols, sender_perm, kernel, ddW, rows, cols, - conv_workspace_zeros(kernel), + conv_workspace_empty(kernel), sender_perm, kernel=kernel, hash=hash,