FlowRes-1 / generate.py
arpecious's picture
Upload 18 files
dbd41fe verified
Raw
History Blame Contribute Delete
2.06 kB
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