Class MxModel

java.lang.Object
ai.djl.BaseModel
ai.djl.mxnet.engine.MxModel
All Implemented Interfaces:
ai.djl.Model, AutoCloseable

public class MxModel extends ai.djl.BaseModel
MxModel is the MXNet implementation of Model.

MxModel contains all the methods in Model to load and process a model. In addition, it provides MXNet Specific functionality, such as getSymbol to obtain the Symbolic graph and getParameters to obtain the parameter NDArrays

  • Field Summary

    Fields inherited from class ai.djl.BaseModel

    artifacts, block, dataType, inputData, manager, modelDir, modelName, properties, wasLoaded
  • Method Summary

    Modifier and Type
    Method
    Description
    void
    void
    load(Path modelPath, String prefix, Map<String,?> options)
    Loads the MXNet model from a specified location.
    ai.djl.training.Trainer
    newTrainer(ai.djl.training.TrainingConfig trainingConfig)

    Methods inherited from class ai.djl.BaseModel

    describeInput, describeOutput, finalize, getArtifact, getArtifact, getArtifactAsStream, getBlock, getDataType, getModelPath, getName, getNDManager, getProperties, getProperty, load, newPredictor, paramPathResolver, readParameters, save, setBlock, setDataType, setModelDir, setProperty, toString

    Methods inherited from class java.lang.Object

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

    Methods inherited from interface ai.djl.Model

    cast, getProperty, load, load, load, newPredictor, quantize
  • Method Details

    • load

      public void load(Path modelPath, String prefix, Map<String,?> options) throws IOException, ai.djl.MalformedModelException
      Loads the MXNet model from a specified location.

      MXNet engine looks for {MODEL_NAME}-symbol.json and {MODEL_NAME}-{EPOCH}.params files in the specified directory. By default, MXNet engine will pick up the latest epoch of the parameter file. However, users can explicitly specify an epoch to be loaded:

       Map<String, String> options = new HashMap<>()
       options.put("epoch", "3");
       model.load(modelPath, "squeezenet", options);
       
      Parameters:
      modelPath - the directory of the model
      prefix - the model file name or path prefix
      options - load model options, see documentation for the specific engine
      Throws:
      IOException - Exception for file loading
      ai.djl.MalformedModelException
    • newTrainer

      public ai.djl.training.Trainer newTrainer(ai.djl.training.TrainingConfig trainingConfig)
      Specified by:
      newTrainer in interface ai.djl.Model
      Overrides:
      newTrainer in class ai.djl.BaseModel
    • getArtifactNames

      public String[] getArtifactNames()
      Specified by:
      getArtifactNames in interface ai.djl.Model
      Overrides:
      getArtifactNames in class ai.djl.BaseModel
    • close

      public void close()
      Specified by:
      close in interface AutoCloseable
      Specified by:
      close in interface ai.djl.Model
      Overrides:
      close in class ai.djl.BaseModel