GSM LatentMAE encoders for the MMD term
Two encoders, for two different latent spaces. Pick by the space your distillation runs in.
| folder | latent space | file | recipe |
|---|---|---|---|
gpt2-base_per-channel/ |
gpt2-base (d=768), per-channel normalization from stats/latent_stats_gpt2_gsm.pt (pos-0 row separate), 512 canvas β yrELF branch mmd_onestep |
dmae_gptbase_lr1e4.pt + latent_mae.py (standalone, no repo imports) |
Drifting-MAE, MSE reconstruction, 4000 steps, eff. batch 192, lr 1e-4, patch 8, mask 0.15. Paired AUC 0.79 |
gpt2-large_ceonly/ |
gpt2-large (d=1280), scalar latent_mean/latent_std, 640 canvas β the earlier mmd-gsm setup |
dmae_ceonly.pt + latent_mae.py + usage.py |
token-CE only, 12000 steps. Details below (original README) |
gpt2-base_per-channel/latent_mae.py documents the exact input contract and the feature map /
loss the no-teacher run used (forward_features, taps (1,3), chunks (1,4) -> (B, 9216);
pseudo-Huber c=0.001, weight 1.0, on the clean real/fake pair). Its normalize_latents is a
verbatim port of load_latent_stats(mode="per_channel"); if you already run
resolve_latent_stats + encode_text you get the same tensors.
(original) # GSM LatentMAE β CE-only encoder (dmae_ceonly)
A frozen feature extractor for the MMD term in ELF/flow-map distillation on GSM. It maps a gpt2-large latent to features that the MMD is computed in, instead of using the teacher's mid-trunk features.
Trained with a learned token head and zero latent MSE β the objective is discrete (predict the token) rather than continuous reconstruction. That is what makes it the best of the encoders we tried.
- arch
plain, 112M params cfg:latent_dim=1280, width=768, depth=8, max_len=640, vocab_size=50257- expects gpt2-large latents, normalised by the run's
latent_mean/latent_std
Measured results
Encoder-quality probe (higher is better):
| encoder | paired | centered AUC |
|---|---|---|
| dmae_ceonly | 0.8250 | 0.9157 |
| dmae_ce_disc | 0.8300 | 0.9139 |
| dmae_lr1e4 (MSE) | β | 0.903 |
| dmae_disc | β | 0.8743 |
Downstream, encoder-only MMD (mmd_weight=0), GSM8K test 1319 rows, 3 seeds,
N=32 / shift 128:
| arm | acc |
|---|---|
| ceonly, taps (1,3) | 29.19 |
| disc | 28.99 |
| ceonly, tap (3) | 28.76 |
| ce+disc | 28.69 |
| ceonly, tap (1) | 28.67 |
| MSE baseline | 28.61 |
Caveats worth knowing before you spend runs on it:
- The whole family sits in a 28.6β29.2 band (sd 0.4β1.3), so the ranking above is within noise. Four quite different encoder objectives produced statistically indistinguishable students.
- Encoder quality does not predict downstream quality.
dmae_discdetects corrupted programs at 0.989 AUC (vs 0.67 for this one) and still lands mid-pack downstream. - Taps matter as much as the encoder: (1,3) > (3) > (1), and deeper taps are clearly worse β (5,6) scored 25.70 vs 28.61 for (1,3).
- On our time-grid recipe this encoder was ~3.7 pts below teacher features. On a fixed-point/resample recipe the gap nearly closed (32.55β33.16). So it is recipe-sensitive.
Use
See usage.py. Two entry points, and picking the wrong one is the most common
mistake:
forward_features(...)->(B, D)pooled, one sample per example. Use with a paired/linear kernel.token_features(...)->(B, S, D)per token. Use with an RBF kernel where samples are a row's tokens β pooled features leave 1 sample per row and the RBF estimator needs >= 2.