CartayaGon commited on
Commit
f91a968
·
verified ·
1 Parent(s): 6999db7

Upload main.py

Browse files
Files changed (1) hide show
  1. main.py +156 -0
main.py ADDED
@@ -0,0 +1,156 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.optim as optim
4
+ import torchvision
5
+ import torchvision.transforms as transforms
6
+ import copy
7
+ import torch.fft
8
+ import torch.nn.functional as F
9
+
10
+ # --- 1. DATA PREPARATION (Split-MNIST) ---
11
+ transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,))])
12
+ trainset = torchvision.datasets.MNIST(root='./data', train=True, download=True, transform=transform)
13
+ testset = torchvision.datasets.MNIST(root='./data', train=False, download=True, transform=transform)
14
+
15
+ def get_split_dataloaders(dataset, classes, batch_size=64):
16
+ indices = [i for i, target in enumerate(dataset.targets) if target in classes]
17
+ subset = torch.utils.data.Subset(dataset, indices)
18
+ return torch.utils.data.DataLoader(subset, batch_size=batch_size, shuffle=True)
19
+
20
+ # Task A: Digits 0-4
21
+ train_loader_A = get_split_dataloaders(trainset, [0, 1, 2, 3, 4])
22
+ test_loader_A = get_split_dataloaders(testset, [0, 1, 2, 3, 4])
23
+
24
+ # Task B: Digits 5-9
25
+ train_loader_B = get_split_dataloaders(trainset, [5, 6, 7, 8, 9])
26
+ test_loader_B = get_split_dataloaders(testset, [5, 6, 7, 8, 9])
27
+
28
+ # --- 2. NEURAL NETWORK ARCHITECTURE (CNN) ---
29
+ class SimpleCNN(nn.Module):
30
+ def __init__(self):
31
+ super(SimpleCNN, self).__init__()
32
+ self.conv1 = nn.Conv2d(1, 16, kernel_size=3, padding=1)
33
+ self.relu = nn.ReLU()
34
+ self.pool = nn.MaxPool2d(2, 2)
35
+ self.conv2 = nn.Conv2d(16, 32, kernel_size=3, padding=1)
36
+ self.fc1 = nn.Linear(32 * 7 * 7, 128)
37
+ self.fc2 = nn.Linear(128, 10)
38
+
39
+ def forward(self, x):
40
+ x = self.pool(self.relu(self.conv1(x)))
41
+ x = self.pool(self.relu(self.conv2(x)))
42
+ x = x.view(-1, 32 * 7 * 7)
43
+ x = self.relu(self.fc1(x))
44
+ x = self.fc2(x)
45
+ return x
46
+
47
+ def evaluate_accuracy(model, dataloader):
48
+ model.eval()
49
+ correct = 0
50
+ total = 0
51
+ with torch.no_grad():
52
+ for images, labels in dataloader:
53
+ outputs = model(images)
54
+ _, predicted = torch.max(outputs.data, 1)
55
+ total += labels.size(0)
56
+ correct += (predicted == labels).sum().item()
57
+ return 100 * correct / total
58
+
59
+ # --- 3. ANASTROPHIC REGULARIZATION ---
60
+ class AnastrophicRegularizer(nn.Module):
61
+ def __init__(self, lambda_reg=1.0, eta_reg=3.0):
62
+ super().__init__()
63
+ self.lambda_reg = lambda_reg
64
+ self.eta_reg = eta_reg
65
+
66
+ def compute_phi(self, w):
67
+ """Calculates Spectral Coherence (Phi) via 1D FFT."""
68
+ fft_w = torch.fft.fft(w.view(-1))
69
+ amplitudes = torch.abs(fft_w)
70
+ phases = torch.angle(fft_w)
71
+
72
+ p_j = (amplitudes ** 2) / (torch.sum(amplitudes ** 2) + 1e-8)
73
+ complex_sum = torch.sum(p_j * torch.exp(1j * phases))
74
+
75
+ return torch.abs(complex_sum)
76
+
77
+ def compute_beta_proxy(self, w, w_prev):
78
+ """Continuous proxy for Anastrophic Beta (BB) measuring harmonic tension."""
79
+ fft_w = torch.fft.fft(w.view(-1))
80
+ fft_prev = torch.fft.fft(w_prev.view(-1))
81
+
82
+ complex_w = torch.view_as_real(fft_w)
83
+ complex_prev = torch.view_as_real(fft_prev)
84
+
85
+ return F.mse_loss(complex_w, complex_prev)
86
+
87
+ def forward(self, model, model_prev):
88
+ loss_ana = 0.0
89
+ for (name, param), (name_prev, param_prev) in zip(model.named_parameters(), model_prev.named_parameters()):
90
+ if 'weight' in name:
91
+ phi = self.compute_phi(param)
92
+ # .detach() is critical to anchor the previous structural state
93
+ beta = self.compute_beta_proxy(param, param_prev.detach())
94
+
95
+ loss_ana += self.lambda_reg * (1 - phi) + self.eta_reg * beta
96
+
97
+ return loss_ana
98
+
99
+ # --- 4. PHASE 1: TRAINING TASK A ---
100
+ model = SimpleCNN()
101
+ criterion = nn.CrossEntropyLoss()
102
+ optimizer = optim.Adam(model.parameters(), lr=0.001)
103
+
104
+ print("--- Starting Task A (Digits 0-4) Training ---")
105
+ model.train()
106
+ for epoch in range(3):
107
+ for images, labels in train_loader_A:
108
+ optimizer.zero_grad()
109
+ outputs = model(images)
110
+ loss = criterion(outputs, labels)
111
+ loss.backward()
112
+ optimizer.step()
113
+
114
+ acc_A = evaluate_accuracy(model, test_loader_A)
115
+ print(f"Accuracy on Task A after training Task A: {acc_A:.2f}%\n")
116
+
117
+ # Freeze the base model to preserve its structural return invariants
118
+ model_A_frozen = copy.deepcopy(model)
119
+ model_A_frozen.eval()
120
+
121
+ # --- 5. PHASE 2: TRAINING TASK B WITH ANASTROPHIC REGULARIZATION ---
122
+ regularizer = AnastrophicRegularizer(lambda_reg=1.0, eta_reg=3.0)
123
+
124
+ print("--- Starting Task B (Digits 5-9) Training with R_ana ---")
125
+ optimizer_B = optim.Adam(model.parameters(), lr=0.001)
126
+
127
+ for epoch in range(3):
128
+ model.train()
129
+ running_loss_class = 0.0
130
+ running_loss_ana = 0.0
131
+
132
+ for images, labels in train_loader_B:
133
+ optimizer_B.zero_grad()
134
+ outputs = model(images)
135
+
136
+ loss_class = criterion(outputs, labels)
137
+
138
+ # Apply Anastrophic Theory to preserve structural relationships
139
+ loss_ana = regularizer(model, model_A_frozen)
140
+
141
+ loss = loss_class + loss_ana
142
+ loss.backward()
143
+ optimizer_B.step()
144
+
145
+ running_loss_class += loss_class.item()
146
+ running_loss_ana += loss_ana.item()
147
+
148
+ print(f"Epoch {epoch+1} | Classification Loss: {running_loss_class/len(train_loader_B):.4f} | R_ana Loss: {running_loss_ana/len(train_loader_B):.4f}")
149
+
150
+ # --- 6. FINAL EVALUATION ---
151
+ print("\n--- Final Results ---")
152
+ acc_B_final = evaluate_accuracy(model, test_loader_B)
153
+ acc_A_final = evaluate_accuracy(model, test_loader_A)
154
+
155
+ print(f"Accuracy on NEW Task B (5-9): {acc_B_final:.2f}%")
156
+ print(f"RETAINED Accuracy on Task A (0-4): {acc_A_final:.2f}%")