From 28eaeeba8b46a8f9cc15d078fcda212ea0364f77 Mon Sep 17 00:00:00 2001 From: RecML authors Date: Fri, 14 Aug 2026 01:09:42 -0700 Subject: [PATCH] Internal change PiperOrigin-RevId: 964551370 --- recml/core/training/keras_trainer.py | 22 +++++++++++++++++----- 1 file changed, 17 insertions(+), 5 deletions(-) diff --git a/recml/core/training/keras_trainer.py b/recml/core/training/keras_trainer.py index c8e121e..6e2e88f 100644 --- a/recml/core/training/keras_trainer.py +++ b/recml/core/training/keras_trainer.py @@ -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. @@ -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) @@ -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 @@ -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 @@ -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: @@ -338,6 +346,8 @@ 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, @@ -345,7 +355,7 @@ def train_and_evaluate(self, task: KerasTask) -> core.Logs: 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) @@ -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) @@ -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, )