update recurrent layers to store hidden state sequences on class instance

This commit is contained in:
Leon Chen
2016-10-13 15:05:56 -04:00
parent e75e3c0148
commit 6ea62ad12e
3 changed files with 12 additions and 11 deletions
+4 -3
View File
@@ -88,7 +88,8 @@ export default class GRU extends Layer {
let tempXH = new Tensor([], [dimHiddenState])
let tempHH = new Tensor([], [dimHiddenState])
let previousHiddenState = new Tensor([], [dimHiddenState])
let hiddenStateSequence = new Tensor([], [x.tensor.shape[0], dimHiddenState])
this.hiddenStateSequence = new Tensor([], [x.tensor.shape[0], dimHiddenState])
const _clearTemp = () => {
const tempTensors = [tempXZ, tempHZ, tempXR, tempHR, tempXH, tempHH]
@@ -128,12 +129,12 @@ export default class GRU extends Layer {
_step()
if (this.returnSequences) {
ops.assign(hiddenStateSequence.tensor.pick(i, null), currentHiddenState.tensor)
ops.assign(this.hiddenStateSequence.tensor.pick(i, null), currentHiddenState.tensor)
}
}
if (this.returnSequences) {
x.tensor = hiddenStateSequence.tensor
x.tensor = this.hiddenStateSequence.tensor
} else {
x.tensor = currentHiddenState.tensor
}
+4 -5
View File
@@ -98,7 +98,8 @@ export default class LSTM extends Layer {
? this.currentHiddenState
: new Tensor([], [dimCandidate])
let previousHiddenState = new Tensor([], [dimCandidate])
let hiddenStateSequence = new Tensor([], [x.tensor.shape[0], dimCandidate])
this.hiddenStateSequence = new Tensor([], [x.tensor.shape[0], dimCandidate])
const _clearTemp = () => {
const tempTensors = [tempXI, tempHI, tempXF, tempHF, tempXO, tempHO, tempXC, tempHC]
@@ -147,13 +148,11 @@ export default class LSTM extends Layer {
_clearTemp()
_step()
if (this.returnSequences) {
ops.assign(hiddenStateSequence.tensor.pick(i, null), currentHiddenState.tensor)
}
ops.assign(this.hiddenStateSequence.tensor.pick(i, null), currentHiddenState.tensor)
}
if (this.returnSequences) {
x.tensor = hiddenStateSequence.tensor
x.tensor = this.hiddenStateSequence.tensor
} else {
x.tensor = currentHiddenState.tensor
}
+4 -3
View File
@@ -66,7 +66,8 @@ export default class SimpleRNN extends Layer {
let tempXH = new Tensor([], [dimHiddenState])
let tempHH = new Tensor([], [dimHiddenState])
let previousHiddenState = new Tensor([], [dimHiddenState])
let hiddenStateSequence = new Tensor([], [x.tensor.shape[0], dimHiddenState])
this.hiddenStateSequence = new Tensor([], [x.tensor.shape[0], dimHiddenState])
const _clearTemp = () => {
const tempTensors = [tempXH, tempHH]
@@ -89,12 +90,12 @@ export default class SimpleRNN extends Layer {
_step()
if (this.returnSequences) {
ops.assign(hiddenStateSequence.tensor.pick(i, null), currentHiddenState.tensor)
ops.assign(this.hiddenStateSequence.tensor.pick(i, null), currentHiddenState.tensor)
}
}
if (this.returnSequences) {
x.tensor = hiddenStateSequence.tensor
x.tensor = this.hiddenStateSequence.tensor
} else {
x.tensor = currentHiddenState.tensor
}