From 76d17b0f48f01229c686eee33c22fe70f4a0b9da Mon Sep 17 00:00:00 2001 From: Cheng Date: Sat, 3 Sep 2022 08:29:40 +0800 Subject: [PATCH] Add java API for xnnpack (#12788) * Add java API for xnnpack * provider option support * a more general interface for creating EP --- .../main/java/ai/onnxruntime/OrtProvider.java | 3 +- .../main/java/ai/onnxruntime/OrtSession.java | 29 +++++++++++++ ...ai_onnxruntime_OrtSession_SessionOptions.c | 41 +++++++++++++++++++ .../java/ai/onnxruntime/InferenceTest.java | 10 +++++ .../xnnpack/xnnpack_execution_provider.cc | 2 +- 5 files changed, 83 insertions(+), 2 deletions(-) diff --git a/java/src/main/java/ai/onnxruntime/OrtProvider.java b/java/src/main/java/ai/onnxruntime/OrtProvider.java index cf5202c432..21ad99fb8b 100644 --- a/java/src/main/java/ai/onnxruntime/OrtProvider.java +++ b/java/src/main/java/ai/onnxruntime/OrtProvider.java @@ -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 valueMap = new HashMap<>(values().length); diff --git a/java/src/main/java/ai/onnxruntime/OrtSession.java b/java/src/main/java/ai/onnxruntime/OrtSession.java index adc38d3aa5..0f4bcc639a 100644 --- a/java/src/main/java/ai/onnxruntime/OrtSession.java +++ b/java/src/main/java/ai/onnxruntime/OrtSession.java @@ -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 providerOptions) throws OrtException { + checkClosed(); + String[] providerOptionKey = new String[providerOptions.size()]; + String[] providerOptionVal = new String[providerOptions.size()]; + int i = 0; + for (Map.Entry 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}. */ diff --git a/java/src/main/native/ai_onnxruntime_OrtSession_SessionOptions.c b/java/src/main/native/ai_onnxruntime_OrtSession_SessionOptions.c index b880d1c4cb..a4b336d745 100644 --- a/java/src/main/native/ai_onnxruntime_OrtSession_SessionOptions.c +++ b/java/src/main/native/ai_onnxruntime_OrtSession_SessionOptions.c @@ -4,6 +4,7 @@ */ #include #include +#include #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); +} diff --git a/java/src/test/java/ai/onnxruntime/InferenceTest.java b/java/src/test/java/ai/onnxruntime/InferenceTest.java index 9c22f77f72..7b11b71c64 100644 --- a/java/src/test/java/ai/onnxruntime/InferenceTest.java +++ b/java/src/test/java/ai/onnxruntime/InferenceTest.java @@ -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 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; diff --git a/onnxruntime/core/providers/xnnpack/xnnpack_execution_provider.cc b/onnxruntime/core/providers/xnnpack/xnnpack_execution_provider.cc index 528aad64e0..42865fe4ba 100644 --- a/onnxruntime/core/providers/xnnpack/xnnpack_execution_provider.cc +++ b/onnxruntime/core/providers/xnnpack/xnnpack_execution_provider.cc @@ -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,