| """Hub-layout adapters around Helion's synchronized PyTorch references.""" |
|
|
| from __future__ import annotations |
|
|
| import torch |
|
|
| from ._helion_reference import ( |
| chunked_linear_attn_reference, |
| naive_recurrent_reference, |
| rel_error, |
| recurrent_step_reference, |
| ) |
|
|
|
|
| def relative_error(a: torch.Tensor | None, b: torch.Tensor | None) -> float: |
| return rel_error(a, b) |
|
|
|
|
| def make_inputs( |
| device: torch.device, |
| *, |
| b: int = 1, |
| t: int = 64, |
| h: int = 2, |
| d: int = 32, |
| dv: int = 32, |
| ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: |
| tensors = [ |
| torch.randn(b, t, h, dim, device=device, dtype=torch.float32) |
| for dim in (d, d, dv) |
| ] |
| return tuple(x.to(torch.bfloat16) for x in tensors) |
|
|
|
|
| def _head_first(x: torch.Tensor) -> torch.Tensor: |
| return x.transpose(1, 2).contiguous() |
|
|
|
|
| def recurrent_reference( |
| q: torch.Tensor, |
| k: torch.Tensor, |
| v: torch.Tensor, |
| *, |
| g: torch.Tensor | None = None, |
| beta: torch.Tensor | None = None, |
| scale: float, |
| ) -> tuple[torch.Tensor, torch.Tensor]: |
| """Run Helion's recurrent reference on time-first Hub inputs.""" |
| qh, kh, vh = (_head_first(x) for x in (q, k, v)) |
| gh = ( |
| _head_first(g) |
| if g is not None |
| else torch.zeros( |
| q.size(0), |
| q.size(2), |
| q.size(1), |
| device=q.device, |
| dtype=torch.float32, |
| ) |
| ) |
| bh = _head_first(beta) if beta is not None else None |
|
|
| output = naive_recurrent_reference( |
| qh, |
| kh, |
| vh, |
| gh.float(), |
| beta=bh, |
| q_scale=scale, |
| ) |
|
|
| state = torch.zeros( |
| q.size(0), |
| q.size(2), |
| q.size(3), |
| v.size(3), |
| device=q.device, |
| dtype=torch.float32, |
| ) |
| for index in range(q.size(1)): |
| decay = gh[:, :, index : index + 1].float().exp() |
| beta_value = bh[:, :, index : index + 1].float() if bh is not None else None |
| _, state = recurrent_step_reference( |
| qh[:, :, index : index + 1].float() * scale, |
| kh[:, :, index : index + 1].float(), |
| vh[:, :, index : index + 1].float(), |
| state, |
| alpha=decay, |
| beta_val=beta_value, |
| ) |
|
|
| return _head_first(output), state |
|
|
|
|
| def chunked_reference( |
| q: torch.Tensor, |
| k: torch.Tensor, |
| v: torch.Tensor, |
| *, |
| g: torch.Tensor | None = None, |
| beta: torch.Tensor | None = None, |
| scale: float, |
| chunk_size: int = 64, |
| ) -> torch.Tensor: |
| """Run Helion's differentiable chunked reference on Hub-layout inputs.""" |
| qh, kh, vh = (_head_first(x) for x in (q, k, v)) |
| gh = ( |
| _head_first(g) |
| if g is not None |
| else torch.zeros( |
| q.size(0), |
| q.size(2), |
| q.size(1), |
| device=q.device, |
| dtype=torch.float32, |
| ) |
| ) |
| bh = _head_first(beta) if beta is not None else None |
| output = chunked_linear_attn_reference( |
| qh * scale, |
| kh, |
| vh, |
| gh, |
| beta=bh, |
| C=chunk_size, |
| ) |
| return _head_first(output) |
|
|
|
|
| def assert_close(actual: torch.Tensor, expected: torch.Tensor) -> None: |
| torch.testing.assert_close(actual.float(), expected.float(), atol=6e-2, rtol=3e-2) |
|
|