Class MxSymbolBlock

java.lang.Object
ai.djl.nn.AbstractBaseBlock
ai.djl.nn.AbstractSymbolBlock
ai.djl.mxnet.engine.MxSymbolBlock
All Implemented Interfaces:
ai.djl.nn.Block, ai.djl.nn.SymbolBlock

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).

  • Field Summary

    Fields inherited from class ai.djl.nn.AbstractBaseBlock

    inputNames, inputShapes, outputDataTypes, version
  • Constructor Summary

    Constructors
    Constructor
    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.
  • Method Summary

    Modifier and Type
    Method
    Description
    ai.djl.util.PairList<String,ai.djl.ndarray.types.Shape>
    ai.djl.util.PairList<String,ai.djl.ndarray.types.Shape>
    protected ai.djl.ndarray.NDList
    forwardInternal(ai.djl.training.ParameterStore parameterStore, ai.djl.ndarray.NDList inputs, boolean training, ai.djl.util.PairList<String,Object> params)
    List<ai.djl.nn.Parameter>
    Returns the list of inputs and parameter NDArrays.
    ai.djl.nn.ParameterList
    Returns the layers' name.
    ai.djl.ndarray.types.Shape[]
    getOutputShapes(ai.djl.ndarray.types.Shape[] inputShapes)
    Returns the Symbolic graph from the model.
    void
    loadParameters(ai.djl.ndarray.NDManager manager, DataInputStream is)
    void
    optimizeFor(String optimization)
    Applies Optimization algorithm for the model.
    void
    void
    void
    setInputNames(List<String> inputNames)
    Sets the names of the input data.

    Methods inherited from class ai.djl.nn.AbstractSymbolBlock

    getChildren

    Methods inherited from class ai.djl.nn.AbstractBaseBlock

    beforeInitialize, cast, clear, forward, forward, forwardInternal, getInputShapes, getOutputDataTypes, getParameters, initialize, initializeChildBlocks, isInitialized, loadMetadata, prepare, readInputShapes, saveInputShapes, saveMetadata, setInitializer, setInitializer, setInitializer, toString

    Methods inherited from class java.lang.Object

    clone, equals, finalize, getClass, hashCode, notify, notifyAll, wait, wait, wait

    Methods inherited from interface ai.djl.nn.Block

    cast, clear, forward, forward, forward, freezeParameters, freezeParameters, getInputShapes, getOutputDataTypes, getOutputShapes, getParameters, initialize, isInitialized, setInitializer, setInitializer, setInitializer
  • Constructor Details

    • MxSymbolBlock

      public MxSymbolBlock(ai.djl.ndarray.NDManager manager, Symbol symbol)
      Constructs a MxSymbolBlock for a Symbol.

      You can create a MxSymbolBlock using Model.load(java.nio.file.Path, String).

      Parameters:
      manager - the manager to use for the block
      symbol - the symbol containing the block's symbolic graph
    • MxSymbolBlock

      public MxSymbolBlock(ai.djl.ndarray.NDManager manager)
      Constructs an empty MxSymbolBlock.
      Parameters:
      manager - the manager to use for the block
  • Method Details

    • setInputNames

      public void setInputNames(List<String> inputNames)
      Sets the names of the input data.
      Parameters:
      inputNames - the names of the input data
    • getAllParameters

      public List<ai.djl.nn.Parameter> getAllParameters()
      Returns the list of inputs and parameter NDArrays.
      Returns:
      the list of inputs and parameter NDArrays
    • getLayerNames

      public List<String> getLayerNames()
      Returns the layers' name.
      Returns:
      a List of String containing the layers' name
    • getSymbol

      public Symbol getSymbol()
      Returns the Symbolic graph from the model.
      Returns:
      a Symbol object
    • optimizeFor

      public void optimizeFor(String optimization)
      Applies Optimization algorithm for the model.
      Parameters:
      optimization - the name of the optimization
    • describeInput

      public ai.djl.util.PairList<String,ai.djl.ndarray.types.Shape> describeInput()
      Specified by:
      describeInput in interface ai.djl.nn.Block
      Overrides:
      describeInput in class ai.djl.nn.AbstractBaseBlock
    • getDirectParameters

      public ai.djl.nn.ParameterList getDirectParameters()
    • describeOutput

      public ai.djl.util.PairList<String,ai.djl.ndarray.types.Shape> describeOutput()
    • forwardInternal

      protected ai.djl.ndarray.NDList forwardInternal(ai.djl.training.ParameterStore parameterStore, ai.djl.ndarray.NDList inputs, boolean training, ai.djl.util.PairList<String,Object> params)
      Specified by:
      forwardInternal in class ai.djl.nn.AbstractBaseBlock
    • getOutputShapes

      public ai.djl.ndarray.types.Shape[] getOutputShapes(ai.djl.ndarray.types.Shape[] inputShapes)
      Specified by:
      getOutputShapes in interface ai.djl.nn.Block
      Overrides:
      getOutputShapes in class ai.djl.nn.AbstractSymbolBlock
    • removeLastBlock

      public void removeLastBlock()
    • saveParameters

      public void saveParameters(DataOutputStream os) throws IOException
      Specified by:
      saveParameters in interface ai.djl.nn.Block
      Overrides:
      saveParameters in class ai.djl.nn.AbstractBaseBlock
      Throws:
      IOException
    • loadParameters

      public void loadParameters(ai.djl.ndarray.NDManager manager, DataInputStream is) throws IOException, ai.djl.MalformedModelException
      Specified by:
      loadParameters in interface ai.djl.nn.Block
      Overrides:
      loadParameters in class ai.djl.nn.AbstractBaseBlock
      Throws:
      IOException
      ai.djl.MalformedModelException