用PyTorch复现AirFormer:手把手教你搭建空气质量预测Transformer(附完整代码)

空气质量预测一直是环境科学和机器学习交叉领域的重要课题。传统方法往往难以捕捉空气污染物在广域范围内的复杂时空关联,而Transformer架构凭借其强大的序列建模能力,为解决这一挑战提供了新思路。本文将深入解析AirFormer论文的核心创新点,并提供一个模块化的PyTorch实现方案,帮助开发者快速掌握这一前沿技术。

1. 环境准备与数据预处理

在开始构建模型前,我们需要配置合适的开发环境并准备空气质量数据集。以下是推荐的环境配置:

# 环境依赖安装
pip install torch==1.12.1+cu113 torchvision==0.13.1+cu113 -f https://download.pytorch.org/whl/torch_stable.html
pip install pandas scikit-learn matplotlib

空气质量数据集通常包含以下关键字段:

  • 站点ID:监测站的唯一标识符
  • 时间戳:观测记录的时间
  • PM2.5/PM10:颗粒物浓度
  • 气象数据:温度、湿度、风速等
  • 地理位置:经度、纬度
import pandas as pd
from sklearn.preprocessing import StandardScaler

def load_air_quality_data(file_path):
    # 加载原始数据
    df = pd.read_csv(file_path, parse_dates=['timestamp'])
    
    # 处理缺失值
    df = df.interpolate(method='linear', limit_direction='both')
    
    # 标准化特征
    numeric_cols = ['PM2.5', 'PM10', 'temperature', 'humidity']
    scaler = StandardScaler()
    df[numeric_cols] = scaler.fit_transform(df[numeric_cols])
    
    # 添加时间特征
    df['hour'] = df['timestamp'].dt.hour
    df['day_of_week'] = df['timestamp'].dt.dayofweek
    
    return df, scaler

提示:实际应用中,建议将数据集划分为训练集(70%)、验证集(15%)和测试集(15%),并确保时间连续性不被破坏。

2. 模型架构设计

AirFormer的核心创新在于其独特的双阶段架构:

2.1 自下而上的确定性阶段

这一阶段包含两个关键组件:

Dartboard Spatial MSA (DS-MSA) 通过区域划分大幅降低计算复杂度:

class DartboardMSA(nn.Module):
    def __init__(self, embed_dim, num_heads, region_size=25):
        super().__init__()
        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads
        self.region_size = region_size
        
        # 可学习的映射矩阵
        self.projection = nn.Parameter(torch.randn(region_size, embed_dim))
        self.qkv = nn.Linear(embed_dim, embed_dim * 3)
        self.out_proj = nn.Linear(embed_dim, embed_dim)
        
    def forward(self, x, station_map):
        """
        x: [batch, time, stations, embed_dim]
        station_map: [stations, region_size] 站点到区域的映射矩阵
        """
        batch, time, stations, _ = x.shape
        
        # 生成区域表示
        region_feat = torch.einsum('sd,btsc->btrc', station_map, x)  # [b,t,regions,c]
        
        # 计算QKV
        q = self.qkv(x)[..., :self.embed_dim]  # [b,t,s,c]
        k = self.qkv(region_feat)[..., self.embed_dim:2*self.embed_dim]  # [b,t,r,c]
        v = self.qkv(region_feat)[..., 2*self.embed_dim:]  # [b,t,r,c]
        
        # 多头注意力计算
        q = q.view(batch, time, stations, self.num_heads, self.head_dim)
        k = k.view(batch, time, self.region_size, self.num_heads, self.head_dim)
        v = v.view(batch, time, self.region_size, self.num_heads, self.head_dim)
        
        attn = torch.einsum('btsnh,btrnh->btsrh', q, k) / (self.head_dim ** 0.5)
        attn = F.softmax(attn, dim=-2)
        output = torch.einsum('btsrh,btrnh->btsnh', attn, v)
        
        output = output.reshape(batch, time, stations, self.embed_dim)
        return self.out_proj(output)

Causal Temporal MSA (CT-MSA) 采用因果注意力机制处理时间序列:

class CausalTemporalMSA(nn.Module):
    def __init__(self, embed_dim, num_heads, window_size=3):
        super().__init__()
        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.window_size = window_size
        
        self.qkv = nn.Linear(embed_dim, embed_dim * 3)
        self.out_proj = nn.Linear(embed_dim, embed_dim)
        
    def forward(self, x):
        batch, time, stations, _ = x.shape
        
        # 非重叠窗口划分
        x = x.view(batch, time // self.window_size, self.window_size, stations, -1)
        
        # 因果掩码
        mask = torch.tril(torch.ones(self.window_size, self.window_size))
        mask = mask.masked_fill(mask == 0, float('-inf'))
        
        # 计算QKV
        qkv = self.qkv(x)
        q, k, v = qkv.chunk(3, dim=-1)
        
        # 多头注意力计算
        q = q.view(batch, -1, self.num_heads, self.head_dim)
        k = k.view(batch, -1, self.num_heads, self.head_dim)
        v = v.view(batch, -1, self.num_heads, self.head_dim)
        
        attn = torch.einsum('btqnh,btknh->btqkh', q, k) / (self.head_dim ** 0.5)
        attn = attn + mask
        attn = F.softmax(attn, dim=-1)
        output = torch.einsum('btqkh,btknh->btqnh', attn, v)
        
        output = output.reshape(batch, time, stations, self.embed_dim)
        return self.out_proj(output)

2.2 自上而下的随机阶段

这一阶段通过潜在变量捕捉不确定性:

class StochasticStage(nn.Module):
    def __init__(self, embed_dim, latent_dim):
        super().__init__()
        self.embed_dim = embed_dim
        self.latent_dim = latent_dim
        
        # 先验网络
        self.prior_net = nn.Sequential(
            nn.Linear(embed_dim, 256),
            nn.ReLU(),
            nn.Linear(256, latent_dim * 2)  # 输出均值和方差
        )
        
        # 后验网络
        self.posterior_net = nn.Sequential(
            nn.Linear(embed_dim * 2, 256),
            nn.ReLU(),
            nn.Linear(256, latent_dim * 2)
        )
        
    def forward(self, deterministic_state, prev_latent=None):
        batch, time, stations, _ = deterministic_state.shape
        
        if prev_latent is None:
            # 初始时间步使用先验
            mu, logvar = self.prior_net(deterministic_state).chunk(2, dim=-1)
        else:
            # 后续时间步使用后验
            combined = torch.cat([deterministic_state, prev_latent], dim=-1)
            mu, logvar = self.posterior_net(combined).chunk(2, dim=-1)
        
        # 重参数化技巧
        std = torch.exp(0.5 * logvar)
        eps = torch.randn_like(std)
        latent = mu + eps * std
        
        return latent, mu, logvar

3. 完整模型集成

将各组件整合为完整的AirFormer模型:

class AirFormer(nn.Module):
    def __init__(self, num_stations, input_dim, embed_dim=128, num_layers=4):
        super().__init__()
        self.num_stations = num_stations
        self.embed_dim = embed_dim
        
        # 输入嵌入层
        self.input_proj = nn.Linear(input_dim, embed_dim)
        
        # 确定性阶段
        self.deterministic_layers = nn.ModuleList([
            nn.ModuleDict({
                'ds_msa': DartboardMSA(embed_dim, num_heads=8),
                'ct_msa': CausalTemporalMSA(embed_dim, num_heads=8),
                'norm1': nn.LayerNorm(embed_dim),
                'norm2': nn.LayerNorm(embed_dim),
                'ffn': nn.Sequential(
                    nn.Linear(embed_dim, embed_dim * 4),
                    nn.GELU(),
                    nn.Linear(embed_dim * 4, embed_dim)
                )
            }) for _ in range(num_layers)
        ])
        
        # 随机阶段
        self.stochastic_layers = nn.ModuleList([
            StochasticStage(embed_dim, latent_dim=64) for _ in range(num_layers)
        ])
        
        # 预测头
        self.predictor = nn.Sequential(
            nn.Linear(embed_dim, embed_dim * 2),
            nn.ReLU(),
            nn.Linear(embed_dim * 2, 1)  # 预测PM2.5浓度
        )
        
    def forward(self, x, station_map):
        # 输入嵌入
        h = self.input_proj(x)  # [b,t,s,c]
        
        # 确定性阶段
        deterministic_states = []
        for layer in self.deterministic_layers:
            # DS-MSA
            h_res = h
            h = layer['norm1'](h)
            h = layer['ds_msa'](h, station_map) + h_res
            
            # CT-MSA
            h_res = h
            h = layer['norm2'](h)
            h = layer['ct_msa'](h) + h_res
            
            # FFN
            h = layer['ffn'](h) + h
            
            deterministic_states.append(h)
        
        # 随机阶段
        latents = []
        kl_loss = 0
        prev_latent = None
        
        for i, (layer, h) in enumerate(zip(self.stochastic_layers, deterministic_states)):
            latent, mu, logvar = layer(h, prev_latent)
            latents.append(latent)
            
            # 计算KL散度
            if prev_latent is not None:
                kl_loss += -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())
            
            prev_latent = latent
        
        # 预测
        prediction = self.predictor(h + latents[-1])
        
        return prediction.squeeze(-1), kl_loss

4. 训练策略与实验设置

4.1 损失函数设计

AirFormer采用复合损失函数:

def composite_loss(pred, target, kl_loss, kl_weight=0.1):
    # L1损失
    l1_loss = F.l1_loss(pred, target)
    
    # KL散度损失
    total_loss = l1_loss + kl_weight * kl_loss
    
    return total_loss, l1_loss, kl_loss

4.2 训练流程

def train_epoch(model, dataloader, optimizer, device):
    model.train()
    total_loss = 0
    
    for batch in dataloader:
        x, y, station_map = batch
        x, y, station_map = x.to(device), y.to(device), station_map.to(device)
        
        optimizer.zero_grad()
        pred, kl_loss = model(x, station_map)
        loss, l1_loss, kl_loss = composite_loss(pred, y, kl_loss)
        
        loss.backward()
        optimizer.step()
        
        total_loss += loss.item()
    
    return total_loss / len(dataloader)

4.3 关键训练参数

参数名称 推荐值 说明
学习率 1e-4 使用Adam优化器
批量大小 32 根据显存调整
训练轮次 100 早停法防止过拟合
嵌入维度 128 特征维度
层数 4 DS/CT-MSA层数
KL权重 0.1 平衡预测与正则化
窗口大小 3 CT-MSA的时间窗口

5. 模型部署与优化建议

在实际部署AirFormer时,可以考虑以下优化策略:

  1. 动态区域划分:根据气象数据(如风向)动态调整DS-MSA的区域权重
  2. 多任务学习:同时预测多种污染物(PM2.5、O3等)提升模型泛化能力
  3. 边缘计算:将模型部署到监测站边缘设备,减少数据传输延迟
  4. 持续学习:定期用新数据微调模型,适应环境变化
# 模型保存与加载
torch.save({
    'model_state_dict': model.state_dict(),
    'optimizer_state_dict': optimizer.state_dict(),
}, 'airformer_model.pth')

checkpoint = torch.load('airformer_model.pth')
model.load_state_dict(checkpoint['model_state_dict'])

注意:实际部署时建议使用TorchScript将模型转换为脚本模式,提高推理效率并支持跨平台部署。

在完成模型训练后,可以通过可视化工具分析注意力权重,验证模型是否学习到了有意义的空间关联模式。例如,可以绘制特定站点在不同风向条件下的注意力分布,验证DS-MSA是否合理关注了上风向区域。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐