diff --git a/docs/design-notes/geglu.md b/docs/design-notes/geglu.md index 142a3eb..8ae1001 100644 --- a/docs/design-notes/geglu.md +++ b/docs/design-notes/geglu.md @@ -407,7 +407,7 @@ Implementation notes: Correctness plan: - Top-level tests cover separate and packed APIs, forward/backward parity, `approximate="tanh"` and `"none"`, fp32/fp16/bf16, odd hidden sizes, preserve-input mode, CPU fallback, invalid inputs, special forward values, and a `hidden=65537` flat-path case. -- Gradcheck is deferred because the Triton CUDA path targets fp32/fp16/bf16, while PyTorch gradcheck expects double precision. Backward parity against PyTorch autograd is the hard gate for the POC. +- `test_geglu_gradcheck_fp64` and `test_geglu_packed_gradcheck_fp64` now run `torch.autograd.gradcheck` through `ForgeGEGLUFunction`/`ForgePackedGEGLUFunction` on small fp64 inputs, the same way RoPE's gradcheck does. `_check_cuda_dtype` allows fp64 for this; fp32/fp16/bf16 remain the intended training/inference dtypes. Benchmark plan: diff --git a/kernels/geglu/geglu.py b/kernels/geglu/geglu.py index 53e94a5..400d5a1 100644 --- a/kernels/geglu/geglu.py +++ b/kernels/geglu/geglu.py @@ -29,9 +29,9 @@ def _check_same_shape_inputs(gate: torch.Tensor, up: torch.Tensor) -> None: def _check_cuda_dtype(x: torch.Tensor) -> None: - """Inputs: a tensor. Outputs: none (raises on unsupported dtype). Logic: the CUDA kernels only handle fp16/bf16/fp32.""" - if x.dtype not in (torch.float16, torch.bfloat16, torch.float32): - raise TypeError(f"CUDA GEGLU supports fp16, bf16, and fp32, got {x.dtype}") + """Inputs: a tensor. Outputs: none (raises on unsupported dtype). Logic: the CUDA kernels handle fp16/bf16/fp32 for training/inference, plus fp64 so gradcheck can run the real kernel.""" + if x.dtype not in (torch.float16, torch.bfloat16, torch.float32, torch.float64): + raise TypeError(f"CUDA GEGLU supports fp16, bf16, fp32, and fp64, got {x.dtype}") def _check_packed_input(gate_up: torch.Tensor) -> None: diff --git a/tests/test_geglu.py b/tests/test_geglu.py index 25d7181..794c431 100644 --- a/tests/test_geglu.py +++ b/tests/test_geglu.py @@ -8,6 +8,8 @@ ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) +from kernels.geglu import ForgeGEGLUFunction +from kernels.geglu import ForgePackedGEGLUFunction from kernels.geglu import geglu from kernels.geglu import geglu_mlp from kernels.geglu import geglu_packed @@ -270,6 +272,51 @@ def test_geglu_mlp_packed_matches_separate_and_reference(use_bias, approximate, torch.testing.assert_close(down_bias_packed.grad, down_bias_ref.grad, atol=atol, rtol=rtol) +@pytest.mark.parametrize("approximate", ["tanh", "none"]) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="Triton GEGLU path requires CUDA") +def test_geglu_gradcheck_fp64(approximate): + torch.manual_seed(0) + shape = (2, 3, 8) + gate = torch.randn(*shape, device=DEVICE, dtype=torch.float64, requires_grad=True) + up = torch.randn(*shape, device=DEVICE, dtype=torch.float64, requires_grad=True) + + def func(gate, up): + return ForgeGEGLUFunction.apply(gate, up, approximate, True) + + assert torch.autograd.gradcheck( + func, + (gate, up), + eps=1e-3, + atol=1e-2, + rtol=1e-2, + nondet_tol=1e-3, + check_undefined_grad=False, + check_batched_grad=False, + ) + + +@pytest.mark.parametrize("approximate", ["tanh", "none"]) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="Triton GEGLU path requires CUDA") +def test_geglu_packed_gradcheck_fp64(approximate): + torch.manual_seed(1) + shape = (2, 3, 8) + gate_up = torch.randn(*shape, device=DEVICE, dtype=torch.float64, requires_grad=True) + + def func(gate_up): + return ForgePackedGEGLUFunction.apply(gate_up, approximate, True) + + assert torch.autograd.gradcheck( + func, + (gate_up,), + eps=1e-3, + atol=1e-2, + rtol=1e-2, + nondet_tol=1e-3, + check_undefined_grad=False, + check_batched_grad=False, + ) + + def test_geglu_rejects_invalid_inputs(): gate = torch.randn(2, 3) up = torch.randn(2, 4)