Skip to content

Commit e6c7331

Browse files
committed
add addcdiv.out, addcmul.out, baddbmm.out, cat.out etc; refactor op name according to native_functions.yaml
1 parent 5708f53 commit e6c7331

52 files changed

Lines changed: 617 additions & 270 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

benchmark/test_attention_perf.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -46,7 +46,7 @@ def torch_flash_attention_forward(
4646
def gems_flash_attention_forward(
4747
q, k, v, scale, is_causal, dropout_p=0.0, return_debug_mask=False, **extra_kwargs
4848
):
49-
return flag_gems.ops.flash_attention_forward(
49+
return flag_gems.ops._flash_attention_forward(
5050
q,
5151
k,
5252
v,

benchmark/test_blas_perf.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -360,7 +360,7 @@ def test_blas_benchmark(op_name, torch_op, input_fn, bench_cls):
360360
)
361361

362362
if op_name == "groupmm":
363-
gems_op = flag_gems.group_mm
363+
gems_op = flag_gems._grouped_mm
364364
bench.set_gems(gems_op)
365365

366366
bench.run()

src/flag_gems/__init__.py

Lines changed: 35 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -29,40 +29,48 @@ def torch_ge(v):
2929

3030
_FULL_CONFIG = (
3131
("_assert_async", _assert_async),
32-
("_flash_attention_forward", flash_attention_forward),
32+
("_flash_attention_forward", _flash_attention_forward),
3333
(
3434
"_functional_sym_constrain_range_for_size",
3535
_functional_sym_constrain_range_for_size,
3636
),
37-
("_grouped_mm", group_mm),
38-
("_log_softmax", log_softmax),
39-
("_log_softmax_backward_data", log_softmax_backward),
37+
("_grouped_mm", _grouped_mm),
38+
("_log_softmax", _log_softmax),
39+
("_log_softmax.out", _log_softmax_out),
40+
("_log_softmax_backward_data", _log_softmax_backward_data),
41+
("_log_softmax_backward_data.out", _log_softmax_backward_data_out),
4042
("_safe_softmax", _safe_softmax),
41-
("_softmax", softmax),
42-
("_softmax_backward_data", softmax_backward),
43+
("_softmax", _softmax),
44+
("_softmax.out", _softmax_out),
45+
("_softmax_backward_data", _softmax_backward_data),
46+
("_softmax_backward_data.out", _softmax_backward_data_out),
4347
(
4448
"_to_copy",
45-
to_copy,
49+
_to_copy,
4650
lambda: version.parse(torch.__version__) >= version.parse("2.4"),
4751
),
4852
("_unique2", _unique2),
4953
("_upsample_bicubic2d_aa", _upsample_bicubic2d_aa),
5054
("_upsample_bicubic2d_aa_backward", _upsample_bicubic2d_aa_backward),
5155
("_upsample_nearest_exact1d", _upsample_nearest_exact1d),
52-
("_weight_norm_interface", weight_norm_interface),
53-
("_weight_norm_interface_backward", weight_norm_interface_backward),
56+
("_weight_norm_interface", _weight_norm_interface),
57+
("_weight_norm_interface_backward", _weight_norm_interface_backward),
5458
("abs", abs),
5559
("abs_", abs_),
5660
("absolute", absolute),
5761
("acos", acos),
5862
("add.Tensor", add),
5963
("add_.Tensor", add_),
6064
("addcdiv", addcdiv),
65+
("addcdiv.out", addcdiv_out),
6166
("addcmul", addcmul),
67+
("addcmul.out", addcmul_out),
6268
("addmv", addmv),
6369
("addmv.out", addmv_out),
6470
("addmm", addmm),
6571
("addmm.out", addmm_out),
72+
("addmm.dtype", addmm_dtype),
73+
("addmm.dtype_out", addmm_dtype_out),
6674
("addr", addr),
6775
("alias_copy", alias_copy),
6876
("all", all),
@@ -92,6 +100,7 @@ def torch_ge(v):
92100
("avg_pool2d", avg_pool2d),
93101
("avg_pool2d_backward", avg_pool2d_backward),
94102
("baddbmm", baddbmm),
103+
("baddbmm.out", baddbmm_out),
95104
("bincount", bincount),
96105
("bitwise_and.Scalar", bitwise_and_scalar),
97106
("bitwise_and.Scalar_Tensor", bitwise_and_scalar_tensor),
@@ -110,6 +119,7 @@ def torch_ge(v):
110119
("bmm", bmm),
111120
("bmm.out", bmm_out),
112121
("cat", cat),
122+
("cat.out", cat_out),
113123
("celu", celu),
114124
("celu_", celu_),
115125
("ceil", ceil),
@@ -364,8 +374,8 @@ def torch_ge(v):
364374
("resolve_neg", resolve_neg),
365375
("rms_norm", rms_norm),
366376
("round", round),
367-
("round.out", round_out),
368377
("round_", round_),
378+
("round.out", round_out),
369379
("rrelu_with_noise_backward", rrelu_with_noise_backward),
370380
("rsqrt", rsqrt),
371381
("rsqrt_", rsqrt_),
@@ -463,6 +473,21 @@ def torch_ge(v):
463473
func_name = fn.__name__ if hasattr(fn, "__name__") else str(fn)
464474
FULL_CONFIG_BY_FUNC.setdefault(func_name, []).append(_item)
465475

476+
# Friendly names for only_enable(include=[...]) when the registered impl is *.out
477+
for _alias, _target in (
478+
("softmax", "_softmax_out"),
479+
("softmax_backward", "_softmax_backward_data_out"),
480+
("log_softmax", "_log_softmax_out"),
481+
("log_softmax_backward", "_log_softmax_backward_data_out"),
482+
("flash_attention_forward", "_flash_attention_forward"),
483+
("group_mm", "_grouped_mm"),
484+
("to_copy", "_to_copy"),
485+
("weight_norm_interface", "_weight_norm_interface"),
486+
("weight_norm_interface_backward", "_weight_norm_interface_backward"),
487+
):
488+
if _target in FULL_CONFIG_BY_FUNC:
489+
FULL_CONFIG_BY_FUNC.setdefault(_alias, []).extend(FULL_CONFIG_BY_FUNC[_target])
490+
466491

467492
def enable(
468493
lib=aten_lib,

src/flag_gems/fused/weight_norm.py

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

88
from flag_gems import runtime
9-
from flag_gems.ops import weight_norm_interface, weight_norm_interface_backward
9+
from flag_gems.ops import _weight_norm_interface, _weight_norm_interface_backward
1010
from flag_gems.runtime import torch_device_fn
1111
from flag_gems.utils import libentry
1212
from flag_gems.utils import triton_lang_extension as tle
@@ -195,7 +195,7 @@ def forward(ctx, v, g, dim=0):
195195
dim = dim % v.ndim
196196
can_use_fused = dim == 0 or dim == v.ndim - 1
197197
if can_use_fused:
198-
output, norm = weight_norm_interface(v, g, dim)
198+
output, norm = _weight_norm_interface(v, g, dim)
199199
else:
200200
output, norm = weight_norm_except_dim(v, g, dim)
201201
ctx.save_for_backward(v, g, norm)
@@ -209,7 +209,7 @@ def backward(ctx, grad):
209209
v, g, norm = ctx.saved_tensors
210210
dim = ctx.dim
211211
if ctx.can_use_fused:
212-
v_grad, g_grad = weight_norm_interface_backward(grad, v, g, norm, dim)
212+
v_grad, g_grad = _weight_norm_interface_backward(grad, v, g, norm, dim)
213213
else:
214214
v_grad, g_grad = weight_norm_except_dim_backward(grad, v, g, norm, dim)
215215
return v_grad, g_grad, None

src/flag_gems/ops/__init__.py

Lines changed: 41 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -7,9 +7,9 @@
77
from flag_gems.ops.absolute import absolute
88
from flag_gems.ops.acos import acos
99
from flag_gems.ops.add import add, add_
10-
from flag_gems.ops.addcdiv import addcdiv
11-
from flag_gems.ops.addcmul import addcmul
12-
from flag_gems.ops.addmm import addmm, addmm_out
10+
from flag_gems.ops.addcdiv import addcdiv, addcdiv_out
11+
from flag_gems.ops.addcmul import addcmul, addcmul_out
12+
from flag_gems.ops.addmm import addmm, addmm_dtype, addmm_dtype_out, addmm_out
1313
from flag_gems.ops.addmv import addmv, addmv_out
1414
from flag_gems.ops.addr import addr
1515
from flag_gems.ops.alias_copy import alias_copy, alias_copy_out
@@ -30,15 +30,15 @@
3030
from flag_gems.ops.atan2 import atan2, atan2_out
3131
from flag_gems.ops.attention import (
3232
ScaleDotProductAttention,
33-
flash_attention_forward,
33+
_flash_attention_forward,
3434
flash_attn_varlen_func,
3535
flash_attn_varlen_opt_func,
3636
scaled_dot_product_attention,
3737
scaled_dot_product_attention_backward,
3838
scaled_dot_product_attention_forward,
3939
)
4040
from flag_gems.ops.avg_pool2d import avg_pool2d, avg_pool2d_backward
41-
from flag_gems.ops.baddbmm import baddbmm
41+
from flag_gems.ops.baddbmm import baddbmm, baddbmm_out
4242
from flag_gems.ops.batch_norm import batch_norm, batch_norm_backward
4343
from flag_gems.ops.bitwise_and import (
4444
bitwise_and_scalar,
@@ -58,7 +58,7 @@
5858
)
5959
from flag_gems.ops.bitwise_right_shift import bitwise_right_shift
6060
from flag_gems.ops.bmm import bmm, bmm_out
61-
from flag_gems.ops.cat import cat
61+
from flag_gems.ops.cat import cat, cat_out
6262
from flag_gems.ops.ceil import ceil, ceil_, ceil_out
6363
from flag_gems.ops.celu import celu, celu_
6464
from flag_gems.ops.clamp import (
@@ -135,7 +135,7 @@
135135
greater_scalar,
136136
greater_scalar_out,
137137
)
138-
from flag_gems.ops.group_gemm import group_mm
138+
from flag_gems.ops.group_gemm import _grouped_mm
139139
from flag_gems.ops.groupnorm import group_norm, group_norm_backward
140140
from flag_gems.ops.gt import gt, gt_scalar
141141
from flag_gems.ops.hardsigmoid import hardsigmoid, hardsigmoid_out
@@ -163,7 +163,12 @@
163163
from flag_gems.ops.log import log
164164
from flag_gems.ops.log1p_ import log1p_
165165
from flag_gems.ops.log_sigmoid import log_sigmoid
166-
from flag_gems.ops.log_softmax import log_softmax, log_softmax_backward
166+
from flag_gems.ops.log_softmax import (
167+
_log_softmax,
168+
_log_softmax_backward_data,
169+
_log_softmax_backward_data_out,
170+
_log_softmax_out,
171+
)
167172
from flag_gems.ops.logaddexp import logaddexp, logaddexp_out
168173
from flag_gems.ops.logical_and import logical_and, logical_and_
169174
from flag_gems.ops.logical_not import logical_not
@@ -269,7 +274,12 @@
269274
from flag_gems.ops.slice_backward import slice_backward
270275
from flag_gems.ops.slice_scatter import slice_scatter
271276
from flag_gems.ops.soft_margin_loss import soft_margin_loss, soft_margin_loss_out
272-
from flag_gems.ops.softmax import softmax, softmax_backward
277+
from flag_gems.ops.softmax import (
278+
_softmax,
279+
_softmax_backward_data,
280+
_softmax_backward_data_out,
281+
_softmax_out,
282+
)
273283
from flag_gems.ops.softplus import softplus
274284
from flag_gems.ops.softshrink import softshrink, softshrink_out
275285
from flag_gems.ops.sort import sort, sort_stable
@@ -286,7 +296,7 @@
286296
from flag_gems.ops.tanh import tanh, tanh_, tanh_backward
287297
from flag_gems.ops.threshold import threshold, threshold_backward
288298
from flag_gems.ops.tile import tile
289-
from flag_gems.ops.to import to_copy
299+
from flag_gems.ops.to import _to_copy
290300
from flag_gems.ops.topk import topk
291301
from flag_gems.ops.trace import trace
292302
from flag_gems.ops.tril import tril, tril_out
@@ -307,8 +317,8 @@
307317
from flag_gems.ops.vstack import vstack
308318
from flag_gems.ops.w8a8_block_fp8_matmul import w8a8_block_fp8_matmul
309319
from flag_gems.ops.weightnorm import (
310-
weight_norm_interface,
311-
weight_norm_interface_backward,
320+
_weight_norm_interface,
321+
_weight_norm_interface_backward,
312322
)
313323
from flag_gems.ops.where import (
314324
where_scalar_other,
@@ -336,8 +346,12 @@
336346
"add",
337347
"add_",
338348
"addcdiv",
349+
"addcdiv_out",
339350
"addcmul",
351+
"addcmul_out",
340352
"addmm",
353+
"addmm_dtype",
354+
"addmm_dtype_out",
341355
"addmm_out",
342356
"addmv",
343357
"addmv_out",
@@ -370,6 +384,7 @@
370384
"avg_pool2d",
371385
"avg_pool2d_backward",
372386
"baddbmm",
387+
"baddbmm_out",
373388
"batch_norm",
374389
"batch_norm_backward",
375390
"bitwise_and_scalar",
@@ -389,6 +404,7 @@
389404
"bmm",
390405
"bmm_out",
391406
"cat",
407+
"cat_out",
392408
"ceil",
393409
"ceil_",
394410
"ceil_out",
@@ -454,7 +470,7 @@
454470
"fill_tensor",
455471
"fill_tensor_",
456472
"fill_tensor_out",
457-
"flash_attention_forward",
473+
"_flash_attention_forward",
458474
"flash_attn_varlen_func",
459475
"flash_attn_varlen_opt_func",
460476
"flip",
@@ -480,7 +496,7 @@
480496
"greater_out",
481497
"greater_scalar",
482498
"greater_scalar_out",
483-
"group_mm",
499+
"_grouped_mm",
484500
"group_norm",
485501
"group_norm_backward",
486502
"gt",
@@ -521,8 +537,10 @@
521537
"linspace",
522538
"log",
523539
"log_sigmoid",
524-
"log_softmax",
525-
"log_softmax_backward",
540+
"_log_softmax",
541+
"_log_softmax_backward_data",
542+
"_log_softmax_backward_data_out",
543+
"_log_softmax_out",
526544
"log1p_",
527545
"logaddexp",
528546
"logaddexp_out",
@@ -658,8 +676,10 @@
658676
"slice_scatter",
659677
"soft_margin_loss",
660678
"soft_margin_loss_out",
661-
"softmax",
662-
"softmax_backward",
679+
"_softmax",
680+
"_softmax_backward_data",
681+
"_softmax_backward_data_out",
682+
"_softmax_out",
663683
"softplus",
664684
"softshrink",
665685
"softshrink_out",
@@ -694,7 +714,7 @@
694714
"threshold",
695715
"threshold_backward",
696716
"tile",
697-
"to_copy",
717+
"_to_copy",
698718
"topk",
699719
"trace",
700720
"tril",
@@ -716,8 +736,8 @@
716736
"vector_norm",
717737
"vstack",
718738
"w8a8_block_fp8_matmul",
719-
"weight_norm_interface",
720-
"weight_norm_interface_backward",
739+
"_weight_norm_interface",
740+
"_weight_norm_interface_backward",
721741
"where_scalar_other",
722742
"where_scalar_self",
723743
"where_self",

src/flag_gems/ops/addcdiv.py

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -16,12 +16,14 @@ def addcdiv_kernel(x, t1, t2, value):
1616
return x + value * (t1 / t2)
1717

1818

19-
def addcdiv(inp, tensor1, tensor2, value=1.0, out=None):
20-
logger.debug("GEMS ADDCDIV FORWARD")
21-
22-
if out is None:
23-
out = torch.empty_like(inp)
24-
19+
def addcdiv_out(inp, tensor1, tensor2, *, value=1.0, out):
20+
logger.debug("GEMS ADDCDIV_OUT")
2521
addcdiv_kernel(inp, tensor1, tensor2, value, out0=out)
26-
2722
return out
23+
24+
25+
def addcdiv(inp, tensor1, tensor2, value=1.0):
26+
"""Functional entry; CUDA may dispatch here without hitting ``addcdiv.out``."""
27+
logger.debug("GEMS ADDCDIV")
28+
out = torch.empty_like(inp)
29+
return addcdiv_out(inp, tensor1, tensor2, value=value, out=out)

src/flag_gems/ops/addcmul.py

Lines changed: 18 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -16,15 +16,21 @@ def addcmul_forward(x, t1, t2, value):
1616
return x + value * t1 * t2
1717

1818

19-
def addcmul(inp, tensor1, tensor2, *, value=1.0, out=None):
20-
logger.debug("GEMS ADDCMUL FORWARD")
21-
if out is not None:
22-
broadcast_shape = torch.broadcast_shapes(
23-
inp.shape, tensor1.shape, tensor2.shape
24-
)
25-
if list(out.shape) != list(broadcast_shape):
26-
out.resize_(broadcast_shape)
27-
addcmul_forward(inp, tensor1, tensor2, value, out0=out)
28-
return out
29-
else:
30-
return addcmul_forward(inp, tensor1, tensor2, value)
19+
def addcmul_out(inp, tensor1, tensor2, *, value=1.0, out):
20+
logger.debug("GEMS ADDCMUL_OUT")
21+
broadcast_shape = torch.broadcast_shapes(inp.shape, tensor1.shape, tensor2.shape)
22+
if list(out.shape) != list(broadcast_shape):
23+
out.resize_(broadcast_shape)
24+
addcmul_forward(inp, tensor1, tensor2, value, out0=out)
25+
return out
26+
27+
28+
def addcmul(inp, tensor1, tensor2, *, value=1.0):
29+
"""Functional entry; keep alongside ``addcmul.out`` for dispatch coverage."""
30+
logger.debug("GEMS ADDCMUL")
31+
broadcast_shape = torch.broadcast_shapes(inp.shape, tensor1.shape, tensor2.shape)
32+
dtype = torch.promote_types(
33+
inp.dtype, torch.promote_types(tensor1.dtype, tensor2.dtype)
34+
)
35+
out = torch.empty(broadcast_shape, device=inp.device, dtype=dtype)
36+
return addcmul_out(inp, tensor1, tensor2, value=value, out=out)

0 commit comments

Comments
 (0)