""" Cortex_2 Chat — just put this script in the same folder as your model files and run: python chat.py Required files in the same folder: - best_model.pt (or any .pt model file) - tokenizer.json - config.json Or if your model has tokenizer+config baked in (new format): - best_model.pt (only this one file needed!) """ import torch import torch.nn.functional as F import json import sys import math 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 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(): """Find the model file in current directory.""" here = Path(".") # Check for .pt files pt_files = list(here.glob("*.pt")) # Priority: best_model.pt > final_model.pt > any other .pt for name in ["best_model.pt", "final_model.pt"]: if name in [f.name for f in pt_files]: return here / name # Any .pt file if pt_files: return pt_files[0] return None def main(): device = torch.device("cpu") # Find model file 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) # Allow override via command line if len(sys.argv) > 1: model_path = Path(sys.argv[1]) print(f"Loading model from: {model_path.name}") # Load checkpoint ckpt = torch.load(model_path, map_location=device, weights_only=False) # Load config & tokenizer if "config" in ckpt and "tokenizer" in ckpt: # New format: everything in one file config = ckpt["config"] from tokenizers import Tokenizer tokenizer = Tokenizer.from_str(ckpt["tokenizer"]) print("Loaded config + tokenizer from checkpoint") else: # Old format: separate files here = model_path.parent config_path = here / "config.json" tokenizer_path = here / "tokenizer.json" if not config_path.exists(): print("config.json not found next to model!") sys.exit(1) if not tokenizer_path.exists(): print("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") # Build and load model 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("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" if is_chat_model: print(" Mode: conversational (dataset_mode=chat)") else: print(" 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)") print("=" * 50) bos_id = tokenizer.token_to_id("") eos_id = tokenizer.token_to_id("") context_length = config["context_length"] # For the chat model we keep the full conversation history as text, # in the same "User: ...\nBot: ..." format used during training. history_lines = [] temperature = 0.8 # Chat loop while True: try: prompt = input("\nYou: ").strip() except (EOFError, KeyboardInterrupt): print("\nBye!") 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 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 number. Example: temp 0.9") else: print(f"Current temperature: {temperature} (example to change: temp 0.9)") continue if not prompt: continue if is_chat_model: # Build the full dialogue text: entire history + new turn + "Bot:" 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) # How many history tokens actually fit in context (before truncation) tokens_before_gen = idx.shape[1] # Truncate from the left if history doesn't fit in the model's context 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 # The tokenizer decodes "User:" as "User :" (a space before # the colon — an artifact of the Whitespace pre-tokenizer), # so we check against the normalized form. partial_text = tokenizer.decode(generated_ids) normalized = partial_text.replace(" :", ":").replace(" ,", ",") if "User:" in normalized: break reply_text = tokenizer.decode(generated_ids) # Trim off anything the model "made up" on behalf of the user. # Normalize the space before ":" and cut on the normalized string, # applying the same cut to both versions. normalized_reply = reply_text.replace(" :", ":") if "User:" in normalized_reply: # Simplest approach: cut on the raw text, also matching "User :". reply_text = reply_text.split("User :")[0].split("User:")[0].strip() else: reply_text = reply_text.strip() print(f"Cortex_2: {reply_text}") # Add the model's reply to history for the next turn history_lines[-1] = f"Bot: {reply_text}" # Show how much of the context window is used (history + generated reply) 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}%)") else: # Legacy mode — plain text continuation (story generation) 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) # temperature 0.8 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()