MCPcopy Create free account
hub / github.com/apache/singa / CharRNN

Class CharRNN

examples/rnn/char_rnn.py:39–89  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

37
38
39class 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
92class Data(object):

Callers 1

trainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected