几何深度学习 几何深度学习(geometric deep learning)是一个统一框架,它揭示出 CNN、Transformer 和 GNN 其实是同一条原理的不同实例:利用对称性(symmetry)。本文件涵盖对称群、群作用、不变性(invariance)、等变性(equivariance)、五大几何域,以及尺度分离。 在这本书里,我们已经学了许多架构:处理图像的 CNN(第 8 章)、处理语言的 Transformer(第 7 章)、以及处理序列决策的强化学习策略(第 6 章)。它们看起来像是为完全不同的问题设计的完全不同的模型。但这里有一个更深层的规律。 几何深度学习揭示出,所有这些架构都是同一个想法的实例:构建能够尊重数据对称性的网络。CNN 利用的是图像中的平移对称性;
几何深度学习(geometric deep learning)是一个统一框架,它揭示出 CNN、Transformer 和 GNN 其实是同一条原理的不同实例:利用对称性(symmetry)。本文件涵盖对称群、群作用、不变性(invariance)、等变性(equivariance)、五大几何域,以及尺度分离。
在这本书里,我们已经学了许多架构:处理图像的 CNN(第 8 章)、处理语言的 Transformer(第 7 章)、以及处理序列决策的强化学习策略(第 6 章)。它们看起来像是为完全不同的问题设计的完全不同的模型。但这里有一个更深层的规律。
几何深度学习揭示出,所有这些架构都是同一个想法的实例:构建能够尊重数据对称性的网络。CNN 利用的是图像中的平移对称性;Transformer 利用的是序列中的排列对称性(注意力不依赖于绝对位置);GNN 利用的是图中的排列对称性。一旦你看清这一点,原本五花八门的架构动物园就变成了一个单一、连贯的框架。
一个对象的**对称性(symmetry)**是让它在变换后保持不变的变换。一个正方形有 8 个对称性:4 个旋转(0°、90°、180°、270°)和 4 个反射。一个圆有无穷多个:绕圆心任意旋转。关键的洞见是——对称性告诉你"什么是无关紧要的",而知道什么是无关紧要的,对学习来说威力巨大。
用机器学习的话说:如果一个任务具有某种对称性,那么无论模型看到的是输入的哪一个"版本",它都应该给出相同的答案。一个猫检测器,不管猫在图像的左上角还是右下角,都应该能正常工作。这就是平移对称性。
对称性被形式化为群(group)。一个群 G 是一组变换,具有四条性质:
这些公理与向量空间(第 1 章)相同,只是把向量换成了变换。两者之间的联系很深:群作用在向量空间上,而这种作用正是神经网络必须尊重的东西。
深度学习中出现的关键群:
**群作用(group action)**描述一个群如何变换数据。如果 G 是一个群、X 是一个数据空间,那么作用 \rho: G \times X \to X 把每个群元素 g 和数据点 x 映射到一个变换后的点 \rho(g, x)。对于图像,平移群通过移动像素坐标来作用;对于图,对称群通过重新标记节点来作用。
给定一个对称群,一个函数可以与它有两种重要的关系:
如果输入被变换后输出不变,则函数 f 对群 G 是**不变(invariant)**的:
例如:平移一幅图像不会改变它的总亮度。图像分类应当是平移不变的:无论猫坐在哪里,类别"猫"都是一样的。
如果变换输入会让输出以相应的方式变换,则函数 f 对 G 是**等变(equivariant)**的:
这种区分很重要:中间层通常应当是等变的(为下游层保留结构信息),而最终输出应当是不变的(答案不应依赖于变换)。CNN 通过堆叠等变卷积层、再在末尾施加全局池化(这是不变的)来实现这一点。
把等变性内建到架构里,比从数据中学到它要高效得多。一个带权重共享的平移等变 CNN,所需的参数远少于一个必须独立学会"位置 (10,10) 处的猫"和"位置 (200,150) 处的猫"的全连接网络。对称性约束让假设空间以指数级缩小。
1. 网格(欧几里得数据):图像、音频频谱图、体数据。底层结构是具有平移对称性的规则网格。对应的群是平移群(可能再加上旋转和反射)。利用这种对称性的架构是 CNN:卷积正是那个对平移等变的运算。跨空间位置共享权重,就是平移等变性的具体体现。
2. 集合(无序集合):点云、粒子系统。对称性是排列不变性:元素的顺序无关紧要。对应的架构是 DeepSets(以及第 8 章的 PointNet):对每个元素施加一个共享函数,再用一个排列不变的运算(求和、求平均或取最大值)聚合。形式化地,f(\{x_1, \ldots, x_n\}) = \phi\left(\sum_i \psi(x_i)\right)。
3. 序列(有序数据):文本、时间序列。序列是一维的网格,但有一点不同:对称性更为微妙。绝对位置可能重要,也可能不重要。RNN 以自回归方式处理序列。带位置编码的 Transformer 可以注意到任意位置,而它的自注意力在加入位置编码之前对排列是等变的。这就是 Transformer 泛化得这么好的原因:它从排列等变出发,再加入刚刚好的位置结构。
4. 图(关系数据):社交网络、分子、知识图谱。对称性是节点的排列:重新标记节点不应改变图的性质。对应的架构是 GNN:在相连节点之间进行消息传递(message passing),使用与节点顺序无关的共享函数。这是本章其余部分的重点。
5. 流形与网格(manifolds and meshes):曲面、三维形状。对称性包括微分同胚(光滑变形)。对应的架构使用内蕴算子(例如 Laplace-Beltrami 算子),它们由曲面几何本身定义,与曲面如何嵌入空间无关。这部分联系到微分几何,与形状分析、球面上的气候建模以及蛋白质表面分析有关。
这个框架的力量在于统一。CNN 是作用在网格图上的 GNN;Transformer 是作用在完全图上的 GNN;DeepSets 是没有边的 GNN。把它们看作同一原理的不同实例,能指导新架构的设计:先识别数据的对称性,再构建一个尊重它的网络。
现实世界的数据在多个尺度上都有结构。一幅图像有细粒度纹理(像素级)、局部模式(边缘、角点)、物体部件(轮子、窗户)和全局结构(整个场景)。一个分子有原子级特征、官能团和整体分子形状。
尺度分离(scale separation)是这样一个原理:这些不同层次的细节可以被分层处理——先捕捉局部结构,再逐步聚合成更粗的表示。这就是粗化(coarsening)或池化(pooling)。
在 CNN 中,池化层(最大池化、平均池化)对空间分辨率做下采样,迫使更高层去捕捉更大尺度的模式。在感受野视角下(第 8 章),更深的层"看到"的图像范围更大。这就是尺度分离在起作用。
在图中,粗化意味着把一组组节点聚成"超节点(supernode)",产生一个保留本质结构的更小的图。这就是图池化,我们会在文件 03 里详细讨论。它与图像池化的类比是直接的:在保留重要特征的同时降低分辨率。
在序列中,分层处理(例如 句子 → 段落 → 文档)在不同时间或语义尺度上捕捉结构。Swin Transformer(第 8 章)通过它的移位窗口层级把这一思想用到图像上。
数学上,粗化定义了一个越来越抽象的表示层级:
在每一层,表示对该层的对称群是等变的。最终的全局表示是不变的,它捕捉输入的本质,而不对无关变换敏感。
正是这个层级,解释了为什么对于结构化数据,深的网络比浅的网络更好:每层增加一层抽象,许多等变层的组合从简单的局部特征构建出复杂的不变特征。
import jax import jax.numpy as jnp # 一维信号和一个简单的滤波器 signal = jnp.array([0, 0, 0, 1, 2, 3, 2, 1, 0, 0, 0], dtype=float) kernel = jnp.array([1, 0, -1], dtype=float) # 先卷积再平移 conv_result = jnp.convolve(signal, kernel, mode="same") shifted_signal = jnp.roll(signal, 3) conv_shifted = jnp.convolve(shifted_signal, kernel, mode="same") shifted_conv = jnp.roll(conv_result, 3) print(f"Conv then shift: {shifted_conv}") print(f"Shift then conv: {conv_shifted}") print(f"Equivariant: {jnp.allclose(shifted_conv, conv_shifted, atol=1e-5)}")
import jax import jax.numpy as jnp # 一个由 4 个向量组成的"集合"(顺序不应有影响) x = jnp.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0], [7.0, 8.0]]) # 简单的共享函数:逐元素平方 psi = lambda v: v ** 2 # 用求和聚合 def deepsets(points): return jnp.sum(jax.vmap(psi)(points), axis=0) # 原始顺序 result1 = deepsets(x) # 打乱顺序 perm = jnp.array([2, 0, 3, 1]) result2 = deepsets(x[perm]) print(f"Original order: {result1}") print(f"Permuted order: {result2}") print(f"Invariant: {jnp.allclose(result1, result2)}")
import jax.numpy as jnp def rot2d(theta): return jnp.array([[jnp.cos(theta), -jnp.sin(theta)], [jnp.sin(theta), jnp.cos(theta)]]) R1 = rot2d(jnp.pi / 6) R2 = rot2d(jnp.pi / 4) R3 = rot2d(jnp.pi / 3) # 封闭性:两个旋转的乘积仍是旋转 R12 = R1 @ R2 print(f"Closure (det=1, orthogonal): det={jnp.linalg.det(R12):.4f}, " f"R^T R = I: {jnp.allclose(R12.T @ R12, jnp.eye(2), atol=1e-5)}") # 结合律 print(f"Associative: {jnp.allclose((R1 @ R2) @ R3, R1 @ (R2 @ R3), atol=1e-5)}") # 单位元 I = rot2d(0.0) print(f"Identity: {jnp.allclose(R1 @ I, R1, atol=1e-5)}") # 逆元 R1_inv = rot2d(-jnp.pi / 6) print(f"Inverse: {jnp.allclose(R1 @ R1_inv, jnp.eye(2), atol=1e-5)}")