BharathLanka/Neural_Network_Based_Language_Model_for_Next_Token_Prediction
class LSTMModel(nn.Module): def _init(self, vocabsize, nembd, nhidden, blocksize, dropout): super(LSTMModel, self).init() self.embedding = nn.Embedding(vocabsize, nembd) self.lstm = nn.LSTM(nembd, nhidden, batchfirst=True) self.fc = nn.Linear(nhidden, vocabsize) self.dropout = nn.Dropout(dropout)
def forward(self, x): x = self.embedding(x) x, _ = self.lstm(x) x = self.fc(self.dropout(x)) return x
Model, optimizer, and loss function initialization
vocabsize = len(wordtoidx) model = LSTMModel(vocabsize, nembd, nhidden, blocksize, dropout).to(device) optimizer = optim.Adam(model.parameters(), lr=learningrate) loss_fn = nn.CrossEntropyLoss()
traininglosses = [] validationlosses = []
Train and Validation Function
def trainmodel(): for epoch in range(maxiters): model.train() totaltrainloss = 0 for batchidx, (x, y) in enumerate(trainloader): x = x.to(device) y = y.to(device)
# Forward pass logits = model(x) logits = logits.view(-1, vocab_size) y = y.view(-1)
# Compute loss loss = loss_fn(logits, y)
# Backpropagation optimizer.zero_grad() loss.backward() optimizer.step()
totaltrainloss += loss.item()
avgtrainloss = totaltrainloss / len(train_loader)
# Validation model.eval() totalvalloss = 0 with torch.nograd(): for valx, valy in valloader: valx = valx.to(device) valy = valy.to(device)
vallogits = model(valx) vallogits = vallogits.view(-1, vocabsize) valy = val_y.view(-1)
valloss = lossfn(vallogits, valy) totalvalloss += val_loss.item()
avgvalloss = totalvalloss / len(val_loader)
# Append loss values to the lists traininglosses.append(avgtrainloss) validationlosses.append(avgvalloss)
print(f'Epoch {epoch + 1}/{maxiters}, Train Loss: {avgtrainloss:.4f}, Validation Loss: {avgval_loss:.4f}')
Step 1: Train the model and collect losses
train_model()
