Kernels:
Trusted publisher
Download tests/reference.py from HelionDSL/linear-attention: direct link, hf CLI and curl.
- Browser
- Download file 3.28 kB
-
https://huggingface.co/kernels/HelionDSL/linear-attention/resolve/v1/tests/reference.py
- Command line
-
hf download hf://HelionDSL/linear-attention@v1/tests/reference.py
-
curl -L -o reference.py https://huggingface.co/kernels/HelionDSL/linear-attention/resolve/v1/tests/reference.py
3.28 kB
| """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) | |