diff --git a/dgf/src/learning/ten_lines/BUILD b/dgf/src/learning/ten_lines/BUILD index dcac82c..bbdeb8b 100644 --- a/dgf/src/learning/ten_lines/BUILD +++ b/dgf/src/learning/ten_lines/BUILD @@ -294,6 +294,8 @@ py_test( ":common", # absl/testing:absltest dep, # absl/testing:parameterized dep, + "//dgf/src/util:log", + # jax dep, # numpy dep, ], ) diff --git a/dgf/src/learning/ten_lines/common.py b/dgf/src/learning/ten_lines/common.py index ec9499f..b35c9f3 100644 --- a/dgf/src/learning/ten_lines/common.py +++ b/dgf/src/learning/ten_lines/common.py @@ -444,3 +444,21 @@ def check_number_of_seeds( f" than the batch size ({batch_size}). Increase the number of" f" validation seed {key}s or decrease the batch size." ) + + +def log_jax_backend(verbose: int = 2) -> None: + """Logs the active JAX backend and issues a warning if running on CPU. + + Args: + verbose: The verbosity level. If >= 2, logs an info message with the + backend. + """ + backend = jax.default_backend() + if verbose >= 2: + log.info("Using %s JAX backend", backend) + + if backend.lower() == "cpu": + log.warning( + "Using CPU JAX backend. Training will be slow. Consider using a GPU" + " or TPU." + ) diff --git a/dgf/src/learning/ten_lines/common_test.py b/dgf/src/learning/ten_lines/common_test.py index 750359e..14f75e3 100644 --- a/dgf/src/learning/ten_lines/common_test.py +++ b/dgf/src/learning/ten_lines/common_test.py @@ -17,6 +17,8 @@ from absl.testing import absltest from absl.testing import parameterized from dgf.src.learning.ten_lines import common +from dgf.src.util import log +import jax import numpy as np @@ -125,6 +127,24 @@ def test_check_number_of_seeds_insufficient_validation_fails(self): batch_size=10, num_training=15, num_validation=5, key="edge" ) + def test_log_jax_backend(self): + with log.capture_logs(log_info=True, log_warning=True) as captured: + common.log_jax_backend(verbose=2) + self.assertTrue( + any( + "Using" in msg.text and "JAX backend" in msg.text + for msg in captured + ) + ) + if jax.default_backend().lower() == "cpu": + self.assertTrue( + any( + msg.severity == log.Severity.WARNING + and "Using CPU JAX backend" in msg.text + for msg in captured + ) + ) + if __name__ == "__main__": absltest.main() diff --git a/dgf/src/learning/ten_lines/link_prediction_train.py b/dgf/src/learning/ten_lines/link_prediction_train.py index 728d256..39c106f 100644 --- a/dgf/src/learning/ten_lines/link_prediction_train.py +++ b/dgf/src/learning/ten_lines/link_prediction_train.py @@ -434,8 +434,7 @@ def train_link_model( if diagnostic_dir is not None: fs.makedirs(diagnostic_dir) - if verbose >= 2: - log.info("Using %s JAX backend", jax.default_backend()) + common.log_jax_backend(verbose) if target_edgeset is None: if len(schema.edge_sets) == 1: diff --git a/dgf/src/learning/ten_lines/node_prediction_train.py b/dgf/src/learning/ten_lines/node_prediction_train.py index 9ae0acf..ae32179 100644 --- a/dgf/src/learning/ten_lines/node_prediction_train.py +++ b/dgf/src/learning/ten_lines/node_prediction_train.py @@ -279,8 +279,7 @@ def train_node_model( if diagnostic_dir is not None: fs.makedirs(diagnostic_dir) - if verbose >= 2: - log.info("Using %s JAX backend", jax.default_backend()) + common.log_jax_backend(verbose) if target_nodeset is None: if len(schema.node_sets) == 1: