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

Julia Flux 0.13及更高版本中Model的定义及自定义非神经网络模型训练适配问题

Julia Flux 0.13及更高版本中Model的定义及自定义非神经网络模型训练适配问题

嗨,我来帮你理清Flux里Model的定义逻辑,以及怎么解决你遇到的矩阵分解训练问题~

首先得明确:Flux里的Model不是单纯的计算函数,它本质是一个包含可训练参数的容器。你之前写的model(x, y) = x * y'只是一段计算逻辑,但没有把需要优化的x、y作为模型的一部分存储起来,这就导致Flux的setup和train!找不到要更新的参数,所以才会弹出那个警告。

下面针对你的矩阵分解问题,一步步给出可行的解决方案:

1. 定义包含可训练参数的模型

Flux支持多种参数容器形式,这里推荐两种简单易用的:

方式一:用NamedTuple(快速上手)

直接把需要优化的x和y打包成一个命名元组,Flux能自动识别其中的可训练参数:

using Flux, Random

dim = 2
A = rand(dim, dim)  # 目标矩阵

# 初始化模型:把可训练的x、y作为模型的属性
model = (x = rand(dim), y = rand(dim))

方式二:自定义结构体(更规范,适合复杂模型)

如果你的模型逻辑更复杂,可以自定义结构体,并用Flux.@functor宏标记,让Flux识别其中的参数:

struct MatrixFactorization
    x::Vector{Float64}
    y::Vector{Float64}
end

# 让Flux能识别这个结构体里的可训练参数
Flux.@functor MatrixFactorization

# 初始化模型
model = MatrixFactorization(rand(dim), rand(dim))

2. 定义损失函数

损失函数需要接收模型作为参数,计算预测值与真实值的误差(返回标量):

# 损失函数:计算x*y'与A的均方误差
loss(model) = sum(abs2, model.x * model.y' - A)

3. 使用Flux训练API

现在就可以正常使用setup和train!了:

# 设置优化器状态(这里用你指定的AdaGrad)
opt_state = Flux.setup(AdaGrad(), model)

# 训练循环(这里我们迭代1000次,每100次打印一次损失)
for epoch in 1:1000
    # 因为我们的损失函数不需要额外输入数据,所以数据迭代器用(()), 每次迭代传空元组
    Flux.train!(loss, model, ((),), opt_state)
    if epoch % 100 == 0
        println("Epoch $epoch, Loss: $(round(loss(model), digits=6))")
    end
end

# 验证结果
println("\n原始矩阵A:")
display(A)
println("\n分解后的近似矩阵x*y':")
display(model.x * model.y')

核心要点总结

  • Flux中的Model必须是可训练参数的容器:只要是Flux能通过Flux.params()识别的结构都可以,比如NamedTuple、Tuple、带@functor的自定义结构体,或者Flux自带的层组合。
  • 损失函数要能接收模型(或模型+数据),返回标量损失值,这样train!才能通过反向传播更新参数。
  • 如果你的问题是监督学习(需要输入数据-标签对),那train!的第三个参数就是数据迭代器(比如zip(X, Y)),损失函数也要对应改成loss(model, x, y) = ...。

备注:内容来源于stack exchange,提问作者Yrogirg

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.22 10:24:52