public abstract class RecurrentBlock extends ParameterBlock
RecurrentBlock is an abstract implementation of recurrent neural networks.
Recurrent neural networks are neural networks with hidden states. They are very popular for natural language processing tasks, and other tasks which involve sequential data.
This [article](http://karpathy.github.io/2015/05/21/rnn-effectiveness/) written by Andrej Karpathy provides a detailed explanation of recurrent neural networks.
Currently, vanilla RNN, LSTM and GRU are implemented, with both multi-layer and bidirectional support.
| Modifier and Type | Class and Description |
|---|---|
static class |
RecurrentBlock.BaseBuilder<T extends RecurrentBlock.BaseBuilder>
The Builder to construct a
RecurrentBlock type of Block. |
| Modifier and Type | Field and Description |
|---|---|
protected NDArray |
beginState |
protected float |
dropRate |
protected int |
gates |
protected java.lang.String |
mode |
protected int |
numDirections |
protected int |
numStackedLayers |
protected java.util.List<Parameter> |
parameters |
protected boolean |
stateOutputs |
protected long |
stateSize |
protected boolean |
useSequenceLength |
inputNames, inputShapes| Constructor and Description |
|---|
RecurrentBlock(RecurrentBlock.BaseBuilder<?> builder)
Creates a
RecurrentBlock object. |
| Modifier and Type | Method and Description |
|---|---|
void |
beforeInitialize(Shape[] inputs)
Performs any action necessary before initialization.
|
NDList |
forward(ParameterStore parameterStore,
NDList inputs,
boolean training,
ai.djl.util.PairList<java.lang.String,java.lang.Object> params)
Applies the operating function of the block once.
|
java.util.List<Parameter> |
getDirectParameters()
Returns a list of all the direct parameters of the block.
|
Shape[] |
getOutputShapes(NDManager manager,
Shape[] inputs)
Returns the expected output shapes of the block for the specified input shapes.
|
Shape |
getParameterShape(java.lang.String name,
Shape[] inputShapes)
Returns the shape of the specified direct parameter of this block given the shapes of the
input to the block.
|
protected boolean |
isBidirectional() |
void |
loadParameters(NDManager manager,
java.io.DataInputStream is)
Loads the parameters from the given input stream.
|
protected NDList |
opInputs(ParameterStore parameterStore,
NDList inputs) |
protected void |
resetBeginStates() |
void |
saveParameters(java.io.DataOutputStream os)
Writes the parameters of the block to the given outputStream.
|
void |
setBeginStates(NDList beginStates)
Sets the initial
NDArray value for the hidden states. |
void |
setStateOutputs(boolean stateOutputs)
Sets the parameter that indicates whether the output must include the hidden states.
|
protected NDList |
updateInputLayoutToTNC(NDList inputs) |
protected void |
validateInputSize(NDList inputs) |
getChildren, initialize, toStringcast, clear, describeInput, getParameters, isInitialized, readInputShapes, saveInputShapes, setInitializer, setInitializerclone, equals, finalize, getClass, hashCode, notify, notifyAll, wait, wait, waitforward, validateLayoutprotected long stateSize
protected float dropRate
protected int numStackedLayers
protected java.lang.String mode
protected boolean useSequenceLength
protected int numDirections
protected int gates
protected boolean stateOutputs
protected NDArray beginState
protected java.util.List<Parameter> parameters
public RecurrentBlock(RecurrentBlock.BaseBuilder<?> builder)
RecurrentBlock object.builder - the Builder that has the necessary configurationsprotected void validateInputSize(NDList inputs)
public void setStateOutputs(boolean stateOutputs)
stateOutputs - whether the output must include the hidden states.public NDList forward(ParameterStore parameterStore, NDList inputs, boolean training, ai.djl.util.PairList<java.lang.String,java.lang.Object> params)
parameterStore - the parameter storeinputs - the input NDListtraining - true for a training forward passparams - optional parameterspublic void setBeginStates(NDList beginStates)
NDArray value for the hidden states.beginStates - the NDArray value for the hidden statesprotected void resetBeginStates()
public Shape[] getOutputShapes(NDManager manager, Shape[] inputs)
manager - an NDManagerinputs - the shapes of the inputspublic java.util.List<Parameter> getDirectParameters()
Parameterpublic void beforeInitialize(Shape[] inputs)
beforeInitialize in class AbstractBlockinputs - the expected shapes of the inputpublic Shape getParameterShape(java.lang.String name, Shape[] inputShapes)
name - the name of the parameterinputShapes - the shapes of the input to the blockpublic void saveParameters(java.io.DataOutputStream os)
throws java.io.IOException
os - the outputstream to save the parameters tojava.io.IOException - if an I/O error occurspublic void loadParameters(NDManager manager, java.io.DataInputStream is) throws java.io.IOException, MalformedModelException
manager - an NDManager to create the parameter arraysis - the inputstream that stream the parameter valuesjava.io.IOException - if an I/O error occursMalformedModelException - if the model file is corrupted or unsupportedprotected boolean isBidirectional()
protected NDList opInputs(ParameterStore parameterStore, NDList inputs)