summaryrefslogtreecommitdiff
path: root/external/dualprop_patches/0009-crossover-isolate-formal-validation-from-test.patch
diff options
context:
space:
mode:
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.patch135
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
+