手把手教你复现Lag-Llama论文实验:从预训练到评估的脚本使用指南
手把手教你复现Lag-Llama论文实验:从预训练到评估的脚本使用指南
Lag-Llama是一个专注于概率时间序列预测的基础模型项目,本文将详细介绍如何使用项目提供的脚本完成从环境准备到模型预训练、微调及评估的全流程实验复现。
📋 实验前准备工作
1. 环境配置与依赖安装
首先需要克隆项目仓库并安装所需依赖:
git clone https://gitcode.com/gh_mirrors/la/lag-llama
cd lag-llama
pip install -r requirements.txt
项目推荐使用Python 3.10.8版本,建议创建新的Anaconda环境以避免依赖冲突。
2. 数据集准备
需要下载非Monash数据集并解压到指定目录:
tar -xvzf nonmonash_datasets.tar.gz -C datasets
图:Lag-Llama模型处理时间序列数据的概念示意图,展示了模型如何学习时间序列中的模式
🔧 预训练模型步骤
预训练脚本位于scripts/pretrain.sh,执行前需要修改第59行的Weights and Biases参数,替换为自己的实体、项目和标签信息。
预训练命令解析
python run.py \
-e $EXP_NAME -d "datasets" --seed $SEED \
-r "experiments/results" \
--batch_size 512 -m 1000 -n 128 \
--wandb_entity "your-entity" --wandb_project "your-project" --wandb_tags "pretraining" \
--all_datasets "australian_electricity_demand" "electricity_hourly" ... \
--test_datasets "weather" "pedestrian_counts" ... \
--num_workers 2 --args_from_dict_path configs/lag_llama.json --search_batch_size \
--lr 0.0001
关键参数说明:
--batch_size: 训练批次大小,默认为512-m 1000: 最大训练轮数-n 128: 每个序列的长度--args_from_dict_path: 模型配置文件路径,默认使用configs/lag_llama.json
执行预训练
chmod +x scripts/pretrain.sh
./scripts/pretrain.sh
预训练过程中会自动创建实验目录结构,并在experiments/seeds/下生成随机种子文件。
🔄 模型微调流程
微调脚本scripts/finetune.sh需要基于已训练的预模型进行,关键配置如下:
微调参数设置
CONFIGPATH="configs/lag_llama.json"
PRETRAINING_EXP_NAME="pretraining_lag_llama"
PERCENTAGE=100 # 可设置为20, 40, 60, 80以复现论文中的不同实验
微调命令解析
python run.py \
-e $EXPERIMENT_NAME -d "datasets" --seed $SEED \
-r "experiments/results" \
--batch_size 512 -m 1000 -n 128 \
--wandb_entity "your-entity" --wandb_project "your-project" \
--num_workers 2 --args_from_dict_path $CONFIGPATH --search_batch_size \
--single_dataset $FINETUNE_DATASET \
--get_ckpt_path_from_experiment_name $PRETRAINING_EXP_NAME --lr 0.00001 \
--use_dataset_prediction_length --num_validation_windows 1 \
--single_dataset_last_k_percentage $PERCENTAGE
执行微调
chmod +x scripts/finetune.sh
./scripts/finetune.sh
脚本会自动读取预训练时生成的种子文件,并对每个目标数据集进行微调。
⚙️ 模型配置详解
模型核心配置文件configs/lag_llama.json包含以下关键参数:
{
"n_layer": 8, // transformer层数
"n_head": 9, // 注意力头数
"n_embd_per_head": 16, // 每个注意力头的嵌入维度
"context_length": 32, // 上下文窗口长度
"dropout": 0.0 // dropout比率
}
可根据硬件条件和具体任务需求调整这些参数,但为了复现论文结果,建议保持默认配置。
📊 实验结果评估
实验结果会保存在experiments/results/目录下,包含模型性能指标和预测可视化结果。主要评估指标包括:
- 均方根误差(RMSE)
- 平均绝对误差(MAE)
- 预测区间覆盖率(PICP)
通过对比不同数据集上的性能表现,可以分析模型在各类时间序列预测任务上的泛化能力。
💡 常见问题解决
-
GPU内存不足:可减小
--batch_size参数或使用--search_batch_size自动搜索合适的批次大小 -
数据集路径错误:确保数据集解压到
datasets/目录,或通过-d参数指定正确路径 -
Weights and Biases错误:检查第59行的wandb参数是否正确配置,或移除相关参数以禁用wandb日志
通过以上步骤,你可以完整复现Lag-Llama论文中的实验,从预训练到微调再到模型评估,深入了解这个概率时间序列预测基础模型的工作原理和性能表现。
更多推荐




所有评论(0)