CoolFace
Modelpublic

xcx0902/tiny_llm

sourceHugging Facemitupdated 1y agoView on Hugging Face
2likes
train.py68 linesDownload Raw Back to root
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