Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 17 additions & 5 deletions recml/core/training/keras_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,10 @@ def create_model(self, **kwargs) -> keras.Model:
A Keras model instance.
"""

def get_custom_callbacks(self) -> list[keras.callbacks.Callback]:
"""Returns the custom callbacks."""
return []

def create_model_for_eval(self, **kwargs) -> keras.Model:
"""Creates a Keras model for evaluation.

Expand Down Expand Up @@ -273,11 +277,13 @@ def train(self, task: KerasTask) -> core.Logs:
**self._maybe_get_model_kws(task, dataset, val=False)
)

callbacks = task.get_custom_callbacks() + self.train_callbacks

history = model.fit(
dataset,
epochs=self._train_epochs + 1,
steps_per_epoch=self._steps_per_loop,
callbacks=self.train_callbacks,
callbacks=callbacks,
initial_epoch=1,
)
model.summary(print_fn=logging.info)
Expand All @@ -296,6 +302,8 @@ def evaluate(self, task: KerasTask) -> core.Logs:
**self._maybe_get_model_kws(task, dataset, val=True)
)

callbacks = task.get_custom_callbacks() + self.eval_callbacks

if keras.backend.backend() == "jax":
[tb_cbk] = [
cbk
Expand All @@ -306,7 +314,7 @@ def evaluate(self, task: KerasTask) -> core.Logs:
history = model.evaluate(
dataset,
steps=self._steps_per_eval,
callbacks=self.eval_callbacks,
callbacks=callbacks,
return_dict=True,
)
epoch_dt = time.time() - epoch_start_time
Expand All @@ -319,7 +327,7 @@ def evaluate(self, task: KerasTask) -> core.Logs:
return model.evaluate(
dataset,
steps=self._steps_per_eval,
callbacks=self.eval_callbacks,
callbacks=callbacks,
)

def train_and_evaluate(self, task: KerasTask) -> core.Logs:
Expand All @@ -338,14 +346,16 @@ def train_and_evaluate(self, task: KerasTask) -> core.Logs:
**self._maybe_get_model_kws(task, train_dataset, val=False)
)

callbacks = task.get_custom_callbacks() + self.train_callbacks

history = model.fit(
train_dataset,
validation_data=eval_dataset,
epochs=self._train_epochs + 1,
steps_per_epoch=self._steps_per_loop,
# Explicitly set to None for deterministic evaluation.
validation_steps=None,
callbacks=self.train_callbacks,
callbacks=callbacks,
initial_epoch=1,
)
model.summary(print_fn=logging.info)
Expand All @@ -371,6 +381,8 @@ def evaluate_continuously(self, task: KerasTask) -> core.Logs | None:
**self._maybe_get_model_kws(task, eval_dataset, val=True),
)

callbacks = task.get_custom_callbacks() + self.eval_callbacks

def timeout_fn() -> bool:
return tf.io.gfile.exists(self._marker_path)

Expand Down Expand Up @@ -425,7 +437,7 @@ def on_test_begin(self, logs: Mapping[str, Any] | None = None):
history = model.evaluate(
eval_dataset,
steps=self._steps_per_eval,
callbacks=[restore_callback] + self.eval_callbacks,
callbacks=[restore_callback] + callbacks,
return_dict=True,
)

Expand Down
Loading