From 266a69aac79e7956d87a15d534db39265c81efb1 Mon Sep 17 00:00:00 2001 From: asglover <140220574+asglover@users.noreply.github.com> Date: Fri, 17 Jul 2026 19:56:28 -0700 Subject: [PATCH 1/3] move hip changes from post process to jinja --- .../openequivariance/_torch/TensorProduct.py | 2 +- .../_torch/TensorProductConv.py | 4 +-- .../_torch/extlib/__init__.py | 8 +---- .../openequivariance/core/LoopUnrollConv.py | 5 ++- .../openequivariance/core/LoopUnrollTP.py | 14 ++++---- .../openequivariance/jax/TensorProduct.py | 2 +- .../openequivariance/jax/TensorProductConv.py | 2 +- .../openequivariance/jax/extlib/__init__.py | 13 ++------ .../openequivariance/templates/jinja_utils.py | 15 ++++++++- .../templates/loop_unroll_batch.cuh | 28 ++++++++-------- .../templates/loop_unroll_conv_atomic.cuh | 28 ++++++++-------- .../templates/loop_unroll_conv_det.cuh | 32 +++++++++---------- .../templates/loop_unroll_tp.cuh | 16 +++++----- .../openequivariance/templates/macros.jinja | 8 ++--- 14 files changed, 86 insertions(+), 91 deletions(-) diff --git a/openequivariance/openequivariance/_torch/TensorProduct.py b/openequivariance/openequivariance/_torch/TensorProduct.py index de587b7b..00fe8c07 100644 --- a/openequivariance/openequivariance/_torch/TensorProduct.py +++ b/openequivariance/openequivariance/_torch/TensorProduct.py @@ -44,7 +44,7 @@ def _init_class(self): self, self.input_args["problem"], dp, - extlib.postprocess_kernel, + extlib.IS_HIP, self.input_args["torch_op"], ) diff --git a/openequivariance/openequivariance/_torch/TensorProductConv.py b/openequivariance/openequivariance/_torch/TensorProductConv.py index c75b7397..d052f909 100644 --- a/openequivariance/openequivariance/_torch/TensorProductConv.py +++ b/openequivariance/openequivariance/_torch/TensorProductConv.py @@ -4,7 +4,7 @@ import torch from openequivariance._torch.extlib import ( - postprocess_kernel, + IS_HIP, DeviceProp, BUILT_EXTENSION, ) @@ -78,7 +78,7 @@ def _init_class(self): self, self.input_args["problem"], dp, - postprocess_kernel, + IS_HIP, idx_dtype=np.int64, torch_op=self.input_args["torch_op"], deterministic=self.input_args["deterministic"], diff --git a/openequivariance/openequivariance/_torch/extlib/__init__.py b/openequivariance/openequivariance/_torch/extlib/__init__.py index 64d3ce38..6378b986 100644 --- a/openequivariance/openequivariance/_torch/extlib/__init__.py +++ b/openequivariance/openequivariance/_torch/extlib/__init__.py @@ -24,13 +24,7 @@ "Only CUDA and HIP backends are supported" ) - -def postprocess_kernel(kernel): - if torch.version.hip: - kernel = kernel.replace("__syncwarp();", "__threadfence_block();") - kernel = kernel.replace("__shfl_down_sync(FULL_MASK,", "__shfl_down(") - kernel = kernel.replace("atomicAdd", "unsafeAtomicAdd") - return kernel +IS_HIP = bool(torch.version.hip) def load_jit_extension(): diff --git a/openequivariance/openequivariance/core/LoopUnrollConv.py b/openequivariance/openequivariance/core/LoopUnrollConv.py index b3c63e66..17869760 100644 --- a/openequivariance/openequivariance/core/LoopUnrollConv.py +++ b/openequivariance/openequivariance/core/LoopUnrollConv.py @@ -20,7 +20,7 @@ def __init__( self, config, dp, - postprocess_kernel, + is_hip, *, idx_dtype: type[np.generic] = np.int64, torch_op: bool = False, @@ -34,7 +34,7 @@ def __init__( if kahan: assert deterministic - env = get_jinja_environment() + env = get_jinja_environment(is_hip=is_hip) template = env.get_template("loop_unroll_conv_atomic.cuh") analysis = filter_and_analyze_problem(config) @@ -207,7 +207,6 @@ def generate_double_backward_schedule(warps_per_block): backward_workspace_offset=self.backward_workspace_offset, double_backwardB_offset=self.double_backwardB_offset, ) - self.jit_kernel = postprocess_kernel(self.jit_kernel) self.kernel_string = json.dumps( { diff --git a/openequivariance/openequivariance/core/LoopUnrollTP.py b/openequivariance/openequivariance/core/LoopUnrollTP.py index 36801405..e2969041 100644 --- a/openequivariance/openequivariance/core/LoopUnrollTP.py +++ b/openequivariance/openequivariance/core/LoopUnrollTP.py @@ -17,10 +17,10 @@ class LoopUnrollTP(TensorProductBase): - def __init__(self, config, dp, postprocess_kernel, torch_op): + def __init__(self, config, dp, is_hip, torch_op): super().__init__(config, torch_op=torch_op) - env = get_jinja_environment() + env = get_jinja_environment(is_hip=is_hip) template = env.get_template("loop_unroll_batch.cuh") analysis = filter_and_analyze_problem(config) @@ -90,12 +90,10 @@ def generate_double_backward_schedule(warps_per_block): except Exception: raise - self.jit_kernel = postprocess_kernel( - template.render( - forward_schedule=self.forward_schedule, - backward_schedule=self.backward_schedule, - double_backward_schedule=self.double_backward_schedule, - ) + self.jit_kernel = template.render( + forward_schedule=self.forward_schedule, + backward_schedule=self.backward_schedule, + double_backward_schedule=self.double_backward_schedule, ) self.kernel_prop = { diff --git a/openequivariance/openequivariance/jax/TensorProduct.py b/openequivariance/openequivariance/jax/TensorProduct.py index 84d75e10..f880544f 100644 --- a/openequivariance/openequivariance/jax/TensorProduct.py +++ b/openequivariance/openequivariance/jax/TensorProduct.py @@ -16,7 +16,7 @@ class TensorProduct(LoopUnrollTP): def __init__(self, problem: TPProblem): dp = extlib.DeviceProp(0) - super().__init__(problem, dp, extlib.postprocess_kernel, torch_op=False) + super().__init__(problem, dp, extlib.IS_HIP, torch_op=False) self.kernel = self.kernel_string self.weight_numel = problem.weight_numel diff --git a/openequivariance/openequivariance/jax/TensorProductConv.py b/openequivariance/openequivariance/jax/TensorProductConv.py index 7101fa00..9234158f 100644 --- a/openequivariance/openequivariance/jax/TensorProductConv.py +++ b/openequivariance/openequivariance/jax/TensorProductConv.py @@ -42,7 +42,7 @@ def __init__( super().__init__( config, dp, - extlib.postprocess_kernel, + extlib.IS_HIP, idx_dtype=np.int32, torch_op=False, deterministic=deterministic, diff --git a/openequivariance/openequivariance/jax/extlib/__init__.py b/openequivariance/openequivariance/jax/extlib/__init__.py index d0965502..c34d72e5 100644 --- a/openequivariance/openequivariance/jax/extlib/__init__.py +++ b/openequivariance/openequivariance/jax/extlib/__init__.py @@ -1,19 +1,10 @@ import jax import openequivariance_extjax as oeq_extjax - -def postprocess_kernel(kernel): - if oeq_extjax.is_hip(): - kernel = kernel.replace("__syncwarp();", "__threadfence_block();") - kernel = kernel.replace("__shfl_down_sync(FULL_MASK,", "__shfl_down(") - kernel = kernel.replace("atomicAdd", "unsafeAtomicAdd") - return kernel - else: - return kernel - +IS_HIP = oeq_extjax.is_hip() platform = "CUDA" -if oeq_extjax.is_hip(): +if IS_HIP: platform = "ROCM" for name, target in oeq_extjax.registrations().items(): diff --git a/openequivariance/openequivariance/templates/jinja_utils.py b/openequivariance/openequivariance/templates/jinja_utils.py index bb326f27..ef7b866b 100644 --- a/openequivariance/openequivariance/templates/jinja_utils.py +++ b/openequivariance/openequivariance/templates/jinja_utils.py @@ -16,7 +16,7 @@ def sizeof(dtype): raise Exception("Provided undefined datatype to sizeof!") -def get_jinja_environment(): +def get_jinja_environment(is_hip=False): env = Environment( loader=PackageLoader("openequivariance"), extensions=["jinja2.ext.do"] ) @@ -24,4 +24,17 @@ def get_jinja_environment(): env.globals["divide"] = divide env.globals["sizeof"] = sizeof env.globals["enumerate"] = enumerate + + # HIP / CUDA intrinsic selection, previously handled by string replacement + # on the rendered kernel (extlib.postprocess_kernel). The HIP spellings + # (including whitespace) match the output of that historical replacement. + env.globals["is_hip"] = is_hip + env.globals["syncwarp"] = "__threadfence_block()" if is_hip else "__syncwarp()" + env.globals["atomic_add"] = "unsafeAtomicAdd" if is_hip else "atomicAdd" + if is_hip: + env.globals["shfl_down"] = lambda val, offset: f"__shfl_down( {val}, {offset})" + else: + env.globals["shfl_down"] = ( + lambda val, offset: f"__shfl_down_sync(FULL_MASK, {val}, {offset})" + ) return env diff --git a/openequivariance/openequivariance/templates/loop_unroll_batch.cuh b/openequivariance/openequivariance/templates/loop_unroll_batch.cuh index 9e775eac..83e0e0d2 100644 --- a/openequivariance/openequivariance/templates/loop_unroll_batch.cuh +++ b/openequivariance/openequivariance/templates/loop_unroll_batch.cuh @@ -43,7 +43,7 @@ forward(size_t num_products, IRREP_T* L1_in, IRREP_T* L2_in, IRREP_T* L3_out, WE {%- endif %} {%- for i, segment in enumerate(forward_schedule.segments) %} { - __syncwarp(); + {{ syncwarp }}; {{ declare_smem_variables(segment, "smem") }} {{ load_ir_segments(segment.L1Map, "l1", "L1_smem", "j") }} {{ load_ir_segments(segment.L2Map, "l2", "L2_smem", "j") }} @@ -53,9 +53,9 @@ forward(size_t num_products, IRREP_T* L1_in, IRREP_T* L2_in, IRREP_T* L3_out, WE ROW_OPERATION({{segment.problem.weight_numel}}, j, weights_smem[j + lane_id] = w[{{segment.weight_offset}} + j + lane_id];) {% endif %} - __syncwarp(); + {{ syncwarp }}; forward_loop_unroll_{{i}}(L1_smem, L2_smem, w, weights_smem, L3_smem, scratch_smem, lane_id); - __syncwarp(); + {{ syncwarp }}; {{ store_ir_segments(segment.L3Map, "l3", "L3_smem", "j") }} } {%- endfor %} @@ -99,7 +99,7 @@ backward(size_t num_products, {{ load_ir_segments(segment.L2Map, "l2_shft", "L2_smem", "j") }} {{ load_ir_segments(segment.L3Map, "l3_shft", "L3_grad_smem", "j") }} - __syncwarp(); + {{ syncwarp }}; {%- if not segment.L1Map.persist_load %} ROW_OPERATION({{segment.L1.dim}}, j, L1_grad_smem[j + lane_id] = 0.0f;) {%- endif %} @@ -113,10 +113,10 @@ backward(size_t num_products, ROW_OPERATION({{segment.problem.weight_numel}}, j, weights_grad_smem[j + lane_id] = 0.0;) {%- endif %} - __syncwarp(); + {{ syncwarp }}; backward_loop_unroll_{{i}}(L1_smem, L2_smem, w, weights_smem, L3_grad_smem, L1_grad_smem, L2_grad_smem, wgrad, weights_grad_smem, scratch_smem, lane_id); - __syncwarp(); + {{ syncwarp }}; IRREP_T* l1_grad_shft = L1_grad + i * {{backward_schedule.L1.dim}} + lane_id; IRREP_T* l2_grad_shft = L2_grad + i * {{backward_schedule.L2.dim}} + lane_id; @@ -134,7 +134,7 @@ backward(size_t num_products, {%- if not tpp.shared_weights %} ROW_OPERATION({{segment.problem.weight_numel}}, j, weights_grad_shft[{{segment.weight_offset}} + j] = weights_grad_smem[j + lane_id];) {%- else %} - ROW_OPERATION({{segment.problem.weight_numel}}, j, atomicAdd(weights_grad_shft + {{segment.weight_offset}} + j, weights_grad_smem[j + lane_id]);) + ROW_OPERATION({{segment.problem.weight_numel}}, j, {{ atomic_add }}(weights_grad_shft + {{segment.weight_offset}} + j, weights_grad_smem[j + lane_id]);) {%- endif %} {%- endif %} } {%- endfor %} @@ -175,7 +175,7 @@ double_backward_A( {%- endif %} {%- for i, segment in enumerate(forward_schedule.segments) %} { - __syncwarp(); + {{ syncwarp }}; {{ declare_smem_variables(segment, "smem") }} ROW_OPERATION({{segment.L3.dim}}, j, L3_smem[j + lane_id] = 0.0f;) WEIGHT_T* w_buffer; @@ -199,9 +199,9 @@ double_backward_A( {{ load_ir_segments_force(segment.L2Map, "l2", "L2_smem", "j") }} {{ load_ir_segments_force(segment.L1Map, "l1_dgrad", "L1_smem", "j") }} } - __syncwarp(); + {{ syncwarp }}; forward_loop_unroll_{{i}}(L1_smem, L2_smem, w_buffer, weights_smem, L3_smem, scratch_smem, lane_id); - __syncwarp(); + {{ syncwarp }}; } {{ store_ir_segments(segment.L3Map, "l3", "L3_smem", "j") }} @@ -257,7 +257,7 @@ double_backward_B( {{ load_ir_segments_force(segment.L2Map, "l2_shft", "L2_smem", "j") }} {{ load_ir_segments_force(segment.L2Map, "l2_original", "L2_dgrad_smem", "j") }} - __syncwarp(); + {{ syncwarp }}; {%- if not segment.L1Map.persist_load %} ROW_OPERATION({{segment.L1.dim}}, j, L1_grad_smem[j + lane_id] = 0.0f;) {%- endif %} @@ -287,10 +287,10 @@ double_backward_B( L2_dgrad_buffer = L2_smem; } - __syncwarp(); + {{ syncwarp }}; double_backward_loop_unroll_{{i}}(L1_smem, L2_buffer, w_buffer, weights_smem, L3_grad_smem, L1_grad_smem, L2_grad_smem, L2_dgrad_buffer, n, wgrad, weights_grad_smem, scratch_smem, lane_id); - __syncwarp(); + {{ syncwarp }}; } IRREP_T* l1_grad_shft = L1_grad + i * {{schedule.L1.dim}} + lane_id; @@ -309,7 +309,7 @@ double_backward_B( {%- if not tpp.shared_weights %} ROW_OPERATION({{segment.problem.weight_numel}}, j, weights_grad_shft[{{segment.weight_offset}} + j] = weights_grad_smem[j + lane_id];) {%- else %} - ROW_OPERATION({{segment.problem.weight_numel}}, j, atomicAdd(weights_grad_shft + {{segment.weight_offset}} + j, weights_grad_smem[j + lane_id]);) + ROW_OPERATION({{segment.problem.weight_numel}}, j, {{ atomic_add }}(weights_grad_shft + {{segment.weight_offset}} + j, weights_grad_smem[j + lane_id]);) {%- endif %} {% endif %} } diff --git a/openequivariance/openequivariance/templates/loop_unroll_conv_atomic.cuh b/openequivariance/openequivariance/templates/loop_unroll_conv_atomic.cuh index 99ebe540..3d461dbc 100644 --- a/openequivariance/openequivariance/templates/loop_unroll_conv_atomic.cuh +++ b/openequivariance/openequivariance/templates/loop_unroll_conv_atomic.cuh @@ -82,7 +82,7 @@ forward(IRREP_T* L1_in, WEIGHT_T* w = weights; {%- endif %} - __syncwarp(); + {{ syncwarp }}; {{ load_ir_segments(segment.L1Map, "l1", "L1_smem", "j") }} {{ load_ir_segments(segment.L2Map, "l2", "L2_smem", "j") }} ROW_OPERATION({{segment.L3.dim}}, j, L3_smem[j + lane_id] = 0.0f;) @@ -91,9 +91,9 @@ forward(IRREP_T* L1_in, ROW_OPERATION({{segment.problem.weight_numel}}, j, weights_smem[j + lane_id] = w[{{segment.weight_offset}} + j + lane_id];) {%- endif %} - __syncwarp(); + {{ syncwarp }}; forward_loop_unroll_{{i}}(L1_smem, L2_smem, w, weights_smem, L3_smem, scratch_smem, lane_id); - __syncwarp(); + {{ syncwarp }}; {{ store_ir_segments(segment.L3Map, "l3", "L3_smem", "j") }} } @@ -143,7 +143,7 @@ backward(IRREP_T* L1_in, IRREP_T* L1_grad, {{ load_ir_segments(segment.L2Map, "l2_shft", "L2_smem", "j") }} {{ load_ir_segments(segment.L3Map, "l3_shft", "L3_grad_smem", "j") }} - __syncwarp(); + {{ syncwarp }}; {%- if not segment.L1Map.persist_load %} ROW_OPERATION({{segment.L1.dim}}, j, L1_grad_smem[j + lane_id] = 0.0f;) {%- endif %} @@ -160,10 +160,10 @@ backward(IRREP_T* L1_in, IRREP_T* L1_grad, IRREP_T* l2_grad_shft = L2_grad + i * {{backward_schedule.L2.dim}} + lane_id; WEIGHT_T* weights_grad_shft = wgrad + lane_id; - __syncwarp(); + {{ syncwarp }}; backward_loop_unroll_{{i}}(L1_smem, L2_smem, w, weights_smem, L3_grad_smem, L1_grad_smem, L2_grad_smem, wgrad, weights_grad_smem, scratch_smem, lane_id); - __syncwarp(); + {{ syncwarp }}; {{ store_ir_segments(segment.L1Map, "l1_grad_shft", "L1_grad_smem", "j") }} {{ store_ir_segments(segment.L2Map, "l2_grad_shft", "L2_grad_smem", "j") }} @@ -172,7 +172,7 @@ backward(IRREP_T* L1_in, IRREP_T* L1_grad, {%- if not tpp.shared_weights %} ROW_OPERATION({{segment.problem.weight_numel}}, j, weights_grad_shft[{{segment.weight_offset}} + j] = weights_grad_smem[j + lane_id];) {%- else %} - ROW_OPERATION({{segment.problem.weight_numel}}, j, atomicAdd(weights_grad_shft + {{segment.weight_offset}} + j, weights_grad_smem[j + lane_id]);) + ROW_OPERATION({{segment.problem.weight_numel}}, j, {{ atomic_add }}(weights_grad_shft + {{segment.weight_offset}} + j, weights_grad_smem[j + lane_id]);) {%- endif %} {%- endif %} } {%- endfor %} @@ -217,7 +217,7 @@ double_backward_A(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, WEIGHT_T* w_dgrad = W_dgrad; {%- endif %} - __syncwarp(); + {{ syncwarp }}; {{ load_ir_segments(segment.L1Map, "l1", "L1_smem", "j") }} {{ load_ir_segments(segment.L2Map, "l2", "L2_smem", "j") }} ROW_OPERATION({{segment.L3.dim}}, j, L3_smem[j + lane_id] = 0.0f;) @@ -241,9 +241,9 @@ double_backward_A(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, {{ load_ir_segments(segment.L1Map, "l1_dgrad", "L1_smem", "j") }} } - __syncwarp(); + {{ syncwarp }}; forward_loop_unroll_{{i}}(L1_smem, L2_smem, w_buffer, weights_smem, L3_smem, scratch_smem, lane_id); - __syncwarp(); + {{ syncwarp }}; } {{ store_ir_segments(segment.L3Map, "l3", "L3_smem", "j") }} @@ -303,7 +303,7 @@ double_backward_B(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, {{ load_ir_segments(segment.L2Map, "l2_shft", "L2_smem", "j") }} {{ load_ir_segments(segment.L2Map, "l2_original", "L2_dgrad_smem", "j") }} - __syncwarp(); + {{ syncwarp }}; {%- if not segment.L1Map.persist_load %} ROW_OPERATION({{segment.L1.dim}}, j, L1_grad_smem[j + lane_id] = 0.0f;) {%- endif %} @@ -333,10 +333,10 @@ double_backward_B(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, L2_dgrad_buffer = L2_smem; } - __syncwarp(); + {{ syncwarp }}; double_backward_loop_unroll_{{i}}(L1_smem, L2_buffer, w_buffer, weights_smem, L3_grad_smem, L1_grad_smem, L2_grad_smem, L2_dgrad_buffer, n, wgrad, weights_grad_smem, scratch_smem, lane_id); - __syncwarp(); + {{ syncwarp }}; } IRREP_T* l1_grad_shft = L1_grad + col * {{schedule.L1.dim}} + lane_id; @@ -355,7 +355,7 @@ double_backward_B(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, {%- if not tpp.shared_weights %} ROW_OPERATION({{segment.problem.weight_numel}}, j, weights_grad_shft[{{segment.weight_offset}} + j] = weights_grad_smem[j + lane_id];) {%- else %} - ROW_OPERATION({{segment.problem.weight_numel}}, j, atomicAdd(weights_grad_shft + {{segment.weight_offset}} + j, weights_grad_smem[j + lane_id]);) + ROW_OPERATION({{segment.problem.weight_numel}}, j, {{ atomic_add }}(weights_grad_shft + {{segment.weight_offset}} + j, weights_grad_smem[j + lane_id]);) {%- endif %} {% endif %} } diff --git a/openequivariance/openequivariance/templates/loop_unroll_conv_det.cuh b/openequivariance/openequivariance/templates/loop_unroll_conv_det.cuh index 01789e2d..f5bb56a4 100644 --- a/openequivariance/openequivariance/templates/loop_unroll_conv_det.cuh +++ b/openequivariance/openequivariance/templates/loop_unroll_conv_det.cuh @@ -152,13 +152,13 @@ forward( ROW_OPERATION({{segment.problem.weight_numel}}, j, weights_smem[j + lane_id] = w[{{segment.weight_offset}} + j + lane_id];) {%- endif %} - __syncwarp(); + {{ syncwarp }}; forward_loop_unroll_{{i}}(L1_smem, L2_smem, w, weights_smem, {{ns.L3_accum}}, scratch_smem, lane_id); - __syncwarp(); + {{ syncwarp }}; {%- if forward_schedule.kahan %} kahanAdd<{{segment.L3.dim}}>(L3_kahan_smem, L3_smem); - __syncwarp(); + {{ syncwarp }}; {%- endif %} bool changeRow = (i < end - 1) && (row != rows[i+1]); @@ -170,7 +170,7 @@ forward( firstSegment = false; } {{ store_ir_segments(segment.L3Map, "dst", "L3_smem", "j") }} - __syncwarp(); + {{ syncwarp }}; ROW_OPERATION({{segment.L3.dim}}, j, L3_smem[j + lane_id] = 0.0f;) } @@ -250,7 +250,7 @@ backward(IRREP_T* L1_in, IRREP_T* L1_grad, {{ load_ir_segments(segment.L2Map, "l2_shft", "L2_smem", "j") }} {{ load_ir_segments(segment.L3Map, "l3_shft", "L3_grad_smem", "j") }} - __syncwarp(); + {{ syncwarp }}; ROW_OPERATION({{segment.L2.dim}}, j, L2_grad_smem[j + lane_id] = 0.0f;) {%- if not backward_schedule.stream_weights %} @@ -262,10 +262,10 @@ backward(IRREP_T* L1_in, IRREP_T* L1_grad, IRREP_T* l2_grad_shft = L2_grad + tperm_idx * {{backward_schedule.L2.dim}} + lane_id; WEIGHT_T* weights_grad_shft = wgrad + lane_id; - __syncwarp(); + {{ syncwarp }}; backward_loop_unroll_{{i}}(L1_smem, L2_smem, w, weights_smem, L3_grad_smem, {{ns.L1_accum}}, L2_grad_smem, wgrad, weights_grad_smem, scratch_smem, lane_id); - __syncwarp(); + {{ syncwarp }}; {%- if backward_schedule.kahan %} kahanAdd<{{segment.L1.dim}}>(L1_kahan_smem, L1_grad_smem); @@ -279,7 +279,7 @@ backward(IRREP_T* L1_in, IRREP_T* L1_grad, firstSegment = false; } {{ store_ir_segments(segment.L1Map, "dst", "L1_grad_smem", "j") }} - __syncwarp(); + {{ syncwarp }}; ROW_OPERATION({{segment.L1.dim}}, j, L1_grad_smem[j + lane_id] = 0.0f;) } @@ -353,7 +353,7 @@ double_backward_A(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, WEIGHT_T* w_dgrad = W_dgrad; {%- endif %} - __syncwarp(); + {{ syncwarp }}; {{ load_ir_segments(segment.L1Map, "l1", "L1_smem", "j") }} {{ load_ir_segments(segment.L2Map, "l2", "L2_smem", "j") }} @@ -376,9 +376,9 @@ double_backward_A(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, {{ load_ir_segments(segment.L1Map, "l1_dgrad", "L1_smem", "j") }} } - __syncwarp(); + {{ syncwarp }}; forward_loop_unroll_{{i}}(L1_smem, L2_smem, w_buffer, weights_smem, {{ns.L3_accum}}, scratch_smem, lane_id); - __syncwarp(); + {{ syncwarp }}; } {%- if forward_schedule.kahan %} @@ -393,7 +393,7 @@ double_backward_A(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, firstSegment = false; } {{ store_ir_segments(segment.L3Map, "dst", "L3_smem", "j") }} - __syncwarp(); + {{ syncwarp }}; ROW_OPERATION({{segment.L3.dim}}, j, L3_smem[j + lane_id] = 0.0f;) } @@ -480,7 +480,7 @@ double_backward_B(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, {{ load_ir_segments(segment.L2Map, "l2_shft", "L2_smem", "j") }} {{ load_ir_segments(segment.L2Map, "l2_original", "L2_dgrad_smem", "j") }} - __syncwarp(); + {{ syncwarp }}; {%- if not segment.L2Map.persist_load %} ROW_OPERATION({{segment.L2.dim}}, j, L2_grad_smem[j + lane_id] = 0.0f;) @@ -507,10 +507,10 @@ double_backward_B(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, L2_dgrad_buffer = L2_smem; } - __syncwarp(); + {{ syncwarp }}; double_backward_loop_unroll_{{i}}(L1_smem, L2_buffer, w_buffer, weights_smem, L3_grad_smem, {{ns.L1_accum}}, L2_grad_smem, L2_dgrad_buffer, n, wgrad, weights_grad_smem, scratch_smem, lane_id); - __syncwarp(); + {{ syncwarp }}; } {%- if backward_schedule.kahan %} @@ -534,7 +534,7 @@ double_backward_B(IRREP_T* L1_in, IRREP_T* L2_in, WEIGHT_T* W, IRREP_T* L3_grad, firstSegment = false; } {{ store_ir_segments(segment.L1Map, "dst", "L1_grad_smem", "j") }} - __syncwarp(); + {{ syncwarp }}; ROW_OPERATION({{segment.L1.dim}}, j, L1_grad_smem[j + lane_id] = 0.0f;) } diff --git a/openequivariance/openequivariance/templates/loop_unroll_tp.cuh b/openequivariance/openequivariance/templates/loop_unroll_tp.cuh index e4c8dabe..52b92191 100644 --- a/openequivariance/openequivariance/templates/loop_unroll_tp.cuh +++ b/openequivariance/openequivariance/templates/loop_unroll_tp.cuh @@ -73,10 +73,10 @@ __device__ __forceinline__ void forward_loop_unroll_{{id}}(IRREP_T* __restrict__ {%- if problem.instructions[k].connection_mode == "uvw" %} {{layout_store(problem.layout, L1[u].mul, L3[w].ir.dim, 'scratch', '0', 'l3_vec', '=', '1.0')}} - __syncwarp(); + {{ syncwarp }}; offset = {{ L3.slices()[w].start}}; matmul_fwd_{{id}}_{{k}}(weights_smem, scratch, L3_smem + offset); - __syncwarp(); + {{ syncwarp }}; #pragma unroll for(int j = 0; j < {{L3[w].ir.dim}}; j++) @@ -220,10 +220,10 @@ __device__ __forceinline__ void forward_loop_unroll_{{id}}(IRREP_T* __restrict__ WEIGHT_T* tmp = weights + {{segment.weight_offset + weight_start}} + k * {{slice_size}} + lane_id; ROW_OPERATION({{slice_size}}, j, weights_smem[j + lane_id] = tmp[j];) - __syncwarp(); + {{ syncwarp }}; offset = {{ L3.slices()[w].start}}; {{matmul_basename}}A_{{id}}_{{k}}(weights_smem, L3_grad_smem + offset, scratch); - __syncwarp(); + {{ syncwarp }}; {{layout_load(problem.layout, L1[u].mul, L3[w].ir.dim, 'scratch', '0', 'l3_grad')}} @@ -250,13 +250,13 @@ __device__ __forceinline__ void forward_loop_unroll_{{id}}(IRREP_T* __restrict__ {{ reg_store(L1[u].mul, L3[w].ir.dim, "scratch", "0", "l3_grad", "=", 1.0) }} - __syncwarp(); + {{ syncwarp }}; {{matmul_basename}}B_{{id}}_{{k}}(L3_grad_smem + offset, scratch, weights_smem); - __syncwarp(); + {{ syncwarp }}; tmp = weights_grad + {{segment.weight_offset + weight_start}} + k * {{slice_size}} + lane_id; {%- if problem.shared_weights %} - ROW_OPERATION({{slice_size}}, j, atomicAdd(tmp + j, weights_smem[j + lane_id]);) + ROW_OPERATION({{slice_size}}, j, {{ atomic_add }}(tmp + j, weights_smem[j + lane_id]);) {%- else %} {%- if double_bwd %} if(n == 0) { @@ -280,7 +280,7 @@ __device__ __forceinline__ void forward_loop_unroll_{{id}}(IRREP_T* __restrict__ } #pragma unroll for (int offset = {{ warp_size // 2}}; offset > 0; offset /= 2) { - l2_grad[j] += __shfl_down_sync(FULL_MASK, l2_grad[j], offset); + l2_grad[j] += {{ shfl_down("l2_grad[j]", "offset") }}; } } diff --git a/openequivariance/openequivariance/templates/macros.jinja b/openequivariance/openequivariance/templates/macros.jinja index f9108822..59727e88 100644 --- a/openequivariance/openequivariance/templates/macros.jinja +++ b/openequivariance/openequivariance/templates/macros.jinja @@ -129,7 +129,7 @@ Keys map to lists of tuples with (name, dtype, num_elements) of each subarray. {%- elif map.storeback_procedure[idx] == "accumulate" %} ROW_OPERATION({{range_len}}, {{loop_var}}, {{glb_ptr_shft}}[{{loop_var}} + {{src_rng.start}}] += {{smem_ptr}}[{{loop_var}} + {{dst_rng.start}} + lane_id];) {%- elif map.storeback_procedure[idx] == "atomic_accumulate" %} - ROW_OPERATION({{range_len}}, {{loop_var}}, atomicAdd({{glb_ptr_shft}} + {{src_rng.start}} + {{loop_var}}, {{smem_ptr}}[{{dst_rng.start}} + lane_id + {{loop_var}}]);) + ROW_OPERATION({{range_len}}, {{loop_var}}, {{ atomic_add }}({{glb_ptr_shft}} + {{src_rng.start}} + {{loop_var}}, {{smem_ptr}}[{{dst_rng.start}} + lane_id + {{loop_var}}]);) {%- endif %} {%- endfor %} {%- elif map.src_views[0].layout == "ir_mul" %} @@ -149,7 +149,7 @@ Keys map to lists of tuples with (name, dtype, num_elements) of each subarray. {%- endfor %} {%- elif map.storeback_procedure[idx] == "atomic_accumulate" %} {%- for i in range(m.src_mul_ir.ir.dim) %} - ROW_OPERATION({{m.src_mul_ir.mul}}, {{loop_var}}, atomicAdd({{glb_ptr_shft}} + {{m.src_view.ir_mul_offset + i * m.src_view.ir_mul_stride}} + {{loop_var}}, {{smem_ptr}}[{{m.dst_rng.start + i * m.src_mul_ir.mul}} + {{loop_var}} + lane_id]);) + ROW_OPERATION({{m.src_mul_ir.mul}}, {{loop_var}}, {{ atomic_add }}({{glb_ptr_shft}} + {{m.src_view.ir_mul_offset + i * m.src_view.ir_mul_stride}} + {{loop_var}}, {{smem_ptr}}[{{m.dst_rng.start + i * m.src_mul_ir.mul}} + {{loop_var}} + lane_id]);) {%- endfor %} {%- endif %} {%- endfor %} @@ -182,7 +182,7 @@ Keys map to lists of tuples with (name, dtype, num_elements) of each subarray. {%- for i in range(dim) %} t_regs[{{i}}] = {{smem_ptr}}[{{offset}} + lane_id * {{dim}} + {{i}}]; {%- endfor %} - __syncwarp(); + {{ syncwarp }}; {%- for i in range(dim) %} {{smem_ptr}}[{{offset}} + lane_id + {{i * mul}}] = t_regs[{{i}}]; {%- endfor %} @@ -201,7 +201,7 @@ Keys map to lists of tuples with (name, dtype, num_elements) of each subarray. {%- for i in range(dim) %} t_regs[{{i}}] = {{smem_ptr}}[{{offset}} + lane_id + {{i * mul}}]; {%- endfor %} - __syncwarp(); + {{ syncwarp }}; {%- for i in range(dim) %} {{smem_ptr}}[{{offset}} + lane_id * {{dim}} + {{i}}] = t_regs[{{i}}]; {%- endfor %} From 8bc7dd21af198a5840300ff9b1cb6fc08b72a395 Mon Sep 17 00:00:00 2001 From: asglover <140220574+asglover@users.noreply.github.com> Date: Fri, 17 Jul 2026 19:56:46 -0700 Subject: [PATCH 2/3] tests - separate commit so removable --- tests/kernel_postprocess_goldens.json | 58 ++++++ tests/kernel_postprocess_jinja_test.py | 252 +++++++++++++++++++++++++ 2 files changed, 310 insertions(+) create mode 100644 tests/kernel_postprocess_goldens.json create mode 100644 tests/kernel_postprocess_jinja_test.py diff --git a/tests/kernel_postprocess_goldens.json b/tests/kernel_postprocess_goldens.json new file mode 100644 index 00000000..32ad0dd4 --- /dev/null +++ b/tests/kernel_postprocess_goldens.json @@ -0,0 +1,58 @@ +{ + "batch_diffdock0_f32-w32": { + "cuda_sha256": "9c43b1f2d88e9ff76cad975afbbf2fde0b6b501b0081c691e2237416869186b2", + "hip_sha256": "a9d543e58c048fd1bc63d963a6ee949822f809227cdc88b7ee227b599935866a" + }, + "batch_diffdock0_f32-w64": { + "cuda_sha256": "0205481eb601a1f973734680c90c2c87a4abce76ce4110b2981d1abd2fee4b9d", + "hip_sha256": "26f5f44e8dc0fb9560a411e7dad4c74cff94542a38cfd3e54ff0fe40596bb6fe" + }, + "batch_mace0_f32-w32": { + "cuda_sha256": "dec44f329030e6399bc961b2cd8affcfcc20b6cc3165aa39a5426c10b18d4c4e", + "hip_sha256": "0c39331080adf7766354ef74806b4c7bd4a5ccda5369e9601a0e616e1f5059b4" + }, + "batch_mace0_f32-w64": { + "cuda_sha256": "5d010cd8a67617996345d80232c527e11405bf281dcadc08965e5c0c71481948", + "hip_sha256": "ab3829a839fe82d7f9c7c86543ec363ed93c87d8b5c07893812aed4a519d434e" + }, + "batch_uvu_f64-w32": { + "cuda_sha256": "31bd53be1fcd63dd153fd5379da4b4a1d7d66aeab8d56f385141205dca9e268d", + "hip_sha256": "a25906aad201448be9592b411d5df5ee0751463f43d47e3fc711c06f23626540" + }, + "batch_uvu_f64-w64": { + "cuda_sha256": "f125225580382cc620c753650f0711fcdc34839e90644a110bc35785aa274450", + "hip_sha256": "1294fcd04604c5b0bc846924d4fcda0f2198c887d6d4384a7aac78b3f33c889b" + }, + "batch_uvw_shared_f32-w32": { + "cuda_sha256": "cd4394fa4421a102591340499087b240045f60ae5ec3cdc6b58514849c53b96b", + "hip_sha256": "2d8fee60c8b181ea78545a957baf625be7e433e0541b4f1a6e77081e3db62ef7" + }, + "batch_uvw_shared_f32-w64": { + "cuda_sha256": "3499c9cffadd3a18f2c894a28a82b058a6aa15b0ecf89f126c53655197e50682", + "hip_sha256": "5400849be2a05443ac9851c7aec2d240f0a6145456472786147cd69bae2ab83a" + }, + "conv_atomic_mace0_f32-w32": { + "cuda_sha256": "aeba7da13bf8eed0fc107a7408dcb95c4ac33fe47f9c6c24b1e887078d5991a9", + "hip_sha256": "e93778792ef88e933fa63e6e2031ccb11407d424db6771c074c5447300fd5490" + }, + "conv_atomic_mace0_f32-w64": { + "cuda_sha256": "c0d30e486731b28e7adf52877463464b2c9819cfc4d205b28187bd09ce47ede8", + "hip_sha256": "30cc73ce23ff8ab1c33dadbc6b5566f3db83931cfd8f46231410a27b9ffe6e3d" + }, + "conv_det_kahan_mace0_f32-w32": { + "cuda_sha256": "1aaa1ab2027de8b2b074d0a9cb84619dcd906a915b9b829572f797d1dbfef8d3", + "hip_sha256": "d7e3ce8e8f7d7d041e62f72a682448b42cc8808eb516765234516a29ab7d73c6" + }, + "conv_det_kahan_mace0_f32-w64": { + "cuda_sha256": "78a11bd39e5fb7151e02c34c1240f51be9560ef0f175fe3f2c73671fdc93ba8c", + "hip_sha256": "ea55a1e70a785aec9d196f5701ced40b1d2fbfb41d13f1582253e6a0cc17a91a" + }, + "conv_det_mace0_f32-w32": { + "cuda_sha256": "008882d9e32fd4e990cf5d02ddd5825c10ec43c2520a3896a2235c168cc5e6f7", + "hip_sha256": "5b1e059408c52e2ba6d73794356d1bf4fe731e1ce231030e22ecd323445ae3bd" + }, + "conv_det_mace0_f32-w64": { + "cuda_sha256": "13823e5b3548a3357f9b99d1fc4a35dddf1c127c8d2376b0195183fa695f16f6", + "hip_sha256": "2ee04cb8cf6a0c1a95e6699eef66b9c189e9d381122fe8eeba8fd08db6000e57" + } +} diff --git a/tests/kernel_postprocess_jinja_test.py b/tests/kernel_postprocess_jinja_test.py new file mode 100644 index 00000000..4c1c7ed9 --- /dev/null +++ b/tests/kernel_postprocess_jinja_test.py @@ -0,0 +1,252 @@ +""" +TEMPORARY test for folding `postprocess_kernel` into the Jinja pipeline. + +`extlib.postprocess_kernel` (both _torch and jax variants) does three string +replacements on the rendered kernel when running on HIP: + + 1. "__syncwarp();" -> "__threadfence_block();" + 2. "__shfl_down_sync(FULL_MASK," -> "__shfl_down(" + 3. "atomicAdd" -> "unsafeAtomicAdd" + +This test: + * renders kernels for a matrix of tensor products / convolutions on CPU + (no GPU required -- rendering is pure Python), + * characterizes exactly what postprocess_kernel changes (and asserts it + catches *every* occurrence, e.g. no "__syncwarp()" without a semicolon + that the string replace would silently miss), + * pins the CUDA render and the postprocessed HIP render as sha256 goldens + in kernel_postprocess_goldens.json, + * once LoopUnrollTP/LoopUnrollConv grow an `is_hip` flag (the "new + process"), verifies the Jinja-rendered HIP kernel is byte-identical to + the old postprocessed output. + +Regenerate goldens with OEQ_REGEN_GOLDENS=1. Delete this file (and the +goldens) once postprocess_kernel is removed. +""" + +import hashlib +import inspect +import json +import os +from pathlib import Path + +import pytest + +# Rendering needs no GPU. If torch is missing or has no CUDA/HIP backend, +# keep openequivariance/__init__.py from importing its torch extension. +try: + import torch + + _TORCH_USABLE = bool(torch.version.cuda or torch.version.hip) +except ImportError: + _TORCH_USABLE = False +if not _TORCH_USABLE: + os.environ["OEQ_NOTORCH"] = "1" + +import numpy as np # noqa: E402 + +from openequivariance.core.e3nn_lite import TPProblem # noqa: E402 +from openequivariance.core.LoopUnrollTP import LoopUnrollTP # noqa: E402 +from openequivariance.core.LoopUnrollConv import LoopUnrollConv # noqa: E402 +from openequivariance.benchmark.problems import ( # noqa: E402 + diffdock_problems, + mace_problems, +) + +GOLDEN_PATH = Path(__file__).parent / "kernel_postprocess_goldens.json" + + +class FakeDeviceProp: + """Stand-in for extlib.DeviceProp so kernels render without a GPU.""" + + def __init__(self, warpsize): + self.warpsize = warpsize + self.maxSharedMemPerBlock = 48 * 1024 + self.multiprocessorCount = 108 + + +def reference_hip_postprocess(kernel): + """Verbatim copy of the HIP branch of extlib.postprocess_kernel.""" + kernel = kernel.replace("__syncwarp();", "__threadfence_block();") + kernel = kernel.replace("__shfl_down_sync(FULL_MASK,", "__shfl_down(") + kernel = kernel.replace("atomicAdd", "unsafeAtomicAdd") + return kernel + + +def _new_process_available(): + return all( + "is_hip" in inspect.signature(cls.__init__).parameters + for cls in (LoopUnrollTP, LoopUnrollConv) + ) + + +_HAS_IS_HIP = _new_process_available() + + +def _uvu_f64_problem(): + return TPProblem( + "32x1e + 8x2e", + "1x1e + 1x2e", + "32x1e + 8x2e", + [(0, 0, 0, "uvu", True), (1, 1, 1, "uvu", True)], + shared_weights=False, + internal_weights=False, + irrep_dtype=np.float64, + weight_dtype=np.float64, + ) + + +def _shared_weight_uvw_problem(): + return TPProblem( + "16x2e", + "4x2e", + "16x2e", + [(0, 0, 0, "uvw", True)], + shared_weights=True, + internal_weights=False, + irrep_dtype=np.float32, + weight_dtype=np.float32, + ) + + +def _render(kind, problem, dp, is_hip): + # Pre-refactor, the third constructor argument was a postprocessing + # callable; post-refactor it is the is_hip flag itself. + if _HAS_IS_HIP: + backend_arg = is_hip + else: + backend_arg = reference_hip_postprocess if is_hip else (lambda k: k) + + if kind == "batch": + return LoopUnrollTP(problem, dp, backend_arg, torch_op=False).jit_kernel + if kind == "conv_atomic": + return LoopUnrollConv( + problem, dp, backend_arg, torch_op=False, deterministic=False + ).jit_kernel + if kind == "conv_det": + return LoopUnrollConv( + problem, dp, backend_arg, torch_op=False, deterministic=True + ).jit_kernel + if kind == "conv_det_kahan": + return LoopUnrollConv( + problem, dp, backend_arg, torch_op=False, deterministic=True, kahan=True + ).jit_kernel + raise ValueError(kind) + + +CASES = { + "batch_mace0_f32": ("batch", lambda: mace_problems()[0]), + "batch_diffdock0_f32": ("batch", lambda: diffdock_problems()[0]), + "batch_uvu_f64": ("batch", _uvu_f64_problem), + "batch_uvw_shared_f32": ("batch", _shared_weight_uvw_problem), + "conv_atomic_mace0_f32": ("conv_atomic", lambda: mace_problems()[0]), + "conv_det_mace0_f32": ("conv_det", lambda: mace_problems()[0]), + "conv_det_kahan_mace0_f32": ("conv_det_kahan", lambda: mace_problems()[0]), +} + +WARP_SIZES = [32, 64] # NVIDIA / AMD CDNA + + +def _case_params(): + return [ + pytest.param(case_id, warpsize, id=f"{case_id}-w{warpsize}") + for case_id in CASES + for warpsize in WARP_SIZES + ] + + +def _render_case(case_id, warpsize, is_hip=False): + kind, problem_fn = CASES[case_id] + return _render(kind, problem_fn(), FakeDeviceProp(warpsize), is_hip) + + +def _sha256(s): + return hashlib.sha256(s.encode()).hexdigest() + + +# --------------------------------------------------------------------------- +# Characterization of the OLD process (passes today). +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("case_id,warpsize", _case_params()) +def test_postprocess_fully_characterized(case_id, warpsize): + """The three string replacements catch every relevant token, and nothing + else in the kernel can be clobbered by them.""" + cuda_kernel = _render_case(case_id, warpsize, is_hip=False) + + # Every __syncwarp appears exactly as "__syncwarp();" -- the string + # replace would silently miss any other spelling. + n_syncwarp = cuda_kernel.count("__syncwarp") + assert n_syncwarp == cuda_kernel.count("__syncwarp();") + + # Every warp shuffle appears exactly as "__shfl_down_sync(FULL_MASK,". + n_shfl = cuda_kernel.count("__shfl_down") + assert n_shfl == cuda_kernel.count("__shfl_down_sync(FULL_MASK,") + + # "atomicAdd" is replaced globally: no identifier may contain it as a + # substring other than the calls themselves ("unsafeAtomicAdd" pre-image + # would double-replace). + assert "unsafeAtomicAdd" not in cuda_kernel + + hip_kernel = reference_hip_postprocess(cuda_kernel) + + assert "__syncwarp" not in hip_kernel + assert "__shfl_down_sync" not in hip_kernel + n_atomic = cuda_kernel.count("atomicAdd") + assert hip_kernel.count("unsafeAtomicAdd") == n_atomic + assert hip_kernel.count("atomicAdd") == 0 # every call became the unsafe one + + assert hip_kernel.count("__threadfence_block();") >= n_syncwarp + + +@pytest.mark.parametrize("case_id,warpsize", _case_params()) +def test_golden_hashes(case_id, warpsize): + """Pin CUDA render + postprocessed HIP render byte-for-byte, so the Jinja + refactor can prove it changes nothing. Regen: OEQ_REGEN_GOLDENS=1.""" + goldens = json.loads(GOLDEN_PATH.read_text()) if GOLDEN_PATH.exists() else {} + key = f"{case_id}-w{warpsize}" + + cuda_kernel = _render_case(case_id, warpsize, is_hip=False) + hip_kernel = reference_hip_postprocess(cuda_kernel) + entry = {"cuda_sha256": _sha256(cuda_kernel), "hip_sha256": _sha256(hip_kernel)} + + if os.environ.get("OEQ_REGEN_GOLDENS") == "1" or key not in goldens: + goldens[key] = entry + GOLDEN_PATH.write_text(json.dumps(goldens, indent=2, sort_keys=True) + "\n") + else: + assert goldens[key] == entry, ( + f"Rendered kernel for {key} deviates from the pinned golden. If the " + "change is intentional, regenerate with OEQ_REGEN_GOLDENS=1." + ) + + +# --------------------------------------------------------------------------- +# Equivalence of the NEW process (skipped until the Jinja pipeline handles HIP). +# --------------------------------------------------------------------------- + + +@pytest.mark.skipif( + not _HAS_IS_HIP, + reason="LoopUnrollTP/LoopUnrollConv do not take is_hip yet (new process not implemented)", +) +@pytest.mark.parametrize("case_id,warpsize", _case_params()) +def test_jinja_hip_matches_postprocess(case_id, warpsize): + """The Jinja-rendered HIP kernel must be byte-identical to what the old + string-replacement postprocess produced, and the CUDA render must be + unchanged relative to the pre-refactor goldens.""" + cuda_new = _render_case(case_id, warpsize, is_hip=False) + hip_new = _render_case(case_id, warpsize, is_hip=True) + + assert hip_new == reference_hip_postprocess(cuda_new), ( + "Jinja HIP render differs from old postprocess_kernel output" + ) + + goldens = json.loads(GOLDEN_PATH.read_text()) + key = f"{case_id}-w{warpsize}" + assert _sha256(cuda_new) == goldens[key]["cuda_sha256"], ( + "CUDA render changed relative to pre-refactor golden" + ) + assert _sha256(hip_new) == goldens[key]["hip_sha256"], ( + "HIP render changed relative to pre-refactor golden" + ) From d66fba15e4cce6022e9e380b489dbba6d16efaed Mon Sep 17 00:00:00 2001 From: austin_glover Date: Mon, 27 Jul 2026 01:30:38 +0000 Subject: [PATCH 3/3] Remove temporary kernel postprocess verification tests These tests existed to verify that folding postprocess_kernel into the Jinja pipeline changed nothing. That has been confirmed: the Jinja HIP render is byte-identical to the old string-replacement output across all 14 parametrizations, and the CUDA render is byte-identical to main. Co-Authored-By: Claude Opus 5 (1M context) --- tests/kernel_postprocess_goldens.json | 58 ------ tests/kernel_postprocess_jinja_test.py | 252 ------------------------- 2 files changed, 310 deletions(-) delete mode 100644 tests/kernel_postprocess_goldens.json delete mode 100644 tests/kernel_postprocess_jinja_test.py diff --git a/tests/kernel_postprocess_goldens.json b/tests/kernel_postprocess_goldens.json deleted file mode 100644 index 32ad0dd4..00000000 --- a/tests/kernel_postprocess_goldens.json +++ /dev/null @@ -1,58 +0,0 @@ -{ - "batch_diffdock0_f32-w32": { - "cuda_sha256": "9c43b1f2d88e9ff76cad975afbbf2fde0b6b501b0081c691e2237416869186b2", - "hip_sha256": "a9d543e58c048fd1bc63d963a6ee949822f809227cdc88b7ee227b599935866a" - }, - "batch_diffdock0_f32-w64": { - "cuda_sha256": "0205481eb601a1f973734680c90c2c87a4abce76ce4110b2981d1abd2fee4b9d", - "hip_sha256": "26f5f44e8dc0fb9560a411e7dad4c74cff94542a38cfd3e54ff0fe40596bb6fe" - }, - "batch_mace0_f32-w32": { - "cuda_sha256": "dec44f329030e6399bc961b2cd8affcfcc20b6cc3165aa39a5426c10b18d4c4e", - "hip_sha256": "0c39331080adf7766354ef74806b4c7bd4a5ccda5369e9601a0e616e1f5059b4" - }, - "batch_mace0_f32-w64": { - "cuda_sha256": "5d010cd8a67617996345d80232c527e11405bf281dcadc08965e5c0c71481948", - "hip_sha256": "ab3829a839fe82d7f9c7c86543ec363ed93c87d8b5c07893812aed4a519d434e" - }, - "batch_uvu_f64-w32": { - "cuda_sha256": "31bd53be1fcd63dd153fd5379da4b4a1d7d66aeab8d56f385141205dca9e268d", - "hip_sha256": "a25906aad201448be9592b411d5df5ee0751463f43d47e3fc711c06f23626540" - }, - "batch_uvu_f64-w64": { - "cuda_sha256": "f125225580382cc620c753650f0711fcdc34839e90644a110bc35785aa274450", - "hip_sha256": "1294fcd04604c5b0bc846924d4fcda0f2198c887d6d4384a7aac78b3f33c889b" - }, - "batch_uvw_shared_f32-w32": { - "cuda_sha256": "cd4394fa4421a102591340499087b240045f60ae5ec3cdc6b58514849c53b96b", - "hip_sha256": "2d8fee60c8b181ea78545a957baf625be7e433e0541b4f1a6e77081e3db62ef7" - }, - "batch_uvw_shared_f32-w64": { - "cuda_sha256": "3499c9cffadd3a18f2c894a28a82b058a6aa15b0ecf89f126c53655197e50682", - "hip_sha256": "5400849be2a05443ac9851c7aec2d240f0a6145456472786147cd69bae2ab83a" - }, - "conv_atomic_mace0_f32-w32": { - "cuda_sha256": "aeba7da13bf8eed0fc107a7408dcb95c4ac33fe47f9c6c24b1e887078d5991a9", - "hip_sha256": "e93778792ef88e933fa63e6e2031ccb11407d424db6771c074c5447300fd5490" - }, - "conv_atomic_mace0_f32-w64": { - "cuda_sha256": "c0d30e486731b28e7adf52877463464b2c9819cfc4d205b28187bd09ce47ede8", - "hip_sha256": "30cc73ce23ff8ab1c33dadbc6b5566f3db83931cfd8f46231410a27b9ffe6e3d" - }, - "conv_det_kahan_mace0_f32-w32": { - "cuda_sha256": "1aaa1ab2027de8b2b074d0a9cb84619dcd906a915b9b829572f797d1dbfef8d3", - "hip_sha256": "d7e3ce8e8f7d7d041e62f72a682448b42cc8808eb516765234516a29ab7d73c6" - }, - "conv_det_kahan_mace0_f32-w64": { - "cuda_sha256": "78a11bd39e5fb7151e02c34c1240f51be9560ef0f175fe3f2c73671fdc93ba8c", - "hip_sha256": "ea55a1e70a785aec9d196f5701ced40b1d2fbfb41d13f1582253e6a0cc17a91a" - }, - "conv_det_mace0_f32-w32": { - "cuda_sha256": "008882d9e32fd4e990cf5d02ddd5825c10ec43c2520a3896a2235c168cc5e6f7", - "hip_sha256": "5b1e059408c52e2ba6d73794356d1bf4fe731e1ce231030e22ecd323445ae3bd" - }, - "conv_det_mace0_f32-w64": { - "cuda_sha256": "13823e5b3548a3357f9b99d1fc4a35dddf1c127c8d2376b0195183fa695f16f6", - "hip_sha256": "2ee04cb8cf6a0c1a95e6699eef66b9c189e9d381122fe8eeba8fd08db6000e57" - } -} diff --git a/tests/kernel_postprocess_jinja_test.py b/tests/kernel_postprocess_jinja_test.py deleted file mode 100644 index 4c1c7ed9..00000000 --- a/tests/kernel_postprocess_jinja_test.py +++ /dev/null @@ -1,252 +0,0 @@ -""" -TEMPORARY test for folding `postprocess_kernel` into the Jinja pipeline. - -`extlib.postprocess_kernel` (both _torch and jax variants) does three string -replacements on the rendered kernel when running on HIP: - - 1. "__syncwarp();" -> "__threadfence_block();" - 2. "__shfl_down_sync(FULL_MASK," -> "__shfl_down(" - 3. "atomicAdd" -> "unsafeAtomicAdd" - -This test: - * renders kernels for a matrix of tensor products / convolutions on CPU - (no GPU required -- rendering is pure Python), - * characterizes exactly what postprocess_kernel changes (and asserts it - catches *every* occurrence, e.g. no "__syncwarp()" without a semicolon - that the string replace would silently miss), - * pins the CUDA render and the postprocessed HIP render as sha256 goldens - in kernel_postprocess_goldens.json, - * once LoopUnrollTP/LoopUnrollConv grow an `is_hip` flag (the "new - process"), verifies the Jinja-rendered HIP kernel is byte-identical to - the old postprocessed output. - -Regenerate goldens with OEQ_REGEN_GOLDENS=1. Delete this file (and the -goldens) once postprocess_kernel is removed. -""" - -import hashlib -import inspect -import json -import os -from pathlib import Path - -import pytest - -# Rendering needs no GPU. If torch is missing or has no CUDA/HIP backend, -# keep openequivariance/__init__.py from importing its torch extension. -try: - import torch - - _TORCH_USABLE = bool(torch.version.cuda or torch.version.hip) -except ImportError: - _TORCH_USABLE = False -if not _TORCH_USABLE: - os.environ["OEQ_NOTORCH"] = "1" - -import numpy as np # noqa: E402 - -from openequivariance.core.e3nn_lite import TPProblem # noqa: E402 -from openequivariance.core.LoopUnrollTP import LoopUnrollTP # noqa: E402 -from openequivariance.core.LoopUnrollConv import LoopUnrollConv # noqa: E402 -from openequivariance.benchmark.problems import ( # noqa: E402 - diffdock_problems, - mace_problems, -) - -GOLDEN_PATH = Path(__file__).parent / "kernel_postprocess_goldens.json" - - -class FakeDeviceProp: - """Stand-in for extlib.DeviceProp so kernels render without a GPU.""" - - def __init__(self, warpsize): - self.warpsize = warpsize - self.maxSharedMemPerBlock = 48 * 1024 - self.multiprocessorCount = 108 - - -def reference_hip_postprocess(kernel): - """Verbatim copy of the HIP branch of extlib.postprocess_kernel.""" - kernel = kernel.replace("__syncwarp();", "__threadfence_block();") - kernel = kernel.replace("__shfl_down_sync(FULL_MASK,", "__shfl_down(") - kernel = kernel.replace("atomicAdd", "unsafeAtomicAdd") - return kernel - - -def _new_process_available(): - return all( - "is_hip" in inspect.signature(cls.__init__).parameters - for cls in (LoopUnrollTP, LoopUnrollConv) - ) - - -_HAS_IS_HIP = _new_process_available() - - -def _uvu_f64_problem(): - return TPProblem( - "32x1e + 8x2e", - "1x1e + 1x2e", - "32x1e + 8x2e", - [(0, 0, 0, "uvu", True), (1, 1, 1, "uvu", True)], - shared_weights=False, - internal_weights=False, - irrep_dtype=np.float64, - weight_dtype=np.float64, - ) - - -def _shared_weight_uvw_problem(): - return TPProblem( - "16x2e", - "4x2e", - "16x2e", - [(0, 0, 0, "uvw", True)], - shared_weights=True, - internal_weights=False, - irrep_dtype=np.float32, - weight_dtype=np.float32, - ) - - -def _render(kind, problem, dp, is_hip): - # Pre-refactor, the third constructor argument was a postprocessing - # callable; post-refactor it is the is_hip flag itself. - if _HAS_IS_HIP: - backend_arg = is_hip - else: - backend_arg = reference_hip_postprocess if is_hip else (lambda k: k) - - if kind == "batch": - return LoopUnrollTP(problem, dp, backend_arg, torch_op=False).jit_kernel - if kind == "conv_atomic": - return LoopUnrollConv( - problem, dp, backend_arg, torch_op=False, deterministic=False - ).jit_kernel - if kind == "conv_det": - return LoopUnrollConv( - problem, dp, backend_arg, torch_op=False, deterministic=True - ).jit_kernel - if kind == "conv_det_kahan": - return LoopUnrollConv( - problem, dp, backend_arg, torch_op=False, deterministic=True, kahan=True - ).jit_kernel - raise ValueError(kind) - - -CASES = { - "batch_mace0_f32": ("batch", lambda: mace_problems()[0]), - "batch_diffdock0_f32": ("batch", lambda: diffdock_problems()[0]), - "batch_uvu_f64": ("batch", _uvu_f64_problem), - "batch_uvw_shared_f32": ("batch", _shared_weight_uvw_problem), - "conv_atomic_mace0_f32": ("conv_atomic", lambda: mace_problems()[0]), - "conv_det_mace0_f32": ("conv_det", lambda: mace_problems()[0]), - "conv_det_kahan_mace0_f32": ("conv_det_kahan", lambda: mace_problems()[0]), -} - -WARP_SIZES = [32, 64] # NVIDIA / AMD CDNA - - -def _case_params(): - return [ - pytest.param(case_id, warpsize, id=f"{case_id}-w{warpsize}") - for case_id in CASES - for warpsize in WARP_SIZES - ] - - -def _render_case(case_id, warpsize, is_hip=False): - kind, problem_fn = CASES[case_id] - return _render(kind, problem_fn(), FakeDeviceProp(warpsize), is_hip) - - -def _sha256(s): - return hashlib.sha256(s.encode()).hexdigest() - - -# --------------------------------------------------------------------------- -# Characterization of the OLD process (passes today). -# --------------------------------------------------------------------------- - - -@pytest.mark.parametrize("case_id,warpsize", _case_params()) -def test_postprocess_fully_characterized(case_id, warpsize): - """The three string replacements catch every relevant token, and nothing - else in the kernel can be clobbered by them.""" - cuda_kernel = _render_case(case_id, warpsize, is_hip=False) - - # Every __syncwarp appears exactly as "__syncwarp();" -- the string - # replace would silently miss any other spelling. - n_syncwarp = cuda_kernel.count("__syncwarp") - assert n_syncwarp == cuda_kernel.count("__syncwarp();") - - # Every warp shuffle appears exactly as "__shfl_down_sync(FULL_MASK,". - n_shfl = cuda_kernel.count("__shfl_down") - assert n_shfl == cuda_kernel.count("__shfl_down_sync(FULL_MASK,") - - # "atomicAdd" is replaced globally: no identifier may contain it as a - # substring other than the calls themselves ("unsafeAtomicAdd" pre-image - # would double-replace). - assert "unsafeAtomicAdd" not in cuda_kernel - - hip_kernel = reference_hip_postprocess(cuda_kernel) - - assert "__syncwarp" not in hip_kernel - assert "__shfl_down_sync" not in hip_kernel - n_atomic = cuda_kernel.count("atomicAdd") - assert hip_kernel.count("unsafeAtomicAdd") == n_atomic - assert hip_kernel.count("atomicAdd") == 0 # every call became the unsafe one - - assert hip_kernel.count("__threadfence_block();") >= n_syncwarp - - -@pytest.mark.parametrize("case_id,warpsize", _case_params()) -def test_golden_hashes(case_id, warpsize): - """Pin CUDA render + postprocessed HIP render byte-for-byte, so the Jinja - refactor can prove it changes nothing. Regen: OEQ_REGEN_GOLDENS=1.""" - goldens = json.loads(GOLDEN_PATH.read_text()) if GOLDEN_PATH.exists() else {} - key = f"{case_id}-w{warpsize}" - - cuda_kernel = _render_case(case_id, warpsize, is_hip=False) - hip_kernel = reference_hip_postprocess(cuda_kernel) - entry = {"cuda_sha256": _sha256(cuda_kernel), "hip_sha256": _sha256(hip_kernel)} - - if os.environ.get("OEQ_REGEN_GOLDENS") == "1" or key not in goldens: - goldens[key] = entry - GOLDEN_PATH.write_text(json.dumps(goldens, indent=2, sort_keys=True) + "\n") - else: - assert goldens[key] == entry, ( - f"Rendered kernel for {key} deviates from the pinned golden. If the " - "change is intentional, regenerate with OEQ_REGEN_GOLDENS=1." - ) - - -# --------------------------------------------------------------------------- -# Equivalence of the NEW process (skipped until the Jinja pipeline handles HIP). -# --------------------------------------------------------------------------- - - -@pytest.mark.skipif( - not _HAS_IS_HIP, - reason="LoopUnrollTP/LoopUnrollConv do not take is_hip yet (new process not implemented)", -) -@pytest.mark.parametrize("case_id,warpsize", _case_params()) -def test_jinja_hip_matches_postprocess(case_id, warpsize): - """The Jinja-rendered HIP kernel must be byte-identical to what the old - string-replacement postprocess produced, and the CUDA render must be - unchanged relative to the pre-refactor goldens.""" - cuda_new = _render_case(case_id, warpsize, is_hip=False) - hip_new = _render_case(case_id, warpsize, is_hip=True) - - assert hip_new == reference_hip_postprocess(cuda_new), ( - "Jinja HIP render differs from old postprocess_kernel output" - ) - - goldens = json.loads(GOLDEN_PATH.read_text()) - key = f"{case_id}-w{warpsize}" - assert _sha256(cuda_new) == goldens[key]["cuda_sha256"], ( - "CUDA render changed relative to pre-refactor golden" - ) - assert _sha256(hip_new) == goldens[key]["hip_sha256"], ( - "HIP render changed relative to pre-refactor golden" - )