diff --git a/tests/mhc/test_functional_mhc_pre.py b/tests/mhc/test_functional_mhc_pre.py new file mode 100644 index 0000000..f442183 --- /dev/null +++ b/tests/mhc/test_functional_mhc_pre.py @@ -0,0 +1,108 @@ +import pytest +import torch +from tile_kernels.modeling.mhc.functional import mhc_pre + + +def generate_mhc_pre_test_data( + n1: int, + mhc_mult: int, + hidden_size: int, + generate_norm_weight: bool, + norm_eps: float = 1e-6, + post_mult_value: float = 1.0, + pre_eps: float = 1e-6, + sinkhorn_eps: float = 1e-6, + sinkhorn_repeat: int = 10, + n_splits: int = 16, +) -> dict[str, torch.Tensor]: + n0 = 1 + mhc_mult3 = mhc_mult * (2 + mhc_mult) + mhc_hidden_size = mhc_mult * hidden_size + device = 'cuda' + + residual = ( + torch.randn((n0, n1, mhc_mult, hidden_size), dtype=torch.float, device=device) + .mul(1 + torch.arange(mhc_mult, device=device).mul(0.01).view(1, 1, -1, 1)) + .bfloat16() + ) + + fn = ( + torch.randn((mhc_mult3, mhc_mult, hidden_size), dtype=torch.float, device=device) + * 1e-4 + * (1 + torch.arange(mhc_mult, device=device).mul(0.01).view(1, -1, 1)) + ).flatten(1, 2) + + scale = torch.randn((3,), dtype=torch.float, device=device) * 0.1 + base = torch.randn((mhc_mult3,), dtype=torch.float, device=device) * 0.1 + + if generate_norm_weight: + norm_weight = torch.randn((mhc_hidden_size,), dtype=torch.float, device=device) * 0.1 + 1.0 + else: + norm_weight = None + + return { + 'residual': residual, + 'fn': fn, + 'scale': scale, + 'base': base, + 'norm_weight': norm_weight, + 'norm_eps': norm_eps, + 'post_mult_value': post_mult_value, + 'pre_eps': pre_eps, + 'sinkhorn_eps': sinkhorn_eps, + 'sinkhorn_repeat': sinkhorn_repeat, + 'n_splits': n_splits, + } + + +@pytest.mark.parametrize('n1', [4096, 8192]) +@pytest.mark.parametrize('hidden_size', [1280, 2560, 7168]) +@pytest.mark.parametrize('generate_norm_weight', [False, True]) +def test_correctness( + n1: int, + hidden_size: int, + generate_norm_weight: bool, +) -> None: + mhc_mult = 4 + + test_data = generate_mhc_pre_test_data( + n1=n1, + mhc_mult=mhc_mult, + hidden_size=hidden_size, + generate_norm_weight=generate_norm_weight, + ) + + train_layer_input, (train_post_mix, train_comb_mix) = mhc_pre( + test_data['residual'], + test_data['fn'], + test_data['scale'], + test_data['base'], + norm_weight=test_data['norm_weight'], + norm_eps=test_data['norm_eps'], + mhc_mult=mhc_mult, + post_mult_value=test_data['post_mult_value'], + pre_eps=test_data['pre_eps'], + sinkhorn_eps=test_data['sinkhorn_eps'], + sinkhorn_repeat=test_data['sinkhorn_repeat'], + n_splits=test_data['n_splits'], + ) + + with torch.no_grad(): + eval_layer_input, (eval_post_mix, eval_comb_mix) = mhc_pre( + test_data['residual'], + test_data['fn'], + test_data['scale'], + test_data['base'], + norm_weight=test_data['norm_weight'], + norm_eps=test_data['norm_eps'], + mhc_mult=mhc_mult, + post_mult_value=test_data['post_mult_value'], + pre_eps=test_data['pre_eps'], + sinkhorn_eps=test_data['sinkhorn_eps'], + sinkhorn_repeat=test_data['sinkhorn_repeat'], + n_splits=test_data['n_splits'], + ) + + torch.testing.assert_close(train_layer_input, eval_layer_input, atol=1e-4, rtol=1e-4) + torch.testing.assert_close(train_post_mix, eval_post_mix, atol=1e-4, rtol=1e-4) + torch.testing.assert_close(train_comb_mix, eval_comb_mix, atol=1e-4, rtol=1e-4) diff --git a/tile_kernels/modeling/mhc/functional.py b/tile_kernels/modeling/mhc/functional.py index ee532ec..f130645 100644 --- a/tile_kernels/modeling/mhc/functional.py +++ b/tile_kernels/modeling/mhc/functional.py @@ -3,7 +3,7 @@ from .ops.expand import expand_to_mhc from .ops.head_compute_mix import mhc_head_compute_mix -from .ops.norm_fn import mhc_pre_norm_fn +from .ops.norm_fn import mhc_pre_norm_fn, mhc_fn_normw_merge from .ops.post import mhc_post from .ops.pre_apply_mix import mhc_pre_apply_mix from .ops.pre_big_fuse import mhc_pre_big_fuse @@ -67,6 +67,9 @@ def mhc_pre( ctx: opaque tuple (post_mix, comb_mix) to pass to mhc_post """ if not torch.is_grad_enabled(): + # mhc_pre_big_fuse does not accept norm_weight as an argument. + # We must pre-fuse norm_weight into fn before calling the kernel. + fn = mhc_fn_normw_merge(fn, norm_weight) post_mix, comb_mix, layer_input = mhc_pre_big_fuse( residual, fn, diff --git a/tile_kernels/modeling/mhc/ops/norm_fn.py b/tile_kernels/modeling/mhc/ops/norm_fn.py index b0ec6cd..38686c1 100644 --- a/tile_kernels/modeling/mhc/ops/norm_fn.py +++ b/tile_kernels/modeling/mhc/ops/norm_fn.py @@ -12,15 +12,22 @@ ) +def _mhc_fn_normw_merge_impl( + mhc_fn: torch.Tensor, + mhc_norm_weight: torch.Tensor, +) -> torch.Tensor: + out_fn = torch.empty_like(mhc_fn) + _mhc_fn_normw_merge_fwd(*mhc_fn.shape)(mhc_fn, mhc_norm_weight, out_fn) + return out_fn + + class _MHCFnNormwMerge(torch.autograd.Function): @staticmethod def forward(ctx: '_MHCFnNormwMerge', fn: torch.Tensor, normw: torch.Tensor) -> torch.Tensor: ctx.fn_main_grad = getattr(fn, 'main_grad', None) ctx.normw_main_grad = getattr(normw, 'main_grad', None) ctx.save_for_backward(fn, normw) - out_fn = torch.empty_like(fn) - _mhc_fn_normw_merge_fwd(*fn.shape)(fn, normw, out_fn) - return out_fn + return _mhc_fn_normw_merge_impl(fn, normw) @staticmethod def backward(ctx: '_MHCFnNormwMerge', out_fn_grad: torch.Tensor) -> tuple[None, None]: @@ -170,6 +177,19 @@ def backward( return x_grad, fn_grad, None, None, None, None +def mhc_fn_normw_merge( + mhc_fn: torch.Tensor, + mhc_norm_weight: torch.Tensor | None, +) -> torch.Tensor: + if mhc_norm_weight is None: + return mhc_fn + + if not torch.is_grad_enabled(): + return _mhc_fn_normw_merge_impl(mhc_fn, mhc_norm_weight) + + return _MHCFnNormwMerge.apply(mhc_fn, mhc_norm_weight) + + def mhc_pre_norm_fn( residual: torch.Tensor, mhc_fn: torch.Tensor, @@ -178,8 +198,7 @@ def mhc_pre_norm_fn( fuse_grad_acc: bool = True, n_splits: int = 16, ) -> torch.Tensor: - if mhc_norm_weight is not None: - mhc_fn = _MHCFnNormwMerge.apply(mhc_fn, mhc_norm_weight) + mhc_fn = mhc_fn_normw_merge(mhc_fn, mhc_norm_weight) return MHCPreNormFn.apply( residual, mhc_fn,