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

在Julia中用GPU运行Lux与NeuralPDE模型遇类型不稳定错误求助

CPU运行正常但GPU上NeuralPDE+Lux模型报错

基于Lux和NeuralPDE的模型在CPU上运行正常,切换到GPU时触发类型不稳定错误,怀疑和实测边界条件有关,代码如下:

using Random
using NeuralPDE, Lux, CUDA, Random
using Optimization
using OptimizationOptimisers
using NNlib
import ModelingToolkit: Interval
using Interpolations

# Measured Boundary Conditions (Arbitrary For Example)
bc1 = 1.0:1:1001.0 .|> Float32
bc2 = 1.0:1:1001.0 .|> Float32
ic1 = zeros(101) .|> Float32
ic2 = zeros(101) .|> Float32;

# Interpolation Functions Registered as Symbolic
itp1 = interpolate(bc1, BSpline(Cubic(Line(OnGrid()))))
up_cond_1_f(t::Float32) = itp1(t)
@register_symbolic up_cond_1_f(t)

itp2 = interpolate(bc2, BSpline(Cubic(Line(OnGrid()))))
up_cond_2_f(t::Float32) = itp2(t)
@register_symbolic up_cond_2_f(t)

itp3 = interpolate(ic1, BSpline(Cubic(Line(OnGrid()))))
init_cond_1_f(x::Float32) = itp3(x)
@register_symbolic init_cond_1_f(x)

itp4 = interpolate(ic2, BSpline(Cubic(Line(OnGrid()))))
init_cond_2_f(x::Float32) = itp4(x)
@register_symbolic init_cond_2_f(x);

# Parameters and differentials
@parameters t, x
@variables u1(..), u2(..)
Dt = Differential(t)
Dx = Differential(x);

# Arbitrary Equations
eqs = [Dt(u1(t, x)) + Dx(u2(t, x)) ~ 0.,
       Dt(u1(t, x)) * u1(t,x) + Dx(u2(t, x)) + 9.81 ~ 0.] 

# Boundary Conditions with Measured Data
bcs = [
       u1(t,1) ~ up_cond_1_f(t),
       u2(t,1) ~ up_cond_2_f(t),
       u1(1,x) ~ init_cond_1_f(x),
       u2(1,x) ~ init_cond_2_f(x)
]

# Space and time domains
domains = [t ∈ Interval(1.0,1001.0),
           x ∈ Interval(1.0,101.0)];

# Neural network
input_ = length(domains)
n = 10
chain = Chain(Dense(input_,n,NNlib.tanh_fast),Dense(n,n,NNlib.tanh_fast),Dense(n,4))

strategy = GridTraining(.25)
ps = Lux.setup(Random.default_rng(), chain)[1]
ps = ps |> Lux.ComponentArray |> gpu .|> Float32

discretization = PhysicsInformedNN(chain,
                                   strategy,
                                   init_params=ps)

# Model Setup
@named pdesystem = PDESystem(eqs,bcs,domains,[t,x],[u1(t, x),u2(t, x)])
prob = discretize(pdesystem,discretization);
sym_prob = symbolic_discretize(pdesystem,discretization);

# Losses and Callbacks
pde_inner_loss_functions = sym_prob.loss_functions.pde_loss_functions
bcs_inner_loss_functions = sym_prob.loss_functions.bc_loss_functions

callback = function (p, l)
    println("loss: ", l)
    println("pde_losses: ", map(l_ -> l_(p), pde_inner_loss_functions))
    println("bcs_losses: ", map(l_ -> l_(p), bcs_inner_loss_functions))
    return false
end;

# Train Model (Throws Error)
res = Optimization.solve(prob,Adam(0.01); callback = callback, maxiters=5000)
phi = discretization.phi;

运行时触发错误:

GPU broadcast resulted in non-concrete element type Union{}.
This probably means that the function you are broadcasting contains an error or type instability.


解决建议

  • 修复插值函数的GPU兼容性
    当前插值对象itp1-itp4仅支持CPU输入,GPU端调用时传入的是CUDA类型数据,导致类型推断失效。可以将输入临时转换为CPU类型再插值(适合快速验证):

    up_cond_1_f(t::Float32) = itp1(CPU(t))
    up_cond_2_f(t::Float32) = itp2(CPU(t))
    init_cond_1_f(x::Float32) = itp3(CPU(x))
    init_cond_2_f(x::Float32) = itp4(CPU(x))
    

    若追求GPU性能,可使用支持CUDA的插值实现,将插值数据和对象迁移到GPU:

    bc1_gpu = bc1 |> gpu
    itp1 = interpolate(bc1_gpu, BSpline(Cubic(Line(OnGrid()))))
    # 其余插值对象同理迁移
    
  • 统一数值类型
    代码中9.81是Float64类型,与其他Float32数据混合会在GPU上引发类型不稳定。将所有常量改为Float32:

    eqs = [Dt(u1(t, x)) + Dx(u2(t, x)) ~ 0f0,
           Dt(u1(t, x)) * u1(t,x) + Dx(u2(t, x)) + 9.81f0 ~ 0f0.] 
    
  • 修正神经网络输出维度
    模型仅需要预测u1和u2两个变量,但当前网络最后一层输出维度为4,维度不匹配会引发后续计算错误。修改为输出维度2:

    chain = Chain(Dense(input_,n,NNlib.tanh_fast),Dense(n,n,NNlib.tanh_fast),Dense(n,2))
    
  • 调整参数初始化顺序
    先将参数转为Float32再迁移到GPU,避免类型转换顺序错误:

    ps = Lux.setup(Random.default_rng(), chain)[1]
    ps = ps |> Lux.ComponentArray |> Float32 |> gpu
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 06:25:24