|
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 |
4 | 3 | from flag_gems.ops._safe_softmax import _safe_softmax |
5 | 4 | from flag_gems.ops._upsample_nearest_exact1d import _upsample_nearest_exact1d |
6 | 5 | from flag_gems.ops.abs import abs, abs_ |
|
25 | 24 | from flag_gems.ops.argmin import argmin |
26 | 25 | from flag_gems.ops.asinh_ import asinh_ |
27 | 26 | 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) |
36 | 33 | from flag_gems.ops.avg_pool2d import avg_pool2d, avg_pool2d_backward |
37 | 34 | from flag_gems.ops.baddbmm import baddbmm |
38 | 35 | 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_) |
46 | 39 | from flag_gems.ops.bitwise_left_shift import bitwise_left_shift |
47 | 40 | 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_) |
55 | 44 | from flag_gems.ops.bitwise_right_shift import bitwise_right_shift |
56 | 45 | from flag_gems.ops.bmm import bmm, bmm_out |
57 | 46 | from flag_gems.ops.cat import cat |
58 | 47 | from flag_gems.ops.ceil import ceil, ceil_, ceil_out |
59 | 48 | 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_) |
68 | 51 | from flag_gems.ops.contiguous import contiguous |
69 | 52 | from flag_gems.ops.conv1d import conv1d |
70 | 53 | from flag_gems.ops.conv2d import conv2d |
71 | 54 | from flag_gems.ops.conv3d import conv3d |
72 | 55 | from flag_gems.ops.conv_depthwise2d import _conv_depthwise2d |
73 | | -from flag_gems.ops.cudnn_convolution import cudnn_convolution |
74 | 56 | from flag_gems.ops.copy import copy, copy_ |
75 | 57 | from flag_gems.ops.cos import cos, cos_ |
76 | 58 | from flag_gems.ops.count_nonzero import count_nonzero |
| 59 | +from flag_gems.ops.cudnn_convolution import cudnn_convolution |
77 | 60 | from flag_gems.ops.cummax import cummax |
78 | 61 | from flag_gems.ops.cummin import cummin |
79 | 62 | from flag_gems.ops.cumsum import cumsum, cumsum_out, normed_cumsum |
80 | 63 | from flag_gems.ops.diag import diag |
81 | 64 | from flag_gems.ops.diag_embed import diag_embed |
82 | 65 | from flag_gems.ops.diagonal import diagonal_backward |
83 | 66 | from flag_gems.ops.digamma_ import digamma_ |
84 | | -from flag_gems.ops.div import ( |
85 | | - div_mode, |
86 | | - div_mode_, |
87 | | - floor_divide, |
88 | | - floor_divide_, |
89 | | - remainder, |
90 | | - remainder_, |
91 | | - true_divide, |
92 | | - true_divide_, |
93 | | - true_divide_out, |
94 | | -) |
| 67 | +from flag_gems.ops.div import (div_mode, div_mode_, floor_divide, |
| 68 | + floor_divide_, remainder, remainder_, |
| 69 | + true_divide, true_divide_, true_divide_out) |
95 | 70 | from flag_gems.ops.dot import dot |
96 | 71 | from flag_gems.ops.dropout import dropout, dropout_backward |
97 | 72 | from flag_gems.ops.elu import elu, elu_, elu_backward |
|
104 | 79 | from flag_gems.ops.exponential_ import exponential_ |
105 | 80 | from flag_gems.ops.eye import eye |
106 | 81 | 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 | | -) |
| 82 | +from flag_gems.ops.fill import (fill_scalar, fill_scalar_, fill_scalar_out, |
| 83 | + fill_tensor, fill_tensor_, fill_tensor_out) |
115 | 84 | from flag_gems.ops.flip import flip |
116 | 85 | from flag_gems.ops.floor_ import floor_ |
117 | 86 | from flag_gems.ops.fmin import fmin, fmin_out |
|
142 | 111 | from flag_gems.ops.kron import kron |
143 | 112 | from flag_gems.ops.layernorm import layer_norm, layer_norm_backward |
144 | 113 | 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_) |
146 | 116 | from flag_gems.ops.lift_fresh_copy import lift_fresh_copy, lift_fresh_copy_out |
147 | 117 | from flag_gems.ops.linspace import linspace |
148 | 118 | from flag_gems.ops.log import log |
|
162 | 132 | from flag_gems.ops.masked_scatter import masked_scatter, masked_scatter_ |
163 | 133 | from flag_gems.ops.masked_select import masked_select |
164 | 134 | 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) |
169 | 137 | from flag_gems.ops.maximum import maximum |
170 | 138 | from flag_gems.ops.mean import mean, mean_dim |
171 | 139 | from flag_gems.ops.min import min, min_dim |
|
179 | 147 | from flag_gems.ops.ne import ne, ne_scalar |
180 | 148 | from flag_gems.ops.neg import neg, neg_ |
181 | 149 | 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) |
188 | 152 | 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) |
195 | 155 | from flag_gems.ops.one_hot import one_hot |
196 | 156 | from flag_gems.ops.ones import ones |
197 | 157 | from flag_gems.ops.ones_like import ones_like |
198 | 158 | 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) |
203 | 161 | from flag_gems.ops.pixel_unshuffle import pixel_unshuffle, pixel_unshuffle_out |
204 | 162 | 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_) |
212 | 166 | from flag_gems.ops.prelu import prelu |
213 | 167 | from flag_gems.ops.prod import prod, prod_dim |
214 | 168 | from flag_gems.ops.quantile import quantile |
|
218 | 172 | from flag_gems.ops.randn_like import randn_like |
219 | 173 | from flag_gems.ops.randperm import randperm |
220 | 174 | 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) |
223 | 179 | from flag_gems.ops.relu import relu, relu_ |
224 | 180 | from flag_gems.ops.relu6 import relu6 |
225 | 181 | 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) |
232 | 187 | from flag_gems.ops.replication_pad3d import replication_pad3d |
233 | 188 | from flag_gems.ops.resolve_conj import resolve_conj |
234 | 189 | 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) |
236 | 192 | from flag_gems.ops.rrelu_with_noise_backward import rrelu_with_noise_backward |
237 | 193 | 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) |
239 | 196 | from flag_gems.ops.scatter import scatter, scatter_ |
240 | 197 | from flag_gems.ops.scatter_add_ import scatter_add_ |
241 | 198 | from flag_gems.ops.select_scatter import select_scatter |
|
248 | 205 | from flag_gems.ops.sinh_ import sinh_ |
249 | 206 | from flag_gems.ops.slice_backward import slice_backward |
250 | 207 | 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) |
252 | 210 | from flag_gems.ops.softmax import softmax, softmax_backward |
253 | 211 | from flag_gems.ops.softplus import softplus |
254 | 212 | from flag_gems.ops.softshrink import softshrink, softshrink_out |
|
284 | 242 | from flag_gems.ops.vector_norm import vector_norm |
285 | 243 | from flag_gems.ops.vstack import vstack |
286 | 244 | 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) |
297 | 249 | from flag_gems.ops.zero import zero, zero_out |
298 | 250 | from flag_gems.ops.zeros import zero_, zeros |
299 | 251 | from flag_gems.ops.zeros_like import zeros_like |
|
0 commit comments