You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在Python Darts中提取RNNModel的训练与验证损失?

在Darts RNNModel(LSTM)中提取训练与验证损失

Darts的RNNModel默认基于PyTorch后端,因此Keras风格的model.history.history["loss"]并不适用。以下是两种可靠的损失提取方案:

方案1:直接通过模型训练器获取(最简方式)

步骤1:拆分训练与验证集

要获取验证损失,必须先拆分数据集:

from darts import TimeSeries
from darts.models import RNNModel
from darts.utils.data import train_test_split

# 替换为你的完整时间序列数据
full_series = TimeSeries.from_dataframe(...)
train, val = train_test_split(full_series, test_size=0.2)

步骤2:训练并提取损失

训练时传入验证集,完成后通过model.trainer.history提取损失序列:

# 初始化LSTM模型
model = RNNModel(
    model="LSTM",
    input_chunk_length=past_samples,  # 替换为你的输入序列长度
    epochs=50,
    batch_size=32
)

# 传入训练集与验证集执行训练
model.fit(series=train, val_series=val)

# 提取训练/验证损失
train_loss = model.trainer.history["train_loss"]
val_loss = model.trainer.history["val_loss"]

# 查看结果
print("各epoch训练损失:", train_loss)
print("各epoch验证损失:", val_loss)

方案2:自定义回调实现灵活记录

如果需要实时保存或处理损失数据,可自定义PyTorch Lightning回调:

from pytorch_lightning.callbacks import Callback

class LossRecorder(Callback):
    def __init__(self):
        self.train_losses = []
        self.val_losses = []

    def on_train_epoch_end(self, trainer, pl_module):
        # 记录当前epoch训练损失
        self.train_losses.append(trainer.callback_metrics["train_loss"].item())

    def on_validation_epoch_end(self, trainer, pl_module):
        # 记录当前epoch验证损失
        self.val_losses.append(trainer.callback_metrics["val_loss"].item())

# 初始化回调实例
loss_recorder = LossRecorder()

# 初始化模型时传入回调
model = RNNModel(
    model="LSTM",
    input_chunk_length=past_samples,
    epochs=50,
    batch_size=32,
    callbacks=[loss_recorder]
)

# 启动训练
model.fit(train, val_series=val)

# 从回调中提取损失
train_loss = loss_recorder.train_losses
val_loss = loss_recorder.val_losses

注意事项

  • 若训练时未传入val_series,则仅能获取train_loss,无法生成验证损失。
  • 不同Darts版本的指标名称可能略有差异,可通过print(model.trainer.history.keys())查看所有可用的日志指标键值。

内容的提问来源于stack exchange,提问作者user23584725

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.28 02:43:24