# Questions on NeuralPDE.jl

**URL:** <https://discourse.julialang.org/t/questions-on-neuralpde-jl/83846>\
**Category:** Modelling & Simulations\
**Created:** [July 6, 2022, 4:34pm UTC](https://discourse.julialang.org/t/questions-on-neuralpde-jl/83846 "2022-07-06T16:34:20Z")\
**Posts on this page:** 6\
**Page:** 1

<div class="post-metadata">

**Author:** ![marcofrancis](https://avatars.discourse-cdn.com/v4/letter/m/bcef8e/32.png) [@marcofrancis](https://discourse.julialang.org/u/marcofrancis)\
**Post date:** [July 6, 2022, 4:34pm UTC](https://discourse.julialang.org/t/questions-on-neuralpde-jl/83846/1 "2022-07-06T16:34:20Z")

</div>

Hi, I’m quite new to Julia and I discovered this amazing package. I have a couple of questions regarding it:  
1 - Should I use Lux or Flux? It seems that Lux is used in almost all examples, but when using the gpu one should use Flux, is this correct? (if I use Lux and use the same syntax I get a warning and the variable “chain” which should contain the NN is of type “Nothing”)  
2 - Suppose I have a Parabolic PDE in 5 dimensions that could be solved also with HighDimPDE.jl, which of the two packages should I use and why?  
3 - Is there any restriction on what packages I can use to optimise/train the NN?  
Thanks a lot 🙂

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [July 6, 2022, 8:03pm UTC](https://discourse.julialang.org/t/questions-on-neuralpde-jl/83846/2 "2022-07-06T20:03:18Z")

</div>

Hey,  
These are the kinds of questions I plan to address in the higher level [https://docs.sciml.ai/dev/](https://docs.sciml.ai/dev/), but I’ll put some short answers here.

> [@marcofrancis](#):
>
> if I use Lux and use the same syntax I get a warning and the variable “chain” which should contain the NN is of type “Nothing”

Can you share more details on this? Someone also that that here for DiffEqFlux on a Lux example:

> <https://github.com/SciML/DiffEqFlux.jl/issues/748>
>
> I am following the tutorial on \[Neural Ordinary Differential Equations\](https://…diffeqflux.sciml.ai/stable/examples/neural\_ode/#Neural-Ordinary-Differential-Equations). 
> 
> After copy-pasting code, there is an error I cannot debug:
> 
> \`\`\`
> ERROR: LoadError: MethodError: objects of type Nothing are not callable
> Stacktrace:
> \[1\] (::DiffEqFlux.var"#dudt\_#133"{NeuralODE{Lux.Chain{NamedTuple{(:layer\_1, :layer\_2, :layer\_3), Tuple{WrappedFunction{var"#1#2"}, Lux.Dense{true, typeof(tanh\_fast), typeof(Lux.glorot\_uniform), typeof(Lux.zeros32)}, Lux.Dense{true, typeof(identity), typeof(Lux.glorot\_uniform), typeof(Lux.zeros32)}}}}, Nothing, Nothing, Tuple{Float32, Float32}, Tuple{Tsit5{typeof(OrdinaryDiffEq.trivial\_limiter!), typeof(OrdinaryDiffEq.trivial\_limiter!), Static.False}}, Base.Pairs{Symbol, StepRangeLen{Float32, Float64, Float64, Int64}, Tuple{Symbol}, NamedTuple{(:saveat,), Tuple{StepRangeLen{Float32, Float64, Float64, Int64}}}}}})(u::Vector{Float32}, p::ComponentVector{Float32}, t::Float32)
> @ DiffEqFlux ~/.julia/packages/DiffEqFlux/7N0N5/src/neural\_de.jl:76
> \[2\] (::ODEFunction{false, DiffEqFlux.var"#dudt\_#133"{NeuralODE{Lux.Chain{NamedTuple{(:layer\_1, :layer\_2, :layer\_3), Tuple{WrappedFunction{var"#1#2"}, Lux.Dense{true, typeof(tanh\_fast), typeof(Lux.glorot\_uniform), typeof(Lux.zeros32)}, Lux.Dense{true, typeof(identity), typeof(Lux.glorot\_uniform), typeof(Lux.zeros32)}}}}, Nothing, Nothing, Tuple{Float32, Float32}, Tuple{Tsit5{typeof(OrdinaryDiffEq.trivial\_limiter!), typeof(OrdinaryDiffEq.trivial\_limiter!), Static.False}}, Base.Pairs{Symbol, StepRangeLen{Float32, Float64, Float64, Int64}, Tuple{Symbol}, NamedTuple{(:saveat,), Tuple{StepRangeLen{Float32, Float64, Float64, Int64}}}}}}, LinearAlgebra.UniformScaling{Bool}, Nothing, typeof(DiffEqFlux.basic\_tgrad), Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, typeof(SciMLBase.DEFAULT\_OBSERVED), Nothing})(::Vector{Float32}, ::Vararg{Any})
> @ SciMLBase ~/.julia/packages/SciMLBase/IJbT7/src/scimlfunctions.jl:1624
> \[3\] initialize!(integrator::OrdinaryDiffEq.ODEIntegrator{Tsit5{typeof(OrdinaryDiffEq.trivial\_limiter!), typeof(OrdinaryDiffEq.trivial\_limiter!), Static.False}, false, Vector{Float32}, Nothing, Float32, ComponentArrays.ComponentVector{Float32, Vector{Float32}, Tuple{ComponentArrays.Axis{(layer\_1 = 1:0, layer\_2 = ViewAxis(1:150, Axis(weight = ViewAxis(1:100, ShapedAxis((50, 2), NamedTuple())), bias = ViewAxis(101:150, ShapedAxis((50, 1), NamedTuple())))), layer\_3 = ViewAxis(151:252, Axis(weight = ViewAxis(1:100, ShapedAxis((2, 50), NamedTuple())), bias = ViewAxis(101:102, ShapedAxis((2, 1), NamedTuple())))))}}}, Float32, Float32, Float32, Float32, Vector{Vector{Float32}}, ODESolution{Float32, 2, Vector{Vector{Float32}}, Nothing, Nothing, Vector{Float32}, Vector{Vector{Vector{Float32}}}, ODEProblem{Vector{Float32}, Tuple{Float32, Float32}, false, ComponentArrays.ComponentVector{Float32, Vector{Float32}, Tuple{ComponentArrays.Axis{(layer\_1 = 1:0, layer\_2 = ViewAxis(1:150, Axis(weight = ViewAxis(1:100, ShapedAxis((50, 2), NamedTuple())), bias = ViewAxis(101:150, ShapedAxis((50, 1), NamedTuple())))), layer\_3 = ViewAxis(151:252, Axis(weight = ViewAxis(1:100, ShapedAxis((2, 50), NamedTuple())), bias = ViewAxis(101:102, ShapedAxis((2, 1), NamedTuple())))))}}}, ODEFunction{false, DiffEqFlux.var"#dudt\_#133"{NeuralODE{Lux.Chain{NamedTuple{(:layer\_1, :layer\_2, :layer\_3), Tuple{WrappedFunction{var"#1#2"}, Lux.Dense{true, typeof(tanh\_fast), typeof(Lux.glorot\_uniform), typeof(Lux.zeros32)}, Lux.Dense{true, typeof(identity), typeof(Lux.glorot\_uniform), typeof(Lux.zeros32)}}}}, Nothing, Nothing, Tuple{Float32, Float32}, Tuple{Tsit5{typeof(OrdinaryDiffEq.trivial\_limiter!), typeof(OrdinaryDiffEq.trivial\_limiter!), Static.False}}, Base.Pairs{Symbol, StepRangeLen{Float32, Float64, Float64, Int64}, Tuple{Symbol}, NamedTuple{(:saveat,), Tuple{StepRangeLen{Float32, Float64, Float64, Int64}}}}}}, LinearAlgebra.UniformScaling{Bool}, Nothing, typeof(DiffEqFlux.basic\_tgrad), Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, typeof(SciMLBase.DEFAULT\_OBSERVED), Nothing}, Base.Pairs{Symbol, Union{}, Tuple{}, NamedTuple{(), Tuple{}}}, SciMLBase.StandardODEProblem}, Tsit5{typeof(OrdinaryDiffEq.trivial\_limiter!), typeof(OrdinaryDiffEq.trivial\_limiter!), Static.False}, OrdinaryDiffEq.InterpolationData{ODEFunction{false, DiffEqFlux.var"#dudt\_#133"{NeuralODE{Lux.Chain{NamedTuple{(:layer\_1, :layer\_2, :layer\_3), Tuple{WrappedFunction{var"#1#2"}, Lux.Dense{true, typeof(tanh\_fast), typeof(Lux.glorot\_uniform), typeof(Lux.zeros32)}, Lux.Dense{true, typeof(identity), typeof(Lux.glorot\_uniform), typeof(Lux.zeros32)}}}}, Nothing, Nothing, Tuple{Float32, Float32}, Tuple{Tsit5{typeof(OrdinaryDiffEq.trivial\_limiter!), typeof(OrdinaryDiffEq.trivial\_limiter!), Static.False}}, Base.Pairs{Symbol, StepRangeLen{Float32, Float64, Float64, Int64}, Tuple{Symbol}, NamedTuple{(:saveat,), Tuple{StepRangeLen{Float32, Float64, Float64, Int64}}}}}}, LinearAlgebra.UniformScaling{Bool}, Nothing, typeof(DiffEqFlux.basic\_tgrad), Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, typeof(SciMLBase.DEFAULT\_OBSERVED), Nothing}, Vector{Vector{Float32}}, Vector{Float32}, Vector{Vector{Vector{Float32}}}, OrdinaryDiffEq.Tsit5ConstantCache{Float32, Float32}}, DiffEqBase.DEStats}, ODEFunction{false, DiffEqFlux.var"#dudt\_#133"{NeuralODE{Lux.Chain{NamedTuple{(:layer\_1, :layer\_2, :layer\_3), Tuple{WrappedFunction{var"#1#2"}, Lux.Dense{true, typeof(tanh\_fast), typeof(Lux.glorot\_uniform), typeof(Lux.zeros32)}, Lux.Dense{true, typeof(identity), typeof(Lux.glorot\_uniform), typeof(Lux.zeros32)}}}}, Nothing, Nothing, Tuple{Float32, Float32}, Tuple{Tsit5{typeof(OrdinaryDiffEq.trivial\_limiter!), typeof(OrdinaryDiffEq.trivial\_limiter!), Static.False}}, Base.Pairs{Symbol, StepRangeLen{Float32, Float64, Float64, Int64}, Tuple{Symbol}, NamedTuple{(:saveat,), Tuple{StepRangeLen{Float32, Float64, Float64, Int64}}}}}}, LinearAlgebra.UniformScaling{Bool}, Nothing, typeof(DiffEqFlux.basic\_tgrad), Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, typeof(SciMLBase.DEFAULT\_OBSERVED), Nothing}, OrdinaryDiffEq.Tsit5ConstantCache{Float32, Float32}, OrdinaryDiffEq.DEOptions{Float32, Float32, Float32, Float32, PIController{Rational{Int64}}, typeof(DiffEqBase.ODE\_DEFAULT\_NORM), typeof(LinearAlgebra.opnorm), Nothing, CallbackSet{Tuple{}, Tuple{}}, typeof(DiffEqBase.ODE\_DEFAULT\_ISOUTOFDOMAIN), typeof(DiffEqBase.ODE\_DEFAULT\_PROG\_MESSAGE), typeof(DiffEqBase.ODE\_DEFAULT\_UNSTABLE\_CHECK), DataStructures.BinaryHeap{Float32, DataStructures.FasterForward}, DataStructures.BinaryHeap{Float32, DataStructures.FasterForward}, Nothing, Nothing, Int64, Tuple{}, StepRangeLen{Float32, Float64, Float64, Int64}, Tuple{}}, Vector{Float32}, Float32, Nothing, OrdinaryDiffEq.DefaultInit}, cache::OrdinaryDiffEq.Tsit5ConstantCache{Float32, Float32})
> @ OrdinaryDiffEq ~/.julia/packages/OrdinaryDiffEq/UG9Mz/src/perform\_step/low\_order\_rk\_perform\_step.jl:569
> \[4\] \_\_init(prob::ODEProblem{Vector{Float32}, Tuple{Float32, Float32}, false, ComponentArrays.ComponentVector{Float32, Vector{Float32}, Tuple{ComponentArrays.Axis{(layer\_1 = 1:0, layer\_2 = ViewAxis(1:150, Axis(weight = ViewAxis(1:100, ShapedAxis((50, 2), NamedTuple())), bias = ViewAxis(101:150, ShapedAxis((50, 1), NamedTuple())))), layer\_3 = ViewAxis(151:252, Axis(weight = ViewAxis(1:100, ShapedAxis((2, 50), NamedTuple())), bias = ViewAxis(101:102, ShapedAxis((2, 1), NamedTuple())))))}}}, ODEFunction{false, DiffEqFlux.var"#dudt\_#133"{NeuralODE{Lux.Chain{NamedTuple{(:layer\_1, :layer\_2, :layer\_3), Tuple{WrappedFunction{var"#1#2"}, Lux.Dense{true, typeof(tanh\_fast), typeof(Lux.glorot\_uniform), typeof(Lux.zeros32)}, Lux.Dense{true, typeof(identity), typeof(Lux.glorot\_uniform), typeof(Lux.zeros32)}}}}, Nothing, Nothing, Tuple{Float32, Float32}, Tuple{Tsit5{typeof(OrdinaryDiffEq.trivial\_limiter!), typeof(OrdinaryDiffEq.trivial\_limiter!), Static.False}}, Base.Pairs{Symbol, StepRangeLen{Float32, Float64, Float64, Int64}, Tuple{Symbol}, NamedTuple{(:saveat,), Tuple{StepRangeLen{Float32, Float64, Float64, Int64}}}}}}, LinearAlgebra.UniformScaling{Bool}, Nothing, typeof(DiffEqFlux.basic\_tgrad), Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, typeof(SciMLBase.DEFAULT\_OBSERVED), Nothing}, Base.Pairs{Symbol, Union{}, Tuple{}, NamedTuple{(), Tuple{}}}, SciMLBase.StandardODEProblem}, alg::Tsit5{typeof(OrdinaryDiffEq.trivial\_limiter!), typeof(OrdinaryDiffEq.trivial\_limiter!), Static.False}, timeseries\_init::Tuple{}, ts\_init::Tuple{}, ks\_init::Tuple{}, recompile::Type{Val{true}}; saveat::StepRangeLen{Float32, Float64, Float64, Int64}, tstops::Tuple{}, d\_discontinuities::Tuple{}, save\_idxs::Nothing, save\_everystep::Bool, save\_on::Bool, save\_start::Bool, save\_end::Nothing, callback::Nothing, dense::Bool, calck::Bool, dt::Float32, dtmin::Nothing, dtmax::Float32, force\_dtmin::Bool, adaptive::Bool, gamma::Rational{Int64}, abstol::Nothing, reltol::Nothing, qmin::Rational{Int64}, qmax::Int64, qsteady\_min::Int64, qsteady\_max::Int64, beta1::Nothing, beta2::Nothing, qoldinit::Rational{Int64}, controller::Nothing, fullnormalize::Bool, failfactor::Int64, maxiters::Int64, internalnorm::typeof(DiffEqBase.ODE\_DEFAULT\_NORM), internalopnorm::typeof(LinearAlgebra.opnorm), isoutofdomain::typeof(DiffEqBase.ODE\_DEFAULT\_ISOUTOFDOMAIN), unstable\_check::typeof(DiffEqBase.ODE\_DEFAULT\_UNSTABLE\_CHECK), verbose::Bool, timeseries\_errors::Bool, dense\_errors::Bool, advance\_to\_tstop::Bool, stop\_at\_next\_tstop::Bool, initialize\_save::Bool, progress::Bool, progress\_steps::Int64, progress\_name::String, progress\_message::typeof(DiffEqBase.ODE\_DEFAULT\_PROG\_MESSAGE), userdata::Nothing, allow\_extrapolation::Bool, initialize\_integrator::Bool, alias\_u0::Bool, alias\_du0::Bool, initializealg::OrdinaryDiffEq.DefaultInit, kwargs::Base.Pairs{Symbol, Union{}, Tuple{}, NamedTuple{(), Tuple{}}})
> @ OrdinaryDiffEq ~/.julia/packages/OrdinaryDiffEq/UG9Mz/src/solve.jl:456
> \[5\] #\_\_solve#502
> @ ~/.julia/packages/OrdinaryDiffEq/UG9Mz/src/solve.jl:4 \[inlined\]
> \[6\] #solve\_call#28
> @ ~/.julia/packages/DiffEqBase/RHAWf/src/solve.jl:429 \[inlined\]
> \[7\] #solve\_up#34
> @ ~/.julia/packages/DiffEqBase/RHAWf/src/solve.jl:767 \[inlined\]
> \[8\] #solve#33
> @ ~/.julia/packages/DiffEqBase/RHAWf/src/solve.jl:752 \[inlined\]
> \[9\] (::NeuralODE{Lux.Chain{NamedTuple{(:layer\_1, :layer\_2, :layer\_3), Tuple{WrappedFunction{var"#1#2"}, Lux.Dense{true, typeof(tanh\_fast), typeof(Lux.glorot\_uniform), typeof(Lux.zeros32)}, Lux.Dense{true, typeof(identity), typeof(Lux.glorot\_uniform), typeof(Lux.zeros32)}}}}, Nothing, Nothing, Tuple{Float32, Float32}, Tuple{Tsit5{typeof(OrdinaryDiffEq.trivial\_limiter!), typeof(OrdinaryDiffEq.trivial\_limiter!), Static.False}}, Base.Pairs{Symbol, StepRangeLen{Float32, Float64, Float64, Int64}, Tuple{Symbol}, NamedTuple{(:saveat,), Tuple{StepRangeLen{Float32, Float64, Float64, Int64}}}}})(x::Vector{Float32}, p::ComponentVector{Float32})
> @ DiffEqFlux ~/.julia/packages/DiffEqFlux/7N0N5/src/neural\_de.jl:80
> \[10\] predict\_neuralode(p::ComponentVector{Float32})
> @ Main ~/Programming/NPDEChaos-combine/NPDEChaos/test/neuralode.jl:24
> \[11\] loss\_neuralode(p::ComponentVector{Float32})
> @ Main ~/Programming/NPDEChaos-combine/NPDEChaos/test/neuralode.jl:34
> \[12\] top-level scope
> @ ~/Programming/NPDEChaos-combine/NPDEChaos/test/neuralode.jl:53
> \[13\] include(fname::String)
> @ Base.MainInclude ./client.jl:451
> \[14\] top-level scope
> @ REPL\[2\]:1
> in expression starting at /home/neo/Programming/NPDEChaos-combine/NPDEChaos/test/neuralode.jl:53
> \`\`\`
> 
> And here are my packages:
> 
> \`\`\`
> \[fbb218c0\] BSON v0.3.5
> \[8e7c35d0\] BlockArrays v0.16.16
> \[052768ef\] CUDA v3.9.1
> \[35d6a980\] ColorSchemes v3.18.0
> \[f68482b8\] Cthulhu v2.6.2
> \[aae7a2af\] DiffEqFlux v1.51.2
> \[0c46a032\] DifferentialEquations v7.1.0
> \[31c24e10\] Distributions v0.25.61
> \[61744808\] DynamicalSystems v2.3.0
> \[587475ba\] Flux v0.13.3
> \[f6369f11\] ForwardDiff v0.10.30
> \[a75be94c\] GalacticOptim v3.4.0
> \[9d3c5eb1\] GalacticOptimJL v0.1.0
> \[033835bb\] JLD2 v0.4.22
> \[682c06a0\] JSON v0.21.3
> \[b2108857\] Lux v0.4.7
> \[23992714\] MAT v0.10.3
> \[eb30cadb\] MLDatasets v0.7.3
> \[961ee093\] ModelingToolkit v8.13.3
> \[315f7962\] NeuralPDE v4.9.0
> \[429524aa\] Optim v1.7.0
> \[7f7a1694\] Optimization v3.6.0
> \[36348300\] OptimizationOptimJL v0.1.1
> \[91a5bcdd\] Plots v1.29.0
> \[d330b81b\] PyPlot v2.10.0
> \[67601950\] Quadrature v2.0.0
> \[8a4e6c94\] QuasiMonteCarlo v0.2.9
> \[295af30f\] Revise v3.3.3
> \[fa074d72\] SeqData v0.1.0 \`https://github.com/maximilian-gelbrecht/SeqData.jl.git#master\`
> \[b8865327\] UnicodePlots v2.12.4
> \[009559a3\] XGBoost v1.5.2
> \[e88e6eb3\] Zygote v0.6.40
> \[2f01184e\] SparseArrays
> \`\`\`
> 
> I think the bug in the first tutorial is important. Anyone can debug it?

But our CI machines and my laptop cannot recreate it. Any computer I have run things on seems fine. So if you can help me hone in on what it could be (maybe it’s something like using Julia v1.6 instead of v1.7? A Mac M1 chip only issue? I don’t know, shots in the dark right now). (@avikpal do you know what it could be?)

> [@marcofrancis](#):
>
> Should I use Lux or Flux? It seems that Lux is used in almost all examples, but when using the gpu one should use Flux, is this correct?

Generally we prefer Lux.jl because it’s simpler to get both correct and fast code with. But we also make sure to support Flux.jl everywhere for compatibility reasons.

ComponentArrays.jl currently has issues with GPUArrays broadcast overloads which is why that’s not used for the GPU examples right now, but the intention is to make all examples Lux-based once that is figured out.

> [@marcofrancis](#):
>
> Suppose I have a Parabolic PDE in 5 dimensions that could be solved also with HighDimPDE.jl, which of the two packages should I use and why?

HighDimPDE.jl. Its methods are much more specialized to a specific type of PDE, and from that it’s much more efficient on those PDEs.

> [@marcofrancis](#):
>
> 3 - Is there any restriction on what packages I can use to optimise/train the NN?

Nope. Optimization.jl is pretty much just a wrapper to all other optimization packages we can get our hands on, so you can just have it call any package in the big list.

[http://optimization.sciml.ai/stable/](http://optimization.sciml.ai/stable/)

If the question is whether you can target it to a different optimization front end, yes it’s possible to take the `OptimizationProblem` and have it call other packages because… that’s what the wrapper solvers are doing 😅. So generally it would just be easier to call the wrapper solver since otherwise you’ll likely be writing similar code. If there’s some wrapper that is missing, it would be nice to PR that and just add it to the interface. It’s generic enough that you can hack it in a few minutes to use something random like SciPy to solve the OptimizationProblem if you really wished, so it definitely doesn’t cover all possible optimization backends.

---

<div class="post-metadata">

**Author:** ![marcofrancis](https://avatars.discourse-cdn.com/v4/letter/m/bcef8e/32.png) [@marcofrancis](https://discourse.julialang.org/u/marcofrancis)\
**Post date:** [July 7, 2022, 7:55am UTC](https://discourse.julialang.org/t/questions-on-neuralpde-jl/83846/3 "2022-07-07T07:55:53Z")

</div>

Cool, thanks a lot for all the answers. Indeed I find the whole SciML ecosystem a bit confusing but super exciting!

Regarding the Nothing issue, I have it if I declare something like:

dim = 2  
inner = 16

chain = Lux.Chain(Dense(dim,inner,Lux.σ),  
Dense(inner,inner,Lux.σ),  
Dense(inner,1)) |\> gpu

I have a RTX2070 mobile so nothing exotic, Julia version is 1.7.2

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [July 7, 2022, 8:29am UTC](https://discourse.julialang.org/t/questions-on-neuralpde-jl/83846/4 "2022-07-07T08:29:38Z")

</div>

> [@marcofrancis](#):
>
> chain = Lux.Chain(Dense(dim,inner,Lux.σ),  
> Dense(inner,inner,Lux.σ),  
> Dense(inner,1)) |\> gpu

Doesn’t make sense. Calling `gpu` on a chain is a Flux idea, not a Lux idea (and this is part of the whole, Lux needs to support GPUs differently, and it doesn’t in the context of vector solves right now because of a ComponentArrays.jl issue).

---

<div class="post-metadata">

**Author:** ![marcofrancis](https://avatars.discourse-cdn.com/v4/letter/m/bcef8e/32.png) [@marcofrancis](https://discourse.julialang.org/u/marcofrancis)\
**Post date:** [July 7, 2022, 8:33am UTC](https://discourse.julialang.org/t/questions-on-neuralpde-jl/83846/5 "2022-07-07T08:33:50Z")

</div>

That’s indeed what I figured. I ended up doing this mistake because Flux and Lux are very similar in spelling and one is used in all the examples with a CPU and the other is used in the GPU example.

---

<div class="post-metadata">

**Author:** ![avikpal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/avikpal/32/6550_2.png) [@avikpal](https://discourse.julialang.org/u/avikpal)\
**Post date:** [July 7, 2022, 3:31pm UTC](https://discourse.julialang.org/t/questions-on-neuralpde-jl/83846/6 "2022-07-07T15:31:17Z")

</div>

> [@ChrisRackauckas](#):
>
> Calling `gpu` on a chain is a Flux idea, not a Lux idea (and this is part of the whole, Lux needs to support GPUs differently,

Sidenote: This will throw a depwarn post v0.4.7 and error in v0.5
