Upload fractus/continuous_engine.py with huggingface_hub
Browse files- fractus/continuous_engine.py +11 -7
fractus/continuous_engine.py
CHANGED
|
@@ -120,11 +120,11 @@ class CTEBlock(nn.Module):
|
|
| 120 |
attn_out = attn_out @ attn.w_out + attn.b_out
|
| 121 |
h = h + attn_out
|
| 122 |
|
| 123 |
-
# Kuramoto:
|
|
|
|
| 124 |
h_kur = self.norm_kur(h)
|
| 125 |
theta = self.kuramoto._encode_from_hidden(h_kur)
|
| 126 |
-
theta =
|
| 127 |
-
theta = torch.remainder(theta, self.kuramoto.TWO_PI)
|
| 128 |
self.kuramoto_phases = theta.detach()
|
| 129 |
|
| 130 |
# MoE: transform the thought, routed by Kuramoto phases.
|
|
@@ -475,7 +475,11 @@ class ContinuousThoughtEngine(nn.Module):
|
|
| 475 |
return output_logits
|
| 476 |
|
| 477 |
def tick_chunk_train(self, observations: torch.Tensor) -> torch.Tensor:
|
| 478 |
-
"""
|
|
|
|
|
|
|
|
|
|
|
|
|
| 479 |
B, C = observations.shape
|
| 480 |
|
| 481 |
obs_vecs = self.observe(observations)
|
|
@@ -485,12 +489,12 @@ class ContinuousThoughtEngine(nn.Module):
|
|
| 485 |
total_lb = torch.tensor(0.0, device=h.device)
|
| 486 |
for blk in self.blocks:
|
| 487 |
h, lb = blk.tick_chunk_core(h)
|
| 488 |
-
total_lb = total_lb + lb
|
| 489 |
self.last_lb_loss = total_lb
|
| 490 |
|
| 491 |
self.thought_state = h[:, -1:, :].detach()
|
| 492 |
-
|
| 493 |
-
return
|
| 494 |
|
| 495 |
def think(self, observations: torch.Tensor, max_ticks: int = 10,
|
| 496 |
confidence_threshold: float = 0.7) -> torch.Tensor:
|
|
|
|
| 120 |
attn_out = attn_out @ attn.w_out + attn.b_out
|
| 121 |
h = h + attn_out
|
| 122 |
|
| 123 |
+
# Kuramoto: ALIGNED with tick_chunk_core (RK4), not single Euler step.
|
| 124 |
+
# Train/gen mismatch was: gen used Euler 0.1, train used full RK4 integrate.
|
| 125 |
h_kur = self.norm_kur(h)
|
| 126 |
theta = self.kuramoto._encode_from_hidden(h_kur)
|
| 127 |
+
theta = self.kuramoto._rk4_integrate(theta)
|
|
|
|
| 128 |
self.kuramoto_phases = theta.detach()
|
| 129 |
|
| 130 |
# MoE: transform the thought, routed by Kuramoto phases.
|
|
|
|
| 475 |
return output_logits
|
| 476 |
|
| 477 |
def tick_chunk_train(self, observations: torch.Tensor) -> torch.Tensor:
|
| 478 |
+
"""Train over the full chunk: logits (B, C, vocab) for dense next-token CE.
|
| 479 |
+
|
| 480 |
+
Stage-2 surgery: supervise every position so the model must chain tokens,
|
| 481 |
+
not only collapse to a single last-position attractor.
|
| 482 |
+
"""
|
| 483 |
B, C = observations.shape
|
| 484 |
|
| 485 |
obs_vecs = self.observe(observations)
|
|
|
|
| 489 |
total_lb = torch.tensor(0.0, device=h.device)
|
| 490 |
for blk in self.blocks:
|
| 491 |
h, lb = blk.tick_chunk_core(h)
|
| 492 |
+
total_lb = total_lb + lb
|
| 493 |
self.last_lb_loss = total_lb
|
| 494 |
|
| 495 |
self.thought_state = h[:, -1:, :].detach()
|
| 496 |
+
logits = self.output_head(h) # (B, C, vocab)
|
| 497 |
+
return logits, total_lb
|
| 498 |
|
| 499 |
def think(self, observations: torch.Tensor, max_ticks: int = 10,
|
| 500 |
confidence_threshold: float = 0.7) -> torch.Tensor:
|