在这里插入图片描述

PyTorch Scala 高校计算机硕士研一课程

章节 10: 进阶神经网络结构 20

虽然像卷积神经网络(CNN)和循环神经网络(RNN)这样的基本网络设计能有效处理许多任务,但某些问题场景需要更专业的结构。本章将着重介绍如何使用PyTorch实现多种进阶神经网络模型。

您将学习重要的现代结构及其实现细节。我们将逐个组件地介绍Transformer模型的构建,包括注意力机制。我们还将使用图神经网络(GNN)处理图结构数据,并运用PyTorch Geometric等库。此外,本章会介绍用于生成任务的归一化流(Normalizing Flows)、用于连续深度建模的神经常微分方程(Neural ODEs),以及针对少样本场景的元学习方法。侧重于理解这些构成要素,并在代码中构建这些复杂的模型。

从组件构建Transformer模型

Transformer模型彻底改变了序列建模,在自然语言处理、计算机视觉及其他方面取得了最先进的成果。其效果源于自注意力机制,该机制使模型在生成输出时,能够衡量不同输入元素的重要性,而无论它们之间的距离如何。使用PyTorch从基本组件构建Transformer模型,有助于充分了解其工作原理。假定熟悉基本的PyTorch模块(nn.Modulenn.Linear等)和深度学习原理。

Transformer架构:概览

最初的Transformer模型由Vaswani等人于2017年在《Attention Is All You Need》中提出,采用编码器-解码器结构,适用于机器翻译等序列到序列的任务。

编码器解码器输入嵌入 +位置编码编码器层 1(多头注意力, 加和归一化, 前馈网络, 加和归一化)输入序列编码器层 N(…)编码器输出(记忆)输出嵌入 +位置编码解码器层 1(带掩码MHA, 加和归一化, 编码器-解码器注意力, 加和归一化, 前馈网络, 加和归一化)记忆 (K, V)解码器层 N(…)目标序列 (右移)线性层Softmax输出概率输入输出

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

编码器-解码器Transformer架构的高级视图。编码器处理输入序列,而解码器则使用编码器的输出和先前生成的目标序列来生成下一个元素。

许多成功模型仅使用编码器堆栈(例如,BERT用于语言理解)或解码器堆栈(例如,GPT用于语言生成)。我们将侧重于实现所有变体共有的核心构成模块。

输入表示:嵌入和位置编码

神经网络处理数字,而非原始文本。因此,第一步是将输入词元(单词、子词或字符)转换为数值向量。

词元嵌入

这通常通过使用嵌入层完成,该嵌入层本质上是一个查找表。词汇表中的每个唯一词元都被分配一个固定大小的稠密向量,dmodeldmode**l。在PyTorch中,这可使用torch.nn.Embedding直接实现。

import torch
import torch.nn as nn
import math

// 示例参数
val vocab_size = 10000 // 词汇表大小
val d_model = 512      // 嵌入维度

val embedding = nn.Embedding(vocab_size, d_model)

// 示例用法:2个序列的批次,长度为10
val input_tokens = torch.randint(0, vocab_size, (2, 10)) // (批次大小, 序列长度)
val input_embeddings = embedding(input_tokens)          // (批次大小, 序列长度, d_model)

println("输入形状:", input_tokens.shape)
println("嵌入形状:", input_embeddings.shape)
位置编码

接下来将要考察的自注意力机制同时处理序列元素。它本质上不考虑词元的顺序或位置。如果没有位置信息,注意力机制在嵌入后会将“the cat sat on the mat”和“the mat sat on the cat”视为相同。

为了解决这个问题,Transformer模型将每个词元的位置信息注入其嵌入中。原始论文提出使用固定的正弦函数:

PE(位置,2i)=sin⁡(位置/100002i/dmodel)P**E(位置,2i)=sin(位置/100002i/dmode**l)PE(位置,2i+1)=cos⁡(位置/100002i/dmodel)P**E(位置,2i+1)=cos(位置/100002i/dmode**l)

这里,posp**os 是词元在序列中的位置,ii 是嵌入向量内的维度索引(0≤2i<dmodel0≤2i<dmode**l)。位置编码的每个维度对应于不同频率的正弦曲线。这种选择使模型能够更容易地获取相对位置信息,因为 PEpos+kPEp**os+k 可以表示为 PEposPEp**os 的线性函数。

或者,可以使用可学习的位置嵌入(类似于词元嵌入,但查找的是位置索引)。我们在这里将实现正弦版本。

class PositionalEncoding extends nn.Module:
    def __init__(self, d_model: Int, dropout: Float = 0.1, max_len: Int = 5000):
        super().__init__()
        val dropout = nn.Dropout(p=dropout)

        // 创建位置索引 (0 到 max_len - 1)
        val position = torch.arange(max_len).unsqueeze(1) // 形状: (max_len, 1)

        // 计算正弦和余弦参数的除数项
        val div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
        // 形状: (d_model / 2)

        // 初始化位置编码矩阵
        val pe = torch.zeros(max_len, d_model) // 形状: (max_len, d_model)

        // 对偶数索引应用sin,对奇数索引应用cos
        pe(:, 0::2) = torch.sin(position * div_term)
        pe(:, 1::2) = torch.cos(position * div_term)

        // 添加批次维度并注册为缓冲区(非模型参数)
        val pe = pe.unsqueeze(0) // 形状: (1, max_len, d_model)
        register_buffer("pe", pe)

    def forward(x: torch.Tensor) -> torch.Tensor:
        """
        参数:
            x: 张量, 形状 [批次大小, 序列长度, d_model]
        返回:
            张量, 形状 [批次大小, 序列长度, d_model]
        """
        // 将位置编码添加到输入嵌入
        // self.pe 形状为 (1, max_len, d_model)。我们取 x 的序列长度范围内的切片。
        // x 的形状为 (批次大小, 序列长度, d_model)
        x = x + pe(:, :x.size(1), :)
        return dropout(x)

// 示例用法
val pos_encoder = PositionalEncoding(d_model, dropout=0.1)
val final_input = pos_encoder(input_embeddings * math.sqrt(d_model)) // 在添加位置编码前对嵌入进行缩放

println("位置编码后的形状:", final_input.shape)
// 注意:原始论文在添加位置编码前,会按 sqrt(d_model) 缩放嵌入。

第一个Transformer层的最终输入是词元嵌入(可选缩放)和位置编码之和。

核心组件:多头自注意力

注意力机制使模型在处理特定元素时能够关注输入序列的相关部分。自注意力将单个序列的不同位置关联起来,以计算该序列的表示。

缩放点积注意力

基本构成块是缩放点积注意力。对于序列中的每个元素,我们计算三个向量:查询(QQ)、键(KK)和值(VV)。这些通常通过将输入嵌入(加上位置编码)乘以可学习的权重矩阵 WQW**Q、WKW**K 和 WVW**V 来获得。

想象您正在处理句子“making Transformer models more interpretable”中的单词“making”。

  • “making”的查询向量提问:“在当前语境中,句子的哪些部分与理解我的意思相关?”
  • 句子中的每个词都生成一个向量,本质上是说:“这是我所代表的信息。”
  • 每个词还会生成一个向量:“如果您认为我相关,这是我将提供的信息。”

一个词(例如,“making”)的查询与另一个词(例如,“Transformer”)的键之间的注意力分数是使用点积计算的。这些分数确定“making”应该对“Transformer”投入多少注意力。

分数通过向量维度(dkd**k)的平方根进行缩放,以防止点积变得过大,这可能使softmax函数饱和并导致梯度消失。然后,softmax函数将这些分数转换为总和为1的概率(权重)。

最后,查询词(“making”)的输出是序列中所有值向量的加权和,其中权重是计算出的概率。

公式如下:

注意力(Q,K,V)=softmax(QKTdk)V注意力(Q,K,V)=softmax(dkQKT)V

其中 QQ、KK 和 VV 是包含序列中所有词元的查询、键和值的矩阵。

多头注意力

多头注意力不同于使用dmodeldmode**l维度的Q、K、V向量进行单次注意力计算,它将输入Q、K、V向量通过hh次(其中hh是头数)不同的可学习线性投影(权重矩阵)投影到维度dqd**q、dkd**k、dvd**v(通常dq=dk=dv=dmodel/hd**q=d**k=d**v=dmode**l/h)。

缩放点积注意力随后独立应用于这些投影版本(每个“头”)。这使模型能够同时关注来自不同表示子空间和不同位置的信息。这就像同时针对输入提出多个不同的问题(查询)一样。

所有hh个头的输出被连接起来,然后通过一个最终线性层(WOW**O),以生成多头注意力层的最终输出。

def scaled_dot_product_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, mask: torch.Tensor = None):
    """计算缩放点积注意力"""
    val d_k = q.size(-1) // 获取最后一个维度(K的嵌入维度)
    // Q与K转置的矩阵乘法: (..., 查询序列长度, d_k) x (..., 键序列长度, d_k) -> (..., 查询序列长度, 键序列长度)
    var scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_k) // (..., 查询序列长度, 键序列长度)

    // 应用掩码(如果提供),将掩码位置设置为一个非常小的数字 (-1e9)
    if mask is not None then 
        scores = scores.masked_fill(mask == 0, -1e9)

    // 应用softmax以获取注意力权重
    val attn_weights = torch.softmax(scores, dim=-1) // (..., 查询序列长度, 键序列长度)

    // 权重与V的矩阵乘法: (..., 查询序列长度, 键序列长度) x (..., 值序列长度, d_v) -> (..., 查询序列长度, d_v)
    // 注意: 键序列长度 == 值序列长度
    val output = torch.matmul(attn_weights, v) // (..., 查询序列长度, d_v)
    return output, attn_weights

class MultiHeadAttention extends nn.Module:
    def __init__( d_model: Int, num_heads: Int):
        super().__init__()
        assert d_model % num_heads == 0, "d_model must be divisible by num_heads"

        val d_model = d_model
        val num_heads = num_heads
        val d_k = d_model // num_heads // 每个头的键/查询维度
        val d_v = d_model // num_heads // 每个头的值维度

        // Q、K、V投影的线性层(应用于所有头)
        val W_q = nn.Linear(d_model, d_model)
        val W_k = nn.Linear(d_model, d_model)
        val W_v = nn.Linear(d_model, d_model)

        // 连接后的最终线性层
        val W_o = nn.Linear(d_model, d_model)

    def split_heads(x: torch.Tensor) : torch.Tensor =
        // 输入 x: (批次大小, 序列长度, d_model)
        val batch_size = x.size(0)
        val seq_len = x.size(1)
        // 重塑为 (批次大小, 序列长度, 头数, d_k)
        x = x.view(batch_size, seq_len, num_heads, d_k)
        // 转置为 (批次大小, 头数, 序列长度, d_k) 以进行注意力计算
        return x.transpose(1, 2)

    def combine_heads(x: torch.Tensor) : torch.Tensor =
        // 输入 x: (批次大小, 头数, 序列长度, d_k)
        val batch_size = x.size(0)
        val seq_len = x.size(2)
        // 转置回 (批次大小, 序列长度, 头数, d_k)
        x = x.transpose(1, 2).contiguous() // 确保转置后内存连续
        // 重塑为 (批次大小, 序列长度, d_model)
        return x.view(batch_size, seq_len, d_model)

    def forward(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, mask: torch.Tensor = None) :torch.Tensor =
        // q, k, v: (批次大小, 序列长度, d_model)
        // 掩码: (批次大小, 1, 查询序列长度, 键序列长度) 或类似的可广播形状

        // 1. 应用线性投影
        q = W_q(q) // (批次大小, 查询序列长度, d_model)
        k = W_k(k) // (批次大小, 键序列长度, d_model)
        v = W_v(v) // (批次大小, 值序列长度, d_model) // 注意: 键序列长度 == 值序列长度

        // 2. 分割成多个头
        q = split_heads(q) // (批次大小, 头数, 查询序列长度, d_k)
        k = split_heads(k) // (批次大小, 头数, 键序列长度, d_k)
        v = split_heads(v) // (批次大小, 头数, 值序列长度, d_k)

        // 3. 应用缩放点积注意力
        // 输出: (批次大小, 头数, 查询序列长度, d_k)
        // 注意力权重: (批次大小, 头数, 查询序列长度, 键序列长度)
        val attention_output, attn_weights = scaled_dot_product_attention(q, k, v, mask)

        // 4. 合并头
        val output = combine_heads(attention_output) // (批次大小, 查询序列长度, d_model)

        // 5. 最终线性层
        val output = W_o(output) // (批次大小, 查询序列长度, d_model)

        return output # 通常我们只需要输出,不需要权重,用于下一层

// 示例用法
// 创建一个多头注意力实例
val mha = MultiHeadAttention(d_model=512, num_heads=8)
// 在自注意力中,Q、K和V最初通常是同一个张量
val query = key = value = final_input // 形状: (批次大小, 序列长度, d_model)
val attention_result = mha(query, key, value, mask=None) // 掩码对填充/解码很重要

println(s"多头注意力输出形状: ${attention_result.shape}")
掩码

掩码在两种情况中不可或缺:

  1. 填充掩码: 为了防止模型关注批次中不同长度序列的填充词元。此掩码在编码器和解码器自注意力中都应用。它通常具有类似 (batch_size, 1, 1, seq_len_k) 的形状,填充位置包含0,否则为1。
  2. 前瞻掩码: 在解码器的注意力层中,这阻止一个位置关注后续位置。这确保位置 ii 的预测只能依赖于小于 ii 的位置的已知输出。这通过掩盖(在softmax前设置为 -1e9)注意力分数矩阵的上三角部分来实现。

加和归一化层

Transformer中的每个子层(如多头注意力或前馈网络)都具有残差连接,然后是层归一化。

残差连接(加和)

子层的输出被添加到子层的输入中:output = x + Sublayer(x)。这种技术借鉴自残差网络(ResNets),有助于缓解深度网络中的梯度消失问题,使梯度在反向传播过程中更直接地流过网络。它也使得训练更深层的模型成为可能。

层归一化(归一化)

层归一化(nn.LayerNorm)独立地对批次中每个单独数据样本(词元)的特征维度(即 dmodeldmode**l 维度)上的激活进行归一化。这与批归一化形成对比,后者对批次维度进行归一化。层归一化在自然语言处理和Transformer模型中通常更受青睐,原因如下:

  • 它不依赖于批次大小,使其在小批次或可变序列长度下也能保持稳定。
  • 它为每个序列元素独立地提供一致的归一化。

加和归一化步骤通常实现为:output = LayerNorm(x + Dropout(Sublayer(x)))。Dropout通常在残差加法和归一化之前应用于子层的输出。

class AddNorm extends nn.Module:
    def __init__(normalized_shape: Int, dropout: Float):
        super().__init__()
        val layer_norm = nn.LayerNorm(normalized_shape)
        val dropout = nn.Dropout(dropout)

    def forward(x: torch.Tensor, sublayer_output: torch.Tensor) :torch.Tensor=
        // 应用残差连接和Dropout,然后是层归一化
        return layer_norm(x + dropout(sublayer_output))

// 示例:在多头注意力后应用加和归一化
val dropout_rate = 0.1
val add_norm1 = AddNorm(d_model, dropout_rate)
// 'final_input' 是MHA层的输入
val normed_attention_output = add_norm1(final_input, attention_result)

println(s"加和归一化输出形状: ${normed_attention_output.shape}")

位置维前馈网络(FFN)

在注意力子层(及其加和归一化)之后,每个位置的表示都通过一个相同且独立的前馈网络(FFN)。该网络通常由两个线性变换和一个非线性激活函数组成,通常是ReLU或GeLU(高斯误差线性单元)。

FFN(x)=max⁡(0,xW1+b1)W2+b2(使用ReLU)FFN(x)=max(0,x**W1+b1)W2+b2(使用ReLU)

维度通常在第一个线性层中增加(例如,到 dff=4×dmodeld**ff=4×dmode**l),然后在第二个层中减少回 dmodeldmode**l。这种FFN使模型能够独立地处理通过注意力在每个位置获取的信息,从而增加了非线性建模能力。

class PositionWiseFeedForward extends nn.Module:
    def __init__(d_model: Int, d_ff: Int, dropout: Float = 0.1):
        super().__init__()
        val linear1 = nn.Linear(d_model, d_ff)
        val activation = nn.ReLU() # 或 nn.GELU()
        val dropout = nn.Dropout(dropout)
        val linear2 = nn.Linear(d_ff, d_model)

    def forward(x: torch.Tensor): torch.Tensor =
        // x: (批次大小, 序列长度, d_model)
        x = linear1(x)     // (批次大小, 序列长度, d_ff)
        x = activation(x)
        x = dropout(x)
        x = linear2(x)     // (批次大小, 序列长度, d_model)
        return x

// 示例用法
val d_ff = d_model * 4 // 常见做法
val ffn = PositionWiseFeedForward(d_model, d_ff, dropout_rate)
val ffn_output = ffn(normed_attention_output)

// 应用第二个加和归一化层
val add_norm2 = AddNorm(d_model, dropout_rate)
// 'normed_attention_output' 是FFN的输入
val encoder_layer_output = add_norm2(normed_attention_output, ffn_output)

println(s"FFN输出形状: ${ffn_output.shape}")
println(s"编码器层输出形状: ${encoder_layer_output.shape}")

层的组合

有了这些组件,我们可以定义一个完整的编码器层和解码器层。

编码器层

一个编码器层包含:

  1. 多头自注意力。
  2. 加和归一化。
  3. 位置维前馈网络。
  4. 加和归一化。
class EncoderLayer extends nn.Module:
    def __init__(d_model: Int, num_heads: Int, d_ff: Int, dropout: Float):
        super().__init__()
        val self_attn = MultiHeadAttention(d_model, num_heads)
        val add_norm1 = AddNorm(d_model, dropout)
        val ffn = PositionWiseFeedForward(d_model, d_ff, dropout)
        val add_norm2 = AddNorm(d_model, dropout)

    def forward(x: torch.Tensor, mask: torch.Tensor): torch.Tensor =
        // 自注意力子层
        val attn_output = self_attn(q=x, k=x, v=x, mask=mask)
        x = add_norm1(x, attn_output) // 残差连接 + 归一化

        // 前馈子层
        val ffn_output = ffn(x)
        x = add_norm2(x, ffn_output) // 残差连接 + 归一化
        return x
解码器层

一个解码器层稍微复杂一些,包含两种注意力机制:

  1. 带掩码的多头自注意力: 使用前瞻掩码,关注解码器的输入序列(目前已生成的目标序列)。
  2. 加和归一化。
  3. 多头编码器-解码器注意力: 关注编码器堆栈的输出(通常称为记忆)。查询(QQ)来自上一个解码器子层的输出,而键(KK)和值(VV)来自编码器输出。这使解码器在生成输出序列时能够考虑输入序列的相关部分。这里可能需要来自编码器输入的填充掩码。
  4. 加和归一化。
  5. 位置维前馈网络。
  6. 加和归一化。
class DecoderLayer extends nn.Module:
    def __init__(d_model: Int, num_heads: Int, d_ff: Int, dropout: Float):
        super().__init__()
        val masked_self_attn = MultiHeadAttention(d_model, num_heads)
        val add_norm1 = AddNorm(d_model, dropout)
        val encoder_decoder_attn = MultiHeadAttention(d_model, num_heads)
        val add_norm2 = AddNorm(d_model, dropout)
        val ffn = PositionWiseFeedForward(d_model, d_ff, dropout)
        val add_norm3 = AddNorm(d_model, dropout)

    def forward(x: torch.Tensor, encoder_output: torch.Tensor,
                look_ahead_mask: torch.Tensor, padding_mask: torch.Tensor) : torch.Tensor=
        // 1. 带掩码的自注意力子层
        // Q=x, K=x, V=x; 使用前瞻掩码
        val self_attn_output = masked_self_attn(q=x, k=x, v=x, mask=look_ahead_mask)
        x = add_norm1(x, self_attn_output)

        // 2. 编码器-解码器注意力子层
        // Q=x (来自上一层), K=编码器输出, V=编码器输出
        // 使用与编码器输出相关的填充掩码
        val enc_dec_attn_output = encoder_decoder_attn(q=x, k=encoder_output, v=encoder_output, mask=padding_mask)
        x = add_norm2(x, enc_dec_attn_output)

        // 3. 前馈子层
        val ffn_output = ffn(x)
        x = add_norm3(x, ffn_output)

        return x

构建完整的Transformer

最终的Transformer模型堆叠多个编码器层(例如 N=6N=6)形成编码器,并堆叠多个解码器层(例如 N=6N=6)形成解码器。nn.ModuleList 对此很方便。

class Transformer extends nn.Module:
    def __init__(num_encoder_layers: Int, num_decoder_layers: Int,
                 d_model: Int, num_heads: Int, d_ff: Int,
                 input_vocab_size: Int, target_vocab_size: Int,
                 max_seq_len: Int, dropout: Float = 0.1): 
        super().__init__()

        val encoder_embedding = nn.Embedding(input_vocab_size, d_model)
        val decoder_embedding = nn.Embedding(target_vocab_size, d_model)
        val positional_encoding = PositionalEncoding(d_model, dropout, max_seq_len)

        val encoder_layers = nn.ModuleList([
            EncoderLayer(d_model, num_heads, d_ff, dropout)
            for _ in range(num_encoder_layers)
        ])
        self.decoder_layers = nn.ModuleList([
            DecoderLayer(d_model, num_heads, d_ff, dropout)
            for _ in range(num_decoder_layers)
        ])

        val final_linear = nn.Linear(d_model, target_vocab_size)
        val d_model = d_model
        val dropout = nn.Dropout(dropout)

    def create_padding_mask(seq: torch.Tensor, pad_token_idx: Int = 0) :torch.Tensor=
        // 序列形状: (批次大小, 序列长度)
        // 输出掩码形状: (批次大小, 1, 1, 序列长度)
        val mask = (seq != pad_token_idx).unsqueeze(1).unsqueeze(2)
        return mask

    def create_look_ahead_mask( size: int): torch.Tensor =
        // 创建一个上三角矩阵用于掩盖未来词元
        // 输出掩码形状: (1, 1, 大小, 大小)
        val mask = torch.triu(torch.ones(size, size), diagonal=1).bool()
        // 我们希望在掩盖处为0,所以我们进行反转(如果注意力中使用0进行掩盖)
        // 或者如果注意力函数期望在掩盖处为True,则按原样返回
        // 假设 scaled_dot_product_attention 使用 masked_fill(mask == 0, -1e9) 或 masked_fill(mask == True, -1e9),请相应调整。
        // 让我们假设后者 (True表示掩码)
        return ~mask.unsqueeze(0).unsqueeze(0) # 在掩盖处设置为False

    def encode(src: torch.Tensor, src_mask: torch.Tensor): torch.Tensor=
        //源: (批次大小, 源序列长度)
        // 源掩码: (批次大小, 1, 1, 源序列长度)
        val src_emb = encoder_embedding(src) * math.sqrt(d_model)
        val src_pos_emb = positional_encoding(src_emb)
        val enc_output = dropout(src_pos_emb)

        for layer in encoder_layers:
            enc_output = layer(enc_output, src_mask)
        return enc_output # (批次大小, 源序列长度, d_model)

    def decode( tgt: torch.Tensor, encoder_output: torch.Tensor,
               look_ahead_mask: torch.Tensor, padding_mask: torch.Tensor) -> torch.Tensor:
        // 目标: (批次大小, 目标序列长度)
        // 编码器输出: (批次大小, 源序列长度, d_model)
        // 前瞻掩码: (批次大小, 1, 目标序列长度, 目标序列长度)
        // 填充掩码: (批次大小, 1, 1, 源序列长度) # 在编码器-解码器注意力中使用
        val tgt_emb = decoder_embedding(tgt) * math.sqrt(d_model)
        val tgt_pos_emb = positional_encoding(tgt_emb)
        val dec_output = dropout(tgt_pos_emb)

        tgt_emb = decoder_embedding(tgt) * math.sqrt(d_model)
        tgt_pos_emb = positional_encoding(tgt_emb)
        dec_output = dropout(dec_output)

        for layer <- decoder_layers:
            dec_output = layer(dec_output, encoder_output, look_ahead_mask, padding_mask)

        return dec_output // (批次大小, 目标序列长度, d_model)

    def forward(src: torch.Tensor, tgt: torch.Tensor): torch.Tensor =
        // 源: (批次大小, 源序列长度)
        // 目标: (批次大小, 目标序列长度) 通常为训练目的而右移
        // 输出: (批次大小, 目标序列长度, 目标词汇表大小)

        val src_padding_mask = create_padding_mask(src)
        val tgt_padding_mask = create_padding_mask(tgt) // 如果目标也有填充,也需要
        val look_ahead_mask = create_look_ahead_mask(tgt.size(1)).to(tgt.device)

        // 将前瞻掩码和目标填充掩码结合用于解码器自注意力
        // 确保两个掩码都可广播: (批次大小, 1, 目标序列长度, 目标序列长度)
        val combined_look_ahead_mask = torch.logical_and(tgt_padding_mask.transpose(-2, -1), look_ahead_mask)

        val encoder_output = encode(src, src_padding_mask)
        val decoder_output = decode(tgt, encoder_output, combined_look_ahead_mask, src_padding_mask)

        // 最终线性投影
        val output = final_linear(decoder_output) // (批次大小, 目标序列长度, 目标词汇表大小)
        return output // 通常在推理/损失计算期间,模型外部接着Softmax

// 示例实例化(参数仅作说明)
val transformer_model = Transformer(
    num_encoder_layers=6, num_decoder_layers=6,
    d_model=512, num_heads=8, d_ff=2048,
    input_vocab_size=10000, target_vocab_size=12000,
    max_seq_len=500, dropout=0.1
)

// 用于形状检查的虚拟输入(假设批次大小为2)
val src_dummy = torch.randint(1, 10000, (2, 100)) // (批次, 源长度)
val tgt_dummy = torch.randint(1, 12000, (2, 120)) // (批次, 目标长度) - 例如右移的目标

// 如果GPU可用,将模型和数据移至GPU
// device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
// transformer_model.to(device)
// src_dummy = src_dummy.to(device)
// tgt_dummy = tgt_dummy.to(device)

// 前向传播
val output_logits = transformer_model(src_dummy, tgt_dummy)
println("最终输出形状 (logits):", output_logits.shape) // 应为 (2, 120, 12000)

通过从这些基本PyTorch模块组装Transformer,您将具体了解信息如何在模型中流动以及注意力机制如何实现语境感知的序列处理。这种基于组件的实现也提供了一个灵活的构建方式,用于试验架构变体或使模型适应不同任务。请记住,高效训练此类模型需要仔细考虑优化、正则化和数据处理,这些主题将在后续章节中介绍。

高级注意力机制

构成Transformer架构核心的标准自注意力机制,计算序列中所有 token 之间的成对互动。这导致复杂度与序列长度 NN 呈平方关系增长,具体为 O(N2⋅d)O(N2⋅d),其中 dd 是模型维度。这种平方级增长对于涉及非常长文档、被视为补丁序列的高分辨率图像或扩展时间序列的应用来说,变得难以承受。为解决这些问题,高级注意力机制旨在应对这些特定局限性,最主要的是与长序列相关的计算和内存需求。

高级注意力机制主要目标是降低这种 O(N2)O(N2) 复杂度到更易于管理的状态,通常是线性或接近线性的(O(N)O(N) 或 O(Nlog⁡N)O(NlogN)),同时试图保持原始注意力公式的建模能力。

稀疏注意力模式

一种方法是使注意力矩阵稀疏。不再是每个 token 都关注其他所有 token,每个 token 只关注一个受限的子集。这种限制通常基于预定义模式。

  • 局部注意力: Token 只关注固定大小的相邻 token 窗口。这能有效捕获局部背景信息,但会遗漏窗口外的长距离依赖关系。滑动窗口注意力是一种常见实现方式。
  • 步进或膨胀注意力: Token 关注固定间隔位置的 token(例如,每隔 kk 个 token)。这能捕获序列中远距离部分的信息,但可能会遗漏不在步进范围内的相邻 token 之间的互动。
  • 组合模式: Longformer 或 BigBird 等更复杂的方法结合了局部注意力、膨胀注意力,有时还会加入一些全局 token(如 [CLS] token),这些 token 关注所有其他 token,也同时被所有其他 token 关注。这试图兼顾两方面优点:局部细节和稀疏的全局背景信息。

实现这些方法通常涉及在 softmax 操作之前精心遮盖注意力分数矩阵,或者使用专门的索引和收集操作来只计算必要的分数。

线性化和高效注意力

另一类方法旨在近似标准注意力机制或重新构建其计算方式,以避免显式计算 N×NN×N 的注意力矩阵 QKTQKT。这些方法通常目标是达到 O(N)O(N) 复杂度。

单个头的标准注意力输出为:

注意力(Q,K,V)=softmax(QKTdk)V注意力(Q,K,V)=softmax(dkQKT)V

线性注意力方法研究近似或重写此公式的方法。例如,如果我们能使用核函数 ϕϕ 来表示 softmax 函数(或其近似),使得 softmax(xiTxj)≈ϕ(xi)Tϕ(xj)softmax(xiTxj)≈ϕ(x**i)T**ϕ(x**j),我们就有可能重写计算方式。

考虑一个不带缩放因子和 softmax 的简化版本:A=QKTVA=QKT**V。这可以重新排序为 A=Q(KTV)A=Q(KTV)。KTVKTV 的计算需要 O(Ndkdv)O(Ndkdv) 时间,乘以 QQ 需要 O(Ndkdv)O(Ndkdv),这导致总体复杂度相对于序列长度 NN 为 O(N)O(N)(假设 dk,dvd**k,d**v 是固定的)。

难点在于纳入 softmax 非线性,同时保持线性复杂度。

  • Performer: 使用基于 Fastfood 算法的随机特征映射来近似 softmax 函数中隐含的高斯核。这实现了注意力机制的线性时间近似。
  • Linformer: 对键 (KK) 和值 (VV) 矩阵应用线性投影,有效地在注意力计算之前降低序列长度维度,从而用低秩矩阵近似完整的注意力矩阵。
  • 其他基于核的方法: 研究不同的核函数 ϕϕ 来近似 softmax 操作,从而实现 Q(KTV)Q(KTV) 的重新排列。

这些方法以准确性换取效率。近似方法的选择会影响模型与标准注意力相比捕获复杂依赖关系的能力。

PyTorch 中的实现考量

尽管你可以从头开始实现稀疏遮盖或核近似,但这可能很复杂,并且需要仔细优化才能达到良好性能。幸运的是,PyTorch 生态系统提供了工具和库:

  • 自定义遮盖: 对于像局部或步进注意力这样的稀疏模式,你通常可以使用 torch.nn.MultiheadAttentiontorch.nn.functional.scaled_dot_product_attention(在较新的 PyTorch 版本中可用)中的 attn_mask 参数。你需要构造一个布尔遮罩,其中 True 表示不应被关注的位置。
  • 专门库: 像 Meta AI 的 xformers 这样的库提供了各种注意力机制的高度优化实现,包括稀疏和内存高效的变体,通常与 CUDA 内核集成以获得最高速度。对于性能要求高的应用,通常建议使用这些库。
import torch
import torch.nn as nn

// 检查 xformers 是否可用于优化注意力
try:
    from xformers.ops import memory_efficient_attention
// 示例用法(API 细节可能有所不同 - 请查阅 xformers 文档)
    // 假设 q, k, v 形状正确(Batch, Seq, Heads, HeadDim 或类似)
    // output = memory_efficient_attention(q, k, v)
    // println("正在使用 xformers memory_efficient_attention")
    XFORMERS_AVAILABLE = true
except ImportError:
    // println("xformers 不可用。需要标准 PyTorch 注意力或手动实现。")
    XFORMERS_AVAILABLE = false

// 在标准 PyTorch 函数式 API 中使用注意力遮罩的示例
// 假设 embed_dim = 64, num_heads = 8, seq_len = 5, batch_size = 2
val embed_dim = 64
val num_heads = 8
val seq_len = 5
val batch_size = 2

// 虚拟输入张量 (Batch, SeqLen, EmbedDim)
val query = torch.randn(batch_size, seq_len, embed_dim)
val key = torch.randn(batch_size, seq_len, embed_dim)
val value = torch.randn(batch_size, seq_len, embed_dim)

// 如果函数需要,为多头注意力重塑形状
// 或在 nn.Module 包装器内处理

// 创建因果遮罩(例如,用于解码器)
// 遮罩需要根据注意力函数设置适当的维度
// 对于 scaled_dot_product_attention,(SeqLen, SeqLen) 遮罩通常是可广播的
val causal_mask_bool = torch.triu(torch.ones(seq_len, seq_len, dtype=torch.bool), diagonal=1)

// 使用 torch.nn.functional.scaled_dot_product_attention (PyTorch 2.0+)
// 注意:此函数在内部处理重塑和缩放
// 它期望布尔遮罩,其中 True 表示“遮盖掉”
try:
    output_sdpa = nn.functional.scaled_dot_product_attention(
        query, value, attn_mask=causal_mask_bool, is_causal=false # 显式遮罩示例
        // 或者使用 is_causal=True 进行自动因果遮罩:
        // output_sdpa = nn.functional.scaled_dot_product_attention(query, key, value, is_causal=True)
    )
    // println("已使用 nn.functional.scaled_dot_product_attention")
catch AttributeError:
    // println("scaled_dot_product_attention 不可用(需要 PyTorch 2.0+)。")
    // 回退到 nn.MultiheadAttention 或手动实现
    pass

// 使用 nn.MultiheadAttention 的示例(需要特定格式的遮罩)
// MHA 期望布尔遮罩为 (Batch * NumHeads, TargetSeqLen, SourceSeqLen) 或 (TargetSeqLen, SourceSeqLen)
val multihead_attn = nn.MultiheadAttention(embed_dim, num_heads, batch_first=true)
// MHA 遮罩:True 表示该位置*将被阻止*关注。
// 创建一个更简单的遮罩用于说明(适用于所有头/批次)
val mha_mask = torch.triu(torch.ones(seq_len, seq_len), diagonal=1).bool()
// attn_output, attn_weights = multihead_attn(query, key, value, attn_mask=mha_mask)
// println("已使用带遮罩的 nn.MultiheadAttention")

上述代码片段说明了你可以在哪里集成来自 xformers 等库的优化注意力,或者标准 PyTorch 函数如何接受注意力遮罩。确切的 API 调用和遮罩形状取决于使用的特定 PyTorch 版本和函数。请始终参考官方文档以获取精确用法。

权衡

选择注意力机制涉及平衡计算效率、内存使用和模型性能。

  • 标准注意力: 表现力最强,捕获所有成对互动,但计算开销大(O(N2)O(N2))。
  • 稀疏注意力: 降低复杂度,适用于局部或预定义全局互动就足够的模式。可能会遗漏不在稀疏模式范围内的重要互动。
  • 线性/高效注意力: 通常能达到 O(N)O(N) 复杂度,非常适合长序列。依赖近似方法,这可能会在需要高度精确长距离依赖关系的任务上,与标准注意力相比略微降低性能。

最佳选择很大程度上取决于具体任务、涉及的序列长度以及可用的计算资源。通常需要通过实验来找到最合适的方案。

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

标准 (O(N2)O(N2)) 与线性 (O(N)O(N)) 注意力机制的计算成本随序列长度增加的理论增长曲线。请注意两个坐标轴都采用对数刻度。线性注意力复杂度以任意常数因子为例进行呈现,以便比较。

此图显示了与线性替代方案相比,标准注意力的成本增长有多快,使得后者对于有效处理长序列必不可少。当你构建更复杂的模型时,理解和应用这些高级注意力机制将对管理计算资源和扩展你的架构有重要作用。

使用 PyTorch Geometric 的图神经网络

许多数据集本质上是关系型的,最适合用图来表示。例子有社交网络、分子结构、引用网络、知识图谱和推荐系统。传统深度学习结构,如 CNN 和 RNN,假定数据结构是网格状或序列状的,这使它们不适合处理图中存在的任意连接。图神经网络 (GNN) 专门设计用于直接处理图结构数据,它们学习到的表示会同时包含节点特征和图的拓扑结构。

PyTorch Geometric (PyG) 是一个功能强大且被广泛采用的库,它建立在 PyTorch 之上,用于开发和应用 GNN。它提供了多种 GNN 层的优化实现、高效的图数据处理以及常见的图基准数据集。本节将指导您如何使用 PyG 来实现和理解不同的 GNN 结构。

在 PyTorch Geometric 中表示图

在构建 GNN 模型之前,我们需要一种标准化的方式来表示图数据。PyG 使用 torch_geometric.data.Data 对象。一个 Data 对象包含描述单个图的各种属性:

  • x: 节点特征矩阵,形状为 [num_nodes, num_node_features]。每行代表一个节点,列代表其特征。
  • edge_index: 图的连接信息,采用 COO (坐标) 格式,形状为 [2, num_edges]。它存储每条边的源节点和目标节点索引。对于从节点 j 到节点 i 的边,其列为 [j, i]。这种表示对于稀疏图是高效的。
  • edge_attr: 边特征矩阵,形状为 [num_edges, num_edge_features]。表示与每条边相关的可选特征。
  • y: 目标标签或值,取决于具体任务。对于节点级任务,形状为 [num_nodes, ...];对于图级任务,形状为 [1, ...]
  • pos: 节点位置特征,形状为 [num_nodes, num_dimensions]。常用于几何深度学习。

下面是创建简单 Data 对象的方法:

import torch
from torch_geometric.data import Data

// 节点特征:3 个节点,每个节点 2 个特征
val x = torch.tensor(Seq(Seq(1, 2), Seq(3, 4), Seq(5, 6)), dtype=torch.float)

// 边:(0 -> 1), (1 -> 0), (1 -> 2), (2 -> 1)
// 表示为源节点和目标节点
val edge_index = torch.tensor(Seq(Seq(0, 1, 1, 2),  // 源节点
                               Seq(1, 0, 2, 1)), // 目标节点
                              dtype=torch.long)

// 可选的边特征:4 条边,每条边 1 个特征
val edge_attr = torch.tensor(Seq(Seq(0.5), Seq(0.5), Seq(0.8), Seq(0.8)), dtype=torch.float)

// 可选的节点标签(例如,用于节点分类)
val y = torch.tensor(Seq(0, 1, 0), dtype=torch.long)

// 创建 Data 对象
val graph_data = Data(x=x, edge_index=edge_index, edge_attr=edge_attr, y=y)

println(graph_data)
// 输出:Data(x=[3, 2], edge_index=[2, 4], edge_attr=[4, 1], y=[3])

PyG 还提供了 torch_geometric.data.Datasettorch_geometric.loader.DataLoader,用于高效处理图集合并创建小批量数据。DataLoader 会自动将不同大小的图整理成更大的批处理对象。

消息传递方法

大多数 GNN 层都基于消息传递原理运行。核心思想是每个节点通过聚合来自其局部邻域的信息,迭代地更新其特征表示(嵌入)。这个过程通常包含对层 ll 中每个节点 ii 的三个步骤:

  1. 消息计算: 每个邻居节点 j∈N(i)j∈N(i) 根据其自身特征 hj(l−1)hj(l−1),以及目标节点特征 hi(l−1)hi(l−1) 和边特征 ej,iej,i,计算一个消息 mj→i(l)mji(l)。

    mj→i(l)=ϕ(l)(hi(l−1),hj(l−1),ej,i)mji(l)=ϕ(l)(hi(l−1),hj(l−1),ej,i)

    其中 ϕ(l)ϕ(l) 是一个可微分的消息函数(例如,一个神经网络)。

  2. 聚合: 节点 ii 使用一个置换不变函数 ⨁⨁(如求和、平均或最大值)来聚合来自其邻居的所有传入消息。

    ai(l)=⨁j∈N(i)mj→i(l)ai(l)=j∈N(i)⨁mji(l)

  3. 更新: 节点 ii 根据其先前的表示 hi(l−1)hi(l−1) 和聚合后的消息 ai(l)ai(l) 来更新其特征向量 hi(l)hi(l)。

    hi(l)=γ(l)(hi(l−1),ai(l))hi(l)=γ(l)(hi(l−1),ai(l))

    其中 γ(l)γ(l) 是一个可微分的更新函数(例如,另一个神经网络或简单地添加聚合消息)。

初始特征 hi(0)hi(0) 通常是输入节点特征 data.x。堆叠多个消息传递层可以使信息在图中传播更远的距离。

层 l-1层 lhᵢ(l-1)hᵢ(l)更新(hᵢ(l-1), 聚合)hⱼ₁(l-1)m₁hⱼ₂(l-1)m₂hⱼ₃(l-1)m₃

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

此图表呈现了更新节点 ii 的消息传递理念。来自邻居 j1,j2,j3j1,j2,j3 的信息(消息 m1,m2,m3m1,m2,m3)被聚合,并与节点的先前状态 hi(l−1)h**i(l−1) 结合,以计算出新状态 hi(l)h**i(l)。

PyG 在其层类中提供了这些步骤的优化实现。

PyTorch Geometric 中常用的 GNN 层

PyG 提供了多种预实现的 GNN 层。让我们看看三个流行的例子:GCN、GraphSAGE 和 GAT。

图卷积网络 (GCN)

GCN 层由 Kipf & Welling (2017) 提出,执行基于谱的图卷积。GCN 层的消息传递更新规则可以简化为:

H(l+1)=σ(D−1/2AD−1/2H(l)W(l))**H**(*l*+1)=*σ*(**D**−1/2A**D**−1/2H(l)W(l))

其中,H(l)H(l) 是层 ll 的节点嵌入矩阵,W(l)W(l) 是一个可训练的权重矩阵,σσ 是一个激活函数(如 ReLU),A=A+I**A**=A+I 是添加了自循环的邻接矩阵,D**D** 是 A**A** 的对角度矩阵。项 D−1/2AD−1/2**D**−1/2A**D**−1/2 表示邻接矩阵的对称归一化。该层平均邻居节点(包括节点自身)的特征,然后应用线性变换,再进行非线性处理。

在 PyG 中,您使用 torch_geometric.nn.GCNConv

import torch.nn.functional as F
import torch_geometric.nn.GCNConv

class SimpleGCN extends torch.nn.Module:
    def __init__(num_node_features: Int, num_classes: Int, hidden_channels: Int):
        super().__init__()
        val conv1 = GCNConv(num_node_features, hidden_channels)
        val conv2 = GCNConv(hidden_channels, num_classes)

    def forward(data: Data):
        val x, edge_index = data.x, data.edge_index

        x = conv1(x, edge_index)
        x = F.relu(x)
        x = F.dropout(x, p=0.5, training=self.training) // 经常使用 Dropout
        x = self.conv2(x, edge_index)

        // 对节点进行分类,通常使用 LogSoftmax
        return F.log_softmax(x, dim=1)
GraphSAGE

GraphSAGE (Hamilton et al., 2017) 专注于学习聚合函数,而不是固定的卷积。它被设计为归纳式的,这意味着它可以在推断时推广到未见过的节点。GraphSAGE 为每个节点采样固定大小的邻域,然后使用平均、最大值或 LSTM 池化等函数聚合邻居特征。

主要步骤包括:

  1. 为节点 ii 采样一个邻域 N(i)N(i)。
  2. 聚合邻居特征:aN(i)(l)=AGGREGATE(l)({hj(l−1)∣j∈N(i)})aN(i)(l)=AGGREGATE(l)({hj(l−1)∣j∈N(i)})
  3. 更新节点 ii 的嵌入:hi(l)=σ(W(l)⋅CONCAT(hi(l−1),aN(i)(l)))hi(l)=σ(W(l)⋅CONCAT(hi(l−1),aN(i)(l)))

PyG 使用 torch_geometric.nn.SAGEConv 实现此功能:

import torch_geometric.nn.SAGEConv

class SimpleGraphSAGE extends torch.nn.Module:
    def __init__(num_node_features: Int, num_classes: Int, hidden_channels: Int):
        super().__init__()
        // 默认聚合器是 'mean'
        val conv1 = SAGEConv(num_node_features, hidden_channels)
        val conv2 = SAGEConv(hidden_channels, num_classes)

    def forward(data: Data):
        val x, edge_index = data.x, data.edge_index

        x = conv1(x, edge_index)
        x = F.relu(x)
        x = F.dropout(x, p=0.5, training=self.training)
        x = self.conv2(x, edge_index)

        // 对节点进行分类,通常使用 LogSoftmax
        return F.log_softmax(x, dim=1)

在创建 SAGEConv 层时,您可以指定聚合器类型(例如,aggr='max'aggr='mean')。

图注意力网络 (GAT)

GAT 层 (Veličković et al., 2018) 引入了注意力机制,使得节点在聚合过程中可以为其邻居分配不同的重要性权重。这使得聚合过程更加灵活,并通常带来更好的性能。

节点 ii 和邻居 jj 之间的注意力系数 eije**ij 基于它们的特征计算,通常使用共享的线性变换和一个注意力机制(例如,一个单层前馈网络):

eij=attention(W(l)hi(l−1),W(l)hj(l−1))e**ij=attention(W(l)hi(l−1),W(l)hj(l−1))

然后,这些系数使用 softmax 函数对节点 ii 的所有邻居进行归一化:

αij=softmaxj(eij)=exp⁡(eij)∑k∈N(i)exp⁡(eik)α**ij=softmaxj(e**ij)=∑k∈N(i)exp(e**ik)exp(e**ij)

聚合后的消息是转换后的邻居特征的加权和:

ai(l)=∑j∈N(i)αijW(l)hj(l−1)ai(l)=j∈N(i)∑αijW**(l)hj(l−1)

更新步骤将此聚合消息与节点自身的特征结合,通常使用拼接后跟激活函数:

hi(l)=σ(ai(l))orhi(l)=σ(CONCAT(hi(l−1),ai(l)))hi(l)=σ(ai(l))orhi(l)=σ(CONCAT(hi(l−1),ai(l)))

GAT 经常使用多头注意力,其中计算多个独立的注意力机制,并将其结果进行拼接或平均。

PyG 使用 torch_geometric.nn.GATConv 实现此功能:

import torch.nn.functional as F
import torch_geometric.nn.GATConv

class SimpleGAT extends torch.nn.Module:
    def __init__(num_node_features: Int, num_classes: Int, hidden_channels: Int, heads: Int = 8):
        super().__init__()
        // 在第一层中使用多头注意力
        val conv1 = GATConv(num_node_features, hidden_channels, heads=heads, dropout=0.6)
        // 多头注意力的输出特征为 heads * hidden_channels
        // 对于最后一层,通常平均各头或使用单头
        val conv2 = GATConv(hidden_channels * heads, num_classes, heads=1, concat=False, dropout=0.6)

    def forward(data: Graph):
        val x, edge_index = data.x, data.edge_index

        x = F.dropout(x, p=0.6, training=self.training) // 对输入特征应用 Dropout
        x = conv1(x, edge_index)
        x = F.elu(x) // ELU 激活在 GAT 中很常见
        x = F.dropout(x, p=0.6, training=self.training)
        x = conv2(x, edge_index)

        return F.log_softmax(x, dim=1)

构建和训练 GNN 模型

使用 PyG 层在 PyTorch 中构建 GNN 遵循标准的 PyTorch 实践。您定义一个继承自 torch.nn.Module 的类,在 __init__ 中初始化 PyG 层,并在 forward 中定义前向传播逻辑。forward 方法通常接受 DataBatch 对象作为输入,并提取 xedge_index,以及可能的 edge_attrbatch 索引。

训练循环也类似于标准的 PyTorch 循环:遍历 DataLoader,执行前向传播,计算损失(例如,用于节点分类的 F.nll_loss,配合 log_softmax),使用 loss.backward() 计算梯度,并使用优化器更新参数。

常见的 GNN 应用

GNN 功能多样,可应用于各种图相关任务:

  1. 节点分类: 预测图中单个节点的标签或属性(例如,对社交网络中的用户进行分类,预测蛋白质功能)。上述示例 (SimpleGCNSimpleGraphSAGESimpleGAT) 均适用于节点分类。
  2. 图分类: 预测整个图的标签或属性(例如,将分子分类为有毒或无毒,对社交群体进行分类)。这需要在 GNN 层之后使用图池化层(例如,torch_geometric.nn.global_mean_poolglobal_max_pool)将节点嵌入聚合成单个图嵌入。
  3. 链接预测: 预测两个节点之间是否存在或将存在边(例如,在社交网络中推荐朋友,预测蛋白质-蛋白质相互作用)。这通常包括学习节点嵌入,然后使用评分函数(例如,点积)对节点嵌入对进行处理。
  4. 图生成: 生成具有期望属性的新图。
  5. 社区检测: 在大型图中识别连接紧密的节点群组。

PyTorch Geometric 为应对这些任务提供了全面的工具。通过结合其优化层、数据处理工具和标准 PyTorch 功能,您可以有效地构建和训练处理复杂图问题的精巧 GNN 模型。请记住,选择合适的 GNN 结构(GCN、GAT、SAGE 或其他)通常取决于您的图数据的具体特征和当前的任务。进行实验和理解每个层的基本原理对于成功应用非常重要。

用于生成建模的归一化流

收藏

与变分自编码器(VAEs)或生成对抗网络(GANs)相比,归一化流为生成建模提供了一种独特的方法。其主要优点在于能够精确计算数据似然,同时在数据空间与具有简单分布(如标准高斯分布)的潜在空间之间定义一个可逆映射。这使得它们特别适用于需要精确密度估计或可逆性有益的场合。

核心思路是学习一个变换 f:X→Zf:X→Z,将复杂的输入数据点 x∈Xx∈X 映射到更简单的潜在变量 z∈Zz∈Z,其中 pZ(z)p**Z(z) 是一个已知、易于处理的概率分布(例如,N(0,I)N(0,I))。由于 ff 被设计为可逆的 (x=f−1(z)x=f−1(z)) 且可微分,我们可以使用概率论中的变量变换定理来关联数据 pX(x)p**X(x) 的密度与潜在变量 pZ(z)p**Z(z) 的密度。

变量变换公式

变量变换公式表明,对于一个将 xx 映射到 zz 的可逆、可微分函数 ff,它们概率密度之间的关系是:

pX(x)=pZ(f(x))∣det⁡(∂f(x)∂xT)∣p**X(x)=p**Z(f(x))det(∂x**Tf(x))

在此,∂f(x)∂xT∂x**Tf(x) 是变换 ff 在 xx 处计算的雅可比矩阵,而 ∣det⁡(⋅)∣∣det(⋅)∣ 表示其行列式的绝对值。

这个公式很重要。它告诉我们,如果能计算 f(x)f(x) 及其雅可比矩阵的行列式,我们就可以使用已知的 pZp**Z 密度,计算任何给定数据点 xx 的精确概率密度 pX(x)p**X(x)。

对于生成建模,我们通常使用对数似然,这避免了数值下溢并简化了计算:

log⁡pX(x)=log⁡pZ(f(x))+log⁡∣det⁡(∂f(x)∂xT)∣logp**X(x)=logp**Z(f(x))+logdet(∂x**Tf(x))

训练归一化流涉及在数据集上最大化此对数似然。这要求变换 ff 具备两个性质:

  1. 它必须易于可逆(f−1f−1 应该可计算)。
  2. 其雅可比行列式必须计算高效。

使用双射函数构建流

复杂的变换 ff 通常通过组合更简单的可逆函数来构建,这些函数常被称为双射函数耦合层:f=fL∘⋯∘f2∘f1f=f**L∘⋯∘f2∘f1。如果每个 fif**i 都是可逆的,并且具有易于计算的雅可比行列式,那么复合函数 ff 就会继承这些性质。

复合函数 ff 的雅可比矩阵是各层雅可比矩阵的乘积:

∂f(x)∂xT=∂fL(zL−1)∂zL−1T…∂f2(z1)∂z1T∂f1(x)∂xT∂x**Tf(x)=∂z**L−1Tf**L(z**L−1)…∂z1Tf2(z1)∂x**Tf1(x)

其中 zi=fi(zi−1)z**i=f**i(z**i−1) 且 z0=xz0=x

由于性质 det⁡(AB)=det⁡(A)det⁡(B)det(A**B)=det(A)det(B),整体雅可比矩阵的对数行列式变成一个和:

log⁡∣det⁡(∂f(x)∂xT)∣=∑i=1Llog⁡∣det⁡(∂fi(zi−1)∂zi−1T)∣logdet(∂x**Tf(x))=i=1∑Llogdet(∂z**i−1Tf**i(z**i−1))

这种组合性使我们能够从更简单的模块构建富有表现力的变换,同时保持计算的可处理性。

这是一个归一化流的图示:

变换 f数据空间x ~ p_X(x)f₁z₁=f₁(x)潜在空间z ~ p_Z(z)(例如,高斯分布)f_Lz_{L-1}=f_L⁻¹(z)x=f₁⁻¹(z₁)f₂z₂=f₂(z₁)z₁=f₂⁻¹(z₂)…z=f_L(z_{L-1})…数据空间 (X) 与潜在空间 (Z) 之间的映射通过一系列可逆变换 (f₁, f₂, …, f_L) 实现。

外链图片转存失败,源站可能有防盗链机制,建议将图片保存下来直接上传

数据空间 XX 通过一系列可逆函数 f1,…,fLf1,…,f**L 变换为简单的潜在空间 ZZ(例如,高斯分布)。逆变换 f−1f−1 允许通过从 pZ(z)p**Z(z) 中抽取 zz 并计算 x=f−1(z)x=f−1(z) 来从 pX(x)p**X(x) 采样。

常见双射函数架构

设计有效的双射函数是归一化流的核心。一种常用且成功的方法包含耦合层

耦合层(例如,RealNVP,NICE)

Real Non-Volume Preserving (RealNVP) 流,基于NICE (Non-linear Independent Components Estimation) 构建,采用巧妙的掩码策略。它们将输入向量 xx 分成两部分,x1x1 和 x2x2。变换按以下方式进行:

  1. 第一部分 x1x1 不变通过:z1=x1z1=x1。
  2. 第二部分 x2x2 使用一个函数(通常是仿射变换)进行变换,其参数由 x1x1 决定。例如,一个仿射耦合层计算:z2=x2⊙exp⁡(s(x1))+t(x1)z2=x2⊙exp(s(x1))+t(x1)在此,⊙⊙ 表示按元素乘法,而 s(⋅)s(⋅)(缩放)和 t(⋅)t(⋅)(平移)是复杂函数,通常作为神经网络实现,以 x1x1 为输入。

逆变换也很简单:

x1=z1x1=z1x2=(z2−t(z1))⊙exp⁡(−s(z1))x2=(z2−t(z1))⊙exp(−s(z1))

此变换的雅可比矩阵是下三角矩阵(或上三角矩阵,取决于哪个部分被变换):

∂z∂xT=(I0∂z2∂x1T∂z2∂x2T)∂x**Tz=(Ix1Tz20∂x2Tz2)

值得注意的是,∂z2∂x2T∂x2Tz2 是一个对角矩阵,其对角元素为 exp⁡(s(x1))exp(s(x1))。因此,行列式就是对角元素的乘积:∏exp⁡(s(x1))=exp⁡(∑s(x1))∏exp(s(x1))=exp(∑s(x1))。对数行列式就是缩放网络输出的总和:∑s(x1)∑s(x1)。

这种结构保证了易于求逆和对数行列式的高效计算。为了确保所有维度都被变换,连续的耦合层通常会交换 x1x1 和 x2x2 的角色或使用不同的掩码。

自回归流(例如,MAF,IAF)

另一类重要的模型是自回归流。在这些模型中,潜在变量的每个维度 ziz**i 仅以前面的维度 z1:i−1z1:i−1 为条件(或 x1:i−1x1:i−1,取决于方向)。

  • 掩码自回归流(MAF): 密度估计高效,但采样速度慢,因为采样需要顺序计算。
  • 逆自回归流(IAF): 采样高效(可并行),但密度估计速度慢。

它们常使用掩码神经网络来强制自回归特性。

在PyTorch中实现归一化流

PyTorch的 torch.distributions 模块提供了构建归一化流的优秀工具。重要组件包括:

  • torch.distributions.Distribution:概率分布的基类(如 Normal)。
  • torch.distributions.Transform:可逆变换的基类。提供了如 AffineTransformExpTransform 等实现。通常会继承此基类定义自定义变换。
  • torch.distributions.TransformedDistribution:通过将一系列 Transform 对象应用于基础分布来创建新分布。它自动处理变量变换的计算。
  • torch.distributions.constraints:用于定义分布的支持域和变换参数的有效性检查。

让我们大致描述一下如何使用耦合层定义一个简单流。您通常会定义一个继承自 torch.distributions.TransformCouplingLayer 类。

import torch
import torch.nn as nn
import torch.distributions as dist

class CouplingLayer extends dist.Transform:
    def __init__(self, input_dim: Int, hidden_dim: Int, mask: Tensor, base_transform_type: String = 'affine'):
        super().__init__()
        val input_dim = input_dim
        // 确保掩码是二元张量(0和1)
        val register_buffer('mask', mask) 

        // 定义计算缩放和平移参数的网络 's_t_network'
        val s_t_network = nn.Sequential(
            nn.Linear(input_dim, hidden_dim),
            nn.ReLU(),
            nn.Linear(hidden_dim, hidden_dim),
            nn.ReLU(),
            nn.Linear(hidden_dim, input_dim * 2) # 输出所有维度的缩放和平移
        )
        
        val bijective = True
        // 对于仿射变换:定义域 = 实数,值域 = 实数
        val domain = dist.constraints.real_vector
        val codomain = dist.constraints.real_vector
        
        // 可选:使用内置变换,如 AffineTransform 
        // if base_transform_type == 'affine' ...(实现细节)

    def _call(self, x):
        """ 应用变换:x -> z """
        val x_masked = x * self.mask
        val s_t_params = self.s_t_network(x_masked)
        // 将输出分为缩放(s)和平移(t)
        // 确保缩放为正,例如使用 tanh 增加稳定性和缩放
        val s = torch.tanh(s_t_params[..., :self.input_dim]) 
        val t = s_t_params[..., self.input_dim:]

        // 仅对未掩码的元素应用变换
        // z = x_masked + (x_unmasked * exp(s) + t) * (1 - mask) 
        val z = self.mask * x + (1 - self.mask) * (x * torch.exp(s) + t)
        return z

    def _inverse(self, z):
        """ 应用逆变换:z -> x """
        val z_masked = z * self.mask
        val s_t_params = self.s_t_network(z_masked)
        val s = torch.tanh(s_t_params[..., :self.input_dim])
        val t = s_t_params[..., self.input_dim:]

        // 仅对未掩码的元素应用逆变换
        // x = z_masked + ((z_unmasked - t) * exp(-s)) * (1 - mask)
        val x = self.mask * z + (1 - self.mask) * ((z - t) * torch.exp(-s))
        return x

    def log_abs_det_jacobian(self, x, z):
        """ 计算 log |det J(x)| """
        val x_masked = x * self.mask
        val s_t_params = self.s_t_network(x_masked)
        val s = torch.tanh(s_t_params[..., :self.input_dim])

        // 对数行列式是变换维度上 's' 的总和
        val log_det_jacobian = (1 - self.mask) * s 
        // 对变换维度求和
        return log_det_jacobian.sum(-1) 

// 示例用法:
val input_dim = 10
val hidden_dim = 64
val num_flows = 5

// 定义基础分布(例如,标准正态分布)
val base_dist = dist.Normal(torch.zeros(input_dim), torch.ones(input_dim))

// 创建掩码(在步骤之间交替是常见的)
val masks = for i <- 0 until num_flows yield {
    val mask = torch.zeros(input_dim)
    mask(i % 2::2) = 1 // 简单的交替掩码
    // 或者存在更复杂的掩码策略
    mask
}
    
// 创建一系列变换(耦合层)
val transforms = for i <- 0 until num_flows yield {
    // 交替哪一部分是恒等变换,哪一部分是变换
    val current_mask = if (i % 2 == 0) masks(i) else (1 - masks(i)) 
    new CouplingLayer(input_dim, hidden_dim, current_mask)
}
    // 可选:在流之间添加置换/激活归一化层
}

// 构建变换后的分布
val flow_dist = dist.TransformedDistribution(base_dist, transforms)

// --- 训练 ---
// optimizer = torch.optim.Adam(flow_dist.parameters(), lr=1e-4)
// 
// for data_batch in dataloader:
//     optimizer.zero_grad()
//     
//     // 计算数据批次的对数概率
val log_prob = flow_dist.log_prob(data_batch) // 形状:[batch_size]
//     
//     // 最大化对数似然 -> 最小化负对数似然
//     loss = -log_prob.mean() 
//     
//     loss.backward()
//     optimizer.step()

// --- 采样 ---
// val n_samples = 64
// val samples = flow_dist.sample(torch.Size([n_samples])) // 形状:[n_samples, input_dim]

优点与考量

优点:

  • 精确似然: 允许直接计算和优化数据似然,这与GAN或VAE(它们使用下界)不同。这对于密度估计和模型比较很有益处。
  • 可逆性: 数据与潜在空间之间的映射是明确可逆的,这对于需要在潜在空间中分析的任务很有用。
  • 高效采样: 某些架构(如IAF)支持并行采样。即使是基于耦合的流,通过顺序应用逆变换,通常也能高效地进行采样。

考量:

  • 架构约束: 双射函数必须可逆且雅可比矩阵易于处理,这限制了可使用的函数类型。与约束较少的模型相比,这可能会限制表现力,尽管复杂的流仍能对复杂的分布进行建模。
  • 计算成本: 即使雅可比行列式易于处理,计算它们也会在训练期间增加计算开销,特别是对于高维数据或深层流。
  • 拓扑结构: 基本流不能改变空间的拓扑结构;它们是微分同胚。这意味着如果从像高斯分布这样简单的连通基础分布开始,它们可能难以完美地建模具有不连通模式的分布。

应用

归一化流已在多个应用场景中得到采用:

  • 生成建模: 生成真实的图像、音频和其他高维数据。
  • 密度估计: 精确建模复杂的概率分布。
  • 变分推断: 在贝叶斯模型中使用流来定义更灵活的近似后验分布,优于像对角高斯这样的简单近似。
  • 强化学习: 建模复杂的策略或价值函数。
  • 表示学习: 潜在空间 zz 可以提供数据 xx 的有意义表示。

总之,归一化流为生成建模和密度估计提供了一个数学上优雅且功能强大的框架。凭借可逆变换和变量变换公式,它们允许进行精确的似然计算,为高级深度学习工具集中特定概率建模任务提供了特有的优点。

Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐