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

Class EncoderRNN

intermediate_source/seq2seq_translation_tutorial.py:333–345  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

331#
332
333class EncoderRNN(nn.Module):
334 def __init__(self, input_size, hidden_size, dropout_p=0.1):
335 super(EncoderRNN, self).__init__()
336 self.hidden_size = hidden_size
337
338 self.embedding = nn.Embedding(input_size, hidden_size)
339 self.gru = nn.GRU(hidden_size, hidden_size, batch_first=True)
340 self.dropout = nn.Dropout(dropout_p)
341
342 def forward(self, input):
343 embedded = self.dropout(self.embedding(input))
344 output, hidden = self.gru(embedded)
345 return output, hidden
346
347######################################################################
348# The Decoder

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected