xcx0902/tiny_llm
2
1import torch
2from SimpleRNN import SimpleRNN
3import json
4from tqdm import tqdm, trange
5
6parameters = json.loads(open("parameter.json").read())
7model_path = parameters["model_path"]
8
9model = torch.load(model_path, weights_only=False)
10with open("vocab.json", "r") as f:
11 chars = json.loads(f.read())
12char_to_idx = {ch: i for i, ch in enumerate(chars)}
13idx_to_char = {i: ch for i, ch in enumerate(chars)}
14print("Loaded pre-trained model.")
15
16input_size = len(chars)
17hidden_size = parameters["hidden_size"]
18output_size = len(chars)
19
20def generate_text(start_text, length):
21 model.eval()
22 hidden = torch.zeros(1, 1, hidden_size)
23 input_seq = torch.tensor([char_to_idx[ch] for ch in start_text])
24
25 generated_text = start_text
26 for _ in trange(length):
27 output, hidden = model(input_seq, hidden)
28 predicted_idx = output.argmax().item()
29 generated_text += idx_to_char[predicted_idx]
30 input_seq = torch.cat((input_seq[1:], torch.tensor([predicted_idx])))
31
32 return generated_text
33
34while True:
35 prompt = input("Ask LLM: ")
36 length = int(input("Length of text: "))
37 print("LLM Output: ", generate_text(prompt, length))
38 