【问题标题】:What does `train!()` do in Flux.jl?`train!()` 在 Flux.jl 中有什么作用?
【发布时间】:2021-06-27 14:50:07
【问题描述】:

在某些机器学习框架中,train 函数本身可能实际上并不进行训练,而只是设置模式(即只是确保模型等已准备好训练)。 Flux 中的 train 函数是这种情况,还是 train!() 函数实际上进行了训练?

【问题讨论】:

    标签: julia flux.jl


    【解决方案1】:

    根据Flux.jl docstrain!() 函数确实进行了实际训练。函数签名看起来像:train!(loss, params, data, opt; cb) 其中:

    对于 data 中的每个数据点 d,通过反向传播计算相对于 params 的损失梯度,并调用优化器 opt. 如果 d 是 loss 的参数元组,则调用 loss(d...),否则调用 loss(d)。 使用关键字参数 cb 给出回调。例如,这将每 10 秒打印一次“training”(使用 Flux.throttle): train!(loss, params, data, opt, cb = throttle(() -> println("training"), 10)) 回调可以调用 Flux.stop 来中断训练循环。 多个优化器和回调可以作为数组传递给 opt 和 cb。

    另一个例子:@epochs 2 Flux.train!(loss, ps, dataset, opt) 我们进行 2 个训练 epoch。您可以在Flux transfer learning tutorial 中找到更多信息。

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 1970-01-01
      • 2021-11-01
      • 2020-03-17
      • 2018-07-20
      • 1970-01-01
      • 2012-06-01
      • 2015-05-02
      • 2015-10-02
      相关资源
      最近更新 更多