FlashMask: Importance and Uncertainty-Guided Selective Mask Learning for Semi-Structured LLM Pruning
Abstract
Semi-structured pruning enables hardware-friendly acceleration of large language models (LLMs), but learning high-quality sparse masks remains costly. Learnable-mask methods refine masks end-to-end with frozen LLM weights, yet they maintain dynamic mask distributions for all weight blocks throughout training, which makes the schedule long and expensive. We propose FlashMask, an importance- and uncertainty-guided selective mask-learning framework. Starting from a SparseGPT prior, FlashMask uses the prior pruning loss to modulate the prior strength of each block and to construct per-matrix candidate pools. After global exploration, it retains dynamic mask refinement only for blocks with small top-two gate-logit margins within these importance-filtered pools, while freezing all remaining masks into fixed sparse weights. Phase B still optimizes the full-model language-modeling objective, but bypasses Gumbel sampling, candidate-pattern mixing, and direct gate-gradient flow for inactive blocks, which both shortens the viable schedule and lowers the cost of each remaining step. The resulting quality–cost trade-off improves in both directions. On LLaMA-2 7B and LLaMA-3 8B under sparsity, FlashMask is consistently cheaper and better than MaskLLM at a matched training budget; most importantly, its -step schedule is already competitive with a -step MaskLLM run while using only about a quarter of the compute ( less). We therefore adopt steps as an efficient default across models.
est. 32% chance this paper gets accepted at ICLR 2027.
What do you think this paper will get?
All positions stay anonymous.