summaryrefslogtreecommitdiff
path: root/sdil
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:18:36 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:18:36 -0500
commitc6af287a6177cda3863e6f25f2bdab8abd611ad8 (patch)
tree1be6b3d37d9cc3d7d54bed7d2851e37fc9bdeb9e /sdil
parentdd705590a6210b6ada988cec0a392f7669c5cb52 (diff)
oral-a: measure training-state apical alignment
Diffstat (limited to 'sdil')
-rw-r--r--sdil/conv.py3
1 files changed, 2 insertions, 1 deletions
diff --git a/sdil/conv.py b/sdil/conv.py
index 40def21..0436525 100644
--- a/sdil/conv.py
+++ b/sdil/conv.py
@@ -753,7 +753,7 @@ def conv_alignment_report(net, x, y, config):
parameters = net.W + net.gamma + net.beta + [net.W_out, net.b_out]
for parameter in parameters:
parameter.requires_grad_(True)
- forward = net.forward(x)
+ forward = net.forward(x, training=True, update_stats=False)
gradients = torch.autograd.grad(
F.cross_entropy(forward["logits"], y), forward["hiddens"])
batch = x.shape[0]
@@ -771,6 +771,7 @@ def conv_alignment_report(net, x, y, config):
return float(F.cosine_similarity(left, right, dim=1).mean())
report = {
+ "normalization_state": "training_batch_stats_without_running_update",
"teaching_negative_gradient_cosine": [
cosine(left, right) for left, right in zip(teaching, negative_gradients)],
"raw_negative_gradient_cosine": [