@@ -70,7 +70,7 @@ def generate_functional_roll_wrapper(
7070 code .writeline ("in0_rank = in0.dim()" )
7171 code .writeline ("in0_shape = list(in0.shape)" )
7272 code .newline ()
73-
73+
7474 code .writeline ("# Normalize dims and shifts to lists" )
7575 code .writeline ("if dims is None:" )
7676 with code .indent ():
@@ -96,20 +96,24 @@ def generate_functional_roll_wrapper(
9696 with code .indent ():
9797 code .writeline ("shifts = [shifts]" )
9898 code .newline ()
99-
99+
100100 code .writeline ("# Normalize negative dimensions" )
101101 code .writeline ("dims = [(d if d >= 0 else d + in0_rank) for d in dims]" )
102102 code .newline ()
103-
103+
104104 code .writeline ("# Validate" )
105105 code .writeline ("assert len(shifts) == len(dims), \\ " )
106- code .writeline (" f'shifts and dimensions must align, got {len(shifts)} shifts and {len(dims)} dims'" )
106+ code .writeline (
107+ " f'shifts and dimensions must align, got {len(shifts)} shifts and {len(dims)} dims'"
108+ )
107109 code .writeline ("for d in dims:" )
108110 with code .indent ():
109111 code .writeline ("assert 0 <= d < in0_rank, \\ " )
110- code .writeline (" f'Dimension out of range (expected to be in range of [0, {in0_rank-1}], but got {d})'" )
112+ code .writeline (
113+ " f'Dimension out of range (expected to be in range of [0, {in0_rank-1}], but got {d})'"
114+ )
111115 code .newline ()
112-
116+
113117 code .writeline ("# Normalize shifts to be within [0, size) for each dimension" )
114118 code .writeline ("normalized_shifts = [0] * in0_rank" )
115119 code .writeline ("for shift, dim in zip(shifts, dims):" )
@@ -120,7 +124,7 @@ def generate_functional_roll_wrapper(
120124 code .writeline ("# Python modulo handles negative shifts correctly" )
121125 code .writeline ("normalized_shifts[dim] = shift % size" )
122126 code .newline ()
123-
127+
124128 code .writeline ("# Check if any rolling is needed" )
125129 code .writeline ("if all(s == 0 for s in normalized_shifts):" )
126130 with code .indent ():
@@ -131,22 +135,22 @@ def generate_functional_roll_wrapper(
131135 code .writeline ("out0 = out0.reshape(original_shape)" )
132136 code .writeline ("return out0" )
133137 code .newline ()
134-
138+
135139 code .writeline ("out0 = torch.empty_like(in0)" )
136140 code .newline ()
137-
141+
138142 output_names : str = output_ref_for_wrapper ()
139143 call_str = (
140144 f"{ output_names } = { destination_passing_func_name } "
141145 f"({ parameter_ref_for_wrapper ()} )"
142146 )
143147 code .writeline (call_str )
144148 code .newline ()
145-
149+
146150 code .writeline ("if original_shape is not None:" )
147151 with code .indent ():
148152 code .writeline ("out0 = out0.reshape(original_shape)" )
149-
153+
150154 return_str = "return out0"
151155 code .writeline (return_str )
152156 code .newline ()
@@ -168,7 +172,7 @@ def generate_destination_passing_roll_wrapper(
168172 code .writeline ("shape = out0.shape" )
169173 code .writeline ("num_tasks = volume(shape)" )
170174 code .newline ()
171-
175+
172176 code .writeline ("# Handle empty tensors" )
173177 code .writeline ("if num_tasks == 0:" )
174178 with code .indent ():
@@ -179,7 +183,9 @@ def generate_destination_passing_roll_wrapper(
179183 code .writeline ("tile_size = min(512, triton.next_power_of_2(num_tasks))" )
180184 code .writeline ("num_warps = 4" )
181185 code .writeline ("num_ctas = min(65535, triton.cdiv(num_tasks, tile_size))" )
182- code .writeline ("tiles_per_cta = triton.cdiv(num_tasks, tile_size * num_ctas)" )
186+ code .writeline (
187+ "tiles_per_cta = triton.cdiv(num_tasks, tile_size * num_ctas)"
188+ )
183189 else :
184190 code .writeline ("num_warps = 1" )
185191 code .writeline ("num_ctas = 1" )
@@ -211,10 +217,12 @@ def generate_destination_passing_roll_wrapper(
211217
212218 shape_args : str = ", " .join (f"shape[{ i } ]" for i in range (rank ))
213219 code .writeline (f"{ shape_args } , # task indexing space" )
214-
215- shifts_args : str = ", " .join (f"normalized_shifts[{ i } ]" for i in range (rank ))
220+
221+ shifts_args : str = ", " .join (
222+ f"normalized_shifts[{ i } ]" for i in range (rank )
223+ )
216224 code .writeline (f"{ shifts_args } , # shifts for each dimension" )
217-
225+
218226 code .writeline ("num_tasks, # num tasks" )
219227 code .writeline ("tiles_per_cta=tiles_per_cta," )
220228 code .writeline ("tile_size=tile_size," )
@@ -294,7 +302,9 @@ def generate_roll_kernel(
294302 code .newline ()
295303
296304 code .writeline ("# loads" )
297- ptrs_expr : str = " + " .join (f"src_i{ j } * in0_stride{ j } " for j in range (rank ))
305+ ptrs_expr : str = " + " .join (
306+ f"src_i{ j } * in0_stride{ j } " for j in range (rank )
307+ )
298308 ptrs_expr : str = f"in0_ptr + { ptrs_expr } "
299309 code .writeline (f"in0 = tl.load({ ptrs_expr } , mask=mask)" )
300310 code .newline ()
@@ -331,7 +341,9 @@ def generate_roll_kernel(
331341 code .newline ()
332342
333343 code .writeline ("# loads" )
334- ptrs_expr : str = " + " .join (f"src_i{ j } * in0_stride{ j } " for j in range (rank ))
344+ ptrs_expr : str = " + " .join (
345+ f"src_i{ j } * in0_stride{ j } " for j in range (rank )
346+ )
335347 ptrs_expr : str = f"in0_ptr + { ptrs_expr } "
336348 code .writeline (f"in0 = tl.load({ ptrs_expr } , mask=mask)" )
337349 code .newline ()
@@ -341,7 +353,9 @@ def generate_roll_kernel(
341353 code .newline ()
342354
343355 code .writeline ("# stores" )
344- ptrs_expr : str = " + " .join (f"i{ j } * out0_stride{ j } " for j in range (rank ))
356+ ptrs_expr : str = " + " .join (
357+ f"i{ j } * out0_stride{ j } " for j in range (rank )
358+ )
345359 ptrs_expr : str = f"out0_ptr + { ptrs_expr } "
346360 code .writeline (f"tl.store({ ptrs_expr } , out0, mask=mask)" )
347361 code .newline ()
@@ -356,8 +370,12 @@ def generate_code(
356370 code : IndentedBuffer ,
357371) -> IndentedBuffer :
358372 code = generate_imports (code )
359- code = generate_functional_roll_wrapper (wrapper_name , destination_passing_func_name , code )
360- code = generate_destination_passing_roll_wrapper (rank , destination_passing_func_name , kernel_name , code )
373+ code = generate_functional_roll_wrapper (
374+ wrapper_name , destination_passing_func_name , code
375+ )
376+ code = generate_destination_passing_roll_wrapper (
377+ rank , destination_passing_func_name , kernel_name , code
378+ )
361379 code = generate_roll_kernel (rank , kernel_name , code )
362380 return code
363381
@@ -408,7 +426,11 @@ def arg_key(self, x, shifts, dims):
408426_roll_func = RollFunction ()
409427
410428
411- def roll (inp : torch .Tensor , shifts : Union [int , Tuple [int , ...]], dims : Union [None , int , Tuple [int , ...]] = None ) -> torch .Tensor :
429+ def roll (
430+ inp : torch .Tensor ,
431+ shifts : Union [int , Tuple [int , ...]],
432+ dims : Union [None , int , Tuple [int , ...]] = None ,
433+ ) -> torch .Tensor :
412434 logger .debug ("GEMS ROLL" )
413435 out = _roll_func (inp , shifts , dims )
414- return out
436+ return out
0 commit comments