Flilax commited on
Commit
430b17b
·
verified ·
1 Parent(s): c6fe063

Upload inference.py

Browse files
Files changed (1) hide show
  1. inference.py +63 -0
inference.py ADDED
@@ -0,0 +1,63 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn.functional as F
3
+ from hrt import ModelConfig, HierarchicalRadialTransformerV7
4
+
5
+ device = "cuda" if torch.cuda.is_available() else "cpu"
6
+
7
+ # 1. Initialize Configuration matching training
8
+ cfg = ModelConfig(
9
+ d_model=768,
10
+ d_ff=3072,
11
+ n_outer_latents=512,
12
+ n_outer_cycles=6,
13
+ n_inner_cycles=8,
14
+ n_center_latents=16,
15
+ routing_k=64,
16
+ n_outer_heads=12,
17
+ n_inner_heads=12,
18
+ n_latent_heads=12,
19
+ vocab_size=257,
20
+ max_seq_len=131072,
21
+ use_qk_norm=True,
22
+ use_rezero=True,
23
+ use_compaction=True,
24
+ use_internalization=True,
25
+ use_jfb=True,
26
+ use_q_cache=True,
27
+ )
28
+
29
+ # 2. Load model & weights
30
+ model = HierarchicalRadialTransformerV7(cfg).to(device)
31
+ weights = torch.load("hrt_v7_148m_weights.pt", map_location=device)
32
+ model.load_state_dict(weights["model"] if "model" in weights else weights)
33
+ model.eval()
34
+
35
+ # 3. Autoregressive Byte-level Generation
36
+ def generate(prompt: str, max_new_bytes: int = 120, temp: float = 0.5, top_k: int = 5):
37
+ prompt_bytes = list(prompt.encode("utf-8"))
38
+ prompt_ids = torch.tensor([prompt_bytes], dtype=torch.long, device=device)
39
+
40
+ with torch.no_grad():
41
+ prompt_emb = model.tok_emb(prompt_ids)
42
+ logits, cache = model._init_generation_cache(prompt_emb)
43
+ out_bytes = list(prompt_bytes)
44
+
45
+ for _ in range(max_new_bytes):
46
+ l = logits / max(temp, 1e-5)
47
+ if top_k > 0:
48
+ v, _ = torch.topk(l, min(top_k, l.size(-1)))
49
+ l[l < v[:, [-1]]] = float("-inf")
50
+
51
+ nxt = torch.multinomial(F.softmax(l, dim=-1), num_samples=1)
52
+ nxt_id = nxt.item()
53
+ if nxt_id == 256: # EOS
54
+ break
55
+
56
+ out_bytes.append(nxt_id)
57
+ nxt_emb = model.tok_emb(nxt)
58
+ logits = model.step_generation(nxt_emb, cache)
59
+
60
+ return bytes(out_bytes).decode("utf-8", errors="replace")
61
+
62
+ # Test completion
63
+ print(generate("def", max_new_bytes=100))