Skip to content

Commit 82a13b3

Browse files
Add fmod operator implementation, tests and benchmark
1 parent 015d315 commit 82a13b3

5 files changed

Lines changed: 135 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
@@ -49,6 +49,7 @@ def get_tflops(self, op, *args, **kwargs):
4949
("pow", torch.pow, FLOAT_DTYPES),
5050
("polar", torch.polar, [torch.float32]),
5151
("floor_divide", torch.floor_divide, INT_DTYPES),
52+
("fmod", torch.fmod, FLOAT_DTYPES),
5253
("remainder", torch.remainder, INT_DTYPES),
5354
("logical_or", torch.logical_or, INT_DTYPES + BOOL_DTYPES),
5455
("logical_and", torch.logical_and, INT_DTYPES + BOOL_DTYPES),

src/flag_gems/__init__.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -165,6 +165,10 @@ def torch_ge(v):
165165
("fill_.Scalar", fill_scalar_),
166166
("fill_.Tensor", fill_tensor_),
167167
("flip", flip),
168+
("fmod.Scalar", fmod),
169+
("fmod.Tensor", fmod),
170+
("fmod_.Scalar", fmod_),
171+
("fmod_.Tensor", fmod_),
168172
("floor_divide", floor_divide),
169173
("floor_divide.Scalar", floor_divide),
170174
("floor_divide_.Scalar", floor_divide_),

src/flag_gems/ops/__init__.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -92,6 +92,7 @@
9292
from flag_gems.ops.eye_m import eye_m
9393
from flag_gems.ops.fill import fill_scalar, fill_scalar_, fill_tensor, fill_tensor_
9494
from flag_gems.ops.flip import flip
95+
from flag_gems.ops.fmod import fmod, fmod_
9596
from flag_gems.ops.full import full
9697
from flag_gems.ops.full_like import full_like
9798
from flag_gems.ops.gather import gather, gather_backward
@@ -350,6 +351,8 @@
350351
"flash_attention_forward",
351352
"flash_attn_varlen_func",
352353
"flip",
354+
"fmod",
355+
"fmod_",
353356
"floor_divide",
354357
"floor_divide_",
355358
"full",

src/flag_gems/ops/fmod.py

Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,51 @@
1+
import logging
2+
3+
import torch
4+
import triton
5+
import triton.language as tl
6+
7+
from flag_gems.utils import pointwise_dynamic
8+
from flag_gems.utils.triton_lang_extension import fmod as _fmod
9+
10+
logger = logging.getLogger(__name__)
11+
12+
13+
@pointwise_dynamic(promotion_methods=[(0, 1, "DEFAULT")])
14+
@triton.jit
15+
def fmod_func(x, y):
16+
# Convert to float32 for computation to avoid libdevice float16/bfloat16 issues
17+
dtype = x.dtype
18+
x_fp32 = x.to(tl.float32)
19+
y_fp32 = y.to(tl.float32)
20+
result = _fmod(x_fp32, y_fp32)
21+
return result.to(dtype)
22+
23+
24+
@pointwise_dynamic(is_tensor=[True, False], promotion_methods=[(0, 1, "DEFAULT")])
25+
@triton.jit
26+
def fmod_func_tensor_scalar(x, y):
27+
# Convert to float32 for computation to avoid libdevice float16/bfloat16 issues
28+
dtype = x.dtype
29+
x_fp32 = x.to(tl.float32)
30+
y_fp32 = y.to(tl.float32)
31+
result = _fmod(x_fp32, y_fp32)
32+
return result.to(dtype)
33+
34+
35+
def fmod(A, B):
36+
logger.debug("GEMS FMOD")
37+
if isinstance(A, torch.Tensor) and isinstance(B, torch.Tensor):
38+
return fmod_func(A, B)
39+
elif isinstance(A, torch.Tensor):
40+
return fmod_func_tensor_scalar(A, B)
41+
else:
42+
# Both scalar - fallback to PyTorch
43+
return torch.fmod(torch.tensor(A), B)
44+
45+
46+
def fmod_(A, B):
47+
logger.debug("GEMS FMOD_")
48+
if isinstance(B, torch.Tensor):
49+
return fmod_func(A, B, out0=A)
50+
else:
51+
return fmod_func_tensor_scalar(A, B, out0=A)

tests/test_binary_pointwise_ops.py

Lines changed: 76 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2166,3 +2166,79 @@ def test_accuracy_addcdiv(shape, dtype):
21662166
res_out = torch.addcdiv(res_inp, t1, t2, value=v)
21672167

21682168
gems_assert_close(res_out, ref_out, dtype)
2169+
2170+
2171+
@pytest.mark.fmod
2172+
@pytest.mark.parametrize("shape", POINTWISE_SHAPES)
2173+
@pytest.mark.parametrize("dtype", FLOAT_DTYPES)
2174+
def test_accuracy_fmod(shape, dtype):
2175+
inp1 = torch.randn(shape, dtype=dtype, device=flag_gems.device)
2176+
inp2 = torch.randn(shape, dtype=dtype, device=flag_gems.device)
2177+
# Avoid division by zero
2178+
inp2 = torch.where(inp2 == 0, torch.ones_like(inp2), inp2)
2179+
ref_inp1 = to_reference(inp1, True)
2180+
ref_inp2 = to_reference(inp2, True)
2181+
2182+
ref_out = torch.fmod(ref_inp1, ref_inp2)
2183+
with flag_gems.use_gems():
2184+
res_out = torch.fmod(inp1, inp2)
2185+
2186+
gems_assert_close(res_out, ref_out, dtype)
2187+
2188+
2189+
@pytest.mark.fmod
2190+
@pytest.mark.parametrize("shape", POINTWISE_SHAPES)
2191+
@pytest.mark.parametrize("scalar", SCALARS)
2192+
@pytest.mark.parametrize("dtype", FLOAT_DTYPES)
2193+
def test_accuracy_fmod_tensor_scalar(shape, scalar, dtype):
2194+
inp1 = torch.randn(shape, dtype=dtype, device=flag_gems.device)
2195+
# Avoid division by zero
2196+
inp2 = scalar if scalar != 0 else 1.0
2197+
ref_inp1 = to_reference(inp1, True)
2198+
2199+
ref_out = torch.fmod(ref_inp1, inp2)
2200+
with flag_gems.use_gems():
2201+
res_out = torch.fmod(inp1, inp2)
2202+
2203+
# Use larger tolerance for fmod with small divisors due to float32 precision limits
2204+
atol = 1e-3 if abs(scalar) < 0.01 else 1e-4
2205+
gems_assert_close(res_out, ref_out, dtype, atol=atol)
2206+
2207+
2208+
@pytest.mark.inplace
2209+
@pytest.mark.fmod_
2210+
@pytest.mark.parametrize("shape", POINTWISE_SHAPES)
2211+
@pytest.mark.parametrize("dtype", FLOAT_DTYPES)
2212+
def test_accuracy_fmod_(shape, dtype):
2213+
inp1 = torch.randn(shape, dtype=dtype, device=flag_gems.device)
2214+
inp2 = torch.randn(shape, dtype=dtype, device=flag_gems.device)
2215+
# Avoid division by zero
2216+
inp2 = torch.where(inp2 == 0, torch.ones_like(inp2), inp2)
2217+
ref_inp1 = to_reference(inp1.clone(), True)
2218+
ref_inp2 = to_reference(inp2, True)
2219+
2220+
ref_out = ref_inp1.fmod_(ref_inp2)
2221+
with flag_gems.use_gems():
2222+
res_out = inp1.fmod_(inp2)
2223+
2224+
gems_assert_close(res_out, ref_out, dtype)
2225+
2226+
2227+
@pytest.mark.inplace
2228+
@pytest.mark.fmod_
2229+
@pytest.mark.parametrize("shape", POINTWISE_SHAPES)
2230+
@pytest.mark.parametrize("scalar", SCALARS)
2231+
@pytest.mark.parametrize("dtype", FLOAT_DTYPES)
2232+
def test_accuracy_fmod_tensor_scalar_(shape, scalar, dtype):
2233+
inp1 = torch.randn(shape, dtype=dtype, device=flag_gems.device)
2234+
# Avoid division by zero
2235+
inp2 = scalar if scalar != 0 else 1.0
2236+
ref_inp1 = to_reference(inp1.clone(), True)
2237+
2238+
ref_out = ref_inp1.fmod_(inp2)
2239+
with flag_gems.use_gems():
2240+
res_out = inp1.fmod_(inp2)
2241+
2242+
# Use larger tolerance for fmod with small divisors due to float32 precision limits
2243+
atol = 1e-3 if abs(scalar) < 0.01 else 1e-4
2244+
gems_assert_close(res_out, ref_out, dtype, atol=atol)

0 commit comments

Comments
 (0)