टेक्स्ट के लिए LSTM मॉडल बनाना
PyBooks में, टीम लगातार नवीनतम तकनीकी प्रगति का लाभ उठाकर यूज़र अनुभव बेहतर करने की कोशिश करती रहती है. इसी दृष्टि के तहत, उन्होंने आपको एक अहम काम सौंपा है. टीम चाहती है कि आप एक और शक्तिशाली टूल की संभावनाएँ जाँचें: LSTM, जो डेटा पैटर्न में मौजूद जटिलताओं को बेहतर तरीके से कैप्चर करने के लिए जाना जाता है. आप पहले वाले ही Newsgroup डेटासेट के साथ काम कर रहे हैं, और उद्देश्य वही है: न्यूज़ आर्टिकल्स को तीन अलग-अलग केटेगरी में क्लासिफाई करना:
rec.autos, sci.med, और comp.graphics.
आपके लिए ये पैकेज लोड कर दिए गए हैं: torch, nn, optim.
यह अभ्यास पाठ्यक्रम का हिस्सा है
PyTorch के साथ टेक्स्ट के लिए डीप लर्निंग
अभ्यास निर्देश
- ज़रूरी पैरामीटर्स के साथ LSTM और linear लेयर्स पूरी करके एक LSTM मॉडल सेट अप करें.
- ज़रूरी पैरामीटर्स के साथ मॉडल को initialize करें.
- ग्रेडिएंट्स को ज़ीरो पर रीसेट करके और इनपुट डेटा
X_train_seqको मॉडल से पास करके LSTM मॉडल को ट्रेन करें. - प्रेडिक्टेड
outputsऔर सही लेबल्स के आधार पर loss कैलकुलेट करें.
इंटरैक्टिव व्यावहारिक अभ्यास
इस अभ्यास को इस नमूना कोड को पूरा करके आज़माएँ।
# Initialize the LSTM and the output layer with parameters
class LSTMModel(nn.Module):
def __init__(self, input_size, hidden_size, num_layers, num_classes):
super(LSTMModel, self).__init__()
self.hidden_size = hidden_size
self.num_layers = num_layers
self.lstm = nn.LSTM(____, ____, ____, batch_first=True)
self.fc = nn.Linear(____, ____)
def forward(self, x):
h0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size)
c0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size)
out, _ = self.lstm(x, (h0, c0))
out = out[:, -1, :]
out = self.fc(out)
return out
# Initialize model with required parameters
lstm_model = LSTMModel(____, ____, ____, ____)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(lstm_model.parameters(), lr=0.01)
# Train the model by passing the correct parameters and zeroing the gradient
for epoch in range(10):
optimizer.____
outputs = lstm_model(____)
loss = criterion(____, y_train_seq)
loss.backward()
optimizer.step()
print(f'Epoch: {epoch+1}, Loss: {loss.item()}')