"""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)