improved context update for JayRnn
This commit is contained in:
@@ -81,12 +81,11 @@ class RnnModel(Model):
|
|||||||
unit = Entity((WIN*vocab_size() + H_SIZE, H_SIZE), EntityParams(do_gaussian_visible=False, do_gaussian_hidden=False), TrainingParams(learning_rate=0.001, momentum=0.9, num_epochs=1000), enable_training=True)
|
unit = Entity((WIN*vocab_size() + H_SIZE, H_SIZE), EntityParams(do_gaussian_visible=False, do_gaussian_hidden=False), TrainingParams(learning_rate=0.001, momentum=0.9, num_epochs=1000), enable_training=True)
|
||||||
self.units.append(unit)
|
self.units.append(unit)
|
||||||
|
|
||||||
def train(self, batch: Mat, status: Status = None):
|
def train(self, vc: Mat, status: Status = None):
|
||||||
c = np.zeros([len(batch), H_SIZE])
|
|
||||||
for i, unit in enumerate(self.units):
|
for i, unit in enumerate(self.units):
|
||||||
vc = concat(batch, c)
|
|
||||||
train(unit, vc, status)
|
train(unit, vc, status)
|
||||||
c = unit.forward(vc)
|
# Update context portion of vc
|
||||||
|
vc[:, WIN*vocab_size():] = unit.forward(vc)
|
||||||
|
|
||||||
def forward_step(self, v_curr: Mat):
|
def forward_step(self, v_curr: Mat):
|
||||||
c = np.zeros([1, H_SIZE])
|
c = np.zeros([1, H_SIZE])
|
||||||
@@ -108,6 +107,10 @@ if __name__ == "__main__":
|
|||||||
model = RnnModel(name='JayRnn', work_dir='results')
|
model = RnnModel(name='JayRnn', work_dir='results')
|
||||||
model.init(0.01)
|
model.init(0.01)
|
||||||
model.load()
|
model.load()
|
||||||
model.train(batch, status=Status())
|
|
||||||
|
# vc contains vis + context
|
||||||
|
# context will be updated after training
|
||||||
|
vc = concat(batch, np.zeros([len(batch), H_SIZE]))
|
||||||
|
model.train(vc, status=Status())
|
||||||
model.save()
|
model.save()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user