11import jax
2- import numpy as np
32import optax
4- from absl import logging
53from jax import Array
64from jax import numpy as jnp
75from jax import random as jr
86from jax ._src .flatten_util import ravel_pytree
9- from tqdm import tqdm
107
118from sbijax ._src ._ne_base import NE
129from sbijax ._src .util .data import as_inference_data
13- from sbijax ._src .util .early_stopping import EarlyStopping
10+ from sbijax ._src .util .train import train_loop
1411from sbijax ._src .util .types import PyTree
1512
1613
@@ -107,68 +104,19 @@ def _fit_model_single_round(
107104 ):
108105 init_key , seed = jr .split (seed )
109106 params = self ._init_params (init_key , ** next (iter (train_iter )))
110- state = optimizer .init (params )
111-
112- @jax .jit
113- def step (params , rng , state , ** batch ):
114- def loss_fn (params , rng , ** batch ):
115- lp = self .model .apply (
116- params ,
117- rng = rng ,
118- method = "loss" ,
119- inputs = batch ["theta" ],
120- context = batch ["y" ],
121- is_training = True ,
122- )
123- return jnp .mean (lp )
124-
125- loss , grads = jax .value_and_grad (loss_fn )(params , rng , ** batch )
126- updates , new_state = optimizer .update (grads , state , params )
127- new_params = optax .apply_updates (params , updates )
128- return loss , new_params , new_state
129-
130- losses = np .zeros ([n_iter , 2 ])
131- early_stop = EarlyStopping (
132- n_early_stopping_delta , n_early_stopping_patience
133- )
134- best_params , best_loss = None , np .inf
135- logging .info ("training model" )
136- for i in tqdm (range (n_iter )):
137- train_loss = 0.0
138- rng_key = jr .fold_in (seed , i )
139- for batch in train_iter :
140- train_key , rng_key = jr .split (rng_key )
141- batch_loss , params , state = step (params , train_key , state , ** batch )
142- train_loss += batch_loss * (
143- batch ["y" ].shape [0 ] / train_iter .num_samples
144- )
145- val_key , rng_key = jr .split (rng_key )
146- validation_loss = self ._validation_loss (val_key , params , val_iter )
147- losses [i ] = jnp .array ([train_loss , validation_loss ])
148-
149- _ , early_stop = early_stop .update (validation_loss )
150- if early_stop .should_stop :
151- logging .info ("early stopping criterion found" )
152- break
153- if validation_loss < best_loss :
154- best_loss = validation_loss
155- best_params = params .copy ()
156-
157- stacked_losses = jnp .vstack (losses )[: (i + 1 ), :]
158- return best_params , stacked_losses
159-
160- def _init_params (self , rng_key , ** init_data ):
161- params = self .model .init (
162- rng_key ,
163- method = "loss" ,
164- inputs = init_data ["theta" ],
165- context = init_data ["y" ],
166- is_training = False ,
167- )
168- return params
169107
170- def _validation_loss (self , rng_key , params , val_iter ):
171108 def loss_fn (params , rng , ** batch ):
109+ lp = self .model .apply (
110+ params ,
111+ rng = rng ,
112+ method = "loss" ,
113+ inputs = batch ["theta" ],
114+ context = batch ["y" ],
115+ is_training = True ,
116+ )
117+ return jnp .mean (lp )
118+
119+ def validation_loss_fn (params , rng , ** batch ):
172120 lp = self .model .apply (
173121 params ,
174122 rng = rng ,
@@ -179,15 +127,28 @@ def loss_fn(params, rng, **batch):
179127 )
180128 return jnp .mean (lp )
181129
182- def body_fn (batch_key , ** batch ):
183- loss = loss_fn (params , batch_key , ** batch )
184- return loss * (batch ["y" ].shape [0 ] / val_iter .num_samples )
130+ return train_loop (
131+ seed ,
132+ params = params ,
133+ optimizer = optimizer ,
134+ loss_fn = loss_fn ,
135+ validation_loss_fn = validation_loss_fn ,
136+ train_iter = train_iter ,
137+ val_iter = val_iter ,
138+ n_iter = n_iter ,
139+ n_early_stopping_patience = n_early_stopping_patience ,
140+ n_early_stopping_delta = n_early_stopping_delta ,
141+ )
185142
186- loss = 0.0
187- for batch in val_iter :
188- val_key , rng_key = jr .split (rng_key )
189- loss += body_fn (val_key , ** batch )
190- return loss
143+ def _init_params (self , rng_key , ** init_data ):
144+ params = self .model .init (
145+ rng_key ,
146+ method = "loss" ,
147+ inputs = init_data ["theta" ],
148+ context = init_data ["y" ],
149+ is_training = False ,
150+ )
151+ return params
191152
192153 # ruff: noqa: D417
193154 def sample_posterior (
0 commit comments