diff --git a/dgf/src/learning/jax/flax_train.py b/dgf/src/learning/jax/flax_train.py index 7538f03..d8203b3 100644 --- a/dgf/src/learning/jax/flax_train.py +++ b/dgf/src/learning/jax/flax_train.py @@ -307,6 +307,8 @@ def train( early_stopping_monitor.EarlyStoppingMonitorConfig | None ) = None, early_stopping_keep_best_param: bool = True, + aot_compile: bool = True, + prefetch: bool | int = 2, ) -> TrainResult: """Trains a Flax module with a flexible and feature-rich training loop. @@ -375,6 +377,11 @@ def train( early_stopping_keep_best_param: If true, and if early stopping is enabled, returns the lowest loss. This option leads to better quality model but consumes more memory. + aot_compile: If True and `train_step` is a jitted function, compiles the + training step ahead of time before the training loop starts, separating + compilation time from training throughput measurements. + prefetch: Number of batches to prefetch asynchronously on a background + thread. If True, defaults to 2. If False or 0, prefetching is disabled. Returns: A `TrainResult` dataclass containing: @@ -419,13 +426,24 @@ def train( pass metric_writer = metric_writers.MultiWriter(writers) + if prefetch: + prefetch_size = 2 if isinstance(prefetch, bool) else int(prefetch) + if prefetch_size > 0: + dataset_iterator = util.prefetch_iterator( + dataset_iterator, prefetch_size=prefetch_size + ) + + first_batch = None if model_params is None: if dummy_data is None: log.info("Generate first batch to initialize model") + first_batch = next(dataset_iterator) if dummy_data_fn is not None: - dummy_data = dummy_data_fn(next(dataset_iterator)) + dummy_data = dummy_data_fn(first_batch) else: - dummy_data = next(dataset_iterator) + dummy_data = first_batch + else: + first_batch = dummy_data with util.print_timer("Create model variables", True): rng_key, model_key = jax.random.split(rng_key, 2) @@ -567,7 +585,19 @@ def run_valid_logs(step: int): if checkpoint_every_n_steps is not None: log.info("Will checkpoint model every %s step(s)", checkpoint_every_n_steps) - log.info("Start training. The first two steps are generally slow.") + if aot_compile and hasattr(train_step, "lower"): + sample_batch = first_batch if first_batch is not None else dummy_data + if sample_batch is not None: + with util.print_timer("Compile train step", True): + rng_key, aot_step_key = jax.random.split(state.rng_key, 2) + train_step = train_step.lower( + state.model_params, + state.opt_state, + sample_batch, + aot_step_key, + ).compile() + + log.info("Start training.") start_time = time.time() pbar = tqdm.tqdm( range(state.step + 1, num_train_steps + 1), @@ -579,6 +609,7 @@ def run_valid_logs(step: int): effective_num_train_steps = 0 for step in pbar: with jax.profiler.StepTraceAnnotation("train", step_num=step): + t_step_start = time.perf_counter() if max_training_time_seconds is not None: elapsed_time = time.time() - start_time if elapsed_time > max_training_time_seconds: @@ -708,6 +739,10 @@ def run_valid_logs(step: int): metric_writer.flush() log.info(f"Final metrics: {list_display_dict}") + + if hasattr(dataset_iterator, "close"): + dataset_iterator.close() + return TrainResult( model_params=state.model_params, opt_state=state.opt_state, diff --git a/dgf/src/learning/jax/flax_train_test.py b/dgf/src/learning/jax/flax_train_test.py index 511fa7e..b7d7b02 100644 --- a/dgf/src/learning/jax/flax_train_test.py +++ b/dgf/src/learning/jax/flax_train_test.py @@ -445,6 +445,29 @@ def train_step(params, opt_state, batch, rng_key): rng_key=jax.random.PRNGKey(42), ) + def test_aot_compile_and_prefetch(self): + @jax.jit + def train_step(params, opt_state, batch, rng_key): + return params, opt_state, {"loss": jnp.array(1.0)} + + model = SimpleModel(hidden_dim=8) + opt = optax.adam(1e-3) + + result = flax_train.train( + model=model, + opt=opt, + train_step=train_step, + dataset_iterator=dataset_iterator(num_steps=None), + dummy_data_fn=lambda x: x["data"], + num_train_steps=5, + rng_key=jax.random.PRNGKey(42), + aot_compile=True, + prefetch=2, + ) + self.assertEqual( + result.model_params["params"]["Dense_0"]["kernel"].shape, (8, 8) + ) + if __name__ == "__main__": absltest.main() diff --git a/dgf/src/learning/ten_lines/common.py b/dgf/src/learning/ten_lines/common.py index 7e8d7d3..4806ee1 100644 --- a/dgf/src/learning/ten_lines/common.py +++ b/dgf/src/learning/ten_lines/common.py @@ -115,6 +115,18 @@ def parse_timeseries_encoder( ) from exc +def enable_fast_compile() -> None: + """Configures XLA for fast compilation during development and testing. + + Disables expensive XLA GPU GEMM autotuning benchmarks. + """ + xla_flags = os.environ.get("XLA_FLAGS", "") + if "--xla_gpu_autotune_level" not in xla_flags: + os.environ["XLA_FLAGS"] = ( + xla_flags + " --xla_gpu_autotune_level=0" + ).strip() + + class TFFunctionInputFormat(enum.Enum): """Input format of a model exported with `to_tensorflow_function`. diff --git a/dgf/src/learning/ten_lines/link_prediction_train.py b/dgf/src/learning/ten_lines/link_prediction_train.py index f6f99da..03cb115 100644 --- a/dgf/src/learning/ten_lines/link_prediction_train.py +++ b/dgf/src/learning/ten_lines/link_prediction_train.py @@ -359,6 +359,7 @@ def train_link_model( source_sampling_plan: sampling_config_lib.SamplingPlan | None = None, target_sampling_plan: sampling_config_lib.SamplingPlan | None = None, early_stopping: bool | int = True, + fast_compile: bool = False, ) -> LinkPredictionModel: """Trains a supervised Graph Neural Network model for edge prediction. @@ -450,6 +451,10 @@ def train_link_model( early_stopping: If True, use early stopping with default parameters (patience=5). If an integer, use early stopping with the given patience. If False, do not use early stopping. + fast_compile: If True, optimizes compilation speed at the expense of slight + runtime execution speed. Useful for fast iteration, interactive debugging, + and unit tests. Specifically, this disables expensive XLA GEMM autotuning + (--xla_gpu_autotune_level=0). Returns: A LinkPredictionModel instance. @@ -460,6 +465,9 @@ def train_link_model( with log.capture_logs() as captured_logs: + if fast_compile: + common.enable_fast_compile() + architecture = common.parse_architecture(architecture) timeseries_encoder = common.parse_timeseries_encoder(timeseries_encoder) begin_train_time = time.time() @@ -750,42 +758,46 @@ def valid_dataset_iterator_fn() -> Iterator[Batch]: yield jax_sample_to_batch(batch) if cache_valid_dataset: - with util.print_timer("Caching validation dataset", verbose >= 1): - if verbose >= 2: - num_examples_to_cache = valid_dataset.num_edge_in_seed_edgeset() - num_batches_to_cache = ( - util.num_batches( - num_examples_to_cache, - batch_size=valid_dataset.batch_size, - drop_remainder=valid_dataset.drop_remainder, - ) - if num_examples_to_cache is not None - else None - ) - if num_valid_steps is not None: - if num_batches_to_cache is None: - num_batches_to_cache = num_valid_steps - else: - num_batches_to_cache = min( - num_batches_to_cache, num_valid_steps - ) - valid_dataset_list = list( - tqdm.tqdm( - valid_dataset_iterator_fn(), - total=num_batches_to_cache, - desc="Caching validation dataset", - ) + num_examples_to_cache = valid_dataset.num_edge_in_seed_edgeset() + num_batches_to_cache = ( + util.num_batches( + num_examples_to_cache, + batch_size=valid_dataset.batch_size, + drop_remainder=valid_dataset.drop_remainder, ) + if num_examples_to_cache is not None + else None + ) + if num_valid_steps is not None: + if num_batches_to_cache is None: + num_batches_to_cache = num_valid_steps else: - valid_dataset_list = list(valid_dataset_iterator_fn()) - - if verbose >= 1: - log.info( - "Number of cache validation batches: %d", len(valid_dataset_list) - ) + num_batches_to_cache = min(num_batches_to_cache, num_valid_steps) + raw_valid_dataset_iterator_fn = valid_dataset_iterator_fn + cached_valid_dataset_list: list[Batch] | None = None def cached_valid_dataset_iterator_fn() -> Iterator[Batch]: - yield from valid_dataset_list + nonlocal cached_valid_dataset_list + if cached_valid_dataset_list is None: + with util.print_timer("Caching validation dataset", verbose >= 1): + if verbose >= 2: + cached_valid_dataset_list = list( + tqdm.tqdm( + raw_valid_dataset_iterator_fn(), + total=num_batches_to_cache, + desc="Caching validation dataset", + ) + ) + else: + cached_valid_dataset_list = list( + raw_valid_dataset_iterator_fn() + ) + if verbose >= 1: + log.info( + "Number of cache validation batches: %d", + len(cached_valid_dataset_list), + ) + yield from cached_valid_dataset_list valid_dataset_iterator_fn = cached_valid_dataset_iterator_fn diff --git a/dgf/src/learning/ten_lines/node_prediction_train.py b/dgf/src/learning/ten_lines/node_prediction_train.py index 75ce30c..40cef75 100644 --- a/dgf/src/learning/ten_lines/node_prediction_train.py +++ b/dgf/src/learning/ten_lines/node_prediction_train.py @@ -47,6 +47,7 @@ from dgf.src.util import log from dgf.src.util import temporal as temporal_util from dgf.src.util import util + import jax import jax.numpy as jnp import jaxtyping @@ -189,6 +190,7 @@ def train_node_model( early_stopping: bool | int = True, evaluate_final_model: bool = True, padding_margin: float = 0.1, + fast_compile: bool = False, ) -> NodePredictionModel: """Trains a supervised Graph Neural Network model for node-level prediction. @@ -286,6 +288,10 @@ def train_node_model( the periodic evaluations done during training. padding_margin: Relative margin added to observed maximum node and edge counts when estimating static graph padding. + fast_compile: If True, optimizes compilation speed at the expense of slight + runtime execution speed. Useful for fast iteration, interactive debugging, + and unit tests. Specifically, this disables expensive XLA GEMM autotuning + (--xla_gpu_autotune_level=0). Returns: A trained `NodePredictionModel` instance. @@ -300,6 +306,9 @@ def train_node_model( with log.capture_logs() as captured_logs: + if fast_compile: + common.enable_fast_compile() + architecture = common.parse_architecture(architecture) timeseries_encoder = common.parse_timeseries_encoder(timeseries_encoder) begin_train_time = time.time() @@ -688,38 +697,49 @@ def valid_dataset_iterator_fn(): num_batches_to_cache = num_valid_steps else: num_batches_to_cache = min(num_batches_to_cache, num_valid_steps) - valid_dataset_list = [] - try: - with util.print_timer("Caching validation dataset", verbose >= 1): - valid_iter = valid_dataset_iterator_fn() - if verbose >= 2: - valid_iter = tqdm.tqdm( - valid_iter, - total=num_batches_to_cache, - desc="Caching validation dataset", + raw_valid_dataset_iterator_fn = valid_dataset_iterator_fn + cached_valid_dataset_list = None + valid_dataset_caching_failed = False + + def cached_valid_dataset_iterator_fn(): + nonlocal cached_valid_dataset_list, valid_dataset_caching_failed + if valid_dataset_caching_failed: + yield from raw_valid_dataset_iterator_fn() + return + if cached_valid_dataset_list is None: + valid_dataset_list = [] + try: + with util.print_timer("Caching validation dataset", verbose >= 1): + valid_iter = raw_valid_dataset_iterator_fn() + if verbose >= 2: + valid_iter = tqdm.tqdm( + valid_iter, + total=num_batches_to_cache, + desc="Caching validation dataset", + ) + for batch in valid_iter: + valid_dataset_list.append(batch) + except (jax.errors.JaxRuntimeError, RuntimeError, ValueError) as e: + valid_dataset_list.clear() + if "RESOURCE_EXHAUSTED" not in str(e): + raise + log.warning( + "Out of device memory while caching validation dataset (%s);" + " falling back to uncached validation dataset.", + e, ) - for batch in valid_iter: - valid_dataset_list.append(batch) - - if verbose >= 1: - log.info( - "Number of cache validation batches: %d", - len(valid_dataset_list), - ) + valid_dataset_caching_failed = True + yield from raw_valid_dataset_iterator_fn() + return + if verbose >= 1: + log.info( + "Number of cache validation batches: %d", + len(valid_dataset_list), + ) + cached_valid_dataset_list = valid_dataset_list + yield from cached_valid_dataset_list - def cached_valid_dataset_iterator_fn(): - yield from valid_dataset_list - - valid_dataset_iterator_fn = cached_valid_dataset_iterator_fn - except (jax.errors.JaxRuntimeError, RuntimeError, ValueError) as e: - valid_dataset_list.clear() - if "RESOURCE_EXHAUSTED" not in str(e): - raise - log.warning( - "Out of device memory while caching validation dataset (%s);" - " falling back to uncached validation dataset.", - e, - ) + valid_dataset_iterator_fn = cached_valid_dataset_iterator_fn def valid_step(params, opt_state, batch: Batch): graph, seed_node_idxs = batch diff --git a/dgf/src/util/util.py b/dgf/src/util/util.py index f7b892d..88281e5 100644 --- a/dgf/src/util/util.py +++ b/dgf/src/util/util.py @@ -17,7 +17,10 @@ from collections.abc import Iterator import contextlib import math +import queue +import threading import time +from typing import Any, TypeVar from dgf.src.util import log import numpy as np @@ -311,3 +314,109 @@ def indent_string(s: str, num_spaces: int = 8) -> str: """Indents a multi-line string by num_spaces.""" indent = " " * num_spaces return s.strip().replace("\n", f"\n{indent}") + + +T = TypeVar("T") + + +class BackgroundPrefetcher(Iterator[T]): + """Prefetches items from an iterator in a background thread. + + Up to `prefetch_size` items are computed ahead of time, so that the consumer + does not wait on the source iterator (e.g. while it loads data from disk). + + Usage example: + + ```python + def load_batches(): + for path in paths: + yield read_batch(path) # Slow I/O. + + for batch in BackgroundPrefetcher(load_batches(), prefetch_size=4): + train_step(batch) # Runs while the next batches are being loaded. + ``` + + Exceptions raised by the source iterator are re-raised in the consumer thread + by `__next__`. Call `close()` to stop the background thread early (e.g. when + breaking out of the loop); it is also called automatically on exhaustion. + """ + + _SENTINEL = object() + + def __init__(self, iterator: Iterator[T], prefetch_size: int = 2): + if prefetch_size <= 0: + raise ValueError(f"`prefetch_size` must be positive, got {prefetch_size}") + self._iterator = iterator + self._prefetch_size = prefetch_size + self._queue: queue.Queue[Any] = queue.Queue(maxsize=prefetch_size) + self._stop_event = threading.Event() + self._thread = threading.Thread(target=self._worker, daemon=True) + self._thread.start() + + def _put(self, item: Any) -> bool: + """Puts `item` in the queue, retrying while the queue is full. + + Returns: + True if the item was enqueued, False if the prefetcher was stopped before + the item could be enqueued. + """ + while not self._stop_event.is_set(): + try: + self._queue.put(item, timeout=0.1) + return True + except queue.Full: + pass + return False + + def _worker(self): + try: + for item in self._iterator: + if not self._put(item): + return + self._put(self._SENTINEL) + except Exception as e: # pylint: disable=broad-exception-caught + self._put(e) + + def __iter__(self) -> Iterator[T]: + return self + + def __next__(self) -> T: + if self._stop_event.is_set(): + raise StopIteration + item = self._queue.get() + if item is self._SENTINEL: + self.close() + raise StopIteration + if isinstance(item, Exception): + self.close() + raise item + return item + + def close(self): + """Stops the background prefetcher thread and drains queued items.""" + self._stop_event.set() + while not self._queue.empty(): + try: + self._queue.get_nowait() + except queue.Empty: + break + + def __del__(self): + self.close() + + +def prefetch_iterator( + iterator: Iterator[T], prefetch_size: int = 2 +) -> Iterator[T]: + """Wraps an iterator to prefetch elements asynchronously on a background thread. + + Args: + iterator: The source iterator. + prefetch_size: The number of elements to prefetch ahead in the buffer. + + Returns: + An iterator that yields prefetched elements. + """ + if prefetch_size <= 0: + return iterator + return BackgroundPrefetcher(iterator, prefetch_size=prefetch_size) diff --git a/dgf/src/util/util_test.py b/dgf/src/util/util_test.py index 1530ec0..274920a 100644 --- a/dgf/src/util/util_test.py +++ b/dgf/src/util/util_test.py @@ -320,5 +320,31 @@ def test_split_temporal_invalid_inputs( ) +class PrefetchIteratorTest(absltest.TestCase): + + def test_prefetch_basic(self): + items = list(range(10)) + prefetched = list(util.prefetch_iterator(iter(items), prefetch_size=2)) + self.assertEqual(prefetched, items) + + def test_prefetch_disabled(self): + items = list(range(5)) + self.assertEqual( + list(util.prefetch_iterator(iter(items), prefetch_size=0)), items + ) + + def test_prefetch_error_propagation(self): + def failing_iterator(): + yield 1 + yield 2 + raise RuntimeError("Failure in iterator") + + it = util.prefetch_iterator(failing_iterator(), prefetch_size=2) + self.assertEqual(next(it), 1) + self.assertEqual(next(it), 2) + with self.assertRaises(RuntimeError): + next(it) + + if __name__ == "__main__": absltest.main()