VADRK155 commited on
Commit
15a0de0
ยท
verified ยท
1 Parent(s): 1e145bd

Upload folder using huggingface_hub

Browse files
Files changed (4) hide show
  1. chat.py +558 -0
  2. config.json +3 -0
  3. cortex-2-code.pt +3 -0
  4. requirements.txt +2 -0
chat.py ADDED
@@ -0,0 +1,558 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn.functional as F
3
+ import json
4
+ import sys
5
+ import math
6
+ import ast
7
+ import os
8
+ import time
9
+ import subprocess
10
+ import tempfile
11
+ from pathlib import Path
12
+
13
+
14
+ class CausalSelfAttention(torch.nn.Module):
15
+ def __init__(self, d_model, n_heads, dropout, context_length):
16
+ super().__init__()
17
+ self.n_heads = n_heads
18
+ self.head_dim = d_model // n_heads
19
+ self.qkv = torch.nn.Linear(d_model, 3 * d_model)
20
+ self.proj = torch.nn.Linear(d_model, d_model)
21
+ self.attn_dropout = torch.nn.Dropout(dropout)
22
+ self.resid_dropout = torch.nn.Dropout(dropout)
23
+ self.register_buffer("mask", torch.tril(torch.ones(context_length, context_length)).unsqueeze(0).unsqueeze(0))
24
+
25
+ def forward(self, x):
26
+ B, T, C = x.shape
27
+ qkv = self.qkv(x)
28
+ q, k, v = qkv.chunk(3, dim=-1)
29
+ q = q.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
30
+ k = k.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
31
+ v = v.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
32
+ attn = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(self.head_dim))
33
+ attn = attn.masked_fill(self.mask[:, :, :T, :T] == 0, float("-inf"))
34
+ attn = F.softmax(attn, dim=-1)
35
+ attn = self.attn_dropout(attn)
36
+ out = attn @ v
37
+ out = out.transpose(1, 2).contiguous().view(B, T, C)
38
+ out = self.proj(out)
39
+ out = self.resid_dropout(out)
40
+ return out
41
+
42
+
43
+ class MLP(torch.nn.Module):
44
+ def __init__(self, d_model, d_ff, dropout):
45
+ super().__init__()
46
+ self.net = torch.nn.Sequential(
47
+ torch.nn.Linear(d_model, d_ff),
48
+ torch.nn.GELU(),
49
+ torch.nn.Linear(d_ff, d_model),
50
+ torch.nn.Dropout(dropout),
51
+ )
52
+
53
+ def forward(self, x):
54
+ return self.net(x)
55
+
56
+
57
+ class TransformerBlock(torch.nn.Module):
58
+ def __init__(self, d_model, n_heads, d_ff, dropout, context_length):
59
+ super().__init__()
60
+ self.ln1 = torch.nn.LayerNorm(d_model)
61
+ self.attn = CausalSelfAttention(d_model, n_heads, dropout, context_length)
62
+ self.ln2 = torch.nn.LayerNorm(d_model)
63
+ self.mlp = MLP(d_model, d_ff, dropout)
64
+
65
+ def forward(self, x):
66
+ x = x + self.attn(self.ln1(x))
67
+ x = x + self.mlp(self.ln2(x))
68
+ return x
69
+
70
+
71
+ def format_instruction(instruction, extra_input=""):
72
+ instruction = (instruction or "").strip()
73
+ extra_input = (extra_input or "").strip()
74
+ if extra_input and extra_input.lower() != "not applicable":
75
+ return f"### Instruction:\n{instruction}\n\n### Input:\n{extra_input}\n\n### Response:\n"
76
+ return f"### Instruction:\n{instruction}\n\n### Response:\n"
77
+
78
+
79
+ _FOREIGN_MARKERS = (
80
+ "#include", "void main", "int main(", "public static void",
81
+ "System.out.println", "console.log", "function ", "</",
82
+ "<?php", "using namespace", "fmt.Println", "package main",
83
+ "fn main", "<html", "<script", "CREATE TABLE", "SELECT ", "=>",
84
+ )
85
+
86
+
87
+ def extract_code(text):
88
+ text = (text or "").strip()
89
+ if "```" not in text:
90
+ return text.strip()
91
+
92
+ def _drop_lang_label(block):
93
+ lines = block.split("\n")
94
+ if lines and lines[0].strip() and len(lines[0].strip()) <= 12 \
95
+ and not any(ch in lines[0] for ch in " \t=()[]{}:;"):
96
+ lines = lines[1:]
97
+ return "\n".join(lines).strip("\n")
98
+
99
+ parts = text.split("```")
100
+ blocks = []
101
+ for i in range(1, len(parts), 2):
102
+ blocks.append(_drop_lang_label(parts[i]))
103
+ if blocks:
104
+ return "\n\n".join(b.strip("\n") for b in blocks).strip()
105
+
106
+ return _drop_lang_label(parts[1]).strip()
107
+
108
+
109
+ def looks_like_python(code):
110
+ head = (code or "")[:3000].lower()
111
+ return not any(m.lower() in head for m in _FOREIGN_MARKERS)
112
+
113
+
114
+ def check_syntax(code):
115
+ if not (code or "").strip():
116
+ return False, "model returned no code (empty response)"
117
+
118
+ try:
119
+ ast.parse(code)
120
+ return True, None
121
+ except (SyntaxError, ValueError) as e:
122
+ if not looks_like_python(code):
123
+ return False, ("this doesn't look like Python code โ€” syntax checking and "
124
+ "execution are only supported for Python")
125
+ if isinstance(e, ValueError):
126
+ return False, f"failed to parse code: {e}"
127
+
128
+ lines = (code or "").splitlines()
129
+ lineno = e.lineno or 1
130
+ offset = e.offset or 1
131
+ out = [f"SyntaxError: {e.msg} (line {lineno}, column {offset})"]
132
+ if 1 <= lineno <= len(lines):
133
+ bad_line = lines[lineno - 1]
134
+ caret_pos = min(max(offset, 1), len(bad_line) + 1) - 1
135
+ out.append(f" {lineno:>4} | {bad_line}")
136
+ out.append(f" | {' ' * caret_pos}^")
137
+ if lineno >= len(lines):
138
+ out.append(" (looks like the code was cut off by the generation limit โ€” "
139
+ "try increasing code_max_new_tokens)")
140
+ return False, "\n".join(out)
141
+
142
+
143
+ def run_python_code(code, timeout=10.0):
144
+ fd, path = tempfile.mkstemp(suffix=".py", prefix="cortex_run_")
145
+ try:
146
+ with os.fdopen(fd, "w", encoding="utf-8") as f:
147
+ f.write(code)
148
+ env = {**os.environ, "PYTHONIOENCODING": "utf-8"}
149
+ proc = subprocess.run(
150
+ [sys.executable, "-u", path],
151
+ stdin=subprocess.DEVNULL,
152
+ capture_output=True,
153
+ text=True,
154
+ encoding="utf-8",
155
+ errors="replace",
156
+ timeout=timeout,
157
+ env=env,
158
+ )
159
+ return proc.returncode, (proc.stdout or "") + (proc.stderr or ""), False
160
+ except subprocess.TimeoutExpired as e:
161
+ partial = ""
162
+ for stream in (e.stdout, e.stderr):
163
+ if not stream:
164
+ continue
165
+ if isinstance(stream, bytes):
166
+ stream = stream.decode("utf-8", "replace")
167
+ partial += stream
168
+ return -1, partial, True
169
+ finally:
170
+ try:
171
+ os.unlink(path)
172
+ except OSError:
173
+ pass
174
+
175
+
176
+ class TinyGPT(torch.nn.Module):
177
+ def __init__(self, config):
178
+ super().__init__()
179
+ self.config = config
180
+ vocab_size = config["tokenizer_vocab_size"] + 10
181
+ self.token_emb = torch.nn.Embedding(vocab_size, config["d_model"])
182
+ self.pos_emb = torch.nn.Embedding(config["context_length"], config["d_model"])
183
+ self.drop = torch.nn.Dropout(config["dropout"])
184
+ self.blocks = torch.nn.ModuleList([
185
+ TransformerBlock(config["d_model"], config["n_heads"], config["d_ff"], config["dropout"], config["context_length"])
186
+ for _ in range(config["n_layers"])
187
+ ])
188
+ self.ln_f = torch.nn.LayerNorm(config["d_model"])
189
+ self.head = torch.nn.Linear(config["d_model"], vocab_size, bias=False)
190
+ self.token_emb.weight = self.head.weight
191
+
192
+ def forward(self, idx, targets=None):
193
+ B, T = idx.shape
194
+ pos = torch.arange(0, T, device=idx.device).unsqueeze(0)
195
+ x = self.token_emb(idx) + self.pos_emb(pos)
196
+ x = self.drop(x)
197
+ for block in self.blocks:
198
+ x = block(x)
199
+ x = self.ln_f(x)
200
+ logits = self.head(x)
201
+ loss = None
202
+ if targets is not None:
203
+ loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1), ignore_index=0)
204
+ return logits, loss
205
+
206
+
207
+ def find_model_file():
208
+ here = Path(".")
209
+
210
+ pt_files = list(here.glob("*.pt"))
211
+
212
+ for name in ["best_model.pt", "final_model.pt"]:
213
+ if name in [f.name for f in pt_files]:
214
+ return here / name
215
+
216
+ if pt_files:
217
+ return pt_files[0]
218
+
219
+ return None
220
+
221
+
222
+ def main():
223
+ device = torch.device("cuda")
224
+
225
+ model_path = find_model_file()
226
+ if model_path is None:
227
+ print("โŒ No .pt model file found! Put this script in the same folder as your model.")
228
+ sys.exit(1)
229
+
230
+ if len(sys.argv) > 1:
231
+ model_path = Path(sys.argv[1])
232
+
233
+ print(f"๐Ÿ“‚ Loading model from: {model_path.name}")
234
+
235
+ ckpt = torch.load(model_path, map_location=device, weights_only=False)
236
+
237
+ if "config" in ckpt and "tokenizer" in ckpt:
238
+ config = ckpt["config"]
239
+ from tokenizers import Tokenizer
240
+ tokenizer = Tokenizer.from_str(ckpt["tokenizer"])
241
+ print("๐Ÿ“ฆ Loaded config + tokenizer from checkpoint")
242
+ else:
243
+ here = model_path.parent
244
+ config_path = here / "config.json"
245
+ tokenizer_path = here / "tokenizer.json"
246
+
247
+ if not config_path.exists():
248
+ print(f"โŒ config.json not found next to model!")
249
+ sys.exit(1)
250
+ if not tokenizer_path.exists():
251
+ print(f"โŒ tokenizer.json not found next to model!")
252
+ sys.exit(1)
253
+
254
+ with open(config_path) as f:
255
+ config = json.load(f)
256
+ from tokenizers import Tokenizer
257
+ tokenizer = Tokenizer.from_file(str(tokenizer_path))
258
+ print("๐Ÿ“ฆ Loaded config + tokenizer from separate files")
259
+
260
+ model = TinyGPT(config).to(device)
261
+ model.load_state_dict(ckpt["model"])
262
+ model.eval()
263
+
264
+ n_params = sum(p.numel() for p in model.parameters())
265
+ step = ckpt.get("step", "?")
266
+ val_loss = ckpt.get("val_loss", "?")
267
+ if isinstance(val_loss, float):
268
+ val_loss = f"{val_loss:.4f}"
269
+
270
+ print(f"โœ… Cortex_2 loaded!")
271
+ print(f" Parameters: {n_params / 1e6:.1f}M")
272
+ print(f" Step: {step}")
273
+ print(f" Val loss: {val_loss}")
274
+ print(f" Device: {device}")
275
+
276
+ dataset_mode = config.get("dataset_mode", "stories")
277
+ is_chat_model = dataset_mode == "chat"
278
+ is_code_model = dataset_mode == "code"
279
+
280
+ if is_chat_model:
281
+ print(f" Mode: ๐Ÿ’ฌ conversational (dataset_mode=chat)")
282
+ elif is_code_model:
283
+ print(f" Mode: ๐Ÿง‘โ€๐Ÿ’ป code (dataset_mode=code)")
284
+ else:
285
+ print(f" Mode: ๐Ÿ“– story completion (dataset_mode=stories)")
286
+
287
+ print()
288
+ print("๐Ÿ’ฌ Type a prompt and press Enter. Type 'quit' to exit.")
289
+ if is_chat_model:
290
+ print(" (type 'reset' to clear conversation history)")
291
+ print(" (type 'temp 0.9' to change temperature, current default: 0.8)")
292
+ if is_code_model:
293
+ print(" Describe a task, e.g.: 'Write a function that reverses a string'.")
294
+ print(" (to set a separate 'Input:', type: task || input)")
295
+ print(" (type 'temp 0.5' to change temperature, current default: 0.5)")
296
+ print()
297
+ print(" Code mode commands:")
298
+ print(" run โ€” run the last generated code")
299
+ print(" save โ€” save the last code to generated_code_NN.py")
300
+ print(" autocheck โ€” auto-regenerate on syntax error")
301
+ print(" timeout N โ€” code execution timeout in seconds")
302
+ print(" After generation the code is syntax-checked, and clean code can be")
303
+ print(" run directly from the chat (y when asked 'Run?').")
304
+ print("=" * 50)
305
+
306
+ bos_id = tokenizer.token_to_id("<bos>")
307
+ eos_id = tokenizer.token_to_id("<eos>")
308
+ context_length = config["context_length"]
309
+
310
+ history_lines = []
311
+ temperature = 0.5 if is_code_model else 0.8
312
+ code_max_new_tokens = 400
313
+ code_top_k = 40
314
+
315
+ last_code = None
316
+ run_timeout = 10.0
317
+ autocheck = True
318
+ max_auto_attempts = 3
319
+
320
+ def generate_code(instruction, extra_input=""):
321
+ text_prompt = format_instruction(instruction, extra_input)
322
+ ids = tokenizer.encode(text_prompt).ids
323
+ idx = torch.tensor([[bos_id] + ids], dtype=torch.long, device=device)
324
+ prompt_len = idx.shape[1]
325
+
326
+ t0 = time.time()
327
+ n_tokens = 0
328
+ with torch.no_grad():
329
+ for _ in range(code_max_new_tokens):
330
+ idx_cond = idx[:, -context_length:]
331
+ logits, _ = model(idx_cond)
332
+ logits = logits[:, -1, :] / temperature
333
+ if code_top_k:
334
+ kth = torch.topk(logits, code_top_k).values[:, -1, None]
335
+ logits = logits.masked_fill(logits < kth, float("-inf"))
336
+ probs = F.softmax(logits, dim=-1)
337
+ next_id = torch.multinomial(probs, num_samples=1)
338
+ idx = torch.cat([idx, next_id], dim=1)
339
+ n_tokens += 1
340
+ if next_id.item() == eos_id:
341
+ break
342
+ print(f" โณ generated {n_tokens} tokens in {time.time() - t0:.1f}s")
343
+ return tokenizer.decode(idx[0, prompt_len:].tolist())
344
+
345
+ def execute_code(code):
346
+ print("โ”€" * 50)
347
+ print(f"โ–ถ Running code (separate process, timeout {run_timeout:.0f}s, stdin closed)...")
348
+ rc, output, timed_out = run_python_code(code, run_timeout)
349
+ if timed_out:
350
+ print(f"โฑ Timeout exceeded ({run_timeout:.0f}s) โ€” process stopped.")
351
+ if output.strip():
352
+ print("๐Ÿ“ค Output before stopping:")
353
+ print(output.rstrip())
354
+ print(" Hint: if the code waits for input(), it will never finish โ€”")
355
+ print(" interactive input is not available when running from chat.")
356
+ elif rc == 0:
357
+ if output.strip():
358
+ print("๐Ÿ“ค Program output:")
359
+ print(output.rstrip())
360
+ else:
361
+ print("๐Ÿ“ค Program finished with no output.")
362
+ print("โœ… Code ran without errors (exit code 0).")
363
+ else:
364
+ if output.strip():
365
+ print("๐Ÿ“ค Program output:")
366
+ print(output.rstrip())
367
+ if "EOFError" in output:
368
+ print(" Hint: the code called input() โ€” input is not available when running from chat.")
369
+ print(f"โŒ Program finished with an error (exit code {rc}).")
370
+ print("โ”€" * 50)
371
+
372
+ # Chat loop
373
+ while True:
374
+ try:
375
+ prompt = input("\nYou: ").strip()
376
+ except (EOFError, KeyboardInterrupt):
377
+ print("\n๐Ÿ‘‹ Bye!")
378
+ break
379
+
380
+ if prompt.lower() == "quit":
381
+ print("๐Ÿ‘‹ Bye!")
382
+ break
383
+ if is_chat_model and prompt.lower() == "reset":
384
+ history_lines = []
385
+ print("๐Ÿ”„ Conversation history cleared.")
386
+ continue
387
+ if (is_chat_model or is_code_model) and prompt.lower().startswith("temp"):
388
+ parts = prompt.split()
389
+ if len(parts) == 2:
390
+ try:
391
+ new_temp = float(parts[1])
392
+ if new_temp <= 0:
393
+ print("โš ๏ธ Temperature must be greater than 0.")
394
+ else:
395
+ temperature = new_temp
396
+ print(f"๐ŸŒก๏ธ Temperature set to: {temperature}")
397
+ except ValueError:
398
+ print("โš ๏ธ Could not parse the number. Example: temp 0.9")
399
+ else:
400
+ print(f"๐ŸŒก๏ธ Current temperature: {temperature} (example to change: temp 0.9)")
401
+ continue
402
+
403
+ if is_code_model and prompt.lower() in ("run", "r"):
404
+ if not last_code:
405
+ print("โš ๏ธ Nothing to run yet โ€” generate some code first.")
406
+ continue
407
+ ok, err = check_syntax(last_code)
408
+ if not ok:
409
+ print(f"โŒ The last code has a syntax error, cannot run it:\n{err}")
410
+ continue
411
+ execute_code(last_code)
412
+ continue
413
+
414
+ if is_code_model and prompt.lower() == "save":
415
+ if not last_code:
416
+ print("โš ๏ธ Nothing to save yet โ€” generate some code first.")
417
+ continue
418
+ n = 1
419
+ while (Path.cwd() / f"generated_code_{n:02d}.py").exists():
420
+ n += 1
421
+ save_path = Path.cwd() / f"generated_code_{n:02d}.py"
422
+ save_path.write_text(last_code, encoding="utf-8")
423
+ print(f"๐Ÿ’พ Code saved: {save_path}")
424
+ continue
425
+
426
+ if is_code_model and prompt.lower().startswith("autocheck"):
427
+ parts = prompt.split()
428
+ if len(parts) == 2 and parts[1].lower() in ("on", "off"):
429
+ autocheck = parts[1].lower() == "on"
430
+ state = "on" if autocheck else "off"
431
+ print(f"๐Ÿ”„ Auto-regenerate on error: {state} (max attempts: {max_auto_attempts})")
432
+ else:
433
+ state = "on" if autocheck else "off"
434
+ print(f"๐Ÿ”„ Auto-regenerate is currently: {state} (example: autocheck off)")
435
+ continue
436
+
437
+ if is_code_model and prompt.lower().startswith("timeout"):
438
+ parts = prompt.split()
439
+ if len(parts) == 2:
440
+ try:
441
+ val = float(parts[1])
442
+ if val <= 0:
443
+ print("โš ๏ธ Timeout must be greater than 0.")
444
+ else:
445
+ run_timeout = val
446
+ print(f"โฑ Code execution timeout: {run_timeout:.0f}s")
447
+ except ValueError:
448
+ print("โš ๏ธ Could not parse the number. Example: timeout 15")
449
+ else:
450
+ print(f"โฑ Current execution timeout: {run_timeout:.0f}s (example: timeout 15)")
451
+ continue
452
+
453
+ if not prompt:
454
+ continue
455
+
456
+ if is_chat_model:
457
+ history_lines.append(f"User: {prompt}")
458
+ history_lines.append("Bot:")
459
+ full_text = "\n".join(history_lines)
460
+
461
+ ids = tokenizer.encode(full_text).ids
462
+ idx = torch.tensor([[bos_id] + ids], dtype=torch.long, device=device)
463
+
464
+ tokens_before_gen = idx.shape[1]
465
+
466
+ if idx.shape[1] > context_length:
467
+ idx = idx[:, -context_length:]
468
+
469
+ generated_ids = []
470
+ with torch.no_grad():
471
+ for _ in range(200):
472
+ idx_cond = idx[:, -context_length:]
473
+ logits, _ = model(idx_cond)
474
+ logits = logits[:, -1, :]
475
+ probs = F.softmax(logits / temperature, dim=-1)
476
+ next_id = torch.multinomial(probs, num_samples=1)
477
+ idx = torch.cat([idx, next_id], dim=1)
478
+ generated_ids.append(next_id.item())
479
+
480
+ if next_id.item() == eos_id:
481
+ break
482
+
483
+ partial_text = tokenizer.decode(generated_ids)
484
+ normalized = partial_text.replace(" :", ":").replace(" ,", ",")
485
+ if "User:" in normalized:
486
+ break
487
+
488
+ reply_text = tokenizer.decode(generated_ids)
489
+ normalized_reply = reply_text.replace(" :", ":")
490
+ if "User:" in normalized_reply:
491
+ cut_pos = normalized_reply.index("User:")
492
+ reply_text = reply_text.split("User :")[0].split("User:")[0].strip()
493
+ else:
494
+ reply_text = reply_text.strip()
495
+
496
+ print(f"Cortex_2: {reply_text}")
497
+
498
+ history_lines[-1] = f"Bot: {reply_text}"
499
+
500
+ tokens_used = min(tokens_before_gen + len(generated_ids), context_length)
501
+ pct = tokens_used / context_length * 100
502
+ print(f"๐Ÿ“Š Context: {tokens_used}/{context_length} tokens ({pct:.1f}%)")
503
+
504
+ elif is_code_model:
505
+ if "||" in prompt:
506
+ instruction, extra_input = prompt.split("||", 1)
507
+ else:
508
+ instruction, extra_input = prompt, ""
509
+ instruction = instruction.strip()
510
+
511
+ code_text = extract_code(generate_code(instruction, extra_input))
512
+ ok, err = check_syntax(code_text)
513
+
514
+ attempt = 1
515
+ while not ok and autocheck and attempt < max_auto_attempts:
516
+ attempt += 1
517
+ print(f"๐Ÿ”„ Attempt {attempt}/{max_auto_attempts}: code has an error, regenerating...")
518
+ code_text = extract_code(generate_code(instruction, extra_input))
519
+ ok, err = check_syntax(code_text)
520
+
521
+ print(f"Cortex_2:\n{code_text}")
522
+ last_code = code_text
523
+
524
+ if ok:
525
+ print("โœ… Syntax: no errors found")
526
+ try:
527
+ ans = input("โ–ถ Run this code? [y/N]: ").strip().lower()
528
+ except (EOFError, KeyboardInterrupt):
529
+ ans = ""
530
+ if ans in ("y", "yes"):
531
+ execute_code(code_text)
532
+ else:
533
+ print(f"โŒ Syntax: error found!\n{err}")
534
+ if not autocheck:
535
+ print(" Hint: enable autocheck on โ€” the chat will try to")
536
+ print(" regenerate the code automatically on error.")
537
+
538
+ else:
539
+ ids = tokenizer.encode(prompt).ids
540
+ idx = torch.tensor([[bos_id] + ids], dtype=torch.long, device=device)
541
+
542
+ with torch.no_grad():
543
+ for _ in range(750):
544
+ idx_cond = idx[:, -context_length:]
545
+ logits, _ = model(idx_cond)
546
+ logits = logits[:, -1, :]
547
+ probs = F.softmax(logits / 0.8, dim=-1)
548
+ next_id = torch.multinomial(probs, num_samples=1)
549
+ idx = torch.cat([idx, next_id], dim=1)
550
+ if next_id.item() == eos_id:
551
+ break
552
+
553
+ text = tokenizer.decode(idx[0].tolist())
554
+ print(f"Cortex_2: {text}")
555
+
556
+
557
+ if __name__ == "__main__":
558
+ main()
config.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ {
2
+ "model": "cortex-2-code"
3
+ }
cortex-2-code.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f204b31eed629b50fd335a329cea553706bd6dd2d67e65deab9e87b835f0201b
3
+ size 495938955
requirements.txt ADDED
@@ -0,0 +1,2 @@
 
 
 
1
+ torch
2
+ tokenizers