DiT

What one GPU reveals about training a text-to-image diffusion transformer

Training a 210M text-to-image DiT from scratch on a single GPU is a bold experiment, and the numbers are worth reading closely.

4 min readMachine Learning
What one GPU reveals about training a text-to-image diffusion transformer
Training a 210M text-to-image DiT from scratch on one GPU: what I measured [P]

Most of what gets published about training large models is either a recipe list or a highlight reel. This training report is different. The author trained a 210M-parameter diffusion transformer from scratch on a single GPU and came back with three measurements that challenge how we usually think about attention, loss curves, and sampling schedules. That is the kind of contribution that moves the field forward, not because it is flashy, but because it is honest about what actually happens under the hood. And in a moment when the community is chasing scale at any cost, there is real value in someone saying, "Here is what I saw when I looked closely." The first finding about learned null attention slots becoming the sink is worth sitting with. Register tokens can absorb excess attention mass, but the data shows that two learned key/value slots appended to cross-attention capture roughly 90% of the attention weight at mid-noise in a middle block, while the EOS token drops to 4%. That is not a small effect. It suggests that the model is not using cross-attention to retrieve content in the way we often assume. It is offloading information into these learned slots and then reading from them later. For anyone building on top of this architecture, that is a practical signal: if you want to inspect or steer what the model is attending to, those slots are where the action is. It also raises a question that deserves more attention: are we designing the right inductive biases for cross-attention, or are we just letting the sink emerge and then calling it a feature? The second measurement is the flow-matching loss as a health signal rather than a quality signal. Training loss moved from 0.805 to 0.754 while FID improved from 33.7 to 27.0 and object accuracy went from 65% to 90%. The key detail is that training and held-out loss stayed equal to the third decimal for 24 epochs. That means the loss is not telling you much about generalization. It is tracking the irreducible variance of the velocity target, not the quality of the samples. If you are monitoring a training run and watching the loss curve, this is a direct warning: do not trust that number. It will not tell you when things are working or when they are about to break. Better signals are needed, and the use of FD-DINOv2 and detector-based object accuracy is a reminder that evaluation has to be tied to the actual task, not the optimization objective. The third point, that the training-time timestep shift is worth more than doubling the number of sampling steps, is the kind of concrete result that saves people weeks of compute. With the final weights, 20 steps with shift 2.8 produced an FID of 27.0, while 50 steps without the shift gave 26.6. That is a negligible difference. But 8 steps with the shift gave 28.4, which is still usable. The shift rule from SD3/RAE, sqrt(32·32·32/4096), is simple to implement and costs nothing at inference time. If you are building a product on top of a diffusion model, this is the difference between a responsive experience and one that feels sluggish. It is also a reminder that the sampling schedule is not a minor detail; it is a first-class hyperparameter that can matter more than the number of steps. The setup is also transparent, which is refreshing. The decision to use five aspect-ratio buckets from step one, flan-t5-base frozen with long and short captions sampled at 50/40/10, and a mix of Pexels, a filtered slice of FLUX-Reason-6M, and COCO with GPT-4V captions is all documented. The code, weights, and demo are linked. That is how research should be shared. For the next phase, the question is which reward to start with for Flow-GRPO: PickScore/HPSv2, a detector-based object reward, or something verifiable like counting. Our take is straightforward. Start with the counting or detector-based reward.

From Machine Learning

I trained a 210M-parameter text-to-image diffusion transformer from scratch (3.5 days, one RTX PRO 6000, 4.2M images at 256²) mainly to understand the recipe end to end. Three measurements came out of it that I have not seen stated plainly elsewhere, so I'm posting those rather than the samples.

Read the original at Machine Learning