Changes
1 changed files (+22/-17)
-
-
@@ -16,9 +16,10 @@ cur.execute('SELECT * FROM statuses')statuses = cur.fetchall() # Get the messages as plain text # Use regex to remove HTML stuff # 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]) text = [re.sub(r'<[^>]*>', '', status[2]) for status in statuses] # print(text[0:100]) class Model(nn.Module):
-
@@ -41,14 +42,14 @@ class Model(nn.Module):dropout=0.2 ) self.fc = nn.Linear(self.lstm_size, n_vocab) def forward(self, x, prev_state): embed = self.embedding(x) output, state = self.lstm(embed, prev_state) logits = self.fc(output) return logits, state 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))
-
@@ -60,14 +61,16 @@ class Dataset(torch.utils.data.Dataset):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.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]))
-
@@ -77,7 +80,7 @@ dataset = Dataset()model = Model(dataset) dataloader = DataLoader(dataset, batch_size=1024) dataloader = DataLoader(dataset, batch_size=256) criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=0.001)
-
@@ -91,11 +94,11 @@ for epoch in range(1):state_h = state_h.detach() state_c = state_c.detach() loss.backward() optimizer.step() print({ 'epoch': epoch, 'batch': batch, 'loss': loss.item() }) print({'epoch': epoch, 'batch': batch, 'loss': loss.item()}) def predict(text, next_words=100):
-
@@ -103,16 +106,18 @@ def predict(text, next_words=100):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() 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') predict('This is a test')
-