Download chat.py from VADRK155/Cortex-2-Code: direct link, hf CLI and curl.
- Browser
- Download file 21.2 kB
-
https://huggingface.co/VADRK155/Cortex-2-Code/resolve/main/chat.py
- Command line
-
hf download hf://VADRK155/Cortex-2-Code/chat.py
-
curl -L -o chat.py https://huggingface.co/VADRK155/Cortex-2-Code/resolve/main/chat.py
21.2 kB
| import torch | |
| import torch.nn.functional as F | |
| import json | |
| import sys | |
| import math | |
| import ast | |
| import os | |
| import time | |
| import subprocess | |
| import tempfile | |
| from pathlib import Path | |
| class CausalSelfAttention(torch.nn.Module): | |
| def __init__(self, d_model, n_heads, dropout, context_length): | |
| super().__init__() | |
| self.n_heads = n_heads | |
| self.head_dim = d_model // n_heads | |
| self.qkv = torch.nn.Linear(d_model, 3 * d_model) | |
| self.proj = torch.nn.Linear(d_model, d_model) | |
| self.attn_dropout = torch.nn.Dropout(dropout) | |
| self.resid_dropout = torch.nn.Dropout(dropout) | |
| self.register_buffer("mask", torch.tril(torch.ones(context_length, context_length)).unsqueeze(0).unsqueeze(0)) | |
| def forward(self, x): | |
| B, T, C = x.shape | |
| qkv = self.qkv(x) | |
| q, k, v = qkv.chunk(3, dim=-1) | |
| q = q.view(B, T, self.n_heads, self.head_dim).transpose(1, 2) | |
| k = k.view(B, T, self.n_heads, self.head_dim).transpose(1, 2) | |
| v = v.view(B, T, self.n_heads, self.head_dim).transpose(1, 2) | |
| attn = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(self.head_dim)) | |
| attn = attn.masked_fill(self.mask[:, :, :T, :T] == 0, float("-inf")) | |
| attn = F.softmax(attn, dim=-1) | |
| attn = self.attn_dropout(attn) | |
| out = attn @ v | |
| out = out.transpose(1, 2).contiguous().view(B, T, C) | |
| out = self.proj(out) | |
| out = self.resid_dropout(out) | |
| return out | |
| class MLP(torch.nn.Module): | |
| def __init__(self, d_model, d_ff, dropout): | |
| super().__init__() | |
| self.net = torch.nn.Sequential( | |
| torch.nn.Linear(d_model, d_ff), | |
| torch.nn.GELU(), | |
| torch.nn.Linear(d_ff, d_model), | |
| torch.nn.Dropout(dropout), | |
| ) | |
| def forward(self, x): | |
| return self.net(x) | |
| class TransformerBlock(torch.nn.Module): | |
| def __init__(self, d_model, n_heads, d_ff, dropout, context_length): | |
| super().__init__() | |
| self.ln1 = torch.nn.LayerNorm(d_model) | |
| self.attn = CausalSelfAttention(d_model, n_heads, dropout, context_length) | |
| self.ln2 = torch.nn.LayerNorm(d_model) | |
| self.mlp = MLP(d_model, d_ff, dropout) | |
| def forward(self, x): | |
| x = x + self.attn(self.ln1(x)) | |
| x = x + self.mlp(self.ln2(x)) | |
| return x | |
| def format_instruction(instruction, extra_input=""): | |
| instruction = (instruction or "").strip() | |
| extra_input = (extra_input or "").strip() | |
| if extra_input and extra_input.lower() != "not applicable": | |
| return f"### Instruction:\n{instruction}\n\n### Input:\n{extra_input}\n\n### Response:\n" | |
| return f"### Instruction:\n{instruction}\n\n### Response:\n" | |
| _FOREIGN_MARKERS = ( | |
| "#include", "void main", "int main(", "public static void", | |
| "System.out.println", "console.log", "function ", "</", | |
| "<?php", "using namespace", "fmt.Println", "package main", | |
| "fn main", "<html", "<script", "CREATE TABLE", "SELECT ", "=>", | |
| ) | |
| def extract_code(text): | |
| text = (text or "").strip() | |
| if "```" not in text: | |
| return text.strip() | |
| def _drop_lang_label(block): | |
| lines = block.split("\n") | |
| if lines and lines[0].strip() and len(lines[0].strip()) <= 12 \ | |
| and not any(ch in lines[0] for ch in " \t=()[]{}:;"): | |
| lines = lines[1:] | |
| return "\n".join(lines).strip("\n") | |
| parts = text.split("```") | |
| blocks = [] | |
| for i in range(1, len(parts), 2): | |
| blocks.append(_drop_lang_label(parts[i])) | |
| if blocks: | |
| return "\n\n".join(b.strip("\n") for b in blocks).strip() | |
| return _drop_lang_label(parts[1]).strip() | |
| def looks_like_python(code): | |
| head = (code or "")[:3000].lower() | |
| return not any(m.lower() in head for m in _FOREIGN_MARKERS) | |
| def check_syntax(code): | |
| if not (code or "").strip(): | |
| return False, "model returned no code (empty response)" | |
| try: | |
| ast.parse(code) | |
| return True, None | |
| except (SyntaxError, ValueError) as e: | |
| if not looks_like_python(code): | |
| return False, ("this doesn't look like Python code โ syntax checking and " | |
| "execution are only supported for Python") | |
| if isinstance(e, ValueError): | |
| return False, f"failed to parse code: {e}" | |
| lines = (code or "").splitlines() | |
| lineno = e.lineno or 1 | |
| offset = e.offset or 1 | |
| out = [f"SyntaxError: {e.msg} (line {lineno}, column {offset})"] | |
| if 1 <= lineno <= len(lines): | |
| bad_line = lines[lineno - 1] | |
| caret_pos = min(max(offset, 1), len(bad_line) + 1) - 1 | |
| out.append(f" {lineno:>4} | {bad_line}") | |
| out.append(f" | {' ' * caret_pos}^") | |
| if lineno >= len(lines): | |
| out.append(" (looks like the code was cut off by the generation limit โ " | |
| "try increasing code_max_new_tokens)") | |
| return False, "\n".join(out) | |
| def run_python_code(code, timeout=10.0): | |
| fd, path = tempfile.mkstemp(suffix=".py", prefix="cortex_run_") | |
| try: | |
| with os.fdopen(fd, "w", encoding="utf-8") as f: | |
| f.write(code) | |
| env = {**os.environ, "PYTHONIOENCODING": "utf-8"} | |
| proc = subprocess.run( | |
| [sys.executable, "-u", path], | |
| stdin=subprocess.DEVNULL, | |
| capture_output=True, | |
| text=True, | |
| encoding="utf-8", | |
| errors="replace", | |
| timeout=timeout, | |
| env=env, | |
| ) | |
| return proc.returncode, (proc.stdout or "") + (proc.stderr or ""), False | |
| except subprocess.TimeoutExpired as e: | |
| partial = "" | |
| for stream in (e.stdout, e.stderr): | |
| if not stream: | |
| continue | |
| if isinstance(stream, bytes): | |
| stream = stream.decode("utf-8", "replace") | |
| partial += stream | |
| return -1, partial, True | |
| finally: | |
| try: | |
| os.unlink(path) | |
| except OSError: | |
| pass | |
| class TinyGPT(torch.nn.Module): | |
| def __init__(self, config): | |
| super().__init__() | |
| self.config = config | |
| vocab_size = config["tokenizer_vocab_size"] + 10 | |
| self.token_emb = torch.nn.Embedding(vocab_size, config["d_model"]) | |
| self.pos_emb = torch.nn.Embedding(config["context_length"], config["d_model"]) | |
| self.drop = torch.nn.Dropout(config["dropout"]) | |
| self.blocks = torch.nn.ModuleList([ | |
| TransformerBlock(config["d_model"], config["n_heads"], config["d_ff"], config["dropout"], config["context_length"]) | |
| for _ in range(config["n_layers"]) | |
| ]) | |
| self.ln_f = torch.nn.LayerNorm(config["d_model"]) | |
| self.head = torch.nn.Linear(config["d_model"], vocab_size, bias=False) | |
| self.token_emb.weight = self.head.weight | |
| def forward(self, idx, targets=None): | |
| B, T = idx.shape | |
| pos = torch.arange(0, T, device=idx.device).unsqueeze(0) | |
| x = self.token_emb(idx) + self.pos_emb(pos) | |
| x = self.drop(x) | |
| for block in self.blocks: | |
| x = block(x) | |
| x = self.ln_f(x) | |
| logits = self.head(x) | |
| loss = None | |
| if targets is not None: | |
| loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1), ignore_index=0) | |
| return logits, loss | |
| def find_model_file(): | |
| here = Path(".") | |
| pt_files = list(here.glob("*.pt")) | |
| for name in ["best_model.pt", "final_model.pt"]: | |
| if name in [f.name for f in pt_files]: | |
| return here / name | |
| if pt_files: | |
| return pt_files[0] | |
| return None | |
| def main(): | |
| device = torch.device("cuda") | |
| model_path = find_model_file() | |
| if model_path is None: | |
| print("โ No .pt model file found! Put this script in the same folder as your model.") | |
| sys.exit(1) | |
| if len(sys.argv) > 1: | |
| model_path = Path(sys.argv[1]) | |
| print(f"๐ Loading model from: {model_path.name}") | |
| ckpt = torch.load(model_path, map_location=device, weights_only=False) | |
| if "config" in ckpt and "tokenizer" in ckpt: | |
| config = ckpt["config"] | |
| from tokenizers import Tokenizer | |
| tokenizer = Tokenizer.from_str(ckpt["tokenizer"]) | |
| print("๐ฆ Loaded config + tokenizer from checkpoint") | |
| else: | |
| here = model_path.parent | |
| config_path = here / "config.json" | |
| tokenizer_path = here / "tokenizer.json" | |
| if not config_path.exists(): | |
| print(f"โ config.json not found next to model!") | |
| sys.exit(1) | |
| if not tokenizer_path.exists(): | |
| print(f"โ tokenizer.json not found next to model!") | |
| sys.exit(1) | |
| with open(config_path) as f: | |
| config = json.load(f) | |
| from tokenizers import Tokenizer | |
| tokenizer = Tokenizer.from_file(str(tokenizer_path)) | |
| print("๐ฆ Loaded config + tokenizer from separate files") | |
| model = TinyGPT(config).to(device) | |
| model.load_state_dict(ckpt["model"]) | |
| model.eval() | |
| n_params = sum(p.numel() for p in model.parameters()) | |
| step = ckpt.get("step", "?") | |
| val_loss = ckpt.get("val_loss", "?") | |
| if isinstance(val_loss, float): | |
| val_loss = f"{val_loss:.4f}" | |
| print(f"โ Cortex_2 loaded!") | |
| print(f" Parameters: {n_params / 1e6:.1f}M") | |
| print(f" Step: {step}") | |
| print(f" Val loss: {val_loss}") | |
| print(f" Device: {device}") | |
| dataset_mode = config.get("dataset_mode", "stories") | |
| is_chat_model = dataset_mode == "chat" | |
| is_code_model = dataset_mode == "code" | |
| if is_chat_model: | |
| print(f" Mode: ๐ฌ conversational (dataset_mode=chat)") | |
| elif is_code_model: | |
| print(f" Mode: ๐งโ๐ป code (dataset_mode=code)") | |
| else: | |
| print(f" Mode: ๐ story completion (dataset_mode=stories)") | |
| print() | |
| print("๐ฌ Type a prompt and press Enter. Type 'quit' to exit.") | |
| if is_chat_model: | |
| print(" (type 'reset' to clear conversation history)") | |
| print(" (type 'temp 0.9' to change temperature, current default: 0.8)") | |
| if is_code_model: | |
| print(" Describe a task, e.g.: 'Write a function that reverses a string'.") | |
| print(" (to set a separate 'Input:', type: task || input)") | |
| print(" (type 'temp 0.5' to change temperature, current default: 0.5)") | |
| print() | |
| print(" Code mode commands:") | |
| print(" run โ run the last generated code") | |
| print(" save โ save the last code to generated_code_NN.py") | |
| print(" autocheck โ auto-regenerate on syntax error") | |
| print(" timeout N โ code execution timeout in seconds") | |
| print(" After generation the code is syntax-checked, and clean code can be") | |
| print(" run directly from the chat (y when asked 'Run?').") | |
| print("=" * 50) | |
| bos_id = tokenizer.token_to_id("<bos>") | |
| eos_id = tokenizer.token_to_id("<eos>") | |
| context_length = config["context_length"] | |
| history_lines = [] | |
| temperature = 0.5 if is_code_model else 0.8 | |
| code_max_new_tokens = 400 | |
| code_top_k = 40 | |
| last_code = None | |
| run_timeout = 10.0 | |
| autocheck = True | |
| max_auto_attempts = 3 | |
| def generate_code(instruction, extra_input=""): | |
| text_prompt = format_instruction(instruction, extra_input) | |
| ids = tokenizer.encode(text_prompt).ids | |
| idx = torch.tensor([[bos_id] + ids], dtype=torch.long, device=device) | |
| prompt_len = idx.shape[1] | |
| t0 = time.time() | |
| n_tokens = 0 | |
| with torch.no_grad(): | |
| for _ in range(code_max_new_tokens): | |
| idx_cond = idx[:, -context_length:] | |
| logits, _ = model(idx_cond) | |
| logits = logits[:, -1, :] / temperature | |
| if code_top_k: | |
| kth = torch.topk(logits, code_top_k).values[:, -1, None] | |
| logits = logits.masked_fill(logits < kth, float("-inf")) | |
| probs = F.softmax(logits, dim=-1) | |
| next_id = torch.multinomial(probs, num_samples=1) | |
| idx = torch.cat([idx, next_id], dim=1) | |
| n_tokens += 1 | |
| if next_id.item() == eos_id: | |
| break | |
| print(f" โณ generated {n_tokens} tokens in {time.time() - t0:.1f}s") | |
| return tokenizer.decode(idx[0, prompt_len:].tolist()) | |
| def execute_code(code): | |
| print("โ" * 50) | |
| print(f"โถ Running code (separate process, timeout {run_timeout:.0f}s, stdin closed)...") | |
| rc, output, timed_out = run_python_code(code, run_timeout) | |
| if timed_out: | |
| print(f"โฑ Timeout exceeded ({run_timeout:.0f}s) โ process stopped.") | |
| if output.strip(): | |
| print("๐ค Output before stopping:") | |
| print(output.rstrip()) | |
| print(" Hint: if the code waits for input(), it will never finish โ") | |
| print(" interactive input is not available when running from chat.") | |
| elif rc == 0: | |
| if output.strip(): | |
| print("๐ค Program output:") | |
| print(output.rstrip()) | |
| else: | |
| print("๐ค Program finished with no output.") | |
| print("โ Code ran without errors (exit code 0).") | |
| else: | |
| if output.strip(): | |
| print("๐ค Program output:") | |
| print(output.rstrip()) | |
| if "EOFError" in output: | |
| print(" Hint: the code called input() โ input is not available when running from chat.") | |
| print(f"โ Program finished with an error (exit code {rc}).") | |
| print("โ" * 50) | |
| # Chat loop | |
| while True: | |
| try: | |
| prompt = input("\nYou: ").strip() | |
| except (EOFError, KeyboardInterrupt): | |
| print("\n๐ Bye!") | |
| break | |
| if prompt.lower() == "quit": | |
| print("๐ Bye!") | |
| break | |
| if is_chat_model and prompt.lower() == "reset": | |
| history_lines = [] | |
| print("๐ Conversation history cleared.") | |
| continue | |
| if (is_chat_model or is_code_model) and prompt.lower().startswith("temp"): | |
| parts = prompt.split() | |
| if len(parts) == 2: | |
| try: | |
| new_temp = float(parts[1]) | |
| if new_temp <= 0: | |
| print("โ ๏ธ Temperature must be greater than 0.") | |
| else: | |
| temperature = new_temp | |
| print(f"๐ก๏ธ Temperature set to: {temperature}") | |
| except ValueError: | |
| print("โ ๏ธ Could not parse the number. Example: temp 0.9") | |
| else: | |
| print(f"๐ก๏ธ Current temperature: {temperature} (example to change: temp 0.9)") | |
| continue | |
| if is_code_model and prompt.lower() in ("run", "r"): | |
| if not last_code: | |
| print("โ ๏ธ Nothing to run yet โ generate some code first.") | |
| continue | |
| ok, err = check_syntax(last_code) | |
| if not ok: | |
| print(f"โ The last code has a syntax error, cannot run it:\n{err}") | |
| continue | |
| execute_code(last_code) | |
| continue | |
| if is_code_model and prompt.lower() == "save": | |
| if not last_code: | |
| print("โ ๏ธ Nothing to save yet โ generate some code first.") | |
| continue | |
| n = 1 | |
| while (Path.cwd() / f"generated_code_{n:02d}.py").exists(): | |
| n += 1 | |
| save_path = Path.cwd() / f"generated_code_{n:02d}.py" | |
| save_path.write_text(last_code, encoding="utf-8") | |
| print(f"๐พ Code saved: {save_path}") | |
| continue | |
| if is_code_model and prompt.lower().startswith("autocheck"): | |
| parts = prompt.split() | |
| if len(parts) == 2 and parts[1].lower() in ("on", "off"): | |
| autocheck = parts[1].lower() == "on" | |
| state = "on" if autocheck else "off" | |
| print(f"๐ Auto-regenerate on error: {state} (max attempts: {max_auto_attempts})") | |
| else: | |
| state = "on" if autocheck else "off" | |
| print(f"๐ Auto-regenerate is currently: {state} (example: autocheck off)") | |
| continue | |
| if is_code_model and prompt.lower().startswith("timeout"): | |
| parts = prompt.split() | |
| if len(parts) == 2: | |
| try: | |
| val = float(parts[1]) | |
| if val <= 0: | |
| print("โ ๏ธ Timeout must be greater than 0.") | |
| else: | |
| run_timeout = val | |
| print(f"โฑ Code execution timeout: {run_timeout:.0f}s") | |
| except ValueError: | |
| print("โ ๏ธ Could not parse the number. Example: timeout 15") | |
| else: | |
| print(f"โฑ Current execution timeout: {run_timeout:.0f}s (example: timeout 15)") | |
| continue | |
| if not prompt: | |
| continue | |
| if is_chat_model: | |
| history_lines.append(f"User: {prompt}") | |
| history_lines.append("Bot:") | |
| full_text = "\n".join(history_lines) | |
| ids = tokenizer.encode(full_text).ids | |
| idx = torch.tensor([[bos_id] + ids], dtype=torch.long, device=device) | |
| tokens_before_gen = idx.shape[1] | |
| if idx.shape[1] > context_length: | |
| idx = idx[:, -context_length:] | |
| generated_ids = [] | |
| with torch.no_grad(): | |
| for _ in range(200): | |
| idx_cond = idx[:, -context_length:] | |
| logits, _ = model(idx_cond) | |
| logits = logits[:, -1, :] | |
| probs = F.softmax(logits / temperature, dim=-1) | |
| next_id = torch.multinomial(probs, num_samples=1) | |
| idx = torch.cat([idx, next_id], dim=1) | |
| generated_ids.append(next_id.item()) | |
| if next_id.item() == eos_id: | |
| break | |
| partial_text = tokenizer.decode(generated_ids) | |
| normalized = partial_text.replace(" :", ":").replace(" ,", ",") | |
| if "User:" in normalized: | |
| break | |
| reply_text = tokenizer.decode(generated_ids) | |
| normalized_reply = reply_text.replace(" :", ":") | |
| if "User:" in normalized_reply: | |
| cut_pos = normalized_reply.index("User:") | |
| reply_text = reply_text.split("User :")[0].split("User:")[0].strip() | |
| else: | |
| reply_text = reply_text.strip() | |
| print(f"Cortex_2: {reply_text}") | |
| history_lines[-1] = f"Bot: {reply_text}" | |
| tokens_used = min(tokens_before_gen + len(generated_ids), context_length) | |
| pct = tokens_used / context_length * 100 | |
| print(f"๐ Context: {tokens_used}/{context_length} tokens ({pct:.1f}%)") | |
| elif is_code_model: | |
| if "||" in prompt: | |
| instruction, extra_input = prompt.split("||", 1) | |
| else: | |
| instruction, extra_input = prompt, "" | |
| instruction = instruction.strip() | |
| code_text = extract_code(generate_code(instruction, extra_input)) | |
| ok, err = check_syntax(code_text) | |
| attempt = 1 | |
| while not ok and autocheck and attempt < max_auto_attempts: | |
| attempt += 1 | |
| print(f"๐ Attempt {attempt}/{max_auto_attempts}: code has an error, regenerating...") | |
| code_text = extract_code(generate_code(instruction, extra_input)) | |
| ok, err = check_syntax(code_text) | |
| print(f"Cortex_2:\n{code_text}") | |
| last_code = code_text | |
| if ok: | |
| print("โ Syntax: no errors found") | |
| try: | |
| ans = input("โถ Run this code? [y/N]: ").strip().lower() | |
| except (EOFError, KeyboardInterrupt): | |
| ans = "" | |
| if ans in ("y", "yes"): | |
| execute_code(code_text) | |
| else: | |
| print(f"โ Syntax: error found!\n{err}") | |
| if not autocheck: | |
| print(" Hint: enable autocheck on โ the chat will try to") | |
| print(" regenerate the code automatically on error.") | |
| else: | |
| ids = tokenizer.encode(prompt).ids | |
| idx = torch.tensor([[bos_id] + ids], dtype=torch.long, device=device) | |
| with torch.no_grad(): | |
| for _ in range(750): | |
| idx_cond = idx[:, -context_length:] | |
| logits, _ = model(idx_cond) | |
| logits = logits[:, -1, :] | |
| probs = F.softmax(logits / 0.8, dim=-1) | |
| next_id = torch.multinomial(probs, num_samples=1) | |
| idx = torch.cat([idx, next_id], dim=1) | |
| if next_id.item() == eos_id: | |
| break | |
| text = tokenizer.decode(idx[0].tolist()) | |
| print(f"Cortex_2: {text}") | |
| if __name__ == "__main__": | |
| main() | |