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

PyTorch中torch.nn.Parameter在Julia Flux中的等价实现

在Flux中实现PyTorch自定义模块的等价功能

完整Flux实现代码

using Flux

struct MyModel
    param::Flux.Param{Matrix{Float32}}
end

# 让Flux识别模型中的可训练参数
Flux.@functor MyModel

# 构造函数,初始化参数
MyModel() = MyModel(Flux.param(randn(Float32, 1, 1)))

# 定义前向传播(Flux中直接重载call方法)
function (m::MyModel)(x)
    x .* m.param
end

# 实例化模型
model = MyModel()

# 查看模型的可训练参数
println(Flux.params(model))
# 输出类似:Params([Float32[0.8764;;]])

# 前向传播
x = randn(Float32, 1, 1)
out = model(x)

# 计算损失
loss = mean(out)

# 反向传播(Flux中用梯度函数计算)
grads = Flux.gradient(() -> mean(model(x)), Flux.params(model))

# 查看参数的梯度
println(grads[model.param])
# 输出类似:Float32[-0.5432;;]

关键等价点说明

  • PyTorch中的nn.Parameter(torch.randn(1, 1))在Flux中对应的是Flux.param(randn(Float32, 1, 1))(或Flux.Param(...),两者等价)。Flux.param会将普通数组包装成可训练参数,自动开启梯度追踪。
  • 在Flux中定义自定义模型时,需要用Flux.@functor宏标注模型结构体,这样Flux才能自动识别并收集其中的可训练参数(对应PyTorch中nn.Module自动管理参数的逻辑)。
  • 前向传播逻辑通过重载模型结构体的call方法实现(即function (m::MyModel)(x)),效果等同于PyTorch中的forward方法。
  • 反向传播在Flux中通过Flux.gradient函数实现,传入损失计算的闭包和要追踪梯度的参数集合,最终可以通过梯度字典获取对应参数的梯度值。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 07:05:05