File size: 8,047 Bytes
07c235f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
"""CognitiveModes: Kuramoto phases as a detector of mental state.

THE INNOVATION. The Kuramoto oscillators aren't just a routing mechanism —
they're a DYNAMICAL SYSTEM whose phase pattern reflects the current "cognitive
mode" of the engine. This module:

    1. Extracts features from the phase vector (synchronization, clustering).
    2. Clusters phase patterns into cognitive modes (UNSUPERVISED — the modes
       emerge from the data, not from external labels).
    3. Lets the engine ADAPT its behavior based on its current mode.

This is what makes Fractus feel ALIVE — it has mental states that change how
it processes information, like a human shifting between focused work and
creative brainstorming.

UNSUPERVISED APPROACH (replaces the original supervised MLP):
    Instead of labelling phases with mode names (which is arbitrary), we collect
    phase features during a training run and cluster them with k-means. The
    clusters that emerge ARE the cognitive modes — defined by their centroids
    in the (synchronization, mean_phase, variance, sin/cos) feature space. Mode
    names are assigned a posteriori by interpreting the cluster characteristics
    (high sync = "focused", low sync = "exploratory", etc.).

Usage:
    modes = CognitiveModes(n_oscillators=8, n_modes=4)
    # Collect phases during training, then fit:
    modes.fit(phase_samples)  # phase_samples: (N_samples, n_oscillators)
    # Classify at runtime:
    mode = modes.classify(phases)  # → {"mode": "cluster_0", "confidence": 0.82, ...}
"""


import torch
import torch.nn as nn


class CognitiveModes(nn.Module):
    """Classify the Kuramoto phase state into cognitive modes via clustering.

    Modes are discovered unsupervised via k-means on phase features. No labels,
    no MLP — the clusters emerge from the structure of the phase space.

    Args:
        n_oscillators: number of Kuramoto oscillators.
        n_modes: number of modes (= k-means clusters).
        mode_names: optional names (assigned after fit by interpretation).
    """

    def __init__(
        self,
        n_oscillators: int = 8,
        n_modes: int = 4,
        mode_names: list = None,
    ):
        super().__init__()
        self.n_oscillators = n_oscillators
        self.n_modes = n_modes
        if mode_names is None:
            mode_names = [f"mode_{i}" for i in range(n_modes)]
        self.mode_names = mode_names[:n_modes]
        self.n_features = 3 + 2 * n_oscillators

        # Centroids: learned via k-means during fit(). Stored as a buffer.
        self.register_buffer("centroids", torch.zeros(n_modes, self.n_features))
        self._fitted = False

    def extract_features(self, phases: torch.Tensor) -> torch.Tensor:
        """Extract cognitive features from the phase vector.

        Args:
            phases: (..., N) oscillator phases in [0, 2Ï€).
        Returns:
            features: (..., 3 + 2*N) feature vector.
        """
        *leading, N = phases.shape
        phases_flat = phases.reshape(-1, N)  # (B, N)

        sin_p = torch.sin(phases_flat)
        cos_p = torch.cos(phases_flat)

        # Feature 1: order parameter r (synchronization degree).
        r = torch.sqrt(cos_p.mean(dim=-1) ** 2 + sin_p.mean(dim=-1) ** 2 + 1e-12)

        # Feature 2: mean phase.
        mean_phase = torch.atan2(sin_p.mean(dim=-1), cos_p.mean(dim=-1))

        # Feature 3: phase variance.
        phase_var = sin_p.var(dim=-1) + cos_p.var(dim=-1)

        # Features 4+: per-oscillator sin/cos.
        osc_features = torch.cat([sin_p, cos_p], dim=-1)  # (B, 2N)

        features = torch.cat([
            r.unsqueeze(-1),
            mean_phase.unsqueeze(-1),
            phase_var.unsqueeze(-1),
            osc_features,
        ], dim=-1)  # (B, 3 + 2N)

        return features.reshape(*leading, features.shape[-1])

    def fit(self, phase_samples: torch.Tensor, n_iters: int = 50) -> dict:
        """Fit k-means on collected phase samples (unsupervised).

        Args:
            phase_samples: (N_samples, n_oscillators) phases collected during training.
            n_iters: k-means iterations.
        Returns:
            dict with cluster info for interpretation.
        """
        features = self.extract_features(phase_samples)  # (N_samples, n_features)
        N = features.shape[0]
        K = self.n_modes

        if N < K:
            # Not enough samples — pad with noise.
            features = torch.cat([features, torch.randn(K - N, self.n_features)], dim=0)
            N = K

        # Initialize centroids: random samples.
        idx = torch.randperm(N)[:K]
        self.centroids = features[idx].clone()

        for _ in range(n_iters):
            # Assign each sample to nearest centroid (cosine distance).
            # Normalize for cosine.
            feat_norm = features / (features.norm(dim=-1, keepdim=True) + 1e-8)
            cent_norm = self.centroids / (self.centroids.norm(dim=-1, keepdim=True) + 1e-8)
            sims = feat_norm @ cent_norm.T  # (N, K) cosine similarity
            assignments = sims.argmax(dim=-1)  # (N,)

            # Update centroids.
            for k in range(K):
                mask = assignments == k
                if mask.any():
                    self.centroids[k] = features[mask].mean(dim=0)

        self._fitted = True

        # Compute cluster statistics for interpretation.
        cluster_info = {}
        for k in range(K):
            mask = assignments == k
            if mask.any():
                cluster_features = features[mask]
                cluster_info[k] = {
                    "size": mask.sum().item(),
                    "mean_sync": cluster_features[:, 0].mean().item(),  # r
                    "mean_var": cluster_features[:, 2].mean().item(),
                }
            else:
                cluster_info[k] = {"size": 0, "mean_sync": 0, "mean_var": 0}
        return cluster_info

    def classify(self, phases: torch.Tensor) -> dict:
        """Classify the current cognitive mode (nearest centroid).

        Args:
            phases: (N,) or (1, N) or (..., N) oscillator phases.
        Returns:
            dict with "mode" (str), "confidence" (float), and "all_modes" (dict).
        """
        if phases.dim() == 1:
            phases = phases.unsqueeze(0)
        features = self.extract_features(phases)  # (1, n_features)

        if not self._fitted:
            # Before fitting, return uniform.
            return {
                "mode": "unfitted",
                "confidence": 1.0 / self.n_modes,
                "all_modes": {name: 1.0 / self.n_modes for name in self.mode_names},
            }

        # Cosine similarity to each centroid.
        feat_norm = features[0] / (features[0].norm() + 1e-8)
        cent_norm = self.centroids / (self.centroids.norm(dim=-1, keepdim=True) + 1e-8)
        sims = cent_norm @ feat_norm  # (K,)
        probs = torch.softmax(sims * 5.0, dim=-1)  # temperature-scaled

        top_idx = probs.argmax(dim=-1).item()
        top_prob = probs[top_idx].item()
        mode_name = self.mode_names[top_idx] if top_idx < len(self.mode_names) else f"mode_{top_idx}"

        all_modes = {
            (self.mode_names[i] if i < len(self.mode_names) else f"mode_{i}"): probs[i].item()
            for i in range(self.n_modes)
        }

        return {
            "mode": mode_name,
            "confidence": top_prob,
            "all_modes": all_modes,
        }

    def label_modes(self, names: list):
        """Assign human-readable names to clusters after fitting (a posteriori).

        Args:
            names: list of n_modes names, in cluster order.
        """
        if len(names) != self.n_modes:
            raise ValueError(f"expected {self.n_modes} names, got {len(names)}")
        self.mode_names = names

    def info(self) -> dict:
        return {
            "n_oscillators": self.n_oscillators,
            "modes": self.mode_names,
            "n_features": self.n_features,
            "fitted": self._fitted,
        }