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

Flux中训练循环与精度计算分离的疑问及PyTorch实现对比

关于Flux MLP-MNIST示例中训练与指标计算分离的疑问解答

示例设计的原因

Flux示例中选择将训练循环与指标计算循环分离,主要有以下几点考量:

  1. 代码模块化与复用性
    单独的loss_and_accuracy函数可同时用于训练集和测试集的指标计算,避免重复编写相同逻辑,让代码更简洁易维护。

  2. 避免梯度计算干扰
    在训练过程中计算精度时,若不小心触发梯度追踪,会额外消耗计算资源甚至影响训练流程。分离的评估循环天然处于无梯度的推理模式(Flux默认不会对非梯度计算操作追踪梯度),无需额外处理就能规避这类问题。

  3. 指标计算的一致性
    用同一个函数计算训练集和测试集的损失与精度,确保两者计算逻辑完全一致,减少重复代码带来的潜在bug。比如示例中统一使用agg=sum累加损失,再除以总样本数得到平均损失,这种逻辑在训练和测试中保持统一。

是否可以合并到同一循环中?

完全可以。你可以参考PyTorch的写法,将梯度更新与指标计算合并到同一个batch循环里,以下是修改后的Flux训练函数示例:

function train(;kws...)
    args = Args(;kws...)
    model = build_model() 
    optimizer = setup(Adam(args.η), model)

    for epoch in 1:args.epochs
        total_loss::Float32 = 0.0f0
        total_correct::Int = 0
        total_samples::Int = 0

        for (X,y) in trainloader
            # 计算梯度并更新模型
            grad = gradient(model) do m
                ŷ = m(X)
                # 和loss_and_accuracy保持一致,用sum聚合批次损失
                loss_func(ŷ, y, agg=sum)
            end
            Optimise.update!(optimizer, model, grad[1])

            # 无梯度模式下计算当前批次指标,避免额外资源消耗
            Flux.no_grad() do
                ŷ = model(X)
                total_loss += loss_func(ŷ, y, agg=sum)
                total_correct += sum(onecold(ŷ) .== onecold(y))
                total_samples += size(X)[end]
            end
        end

        # 计算平均损失与精度
        train_loss = total_loss / total_samples
        train_acc = total_correct / total_samples
        test_loss, test_acc = loss_and_accuracy(testloader, model)
        
        println("Epoch : $epoch")
        println("training loss : $train_loss , training acc : $train_acc")
        println("testing loss : $test_loss, testing acc : $test_acc")
    end
end

注意事项

  • 使用Flux.no_grad()包裹指标计算部分,避免不必要的梯度追踪,提升运行效率。
  • 损失累加需和loss_and_accuracy保持一致的聚合方式(比如均使用agg=sum),确保最终平均损失计算准确。
  • 合并后的代码在效率上略高(减少一次训练集遍历),但逻辑会稍复杂,可根据自身需求选择合适的实现风格。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 23:25:06