| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| """This file contains function that coordinates the training and evaluation of the model.""" |
|
|
| import functools |
| import math |
| import os |
|
|
| from absl import logging |
| from clu import metric_writers |
| from flax import jax_utils |
| from flax import linen as nn |
| from flax.training import checkpoints |
| import jax |
| from jax import random |
| import jax.numpy as jnp |
| import numpy as np |
| import tensorflow as tf |
| import wandb |
|
|
| from train import data |
| from train import evaluater |
| from train import model |
| from train import trainer |
|
|
|
|
|
|
| def log_hyperparams_tb( |
| config, model_config, initial_variables, tf_summary_writer |
| ): |
| """Log hyperparameters to TensorBoard. |
| |
| Args: |
| config: experiment's ConfigDict |
| model_config: model's ConfigDict |
| initial_variables: initial hyperparameter values |
| tf_summary_writer: SummaryWriter object. |
| |
| Returns: |
| The SummaryWriter object and the config. |
| """ |
| |
| config.num_model_parameters = sum( |
| x.size for x in jax.tree_util.tree_leaves(initial_variables) |
| ) |
|
|
| |
| config_hyperparameters = [ |
| tf.convert_to_tensor([k, str(v)]) for k, v in config.items() |
| ] |
| model_config_hyperparameters = [ |
| tf.convert_to_tensor([k, str(v)]) |
| for k, v in model_config.__dict__.items() |
| ] |
|
|
| |
| with tf_summary_writer.as_default(): |
| tf.summary.text( |
| "Model hyperparameters", tf.stack(model_config_hyperparameters), step=0 |
| ) |
| tf.summary.text( |
| "Config hyperparameters", tf.stack(config_hyperparameters), step=0 |
| ) |
|
|
| return tf_summary_writer, config |
|
|
|
|
|
|
| def train_and_evaluate(config, workdir): |
| """The training and evaluation loops for the model. |
| |
| Args: |
| config: experiment's config dictionary. |
| workdir: directory to use for logging. |
| """ |
| |
| workdir = os.path.abspath(workdir) |
|
|
| logging.info("Creating training and evaluator dataset iterator") |
| curriculum = data.CurriculumState( |
| stage=int(getattr(config, "curriculum_start_stage", 1)), |
| max_stage=int(getattr(config, "curriculum_max_stage", 6))) |
| train_data_iter = data.create_iter( |
| config, config.minibatch_size, train=True, curriculum=curriculum) |
| eval_data_iter = data.create_iter(config, config.minibatch_size, train=False) |
|
|
| logging.info("Finished creating training dataset iterator") |
|
|
| model_config = model.TransformerConfig( |
| dtype=config.dtype, |
| vocab_size=config.vocab_size, |
| seq_len=config.seq_len, |
| num_heads=config.num_heads, |
| num_layers=config.num_layers, |
| emb_dim=config.emb_dim, |
| qkv_dim=config.qkv_dim, |
| mlp_dim=config.mlp_dim, |
| dropout_rate=config.dropout_rate, |
| attention_dropout_rate=config.attention_dropout_rate, |
| deterministic=False, |
| num_latent_slots=int(config.num_latent_slots), |
| inject_latents=bool(int(getattr(config, "recurrent_latent", 1))), |
| ) |
|
|
| logging.info("train_config: %s", str(model_config.__dict__)) |
| print(str(model_config.__dict__), flush=True) |
|
|
| rng = jax.random.PRNGKey(config.seed) |
| rng, init_rng, inference_rng = random.split(rng, num=3) |
|
|
| |
| |
| rng, dropout_rng = jax.random.split(rng) |
| input_shape = (config.minibatch_size, config.seq_len) |
| net = model.TransformerLMHeadModel(model_config) |
| rng_keys = {"params": init_rng, "dropout": dropout_rng} |
| K = int(config.num_latent_slots) |
| dummy_latents = jnp.zeros( |
| (config.minibatch_size, K, config.emb_dim), model_config.dtype) |
| dummy_positions = jnp.zeros((config.minibatch_size, K), jnp.int32) |
| dummy_active = jnp.zeros((config.minibatch_size, K), bool) |
| sample_out, initial_variables = jax.jit( |
| net.init_with_output |
| )(rng_keys, jnp.ones(input_shape, jnp.int32), dummy_latents, |
| dummy_positions, dummy_active) |
| |
| state, lr_scheduler_fn = trainer.get_state(config, net, initial_variables) |
| |
| |
| |
| start_step = 0 |
| if config.resume_training: |
| state = checkpoints.restore_checkpoint(config.ckpt_loc, state) |
| start_step = int(state.step) |
| print("----------Restored model from", config.ckpt_loc, |
| f"at step {start_step}-----------") |
|
|
| writer = metric_writers.create_default_writer( |
| workdir, asynchronous=False, just_logging=(jax.process_index() > 0)) |
| tf_summary_writer = tf.summary.create_file_writer(workdir) |
|
|
| logging.info("config: %s", str(config.__dict__)) |
| state = jax_utils.replicate(state) |
|
|
| dropout_rngs = jax.random.split(rng, jax.local_device_count()) |
|
|
| def make_p_train_step(num_passes): |
| return jax.pmap( |
| functools.partial( |
| trainer.train_step, |
| config=model_config, |
| hyperparams=config, |
| learning_rate_fn=lr_scheduler_fn, |
| num_passes=num_passes), |
| axis_name="batch", |
| donate_argnums=(0,)) |
|
|
| |
| |
| |
| |
| recurrent = bool(int(getattr(config, "recurrent_latent", 1))) |
| passes_per_stage = int(getattr(config, "passes_per_stage", 1)) |
| def passes_for(stage): |
| return min(passes_per_stage * stage, K) if recurrent else 0 |
| p_train_step = make_p_train_step(passes_for(curriculum.stage)) |
|
|
| p_eval_step = jax.pmap(functools.partial(evaluater.eval_step, |
| config=model_config.replace(deterministic=True)), |
| axis_name="batch") |
| |
| hooks, report_progress, train_metrics = trainer.get_metrics_report_progress( |
| config, workdir, writer) |
|
|
| tf_summary_writer, config = log_hyperparams_tb( |
| config, model_config, initial_variables, tf_summary_writer |
| ) |
|
|
| promote_threshold = float(getattr(config, "promote_acc_threshold", 0.85)) |
| promote_loc_threshold = float(getattr(config, "promote_loc_threshold", 0.70)) |
| promote_patience = int(getattr(config, "promote_patience_steps", 8000)) |
| min_stage_steps = int(getattr(config, "min_stage_steps", 2000)) |
| plateau_steps = int(getattr(config, "plateau_steps", 0)) |
| plateau_delta = float(getattr(config, "plateau_delta", 0.005)) |
| instance_mode = bool(getattr(config, "instance_dir", None)) |
| per_stage_inset = {} |
| |
| |
| stage_best_acc = -1.0 |
| stage_best_step = start_step |
| ckpt_keep = int(getattr(config, "ckpt_keep", 100)) |
| |
| stage_ckpt_dir = os.path.join(workdir, "stage_ckpts") |
| stage_started_at = start_step |
|
|
| with metric_writers.ensure_flushes(writer): |
| for step in range(start_step, config.max_steps): |
| if step%10000 == 0: |
| print("Step:", step, flush=True) |
|
|
| state, metrics = trainer.train_one_step(p_train_step, config, state, |
| step, dropout_rngs, train_data_iter) |
| |
| for h in hooks: |
| h(step) |
|
|
| if math.isnan(metrics["loss"][0]): |
| print("The loss function became nan: This might be due to the choice of hyperparameters.") |
| break |
|
|
| if step % config.eval_every_steps == 0: |
| eval_metrics = evaluater.get_eval_metrics( |
| state, eval_data_iter, p_eval_step, config) |
| per_level = eval_metrics.pop("per_level_acc") |
| per_slot = eval_metrics.pop("per_slot_acc", {}) |
| per_slot_ch = eval_metrics.pop("per_slot_acc_changed", {}) |
| per_stage_inset = eval_metrics.pop("per_stage_inset_acc", {}) |
| per_bin = eval_metrics.pop("per_bin_acc", {}) |
| def _m(key): |
| v = eval_metrics.get(key, []) |
| return round(float(np.mean(v)), 4) if len(v) else -1.0 |
| print(step, |
| "stage", curriculum.stage, |
| "loss", round(float(metrics["loss"].mean()), 4), |
| "ce", round(float(metrics["ce_loss"].mean()), 4), |
| "aux_bce", round(float(metrics["aux_loss"].mean()), 4), |
| "| val_acc", _m("acc"), |
| "loc_acc", _m("loc_acc"), |
| "val|loc", _m("val_given_loc_acc"), |
| "| cand_bit_acc", _m("cand_bit_acc"), |
| "cand_set_acc", _m("cand_set_acc"), |
| "cand_set_chg", _m("cand_set_acc_changed"), |
| "| inset", round( |
| float(per_stage_inset.get(curriculum.stage, -1.0)), 4), |
| "promote_need inset>={:.2f} loc>={:.2f}".format( |
| promote_threshold, promote_loc_threshold), |
| "| depth_acc", round( |
| float(per_slot_ch.get(curriculum.stage, -1.0)), 4), |
| flush=True) |
| with tf_summary_writer.as_default(): |
| tf.summary.scalar("loss", metrics["loss"].mean(), step=step) |
| tf.summary.scalar("ce_loss", metrics["ce_loss"].mean(), step=step) |
| tf.summary.scalar("aux_bce_loss", metrics["aux_loss"].mean(), step=step) |
| tf.summary.scalar( |
| "learning rate", metrics["learning_rate"].mean(), step=step |
| ) |
| tf.summary.scalar("curriculum_stage", curriculum.stage, step=step) |
|
|
| log_dict = {'loss': metrics["loss"].mean(), 'learning rate': metrics["learning_rate"].mean()} |
|
|
| for key in eval_metrics.keys(): |
| vals = eval_metrics[key] |
| if not vals: |
| continue |
| tf.summary.scalar( |
| "eval_" + key, np.array(vals).mean(), step=step |
| ) |
| log_dict[ "eval_" + key ] = np.array(vals).mean() |
|
|
| for lvl, v in per_level.items(): |
| if v >= 0: |
| tf.summary.scalar(f"eval_acc_level_{lvl}", v, step=step) |
| log_dict[f"eval_acc_level_{lvl}"] = v |
|
|
| for s, v in per_slot.items(): |
| if v >= 0: |
| tf.summary.scalar(f"eval_cand_depth_{s}", v, step=step) |
| log_dict[f"eval_cand_depth_{s}"] = v |
| for s, v in per_slot_ch.items(): |
| if v >= 0: |
| tf.summary.scalar(f"eval_cand_depth_chg_{s}", v, step=step) |
|
|
| if config.use_wandb: wandb.log(log_dict, step=step) |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| rounds_curric = str( |
| getattr(config, "data_curriculum", "none")) == "rounds" |
| has_frontier = int(config.num_latent_slots) > 0 or rounds_curric |
| if has_frontier and curriculum.stage < curriculum.max_stage: |
| t = curriculum.stage |
| loc_now = _m("loc_acc") |
| if instance_mode: |
| |
| |
| |
| |
| |
| |
| |
| inset_now = per_stage_inset.get(t, -1.0) |
| frontier_acc = float(inset_now) if inset_now >= 0 else -1.0 |
| elif int(config.num_latent_slots) > 0: |
| frontier_acc = per_slot_ch.get(t, -1.0) |
| if frontier_acc < 0: |
| frontier_acc = per_slot.get(t, -1.0) |
| else: |
| frontier_acc = per_bin.get(t, -1.0) |
| if frontier_acc < 0: |
| |
| |
| |
| seen = [v for b, v in per_bin.items() |
| if b <= t and v >= 0] |
| frontier_acc = (float(np.mean(seen)) if seen |
| else -1.0) |
| steps_in_stage = step - stage_started_at |
| |
| |
| |
| |
| |
| measured = frontier_acc >= 0 |
| if measured and frontier_acc > stage_best_acc + plateau_delta: |
| stage_best_acc = frontier_acc |
| stage_best_step = step |
| loc_ready = (not instance_mode) or loc_now >= promote_loc_threshold |
| hit_threshold = (measured and loc_ready |
| and frontier_acc >= promote_threshold) |
| stalled = (measured and loc_ready and plateau_steps > 0 |
| and (step - stage_best_step) >= plateau_steps) |
| patience_over = steps_in_stage >= promote_patience |
| if steps_in_stage >= min_stage_steps and ( |
| hit_threshold or stalled or patience_over): |
| reason = ("threshold" if hit_threshold |
| else "plateau" if stalled else "patience") |
| curriculum.stage += 1 |
| stage_started_at = step |
| stage_best_acc = -1.0 |
| stage_best_step = step |
| p_train_step = make_p_train_step(passes_for(curriculum.stage)) |
| what = ("round-bin" if int(config.num_latent_slots) == 0 |
| else "depth") |
| print(f"[curriculum] step {step}: promote to stage " |
| f"{curriculum.stage} ({reason}; graduated {what} " |
| f"{t} acc={frontier_acc:.3f} after " |
| f"{steps_in_stage} steps); " |
| f"latent passes={curriculum.stage}, " |
| f"pool/snapshots now 1..{curriculum.stage}", |
| flush=True) |
| if config.save_checkpoint: |
| unrep_state = jax_utils.unreplicate(state) |
| |
| checkpoints.save_checkpoint_multiprocess( |
| workdir, unrep_state, step, |
| keep=ckpt_keep, overwrite=True) |
| |
| |
| |
| checkpoints.save_checkpoint_multiprocess( |
| stage_ckpt_dir, unrep_state, step, |
| keep=100, overwrite=True, |
| prefix=f"stage{curriculum.stage}_") |
|
|
| if config.save_checkpoint and step > 0 and step % config.save_every_steps == 0: |
| checkpoints.save_checkpoint_multiprocess( |
| workdir, jax_utils.unreplicate(state), step, |
| keep=ckpt_keep, overwrite=True |
| ) |
|
|
| |
| if config.save_checkpoint: |
| checkpoints.save_checkpoint_multiprocess( |
| workdir, jax_utils.unreplicate(state), config.max_steps, |
| keep=ckpt_keep, overwrite=True) |
|
|
| |