summaryrefslogtreecommitdiff
path: root/external/dualprop_patches/0010-crossover-add-paper-PEPITA-learning-schedule.patch
diff options
context:
space:
mode:
Diffstat (limited to 'external/dualprop_patches/0010-crossover-add-paper-PEPITA-learning-schedule.patch')
-rw-r--r--external/dualprop_patches/0010-crossover-add-paper-PEPITA-learning-schedule.patch110
1 files changed, 110 insertions, 0 deletions
diff --git a/external/dualprop_patches/0010-crossover-add-paper-PEPITA-learning-schedule.patch b/external/dualprop_patches/0010-crossover-add-paper-PEPITA-learning-schedule.patch
new file mode 100644
index 0000000..73f17d1
--- /dev/null
+++ b/external/dualprop_patches/0010-crossover-add-paper-PEPITA-learning-schedule.patch
@@ -0,0 +1,110 @@
+From 17fc860412aa78b16b81f7d2476644b6650b2bd7 Mon Sep 17 00:00:00 2001
+From: YurenHao0426 <Blackhao0426@gmail.com>
+Date: Mon, 27 Jul 2026 13:26:52 -0500
+Subject: [PATCH 10/19] crossover: add paper PEPITA learning schedule
+
+---
+ config/cli_config.py | 8 ++++++++
+ src/training_utils.py | 26 +++++++++++++++++++-------
+ train.py | 10 ++++++++--
+ 3 files changed, 35 insertions(+), 9 deletions(-)
+
+diff --git a/config/cli_config.py b/config/cli_config.py
+index 67a04fe..b38ff7b 100644
+--- a/config/cli_config.py
++++ b/config/cli_config.py
+@@ -24,6 +24,11 @@ parser.add_argument('--decay-epochs', default=None, type=int, help='')
+
+ parser.add_argument('--warmup-epochs', default=0, type=int, help='')
+
++parser.add_argument(
++ '--optimizer-schedule', default='author',
++ choices=['author', 'pepita'],
++ help='Author warmup/cosine schedule or PEPITA epoch-60/90 drops.')
++
+ parser.add_argument('--momentum', default=0.9, type=float, const=None, action='store', nargs='?', help='')
+
+ parser.add_argument('--weight-decay', default=5e-4, type=float, help='')
+@@ -127,6 +132,9 @@ datasets = dict(fashionmnist=get_fashionmnist, mnist=get_mnist, svhn=get_svhn, c
+ parser.add_argument('--dataset', '-d', choices=datasets.keys(), default='cifar10', help='')
+
+ config = parser.parse_args()
++if (config.optimizer_schedule == "pepita"
++ and config.learning_algorithm != "pepita"):
++ parser.error("--optimizer-schedule pepita is restricted to PEPITA")
+
+ # Note for comparisson to previous work we pad MNIST and FMNIST to size (32,32,1)
+ imagedims = {"fashionmnist": (28,28,1), "mnist": (32,32,1), "svhn": (32,32,3), "cifar10": (32,32,3), "cifar100": (32,32,3), "imagenet_32x32": (32,32,3)}
+diff --git a/src/training_utils.py b/src/training_utils.py
+index ac05857..271a81e 100644
+--- a/src/training_utils.py
++++ b/src/training_utils.py
+@@ -197,18 +197,30 @@ def get_imagenet_32x32(dtype, percent_train=95, percent_val=5):
+ test_ds['image'] = (test_ds['image'] - mean_data)/std_data
+ return train_ds, val_ds, test_ds
+
+-def create_train_state(rng, model, image_dims, lr, wlr, lrf, momentum, weight_decay, num_epochs, warmup_epochs, decay_epochs, steps_per_epoch):
++def create_train_state(rng, model, image_dims, lr, wlr, lrf, momentum,
++ weight_decay, num_epochs, warmup_epochs, decay_epochs,
++ steps_per_epoch, schedule_mode="author"):
+ """Creates initial `TrainState`."""
+ w, h, ch = image_dims
+ x = jnp.ones([1, w, h, ch])
+ params = model.init(rng, x)['params']
+
+- schedule = optax.warmup_cosine_decay_schedule(
+- init_value=wlr,
+- peak_value=lr,
+- warmup_steps=warmup_epochs*steps_per_epoch,
+- decay_steps = (decay_epochs)*steps_per_epoch,
+- end_value=lrf)
++ if schedule_mode == "author":
++ schedule = optax.warmup_cosine_decay_schedule(
++ init_value=wlr,
++ peak_value=lr,
++ warmup_steps=warmup_epochs * steps_per_epoch,
++ decay_steps=decay_epochs * steps_per_epoch,
++ end_value=lrf)
++ elif schedule_mode == "pepita":
++ schedule = optax.piecewise_constant_schedule(
++ init_value=lr,
++ boundaries_and_scales={
++ 60 * steps_per_epoch: 0.1,
++ 90 * steps_per_epoch: 0.1,
++ })
++ else:
++ raise ValueError(f"unknown optimizer schedule: {schedule_mode}")
+
+ tx = optax.chain(
+ optax.add_decayed_weights(weight_decay=weight_decay, mask=None),
+diff --git a/train.py b/train.py
+index 0b23304..2a03ada 100644
+--- a/train.py
++++ b/train.py
+@@ -49,7 +49,12 @@ for experiment_index, seed in enumerate(config.seeds):
+ # Define model
+ steps_per_epoch = len(config.train_ds['image']) // config.batch_size
+
+- state = create_train_state(init_rng, config.model, config.image_dims, config.learning_rate, config.warmup_learning_rate, config.learning_rate_final, config.momentum, config.weight_decay, config.num_epochs, config.warmup_epochs, config.decay_epochs, steps_per_epoch)
++ state = create_train_state(
++ init_rng, config.model, config.image_dims, config.learning_rate,
++ config.warmup_learning_rate, config.learning_rate_final,
++ config.momentum, config.weight_decay, config.num_epochs,
++ config.warmup_epochs, config.decay_epochs, steps_per_epoch,
++ schedule_mode=config.optimizer_schedule)
+ local_feedback = create_local_feedback(
+ jax.random.PRNGKey(config.feedback_seed), config.model, state.params,
+ config.image_dims, config.learning_algorithm, config.num_classes,
+@@ -61,7 +66,8 @@ for experiment_index, seed in enumerate(config.seeds):
+ config.image_dims, config.learning_rate,
+ config.warmup_learning_rate, config.learning_rate_final,
+ config.momentum, config.weight_decay, config.num_epochs,
+- config.warmup_epochs, config.decay_epochs, steps_per_epoch)
++ config.warmup_epochs, config.decay_epochs, steps_per_epoch,
++ schedule_mode=config.optimizer_schedule)
+ sdil_auxiliary = None
+ sdil_initialization = None
+ if config.learning_algorithm == "sdil":
+--
+2.54.0
+