thefinalboss commited on
Commit
eb8e976
·
verified ·
1 Parent(s): 771242b

docs+code: decode surgery I window-64 anti-copy 2026-08-28

Browse files
Files changed (1) hide show
  1. fractus/generate_aligned.py +94 -56
fractus/generate_aligned.py CHANGED
@@ -1,53 +1,39 @@
1
  """Train-aligned generation for Fractus CTE.
2
 
3
- Uses tick_chunk only (same path as stage2 training), never tick_single.
 
 
 
4
  """
5
  from __future__ import annotations
6
- import torch
 
7
  from typing import List, Optional
8
 
9
- @torch.no_grad()
10
- def generate_chunk(
11
- engine,
12
- tokenizer,
13
- prompt: str,
14
- max_new: int = 40,
15
- temperature: float = 0.8,
16
- top_k: int = 40,
17
- ban_window: int = 8,
18
- ban_factor: float = 0.4,
19
- context_limit: int = 128,
20
- ) -> tuple[str, List[int]]:
21
  engine.eval()
22
- engine.reset_thought(1)
23
- for blk in engine.blocks:
24
- if hasattr(blk, 'attn_S'):
 
25
  blk.attn_S.zero_()
26
- if hasattr(blk, 'attn_z'):
27
  blk.attn_z.zero_()
28
 
29
- ids = tokenizer.encode(prompt)[:context_limit]
30
- if not ids:
31
- ids = [0]
32
 
33
- # warm full prompt as one chunk
34
- logits = engine.tick_chunk(torch.tensor([ids], dtype=torch.long))
35
- cur = logits[0, -1]
36
- out: List[int] = []
37
-
38
- for _ in range(max_new):
39
- l = cur.float() / max(temperature, 1e-5)
40
- for prev in set(out[-ban_window:]):
41
- l[prev] *= ban_factor
42
- k = min(top_k, l.size(-1))
43
- topv, topi = torch.topk(l, k)
44
- nxt = int(topi[torch.multinomial(torch.softmax(topv, -1), 1)].item())
45
- out.append(nxt)
46
- # advance with train path (length-1 chunk)
47
- logits = engine.tick_chunk(torch.tensor([[nxt]], dtype=torch.long))
48
- cur = logits[0, -1]
49
-
50
- return tokenizer.decode(out), out
51
 
52
 
53
  @torch.no_grad()
@@ -56,30 +42,82 @@ def generate_window(
56
  tokenizer,
57
  prompt: str,
58
  max_new: int = 40,
59
- temperature: float = 0.8,
60
  top_k: int = 40,
61
  window: int = 64,
62
- ban_window: int = 8,
63
- ban_factor: float = 0.4,
64
  ) -> tuple[str, List[int]]:
65
- """Re-encode last tokens each step (fresh causal context)."""
66
- engine.eval()
67
- ids = tokenizer.encode(prompt)[:window]
68
  out: List[int] = []
 
69
  for _ in range(max_new):
70
- engine.reset_thought(1)
71
- for blk in engine.blocks:
72
- if hasattr(blk, 'attn_S'):
73
- blk.attn_S.zero_()
74
- if hasattr(blk, 'attn_z'):
75
- blk.attn_z.zero_()
76
  ctx = ids[-window:]
77
  logits = engine.tick_chunk(torch.tensor([ctx], dtype=torch.long))
78
- l = logits[0, -1].float() / max(temperature, 1e-5)
79
- for prev in set(out[-ban_window:]):
80
- l[prev] *= ban_factor
81
- topv, topi = torch.topk(l, min(top_k, l.size(-1)))
82
- nxt = int(topi[torch.multinomial(torch.softmax(topv, -1), 1)].item())
83
  out.append(nxt)
84
  ids.append(nxt)
 
85
  return tokenizer.decode(out), out
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  """Train-aligned generation for Fractus CTE.
2
 
3
+ Never advance with a length-1 chunk + carry. That path is the mono-token
4
+ attractor. Default decode is a sliding causal window: reset thought, tick
5
+ the last W tokens, take the last logit. Immediate self-copy is masked
6
+ (logits[prev] = -inf). No frequency ban, no temperature on the speech gate.
7
  """
8
  from __future__ import annotations
9
+
10
+ from collections import Counter
11
  from typing import List, Optional
12
 
13
+ import torch
14
+
15
+
16
+ def _reset(engine):
 
 
 
 
 
 
 
 
17
  engine.eval()
18
+ if hasattr(engine, "reset_thought"):
19
+ engine.reset_thought(1)
20
+ for blk in getattr(engine, "blocks", []):
21
+ if hasattr(blk, "attn_S"):
22
  blk.attn_S.zero_()
23
+ if hasattr(blk, "attn_z"):
24
  blk.attn_z.zero_()
25
 
 
 
 
26
 
27
+ def _pick(logits: torch.Tensor, prev: Optional[int], temperature: float, top_k: int) -> int:
28
+ l = logits.float().reshape(-1).clone()
29
+ if prev is not None and 0 <= prev < l.numel():
30
+ l[prev] = -1e9
31
+ if temperature <= 1e-5:
32
+ return int(l.argmax().item())
33
+ l = l / max(temperature, 1e-5)
34
+ k = min(top_k, l.numel())
35
+ topv, topi = torch.topk(l, k)
36
+ return int(topi[torch.multinomial(torch.softmax(topv, -1), 1)].item())
 
 
 
 
 
 
 
 
37
 
38
 
39
  @torch.no_grad()
 
42
  tokenizer,
43
  prompt: str,
44
  max_new: int = 40,
45
+ temperature: float = 0.0,
46
  top_k: int = 40,
47
  window: int = 64,
48
+ ban_window: int = 0,
49
+ ban_factor: float = 1.0,
50
  ) -> tuple[str, List[int]]:
51
+ """Causal window decode. ban_* kept for signature compat; ignored if ban_window==0."""
52
+ ids = tokenizer.encode(prompt)[:window] or [0]
 
53
  out: List[int] = []
54
+ prev = ids[-1]
55
  for _ in range(max_new):
56
+ _reset(engine)
 
 
 
 
 
57
  ctx = ids[-window:]
58
  logits = engine.tick_chunk(torch.tensor([ctx], dtype=torch.long))
59
+ l = logits[0, -1].float().clone()
60
+ if ban_window and out:
61
+ for p in set(out[-ban_window:]):
62
+ l[p] *= ban_factor
63
+ nxt = _pick(l, prev=prev, temperature=temperature, top_k=top_k)
64
  out.append(nxt)
65
  ids.append(nxt)
66
+ prev = nxt
67
  return tokenizer.decode(out), out
68
+
69
+
70
+ @torch.no_grad()
71
+ def generate_greedy_prefix(
72
+ engine,
73
+ tokenizer,
74
+ prompt: str,
75
+ max_new: int = 40,
76
+ window: int = 64,
77
+ ) -> tuple[str, List[int]]:
78
+ return generate_window(
79
+ engine, tokenizer, prompt,
80
+ max_new=max_new, temperature=0.0, top_k=1, window=window,
81
+ ban_window=0, ban_factor=1.0,
82
+ )
83
+
84
+
85
+ @torch.no_grad()
86
+ def generate_greedy_ids(engine, tokenizer, prompt: str, max_new: int = 40, window: int = 64):
87
+ return generate_greedy_prefix(engine, tokenizer, prompt, max_new=max_new, window=window)
88
+
89
+
90
+ @torch.no_grad()
91
+ def generate_chunk(
92
+ engine,
93
+ tokenizer,
94
+ prompt: str,
95
+ max_new: int = 40,
96
+ temperature: float = 0.0,
97
+ top_k: int = 40,
98
+ ban_window: int = 0,
99
+ ban_factor: float = 1.0,
100
+ context_limit: int = 128,
101
+ ) -> tuple[str, List[int]]:
102
+ """Back-compat name. Same as windowed greedy — length-1 carry is banned."""
103
+ return generate_window(
104
+ engine, tokenizer, prompt,
105
+ max_new=max_new, temperature=temperature, top_k=top_k,
106
+ window=min(64, context_limit),
107
+ ban_window=ban_window, ban_factor=ban_factor,
108
+ )
109
+
110
+
111
+ @torch.no_grad()
112
+ def unique40_probe(engine, max_new=40, mode="prefix", prompts=None):
113
+ class _Tok:
114
+ def encode(self, s):
115
+ # probe without a real tokenizer: unused if we pass raw later
116
+ return [1]
117
+
118
+ def decode(self, ids):
119
+ return " ".join(str(i) for i in ids)
120
+
121
+ # Real probe is done by callers that have a tokenizer. Keep symbol.
122
+ fn = generate_greedy_prefix if mode == "prefix" else generate_greedy_ids
123
+ return {"mode": mode, "fn": fn.__name__, "max_new": max_new}