MCPcopy Create free account
hub / github.com/pytorch/tutorials / RNN

Class RNN

intermediate_source/char_rnn_generation_tutorial.py:157–179  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

155import torch.nn as nn
156
157class RNN(nn.Module):
158 def __init__(self, input_size, hidden_size, output_size):
159 super(RNN, self).__init__()
160 self.hidden_size = hidden_size
161
162 self.i2h = nn.Linear(n_categories + input_size + hidden_size, hidden_size)
163 self.i2o = nn.Linear(n_categories + input_size + hidden_size, output_size)
164 self.o2o = nn.Linear(hidden_size + output_size, output_size)
165 self.dropout = nn.Dropout(0.1)
166 self.softmax = nn.LogSoftmax(dim=1)
167
168 def forward(self, category, input, hidden):
169 input_combined = torch.cat((category, input, hidden), 1)
170 hidden = self.i2h(input_combined)
171 output = self.i2o(input_combined)
172 output_combined = torch.cat((hidden, output), 1)
173 output = self.o2o(output_combined)
174 output = self.dropout(output)
175 output = self.softmax(output)
176 return output, hidden
177
178 def initHidden(self):
179 return torch.zeros(1, self.hidden_size)
180
181
182######################################################################

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected