mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-24 19:43:35 +00:00
Add java API for xnnpack (#12788)
* Add java API for xnnpack * provider option support * a more general interface for creating EP
This commit is contained in:
parent
295bd26980
commit
76d17b0f48
5 changed files with 83 additions and 2 deletions
|
|
@ -23,7 +23,8 @@ public enum OrtProvider {
|
|||
ACL("ACLExecutionProvider"),
|
||||
ARM_NN("ArmNNExecutionProvider"),
|
||||
ROCM("ROCMExecutionProvider"),
|
||||
CORE_ML("CoreMLExecutionProvider");
|
||||
CORE_ML("CoreMLExecutionProvider"),
|
||||
XNNPACK("XnnpackExecutionProvider");
|
||||
|
||||
private static final Map<String, OrtProvider> valueMap = new HashMap<>(values().length);
|
||||
|
||||
|
|
|
|||
|
|
@ -987,6 +987,27 @@ public class OrtSession implements AutoCloseable {
|
|||
addCoreML(OnnxRuntime.ortApiHandle, nativeHandle, OrtFlags.aggregateToInt(flags));
|
||||
}
|
||||
|
||||
/**
|
||||
* Adds Xnnpack as an execution backend. Needs to list all options here if a new option
|
||||
* supported. current supported options: {}
|
||||
*
|
||||
* @param providerOptions options pass to XNNPACK EP for initialization.
|
||||
* @throws OrtException If there was an error in native code.
|
||||
*/
|
||||
public void addXnnpack(Map<String, String> providerOptions) throws OrtException {
|
||||
checkClosed();
|
||||
String[] providerOptionKey = new String[providerOptions.size()];
|
||||
String[] providerOptionVal = new String[providerOptions.size()];
|
||||
int i = 0;
|
||||
for (Map.Entry<String, String> entry : providerOptions.entrySet()) {
|
||||
providerOptionKey[i] = entry.getKey();
|
||||
providerOptionVal[i] = entry.getValue();
|
||||
i++;
|
||||
}
|
||||
addExecutionProvider(
|
||||
OnnxRuntime.ortApiHandle, nativeHandle, "XNNPACK", providerOptionKey, providerOptionVal);
|
||||
}
|
||||
|
||||
private native void setExecutionMode(long apiHandle, long nativeHandle, int mode)
|
||||
throws OrtException;
|
||||
|
||||
|
|
@ -1098,6 +1119,14 @@ public class OrtSession implements AutoCloseable {
|
|||
|
||||
private native void addCoreML(long apiHandle, long nativeHandle, int coreMLFlags)
|
||||
throws OrtException;
|
||||
|
||||
private native void addExecutionProvider(
|
||||
long apiHandle,
|
||||
long nativeHandle,
|
||||
String epName,
|
||||
String[] providerOptionKey,
|
||||
String[] providerOptionVal)
|
||||
throws OrtException;
|
||||
}
|
||||
|
||||
/** Used to control logging and termination of a call to {@link OrtSession#run}. */
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@
|
|||
*/
|
||||
#include <jni.h>
|
||||
#include <string.h>
|
||||
#include <stdlib.h>
|
||||
#include "onnxruntime/core/session/onnxruntime_c_api.h"
|
||||
#include "OrtJniUtil.h"
|
||||
#include "ai_onnxruntime_OrtSession_SessionOptions.h"
|
||||
|
|
@ -615,3 +616,43 @@ JNIEXPORT void JNICALL Java_ai_onnxruntime_OrtSession_00024SessionOptions_addROC
|
|||
throwOrtException(jniEnv,convertErrorCode(ORT_INVALID_ARGUMENT),"This binary was not compiled with ROCM support.");
|
||||
#endif
|
||||
}
|
||||
|
||||
/*
|
||||
* Class:: ai_onnxruntime_OrtSession_SessionOptions
|
||||
* Method: addExecutionProvider
|
||||
* Signature: (JILjava/lang/String)V
|
||||
*/
|
||||
JNIEXPORT void JNICALL Java_ai_onnxruntime_OrtSession_00024SessionOptions_addExecutionProvider(
|
||||
JNIEnv* jniEnv, jobject jobj, jlong apiHandle, jlong optionsHandle,
|
||||
jstring jepName, jobjectArray configKeyArr, jobjectArray configValueArr) {
|
||||
(void)jobj;
|
||||
|
||||
const char* epName = (*jniEnv)->GetStringUTFChars(jniEnv, jepName, NULL);
|
||||
const OrtApi* api = (const OrtApi*)apiHandle;
|
||||
OrtSessionOptions* options = (OrtSessionOptions*)optionsHandle;
|
||||
int keyCount = (*jniEnv)->GetArrayLength(jniEnv, configKeyArr);
|
||||
|
||||
const char** keyArray = (const char**)malloc(keyCount * sizeof(const char*));
|
||||
const char** valueArray = (const char**)malloc(keyCount * sizeof(const char*));
|
||||
jstring* jkeyArray = (jstring*)malloc(keyCount * sizeof(jstring));
|
||||
jstring* jvalueArray = (jstring*)malloc(keyCount * sizeof(jstring));
|
||||
|
||||
for (int i = 0; i < keyCount; i++) {
|
||||
jkeyArray[i] = (jstring)((*jniEnv)->GetObjectArrayElement(jniEnv, configKeyArr, i));
|
||||
jvalueArray[i] = (jstring)((*jniEnv)->GetObjectArrayElement(jniEnv, configValueArr, i));
|
||||
keyArray[i] = (*jniEnv)->GetStringUTFChars(jniEnv, jkeyArray[i], NULL);
|
||||
valueArray[i] = (*jniEnv)->GetStringUTFChars(jniEnv, jvalueArray[i], NULL);
|
||||
}
|
||||
|
||||
checkOrtStatus(jniEnv, api, api->SessionOptionsAppendExecutionProvider(options, epName, keyArray, valueArray, keyCount));
|
||||
|
||||
for (int i = 0; i < keyCount; i++) {
|
||||
(*jniEnv)->ReleaseStringUTFChars(jniEnv, jkeyArray[i], keyArray[i]);
|
||||
(*jniEnv)->ReleaseStringUTFChars(jniEnv, jvalueArray[i], valueArray[i]);
|
||||
}
|
||||
(*jniEnv)->ReleaseStringUTFChars(jniEnv, jepName, epName);
|
||||
free((void*)keyArray);
|
||||
free((void*)valueArray);
|
||||
free((void*)jkeyArray);
|
||||
free((void*)jvalueArray);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ import java.nio.file.Path;
|
|||
import java.nio.file.Paths;
|
||||
import java.util.ArrayList;
|
||||
import java.util.Arrays;
|
||||
import java.util.Collections;
|
||||
import java.util.EnumSet;
|
||||
import java.util.HashMap;
|
||||
import java.util.HashSet;
|
||||
|
|
@ -617,6 +618,12 @@ public class InferenceTest {
|
|||
runProvider(OrtProvider.DNNL);
|
||||
}
|
||||
|
||||
@Test
|
||||
@EnabledIfSystemProperty(named = "USE_XNNPACK", matches = "1")
|
||||
public void testXNNPACK() throws OrtException {
|
||||
runProvider(OrtProvider.XNNPACK);
|
||||
}
|
||||
|
||||
private void runProvider(OrtProvider provider) throws OrtException {
|
||||
EnumSet<OrtProvider> providers = OrtEnvironment.getAvailableProviders();
|
||||
assertTrue(providers.size() > 1);
|
||||
|
|
@ -1515,6 +1522,9 @@ public class InferenceTest {
|
|||
case CORE_ML:
|
||||
options.addCoreML();
|
||||
break;
|
||||
case XNNPACK:
|
||||
options.addXnnpack(Collections.emptyMap());
|
||||
break;
|
||||
case NUPHAR:
|
||||
options.addNuphar(true, "");
|
||||
break;
|
||||
|
|
|
|||
|
|
@ -149,7 +149,7 @@ static void AddComputeCapabilityForNodeUnit(const NodeUnit& node_unit,
|
|||
}
|
||||
|
||||
// The first call to add compute capability in GetCapability, we just tell this all nodes in nodeunit
|
||||
// are supported by Xnnapck EP as long as it's target node is supported.
|
||||
// are supported by Xnnpack EP as long as it's target node is supported.
|
||||
// One node in one sub_graph separately
|
||||
static void AddComputeCapabilityForEachNodeInNodeUnit(
|
||||
const NodeUnit& node_unit,
|
||||
|
|
|
|||
Loading…
Reference in a new issue