在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]

解决方案

  1. 确保你使用的是PyTorch 1.10.2
  2. 检查数据加载器的输出维度
  3. 必要时调整模型的输入处理代码

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

常见问题及解决方案

  1. CUDA内存不足

    • 减小batch_size
    • 使用--device 'cpu'暂时切换到CPU模式测试
  2. NaN损失值

    • 降低学习率(--learning_rate)
    • 增加--weight_decay
  3. 训练不收敛

    • 检查数据标准化是否正确
    • 尝试不同的学习率调度策略

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. 高级技巧与优化建议

一旦基础版本能正常运行,你可以尝试以下优化:

  1. 自适应学习率调整: 在训练循环中添加学习率衰减:

    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
    
  2. 梯度裁剪: 修改engine.py中的clip值,防止梯度爆炸:

    def __init__(self, ...):
        self.clip = 3  # 原为5,可以尝试更小的值
    
  3. 模型架构调整

    • 尝试不同的--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版本环境,以避免兼容性问题。

Logo

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

更多推荐