几何深度学习


文档摘要

几何深度学习 几何深度学习(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 是一组变换,具有四条性质:

    • 封闭性(closure):把两个变换组合起来,得到的仍然是集合里的某个变换。先旋转 90° 再旋转 90° 得到 180°,它仍在集合里。
    • 结合律(associativity)(g_1 \circ g_2) \circ g_3 = g_1 \circ (g_2 \circ g_3)。分组顺序无关紧要(回想第 2 章里矩阵乘法的结合律)。
    • 单位元(identity):存在一个"什么都不做"的变换 e,使得 e \circ g = g \circ e = g
    • 逆元(inverse):每个变换都有一个撤销它的变换:g \circ g^{-1} = e
  • 这些公理与向量空间(第 1 章)相同,只是把向量换成了变换。两者之间的联系很深:群作用在向量空间上,而这种作用正是神经网络必须尊重的东西。

  • 深度学习中出现的关键群:

    • 平移群(translation group) (\mathbb{R}^n, +):平移一幅图像或信号。这是 CNN 所利用的对称性。
    • 对称群(symmetric group) S_nn 个元素的所有排列。这是 GNN 和 Transformer 所利用的对称性(重排节点或 token 不应改变结果)。
    • 旋转群(rotation group) SO(n)n 维空间中的所有旋转。SO(2) 是平面上的旋转,SO(3) 是三维空间中的旋转(对分子和三维视觉任务至关重要)。
    • 欧几里得群(Euclidean group) E(n):所有旋转、反射和平移。物理空间的对称性。
    • 特殊欧几里得群(special Euclidean group) SE(n):旋转和平移(不含反射)。刚体运动的对称性。
  • **群作用(group action)**描述一个群如何变换数据。如果 G 是一个群、X 是一个数据空间,那么作用 \rho: G \times X \to X 把每个群元素 g 和数据点 x 映射到一个变换后的点 \rho(g, x)。对于图像,平移群通过移动像素坐标来作用;对于图,对称群通过重新标记节点来作用。

不变性与等变性

  • 给定一个对称群,一个函数可以与它有两种重要的关系:

  • 如果输入被变换后输出不变,则函数 f 对群 G 是**不变(invariant)**的:

f(\rho(g, x)) = f(x) \quad \text{for all } g \in G
  • 例如:平移一幅图像不会改变它的总亮度。图像分类应当是平移不变的:无论猫坐在哪里,类别"猫"都是一样的。

  • 如果变换输入会让输出以相应的方式变换,则函数 fG 是**等变(equivariant)**的:

f(\rho_{\text{in}}(g, x)) = \rho_{\text{out}}(g, f(x)) \quad \text{for all } g \in G
  • 例如:如果你把图像向右平移 5 个像素,CNN 中的特征图也会向右平移 5 个像素。卷积运算是平移等变的:它保持空间关系。目标检测应当是等变的:如果猫移动了,边界框也应该跟着移动。

不变性:无论怎样变换,输出都保持不变。等变性:输出随之作相应的变换

  • 这种区分很重要:中间层通常应当是等变的(为下游层保留结构信息),而最终输出应当是不变的(答案不应依赖于变换)。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 章)通过它的移位窗口层级把这一思想用到图像上。

  • 数学上,粗化定义了一个越来越抽象的表示层级

x \xrightarrow{\text{local features}} h^{(1)} \xrightarrow{\text{coarsen}} h^{(2)} \xrightarrow{\text{coarsen}} \cdots \xrightarrow{\text{global}} y
  • 在每一层,表示对该层的对称群是等变的。最终的全局表示是不变的,它捕捉输入的本质,而不对无关变换敏感。

  • 正是这个层级,解释了为什么对于结构化数据,深的网络比浅的网络更好:每层增加一层抽象,许多等变层的组合从简单的局部特征构建出复杂的不变特征。

编程练习(使用 CoLab 或 notebook)

  1. 验证卷积的平移等变性。对一幅图像做卷积,然后平移图像再做卷积,检查两次输出是不是彼此平移后的版本。
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)}")
  1. 验证 DeepSets 风格聚合的排列不变性。对集合里的每个元素施加一个共享函数,把结果求和,并检查无论元素顺序如何,输出都相同。
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)}")
  1. 探索群结构。通过检验封闭性、结合律、单位元和逆元,验证二维旋转矩阵构成一个群。
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)}")

作者与出处
原作者: HenryNdubuaku
来源:HenryNdubuaku
许可证:Apache-2.0
整理: 灏天文库整理
由灏天文库结构化整理,提供目录导航、全文检索与在线阅读,便于系统化学习
发布者: 作者: HenryNdubuaku 转发
评论区 (0)
U