Skip to content

Commit 4d09e3f

Browse files
committed
move optimizer in MLJ models to sol_kwargs
1 parent 1d116b6 commit 4d09e3f

3 files changed

Lines changed: 25 additions & 43 deletions

File tree

examples/usage.jl

Lines changed: 5 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -80,18 +80,16 @@ if !isfile(icnf_mach_fn)
8080
df = DataFrame(permutedims(r), :auto)
8181
model = ICNFModel(;
8282
icnf,
83-
optimizers = (
84-
OptimiserChain(
85-
WeightDecay(; lambda = 1.0e-4),
86-
ClipNorm(10.0, 2.0; throw = true),
87-
Adam(; eta = 0.001, beta = (0.9, 0.999), epsilon = 1.0e-8),
88-
),
89-
),
9083
batchsize = 1024,
9184
adtype = AutoZygote(),
9285
sol_kwargs = (;
9386
epochs = 300,
9487
callback = opt_callback,
88+
alg = OptimiserChain(
89+
WeightDecay(; lambda = 1.0e-4),
90+
ClipNorm(10.0, 2.0; throw = true),
91+
Adam(; eta = 0.001, beta = (0.9, 0.999), epsilon = 1.0e-8),
92+
),
9593
progress = true,
9694
verbose = Detailed(),
9795
), # pass to the solver

src/exts/mlj_ext/core_cond_icnf.jl

Lines changed: 10 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
11
mutable struct CondICNFModel{AICNF <: AbstractICNF} <: MLJICNF{AICNF}
22
icnf::AICNF
33
loss::Function
4-
optimizers::Tuple
54
batchsize::Int
65
adtype::ADTypes.AbstractADType
76
sol_kwargs::NamedTuple
@@ -10,8 +9,12 @@ end
109
function CondICNFModel(;
1110
icnf::AbstractICNF = ICNF(),
1211
loss::Function = loss,
13-
optimizers::Tuple = (
14-
Optimisers.OptimiserChain(
12+
batchsize::Int = 1024,
13+
adtype::ADTypes.AbstractADType = ADTypes.AutoZygote(),
14+
sol_kwargs::NamedTuple = (;
15+
epochs = 300,
16+
callback = make_opt_callback(64),
17+
alg = Optimisers.OptimiserChain(
1518
Optimisers.WeightDecay(; lambda = convert(eltype(icnf), 1.0e-4)),
1619
Optimisers.ClipNorm(
1720
convert(eltype(icnf), 10.0),
@@ -24,17 +27,11 @@ function CondICNFModel(;
2427
epsilon = convert(eltype(icnf), 1.0e-8),
2528
),
2629
),
27-
),
28-
batchsize::Int = 1024,
29-
adtype::ADTypes.AbstractADType = ADTypes.AutoZygote(),
30-
sol_kwargs::NamedTuple = (;
31-
epochs = 300,
32-
callback = make_opt_callback(64),
3330
progress = true,
3431
verbose = SciMLLogging.Detailed(),
3532
),
3633
)
37-
return CondICNFModel(icnf, loss, optimizers, batchsize, adtype, sol_kwargs)
34+
return CondICNFModel(icnf, loss, batchsize, adtype, sol_kwargs)
3835
end
3936

4037
function MLJModelInterface.fit(model::CondICNFModel, verbosity, XY)
@@ -59,17 +56,12 @@ function MLJModelInterface.fit(model::CondICNFModel, verbosity, XY)
5956
data;
6057
model.sol_kwargs...,
6158
)
62-
res_stats = Any[]
63-
for opt in model.optimizers
64-
optprob_re = SciMLBase.remake(optprob; u0 = ps)
65-
res = SciMLBase.solve(optprob_re, opt; model.sol_kwargs...)
66-
ps .= res.u
67-
push!(res_stats, res.stats)
68-
end
59+
res = SciMLBase.solve(optprob; model.sol_kwargs...)
60+
ps .= res.u
6961

7062
fitresult = (ps, st)
7163
cache = nothing
72-
report = (stats = res_stats,)
64+
report = (stats = res.stats,)
7365
return (fitresult, cache, report)
7466
end
7567

src/exts/mlj_ext/core_icnf.jl

Lines changed: 10 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
11
mutable struct ICNFModel{AICNF <: AbstractICNF} <: MLJICNF{AICNF}
22
icnf::AICNF
33
loss::Function
4-
optimizers::Tuple
54
batchsize::Int
65
adtype::ADTypes.AbstractADType
76
sol_kwargs::NamedTuple
@@ -10,8 +9,12 @@ end
109
function ICNFModel(;
1110
icnf::AbstractICNF = ICNF(),
1211
loss::Function = loss,
13-
optimizers::Tuple = (
14-
Optimisers.OptimiserChain(
12+
batchsize::Int = 1024,
13+
adtype::ADTypes.AbstractADType = ADTypes.AutoZygote(),
14+
sol_kwargs::NamedTuple = (;
15+
epochs = 300,
16+
callback = make_opt_callback(64),
17+
alg = Optimisers.OptimiserChain(
1518
Optimisers.WeightDecay(; lambda = convert(eltype(icnf), 1.0e-4)),
1619
Optimisers.ClipNorm(
1720
convert(eltype(icnf), 10.0),
@@ -24,17 +27,11 @@ function ICNFModel(;
2427
epsilon = convert(eltype(icnf), 1.0e-8),
2528
),
2629
),
27-
),
28-
batchsize::Int = 1024,
29-
adtype::ADTypes.AbstractADType = ADTypes.AutoZygote(),
30-
sol_kwargs::NamedTuple = (;
31-
epochs = 300,
32-
callback = make_opt_callback(64),
3330
progress = true,
3431
verbose = SciMLLogging.Detailed(),
3532
),
3633
)
37-
return ICNFModel(icnf, loss, optimizers, batchsize, adtype, sol_kwargs)
34+
return ICNFModel(icnf, loss, batchsize, adtype, sol_kwargs)
3835
end
3936

4037
function MLJModelInterface.fit(model::ICNFModel, verbosity, X)
@@ -56,17 +53,12 @@ function MLJModelInterface.fit(model::ICNFModel, verbosity, X)
5653
data;
5754
model.sol_kwargs...,
5855
)
59-
res_stats = Any[]
60-
for opt in model.optimizers
61-
optprob_re = SciMLBase.remake(optprob; u0 = ps)
62-
res = SciMLBase.solve(optprob_re, opt; model.sol_kwargs...)
63-
ps .= res.u
64-
push!(res_stats, res.stats)
65-
end
56+
res = SciMLBase.solve(optprob; model.sol_kwargs...)
57+
ps .= res.u
6658

6759
fitresult = (ps, st)
6860
cache = nothing
69-
report = (stats = res_stats,)
61+
report = (stats = res.stats,)
7062
return (fitresult, cache, report)
7163
end
7264

0 commit comments

Comments
 (0)