quinnlue's picture
MAE cross-objective mixture-ranking campaign
ffdcfe7 verified
Raw
History Blame Contribute Delete
9 kB
"""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