linear-attention / tests /reference.py
Tarindu Jayatilaka
Add linear attention examples and tests
ffa8327
Raw
History Blame
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)