summaryrefslogtreecommitdiff
path: root/external/dualprop_patches/0010-crossover-add-paper-PEPITA-learning-schedule.patch
blob: 73f17d182871a1e871c6183ee0747a2f111eecb0 (plain)
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