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
|
#!/usr/bin/env python3
"""Deterministic equation audits for matched ResNet crossover adapters."""
import json
import os
import sys
import torch
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from sdil.conv import CIFARLocalResNet
from sdil.conv_crossover import CIFARDualPropResNet
def relative_error(actual, expected):
numerator = torch.linalg.vector_norm(actual - expected)
denominator = torch.linalg.vector_norm(expected).clamp_min(1e-30)
return float(numerator / denominator)
def audit_dualprop():
common = dict(
depth=8, base_width=2, seed=71, normalization="batchnorm",
residual_scale=1.0, dtype=torch.float64)
reference = CIFARLocalResNet(**common)
net = CIFARDualPropResNet(
**common, alpha=0.0, dp_beta=0.1, inference_passes=2)
generator = torch.Generator().manual_seed(72)
image = torch.randn(3, 3, 32, 32, generator=generator,
dtype=torch.float64)
labels = torch.tensor([0, 3, 7])
one_hot = torch.nn.functional.one_hot(labels, 10).to(torch.float64)
with torch.no_grad():
reference_output = reference.forward(
image, training=True, update_stats=False)["logits"]
net_output = net.forward(
image, training=True, update_stats=False)["logits"]
clean = net.forward(
image, return_cache=True, training=True, update_stats=False)
plus, minus = net.infer_dual_states(image, one_hot, clean)
forward_error = float(torch.max(torch.abs(
reference_output - net_output)))
parameters = (
net.W + net.gamma + net.beta + [net.W_out, net.b_out])
for parameter in parameters:
parameter.requires_grad_(True)
alpha = net.dp_alpha
beta = net.dp_beta
states = [
(alpha * positive + (1.0 - alpha) * negative).detach()
for positive, negative in zip(plus[:-1], minus[:-1])]
deltas = [
((positive - negative) / beta).detach()
for positive, negative in zip(plus, minus)]
objective = image.new_zeros(())
for index in range(net.n_hidden):
prediction, _ = net._node_prediction(index, states, image)
objective -= torch.sum(deltas[index] * prediction) / image.shape[0]
features = states[-1].mean(dim=(2, 3))
output_prediction = features @ net.W_out.t() + net.b_out
objective -= torch.sum(deltas[-1] * output_prediction) / image.shape[0]
gradients = torch.autograd.grad(objective, parameters)
(directions, gamma_directions, beta_directions,
output_weight, output_bias) = net.dualprop_ascent_directions(
image, plus, minus)
actual = (
directions + gamma_directions + beta_directions
+ [output_weight, output_bias])
errors = [
relative_error(direction, -gradient)
for direction, gradient in zip(actual, gradients)]
for parameter in parameters:
parameter.requires_grad_(False)
if forward_error >= 1e-12 or max(errors) >= 2e-12:
raise AssertionError({
"forward_error": forward_error,
"direction_errors": errors,
})
return {
"matched_forward_max_absolute_error": forward_error,
"contrastive_direction_max_relative_error": max(errors),
"num_audited_parameter_tensors": len(errors),
"uses_symmetric_forward_edge_transposes": True,
"uses_reverse_task_loss_graph": False,
}
def main():
print(json.dumps({"dualprop": audit_dualprop()},
indent=2, sort_keys=True))
if __name__ == "__main__":
main()
|