Changes
1 changed files (+37/-0)
-
-
@@ -2,6 +2,9 @@import re import psycopg2 import torch import torch.nn as nn # Fetch messages from database since it's way faster than using the API conn = psycopg2.connect(dbname="mastodon_production")
-
@@ -10,6 +13,40 @@ cur.execute('SELECT * FROM statuses')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]) # https://closeheat.com/blog/pytorch-lstm-text-generation-tutorial class Model(nn.Module): def __init__(self, dataset): super(Model, self).__init__() self.lstm_size = 128 self.embedding_dim = 128 self.num_layers = 3 n_vocab = len(dataset.uniq_words) self.embedding = nn.Embedding( num_embeddings=n_vocab, embedding_dim=self.embedding_dim ) self.lstm = nn.LSTM( input_size=self.lstm_size, hidden_size=self.lstm_size, num_layers=self.num_layers, 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))
-