import torch from model import MiniTransformer from config import * device = "cuda" if torch.cuda.is_available() else "cpu" model = MiniTransformer().to(device) model.load_state_dict( torch.load("mini.pt", map_location=device, weights_only=True) ) model.eval() print("Model loaded. Enter your prompts below. Press Ctrl+C to exit.") context_tensor = None while True: try: prompt = input("\nUser: ") if not prompt.strip(): continue # Append special separator for chat if desired new_tokens = [ord(c) % 256 for c in prompt + "\nBot: "] new_tensor = torch.tensor([new_tokens], dtype=torch.long, device=device) if context_tensor is None: context_tensor = new_tensor else: context_tensor = torch.cat([context_tensor, new_tensor], dim=1) print("Bot: ", end="", flush=True) for _ in range(200): x_crop = context_tensor[:, -BLOCK_SIZE:] with torch.no_grad(): logits = model(x_crop) # Apply temperature scaling to soften probability distributions temperature = 0.8 logits_scaled = logits[:, -1, :] / temperature # Probabilistic sampling to prevent deterministic repetition probs = torch.softmax(logits_scaled, dim=-1) next_token = torch.multinomial(probs, num_samples=1) context_tensor = torch.cat([context_tensor, next_token], dim=1) # Print the generated character dynamically char = chr(int(next_token.item())) print(char, end="", flush=True) # Optional: stop generating if the model outputs a newline # if char == '\n': # break print() # Newline after generation completes except KeyboardInterrupt: print("\nExiting chat loop...") break except EOFError: break