thefinalboss commited on
Commit
07c235f
·
verified ·
1 Parent(s): cee9d4a

Upload fractus/cognitive_modes.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. fractus/cognitive_modes.py +212 -0
fractus/cognitive_modes.py ADDED
@@ -0,0 +1,212 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """CognitiveModes: Kuramoto phases as a detector of mental state.
2
+
3
+ THE INNOVATION. The Kuramoto oscillators aren't just a routing mechanism —
4
+ they're a DYNAMICAL SYSTEM whose phase pattern reflects the current "cognitive
5
+ mode" of the engine. This module:
6
+
7
+ 1. Extracts features from the phase vector (synchronization, clustering).
8
+ 2. Clusters phase patterns into cognitive modes (UNSUPERVISED — the modes
9
+ emerge from the data, not from external labels).
10
+ 3. Lets the engine ADAPT its behavior based on its current mode.
11
+
12
+ This is what makes Fractus feel ALIVE — it has mental states that change how
13
+ it processes information, like a human shifting between focused work and
14
+ creative brainstorming.
15
+
16
+ UNSUPERVISED APPROACH (replaces the original supervised MLP):
17
+ Instead of labelling phases with mode names (which is arbitrary), we collect
18
+ phase features during a training run and cluster them with k-means. The
19
+ clusters that emerge ARE the cognitive modes — defined by their centroids
20
+ in the (synchronization, mean_phase, variance, sin/cos) feature space. Mode
21
+ names are assigned a posteriori by interpreting the cluster characteristics
22
+ (high sync = "focused", low sync = "exploratory", etc.).
23
+
24
+ Usage:
25
+ modes = CognitiveModes(n_oscillators=8, n_modes=4)
26
+ # Collect phases during training, then fit:
27
+ modes.fit(phase_samples) # phase_samples: (N_samples, n_oscillators)
28
+ # Classify at runtime:
29
+ mode = modes.classify(phases) # → {"mode": "cluster_0", "confidence": 0.82, ...}
30
+ """
31
+
32
+
33
+ import torch
34
+ import torch.nn as nn
35
+
36
+
37
+ class CognitiveModes(nn.Module):
38
+ """Classify the Kuramoto phase state into cognitive modes via clustering.
39
+
40
+ Modes are discovered unsupervised via k-means on phase features. No labels,
41
+ no MLP — the clusters emerge from the structure of the phase space.
42
+
43
+ Args:
44
+ n_oscillators: number of Kuramoto oscillators.
45
+ n_modes: number of modes (= k-means clusters).
46
+ mode_names: optional names (assigned after fit by interpretation).
47
+ """
48
+
49
+ def __init__(
50
+ self,
51
+ n_oscillators: int = 8,
52
+ n_modes: int = 4,
53
+ mode_names: list = None,
54
+ ):
55
+ super().__init__()
56
+ self.n_oscillators = n_oscillators
57
+ self.n_modes = n_modes
58
+ if mode_names is None:
59
+ mode_names = [f"mode_{i}" for i in range(n_modes)]
60
+ self.mode_names = mode_names[:n_modes]
61
+ self.n_features = 3 + 2 * n_oscillators
62
+
63
+ # Centroids: learned via k-means during fit(). Stored as a buffer.
64
+ self.register_buffer("centroids", torch.zeros(n_modes, self.n_features))
65
+ self._fitted = False
66
+
67
+ def extract_features(self, phases: torch.Tensor) -> torch.Tensor:
68
+ """Extract cognitive features from the phase vector.
69
+
70
+ Args:
71
+ phases: (..., N) oscillator phases in [0, 2π).
72
+ Returns:
73
+ features: (..., 3 + 2*N) feature vector.
74
+ """
75
+ *leading, N = phases.shape
76
+ phases_flat = phases.reshape(-1, N) # (B, N)
77
+
78
+ sin_p = torch.sin(phases_flat)
79
+ cos_p = torch.cos(phases_flat)
80
+
81
+ # Feature 1: order parameter r (synchronization degree).
82
+ r = torch.sqrt(cos_p.mean(dim=-1) ** 2 + sin_p.mean(dim=-1) ** 2 + 1e-12)
83
+
84
+ # Feature 2: mean phase.
85
+ mean_phase = torch.atan2(sin_p.mean(dim=-1), cos_p.mean(dim=-1))
86
+
87
+ # Feature 3: phase variance.
88
+ phase_var = sin_p.var(dim=-1) + cos_p.var(dim=-1)
89
+
90
+ # Features 4+: per-oscillator sin/cos.
91
+ osc_features = torch.cat([sin_p, cos_p], dim=-1) # (B, 2N)
92
+
93
+ features = torch.cat([
94
+ r.unsqueeze(-1),
95
+ mean_phase.unsqueeze(-1),
96
+ phase_var.unsqueeze(-1),
97
+ osc_features,
98
+ ], dim=-1) # (B, 3 + 2N)
99
+
100
+ return features.reshape(*leading, features.shape[-1])
101
+
102
+ def fit(self, phase_samples: torch.Tensor, n_iters: int = 50) -> dict:
103
+ """Fit k-means on collected phase samples (unsupervised).
104
+
105
+ Args:
106
+ phase_samples: (N_samples, n_oscillators) phases collected during training.
107
+ n_iters: k-means iterations.
108
+ Returns:
109
+ dict with cluster info for interpretation.
110
+ """
111
+ features = self.extract_features(phase_samples) # (N_samples, n_features)
112
+ N = features.shape[0]
113
+ K = self.n_modes
114
+
115
+ if N < K:
116
+ # Not enough samples — pad with noise.
117
+ features = torch.cat([features, torch.randn(K - N, self.n_features)], dim=0)
118
+ N = K
119
+
120
+ # Initialize centroids: random samples.
121
+ idx = torch.randperm(N)[:K]
122
+ self.centroids = features[idx].clone()
123
+
124
+ for _ in range(n_iters):
125
+ # Assign each sample to nearest centroid (cosine distance).
126
+ # Normalize for cosine.
127
+ feat_norm = features / (features.norm(dim=-1, keepdim=True) + 1e-8)
128
+ cent_norm = self.centroids / (self.centroids.norm(dim=-1, keepdim=True) + 1e-8)
129
+ sims = feat_norm @ cent_norm.T # (N, K) cosine similarity
130
+ assignments = sims.argmax(dim=-1) # (N,)
131
+
132
+ # Update centroids.
133
+ for k in range(K):
134
+ mask = assignments == k
135
+ if mask.any():
136
+ self.centroids[k] = features[mask].mean(dim=0)
137
+
138
+ self._fitted = True
139
+
140
+ # Compute cluster statistics for interpretation.
141
+ cluster_info = {}
142
+ for k in range(K):
143
+ mask = assignments == k
144
+ if mask.any():
145
+ cluster_features = features[mask]
146
+ cluster_info[k] = {
147
+ "size": mask.sum().item(),
148
+ "mean_sync": cluster_features[:, 0].mean().item(), # r
149
+ "mean_var": cluster_features[:, 2].mean().item(),
150
+ }
151
+ else:
152
+ cluster_info[k] = {"size": 0, "mean_sync": 0, "mean_var": 0}
153
+ return cluster_info
154
+
155
+ def classify(self, phases: torch.Tensor) -> dict:
156
+ """Classify the current cognitive mode (nearest centroid).
157
+
158
+ Args:
159
+ phases: (N,) or (1, N) or (..., N) oscillator phases.
160
+ Returns:
161
+ dict with "mode" (str), "confidence" (float), and "all_modes" (dict).
162
+ """
163
+ if phases.dim() == 1:
164
+ phases = phases.unsqueeze(0)
165
+ features = self.extract_features(phases) # (1, n_features)
166
+
167
+ if not self._fitted:
168
+ # Before fitting, return uniform.
169
+ return {
170
+ "mode": "unfitted",
171
+ "confidence": 1.0 / self.n_modes,
172
+ "all_modes": {name: 1.0 / self.n_modes for name in self.mode_names},
173
+ }
174
+
175
+ # Cosine similarity to each centroid.
176
+ feat_norm = features[0] / (features[0].norm() + 1e-8)
177
+ cent_norm = self.centroids / (self.centroids.norm(dim=-1, keepdim=True) + 1e-8)
178
+ sims = cent_norm @ feat_norm # (K,)
179
+ probs = torch.softmax(sims * 5.0, dim=-1) # temperature-scaled
180
+
181
+ top_idx = probs.argmax(dim=-1).item()
182
+ top_prob = probs[top_idx].item()
183
+ mode_name = self.mode_names[top_idx] if top_idx < len(self.mode_names) else f"mode_{top_idx}"
184
+
185
+ all_modes = {
186
+ (self.mode_names[i] if i < len(self.mode_names) else f"mode_{i}"): probs[i].item()
187
+ for i in range(self.n_modes)
188
+ }
189
+
190
+ return {
191
+ "mode": mode_name,
192
+ "confidence": top_prob,
193
+ "all_modes": all_modes,
194
+ }
195
+
196
+ def label_modes(self, names: list):
197
+ """Assign human-readable names to clusters after fitting (a posteriori).
198
+
199
+ Args:
200
+ names: list of n_modes names, in cluster order.
201
+ """
202
+ if len(names) != self.n_modes:
203
+ raise ValueError(f"expected {self.n_modes} names, got {len(names)}")
204
+ self.mode_names = names
205
+
206
+ def info(self) -> dict:
207
+ return {
208
+ "n_oscillators": self.n_oscillators,
209
+ "modes": self.mode_names,
210
+ "n_features": self.n_features,
211
+ "fitted": self._fitted,
212
+ }