diff --git a/src/layers/recurrent/GRU.js b/src/layers/recurrent/GRU.js index 9a8657c..99079ba 100644 --- a/src/layers/recurrent/GRU.js +++ b/src/layers/recurrent/GRU.js @@ -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 } diff --git a/src/layers/recurrent/LSTM.js b/src/layers/recurrent/LSTM.js index a810352..09218b4 100644 --- a/src/layers/recurrent/LSTM.js +++ b/src/layers/recurrent/LSTM.js @@ -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 } diff --git a/src/layers/recurrent/SimpleRNN.js b/src/layers/recurrent/SimpleRNN.js index 65c69a7..61dc9f7 100644 --- a/src/layers/recurrent/SimpleRNN.js +++ b/src/layers/recurrent/SimpleRNN.js @@ -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 }