thefinalboss commited on
Commit
a834d38
·
verified ·
1 Parent(s): 7764b77

Upload fractus/continuous_engine.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. 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: advance phases by one Euler step.
 
124
  h_kur = self.norm_kur(h)
125
  theta = self.kuramoto._encode_from_hidden(h_kur)
126
- theta = theta + 0.1 * self.kuramoto._derivative(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
- """Fast training: head on LAST position only. Returns logits (B, vocab)."""
 
 
 
 
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 # keep graph — load-balance must train
489
  self.last_lb_loss = total_lb
490
 
491
  self.thought_state = h[:, -1:, :].detach()
492
- last_logits = self.output_head(h[:, -1, :])
493
- return last_logits, total_lb
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: