Skip to content

Commit af1ae16

Browse files
authored
refactor: extract shared train_loop for neural estimators (#66)
1 parent 41c06f8 commit af1ae16

8 files changed

Lines changed: 324 additions & 389 deletions

File tree

sbijax/_src/fmpe.py

Lines changed: 33 additions & 72 deletions
Original file line numberDiff line numberDiff line change
@@ -1,16 +1,13 @@
11
import jax
2-
import numpy as np
32
import optax
4-
from absl import logging
53
from jax import Array
64
from jax import numpy as jnp
75
from jax import random as jr
86
from jax._src.flatten_util import ravel_pytree
9-
from tqdm import tqdm
107

118
from sbijax._src._ne_base import NE
129
from 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
1411
from 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(

sbijax/_src/nass.py

Lines changed: 15 additions & 56 deletions
Original file line numberDiff line numberDiff line change
@@ -1,17 +1,13 @@
1-
from functools import partial
21
from typing import Any
32

43
import jax
5-
import numpy as np
64
import optax
7-
from absl import logging
85
from jax import numpy as jnp
96
from jax import random as jr
10-
from tqdm import tqdm
117

128
from sbijax._src._ne_base import NE
139
from sbijax._src.util.dataloader import as_numpy_iterator_from_slices
14-
from sbijax._src.util.early_stopping import EarlyStopping
10+
from sbijax._src.util.train import train_loop
1511

1612

1713
def _jsd_summary_loss(params, rng, apply_fn, **batch):
@@ -136,7 +132,6 @@ def _summarize(batch):
136132

137133
return ret_summaries
138134

139-
# pylint: disable=undefined-loop-variable
140135
def _fit_summary_net(
141136
self,
142137
rng_key,
@@ -148,62 +143,26 @@ def _fit_summary_net(
148143
):
149144
init_key, rng_key = jr.split(rng_key)
150145
params = self._init_summary_net_params(init_key, **next(iter(train_iter)))
151-
state = optimizer.init(params)
152-
loss_fn = jax.jit(partial(_jsd_summary_loss, apply_fn=self.model.apply))
153146

154-
@jax.jit
155-
def step(rng, params, state, **batch):
156-
loss, grads = jax.value_and_grad(loss_fn)(params, rng, **batch)
157-
updates, new_state = optimizer.update(grads, state, params)
158-
new_params = optax.apply_updates(params, updates)
159-
return loss, new_params, new_state
160-
161-
losses = np.zeros([n_iter, 2])
162-
early_stop = EarlyStopping(1e-3, n_early_stopping_patience)
163-
best_params, best_loss = None, np.inf
164-
logging.info("training summary net")
165-
for i in tqdm(range(n_iter)):
166-
train_loss = 0.0
167-
epoch_key, rng_key = jr.split(rng_key)
168-
for j, batch in enumerate(train_iter):
169-
batch_loss, params, state = step(
170-
jr.fold_in(epoch_key, j), params, state, **batch
171-
)
172-
train_loss += batch_loss * (
173-
batch["y"].shape[0] / train_iter.num_samples
174-
)
175-
val_key, rng_key = jr.split(rng_key)
176-
validation_loss = self._summary_validation_loss(params, val_key, val_iter)
177-
losses[i] = jnp.array([train_loss, validation_loss])
178-
179-
_, early_stop = early_stop.update(validation_loss)
180-
if early_stop.should_stop:
181-
logging.info("early stopping criterion found")
182-
break
183-
if validation_loss < best_loss:
184-
best_loss = validation_loss
185-
best_params = params.copy()
186-
187-
stacked_losses = jnp.vstack(losses)[: (i + 1), :]
188-
return best_params, stacked_losses
147+
def loss_fn(params, rng, **batch):
148+
return _jsd_summary_loss(params, rng, self.model.apply, **batch)
149+
150+
return train_loop(
151+
rng_key,
152+
params=params,
153+
optimizer=optimizer,
154+
loss_fn=loss_fn,
155+
validation_loss_fn=loss_fn,
156+
train_iter=train_iter,
157+
val_iter=val_iter,
158+
n_iter=n_iter,
159+
n_early_stopping_patience=n_early_stopping_patience,
160+
)
189161

190162
def _init_summary_net_params(self, rng_key, **init_data):
191163
params = self.model.init(rng_key, method="forward", **init_data)
192164
return params
193165

194-
def _summary_validation_loss(self, params, rng_key, val_iter):
195-
loss_fn = jax.jit(partial(_jsd_summary_loss, apply_fn=self.model.apply))
196-
197-
def body_fn(batch_key, **batch):
198-
loss = loss_fn(params, batch_key, **batch)
199-
return loss * (batch["y"].shape[0] / val_iter.num_samples)
200-
201-
losses = 0.0
202-
for batch in val_iter:
203-
batch_key, rng_key = jr.split(rng_key)
204-
losses += body_fn(batch_key, **batch)
205-
return losses
206-
207166
def simulate_data(
208167
self,
209168
rng_key,

sbijax/_src/nasss.py

Lines changed: 16 additions & 58 deletions
Original file line numberDiff line numberDiff line change
@@ -1,14 +1,9 @@
1-
from functools import partial
2-
31
import jax
4-
import numpy as np
5-
import optax
6-
from absl import logging
72
from jax import numpy as jnp
83
from jax import random as jr
94

105
from sbijax._src.nass import NASS
11-
from sbijax._src.util.early_stopping import EarlyStopping
6+
from sbijax._src.util.train import train_loop
127

138

149
def _sample_unit_sphere(rng_key, n, dim):
@@ -94,7 +89,6 @@ class NASSS(NASS):
9489
def __init__(self, model_fns, summary_net):
9590
super().__init__(model_fns, summary_net)
9691

97-
# pylint: disable=undefined-loop-variable
9892
def _fit_summary_net(
9993
self,
10094
rng_key,
@@ -106,54 +100,18 @@ def _fit_summary_net(
106100
):
107101
init_key, rng_key = jr.split(rng_key)
108102
params = self._init_summary_net_params(init_key, **next(iter(train_iter)))
109-
state = optimizer.init(params)
110-
loss_fn = jax.jit(partial(_jsd_summary_loss, apply_fn=self.model.apply))
111-
112-
@jax.jit
113-
def step(rng, params, state, **batch):
114-
loss, grads = jax.value_and_grad(loss_fn)(params, rng, **batch)
115-
updates, new_state = optimizer.update(grads, state, params)
116-
new_params = optax.apply_updates(params, updates)
117-
return loss, new_params, new_state
118-
119-
losses = np.zeros([n_iter, 2])
120-
early_stop = EarlyStopping(1e-3, n_early_stopping_patience)
121-
best_params, best_loss = None, np.inf
122-
logging.info("training summary net")
123-
for i in range(n_iter):
124-
train_loss = 0.0
125-
epoch_key, rng_key = jr.split(rng_key)
126-
for j, batch in enumerate(train_iter):
127-
batch_loss, params, state = step(
128-
jr.fold_in(epoch_key, j), params, state, **batch
129-
)
130-
train_loss += batch_loss * (
131-
batch["y"].shape[0] / train_iter.num_samples
132-
)
133-
val_key, rng_key = jr.split(rng_key)
134-
validation_loss = self._summary_validation_loss(params, val_key, val_iter)
135-
losses[i] = jnp.array([train_loss, validation_loss])
136-
137-
_, early_stop = early_stop.update(validation_loss)
138-
if early_stop.should_stop:
139-
logging.info("early stopping criterion found")
140-
break
141-
if validation_loss < best_loss:
142-
best_loss = validation_loss
143-
best_params = params.copy()
144-
145-
stacked_losses = jnp.vstack(losses)[: (i + 1), :]
146-
return best_params, stacked_losses
147-
148-
def _summary_validation_loss(self, params, rng_key, val_iter):
149-
loss_fn = jax.jit(partial(_jsd_summary_loss, apply_fn=self.model.apply))
150-
151-
def body_fn(batch_key, **batch):
152-
loss = loss_fn(params, batch_key, **batch)
153-
return loss * (batch["y"].shape[0] / val_iter.num_samples)
154-
155-
losses = 0.0
156-
for batch in val_iter:
157-
batch_key, rng_key = jr.split(rng_key)
158-
losses += body_fn(batch_key, **batch)
159-
return losses
103+
104+
def loss_fn(params, rng, **batch):
105+
return _jsd_summary_loss(params, rng, self.model.apply, **batch)
106+
107+
return train_loop(
108+
rng_key,
109+
params=params,
110+
optimizer=optimizer,
111+
loss_fn=loss_fn,
112+
validation_loss_fn=loss_fn,
113+
train_iter=train_iter,
114+
val_iter=val_iter,
115+
n_iter=n_iter,
116+
n_early_stopping_patience=n_early_stopping_patience,
117+
)

0 commit comments

Comments
 (0)