forked from flagos-ai/FlagGems
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy path__init__.py
More file actions
807 lines (806 loc) · 22.6 KB
/
Copy path__init__.py
File metadata and controls
807 lines (806 loc) · 22.6 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
from flag_gems.ops._functional_sym_constrain_range_for_size import (
_functional_sym_constrain_range_for_size,
)
from flag_gems.ops._is_all_true import _is_all_true
from flag_gems.ops._safe_softmax import _safe_softmax
from flag_gems.ops._upsample_nearest_exact1d import _upsample_nearest_exact1d
from flag_gems.ops.abs import abs, abs_
from flag_gems.ops.absolute import absolute
from flag_gems.ops.acos import acos
from flag_gems.ops.act_quant import act_quant_triton
from flag_gems.ops.add import add, add_
from flag_gems.ops.addcdiv import addcdiv, addcdiv_out
from flag_gems.ops.addcmul import addcmul, addcmul_out
from flag_gems.ops.addmm import addmm, addmm_dtype, addmm_dtype_out, addmm_out
from flag_gems.ops.addmv import addmv, addmv_out
from flag_gems.ops.addr import addr
from flag_gems.ops.alias_copy import alias_copy, alias_copy_out
from flag_gems.ops.all import all, all_dim, all_dims
from flag_gems.ops.amax import amax
from flag_gems.ops.aminmax import aminmax
from flag_gems.ops.angle import angle
from flag_gems.ops.any import any, any_dim, any_dims
from flag_gems.ops.arange import arange, arange_start
from flag_gems.ops.arcsinh import arcsinh, arcsinh_out
from flag_gems.ops.arcsinh_ import arcsinh_
from flag_gems.ops.arctanh_ import arctanh_
from flag_gems.ops.argmax import argmax
from flag_gems.ops.argmin import argmin
from flag_gems.ops.asinh import asinh, asinh_out
from flag_gems.ops.asinh_ import asinh_
from flag_gems.ops.assert_async import _assert_async
from flag_gems.ops.atan import atan, atan_
from flag_gems.ops.atan2 import atan2, atan2_out
from flag_gems.ops.attention import (
ScaleDotProductAttention,
flash_attention_forward,
flash_attn_varlen_func,
flash_attn_varlen_opt_func,
scaled_dot_product_attention,
scaled_dot_product_attention_backward,
scaled_dot_product_attention_forward,
)
from flag_gems.ops.avg_pool2d import avg_pool2d, avg_pool2d_backward
from flag_gems.ops.avg_pool3d import avg_pool3d, avg_pool3d_backward
from flag_gems.ops.baddbmm import baddbmm, baddbmm_out
from flag_gems.ops.batch_norm import batch_norm, batch_norm_backward
from flag_gems.ops.bernoulli_ import bernoulli_
from flag_gems.ops.bitwise_and import (
bitwise_and_scalar,
bitwise_and_scalar_,
bitwise_and_scalar_tensor,
bitwise_and_tensor,
bitwise_and_tensor_,
)
from flag_gems.ops.bitwise_left_shift import bitwise_left_shift
from flag_gems.ops.bitwise_not import bitwise_not, bitwise_not_
from flag_gems.ops.bitwise_or import (
bitwise_or_scalar,
bitwise_or_scalar_,
bitwise_or_scalar_tensor,
bitwise_or_tensor,
bitwise_or_tensor_,
)
from flag_gems.ops.bitwise_right_shift import bitwise_right_shift
from flag_gems.ops.bmm import bmm, bmm_out
from flag_gems.ops.cat import cat, cat_out
from flag_gems.ops.ceil import ceil, ceil_, ceil_out
from flag_gems.ops.celu import celu, celu_
from flag_gems.ops.clamp import (
clamp,
clamp_,
clamp_min,
clamp_min_,
clamp_tensor,
clamp_tensor_,
)
from flag_gems.ops.clip import clip, clip_
from flag_gems.ops.conj_physical import conj_physical
from flag_gems.ops.contiguous import contiguous
from flag_gems.ops.conv1d import conv1d
from flag_gems.ops.conv2d import conv2d
from flag_gems.ops.conv3d import conv3d
from flag_gems.ops.conv_depthwise2d import _conv_depthwise2d
from flag_gems.ops.copy import copy, copy_
from flag_gems.ops.copysign import copysign, copysign_out
from flag_gems.ops.cos import cos, cos_
from flag_gems.ops.cosh import cosh, cosh_, cosh_out
from flag_gems.ops.count_nonzero import count_nonzero
from flag_gems.ops.cummax import cummax
from flag_gems.ops.cummin import cummin
from flag_gems.ops.cumsum import cumsum, cumsum_out, normed_cumsum
from flag_gems.ops.diag import diag
from flag_gems.ops.diag_embed import diag_embed
from flag_gems.ops.diagonal import diagonal_backward
from flag_gems.ops.digamma_ import digamma_
from flag_gems.ops.div import (
div_mode,
div_mode_,
floor_divide,
floor_divide_,
remainder,
remainder_,
true_divide,
true_divide_,
true_divide_out,
)
from flag_gems.ops.dot import dot
from flag_gems.ops.dropout import dropout, dropout_backward
from flag_gems.ops.einsum import einsum
from flag_gems.ops.elu import elu, elu_, elu_backward
from flag_gems.ops.embedding import embedding, embedding_backward
from flag_gems.ops.embedding_dense_backward import embedding_dense_backward
from flag_gems.ops.eq import eq, eq_scalar, equal
from flag_gems.ops.erf import erf, erf_
from flag_gems.ops.exp import exp, exp_, exp_out
from flag_gems.ops.exp2 import exp2, exp2_
from flag_gems.ops.expm1 import expm1, expm1_, expm1_out
from flag_gems.ops.exponential_ import exponential_
from flag_gems.ops.eye import eye
from flag_gems.ops.eye_m import eye_m
from flag_gems.ops.fill import (
fill_scalar,
fill_scalar_,
fill_scalar_out,
fill_tensor,
fill_tensor_,
fill_tensor_out,
)
from flag_gems.ops.flip import flip
from flag_gems.ops.floor_ import floor_
from flag_gems.ops.fmin import fmin, fmin_out
from flag_gems.ops.fp8_matmul import fp8_matmul
from flag_gems.ops.full import full
from flag_gems.ops.full_like import full_like
from flag_gems.ops.gather import gather, gather_backward
from flag_gems.ops.gcd import gcd, gcd_out
from flag_gems.ops.ge import ge, ge_scalar
from flag_gems.ops.gelu import gelu, gelu_, gelu_backward
from flag_gems.ops.get_paged_mqa_logits_metadata import get_paged_mqa_logits_metadata
from flag_gems.ops.get_scheduler_metadata import get_scheduler_metadata
from flag_gems.ops.glu import glu, glu_backward
from flag_gems.ops.greater import (
greater,
greater_out,
greater_scalar,
greater_scalar_out,
)
from flag_gems.ops.grid_sample import grid_sample
from flag_gems.ops.group_gemm import group_mm
from flag_gems.ops.groupnorm import group_norm, group_norm_backward
from flag_gems.ops.gt import gt, gt_scalar
from flag_gems.ops.hadamard_transform import hadamard_transform
from flag_gems.ops.hardsigmoid import hardsigmoid, hardsigmoid_out
from flag_gems.ops.hardswish_ import hardswish_
from flag_gems.ops.hstack import hstack
from flag_gems.ops.hypot import hypot, hypot_out
from flag_gems.ops.i0 import i0, i0_out
from flag_gems.ops.i0_ import i0_
from flag_gems.ops.index import index
from flag_gems.ops.index_add import index_add, index_add_
from flag_gems.ops.index_put import _index_put_impl_, index_put, index_put_
from flag_gems.ops.index_select import index_select
from flag_gems.ops.isclose import allclose, isclose
from flag_gems.ops.isfinite import isfinite
from flag_gems.ops.isin import isin
from flag_gems.ops.isinf import isinf
from flag_gems.ops.isnan import isnan
from flag_gems.ops.isneginf import isneginf, isneginf_out
from flag_gems.ops.kron import kron
from flag_gems.ops.layernorm import layer_norm, layer_norm_backward
from flag_gems.ops.le import le, le_scalar
from flag_gems.ops.leaky_relu import leaky_relu, leaky_relu_, leaky_relu_out
from flag_gems.ops.lerp import lerp_scalar, lerp_scalar_, lerp_tensor, lerp_tensor_
from flag_gems.ops.lift_fresh_copy import lift_fresh_copy, lift_fresh_copy_out
from flag_gems.ops.linspace import linspace
from flag_gems.ops.log import log
from flag_gems.ops.log1p_ import log1p_
from flag_gems.ops.log10 import log10, log10_, log10_out
from flag_gems.ops.log_sigmoid import log_sigmoid
from flag_gems.ops.log_softmax import (
log_softmax,
log_softmax_backward,
log_softmax_backward_out,
log_softmax_out,
)
from flag_gems.ops.logaddexp import logaddexp, logaddexp_out
from flag_gems.ops.logical_and import logical_and, logical_and_
from flag_gems.ops.logical_not import logical_not
from flag_gems.ops.logical_or import logical_or, logical_or_
from flag_gems.ops.logical_xor import logical_xor
from flag_gems.ops.logit import logit, logit_out
from flag_gems.ops.logit_ import logit_
from flag_gems.ops.logspace import logspace
from flag_gems.ops.lt import lt, lt_scalar
from flag_gems.ops.margin_ranking_loss import margin_ranking_loss
from flag_gems.ops.masked_fill import masked_fill, masked_fill_
from flag_gems.ops.masked_scatter import masked_scatter, masked_scatter_
from flag_gems.ops.masked_select import masked_select
from flag_gems.ops.max import max, max_dim
from flag_gems.ops.max_pool2d_with_indices import (
max_pool2d_backward,
max_pool2d_with_indices,
)
from flag_gems.ops.max_pool3d_with_indices import (
max_pool3d_backward,
max_pool3d_with_indices,
)
from flag_gems.ops.maximum import maximum
from flag_gems.ops.mean import mean, mean_dim
from flag_gems.ops.min import min, min_dim
from flag_gems.ops.minimum import minimum
from flag_gems.ops.mm import mm, mm_out
from flag_gems.ops.mse_loss import mse_loss
from flag_gems.ops.mul import mul, mul_
from flag_gems.ops.multinomial import multinomial
from flag_gems.ops.mv import mv
from flag_gems.ops.nan_to_num import nan_to_num
from flag_gems.ops.ne import ne, ne_scalar
from flag_gems.ops.neg import neg, neg_
from flag_gems.ops.new_full import new_full
from flag_gems.ops.nll_loss_nd import nll_loss_nd_backward, nll_loss_nd_forward
from flag_gems.ops.nllloss import (
nll_loss2d_backward,
nll_loss2d_forward,
nll_loss_backward,
nll_loss_forward,
)
from flag_gems.ops.nonzero import nonzero
from flag_gems.ops.normal import (
normal_,
normal_float_tensor,
normal_tensor_float,
normal_tensor_tensor,
)
from flag_gems.ops.one_hot import one_hot
from flag_gems.ops.ones import ones
from flag_gems.ops.ones_like import ones_like
from flag_gems.ops.pad import constant_pad_nd, pad
from flag_gems.ops.per_token_group_quant_fp8 import (
SUPPORTED_FP8_DTYPE,
per_token_group_quant_fp8,
)
from flag_gems.ops.pixel_shuffle import pixel_shuffle
from flag_gems.ops.pixel_unshuffle import pixel_unshuffle, pixel_unshuffle_out
from flag_gems.ops.polar import polar
from flag_gems.ops.pow import (
pow_scalar,
pow_tensor_scalar,
pow_tensor_scalar_,
pow_tensor_tensor,
pow_tensor_tensor_,
)
from flag_gems.ops.prelu import prelu
from flag_gems.ops.prod import prod, prod_dim
from flag_gems.ops.quantile import quantile
from flag_gems.ops.rand import rand
from flag_gems.ops.rand_like import rand_like
from flag_gems.ops.randn import randn
from flag_gems.ops.randn_like import randn_like
from flag_gems.ops.randperm import randperm
from flag_gems.ops.reciprocal import reciprocal, reciprocal_
from flag_gems.ops.reflection_pad1d import reflection_pad1d, reflection_pad1d_out
from flag_gems.ops.reflection_pad2d import reflection_pad2d, reflection_pad2d_out
from flag_gems.ops.relu import relu, relu_
from flag_gems.ops.relu6 import relu6
from flag_gems.ops.repeat import repeat
from flag_gems.ops.repeat_interleave import (
repeat_interleave_self_int,
repeat_interleave_self_tensor,
repeat_interleave_tensor,
)
from flag_gems.ops.replication_pad1d import replication_pad1d, replication_pad1d_out
from flag_gems.ops.replication_pad3d import replication_pad3d
from flag_gems.ops.resolve_conj import resolve_conj
from flag_gems.ops.resolve_neg import resolve_neg
from flag_gems.ops.rms_norm import rms_norm, rms_norm_backward, rms_norm_forward
from flag_gems.ops.roll import roll
from flag_gems.ops.round import round, round_, round_out
from flag_gems.ops.rrelu_with_noise_backward import rrelu_with_noise_backward
from flag_gems.ops.rsqrt import rsqrt, rsqrt_
from flag_gems.ops.scaled_softmax import scaled_softmax_backward, scaled_softmax_forward
from flag_gems.ops.scatter import scatter, scatter_
from flag_gems.ops.scatter_add_ import scatter_add_
from flag_gems.ops.scatter_reduce_ import scatter_reduce_
from flag_gems.ops.select_backward import select_backward
from flag_gems.ops.select_scatter import select_scatter
from flag_gems.ops.selu import selu
from flag_gems.ops.selu_ import selu_
from flag_gems.ops.sgn_ import sgn_
from flag_gems.ops.sigmoid import sigmoid, sigmoid_, sigmoid_backward
from flag_gems.ops.signbit import signbit, signbit_out
from flag_gems.ops.silu import silu, silu_, silu_backward
from flag_gems.ops.sin import sin, sin_
from flag_gems.ops.sinh_ import sinh_
from flag_gems.ops.slice_backward import slice_backward
from flag_gems.ops.slice_scatter import slice_scatter
from flag_gems.ops.soft_margin_loss import soft_margin_loss, soft_margin_loss_out
from flag_gems.ops.softmax import (
softmax,
softmax_backward,
softmax_backward_out,
softmax_out,
)
from flag_gems.ops.softplus import softplus
from flag_gems.ops.softshrink import softshrink, softshrink_out
from flag_gems.ops.sort import sort, sort_stable
from flag_gems.ops.special_i0e import special_i0e, special_i0e_out
from flag_gems.ops.special_i1 import special_i1, special_i1_out
from flag_gems.ops.sqrt import sqrt, sqrt_
from flag_gems.ops.square import square, square_, square_out
from flag_gems.ops.stack import stack
from flag_gems.ops.std import std
from flag_gems.ops.sub import sub, sub_
from flag_gems.ops.sum import sum, sum_dim, sum_dim_out, sum_out
from flag_gems.ops.t_copy import t_copy, t_copy_out
from flag_gems.ops.tan import tan, tan_
from flag_gems.ops.tanh import tanh, tanh_, tanh_backward
from flag_gems.ops.threshold import threshold, threshold_backward
from flag_gems.ops.tile import tile
from flag_gems.ops.to import to_copy
from flag_gems.ops.topk import topk
from flag_gems.ops.trace import trace
from flag_gems.ops.tril import tril, tril_out
from flag_gems.ops.triu import triu, triu_
from flag_gems.ops.unfold_backward import unfold_backward
from flag_gems.ops.uniform import uniform_
from flag_gems.ops.unique import _unique2
from flag_gems.ops.unique_consecutive import unique_consecutive
from flag_gems.ops.upsample_bicubic2d import upsample_bicubic2d
from flag_gems.ops.upsample_bicubic2d_aa import _upsample_bicubic2d_aa
from flag_gems.ops.upsample_bicubic2d_aa_backward import _upsample_bicubic2d_aa_backward
from flag_gems.ops.upsample_linear1d import upsample_linear1d
from flag_gems.ops.upsample_nearest1d import upsample_nearest1d
from flag_gems.ops.upsample_nearest2d import upsample_nearest2d
from flag_gems.ops.upsample_nearest3d import upsample_nearest3d
from flag_gems.ops.var import var, var_correction, var_dim
from flag_gems.ops.var_mean import var_mean
from flag_gems.ops.vdot import vdot
from flag_gems.ops.vector_norm import vector_norm
from flag_gems.ops.vstack import vstack
from flag_gems.ops.w8a8_block_fp8_matmul import w8a8_block_fp8_matmul
from flag_gems.ops.weightnorm import (
weight_norm_interface,
weight_norm_interface_backward,
)
from flag_gems.ops.where import (
where_scalar_other,
where_scalar_self,
where_self,
where_self_out,
)
from flag_gems.ops.zero import zero, zero_out
from flag_gems.ops.zeros import zero_, zeros
from flag_gems.ops.zeros_like import zeros_like
__all__ = [
"_assert_async",
"_conv_depthwise2d",
"_functional_sym_constrain_range_for_size",
"_index_put_impl_",
"_is_all_true",
"_safe_softmax",
"_unique2",
"_upsample_bicubic2d_aa",
"_upsample_bicubic2d_aa_backward",
"_upsample_nearest_exact1d",
"abs",
"abs_",
"absolute",
"act_quant_triton",
"acos",
"add",
"add_",
"addcdiv",
"addcdiv_out",
"addcmul",
"addcmul_out",
"addmm",
"addmm_dtype",
"addmm_dtype_out",
"addmm_out",
"addmv",
"addmv_out",
"addr",
"alias_copy",
"alias_copy_out",
"all",
"all_dim",
"all_dims",
"allclose",
"amax",
"aminmax",
"angle",
"any",
"any_dim",
"any_dims",
"arange",
"arange_start",
"arcsinh",
"arcsinh_out",
"arctanh_",
"arcsinh_",
"argmax",
"argmin",
"asinh",
"asinh_",
"asinh_out",
"atan",
"atan_",
"atan2",
"atan2_out",
"avg_pool2d",
"avg_pool2d_backward",
"avg_pool3d",
"avg_pool3d_backward",
"baddbmm",
"baddbmm_out",
"batch_norm",
"batch_norm_backward",
"bernoulli_",
"bitwise_and_scalar",
"bitwise_and_scalar_",
"bitwise_and_scalar_tensor",
"bitwise_and_tensor",
"bitwise_and_tensor_",
"bitwise_left_shift",
"bitwise_not",
"bitwise_not_",
"bitwise_or_scalar",
"bitwise_or_scalar_",
"bitwise_or_scalar_tensor",
"bitwise_or_tensor",
"bitwise_or_tensor_",
"bitwise_right_shift",
"bmm",
"bmm_out",
"cat",
"cat_out",
"ceil",
"ceil_",
"ceil_out",
"celu",
"celu_",
"clamp",
"clamp_",
"clamp_min",
"clamp_min_",
"clamp_tensor",
"clamp_tensor_",
"clip",
"clip_",
"constant_pad_nd",
"contiguous",
"conv1d",
"conv2d",
"conv3d",
"copy",
"copy_",
"copysign",
"copysign_out",
"cos",
"cos_",
"cosh",
"cosh_",
"cosh_out",
"count_nonzero",
"cummax",
"cummin",
"cumsum",
"cumsum_out",
"conj_physical",
"diag",
"diag_embed",
"diagonal_backward",
"digamma_",
"div_mode",
"div_mode_",
"dot",
"dropout",
"dropout_backward",
"einsum",
"elu",
"elu_",
"elu_backward",
"embedding",
"embedding_backward",
"embedding_dense_backward",
"eq",
"eq_scalar",
"equal",
"erf",
"erf_",
"exp",
"exp_",
"exp_out",
"exp2",
"exp2_",
"expm1",
"expm1_",
"expm1_out",
"exponential_",
"eye",
"eye_m",
"fill_scalar",
"fill_scalar_",
"fill_scalar_out",
"fill_tensor",
"fill_tensor_",
"fill_tensor_out",
"flash_attention_forward",
"flash_attn_varlen_func",
"flash_attn_varlen_opt_func",
"flip",
"floor_",
"floor_divide",
"floor_divide_",
"fmin",
"fmin_out",
"full",
"full_like",
"gather",
"gather_backward",
"gcd",
"gcd_out",
"ge",
"ge_scalar",
"gelu",
"gelu_",
"gelu_backward",
"get_paged_mqa_logits_metadata",
"get_scheduler_metadata",
"glu",
"glu_backward",
"grid_sample",
"greater",
"greater_out",
"greater_scalar",
"greater_scalar_out",
"group_mm",
"group_norm",
"group_norm_backward",
"gt",
"gt_scalar",
"hadamard_transform",
"hardsigmoid",
"hardsigmoid_out",
"hardswish_",
"hstack",
"hypot",
"hypot_out",
"i0",
"i0_out",
"i0_",
"index",
"index_add",
"index_add_",
"index_put",
"index_put_",
"index_select",
"isclose",
"isfinite",
"isin",
"isinf",
"isnan",
"isneginf",
"isneginf_out",
"kron",
"layer_norm",
"layer_norm_backward",
"leaky_relu",
"leaky_relu_",
"leaky_relu_out",
"le",
"le_scalar",
"lerp_scalar",
"lerp_scalar_",
"lerp_tensor",
"lerp_tensor_",
"lift_fresh_copy",
"lift_fresh_copy_out",
"linspace",
"log",
"log10",
"log10_",
"log10_out",
"log_sigmoid",
"log_softmax",
"log_softmax_backward",
"log_softmax_backward_out",
"log_softmax_out",
"log1p_",
"logaddexp",
"logaddexp_out",
"logical_and",
"logical_and_",
"logical_not",
"logical_or",
"logical_or_",
"logical_xor",
"logit",
"logit_out",
"logit_",
"logspace",
"lt",
"lt_scalar",
"margin_ranking_loss",
"masked_fill",
"masked_fill_",
"masked_scatter",
"masked_scatter_",
"masked_select",
"max",
"max_dim",
"max_pool2d_with_indices",
"max_pool2d_backward",
"max_pool3d_with_indices",
"max_pool3d_backward",
"maximum",
"mean",
"mean_dim",
"min",
"min_dim",
"minimum",
"mm",
"mm_out",
"mse_loss",
"mul",
"mul_",
"multinomial",
"mv",
"nan_to_num",
"ne",
"ne_scalar",
"neg",
"neg_",
"new_full",
"nll_loss_backward",
"nll_loss_forward",
"nll_loss2d_backward",
"nll_loss2d_forward",
"nll_loss_nd_forward",
"nll_loss_nd_backward",
"nonzero",
"normal_float_tensor",
"normal_tensor_float",
"normal_tensor_tensor",
"normal_",
"normed_cumsum",
"ones",
"ones_like",
"one_hot",
"pad",
"per_token_group_quant_fp8",
"pixel_shuffle",
"pixel_unshuffle",
"pixel_unshuffle_out",
"polar",
"pow_scalar",
"pow_tensor_scalar",
"pow_tensor_scalar_",
"pow_tensor_tensor",
"pow_tensor_tensor_",
"prelu",
"prod",
"prod_dim",
"quantile",
"rand",
"rand_like",
"randn",
"randn_like",
"randperm",
"reciprocal",
"reciprocal_",
"reflection_pad2d",
"reflection_pad2d_out",
"reflection_pad1d",
"reflection_pad1d_out",
"relu",
"relu_",
"relu6",
"remainder",
"remainder_",
"repeat",
"repeat_interleave_self_int",
"repeat_interleave_self_tensor",
"repeat_interleave_tensor",
"replication_pad1d",
"replication_pad1d_out",
"replication_pad3d",
"resolve_conj",
"resolve_neg",
"rms_norm",
"rms_norm_backward",
"rms_norm_forward",
"roll",
"round",
"round_",
"round_out",
"rrelu_with_noise_backward",
"rsqrt",
"rsqrt_",
"scaled_dot_product_attention",
"scaled_dot_product_attention_backward",
"scaled_dot_product_attention_forward",
"scaled_softmax_backward",
"scaled_softmax_forward",
"scatter",
"scatter_",
"scatter_add_",
"scatter_reduce_",
"select_backward",
"select_scatter",
"selu",
"selu_",
"sgn_",
"sigmoid",
"sigmoid_",
"sigmoid_backward",
"signbit",
"signbit_out",
"silu",
"silu_",
"silu_backward",
"sin",
"sin_",
"sinh_",
"slice_backward",
"slice_scatter",
"soft_margin_loss",
"soft_margin_loss_out",
"softmax",
"softmax_backward",
"softmax_backward_out",
"softmax_out",
"softplus",
"softshrink",
"softshrink_out",
"sort",
"sort_stable",
"special_i1",
"special_i1_out",
"special_i0e",
"special_i0e_out",
"sqrt",
"sqrt_",
"square",
"square_",
"square_out",
"stack",
"std",
"sub",
"sub_",
"sum",
"sum_dim",
"sum_dim_out",
"sum_out",
"ScaleDotProductAttention",
"SUPPORTED_FP8_DTYPE",
"t_copy",
"t_copy_out",
"tan",
"tan_",
"tanh",
"tanh_",
"tanh_backward",
"threshold",
"threshold_backward",
"tile",
"to_copy",
"topk",
"trace",
"tril",
"tril_out",
"triu",
"triu_",
"true_divide",
"true_divide_",
"true_divide_out",
"unfold_backward",
"uniform_",
"unique_consecutive",
"upsample_bicubic2d",
"upsample_linear1d",
"upsample_nearest1d",
"upsample_nearest2d",
"upsample_nearest3d",
"var_mean",
"var",
"var_correction",
"var_dim",
"vdot",
"vector_norm",
"vstack",
"fp8_matmul",
"w8a8_block_fp8_matmul",
"weight_norm_interface",
"weight_norm_interface_backward",
"where_scalar_other",
"where_scalar_self",
"where_self",
"where_self_out",
"zero",
"zero_out",
"zero_",
"zeros",
"zeros_like",
]