From e30530d9ea214f2e540a3d21ab444d1b633c0de4 Mon Sep 17 00:00:00 2001 From: Guoyu Wang <62914304+gwang-msft@users.noreply.github.com> Date: Tue, 22 Sep 2020 14:51:39 -0700 Subject: [PATCH] Add java API for AddSessionConfigEntry (#5241) * Add session option config entry API for java * Java format * Add extra test verification * Address PR comments * Update comments Co-authored-by: gwang0000 <62914304+gwang0000@users.noreply.github.com> --- .../main/java/ai/onnxruntime/OrtSession.java | 27 +++++++++++++++++++ ...ai_onnxruntime_OrtSession_SessionOptions.c | 17 ++++++++++++ .../java/ai/onnxruntime/InferenceTest.java | 11 ++++++++ 3 files changed, 55 insertions(+) diff --git a/java/src/main/java/ai/onnxruntime/OrtSession.java b/java/src/main/java/ai/onnxruntime/OrtSession.java index 8a30124187..96763984b5 100644 --- a/java/src/main/java/ai/onnxruntime/OrtSession.java +++ b/java/src/main/java/ai/onnxruntime/OrtSession.java @@ -7,6 +7,7 @@ package ai.onnxruntime; import java.io.IOException; import java.util.ArrayList; import java.util.Arrays; +import java.util.Collections; import java.util.Iterator; import java.util.LinkedHashMap; import java.util.LinkedHashSet; @@ -495,12 +496,15 @@ public class OrtSession implements AutoCloseable { private final List customLibraryHandles; + private Map configEntries; + private boolean closed = false; /** Create an empty session options. */ public SessionOptions() { nativeHandle = createOptions(OnnxRuntime.ortApiHandle); customLibraryHandles = new ArrayList<>(); + configEntries = new LinkedHashMap(); } /** Closes the session options, releasing any memory acquired. */ @@ -675,6 +679,25 @@ public class OrtSession implements AutoCloseable { customLibraryHandles.add(customHandle); } + /** + * Adds a single session configuration entry as a pair of strings. + * + * @param configKey The config key string. + * @param configValue The config value string. + * @throws OrtException If there was an error in native code. + */ + public void addConfigEntry(String configKey, String configValue) throws OrtException { + checkClosed(); + addConfigEntry(OnnxRuntime.ortApiHandle, nativeHandle, configKey, configValue); + configEntries.put(configKey, configValue); + } + + /** Returns an unmodifiable view of the map contains all session configuration entries. */ + public Map getConfigEntries() { + checkClosed(); + return Collections.unmodifiableMap(configEntries); + } + /** * Add CUDA as an execution backend, using device 0. * @@ -844,6 +867,10 @@ public class OrtSession implements AutoCloseable { private native void closeOptions(long apiHandle, long nativeHandle); + private native void addConfigEntry( + long apiHandle, long nativeHandle, String configKey, String configValue) + throws OrtException; + /* * To use additional providers, you must build ORT with the extra providers enabled. Then call one of these * functions to enable them in the session: diff --git a/java/src/main/native/ai_onnxruntime_OrtSession_SessionOptions.c b/java/src/main/native/ai_onnxruntime_OrtSession_SessionOptions.c index 4c89debf15..536a708ff5 100644 --- a/java/src/main/native/ai_onnxruntime_OrtSession_SessionOptions.c +++ b/java/src/main/native/ai_onnxruntime_OrtSession_SessionOptions.c @@ -295,6 +295,23 @@ JNIEXPORT void JNICALL Java_ai_onnxruntime_OrtSession_00024SessionOptions_closeC (*jniEnv)->ReleaseLongArrayElements(jniEnv,libraryHandles,handles,JNI_ABORT); } +/* + * Class: ai_onnxruntime_OrtSession_SessionOptions + * Method: addConfigEntry + * Signature: (JJLjava/lang/String;Ljava/lang/String;)V + */ +JNIEXPORT void JNICALL Java_ai_onnxruntime_OrtSession_00024SessionOptions_addConfigEntry + (JNIEnv * jniEnv, jobject jobj, jlong apiHandle, jlong optionsHandle, jstring configKey, jstring configValue) { + (void) jobj; // Required JNI parameters not needed by functions which don't need to access their host object. + const OrtApi* api = (const OrtApi*)apiHandle; + OrtSessionOptions* options = (OrtSessionOptions*) optionsHandle; + const char* configKeyStr = (*jniEnv)->GetStringUTFChars(jniEnv, configKey, NULL); + const char* configValueStr = (*jniEnv)->GetStringUTFChars(jniEnv, configValue, NULL); + checkOrtStatus(jniEnv,api,api->AddSessionConfigEntry(options, configKeyStr, configValueStr)); + (*jniEnv)->ReleaseStringUTFChars(jniEnv, configKey, configKeyStr); + (*jniEnv)->ReleaseStringUTFChars(jniEnv, configValue, configValueStr); +} + /* * Class: ai_onnxruntime_OrtSession_SessionOptions * Method: addCPU diff --git a/java/src/test/java/ai/onnxruntime/InferenceTest.java b/java/src/test/java/ai/onnxruntime/InferenceTest.java index d4b6477d5b..4c29d65cf2 100644 --- a/java/src/test/java/ai/onnxruntime/InferenceTest.java +++ b/java/src/test/java/ai/onnxruntime/InferenceTest.java @@ -852,6 +852,17 @@ public class InferenceTest { options.setLoggerId("monkeys"); options.setSessionLogLevel(OrtLoggingLevel.ORT_LOGGING_LEVEL_FATAL); options.setSessionLogVerbosityLevel(5); + Map configEntries = options.getConfigEntries(); + assertTrue(configEntries.isEmpty()); + options.addConfigEntry("key", "value"); + assertEquals("value", configEntries.get("key")); + try { + options.addConfigEntry("", "invalid key"); + fail("Add config entry with empty key should have failed"); + } catch (OrtException e) { + assertTrue(e.getMessage().contains("Config key is empty")); + assertEquals(OrtException.OrtErrorCode.ORT_INVALID_ARGUMENT, e.getCode()); + } try (OrtSession session = env.createSession(modelPath, options)) { String inputName = session.getInputNames().iterator().next(); Map container = new HashMap<>();