1. 项目概述:当神经微分方程遇见天文时间序列
天文观测数据可能是最典型的"不守规矩"时间序列——望远镜受地球自转限制导致采样间隔不规则,云层干扰造成数据缺失,不同波段观测设备产生异步时间戳。传统RNN/LSTM在等间隔插值过程中会损失关键特征,而今年ICLR最佳论文提名的随机延迟神经微分方程(SDDE)恰好能建模这种具有随机延迟和缺失的序列。我在处理凌星系外行星巡天卫星(TESS)光变曲线时,发现SDDE对掩星事件预测的RMSE比Transformer低37%。
2. 核心技术拆解:SDDE如何驯服不规则序列
2.1 微分方程框架设计
核心采用Itô随机微分方程形式:
dX_t = f(X_t,X_{t-τ(t)},t)dt + g(X_t,t)dW_t其中τ(t)是随机延迟过程,我们使用Gamma过程模拟天文观测间隔:
class GammaDelay(nn.Module): def __init__(self, shape=2.0, scale=1.0): self.dist = torch.distributions.Gamma(shape, scale) def forward(self, t): return self.dist.sample() * t # 延迟时间与当前时间正相关2.2 非参数化时间编码
为避免人工设计时间特征,采用神经过程编码器:
class TimeEncoder(nn.Module): def __init__(self, hidden_dim=64): self.rff = RandomFourierFeatures(hidden_dim) # 随机傅里叶特征 def forward(self, delta_t): return self.rff(torch.log1p(delta_t)) # 对数变换处理长尾分布2.3 记忆压缩与回溯
设计可微分的历史缓存模块解决长程依赖:
class MemoryBank(nn.Module): def __init__(self, capacity=1000): self.queue = DifferentiablePriorityQueue(capacity) def update(self, t, x): self.queue.insert(t, x) # 按时间戳排序存储 def query(self, t): return self.queue.range_query(t-τ, t) # 支持子序列微分3. 天文场景下的特殊处理
3.1 多尺度特征提取
针对行星凌日信号(小时级)和恒星活动(天级)的混合周期:
class MultiScaleBlock(nn.Module): def __init__(self): self.conv1d = nn.ModuleList([ nn.Conv1d(1, 16, kernel_size=24, stride=4), # 短周期 nn.Conv1d(1, 16, kernel_size=720, stride=60) # 长周期 ]) def forward(self, x): return torch.cat([conv(x) for conv in self.conv1d], dim=1)3.2 不确定性量化
使用贝叶斯神经网络输出预测区间:
class BayesianOutput(nn.Module): def __init__(self, hidden_dim): self.mu = nn.Linear(hidden_dim, 1) self.logvar = nn.Linear(hidden_dim, 1) def forward(self, x): return self.mu(x), torch.exp(self.logvar(x)) # 高斯分布参数4. 实战:TESS光变曲线预测
4.1 数据预处理流程
- 原始FITS文件解析
from astropy.io import fits hdul = fits.open("tess2020009025919-s0001-0000000269555394-0131-s_lc.fits") flux = hdul[1].data['PDCSAP_FLUX'] time = hdul[1].data['TIME']- 异常值处理(宇宙射线干扰)
mad = torch.median(torch.abs(flux - torch.median(flux))) valid_mask = (flux < 3*mad) & (flux > -3*mad)- 非均匀采样对齐
def irregular_resample(t, x, new_t): kernel = EpanechnikovKernel(bandwidth=0.1) weights = kernel(t - new_t.unsqueeze(1)) return (weights * x).sum(dim=1) / weights.sum(dim=1)4.2 训练技巧
- 采用课程学习策略:先训练规则采样子序列,逐步引入随机缺失
- 损失函数组合:
loss = nn.GaussianNLLLoss()(pred, target, var) + 0.1*kl_divergence- 学习率调度:
scheduler = torch.optim.lr_scheduler.OneCycleLR( optimizer, max_lr=1e-3, steps_per_epoch=100, epochs=50)5. 性能对比与误差分析
| 模型 | RMSE | MAE | 推理速度(样本/秒) |
|---|---|---|---|
| LSTM+插值 | 0.142 | 0.098 | 1200 |
| Transformer | 0.118 | 0.085 | 800 |
| Neural ODE | 0.105 | 0.072 | 350 |
| 本方案(SDDE) | 0.074 | 0.051 | 280 |
典型误差案例:
- 恒星耀斑爆发时刻预测偏差较大(延迟5-10分钟)
- 双星系统掩星事件有时会出现双峰误判
6. 工程落地优化
6.1 内存效率提升
- 历史缓存的分块加载:
class ChunkedMemory: def __init__(self, chunk_size=100): self.chunks = [torch.empty(chunk_size)] def insert(self, t, x): if len(self.chunks[-1]) >= chunk_size: self.chunks.append(torch.empty(chunk_size)) self.chunks[-1][-1] = (t, x)6.2 实时推理加速
- 使用TorchScript导出模型:
traced_model = torch.jit.script(model) traced_model.save("sdde_astronomy.pt")- 选择性历史回溯:
def get_relevant_history(t, τ): # 只加载[t-2τ, t]时间窗的数据 return memory.query(t-2*τ, t)在天文台部署时发现,当处理亚秒级采样数据时,需要特别注意GPU显存管理。我们最终采用流式处理+显存预分配的方案,使得RTX 3090上能稳定处理1kHz采样率的太阳射电爆发数据。