Skip to content
Merged
Show file tree
Hide file tree
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
41 changes: 38 additions & 3 deletions dgf/src/learning/jax/flax_train.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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),
Expand All @@ -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:
Expand Down Expand Up @@ -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,
Expand Down
23 changes: 23 additions & 0 deletions dgf/src/learning/jax/flax_train_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
12 changes: 12 additions & 0 deletions dgf/src/learning/ten_lines/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -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`.

Expand Down
76 changes: 44 additions & 32 deletions dgf/src/learning/ten_lines/link_prediction_train.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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.
Expand All @@ -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()
Expand Down Expand Up @@ -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

Expand Down
80 changes: 50 additions & 30 deletions dgf/src/learning/ten_lines/node_prediction_train.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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.
Expand All @@ -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()
Expand Down Expand Up @@ -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
Expand Down
Loading
Loading