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:
Cheng 2022-09-03 08:29:40 +08:00 committed by GitHub
parent 295bd26980
commit 76d17b0f48
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
5 changed files with 83 additions and 2 deletions

View file

@ -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);

View file

@ -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}. */

View file

@ -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);
}

View file

@ -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;

View file

@ -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,