diff --git a/pyproject.toml b/pyproject.toml index fe7c31a5f4..fcb377e607 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -476,7 +476,6 @@ name = 'gridtools' url = 'https://gridtools.github.io/pypi/' # Add the uv source below to pull dace from the gridtools index instead of PyPI: -# dace = {index = "gridtools"} [tool.uv.sources] atlas4py = {index = "test.pypi"} diff --git a/src/gt4py/cartesian/gtc/dace/oir_to_tasklet.py b/src/gt4py/cartesian/gtc/dace/oir_to_tasklet.py index 2329128d70..6a7c71047e 100644 --- a/src/gt4py/cartesian/gtc/dace/oir_to_tasklet.py +++ b/src/gt4py/cartesian/gtc/dace/oir_to_tasklet.py @@ -226,38 +226,38 @@ def visit_NativeFunction(self, node: common.NativeFunction, **_kwargs: Any) -> s common.NativeFunction.ABS: "abs", common.NativeFunction.MIN: "min", common.NativeFunction.MAX: "max", - common.NativeFunction.MOD: "fmod", + common.NativeFunction.MOD: "dace.math.fmod", common.NativeFunction.SIN: "dace.math.sin", common.NativeFunction.COS: "dace.math.cos", common.NativeFunction.TAN: "dace.math.tan", - common.NativeFunction.ARCSIN: "asin", - common.NativeFunction.ARCCOS: "acos", - common.NativeFunction.ARCTAN: "atan", + common.NativeFunction.ARCSIN: "dace.math.asin", + common.NativeFunction.ARCCOS: "dace.math.acos", + common.NativeFunction.ARCTAN: "dace.math.atan", common.NativeFunction.SINH: "dace.math.sinh", common.NativeFunction.COSH: "dace.math.cosh", common.NativeFunction.TANH: "dace.math.tanh", - common.NativeFunction.ARCSINH: "asinh", - common.NativeFunction.ARCCOSH: "acosh", - common.NativeFunction.ARCTANH: "atanh", + common.NativeFunction.ARCSINH: "dace.math.asinh", + common.NativeFunction.ARCCOSH: "dace.math.acosh", + common.NativeFunction.ARCTANH: "dace.math.atanh", common.NativeFunction.SQRT: "dace.math.sqrt", common.NativeFunction.POW: "dace.math.pow", common.NativeFunction.EXP: "dace.math.exp", common.NativeFunction.LOG: "dace.math.log", - common.NativeFunction.LOG10: "log10", - common.NativeFunction.GAMMA: "tgamma", - common.NativeFunction.CBRT: "cbrt", + common.NativeFunction.LOG10: "dace.math.log10", + common.NativeFunction.GAMMA: "dace.math.tgamma", + common.NativeFunction.CBRT: "dace.math.cbrt", common.NativeFunction.ISFINITE: "isfinite", common.NativeFunction.ISINF: "isinf", common.NativeFunction.ISNAN: "isnan", common.NativeFunction.FLOOR: "dace.math.ifloor", - common.NativeFunction.CEIL: "ceil", - common.NativeFunction.TRUNC: "trunc", + common.NativeFunction.CEIL: "dace.math.ceil", + common.NativeFunction.TRUNC: "dace.math.trunc", common.NativeFunction.INT32: "dace.int32", common.NativeFunction.INT64: "dace.int64", common.NativeFunction.FLOAT32: "dace.float32", common.NativeFunction.FLOAT64: "dace.float64", - common.NativeFunction.ERF: "erf", - common.NativeFunction.ERFC: "erfc", + common.NativeFunction.ERF: "dace.math.erf", + common.NativeFunction.ERFC: "dace.math.erfc", common.NativeFunction.ROUND: "nearbyint", common.NativeFunction.ROUND_AWAY_FROM_ZERO: "round", } diff --git a/tests/cartesian_tests/unit_tests/test_gtc/dace/test_oir_to_tasklet.py b/tests/cartesian_tests/unit_tests/test_gtc/dace/test_oir_to_tasklet.py index 1ecb7659ee..980727fa3c 100644 --- a/tests/cartesian_tests/unit_tests/test_gtc/dace/test_oir_to_tasklet.py +++ b/tests/cartesian_tests/unit_tests/test_gtc/dace/test_oir_to_tasklet.py @@ -139,3 +139,22 @@ def test_integer_power_of_integer() -> None: tasklet_code = visitor.visit_NativeFuncCall(pow_call, ctx=fake_context, is_target=False) assert "ipow" not in tasklet_code + + +@pytest.mark.parametrize( + "arg", + [ + oir.Literal(value="2", dtype=common.DataType.FLOAT32), + oir.Literal(value="2", dtype=common.DataType.FLOAT64), + ], +) +def test_log10_respects_floating_point_precision(arg: oir.Literal) -> None: + log10_call = oir.NativeFuncCall(func=common.NativeFunction.LOG10, args=[arg]) + + visitor = oir_to_tasklet.OIRToTasklet() + fake_context = oir_to_tasklet.Context( + code="asdf", targets=set(), inputs={}, outputs={}, tree=None, scope=None + ) + tasklet_code = visitor.visit_NativeFuncCall(log10_call, ctx=fake_context, is_target=False) + + assert "dace.math.log10" in tasklet_code