From f07f1e9478f6ad8b573992a1d0e58ca35f6f1d99 Mon Sep 17 00:00:00 2001 From: Paul Fuchs Date: Sun, 13 Sep 2026 13:55:11 +0200 Subject: [PATCH 1/2] Add receiver-streaming UVU convolution kernel --- docs/api.rst | 18 +- .../core/FactorizedComputationSchedule.py | 452 +++++++++ .../extension/factorized_projected.hpp | 89 ++ .../jax/LoopUnrollTensorProductConv.py | 234 +++++ .../jax/StreamingTensorProductConv.py | 283 ++++++ .../openequivariance/jax/TensorProductConv.py | 396 +++++--- .../openequivariance/jax/__init__.py | 20 +- .../openequivariance/jax/ffi_targets.py | 5 +- .../jax/jvp/factorized_projected_prim.py | 905 ++++++++++++++++++ .../templates/factorized_projected.cuh | 474 +++++++++ .../openequivariance/templates/jinja_utils.py | 28 +- openequivariance_extjax/CMakeLists.txt | 1 + openequivariance_extjax/src/ffi_handlers.cpp | 323 +++++++ tests/factorized_computation_schedule_test.py | 120 +++ tests/jax_ffi_abi_test.py | 4 +- .../jax_tensor_product_conv_dispatch_test.py | 524 ++++++++++ tests/vmap_test.py | 100 ++ 17 files changed, 3812 insertions(+), 164 deletions(-) create mode 100644 openequivariance/openequivariance/core/FactorizedComputationSchedule.py create mode 100644 openequivariance/openequivariance/extension/factorized_projected.hpp create mode 100644 openequivariance/openequivariance/jax/LoopUnrollTensorProductConv.py create mode 100644 openequivariance/openequivariance/jax/StreamingTensorProductConv.py create mode 100644 openequivariance/openequivariance/jax/jvp/factorized_projected_prim.py create mode 100644 openequivariance/openequivariance/templates/factorized_projected.cuh create mode 100644 tests/factorized_computation_schedule_test.py create mode 100644 tests/jax_tensor_product_conv_dispatch_test.py diff --git a/docs/api.rst b/docs/api.rst index 15b1aec9..162bce72 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -42,6 +42,12 @@ The JAX API consists of ``TensorProduct`` and ``TensorProductConv`` classes that behave identically to their PyTorch counterparts. These classes do not conform exactly to the e3nn-jax API, but perform the same computation. +JAX ``TensorProductConv`` uses the established loop-unroll implementation by +default. Select ``mode="streaming"`` to require receiver streaming, or +``mode="auto"`` to use it when the problem is supported. Streaming accepts +padded tail edges and optional receiver offsets, but vmapped calls must share +one static topology. + If you plan to use ``oeq.jax`` without PyTorch installed, you need to set ``OEQ_NOTORCH=1`` in your local environment (within Python, ``os.environ["OEQ_NOTORCH"] = 1``). For the moment, we require this to avoid @@ -54,6 +60,16 @@ breaking the PyTorch version of OpenEquivariance. :exclude-members: .. autoclass:: openequivariance.jax.TensorProductConv + :members: forward, reorder_weights_from_e3nn, reorder_weights_to_e3nn, implementation, uses_streaming_kernel + :undoc-members: + :exclude-members: + +.. autoclass:: openequivariance.jax.LoopUnrollTensorProductConv + :members: forward, reorder_weights_from_e3nn, reorder_weights_to_e3nn + :undoc-members: + :exclude-members: + +.. autoclass:: openequivariance.jax.StreamingTensorProductConv :members: forward, reorder_weights_from_e3nn, reorder_weights_to_e3nn :undoc-members: :exclude-members: @@ -76,4 +92,4 @@ both packages. .. autoclass:: openequivariance.Irreps :members: - :undoc-members: \ No newline at end of file + :undoc-members: diff --git a/openequivariance/openequivariance/core/FactorizedComputationSchedule.py b/openequivariance/openequivariance/core/FactorizedComputationSchedule.py new file mode 100644 index 00000000..06834cd7 --- /dev/null +++ b/openequivariance/openequivariance/core/FactorizedComputationSchedule.py @@ -0,0 +1,452 @@ +"""Build schedules for receiver-streaming convolutions.""" + +from dataclasses import dataclass +from enum import IntEnum +from functools import cache + +import numpy as np + +from openequivariance.core.e3nn_lite import TPProblem, wigner_3j +from openequivariance.core.utils import calc_weight_offsets, hash_str_64 +from openequivariance.templates.jinja_utils import ( + cpp_scalar_type, + get_jinja_environment, +) + + +class Input(IntEnum): + """Floating operands whose derivative buffers may be inactive.""" + + X = 0 + SH = 1 + W = 2 + + +@dataclass(frozen=True, slots=True) +class FactorizedLaunchConfig: + """Runtime-grid launch geometry shared by generated streaming kernels.""" + + num_threads: int = 128 + logical_cohort_width: int = 32 + shared_memory_bytes: int = 0 + + def __post_init__(self): + """Validate launch geometry required by the generated kernels.""" + if self.num_threads <= 0: + raise ValueError("the launch must contain at least one thread") + if self.logical_cohort_width != 32: + raise ValueError("the generated reductions require 32-thread cohorts") + if self.num_threads % self.logical_cohort_width: + raise ValueError("the thread count must contain complete logical cohorts") + if self.shared_memory_bytes < 0: + raise ValueError("the shared-memory allocation cannot be negative") + + +@dataclass(frozen=True, slots=True) +class FactorizedAccumulatorSlot: + """One register accumulator and its packed-array address.""" + + name: str + array_index: int + irrep_dim: int + + +@dataclass(frozen=True, slots=True) +class FactorizedCoupling: + """One nonzero component coupling.""" + + input_component: int + edge_component: int + output_component: int + coefficient: float + + +@dataclass(frozen=True, slots=True) +class FactorizedCouplingPath: + """One weighted sparse coupling path.""" + + input_start: int + edge_start: int + output_start: int + weight_start: int + input_irrep_dim: int + edge_irrep_dim: int + output_irrep_dim: int + edge_mul: int + couplings: tuple[FactorizedCoupling, ...] + + +@dataclass(frozen=True, slots=True) +class FactorizedScheduledCoupling: + """One component coupling with its input-gradient accumulator.""" + + input_component: int + edge_component: int + coefficient: float + input_accumulator: FactorizedAccumulatorSlot + + +@dataclass(frozen=True, slots=True) +class FactorizedPathOutput: + """Sparse terms and accumulator for one output component.""" + + output_component: int + accumulator: FactorizedAccumulatorSlot + terms: tuple[FactorizedScheduledCoupling, ...] + + +@dataclass(frozen=True, slots=True) +class FactorizedComputationPath: + """Static offsets and sparse couplings for one weighted instruction.""" + + input_start: int + edge_start: int + weight_start: int + input_irrep_dim: int + edge_irrep_dim: int + edge_mul: int + outputs: tuple[FactorizedPathOutput, ...] + + +@dataclass(frozen=True, slots=True) +class FactorizedKernel: + """Rendered source and FFI attributes for one kernel specialization.""" + + jit_kernel: str + hash: int + ffi_attributes: dict[str, object] + + +@dataclass(frozen=True, slots=True, eq=False) +class FactorizedComputationSchedule: + """Static sparse schedule for a receiver-owned generated convolution. + + Paths retain the public flattened input, edge, output, and per-edge weight + layouts of the source :class:`~openequivariance.core.e3nn_lite.TPProblem`. + Output and sender-gradient slots are fully resolved before rendering, so + the CUDA or HIP template only executes this schedule. Graph size and edge + count remain runtime values because the forward pass follows CSR rows. + """ + + paths: tuple[FactorizedComputationPath, ...] + output_slots: tuple[FactorizedAccumulatorSlot, ...] + input_gradient_slots: tuple[FactorizedAccumulatorSlot, ...] + input_dim: int + edge_dim: int + output_dim: int + weight_numel: int + channels: int + launch_config: FactorizedLaunchConfig = FactorizedLaunchConfig() + layout: str = "mul_ir" + + @cache + def kernel( + self, + dtype: object, + *, + is_hip: bool, + forward_jvp_active: tuple[bool, bool, bool] = (True, True, True), + backward_jvp_active: tuple[bool, bool, bool, bool] = (True, True, True, True), + ) -> FactorizedKernel: + """Render and cache one CUDA or HIP derivative specialization.""" + if len(forward_jvp_active) != len(Input): + raise ValueError("forward JVP activity must contain X, SH, and W") + if len(backward_jvp_active) != len(Input) + 1: + raise ValueError("backward JVP activity must also contain dout") + source = ( + get_jinja_environment(is_hip=is_hip) + .get_template("factorized_projected.cuh") + .render( + scalar=cpp_scalar_type(dtype), + schedule=self, + Input=Input, + forward_jvp_active=forward_jvp_active, + backward_jvp_active=backward_jvp_active, + backward_jvp_dout_active=backward_jvp_active[len(Input)], + ) + ) + launch = self.launch_config + cache_source = ( + f"{source}\0{launch.num_threads}\0{launch.logical_cohort_width}" + f"\0{launch.shared_memory_bytes}" + ) + kernel_hash = hash_str_64(cache_source) + attributes = { + "source": source, + "hash": kernel_hash, + "input_dim": self.input_dim, + "edge_dim": self.edge_dim, + "weight_dim": self.weight_numel, + "output_dim": self.output_dim, + "channels": self.channels, + "num_threads": launch.num_threads, + "logical_cohort_width": launch.logical_cohort_width, + "shared_memory_bytes": launch.shared_memory_bytes, + } + return FactorizedKernel(source, kernel_hash, attributes) + + def weight_reordering_info(self, weights_in, has_batch_dim: bool): + """Describe canonical-to-native weight transposes for OEQ utilities.""" + batch_dim = weights_in.shape[0] + specs = [] + for path in self.paths: + start = path.weight_start + stop = start + self.channels * path.edge_mul + parent_shape = [self.channels, path.edge_mul] + child_shape = list(parent_shape) + parent_range = [slice(start, stop)] + child_range = [slice(start, stop)] + weights_subrange = [slice(None), slice(None)] + transpose_perm = [1, 0] + reshape_size = [-1] + if has_batch_dim: + parent_shape = [batch_dim] + parent_shape + child_shape = [batch_dim] + child_shape + parent_range.insert(0, slice(0, batch_dim)) + child_range.insert(0, slice(0, batch_dim)) + weights_subrange.insert(0, slice(0, batch_dim)) + transpose_perm = [0, 2, 1] + reshape_size = [batch_dim, -1] + specs.append( + { + "parent_range": tuple(parent_range), + "parent_shape": parent_shape, + "weights_subrange": tuple(weights_subrange), + "child_range": tuple(child_range), + "child_shape": child_shape, + "transpose_perm": transpose_perm, + "reshape_size": reshape_size, + "transpose_child_shape": [ + child_shape[index] for index in transpose_perm + ], + } + ) + return specs + + +def factorized_schedule_from_problem( + problem: TPProblem, +) -> FactorizedComputationSchedule: + """Lower a supported tensor-product problem to a sparse streaming schedule. + + Each instruction in a homogeneous weighted ``uvu`` problem becomes one + execution path. + Channel-wise UVW paths are treated as UVU. + + :param problem: External-weight, unshared ``TPProblem`` in ``"mul_ir"`` + layout with homogeneous weighted ``"uvu"`` instructions, or an exactly + reducible scalar-multiplicity ``"uvw"`` instruction set. Instructions + using the two connection modes cannot be mixed in one problem. + :return: A static sparse execution schedule preserving external layouts. + :raises ValueError: If the layout, instructions, channel structure, or + flat weight coverage is unsupported. + """ + if problem.layout != "mul_ir": + raise ValueError("generated factorized convolution requires mul_ir layout") + if problem.shared_weights or problem.internal_weights: + raise ValueError( + "generated factorized convolution requires external unshared weights" + ) + + declared_modes = { + instruction.connection_mode for instruction in problem.instructions + } + if not declared_modes.issubset({"uvu", "uvw"}): + raise ValueError( + "generated convolution supports homogeneous weighted uvu or uvw paths" + ) + if len(declared_modes) != 1: + raise ValueError( + "generated convolution requires all paths to use one connection mode" + ) + declared_connection_mode = next(iter(declared_modes), None) + reduce_scalar_uvw = declared_connection_mode == "uvw" and all( + problem.irreps_in1[instruction.i_in1].mul == 1 + and problem.irreps_out[instruction.i_out].mul == 1 + for instruction in problem.instructions + ) + if declared_connection_mode == "uvw" and not reduce_scalar_uvw: + raise ValueError( + "receiver-streaming generated convolution supports UVW only when " + "every path has one input and one output multiplicity" + ) + + # Extract the sparse component couplings in instruction order. + input_slices = problem.irreps_in1.slices() + edge_slices = problem.irreps_in2.slices() + output_slices = problem.irreps_out.slices() + weight_offsets = calc_weight_offsets(problem) + coupling_paths = [] + referenced_outputs = set() + channel_count = None + + for path_index, instruction in enumerate(problem.instructions): + if not instruction.has_weight: + raise ValueError("generated factorized convolution requires weighted paths") + + input_mul_ir = problem.irreps_in1[instruction.i_in1] + edge_mul_ir = problem.irreps_in2[instruction.i_in2] + output_mul_ir = problem.irreps_out[instruction.i_out] + referenced_outputs.add(instruction.i_out) + if input_mul_ir.mul != output_mul_ir.mul: + raise ValueError("uvu input and output channel multiplicities must match") + if channel_count is None: + channel_count = input_mul_ir.mul + elif input_mul_ir.mul != channel_count: + raise ValueError("generated uvu convolution requires uniform channels") + + cg = np.asarray( + wigner_3j(input_mul_ir.ir.l, edge_mul_ir.ir.l, output_mul_ir.ir.l), + dtype=np.float64, + ) * float(instruction.path_weight) + nonzero = np.nonzero(cg) + if len(nonzero[0]) == 0: + raise ValueError(f"instruction {path_index} has an empty coupling tensor") + + couplings = [] + for input_component, edge_component, output_component in zip( + *nonzero, strict=True + ): + couplings.append( + FactorizedCoupling( + input_component=int(input_component), + edge_component=int(edge_component), + output_component=int(output_component), + coefficient=float( + cg[input_component, edge_component, output_component] + ), + ) + ) + + coupling_paths.append( + FactorizedCouplingPath( + input_start=input_slices[instruction.i_in1].start, + edge_start=edge_slices[instruction.i_in2].start, + output_start=output_slices[instruction.i_out].start, + weight_start=weight_offsets[path_index], + input_irrep_dim=input_mul_ir.ir.dim, + edge_irrep_dim=edge_mul_ir.ir.dim, + output_irrep_dim=output_mul_ir.ir.dim, + edge_mul=edge_mul_ir.mul, + couplings=tuple(couplings), + ) + ) + + if channel_count is None: + raise ValueError("generated factorized convolution requires at least one path") + missing_outputs = [ + output + for output, mul_ir in enumerate(problem.irreps_out) + if mul_ir.mul * mul_ir.ir.dim > 0 and output not in referenced_outputs + ] + if missing_outputs: + raise ValueError( + "factorized schedule does not reference positive-dimensional output " + f"irrep blocks {missing_outputs}" + ) + covered_weight_count = sum( + np.prod(instruction.path_shape) + for instruction in problem.instructions + if instruction.has_weight + ) + if covered_weight_count != problem.weight_numel: + raise ValueError( + "factorized schedule does not cover the complete weight vector" + ) + + # Group couplings by output component and assign their accumulator slots. + output_slots = [] + input_gradient_slots = [] + output_slot_indices = {} + input_gradient_slot_indices = {} + + def intern_slot(slots, slot_indices, prefix, start, component, irrep_dim): + key = (start, component, irrep_dim) + existing = slot_indices.get(key) + if existing is not None: + return slots[existing] + slot_indices[key] = len(slots) + slot = FactorizedAccumulatorSlot( + name=f"{prefix}_{start}_{component}", + array_index=start + component, + irrep_dim=irrep_dim, + ) + slots.append(slot) + return slot + + paths = [] + for coupling_path in coupling_paths: + outputs = [] + for output_component in range(coupling_path.output_irrep_dim): + scheduled_terms = [] + for coupling in coupling_path.couplings: + if coupling.output_component != output_component: + continue + scheduled_terms.append( + FactorizedScheduledCoupling( + input_component=coupling.input_component, + edge_component=coupling.edge_component, + coefficient=coupling.coefficient, + input_accumulator=intern_slot( + input_gradient_slots, + input_gradient_slot_indices, + "input_gradient", + coupling_path.input_start, + coupling.input_component, + coupling_path.input_irrep_dim, + ), + ) + ) + + outputs.append( + FactorizedPathOutput( + output_component=output_component, + accumulator=intern_slot( + output_slots, + output_slot_indices, + "output", + coupling_path.output_start, + output_component, + coupling_path.output_irrep_dim, + ), + terms=tuple(scheduled_terms), + ) + ) + + paths.append( + FactorizedComputationPath( + input_start=coupling_path.input_start, + edge_start=coupling_path.edge_start, + weight_start=coupling_path.weight_start, + input_irrep_dim=coupling_path.input_irrep_dim, + edge_irrep_dim=coupling_path.edge_irrep_dim, + edge_mul=coupling_path.edge_mul, + outputs=tuple(outputs), + ) + ) + + return FactorizedComputationSchedule( + paths=tuple(paths), + output_slots=tuple(output_slots), + input_gradient_slots=tuple(input_gradient_slots), + input_dim=problem.irreps_in1.dim, + edge_dim=problem.irreps_in2.dim, + output_dim=problem.irreps_out.dim, + weight_numel=problem.weight_numel, + channels=channel_count, + ) + + +__all__ = [ + "FactorizedAccumulatorSlot", + "FactorizedCoupling", + "FactorizedCouplingPath", + "FactorizedComputationPath", + "FactorizedComputationSchedule", + "FactorizedKernel", + "FactorizedLaunchConfig", + "FactorizedPathOutput", + "FactorizedScheduledCoupling", + "Input", + "factorized_schedule_from_problem", +] diff --git a/openequivariance/openequivariance/extension/factorized_projected.hpp b/openequivariance/openequivariance/extension/factorized_projected.hpp new file mode 100644 index 00000000..04ee551a --- /dev/null +++ b/openequivariance/openequivariance/extension/factorized_projected.hpp @@ -0,0 +1,89 @@ +#pragma once + +#include +#include +#include + +template +class __attribute__ ((visibility ("default"))) JITFactorizedProjectedImpl { +public: + JIT_IMPL jit; + + static std::vector kernel_entry_points() { + return { + "oeq_projected_forward", "oeq_projected_forward_jvp", + "oeq_projected_backward", + "oeq_projected_backward_jvp"}; + } + + static std::vector> kernel_template_parameters() { + return std::vector>(kernel_entry_points().size()); + } + + JITFactorizedProjectedImpl( + std::string source, int64_t num_threads, + int64_t logical_cohort_width, int64_t shared_memory_bytes) : + jit(std::move(source)), + num_threads_(num_threads), + logical_cohort_width_(logical_cohort_width), + shared_memory_bytes_(shared_memory_bytes) { + jit.compile(kernel_entry_points(), kernel_template_parameters()); + } + + void forward( + int64_t node_count, int64_t edge_count, int64_t channels, + void* x, void* sh, void* weights, void* senders, void* row_ptr, + void* out, Stream stream) { + void* args[] = { + &node_count, &edge_count, &x, &sh, &weights, &senders, &row_ptr, &out}; + execute(0, node_count * channels, args, stream); + } + + void forward_jvp( + int64_t node_count, int64_t edge_count, int64_t channels, + void* x, void* sh, void* weights, void* senders, void* row_ptr, + void* tx, void* tsh, void* tweights, void* out, Stream stream) { + void* args[] = { + &node_count, &edge_count, &x, &sh, &weights, &senders, &row_ptr, + &tx, &tsh, &tweights, &out}; + execute(1, node_count * channels, args, stream); + } + + void backward( + int64_t node_count, int64_t edge_count, + void* x, void* sh, void* weights, void* senders, void* receivers, + void* dout, void* dx, void* dsh, void* dweights, Stream stream) { + void* args[] = { + &node_count, &edge_count, &x, &sh, &weights, &senders, &receivers, + &dout, &dx, &dsh, &dweights}; + execute(2, edge_count * logical_cohort_width_, args, stream); + } + + void backward_jvp( + int64_t node_count, int64_t edge_count, + void* x, void* sh, void* weights, void* senders, void* receivers, + void* dout, void* tx, void* tsh, void* tweights, void* tdout, + void* tdx, void* tdsh, void* tdweights, Stream stream) { + void* args[] = { + &node_count, &edge_count, &x, &sh, &weights, &senders, &receivers, &dout, + &tx, &tsh, &tweights, &tdout, &tdx, &tdsh, &tdweights}; + execute(3, edge_count * logical_cohort_width_, args, stream); + } + +private: + int64_t num_threads_; + int64_t logical_cohort_width_; + int64_t shared_memory_bytes_; + + void execute(int kernel_index, int64_t work_items, void* args[], Stream stream) { + if (work_items == 0) + return; + const int64_t blocks = + (work_items + num_threads_ - 1) / num_threads_; + jit.execute( + kernel_index, args, + with_stream( + KernelLaunchConfig(blocks, num_threads_, shared_memory_bytes_), + stream)); + } +}; diff --git a/openequivariance/openequivariance/jax/LoopUnrollTensorProductConv.py b/openequivariance/openequivariance/jax/LoopUnrollTensorProductConv.py new file mode 100644 index 00000000..163825df --- /dev/null +++ b/openequivariance/openequivariance/jax/LoopUnrollTensorProductConv.py @@ -0,0 +1,234 @@ +import jax +import jax.numpy as jnp +import numpy as np +from typing import Optional +from openequivariance.jax import extlib + + +from openequivariance.core.e3nn_lite import TPProblem +from openequivariance.core.LoopUnrollConv import LoopUnrollConv +from openequivariance.jax.utils import reorder_jax + +from openequivariance.core.logging import getLogger +from openequivariance.jax.jvp import conv_prim +from openequivariance.jax.vjp import conv_func + + +logger = getLogger() + + +class LoopUnrollTensorProductConv(LoopUnrollConv): + r"""Apply the established loop-unroll convolution with JAX. + + :param problem: Specification of the tensor product. + :param deterministic: if ``False``, uses atomics for the convolution. If ``True``, uses a deterministic + fixup-based algorithm. `Default`: ``False``. + :param kahan: If ``True``, uses Kahan summation to improve accuracy during aggregation. To use this option, + the input tensors must be in float32 precision AND you must set ``deterministic=True``. *Default*: ``False``. + """ + + def __init__( + self, + config: TPProblem, + deterministic: bool = False, + kahan: bool = False, + requires_jvp: bool = True, + ): + dp = extlib.DeviceProp(0) + self.requires_jvp = requires_jvp + super().__init__( + config, + dp, + extlib.IS_HIP, + idx_dtype=np.int32, + torch_op=False, + deterministic=deterministic, + kahan=kahan, + ) + + self.kernel = self.kernel_string + 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( + self, + X: jax.numpy.ndarray, + Y: jax.numpy.ndarray, + W: jax.numpy.ndarray, + rows: jax.numpy.ndarray, + cols: jax.numpy.ndarray, + sender_perm: Optional[jax.numpy.ndarray] = None, + *, + indices_are_sorted: bool = False, + row_ptr: Optional[jax.numpy.ndarray] = None, + ) -> jax.numpy.ndarray: + """Apply the loop-unroll convolution.""" + if not isinstance(indices_are_sorted, bool): + raise TypeError("indices_are_sorted must be a Python bool") + del row_ptr + # Preserve rows, cols, and sender_perm exactly. Changing endpoints here + # would invalidate a deterministic sender permutation. + if not self.deterministic: + sender_perm = self.dummy_transpose_perm + else: + assert sender_perm is not None, ( + "Must provide sender_perm for deterministic convolutions." + ) + + func = conv_prim.conv_fwd_p.bind + + if not self.requires_jvp: + func = conv_func.forward + + return func( + X, + Y, + W, + rows, + cols, + self.workspace, + sender_perm, + L3_dim=self.L3_dim, + kernel=self.kernel, + hash=self.hash, + ) + + def __call__( + self, + X: jax.numpy.ndarray, + Y: jax.numpy.ndarray, + W: jax.numpy.ndarray, + rows: jax.numpy.ndarray, + cols: jax.numpy.ndarray, + sender_perm: Optional[jax.numpy.ndarray] = None, + *, + indices_are_sorted: bool = False, + row_ptr: Optional[jax.numpy.ndarray] = None, + ) -> jax.numpy.ndarray: + return self.forward( + X, + Y, + W, + rows, + cols, + sender_perm, + indices_are_sorted=indices_are_sorted, + row_ptr=row_ptr, + ) + + def reorder_weights_from_e3nn(self, weights, has_batch_dim=True): + return reorder_jax(self.forward_schedule, weights, "forward", has_batch_dim) + + def reorder_weights_to_e3nn(self, weights, has_batch_dim=True): + return reorder_jax(self.forward_schedule, weights, "backward", has_batch_dim) + + def forward_cpu(self, L1_in, L2_in, weights, L3_out, graph): + rows = graph.rows.astype(np.int32) + cols = graph.cols.astype(np.int32) + sender_perm = graph.transpose_perm.astype(np.int32) + weights = self.reorder_weights_from_e3nn( + weights, has_batch_dim=not self.config.shared_weights + ) + + jit_fwd = jax.jit(self.forward) + result = jit_fwd( + jax.numpy.asarray(L1_in), + jax.numpy.asarray(L2_in), + jax.numpy.asarray(weights), + jax.numpy.asarray(rows), + jax.numpy.asarray(cols), + jax.numpy.asarray(sender_perm), + ) + L3_out[:] = np.asarray(result) + + def backward_cpu( + self, + L1_in, + L1_grad, + L2_in, + L2_grad, + L3_grad, + weights, + weights_grad, + graph, + ): + rows = graph.rows.astype(np.int32) + cols = graph.cols.astype(np.int32) + sender_perm = graph.transpose_perm.astype(np.int32) + weights = self.reorder_weights_from_e3nn( + weights, has_batch_dim=not self.config.shared_weights + ) + + backward_fn = jax.jit( + jax.vjp( + lambda X, Y, W: self.forward( + X, + Y, + W, + jax.numpy.asarray(rows), + jax.numpy.asarray(cols), + jax.numpy.asarray(sender_perm), + ), + jax.numpy.asarray(L1_in), + jax.numpy.asarray(L2_in), + jax.numpy.asarray(weights), + )[1] + ) + + L1_grad_jax, L2_grad_jax, weights_grad_jax = backward_fn( + jax.numpy.asarray(L3_grad) + ) + L1_grad[:] = np.asarray(L1_grad_jax) + L2_grad[:] = np.asarray(L2_grad_jax) + weights_grad[:] = np.asarray(weights_grad_jax) + weights_grad[:] = self.reorder_weights_to_e3nn( + weights_grad, has_batch_dim=not self.config.shared_weights + ) + + def double_backward_cpu( + self, in1, in2, out_grad, weights, weights_dgrad, in1_dgrad, in2_dgrad, graph + ): + in1_jax = jax.numpy.asarray(in1) + in2_jax = jax.numpy.asarray(in2) + weights_jax = jax.numpy.asarray(weights) + out_grad_jax = jax.numpy.asarray(out_grad) + in1_dgrad_jax = jax.numpy.asarray(in1_dgrad) + in2_dgrad_jax = jax.numpy.asarray(in2_dgrad) + weights_dgrad_jax = jax.numpy.asarray(weights_dgrad) + + rows_jax = jax.numpy.asarray(graph.rows.astype(self.idx_dtype)) + cols_jax = jax.numpy.asarray(graph.cols.astype(self.idx_dtype)) + sender_perm_jax = jax.numpy.asarray(graph.transpose_perm.astype(self.idx_dtype)) + + in1_grad, in2_grad, weights_grad, out_dgrad = jax.jit( + jax.vjp( + lambda x, y, w, o: jax.vjp( + lambda a, b, c: self.forward( + a, b, c, rows_jax, cols_jax, sender_perm_jax + ), + x, + y, + w, + )[1](o), + in1_jax, + in2_jax, + weights_jax, + out_grad_jax, + )[1] + )((in1_dgrad_jax, in2_dgrad_jax, weights_dgrad_jax)) + + return ( + np.asarray(in1_grad), + np.asarray(in2_grad), + np.asarray(weights_grad), + np.asarray(out_dgrad), + ) + + +__all__ = ["LoopUnrollTensorProductConv"] diff --git a/openequivariance/openequivariance/jax/StreamingTensorProductConv.py b/openequivariance/openequivariance/jax/StreamingTensorProductConv.py new file mode 100644 index 00000000..735cbb81 --- /dev/null +++ b/openequivariance/openequivariance/jax/StreamingTensorProductConv.py @@ -0,0 +1,283 @@ +"""Receiver-streaming tensor-product convolution implementation.""" + +from typing import NamedTuple, Optional + +import jax +import jax.numpy as jnp +import numpy as np + +from openequivariance.core.e3nn_lite import TPProblem +from openequivariance.core.FactorizedComputationSchedule import ( + factorized_schedule_from_problem, +) +from openequivariance.jax.utils import reorder_jax + + +class StreamingUnavailableError(ValueError): + """Raised when a convolution cannot use receiver-row streaming.""" + + +class StreamingConvTopology(NamedTuple): + """Store receiver-sorted topology that can be reused by streaming calls. + + Invalid padded edges are sorted after valid rows and assigned a sentinel + receiver. Forward traversal excludes the sentinel suffix, while reverse + kernels discard it before reading floating-point operands. + """ + + order: jax.Array + receivers: jax.Array + senders: jax.Array + row_ptr: jax.Array + valid_edges: jax.Array + + +def streaming_support(config: TPProblem) -> tuple[bool, str]: + """Report whether ``config`` can use streaming without constructing a kernel. + + The predicate only inspects immutable problem metadata and lowers the + lightweight static schedule. It has no CUDA, JAX compilation, or other runtime + side effects. + """ + if config.shared_weights: + return False, "receiver streaming requires unshared per-edge weights" + if config.internal_weights: + return False, "receiver streaming requires externally supplied weights" + if config.layout not in ("mul_ir", "ir_mul"): + return False, "receiver streaming supports only mul_ir and ir_mul layouts" + if config.irrep_dtype != config.weight_dtype: + return False, "receiver streaming requires matching irrep and weight dtypes" + if config.irrep_dtype not in (np.float32, np.float64): + return False, "receiver streaming supports only float32 and float64" + native_config = config.clone() + native_config.layout = "mul_ir" + try: + factorized_schedule_from_problem(native_config) + except ValueError as error: + return False, str(error) + return True, "" + + +class StreamingTensorProductConv: + """Apply a full external-weight convolution with receiver-owned accumulation. + + The operator always accepts ``X[N, input_dim]``, ``Y[E, edge_dim]``, and + native-order ``W[E, weight_numel]``. The public ``TensorProductConv`` + selects this implementation only for supported configurations. + + Feature layouts ``"mul_ir"`` and ``"ir_mul"`` are both accepted. Both + implementations execute in ``"mul_ir"`` and use differentiable boundary + transposes when the public problem uses ``"ir_mul"``. + """ + + def __init__(self, config: TPProblem): + """Create a streaming convolution for a previously validated problem.""" + if config.shared_weights: + raise StreamingUnavailableError( + "The streaming schedule requires unshared per-edge weights." + ) + if config.internal_weights: + raise StreamingUnavailableError( + "The streaming schedule requires externally supplied weights." + ) + if config.layout not in ("mul_ir", "ir_mul"): + raise StreamingUnavailableError( + "The streaming schedule supports only mul_ir and ir_mul layouts." + ) + if config.irrep_dtype != config.weight_dtype: + raise StreamingUnavailableError( + "The streaming schedule requires matching irrep and weight dtypes." + ) + if config.irrep_dtype not in (np.float32, np.float64): + raise StreamingUnavailableError( + "The streaming schedule supports only float32 and float64." + ) + + self.config = config + self.weight_numel = config.weight_numel + self.public_layout = config.layout + native_config = config.clone() + native_config.layout = "mul_ir" + self._native_config = native_config + try: + self.schedule = factorized_schedule_from_problem(native_config) + except ValueError as error: + raise StreamingUnavailableError(str(error)) from error + + @property + def uses_streaming_kernel(self) -> bool: + """Return whether this object uses the receiver-streaming kernel.""" + return True + + @property + def implementation(self) -> str: + """Return the implementation name for this static problem.""" + return "streaming" + + def reorder_weights_from_e3nn( + self, weights: jax.Array, has_batch_dim: bool = True + ) -> jax.Array: + """Convert canonical e3nn weights to the selected OEQ weight order.""" + weights = jnp.asarray(weights) + if weights.ndim < 1 or weights.shape[-1] != self.weight_numel: + raise ValueError(f"weights must end in dimension {self.weight_numel}") + return reorder_jax(self.schedule, weights, "forward", has_batch_dim) + + def reorder_weights_to_e3nn( + self, weights: jax.Array, has_batch_dim: bool = True + ) -> jax.Array: + """Convert selected OEQ-order weights to canonical e3nn order.""" + weights = jnp.asarray(weights) + if weights.ndim < 1 or weights.shape[-1] != self.weight_numel: + raise ValueError(f"weights must end in dimension {self.weight_numel}") + return reorder_jax(self.schedule, weights, "backward", has_batch_dim) + + @staticmethod + def _prepare_topology( + rows: jax.Array, + cols: jax.Array, + num_nodes: int, + *, + indices_are_sorted: bool = False, + ) -> StreamingConvTopology: + """Prepare receiver rows and offsets for one streaming call. + + When the caller has already sorted the edge operands by receiver, this + method reuses their order and only builds the offsets. Otherwise it + creates the stable permutation that places valid edges first in + nondecreasing receiver order. + """ + if num_nodes < 0: + raise ValueError("num_nodes must be non-negative") + if num_nodes == 0 and rows.shape[0] != 0: + raise ValueError("N=0 requires E=0 for streaming convolution") + if rows.ndim != 1 or cols.ndim != 1 or rows.shape != cols.shape: + raise ValueError("rows and cols must be equal-length vectors") + if rows.dtype != jnp.int32 or cols.dtype != jnp.int32: + raise ValueError("streaming convolution topology must use int32 indices") + valid_edges = ( + (rows >= 0) & (rows < num_nodes) & (cols >= 0) & (cols < num_nodes) + ) + receivers = jnp.where(valid_edges, rows, num_nodes) + order = ( + jnp.arange(rows.shape[0], dtype=jnp.int32) + if indices_are_sorted + else jnp.argsort(receivers, stable=True).astype(jnp.int32) + ) + edges_per_receiver = jnp.bincount( + jnp.where(valid_edges, rows, 0), + weights=valid_edges.astype(rows.dtype), + length=num_nodes, + ).astype(rows.dtype) + row_ptr = jnp.concatenate( + (jnp.zeros((1,), dtype=rows.dtype), jnp.cumsum(edges_per_receiver)) + ) + return StreamingConvTopology( + order=order, + receivers=receivers[order], + senders=cols[order], + row_ptr=row_ptr, + valid_edges=valid_edges[order], + ) + + @staticmethod + def _validate_row_ptr( + row_ptr: jax.Array, rows: jax.Array, cols: jax.Array, num_nodes: int + ) -> None: + """Validate the shape and dtype of cached receiver offsets.""" + del rows, cols + if row_ptr.dtype != jnp.int32 or row_ptr.shape != (num_nodes + 1,): + raise ValueError("row_ptr must be an int32 vector with shape [N + 1]") + + def _validate_native_operands( + self, + X: jax.Array, + Y: jax.Array, + W: jax.Array, + rows: jax.Array, + cols: jax.Array, + ) -> tuple[jax.Array, jax.Array, jax.Array]: + """Validate the public receiver-streaming operands.""" + X, Y, W = map(jnp.asarray, (X, Y, W)) + if X.ndim != 2 or X.shape[1] != self.schedule.input_dim: + raise ValueError(f"X must have shape [N, {self.schedule.input_dim}]") + if Y.ndim != 2 or Y.shape[1] != self.schedule.edge_dim: + raise ValueError(f"Y must have shape [E, {self.schedule.edge_dim}]") + if W.ndim != 2 or W.shape != (Y.shape[0], self.weight_numel): + raise ValueError(f"W must have shape [E, {self.weight_numel}]") + if X.dtype != Y.dtype or X.dtype != W.dtype: + raise ValueError("X, Y, and W must have matching dtypes") + expected_dtype = jnp.dtype(self.config.irrep_dtype) + if X.dtype != expected_dtype: + raise ValueError( + f"X, Y, and W must have the configured dtype {expected_dtype}" + ) + if X.shape[0] == 0 and Y.shape[0] != 0: + raise ValueError("N=0 requires E=0 for streaming convolution") + if rows.ndim != 1 or cols.ndim != 1 or rows.shape != cols.shape: + raise ValueError("rows and cols must be equal-length vectors") + if rows.shape[0] != Y.shape[0]: + raise ValueError("topology and edge feature counts must agree") + if rows.dtype != jnp.int32 or cols.dtype != jnp.int32: + raise ValueError("streaming convolution topology must use int32 indices") + return X, Y, W + + def forward( + self, + X: jax.Array, + Y: jax.Array, + W: jax.Array, + rows: Optional[jax.Array] = None, + cols: Optional[jax.Array] = None, + *, + indices_are_sorted: bool = False, + row_ptr: Optional[jax.Array] = None, + ) -> jax.Array: + """Apply receiver streaming with the public sorted-index contract.""" + if not isinstance(indices_are_sorted, bool): + raise TypeError("indices_are_sorted must be a Python bool") + if row_ptr is not None and not indices_are_sorted: + raise ValueError("row_ptr requires indices_are_sorted=True") + X, Y, W = self._validate_native_operands(X, Y, W, rows, cols) + if Y.shape[0] == 0: + return jnp.zeros((X.shape[0], self.schedule.output_dim), dtype=X.dtype) + topology = self._prepare_topology( + rows, cols, X.shape[0], indices_are_sorted=indices_are_sorted + ) + if row_ptr is not None: + self._validate_row_ptr(row_ptr, rows, cols, X.shape[0]) + topology = topology._replace(row_ptr=row_ptr) + if not indices_are_sorted: + Y, W = Y[topology.order], W[topology.order] + if self.public_layout == "ir_mul": + from openequivariance.jax import transpose_irreps + + X = transpose_irreps(X, self.config.irreps_in1, "ir_mul", "mul_ir") + Y = transpose_irreps(Y, self.config.irreps_in2, "ir_mul", "mul_ir") + from openequivariance.jax.jvp.factorized_projected_prim import ( + factorized_projected, + ) + + output = factorized_projected( + self.schedule, + X, + Y, + W, + topology.senders, + topology.receivers, + topology.row_ptr, + ) + if self.public_layout == "ir_mul": + output = transpose_irreps( + output, self.config.irreps_out, "mul_ir", "ir_mul" + ) + return output + + __call__ = forward + + +__all__ = [ + "StreamingUnavailableError", + "StreamingTensorProductConv", + "streaming_support", +] diff --git a/openequivariance/openequivariance/jax/TensorProductConv.py b/openequivariance/openequivariance/jax/TensorProductConv.py index 9234158f..7c3c78f7 100644 --- a/openequivariance/openequivariance/jax/TensorProductConv.py +++ b/openequivariance/openequivariance/jax/TensorProductConv.py @@ -1,133 +1,213 @@ -import jax -import jax.numpy as jnp -import numpy as np +"""Public JAX tensor-product convolution selection.""" + from typing import Optional -from openequivariance.jax import extlib +import jax +import numpy as np from openequivariance.core.e3nn_lite import TPProblem -from openequivariance.core.LoopUnrollConv import LoopUnrollConv -from openequivariance.jax.utils import reorder_jax - -from openequivariance.core.logging import getLogger -from openequivariance.jax.jvp import conv_prim -from openequivariance.jax.vjp import conv_func +from openequivariance.core.utils import transpose_irrep_layout +from openequivariance.jax.LoopUnrollTensorProductConv import ( + LoopUnrollTensorProductConv, +) +from openequivariance.jax.StreamingTensorProductConv import ( + StreamingTensorProductConv, + StreamingUnavailableError, + streaming_support, +) -logger = getLogger() +class TensorProductConv: + r"""Apply a tensor-product convolution with an explicit implementation mode. + ``mode="standard"`` preserves the established loop-unroll implementation. + ``mode="auto"`` uses receiver-row streaming convolution for supported + external unshared UVU problems and exact ``[1, V, 1]`` UVW reductions. It + selects the established loop-unroll implementation for other problems. + ``mode="streaming"`` requires receiver-row streaming. ``mode="standard"`` + always selects the established loop-unroll implementation. + Both public feature layouts are accepted. The standard implementation uses + differentiable boundary transposes when an ``"ir_mul"`` problem reaches + its native ``"mul_ir"`` kernel. -class TensorProductConv(LoopUnrollConv): - r""" - Identical to ``oeq.torch.TensorProductConv`` with functionality in JAX, with one - key difference: integer arrays passed to this function must have dtype - ``np.int32`` (as opposed to ``np.int64`` in the PyTorch version). - - :param problem: Specification of the tensor product. - :param deterministic: if ``False``, uses atomics for the convolution. If ``True``, uses a deterministic - fixup-based algorithm. `Default`: ``False``. - :param kahan: If ``True``, uses Kahan summation to improve accuracy during aggregation. To use this option, - the input tensors must be in float32 precision AND you must set ``deterministic=True``. *Default*: ``False``. + :param config: Specification of the tensor product. + :param deterministic: Request deterministic aggregation. This selects the + standard implementation in automatic mode and is unavailable in + streaming mode. + :param kahan: Request Kahan summation. This selects the standard + implementation in automatic mode and is unavailable in streaming mode. + :param requires_jvp: Preserve JVP selection for the standard implementation. + :param mode: One of ``"auto"``, ``"streaming"``, or ``"standard"``. """ + _MODES = frozenset(("auto", "streaming", "standard")) + def __init__( self, config: TPProblem, deterministic: bool = False, kahan: bool = False, requires_jvp: bool = True, + mode: str = "standard", ): - dp = extlib.DeviceProp(0) - self.requires_jvp = requires_jvp - super().__init__( - config, - dp, - extlib.IS_HIP, - idx_dtype=np.int32, - torch_op=False, - deterministic=deterministic, - kahan=kahan, - ) + """Create a convolution and select one static implementation.""" + if mode not in self._MODES: + choices = ", ".join(sorted(self._MODES)) + raise ValueError(f"mode must be one of {choices}. Received {mode!r}") - self.kernel = self.kernel_string + self.config = config + self.mode = mode + self.deterministic = deterministic + self.kahan = kahan + self.requires_jvp = requires_jvp 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) + self._standard_layout_adapter = False + supports_streaming, reason = streaming_support(config) + if deterministic: + supports_streaming = False + reason = ( + "deterministic aggregation is provided by the standard implementation" + ) + elif kahan: + supports_streaming = False + reason = "Kahan summation is provided by the standard implementation" - def forward( - self, - X: jax.numpy.ndarray, - Y: jax.numpy.ndarray, - W: jax.numpy.ndarray, - rows: jax.numpy.ndarray, - cols: jax.numpy.ndarray, - sender_perm: Optional[jax.numpy.ndarray] = None, - ) -> jax.numpy.ndarray: - if not self.deterministic: - sender_perm = self.dummy_transpose_perm + if mode == "streaming" and not supports_streaming: + raise StreamingUnavailableError( + f"mode='streaming' is unavailable: {reason}." + ) + if mode != "standard" and supports_streaming: + self._impl: StreamingTensorProductConv | LoopUnrollTensorProductConv = ( + StreamingTensorProductConv(config) + ) else: - assert sender_perm is not None, ( - "Must provide sender_perm for deterministic convolutions." + standard_config = config + if config.layout == "ir_mul": + # The established loop-unroll kernels are native ``mul_ir`` + # kernels. Preserve the public layout at this boundary + # instead of making their long-standing schedule interpret a + # different memory order. + standard_config = config.clone() + standard_config.layout = "mul_ir" + self._standard_layout_adapter = True + self._impl = LoopUnrollTensorProductConv( + standard_config, + deterministic=deterministic, + kahan=kahan, + requires_jvp=requires_jvp, ) - func = conv_prim.conv_fwd_p.bind + @property + def uses_streaming_kernel(self) -> bool: + """Return whether this object selected receiver-row streaming.""" + return isinstance(self._impl, StreamingTensorProductConv) - if not self.requires_jvp: - func = conv_func.forward + @property + def implementation(self) -> str: + """Return ``"streaming"`` or ``"standard"`` for this object.""" + return "streaming" if self.uses_streaming_kernel else "standard" - return func( + @property + def L3_dim(self) -> int: + """Return the flattened output feature dimension.""" + return self.config.irreps_out.dim + + def reorder_weights_from_e3nn(self, weights, has_batch_dim: bool = True): + """Convert canonical e3nn weights to the selected OEQ weight order.""" + return self._impl.reorder_weights_from_e3nn(weights, has_batch_dim) + + def reorder_weights_to_e3nn(self, weights, has_batch_dim: bool = True): + """Convert selected OEQ-order weights to canonical e3nn order.""" + return self._impl.reorder_weights_to_e3nn(weights, has_batch_dim) + + def forward( + self, + X: jax.Array, + Y: jax.Array, + W: jax.Array, + rows: Optional[jax.Array] = None, + cols: Optional[jax.Array] = None, + sender_perm: Optional[jax.Array] = None, + *, + indices_are_sorted: bool = False, + row_ptr: Optional[jax.Array] = None, + ) -> jax.Array: + """Apply the selected convolution implementation. + + When ``indices_are_sorted`` is true, ``rows``, ``cols``, ``Y``, and + ``W`` must already share receiver-row order. Streaming then skips its + topology sort and operand gathers. The standard implementation simply + consumes that aligned order. A deterministic ``sender_perm`` must use + the same order as the supplied edge operands. For streaming, valid + edges must form a nondecreasing receiver-row prefix and invalid padding + must be a tail. Under JIT this is a trusted contract, so it adds no + runtime validation. The standard implementation requires every + endpoint to be valid. ``row_ptr`` optionally reuses the cached receiver + offsets for a sorted streaming call. + """ + if not isinstance(indices_are_sorted, bool): + raise TypeError("indices_are_sorted must be a Python bool") + if row_ptr is not None and not indices_are_sorted: + raise ValueError("row_ptr requires indices_are_sorted=True") + if rows is None or cols is None: + raise ValueError("rows and cols are required for convolution") + if self.uses_streaming_kernel: + if sender_perm is not None: + raise StreamingUnavailableError( + "Receiver streaming does not accept sender_perm. Select " + "mode='standard' for deterministic aggregation." + ) + assert isinstance(self._impl, StreamingTensorProductConv) + return self._impl.forward( + X, + Y, + W, + rows, + cols, + indices_are_sorted=indices_are_sorted, + row_ptr=row_ptr, + ) + assert isinstance(self._impl, LoopUnrollTensorProductConv) + if getattr(self, "_standard_layout_adapter", False): + from openequivariance.jax import transpose_irreps + + X = transpose_irreps(X, self.config.irreps_in1, "ir_mul", "mul_ir") + Y = transpose_irreps(Y, self.config.irreps_in2, "ir_mul", "mul_ir") + output = self._impl.forward( X, Y, W, rows, cols, - self.workspace, sender_perm, - L3_dim=self.L3_dim, - kernel=self.kernel, - hash=self.hash, + indices_are_sorted=indices_are_sorted, ) - - def __call__( - self, - X: jax.numpy.ndarray, - Y: jax.numpy.ndarray, - W: jax.numpy.ndarray, - rows: jax.numpy.ndarray, - cols: jax.numpy.ndarray, - sender_perm: Optional[jax.numpy.ndarray] = None, - ) -> jax.numpy.ndarray: - return self.forward(X, Y, W, rows, cols, sender_perm) - - def reorder_weights_from_e3nn(self, weights, has_batch_dim=True): - return reorder_jax(self.forward_schedule, weights, "forward", has_batch_dim) - - def reorder_weights_to_e3nn(self, weights, has_batch_dim=True): - return reorder_jax(self.forward_schedule, weights, "backward", has_batch_dim) + if getattr(self, "_standard_layout_adapter", False): + output = transpose_irreps( + output, self.config.irreps_out, "mul_ir", "ir_mul" + ) + return output def forward_cpu(self, L1_in, L2_in, weights, L3_out, graph): - rows = graph.rows.astype(np.int32) - cols = graph.cols.astype(np.int32) - sender_perm = graph.transpose_perm.astype(np.int32) - weights = self.reorder_weights_from_e3nn( - weights, has_batch_dim=not self.config.shared_weights + """Run the standard CPU helper when the standard implementation is selected.""" + if self.uses_streaming_kernel: + raise StreamingUnavailableError( + "Receiver streaming does not provide the CPU helper API." + ) + assert isinstance(self._impl, LoopUnrollTensorProductConv) + if not self._standard_layout_adapter: + return self._impl.forward_cpu(L1_in, L2_in, weights, L3_out, graph) + native_output = np.empty_like(L3_out) + self._impl.forward_cpu( + transpose_irrep_layout(L1_in, self.config.irreps_in1, "ir_mul", "mul_ir"), + transpose_irrep_layout(L2_in, self.config.irreps_in2, "ir_mul", "mul_ir"), + weights, + native_output, + graph, ) - - jit_fwd = jax.jit(self.forward) - result = jit_fwd( - jax.numpy.asarray(L1_in), - jax.numpy.asarray(L2_in), - jax.numpy.asarray(weights), - jax.numpy.asarray(rows), - jax.numpy.asarray(cols), - jax.numpy.asarray(sender_perm), + L3_out[...] = transpose_irrep_layout( + native_output, self.config.irreps_out, "mul_ir", "ir_mul" ) - L3_out[:] = np.asarray(result) def backward_cpu( self, @@ -140,74 +220,78 @@ def backward_cpu( weights_grad, graph, ): - rows = graph.rows.astype(np.int32) - cols = graph.cols.astype(np.int32) - sender_perm = graph.transpose_perm.astype(np.int32) - weights = self.reorder_weights_from_e3nn( - weights, has_batch_dim=not self.config.shared_weights - ) - - backward_fn = jax.jit( - jax.vjp( - lambda X, Y, W: self.forward( - X, - Y, - W, - jax.numpy.asarray(rows), - jax.numpy.asarray(cols), - jax.numpy.asarray(sender_perm), - ), - jax.numpy.asarray(L1_in), - jax.numpy.asarray(L2_in), - jax.numpy.asarray(weights), - )[1] + """Run the standard backward CPU helper when it is selected.""" + if self.uses_streaming_kernel: + raise StreamingUnavailableError( + "Receiver streaming does not provide the CPU helper API." + ) + assert isinstance(self._impl, LoopUnrollTensorProductConv) + if not self._standard_layout_adapter: + return self._impl.backward_cpu( + L1_in, L1_grad, L2_in, L2_grad, L3_grad, weights, weights_grad, graph + ) + native_l1_grad = np.empty_like(L1_grad) + native_l2_grad = np.empty_like(L2_grad) + self._impl.backward_cpu( + transpose_irrep_layout(L1_in, self.config.irreps_in1, "ir_mul", "mul_ir"), + native_l1_grad, + transpose_irrep_layout(L2_in, self.config.irreps_in2, "ir_mul", "mul_ir"), + native_l2_grad, + transpose_irrep_layout(L3_grad, self.config.irreps_out, "ir_mul", "mul_ir"), + weights, + weights_grad, + graph, ) - - L1_grad_jax, L2_grad_jax, weights_grad_jax = backward_fn( - jax.numpy.asarray(L3_grad) + L1_grad[...] = transpose_irrep_layout( + native_l1_grad, self.config.irreps_in1, "mul_ir", "ir_mul" ) - L1_grad[:] = np.asarray(L1_grad_jax) - L2_grad[:] = np.asarray(L2_grad_jax) - weights_grad[:] = np.asarray(weights_grad_jax) - weights_grad[:] = self.reorder_weights_to_e3nn( - weights_grad, has_batch_dim=not self.config.shared_weights + L2_grad[...] = transpose_irrep_layout( + native_l2_grad, self.config.irreps_in2, "mul_ir", "ir_mul" ) def double_backward_cpu( self, in1, in2, out_grad, weights, weights_dgrad, in1_dgrad, in2_dgrad, graph ): - in1_jax = jax.numpy.asarray(in1) - in2_jax = jax.numpy.asarray(in2) - weights_jax = jax.numpy.asarray(weights) - out_grad_jax = jax.numpy.asarray(out_grad) - in1_dgrad_jax = jax.numpy.asarray(in1_dgrad) - in2_dgrad_jax = jax.numpy.asarray(in2_dgrad) - weights_dgrad_jax = jax.numpy.asarray(weights_dgrad) - - rows_jax = jax.numpy.asarray(graph.rows.astype(self.idx_dtype)) - cols_jax = jax.numpy.asarray(graph.cols.astype(self.idx_dtype)) - sender_perm_jax = jax.numpy.asarray(graph.transpose_perm.astype(self.idx_dtype)) - - in1_grad, in2_grad, weights_grad, out_dgrad = jax.jit( - jax.vjp( - lambda x, y, w, o: jax.vjp( - lambda a, b, c: self.forward( - a, b, c, rows_jax, cols_jax, sender_perm_jax - ), - x, - y, - w, - )[1](o), - in1_jax, - in2_jax, - weights_jax, - out_grad_jax, - )[1] - )((in1_dgrad_jax, in2_dgrad_jax, weights_dgrad_jax)) - + """Run the standard double-backward CPU helper when it is selected.""" + if self.uses_streaming_kernel: + raise StreamingUnavailableError( + "Receiver streaming does not provide the CPU helper API." + ) + assert isinstance(self._impl, LoopUnrollTensorProductConv) + if not self._standard_layout_adapter: + return self._impl.double_backward_cpu( + in1, in2, out_grad, weights, weights_dgrad, in1_dgrad, in2_dgrad, graph + ) + in1_grad, in2_grad, weights_grad, out_dgrad = self._impl.double_backward_cpu( + transpose_irrep_layout(in1, self.config.irreps_in1, "ir_mul", "mul_ir"), + transpose_irrep_layout(in2, self.config.irreps_in2, "ir_mul", "mul_ir"), + transpose_irrep_layout( + out_grad, self.config.irreps_out, "ir_mul", "mul_ir" + ), + weights, + weights_dgrad, + transpose_irrep_layout( + in1_dgrad, self.config.irreps_in1, "ir_mul", "mul_ir" + ), + transpose_irrep_layout( + in2_dgrad, self.config.irreps_in2, "ir_mul", "mul_ir" + ), + graph, + ) return ( - np.asarray(in1_grad), - np.asarray(in2_grad), - np.asarray(weights_grad), - np.asarray(out_dgrad), + transpose_irrep_layout( + in1_grad, self.config.irreps_in1, "mul_ir", "ir_mul" + ), + transpose_irrep_layout( + in2_grad, self.config.irreps_in2, "mul_ir", "ir_mul" + ), + weights_grad, + transpose_irrep_layout( + out_dgrad, self.config.irreps_out, "mul_ir", "ir_mul" + ), ) + + __call__ = forward + + +__all__ = ["TensorProductConv"] diff --git a/openequivariance/openequivariance/jax/__init__.py b/openequivariance/openequivariance/jax/__init__.py index c26606b6..b5089a29 100644 --- a/openequivariance/openequivariance/jax/__init__.py +++ b/openequivariance/openequivariance/jax/__init__.py @@ -1,3 +1,5 @@ +"""JAX tensor-product and convolution interfaces.""" + import jax import jax.numpy as jnp @@ -6,6 +8,14 @@ from openequivariance.jax.TensorProductConv import ( TensorProductConv as TensorProductConv, ) +from openequivariance.jax.LoopUnrollTensorProductConv import ( + LoopUnrollTensorProductConv as LoopUnrollTensorProductConv, +) +from openequivariance.jax.StreamingTensorProductConv import ( + StreamingTensorProductConv as StreamingTensorProductConv, + StreamingUnavailableError as StreamingUnavailableError, + streaming_support as streaming_support, +) def transpose_irreps( @@ -78,4 +88,12 @@ def transpose_irreps( return out -__all__ = ["TensorProduct", "TensorProductConv", "transpose_irreps"] +__all__ = [ + "TensorProduct", + "TensorProductConv", + "LoopUnrollTensorProductConv", + "StreamingTensorProductConv", + "StreamingUnavailableError", + "streaming_support", + "transpose_irreps", +] diff --git a/openequivariance/openequivariance/jax/ffi_targets.py b/openequivariance/openequivariance/jax/ffi_targets.py index 8387a33c..a68795c4 100644 --- a/openequivariance/openequivariance/jax/ffi_targets.py +++ b/openequivariance/openequivariance/jax/ffi_targets.py @@ -1,4 +1,4 @@ -"""Names of JAX FFI handlers provided by the native extension.""" +"""Names of JAX FFI kernel targets.""" TENSOR_PRODUCT_TARGETS = ( "tp_forward", @@ -12,4 +12,5 @@ "conv_double_backward", ) -FFI_TARGETS = TENSOR_PRODUCT_TARGETS + CONVOLUTION_TARGETS +FACTORIZED_TARGET = "factorized_projected" +FFI_TARGETS = TENSOR_PRODUCT_TARGETS + CONVOLUTION_TARGETS + (FACTORIZED_TARGET,) diff --git a/openequivariance/openequivariance/jax/jvp/factorized_projected_prim.py b/openequivariance/openequivariance/jax/jvp/factorized_projected_prim.py new file mode 100644 index 00000000..c4ef2442 --- /dev/null +++ b/openequivariance/openequivariance/jax/jvp/factorized_projected_prim.py @@ -0,0 +1,905 @@ +"""AD primitives for generated preprojected factorized kernels.""" + +import jax +import jax.numpy as jnp +import numpy as np +from jax.extend import core +from jax.interpreters import ad, batching, mlir + +from openequivariance.core.FactorizedComputationSchedule import Input +from openequivariance.jax.ffi_targets import FACTORIZED_TARGET +from openequivariance.jax.extlib import IS_HIP + +FACTORIZED_FORWARD = 0 +FACTORIZED_BACKWARD = 2 + + +def factorized_forward_jvp_operation(active: tuple[bool, bool, bool]) -> int: + """Encode the active forward-JVP tangent buffers for the native ABI.""" + return 16 + sum(1 << index for index, value in enumerate(active) if value) + + +def factorized_backward_jvp_operation( + active: tuple[bool, bool, bool, bool], +) -> int: + """Encode the active backward-JVP tangent buffers for the native ABI.""" + mask = sum(1 << index for index, value in enumerate(active) if value) + if mask == 0: + raise ValueError("factorized backward JVP requires an active tangent") + return 32 + mask + + +def _materialize_symbolic_zero(value): + """Materialize a symbolic AD zero only after its activity was classified. + + Lowerings omit inactive operands, so callers must not materialize a zero + merely to satisfy a generated-kernel argument list. + """ + if ad.is_undefined_primal(value) or type(value) is ad.Zero: + return jnp.zeros(value.aval.shape, value.aval.dtype) + return value + + +def _generated_attrs(schedule, dtype, **source_options): + """Return FFI attributes generated by the static computation schedule.""" + specialization = schedule.kernel(dtype, is_hip=IS_HIP, **source_options) + return dict(specialization.ffi_attributes) + + +def _ffi_lowering(target, operand_indices, *, operation, source_options=None): + """Lower directly to FFI without recursively tracing the Python impl. + + ``mlir.lower_fun`` is unsuitable here: the implementation constructs an + ``ffi_call`` whose internal lowering cache can retain the surrounding + primitive tracer during the nested trace. Building the FFI custom call + directly is both simpler and safe for repeated JIT/export traces. + """ + source_options = source_options or (lambda params: {}) + + def lowering(ctx, *operands, schedule, **params): + attributes = _generated_attrs( + schedule, + ctx.avals_in[0].dtype, + **source_options(params), + ) + attributes["operation"] = ( + operation(params) if callable(operation) else operation + ) + rule = jax.ffi.ffi_lowering(target) + indices = ( + operand_indices(params) if callable(operand_indices) else operand_indices + ) + selected = tuple(operands[index] for index in indices) + return rule( + ctx.replace(avals_in=tuple(ctx.avals_in[index] for index in indices)), + *selected, + **attributes, + ) + + return lowering + + +def _register_generated_lowering( + primitive, operand_indices, *, operation, source_options=None +): + """Register the identical generated FFI lowering on both GPU platforms.""" + for platform in ("cuda", "rocm"): + mlir.register_lowering( + primitive, + _ffi_lowering( + FACTORIZED_TARGET, + operand_indices, + operation=operation, + source_options=source_options, + ), + platform=platform, + ) + + +def _validate(schedule, x, sh, weights, senders, receivers, row_ptr, dout=None): + e = senders.shape[0] + if x.ndim != 2 or x.shape[1] != schedule.input_dim: + raise ValueError(f"x must have shape [N, {schedule.input_dim}]") + if sh.shape != (e, schedule.edge_dim) or weights.shape != ( + e, + schedule.weight_numel, + ): + raise ValueError("sh/weights must have E rows and schedule feature widths") + if receivers.shape != senders.shape or senders.dtype != np.dtype(np.int32): + raise ValueError("senders/receivers must be equal int32 vectors") + if receivers.dtype != np.dtype(np.int32) or row_ptr.dtype != np.dtype(np.int32): + raise ValueError("topology must use int32") + if row_ptr.shape != (x.shape[0] + 1,): + raise ValueError("receiver row pointer must have shape [N+1]") + if sh.dtype != x.dtype or weights.dtype != x.dtype: + raise ValueError("floating operand dtypes must match") + if dout is not None and ( + dout.shape != (x.shape[0], schedule.output_dim) or dout.dtype != x.dtype + ): + raise ValueError("dout shape/dtype mismatch") + + +dbwd_p = core.Primitive("factorized_projected_q_dbwd") +dbwd_p.multiple_results = True +fwd_jvp_p = core.Primitive("factorized_projected_fwd_jvp") + + +def _batch_factorized(primitive, value_kinds, result_kinds, tangent_start=None): + """Flatten a shared-topology batch into one generated convolution call. + + Valid graph blocks are packed before one global padded suffix. Graph-local + indices are offset before flattening, while invalid endpoints are mapped to + the single out-of-bounds node index of the flattened batch. This preserves + receiver order and lets the generated kernels discard padding immediately. + """ + + def rule(args, axes, *, schedule, **params): + value_count = len(value_kinds) + topology_axes = axes[value_count:] + if any(axis is not None for axis in topology_axes): + raise NotImplementedError( + "vmap over receiver topology is unsupported. Map X, Y, W, " + "and derivative arrays while keeping rows and cols static." + ) + sizes = [ + value.shape[axis] + for value, axis in zip(args[:value_count], axes[:value_count]) + if axis is not None + ] + if not sizes: + result = primitive.bind(*args, schedule=schedule, **params) + return result, (0,) * len(result) if primitive.multiple_results else 0 + batch_size = sizes[0] + if any(size != batch_size for size in sizes[1:]): + raise ValueError("vmap batch dimensions must have matching sizes") + + def leading_axis(value, axis): + if axis is None: + return jnp.broadcast_to(value, (batch_size,) + value.shape) + return jnp.moveaxis(value, axis, 0) + + active = params.get("active") + values = [] + for index, (value, axis, kind) in enumerate( + zip(args[:value_count], axes[:value_count], value_kinds) + ): + is_active = ( + tangent_start is None + or index < tangent_start + or active[index - tangent_start] + ) + values.append(value if not is_active else leading_axis(value, axis)) + + x = values[0] + nodes, edges = x.shape[1], args[value_count].shape[0] + node_offsets = jnp.arange(batch_size, dtype=jnp.int32) * nodes + senders, receivers, row_ptr = args[value_count:] + edge_indices = jnp.arange(edges, dtype=jnp.int32) + valid_count = row_ptr[-1] + valid_edges = edge_indices < valid_count + graph_indices = jnp.arange(batch_size, dtype=jnp.int32)[:, None] + destination = jnp.where( + valid_edges[None, :], + graph_indices * valid_count + edge_indices[None, :], + batch_size * valid_count + + graph_indices * (edges - valid_count) + + edge_indices[None, :] + - valid_count, + ).reshape(-1) + + def pack_edges(value): + flat = value.reshape((-1,) + value.shape[2:]) + return jnp.zeros_like(flat).at[destination].set(flat) + + sentinel = jnp.asarray(batch_size * nodes, dtype=jnp.int32) + senders = pack_edges( + jnp.where( + valid_edges[None, :], + senders[None, :] + node_offsets[:, None], + sentinel, + ) + ) + receivers = pack_edges( + jnp.where( + valid_edges[None, :], + receivers[None, :] + node_offsets[:, None], + sentinel, + ) + ) + row_offsets = jnp.arange(batch_size, dtype=jnp.int32) * valid_count + row_ptr = jnp.concatenate( + ( + (row_offsets[:, None] + row_ptr[None, :-1]).reshape(-1), + jnp.asarray((batch_size * valid_count,), dtype=jnp.int32), + ) + ) + + flat_values = [] + for index, (value, kind) in enumerate(zip(values, value_kinds)): + is_active = ( + tangent_start is None + or index < tangent_start + or active[index - tangent_start] + ) + if not is_active: + flat_values.append(value) + elif kind == "node": + flat_values.append(value.reshape(-1, value.shape[-1])) + else: + flat_values.append(pack_edges(value)) + + result = primitive.bind( + *flat_values, senders, receivers, row_ptr, schedule=schedule, **params + ) + + def restore(value, kind): + if kind == "node": + return value.reshape(batch_size, nodes, value.shape[-1]) + return value[destination].reshape(batch_size, edges, value.shape[-1]) + + if primitive.multiple_results: + return ( + tuple( + restore(value, kind) for value, kind in zip(result, result_kinds) + ), + (0,) * len(result_kinds), + ) + return restore(result, result_kinds[0]), 0 + + return rule + + +def _fixed_generated_primitive( + name, operation, ffi_indices, result_indices, *, dout_index=None +): + """Create one fixed-arity projected primitive from its native ABI descriptor. + + The descriptor keeps the public primitive name, operation code, result + shape, and native operand order together. It removes repeated plumbing + without merging the distinct JAX primitives used by AD and batching rules. + """ + primitive = core.Primitive(name) + primitive.multiple_results = isinstance(result_indices, tuple) + + def validate(args, schedule): + x, sh, weights = args[:3] + senders, receivers, row_ptr = args[-3:] + dout = None if dout_index is None else args[dout_index] + _validate(schedule, x, sh, weights, senders, receivers, row_ptr, dout) + + def specs(args, schedule, array_type): + if result_indices == "output": + values = (((args[0].shape[0], schedule.output_dim), args[0].dtype),) + else: + indices = ( + result_indices + if isinstance(result_indices, tuple) + else (result_indices,) + ) + values = tuple((args[index].shape, args[index].dtype) for index in indices) + return tuple(array_type(shape, dtype) for shape, dtype in values) + + def impl(*args, schedule): + validate(args, schedule) + attributes = _generated_attrs(schedule, args[0].dtype) + shapes = specs(args, schedule, jax.ShapeDtypeStruct) + return jax.ffi.ffi_call( + FACTORIZED_TARGET, shapes if primitive.multiple_results else shapes[0] + )( + *(args[index] for index in ffi_indices), + operation=operation, + **attributes, + ) + + def abstract(*args, schedule): + validate(args, schedule) + values = specs(args, schedule, jax.core.ShapedArray) + return values if primitive.multiple_results else values[0] + + primitive.def_impl(impl) + primitive.def_abstract_eval(abstract) + _register_generated_lowering(primitive, ffi_indices, operation=operation) + return primitive + + +fwd_p = _fixed_generated_primitive( + "factorized_projected_fwd", + FACTORIZED_FORWARD, + (Input.X.value, Input.SH.value, Input.W.value, 3, 5), + "output", +) +# B_q returns the complete gradient with q=(x, sh, preprojected weights). +bwd_p = _fixed_generated_primitive( + "factorized_projected_q_bwd", + FACTORIZED_BACKWARD, + (Input.X.value, Input.SH.value, Input.W.value, 4, 5, 3), + (0, 1, 2), + dout_index=3, +) + + +def _dbwd_impl( + x, + sh, + weights, + dout, + tx, + tsh, + tweights, + tdout, + senders, + receivers, + row_ptr, + *, + schedule, + active, +): + _validate(schedule, x, sh, weights, senders, receivers, row_ptr, dout) + attributes = _generated_attrs(schedule, x.dtype, backward_jvp_active=active) + shapes = tuple(jax.ShapeDtypeStruct(a.shape, a.dtype) for a in (x, sh, weights)) + operands = (x, sh, weights, senders, receivers, dout) + tuple( + value + for value, requested in zip((tx, tsh, tweights, tdout), active) + if requested + ) + return jax.ffi.ffi_call(FACTORIZED_TARGET, shapes)( + *operands, + operation=factorized_backward_jvp_operation(active), + **attributes, + ) + + +def _dbwd_abstract( + x, + sh, + weights, + dout, + tx, + tsh, + tweights, + tdout, + senders, + receivers, + row_ptr, + *, + schedule, + active, +): + _validate(schedule, x, sh, weights, senders, receivers, row_ptr, dout) + for a, b, requested in zip( + (x, sh, weights, dout), (tx, tsh, tweights, tdout), active + ): + if requested and (a.shape != b.shape or a.dtype != b.dtype): + raise ValueError("tangent shape/dtype mismatch") + return tuple(jax.core.ShapedArray(a.shape, a.dtype) for a in (x, sh, weights)) + + +dbwd_p.def_impl(_dbwd_impl) +dbwd_p.def_abstract_eval(_dbwd_abstract) +_register_generated_lowering( + dbwd_p, + lambda params: ( + (Input.X.value, Input.SH.value, Input.W.value, 8, 9, 3) + + tuple(4 + index for index, value in enumerate(params["active"]) if value) + ), + operation=lambda params: factorized_backward_jvp_operation(params["active"]), + source_options=lambda params: {"backward_jvp_active": params["active"]}, +) + + +def _fwd_jvp_impl( + x, sh, weights, tx, tsh, tweights, senders, receivers, row_ptr, *, schedule, active +): + if not any(active): + return jnp.zeros((x.shape[0], schedule.output_dim), dtype=x.dtype) + attributes = _generated_attrs(schedule, x.dtype, forward_jvp_active=active) + shape = jax.ShapeDtypeStruct((x.shape[0], schedule.output_dim), x.dtype) + operands = (x, sh, weights, senders, row_ptr) + tuple( + value for value, requested in zip((tx, tsh, tweights), active) if requested + ) + return jax.ffi.ffi_call(FACTORIZED_TARGET, shape)( + *operands, + operation=factorized_forward_jvp_operation(active), + **attributes, + ) + + +def _fwd_jvp_abstract( + x, sh, weights, tx, tsh, tweights, senders, receivers, row_ptr, *, schedule, active +): + _validate(schedule, x, sh, weights, senders, receivers, row_ptr) + return jax.core.ShapedArray((x.shape[0], schedule.output_dim), x.dtype) + + +fwd_jvp_p.def_impl(_fwd_jvp_impl) +fwd_jvp_p.def_abstract_eval(_fwd_jvp_abstract) +_register_generated_lowering( + fwd_jvp_p, + lambda params: ( + (Input.X.value, Input.SH.value, Input.W.value, 6, 8) + + tuple( + len(Input) + index for index, value in enumerate(params["active"]) if value + ) + ), + operation=lambda params: factorized_forward_jvp_operation(params["active"]), + source_options=lambda params: {"forward_jvp_active": params["active"]}, +) + + +def _fwd_jvp_rule(primals, tangents, *, schedule): + x, sh, weights, senders, receivers, row_ptr = primals + active = tuple(type(tangent) is not ad.Zero for tangent in tangents[:3]) + tx, tsh, tw = ( + _materialize_symbolic_zero(tangent) + if is_active + else jnp.zeros((), dtype=primal.dtype) + for primal, tangent, is_active in zip((x, sh, weights), tangents[:3], active) + ) + return ( + fwd_p.bind(*primals, schedule=schedule), + fwd_jvp_p.bind( + x, + sh, + weights, + tx, + tsh, + tw, + senders, + receivers, + row_ptr, + schedule=schedule, + active=active, + ), + ) + + +ad.primitive_jvps[fwd_p] = _fwd_jvp_rule + + +def _fwd_jvp_jvp_rule(primals, tangents, *, schedule, active): + """Differentiate the trilinear projected forward-JVP exactly. + + For trilinear ``F``, the first JVP is the sum of ``F(tx, sh, w)``, + ``F(x, tsh, w)``, and ``F(x, sh, tw)``. Differentiating each of its three + operands gives a 3-by-3 product rule of generated forward-JVP calls. + """ + x, sh, weights, tx, tsh, tw, senders, receivers, row_ptr = primals + dx, dsh, dw, dtx, dtsh, dtw, _, _, _ = tangents + + # ``active`` describes (tx, tsh, tw) in the inner JVP. ``q_active`` marks + # outer derivatives of (x, sh, weights), while ``t_active`` marks outer + # derivatives of those inner tangents. Classify them before materialising + # zeros: inactive scalar placeholders must never reach a generated kernel. + q_active = tuple(type(tangent) is not ad.Zero for tangent in (dx, dsh, dw)) + t_active = tuple( + requested and type(tangent) is not ad.Zero + for requested, tangent in zip(active, (dtx, dtsh, dtw)) + ) + zero_x, zero_sh, zero_w = ( + jnp.zeros_like(x), + jnp.zeros_like(sh), + jnp.zeros_like(weights), + ) + tx = _materialize_symbolic_zero(tx) if active[0] else zero_x + tsh = _materialize_symbolic_zero(tsh) if active[1] else zero_sh + tw = _materialize_symbolic_zero(tw) if active[2] else zero_w + dx = _materialize_symbolic_zero(dx) if q_active[0] else zero_x + dsh = _materialize_symbolic_zero(dsh) if q_active[1] else zero_sh + dw = _materialize_symbolic_zero(dw) if q_active[2] else zero_w + dtx = _materialize_symbolic_zero(dtx) if t_active[0] else zero_x + dtsh = _materialize_symbolic_zero(dtsh) if t_active[1] else zero_sh + dtw = _materialize_symbolic_zero(dtw) if t_active[2] else zero_w + + # Rows select the tangent operand in the first JVP. Columns differentiate + # one operand of that trilinear summand. Row-major order defines the + # generated product-rule summand order. + base_operands = (x, sh, weights) + tangent_operands = (tx, tsh, tw) + primal_directions = (dx, dsh, dw) + tangent_directions = (dtx, dtsh, dtw) + zero_directions = (zero_x, zero_sh, zero_w) + product_rule = [] + for tangent_axis in range(3): + values = list(base_operands) + values[tangent_axis] = tangent_operands[tangent_axis] + for derivative_axis in range(3): + directions = list(zero_directions) + if derivative_axis == tangent_axis: + enabled = t_active[tangent_axis] + directions[derivative_axis] = tangent_directions[derivative_axis] + else: + enabled = active[tangent_axis] and q_active[derivative_axis] + directions[derivative_axis] = primal_directions[derivative_axis] + term_active = tuple(axis == derivative_axis for axis in range(3)) + product_rule.append( + (enabled, tuple(values), tuple(directions), term_active) + ) + terms = [ + fwd_jvp_p.bind( + *values, + *directions, + senders, + receivers, + row_ptr, + schedule=schedule, + active=term_active, + ) + for enabled, values, directions, term_active in product_rule + if enabled + ] + tangent = ( + sum(terms[1:], terms[0]) + if terms + else jnp.zeros((x.shape[0], schedule.output_dim), dtype=x.dtype) + ) + return ( + fwd_jvp_p.bind( + x, + sh, + weights, + tx, + tsh, + tw, + senders, + receivers, + row_ptr, + schedule=schedule, + active=active, + ), + tangent, + ) + + +ad.primitive_jvps[fwd_jvp_p] = _fwd_jvp_jvp_rule + + +def _fwd_jvp_transpose( + ct, x, sh, weights, tx, tsh, tw, senders, receivers, row_ptr, *, schedule, active +): + if any(ad.is_undefined_primal(value) for value in (x, sh, weights)): + raise NotImplementedError( + "factorized_projected_fwd_jvp transpose requires defined coefficients" + ) + undefined = tuple(ad.is_undefined_primal(value) for value in (tx, tsh, tw)) + if any( + is_undefined and not requested + for is_undefined, requested in zip(undefined, active) + ): + raise NotImplementedError( + "inactive factorized_projected_fwd_jvp directions must be defined" + ) + if type(ct) is ad.Zero: + return ( + None, + None, + None, + *( + jnp.zeros_like(_materialize_symbolic_zero(value)) + if is_undefined + else None + for value, is_undefined in zip((tx, tsh, tw), undefined) + ), + None, + None, + None, + ) + x, sh, weights, ct = map(_materialize_symbolic_zero, (x, sh, weights, ct)) + grads = bwd_p.bind( + x, sh, weights, ct, senders, receivers, row_ptr, schedule=schedule + ) + tangent_grads = tuple( + grad if is_undefined else None for grad, is_undefined in zip(grads, undefined) + ) + return (None, None, None, *tangent_grads, None, None, None) + + +ad.primitive_transposes[fwd_jvp_p] = _fwd_jvp_transpose + + +def _bwd_jvp_rule(primals, tangents, *, schedule): + x, sh, weights, dout, senders, receivers, row_ptr = primals + active = tuple(type(tangent) is not ad.Zero for tangent in tangents[:4]) + primal = bwd_p.bind(*primals, schedule=schedule) + if not any(active): + return primal, tuple(jnp.zeros_like(value) for value in primal) + tx, tsh, tw, tdout = ( + _materialize_symbolic_zero(tangent) + if is_active + else jnp.zeros((), dtype=primal.dtype) + for primal, tangent, is_active in zip( + (x, sh, weights, dout), tangents[:4], active + ) + ) + return ( + primal, + dbwd_p.bind( + x, + sh, + weights, + dout, + tx, + tsh, + tw, + tdout, + senders, + receivers, + row_ptr, + schedule=schedule, + active=active, + ), + ) + + +ad.primitive_jvps[bwd_p] = _bwd_jvp_rule + + +def _bwd_transpose( + cotangents, x, sh, weights, dout, senders, receivers, row_ptr, *, schedule +): + """Transpose one differentiated input of the projected backward. + + This is the Hessian symmetry identity for ``B_q = grad_q ``. + Higher-order transforms only need one unknown primal at a time. Rejecting + wider requests avoids pretending that an incomplete custom rule handles a + general reverse sweep. + """ + primals = (x, sh, weights, dout) + undefined = tuple(ad.is_undefined_primal(value) for value in primals) + if sum(undefined) != 1: + raise NotImplementedError( + "factorized_projected_q_bwd transpose requires one undefined primal" + ) + active = tuple(type(value) is not ad.Zero for value in cotangents) + if not any(active): + zeros = tuple( + jnp.zeros_like(_materialize_symbolic_zero(value)) for value in primals + ) + return ( + *(zero if requested else None for zero, requested in zip(zeros, undefined)), + None, + None, + None, + ) + x, sh, weights, dout = map(_materialize_symbolic_zero, primals) + cx, csh, cw = ( + _materialize_symbolic_zero(value) + if requested + else jnp.zeros((), dtype=primal.dtype) + for value, primal, requested in zip(cotangents, (x, sh, weights), active) + ) + q_grads = dbwd_p.bind( + x, + sh, + weights, + dout, + cx, + csh, + cw, + jnp.zeros_like(dout), + senders, + receivers, + row_ptr, + schedule=schedule, + active=(*active, False), + ) + dout_grad = fwd_jvp_p.bind( + x, + sh, + weights, + cx, + csh, + cw, + senders, + receivers, + row_ptr, + schedule=schedule, + active=active, + ) + grads = (*q_grads, dout_grad) + return ( + *(grad if requested else None for grad, requested in zip(grads, undefined)), + None, + None, + None, + ) + + +ad.primitive_transposes[bwd_p] = _bwd_transpose + + +def _dbwd_jvp_rule(primals, tangents, *, schedule, active): + """Differentiate ``DB_q(z)[v]`` using its multilinear cross terms.""" + x, sh, weights, dout, tx, tsh, tw, tdout, senders, receivers, row_ptr = primals + dx, dsh, dw, ddout, dtx, dtsh, dtw, dtdout, _, _, _ = tangents + z = (x, sh, weights, dout) + dz = (dx, dsh, dw, ddout) + v = (tx, tsh, tw, tdout) + dv = (dtx, dtsh, dtw, dtdout) + z_active = tuple(type(tangent) is not ad.Zero for tangent in dz) + dv_active = tuple( + requested and type(tangent) is not ad.Zero + for requested, tangent in zip(active, dv) + ) + zero_v = tuple(jnp.zeros((), dtype=value.dtype) for value in z) + primal = dbwd_p.bind(*primals, schedule=schedule, active=active) + if any(dv_active): + dv_term = dbwd_p.bind( + x, + sh, + weights, + dout, + *( + _materialize_symbolic_zero(tangent) if requested else zero + for tangent, requested, zero in zip(dv, dv_active, zero_v) + ), + senders, + receivers, + row_ptr, + schedule=schedule, + active=dv_active, + ) + else: + dv_term = tuple(jnp.zeros_like(value) for value in primal) + + tangent = list(dv_term) + for i, (zi, dzi) in enumerate(zip(z, dz)): + if not z_active[i]: + continue + for j, vj in enumerate(v): + if i == j or not active[j]: + continue + cross_inputs = list(z) + cross_inputs[i] = _materialize_symbolic_zero(dzi) + cross_inputs[j] = _materialize_symbolic_zero(vj) + cross = bwd_p.bind( + *cross_inputs, senders, receivers, row_ptr, schedule=schedule + ) + for k, value in enumerate(cross): + if k != i and k != j: + tangent[k] = tangent[k] + value + return primal, tuple(tangent) + + +ad.primitive_jvps[dbwd_p] = _dbwd_jvp_rule + + +def _dbwd_transpose( + cotangents, + x, + sh, + weights, + dout, + tx, + tsh, + tweights, + tdout, + senders, + receivers, + row_ptr, + *, + schedule, + active, +): + """Transpose the mixed derivative using Hessian symmetry. + + The spatial backward is the gradient of ``vdot(F, dout)`` with respect + to ``(x, sh, weights)``. Its derivative is therefore symmetric in those + three tangent slots. The remaining ``dout`` cotangent is the forward JVP + in the received cotangent direction. This closes reverse-over-forward + force-loss differentiation without a materialized reference expansion. + """ + if any(ad.is_undefined_primal(value) for value in (x, sh, weights, dout)): + raise NotImplementedError( + "factorized_projected_q_dbwd transpose requires defined coefficients" + ) + undefined = tuple( + ad.is_undefined_primal(value) for value in (tx, tsh, tweights, tdout) + ) + if any( + is_undefined and not requested + for is_undefined, requested in zip(undefined, active) + ): + raise NotImplementedError( + "inactive factorized_projected_q_dbwd directions must be defined" + ) + active_cotangents = tuple(type(value) is not ad.Zero for value in cotangents) + if not any(active_cotangents): + zeros = tuple( + jnp.zeros_like(_materialize_symbolic_zero(value)) + for value in (tx, tsh, tweights, tdout) + ) + return ( + None, + None, + None, + None, + *(zero if requested else None for zero, requested in zip(zeros, undefined)), + None, + None, + None, + ) + x, sh, weights, dout = map(_materialize_symbolic_zero, (x, sh, weights, dout)) + cx, csh, cw = ( + _materialize_symbolic_zero(value) + if requested + else jnp.zeros((), dtype=primal.dtype) + for value, primal, requested in zip( + cotangents, (x, sh, weights), active_cotangents + ) + ) + zero_dout = jnp.zeros_like(dout) + q_grads = dbwd_p.bind( + x, + sh, + weights, + dout, + cx, + csh, + cw, + zero_dout, + senders, + receivers, + row_ptr, + schedule=schedule, + active=(*active_cotangents, False), + ) + dout_grad = fwd_jvp_p.bind( + x, + sh, + weights, + cx, + csh, + cw, + senders, + receivers, + row_ptr, + schedule=schedule, + active=active_cotangents, + ) + tangent_grads = (q_grads[0], q_grads[1], q_grads[2], dout_grad) + tangent_grads = tuple( + value if is_undefined else None + for value, is_undefined in zip(tangent_grads, undefined) + ) + return ( + None, + None, + None, + None, + *tangent_grads, + None, + None, + None, + ) + + +ad.primitive_transposes[dbwd_p] = _dbwd_transpose + + +# The generated kernels operate on one graph topology at a time. These rules +# preserve that contract while allowing JAX to map feature and weight batches. +batching.primitive_batchers[fwd_p] = _batch_factorized( + fwd_p, ("node", "edge", "edge"), ("node",) +) +batching.primitive_batchers[bwd_p] = _batch_factorized( + bwd_p, ("node", "edge", "edge", "node"), ("node", "edge", "edge") +) +batching.primitive_batchers[dbwd_p] = _batch_factorized( + dbwd_p, + ("node", "edge", "edge", "node", "node", "edge", "edge", "node"), + ("node", "edge", "edge"), + tangent_start=4, +) +batching.primitive_batchers[fwd_jvp_p] = _batch_factorized( + fwd_jvp_p, + ("node", "edge", "edge", "node", "edge", "edge"), + ("node",), + tangent_start=3, +) + + +def factorized_projected(schedule, x, sh, weights, senders, receivers, row_ptr): + return fwd_p.bind(x, sh, weights, senders, receivers, row_ptr, schedule=schedule) + + +__all__ = ["factorized_projected"] diff --git a/openequivariance/openequivariance/templates/factorized_projected.cuh b/openequivariance/openequivariance/templates/factorized_projected.cuh new file mode 100644 index 00000000..9507ddef --- /dev/null +++ b/openequivariance/openequivariance/templates/factorized_projected.cuh @@ -0,0 +1,474 @@ +// Receiver-owned forward kernels and edge-owned reverse kernels generated from +// sparse coupling paths. +using int64_t = signed long long; +using int32_t = int; +using scalar_t = {{ scalar }}; +using acc_t = {{ scalar }}; + +constexpr int INPUT_DIM = {{ schedule.input_dim }}; +constexpr int EDGE_DIM = {{ schedule.edge_dim }}; +constexpr int OUTPUT_DIM = {{ schedule.output_dim }}; +constexpr int WEIGHT_DIM = {{ schedule.weight_numel }}; +constexpr int CHANNEL_DIM = {{ schedule.channels }}; +constexpr int LOGICAL_GROUP_SIZE = + {{ schedule.launch_config.logical_cohort_width }}; +constexpr unsigned int FULL_MASK = 0xffffffffu; + +struct EdgeRange { + int64_t begin; + int64_t end; +}; + +__device__ __forceinline__ EdgeRange safe_edge_range( + const int32_t* row_ptr, + const int64_t node, + const int64_t edge_count) { + const int64_t begin = row_ptr[node] < 0 ? 0 : row_ptr[node]; + int64_t end = row_ptr[node + 1] < begin ? begin : row_ptr[node + 1]; + end = end > edge_count ? edge_count : end; + return {begin, end}; +} + +{% macro value_or_zero(active, value) -%} + {%- if active -%} + {{ value }} + {%- else -%} + scalar_t(0) + {%- endif -%} +{%- endmacro %} + +{# Selectively active tangent reads. Inactive operands render as zero. #} +{% macro value_if_active(input, index, forward) -%} + {%- if forward -%} + {%- if forward_jvp_active[input.value] -%} + {%- if input == Input.X -%} + tx[{{ index }}] + {%- elif input == Input.SH -%} + tsh[{{ index }}] + {%- elif input == Input.W -%} + tweights[{{ index }}] + {%- endif -%} + {%- else -%} + scalar_t(0) + {%- endif -%} + {%- else -%} + {%- if backward_jvp_active[input.value] -%} + {%- if input == Input.X -%} + tx[{{ index }}] + {%- elif input == Input.SH -%} + tsh[{{ index }}] + {%- elif input == Input.W -%} + tweights[{{ index }}] + {%- endif -%} + {%- else -%} + scalar_t(0) + {%- endif -%} + {%- endif -%} +{%- endmacro %} + + +{% macro angular_contraction(name, path, output, edge_channel, tangent=false) -%} + {%- if output.terms | length %} + {%- for term in output.terms %} + const int64_t {{ name }}_input_index_{{ loop.index0 }} = + input_base + {{ path.input_start + term.input_component }} + + channel * {{ path.input_irrep_dim }}; + const int64_t {{ name }}_edge_index_{{ loop.index0 }} = + edge_base + {{ path.edge_start + term.edge_component }} + + {{ edge_channel }} * {{ path.edge_irrep_dim }}; + const scalar_t {{ name }}_input_value_{{ loop.index0 }} = + x[{{ name }}_input_index_{{ loop.index0 }}]; + const scalar_t {{ name }}_edge_value_{{ loop.index0 }} = + sh[{{ name }}_edge_index_{{ loop.index0 }}]; + {%- if tangent %} + const scalar_t {{ name }}_tangent_input_value_{{ loop.index0 }} = + {{ value_if_active( + Input.X, name ~ "_input_index_" ~ loop.index0, true + ) }}; + const scalar_t {{ name }}_tangent_edge_value_{{ loop.index0 }} = + {{ value_if_active( + Input.SH, name ~ "_edge_index_" ~ loop.index0, true + ) }}; + {%- endif %} + {%- endfor %} + + const scalar_t {{ name }} = + {%- for term in output.terms %} + {%- if not loop.first %} + + + {%- endif %} + + {%- if tangent %} + ({{ name }}_tangent_input_value_{{ loop.index0 }} + * {{ name }}_edge_value_{{ loop.index0 }} + + {{ name }}_input_value_{{ loop.index0 }} + * {{ name }}_tangent_edge_value_{{ loop.index0 }}) + {%- else %} + {{ name }}_input_value_{{ loop.index0 }} + * {{ name }}_edge_value_{{ loop.index0 }} + {%- endif %} + * {{ cpp_scalar_literal(term.coefficient, scalar) }} + {%- endfor %} + ; + {%- else %} + const scalar_t {{ name }} = scalar_t(0); + {%- endif %} +{%- endmacro %} + +// One thread owns (receiver node, channel), scans that receiver's CSR edges, +// and directly writes all of its output components. +extern "C" __global__ void oeq_projected_forward( + int64_t node_count, + int64_t edge_count, + const scalar_t* x, + const scalar_t* sh, + const scalar_t* weights, + const int32_t* senders, + const int32_t* row_ptr, + scalar_t* out) { + + const int64_t node_channel_index = + int64_t(blockIdx.x) * blockDim.x + threadIdx.x; + const int64_t node_channel_count = node_count * CHANNEL_DIM; + if (node_channel_index >= node_channel_count) return; + + const int64_t node = node_channel_index / CHANNEL_DIM; + const int channel = int(node_channel_index - node * CHANNEL_DIM); + const int64_t output_base = node * OUTPUT_DIM; + + {%- for slot in schedule.output_slots %} + scalar_t {{ slot.name }} = scalar_t(0); + {%- endfor %} + + const EdgeRange edge_range = + safe_edge_range(row_ptr, node, edge_count); + for (int64_t e = edge_range.begin; e < edge_range.end; ++e) { + const int32_t sender = senders[e]; + if (sender < 0 || sender >= node_count) continue; + + const int64_t input_base = int64_t(sender) * INPUT_DIM; + const int64_t edge_base = e * EDGE_DIM; + const int64_t weight_base = e * WEIGHT_DIM; + + {% for path in schedule.paths %} + {% for edge_channel in range(path.edge_mul) %} { + const int weight_index = + {{ path.weight_start + edge_channel * schedule.channels }} + channel; + const int64_t radial_weight_index = weight_base + weight_index; + const scalar_t radial_weight = weights[radial_weight_index]; + + {% for output in path.outputs %} { + {{ angular_contraction( + "angular_value", path, output, edge_channel + ) }} + {{ output.accumulator.name }} += + radial_weight * angular_value; + } {%- endfor %} + } {%- endfor %} + {%- endfor %} + } + {% for slot in schedule.output_slots %} { + const int64_t output_index = + output_base + {{ slot.array_index }} + + channel * {{ slot.irrep_dim }}; + + out[output_index] = {{ slot.name }}; + } {%- endfor %} +} + +// The same (receiver node, channel) ownership computes the forward tangent. +extern "C" __global__ void oeq_projected_forward_jvp( + int64_t node_count, + int64_t edge_count, + const scalar_t* x, + const scalar_t* sh, + const scalar_t* weights, + const int32_t* senders, + const int32_t* row_ptr, + const scalar_t* tx, + const scalar_t* tsh, + const scalar_t* tweights, + scalar_t* out) { + + const int64_t node_channel_index = + int64_t(blockIdx.x) * blockDim.x + threadIdx.x; + const int64_t node_channel_count = node_count * CHANNEL_DIM; + if (node_channel_index >= node_channel_count) return; + + const int64_t node = node_channel_index / CHANNEL_DIM; + const int channel = int(node_channel_index - node * CHANNEL_DIM); + const int64_t output_base = node * OUTPUT_DIM; + + {%- for slot in schedule.output_slots %} + scalar_t tangent_{{ slot.name }} = scalar_t(0); + {%- endfor %} + + const EdgeRange edge_range = + safe_edge_range(row_ptr, node, edge_count); + for (int64_t e = edge_range.begin; e < edge_range.end; ++e) { + const int32_t sender = senders[e]; + if (sender < 0 || sender >= node_count) continue; + + const int64_t input_base = int64_t(sender) * INPUT_DIM; + const int64_t edge_base = e * EDGE_DIM; + const int64_t weight_base = e * WEIGHT_DIM; + {% for path in schedule.paths %} + {% for edge_channel in range(path.edge_mul) %} { + const int weight_index = + {{ path.weight_start + edge_channel * schedule.channels }} + channel; + const int64_t radial_weight_index = weight_base + weight_index; + const scalar_t radial_weight = weights[radial_weight_index]; + const scalar_t tangent_weight = + {{ value_if_active(Input.W, "radial_weight_index", true) }}; + + {% for output in path.outputs %} { + {{ angular_contraction( + "angular_value", path, output, edge_channel + ) }} + {{ angular_contraction( + "tangent_angular_value", path, output, edge_channel, + tangent=true + ) }} + tangent_{{ output.accumulator.name }} += + tangent_weight * angular_value + + radial_weight * tangent_angular_value; + } {%- endfor %} + } {%- endfor %} + {%- endfor %} + } + {% for slot in schedule.output_slots %} { + const int64_t output_index = + output_base + {{ slot.array_index }} + + channel * {{ slot.irrep_dim }}; + + out[output_index] = tangent_{{ slot.name }}; + } {%- endfor %} +} + +// One logical thread group owns an edge, including on HIP wavefront-64. +// Each lane owns channels separated by LOGICAL_GROUP_SIZE. +extern "C" __global__ void oeq_projected_backward( + int64_t node_count, + int64_t edge_count, + const scalar_t* x, + const scalar_t* sh, + const scalar_t* weights, + const int32_t* senders, + const int32_t* receivers, + const scalar_t* dout, + scalar_t* dx, + scalar_t* dsh, + scalar_t* dweights) { + + const int lane = threadIdx.x % LOGICAL_GROUP_SIZE; + const int logical_group = threadIdx.x / LOGICAL_GROUP_SIZE; + const int64_t e = + int64_t(blockIdx.x) * (blockDim.x / LOGICAL_GROUP_SIZE) + logical_group; + if (e >= edge_count) return; + + const int32_t sender = senders[e]; + const int32_t receiver = receivers[e]; + if (sender < 0 || sender >= node_count + || receiver < 0 || receiver >= node_count) + return; + + const int64_t input_base = int64_t(sender) * INPUT_DIM; + const int64_t output_base = int64_t(receiver) * OUTPUT_DIM; + const int64_t edge_base = e * EDGE_DIM; + const int64_t weight_base = e * WEIGHT_DIM; + + acc_t edge_gradient[EDGE_DIM] = {acc_t(0)}; + for (int channel = lane; channel < CHANNEL_DIM; + channel += LOGICAL_GROUP_SIZE) { + {%- for slot in schedule.input_gradient_slots %} + acc_t {{ slot.name }} = acc_t(0); + {%- endfor %} + + // Path weight intervals are disjoint. This lane directly owns each + // dweights[e, weight] element generated below. + {% for path in schedule.paths %} + {% for edge_channel in range(path.edge_mul) %} { + const int weight_index = + {{ path.weight_start + edge_channel * schedule.channels }} + channel; + const int64_t radial_weight_index = weight_base + weight_index; + const scalar_t radial_weight = weights[radial_weight_index]; + acc_t weight_gradient = acc_t(0); + + {% for output in path.outputs %} { + const int64_t output_index = + output_base + {{ output.accumulator.array_index }} + + channel * {{ output.accumulator.irrep_dim }}; + const scalar_t output_gradient = dout[output_index]; + + // Accumulate weight, sender-feature, and edge-feature adjoints. + {% for term in output.terms %} { + const int64_t input_index = + input_base + {{ path.input_start + term.input_component }} + + channel * {{ path.input_irrep_dim }}; + const int edge_component_index = + {{ path.edge_start + term.edge_component }} + + {{ edge_channel }} * {{ path.edge_irrep_dim }}; + + const int64_t edge_index = edge_base + edge_component_index; + const scalar_t input_value = x[input_index]; + const scalar_t edge_value = sh[edge_index]; + + weight_gradient += output_gradient * input_value + * edge_value * {{ cpp_scalar_literal(term.coefficient, scalar) }}; + {{ term.input_accumulator.name }} += output_gradient + * radial_weight * edge_value * {{ cpp_scalar_literal(term.coefficient, scalar) }}; + edge_gradient[edge_component_index] += + output_gradient * radial_weight * input_value + * {{ cpp_scalar_literal(term.coefficient, scalar) }}; + } {%- endfor %} + } {%- endfor %} + dweights[radial_weight_index] = scalar_t(weight_gradient); + } {%- endfor %} + {%- endfor %} + + // Paths are combined locally for this (edge, channel, input component), + // but different receiver-owned edges can contribute to the same sender. + {%- for slot in schedule.input_gradient_slots %} + {{ atomic_add }}( + dx + input_base + {{ slot.array_index }} + + channel * {{ slot.irrep_dim }}, + scalar_t({{ slot.name }})); + {%- endfor %} + } + + // Reduce within the logical 32-thread edge owner. Lane 0 writes dsh[e, :]. + for (int offset = 16; offset > 0; offset >>= 1) + for (int component = 0; component < EDGE_DIM; ++component) + edge_gradient[component] += + {{ shfl_down_32("edge_gradient[component]", "offset") }}; + + if (lane == 0) + for (int component = 0; component < EDGE_DIM; ++component) + dsh[edge_base + component] = scalar_t(edge_gradient[component]); +} + +// Mixed derivative of the edge-owned reverse pass, with the same logical +// 32-thread group and channel ownership as oeq_projected_backward. +extern "C" __global__ void oeq_projected_backward_jvp( + int64_t node_count, + int64_t edge_count, + const scalar_t* x, + const scalar_t* sh, + const scalar_t* weights, + const int32_t* senders, + const int32_t* receivers, + const scalar_t* dout, + const scalar_t* tx, + const scalar_t* tsh, + const scalar_t* tweights, + const scalar_t* tdout, + scalar_t* tdx, + scalar_t* tdsh, + scalar_t* tdweights) { + + const int lane = threadIdx.x % LOGICAL_GROUP_SIZE; + const int logical_group = threadIdx.x / LOGICAL_GROUP_SIZE; + const int64_t e = + int64_t(blockIdx.x) * (blockDim.x / LOGICAL_GROUP_SIZE) + logical_group; + if (e >= edge_count) return; + + const int32_t sender = senders[e]; + const int32_t receiver = receivers[e]; + if (sender < 0 || sender >= node_count + || receiver < 0 || receiver >= node_count) + return; + + const int64_t input_base = int64_t(sender) * INPUT_DIM; + const int64_t output_base = int64_t(receiver) * OUTPUT_DIM; + const int64_t edge_base = e * EDGE_DIM; + const int64_t weight_base = e * WEIGHT_DIM; + + acc_t tangent_edge_gradient[EDGE_DIM] = {acc_t(0)}; + for (int channel = lane; channel < CHANNEL_DIM; + channel += LOGICAL_GROUP_SIZE) { + {%- for slot in schedule.input_gradient_slots %} + acc_t {{ slot.name }} = acc_t(0); + {%- endfor %} + + // Path weight intervals are disjoint. This lane directly owns each + // tdweights[e, weight] element generated below. + {% for path in schedule.paths %} + {% for edge_channel in range(path.edge_mul) %} { + const int weight_index = + {{ path.weight_start + edge_channel * schedule.channels }} + channel; + const int64_t radial_weight_index = weight_base + weight_index; + const scalar_t radial_weight = weights[radial_weight_index]; + const scalar_t tangent_weight = + {{ value_if_active(Input.W, "radial_weight_index", false) }}; + + acc_t tangent_weight_gradient = acc_t(0); + {% for output in path.outputs %} { + const int64_t output_index = + output_base + {{ output.accumulator.array_index }} + + channel * {{ output.accumulator.irrep_dim }}; + const scalar_t output_gradient = dout[output_index]; + const scalar_t tangent_output_gradient = + {{ value_or_zero( + backward_jvp_dout_active, "tdout[output_index]") }}; + + {% for term in output.terms %} { + const int64_t input_index = + input_base + {{ path.input_start + term.input_component }} + + channel * {{ path.input_irrep_dim }}; + const int edge_component_index = + {{ path.edge_start + term.edge_component }} + + {{ edge_channel }} * {{ path.edge_irrep_dim }}; + const int64_t edge_index = edge_base + edge_component_index; + const scalar_t input_value = x[input_index]; + const scalar_t edge_value = sh[edge_index]; + const scalar_t tangent_input_value = + {{ value_if_active(Input.X, "input_index", false) }}; + const scalar_t tangent_edge_value = + {{ value_if_active(Input.SH, "edge_index", false) }}; + + {{ term.input_accumulator.name }} += + (tangent_output_gradient * radial_weight * edge_value + + output_gradient * tangent_weight * edge_value + + output_gradient * radial_weight * tangent_edge_value) + * {{ cpp_scalar_literal(term.coefficient, scalar) }}; + tangent_edge_gradient[edge_component_index] += + (tangent_output_gradient * radial_weight * input_value + + output_gradient * tangent_weight * input_value + + output_gradient * radial_weight * tangent_input_value) + * {{ cpp_scalar_literal(term.coefficient, scalar) }}; + tangent_weight_gradient += + (tangent_output_gradient * input_value * edge_value + + output_gradient * tangent_input_value * edge_value + + output_gradient * input_value * tangent_edge_value) + * {{ cpp_scalar_literal(term.coefficient, scalar) }}; + + } {%- endfor %} + } {%- endfor %} + + tdweights[radial_weight_index] = + scalar_t(tangent_weight_gradient); + + } {%- endfor %} + {%- endfor %} + + // Paths are combined locally for this (edge, channel, input component), + // but different receiver-owned edges can contribute to the same sender. + {%- for slot in schedule.input_gradient_slots %} + {{ atomic_add }}( + tdx + input_base + {{ slot.array_index }} + + channel * {{ slot.irrep_dim }}, + scalar_t({{ slot.name }})); + {%- endfor %} + } + + // Reduce within the logical 32-thread edge owner. Lane 0 writes tdsh[e, :]. + for (int offset = 16; offset > 0; offset >>= 1) + for (int component = 0; component < EDGE_DIM; ++component) + tangent_edge_gradient[component] += + {{ shfl_down_32("tangent_edge_gradient[component]", "offset") }}; + + if (lane == 0) + for (int component = 0; component < EDGE_DIM; ++component) + tdsh[edge_base + component] = + scalar_t(tangent_edge_gradient[component]); +} diff --git a/openequivariance/openequivariance/templates/jinja_utils.py b/openequivariance/openequivariance/templates/jinja_utils.py index b4ede597..c8b10f98 100644 --- a/openequivariance/openequivariance/templates/jinja_utils.py +++ b/openequivariance/openequivariance/templates/jinja_utils.py @@ -1,3 +1,4 @@ +import numpy as np from jinja2 import Environment, PackageLoader @@ -16,6 +17,24 @@ def sizeof(dtype): raise Exception("Provided undefined datatype to sizeof!") +def cpp_scalar_type(dtype): + """Return the C++ scalar spelling for a supported generated dtype.""" + dtype = np.dtype(dtype) + if dtype == np.dtype(np.float32): + return "float" + if dtype == np.dtype(np.float64): + return "double" + raise ValueError("generated kernels support f32 and f64 scalar types") + + +def cpp_scalar_literal(value, scalar): + """Return an exact C++ literal for a generated ``float`` or ``double``.""" + if scalar not in ("float", "double"): + raise ValueError(f"Unsupported generated scalar type: {scalar}") + suffix = "f" if scalar == "float" else "" + return f"static_cast({float(value).hex()}{suffix})" + + def get_jinja_environment(is_hip=False): env = Environment( loader=PackageLoader("openequivariance"), extensions=["jinja2.ext.do"] @@ -24,6 +43,7 @@ def get_jinja_environment(is_hip=False): env.globals["divide"] = divide env.globals["sizeof"] = sizeof env.globals["enumerate"] = enumerate + env.globals["cpp_scalar_literal"] = cpp_scalar_literal env.globals["is_hip"] = is_hip env.globals["syncwarp"] = ( @@ -35,8 +55,12 @@ def get_jinja_environment(is_hip=False): if is_hip: env.globals["shfl_down"] = lambda val, offset: f"__shfl_down( {val}, {offset})" + env.globals["shfl_down_32"] = lambda val, offset: ( + f"__shfl_down( {val}, {offset}, 32)" + ) else: - env.globals["shfl_down"] = ( - lambda val, offset: f"__shfl_down_sync(FULL_MASK, {val}, {offset})" + env.globals["shfl_down"] = lambda val, offset: ( + f"__shfl_down_sync(FULL_MASK, {val}, {offset})" ) + env.globals["shfl_down_32"] = env.globals["shfl_down"] return env diff --git a/openequivariance_extjax/CMakeLists.txt b/openequivariance_extjax/CMakeLists.txt index 52b38c5a..3ec96892 100644 --- a/openequivariance_extjax/CMakeLists.txt +++ b/openequivariance_extjax/CMakeLists.txt @@ -59,6 +59,7 @@ set(OEQ_JAX_SOURCES set(OEQ_JAX_HEADERS src/ffi_handler_table.h ${HEADER_DIR}/convolution.hpp + ${HEADER_DIR}/factorized_projected.hpp ${HEADER_DIR}/tensorproducts.hpp ${HEADER_DIR}/backend/backend_cuda.hpp ${HEADER_DIR}/backend/backend_hip.hpp diff --git a/openequivariance_extjax/src/ffi_handlers.cpp b/openequivariance_extjax/src/ffi_handlers.cpp index fb3a702e..eb60cfa8 100644 --- a/openequivariance_extjax/src/ffi_handlers.cpp +++ b/openequivariance_extjax/src/ffi_handlers.cpp @@ -32,6 +32,7 @@ using json = json11::Json; #include "tensorproducts.hpp" #include "convolution.hpp" +#include "factorized_projected.hpp" xla::ffi::DataType enum_to_xla_dtype(int64_t i){ switch(i) { @@ -85,6 +86,10 @@ inline void* data_ptr(ffi::AnyBuffer &buffer) { return buffer.untyped_data(); } +inline void* data_ptr(const ffi::AnyBuffer &buffer) { + return const_cast(buffer.untyped_data()); +} + inline void* data_ptr(ffi::Result &buffer) { return data_ptr(*buffer); } @@ -174,6 +179,9 @@ std::unordered_map>, KernelProp >> conv_cache; + +std::unordered_map>> + factorized_projected_cache; std::mutex mut; template @@ -245,6 +253,19 @@ std::pair*, KernelProp> return {cached.first.get(), cached.second}; } +JITFactorizedProjectedImpl* compile_factorized_projected_with_caching( + std::string_view source, int64_t hash, int64_t num_threads, + int64_t logical_cohort_width, int64_t shared_memory_bytes) { + auto& cached = find_or_compile_cached( + factorized_projected_cache, hash, [&] { + return std::make_unique>( + std::string(source), num_threads, logical_cohort_width, + shared_memory_bytes); + }); + return cached.get(); +} + + inline void check_tensor(const ffi::AnyBuffer &buffer, std::initializer_list expected_shape, xla::ffi::DataType expected_dtype, @@ -741,6 +762,307 @@ XLA_FFI_DEFINE_HANDLER_SYMBOL( .Attr("hash"), {xla::ffi::Traits::kCmdBufferCompatible}); +// ------------------- Generated factorized convolution ------------------- +void validate_projected_inputs(ffi::AnyBuffer &x, ffi::AnyBuffer &sh, + ffi::AnyBuffer &senders, int64_t input_dim, + int64_t edge_dim) { + if (x.dimensions().size() != 2 || sh.dimensions().size() != 2) { + throw std::logic_error("projected factorized inputs must have rank two"); + } + if (x.element_type() != xla::ffi::DataType::F32 && + x.element_type() != xla::ffi::DataType::F64) { + throw std::logic_error("projected factorized kernels support only f32 and f64"); + } + const int64_t edge_count = sh.dimensions()[0]; + check_tensor(x, {x.dimensions()[0], input_dim}, x.element_type(), "x"); + check_tensor(sh, {edge_count, edge_dim}, x.element_type(), "sh"); + check_tensor(senders, {edge_count}, xla::ffi::DataType::S32, "senders"); +} + +ffi::Error factorized_projected_forward_impl( + ffi::AnyBuffer x, ffi::AnyBuffer sh, ffi::AnyBuffer weights, ffi::AnyBuffer senders, + ffi::AnyBuffer row_ptr, ffi::Result out, stream_t stream, + std::string_view source, int64_t hash, int64_t channels, + int64_t input_dim, int64_t edge_dim, int64_t weight_dim, int64_t output_dim, + int64_t num_threads, int64_t logical_cohort_width, + int64_t shared_memory_bytes) { + validate_projected_inputs(x, sh, senders, input_dim, edge_dim); + const int64_t node_count = x.dimensions()[0], edge_count = sh.dimensions()[0]; + check_tensor(weights, {edge_count, weight_dim}, x.element_type(), "weights"); + check_tensor(row_ptr, {node_count + 1}, xla::ffi::DataType::S32, "row_ptr"); + check_tensor(*out, {node_count, output_dim}, x.element_type(), "out"); + auto* jit_kernel = compile_factorized_projected_with_caching( + source, hash, num_threads, logical_cohort_width, shared_memory_bytes); + jit_kernel->forward( + node_count, edge_count, channels, data_ptr(x), data_ptr(sh), data_ptr(weights), + data_ptr(senders), data_ptr(row_ptr), data_ptr(out), stream); + return ffi::Error::Success(); +} + +ffi::Error factorized_projected_forward_jvp_impl( + ffi::AnyBuffer x, ffi::AnyBuffer sh, ffi::AnyBuffer weights, ffi::AnyBuffer senders, + ffi::AnyBuffer row_ptr, const ffi::AnyBuffer* tx, const ffi::AnyBuffer* tsh, + const ffi::AnyBuffer* tweights, ffi::Result out, stream_t stream, + std::string_view source, int64_t hash, int64_t channels, + int64_t input_dim, int64_t edge_dim, int64_t weight_dim, int64_t output_dim, + int64_t num_threads, int64_t logical_cohort_width, + int64_t shared_memory_bytes) { + validate_projected_inputs(x, sh, senders, input_dim, edge_dim); + const int64_t node_count = x.dimensions()[0], edge_count = sh.dimensions()[0]; + check_tensor(weights, {edge_count, weight_dim}, x.element_type(), "weights"); + check_tensor(row_ptr, {node_count + 1}, xla::ffi::DataType::S32, "row_ptr"); + if (tx != nullptr) check_tensor(*tx, {node_count, input_dim}, x.element_type(), "tx"); + if (tsh != nullptr) check_tensor(*tsh, {edge_count, edge_dim}, x.element_type(), "tsh"); + if (tweights != nullptr) { + check_tensor(*tweights, {edge_count, weight_dim}, x.element_type(), "tweights"); + } + check_tensor(*out, {node_count, output_dim}, x.element_type(), "out"); + auto* jit_kernel = compile_factorized_projected_with_caching( + source, hash, num_threads, logical_cohort_width, shared_memory_bytes); + jit_kernel->forward_jvp( + node_count, edge_count, channels, data_ptr(x), data_ptr(sh), data_ptr(weights), + data_ptr(senders), data_ptr(row_ptr), tx == nullptr ? nullptr : data_ptr(*tx), + tsh == nullptr ? nullptr : data_ptr(*tsh), + tweights == nullptr ? nullptr : data_ptr(*tweights), data_ptr(out), stream); + return ffi::Error::Success(); +} + +ffi::Error factorized_projected_backward_impl( + ffi::AnyBuffer x, ffi::AnyBuffer sh, ffi::AnyBuffer weights, ffi::AnyBuffer senders, + ffi::AnyBuffer receivers, ffi::AnyBuffer dout, ffi::Result dx, + ffi::Result dsh, ffi::Result dweights, + stream_t stream, std::string_view source, int64_t hash, + int64_t channels, int64_t input_dim, int64_t edge_dim, int64_t weight_dim, + int64_t output_dim, int64_t num_threads, int64_t logical_cohort_width, + int64_t shared_memory_bytes) { + validate_projected_inputs(x, sh, senders, input_dim, edge_dim); + const int64_t node_count = x.dimensions()[0], edge_count = sh.dimensions()[0]; + check_tensor(weights, {edge_count, weight_dim}, x.element_type(), "weights"); + check_tensor(receivers, {edge_count}, xla::ffi::DataType::S32, "receivers"); + check_tensor(dout, {node_count, output_dim}, x.element_type(), "dout"); + check_tensor(*dx, {node_count, input_dim}, x.element_type(), "dx"); + check_tensor(*dsh, {edge_count, edge_dim}, x.element_type(), "dsh"); + check_tensor(*dweights, {edge_count, weight_dim}, x.element_type(), "dweights"); + zero_buffer(*dx, stream); + zero_buffer(*dsh, stream); + zero_buffer(*dweights, stream); + auto* jit_kernel = compile_factorized_projected_with_caching( + source, hash, num_threads, logical_cohort_width, shared_memory_bytes); + jit_kernel->backward( + node_count, edge_count, data_ptr(x), data_ptr(sh), data_ptr(weights), data_ptr(senders), + data_ptr(receivers), data_ptr(dout), data_ptr(dx), data_ptr(dsh), + data_ptr(dweights), stream); + return ffi::Error::Success(); +} + +ffi::Error factorized_projected_backward_jvp_impl( + ffi::AnyBuffer x, ffi::AnyBuffer sh, ffi::AnyBuffer weights, ffi::AnyBuffer senders, + ffi::AnyBuffer receivers, ffi::AnyBuffer dout, const ffi::AnyBuffer* tx, + const ffi::AnyBuffer* tsh, const ffi::AnyBuffer* tweights, + const ffi::AnyBuffer* tdout, + ffi::Result tdx, ffi::Result tdsh, + ffi::Result tdweights, stream_t stream, + std::string_view source, int64_t hash, int64_t channels, int64_t input_dim, + int64_t edge_dim, int64_t weight_dim, int64_t output_dim, + int64_t num_threads, int64_t logical_cohort_width, + int64_t shared_memory_bytes) { + validate_projected_inputs(x, sh, senders, input_dim, edge_dim); + const int64_t node_count = x.dimensions()[0], edge_count = sh.dimensions()[0]; + check_tensor(weights, {edge_count, weight_dim}, x.element_type(), "weights"); + check_tensor(receivers, {edge_count}, xla::ffi::DataType::S32, "receivers"); + check_tensor(dout, {node_count, output_dim}, x.element_type(), "dout"); + if (tx != nullptr) check_tensor(*tx, {node_count, input_dim}, x.element_type(), "tx"); + if (tsh != nullptr) check_tensor(*tsh, {edge_count, edge_dim}, x.element_type(), "tsh"); + if (tweights != nullptr) { + check_tensor(*tweights, {edge_count, weight_dim}, x.element_type(), "tweights"); + } + if (tdout != nullptr) { + check_tensor(*tdout, {node_count, output_dim}, x.element_type(), "tdout"); + } + check_tensor(*tdx, {node_count, input_dim}, x.element_type(), "tdx"); + check_tensor(*tdsh, {edge_count, edge_dim}, x.element_type(), "tdsh"); + check_tensor(*tdweights, {edge_count, weight_dim}, x.element_type(), "tdweights"); + zero_buffer(*tdx, stream); + zero_buffer(*tdsh, stream); + zero_buffer(*tdweights, stream); + auto* jit_kernel = compile_factorized_projected_with_caching( + source, hash, num_threads, logical_cohort_width, shared_memory_bytes); + jit_kernel->backward_jvp( + node_count, edge_count, data_ptr(x), data_ptr(sh), data_ptr(weights), data_ptr(senders), + data_ptr(receivers), data_ptr(dout), tx == nullptr ? nullptr : data_ptr(*tx), + tsh == nullptr ? nullptr : data_ptr(*tsh), + tweights == nullptr ? nullptr : data_ptr(*tweights), + tdout == nullptr ? nullptr : data_ptr(*tdout), data_ptr(tdx), data_ptr(tdsh), + data_ptr(tdweights), stream); + return ffi::Error::Success(); +} + + +struct GeneratedBuffers { + std::vector args; + std::vector> rets; +}; + +ffi::ErrorOr decode_generated_buffers( + ffi::RemainingArgs args, ffi::RemainingRets rets, size_t expected_args, + size_t expected_rets, std::string_view family, int64_t operation) { + if (args.size() != expected_args || rets.size() != expected_rets) { + return ffi::Unexpected(ffi::Error::InvalidArgument( + std::string(family) + " operation " + std::to_string(operation) + + " received an unexpected number of buffers")); + } + GeneratedBuffers buffers; + buffers.args.reserve(expected_args); + buffers.rets.reserve(expected_rets); + for (size_t index = 0; index < expected_args; ++index) { + auto value = args.get(index); + if (!value) return ffi::Unexpected(value.error()); + buffers.args.push_back(*value); + } + for (size_t index = 0; index < expected_rets; ++index) { + auto value = rets.get(index); + if (!value) return ffi::Unexpected(value.error()); + buffers.rets.push_back(*value); + } + return buffers; +} + +ffi::Error factorized_projected_execute_impl( + ffi::RemainingArgs args, ffi::RemainingRets rets, stream_t stream, + std::string_view source, + int64_t hash, int64_t operation, int64_t channels, int64_t input_dim, + int64_t edge_dim, int64_t weight_dim, int64_t output_dim, + int64_t num_threads, int64_t logical_cohort_width, + int64_t shared_memory_bytes) { + constexpr std::string_view kFamily = "factorized_projected"; + const auto active_count = [](int64_t mask) { + size_t count = 0; + for (; mask != 0; mask >>= 1) count += static_cast(mask & 1); + return count; + }; + if (operation >= 17 && operation <= 23) { + const int64_t mask = operation - 16; + auto buffers = decode_generated_buffers( + args, rets, 5 + active_count(mask), 1, kFamily, operation); + if (!buffers) return buffers.error(); + auto& a = buffers->args; + size_t index = 5; + const auto next = [&](int64_t bit) -> const ffi::AnyBuffer* { + return mask & bit ? &a[index++] : nullptr; + }; + const auto* tx = next(1); + const auto* tsh = next(2); + const auto* tweights = next(4); + return factorized_projected_forward_jvp_impl( + a[0], a[1], a[2], a[3], a[4], tx, tsh, tweights, + buffers->rets[0], stream, source, hash, channels, + input_dim, edge_dim, weight_dim, output_dim, num_threads, + logical_cohort_width, shared_memory_bytes); + } + if (operation >= 33 && operation <= 47) { + const int64_t mask = operation - 32; + auto buffers = decode_generated_buffers( + args, rets, 6 + active_count(mask), 3, kFamily, operation); + if (!buffers) return buffers.error(); + auto& a = buffers->args; + size_t index = 6; + const auto next = [&](int64_t bit) -> const ffi::AnyBuffer* { + return mask & bit ? &a[index++] : nullptr; + }; + const auto* tx = next(1); + const auto* tsh = next(2); + const auto* tweights = next(4); + const auto* tdout = next(8); + return factorized_projected_backward_jvp_impl( + a[0], a[1], a[2], a[3], a[4], a[5], tx, tsh, tweights, tdout, + buffers->rets[0], buffers->rets[1], buffers->rets[2], stream, + source, hash, channels, input_dim, edge_dim, + weight_dim, output_dim, num_threads, logical_cohort_width, + shared_memory_bytes); + } + size_t expected_args; + size_t expected_rets; + switch (operation) { + case 0: + expected_args = 5; + expected_rets = 1; + break; + case 2: + expected_args = 6; + expected_rets = 3; + break; + default: + return ffi::Error::InvalidArgument("unknown factorized projected operation"); + } + auto buffers = decode_generated_buffers( + args, rets, expected_args, expected_rets, kFamily, operation); + if (!buffers) return buffers.error(); + auto& a = buffers->args; + auto& r = buffers->rets; + switch (operation) { + case 0: + return factorized_projected_forward_impl( + a[0], a[1], a[2], a[3], a[4], r[0], stream, source, hash, + channels, input_dim, edge_dim, weight_dim, output_dim, + num_threads, logical_cohort_width, shared_memory_bytes); + case 2: + return factorized_projected_backward_impl( + a[0], a[1], a[2], a[3], a[4], a[5], r[0], r[1], r[2], stream, source, hash, channels, input_dim, edge_dim, weight_dim, + output_dim, num_threads, logical_cohort_width, + shared_memory_bytes); + } + return ffi::Error::Internal("unreachable factorized projected operation"); +} + + +ffi::Error factorized_projected_initialize_impl( + ffi::RemainingArgs, ffi::RemainingRets, stream_t, std::string_view source, + int64_t hash, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, + int64_t num_threads, int64_t logical_cohort_width, + int64_t shared_memory_bytes) { + compile_factorized_projected_with_caching( + source, hash, num_threads, logical_cohort_width, shared_memory_bytes); + return ffi::Error::Success(); +} + +#define OEQ_GENERATED_INITIALIZE_CONTEXTS \ + .RemainingArgs() \ + .RemainingRets() \ + .Ctx>() + +#define OEQ_GENERATED_ATTRIBUTES \ + .Attr("source") \ + .Attr("hash") \ + .Attr("operation") + +#define OEQ_PROJECTED_ATTRIBUTES \ + OEQ_GENERATED_ATTRIBUTES \ + .Attr("channels") \ + .Attr("input_dim") \ + .Attr("edge_dim") \ + .Attr("weight_dim") \ + .Attr("output_dim") \ + .Attr("num_threads") \ + .Attr("logical_cohort_width") \ + .Attr("shared_memory_bytes") + +XLA_FFI_DEFINE_HANDLER_SYMBOL( + factorized_projected_initialize, factorized_projected_initialize_impl, + ffi::Ffi::Bind() + OEQ_GENERATED_INITIALIZE_CONTEXTS OEQ_PROJECTED_ATTRIBUTES); + +XLA_FFI_DEFINE_HANDLER_SYMBOL( + factorized_projected, factorized_projected_execute_impl, + ffi::Ffi::Bind() + .RemainingArgs() + .RemainingRets() + .Ctx>() OEQ_PROJECTED_ATTRIBUTES, + {xla::ffi::Traits::kCmdBufferCompatible}); + +#undef OEQ_PROJECTED_ATTRIBUTES +#undef OEQ_GENERATED_ATTRIBUTES +#undef OEQ_GENERATED_INITIALIZE_CONTEXTS + namespace { #define OEQ_FFI_HANDLER(NAME, INITIALIZE) \ @@ -754,6 +1076,7 @@ const OeqFfiHandler kFfiHandlers[] = { OEQ_FFI_HANDLER(conv_forward, conv_initialize), OEQ_FFI_HANDLER(conv_backward, conv_initialize), OEQ_FFI_HANDLER(conv_double_backward, conv_initialize), + OEQ_FFI_HANDLER(factorized_projected, factorized_projected_initialize), }; #undef OEQ_FFI_HANDLER diff --git a/tests/factorized_computation_schedule_test.py b/tests/factorized_computation_schedule_test.py new file mode 100644 index 00000000..5eeaf1ec --- /dev/null +++ b/tests/factorized_computation_schedule_test.py @@ -0,0 +1,120 @@ +"""Tests for static receiver-streaming computation schedules.""" + +from dataclasses import replace + +import numpy as np +import pytest + +from openequivariance.core.FactorizedComputationSchedule import ( + FactorizedLaunchConfig, + Input, + factorized_schedule_from_problem, +) +from openequivariance.core.e3nn_lite import Irreps, TPProblem, wigner_3j + + +def _problem(mode="uvu", edge_mul=1, channels=4): + return TPProblem( + Irreps(f"{channels}x1e"), + Irreps(f"{edge_mul}x1e"), + Irreps(f"{channels}x1e"), + [(0, 0, 0, mode, True)], + shared_weights=False, + internal_weights=False, + irrep_dtype=np.float64, + weight_dtype=np.float64, + ) + + +def test_factorized_schedule_preserves_layout_couplings_and_geometry(): + """Lowering resolves the same normalized CG entries and launch geometry.""" + problem = _problem() + schedule = factorized_schedule_from_problem(problem) + path = schedule.paths[0] + assert (schedule.input_dim, schedule.edge_dim, schedule.output_dim) == ( + problem.irreps_in1.dim, + problem.irreps_in2.dim, + problem.irreps_out.dim, + ) + assert schedule.weight_numel == problem.weight_numel + assert schedule.launch_config.num_threads == 128 + assert schedule.launch_config.logical_cohort_width == 32 + assert schedule.launch_config.shared_memory_bytes == 0 + assert (Input.X.value, Input.SH.value, Input.W.value) == (0, 1, 2) + + output_irrep_dim = problem.irreps_out[0].ir.dim + actual = np.zeros((path.input_irrep_dim, path.edge_irrep_dim, output_irrep_dim)) + for output in path.outputs: + for term in output.terms: + actual[ + term.input_component, term.edge_component, output.output_component + ] = term.coefficient + np.testing.assert_array_equal( + actual, wigner_3j(1, 1, 1) * problem.instructions[0].path_weight + ) + assert len(schedule.output_slots) == output_irrep_dim + assert len(schedule.input_gradient_slots) == path.input_irrep_dim + assert sum(len(output.terms) for output in path.outputs) == np.count_nonzero(actual) + + +def test_factorized_schedule_caches_complete_ffi_specializations(): + """Create source and FFI metadata together once per specialization.""" + schedule = factorized_schedule_from_problem(_problem()) + kernel = schedule.kernel(np.float32, is_hip=False) + assert kernel is schedule.kernel(np.float32, is_hip=False) + assert kernel.ffi_attributes["source"] == kernel.jit_kernel + assert kernel.ffi_attributes["hash"] == kernel.hash + assert kernel.ffi_attributes["input_dim"] == schedule.input_dim + assert kernel.ffi_attributes["edge_dim"] == schedule.edge_dim + assert kernel.ffi_attributes["weight_dim"] == schedule.weight_numel + assert kernel.ffi_attributes["output_dim"] == schedule.output_dim + assert kernel.ffi_attributes["num_threads"] == schedule.launch_config.num_threads + assert ( + kernel.ffi_attributes["logical_cohort_width"] + == schedule.launch_config.logical_cohort_width + ) + assert ( + kernel.ffi_attributes["shared_memory_bytes"] + == schedule.launch_config.shared_memory_bytes + ) + + tuned_schedule = replace( + schedule, launch_config=FactorizedLaunchConfig(num_threads=256) + ) + tuned_kernel = tuned_schedule.kernel(np.float32, is_hip=False) + assert tuned_kernel.jit_kernel == kernel.jit_kernel + assert tuned_kernel.hash != kernel.hash + + +def test_factorized_schedule_reduces_only_exact_scalar_uvw(): + """Only UVW paths with one input and output multiplicity reduce to UVU.""" + reducible = TPProblem( + Irreps("1x1e"), + Irreps("2x0e"), + Irreps("1x1e"), + [(0, 0, 0, "uvw", True)], + shared_weights=False, + internal_weights=False, + ) + assert ( + factorized_schedule_from_problem(reducible).weight_numel + == reducible.weight_numel + ) + with pytest.raises(ValueError, match="one input and one output"): + factorized_schedule_from_problem(_problem(mode="uvw")) + + +def test_factorized_schedule_rejects_unreferenced_output_block(): + """Reject an output irrep block with no producing instruction.""" + problem = TPProblem( + Irreps("4x1e"), + Irreps("1x0e"), + Irreps("4x1e + 4x0e"), + [(0, 0, 0, "uvu", True)], + shared_weights=False, + internal_weights=False, + irrep_dtype=np.float64, + weight_dtype=np.float64, + ) + with pytest.raises(ValueError, match="output irrep blocks \\[1\\]"): + factorized_schedule_from_problem(problem) diff --git a/tests/jax_ffi_abi_test.py b/tests/jax_ffi_abi_test.py index e483457a..f8e72d8f 100644 --- a/tests/jax_ffi_abi_test.py +++ b/tests/jax_ffi_abi_test.py @@ -64,7 +64,7 @@ def test_exported_handler_table_matches_manifest(with_jax): ] assert table.abi_version == 2 assert tuple(names) == FFI_TARGETS - assert len(names) == len(set(names)) == 6 + assert len(names) == len(set(names)) == 7 assert all( not table.handlers[index].instantiate for index in range(table.handler_count) ) @@ -87,7 +87,7 @@ def test_nanobind_registrations_match_handler_table(with_jax): ext, table = _handler_table() registrations = ext.registrations() assert tuple(registrations) == FFI_TARGETS - assert len(registrations) == table.handler_count == 6 + assert len(registrations) == table.handler_count == 7 for index, name in enumerate(FFI_TARGETS): handler = table.handlers[index] registration = registrations[name] diff --git a/tests/jax_tensor_product_conv_dispatch_test.py b/tests/jax_tensor_product_conv_dispatch_test.py new file mode 100644 index 00000000..830cb88b --- /dev/null +++ b/tests/jax_tensor_product_conv_dispatch_test.py @@ -0,0 +1,524 @@ +"""Numerical coverage for public receiver-streaming convolution selection.""" + +import numpy as np +import pytest + + +def test_streaming_mode_requires_an_eligible_problem(): + """Reject irreducible UVW and aggregation options in streaming mode.""" + from openequivariance.core.e3nn_lite import TPProblem + from openequivariance.jax import ( + StreamingUnavailableError, + TensorProductConv, + streaming_support, + ) + + irreducible = TPProblem( + "2x0e", + "1x0e", + "2x0e", + [(0, 0, 0, "uvw", True)], + shared_weights=False, + internal_weights=False, + ) + supported, reason = streaming_support(irreducible) + assert not supported + assert "UVW" in reason + with pytest.raises(StreamingUnavailableError, match="mode='streaming'"): + TensorProductConv(irreducible, mode="streaming") + + eligible = TPProblem( + "1x0e", + "2x0e", + "1x0e", + [(0, 0, 0, "uvw", True)], + shared_weights=False, + internal_weights=False, + ) + with pytest.raises(StreamingUnavailableError, match="deterministic"): + TensorProductConv(eligible, deterministic=True, mode="streaming") + + mixed_dtype = TPProblem( + "1x0e", + "1x0e", + "1x0e", + [(0, 0, 0, "uvu", True)], + shared_weights=False, + internal_weights=False, + irrep_dtype=np.float32, + weight_dtype=np.float64, + ) + supported, reason = streaming_support(mixed_dtype) + assert not supported + assert "matching irrep and weight dtypes" in reason + with pytest.raises(StreamingUnavailableError, match="matching irrep"): + TensorProductConv(mixed_dtype, mode="streaming") + + +def test_standard_sorted_indices_keep_endpoints_and_sender_permutation(): + """Keep valid sorted endpoints and the deterministic sender permutation.""" + import jax.numpy as jnp + + from openequivariance.jax import ( + LoopUnrollTensorProductConv, + TensorProductConv, + ) + from openequivariance.jax.StreamingTensorProductConv import ( + StreamingTensorProductConv, + ) + + class FakeLoopUnroll(LoopUnrollTensorProductConv): + """Provide loop-unroll aggregation without requiring CUDA in this unit test.""" + + def __init__(self): + """Avoid constructing a CUDA kernel for this selector unit test.""" + self.last_sender_perm = None + + def forward(self, X, Y, W, rows, cols, sender_perm=None, **kwargs): + del kwargs + self.last_sender_perm = sender_perm + values = X[cols, :1] * Y[:, :1] * W[:, :1] + return jnp.zeros((X.shape[0], 1), dtype=X.dtype).at[rows].add(values) + + selector = object.__new__(TensorProductConv) + loop = FakeLoopUnroll() + selector._impl = loop + selector.deterministic = False + selector.kahan = False + X = jnp.asarray([[2.0], [3.0], [5.0]]) + Y = jnp.asarray([[7.0], [11.0], [13.0]]) + W = jnp.asarray([[17.0], [19.0], [23.0]]) + rows = jnp.asarray([1, 2, 0], dtype=jnp.int32) + cols = jnp.asarray([1, 0, 2], dtype=jnp.int32) + topology = StreamingTensorProductConv._prepare_topology(rows, cols, X.shape[0]) + + ordinary = jnp.asarray( + [[5.0 * 13.0 * 23.0], [3.0 * 7.0 * 17.0], [2.0 * 11.0 * 19.0]] + ) + sorted_output = selector.forward( + X, + Y[topology.order], + W[topology.order], + topology.receivers, + topology.senders, + indices_are_sorted=True, + ) + np.testing.assert_allclose(sorted_output, ordinary) + np.testing.assert_allclose(selector.forward(X, Y, W, rows, cols), ordinary) + + selector.deterministic = True + np.testing.assert_allclose( + selector.forward( + X, + Y[topology.order], + W[topology.order], + topology.receivers, + topology.senders, + sender_perm=jnp.arange(Y.shape[0], dtype=jnp.int32), + indices_are_sorted=True, + ), + ordinary, + ) + np.testing.assert_array_equal( + loop.last_sender_perm, jnp.arange(Y.shape[0], dtype=jnp.int32) + ) + + np.testing.assert_allclose( + selector.forward( + X, + Y[topology.order], + W[topology.order], + topology.receivers, + topology.senders, + indices_are_sorted=True, + row_ptr=topology.row_ptr, + ), + ordinary, + ) + + +def test_streaming_row_ptr_requires_expected_shape_and_dtype(): + """Reject cached receiver offsets with incompatible metadata.""" + import jax.numpy as jnp + + from openequivariance.jax.StreamingTensorProductConv import ( + StreamingTensorProductConv, + ) + + rows = jnp.asarray([0, 0, 2], dtype=jnp.int32) + cols = jnp.asarray([1, 2, 0], dtype=jnp.int32) + correct = jnp.asarray([0, 2, 2, 3], dtype=jnp.int32) + StreamingTensorProductConv._validate_row_ptr(correct, rows, cols, 3) + with pytest.raises(ValueError, match=r"shape \[N \+ 1\]"): + StreamingTensorProductConv._validate_row_ptr(correct[:-1], rows, cols, 3) + with pytest.raises(ValueError, match="int32"): + StreamingTensorProductConv._validate_row_ptr( + correct.astype(jnp.float32), rows, cols, 3 + ) + + +@pytest.fixture(scope="module") +def gpu_context(with_jax): + """Provide JAX only when the explicitly requested GPU backend is ready.""" + if not with_jax: + pytest.skip("requires --jax") + import jax + import jax.numpy as jnp + + if jax.default_backend() != "gpu": + pytest.skip("requires a JAX GPU backend") + return jax, jnp + + +def _scalar_case(gpu_context, mode): + """Create a V=2 scalar-output case and its standard OEQ reference.""" + jax, jnp = gpu_context + from openequivariance.core.e3nn_lite import TPProblem + from openequivariance.jax.TensorProductConv import TensorProductConv + + problem = TPProblem( + "1x1e", + "2x1e", + "1x1e", + [(0, 0, 0, mode, True)], + shared_weights=False, + internal_weights=False, + ) + generated = TensorProductConv(problem, mode="auto") + assert generated.uses_streaming_kernel + reference = TensorProductConv(problem, requires_jvp=True, mode="standard") + keys = jax.random.split(jax.random.key(728), 7) + dtype = problem.irrep_dtype + x = jax.random.normal(keys[0], (5, problem.irreps_in1.dim), dtype=dtype) + y = jax.random.normal(keys[1], (8, problem.irreps_in2.dim), dtype=dtype) + w = jax.random.normal(keys[2], (8, problem.weight_numel), dtype=dtype) + tangents = tuple( + jax.random.normal(key, operand.shape, dtype=dtype) + for key, operand in zip(keys[3:6], (x, y, w), strict=True) + ) + return ( + generated, + reference, + (x, y, w), + tangents, + jax.random.normal(keys[6], (5, problem.irreps_out.dim), dtype=dtype), + jnp.array([3, 0, 4, 1, 0, 2, 3, 1], jnp.int32), + jnp.array([0, 2, 1, 3, 4, 0, 2, 4], jnp.int32), + ) + + +def _assert_tree_close(actual, expected, *, atol=2e-3, rtol=2e-3): + """Compare matching JAX array tuples using float32-safe tolerances.""" + actual = actual if isinstance(actual, tuple) else (actual,) + expected = expected if isinstance(expected, tuple) else (expected,) + for got, want in zip(actual, expected, strict=True): + np.testing.assert_allclose(got, want, atol=atol, rtol=rtol) + + +@pytest.mark.parametrize("mode", ("uvu", "uvw")) +def test_streaming_uvu_and_scalar_uvw_match_standard_through_hvp(gpu_context, mode): + """Match standard values and AD for UVU and exact [1,V,1] UVW lowering.""" + jax, jnp = gpu_context + generated, standard, values, tangents, dout, rows, cols = _scalar_case( + gpu_context, mode + ) + + def native(*args): + return generated(*args, rows, cols) + + def standard_operator(*args): + return standard(*args, rows, cols) + + _assert_tree_close(jax.jit(native)(*values), jax.jit(standard_operator)(*values)) + _assert_tree_close( + jax.jvp(native, values, tangents)[1], + jax.jvp(standard_operator, values, tangents)[1], + ) + + def energy(operator, *args): + return jnp.vdot(operator(*args), dout) + + native_grad = jax.grad(lambda *args: energy(native, *args), (0, 1, 2)) + standard_grad = jax.grad(lambda *args: energy(standard_operator, *args), (0, 1, 2)) + _assert_tree_close(native_grad(*values), standard_grad(*values)) + _assert_tree_close( + jax.jvp(native_grad, values, tangents)[1], + jax.jvp(standard_grad, values, tangents)[1], + atol=4e-3, + rtol=4e-3, + ) + + +@pytest.mark.parametrize("active_input", range(3)) +def test_streaming_single_input_derivatives_match_standard(gpu_context, active_input): + """Match JVP and backward-JVP results with one active primal input.""" + jax, jnp = gpu_context + generated, standard, values, tangents, dout, rows, cols = _scalar_case( + gpu_context, "uvu" + ) + + def operator(implementation, arguments): + return implementation(*arguments, rows, cols) + + def varying_forward(implementation, value): + arguments = list(values) + arguments[active_input] = value + return operator(implementation, arguments) + + _assert_tree_close( + jax.jvp( + lambda value: varying_forward(generated, value), + (values[active_input],), + (tangents[active_input],), + )[1], + jax.jvp( + lambda value: varying_forward(standard, value), + (values[active_input],), + (tangents[active_input],), + )[1], + ) + + def gradients(implementation, arguments): + return jax.grad( + lambda *operands: jnp.vdot(operator(implementation, operands), dout), + (0, 1, 2), + )(*arguments) + + def varying_gradients(implementation, value): + arguments = list(values) + arguments[active_input] = value + return gradients(implementation, arguments) + + _assert_tree_close( + jax.jvp( + lambda value: varying_gradients(generated, value), + (values[active_input],), + (tangents[active_input],), + )[1], + jax.jvp( + lambda value: varying_gradients(standard, value), + (values[active_input],), + (tangents[active_input],), + )[1], + atol=4e-3, + rtol=4e-3, + ) + + +def test_irreducible_uvw_uses_standard_path_and_ir_mul_layout(monkeypatch): + """Delegate irreducible UVW through the established native layout.""" + import importlib + import jax.numpy as jnp + + from openequivariance.core.e3nn_lite import TPProblem + + module = importlib.import_module("openequivariance.jax.TensorProductConv") + + class FakeLoopUnroll: + """Record the native configuration without constructing a GPU kernel.""" + + def __init__(self, config, **kwargs): + del kwargs + self.config = config + + def forward(self, X, Y, W, rows, cols, sender_perm, **kwargs): + del Y, W, rows, cols, sender_perm, kwargs + return X + + def forward_cpu(self, X, Y, W, output, graph): + del Y, W, graph + output[...] = X + + monkeypatch.setattr(module, "LoopUnrollTensorProductConv", FakeLoopUnroll) + + problem = TPProblem( + "2x1e", + "1x0e", + "2x1e", + [(0, 0, 0, "uvw", True)], + shared_weights=False, + internal_weights=False, + layout="ir_mul", + ) + operator = module.TensorProductConv(problem, mode="auto") + assert operator.implementation == "standard" + assert operator._impl.config.layout == "mul_ir" + x = jnp.arange(6, dtype=jnp.float32).reshape(1, 6) + np.testing.assert_array_equal( + operator( + x, + jnp.ones((1, 1), dtype=jnp.float32), + jnp.ones((1, problem.weight_numel), dtype=jnp.float32), + jnp.zeros((1,), dtype=jnp.int32), + jnp.zeros((1,), dtype=jnp.int32), + ), + x, + ) + output = np.empty_like(np.asarray(x)) + operator.forward_cpu( + np.asarray(x), + np.ones((1, 1), dtype=np.float32), + np.ones((1, problem.weight_numel), dtype=np.float32), + output, + object(), + ) + np.testing.assert_array_equal(output, x) + + +def test_native_weight_permutation_round_trips_multiple_paths(): + """Round-trip canonical weights through nontrivial native UVU ordering.""" + import jax.numpy as jnp + + from openequivariance.core.e3nn_lite import TPProblem + from openequivariance.jax.TensorProductConv import TensorProductConv + + problem = TPProblem( + "2x0e + 2x1e", + "2x0e + 2x1e", + "2x0e + 2x1e", + [(0, 0, 0, "uvu", True), (1, 1, 1, "uvu", True)], + shared_weights=False, + internal_weights=False, + ) + operator = TensorProductConv(problem, mode="streaming") + canonical = jnp.arange(3 * problem.weight_numel, dtype=jnp.float32).reshape( + 3, problem.weight_numel + ) + np.testing.assert_array_equal( + operator.reorder_weights_to_e3nn(operator.reorder_weights_from_e3nn(canonical)), + canonical, + ) + + +def test_multipath_streaming_matches_standard_through_hvp(gpu_context): + """Match forward, full VJP, and HVP results for multiple UVU paths.""" + jax, jnp = gpu_context + from openequivariance.core.e3nn_lite import TPProblem + from openequivariance.jax.TensorProductConv import TensorProductConv + + problem = TPProblem( + "2x0e + 2x1e", + "1x0e + 1x1e", + "2x0e + 2x1e + 2x2e", + [ + (0, 0, 0, "uvu", True), + (0, 1, 1, "uvu", True), + (1, 0, 1, "uvu", True), + (1, 1, 0, "uvu", True), + (1, 1, 2, "uvu", True), + ], + shared_weights=False, + internal_weights=False, + ) + streaming = TensorProductConv(problem, mode="streaming") + standard = TensorProductConv(problem, mode="standard") + keys = jax.random.split(jax.random.key(912), 8) + rows = jnp.asarray([2, 0, 3, 1, 0, 2, 3, 1], dtype=jnp.int32) + cols = jnp.asarray([0, 2, 1, 3, 1, 3, 2, 0], dtype=jnp.int32) + x = jax.random.normal(keys[0], (4, problem.irreps_in1.dim), dtype=np.float32) + sh = jax.random.normal(keys[1], (8, problem.irreps_in2.dim), dtype=np.float32) + weights = jax.random.normal(keys[2], (8, problem.weight_numel), dtype=np.float32) + tangent_weights = jax.random.normal( + keys[5], (8, problem.weight_numel), dtype=np.float32 + ) + streaming_values = (x, sh, streaming.reorder_weights_from_e3nn(weights)) + standard_values = (x, sh, standard.reorder_weights_from_e3nn(weights)) + streaming_tangents = ( + jax.random.normal(keys[3], x.shape, dtype=np.float32), + jax.random.normal(keys[4], sh.shape, dtype=np.float32), + streaming.reorder_weights_from_e3nn(tangent_weights), + ) + standard_tangents = ( + streaming_tangents[0], + streaming_tangents[1], + standard.reorder_weights_from_e3nn(tangent_weights), + ) + dout = jax.random.normal(keys[6], (4, problem.irreps_out.dim), dtype=np.float32) + + def operator(implementation, *operands): + return implementation(*operands, rows, cols) + + _assert_tree_close( + operator(streaming, *streaming_values), + operator(standard, *standard_values), + ) + + def gradient(implementation, operands): + return jax.grad( + lambda *values: jnp.vdot(operator(implementation, *values), dout), + (0, 1, 2), + )(*operands) + + streaming_gradient = gradient(streaming, streaming_values) + standard_gradient = gradient(standard, standard_values) + _assert_tree_close( + ( + *streaming_gradient[:2], + streaming.reorder_weights_to_e3nn(streaming_gradient[2]), + ), + ( + *standard_gradient[:2], + standard.reorder_weights_to_e3nn(standard_gradient[2]), + ), + ) + streaming_hvp = jax.jvp( + lambda *values: gradient(streaming, values), + streaming_values, + streaming_tangents, + )[1] + standard_hvp = jax.jvp( + lambda *values: gradient(standard, values), + standard_values, + standard_tangents, + )[1] + _assert_tree_close( + (*streaming_hvp[:2], streaming.reorder_weights_to_e3nn(streaming_hvp[2])), + (*standard_hvp[:2], standard.reorder_weights_to_e3nn(standard_hvp[2])), + atol=4e-3, + rtol=4e-3, + ) + + +def test_shared_weights_use_the_loop_unroll_fallback(monkeypatch): + """Keep the established shared-weight interface on the standard path.""" + import importlib + import jax.numpy as jnp + + from openequivariance.core.e3nn_lite import TPProblem + + module = importlib.import_module("openequivariance.jax.TensorProductConv") + + class FakeLoopUnroll: + """Record the fallback call without constructing a GPU kernel.""" + + def __init__(self, config, **kwargs): + del kwargs + self.config = config + self.weights = None + + def forward(self, X, Y, W, rows, cols, sender_perm, **kwargs): + del X, Y, rows, cols, sender_perm, kwargs + self.weights = W + return W + + monkeypatch.setattr(module, "LoopUnrollTensorProductConv", FakeLoopUnroll) + problem = TPProblem( + "1x0e", + "1x0e", + "1x0e", + [(0, 0, 0, "uvu", True)], + shared_weights=True, + internal_weights=False, + ) + operator = module.TensorProductConv(problem) + shared_weights = jnp.ones((problem.weight_numel,), dtype=jnp.float32) + result = operator( + jnp.ones((2, 1), dtype=jnp.float32), + jnp.ones((3, 1), dtype=jnp.float32), + shared_weights, + jnp.array([0, 1, 0], dtype=jnp.int32), + jnp.array([1, 0, 1], dtype=jnp.int32), + ) + assert operator.implementation == "standard" + assert operator._impl.config.shared_weights + np.testing.assert_array_equal(result, shared_weights) diff --git a/tests/vmap_test.py b/tests/vmap_test.py index 4aad8942..5de5ddc2 100644 --- a/tests/vmap_test.py +++ b/tests/vmap_test.py @@ -75,3 +75,103 @@ def test_vmap_bcast_XW(ctx): (None, 0, None, None, None), (ctx["X"][0], ctx["Y"], ctx["W"][0], ctx["r"], ctx["c"]), ) + + +def test_vmap_streaming_preserves_padded_sentinel(with_jax): + """Match independent streaming calls with one padded batched launch.""" + if not with_jax: + pytest.skip("Skipping JAX tests") + os.environ["OEQ_NOTORCH"] = "1" + import jax + import jax.numpy as jnp + import openequivariance as oeq + + problem = oeq.TPProblem( + oeq.Irreps("2x0e"), + oeq.Irreps("1x0e"), + oeq.Irreps("2x0e"), + [(0, 0, 0, "uvu", True)], + shared_weights=False, + internal_weights=False, + ) + conv = oeq.jax.TensorProductConv(problem, mode="streaming") + batch_size, nodes, edges = 3, 3, 6 + keys = jax.random.split(jax.random.key(17), 4) + dtype = problem.irrep_dtype + x = jax.random.normal( + keys[0], (batch_size, nodes, problem.irreps_in1.dim), dtype=dtype + ) + y = jax.random.normal( + keys[1], (batch_size, edges, problem.irreps_in2.dim), dtype=dtype + ) + weights = jax.random.normal( + keys[2], (batch_size, edges, problem.weight_numel), dtype=dtype + ) + receivers = jnp.asarray((0, 1, 1, 2, nodes, nodes), dtype=jnp.int32) + senders = jnp.asarray((1, 0, 2, 1, nodes, nodes), dtype=jnp.int32) + row_ptr = jnp.asarray((0, 1, 3, 4), dtype=jnp.int32) + + def apply(a, b, w): + return conv( + a, + b, + w, + receivers, + senders, + indices_are_sorted=True, + row_ptr=row_ptr, + ) + + mapped = jax.jit(jax.vmap(apply))(x, y, weights) + expected = jnp.stack([apply(x[i], y[i], weights[i]) for i in range(batch_size)]) + assert jnp.allclose(mapped, expected, rtol=1e-5, atol=1e-5) + + cotangent = jax.random.normal(keys[3], mapped.shape, dtype=dtype) + + def loss(a, b, w, dout): + return jnp.vdot(apply(a, b, w), dout) + + gradient = jax.grad(loss, argnums=(0, 1, 2)) + mapped_grad = jax.jit(jax.vmap(gradient))(x, y, weights, cotangent) + expected_grad = tuple( + jnp.stack( + [ + gradient(x[i], y[i], weights[i], cotangent[i])[argument] + for i in range(batch_size) + ] + ) + for argument in range(3) + ) + for actual, reference in zip(mapped_grad, expected_grad): + assert jnp.allclose(actual, reference, rtol=1e-5, atol=1e-5) + assert jnp.all(mapped_grad[1][:, 4:] == 0) + assert jnp.all(mapped_grad[2][:, 4:] == 0) + + def hvp(a, b, w, dout, da, db, dw): + return jax.jvp( + lambda pa, pb, pw: gradient(pa, pb, pw, dout), + (a, b, w), + (da, db, dw), + )[1] + + tangents = (0.1 * x, 0.1 * y, 0.1 * weights) + mapped_hvp = jax.jit(jax.vmap(hvp))(x, y, weights, cotangent, *tangents) + expected_hvp = tuple( + jnp.stack( + [ + hvp( + x[i], + y[i], + weights[i], + cotangent[i], + tangents[0][i], + tangents[1][i], + tangents[2][i], + )[argument] + for i in range(batch_size) + ] + ) + for argument in range(3) + ) + for actual, reference in zip(mapped_hvp, expected_hvp): + assert jnp.allclose(actual, reference, rtol=1e-5, atol=1e-5) From 1966a047e07907015b904896577578589c985a76 Mon Sep 17 00:00:00 2001 From: Paul Fuchs Date: Tue, 15 Sep 2026 13:25:27 +0200 Subject: [PATCH 2/2] Clarify streaming convolution documentation --- docs/api.rst | 6 ++--- .../core/FactorizedComputationSchedule.py | 10 +++++++- .../openequivariance/jax/TensorProductConv.py | 23 ++++++++++--------- 3 files changed, 24 insertions(+), 15 deletions(-) diff --git a/docs/api.rst b/docs/api.rst index 162bce72..979a6274 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -44,9 +44,9 @@ do not conform exactly to the e3nn-jax API, but perform the same computation. JAX ``TensorProductConv`` uses the established loop-unroll implementation by default. Select ``mode="streaming"`` to require receiver streaming, or -``mode="auto"`` to use it when the problem is supported. Streaming accepts -padded tail edges and optional receiver offsets, but vmapped calls must share -one static topology. +``mode="auto"`` to use it when the problem is supported. Streaming supports +trailing padded edges represented by out-of-bounds node indices. Receiver row +pointers may optionally be provided for receiver-sorted edges. If you plan to use ``oeq.jax`` without PyTorch installed, you need to set ``OEQ_NOTORCH=1`` in your local environment (within Python, diff --git a/openequivariance/openequivariance/core/FactorizedComputationSchedule.py b/openequivariance/openequivariance/core/FactorizedComputationSchedule.py index 06834cd7..5c921eab 100644 --- a/openequivariance/openequivariance/core/FactorizedComputationSchedule.py +++ b/openequivariance/openequivariance/core/FactorizedComputationSchedule.py @@ -1,4 +1,12 @@ -"""Build schedules for receiver-streaming convolutions.""" +"""Build schedules for receiver-streaming convolutions. + +This implementation adapts the streaming concept and computational schedule +proposed by Chorošajev and Bény [CB2026]_. + +.. [CB2026] Chorošajev and Bény, *Sobek: Streaming Equivariant Tensor Product + Convolutions*, arXiv (2026). + https://doi.org/10.48550/arXiv.2607.18074 +""" from dataclasses import dataclass from enum import IntEnum diff --git a/openequivariance/openequivariance/jax/TensorProductConv.py b/openequivariance/openequivariance/jax/TensorProductConv.py index 7c3c78f7..26d46a60 100644 --- a/openequivariance/openequivariance/jax/TensorProductConv.py +++ b/openequivariance/openequivariance/jax/TensorProductConv.py @@ -18,17 +18,14 @@ class TensorProductConv: - r"""Apply a tensor-product convolution with an explicit implementation mode. + r"""Apply a tensor-product convolution to a directed graph. - ``mode="standard"`` preserves the established loop-unroll implementation. - ``mode="auto"`` uses receiver-row streaming convolution for supported - external unshared UVU problems and exact ``[1, V, 1]`` UVW reductions. It - selects the established loop-unroll implementation for other problems. - ``mode="streaming"`` requires receiver-row streaming. ``mode="standard"`` - always selects the established loop-unroll implementation. - Both public feature layouts are accepted. The standard implementation uses - differentiable boundary transposes when an ``"ir_mul"`` problem reaches - its native ``"mul_ir"`` kernel. + The mode selects between the general standard convolution, which supports + all tensor-product operations, and a more efficient streaming convolution + for UVU operations. The streaming implementation is based on the + computational schedule proposed by Chorošajev and Bény [CB2026]_. + ``mode="auto"`` uses streaming when supported and otherwise selects the + standard convolution. :param config: Specification of the tensor product. :param deterministic: Request deterministic aggregation. This selects the @@ -36,8 +33,12 @@ class TensorProductConv: streaming mode. :param kahan: Request Kahan summation. This selects the standard implementation in automatic mode and is unavailable in streaming mode. - :param requires_jvp: Preserve JVP selection for the standard implementation. + :param requires_jvp: Enable JVP support for the standard implementation. :param mode: One of ``"auto"``, ``"streaming"``, or ``"standard"``. + + .. [CB2026] Chorošajev and Bény, *Sobek: Streaming Equivariant Tensor + Product Convolutions*, arXiv (2026). + https://doi.org/10.48550/arXiv.2607.18074 """ _MODES = frozenset(("auto", "streaming", "standard"))