Skip to content

Commit d8ffede

Browse files
committed
Add rsub operator implementation, tests and benchmark
1 parent c322437 commit d8ffede

5 files changed

Lines changed: 71 additions & 0 deletions

File tree

benchmark/test_binary_pointwise_perf.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -75,6 +75,7 @@ def get_tflops(self, op, *args, **kwargs):
7575
("div", torch.div, FLOAT_DTYPES + COMPLEX_DTYPES),
7676
("mul", torch.mul, FLOAT_DTYPES + COMPLEX_DTYPES),
7777
("sub", torch.sub, FLOAT_DTYPES + COMPLEX_DTYPES),
78+
("rsub", torch.rsub, FLOAT_DTYPES),
7879
("pow", torch.pow, FLOAT_DTYPES),
7980
("polar", torch.polar, [torch.float32]),
8081
("floor_divide", torch.floor_divide, INT_DTYPES),

src/flag_gems/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -377,6 +377,8 @@ def torch_ge(v):
377377
("rrelu_with_noise_backward", rrelu_with_noise_backward),
378378
("rsqrt", rsqrt),
379379
("rsqrt_", rsqrt_),
380+
("rsub.Scalar", rsub_scalar),
381+
("rsub.Tensor", rsub_tensor),
380382
("scaled_softmax_backward", scaled_softmax_backward),
381383
("scaled_softmax_forward", scaled_softmax_forward),
382384
("scatter.reduce", scatter),

src/flag_gems/ops/__init__.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -256,6 +256,7 @@
256256
from flag_gems.ops.round import round, round_, round_out
257257
from flag_gems.ops.rrelu_with_noise_backward import rrelu_with_noise_backward
258258
from flag_gems.ops.rsqrt import rsqrt, rsqrt_
259+
from flag_gems.ops.rsub import rsub_scalar, rsub_tensor
259260
from flag_gems.ops.scaled_softmax import scaled_softmax_backward, scaled_softmax_forward
260261
from flag_gems.ops.scatter import scatter, scatter_
261262
from flag_gems.ops.scatter_add_ import scatter_add_
@@ -640,6 +641,8 @@
640641
"rrelu_with_noise_backward",
641642
"rsqrt",
642643
"rsqrt_",
644+
"rsub_scalar",
645+
"rsub_tensor",
643646
"scaled_dot_product_attention",
644647
"scaled_dot_product_attention_backward",
645648
"scaled_dot_product_attention_forward",

src/flag_gems/ops/rsub.py

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,31 @@
1+
import logging
2+
3+
import triton
4+
5+
from flag_gems.utils import pointwise_dynamic
6+
7+
logger = logging.getLogger(__name__)
8+
9+
10+
@pointwise_dynamic(is_tensor=[True, True, False], promotion_methods=[(0, 1, "DEFAULT")])
11+
@triton.jit
12+
def rsub_func(x, y, alpha):
13+
return y - x * alpha
14+
15+
16+
@pointwise_dynamic(
17+
is_tensor=[True, False, False], promotion_methods=[(0, 1, "DEFAULT")]
18+
)
19+
@triton.jit
20+
def rsub_func_tensor_scalar(x, y, alpha):
21+
return y - x * alpha
22+
23+
24+
def rsub_tensor(A, B, *, alpha=1):
25+
logger.debug("GEMS RSUB TENSOR")
26+
return rsub_func(A, B, alpha)
27+
28+
29+
def rsub_scalar(A, B, alpha=1):
30+
logger.debug("GEMS RSUB SCALAR")
31+
return rsub_func_tensor_scalar(A, B, alpha)

tests/test_binary_pointwise_ops.py

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1864,6 +1864,40 @@ def test_accuracy_sub_scalar_scalar(dtype):
18641864
gems_assert_close(res_out, ref_out, dtype)
18651865

18661866

1867+
@pytest.mark.rsub
1868+
@pytest.mark.parametrize("shape", POINTWISE_SHAPES)
1869+
@pytest.mark.parametrize("alpha", SCALARS)
1870+
@pytest.mark.parametrize("dtype", FLOAT_DTYPES)
1871+
def test_accuracy_rsub(shape, alpha, dtype):
1872+
inp1 = torch.randn(shape, dtype=dtype, device=flag_gems.device)
1873+
inp2 = torch.randn(shape, dtype=dtype, device=flag_gems.device)
1874+
ref_inp1 = to_reference(inp1, True)
1875+
ref_inp2 = to_reference(inp2, True)
1876+
1877+
ref_out = torch.rsub(ref_inp1, ref_inp2, alpha=alpha)
1878+
with flag_gems.use_gems():
1879+
res_out = torch.rsub(inp1, inp2, alpha=alpha)
1880+
1881+
gems_assert_close(res_out, ref_out, dtype)
1882+
1883+
1884+
@pytest.mark.rsub
1885+
@pytest.mark.parametrize("shape", POINTWISE_SHAPES)
1886+
@pytest.mark.parametrize("scalar", SCALARS)
1887+
@pytest.mark.parametrize("alpha", SCALARS)
1888+
@pytest.mark.parametrize("dtype", FLOAT_DTYPES)
1889+
def test_accuracy_rsub_tensor_scalar(shape, scalar, alpha, dtype):
1890+
inp1 = torch.randn(shape, dtype=dtype, device=flag_gems.device)
1891+
inp2 = scalar
1892+
ref_inp1 = to_reference(inp1, True)
1893+
1894+
ref_out = torch.rsub(ref_inp1, inp2, alpha=alpha)
1895+
with flag_gems.use_gems():
1896+
res_out = torch.rsub(inp1, inp2, alpha=alpha)
1897+
1898+
gems_assert_close(res_out, ref_out, dtype)
1899+
1900+
18671901
@pytest.mark.where
18681902
@pytest.mark.parametrize("shape", POINTWISE_SHAPES)
18691903
@pytest.mark.parametrize("dtype", FLOAT_DTYPES)

0 commit comments

Comments
 (0)