Class JnaUtils

java.lang.Object
ai.djl.mxnet.jna.JnaUtils

public final class JnaUtils extends Object
A class containing utilities to interact with the MXNet Engine's Java Native Access (JNA) layer.
  • Field Details

    • REFS

      public static final ObjectPool<com.sun.jna.ptr.PointerByReference> REFS
  • Method Details

    • getVersion

      public static int getVersion()
    • getAllOpNames

      public static Set<String> getAllOpNames()
    • getNdArrayFunctions

      public static Map<String,FunctionInfo> getNdArrayFunctions()
    • op

      public static FunctionInfo op(String opName)
    • getGpuCount

      public static int getGpuCount()
    • getGpuMemory

      public static long[] getGpuMemory(ai.djl.Device device)
    • getFeatures

      public static Set<String> 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

      public static void ndArraySyncCopyFromNdArray(MxNDArray dest, MxNDArray src, int location)
    • 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

      public static void syncCopyFromCPU(com.sun.jna.Pointer ndArray, Buffer data, int len)
    • 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

      public static void setNumpyMode(JnaUtils.NumpyMode mode)
    • getGradient

      public static com.sun.jna.Pointer getGradient(com.sun.jna.Pointer handle)
    • parameterStoreCreate

      public static com.sun.jna.Pointer parameterStoreCreate(String type)
    • 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

      public static void parameterStorePushPull(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)
    • 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

      public static String[] listSymbolOutputs(com.sun.jna.Pointer symbol)
    • freeSymbol

      public static void freeSymbol(com.sun.jna.Pointer symbol)
    • listSymbolNames

      public static String[] listSymbolNames(com.sun.jna.Pointer symbol)
    • listSymbolArguments

      public static String[] listSymbolArguments(com.sun.jna.Pointer symbol)
    • listSymbolAuxiliaryStates

      public static String[] listSymbolAuxiliaryStates(com.sun.jna.Pointer symbol)
    • getSymbolInternals

      public static com.sun.jna.Pointer getSymbolInternals(com.sun.jna.Pointer symbol)
    • createSymbolFromFile

      public static com.sun.jna.Pointer createSymbolFromFile(String path)
    • createSymbolFromString

      public static com.sun.jna.Pointer createSymbolFromString(String json)
    • getSymbolString

      public static String getSymbolString(com.sun.jna.Pointer symbol)
    • inferShape

      public static List<List<ai.djl.ndarray.types.Shape>> inferShape(Symbol symbol, ai.djl.util.PairList<String,ai.djl.ndarray.types.Shape> args)
    • loadLib

      public static void loadLib(String path, boolean verbose)
    • optimizeFor

      public static com.sun.jna.Pointer optimizeFor(Symbol current, String backend, ai.djl.Device device)
    • createCachedOp

      public static CachedOp createCachedOp(MxSymbolBlock block, MxNDManager manager, boolean training)
      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 - the MxSymbolBlock that loaded in the backend
      manager - the NDManager used to create NDArray
      training - 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)