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

