From a7541f97530b4d0668a5147460dc55cee96f03aa Mon Sep 17 00:00:00 2001 From: Adam Pocock Date: Fri, 21 Feb 2020 16:13:02 -0500 Subject: [PATCH] [Java] Fix for incorrect input and output lengths in run call (#3064) --- java/build.gradle | 3 +- .../main/java/ai/onnxruntime/OrtSession.java | 4 +- .../java/ai/onnxruntime/InferenceTest.java | 228 +++++++++++++++++- java/testdata/partial-inputs-test-2.onnx | Bin 0 -> 250 bytes java/testdata/partial-inputs-test.onnx | Bin 0 -> 555 bytes 5 files changed, 229 insertions(+), 6 deletions(-) create mode 100644 java/testdata/partial-inputs-test-2.onnx create mode 100644 java/testdata/partial-inputs-test.onnx diff --git a/java/build.gradle b/java/build.gradle index 634ecc4776..2a46c23124 100644 --- a/java/build.gradle +++ b/java/build.gradle @@ -47,7 +47,8 @@ sourceSets.test { // add test resource files resources.srcDirs += [ "${rootProject.projectDir}/../csharp/testdata", - "${rootProject.projectDir}/../onnxruntime/test/testdata" + "${rootProject.projectDir}/../onnxruntime/test/testdata", + "${rootProject.projectDir}/../java/testdata" ] if (cmakeBuildDir != null) { // add compiled native libs diff --git a/java/src/main/java/ai/onnxruntime/OrtSession.java b/java/src/main/java/ai/onnxruntime/OrtSession.java index 7e2b56f079..2cd21cd3bf 100644 --- a/java/src/main/java/ai/onnxruntime/OrtSession.java +++ b/java/src/main/java/ai/onnxruntime/OrtSession.java @@ -251,9 +251,9 @@ public class OrtSession implements AutoCloseable { allocator.handle, inputNamesArray, inputHandles, - numInputs, + inputNamesArray.length, outputNamesArray, - numOutputs); + outputNamesArray.length); return new Result(outputNamesArray, outputValues); } else { throw new IllegalStateException("Trying to score a closed OrtSession."); diff --git a/java/src/test/java/ai/onnxruntime/InferenceTest.java b/java/src/test/java/ai/onnxruntime/InferenceTest.java index 9ca409c852..c3ed213329 100644 --- a/java/src/test/java/ai/onnxruntime/InferenceTest.java +++ b/java/src/test/java/ai/onnxruntime/InferenceTest.java @@ -40,6 +40,7 @@ import java.util.Set; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; import java.util.concurrent.TimeUnit; +import java.util.function.BiFunction; import java.util.regex.Pattern; import java.util.stream.Collectors; import java.util.stream.Stream; @@ -51,10 +52,10 @@ import org.junit.jupiter.api.Test; public class InferenceTest { private static final Pattern LOAD_PATTERN = Pattern.compile("[,\\[\\] ]"); - private static String propertiesFile = "Properties.txt"; + private static final String propertiesFile = "Properties.txt"; - private static Pattern inputPBPattern = Pattern.compile("input_*.pb"); - private static Pattern outputPBPattern = Pattern.compile("output_*.pb"); + private static final Pattern inputPBPattern = Pattern.compile("input_*.pb"); + private static final Pattern outputPBPattern = Pattern.compile("output_*.pb"); private static Path getResourcePath(String path) { return new File(InferenceTest.class.getResource(path).getFile()).toPath(); @@ -111,6 +112,227 @@ public class InferenceTest { } } + @Test + public void morePartialInputsTest() throws OrtException { + String modelPath = getResourcePath("/partial-inputs-test-2.onnx").toString(); + try (OrtEnvironment env = OrtEnvironment.getEnvironment("partialInputs"); + OrtSession.SessionOptions options = new SessionOptions(); + OrtSession session = env.createSession(modelPath, options)) { + assertNotNull(session); + assertEquals(3, session.getNumInputs()); + assertEquals(1, session.getNumOutputs()); + + // Input and output collections. + Map inputMap = new HashMap<>(); + Set requestedOutputs = new HashSet<>(); + + BiFunction unwrapFunc = + (r, s) -> { + try { + return ((float[]) r.get(s).get().getValue())[0]; + } catch (OrtException e) { + return Float.NaN; + } + }; + + // Graph has three scalar inputs, a, b, c, and a single output, ab. + OnnxTensor a = OnnxTensor.createTensor(env, new float[] {2.0f}); + OnnxTensor b = OnnxTensor.createTensor(env, new float[] {3.0f}); + OnnxTensor c = OnnxTensor.createTensor(env, new float[] {5.0f}); + + // Request all outputs, supply all inputs + inputMap.put("a:0", a); + inputMap.put("b:0", b); + inputMap.put("c:0", c); + requestedOutputs.add("ab:0"); + try (Result r = session.run(inputMap, requestedOutputs)) { + assertEquals(1, r.size()); + float abVal = unwrapFunc.apply(r, "ab:0"); + assertEquals(6.0f, abVal, 1e-10); + } + + // Don't specify an output, expect all of them returned. + try (Result r = session.run(inputMap)) { + assertEquals(1, r.size()); + float abVal = unwrapFunc.apply(r, "ab:0"); + assertEquals(6.0f, abVal, 1e-10); + } + + inputMap.clear(); + requestedOutputs.clear(); + + // Request single output ab, supply required inputs + inputMap.put("a:0", a); + inputMap.put("b:0", b); + requestedOutputs.add("ab:0"); + try (Result r = session.run(inputMap, requestedOutputs)) { + assertEquals(1, r.size()); + float abVal = unwrapFunc.apply(r, "ab:0"); + assertEquals(6.0f, abVal, 1e-10); + } + inputMap.clear(); + requestedOutputs.clear(); + + // Request output but don't supply the inputs + inputMap.put("c:0", c); + requestedOutputs.add("ab:0"); + try (Result r = session.run(inputMap, requestedOutputs)) { + fail("Expected to throw OrtException due to incorrect inputs"); + } catch (OrtException e) { + // System.out.println(e.getMessage()); + // pass + } + inputMap.clear(); + requestedOutputs.clear(); + + // Request output but don't supply all the inputs + inputMap.put("b:0", b); + requestedOutputs.add("ab:0"); + try (Result r = session.run(inputMap, requestedOutputs)) { + fail("Expected to throw OrtException due to incorrect inputs"); + } catch (OrtException e) { + // System.out.println(e.getMessage()); + // pass + } + } + } + + @Test + public void partialInputsTest() throws OrtException { + String modelPath = getResourcePath("/partial-inputs-test.onnx").toString(); + try (OrtEnvironment env = OrtEnvironment.getEnvironment("partialInputs"); + OrtSession.SessionOptions options = new SessionOptions(); + OrtSession session = env.createSession(modelPath, options)) { + assertNotNull(session); + assertEquals(3, session.getNumInputs()); + assertEquals(3, session.getNumOutputs()); + + // Input and output collections. + Map inputMap = new HashMap<>(); + Set requestedOutputs = new HashSet<>(); + + BiFunction unwrapFunc = + (r, s) -> { + try { + return ((float[]) r.get(s).get().getValue())[0]; + } catch (OrtException e) { + return Float.NaN; + } + }; + + // Graph has three scalar inputs, a, b, c, and three outputs, ab, bc, ab + bc. + OnnxTensor a = OnnxTensor.createTensor(env, new float[] {2.0f}); + OnnxTensor b = OnnxTensor.createTensor(env, new float[] {3.0f}); + OnnxTensor c = OnnxTensor.createTensor(env, new float[] {5.0f}); + + // Request all outputs, supply all inputs + inputMap.put("a:0", a); + inputMap.put("b:0", b); + inputMap.put("c:0", c); + requestedOutputs.add("ab:0"); + requestedOutputs.add("bc:0"); + requestedOutputs.add("abc:0"); + try (Result r = session.run(inputMap, requestedOutputs)) { + assertEquals(3, r.size()); + float abVal = unwrapFunc.apply(r, "ab:0"); + assertEquals(6.0f, abVal, 1e-10); + float bcVal = unwrapFunc.apply(r, "bc:0"); + assertEquals(15.0f, bcVal, 1e-10); + float abcVal = unwrapFunc.apply(r, "abc:0"); + assertEquals(21.0f, abcVal, 1e-10); + } + + // Don't specify an output, expect all of them returned. + try (Result r = session.run(inputMap)) { + assertEquals(3, r.size()); + float abVal = unwrapFunc.apply(r, "ab:0"); + assertEquals(6.0f, abVal, 1e-10); + float bcVal = unwrapFunc.apply(r, "bc:0"); + assertEquals(15.0f, bcVal, 1e-10); + float abcVal = unwrapFunc.apply(r, "abc:0"); + assertEquals(21.0f, abcVal, 1e-10); + } + + inputMap.clear(); + requestedOutputs.clear(); + + // Request single output ab, supply all inputs + inputMap.put("a:0", a); + inputMap.put("b:0", b); + inputMap.put("c:0", c); + requestedOutputs.add("ab:0"); + try (Result r = session.run(inputMap, requestedOutputs)) { + assertEquals(1, r.size()); + float abVal = unwrapFunc.apply(r, "ab:0"); + assertEquals(6.0f, abVal, 1e-10); + } + inputMap.clear(); + requestedOutputs.clear(); + + // Request single output abc, supply all inputs + inputMap.put("a:0", a); + inputMap.put("b:0", b); + inputMap.put("c:0", c); + requestedOutputs.add("abc:0"); + try (Result r = session.run(inputMap, requestedOutputs)) { + assertEquals(1, r.size()); + float abcVal = unwrapFunc.apply(r, "abc:0"); + assertEquals(21.0f, abcVal, 1e-10); + } + inputMap.clear(); + requestedOutputs.clear(); + + /* The native library does all the computations, rather than the requested subset. + * Leaving these tests commented out until it's fixed. + // Request single output ab, supply required inputs + inputMap.put("a:0",a); + inputMap.put("b:0",b); + requestedOutputs.add("ab:0"); + try (Result r = session.run(inputMap,requestedOutputs)) { + assertEquals(1,r.size()); + float abVal = unwrapFunc.apply(r,"ab:0"); + assertEquals(6.0f,abVal,1e-10); + } + inputMap.clear(); + requestedOutputs.clear(); + + // Request single output bc, supply required inputs + inputMap.put("b:0",b); + inputMap.put("c:0",c); + requestedOutputs.add("bc:0"); + try (Result r = session.run(inputMap,requestedOutputs)) { + assertEquals(1,r.size()); + float bcVal = unwrapFunc.apply(r,"bc:0"); + assertEquals(15.0f,bcVal,1e-10); + } + inputMap.clear(); + requestedOutputs.clear(); + */ + + // Request output but don't supply the inputs + inputMap.put("c:0", c); + requestedOutputs.add("ab:0"); + try (Result r = session.run(inputMap, requestedOutputs)) { + fail("Expected to throw OrtException due to incorrect inputs"); + } catch (OrtException e) { + // System.out.println(e.getMessage()); + // pass + } + inputMap.clear(); + requestedOutputs.clear(); + + // Request output but don't supply all the inputs + inputMap.put("b:0", b); + requestedOutputs.add("ab:0"); + try (Result r = session.run(inputMap, requestedOutputs)) { + fail("Expected to throw OrtException due to incorrect inputs"); + } catch (OrtException e) { + // System.out.println(e.getMessage()); + // pass + } + } + } + @Test public void createSessionFromByteArray() throws IOException, OrtException { Path modelPath = getResourcePath("/squeezenet.onnx"); diff --git a/java/testdata/partial-inputs-test-2.onnx b/java/testdata/partial-inputs-test-2.onnx new file mode 100644 index 0000000000000000000000000000000000000000..ebe37336d2ca743ad1a4a45bc2f4a7a7ed319808 GIT binary patch literal 250 zcmd;J6Jjq(Gs@4)tB_(f)HBsHv3khJrNzaZXl1~~oMdGnB%GKOUzAuLpI=&1P+Afn zA8%=8AjOoJq{Qr7nq$Sl<;I0gg%C?3P_vXQP;+`wVnGH}dvUyHN@`w7W=UmyyrF>- z2aIRM0Cz@^XhC98NoHb>Ze||P!eZT$)Z!9dqbNZx=47CAxVSht7=>84m^c_gLLe8S V2?1S>ER-a~1$GOvm=lu#2LR^wJK+ET literal 0 HcmV?d00001 diff --git a/java/testdata/partial-inputs-test.onnx b/java/testdata/partial-inputs-test.onnx new file mode 100644 index 0000000000000000000000000000000000000000..7675916dd59f6a26fc65e5d3cef0458a2be5b0d6 GIT binary patch literal 555 zcmZ`$O;5r=5S4Dx&75%SsliK+B#@$jW^bC9XrdR7UN*BWU=vcBb{q6>_=}wtFm1CO z_wAeaF~iTrg<0Kf^ZYC9Pbc%qO#b*V0;XjQERnGYbfQY!scmhF+9;*wG7deRMC`5J z$TN5X7en}(hQlEZuS+aG595`3Nte0F%(qgDh#wy$LzZYQ$yWrZ+m^T167|vY6OiY-Bo{tXsc=YNr(|YLng4^l2L+ZU z!_7F$Y4z3EEGzSIxjzd4=RM(r$9opCYwJL?&L*S<`~~S^wqy$nVBfq6K6(TG3n$^0 A8vp