Changes
1 changed files (+69/-3)
-
-
@@ -1,9 +1,12 @@#!/usr/bin/python3 import re from collections import Counter import psycopg2 import torch import torch.nn as nn from torch import nn, optim from torch.utils.data import DataLoader # Fetch messages from database since it's way faster than using the API
-
@@ -15,10 +18,9 @@ statuses = cur.fetchall()# Get the messages as plain text # TODO: Remove punctuation and other junk text = [re.sub(r'<[^>]*>', '', status[2]) for status in statuses] # Use regex to remove HTML stuff print(text[0:100]) #print(text[0:100]) # https://closeheat.com/blog/pytorch-lstm-text-generation-tutorial class Model(nn.Module): def __init__(self, dataset): super(Model, self).__init__()
-
@@ -50,3 +52,67 @@ class Model(nn.Module):def init_state(self, sequence_length): return (torch.zeros(self.num_layers, sequence_length, self.lstm_size), torch.zeros(self.num_layers, sequence_length, self.lstm_size)) class Dataset(torch.utils.data.Dataset): def __init__(self): self.words = [word for message in text for word in message.split()] self.word_counts = Counter(self.words) self.uniq_words = sorted(self.word_counts, key=self.word_counts.get) self.index_to_word = {index: word for index, word in enumerate(self.uniq_words)} self.word_to_index = {word: index for index, word in enumerate(self.uniq_words)} self.words_indexes = [self.word_to_index[w] for w in self.words] def __len__(self): return len(self.words_indexes) - 4 def __getitem__(self, index): return (torch.tensor(self.words_indexes[index:index+4]), torch.tensor(self.words_indexes[index+1:index+4+1])) dataset = Dataset() model = Model(dataset) dataloader = DataLoader(dataset, batch_size=1024) criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=0.001) for epoch in range(1): state_h, state_c = model.init_state(4) for batch, (x, y) in enumerate(dataloader): optimizer.zero_grad() y_pred, (state_h, state_c) = model(x, (state_h, state_c)) loss = criterion(y_pred.transpose(1, 2), y) state_h = state_h.detach() state_c = state_c.detach() loss.backward() optimizer.step() print({ 'epoch': epoch, 'batch': batch, 'loss': loss.item() }) def predict(text, next_words=100): model.eval() words = text.split(' ') state_h, state_c = model.init_state(len(words)) for i in range(0, next_words): x = torch.tensor([[dataset.word_to_index[w] for w in words[i:]]]) y_pred, (state_h, state_c) = model(x, (state_h, state_c)) last_word_logits = y_pred[0][-1] p = torch.nn.functional.softmax(last_word_logits, dim=0).detach().numpy() word_index = np.random.choice(len(last_word_logits), p=p) words.append(dataset.index_to_word[word_index]) return words predict('This is a test')
-