Package ai.djl.mxnet.jna
Class JnaUtils
java.lang.Object
ai.djl.mxnet.jna.JnaUtils
A class containing utilities to interact with the MXNet Engine's Java Native Access (JNA) layer.
-
Nested Class Summary
Nested ClassesModifier and TypeClassDescriptionstatic enumAn enum that enumerates the statuses of numpy mode. -
Field Summary
Fields -
Method Summary
Modifier and TypeMethodDescriptionstatic voidautogradBackward(ai.djl.ndarray.NDList array, int retainGraph) static voidautogradBackwardExecute(int numOutput, ai.djl.ndarray.NDList array, ai.djl.ndarray.NDArray outgrad, int numVariables, com.sun.jna.Pointer varHandles, int retainGraph, int createGraph, int isTrain, com.sun.jna.Pointer gradHandles, com.sun.jna.Pointer gradSparseFormat) static com.sun.jna.PointerautogradGetSymbol(ai.djl.ndarray.NDArray array) static booleanstatic booleanstatic voidautogradMarkVariables(int numVar, com.sun.jna.Pointer varHandles, IntBuffer reqsArray, com.sun.jna.Pointer gradHandles) static booleanautogradSetIsRecording(boolean isRecording) static booleanautogradSetTraining(boolean isTraining) static MxNDArray[]cachedOpInvoke(MxNDManager manager, com.sun.jna.Pointer cachedOpHandle, MxNDArray[] inputs) static voidcheckCall(int ret) static CachedOpcreateCachedOp(MxSymbolBlock block, MxNDManager manager, boolean training) Creates cached op flags.static com.sun.jna.PointercreateNdArray(ai.djl.Device device, ai.djl.ndarray.types.Shape shape, ai.djl.ndarray.types.DataType dtype, int size, boolean delayedAlloc) static com.sun.jna.PointercreateSparseNdArray(ai.djl.ndarray.types.SparseFormat fmt, ai.djl.Device device, ai.djl.ndarray.types.Shape shape, ai.djl.ndarray.types.DataType dtype, ai.djl.ndarray.types.DataType[] auxDTypes, ai.djl.ndarray.types.Shape[] auxShapes, boolean delayedAlloc) static com.sun.jna.PointercreateSymbolFromFile(String path) static com.sun.jna.PointercreateSymbolFromString(String json) static com.sun.jna.PointerdetachGradient(com.sun.jna.Pointer handle) static voidfreeCachedOp(com.sun.jna.Pointer handle) static voidfreeNdArray(com.sun.jna.Pointer ndArray) static voidfreeSymbol(com.sun.jna.Pointer symbol) static ai.djl.ndarray.types.DataTypegetDataType(com.sun.jna.Pointer ndArray) static ai.djl.DevicegetDevice(com.sun.jna.Pointer ndArray) static intstatic long[]getGpuMemory(ai.djl.Device device) static com.sun.jna.PointergetGradient(com.sun.jna.Pointer handle) static Map<String,FunctionInfo> static ai.djl.ndarray.types.ShapegetShape(com.sun.jna.Pointer ndArray) static ai.djl.ndarray.types.SparseFormatgetStorageType(com.sun.jna.Pointer ndArray) static com.sun.jna.PointergetSymbolInternals(com.sun.jna.Pointer symbol) static com.sun.jna.PointergetSymbolOutput(com.sun.jna.Pointer symbol, int index) static StringgetSymbolString(com.sun.jna.Pointer symbol) static intstatic ai.djl.util.PairList<com.sun.jna.Pointer,ai.djl.ndarray.types.SparseFormat> imperativeInvoke(com.sun.jna.Pointer function, ai.djl.ndarray.NDArray[] src, ai.djl.ndarray.NDArray[] dest, ai.djl.util.PairList<String, ?> params) inferShape(Symbol symbol, ai.djl.util.PairList<String, ai.djl.ndarray.types.Shape> args) static intstatic String[]listSymbolArguments(com.sun.jna.Pointer symbol) static String[]listSymbolAuxiliaryStates(com.sun.jna.Pointer symbol) static String[]listSymbolNames(com.sun.jna.Pointer symbol) static String[]listSymbolOutputs(com.sun.jna.Pointer symbol) static voidstatic ai.djl.ndarray.NDListloadNdArray(MxNDManager manager, Path path, ai.djl.Device device) static voidndArraySyncCopyFromNdArray(MxNDArray dest, MxNDArray src, int location) static FunctionInfostatic com.sun.jna.PointeroptimizeFor(Symbol current, String backend, ai.djl.Device device) static voidparameterStoreClose(com.sun.jna.Pointer handle) static com.sun.jna.PointerparameterStoreCreate(String type) static voidparameterStoreInit(com.sun.jna.Pointer handle, int num, String[] keys, ai.djl.ndarray.NDList vals) static voidparameterStorePull(com.sun.jna.Pointer handle, int num, int[] keys, ai.djl.ndarray.NDList vals, int priority) static voidparameterStorePull(com.sun.jna.Pointer handle, int num, String[] keys, ai.djl.ndarray.NDList vals, int priority) static voidparameterStorePush(com.sun.jna.Pointer handle, int num, String[] keys, ai.djl.ndarray.NDList vals, int priority) static voidparameterStorePushPull(com.sun.jna.Pointer handle, int inputNum, String[] inputKeys, int outputNum, String[] outputKey, ai.djl.ndarray.NDList inputs, ai.djl.ndarray.NDList outputs, int priority) static voidparameterStoreSetUpdater(com.sun.jna.Pointer handle, MxnetLibrary.MXKVStoreUpdater updater, MxnetLibrary.MXKVStoreStrUpdater stringUpdater, com.sun.jna.Pointer updaterHandle) static voidparameterStoreSetUpdater(com.sun.jna.Pointer handle, MxnetLibrary.MXKVStoreUpdater updater, com.sun.jna.Pointer updaterHandle) static intrandomSeed(int seed) static voidstatic voidsyncCopyFromCPU(com.sun.jna.Pointer ndArray, Buffer data, int len) static voidsyncCopyToCPU(com.sun.jna.Pointer ndArray, com.sun.jna.Pointer data, int len) static voidwaitAll()static voidwaitToRead(com.sun.jna.Pointer ndArray) static voidwaitToWrite(com.sun.jna.Pointer ndArray)
-
Field Details
-
REFS
-
-
Method Details
-
getVersion
public static int getVersion() -
getAllOpNames
-
getNdArrayFunctions
-
op
-
getGpuCount
public static int getGpuCount() -
getGpuMemory
public static long[] getGpuMemory(ai.djl.Device device) -
getFeatures
-
randomSeed
public static int randomSeed(int seed) -
createNdArray
public static com.sun.jna.Pointer createNdArray(ai.djl.Device device, ai.djl.ndarray.types.Shape shape, ai.djl.ndarray.types.DataType dtype, int size, boolean delayedAlloc) -
createSparseNdArray
public static com.sun.jna.Pointer createSparseNdArray(ai.djl.ndarray.types.SparseFormat fmt, ai.djl.Device device, ai.djl.ndarray.types.Shape shape, ai.djl.ndarray.types.DataType dtype, ai.djl.ndarray.types.DataType[] auxDTypes, ai.djl.ndarray.types.Shape[] auxShapes, boolean delayedAlloc) -
ndArraySyncCopyFromNdArray
-
loadNdArray
public static ai.djl.ndarray.NDList loadNdArray(MxNDManager manager, Path path, ai.djl.Device device) -
freeNdArray
public static void freeNdArray(com.sun.jna.Pointer ndArray) -
waitToRead
public static void waitToRead(com.sun.jna.Pointer ndArray) -
waitToWrite
public static void waitToWrite(com.sun.jna.Pointer ndArray) -
waitAll
public static void waitAll() -
syncCopyToCPU
public static void syncCopyToCPU(com.sun.jna.Pointer ndArray, com.sun.jna.Pointer data, int len) -
syncCopyFromCPU
-
imperativeInvoke
public static ai.djl.util.PairList<com.sun.jna.Pointer,ai.djl.ndarray.types.SparseFormat> imperativeInvoke(com.sun.jna.Pointer function, ai.djl.ndarray.NDArray[] src, ai.djl.ndarray.NDArray[] dest, ai.djl.util.PairList<String, ?> params) -
getStorageType
public static ai.djl.ndarray.types.SparseFormat getStorageType(com.sun.jna.Pointer ndArray) -
getDevice
public static ai.djl.Device getDevice(com.sun.jna.Pointer ndArray) -
getShape
public static ai.djl.ndarray.types.Shape getShape(com.sun.jna.Pointer ndArray) -
getDataType
public static ai.djl.ndarray.types.DataType getDataType(com.sun.jna.Pointer ndArray) -
autogradSetIsRecording
public static boolean autogradSetIsRecording(boolean isRecording) -
autogradSetTraining
public static boolean autogradSetTraining(boolean isTraining) -
autogradIsRecording
public static boolean autogradIsRecording() -
autogradIsTraining
public static boolean autogradIsTraining() -
autogradMarkVariables
public static void autogradMarkVariables(int numVar, com.sun.jna.Pointer varHandles, IntBuffer reqsArray, com.sun.jna.Pointer gradHandles) -
autogradBackward
public static void autogradBackward(ai.djl.ndarray.NDList array, int retainGraph) -
autogradBackwardExecute
public static void autogradBackwardExecute(int numOutput, ai.djl.ndarray.NDList array, ai.djl.ndarray.NDArray outgrad, int numVariables, com.sun.jna.Pointer varHandles, int retainGraph, int createGraph, int isTrain, com.sun.jna.Pointer gradHandles, com.sun.jna.Pointer gradSparseFormat) -
autogradGetSymbol
public static com.sun.jna.Pointer autogradGetSymbol(ai.djl.ndarray.NDArray array) -
isNumpyMode
public static int isNumpyMode() -
setNumpyMode
-
getGradient
public static com.sun.jna.Pointer getGradient(com.sun.jna.Pointer handle) -
parameterStoreCreate
-
parameterStoreClose
public static void parameterStoreClose(com.sun.jna.Pointer handle) -
parameterStoreInit
public static void parameterStoreInit(com.sun.jna.Pointer handle, int num, String[] keys, ai.djl.ndarray.NDList vals) -
parameterStorePush
public static void parameterStorePush(com.sun.jna.Pointer handle, int num, String[] keys, ai.djl.ndarray.NDList vals, int priority) -
parameterStorePull
public static void parameterStorePull(com.sun.jna.Pointer handle, int num, int[] keys, ai.djl.ndarray.NDList vals, int priority) -
parameterStorePull
public static void parameterStorePull(com.sun.jna.Pointer handle, int num, String[] keys, ai.djl.ndarray.NDList vals, int priority) -
parameterStorePushPull
-
parameterStoreSetUpdater
public static void parameterStoreSetUpdater(com.sun.jna.Pointer handle, MxnetLibrary.MXKVStoreUpdater updater, MxnetLibrary.MXKVStoreStrUpdater stringUpdater, com.sun.jna.Pointer updaterHandle) -
parameterStoreSetUpdater
public static void parameterStoreSetUpdater(com.sun.jna.Pointer handle, MxnetLibrary.MXKVStoreUpdater updater, com.sun.jna.Pointer updaterHandle) -
detachGradient
public static com.sun.jna.Pointer detachGradient(com.sun.jna.Pointer handle) -
getSymbolOutput
public static com.sun.jna.Pointer getSymbolOutput(com.sun.jna.Pointer symbol, int index) -
listSymbolOutputs
-
freeSymbol
public static void freeSymbol(com.sun.jna.Pointer symbol) -
listSymbolNames
-
listSymbolArguments
-
listSymbolAuxiliaryStates
-
getSymbolInternals
public static com.sun.jna.Pointer getSymbolInternals(com.sun.jna.Pointer symbol) -
createSymbolFromFile
-
createSymbolFromString
-
getSymbolString
-
inferShape
-
loadLib
-
optimizeFor
-
createCachedOp
Creates cached op flags.data_indices : [0, 2, 4] Used to label input location, param_indices : [1, 3] Used to label param location
- Parameters:
block- theMxSymbolBlockthat loaded in the backendmanager- the NDManager used to create NDArraytraining- true if CachedOp is created to forward in traning otherwise, false- Returns:
- a CachedOp for inference
-
freeCachedOp
public static void freeCachedOp(com.sun.jna.Pointer handle) -
cachedOpInvoke
public static MxNDArray[] cachedOpInvoke(MxNDManager manager, com.sun.jna.Pointer cachedOpHandle, MxNDArray[] inputs) -
checkCall
public static void checkCall(int ret)
-