File size: 8,967 Bytes
41c4bbc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
"""Online trainer for the ContinuousThoughtEngine.

THE TRAINING BREAKTHROUGH. No batches. No BPTT. One observation at a time,
one gradient at a time. The model learns as it "sees" data, like a human.

    for each token in the data stream:
        1. Feed the token to the engine (tick).
        2. The engine produces a prediction + confidence.
        3. Compute the loss (was the prediction right?).
        4. Backward + step IMMEDIATELY (online SGD, 1 sample at a time).
        5. The thought state is carried forward (detached — no BPTT).

WHY THIS IS FAST:
    - Each step processes ONE token (not B×L).
    - The forward is tiny (1 token, 1 tick).
    - The backward is tiny (1 sample).
    - No batching, no padding, no sequence masking.

This is the training method that makes the Continuous Thought Engine
trainable on ANY CPU, because the per-step cost is minimal.
"""

import torch
import torch.nn as nn
import torch.nn.functional as F


class OnlineTrainer:
    """Online SGD trainer for the ContinuousThoughtEngine.

    Args:
        engine:           a ContinuousThoughtEngine.
        lr:               learning rate.
        weight_decay:     AdamW weight decay.
        confidence_threshold: the engine emits output when confidence exceeds this.
    """

    def __init__(
        self,
        engine,
        lr: float = 1e-3,
        weight_decay: float = 0.01,
    ):
        self.engine = engine
        self.optimizer = torch.optim.AdamW(engine.parameters(), lr=lr,
                                           weight_decay=weight_decay)
        self.step_count = 0
        self.losses = []

    def train_on_stream(self, token_ids: torch.Tensor, max_ticks: int = 3) -> dict:
        """Train on a stream of tokens, one at a time (pure online, 1 backward/token).

        token_ids: (L,) a 1D tensor of token ids (the data stream).
        max_ticks: max thinking ticks per token.

        Returns a dict with average loss, accuracy, and steps.
        """
        self.engine.train()
        self.engine.reset_thought(batch_size=1)

        total_loss = 0.0
        correct = 0
        total = 0

        for t in range(len(token_ids) - 1):
            obs = token_ids[t:t + 1]  # (1,) current token
            target = token_ids[t + 1]  # scalar, next token

            # Think: tick until confidence or max_ticks.
            for tick in range(max_ticks):
                logits, conf = self.engine.tick(obs if tick == 0 else None)
                if conf.item() > 0.5:
                    break

            # Online loss: did we predict the next token?
            loss = F.cross_entropy(logits, target.unsqueeze(0))

            # Immediate backward + step (online SGD).
            self.optimizer.zero_grad()
            loss.backward()
            torch.nn.utils.clip_grad_norm_(self.engine.parameters(), 1.0)
            self.optimizer.step()

            total_loss += loss.item()
            pred = logits.argmax(dim=-1).item()
            if pred == target.item():
                correct += 1
            total += 1
            self.step_count += 1
            self.losses.append(loss.item())

        return {
            "avg_loss": total_loss / max(total, 1),
            "accuracy": correct / max(total, 1),
            "steps": total,
        }

    def train_on_stream_minibatch(self, token_ids: torch.Tensor, max_ticks: int = 2,
                                   accum_steps: int = 16) -> dict:
        """Train on a stream with mini-batch gradient accumulation.

        Accumulates the loss over `accum_steps` tokens, then does ONE backward
        + optimizer step. This is 10-16× faster than train_on_stream (which
        does 1 backward per token) because the Python/autograd overhead is
        amortized over N tokens.

        The thought state is still carried forward (detached between backward
        steps), preserving the continuous-reasoning paradigm.

        Args:
            token_ids:    (L,) 1D tensor of token ids.
            max_ticks:    max thinking ticks per token.
            accum_steps:  tokens per backward pass (16 = 16× fewer backward calls).
        """
        self.engine.train()
        self.engine.reset_thought(batch_size=1)

        total_loss = 0.0
        correct = 0
        total = 0
        accum_loss = torch.tensor(0.0, requires_grad=False)

        for t in range(len(token_ids) - 1):
            obs = token_ids[t:t + 1]
            target = token_ids[t + 1]

            # Think (1 tick per token for speed).
            logits, conf = self.engine.tick(obs)

            # Per-token loss.
            loss = F.cross_entropy(logits, target.unsqueeze(0))
            accum_loss = accum_loss + loss

            total_loss += loss.item()
            pred = logits.argmax(dim=-1).item()
            if pred == target.item():
                correct += 1
            total += 1

            # Backward every accum_steps tokens.
            if (t + 1) % accum_steps == 0:
                self.optimizer.zero_grad()
                avg_loss = accum_loss / accum_steps
                avg_loss.backward()
                torch.nn.utils.clip_grad_norm_(self.engine.parameters(), 1.0)
                self.optimizer.step()
                self.step_count += 1
                self.losses.append(total_loss / total)
                accum_loss = torch.tensor(0.0, requires_grad=False)

        # Final partial accumulation.
        if total % accum_steps != 0 and isinstance(accum_loss, torch.Tensor) and accum_loss.requires_grad:
            self.optimizer.zero_grad()
            (accum_loss / (total % accum_steps)).backward()
            self.optimizer.step()
            self.step_count += 1

        return {
            "avg_loss": total_loss / max(total, 1),
            "accuracy": correct / max(total, 1),
            "steps": total,
            "optimizer_steps": self.step_count,
        }

    def train_on_stream_chunked(self, token_ids: torch.Tensor,
                                chunk_len: int = 16) -> dict:
        """Train using chunk-based processing (16x fewer forward passes).

        Splits the stream into chunks of `chunk_len` tokens. Each chunk is
        processed in ONE forward pass (tick_chunk), then ONE backward.
        This is the FASTEST training mode — the forward/backward overhead
        is amortized over chunk_len tokens.

        The thought state (S,z) is carried between chunks (detached).

        Args:
            token_ids:  (L,) 1D tensor.
            chunk_len:  tokens per chunk (16 = 16× fewer forward passes).
        """
        self.engine.train()
        self.engine.reset_thought(batch_size=1)
        vocab = self.engine.vocab_size

        total_loss = 0.0
        correct = 0
        total = 0

        for start in range(0, len(token_ids) - chunk_len - 1, chunk_len):
            chunk = token_ids[start:start + chunk_len].unsqueeze(0)  # (1, C)
            target = token_ids[start + 1:start + chunk_len + 1].unsqueeze(0)  # (1, C)

            # One forward over the whole chunk.
            logits = self.engine.tick_chunk(chunk)  # (1, C, vocab)

            # Cross-entropy on all positions.
            loss = F.cross_entropy(logits.reshape(-1, vocab), target.reshape(-1))

            # One backward + step per chunk.
            self.optimizer.zero_grad()
            loss.backward()
            torch.nn.utils.clip_grad_norm_(self.engine.parameters(), 1.0)
            self.optimizer.step()
            self.step_count += 1

            total_loss += loss.item() * chunk_len
            preds = logits.argmax(dim=-1)
            correct += (preds == target).sum().item()
            total += chunk_len
            self.losses.append(loss.item())

        return {
            "avg_loss": total_loss / max(total, 1),
            "accuracy": correct / max(total, 1),
            "steps": total,
            "optimizer_steps": self.step_count,
        }

    def train_step_batch(self, input_ids: torch.Tensor, target_ids: torch.Tensor,
                         max_ticks: int = 3) -> dict:
        """Train on a small batch using the think() method.

        input_ids:  (B, L) token ids.
        target_ids: (B, L) next-token targets.
        """
        self.engine.train()
        self.engine.reset_thought(batch_size=input_ids.shape[0])

        # Use think() to process the whole sequence.
        logits = self.engine.think(input_ids, max_ticks=max_ticks, confidence_threshold=0.5)
        # The logits are (B, L, vocab) — but think() only produces output when
        # confident. For training we compute loss on ALL positions.
        loss = F.cross_entropy(
            logits.reshape(-1, self.engine.vocab_size),
            target_ids.reshape(-1),
        )

        self.optimizer.zero_grad()
        loss.backward()
        torch.nn.utils.clip_grad_norm_(self.engine.parameters(), 1.0)
        self.optimizer.step()
        self.step_count += 1

        return {"loss": loss.item()}