Skip to content

Commit adba6e1

Browse files
committed
Added Clipnorm to the MC gradient descent
1 parent 115ea74 commit adba6e1

2 files changed

Lines changed: 29 additions & 9 deletions

File tree

colibri/monte_carlo_fit.py

Lines changed: 19 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616

1717
from colibri.data_batch import data_batches
1818
from colibri.mc_utils import write_exportgrid_mc
19+
from colibri.optax_optimizer import optimizer_provider
1920

2021
log = 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)

colibri/optax_optimizer.py

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@
1414

1515

1616
def optimizer_provider(
17-
optimizer="adam", optimizer_settings={}
17+
optimizer="adam", optimizer_settings={}, clipnorm=6.073e-6
1818
) -> optax._src.base.GradientTransformationExtraArgs:
1919
"""
2020
Define the optimizer.
@@ -27,6 +27,9 @@ def optimizer_provider(
2727
optimizer_settings : dict, default = {}
2828
Dictionary containing the optimizer settings.
2929
30+
clipnorm : float, optional
31+
If not None, gradients will be clipped to this norm.
32+
3033
Returns
3134
-------
3235
optax._src.base.GradientTransformationExtraArgs
@@ -37,9 +40,13 @@ def optimizer_provider(
3740
if not "learning_rate" in optimizer_settings.keys():
3841
optimizer_settings["learning_rate"] = 5e-4
3942

40-
opt = getattr(optax, optimizer)
43+
opt = getattr(optax, optimizer)(**optimizer_settings)
44+
45+
if clipnorm is not None:
46+
log.info(f"Using clipnorm: {clipnorm}")
47+
return optax.chain(optax.clip_by_global_norm(clipnorm), opt)
4148

42-
return opt(**optimizer_settings)
49+
return opt
4350

4451

4552
def early_stopper(

0 commit comments

Comments
 (0)