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
108 changes: 108 additions & 0 deletions tests/mhc/test_functional_mhc_pre.py
Original file line number Diff line number Diff line change
@@ -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)
Comment thread
GenTang marked this conversation as resolved.
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,
Comment thread
GenTang marked this conversation as resolved.
)

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'],
)

Comment thread
GenTang marked this conversation as resolved.
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'],
)
Comment thread
GenTang marked this conversation as resolved.

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)
5 changes: 4 additions & 1 deletion tile_kernels/modeling/mhc/functional.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down
29 changes: 24 additions & 5 deletions tile_kernels/modeling/mhc/ops/norm_fn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down