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

Upload infer.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. infer.py +385 -0
infer.py ADDED
@@ -0,0 +1,385 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ CogNet Inference Engine for Next.js API
3
+ =======================================
4
+ Loads trained CogNet model and CharTokenizer, supports:
5
+ - generate: text generation with temperature/top-k sampling
6
+ - analyze: logits analysis, entropy, top predictions
7
+ - inspect: model architecture details
8
+ - info: model info without loading weights
9
+ """
10
+
11
+ import json
12
+ import math
13
+ import os
14
+ import sys
15
+ from typing import Any, Dict, List, Optional
16
+
17
+ import torch
18
+ import torch.nn.functional as F
19
+
20
+ # Import from same directory
21
+ sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
22
+ from cognet_1b import CogNet1B
23
+
24
+
25
+ # ─── Model Config (matches training) ────────────────────────────────────────
26
+
27
+ MODEL_CONFIG = {
28
+ 'vocab_size': 136,
29
+ 'hidden_dim': 512,
30
+ 'num_blocks': 6,
31
+ 'num_channels': 6,
32
+ 'channel_dim': 128,
33
+ 'ff_dim': 1024,
34
+ 'routing_iters': 1,
35
+ 'max_adaptive_steps': 2,
36
+ 'max_seq_len': 192,
37
+ 'working_slots': 32,
38
+ 'episodic_slots': 64,
39
+ 'semantic_slots': 128,
40
+ 'key_dim': 256,
41
+ 'dropout': 0.1,
42
+ }
43
+
44
+ CKPT_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'checkpoints')
45
+ TOKENIZER_PATH = os.path.join(CKPT_DIR, 'tokenizer_v3.json')
46
+ BEST_MODEL_PATH = os.path.join(CKPT_DIR, 'cognet_best.pt')
47
+ LATEST_MODEL_PATH = os.path.join(CKPT_DIR, 'cognet_latest.pt')
48
+
49
+
50
+ # ─── CharTokenizer (standalone, no import needed from train_pipeline) ───────
51
+
52
+ class CharTokenizer:
53
+ """Character-level tokenizer: printable ASCII + French accents + newline/tab."""
54
+
55
+ def __init__(self):
56
+ self.chars = sorted(set(
57
+ [chr(i) for i in range(32, 127)]
58
+ + list('àâäéèêëïîôùûüÿçœæÀÂÄÉÈÊËÏÎÔÙÛÜŸÇŒÆ')
59
+ + list('ëßñ¿«»')
60
+ + ['\t', '\n']
61
+ ))
62
+ self.char_to_id = {c: i for i, c in enumerate(self.chars)}
63
+ self.id_to_char = {i: c for i, c in enumerate(self.chars)}
64
+ self.vocab_size = len(self.chars)
65
+
66
+ def encode(self, text: str) -> List[int]:
67
+ return [self.char_to_id.get(c, self.char_to_id.get(' ', 0)) for c in text]
68
+
69
+ def decode(self, ids: List[int]) -> str:
70
+ return ''.join(self.id_to_char.get(i, ' ') for i in ids)
71
+
72
+ def save(self, path: str):
73
+ with open(path, 'w', encoding='utf-8') as f:
74
+ json.dump({
75
+ 'chars': self.chars,
76
+ 'vocab_size': self.vocab_size,
77
+ }, f, ensure_ascii=False, indent=2)
78
+
79
+ @classmethod
80
+ def load(cls, path: str) -> 'CharTokenizer':
81
+ tok = cls.__new__(cls)
82
+ with open(path, 'r', encoding='utf-8') as f:
83
+ data = json.load(f)
84
+ tok.chars = data['chars']
85
+ tok.char_to_id = {c: i for i, c in enumerate(tok.chars)}
86
+ tok.id_to_char = {i: c for i, c in enumerate(tok.chars)}
87
+ tok.vocab_size = data['vocab_size']
88
+ return tok
89
+
90
+
91
+ # ─── JSON Helpers ────────────────────────────────────────────────────────────
92
+
93
+ def sanitize_for_json(obj: Any) -> Any:
94
+ """Replace NaN/Inf with None for JSON serialization."""
95
+ if isinstance(obj, float):
96
+ if math.isnan(obj) or math.isinf(obj):
97
+ return None
98
+ return obj
99
+ if isinstance(obj, dict):
100
+ return {k: sanitize_for_json(v) for k, v in obj.items()}
101
+ if isinstance(obj, list):
102
+ return [sanitize_for_json(v) for v in obj]
103
+ return obj
104
+
105
+
106
+ # ─── Model Cache ─────────────────────────────────────────────────────────────
107
+
108
+ _model_cache: Dict[str, Any] = {
109
+ 'model': None,
110
+ 'tokenizer': None,
111
+ 'device': None,
112
+ 'loaded': False,
113
+ }
114
+
115
+
116
+ def load_model_and_tokenizer() -> tuple:
117
+ """Load model and tokenizer with caching."""
118
+ if _model_cache['loaded']:
119
+ return _model_cache['model'], _model_cache['tokenizer'], _model_cache['device']
120
+
121
+ device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
122
+
123
+ # Load tokenizer
124
+ if not os.path.exists(TOKENIZER_PATH):
125
+ raise FileNotFoundError(
126
+ f"Tokenizer not found at {TOKENIZER_PATH}. "
127
+ "Run train_pipeline.py first to create it."
128
+ )
129
+ tokenizer = CharTokenizer.load(TOKENIZER_PATH)
130
+
131
+ # Update vocab_size from tokenizer
132
+ config = dict(MODEL_CONFIG)
133
+ config['vocab_size'] = tokenizer.vocab_size
134
+
135
+ # Create model
136
+ model = CogNet1B(**config).to(device)
137
+
138
+ # Load weights (prefer best, then latest)
139
+ model_path = BEST_MODEL_PATH if os.path.exists(BEST_MODEL_PATH) else LATEST_MODEL_PATH
140
+ if model_path and os.path.exists(model_path):
141
+ ckpt = torch.load(model_path, map_location=device, weights_only=False)
142
+ model.load_state_dict(ckpt['model_state_dict'])
143
+ step = ckpt.get('metrics', {}).get('step', '?')
144
+ print(f"Loaded model from {model_path} (step={step})")
145
+ else:
146
+ print("WARNING: No trained weights found. Using random initialization.")
147
+
148
+ model.eval()
149
+
150
+ # Cache
151
+ _model_cache['model'] = model
152
+ _model_cache['tokenizer'] = tokenizer
153
+ _model_cache['device'] = device
154
+ _model_cache['loaded'] = True
155
+
156
+ return model, tokenizer, device
157
+
158
+
159
+ # ─── Action Handlers ─────────────────────────────────────────────────────────
160
+
161
+ def handle_generate(prompt: str, max_tokens: int = 100,
162
+ temperature: float = 0.8, top_k: int = 20) -> Dict:
163
+ """Generate text from a prompt."""
164
+ model, tokenizer, device = load_model_and_tokenizer()
165
+
166
+ # Encode prompt
167
+ ids = tokenizer.encode(prompt)
168
+ if len(ids) == 0:
169
+ ids = [0]
170
+
171
+ input_ids = torch.tensor([ids], dtype=torch.long, device=device)
172
+
173
+ # Generate
174
+ with torch.no_grad():
175
+ output_ids = model.generate(
176
+ input_ids,
177
+ max_new_tokens=max_tokens,
178
+ temperature=temperature,
179
+ top_k=top_k,
180
+ )
181
+
182
+ # Decode
183
+ generated_ids = output_ids[0].tolist()
184
+ generated_text = tokenizer.decode(generated_ids)
185
+ new_text = tokenizer.decode(generated_ids[len(ids):])
186
+
187
+ # Token details
188
+ token_details = []
189
+ for i, tid in enumerate(generated_ids):
190
+ char = tokenizer.decode([tid])
191
+ token_details.append({
192
+ 'id': tid,
193
+ 'char': char,
194
+ 'position': i,
195
+ })
196
+
197
+ return sanitize_for_json({
198
+ 'action': 'generate',
199
+ 'prompt': prompt,
200
+ 'generated_text': generated_text,
201
+ 'new_text': new_text,
202
+ 'token_details': token_details,
203
+ 'num_tokens': len(generated_ids),
204
+ 'temperature': temperature,
205
+ 'top_k': top_k,
206
+ })
207
+
208
+
209
+ def handle_analyze(prompt: str) -> Dict:
210
+ """Analyze logits, entropy, and top predictions."""
211
+ model, tokenizer, device = load_model_and_tokenizer()
212
+
213
+ ids = tokenizer.encode(prompt)
214
+ if len(ids) == 0:
215
+ ids = [0]
216
+
217
+ input_ids = torch.tensor([ids], dtype=torch.long, device=device)
218
+
219
+ with torch.no_grad():
220
+ result = model(input_ids, return_stats=True)
221
+ logits = result['logits']
222
+
223
+ # Analyze last token's predictions
224
+ last_logits = logits[0, -1, :] # (vocab_size,)
225
+ probs = F.softmax(last_logits, dim=-1)
226
+
227
+ # Entropy
228
+ entropy = -(probs * (probs + 1e-10).log()).sum().item()
229
+
230
+ # Top 10 predictions
231
+ topk_vals, topk_ids = torch.topk(probs, min(10, probs.size(0)))
232
+ top_predictions = []
233
+ for prob, tid in zip(topk_vals.tolist(), topk_ids.tolist()):
234
+ top_predictions.append({
235
+ 'token_id': tid,
236
+ 'char': tokenizer.decode([tid]),
237
+ 'probability': prob,
238
+ })
239
+
240
+ # Per-position entropy
241
+ all_probs = F.softmax(logits[0], dim=-1)
242
+ pos_entropy = (-(all_probs * (all_probs + 1e-10).log()).sum(dim=-1)).tolist()
243
+
244
+ # Stats
245
+ stats = result.get('stats', {})
246
+ stats_summary = {}
247
+ for k, v in stats.items():
248
+ if isinstance(v, torch.Tensor):
249
+ v = v.item()
250
+ if isinstance(v, float) and (math.isnan(v) or math.isinf(v)):
251
+ v = None
252
+ stats_summary[k] = v
253
+
254
+ return sanitize_for_json({
255
+ 'action': 'analyze',
256
+ 'prompt': prompt,
257
+ 'prompt_length': len(ids),
258
+ 'entropy': entropy,
259
+ 'top_predictions': top_predictions,
260
+ 'per_position_entropy': pos_entropy,
261
+ 'model_stats': stats_summary,
262
+ })
263
+
264
+
265
+ def handle_inspect() -> Dict:
266
+ """Return model architecture details."""
267
+ model, tokenizer, device = load_model_and_tokenizer()
268
+
269
+ params = model.count_parameters()
270
+ complexity = model.get_complexity_analysis()
271
+
272
+ # Layer details
273
+ layers = []
274
+ for i, block in enumerate(model.blocks):
275
+ layer_params = sum(p.numel() for p in block.parameters())
276
+ layers.append({
277
+ 'block_index': i,
278
+ 'parameters': layer_params,
279
+ 'components': ['CognitiveRouter', 'SharedHierarchicalMemory',
280
+ 'AdaptiveComputationBlock', 'CompositionalReasoner'],
281
+ })
282
+
283
+ return sanitize_for_json({
284
+ 'action': 'inspect',
285
+ 'architecture': 'CogNet (Non-Transformer)',
286
+ 'total_parameters': params['total'],
287
+ 'trainable_parameters': params['trainable'],
288
+ 'config': {
289
+ 'vocab_size': model.vocab_size,
290
+ 'hidden_dim': model.hidden_dim,
291
+ 'num_blocks': model.num_blocks,
292
+ 'num_channels': model.num_channels,
293
+ 'channel_dim': model.channel_dim,
294
+ 'ff_dim': model.ff_dim,
295
+ 'max_seq_len': model.max_seq_len,
296
+ 'tokenizer_vocab_size': tokenizer.vocab_size,
297
+ },
298
+ 'complexity_analysis': complexity,
299
+ 'layers': layers,
300
+ 'device': str(device),
301
+ })
302
+
303
+
304
+ def handle_info() -> Dict:
305
+ """Return model info without loading weights."""
306
+ config = dict(MODEL_CONFIG)
307
+
308
+ # Check what's available
309
+ has_tokenizer = os.path.exists(TOKENIZER_PATH)
310
+ has_best = os.path.exists(BEST_MODEL_PATH)
311
+ has_latest = os.path.exists(LATEST_MODEL_PATH)
312
+
313
+ # Estimate param count without loading
314
+ model = CogNet1B(**config)
315
+ params = model.count_parameters()
316
+
317
+ # Check checkpoint info if available
318
+ checkpoint_info = {}
319
+ if has_best:
320
+ try:
321
+ ckpt = torch.load(BEST_MODEL_PATH, map_location='cpu', weights_only=False)
322
+ checkpoint_info['best'] = {
323
+ 'step': ckpt.get('metrics', {}).get('step', None),
324
+ 'val_loss': ckpt.get('metrics', {}).get('val_loss', None),
325
+ 'val_ppl': ckpt.get('metrics', {}).get('val_ppl', None),
326
+ }
327
+ except Exception:
328
+ checkpoint_info['best'] = {'error': 'Could not read checkpoint'}
329
+ if has_latest:
330
+ try:
331
+ ckpt = torch.load(LATEST_MODEL_PATH, map_location='cpu', weights_only=False)
332
+ checkpoint_info['latest'] = {
333
+ 'step': ckpt.get('metrics', {}).get('step', None),
334
+ }
335
+ except Exception:
336
+ checkpoint_info['latest'] = {'error': 'Could not read checkpoint'}
337
+
338
+ return sanitize_for_json({
339
+ 'action': 'info',
340
+ 'model_name': 'CogNet',
341
+ 'architecture': 'Non-Transformer (Cognitive Routing)',
342
+ 'estimated_parameters': params['total'],
343
+ 'config': config,
344
+ 'files': {
345
+ 'tokenizer': has_tokenizer,
346
+ 'best_checkpoint': has_best,
347
+ 'latest_checkpoint': has_latest,
348
+ },
349
+ 'checkpoint_info': checkpoint_info,
350
+ })
351
+
352
+
353
+ # ─── CLI Entry Point ─────────────────────────────────────────────────────────
354
+
355
+ def main():
356
+ import argparse
357
+ parser = argparse.ArgumentParser(description='CogNet Inference Engine')
358
+ parser.add_argument('action', choices=['generate', 'analyze', 'inspect', 'info'],
359
+ help='Action to perform')
360
+ parser.add_argument('--prompt', type=str, default='The ',
361
+ help='Prompt text (for generate/analyze)')
362
+ parser.add_argument('--max-tokens', type=int, default=100,
363
+ help='Max tokens to generate')
364
+ parser.add_argument('--temperature', type=float, default=0.8,
365
+ help='Sampling temperature')
366
+ parser.add_argument('--top-k', type=int, default=20,
367
+ help='Top-k sampling')
368
+
369
+ args = parser.parse_args()
370
+
371
+ if args.action == 'generate':
372
+ result = handle_generate(args.prompt, args.max_tokens,
373
+ args.temperature, args.top_k)
374
+ elif args.action == 'analyze':
375
+ result = handle_analyze(args.prompt)
376
+ elif args.action == 'inspect':
377
+ result = handle_inspect()
378
+ elif args.action == 'info':
379
+ result = handle_info()
380
+
381
+ print(json.dumps(result, indent=2, ensure_ascii=False))
382
+
383
+
384
+ if __name__ == '__main__':
385
+ main()