1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
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
|