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 %}