在Google Colab上跑通Graph WaveNet:从环境配置到成功训练(Python 3.6 + PyTorch 1.10.2避坑指南)
在Google Colab上跑通Graph WaveNet:从环境配置到成功训练(Python 3.6 + PyTorch 1.10.2避坑指南)
当你在Google Colab上第一次尝试运行Graph WaveNet时,可能会遇到各种令人沮丧的错误。最常见的就是那个让人摸不着头脑的Expected 2D (unbatched) or 3D (batched) input to conv1d, but got input of size: [64, 32, 207, 13]。别担心,这不是你的问题,而是环境配置的陷阱。本文将带你一步步避开这些坑,从零开始搭建正确的运行环境,直到成功训练模型。
1. 为什么Python 3.6和PyTorch 1.10.2是关键
Graph WaveNet的代码对PyTorch版本极其敏感。最新版本的PyTorch往往会导致各种维度不匹配的错误,这就是为什么我们需要特定的环境组合:
- Python 3.6:这个版本与PyTorch 1.10.2兼容性最佳
- PyTorch 1.10.2:修复了早期版本的bug,同时保持了Graph WaveNet所需的张量操作行为
# 验证当前Python版本
!python --version
# 应该显示Python 3.6.x
如果你看到的是更高版本,就需要降级。在Colab上降级Python不像本地环境那么简单,因为Colab默认使用系统Python。
2. 在Colab上配置Python 3.6环境
Colab默认使用较新的Python版本,我们需要通过Miniconda来安装和管理Python 3.6:
%%bash
MINICONDA_INSTALLER_SCRIPT=Miniconda3-4.5.4-Linux-x86_64.sh
MINICONDA_PREFIX=/usr/local
wget https://repo.continuum.io/miniconda/$MINICONDA_INSTALLER_SCRIPT
chmod +x $MINICONDA_INSTALLER_SCRIPT
./$MINICONDA_INSTALLER_SCRIPT -b -f -p $MINICONDA_PREFIX
安装完成后,设置Python 3.6环境:
%%bash
conda install --channel defaults conda python=3.6 --yes
conda update --channel defaults --all --yes
验证安装:
!python --version
# 应该显示Python 3.6.13
3. 安装PyTorch 1.10.2及其他依赖
现在可以安装特定版本的PyTorch了:
!pip install torch==1.10.2
同时安装其他必要的依赖:
!pip install -q ipykernel numpy pandas argparse matplotlib seaborn tables scipy
验证PyTorch版本:
import torch
print(torch.__version__) # 应该显示1.10.2
4. 准备Graph WaveNet代码和数据
从GitHub克隆代码库:
!git clone https://github.com/nnzhan/Graph-WaveNet.git
%cd Graph-WaveNet
Graph WaveNet需要DCRNN项目中的交通数据。我们需要额外克隆这个仓库:
!git clone https://github.com/liyaguang/DCRNN.git
数据准备步骤:
# 生成训练数据
!python generate_training_data.py --output_dir=data/METR-LA --traffic_df_filename=../DCRNN/data/metr-la.h5
5. 解决常见的张量维度错误
当你第一次运行训练脚本时,很可能会遇到维度不匹配的错误。这是因为PyTorch不同版本对卷积操作的处理方式不同。
典型的错误信息:
RuntimeError: Expected 2D (unbatched) or 3D (batched) input to conv1d, but got input of size: [64, 32, 207, 13]
解决方案:
- 确保你使用的是PyTorch 1.10.2
- 检查数据加载器的输出维度
- 必要时调整模型的输入处理代码
在train.py中,找到数据加载部分,确保张量转置正确:
trainx = torch.Tensor(x).to(device)
trainx = trainx.transpose(1, 3) # 关键步骤:调整维度顺序
6. 完整训练命令与参数解析
正确的训练命令应该包含所有必要的参数:
!python train.py \
--adjdata '../DCRNN/data/sensor_graph/adj_mx.pkl' \
--device 'cuda:0' \
--data 'data/METR-LA' \
--adjtype 'doubletransition' \
--gcn_bool \
--addaptadj \
--seq_length 12 \
--nhid 32 \
--batch_size 64 \
--learning_rate 0.001 \
--dropout 0.3 \
--weight_decay 0.0001 \
--epochs 100
关键参数说明:
| 参数 | 说明 | 推荐值 |
|---|---|---|
--adjdata |
邻接矩阵文件路径 | '../DCRNN/data/sensor_graph/adj_mx.pkl' |
--gcn_bool |
是否使用图卷积层 | 建议启用 |
--addaptadj |
是否添加自适应邻接矩阵 | 建议启用 |
--seq_length |
输入序列长度 | 12 |
--nhid |
隐藏层维度 | 32 |
--batch_size |
批量大小 | 64 |
7. 训练过程监控与问题排查
训练开始后,你应该看到类似下面的输出:
Iter: 000, Train Loss: 1.3524, Train MAPE: 8.7243, Train RMSE: 2.5123
Epoch: 001, Inference Time: 0.8421 secs
Epoch: 001, Train Loss: 0.8732, Train MAPE: 5.6421, Train RMSE: 1.8732, Valid Loss: 0.7921, Valid MAPE: 5.1234, Valid RMSE: 1.7321, Training Time: 45.32/epoch
常见问题及解决方案:
-
CUDA内存不足:
- 减小
batch_size - 使用
--device 'cpu'暂时切换到CPU模式测试
- 减小
-
NaN损失值:
- 降低学习率(
--learning_rate) - 增加
--weight_decay值
- 降低学习率(
-
训练不收敛:
- 检查数据标准化是否正确
- 尝试不同的学习率调度策略
8. 模型评估与结果解读
训练完成后,模型会自动在测试集上评估。典型的输出如下:
Evaluate best model on test data for horizon 1, Test MAE: 1.3524, Test MAPE: 3.7243, Test RMSE: 2.5123
...
On average over 12 horizons, Test MAE: 1.8732, Test MAPE: 4.6421, Test RMSE: 2.8732
指标解释:
- MAE (Mean Absolute Error):平均绝对误差,越小越好
- MAPE (Mean Absolute Percentage Error):平均绝对百分比误差,小于5%通常算不错
- RMSE (Root Mean Square Error):均方根误差,对大误差更敏感
9. 高级技巧与优化建议
一旦基础版本能正常运行,你可以尝试以下优化:
-
自适应学习率调整: 在训练循环中添加学习率衰减:
if i % 10 == 0: lr = max(0.000002, args.learning_rate * (0.1 ** (i // 10))) for g in engine.optimizer.param_groups: g['lr'] = lr -
梯度裁剪: 修改
engine.py中的clip值,防止梯度爆炸:def __init__(self, ...): self.clip = 3 # 原为5,可以尝试更小的值 -
模型架构调整:
- 尝试不同的
--nhid值(如64) - 调整
--seq_length以适应更长的时间依赖
- 尝试不同的
10. 保存与部署训练好的模型
训练完成后,最佳模型会自动保存。你也可以手动保存特定epoch的模型:
torch.save(engine.model.state_dict(), 'custom_model_name.pth')
加载模型进行预测:
model.load_state_dict(torch.load('custom_model_name.pth'))
model.eval()
with torch.no_grad():
predictions = model(test_data)
记住,部署时也需要保持相同的Python和PyTorch版本环境,以避免兼容性问题。
更多推荐


所有评论(0)