MaskAlign: Token-Subset Representation Alignment for Efficient Diffusion Training
Abstract
Representation alignment with pretrained vision models has recently shown strong potential for accelerating diffusion transformer training. By aligning intermediate diffusion features with representations of clean images from pretrained vision encoders, existing methods improve convergence and generation quality. Recent studies further suggest that the effectiveness of representation alignment is closely related to spatial structure. In this paper, we study representation alignment across different samples. We find that, when representation alignment is applied to all tokens, tokens with large alignment gradient norms exhibit a stable spatial preference across different samples, suggesting that the alignment objective does not affect all tokens uniformly and may encourage the model to fit spatial patterns associated with the encoder rather than visual features specific to each image, which may interfere with the learning of representations that reflect image content. To address this issue, we propose MaskAlign, a representation alignment method that applies alignment to randomly sampled token subsets during training. By exposing the model to different token subsets across iterations, MaskAlign reduces the repeated reinforcement of stable spatial preferences and encourages the model to learn features that are more closely related to individual images. To mitigate the information loss caused by directly dropping tokens, we further introduce a lightweight token mixing block that shares information across tokens before masking. Experiments on ImageNet show that MaskAlign consistently improves training convergence and generation quality. On SiT-XL/2, applying MaskAlign to REPA reduces FID from 7.9 to 5.5 at 400K iterations and from 6.4 to 4.9 at 1M iterations. When applied to REG, MaskAlign reduces FID from 3.4 to 2.8 at 400K iterations and from 2.7 to 2.4 at 1M iterations, while reducing training time per step by 11.7%. For REG, the improvement also holds when FID is computed using DINOv2-L, CLIP-L, and VGG16 features, both without and with classifier-free guidance.
est. 32% chance this paper gets accepted at ICLR 2027.
What do you think this paper will get?
All positions stay anonymous.