Skip to content
Draft
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
5 changes: 5 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
@@ -1,5 +1,10 @@
## Latest Changes

### Unreleased
**Changed**:
- The deterministic convolution workspace is now allocated by the framework (PyTorch / JAX) each call

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Mm - it was always allocated by the framework (albeit at the Python level) - more accurate to state that it's allocated inside the extension at the C++ level.

- Atomic convolutions no longer compile or launch fixup kernels.

### v0.7.0 (2026-09-10)
**Added**:
- Public XLA FFI registration provider
Expand Down
27 changes: 2 additions & 25 deletions openequivariance/openequivariance/_torch/TensorProductConv.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,8 +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
self.kernel = string_to_tensor(self.kernel_string)
Expand Down Expand Up @@ -180,18 +178,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
Expand Down Expand Up @@ -282,9 +271,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")
Expand All @@ -297,7 +284,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)
Expand All @@ -315,7 +301,6 @@ def fake_double_backward(
w_dgrad,
rows,
cols,
workspace_buffer,
transpose_perm=None,
):
return [
Expand Down Expand Up @@ -345,7 +330,6 @@ def setup_context(ctx, inputs, output):
ctx.L3_dim,
ctx.rows,
ctx.cols,
ctx.workspace_buffer,
ctx.sender_perm,
) = inputs

Expand All @@ -359,10 +343,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
Expand All @@ -378,7 +361,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
Expand All @@ -396,7 +378,6 @@ def double_backward(ctx, E, F, G):
G,
ctx.rows,
ctx.cols,
ctx.workspace_buffer,
ctx.sender_perm,
)
return (
Expand All @@ -409,7 +390,6 @@ def double_backward(ctx, E, F, G):
None,
None,
None,
None,
)

torch.library.register_autograd(
Expand All @@ -431,7 +411,6 @@ def setup_context_triple_backward(ctx, inputs, output):
ctx.W_dgrad,
ctx.rows,
ctx.cols,
ctx.workspace_buffer,
ctx.sender_perm,
) = inputs

Expand All @@ -444,7 +423,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,
)

Expand Down Expand Up @@ -546,7 +524,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(
Expand Down
1 change: 0 additions & 1 deletion openequivariance/openequivariance/core/ConvolutionBase.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
2 changes: 1 addition & 1 deletion openequivariance/openequivariance/core/LoopUnrollConv.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
59 changes: 38 additions & 21 deletions openequivariance/openequivariance/extension/convolution.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<string> kernels = {"forward", "backward", "fixup_forward", "fixup_backward", "double_backward_A", "double_backward_B", "fixup_double_backwardB"};
jit.compile(kernels, {{}, {}, {}, {}, {}, {}, {}}, opt_level);
vector<string> 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<vector<int>>(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);

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Cool, thanks for making this an enum

}

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);
}
}

Expand All @@ -73,7 +89,8 @@ class __attribute__ ((visibility ("default"))) JITConvImpl {
dbl_bwd_dict["num_threads"],
dbl_bwd_dict["smem"]
),
static_cast<int>(kernel_dims["opt_level"])) { }
static_cast<int>(kernel_dims["opt_level"]),
kernel_dims["deterministic"] != 0) { }

void exec_conv(
void* L1_in,
Expand All @@ -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<uint64_t>(workspace) != 0) {
if(deterministic) {
void *fixup_args[] = {&workspace, &L3_out};

KernelLaunchConfig fixup_config(
Expand All @@ -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);
}
}

Expand All @@ -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<uint64_t>(workspace) != 0) {
if(deterministic) {
void *fixup_args[] = {&workspace, &L1_grad};

KernelLaunchConfig fixup_config(
Expand All @@ -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);
}
}

Expand All @@ -150,28 +167,28 @@ 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<uint64_t>(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,
forward_config_ref.num_threads,
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<uint64_t>(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,
double_backward_config_ref.num_threads,
0
);
fixup_config.hStream = stream;
jit.execute(6, fixup_args, fixup_config);
jit.execute(FIXUP_DOUBLE_BACKWARD_B, fixup_args, fixup_config);
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,10 @@ Tensor tensor_zeros_like(const Tensor &ref, const std::vector<int64_t> &sizes) {
return torch::zeros(sizes, ref.options());
}

Tensor tensor_empty_bytes(const Tensor &ref, int64_t nbytes) {

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Seems like this function is unused?

return torch::empty({nbytes}, ref.options().dtype(torch::kByte));
}

void tensor_zero_(Tensor &tensor) {
tensor.zero_();
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,12 @@ Tensor tensor_zeros_like(const Tensor &ref, const std::vector<int64_t> &sizes) {
return out;
}

Tensor tensor_empty_bytes(const Tensor &ref, int64_t nbytes) {
std::vector<int64_t> sizes = {nbytes};
auto sizes_ref = torch::headeronly::IntHeaderOnlyArrayRef(sizes.data(), sizes.size());
return torch::stable::new_empty(ref, sizes_ref, kByte);
}

void tensor_zero_(Tensor &tensor) {
torch::stable::zero_(tensor);
}
Expand Down
Loading
Loading