mirror of
https://github.com/wassname/keras-js.git
synced 2026-09-10 12:15:12 +08:00
update recurrent layers to store hidden state sequences on class instance
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user