Factorized embeddings and GDN-2 recurrence on single-GPU budgets (~125M)

#1
by AndrewThompson1233 - opened

Hi Yang,

Really impressive work squeezing a from-scratch 125M pretraining run onto a single RTX 4090 in ~40 hours. Transparently reporting negative results (the DPO and DeepSeek distillation plateau) alongside the positive interpolation results is rare and refreshing to see.

I've been working on a similar ~100M parameter budget constraint in an open architecture called Maba (101M reference model):
https://huggingface.co/AndrewThompson1233/maba-v1-architecture

A couple of architectural observations that might be relevant for future single-GPU runs:

  • Factorized embeddings: with a 32k vocab and 576 hidden dim, standard tied embeddings take 18.4M parameters (14.8% of your total budget). Using a low-rank bottleneck (rank 128) drops that tax to ~4.3%, letting you redirect roughly 13M parameters straight into deeper/wider reasoning layers without increasing FLOPs.
  • Hybrid GDN-2 + GQA: replacing 75% of attention layers with Gated DeltaNet linear recurrence keeps the recurrent state O(1) and significantly reduces memory bandwidth bottlenecks during both single-GPU pretraining and generation.
  • Native MTP (k=2): speculative decoding heads built into the backbone to boost single-card inference without needing an external draft model.

Curious: during the 40-hour run on the 4090, what kind of MFU / token throughput were you averaging on the 2,048 sequence length, and did you hit memory bandwidth walls during the 30-layer GQA backward pass?

Best,
Andrew

Hi Andrew,

Many thanks — I really appreciate you taking the time to read through the project in that much detail.

Your three suggestions are very valuable to the kind of constrained small-model regime I'm interested in. The factorized-embedding point in particular made me realize how much of the 125M parameter budget is being spent on the vocabulary interface rather than the backbone. The GDN-2/GQA hybrid and native MTP directions are also things I hadn't explored in PetitGPT research-v1.

I'm freezing that release for now rather than retrofitting new architectural ideas into it, but I'm definitely going to study Maba more closely as I think about future small-model experiments. If I end up building on or directly comparing against these ideas, I'll of course cite/acknowledge Maba appropriately.

On throughput and MFU, I went back through the retained logs. Across the two production stages, the ~13B-token run took about 39 hours of active training time. Logged throughput was very stable at about 93.5k tokens/s on average, with a median of about 93.8k tokens/s.

I didn't log MFU directly. Using an explicit FLOP estimate for this architecture, I get roughly 0.96 GFLOP/token for training, corresponding to about 89.7 TFLOP/s, or an estimated ~54% MFU against the RTX 4090's 165.2 TFLOP/s dense BF16 Tensor Core peak with FP32 accumulation. I would treat that as a back-calculated estimate rather than a measured MFU.

With micro-batch size 8 and torch.compile enabled, peak reserved VRAM was about 19.36 GiB.

On the bandwidth question, though, I can't give you a definitive answer. I never ran Nsight Systems/Compute or collected DRAM-bandwidth, SM-occupancy, Tensor-Core-activity, or roofline counters. The batch-size scaling is consistent with the smaller GEMMs/kernels being under-utilized at small batches, but it isn't enough to tell me whether DRAM bandwidth was actually the dominant bottleneck during backward.

One implementation detail I would definitely revisit is GQA: this version explicitly expands the 3 KV heads to 9 with repeat_interleave before SDPA. That's a plausible source of avoidable memory traffic, but I never profiled its actual cost.

The GDN-2 direction is especially interesting in that context. At roughly the 100M scale and ~2k context, do you have a measured training-throughput comparison between the 75% GDN-2 / 25% GQA hybrid and an otherwise comparable pure-GQA baseline? I'd also be curious whether the rank-128 factorized embedding showed any measurable quality tradeoff once you reinvested the saved parameters into the backbone.

And for the native k=2 MTP head, have you measured the realized decoding speedup on a single GPU relative to ordinary autoregressive decoding?

Best,
Yang

Hi Yang,

Thanks for sharing the exact breakdown! ~93.5k tok/s at ~54% MFU on a single 4090 with torch.compile is a great reference point. Also, regarding GQA: PyTorch's native SDPA supports GQA/MQA broadcasting out of the box without needing an explicit repeat_interleave, which should save you some avoidable DRAM bandwidth.

To your questions — we actually evaluated this empirically on the 101M checkpoint trained on TinyStories (open weights and full breakdown are up on AndrewThompson1233/maba-101m):

  1. Quality tradeoff of rank-128 factorized embeddings
    There was no penalty — in fact, it took Rank 1 overall. With the ~13M parameters reinvested into the backbone and combined with 2-pass block recycling (40 effective layers from 20 physical blocks), Maba v1.1 achieved a validation perplexity of 357.34 (vs 382.84 for pure-GQA MiniCPM5 and 470.51 for Qwen 3.8) and led on ARC-Easy (26.80%) and reasoning margin (+0.2237). At ~100M, shrinking the vocab bottleneck to 4.3% is an unambiguous net win.

  2. Throughput: 75% GDN-2 / 25% GQA vs Pure GQA
    On single-GPU generation (NVIDIA L4 bfloat16), Maba v1.1 clocked 394.2 tok/s compared to 278.2 tok/s for the pure-GQA MiniCPM5 baseline (+41.7% throughput), while slashing active 4k KV-cache from 42.0 MB down to 10.0 MB (-76.2%). The constant O(1) recurrent state takes strictly 1.17 MB and never expands.

  3. Native MTP (k=2) Speedup
    The auxiliary MTP head costs only 492k parameters (~0.49% of the budget). Because speculative verification happens in the same forward graph without swapping KV caches or orchestrating a separate draft worker, it delivers ~1.5x to 1.7x empirical wall-clock generation speedup on low-entropy / structured sequences.

Feel free to poke around maba-v1-architecture and maba-101m if you want to inspect the safetensors or the layer configs!

Best,

Andrew

Hi Andrew,

Thanks — this is extremely helpful, especially having the numbers from a matched ~100M-parameter setting.

And thanks for pointing out the native SDPA GQA path. I had kept the explicit repeat_interleave implementation from an earlier version of the model, so using enable_gqa=True with the fused SDPA path is definitely something I would change in a future implementation rather than materializing the KV heads myself.

The factorized-embedding result is particularly interesting to me. At this scale, recovering on the order of 10M+ parameters from the vocabulary interface is a very large architectural budget, so I'd like to understand that tradeoff much better rather than simply treating full-width embeddings as a default.

The GDN-2 numbers are also a strong motivation for me to look more seriously at hybrid recurrent/attention architectures. I'll spend some time going through both maba-v1-architecture and the 101M checkpoint/configs. For a future small-model experiment, a controlled pure-GQA vs hybrid comparison — along with a factorized-embedding ablation — would be a very interesting direction.

The MTP result is also impressive. I hadn't seriously considered integrating speculative decoding into the backbone itself before this discussion.

Thanks again for sharing all of this. If I end up building on or comparing against these ideas in a future project, I'll make sure to reference Maba appropriately.

Best,
Yang

Hi Yang,

I came across your Linkedin post about PetitGPT release and training the 124.6M-parameter model from scratch on a single RTX 4090. I liked that the release includes not only the weights but also the tokenizer, native inference implementation, checksums, provenance, and documented limitations.

I maintain an open-source project called kaggle-vllm, which provides a CUDA-first runtime for upstream vLLM on Kaggle's dual NVIDIA Tesla T4 GPUs (SM75). I was curious about how portable PetitGPT's released checkpoint would be to that environment, so I ran a small compatibility experiment against your pinned research-v1 release rather than modifying the checkpoint.

I was able to verify the complete PetitGPT release with its published SHA256 checksums and run the native PyTorch fp32_math inference path successfully on a Kaggle Tesla T4. The Kaggle environment in the experiment was Python 3.12.13, PyTorch 2.10.0+cu128, CUDA 12.8, with 2× Tesla T4 GPUs available. The actual PetitGPT native smoke test was deliberately restricted to GPU 0, so I am not claiming that part as dual-GPU inference.

I then tested the exact same checkpoint through the vLLM runtime provided by kaggle-vllm. I used TP=1 because PetitGPT has 9 query attention heads, so TP=2 cannot evenly partition the heads. The direct vLLM load currently stops before CUDA model execution: Transformers/vLLM cannot recognize the released config because it does not expose a registered model_type/model integration. That matches the compatibility limitation you document in the model card.

So the result is actually fairly clean:

PetitGPT native PyTorch on Tesla T4: works.
kaggle-vllm CUDA runtime on dual T4/SM75: works.
Direct PetitGPT → vLLM model construction: currently needs a proper PetitGPT integration layer.

I deliberately did not relabel PetitGPT as Llama or modify your config just to force it through vLLM, because I wanted to preserve the architecture and checkpoint semantics. I also did not generate a vLLM sharded_state, since that should only happen after successful vLLM model construction.

I have documented the experiment in a reproducible Kaggle notebook in my kaggle-vllm repository, including the pinned PetitGPT revision, architecture checks, native inference result, exact vLLM loader traceback, and machine-readable compatibility report:

Notebook: https://github.com/kaggle-vllm/kaggle-vllm/blob/main/kaggle-notebooks
kaggle-vllm: https://github.com/kaggle-vllm/kaggle-vllm
Your PetitGPT release used in the experiment: https://huggingface.co/yqi0/petitgpt
Google Drive: https://drive.google.com/file/d/1wiui9SQWa3m5RVjIqIrdKx2f4r6q3Tns/view?usp=sharing

One interesting next step would be a faithful PetitGPT adapter/model registration for vLLM—ideally without changing the original checkpoint format—followed by TP=1 T4 inference measurements and only then a vLLM-native sharded-state save/reload experiment.

If that direction is interesting to you, I would be happy to compare notes on the model implementation and checkpoint mapping. I wanted to share the current result first because it gives a fairly precise boundary between CUDA/runtime portability and model-framework integration.

You may download the zip file from my Google Drive link with all the necessary output I have generated through my Python SDK kaggle-vllm in Kaggle platform.

Kindly, go through the Google drive file and give me your feedback.

Best,
Mohammad Waqas

Hi Mohammad,

Thank you very much for the thoughtful experiment and for sharing the notebook and evidence bundle. I really appreciate the care you took to preserve the released checkpoint and distinguish the native single-T4 result from the vLLM loading attempt.

The reported fp32_math run on a T4 is useful portability feedback. Your results also help clarify the current integration boundary: model/config recognition is the first blocker for direct vLLM loading, while PetitGPT execution through vLLM remains untested.

A faithful vLLM integration sounds like an interesting direction. My main suggestion would be to check it against the native implementation before benchmarking, using identical input token IDs and carefully checking the weight mapping, with numerical-precision differences taken into account.

My bandwidth for new development is limited at the moment, but I’m really glad you shared this work. It’s encouraging to see PetitGPT being explored on other hardware, and this is exactly the kind of concrete feedback that makes sharing the full implementation worthwhile.

Thanks again for taking the time to test it and document the results so carefully.

All the best,
Yang

Sign up or log in to comment