thefinalboss commited on
Commit
a3a842c
·
verified ·
1 Parent(s): 50c3406

Upload fractus/decode_surgery.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. fractus/decode_surgery.py +64 -0
fractus/decode_surgery.py ADDED
@@ -0,0 +1,64 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Decode-time surgery for Fractus CTE.
2
+
3
+ Breaks single-token and short-cycle attractors without changing weights.
4
+ Techniques: phase noise, thought noise, recent-token bans, frequency penalty,
5
+ cycle detection with hard scramble, periodic forced escape tokens.
6
+ """
7
+ import math
8
+ import random
9
+ import torch
10
+
11
+ def generate_with_surgery(engine, tokenizer, prompt, max_new=48, temperature=1.15,
12
+ phase_noise=0.55, thought_noise=0.08, ban_window=20,
13
+ escape_every=4, freq_penalty=1.2):
14
+ engine.eval()
15
+ for blk in engine.blocks:
16
+ if hasattr(blk, 'moe'):
17
+ blk.moe.temperature = max(getattr(blk.moe, 'temperature', 1.0), 3.0)
18
+ engine.reset_thought(1)
19
+ ids = tokenizer.encode(prompt)[:64]
20
+ with torch.no_grad():
21
+ for t in ids:
22
+ engine.tick(torch.tensor([t]))
23
+ out = []
24
+ cur = ids[-1] if ids else 0
25
+ freq = {}
26
+ for step in range(max_new):
27
+ for blk in engine.blocks:
28
+ if hasattr(blk, 'kuramoto_phases'):
29
+ blk.kuramoto_phases = torch.remainder(
30
+ blk.kuramoto_phases + phase_noise * torch.randn_like(blk.kuramoto_phases),
31
+ 2 * math.pi)
32
+ engine.thought_state = 0.9 * engine.thought_state + thought_noise * torch.randn_like(engine.thought_state)
33
+ if step > 0 and escape_every and step % escape_every == 0:
34
+ banned = set(out[-ban_window:])
35
+ esc = random.randint(0, 50256)
36
+ for _ in range(40):
37
+ esc = random.randint(0, 50256)
38
+ if esc not in banned and esc != 50256:
39
+ break
40
+ engine.tick(torch.tensor([esc]))
41
+ out.append(esc)
42
+ freq[esc] = freq.get(esc, 0) + 1
43
+ cur = esc
44
+ for blk in engine.blocks:
45
+ if hasattr(blk, 'kuramoto_phases'):
46
+ blk.kuramoto_phases = torch.rand_like(blk.kuramoto_phases) * 2 * math.pi
47
+ continue
48
+ logits, _ = engine.tick(torch.tensor([cur]))
49
+ l = logits[0].float().clone()
50
+ for prev in out[-ban_window:]:
51
+ l[prev] = -1e9
52
+ for tid, c in freq.items():
53
+ l[tid] -= freq_penalty * c
54
+ topv, topi = torch.topk(l / max(temperature, 1e-5), 150)
55
+ mask = torch.isfinite(topv)
56
+ topv, topi = topv[mask], topi[mask]
57
+ if topv.numel() == 0:
58
+ nxt = int(torch.argmax(logits[0]).item())
59
+ else:
60
+ nxt = int(topi[torch.multinomial(torch.softmax(topv, -1), 1)].item())
61
+ out.append(nxt)
62
+ freq[nxt] = freq.get(nxt, 0) + 1
63
+ cur = nxt
64
+ return tokenizer.decode(out), out