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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 10:50:25