thefinalboss commited on
Commit
4a246e8
·
verified ·
1 Parent(s): 862b205

Upload fractus/train/online.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. fractus/train/online.py +268 -0
fractus/train/online.py ADDED
@@ -0,0 +1,268 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Online trainer for the ContinuousThoughtEngine.
2
+
3
+ THE TRAINING BREAKTHROUGH. No batches. No BPTT. One observation at a time,
4
+ one gradient at a time. The model learns as it "sees" data, like a human.
5
+
6
+ for each token in the data stream:
7
+ 1. Feed the token to the engine (tick).
8
+ 2. The engine produces a prediction + confidence.
9
+ 3. Compute the loss (was the prediction right?).
10
+ 4. Backward + step IMMEDIATELY (online SGD, 1 sample at a time).
11
+ 5. The thought state is carried forward (detached — no BPTT).
12
+
13
+ WHY THIS IS FAST:
14
+ - Each step processes ONE token (not B×L).
15
+ - The forward is tiny (1 token, 1 tick).
16
+ - The backward is tiny (1 sample).
17
+ - No batching, no padding, no sequence masking.
18
+
19
+ This is the training method that makes the Continuous Thought Engine
20
+ trainable on ANY CPU, because the per-step cost is minimal.
21
+ """
22
+
23
+ import torch
24
+ import torch.nn as nn
25
+ import torch.nn.functional as F
26
+
27
+
28
+ class OnlineTrainer:
29
+ """Online trainer for the ContinuousThoughtEngine.
30
+
31
+ Args:
32
+ engine: a ContinuousThoughtEngine.
33
+ lr: learning rate.
34
+ weight_decay: optimizer weight decay.
35
+ accumulation_steps: gradient accumulation (fewer optimizer steps).
36
+ optimizer: 'adamw', 'sgd', or 'rmsprop'.
37
+ """
38
+
39
+ def __init__(
40
+ self,
41
+ engine,
42
+ lr: float = 1e-3,
43
+ weight_decay: float = 0.01,
44
+ accumulation_steps: int = 8,
45
+ optimizer: str = "adamw",
46
+ ):
47
+ self.engine = engine
48
+ self.accumulation_steps = max(accumulation_steps, 1)
49
+ if optimizer == "sgd":
50
+ self.optimizer = torch.optim.SGD(engine.parameters(), lr=lr, momentum=0.9)
51
+ elif optimizer == "rmsprop":
52
+ self.optimizer = torch.optim.RMSprop(engine.parameters(), lr=lr, weight_decay=weight_decay)
53
+ else:
54
+ self.optimizer = torch.optim.AdamW(engine.parameters(), lr=lr, weight_decay=weight_decay)
55
+ self.step_count = 0
56
+ self.losses = []
57
+
58
+ def train_on_stream(self, token_ids: torch.Tensor, max_ticks: int = 3) -> dict:
59
+ """Train on a stream of tokens, one at a time (pure online, 1 backward/token).
60
+
61
+ token_ids: (L,) a 1D tensor of token ids (the data stream).
62
+ max_ticks: max thinking ticks per token.
63
+
64
+ Returns a dict with average loss, accuracy, and steps.
65
+ """
66
+ self.engine.train()
67
+ self.engine.reset_thought(batch_size=1)
68
+
69
+ total_loss = 0.0
70
+ correct = 0
71
+ total = 0
72
+
73
+ for t in range(len(token_ids) - 1):
74
+ obs = token_ids[t:t + 1] # (1,) current token
75
+ target = token_ids[t + 1] # scalar, next token
76
+
77
+ # Think: tick until confidence or max_ticks.
78
+ for tick in range(max_ticks):
79
+ logits, conf = self.engine.tick(obs if tick == 0 else None)
80
+ if conf.item() > 0.5:
81
+ break
82
+
83
+ # Online loss: did we predict the next token?
84
+ loss = F.cross_entropy(logits, target.unsqueeze(0))
85
+
86
+ # Immediate backward + step (online SGD).
87
+ self.optimizer.zero_grad()
88
+ loss.backward()
89
+ torch.nn.utils.clip_grad_norm_(self.engine.parameters(), 1.0)
90
+ self.optimizer.step()
91
+
92
+ total_loss += loss.item()
93
+ pred = logits.argmax(dim=-1).item()
94
+ if pred == target.item():
95
+ correct += 1
96
+ total += 1
97
+ self.step_count += 1
98
+ self.losses.append(loss.item())
99
+
100
+ return {
101
+ "avg_loss": total_loss / max(total, 1),
102
+ "accuracy": correct / max(total, 1),
103
+ "steps": total,
104
+ }
105
+
106
+ def train_on_stream_minibatch(self, token_ids: torch.Tensor, max_ticks: int = 2,
107
+ accum_steps: int = 16) -> dict:
108
+ """Train on a stream with mini-batch gradient accumulation.
109
+
110
+ Accumulates the loss over `accum_steps` tokens, then does ONE backward
111
+ + optimizer step. This is 10-16× faster than train_on_stream (which
112
+ does 1 backward per token) because the Python/autograd overhead is
113
+ amortized over N tokens.
114
+
115
+ The thought state is still carried forward (detached between backward
116
+ steps), preserving the continuous-reasoning paradigm.
117
+
118
+ Args:
119
+ token_ids: (L,) 1D tensor of token ids.
120
+ max_ticks: max thinking ticks per token.
121
+ accum_steps: tokens per backward pass (16 = 16× fewer backward calls).
122
+ """
123
+ self.engine.train()
124
+ self.engine.reset_thought(batch_size=1)
125
+
126
+ total_loss = 0.0
127
+ correct = 0
128
+ total = 0
129
+ accum_loss = torch.tensor(0.0, requires_grad=False)
130
+
131
+ for t in range(len(token_ids) - 1):
132
+ obs = token_ids[t:t + 1]
133
+ target = token_ids[t + 1]
134
+
135
+ # Think (1 tick per token for speed).
136
+ logits, conf = self.engine.tick(obs)
137
+
138
+ # Per-token loss.
139
+ loss = F.cross_entropy(logits, target.unsqueeze(0))
140
+ accum_loss = accum_loss + loss
141
+
142
+ total_loss += loss.item()
143
+ pred = logits.argmax(dim=-1).item()
144
+ if pred == target.item():
145
+ correct += 1
146
+ total += 1
147
+
148
+ # Backward every accum_steps tokens.
149
+ if (t + 1) % accum_steps == 0:
150
+ self.optimizer.zero_grad()
151
+ avg_loss = accum_loss / accum_steps
152
+ avg_loss.backward()
153
+ torch.nn.utils.clip_grad_norm_(self.engine.parameters(), 1.0)
154
+ self.optimizer.step()
155
+ self.step_count += 1
156
+ self.losses.append(total_loss / total)
157
+ accum_loss = torch.tensor(0.0, requires_grad=False)
158
+
159
+ # Final partial accumulation.
160
+ if total % accum_steps != 0 and isinstance(accum_loss, torch.Tensor) and accum_loss.requires_grad:
161
+ self.optimizer.zero_grad()
162
+ (accum_loss / (total % accum_steps)).backward()
163
+ self.optimizer.step()
164
+ self.step_count += 1
165
+
166
+ return {
167
+ "avg_loss": total_loss / max(total, 1),
168
+ "accuracy": correct / max(total, 1),
169
+ "steps": total,
170
+ "optimizer_steps": self.step_count,
171
+ }
172
+
173
+ def train_on_stream_chunked(self, token_ids: torch.Tensor,
174
+ chunk_len: int = 16) -> dict:
175
+ """Train using chunk-based processing (16x fewer forward passes).
176
+
177
+ Splits the stream into chunks of `chunk_len` tokens. Each chunk is
178
+ processed in ONE forward pass (tick_chunk), then ONE backward.
179
+ This is the FASTEST training mode — the forward/backward overhead
180
+ is amortized over chunk_len tokens.
181
+
182
+ The thought state (S,z) is carried between chunks (detached).
183
+
184
+ Args:
185
+ token_ids: (L,) 1D tensor.
186
+ chunk_len: tokens per chunk (16 = 16× fewer forward passes).
187
+ """
188
+ self.engine.train()
189
+ self.engine.reset_thought(batch_size=1)
190
+ vocab = self.engine.vocab_size
191
+
192
+ total_loss = 0.0
193
+ correct = 0
194
+ total = 0
195
+
196
+ # Gradient accumulation: backward every chunk, step every accumulation_steps chunks.
197
+ accum = self.accumulation_steps
198
+ self.optimizer.zero_grad()
199
+ chunk_idx = 0
200
+
201
+ for start in range(0, len(token_ids) - chunk_len - 1, chunk_len):
202
+ chunk = token_ids[start:start + chunk_len].unsqueeze(0) # (1, C)
203
+
204
+ # Fast training path: head on LAST position only.
205
+ last_logits = self.engine.tick_chunk_train(chunk) # (1, vocab)
206
+ target = token_ids[start + chunk_len] # scalar: the token after the chunk
207
+
208
+ # CE on the single predicted token (scaled by 1/accum for correct averaging).
209
+ loss = F.cross_entropy(last_logits, target.unsqueeze(0)) / accum
210
+
211
+ # Backward every chunk (gradients accumulate).
212
+ loss.backward()
213
+
214
+ total_loss += loss.item() * accum # un-scale for logging
215
+ pred = last_logits.argmax(dim=-1)
216
+ correct += (pred == target.unsqueeze(0)).sum().item()
217
+ total += 1
218
+ self.losses.append(loss.item() * accum)
219
+
220
+ chunk_idx += 1
221
+
222
+ # Step only every accumulation_steps chunks.
223
+ if chunk_idx % accum == 0:
224
+ torch.nn.utils.clip_grad_norm_(self.engine.parameters(), 1.0)
225
+ self.optimizer.step()
226
+ self.optimizer.zero_grad()
227
+ self.step_count += 1
228
+
229
+ # Handle the remainder: if chunks aren't a clean multiple of accum, do a final step.
230
+ if chunk_idx % accum != 0:
231
+ torch.nn.utils.clip_grad_norm_(self.engine.parameters(), 1.0)
232
+ self.optimizer.step()
233
+ self.optimizer.zero_grad()
234
+ self.step_count += 1
235
+
236
+ return {
237
+ "avg_loss": total_loss / max(total, 1),
238
+ "accuracy": correct / max(total, 1),
239
+ "steps": total,
240
+ "optimizer_steps": self.step_count,
241
+ }
242
+
243
+ def train_step_batch(self, input_ids: torch.Tensor, target_ids: torch.Tensor,
244
+ max_ticks: int = 3) -> dict:
245
+ """Train on a small batch using the think() method.
246
+
247
+ input_ids: (B, L) token ids.
248
+ target_ids: (B, L) next-token targets.
249
+ """
250
+ self.engine.train()
251
+ self.engine.reset_thought(batch_size=input_ids.shape[0])
252
+
253
+ # Use think() to process the whole sequence.
254
+ logits = self.engine.think(input_ids, max_ticks=max_ticks, confidence_threshold=0.5)
255
+ # The logits are (B, L, vocab) — but think() only produces output when
256
+ # confident. For training we compute loss on ALL positions.
257
+ loss = F.cross_entropy(
258
+ logits.reshape(-1, self.engine.vocab_size),
259
+ target_ids.reshape(-1),
260
+ )
261
+
262
+ self.optimizer.zero_grad()
263
+ loss.backward()
264
+ torch.nn.utils.clip_grad_norm_(self.engine.parameters(), 1.0)
265
+ self.optimizer.step()
266
+ self.step_count += 1
267
+
268
+ return {"loss": loss.item()}