thefinalboss commited on
Commit
16d4393
·
verified ·
1 Parent(s): 6a269fe

fix generate_window kwargs force_second carry 2026-08-29b

Browse files
Files changed (1) hide show
  1. fractus/generate_aligned.py +42 -74
fractus/generate_aligned.py CHANGED
@@ -1,18 +1,8 @@
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"):
@@ -23,7 +13,6 @@ def _reset(engine):
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():
@@ -31,11 +20,10 @@ def _pick(logits: torch.Tensor, prev: Optional[int], temperature: float, top_k:
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()
40
  def generate_window(
41
  engine,
@@ -44,80 +32,60 @@ def generate_window(
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}
 
1
+ """Train-aligned generation for Fractus CTE. No frequency ban / scramble."""
 
 
 
 
 
 
2
  from __future__ import annotations
 
 
3
  from typing import List, Optional
 
4
  import torch
5
 
 
6
  def _reset(engine):
7
  engine.eval()
8
  if hasattr(engine, "reset_thought"):
 
13
  if hasattr(blk, "attn_z"):
14
  blk.attn_z.zero_()
15
 
 
16
  def _pick(logits: torch.Tensor, prev: Optional[int], temperature: float, top_k: int) -> int:
17
  l = logits.float().reshape(-1).clone()
18
  if prev is not None and 0 <= prev < l.numel():
 
20
  if temperature <= 1e-5:
21
  return int(l.argmax().item())
22
  l = l / max(temperature, 1e-5)
23
+ k = min(max(1, top_k), l.numel())
24
  topv, topi = torch.topk(l, k)
25
  return int(topi[torch.multinomial(torch.softmax(topv, -1), 1)].item())
26
 
 
27
  @torch.no_grad()
28
  def generate_window(
29
  engine,
 
32
  max_new: int = 40,
33
  temperature: float = 0.0,
34
  top_k: int = 40,
35
+ window: int = 128,
36
+ force_second: bool = False,
37
+ carry: bool = False,
38
  ) -> tuple[str, List[int]]:
 
39
  ids = tokenizer.encode(prompt)[:window] or [0]
40
  out: List[int] = []
41
  prev = ids[-1]
42
+ if carry:
43
  _reset(engine)
44
+ logits = engine.tick_chunk(torch.tensor([ids], dtype=torch.long))
45
+ cur = logits[0, -1]
46
+ for i in range(max_new):
47
+ if force_second and i == 0:
48
+ l = cur.float().clone()
49
+ l[int(l.argmax())] = -1e9
50
+ nxt = _pick(l, prev=prev, temperature=temperature, top_k=top_k)
51
+ else:
52
+ nxt = _pick(cur, prev=prev, temperature=temperature, top_k=top_k)
53
+ out.append(nxt)
54
+ ids.append(nxt)
55
+ logits = engine.tick_chunk(torch.tensor([[nxt]], dtype=torch.long))
56
+ cur = logits[0, -1]
57
+ prev = nxt
58
+ else:
59
+ for i in range(max_new):
60
+ _reset(engine)
61
+ ctx = ids[-window:]
62
+ logits = engine.tick_chunk(torch.tensor([ctx], dtype=torch.long))
63
+ cur = logits[0, -1]
64
+ if force_second and i == 0:
65
+ l = cur.float().clone()
66
+ l[int(l.argmax())] = -1e9
67
+ nxt = _pick(l, prev=prev, temperature=temperature, top_k=top_k)
68
+ else:
69
+ nxt = _pick(cur, prev=prev, temperature=temperature, top_k=top_k)
70
+ out.append(nxt)
71
+ ids.append(nxt)
72
+ prev = nxt
73
  return tokenizer.decode(out), out
74
 
 
75
  @torch.no_grad()
76
+ def generate_greedy_prefix(engine, tokenizer, prompt: str, max_new: int = 40, window: int = 128):
77
+ return generate_window(engine, tokenizer, prompt, max_new=max_new, temperature=0.0, top_k=1, window=window, carry=False)
 
 
 
 
 
 
 
 
 
 
 
78
 
79
  @torch.no_grad()
80
+ def generate_greedy_ids(engine, tokenizer, prompt: str, max_new: int = 40, window: int = 128):
81
  return generate_greedy_prefix(engine, tokenizer, prompt, max_new=max_new, window=window)
82
 
 
83
  @torch.no_grad()
84
+ def generate_chunk(engine, tokenizer, prompt: str, max_new: int = 40, **kw):
85
+ return generate_window(engine, tokenizer, prompt, max_new=max_new, window=128, **kw)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
86
 
87
  @torch.no_grad()
88
  def unique40_probe(engine, max_new=40, mode="prefix", prompts=None):
89
+ return {"mode": mode, "max_new": max_new}
 
 
 
 
 
 
90
 
91
+ # space-sync 2026-08-29b force_second carry