Merged `main` at `e094f0860`. Main's expert parameters remain in the
transformer block's FSDP group; prefetching retains that behavior and
Qwen3.8's final hyper-connection mixer handling. Adds Qwen3.8 Flash Next
text training for `qwen4_exp` and `qwen4_exp_text`: hyper-connection
residual streams, Gated DeltaNet, indexed sparse attention, PLE N-gram
embeddings, and sigmoid-gated MoE, with HF checkpoint conversion.
## Head parallelism
PLE uses ordinary `nn.Embedding` parameters. `EmbeddingParallel`, an
external `RowwiseParallel` style, shards vocabulary rows across
`dp_shard_cp`, independently of EP. FSDP handles precision, replication,
and parameter lifetime on the complementary mesh. Packed token lengths
can differ across data ranks, so the style pads before collectives and
crops local outputs. The requested TODO marks the host synchronization
needed to size this padding.
When the head mesh has size one, the embedding follows ordinary FSDP and
no extra head dimension is created. The extra head axis is appended
after the existing world-mesh dimensions, preserving their indices.
Optimizer code is unchanged. Dion indexes each parameter's DTensor mesh,
not the root mesh; its momentum tensors retain the parameter placements.
Four-rank CPU/Gloo checks against main passed for six
DP/EP/CP/replication layouts, including momentum local shard sizes and
head-parallel state. Qwen3.5 and the shared MoE implementation are
unchanged; Qwen3.8's gated MoE lives in its own model package. Skills
and unrelated comments are unchanged.
## Required integration changes
- **Gradient clipping:** TorchTitan separates dense and EP gradients but
cannot combine the additional head mesh: its `aten.stack` raises “All
operands … must have the same mesh.” Resolve norms separately on each
DTensor mesh, then combine them into one global L2 norm and apply the
existing clipping formula. This is the same clipping mathematics; no
threshold, optimizer, `optimization_dtype`, or `reduce_dtype` changes.
The two-rank CPU/Gloo reproduction fails with the original helper and
matches the unsharded norm reference after mesh reductions. B300
gradient/update parity passed in SLURM 1440.
- **NCCL buffers:** a PLE-containing layer requires a 100.26 GiB
concatenation under the original protocol. Bound same-dtype groups to
512 MiB and send larger individual tensors without a concatenation copy.
Small groups retain the same collective count. Prior two-GPU A/B receive
medians were 0.365/0.367 ms at 128 MiB, 1.032/1.033 ms at 512 MiB, and
2.027/1.707 ms for a 1 GiB tensor plus small tensors (original/bounded;
SLURM 1347).
- **vLLM compatibility:** filter PLE checkpoint pieces with no overlap
on the local TP rank before vLLM's deferred loader buffers them. The
original loader still handles valid local pieces and malformed shards.
The generic reload machinery is unchanged. Real Qwen PLE reload values,
storage identity, and queue behavior passed on four GPUs (SLURM 1352).
- **Model setup:** register the config/model, freeze its sparse indexer,
wrap its final hyper-connection mixer with FSDP, and register its
indexed-attention operators with activation checkpointing. `vlm.py` only
supplies composite-config text-backbone lookup; this PR adds text
training, not a vision implementation.
## Validation
The merge with `e094f0860` was checked for unresolved conflicts and
whitespace errors; tests were not rerun. The results below remain from
their stated pre-merge commits/jobs.
SLURM **1442** completed successfully at cleanup commit `e50e42316`
(exit `0:0`, 3269 seconds including startup/shutdown). The prerequisite
four-B300 mesh/Muon-state/vLLM-reload check **1440** also passed.
- Official Qwen3.8 Flash Next checkpoint; 16 trainer B300 GPUs and TP8
inference.
- **20 math RL steps, batch 64**, group 4, `math-env` /
`PrimeIntellect/Hendrycks-Math` training split, debug algorithm, default
policy lag 8.
- **GPU memory utilization 0.85**, CUDA graphs and default vLLM
batching. Context 4096, completion limit 3072, thinking disabled;
SignSGD at `1e-6`, BF16 router projections on both sides, unchanged FP32
optimization/reduction defaults.
- GPU-resident PLE on vLLM 0.30 with explicit
`engram_config.cpu_offload=false`. NCCL transfers, no checkpoint saves
or filesystem weight transfer. `FLA_DISABLE_TENSOR_CACHE=1` is explicit
in the run environment; launcher defaults are unchanged.
- All 20 mean mismatch KL values are below **0.015**. Maximum
**0.001570820**; mean across steps **0.000701912**.
- Median warmed-up step **56.1s**, median broadcast **41.5s**.
- SLURM 1440 checked dense/EP/head norms, clipped gradients and
parameter updates against full-tensor references, plus uneven token
counts, replication/CP topologies, FSDP CPU offload, Muon momentum shard
sizes against main, and vLLM 0.30 PLE reloads. The original clipping
helper reproduced the mixed-mesh error; dense/EP-only norms matched the
original helper.
- Four config/checkpoint conversion cases, unchanged single-rank Muon
construction, Ruff and whitespace checks passed.
- Artifacts: `outputs/qwen38-validation/runs/qwen38-math20-main-gpu/`
(`validation-summary.json`, `mismatch-kl.md`, resolved configs and
logs); mesh log `outputs/qwen38-validation/rebased-head-1440.log`.
| Step | Mean mismatch KL |
| --- | ---: |
| 1 | 0.000282296 |
| 2 | 0.000228267 |
| 3 | 0.000326421 |
| 4 | 0.000275441 |
| 5 | 0.000435869 |
| 6 | 0.000398891 |
| 7 | 0.000725010 |
| 8 | 0.000521352 |
| 9 | 0.001052707 |
| 10 | 0.000689109 |
| 11 | 0.000823852 |
| 12 | 0.000369604 |
| 13 | 0.001008985 |
| 14 | 0.000577958 |
| 15 | 0.001168500 |
| 16 | 0.000841276 |
| 17 | 0.001036190 |
| 18 | 0.000672617 |
| 19 | 0.001570820 |
| 20 | 0.001033083 |
## Kernel dependency
[prime-kernels
#8](https://github.com/PrimeIntellect-ai/prime-kernels/pull/8) is merged
and released as
[v0.1.0-6f48176](https://github.com/PrimeIntellect-ai/prime-kernels/releases/tag/v0.1.0-6f48176).
Both x86_64 and aarch64 wheel URLs and SHA256 hashes are pinned in
`pyproject.toml` / `uv.lock`; the submodule matches release commit
`6f48176`. The release tree is identical to the `7bd1e89` source used in
the 20-step math run above, and the indexed-attention Python files in
both downloaded wheels match that source byte for byte.
The release wheel installed successfully with `uv sync --all-extras
--all-packages --frozen`; `uv lock --check` and `git diff --check`
passed. B300 validation of the installed x86_64 release wheel (SLURM
1468) passed forward/backward numerical comparisons, selection and tie
handling, fullgraph forward/backward compilation, and selective
activation checkpointing. The 20-step KL table above remains from SLURM
1442; it was not rerun for this wheel-only update.
The temporary Qwen overlap shim can be removed when pinned vLLM includes
the overlap check in `Qwen4ExpNGramEmbedding.load_weights`.
<!-- CURSOR_SUMMARY -->
---
> [!NOTE]
> **Medium Risk**
> Large new model and head-parallel embedding path, plus mesh-aware
gradient clipping and NCCL bucketing that can affect any run using these
parallelism modes or very large PLE layers.
>
> **Overview**
> Adds a **native trainer implementation** for Qwen3.8 Flash Next
(`qwen4_exp` / `qwen4_exp_text`): hyper-connection residual streams,
Gated DeltaNet linear layers, **indexed sparse attention** (via
`prime_kernels`), PLE N-gram embeddings, and sigmoid-gated MoE, plus
HF↔Prime checkpoint conversion (including sharded `ngram_embedding`
weights).
>
> **Distributed setup** introduces vocabulary-parallel
**`EmbeddingParallel`** for PLE embeddings on a new **head** mesh
(`dp_shard_mod_head`), FSDP on those tables, freezing of
`SparseAttentionIndexer`, and FSDP/prefetch handling for the final
**`hyper_connection_mixer`**. Activation checkpointing now saves
**`prime_kernels`** indexed-attention ops.
>
> **Training/inference plumbing:** **`clip_grad_norm_`** computes norms
per DTensor mesh then combines them (replacing TorchTitan’s EP-only
helper). NCCL weight send/receive uses **`iter_tensor_buckets`**
(512 MiB) to avoid huge single concatenations. vLLM gets a **PLE shard
filter** patch so non-local TP pieces are skipped before load.
>
> Configs/models are registered in AutoConfig/custom LM/VLM maps;
`qwen4_exp` is listed in the VLM registry for composite configs (text
training only in this PR). Unit test covers config and checkpoint
round-trip.
>
> <sup>Reviewed by [Cursor Bugbot](https://cursor.com/bugbot) for commit
e50e423169964b8e8a800a610cf8e77d4c56053f. Bugbot is set up for automated
code reviews on this repo. Configure
[here](https://www.cursor.com/dashboard/bugbot).</sup>
<!-- /CURSOR_SUMMARY -->
---------
Co-authored-by: matejsirovatka <matejsirovatka@users.noreply.github.com>