Search arXiv⌕ Search

arXiv subjects

Qiaozhe Zhang

Publications and source records attributed to Qiaozhe Zhang.

4 recordsLinked to original sources

Beyond the Matrix Sign: Quadratic Spectral Descent

Muon emerges as a strong competitor of the AdamW for LLM pretraining, because the matrix-wise update it employs can potentially incur smaller second-order penalty than the once dominating AdamW, which performs coordinate-wise update. However, the spectral flattening procedure in Muon is quite debatable since it discards the spectral amplitude information totally. To seek for better spectral allocation (and the associated spectral subspace), we propose to solve the quadratic model of loss function under the spectral norm constraint \textit{directly} (i.e., in a genuinely Newtonian way) and thus obtaining the Quadratic Spectral Descent (QSD) algorithm. In contrast, many existing curvature-aware methods either exploit the second-order information in an \textit{implicit} way by changing the weight update geometry (such as Mousse, FISMO) or rely on strong assumptions (such as the weight displacement isotropy assumption in Newton-Muon). QSD's potential advantage over these methods is best illustrated in the isotropic curvature scenario, where Mousse, FISMO and Newton-Muon all reduce to Muon while the spectral allocation in QSD is still \textit{non-flat} (since the spectral allocation in QSD depends on the \textit{gradient to curvature ratio}). Meanwhile, to control the complexity of QSD, we employ inversion-free K-FAC and \textit{online} Frank-Wolfe update which is essentially a matrix sign operator. Overall, the complexity increase can be rather mild. Experiments on GPT pre-training show that QSD consistently improves validation loss over Muon and recent Muon variants, while achieving up to an $8.49\%$ wall-clock speedup at matched validation loss.

cs.LG↗

Rényi Sharpness: A Novel Sharpness that Strongly Correlates with Generalization

Sharpness (of the loss minima) is widely believed to be a good indicator of generalization of neural networks. Unfortunately, the correlation between existing sharpness measures and generalization is not as strong as expected, and sometimes even contradiction occurs. To address this problem, a key observation in this paper is: what really matters for generalization is the average spread (or unevenness) of the spectrum of loss Hessian $\mathbf{H}$. For this reason, conventional sharpness measures, such as trace sharpness $\operatorname{tr}(\mathbf{H})$, which cares about the average value of the spectrum, or max-eigenvalue sharpness $λ_{\max}(\mathbf{H})$, which concerns the maximum spread of the spectrum, are not sufficient to well predict generalization. To characterize the average spread of the Hessian spectrum, we leverage the notion of Rényi entropy in information theory, which captures the unevenness of a probability vector and can thus be extended to a general non-negative vector, such as the Hessian spectrum at loss minima. Specifically, we propose Rényi sharpness, defined as the negative of the Rényi entropy of loss Hessian $\mathbf{H}$. Extensive experiments demonstrate that Rényi sharpness exhibits strong and consistent correlation with generalization in various scenarios. Moreover, two generalization bounds with respect to Rényi sharpness are established by exploiting its desirable reparametrization invariance property. Finally, as an initial attempt to exploit Rényi sharpness for regularization, Rényi Sharpness Aware Minimization (RSAM) is proposed, where a variant of Rényi sharpness is used as the regularizer. RSAM is competitive with state-of-the-art SAM algorithms and far better than conventional SAM based on max-eigenvalue sharpness.

cs.LG↗

How Sparse Can We Prune A Deep Network: A Fundamental Limit Perspective

Network pruning is a commonly used measure to alleviate the storage and computational burden of deep neural networks. However, the fundamental limit of network pruning is still lacking. To close the gap, in this work we'll take a first-principles approach, i.e. we'll directly impose the sparsity constraint on the loss function and leverage the framework of statistical dimension in convex geometry, thus enabling us to characterize the sharp phase transition point, which can be regarded as the fundamental limit of the pruning ratio. Through this limit, we're able to identify two key factors that determine the pruning ratio limit, namely, weight magnitude and network sharpness. Generally speaking, the flatter the loss landscape or the smaller the weight magnitude, the smaller pruning ratio. Moreover, we provide efficient countermeasures to address the challenges in the computation of the pruning limit, which mainly involves the accurate spectrum estimation of a large-scale and non-positive Hessian matrix. Moreover, through the lens of the pruning ratio threshold, we can also provide rigorous interpretations on several heuristics in existing pruning algorithms. Extensive experiments are performed which demonstrate that our theoretical pruning ratio threshold coincides very well with the experiments. All codes are available at: https://github.com/QiaozheZhang/Global-One-shot-Pruning

stat.ML↗

Multi-level Multiple Instance Learning with Transformer for Whole Slide Image Classification

Whole slide image (WSI) refers to a type of high-resolution scanned tissue image, which is extensively employed in computer-assisted diagnosis (CAD). The extremely high resolution and limited availability of region-level annotations make employing deep learning methods for WSI-based digital diagnosis challenging. Recently integrating multiple instance learning (MIL) and Transformer for WSI analysis shows very promising results. However, designing effective Transformers for this weakly-supervised high-resolution image analysis is an underexplored yet important problem. In this paper, we propose a Multi-level MIL (MMIL) scheme by introducing a hierarchical structure to MIL, which enables efficient handling of MIL tasks involving a large number of instances. Based on MMIL, we instantiated MMIL-Transformer, an efficient Transformer model with windowed exact self-attention for large-scale MIL tasks. To validate its effectiveness, we conducted a set of experiments on WSI classification tasks, where MMIL-Transformer demonstrate superior performance compared to existing state-of-the-art methods, i.e., 96.80% test AUC and 97.67% test accuracy on the CAMELYON16 dataset, 99.04% test AUC and 94.37% test accuracy on the TCGA-NSCLC dataset, respectively. All code and pre-trained models are available at: https://github.com/hustvl/MMIL-Transformer

cs.CV↗