Skip to content

Commit 5beb3b8

Browse files
committed
fix: split fmod into separate tensor/scalar entries, fix typo in test name
1 parent 466402c commit 5beb3b8

4 files changed

Lines changed: 26 additions & 24 deletions

File tree

src/flag_gems/__init__.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -215,10 +215,10 @@ def torch_ge(v):
215215
("floor_divide_.Tensor", floor_divide_),
216216
("fmin", fmin),
217217
("fmin.out", fmin_out),
218-
("fmod.Scalar", fmod),
219-
("fmod.Tensor", fmod),
220-
("fmod_.Scalar", fmod_),
221-
("fmod_.Tensor", fmod_),
218+
("fmod.Scalar", fmod_scalar),
219+
("fmod.Tensor", fmod_tensor),
220+
("fmod_.Scalar", fmod_scalar_),
221+
("fmod_.Tensor", fmod_tensor_),
222222
("full", full),
223223
("full_like", full_like),
224224
("gather", gather),

src/flag_gems/ops/__init__.py

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -125,7 +125,7 @@
125125
from flag_gems.ops.flip import flip
126126
from flag_gems.ops.floor_ import floor_
127127
from flag_gems.ops.fmin import fmin, fmin_out
128-
from flag_gems.ops.fmod import fmod, fmod_
128+
from flag_gems.ops.fmod import fmod_scalar, fmod_scalar_, fmod_tensor, fmod_tensor_
129129
from flag_gems.ops.full import full
130130
from flag_gems.ops.full_like import full_like
131131
from flag_gems.ops.gather import gather, gather_backward
@@ -480,8 +480,10 @@
480480
"floor_divide_",
481481
"fmin",
482482
"fmin_out",
483-
"fmod",
484-
"fmod_",
483+
"fmod_scalar",
484+
"fmod_scalar_",
485+
"fmod_tensor",
486+
"fmod_tensor_",
485487
"full",
486488
"full_like",
487489
"gather",

src/flag_gems/ops/fmod.py

Lines changed: 16 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,5 @@
11
import logging
22

3-
import torch
43
import triton
54
import triton.language as tl
65

@@ -32,20 +31,21 @@ def fmod_func_tensor_scalar(x, y):
3231
return result.to(dtype)
3332

3433

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)
34+
def fmod_tensor(A, B):
35+
logger.debug("GEMS FMOD TENSOR")
36+
return fmod_func(A, B)
4437

4538

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)
39+
def fmod_scalar(A, B):
40+
logger.debug("GEMS FMOD SCALAR")
41+
return fmod_func_tensor_scalar(A, B)
42+
43+
44+
def fmod_tensor_(A, B):
45+
logger.debug("GEMS FMOD_ TENSOR")
46+
return fmod_func(A, B, out0=A)
47+
48+
49+
def fmod_scalar_(A, B):
50+
logger.debug("GEMS FMOD_ SCALAR")
51+
return fmod_func_tensor_scalar(A, B, out0=A)

tests/test_fmod.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -61,7 +61,7 @@ def test_fmod_tensor_inplace(shape, dtype):
6161
@pytest.mark.parametrize("shape", POINTWISE_SHAPES)
6262
@pytest.mark.parametrize("scalar", SCALARS)
6363
@pytest.mark.parametrize("dtype", FLOAT_DTYPES)
64-
def testy_fmod_scalar_inplace(shape, scalar, dtype):
64+
def test_fmod_scalar_inplace(shape, scalar, dtype):
6565
inp1 = torch.randn(shape, dtype=dtype, device=flag_gems.device)
6666
inp2 = scalar if scalar != 0 else 1.0
6767
ref_inp1 = to_reference(inp1.clone(), True)

0 commit comments

Comments
 (0)