File size: 8,999 Bytes
ffdcfe7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
"""Masked-autoencoder objective over the encoder EAT already trains.

The point of this module is what it does *not* change. `MAEPretrainer` builds
its student from `eatmap.model.Encoder` verbatim -- same patch embedding, same
blocks, same cls token, same sincos positions -- so an MAE export loads through
`probe.load_encoder` and `lora.inject_lora` with no adaptation, and the two
objectives differ in the pretraining loss and nothing else. That is the whole
design: a cross-objective ranking comparison is only interpretable if
architecture, data, readout and mixture pool are held identical.

The call signature mirrors `EATPretrainer` exactly -- `masker(batch, device)`
returning `(ids_keep, ids_restore, mask)` and `forward(spec, keep, restore,
mask)` returning `(total, frame, utterance)` -- so `runner.py` needs a dispatch
line and no other change. MAE has no utterance branch, so that term is reported
as a constant zero rather than dropped, keeping the events schema stable across
objectives.

Two deliberate differences from `EATPretrainer`:

* **No EMA teacher.** MAE regresses the input, not a moving target, so
  `update_teacher` is a no-op and there is no second encoder to carry. This
  makes an MAE step markedly cheaper than an EAT step at equal `clone_batch`.
* **`clone_batch` is 1.** EAT amortizes an expensive full-grid teacher pass
  across 16 masked clones; with no teacher there is nothing to amortize, and
  clones would just be correlated gradient samples. The field is still honoured
  if set, so the choice stays ablatable.
"""

from __future__ import annotations

import torch
import torch.nn.functional as F
from torch import nn

from .model import Encoder, TransformerBlock


def patchify(spectrogram: torch.Tensor, patch_time: int, patch_freq: int) -> torch.Tensor:
    """Cut a spectrogram into the same patch order the encoder's Conv2d produces.

    `Encoder.patches` is `Conv2d(stride=(pt, pf)).flatten(2).transpose(1, 2)`,
    which walks the grid row-major over (time, frequency). The reconstruction
    target has to agree with that ordering token for token or the loss is
    computed against permuted patches -- a bug that trains to a plausible-looking
    loss curve and a useless encoder, so it is asserted in the smoke test.

    Returns `(batch, num_patches, patch_time * patch_freq)`.
    """
    if spectrogram.ndim != 4 or spectrogram.shape[1] != 1:
        raise ValueError(f"expected (B, 1, T, F), got {tuple(spectrogram.shape)}")
    batch, _, frames, mels = spectrogram.shape
    if frames % patch_time or mels % patch_freq:
        raise ValueError(f"grid {frames}x{mels} is not divisible by patch {patch_time}x{patch_freq}")
    time_patches, freq_patches = frames // patch_time, mels // patch_freq
    x = spectrogram.reshape(batch, 1, time_patches, patch_time, freq_patches, patch_freq)
    x = x.permute(0, 2, 4, 3, 5, 1)  # (B, T', F', pt, pf, 1)
    return x.reshape(batch, time_patches * freq_patches, patch_time * patch_freq)


class RandomMasker(nn.Module):
    """Uniform random masking at a fixed ratio, matching InverseBlockMasker's ABI.

    MAE's masking is deliberately unstructured: He et al. find random masking at
    a high ratio beats block masking for reconstruction pretraining, and it is
    the choice the Dasheng-style audio MAE inherits. Using EAT's inverse block
    masker here would confound the objective contrast with a masking contrast.
    """

    def __init__(self, grid: tuple[int, int], mask_prob: float):
        super().__init__()
        self.grid = grid
        self.patches = grid[0] * grid[1]
        self.visible = int(self.patches * (1.0 - mask_prob))
        if self.visible < 1:
            raise ValueError(f"mask_prob {mask_prob} leaves no visible patches on {grid}")
        # runner.py records this for the events log; constant for random masking.
        self.centers_per_mask = 0

    def forward(self, batch: int, device: torch.device, generator: torch.Generator | None = None):
        scores = torch.rand((batch, self.patches), device=device, generator=generator)
        ids_shuffle = scores.argsort(dim=1)
        ids_keep = ids_shuffle[:, : self.visible]
        mask = torch.ones((batch, self.patches), dtype=torch.bool, device=device)
        mask.scatter_(1, ids_keep, False)
        return ids_keep, ids_shuffle.argsort(dim=1), mask


class MAEDecoder(nn.Module):
    """Asymmetric transformer decoder: narrow, shallow, discarded after training.

    Built from the same `TransformerBlock` as the encoder so there is one
    attention implementation in the codebase rather than two that can drift.
    Positions are learnable here rather than sincos because the decoder sees a
    restored full-length sequence whose tokens are a mix of encoded patches and
    a shared mask token, and the mask token needs a position signal it can use.
    """

    def __init__(self, embed_dim: int, decoder_dim: int, depth: int, num_heads: int,
                 mlp_ratio: float, num_patches: int, patch_values: int):
        super().__init__()
        self.decoder_dim = decoder_dim
        self.project_in = nn.Linear(embed_dim, decoder_dim)
        self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_dim))
        self.pos_embed = nn.Parameter(torch.zeros(1, num_patches, decoder_dim))
        self.blocks = nn.ModuleList(
            [TransformerBlock(decoder_dim, num_heads, mlp_ratio) for _ in range(depth)]
        )
        self.norm = nn.LayerNorm(decoder_dim, eps=1e-6)
        self.project_out = nn.Linear(decoder_dim, patch_values)
        nn.init.normal_(self.mask_token, std=0.02)
        nn.init.normal_(self.pos_embed, std=0.02)

    def forward(self, visible: torch.Tensor, ids_restore: torch.Tensor) -> torch.Tensor:
        x = self.project_in(visible)
        batch, kept, dim = x.shape
        num_patches = ids_restore.shape[1]
        mask_tokens = self.mask_token.expand(batch, num_patches - kept, dim)
        x = torch.cat((x, mask_tokens), dim=1)
        x = x.gather(1, ids_restore.unsqueeze(-1).expand(-1, -1, dim))
        x = x + self.pos_embed.to(x.dtype)
        for block in self.blocks:
            x, _ = block(x)
        return self.project_out(self.norm(x))


class MAEPretrainer(nn.Module):
    """Student encoder + asymmetric decoder, trained to reconstruct masked patches."""

    def __init__(self, config):
        super().__init__()
        m = config.model
        o = config.objective
        grid = config.patch_grid
        self.config = config
        self.embed_dim = m.embed_dim
        self.clone_batch = o.clone_batch
        self.num_patches = config.num_patches
        self.visible_patches = config.visible_patches
        self.norm_pix = o.mae_norm_pix
        self.patch_time, self.patch_freq = m.patch_time, m.patch_freq

        self.student = Encoder(
            m.embed_dim, m.depth, m.num_heads, m.mlp_ratio, m.patch_time, m.patch_freq, grid
        )
        self.masker = RandomMasker(grid, o.mask_prob)
        self.decoder = MAEDecoder(
            m.embed_dim, o.mae_decoder_dim, o.mae_decoder_depth, o.mae_decoder_heads,
            m.mlp_ratio, config.num_patches, m.patch_time * m.patch_freq,
        )

    def forward(self, spectrogram, ids_keep, ids_restore, mask):
        dim = self.embed_dim

        patches = self.student.patches(spectrogram) + self.student.pos_embed
        if self.clone_batch > 1:
            patches = patches.repeat_interleave(self.clone_batch, dim=0)
        x = patches.gather(1, ids_keep.unsqueeze(-1).expand(-1, -1, dim))
        x = torch.cat((self.student.cls_token.expand(x.shape[0], -1, -1), x), dim=1)
        for block in self.student.blocks:
            x, _ = block(x)

        # Drop the cls token before decoding: it carries no position and the
        # decoder's restore index is defined over patch tokens only.
        prediction = self.decoder(x[:, 1:], ids_restore)

        target = patchify(spectrogram, self.patch_time, self.patch_freq)
        if self.clone_batch > 1:
            target = target.repeat_interleave(self.clone_batch, dim=0)
        if self.norm_pix:
            # Per-patch standardization. Without it the loss is dominated by the
            # loudest patches and the encoder spends capacity on gain rather than
            # structure; this is the `norm_pix_loss` of He et al. and the
            # `per_patch` setting the Dasheng-style recipe uses.
            mean = target.mean(dim=-1, keepdim=True)
            var = target.var(dim=-1, keepdim=True, unbiased=False)
            target = (target - mean) * torch.rsqrt(var + 1e-6)

        frame_loss = F.mse_loss(prediction[mask].float(), target[mask].float())
        utterance_loss = torch.zeros((), device=frame_loss.device, dtype=frame_loss.dtype)
        return frame_loss, frame_loss.detach(), utterance_loss.detach()

    @torch.no_grad()
    def update_teacher(self, decay: float) -> None:
        """No-op. MAE regresses the input; there is no EMA target to advance."""
        return