| 37 | |
| 38 | |
| 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) |
| 55 | self.cx.to_device(dev) |
| 56 | self.hx.set_value(0.0) |
| 57 | self.cx.set_value(0.0) |
| 58 | |
| 59 | def initialize(self, inputs): |
| 60 | batchsize = inputs[0].shape[0] |
| 61 | self.hx = tensor.Tensor((batchsize, self.hidden_size)) |
| 62 | self.cx = tensor.Tensor((batchsize, self.hidden_size)) |
| 63 | self.reset_states(inputs[0].device) |
| 64 | |
| 65 | def forward(self, inputs): |
| 66 | x, hx, cx = self.rnn(inputs, (self.hx, self.cx)) |
| 67 | self.hx.copy_data(hx) |
| 68 | self.cx.copy_data(cx) |
| 69 | x = self.cat(x) |
| 70 | x = self.reshape1(x, (-1, self.hidden_size)) |
| 71 | return self.dense(x) |
| 72 | |
| 73 | def train_one_batch(self, x, y): |
| 74 | out = self.forward(x) |
| 75 | y = self.reshape2(y, (-1, 1)) |
| 76 | loss = self.softmax_cross_entropy(out, y) |
| 77 | self.optimizer(loss) |
| 78 | return out, loss |
| 79 | |
| 80 | def get_states(self): |
| 81 | ret = super().get_states() |
| 82 | ret[self.hx.name] = self.hx |
| 83 | ret[self.cx.name] = self.cx |
| 84 | return ret |
| 85 | |
| 86 | def set_states(self, states): |
| 87 | self.hx.copy_from(states[self.hx.name]) |
| 88 | self.hx.copy_from(states[self.hx.name]) |
| 89 | super().set_states(states) |
| 90 | |
| 91 | |
| 92 | class Data(object): |