Skip to content

Commit 13256df

Browse files
committed
Pre-commit fixed
1 parent 6c56d63 commit 13256df

3 files changed

Lines changed: 47 additions & 25 deletions

File tree

benchmark/test_roll_perf.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,7 @@
33
import pytest
44
import torch
55

6-
import flag_gems
6+
import flag_gems # noqa: F401
77
from benchmark.attri_util import DEFAULT_METRICS, FLOAT_DTYPES
88
from benchmark.performance_utils import Benchmark, generate_tensor_input
99

src/flag_gems/__init__.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -345,7 +345,7 @@ def torch_ge(v):
345345
("resolve_conj", resolve_conj),
346346
("resolve_neg", resolve_neg),
347347
("rms_norm", rms_norm),
348-
("roll", roll),
348+
("roll", roll),
349349
("rrelu_with_noise_backward", rrelu_with_noise_backward),
350350
("rsqrt", rsqrt),
351351
("rsqrt_", rsqrt_),

src/flag_gems/ops/roll.py

Lines changed: 45 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)