使用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
相关产品推荐
相关产品推荐

