1616
1717from colibri .data_batch import data_batches
1818from colibri .mc_utils import write_exportgrid_mc
19+ from colibri .optax_optimizer import optimizer_provider
1920
2021log = logging .getLogger (__name__ )
2122
@@ -52,10 +53,12 @@ def monte_carlo_fit(
5253 len_trval_data ,
5354 pdf_model ,
5455 mc_initial_parameters ,
55- optimizer_provider ,
5656 early_stopper ,
5757 max_epochs ,
5858 FIT_XGRID ,
59+ optimizer = "adam" ,
60+ optimizer_settings = {},
61+ clipnorm = 6.073e-6 ,
5962 batch_size = None ,
6063 batch_seed = 1 ,
6164 alpha = 1e-7 ,
@@ -86,9 +89,6 @@ def monte_carlo_fit(
8689 mc_initial_parameters: jnp.array
8790 Initial parameters for the Monte Carlo fit.
8891
89- optimizer_provider: optax._src.base.GradientTransformationExtraArgs
90- Optax optimizer.
91-
9292 early_stopper: flax.training.early_stopping.EarlyStopping
9393 Early stopping criteria.
9494
@@ -99,6 +99,15 @@ def monte_carlo_fit(
9999 xgrid of the theory, computed by a production rule by taking
100100 the sorted union of the xgrids of the datasets entering the fit.
101101
102+ optimizer: str, default="adam"
103+ The optimizer to use.
104+
105+ optimizer_settings: dict, default={}
106+ A dictionary with the settings for the optimizer.
107+
108+ clipnorm: float, optional
109+ If not None, gradients will be clipped to this norm.
110+
102111 batch_size: int, default is None which sets it to the full size of data
103112 Size of batches during training.
104113
@@ -121,6 +130,10 @@ def monte_carlo_fit(
121130
122131 pred_and_pdf = pdf_model .pred_and_pdf_func (FIT_XGRID , forward_map = _pred_data )
123132
133+ optim = optimizer_provider (
134+ optimizer = optimizer , optimizer_settings = optimizer_settings , clipnorm = clipnorm
135+ )
136+
124137 @jax .jit
125138 def loss_training (
126139 parameters ,
@@ -173,7 +186,7 @@ def step(
173186 alpha ,
174187 lambda_positivity ,
175188 )
176- updates , opt_state = optimizer_provider .update (grads , opt_state , params )
189+ updates , opt_state = optim .update (grads , opt_state , params )
177190 params = optax .apply_updates (params , updates )
178191 return params , opt_state , loss_value
179192
@@ -189,7 +202,7 @@ def step(
189202 loss = []
190203 val_loss = []
191204
192- opt_state = optimizer_provider .init (mc_initial_parameters )
205+ opt_state = optim .init (mc_initial_parameters )
193206 parameters = mc_initial_parameters .copy ()
194207
195208 data_batch = data_batches (len_tr_idx , batch_size , batch_seed )
0 commit comments