From 17fc860412aa78b16b81f7d2476644b6650b2bd7 Mon Sep 17 00:00:00 2001 From: YurenHao0426 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