From d79772f0575c9432aca88e73094fefcc1a44e451 Mon Sep 17 00:00:00 2001 From: chenyu Date: Tue, 4 Aug 2026 19:30:43 -0400 Subject: [PATCH] fix pow on extreme inputs (#17397) * fix pow on extreme inputs * WEBGPU --- test/backend/test_ops.py | 11 +++++++++++ tinygrad/codegen/decomp/transcendental.py | 8 ++++---- tinygrad/mixin/gradient.py | 2 +- 3 files changed, 16 insertions(+), 5 deletions(-) diff --git a/test/backend/test_ops.py b/test/backend/test_ops.py index 520ed9d793..ae2c81ee15 100644 --- a/test/backend/test_ops.py +++ b/test/backend/test_ops.py @@ -728,6 +728,17 @@ class TestOps(unittest.TestCase): else: self.assertAlmostEqual(tiny_out, torch_out, msg=f"{x}, {c}") + def test_pow_neg_inf_frac_exponent(self): + # pow(-inf, 0.3) is +inf, so the gradient 0.3*pow(-inf, -0.7) is 0, never nan + helper_test_op(None, lambda x: x**0.3, vals=[[-math.inf]]) + # is_odd truncates, so it calls 3.3 odd: only the non_int guard keeps pow(-inf, 3.3) from negating to -inf + helper_test_op(None, lambda x: x**3.3, vals=[[-math.inf]]) + + def test_pow_zero_exponent(self): + # x ** 0 is the constant 1 for every x, so the gradient with respect to the base is 0, never nan + # TODO: nan ** 0, failed on WEBGPU + helper_test_op(None, lambda x,y: x**y, vals=[[-math.inf, math.inf, 0.0], [0.0, 0.0, 0.0]]) + def test_pow_zero_tensor(self): helper_test_op(None, lambda x,y: x**y, vals=[[0.0], [0.0]]) # TODO: fix WEBGPU diff --git a/tinygrad/codegen/decomp/transcendental.py b/tinygrad/codegen/decomp/transcendental.py index b2771acb37..e4e66fbd5a 100644 --- a/tinygrad/codegen/decomp/transcendental.py +++ b/tinygrad/codegen/decomp/transcendental.py @@ -257,12 +257,12 @@ def xlog2(d:UOp) -> UOp: def xpow(base:UOp, exponent:UOp) -> UOp: # start with b ** e = exp2(e * log2(b)) ret = (base < 0).where(-base, base).log2().mul(exponent).exp2() - # negative base: nan for non-integer exponent, negate for odd integer exponent + # negative base: nan for non-integer exponent, negate for odd integer exponent. -inf is never nan, it stays |base| ** exponent non_int = exponent != exponent.cast(dtypes.int32).cast(exponent.dtype) is_odd = (exponent < 0).where(-exponent, exponent).cast(dtypes.int32).mod(2).cast(dtypes.bool) - neg_base = non_int.where(ret.const_like(math.nan), is_odd.where(-ret, ret)) - # fix 0 ** 0 = 1 - return (base.eq(0) & exponent.eq(0)).where(ret.const_like(1), (base < 0).where(neg_base, ret)) + neg_base = non_int.where(base.ne(-math.inf).where(ret.const_like(math.nan), ret), is_odd.where(-ret, ret)) + # x ** 0 = 1, including 0 ** 0 and inf ** 0 + return exponent.eq(0).where(ret.const_like(1), (base < 0).where(neg_base, ret)) @functools.cache def get_transcendental_patterns(ops:tuple[Ops, ...], force_transcendental:bool) -> PatternMatcher: diff --git a/tinygrad/mixin/gradient.py b/tinygrad/mixin/gradient.py index 57d5e0dc41..93cc63843b 100644 --- a/tinygrad/mixin/gradient.py +++ b/tinygrad/mixin/gradient.py @@ -53,7 +53,7 @@ pm_gradient = PatternMatcher([ (UPat((Ops.CMPLT, Ops.CMPNE)), lambda: (None, None)), (UPat(Ops.ADD), lambda ctx: (ctx, ctx)), (UPat(Ops.POW, name="ret", src=(UPat.var("b"), UPat.var("e"))), lambda ctx, ret, b, e: - (ctx * (b.eq(0)&e.eq(0)).where(e, e*b.pow(e-1)), ctx * b.eq(0).where((e<0).where(ret.const_like(-math.inf), 0), ret*b.log2()*math.log(2.0)))), + (ctx * e.eq(0).where(e, e*b.pow(e-1)), ctx * b.eq(0).where((e<0).where(ret.const_like(-math.inf), 0), ret*b.log2()*math.log(2.0)))), (UPat(Ops.MAX, src=(UPat.var("x"), UPat.var("y"))), lambda ctx, x, y: ((x>y).where(ctx, (x.eq(y)).where(ctx * 0.5, 0)), (x