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

使用Flux.jl获取每轮训练损失值是否必须编写自定义训练循环?

Flux.jl通过@epoch与train!获取训练指标的方法

核心结论

完全可以通过@epoch宏搭配train!的回调机制获取每轮训练的损失、准确率等指标,无需自行编写完整自定义训练循环。

实现逻辑

train!函数的cb参数支持传入自定义回调函数,每完成一个批次的训练就会自动触发该回调,我们可以在回调中完成批次级的指标统计,再在@epoch的循环体内完成轮次级的指标聚合与输出。

代码示例

using Flux
using Flux: onecold, @epoch

# 前置准备:已提前定义模型model、损失函数loss、优化器opt、训练数据集train_data
# 初始化统计变量
train_loss = 0.0
correct_num = 0
total_num = 0

# 自定义回调:每批次训练完成后统计指标
function batch_metrics_cb(batch)
    x, y = batch
    y_pred = model(x)
    batch_size = size(x, ndims(x))
    # 累加损失与样本数
    global train_loss += loss(y_pred, y) * batch_size
    # 累加正确预测样本数
    global correct_num += sum(onecold(y_pred) .== onecold(y))
    global total_num += batch_size
end

# 主训练循环:训练10轮
@epoch 10 begin
    # 每轮开始前重置统计变量
    global train_loss = 0.0
    global correct_num = 0
    global total_num = 0
    # 传入回调启动训练
    Flux.train!(loss, Flux.params(model), train_data, opt, cb = batch_metrics_cb)
    # 轮次结束后计算并打印指标
    avg_loss = round(train_loss / total_num, digits=4)
    train_acc = round(correct_num / total_num * 100, digits=2)
    println("Epoch $epoch | 平均损失: $avg_loss | 训练准确率: $train_acc%")
end

扩展说明

  • 若需要按批次打印指标,可直接将打印逻辑写入回调函数中,还可添加条件控制打印频率,比如每10个批次输出一次
  • 若需要统计验证集指标,可在每轮训练结束后新增验证逻辑,直接调用模型跑验证集计算对应指标即可

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 16:45:04