public class MxSymbolBlock
extends ai.djl.nn.AbstractSymbolBlock
MxSymbolBlock is the MXNet implementation of SymbolBlock.
You can create a MxSymbolBlock using Model.load(java.nio.file.Path,
String).
| Constructor and Description |
|---|
MxSymbolBlock(ai.djl.ndarray.NDManager manager)
Constructs an empty
MxSymbolBlock. |
MxSymbolBlock(ai.djl.ndarray.NDManager manager,
Symbol symbol)
Constructs a
MxSymbolBlock for a Symbol. |
| Modifier and Type | Method and Description |
|---|---|
ai.djl.util.PairList<java.lang.String,ai.djl.ndarray.types.Shape> |
describeInput() |
ai.djl.util.PairList<java.lang.String,ai.djl.ndarray.types.Shape> |
describeOutput() |
protected ai.djl.ndarray.NDList |
forwardInternal(ai.djl.training.ParameterStore parameterStore,
ai.djl.ndarray.NDList inputs,
boolean training,
ai.djl.util.PairList<java.lang.String,java.lang.Object> params) |
java.util.List<ai.djl.nn.Parameter> |
getAllParameters()
Returns the list of inputs and parameter NDArrays.
|
java.util.List<java.lang.String> |
getLayerNames()
Returns the layers' name.
|
ai.djl.ndarray.types.Shape[] |
getOutputShapes(ai.djl.ndarray.types.Shape[] inputShapes) |
Symbol |
getSymbol()
Returns the Symbolic graph from the model.
|
void |
loadParameters(ai.djl.ndarray.NDManager manager,
java.io.DataInputStream is) |
void |
optimizeFor(java.lang.String optimization)
Applies Optimization algorithm for the model.
|
void |
removeLastBlock() |
void |
saveParameters(java.io.DataOutputStream os) |
void |
setInputNames(java.util.List<java.lang.String> inputNames)
Sets the names of the input data.
|
addChildBlock, addParameter, beforeInitialize, cast, clear, forward, forward, forwardInternal, getChildren, getDirectParameters, getParameters, initialize, initializeChildBlocks, isInitialized, loadMetadata, prepare, readInputShapes, saveInputShapes, saveMetadata, setInitializer, setInitializer, setInitializer, toStringpublic MxSymbolBlock(ai.djl.ndarray.NDManager manager,
Symbol symbol)
MxSymbolBlock for a Symbol.
You can create a MxSymbolBlock using Model.load(java.nio.file.Path,
String).
manager - the manager to use for the blocksymbol - the symbol containing the block's symbolic graphpublic MxSymbolBlock(ai.djl.ndarray.NDManager manager)
MxSymbolBlock.manager - the manager to use for the blockpublic void setInputNames(java.util.List<java.lang.String> inputNames)
inputNames - the names of the input datapublic java.util.List<ai.djl.nn.Parameter> getAllParameters()
public java.util.List<java.lang.String> getLayerNames()
public Symbol getSymbol()
Symbol objectpublic void optimizeFor(java.lang.String optimization)
optimization - the name of the optimizationpublic ai.djl.util.PairList<java.lang.String,ai.djl.ndarray.types.Shape> describeInput()
describeInput in interface ai.djl.nn.BlockdescribeInput in class ai.djl.nn.AbstractBlockpublic ai.djl.util.PairList<java.lang.String,ai.djl.ndarray.types.Shape> describeOutput()
protected ai.djl.ndarray.NDList forwardInternal(ai.djl.training.ParameterStore parameterStore,
ai.djl.ndarray.NDList inputs,
boolean training,
ai.djl.util.PairList<java.lang.String,java.lang.Object> params)
forwardInternal in class ai.djl.nn.AbstractBlockpublic ai.djl.ndarray.types.Shape[] getOutputShapes(ai.djl.ndarray.types.Shape[] inputShapes)
getOutputShapes in interface ai.djl.nn.BlockgetOutputShapes in class ai.djl.nn.AbstractSymbolBlockpublic void removeLastBlock()
public void saveParameters(java.io.DataOutputStream os)
throws java.io.IOException
saveParameters in interface ai.djl.nn.BlocksaveParameters in class ai.djl.nn.AbstractBlockjava.io.IOExceptionpublic void loadParameters(ai.djl.ndarray.NDManager manager,
java.io.DataInputStream is)
throws java.io.IOException,
ai.djl.MalformedModelException
loadParameters in interface ai.djl.nn.BlockloadParameters in class ai.djl.nn.AbstractBlockjava.io.IOExceptionai.djl.MalformedModelException