使用Flux训练含Attention层的自定义模型时遇训练错误求助
解决Flux自定义Attention层训练时的错误和警告
问题1:trainable返回类型警告
警告提示trainable(x) should now return a NamedTuple with the field names, not a Tuple,新版Flux要求trainable方法返回**命名元组(NamedTuple)**而非普通元组,确保优化器能正确识别可训练参数的名称。
修复方法
将原代码中的:
Flux.trainable(a_net::AttentionNet) = (a_net.embedding, a_net.attention, a_net.fc_output,)
修改为:
Flux.trainable(a_net::AttentionNet) = (embedding=a_net.embedding, attention=a_net.attention, fc_output=a_net.fc_output)
问题2:自动微分的MethodError
错误no method matching +(::Base.RefValue{Any}, ::NamedTuple...)源于Attention层前向传播中使用了列表推导式和列表求和操作,Zygote对这类循环/列表操作的梯度追踪支持不佳,导致梯度结构异常。
修复方法
将Attention层的前向传播改为张量化操作,避免循环和列表,让Zygote能正确追踪梯度:
function (a::Attention)(inputs) # 将输入列表堆叠为形状 (embedding_dim, num_vehicles, batch_size) 的张量 inputs_stacked = cat(inputs..., dims=2) # 对所有输入批量应用W变换 W_out = a.W(inputs_stacked) # 计算注意力权重并激活 alphas = sigmoid.(a.v(W_out)) # 加权求和后压缩维度 output = sum(alphas .* inputs_stacked, dims=2) return dropdims(output, dims=2) end
同时修正Attention构造函数中的拼写错误:vehile_embedding_dim改为vehicle_embedding_dim(不影响运行但更规范)。
完整修改后的代码
using Flux struct Attention W v end # 修正拼写错误:vehile_embedding_dim → vehicle_embedding_dim Attention(vehicle_embedding_dim::Integer) = Attention( Dense(vehicle_embedding_dim => vehicle_embedding_dim, tanh), Dense(vehicle_embedding_dim, 1, bias=false, init=Flux.zeros32) ) function (a::Attention)(inputs) # 张量化操作替代列表循环 inputs_stacked = cat(inputs..., dims=2) W_out = a.W(inputs_stacked) alphas = sigmoid.(a.v(W_out)) output = sum(alphas .* inputs_stacked, dims=2) return dropdims(output, dims=2) end Flux.@functor Attention struct AttentionNet embedding attention fc_output vehicle_num::Integer vehicle_dim::Integer end AttentionNet(vehicle_num::Integer, vehicle_dim::Integer, embedding_dim::Integer) = AttentionNet( Dense(vehicle_dim+1 => embedding_dim, relu), Attention(embedding_dim), Dense(1+embedding_dim => 1), vehicle_num, vehicle_dim ) function (a_net::AttentionNet)(x) time_idx = x[[1], :] vehicle_states = [x[2+a_net.vehicle_dim*(i-1):2+a_net.vehicle_dim*i-1, :] for i in 1:a_net.vehicle_num] vehicle_states = [vcat(time_idx, vehicle_state) for vehicle_state in vehicle_states] vehicle_embeddings = a_net.embedding.(vehicle_states) attention_output = a_net.attention(vehicle_embeddings) x = a_net.fc_output(vcat(time_idx, attention_output)) return x end Flux.@functor AttentionNet # 修改为返回NamedTuple Flux.trainable(a_net::AttentionNet) = (embedding=a_net.embedding, attention=a_net.attention, fc_output=a_net.fc_output) # 测试代码 fake_inputs = rand(22, 640) fake_outputs = rand(1, 640) a_net = AttentionNet(3, 7, 64)|> gpu opt = Adam(.01) opt_state = Flux.setup(opt, a_net) data = Flux.DataLoader((fake_inputs, fake_outputs)|>gpu, batchsize=32, shuffle=true) Flux.train!(a_net, data, opt_state) do m, x, y Flux.mse(m(x), y) end
内容的提问来源于stack exchange,提问作者PokeLu
相关产品推荐
相关产品推荐

