xcx0902/tiny_llm
2
1import torch
2import torch.nn as nn
3import torch.optim as optim
4from SimpleRNN import SimpleRNN
5import os
6import json
7from tqdm import tqdm, trange
8import time
9
10training_text = open("train_data.txt", encoding="utf-8").read()
11chars = sorted(list(set(training_text))) # Unique characters
12char_to_idx = {ch: i for i, ch in enumerate(chars)}
13idx_to_char = {i: ch for i, ch in enumerate(chars)}
14
15parameters = json.loads(open("parameter.json").read())
16input_size = len(chars)
17hidden_size = parameters["hidden_size"]
18output_size = len(chars)
19sequence_length = parameters["sequence_length"]
20epochs = 1000
21learning_rate = parameters["learning_rate"]
22model_path = parameters["model_path"]
23
24train_data = []
25for i in range(len(training_text) - sequence_length):
26 input_seq = training_text[i : i + sequence_length]
27 target_char = training_text[i + sequence_length]
28 train_data.append((torch.tensor([char_to_idx[ch] for ch in input_seq]), char_to_idx[target_char]))
29
30if os.path.exists(model_path):
31 model = torch.load(model_path, weights_only=False)
32 print("Loaded pre-trained model. Continue training...")
33else:
34 print("Training new model...")
35 model = SimpleRNN(input_size, hidden_size, output_size)
36
37criterion = nn.CrossEntropyLoss()
38optimizer = optim.Adam(model.parameters(), lr=learning_rate)
39for epoch in range(epochs):
40 try:
41 total_loss = 0
42 hidden = torch.zeros(1, 1, hidden_size)
43
44 pbar = tqdm(train_data, desc=f"Epoch={epoch}, Loss=N/A")
45 count = 0
46 for input_seq, target in pbar:
47 count += 1
48 optimizer.zero_grad()
49 output, hidden = model(input_seq, hidden.detach())
50 loss = criterion(output, torch.tensor([target]))
51 loss.backward()
52 optimizer.step()
53 total_loss += loss.item()
54 pbar.desc = f"Epoch={epoch}, Loss={total_loss / count:.12f}"
55
56 pbar.close()
57 time.sleep(1)
58 except KeyboardInterrupt:
59 break
60
61hidden = torch.zeros(1, 1, hidden_size)
62output, hidden = model(input_seq, hidden.detach())
63
64torch.save(model, model_path)
65with open("vocab.json", "w") as f:
66 f.write(json.dumps(chars))
67print("Model saved.")
68 