diff options
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.patch | 110 |
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 + |
