thefinalboss commited on
Commit
4715d0c
·
verified ·
1 Parent(s): f556dc6

Upload fractus/generate_aligned.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. fractus/generate_aligned.py +85 -0
fractus/generate_aligned.py ADDED
@@ -0,0 +1,85 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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()
54
+ def generate_window(
55
+ engine,
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