LeCun连续转发:攻克「表征坍塌」核心难题新智元

7/28/2026

JEPA世界模型的底层是Yann LeCun自2017年起持续倡导的自监督学习(Self-Supervised Learning, SSL)。

SSL 无需人工标注即可从海量数据中学习通用表征,但普遍面临一个核心难题——表征坍塌(representation collapse):模型倾向于把不同输入映射到相同或极少数几个向量上,看似完成了训练,实则未学到有判别力的表征。

为抑制坍塌,主流方法大多依赖一系列启发式技巧(EMA、教师-学生网络、停止梯度、冻结层等)。这些技巧使训练变得脆弱、难以调参,也削弱了方法的可解释性与可扩展性。

另一条路线是通过正则项直接约束表征分布。

LeCun团队提出的VICReg将学习目标拆为方差、不变性、协方差三项,用协方差约束各维度之间的相关性;但协方差仅刻画二阶统计量,无法区分「均值、方差相同,而分布形状迥异」的两种表征。

其后提出的SIGReg基于Cramér–Wold定理,用sketching技术将整个嵌入分布对齐到标准高斯,从而约束完整的分布形状。

然而SIGReg仍存在两个关键缺陷:

坍塌时梯度消失:当表征开始坍塌时,SIGReg的梯度随之衰减——坍塌越严重、修正信号越弱,模型难以自行恢复;

尺度与形状耦合:未将「幅度大小(尺度)」与「分布形态(形状)」两个独立属性分离,二者在优化中相互干扰,导致在长尾、低质量、低秩数据上适配性较差。

也就是说,在模型最需要梯度信号来逃离坍塌状态时,SIGReg的梯度恰恰趋近于消失。

这正是VISReg要解决的核心问题。

近日,自监督学习新工作VISReg(Variance-Invariance-Sketching Regularization)获图灵奖得主Yann LeCun连续转发并给予高度认可——他在转发时评价道「VICReg begat SIGReg which begat VISReg」(VICReg孕育了SIGReg,SIGReg又孕育了VISReg),一句话点明了这条正则化路线的技术传承。

能获得LeCun如此认可,VISReg究竟强在哪里?

答案在于:它精准命中了LeCun长期押注的JEPA世界模型的核心难题——表征坍塌(representation collapse)。

论文链接:https://arxiv.org/abs/2606.02572

代码 / 预训练权重:https://github.com/HaiyuWu/visreg

项目主页:https://haiyuwu.github.io/visreg/

VISReg将防止坍塌的正则项解耦为「尺度」与「形状」两个独立目标,在不依赖任何启发式训练技巧、也不依赖海量数据的前提下,于15个数据集上综合表现超过7种主流自监督学习方法;其中仅用约1/10的训练数据,即在分布外(OOD)基准上追平DINOv2。

图 2:不同正则方法在表征坍缩各阶段的梯度幅值‖∇ℒ‖模拟。VISReg在坍缩状态下仍能保持强梯度,而SIGReg的梯度几近消失。

VISReg对VICReg与SIGReg取长补短:保留VICReg的方差项来控制尺度,同时用基于切片Wasserstein距离(Sliced Wasserstein Distance, SWD)的 sketching 目标替代协方差项来控制形状,并通过停止梯度将二者彻底解耦。整个正则目标由三部分组成。

尺度正则(Scale Regularization)

第一部分约束每一维的方差,防止幅值坍缩:

其关键性质在于:当模型坍缩时,该项的梯度趋近于一个常数,从而保证模型能够稳定地恢复——这恰好弥补了SIGReg梯度消失的缺陷。

形状正则(Shape Regularization)

第二部分先归一化以消除尺度影响,再单独约束形状。关键的一步是带「停止梯度」(stop-gradient, sg)的归一化:

这里对标准差 σσ 施加停止梯度,使得形状损失的优化不会反过来改变尺度——这正是「尺度」与「形状」两个目标真正解耦、互不干扰的机理所在。

仅需约15行PyTorch代码

该正则目标在实现上非常轻量,核心逻辑只需约15行:

def visreg(z, K=64):

# 1. 中心化损失

mu = z.mean(dim=0)

L_center = mu.pow(2).mean()

# 2. 尺度损失

z_cent = z - mu

std = z_cent.std(dim=0, unbiased=False)

L_scale = (1.0 - std).pow(2).mean()

# 3. 形状损失:切片 Wasserstein 距离

z_norm = z_cent / (std.detach())

W = torch.randn(D, K)

W /= W.norm(p=2, dim=0)

p_sorted = torch.sort(z_norm @ W, dim=0).values

u = torch.arange(1, N+1) / (N+1)

target = Normal(0, 1).icdf(u)

L_shape = (p_sorted - target).pow(2).mean()

return L_scale + L_shape + L_center

计算复杂度与扩展性

Scroll for more