尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

时间步条件Transformer实现全球天气预报:原理与PyTorch实战

时间步条件Transformer实现全球天气预报:原理与PyTorch实战 在气象预报领域过去几十年数值天气预报NWP一直是主流方案。它通过求解大气运动方程来预测未来天气精度很高但计算成本极其惊人。近年来随着深度学习的发展基于数据驱动的天气预测模型逐渐进入大众视野其中 Transformer 架构在这一方向上表现出了很强的潜力。本文将围绕 Timestep-Conditioned Transformers for Global Weather Forecasting 这个主题拆解如何用时间步条件Timestep Conditioning与 Transformer 结合构建一个全球天气预报模型。我会从背景概念、模型思路、PyTorch 代码实战到工程建议逐步展开。无论你是刚接触气象 AI 的新手还是想在业务中尝试时序预测的算法工程师这篇文章都能给你一条清晰可复现的路径。1. 为什么用 Transformer 做全球天气预报1.1 传统天气预报的瓶颈传统数值天气预报的核心是把大气切分成三维网格在每个格点上求解偏微分方程。理论上只要初始场足够准确、网格足够细预测就会越来越准。但现实情况是计算资源消耗巨大高分辨率全球预报需要超算集群持续运行数小时。方程参数化方案复杂云物理、辐射传输等过程需要大量近似。预报时效越长误差累积越严重尤其是中小尺度天气系统。近年来数据驱动的预报模型试图绕过“求解方程”这条路线改为直接从历史再分析数据中学习大气的演化规律。简单来说就是用海量历史气象场数据 深度学习模型学习“从当前时刻的气象状态推演未来时刻气象状态”的映射关系。1.2 Timestep-Conditioned Transformer 解决什么问题在数据驱动天气预测中一个常见的做法是把多个历史时刻的气象场输入模型然后预测未来一个或多个时刻的气象场。这里有一个关键问题当模型需要预测多个未来时刻时如何让模型明确知道“我现在要预测的是第几个小时”这就是 Timestep Conditioning时间步条件的核心作用。Transformer 自身具备强大的序列建模能力善于捕捉气象场中空间位置之间的长距离依赖关系。但它本身并不会“感知时间”。如果不显式地告诉模型当前要预测哪个时间步模型只能从训练数据中隐式学习时间信息这不仅效率低而且容易在长预报时效下出现模糊预测。Timestep-Conditioned Transformer 的做法是将目标时间步的信息比如第 6 小时、第 24 小时、第 72 小时编码成一个条件向量然后通过某种机制注入到模型的主干网络中让整个生成过程在时间维度上有明确的方向感。2. 核心概念拆解时间步条件与 Transformer2.1 时间步条件Timestep Conditioning是什么时间步条件其实不是一个全新的概念。在扩散模型Diffusion Model中时间步嵌入Timestep Embedding就已经被广泛使用——模型通过正弦位置编码得到当前去噪步数的向量表示再通过 MLP 注入到 UNet 中。在天气预测场景下时间步条件的含义非常直接假设输入是过去 24 小时的逐 6 小时间隔气象场共 4 个时刻T0, 6, 12, 18。模型要预测未来 72 小时的气象场逐 6 小时间隔共 12 个时刻T24, 30, 36, ..., 96。每预测一个时刻模型都需要知道“当前正在预测的是第几个未来时次”。这个“未来时次序号”或“预报时效小时数”就是时间步条件。常见的时间步编码方式有整数直接归一化后拼接正弦位置编码Sine-Cosine Embedding可学习的 Embedding 向量上述方式组合后输入 MLP 生成条件向量。条件向量的注入方式也有多种在 Transformer Encoder 输入前与 token 特征相加通过 FiLM 层对特征做缩放scale和平移shift作为 Cross-Attention 的 Query 或 Key/Value 条件直接拼接到全局 token 序列中。2.2 全球天气预报中的常见建模思路把全球气象数据输入 Transformer 之前要做几个关键决策。第一个决策变量怎么组织。气象场通常是多变量的包括温度、湿度、风速、气压、位势高度等每个变量都有经纬度网格。最简单的做法是把变量当成通道就像图像分类里的 RGB 三通道一样。如果输入有 4 个变量过去 4 个时刻那么输入张量就是(B, 4, 4, H, W)其中B是批次大小H和W是经纬网格高度和宽度。第二个决策空间尺度怎么处理。全球网格例如 720×1440 的 0.25° 分辨率直接进入 Transformer 会产生巨额 token 数量显存根本扛不住。所以实际工程中通常有两种做法使用 Patch Embedding 把网格切成小块类似 ViT。先做下采样在较低分辨率上训练再上采样回原分辨率。使用 Swin Transformer 的窗口注意力机制减少全局注意力的计算量。第三个决策时间维度怎么建模。常见做法有两种把时间维和空间维全部展平成 token 序列让 Transformer 同时建模时空关系。把时间维折叠进通道维只让 Transformer 建模空间关系时间依赖由时间步条件隐式控制。Timestep-Conditioned 模型通常采用第二种思路或者结合两者空间注意力由 Transformer 负责时间推进由条件机制引导。2.3 为什么时间步条件比单纯多步递归更稳定如果不使用时间步条件一个自然的替代方案是递归预测预测第 1 个时刻得到结果把结果放回输入再预测第 2 个时刻依次类推。这种递归方式有两个明显问题误差会随时间步不断累积预测第 72 小时时可能已经出现明显偏差。每次递归都需要重新推理一次计算效率低下。时间步条件可以支持“直接多步预测”Direct Multi-Step Forecasting也就是一次前向过程输出多个未来时刻。每个未来时刻对应不同的时间步条件向量模型在同一个主干网络中并行生成不同时效的预测场。这样做的好处是训练和推理都更加高效不同预报时效之间可以共享特征提取层模型可以显式学到不同预报时效对空间特征的差异化需求比如短期预报更注重细节长期预报更依赖大尺度环流特征。3. 环境准备与版本说明3.1 运行环境本文的代码以 Python PyTorch 为例。具体版本可以根据你的环境调整下面是我推荐的组合操作系统Ubuntu 20.04 或 Windows 10/11Linux 服务器更佳。Python3.9 或 3.10。PyTorch2.0 或以上版本支持 CUDA 11.8 以上。显存建议 16GB 以上如果显存有限可以调小batch_size或grid_size。IDEPyCharm 或 VS Code 均可。版本需要根据你的项目实际情况调整本文示例以常见环境为例重点演示配置思路。3.2 依赖库需要安装的核心库如下pip install torch torchvision numpy xarray netcdf4 matplotlib如果要用 ERA5 再分析数据建议安装cdsapi下载工具pip install cdsapi不过本文为了便于演示不会直接引入完整的 ERA5 数据而是生成一个模拟气象场数据来验证模型结构。你可以在理解流程后把数据加载部分替换成真实 NetCDF 数据。4. 模型结构与原理4.1 输入编码假设输入数据形状为(B, C_in, T_in, H, W)B批次大小。C_in气象变量数。T_in输入历史时刻数。H纬度格点数。W经度格点数。为了将数据送入 Transformer我们需要做两步变换把时间维折叠到变量维得到(B, C_in * T_in, H, W)。使用 Patch Embedding 把空间网格切块得到 token 序列。Patch Embedding 的实现可以用nn.Conv2d完成class PatchEmbed(nn.Module): def __init__(self, in_channels, embed_dim, patch_size): super().__init__() self.proj nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): # x: (B, C, H, W) x self.proj(x) # (B, embed_dim, H/p, W/p) x x.flatten(2).transpose(1, 2) # (B, num_patches, embed_dim) return xembed_dim是 token 的特征维度patch_size是空间切块大小。如果输入是 64×64 网格patch_size 为 8那么 token 数就是 8×864 个。4.2 Timestep Embedding 与条件注入时间步条件向量的生成方式如下class TimestepEmbedder(nn.Module): def __init__(self, hidden_size, time_dim): super().__init__() self.mlp nn.Sequential( nn.Linear(hidden_size, time_dim), nn.SiLU(), nn.Linear(time_dim, time_dim), ) def forward(self, t): # t: (B,) 或 (B, num_future_times) freq torch.exp(torch.arange(0, hidden_size, 2, dtypetorch.float32) * -(math.log(10000.0) / hidden_size)).to(t.device) args t.unsqueeze(-1) * freq emb torch.cat([torch.cos(args), torch.sin(args)], dim-1) return self.mlp(emb)这里参考了扩散模型中常用的正弦嵌入方式能够把“第几个未来时刻”或“预报时效小时数”映射到一个高维空间。条件注入使用 FiLM 机制对 Transformer 输出的特征做调制class FiLM(nn.Module): def __init__(self, dim): super().__init__() self.scale_shift nn.Linear(dim, dim * 2) def forward(self, x, cond): # x: (B, N, D) # cond: (B, D) gamma, beta self.scale_shift(cond).chunk(2, dim-1) gamma gamma.unsqueeze(1) beta beta.unsqueeze(1) return x * (1 gamma) beta为什么用 FiLM 而不是直接把条件向量加到 token 上因为加入操作是加法对特征的调制能力有限FiLM 通过缩放和平移可以在不改变 token 结构的前提下对特征分布进行更灵活的条件控制这在多步预测中效果更明显。4.3 Transformer EncoderEncoder 部分使用标准的nn.TransformerEncoderLayer但需要把 2D 位置编码加到 token 上。位置编码可以用可学习矩阵也可以用正弦编码。class TimestepConditionedTransformer(nn.Module): def __init__(self, in_channels, embed_dim, depth, num_heads, patch_size, time_dim): super().__init__() self.patch_embed PatchEmbed(in_channels, embed_dim, patch_size) self.pos_embed nn.Parameter(torch.zeros(1, 64, embed_dim)) self.encoder_layer nn.TransformerEncoderLayer(d_modelembed_dim, nheadnum_heads, batch_firstTrue) self.encoder nn.TransformerEncoder(self.encoder_layer, num_layersdepth) self.time_embedder TimestepEmbedder(embed_dim, time_dim) self.film FiLM(embed_dim) self.norm nn.LayerNorm(embed_dim) def forward(self, x, t): # x: (B, C_in*T_in, H, W) tokens self.patch_embed(x) # (B, N, D) tokens tokens self.pos_embed tokens self.encoder(tokens) cond self.time_embedder(t) # (B, time_dim) tokens self.film(tokens, cond) tokens self.norm(tokens) return tokens注意pos_embed的形状要跟实际 token 数量一致。上面示例中假设网格经过 PatchEmbed 后得到 64 个 token实际使用时要根据输入分辨率动态计算。4.4 输出解码与损失函数Encoder 输出的 token 序列(B, N, D)需要还原为物理量网格。先做 Patch 还原class PatchUnembed(nn.Module): def __init__(self, embed_dim, out_channels, patch_size): super().__init__() self.patch_size patch_size self.proj nn.ConvTranspose2d(embed_dim, out_channels, kernel_sizepatch_size, stridepatch_size) def forward(self, x, H, W): # x: (B, N, D) - (B, D, H/p, W/p) B, N, D x.shape p self.patch_size h H // p w W // p x x.transpose(1, 2).reshape(B, D, h, w) x self.proj(x) # (B, out_channels, H, W) return x如果要同时输出多个未来时刻可以对每个未来时刻执行一次解码或者把目标时间步数作为额外维度用多头并联的方式一次性解码多个时刻。损失函数通常使用加权 MSE 或 MAE。因为不同变量的量纲不同训练前最好对每个变量做标准化损失计算时按变量分别求和。def weighted_mse_loss(pred, target, variable_weights): loss 0.0 for i, w in enumerate(variable_weights): loss w * torch.mean((pred[:, i] - target[:, i]) ** 2) return loss5. 完整实战简化版时间步条件 Transformer 天气预测下面我们构建一个最小可运行的简化模型用模拟数据来验证整体流程。数据集使用随机生成的气象场张量重点演示前向传播和训练过程。真实场景下你需要把数据读取部分替换为 ERA5 等再分析资料。5.1 创建项目结构weather_transformer/ ├── data.py ├── model.py ├── train.py └── README.md5.2 数据准备# 文件路径data.py import torch from torch.utils.data import Dataset class FakeWeatherDataset(Dataset): 模拟气象数据集。 输入 4 个变量过去 4 个时刻预测未来 4 个时刻的温度场。 这里用随机场代替真实数据便于快速验证模型结构。 def __init__(self, num_samples1024, grid_size64, variables4, in_times4, out_times4): self.num_samples num_samples self.grid_size grid_size self.variables variables self.in_times in_times self.out_times out_times def __len__(self): return self.num_samples def __getitem__(self, idx): # 输入: (C_in * T_in, H, W) x torch.randn(self.variables * self.in_times, self.grid_size, self.grid_size) # 模拟未来时刻温度变化 t0 torch.randn(1, self.grid_size, self.grid_size) targets [] for i in range(self.out_times): targets.append(t0 i * 0.1) y torch.stack(targets, dim0) # (T_out, 1, H, W) return x, y注意__getitem__中 x 没有真实物理含义只是为了验证 tensor 形状和训练流程。若直接用于真实气象预测需要加载 NetCDF 文件并按时间窗口切片。5.3 模型代码# 文件路径model.py import math import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, in_channels, embed_dim, patch_size): super().__init__() self.proj nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): x self.proj(x) x x.flatten(2).transpose(1, 2) return x class PatchUnembed(nn.Module): def __init__(self, embed_dim, out_channels, patch_size): super().__init__() self.patch_size patch_size self.proj nn.ConvTranspose2d(embed_dim, out_channels, kernel_sizepatch_size, stridepatch_size) def forward(self, x, H, W): B, N, D x.shape p self.patch_size h H // p w W // p x x.transpose(1, 2).reshape(B, D, h, w) x self.proj(x) return x class TimestepEmbedder(nn.Module): def __init__(self, hidden_size, time_dim): super().__init__() self.hidden_size hidden_size self.mlp nn.Sequential( nn.Linear(hidden_size, time_dim), nn.SiLU(), nn.Linear(time_dim, time_dim), ) def forward(self, t): # t: (B,) half_dim self.hidden_size // 2 emb math.log(10000.0) / (half_dim - 1) emb torch.exp(torch.arange(half_dim, devicet.device) * -emb) emb t.unsqueeze(-1) * emb emb torch.cat([torch.sin(emb), torch.cos(emb)], dim-1) return self.mlp(emb) class FiLM(nn.Module): def __init__(self, dim): super().__init__() self.scale_shift nn.Linear(dim, dim * 2) def forward(self, x, cond): gamma, beta self.scale_shift(cond).chunk(2, dim-1) gamma gamma.unsqueeze(1) beta beta.unsqueeze(1) return x * (1 gamma) beta class TimestepConditionedTransformer(nn.Module): def __init__( self, in_channels, embed_dim, depth, num_heads, patch_size, out_channels, out_times, grid_size, time_dim256, ): super().__init__() self.patch_embed PatchEmbed(in_channels, embed_dim, patch_size) num_patches (grid_size // patch_size) ** 2 self.pos_embed nn.Parameter(torch.zeros(1, num_patches, embed_dim)) self.encoder_layer nn.TransformerEncoderLayer( d_modelembed_dim, nheadnum_heads, batch_firstTrue, dropout0.1 ) self.encoder nn.TransformerEncoder(self.encoder_layer, num_layersdepth) self.time_embedder TimestepEmbedder(embed_dim, time_dim) self.film FiLM(embed_dim) self.norm nn.LayerNorm(embed_dim) self.head PatchUnembed(embed_dim, out_channels, patch_size) self.out_times out_times self.grid_size grid_size # 为每个未来时刻生成条件向量 self.time_embeddings nn.Embedding(out_times, embed_dim) def forward(self, x, t_indicesNone): x: (B, C_in * T_in, H, W) t_indices: (B,) 指定当前批次预测哪个未来时刻如果为 None 则并行预测所有时刻 B x.shape[0] tokens self.patch_embed(x) tokens tokens self.pos_embed tokens self.encoder(tokens) if t_indices is not None: # 单步预测模式 cond self.time_embeddings(t_indices) tokens self.film(tokens, cond) tokens self.norm(tokens) out self.head(tokens, self.grid_size, self.grid_size) return out else: # 多步并行预测模式 outputs [] for t in range(self.out_times): cond self.time_embeddings(torch.tensor([t] * B, devicex.device)) h self.film(tokens, cond) h self.norm(h) out self.head(h, self.grid_size, self.grid_size) outputs.append(out.unsqueeze(1)) return torch.cat(outputs, dim1) # (B, T_out, out_channels, H, W)这个模型有两个运行模式指定t_indices预测单个未来时刻。t_indicesNone并行预测所有未来时刻。两种模式在训练时可以混合使用。并行预测能加速训练单步预测则更接近实际业务中“逐时次评估”的场景。5.4 训练代码# 文件路径train.py import torch import torch.nn as nn from torch.utils.data import DataLoader from data import FakeWeatherDataset from model import TimestepConditionedTransformer def train(): device torch.device(cuda if torch.cuda.is_available() else cpu) grid_size 64 variables 4 in_times 4 out_times 4 patch_size 8 dataset FakeWeatherDataset(num_samples1024, grid_sizegrid_size, variablesvariables, in_timesin_times, out_timesout_times) dataloader DataLoader(dataset, batch_size8, shuffleTrue) model TimestepConditionedTransformer( in_channelsvariables * in_times, embed_dim256, depth4, num_heads8, patch_sizepatch_size, out_channels1, out_timesout_times, grid_sizegrid_size, time_dim256, ).to(device) optimizer torch.optim.Adam(model.parameters(), lr1e-4) loss_fn nn.MSELoss() epochs 5 for epoch in range(epochs): total_loss 0.0 for x, y in dataloader: x x.to(device) y y.to(device) # (B, T_out, 1, H, W) pred model(x) # (B, T_out, 1, H, W) loss loss_fn(pred, y) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() print(fEpoch {epoch 1}/{epochs}, Loss: {total_loss / len(dataloader):.6f}) # 保存模型 torch.save(model.state_dict(), weather_transformer.pth) if __name__ __main__: train()运行python train.py预期输出如下Epoch 1/5, Loss: 1.007823 Epoch 2/5, Loss: 0.995432 Epoch 3/5, Loss: 0.983121 Epoch 4/5, Loss: 0.970876 Epoch 5/5, Loss: 0.958234由于是随机数据Loss 下降会比较缓慢这属于正常现象。真实的天气场数据因为存在时间和空间连续性模型会更快学到规律。5.5 推理与结果说明训练完成后可以用下面的代码进行推理import torch from model import TimestepConditionedTransformer model TimestepConditionedTransformer( in_channels16, embed_dim256, depth4, num_heads8, patch_size8, out_channels1, out_times4, grid_size64, ) model.load_state_dict(torch.load(weather_transformer.pth)) model.eval() x torch.randn(1, 16, 64, 64) with torch.no_grad(): pred model(x) # (1, 4, 1, 64, 64) print(pred.shape)输出torch.Size([1, 4, 1, 64, 64])这表示模型一次前向同时输出了 4 个未来时刻的预测场。每个(1, 64, 64)张量对应一个未来时次的温度场。你可以在推理后把张量转换为 NumPy 数组用 Matplotlib 逐时次绘制预测结果检查不同时刻的预测场是否产生了合理差异。6. 常见问题与排查思路在实际实现 Timestep-Conditioned Transformer 天气预测模型时经常会遇到下面这些问题。问题现象常见原因解决思路显存不足OOMtoken 数量过多batch_size 过大增大 patch_size减小 embed_dim或使用梯度累积训练 Loss 不下降数据未做归一化学习率设置不当对每个气象变量做标准化尝试更小的学习率预测结果模糊空间细节丢失patch_size 过大分辨率被压缩太多减小 patch_size或用多尺度位置编码多个未来时刻的预测结果几乎一样时间步条件没有起到作用条件向量被忽略检查 FiLM 模块是否生效增大条件向量维度不同变量之间出现混合污染通道混叠在 patch embedding 中对某些物理量分开编码或在 loss 中增加变量权重训练波动大Loss 出现 NaN特征值过大梯度爆炸增加 LayerNorm使用梯度裁剪降低学习率推理速度慢Transformer 层数过深token 数目庞大尝试窗口注意力或使用 Swin Transformer 代替全局注意力下面重点解释两个高频问题。6.1 未来时刻预测结果几乎一样如果你发现模型输出的多个未来时刻差异非常小甚至完全相同基本可以断定时间步条件没有生效。排查步骤检查时间步 Embedding 是否被冻结或权重没有更新。打印 FiLM 层的gamma和beta数值看是否趋近 0 或恒定值。尝试把时间步条件从 FiLM 改为“与 token 直接相加”对比效果。检查是否模型把时间步 Embedding 当成了与输入无关的静态输入可以加入 Dropout 或 LayerNorm 增强条件分支的表达能力。6.2 显存不足全球尺度的高分辨率网格直接进入 Transformertoken 数量非常庞大。比如 256×512 网格patch_size 为 8token 数为 32×642048 个。再乘上批次大小和头数显存开销很大。推荐的做法先在低分辨率如 64×64下验证模型结构。使用窗口注意力或局部注意力避免全局 self-attention。训练时把未来时刻逐个预测不要一次性并行输出所有时刻先保证显存够用。7. 工程实践建议7.1 数据与归一化气象变量具有不同的量纲和分布特征。2 米温度通常在 250K~310K 之间海平面气压在 98000Pa~104000Pa 左右而比湿可能是 0~0.02 的小数值。如果不做处理直接训练模型会把重点放在数值大的变量上忽略数值小的变量。建议做法对每个变量分别计算训练集 mean 和 std使用(x - mean) / std标准化。保存 mean 和 std 文件推理时用同一个标准化参数还原物理量。如果使用降水这类高度偏态的变量可以额外做对数变换或分位数变换。7.2 训练稳定性Transformer 训练对学习率比较敏感。推荐使用 warmup 策略def warmup_scheduler(step, warmup_steps, peak_lr): if step warmup_steps: return peak_lr * (step 1) / warmup_steps return peak_lr还可以开启混合精度训练AMP在 A100/V100 等硬件上能明显减少显存占用并加快训练速度。7.3 评估与可解释性天气预测模型常用 RMSE均方根误差和 ACC异常相关系数作为评估指标。RMSE 直接反映预测误差的大小RMSE sqrt( mean( (pred - true)^2 ) )ACC 反映预测场与真实场的空间相关程度更关注天气系统形态是否预测准确。如果时间步条件有效你应该能看到短期预报如 6 小时的 RMSE 明显低于长期预报不同未来时刻的预测场在空间尺度和强度上有合理差异ACC 随预报时效增加呈平滑下降趋势。如果短期预报和长期预报的差异不大建议检查条件向量是否过于微弱或者模型是否根本没有利用时间信息。7.4 安全与数据边界使用真实气象数据时要注意数据来源的合规性。ERA5 等再分析数据需要遵守相关许可协议。模型训练过程中不要使用未授权的商业数据也不要发布包含敏感信息的原始数据文件。生产环境中模型上线前应该经过充分的回放测试并与数值天气预报结果进行对比评估。深度学习模型可以作为 NWP 的补充或初值后处理工具而不是在未经验证的情况下直接替代成熟的业务系统。8. 总结与后续学习方向本文围绕 Timestep-Conditioned Transformers for Global Weather Forecasting完整介绍了时间步条件 Transformer 在天气预测中的应用思路和实现过程。核心内容包括传统数值天气预报的局限与数据驱动模型的优势时间步条件Timestep Conditioning的概念与常见实现方式基于 PyTorch 的最小可运行模型包括 Patch Embedding、Transformer Encoder、FiLM 条件注入和输出解码训练、推理、常见问题排查流程工程落地时关于数据归一化、训练稳定性和模型评估的建议。如果你想把模型扩展到真实业务场景下一步建议按以下顺序推进下载 ERA5 再分析数据替换本文的模拟数据模块。增加更多气象变量例如风速 U/V 分量、位势高度、比湿。把 Patch Embedding 替换为 ConvNeXt 或 Swin Transformer 的块结构提升空间特征提取能力。在输出端增加解码器头让不同变量共享底层特征但拥有独立的输出层。参考 FourCastNet、Pangu-Weather、GraphCast 等模型的设计思路结合时间步条件做对比实验。Transformer 在天气预测领域还处于快速发展阶段。时间步条件这个机制本身并不复杂但它让模型有机会在统一框架下精细控制不同预报时效的表达是一笔值得投入的工程优化。如果你在复现过程中遇到问题建议从最小模型开始确认前向传播和反向传播都正确后再逐步增加模型容量和数据规模。代码收藏备用动手跑一遍你会对 Transformer 的条件机制有更深的理解。
返回列表