thefinalboss commited on
Commit
86ecf8e
·
verified ·
1 Parent(s): d2c7e8f

Upload chat_infer.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. chat_infer.py +73 -0
chat_infer.py ADDED
@@ -0,0 +1,73 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Fast CogNet inference for chat API. Called by Next.js API route."""
3
+ import sys, os, json, time
4
+ sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
5
+
6
+ import torch
7
+ from cognet_1b import CogNet1B
8
+ from infer import CharTokenizer
9
+
10
+ # Globals - load once, reuse across calls
11
+ _model = None
12
+ _tokenizer = None
13
+
14
+ def load_model():
15
+ global _model, _tokenizer
16
+ if _model is not None:
17
+ return _model, _tokenizer
18
+
19
+ ckpt_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'checkpoints')
20
+ _tokenizer = CharTokenizer.load(os.path.join(ckpt_dir, 'tokenizer_v3.json'))
21
+ _model = CogNet1B(
22
+ vocab_size=_tokenizer.vocab_size, hidden_dim=512, num_blocks=6,
23
+ num_channels=6, channel_dim=128, ff_dim=1024, routing_iters=1,
24
+ max_adaptive_steps=2, max_seq_len=192, working_slots=32,
25
+ episodic_slots=64, semantic_slots=128, key_dim=256, dropout=0.1
26
+ )
27
+ ckpt_path = os.path.join(ckpt_dir, 'cognet_best.pt')
28
+ if os.path.exists(ckpt_path):
29
+ ckpt = torch.load(ckpt_path, map_location='cpu', weights_only=False)
30
+ state = ckpt['model_state_dict']
31
+ # Handle FP16 weights
32
+ fp16_state = {k: v.float() if v.dtype == torch.float16 else v for k, v in state.items()}
33
+ _model.load_state_dict(fp16_state)
34
+ step = ckpt.get('metrics', {}).get('step', '?')
35
+ sys.stderr.write(f'Model loaded (step={step})\n')
36
+ else:
37
+ sys.stderr.write('WARNING: No checkpoint found, using random weights\n')
38
+ _model.eval()
39
+ return _model, _tokenizer
40
+
41
+ def main():
42
+ prompt = sys.argv[1] if len(sys.argv) > 1 else "Hello"
43
+ max_tokens = int(sys.argv[2]) if len(sys.argv) > 2 else 80
44
+ temperature = float(sys.argv[3]) if len(sys.argv) > 3 else 0.7
45
+ top_k = int(sys.argv[4]) if len(sys.argv) > 4 else 30
46
+
47
+ model, tokenizer = load_model()
48
+
49
+ ids = tokenizer.encode(prompt)
50
+ if not ids:
51
+ ids = [0]
52
+ input_ids = torch.tensor([ids], dtype=torch.long)
53
+
54
+ t0 = time.time()
55
+ with torch.no_grad():
56
+ gen = model.generate(input_ids, max_new_tokens=max_tokens, temperature=temperature, top_k=top_k)
57
+ elapsed = time.time() - t0
58
+
59
+ generated_ids = gen[0].tolist()
60
+ generated_text = tokenizer.decode(generated_ids)
61
+ new_text = tokenizer.decode(generated_ids[len(ids):])
62
+
63
+ result = {
64
+ 'generated_text': generated_text,
65
+ 'new_text': new_text,
66
+ 'num_tokens': len(generated_ids),
67
+ 'inference_time_ms': round(elapsed * 1000),
68
+ 'prompt': prompt,
69
+ }
70
+ print(json.dumps(result, ensure_ascii=False))
71
+
72
+ if __name__ == '__main__':
73
+ main()