11mutable 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
109function 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)
3835end
3936
4037function 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)
7466end
7567
0 commit comments