public class MxModel
extends java.lang.Object
implements ai.djl.Model
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
| Modifier and Type | Method and Description |
|---|---|
void |
cast(ai.djl.ndarray.types.DataType dataType) |
void |
close() |
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 void |
finalize() |
java.net.URL |
getArtifact(java.lang.String artifactName) |
<T> T |
getArtifact(java.lang.String name,
java.util.function.Function<java.io.InputStream,T> function) |
java.io.InputStream |
getArtifactAsStream(java.lang.String name) |
java.lang.String[] |
getArtifactNames() |
ai.djl.nn.Block |
getBlock() |
ai.djl.ndarray.types.DataType |
getDataType() |
java.lang.String |
getName() |
ai.djl.ndarray.NDManager |
getNDManager() |
java.lang.String |
getProperty(java.lang.String key) |
void |
load(java.nio.file.Path modelPath,
java.lang.String modelName,
java.util.Map<java.lang.String,java.lang.String> options)
Loads the MXNet model from a specified location.
|
<I,O> ai.djl.inference.Predictor<I,O> |
newPredictor(ai.djl.translate.Translator<I,O> translator) |
ai.djl.training.Trainer |
newTrainer(ai.djl.training.TrainingConfig trainingConfig) |
void |
save(java.nio.file.Path modelPath,
java.lang.String modelName) |
void |
setBlock(ai.djl.nn.Block block) |
void |
setDataType(ai.djl.ndarray.types.DataType dataType) |
void |
setProperty(java.lang.String key,
java.lang.String value) |
java.lang.String |
toString() |
public void load(java.nio.file.Path modelPath,
java.lang.String modelName,
java.util.Map<java.lang.String,java.lang.String> options)
throws java.io.IOException,
ai.djl.MalformedModelException
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);
load in interface ai.djl.ModelmodelPath - the directory of the modelmodelName - the name/prefix of the modeloptions - load model options, see documentation for the specific enginejava.io.IOException - Exception for file loadingai.djl.MalformedModelExceptionpublic void save(java.nio.file.Path modelPath,
java.lang.String modelName)
throws java.io.IOException
save in interface ai.djl.Modeljava.io.IOExceptionpublic ai.djl.nn.Block getBlock()
getBlock in interface ai.djl.Modelpublic void setBlock(ai.djl.nn.Block block)
setBlock in interface ai.djl.Modelpublic java.lang.String getName()
getName in interface ai.djl.Modelpublic ai.djl.training.Trainer newTrainer(ai.djl.training.TrainingConfig trainingConfig)
newTrainer in interface ai.djl.Modelpublic <I,O> ai.djl.inference.Predictor<I,O> newPredictor(ai.djl.translate.Translator<I,O> translator)
newPredictor in interface ai.djl.Modelpublic void setDataType(ai.djl.ndarray.types.DataType dataType)
setDataType in interface ai.djl.Modelpublic ai.djl.ndarray.types.DataType getDataType()
getDataType in interface ai.djl.Modelpublic void cast(ai.djl.ndarray.types.DataType dataType)
cast in interface ai.djl.Modelpublic ai.djl.util.PairList<java.lang.String,ai.djl.ndarray.types.Shape> describeInput()
describeInput in interface ai.djl.Modelpublic ai.djl.util.PairList<java.lang.String,ai.djl.ndarray.types.Shape> describeOutput()
describeOutput in interface ai.djl.Modelpublic java.lang.String[] getArtifactNames()
getArtifactNames in interface ai.djl.Modelpublic <T> T getArtifact(java.lang.String name,
java.util.function.Function<java.io.InputStream,T> function)
throws java.io.IOException
getArtifact in interface ai.djl.Modeljava.io.IOExceptionpublic java.net.URL getArtifact(java.lang.String artifactName)
throws java.io.IOException
getArtifact in interface ai.djl.Modeljava.io.IOExceptionpublic java.io.InputStream getArtifactAsStream(java.lang.String name)
throws java.io.IOException
getArtifactAsStream in interface ai.djl.Modeljava.io.IOExceptionpublic ai.djl.ndarray.NDManager getNDManager()
getNDManager in interface ai.djl.Modelpublic void setProperty(java.lang.String key,
java.lang.String value)
setProperty in interface ai.djl.Modelpublic java.lang.String getProperty(java.lang.String key)
getProperty in interface ai.djl.Modelpublic void close()
close in interface ai.djl.Modelclose in interface java.lang.AutoCloseableprotected void finalize()
throws java.lang.Throwable
finalize in class java.lang.Objectjava.lang.Throwablepublic java.lang.String toString()
toString in class java.lang.Object