CoolFace
Datasetpublic

ysn-rfd/text-dataset-tiny-code-script-py-format

USED of tahamajs/medicine_ds_persian for .parquet file USED of Alijafarixcs2/persian-it-llama2-2k for .parquet file USED of Abirate/english_quotes for .jsonl file NEW FILES (05/12/2025) NEW FILES (12/26/2025) NEW FILES (02/15/2026)

sourceHugging Faceapache-2.0updated 4mo agoView on Hugging Face
3likes1.6kdownloads
1import torch
2import torch.nn as nn
3from torch.utils.data import Dataset, DataLoader, random_split
4import numpy as np
5from tqdm import tqdm
6
7# Configuration
8CONFIG = {
9    "FILE_PATH": 'dataset.txt',
10    "SEQ_LENGTH": 32,          # Increased sequence length
11    "BATCH_SIZE": 8,          # Increased batch size
12    "EPOCHS": 1,
13    "EMBEDDING_DIM": 64,      # Deeper embedding layer
14    "HIDDEN_DIM": 64,         # Larger hidden dimension
15    "NUM_LAYERS": 1,           # More LSTM layers
16    "BIDIRECTIONAL": False,    # Optional bidirectionality
17    "DROPOUT": 0.3,
18    "LEARNING_RATE": 0.01,
19    "CLIP_GRAD": 1.0,          # Gradient clipping
20    "LR_GAMMA": 0.9,           # Learning rate decay
21    "VAL_SPLIT": 0.1,          # Validation split
22    "EARLY_STOP_PATIENCE": 3,  # Early stopping patience
23    "MODEL_SAVE_PATH": "char_lm_advanced.pth",
24    "TEMPERATURE": 0.7,
25    "TOP_K": 10,
26    "TOP_P": 0.95
27}
28
29# Check for GPU
30device = torch.device("cpu")
31
32# Read and process text
33with open(CONFIG["FILE_PATH"], 'r', encoding='utf-8') as f:
34    text = f.read()
35
36chars = sorted(set(text))
37vocab_size = len(chars)
38char_to_idx = {ch: i for i, ch in enumerate(chars)}
39idx_to_char = {i: ch for i, ch in enumerate(chars)}
40encoded_text = np.array([char_to_idx[ch] for ch in text])
41
42# Dataset Class
43class TextDataset(Dataset):
44    def __init__(self, data, seq_length):
45        self.data = data
46        self.seq_length = seq_length
47    
48    def __len__(self):
49        return len(self.data) - self.seq_length - 1
50    
51    def __getitem__(self, idx):
52        x = self.data[idx:idx+self.seq_length]
53        y = self.data[idx+1:idx+self.seq_length+1]
54        return torch.from_numpy(x).long(), torch.from_numpy(y).long()
55
56# Splitting dataset
57full_dataset = TextDataset(encoded_text, CONFIG["SEQ_LENGTH"])
58val_size = int(len(full_dataset) * CONFIG["VAL_SPLIT"])
59train_size = len(full_dataset) - val_size
60train_dataset, val_dataset = random_split(full_dataset, [train_size, val_size])
61
62train_loader = DataLoader(train_dataset, batch_size=CONFIG["BATCH_SIZE"], shuffle=True)
63val_loader = DataLoader(val_dataset, batch_size=CONFIG["BATCH_SIZE"])
64
65# Advanced LSTM Model
66class CharLM(nn.Module):
67    def __init__(self):
68        super(CharLM, self).__init__()
69        self.embedding = nn.Embedding(vocab_size, CONFIG["EMBEDDING_DIM"])
70        self.lstm = nn.LSTM(
71            CONFIG["EMBEDDING_DIM"], CONFIG["HIDDEN_DIM"], CONFIG["NUM_LAYERS"],
72            dropout=CONFIG["DROPOUT"], bidirectional=CONFIG["BIDIRECTIONAL"], batch_first=True
73        )
74        self.layer_norm = nn.LayerNorm(CONFIG["HIDDEN_DIM"])
75        self.fc = nn.Linear(CONFIG["HIDDEN_DIM"], vocab_size)
76        self.dropout = nn.Dropout(CONFIG["DROPOUT"])
77        self.init_weights()
78
79    def init_weights(self):
80        nn.init.xavier_uniform_(self.embedding.weight)
81        for name, param in self.lstm.named_parameters():
82            if 'weight' in name:
83                nn.init.xavier_uniform_(param)
84            elif 'bias' in name:
85                param.data.fill_(0)
86
87    def forward(self, x, hidden=None):
88        x = self.embedding(x)
89        out, hidden = self.lstm(x, hidden)
90        out = self.layer_norm(out)
91        out = self.dropout(out)
92        out = self.fc(out)
93        return out, hidden
94
95model = CharLM().to(device)
96criterion = nn.CrossEntropyLoss()
97optimizer = torch.optim.AdamW(model.parameters(), lr=CONFIG["LEARNING_RATE"])
98scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=1, gamma=CONFIG["LR_GAMMA"])
99scaler = torch.cuda.amp.GradScaler()
100
101# Training with Mixed Precision & Early Stopping
102best_val_loss = float('inf')
103patience_counter = 0
104
105for epoch in range(CONFIG["EPOCHS"]):
106    model.train()
107    train_loss = 0
108    progress_bar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{CONFIG["EPOCHS"]}')
109    
110    for inputs, targets in progress_bar:
111        inputs, targets = inputs.to(device), targets.to(device)
112        optimizer.zero_grad()
113        
114        with torch.cuda.amp.autocast():
115            outputs, _ = model(inputs)
116            loss = criterion(outputs.reshape(-1, vocab_size), targets.reshape(-1))
117        
118        scaler.scale(loss).backward()
119        scaler.unscale_(optimizer)
120        nn.utils.clip_grad_norm_(model.parameters(), CONFIG["CLIP_GRAD"])
121        scaler.step(optimizer)
122        scaler.update()
123        
124        train_loss += loss.item()
125        progress_bar.set_postfix({'loss': loss.item()})
126    
127    # Validation
128    model.eval()
129    val_loss = 0
130    with torch.no_grad():
131        for inputs, targets in val_loader:
132            inputs, targets = inputs.to(device), targets.to(device)
133            outputs, _ = model(inputs)
134            loss = criterion(outputs.reshape(-1, vocab_size), targets.reshape(-1))
135            val_loss += loss.item()
136    
137    avg_train_loss = train_loss / len(train_loader)
138    avg_val_loss = val_loss / len(val_loader)
139    print(f'Epoch {epoch+1} | Train Loss: {avg_train_loss:.4f} | Val Loss: {avg_val_loss:.4f}')
140    
141    if avg_val_loss < best_val_loss:
142        best_val_loss = avg_val_loss
143        torch.save(model.state_dict(), CONFIG["MODEL_SAVE_PATH"])
144        patience_counter = 0
145    else:
146        patience_counter += 1
147        if patience_counter >= CONFIG["EARLY_STOP_PATIENCE"]:
148            print("Early stopping triggered")
149            break
150    
151    scheduler.step()
152
153print(f'Model saved to {CONFIG["MODEL_SAVE_PATH"]} with best val loss: {best_val_loss:.4f}')
154