Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion docs/design-notes/geglu.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:

Expand Down
6 changes: 3 additions & 3 deletions kernels/geglu/geglu.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
47 changes: 47 additions & 0 deletions tests/test_geglu.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down