| 39 | class CharRNN(model.Model): |
| 40 | |
| 41 | def __init__(self, vocab_size, hidden_size=32): |
| 42 | super(CharRNN, self).__init__() |
| 43 | self.rnn = layer.LSTM(vocab_size, hidden_size) |
| 44 | self.cat = layer.Cat() |
| 45 | self.reshape1 = layer.Reshape() |
| 46 | self.dense = layer.Linear(hidden_size, vocab_size) |
| 47 | self.reshape2 = layer.Reshape() |
| 48 | self.softmax_cross_entropy = layer.SoftMaxCrossEntropy() |
| 49 | self.optimizer = opt.SGD(0.01) |
| 50 | self.hidden_size = hidden_size |
| 51 | self.vocab_size = vocab_size |
| 52 | |
| 53 | def reset_states(self, dev): |
| 54 | self.hx.to_device(dev) |