Skip to content

Commit a5ccfe4

Browse files
committed
fix: codestyle fixes from black and isort
1 parent 6a71a3f commit a5ccfe4

5 files changed

Lines changed: 94 additions & 147 deletions

File tree

benchmark/test_special_perf.py

Lines changed: 14 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -6,19 +6,15 @@
66
import triton
77

88
import flag_gems
9-
from benchmark.attri_util import BOOL_DTYPES, FLOAT_DTYPES, INT_DTYPES, BenchLevel
10-
from benchmark.performance_utils import (
11-
Benchmark,
12-
Config,
13-
GenericBenchmark,
14-
GenericBenchmark2DOnly,
15-
GenericBenchmark4DOnly,
16-
GenericBenchmarkExcluse1D,
17-
GenericBenchmarkExcluse3D,
18-
SkipVersion,
19-
generate_tensor_input,
20-
vendor_name,
21-
)
9+
from benchmark.attri_util import (BOOL_DTYPES, FLOAT_DTYPES, INT_DTYPES,
10+
BenchLevel)
11+
from benchmark.performance_utils import (Benchmark, Config, GenericBenchmark,
12+
GenericBenchmark2DOnly,
13+
GenericBenchmark4DOnly,
14+
GenericBenchmarkExcluse1D,
15+
GenericBenchmarkExcluse3D,
16+
SkipVersion, generate_tensor_input,
17+
vendor_name)
2218

2319

2420
class GroupedTopKBenchmark(Benchmark):
@@ -1330,6 +1326,11 @@ def test_perf_t_copy():
13301326
bench = TCopyBenchmark(
13311327
op_name="t_copy",
13321328
torch_op=torch.ops.aten.t_copy,
1329+
dtypes=FLOAT_DTYPES,
1330+
)
1331+
bench.run()
1332+
1333+
13331334
@pytest.mark.feature_dropout
13341335
def test_perf_feature_dropout():
13351336
"""Benchmark for feature_dropout operation."""

src/flag_gems/__init__.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,8 @@
88
from flag_gems.config import aten_patch_list, resolve_user_setting
99
from flag_gems.experimental_ops import * # noqa: F403
1010
from flag_gems.fused import * # noqa: F403
11-
from flag_gems.logging_utils import setup_flaggems_logging, teardown_flaggems_logging
11+
from flag_gems.logging_utils import (setup_flaggems_logging,
12+
teardown_flaggems_logging)
1213
from flag_gems.modules import * # noqa: F403
1314
from flag_gems.ops import * # noqa: F403
1415
from flag_gems.patches import * # noqa: F403

src/flag_gems/ops/__init__.py

Lines changed: 54 additions & 102 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,5 @@
1-
from flag_gems.ops._functional_sym_constrain_range_for_size import (
2-
_functional_sym_constrain_range_for_size,
3-
)
1+
from flag_gems.ops._functional_sym_constrain_range_for_size import \
2+
_functional_sym_constrain_range_for_size
43
from flag_gems.ops._safe_softmax import _safe_softmax
54
from flag_gems.ops._upsample_nearest_exact1d import _upsample_nearest_exact1d
65
from flag_gems.ops.abs import abs, abs_
@@ -25,46 +24,30 @@
2524
from flag_gems.ops.argmin import argmin
2625
from flag_gems.ops.asinh_ import asinh_
2726
from flag_gems.ops.atan import atan, atan_
28-
from flag_gems.ops.attention import (
29-
ScaleDotProductAttention,
30-
flash_attention_forward,
31-
flash_attn_varlen_func,
32-
scaled_dot_product_attention,
33-
scaled_dot_product_attention_backward,
34-
scaled_dot_product_attention_forward,
35-
)
27+
from flag_gems.ops.attention import (ScaleDotProductAttention,
28+
flash_attention_forward,
29+
flash_attn_varlen_func,
30+
scaled_dot_product_attention,
31+
scaled_dot_product_attention_backward,
32+
scaled_dot_product_attention_forward)
3633
from flag_gems.ops.avg_pool2d import avg_pool2d, avg_pool2d_backward
3734
from flag_gems.ops.baddbmm import baddbmm
3835
from flag_gems.ops.batch_norm import batch_norm, batch_norm_backward
39-
from flag_gems.ops.bitwise_and import (
40-
bitwise_and_scalar,
41-
bitwise_and_scalar_,
42-
bitwise_and_scalar_tensor,
43-
bitwise_and_tensor,
44-
bitwise_and_tensor_,
45-
)
36+
from flag_gems.ops.bitwise_and import (bitwise_and_scalar, bitwise_and_scalar_,
37+
bitwise_and_scalar_tensor,
38+
bitwise_and_tensor, bitwise_and_tensor_)
4639
from flag_gems.ops.bitwise_left_shift import bitwise_left_shift
4740
from flag_gems.ops.bitwise_not import bitwise_not, bitwise_not_
48-
from flag_gems.ops.bitwise_or import (
49-
bitwise_or_scalar,
50-
bitwise_or_scalar_,
51-
bitwise_or_scalar_tensor,
52-
bitwise_or_tensor,
53-
bitwise_or_tensor_,
54-
)
41+
from flag_gems.ops.bitwise_or import (bitwise_or_scalar, bitwise_or_scalar_,
42+
bitwise_or_scalar_tensor,
43+
bitwise_or_tensor, bitwise_or_tensor_)
5544
from flag_gems.ops.bitwise_right_shift import bitwise_right_shift
5645
from flag_gems.ops.bmm import bmm, bmm_out
5746
from flag_gems.ops.cat import cat
5847
from flag_gems.ops.ceil import ceil, ceil_, ceil_out
5948
from flag_gems.ops.celu import celu, celu_
60-
from flag_gems.ops.clamp import (
61-
clamp,
62-
clamp_,
63-
clamp_min,
64-
clamp_min_,
65-
clamp_tensor,
66-
clamp_tensor_,
67-
)
49+
from flag_gems.ops.clamp import (clamp, clamp_, clamp_min, clamp_min_,
50+
clamp_tensor, clamp_tensor_)
6851
from flag_gems.ops.contiguous import contiguous
6952
from flag_gems.ops.conv1d import conv1d
7053
from flag_gems.ops.conv2d import conv2d
@@ -80,17 +63,9 @@
8063
from flag_gems.ops.diag_embed import diag_embed
8164
from flag_gems.ops.diagonal import diagonal_backward
8265
from flag_gems.ops.digamma_ import digamma_
83-
from flag_gems.ops.div import (
84-
div_mode,
85-
div_mode_,
86-
floor_divide,
87-
floor_divide_,
88-
remainder,
89-
remainder_,
90-
true_divide,
91-
true_divide_,
92-
true_divide_out,
93-
)
66+
from flag_gems.ops.div import (div_mode, div_mode_, floor_divide,
67+
floor_divide_, remainder, remainder_,
68+
true_divide, true_divide_, true_divide_out)
9469
from flag_gems.ops.dot import dot
9570
from flag_gems.ops.dropout import dropout, dropout_backward
9671
from flag_gems.ops.elu import elu, elu_, elu_backward
@@ -101,17 +76,11 @@
10176
from flag_gems.ops.exp import exp, exp_, exp_out
10277
from flag_gems.ops.exp2 import exp2, exp2_
10378
from flag_gems.ops.exponential_ import exponential_
104-
from flag_gems.ops.feature_dropout import feature_dropout, feature_dropout_
10579
from flag_gems.ops.eye import eye
10680
from flag_gems.ops.eye_m import eye_m
107-
from flag_gems.ops.fill import (
108-
fill_scalar,
109-
fill_scalar_,
110-
fill_scalar_out,
111-
fill_tensor,
112-
fill_tensor_,
113-
fill_tensor_out,
114-
)
81+
from flag_gems.ops.feature_dropout import feature_dropout, feature_dropout_
82+
from flag_gems.ops.fill import (fill_scalar, fill_scalar_, fill_scalar_out,
83+
fill_tensor, fill_tensor_, fill_tensor_out)
11584
from flag_gems.ops.flip import flip
11685
from flag_gems.ops.floor_ import floor_
11786
from flag_gems.ops.fmin import fmin, fmin_out
@@ -142,7 +111,8 @@
142111
from flag_gems.ops.kron import kron
143112
from flag_gems.ops.layernorm import layer_norm, layer_norm_backward
144113
from flag_gems.ops.le import le, le_scalar
145-
from flag_gems.ops.lerp import lerp_scalar, lerp_scalar_, lerp_tensor, lerp_tensor_
114+
from flag_gems.ops.lerp import (lerp_scalar, lerp_scalar_, lerp_tensor,
115+
lerp_tensor_)
146116
from flag_gems.ops.lift_fresh_copy import lift_fresh_copy, lift_fresh_copy_out
147117
from flag_gems.ops.linspace import linspace
148118
from flag_gems.ops.log import log
@@ -162,10 +132,8 @@
162132
from flag_gems.ops.masked_scatter import masked_scatter, masked_scatter_
163133
from flag_gems.ops.masked_select import masked_select
164134
from flag_gems.ops.max import max, max_dim
165-
from flag_gems.ops.max_pool2d_with_indices import (
166-
max_pool2d_backward,
167-
max_pool2d_with_indices,
168-
)
135+
from flag_gems.ops.max_pool2d_with_indices import (max_pool2d_backward,
136+
max_pool2d_with_indices)
169137
from flag_gems.ops.maximum import maximum
170138
from flag_gems.ops.mean import mean, mean_dim
171139
from flag_gems.ops.min import min, min_dim
@@ -179,36 +147,22 @@
179147
from flag_gems.ops.ne import ne, ne_scalar
180148
from flag_gems.ops.neg import neg, neg_
181149
from flag_gems.ops.nll_loss_nd import nll_loss_nd_backward, nll_loss_nd_forward
182-
from flag_gems.ops.nllloss import (
183-
nll_loss2d_backward,
184-
nll_loss2d_forward,
185-
nll_loss_backward,
186-
nll_loss_forward,
187-
)
150+
from flag_gems.ops.nllloss import (nll_loss2d_backward, nll_loss2d_forward,
151+
nll_loss_backward, nll_loss_forward)
188152
from flag_gems.ops.nonzero import nonzero
189-
from flag_gems.ops.normal import (
190-
normal_,
191-
normal_float_tensor,
192-
normal_tensor_float,
193-
normal_tensor_tensor,
194-
)
153+
from flag_gems.ops.normal import (normal_, normal_float_tensor,
154+
normal_tensor_float, normal_tensor_tensor)
195155
from flag_gems.ops.one_hot import one_hot
196156
from flag_gems.ops.ones import ones
197157
from flag_gems.ops.ones_like import ones_like
198158
from flag_gems.ops.pad import constant_pad_nd, pad
199-
from flag_gems.ops.per_token_group_quant_fp8 import (
200-
SUPPORTED_FP8_DTYPE,
201-
per_token_group_quant_fp8,
202-
)
159+
from flag_gems.ops.per_token_group_quant_fp8 import (SUPPORTED_FP8_DTYPE,
160+
per_token_group_quant_fp8)
203161
from flag_gems.ops.pixel_unshuffle import pixel_unshuffle, pixel_unshuffle_out
204162
from flag_gems.ops.polar import polar
205-
from flag_gems.ops.pow import (
206-
pow_scalar,
207-
pow_tensor_scalar,
208-
pow_tensor_scalar_,
209-
pow_tensor_tensor,
210-
pow_tensor_tensor_,
211-
)
163+
from flag_gems.ops.pow import (pow_scalar, pow_tensor_scalar,
164+
pow_tensor_scalar_, pow_tensor_tensor,
165+
pow_tensor_tensor_)
212166
from flag_gems.ops.prelu import prelu
213167
from flag_gems.ops.prod import prod, prod_dim
214168
from flag_gems.ops.quantile import quantile
@@ -218,24 +172,27 @@
218172
from flag_gems.ops.randn_like import randn_like
219173
from flag_gems.ops.randperm import randperm
220174
from flag_gems.ops.reciprocal import reciprocal, reciprocal_
221-
from flag_gems.ops.reflection_pad1d import reflection_pad1d, reflection_pad1d_out
222-
from flag_gems.ops.reflection_pad2d import reflection_pad2d, reflection_pad2d_out
175+
from flag_gems.ops.reflection_pad1d import (reflection_pad1d,
176+
reflection_pad1d_out)
177+
from flag_gems.ops.reflection_pad2d import (reflection_pad2d,
178+
reflection_pad2d_out)
223179
from flag_gems.ops.relu import relu, relu_
224180
from flag_gems.ops.relu6 import relu6
225181
from flag_gems.ops.repeat import repeat
226-
from flag_gems.ops.repeat_interleave import (
227-
repeat_interleave_self_int,
228-
repeat_interleave_self_tensor,
229-
repeat_interleave_tensor,
230-
)
231-
from flag_gems.ops.replication_pad1d import replication_pad1d, replication_pad1d_out
182+
from flag_gems.ops.repeat_interleave import (repeat_interleave_self_int,
183+
repeat_interleave_self_tensor,
184+
repeat_interleave_tensor)
185+
from flag_gems.ops.replication_pad1d import (replication_pad1d,
186+
replication_pad1d_out)
232187
from flag_gems.ops.replication_pad3d import replication_pad3d
233188
from flag_gems.ops.resolve_conj import resolve_conj
234189
from flag_gems.ops.resolve_neg import resolve_neg
235-
from flag_gems.ops.rms_norm import rms_norm, rms_norm_backward, rms_norm_forward
190+
from flag_gems.ops.rms_norm import (rms_norm, rms_norm_backward,
191+
rms_norm_forward)
236192
from flag_gems.ops.rrelu_with_noise_backward import rrelu_with_noise_backward
237193
from flag_gems.ops.rsqrt import rsqrt, rsqrt_
238-
from flag_gems.ops.scaled_softmax import scaled_softmax_backward, scaled_softmax_forward
194+
from flag_gems.ops.scaled_softmax import (scaled_softmax_backward,
195+
scaled_softmax_forward)
239196
from flag_gems.ops.scatter import scatter, scatter_
240197
from flag_gems.ops.scatter_add_ import scatter_add_
241198
from flag_gems.ops.select_scatter import select_scatter
@@ -248,7 +205,8 @@
248205
from flag_gems.ops.sinh_ import sinh_
249206
from flag_gems.ops.slice_backward import slice_backward
250207
from flag_gems.ops.slice_scatter import slice_scatter
251-
from flag_gems.ops.soft_margin_loss import soft_margin_loss, soft_margin_loss_out
208+
from flag_gems.ops.soft_margin_loss import (soft_margin_loss,
209+
soft_margin_loss_out)
252210
from flag_gems.ops.softmax import softmax, softmax_backward
253211
from flag_gems.ops.softplus import softplus
254212
from flag_gems.ops.softshrink import softshrink, softshrink_out
@@ -284,16 +242,10 @@
284242
from flag_gems.ops.vector_norm import vector_norm
285243
from flag_gems.ops.vstack import vstack
286244
from flag_gems.ops.w8a8_block_fp8_matmul import w8a8_block_fp8_matmul
287-
from flag_gems.ops.weightnorm import (
288-
weight_norm_interface,
289-
weight_norm_interface_backward,
290-
)
291-
from flag_gems.ops.where import (
292-
where_scalar_other,
293-
where_scalar_self,
294-
where_self,
295-
where_self_out,
296-
)
245+
from flag_gems.ops.weightnorm import (weight_norm_interface,
246+
weight_norm_interface_backward)
247+
from flag_gems.ops.where import (where_scalar_other, where_scalar_self,
248+
where_self, where_self_out)
297249
from flag_gems.ops.zero import zero, zero_out
298250
from flag_gems.ops.zeros import zero_, zeros
299251
from flag_gems.ops.zeros_like import zeros_like

src/flag_gems/ops/feature_dropout.py

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,8 @@
55
import triton.language as tl
66

77
from flag_gems.runtime import torch_device_fn
8-
from flag_gems.utils.random_utils import philox_backend_seed_offset, uint_to_uniform_float
8+
from flag_gems.utils.random_utils import (philox_backend_seed_offset,
9+
uint_to_uniform_float)
910

1011
logger = logging.getLogger(__name__)
1112

@@ -129,7 +130,9 @@ def feature_dropout(input, p, train=True):
129130
return torch.zeros_like(input)
130131

131132
if input.ndim < 2:
132-
raise RuntimeError("Feature dropout requires at least 2 dimensions in the input")
133+
raise RuntimeError(
134+
"Feature dropout requires at least 2 dimensions in the input"
135+
)
133136

134137
assert 0.0 < p < 1.0, "p must be in (0, 1)"
135138

@@ -191,7 +194,9 @@ def feature_dropout_(input, p, train=True):
191194
return input
192195

193196
if input.ndim < 2:
194-
raise RuntimeError("Feature dropout requires at least 2 dimensions in the input")
197+
raise RuntimeError(
198+
"Feature dropout requires at least 2 dimensions in the input"
199+
)
195200

196201
assert 0.0 < p < 1.0, "p must be in (0, 1)"
197202

0 commit comments

Comments
 (0)