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 context window
11    "BATCH_SIZE": 512,          # Increased batch size
12    "EPOCHS": 20,
13    "EMBEDDING_DIM": 64,
14    "HIDDEN_DIM": 64,
15    "NUM_LAYERS": 1,           # Multi-layer LSTM
16    "DROPOUT": 0.1,
17    "LEARNING_RATE": 0.01,
18    "CLIP_GRAD": 1.0,          # Gradient clipping
19    "LR_GAMMA": 0.95,          # Learning rate decay
20    "VAL_SPLIT": 0.1,          # Validation split
21    "EARLY_STOP_PATIENCE": 3,  # Early stopping patience
22    "MODEL_SAVE_PATH": "char_lm_model.pth",
23    "TEMPERATURE": 0.7,
24    "TOP_K": 5,
25    "TOP_P": 0.95
26}
27
28# Read and process text
29with open(CONFIG["FILE_PATH"], 'r', encoding='utf-8') as f:
30    text = f.read()
31
32# Vocabulary setup
33chars = sorted(list(set(text)))
34vocab_size = len(chars)
35char_to_idx = {ch: i for i, ch in enumerate(chars)}
36idx_to_char = {i: ch for i, ch in enumerate(chars)}
37
38# Encode text
39encoded_text = np.array([char_to_idx[ch] for ch in text])
40
41# Dataset class with train-val split
42class TextDataset(Dataset):
43    def __init__(self, data, seq_length):
44        self.data = data
45        self.seq_length = seq_length
46        
47    def __len__(self):
48        return len(self.data) - self.seq_length - 1
49    
50    def __getitem__(self, idx):
51        x = self.data[idx:idx+self.seq_length]
52        y = self.data[idx+1:idx+self.seq_length+1]
53        return torch.from_numpy(x).long(), torch.from_numpy(y).long()
54
55dataset = TextDataset(encoded_text, CONFIG["SEQ_LENGTH"])
56val_size = int(len(dataset) * CONFIG["VAL_SPLIT"])
57train_size = len(dataset) - val_size
58train_dataset, val_dataset = random_split(dataset, [train_size, val_size])
59
60train_loader = DataLoader(train_dataset, batch_size=CONFIG["BATCH_SIZE"], shuffle=True)
61val_loader = DataLoader(val_dataset, batch_size=CONFIG["BATCH_SIZE"])
62
63# Advanced Model architecture with LSTM and dropout
64class CharLM(nn.Module):
65    def __init__(self):
66        super(CharLM, self).__init__()
67        self.embedding = nn.Embedding(vocab_size, CONFIG["EMBEDDING_DIM"])
68        self.lstm = nn.LSTM(
69            CONFIG["EMBEDDING_DIM"], 
70            CONFIG["HIDDEN_DIM"],
71            num_layers=CONFIG["NUM_LAYERS"],
72            dropout=CONFIG["DROPOUT"] if CONFIG["NUM_LAYERS"] > 1 else 0,
73            batch_first=True
74        )
75        self.dropout = nn.Dropout(CONFIG["DROPOUT"])
76        self.fc = nn.Linear(CONFIG["HIDDEN_DIM"], vocab_size)
77        
78        self.init_weights()
79        
80    def init_weights(self):
81        # Initialize weights for better convergence
82        nn.init.xavier_uniform_(self.embedding.weight)
83        for name, param in self.lstm.named_parameters():
84            if 'weight_ih' in name:
85                nn.init.xavier_uniform_(param.data)
86            elif 'weight_hh' in name:
87                nn.init.orthogonal_(param.data)
88            elif 'bias' in name:
89                param.data.fill_(0)
90        
91    def forward(self, x, hidden=None):
92        x = self.embedding(x)
93        out, hidden = self.lstm(x, hidden)
94        out = self.dropout(out)
95        out = self.fc(out)
96        return out, hidden
97
98model = CharLM()
99criterion = nn.CrossEntropyLoss()
100optimizer = torch.optim.Adam(model.parameters(), lr=CONFIG["LEARNING_RATE"])
101scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=1, gamma=CONFIG["LR_GAMMA"])
102
103# Training loop with validation and early stopping
104best_val_loss = float('inf')
105patience_counter = 0
106
107for epoch in range(CONFIG["EPOCHS"]):
108    model.train()
109    train_loss = 0
110    progress_bar = tqdm(train_loader, desc=f'Epoch {epoch+1}/{CONFIG["EPOCHS"]}')
111    
112    for inputs, targets in progress_bar:
113        optimizer.zero_grad()
114        outputs, _ = model(inputs)
115        loss = criterion(outputs.reshape(-1, vocab_size), targets.reshape(-1))
116        loss.backward()
117        nn.utils.clip_grad_norm_(model.parameters(), CONFIG["CLIP_GRAD"])
118        optimizer.step()
119        train_loss += loss.item()
120        progress_bar.set_postfix({'loss': loss.item()})
121    
122    # Validation phase
123    model.eval()
124    val_loss = 0
125    with torch.no_grad():
126        for inputs, targets in val_loader:
127            outputs, _ = model(inputs)
128            loss = criterion(outputs.reshape(-1, vocab_size), targets.reshape(-1))
129            val_loss += loss.item()
130    
131    avg_train_loss = train_loss / len(train_loader)
132    avg_val_loss = val_loss / len(val_loader)
133    print(f'Epoch {epoch+1} | Train Loss: {avg_train_loss:.4f} | Val Loss: {avg_val_loss:.4f}')
134    
135    # Early stopping and checkpointing
136    if avg_val_loss < best_val_loss:
137        best_val_loss = avg_val_loss
138        torch.save(model.state_dict(), CONFIG["MODEL_SAVE_PATH"])
139        patience_counter = 0
140    else:
141        patience_counter += 1
142        if patience_counter >= CONFIG["EARLY_STOP_PATIENCE"]:
143            print("Early stopping triggered")
144            break
145    
146    scheduler.step()
147
148print(f'Best model saved to {CONFIG["MODEL_SAVE_PATH"]} with validation loss: {best_val_loss:.4f}')
149
150# Advanced Text Generation with multiple sampling methods
151def generate_text(model, start_str, length=200, temperature=CONFIG["TEMPERATURE"], 
152                 top_k=CONFIG["TOP_K"], top_p=CONFIG["TOP_P"]):
153    """
154    Generate text with temperature scaling, top-k, and nucleus (top-p) sampling
155    """
156    model.eval()
157    chars = list(start_str)
158    input_seq = torch.tensor([char_to_idx[ch] for ch in chars]).unsqueeze(0)
159    hidden = None
160    
161    with torch.no_grad():
162        for _ in tqdm(range(length), desc="Generating text"):
163            outputs, hidden = model(input_seq, hidden)
164            logits = outputs[0, -1] / temperature
165            
166            # Apply top-k filtering
167            if top_k > 0:
168                top_vals, top_idx = torch.topk(logits, top_k)
169                logits[logits < top_vals[-1]] = -float('Inf')
170            
171            # Apply nucleus (top-p) filtering
172            if top_p > 0:
173                sorted_logits, sorted_indices = torch.sort(logits, descending=True)
174                cumulative_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
175                sorted_indices_to_remove = cumulative_probs > top_p
176                sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
177                sorted_indices_to_remove[..., 0] = 0
178                indices_to_remove = sorted_indices[sorted_indices_to_remove]
179                logits[indices_to_remove] = -float('Inf')
180
181            probs = torch.softmax(logits, dim=-1)
182            next_char = torch.multinomial(probs, num_samples=1).item()
183            chars.append(idx_to_char[next_char])
184            input_seq = torch.tensor([[next_char]])
185    
186    return ''.join(chars)
187
188# Generation examples with different parameters
189print("\nConservative sampling (temperature=0.5):")
190print(generate_text(model, "The ", temperature=0.5))
191
192print("\nCreative sampling (temperature=1.2, top_p=0.9):")
193print(generate_text(model, "Once ", temperature=1.2, top_p=0.9))
194
195print("\nTop-k sampling (k=5):")
196print(generate_text(model, "In ", top_k=5))
197
198print("\nCombined sampling (temp=0.7, top_k=3, top_p=0.9):")
199print(generate_text(model, "Artificial is ", temperature=0.7, top_k=3, top_p=0.9))