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
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
|
From cd932728011d5fef91109dd24df8a5f2dfb9e6e7 Mon Sep 17 00:00:00 2001
From: YurenHao0426 <Blackhao0426@gmail.com>
Date: Mon, 27 Jul 2026 12:50:47 -0500
Subject: [PATCH 01/19] crossover: separate author diagnostics from timed
training
---
config/cli_config.py | 14 +++++++++++++-
src/training_utils.py | 42 ++++++++++++++++++++++++++++++------------
train.py | 27 ++++++++++++++++++++-------
3 files changed, 63 insertions(+), 20 deletions(-)
diff --git a/config/cli_config.py b/config/cli_config.py
index 0aec5f6..c1b23ba 100644
--- a/config/cli_config.py
+++ b/config/cli_config.py
@@ -44,6 +44,18 @@ parser.add_argument('--model', default='VGG16', choices=['VGG16', 'VGGlike', 'CN
parser.add_argument('--learning-algorithm', default='dualprop-lagr-ff', choices=['backprop', 'dualprop-lagr-ff', 'dualprop-raovr-ff', 'dualprop-raovr-dampened-ff'])
+parser.add_argument(
+ '--gradient-diagnostics', default='full', choices=['none', 'full'],
+ help=('Compute the exact BP reference gradient and layerwise cosine on '
+ 'every training minibatch. The author-compatible default is full; '
+ 'use none for diagnostic-free timed crossover runs.'))
+
+parser.add_argument(
+ '--spectral-diagnostics', default='full', choices=['none', 'full'],
+ help=('Run the author power-iteration L/gamma probes after every '
+ 'validation evaluation. The author-compatible default is full; '
+ 'use none for diagnostic-free timed crossover runs.'))
+
dtypes = {'bfloat16': jnp.bfloat16, 'float16':jnp.float16, 'float32':jnp.float32}
parser.add_argument('--dtype', default='float32', choices=dtypes.keys())
parser.add_argument('--param-dtype', default='float32', choices=['bfloat16', 'float16', 'float32'])
@@ -120,4 +132,4 @@ config.model = modeltype[config.learning_algorithm](loss_func, Conv, Dense, acti
)
# Load datasets
-config.train_ds, config.val_ds, config.test_ds = datasets[config.dataset](config.dtype, config.percent_train, config.percent_val)
\ No newline at end of file
+config.train_ds, config.val_ds, config.test_ds = datasets[config.dataset](config.dtype, config.percent_train, config.percent_val)
diff --git a/src/training_utils.py b/src/training_utils.py
index 7152451..2fbfccd 100644
--- a/src/training_utils.py
+++ b/src/training_utils.py
@@ -240,7 +240,8 @@ def to_float16(ptree):
def to_float32(ptree):
return tree_map(lambda x: x.astype(jnp.float32), ptree)
-def train_epoch(state, train_ds, batch_size, rng, augmentation_on, learning_algorithm, num_classes):
+def train_epoch(state, train_ds, batch_size, rng, augmentation_on,
+ learning_algorithm, num_classes, gradient_diagnostics=True):
"""Train for a single epoch."""
t0 = time.time()
train_ds_size = len(train_ds['image'])
@@ -261,7 +262,9 @@ def train_epoch(state, train_ds, batch_size, rng, augmentation_on, learning_algo
# image = vmap_augment_train_imagenet(image, batch_rng)
if learning_algorithm != "backprop":
- state, metrics = train_step(state, image, labels_onehot, labels, batch_rng, inf_rng, augmentation_on)
+ state, metrics = train_step(
+ state, image, labels_onehot, labels, batch_rng, inf_rng,
+ augmentation_on, gradient_diagnostics)
elif learning_algorithm == "backprop":
state, metrics = train_step_bp(state, image, labels_onehot, labels, batch_rng, inf_rng, augmentation_on)
batch_metrics.append(metrics)
@@ -279,7 +282,8 @@ def train_epoch(state, train_ds, batch_size, rng, augmentation_on, learning_algo
return state, epoch_metrics_np, runtime
@jax.jit
-def train_step(state, image, labels_onehot, labels, batch_rng, inf_rng, augmentation_on):
+def train_step(state, image, labels_onehot, labels, batch_rng, inf_rng,
+ augmentation_on, gradient_diagnostics):
"""Train for a single step."""
# batch_rng = jax.random.split(batch_rng, batch['image'].shape[0])
@@ -298,8 +302,9 @@ def train_step(state, image, labels_onehot, labels, batch_rng, inf_rng, augmenta
inf_rng, _ = jax.random.split(inf_rng)
metrics = compute_metrics(image=image, labels_onehot=labels_onehot, labels=labels, state=state)
- get_ref_grad_angle = True
- metrics = jax.lax.cond(get_ref_grad_angle, ref_grad_and_angle, no_ref_grad_and_angle, state, grads, image, labels_onehot, metrics)
+ metrics = jax.lax.cond(
+ gradient_diagnostics, ref_grad_and_angle, no_ref_grad_and_angle,
+ state, grads, image, labels_onehot, metrics)
# The optimizer may modify grads, so we need to compare grads and ref_grads before performing the gradient step.
state = state.apply_gradients(grads=grads)
@@ -340,7 +345,8 @@ def eval_step(state, params, image, labels_onehot, labels, inf_rng):
metrics = compute_metrics(image=image, labels_onehot=labels_onehot, labels=labels, state=state)
return metrics
-def eval_model(state, params, test_ds, batch_size, num_classes, eval_rng):
+def eval_model(state, params, test_ds, batch_size, num_classes, eval_rng,
+ spectral_diagnostics=True):
t0 = time.time()
test_ds_size = len(test_ds['image'])
steps = test_ds_size // batch_size
@@ -365,11 +371,23 @@ def eval_model(state, params, test_ds, batch_size, num_classes, eval_rng):
runtime = time.time() - t0
- # dummy states, used by get_L_and_gamma to infer correct array shape when generating random arrays
- sdummy = state.apply_fn({'params': params}, batch["image"], method='make_predictions')
- eval_rng, _ = jax.random.split(eval_rng)
- L10, gamma10 = state.apply_fn({'params': state.params}, s=sdummy, rng_key=eval_rng, numiter=10, method='get_L_and_gamma')
- L20, gamma20 = state.apply_fn({'params': state.params}, s=sdummy, rng_key=eval_rng, numiter=20, method='get_L_and_gamma')
+ if spectral_diagnostics:
+ # Dummy states infer the array shapes used by the author's power
+ # iterations. These diagnostics are expensive and are deliberately
+ # optional in timed crossover runs.
+ sdummy = state.apply_fn(
+ {'params': params}, batch["image"], method='make_predictions')
+ eval_rng, _ = jax.random.split(eval_rng)
+ L10, gamma10 = state.apply_fn(
+ {'params': state.params}, s=sdummy, rng_key=eval_rng, numiter=10,
+ method='get_L_and_gamma')
+ L20, gamma20 = state.apply_fn(
+ {'params': state.params}, s=sdummy, rng_key=eval_rng, numiter=20,
+ method='get_L_and_gamma')
+ else:
+ count = len(state.params)
+ L10 = L20 = gamma10 = gamma20 = [
+ jnp.asarray(jnp.nan, dtype=jnp.float32) for _ in range(count)]
return summary['loss'], summary['accuracy'], summary['top5accuracy'], runtime, L10, L20, gamma10, gamma20
@@ -470,4 +488,4 @@ def plot_L_or_gamma(L20, L10, ylabel, save_path):
plt.tight_layout()
plt.savefig(save_path)
plt.close()
- return
\ No newline at end of file
+ return
diff --git a/train.py b/train.py
index bd7f6b7..139fd6c 100644
--- a/train.py
+++ b/train.py
@@ -71,7 +71,10 @@ for experiment_index, seed in enumerate(config.seeds):
# Run an optimization step over a training batch
# last augument turns off data augmentation for mnist
augmentation_on = (config.dataset!="mnist") and (config.dataset!="fashionmnist")
- state, epoch_metrics, train_time = train_epoch(state, config.train_ds, config.batch_size, input_rng, augmentation_on, config.learning_algorithm, config.num_classes)
+ state, epoch_metrics, train_time = train_epoch(
+ state, config.train_ds, config.batch_size, input_rng,
+ augmentation_on, config.learning_algorithm, config.num_classes,
+ gradient_diagnostics=(config.gradient_diagnostics == "full"))
loginfo_and_print('train: \tloss: %.4f, \taccuracy: %.4f, \truntime: %.4f' % (epoch_metrics["loss"], epoch_metrics["accuracy"], train_time))
hist['train_loss'][epoch-1], hist['train_accuracy'][epoch-1], hist['train_time'][epoch-1] = epoch_metrics["loss"], epoch_metrics["accuracy"], train_time
@@ -83,7 +86,10 @@ for experiment_index, seed in enumerate(config.seeds):
# Evaluate on the validation set after each training epoch
rng, input_rng = jax.random.split(rng)
- val_loss, val_accuracy, val_top5_accuracy, val_time, L10, L20, gamma10, gamma20 = eval_model(state, state.params, config.val_ds, config.batch_size, config.num_classes, input_rng)
+ val_loss, val_accuracy, val_top5_accuracy, val_time, L10, L20, gamma10, gamma20 = eval_model(
+ state, state.params, config.val_ds, config.batch_size,
+ config.num_classes, input_rng,
+ spectral_diagnostics=(config.spectral_diagnostics == "full"))
loginfo_and_print('val: \tloss: %.4f, \taccuracy: %.4f, \ttop5_accuracy: %.4f, \truntime: %.4f' % (val_loss, val_accuracy, val_top5_accuracy, val_time))
loginfo_and_print(f"L20: {[np.round(Li.item(), decimals=4) for Li in L20]}")
loginfo_and_print(f"gamma20: {[np.round(gi.item(), decimals=4) for gi in gamma20]}")
@@ -108,18 +114,25 @@ for experiment_index, seed in enumerate(config.seeds):
loginfo_and_print(f"\n====Loading model with best validation accuracy (epoch {best_epoch})====")
best_state = checkpoints.restore_checkpoint(ckpt_dir=CKPT_DIR, target=state)
rng, input_rng = jax.random.split(rng)
- test_loss, test_accuracy, test_top5accuracy, test_time, _, _, _, _ = eval_model(best_state, best_state.params, config.test_ds, config.batch_size, config.num_classes, input_rng)
+ test_loss, test_accuracy, test_top5accuracy, test_time, _, _, _, _ = eval_model(
+ best_state, best_state.params, config.test_ds, config.batch_size,
+ config.num_classes, input_rng,
+ spectral_diagnostics=(config.spectral_diagnostics == "full"))
hist['test_loss'], hist['test_accuracy'], hist['test_top5accuracy'], hist['test_time'] = test_loss, test_accuracy, test_top5accuracy, test_time
loginfo_and_print('test: \tloss: %.4f, \taccuracy: %.4f, \ttop5_accuracy: %.4f, \truntime: %.4f' % (test_loss, test_accuracy, test_top5accuracy, test_time))
- if config.learning_algorithm != "backprop":
+ if (config.learning_algorithm != "backprop"
+ and config.gradient_diagnostics == "full"):
# Use color_norm=LogNorm(clip=True) for logscale plot
heatmap_grads_epochs(hist["grad_cos_sim_epochs"], outpath+"grad_angle_epochs.pdf", True, color_norm=None)
#Grad angle across batches in the first epoch
first_N = 100
heatmap_grads_batches(hist["grad_cos_sim_batches"][:,0:first_N], outpath+f"grad_angle_first_{first_N}_batches.pdf", True, color_norm=None)
- plot_L_or_gamma(hist["L20"], hist["L10"], "L", outpath+"L.pdf")
- plot_L_or_gamma(hist["gamma20"], hist["gamma10"], r"$\gamma$", outpath+"gamma.pdf")
+ if config.spectral_diagnostics == "full":
+ plot_L_or_gamma(hist["L20"], hist["L10"], "L", outpath+"L.pdf")
+ plot_L_or_gamma(
+ hist["gamma20"], hist["gamma10"], r"$\gamma$",
+ outpath+"gamma.pdf")
- np.save(outpath+"hist.npy", hist)
\ No newline at end of file
+ np.save(outpath+"hist.npy", hist)
--
2.54.0
|