VADRK155's picture
Upload folder using huggingface_hub
f1890d3 verified
Raw History Blame Contribute Delete
12 kB
"""
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("<bos>")
eos_id = tokenizer.token_to_id("<eos>")
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()