用PyTorch复现AirFormer:手把手教你搭建空气质量预测Transformer(附完整代码)
·
用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时,可以考虑以下优化策略:
- 动态区域划分:根据气象数据(如风向)动态调整DS-MSA的区域权重
- 多任务学习:同时预测多种污染物(PM2.5、O3等)提升模型泛化能力
- 边缘计算:将模型部署到监测站边缘设备,减少数据传输延迟
- 持续学习:定期用新数据微调模型,适应环境变化
# 模型保存与加载
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是否合理关注了上风向区域。
更多推荐


所有评论(0)