diff options
Diffstat (limited to 'external/dualprop_patches/0009-crossover-isolate-formal-validation-from-test.patch')
| -rw-r--r-- | external/dualprop_patches/0009-crossover-isolate-formal-validation-from-test.patch | 135 |
1 files changed, 135 insertions, 0 deletions
diff --git a/external/dualprop_patches/0009-crossover-isolate-formal-validation-from-test.patch b/external/dualprop_patches/0009-crossover-isolate-formal-validation-from-test.patch new file mode 100644 index 0000000..03a826a --- /dev/null +++ b/external/dualprop_patches/0009-crossover-isolate-formal-validation-from-test.patch @@ -0,0 +1,135 @@ +From 9fc597c20d6b758ce6b0c3a326b8bc1b9262b0d0 Mon Sep 17 00:00:00 2001 +From: YurenHao0426 <Blackhao0426@gmail.com> +Date: Mon, 27 Jul 2026 13:25:10 -0500 +Subject: [PATCH 09/19] crossover: isolate formal validation from test + +--- + config/cli_config.py | 10 ++++++++++ + train.py | 43 ++++++++++++++++++++++++++----------------- + train_ff.py | 22 +++++++++++++++------- + 3 files changed, 51 insertions(+), 24 deletions(-) + +diff --git a/config/cli_config.py b/config/cli_config.py +index 4979b37..67a04fe 100644 +--- a/config/cli_config.py ++++ b/config/cli_config.py +@@ -100,6 +100,16 @@ parser.add_argument( + 'validation evaluation. The author-compatible default is full; ' + 'use none for diagnostic-free timed crossover runs.')) + ++parser.add_argument( ++ '--test-policy', default='author', choices=['author', 'none'], ++ help=('author restores the best-validation checkpoint and evaluates test; ' ++ 'none keeps test completely untouched for formal validation runs.')) ++ ++parser.add_argument( ++ '--early-stop-policy', default='author', choices=['author', 'none'], ++ help=('author retains the original severe-validation-drop stop; none runs ' ++ 'the full schedule unless a nonfinite training loss occurs.')) ++ + dtypes = {'bfloat16': jnp.bfloat16, 'float16':jnp.float16, 'float32':jnp.float32} + parser.add_argument('--dtype', default='float32', choices=dtypes.keys()) + parser.add_argument('--param-dtype', default='float32', choices=['bfloat16', 'float16', 'float32']) +diff --git a/train.py b/train.py +index b1df1b0..0b23304 100644 +--- a/train.py ++++ b/train.py +@@ -43,7 +43,7 @@ for experiment_index, seed in enumerate(config.seeds): + print(msg) + + loginfo_and_print(f"\n\tStarting experiment {experiment_index+1}/{len(config.seeds)}. Current seed is {seed}") +- loginfo_and_print(f"\model settings:\n{config.model}") ++ loginfo_and_print(f"\nmodel settings:\n{config.model}") + rng = jax.random.PRNGKey(seed) + rng, init_rng = jax.random.split(rng) + # Define model +@@ -179,23 +179,32 @@ for experiment_index, seed in enumerate(config.seeds): + loginfo_and_print("NaN or Inf encountered. Terminating training loop early") + break + if (epoch>5) and (hist["val_accuracy"][epoch-1] < 0.5*best_accuracy): +- loginfo_and_print("Terminating training early as validation accuracy has drastically dropped") +- break +- +- loginfo_and_print(f"\n====Loading model with best validation accuracy (epoch {best_epoch})====") +- best_state = checkpoints.restore_checkpoint(ckpt_dir=CKPT_DIR, target=state) +- rng, input_rng = jax.random.split(rng) +- if config.learning_algorithm == "ep": +- test_loss, test_accuracy, test_top5accuracy, test_time, _, _, _, _ = eval_ep_model( +- best_state, best_state.params, config.test_ds, config.batch_size, +- config.num_classes, config.ep_free_steps, config.ep_dt) ++ if config.early_stop_policy == "author": ++ loginfo_and_print("Terminating training early as validation accuracy has drastically dropped") ++ break ++ ++ hist["best_validation_accuracy"] = best_accuracy ++ hist["best_epoch"] = best_epoch ++ hist["epochs_completed"] = epoch ++ if config.test_policy == "author": ++ loginfo_and_print(f"\n====Loading model with best validation accuracy (epoch {best_epoch})====") ++ best_state = checkpoints.restore_checkpoint(ckpt_dir=CKPT_DIR, target=state) ++ rng, input_rng = jax.random.split(rng) ++ if config.learning_algorithm == "ep": ++ test_loss, test_accuracy, test_top5accuracy, test_time, _, _, _, _ = eval_ep_model( ++ best_state, best_state.params, config.test_ds, ++ config.batch_size, config.num_classes, config.ep_free_steps, ++ config.ep_dt) ++ else: ++ test_loss, test_accuracy, test_top5accuracy, test_time, _, _, _, _ = eval_model( ++ best_state, best_state.params, config.test_ds, ++ config.batch_size, config.num_classes, input_rng, ++ spectral_diagnostics=( ++ config.spectral_diagnostics == "full")) ++ hist['test_loss'], hist['test_accuracy'], hist['test_top5accuracy'], hist['test_time'] = test_loss, test_accuracy, test_top5accuracy, test_time ++ loginfo_and_print('test: \tloss: %.4f, \taccuracy: %.4f, \ttop5_accuracy: %.4f, \truntime: %.4f' % (test_loss, test_accuracy, test_top5accuracy, test_time)) + else: +- test_loss, test_accuracy, test_top5accuracy, test_time, _, _, _, _ = eval_model( +- best_state, best_state.params, config.test_ds, config.batch_size, +- config.num_classes, input_rng, +- spectral_diagnostics=(config.spectral_diagnostics == "full")) +- hist['test_loss'], hist['test_accuracy'], hist['test_top5accuracy'], hist['test_time'] = test_loss, test_accuracy, test_top5accuracy, test_time +- loginfo_and_print('test: \tloss: %.4f, \taccuracy: %.4f, \ttop5_accuracy: %.4f, \truntime: %.4f' % (test_loss, test_accuracy, test_top5accuracy, test_time)) ++ loginfo_and_print("\n====Test evaluation disabled by validation-only policy====") + + if (config.learning_algorithm != "backprop" + and config.gradient_diagnostics == "full"): +diff --git a/train_ff.py b/train_ff.py +index 3839dd3..ada45a3 100644 +--- a/train_ff.py ++++ b/train_ff.py +@@ -75,9 +75,12 @@ for experiment_index, seed in enumerate(config.seeds): + validation_accuracy, validation_time = eval_ff_model( + state, config.val_ds, config.batch_size, config.num_classes, + config.ff_score_from_layer) +- test_accuracy, test_time = eval_ff_model( +- state, config.test_ds, config.batch_size, config.num_classes, +- config.ff_score_from_layer) ++ test_accuracy = np.nan ++ test_time = np.nan ++ if config.test_policy == "author": ++ test_accuracy, test_time = eval_ff_model( ++ state, config.test_ds, config.batch_size, config.num_classes, ++ config.ff_score_from_layer) + history["final"] = { + "validation_accuracy": validation_accuracy, + "validation_time": validation_time, +@@ -85,8 +88,13 @@ for experiment_index, seed in enumerate(config.seeds): + "test_time": test_time, + "train_and_eval_wall": time.time() - started, + } +- print( +- f"final val_accuracy={validation_accuracy:.3f}% " +- f"test_accuracy={test_accuracy:.3f}%", +- flush=True) ++ if config.test_policy == "author": ++ print( ++ f"final val_accuracy={validation_accuracy:.3f}% " ++ f"test_accuracy={test_accuracy:.3f}%", ++ flush=True) ++ else: ++ print( ++ f"final val_accuracy={validation_accuracy:.3f}% test=disabled", ++ flush=True) + np.save(os.path.join(outpath, "hist.npy"), history) +-- +2.54.0 + |
