| 155 | import torch.nn as nn |
| 156 | |
| 157 | class 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 | ###################################################################### |
no outgoing calls
no test coverage detected