File size: 3,276 Bytes
ffa8327
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
"""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)