# MUDD-gated & Lightweight DC

This record tracks three iterations of the same idea:

- `this_pr_v1`: initial MUDD gate + lightweight layer-10 DC on top of the `2026-05-07_XSAGatedLayers` short-track baseline, plus cleanup that removes some redundant XSA and `attn_gate_w` attention-gate paths.
- `this_pr_v2`: makes DC gate-generated by the MUDD gate and cuts the schedule to `1315 + 10 = 1325` steps.
- `this_pr_v3`: rebases the idea onto the latest upstream base with the bigram sign trick and FP8 MLP changes, then adjusts the bigram configuration so MUDD gate can work well with that base.

The current submitted version is `this_pr_v3`.

## Results

Older v1/v2 runs were based on the `2026-05-07_XSAGatedLayers` short-track baseline:

```text
                     Runs  Steps  Time mean  Time sd  Time delta  Loss mean  Loss sd  Loss p
xsa-baseline-s1385     10   1385    81.2187   0.0681      0.0000    3.27841  0.00133  2.17e-03
xsa-baseline-s1335      9   1335    78.5576   0.0401     -2.6611    3.28477  0.00168  1.00e+00
this_pr_v1             11   1335    78.9669   0.0425     -2.2518    3.27714  0.00123  8.15e-06
this_pr_v2             10   1325    77.7050   0.0297     -3.5137    3.27808  0.00098  7.76e-05
```

v3 is based on the newer upstream bigram-sign + FP8-MLP base:

```text
                     Runs  Steps  Time mean  Time sd  Time delta  Loss mean  Loss sd  Loss p
upstream-sign-fp8       -   1390    79.2000      -        0           -        -         -
this_pr_v3             10   1290    75.5199   0.0487     3.6801    3.27851  0.00097  4.53e-04
```

In the first table, `Time delta` is relative to `xsa-baseline-s1385`. In the second table, `Time delta` is relative to `upstream-sign-fp8` and is shown as positive speedup. `Loss p` is:

```python
scipy.stats.ttest_1samp(losses, 3.28, alternative="less").pvalue
```

`this_pr_v3` passes the `p < 0.01` target rule with `p = 4.53e-04`.

The `upstream-sign-fp8` row is the copied upstream result for the latest bigram-sign + FP8-MLP base. It was not rerun in this folder, so it is only a timing reference and does not have repeated loss logs here. Compared with that timing, `this_pr_v3` is **3.68s faster** (`79.2000s -> 75.5199s`).

Compared with `this_pr_v2` for context, v3 is **2.19s faster** while giving up only `+0.00043` mean loss, which is not significant in these samples (two-sided Welch `p = 0.336`). This is not a strict apples-to-apples base comparison because v2 predates the upstream bigram-sign / FP8-MLP base. The size of the speedup lines up with the rebase benefit: roughly **0.5s** from FP8 MLP plus roughly **1.5s** from the bigram trick, so the v2 -> v3 timing gain should mostly be read as the upstream rebase benefit rather than a new MUDD/DC speed improvement.

## v1: MUDD Gate + Lightweight DC

v1 combines two model additions with cleanup of redundant paths:

- Add a lightweight layer-10 DC correction after the base FA3 attention output.
- Add a small MUDD-style gate stack that generates per-token coefficients for selected XSA gates, attention gates, x0/bigram injections, and the layer-3 to layer-6 skip.

The cleanup is based on the idea that DC, XSA, and `attn_gate_w` all modulate the attention stream. Once layer 10 gets the DC correction, some of the older XSA / attention-gate modulation looked redundant, so v1 keeps only a smaller gated subset and lets the MUDD gate generate the active coefficients.

The v1 MUDD gate layout is:

- pre gate from `x0`: XSA for layers `1, 3, 4`, attention gate for layer `3`, and x0/bigram injection gates for layers `0..5`.
- post gate generated at layer `6`: XSA for layer `7`, attention gate for layer `10`, x0/bigram injection gates for layers `7..10`, and the layer-3 to layer-6 skip coefficient.

After this cleanup, the active XSA layers are `1, 3, 4, 7`. The older XSA use on layers `8` and `10` is removed from the active path. Active attention gates come from the MUDD gate on layers `3` and `10`. Fixed `x0_lambdas` / `bigram_lambdas`, static `xsa_alphas`, and standalone `skip_gate` are likewise no longer the active control paths.

The DC path is deliberately narrow: post-only, no pre-composition, no dense-dense term, only on the final non-paired attention layer. The base attention still comes from FA3; the custom Triton kernel recomputes the local-window QK softmax and adds the DC correction.

v1 cuts the schedule from `1375 + 10 = 1385` to `1315 + 20 = 1335` while improving mean validation loss. Compared with the matched `xsa-baseline-s1335`, v1 improves mean loss by about **0.00763** (two-sided Welch `p = 1.44e-08`). Known isolated experiments suggest DC alone accounts for about **0.0046** of that gap, with the remaining improvement coming from the MUDD gate and cleanup.

## v2: Gate-Generated DC

v2 keeps the same basic idea but removes the standalone DC parameter path. The layer-10 DC tensors now come directly from the MUDD gate:

```text
dc_weights[10] = (post_gate[..., 29:35], post_gate[..., 35:41])
```

`dc_gate` only validates shape, RMS-normalizes `w1` across heads, and returns contiguous `(post_w1, post_w2)`.

The MUDD gate layout is narrowed and made explicit:

- pre gate from `x0`: XSA for layers `1, 3`, attention gate for layer `3`, and x0/bigram injection gates for layers `0..3`.
- post gate at the start of layer `4`: XSA for layers `4, 7`, attention gate for layer `10`, x0/bigram injection gates for layers `4, 5, 7, 8, 9`, the layer-3 skip coefficient, and the layer-10 DC tensors.
- final-layer x0/bigram injection stays in the existing layer-10 MUDD block via `mu[10]` and `mu[11]`.

v2 uses `1315 + 10 = 1325` steps and passes the target with `p = 7.76e-05`. Compared with v1, it trades `+0.00094` mean loss for **1.26s** less wall-clock; the loss gap is not significant at 5% in this sample (two-sided Welch `p = 0.066`).

## v3: Bigram-Compatible Rebase

v3 rebases the MUDD-gate + gate-generated DC idea onto the latest upstream base with the bigram sign trick and FP8 MLP changes. The main issue after rebasing was the bigram path: upstream used `bigram_dim = 192` with `bigram_vocab_size = 50304 * 15`. With that setting, the MUDD gate did not recover the gain observed in v2, likely because the bigram dimension / vocabulary tradeoff did not let the MUDD-generated bigram gates make full use of the bigram embedding path.

v3 uses:

```python
bigram_dim = 768
bigram_vocab_size = 50304 * 15 // 2
```

This is a loss/speed tradeoff for the gated bigram path: enough width for the MUDD gate to help, but a smaller vocabulary than the original `768`-dim bigram embedding would use.

v3 also makes `_mudd_gate_scale` a learnable parameter initialized at `0.1`. This was added because training showed some instability with the MUDD gate; letting the model tune the global gate scale made the gated path easier to stabilize.

The submitted v3 schedule is `1275 + 15 = 1290` steps.

## Files

- `xsa-baseline-s1385/`: 10 baseline logs at `1375 + 10 = 1385` total steps.
- `xsa-baseline-s1335/`: 9 matched-step XSA baseline logs at `1315 + 20 = 1335` total steps.
- `this_pr_v1/`: 11 v1 logs at `1315 + 20 = 1335` total steps.
- `this_pr_v2/`: 10 v2 logs at `1315 + 10 = 1325` total steps.
- `this_pr_v3/`: 10 complete v3 logs at `1275 + 15 = 1290` total steps. Four additional files in this folder are incomplete warmup/interrupted logs and are excluded from the statistics.
- Root `train_gpt.py`: current submitted training code.
- Root `dc_triton_kernels.py`: DC correction kernel included in each PR log.

Timing environment from the logs: 8x NVIDIA H100 80GB HBM3, Python 3.12.3, PyTorch `2.10.0+cu128`, Triton `3.6.0`, driver `580.159.03`.