| 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 |
| |
| |
| 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) |
| |
| |
| temperature = 0.8 |
| logits_scaled = logits[:, -1, :] / temperature |
| |
| |
| 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) |
| |
| |
| char = chr(int(next_token.item())) |
| print(char, end="", flush=True) |
| |
| |
| |
| |
| |
| print() |
| |
| except KeyboardInterrupt: |
| print("\nExiting chat loop...") |
| break |
| except EOFError: |
| break |
|
|