Add onnx-djl module: DistilBERT sentiment with ONNX Runtime and DJL, batching and a Spring Boot endpoint with measured latency

Co-Authored-By: Claude Sonnet 5.5 <[email protected]>
Claude-Session: https://claude.ai/code/session_01G8ikz8xdWuTP5yun8DZ1hk
This commit is contained in:
Claude
2026-10-11 09:33:36 +00:00
parent fdff2b7388
commit 280cbe7d56
20 changed files with 805 additions and 0 deletions
+11
View File
@@ -8,6 +8,7 @@ facts, so the build fails when a claim stops being true.
|---|---|---|
| `embabel` | Embabel: Goal-Oriented AI Agents on the JVM | actions, goals, conditions and cost-based planning with GOAP, compared with plain Spring AI |
| `a2a` | A2A Protocol in Java: Agents That Talk to Each Other | agent card, tasks, streaming, input-required and agent-to-agent calls with the A2A Java SDK |
| `onnx-djl` | Running ONNX Models in Java with ONNX Runtime and DJL | a DistilBERT sentiment model: DJL tokenizer, ONNX Runtime directly and through DJL's engine, batching, and a Spring Boot endpoint with latency numbers |
| `jlama` | Local LLM Inference in Pure Java with Jlama | load and run a 4-bit TinyLlama inside the JVM on the Vector API, tokens per second, streaming, and the same measurement against Ollama |
## Versions (verified on Maven Central, 2026-10-11)
@@ -18,6 +19,10 @@ facts, so the build fails when a claim stops being true.
| A2A Java SDK `io.github.a2asdk` | 1.0.0.Alpha3 | alpha; targets A2A protocol 1.0. The newest stable release, 0.3.3.Final, targets the 0.3 protocol and its API differs |
| Jlama `jlama-core` (`com.github.tjake`) | 0.8.4 | needs `--add-modules jdk.incubator.vector` at compile and run time |
| Model | `tjake/TinyLlama-1.1B-Chat-v1.0-Jlama-Q4` | 1.1 GB, downloaded by Jlama into `jlama/models/` on first run (git-ignored) |
| ONNX Runtime `com.microsoft.onnxruntime:onnxruntime` | 1.31.0 | overrides the 1.21.1 that DJL's engine declares |
| DJL (`ai.djl`) | 0.38.0 | `api`, `huggingface:tokenizers`, `onnxruntime:onnxruntime-engine` |
| Spring Boot | 4.1.1 | `onnx-djl` endpoint only |
| Model | `distilbert/distilbert-base-uncased-finetuned-sst-2-english` | `onnx/model.onnx` (268 MB), `onnx/tokenizer.json`, `config.json` in `onnx-djl/models/sst2/` (git-ignored) |
| JDK | 25 LTS | |
| JUnit | 6.1.3 | |
@@ -46,3 +51,9 @@ mvn test -pl jlama -Dollama.url=http://127.0.0.1:11434
```
The committed numbers in `jlama/output/` come from one 2-vCPU cloud machine, not a laptop. Run it on yours and expect different figures.
## The `onnx-djl` module needs a model on disk
Download four files from `https://huggingface.co/distilbert/distilbert-base-uncased-finetuned-sst-2-english/resolve/main/` into `onnx-djl/models/sst2/`:
`onnx/model.onnx`, `onnx/tokenizer.json`, `onnx/tokenizer_config.json` and `config.json`, then run `mvn test -pl onnx-djl`.
Each test class runs in its own JVM (`reuseForks=false`) because only one creator of ONNX Runtime's shared environment may exist per JVM.
+1
View File
@@ -0,0 +1 @@
models/
@@ -0,0 +1,19 @@
Model: distilbert-base-uncased-finetuned-sst-2-english, ONNX export, 267 MB
ONNX Runtime 1.31.0
Inputs of the ONNX graph:
input_ids TensorInfo(javaType=INT64,onnxType=ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64,shape=[-1, -1],dimNames=[batch_size,sequence_length])
attention_mask TensorInfo(javaType=INT64,onnxType=ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64,shape=[-1, -1],dimNames=[batch_size,sequence_length])
Outputs:
logits TensorInfo(javaType=FLOAT,onnxType=ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT,shape=[-1, 2],dimNames=[batch_size,""])
Tokenizing "I loved this film.":
tokens: [[CLS], i, loved, this, film, ., [SEP]]
input_ids: [101, 1045, 3866, 2023, 2143, 1012, 102]
attention_mask: [1, 1, 1, 1, 1, 1, 1]
Two texts in one batch are padded to the same length:
[[CLS], i, loved, this, film, ., [SEP]]
mask [1, 1, 1, 1, 1, 1, 1]
[[CLS], bad, ., [SEP], [PAD], [PAD], [PAD]]
mask [1, 1, 1, 1, 0, 0, 0]
+8
View File
@@ -0,0 +1,8 @@
ONNX Runtime directly, DJL tokenizer, one batch of 6 texts.
POSITIVE 1.00 I loved this film, the acting was wonderful.
NEGATIVE 1.00 The service was slow and the food arrived cold.
POSITIVE 1.00 It works.
POSITIVE 1.00 Not bad at all.
NEGATIVE 1.00 I wanted to love it, but it was a mess.
NEGATIVE 0.95 The build finished in four minutes.
+10
View File
@@ -0,0 +1,10 @@
DJL engine: OnnxRuntime 1.21.1
direct POSITIVE 0.9999 | DJL POSITIVE 0.9999 | difference 0.00e+00
direct NEGATIVE 0.9998 | DJL NEGATIVE 0.9998 | difference 0.00e+00
direct POSITIVE 0.9999 | DJL POSITIVE 0.9999 | difference 0.00e+00
direct POSITIVE 0.9993 | DJL POSITIVE 0.9993 | difference 0.00e+00
direct NEGATIVE 0.9996 | DJL NEGATIVE 0.9996 | difference 0.00e+00
direct NEGATIVE 0.9491 | DJL NEGATIVE 0.9491 | difference 0.00e+00
Same labels: true, largest score difference: 0.00e+00
+9
View File
@@ -0,0 +1,9 @@
Padding does not change the answers: 32 texts one by one versus one batch of 32.
same labels: true, largest score difference 0.00e+00
Time to classify 32 texts, median of 5 runs after a warm-up, by batch size:
batch size 1: 230 ms in total, 7.2 ms per text
batch size 4: 204 ms in total, 6.4 ms per text
batch size 8: 166 ms in total, 5.2 ms per text
batch size 16: 142 ms in total, 4.4 ms per text
batch size 32: 134 ms in total, 4.2 ms per text
+8
View File
@@ -0,0 +1,8 @@
The same two things started in one JVM, in both orders (a separate JVM for each).
direct-first (ONNX Runtime created by my code, then DJL):
Caused by: java.lang.IllegalStateException: Tried to specify the thread pool when creating an OrtEnvironment, but one already exists.
djl-first (ONNX Runtime created by DJL, then my code):
both started
+16
View File
@@ -0,0 +1,16 @@
POST /sentiment
{"texts":["I loved this film, the acting was wonderful.","The service was slow and the food arrived cold."]}
HTTP 200
{
"results" : [ {
"label" : "POSITIVE",
"score" : 0.9998835686297906
}, {
"label" : "NEGATIVE",
"score" : 0.9997554818864473
} ]
}
Empty list: HTTP 400
65 texts (limit is 64): HTTP 400
+8
View File
@@ -0,0 +1,8 @@
Round trip over HTTP on the same machine (client and server in one JVM), median of 20 requests.
1 text(s) per request: 17.6 ms per request, 17.6 ms per text
4 text(s) per request: 31.7 ms per request, 7.9 ms per text
16 text(s) per request: 80.0 ms per request, 5.0 ms per text
64 text(s) per request: 297.2 ms per request, 4.6 ms per text
100 requests of 1 text: one client after another 1.7 s (59 requests/s), 4 clients at once 1.4 s (70 requests/s)
+20
View File
@@ -0,0 +1,20 @@
mvn dependency:tree for the onnx-djl module (first level and the ONNX Runtime entries):
ai.djl:api:jar:0.38.0:compile
ai.djl.huggingface:tokenizers:jar:0.38.0:compile
ai.djl.onnxruntime:onnxruntime-engine:jar:0.38.0:compile
com.microsoft.onnxruntime:onnxruntime:jar:1.31.0:compile
org.springframework.boot:spring-boot-starter-webmvc:jar:4.1.1:compile
org.slf4j:slf4j-simple:jar:2.0.17:test
org.springframework.boot:spring-boot-starter-test:jar:4.1.1:test
org.assertj:assertj-core:jar:3.27.7:test
ONNX Runtime and DJL entries anywhere in the tree:
ai.djl.huggingface:tokenizers:jar:0.38.0:compile
ai.djl.onnxruntime:onnxruntime-engine:jar:0.38.0:compile
ai.djl:api:jar:0.38.0:compile
com.microsoft.onnxruntime:onnxruntime:jar:1.31.0:compile
Verbose tree lines about ONNX Runtime (-Dverbose):
+- ai.djl.onnxruntime:onnxruntime-engine:jar:0.38.0:compile
| \- (com.microsoft.onnxruntime:onnxruntime:jar:1.21.1:compile - omitted for conflict with 1.31.0)
+- com.microsoft.onnxruntime:onnxruntime:jar:1.31.0:compile (scope not updated to compile)
+85
View File
@@ -0,0 +1,85 @@
<?xml version="1.0" encoding="UTF-8"?>
<project xmlns="http://maven.apache.org/POM/4.0.0"
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd">
<modelVersion>4.0.0</modelVersion>
<parent>
<groupId>com.ankurm.agents</groupId>
<artifactId>java-ai-agents</artifactId>
<version>1.0.0</version>
</parent>
<artifactId>onnx-djl</artifactId>
<dependencyManagement>
<dependencies>
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-dependencies</artifactId>
<version>${spring-boot.version}</version>
<type>pom</type>
<scope>import</scope>
</dependency>
</dependencies>
</dependencyManagement>
<build>
<plugins>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-surefire-plugin</artifactId>
<configuration>
<!-- one JVM per test class: only one creator of the shared OrtEnvironment may exist per JVM (see test 05) -->
<reuseForks>false</reuseForks>
<systemPropertyVariables>
<sentiment.model-dir>${project.basedir}/models/sst2</sentiment.model-dir>
</systemPropertyVariables>
</configuration>
</plugin>
</plugins>
</build>
<dependencies>
<dependency>
<groupId>ai.djl</groupId>
<artifactId>api</artifactId>
<version>${djl.version}</version>
</dependency>
<dependency>
<groupId>ai.djl.huggingface</groupId>
<artifactId>tokenizers</artifactId>
<version>${djl.version}</version>
</dependency>
<dependency>
<groupId>ai.djl.onnxruntime</groupId>
<artifactId>onnxruntime-engine</artifactId>
<version>${djl.version}</version>
</dependency>
<dependency>
<groupId>com.microsoft.onnxruntime</groupId>
<artifactId>onnxruntime</artifactId>
<version>${onnxruntime.version}</version>
</dependency>
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-webmvc</artifactId>
</dependency>
<dependency>
<groupId>org.slf4j</groupId>
<artifactId>slf4j-simple</artifactId>
<version>${slf4j.version}</version>
<scope>test</scope>
</dependency>
<dependency>
<groupId>org.springframework.boot</groupId>
<artifactId>spring-boot-starter-test</artifactId>
<scope>test</scope>
</dependency>
<dependency>
<groupId>org.assertj</groupId>
<artifactId>assertj-core</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>
@@ -0,0 +1,94 @@
package com.ankurm.agents.onnx;
import ai.djl.MalformedModelException;
import ai.djl.huggingface.tokenizers.Encoding;
import ai.djl.huggingface.tokenizers.HuggingFaceTokenizer;
import ai.djl.inference.Predictor;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.NDManager;
import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelNotFoundException;
import ai.djl.repository.zoo.ZooModel;
import ai.djl.translate.NoBatchifyTranslator;
import ai.djl.translate.TranslateException;
import ai.djl.translate.TranslatorContext;
import java.io.IOException;
import java.nio.file.Path;
import java.util.Arrays;
import java.util.List;
/** The same model, loaded through DJL's ONNX Runtime engine instead of calling ONNX Runtime directly. */
public final class DjlSentiment implements AutoCloseable {
private final ZooModel<String[], Sentiment[]> model;
private final Predictor<String[], Sentiment[]> predictor;
public DjlSentiment(Path modelDir) throws IOException, ModelNotFoundException, MalformedModelException {
HuggingFaceTokenizer tokenizer = OnnxSentiment.tokenizer(modelDir);
Criteria<String[], Sentiment[]> criteria = Criteria.builder()
.setTypes(String[].class, Sentiment[].class)
.optModelPath(modelDir.resolve("onnx"))
.optModelName("model")
.optEngine("OnnxRuntime")
.optTranslator(new BatchTranslator(tokenizer))
.build();
this.model = criteria.loadModel();
this.predictor = model.newPredictor();
}
public List<Sentiment> classify(List<String> texts) {
try {
return Arrays.asList(predictor.predict(texts.toArray(new String[0])));
} catch (TranslateException e) {
throw new IllegalStateException(e);
}
}
public String engineName() {
return model.getNDManager().getEngine().getEngineName() + " " + model.getNDManager().getEngine().getVersion();
}
@Override
public void close() {
predictor.close();
model.close();
}
/** Turns texts into the two input arrays the model expects, and the logits back into labels. */
static final class BatchTranslator implements NoBatchifyTranslator<String[], Sentiment[]> {
private final HuggingFaceTokenizer tokenizer;
BatchTranslator(HuggingFaceTokenizer tokenizer) {
this.tokenizer = tokenizer;
}
@Override
public NDList processInput(TranslatorContext ctx, String[] texts) {
Encoding[] encodings = tokenizer.batchEncode(texts);
long[][] ids = new long[encodings.length][];
long[][] mask = new long[encodings.length][];
for (int i = 0; i < encodings.length; i++) {
ids[i] = encodings[i].getIds();
mask[i] = encodings[i].getAttentionMask();
}
NDManager manager = ctx.getNDManager();
NDArray idArray = manager.create(ids);
idArray.setName("input_ids");
NDArray maskArray = manager.create(mask);
maskArray.setName("attention_mask");
return new NDList(idArray, maskArray);
}
@Override
public Sentiment[] processOutput(TranslatorContext ctx, NDList list) {
float[] flat = list.get(0).toFloatArray();
Sentiment[] out = new Sentiment[flat.length / 2];
for (int i = 0; i < out.length; i++) {
out[i] = OnnxSentiment.softmaxWinner(new float[]{flat[2 * i], flat[2 * i + 1]});
}
return out;
}
}
}
@@ -0,0 +1,22 @@
package com.ankurm.agents.onnx;
import java.nio.file.Path;
/**
* Starts ONNX Runtime directly and through DJL in one JVM, in the order given on the command line.
* A test runs it twice, once per order, to show that the order matters.
*/
public final class EngineOrder {
public static void main(String[] args) throws Exception {
Path dir = Path.of(args[1]);
if (args[0].equals("direct-first")) {
try (OnnxSentiment direct = new OnnxSentiment(dir); DjlSentiment djl = new DjlSentiment(dir)) {
System.out.println("both started");
}
} else {
try (DjlSentiment djl = new DjlSentiment(dir); OnnxSentiment direct = new OnnxSentiment(dir)) {
System.out.println("both started");
}
}
}
}
@@ -0,0 +1,84 @@
package com.ankurm.agents.onnx;
import ai.djl.huggingface.tokenizers.Encoding;
import ai.djl.huggingface.tokenizers.HuggingFaceTokenizer;
import ai.onnxruntime.OnnxTensor;
import ai.onnxruntime.OrtEnvironment;
import ai.onnxruntime.OrtException;
import ai.onnxruntime.OrtSession;
import java.io.IOException;
import java.nio.file.Path;
import java.util.List;
import java.util.Map;
/**
* Sentiment classification with a DistilBERT model exported to ONNX. DJL does the tokenizing, ONNX Runtime's own
* Java API runs the model. A batch of texts goes through the model in one call.
*/
public final class OnnxSentiment implements AutoCloseable {
/** Order matches config.json: id2label {0: NEGATIVE, 1: POSITIVE}. */
static final String[] LABELS = {"NEGATIVE", "POSITIVE"};
private final OrtEnvironment env = OrtEnvironment.getEnvironment();
private final OrtSession session;
private final HuggingFaceTokenizer tokenizer;
public OnnxSentiment(Path modelDir) throws IOException, OrtException {
this.tokenizer = tokenizer(modelDir);
this.session = env.createSession(modelDir.resolve("onnx/model.onnx").toString(), new OrtSession.SessionOptions());
}
/** Pads every text in a batch to the longest one and cuts anything over 128 tokens. */
static HuggingFaceTokenizer tokenizer(Path modelDir) throws IOException {
return HuggingFaceTokenizer.builder()
.optTokenizerPath(modelDir.resolve("onnx/tokenizer.json"))
.optPadding(true)
.optTruncation(true)
.optMaxLength(128)
.build();
}
public HuggingFaceTokenizer tokenizer() {
return tokenizer;
}
public OrtSession session() {
return session;
}
public List<Sentiment> classify(List<String> texts) {
Encoding[] encodings = tokenizer.batchEncode(texts.toArray(new String[0]));
int rows = encodings.length;
int cols = encodings[0].getIds().length;
long[][] ids = new long[rows][];
long[][] mask = new long[rows][];
for (int i = 0; i < rows; i++) {
ids[i] = encodings[i].getIds();
mask[i] = encodings[i].getAttentionMask();
}
try (OnnxTensor idTensor = OnnxTensor.createTensor(env, ids);
OnnxTensor maskTensor = OnnxTensor.createTensor(env, mask);
OrtSession.Result result = session.run(Map.of("input_ids", idTensor, "attention_mask", maskTensor))) {
float[][] logits = (float[][]) result.get(0).getValue();
return java.util.stream.Stream.of(logits).map(OnnxSentiment::softmaxWinner).toList();
} catch (OrtException e) {
throw new IllegalStateException(e);
}
}
static Sentiment softmaxWinner(float[] logits) {
double max = Math.max(logits[0], logits[1]);
double e0 = Math.exp(logits[0] - max);
double e1 = Math.exp(logits[1] - max);
int best = e1 > e0 ? 1 : 0;
return new Sentiment(LABELS[best], Math.max(e0, e1) / (e0 + e1));
}
@Override
public void close() throws OrtException {
session.close();
tokenizer.close();
}
}
@@ -0,0 +1,5 @@
package com.ankurm.agents.onnx;
/** One classification result: the winning label and the softmax probability behind it. */
public record Sentiment(String label, double score) {
}
@@ -0,0 +1,55 @@
package com.ankurm.agents.onnx;
import ai.onnxruntime.OrtException;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.boot.SpringApplication;
import org.springframework.boot.autoconfigure.SpringBootApplication;
import org.springframework.context.annotation.Bean;
import org.springframework.web.bind.annotation.PostMapping;
import org.springframework.web.bind.annotation.RequestBody;
import org.springframework.web.bind.annotation.RestController;
import org.springframework.web.server.ResponseStatusException;
import org.springframework.http.HttpStatus;
import java.io.IOException;
import java.nio.file.Path;
import java.util.List;
/** A sentiment endpoint: POST /sentiment with {"texts": [...]} classifies the whole batch in one model call. */
@SpringBootApplication
public class SentimentApplication {
public static void main(String[] args) {
SpringApplication.run(SentimentApplication.class, args);
}
/** One session for the whole application; ONNX Runtime sessions can be used from several threads at once. */
@Bean(destroyMethod = "close")
OnnxSentiment onnxSentiment(@Value("${sentiment.model-dir}") String modelDir) throws IOException, OrtException {
return new OnnxSentiment(Path.of(modelDir));
}
record Request(List<String> texts) {
}
record Response(List<Sentiment> results) {
}
@RestController
static class SentimentController {
private static final int MAX_BATCH = 64;
private final OnnxSentiment model;
SentimentController(OnnxSentiment model) {
this.model = model;
}
@PostMapping("/sentiment")
Response classify(@RequestBody Request request) {
if (request.texts() == null || request.texts().isEmpty() || request.texts().size() > MAX_BATCH) {
throw new ResponseStatusException(HttpStatus.BAD_REQUEST, "send between 1 and " + MAX_BATCH + " texts");
}
return new Response(model.classify(request.texts()));
}
}
}
@@ -0,0 +1,198 @@
package com.ankurm.agents.onnx;
import ai.djl.huggingface.tokenizers.Encoding;
import ai.onnxruntime.NodeInfo;
import ai.onnxruntime.OrtEnvironment;
import org.junit.jupiter.api.AfterAll;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.TestInstance;
import java.nio.file.Files;
import java.nio.file.Path;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.List;
import java.util.Locale;
import java.util.Map;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.within;
@TestInstance(TestInstance.Lifecycle.PER_CLASS)
class OnnxDjlTest {
static final Path MODEL_DIR = Path.of(System.getProperty("sentiment.model-dir", "models/sst2"));
static final List<String> SENTENCES = List.of(
"I loved this film, the acting was wonderful.",
"The service was slow and the food arrived cold.",
"It works.",
"Not bad at all.",
"I wanted to love it, but it was a mess.",
"The build finished in four minutes.");
OnnxSentiment onnx;
DjlSentiment djl;
@BeforeAll
void load() throws Exception {
// DJL first: its engine creates the shared OrtEnvironment, and a second creator fails (see the engine-order test).
djl = new DjlSentiment(MODEL_DIR);
onnx = new OnnxSentiment(MODEL_DIR);
}
@AfterAll
void close() throws Exception {
onnx.close();
djl.close();
}
static String f4(double d) {
return String.format(Locale.ROOT, "%.4f", d);
}
static String f2(double d) {
return String.format(Locale.ROOT, "%.2f", d);
}
// ------------------------------------------------------------------ 1. what the model expects
@Test
void showsWhatGoesIntoTheModelAndWhatComesOut() throws Exception {
StringBuilder sb = new StringBuilder("Model: distilbert-base-uncased-finetuned-sst-2-english, ONNX export, "
+ Files.size(MODEL_DIR.resolve("onnx/model.onnx")) / 1_000_000 + " MB\n");
sb.append("ONNX Runtime ").append(OrtEnvironment.getEnvironment().getVersion()).append('\n');
sb.append("\nInputs of the ONNX graph:\n");
for (Map.Entry<String, NodeInfo> e : onnx.session().getInputInfo().entrySet()) {
sb.append(" ").append(e.getKey()).append(" ").append(e.getValue().getInfo()).append('\n');
}
sb.append("Outputs:\n");
for (Map.Entry<String, NodeInfo> e : onnx.session().getOutputInfo().entrySet()) {
sb.append(" ").append(e.getKey()).append(" ").append(e.getValue().getInfo()).append('\n');
}
String text = "I loved this film.";
Encoding enc = onnx.tokenizer().encode(text);
sb.append("\nTokenizing \"").append(text).append("\":\n");
sb.append(" tokens: ").append(Arrays.toString(enc.getTokens())).append('\n');
sb.append(" input_ids: ").append(Arrays.toString(enc.getIds())).append('\n');
sb.append(" attention_mask: ").append(Arrays.toString(enc.getAttentionMask())).append('\n');
Encoding[] batch = onnx.tokenizer().batchEncode(new String[]{"I loved this film.", "Bad."});
sb.append("\nTwo texts in one batch are padded to the same length:\n");
for (Encoding e : batch) {
sb.append(" ").append(Arrays.toString(e.getTokens())).append("\n mask ").append(Arrays.toString(e.getAttentionMask())).append('\n');
}
Transcript.write("01-model-and-tokenizer.txt", sb.toString());
assertThat(onnx.session().getInputNames()).containsExactlyInAnyOrder("input_ids", "attention_mask");
assertThat(batch[0].getIds()).hasSameSizeAs(batch[1].getIds());
assertThat(batch[1].getAttentionMask()).contains(0L);
}
// ------------------------------------------------------------------ 2. classification
@Test
void classifiesSentences() {
List<Sentiment> results = onnx.classify(SENTENCES);
StringBuilder sb = new StringBuilder("ONNX Runtime directly, DJL tokenizer, one batch of " + SENTENCES.size() + " texts.\n\n");
for (int i = 0; i < SENTENCES.size(); i++) {
sb.append(String.format(Locale.ROOT, " %-9s %s %s%n", results.get(i).label(), f2(results.get(i).score()), SENTENCES.get(i)));
}
Transcript.write("02-classification.txt", sb.toString());
assertThat(results.get(0).label()).isEqualTo("POSITIVE");
assertThat(results.get(1).label()).isEqualTo("NEGATIVE");
}
// ------------------------------------------------------------------ 3. the DJL engine gives the same answers
@Test
void theDjlEngineGivesTheSameAnswers() {
List<Sentiment> direct = onnx.classify(SENTENCES);
List<Sentiment> viaDjl = djl.classify(SENTENCES);
StringBuilder sb = new StringBuilder("DJL engine: " + djl.engineName() + "\n\n");
double worst = 0;
for (int i = 0; i < SENTENCES.size(); i++) {
double diff = Math.abs(direct.get(i).score() - viaDjl.get(i).score());
worst = Math.max(worst, diff);
sb.append(String.format(Locale.ROOT, " direct %-9s %s | DJL %-9s %s | difference %.2e%n", direct.get(i).label(),
f4(direct.get(i).score()), viaDjl.get(i).label(), f4(viaDjl.get(i).score()), diff));
}
sb.append(String.format(Locale.ROOT, "%nSame labels: %s, largest score difference: %.2e%n",
direct.stream().map(Sentiment::label).toList().equals(viaDjl.stream().map(Sentiment::label).toList()), worst));
Transcript.write("03-djl-engine.txt", sb.toString());
assertThat(viaDjl.stream().map(Sentiment::label).toList()).isEqualTo(direct.stream().map(Sentiment::label).toList());
for (int i = 0; i < direct.size(); i++) {
assertThat(viaDjl.get(i).score()).isCloseTo(direct.get(i).score(), within(1e-4));
}
}
// ------------------------------------------------------------------ 4. batching
@Test
void batchingGivesTheSameAnswersAndCostsLessPerText() {
List<String> texts = new ArrayList<>();
for (int i = 0; i < 32; i++) texts.add(SENTENCES.get(i % SENTENCES.size()));
List<Sentiment> single = new ArrayList<>();
for (String t : texts) single.add(onnx.classify(List.of(t)).get(0));
List<Sentiment> batched = onnx.classify(texts);
double worst = 0;
for (int i = 0; i < texts.size(); i++) worst = Math.max(worst, Math.abs(single.get(i).score() - batched.get(i).score()));
StringBuilder sb = new StringBuilder("Padding does not change the answers: 32 texts one by one versus one batch of 32.\n");
sb.append(String.format(Locale.ROOT, " same labels: %s, largest score difference %.2e%n%n",
single.stream().map(Sentiment::label).toList().equals(batched.stream().map(Sentiment::label).toList()), worst));
sb.append("Time to classify 32 texts, median of 5 runs after a warm-up, by batch size:\n");
long baseline = 0;
for (int size : new int[]{1, 4, 8, 16, 32}) {
long ms = medianMillis(() -> {
for (int i = 0; i < texts.size(); i += size) onnx.classify(texts.subList(i, i + size));
});
if (size == 1) baseline = ms;
sb.append(String.format(Locale.ROOT, " batch size %2d: %5d ms in total, %5.1f ms per text%n", size, ms, ms / 32.0));
}
Transcript.write("04-batching.txt", sb.toString());
assertThat(single.stream().map(Sentiment::label).toList()).isEqualTo(batched.stream().map(Sentiment::label).toList());
assertThat(worst).isLessThan(1e-3);
assertThat(baseline).isPositive();
}
static long medianMillis(Runnable r) {
r.run();
long[] t = new long[5];
for (int i = 0; i < t.length; i++) {
long t0 = System.nanoTime();
r.run();
t[i] = (System.nanoTime() - t0) / 1_000_000;
}
Arrays.sort(t);
return t[t.length / 2];
}
// ------------------------------------------------------------------ 5. who creates the ONNX Runtime environment first
@Test
void theOrderOfStartingTheTwoMatters() throws Exception {
StringBuilder sb = new StringBuilder("The same two things started in one JVM, in both orders (a separate JVM for each).\n\n");
String[] results = new String[2];
String[] orders = {"direct-first", "djl-first"};
for (int i = 0; i < 2; i++) {
Process p = new ProcessBuilder(Path.of(System.getProperty("java.home"), "bin", "java").toString(),
"-cp", System.getProperty("java.class.path"), EngineOrder.class.getName(), orders[i], MODEL_DIR.toAbsolutePath().toString())
.redirectErrorStream(true).start();
String out = new String(p.getInputStream().readAllBytes());
p.waitFor();
String key = out.lines().filter(l -> l.contains("IllegalStateException") || l.equals("both started")).findFirst().orElse("(no result line)");
results[i] = key.strip();
sb.append(orders[i]).append(" (ONNX Runtime ").append(i == 0 ? "created by my code, then DJL" : "created by DJL, then my code")
.append("):\n ").append(results[i]).append("\n\n");
}
Transcript.write("05-engine-order.txt", sb.toString());
assertThat(results[1]).isEqualTo("both started");
assertThat(results[0]).contains("one already exists");
}
}
@@ -0,0 +1,109 @@
package com.ankurm.agents.onnx;
import tools.jackson.databind.JsonNode;
import tools.jackson.databind.json.JsonMapper;
import org.junit.jupiter.api.Test;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.boot.test.context.SpringBootTest;
import java.net.URI;
import java.net.http.HttpClient;
import java.net.http.HttpRequest;
import java.net.http.HttpResponse;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.List;
import java.util.Locale;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
import java.util.concurrent.Future;
import static org.assertj.core.api.Assertions.assertThat;
@SpringBootTest(classes = SentimentApplication.class, webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT)
class SentimentEndpointTest {
@Value("${local.server.port}")
int port;
final HttpClient http = HttpClient.newHttpClient();
final JsonMapper json = JsonMapper.builder().build();
HttpResponse<String> post(String body) throws Exception {
return http.send(HttpRequest.newBuilder(URI.create("http://127.0.0.1:" + port + "/sentiment"))
.header("Content-Type", "application/json").POST(HttpRequest.BodyPublishers.ofString(body)).build(), HttpResponse.BodyHandlers.ofString());
}
String body(int n) throws Exception {
List<String> texts = new ArrayList<>();
for (int i = 0; i < n; i++) texts.add(OnnxDjlTest.SENTENCES.get(i % OnnxDjlTest.SENTENCES.size()));
return json.writeValueAsString(java.util.Map.of("texts", texts));
}
static long median(long[] v) {
long[] c = v.clone();
Arrays.sort(c);
return c[c.length / 2];
}
long medianMs(String body, int runs) throws Exception {
long[] t = new long[runs];
for (int i = 0; i < runs; i++) {
long t0 = System.nanoTime();
assertThat(post(body).statusCode()).isEqualTo(200);
t[i] = (System.nanoTime() - t0) / 1000;
}
return median(t);
}
@Test
void anEndpointClassifiesABatchAndRejectsBadInput() throws Exception {
String req = "{\"texts\":[\"I loved this film, the acting was wonderful.\",\"The service was slow and the food arrived cold.\"]}";
HttpResponse<String> ok = post(req);
HttpResponse<String> empty = post("{\"texts\":[]}");
HttpResponse<String> tooMany = post(body(65));
JsonNode tree = json.readTree(ok.body());
StringBuilder sb = new StringBuilder("POST /sentiment\n").append(req).append("\n\nHTTP ").append(ok.statusCode()).append('\n')
.append(json.writerWithDefaultPrettyPrinter().writeValueAsString(tree)).append("\n\n");
sb.append("Empty list: HTTP ").append(empty.statusCode()).append('\n');
sb.append("65 texts (limit is 64): HTTP ").append(tooMany.statusCode()).append('\n');
Transcript.write("06-endpoint.txt", sb.toString());
assertThat(tree.get("results")).hasSize(2);
assertThat(tree.get("results").get(0).get("label").asText()).isEqualTo("POSITIVE");
assertThat(empty.statusCode()).isEqualTo(400);
assertThat(tooMany.statusCode()).isEqualTo(400);
}
@Test
void measuresLatencyAndThroughputOverHttp() throws Exception {
medianMs(body(4), 10); // warm-up: the JIT and Tomcat's threads
StringBuilder sb = new StringBuilder("Round trip over HTTP on the same machine (client and server in one JVM), median of 20 requests.\n\n");
for (int n : new int[]{1, 4, 16, 64}) {
long us = medianMs(body(n), 20);
sb.append(String.format(Locale.ROOT, " %2d text(s) per request: %6.1f ms per request, %5.1f ms per text%n", n, us / 1000.0, us / 1000.0 / n));
}
int threads = 4, perThread = 25;
String one = body(1);
long t0 = System.nanoTime();
for (int i = 0; i < threads * perThread; i++) post(one);
double seq = (System.nanoTime() - t0) / 1e9;
ExecutorService pool = Executors.newFixedThreadPool(threads);
t0 = System.nanoTime();
List<Future<?>> fs = new ArrayList<>();
for (int t = 0; t < threads; t++) fs.add(pool.submit(() -> {
for (int i = 0; i < perThread; i++) post(one);
return null;
}));
for (Future<?> f : fs) f.get();
double par = (System.nanoTime() - t0) / 1e9;
pool.shutdown();
sb.append(String.format(Locale.ROOT, "%n100 requests of 1 text: one client after another %.1f s (%.0f requests/s), 4 clients at once %.1f s (%.0f requests/s)%n",
seq, 100 / seq, par, 100 / par));
Transcript.write("07-http-latency.txt", sb.toString());
assertThat(par).isPositive();
}
}
@@ -0,0 +1,39 @@
package com.ankurm.agents.onnx;
import java.io.IOException;
import java.nio.file.Files;
import java.nio.file.Path;
import java.util.regex.Pattern;
/** Writes what a test observed to output/NN-name.txt so every figure in the post comes from a file. */
final class Transcript {
private static final Path DIR = Path.of("output");
private static final Pattern UUID = Pattern.compile("[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}");
private static final Pattern PORT = Pattern.compile("127\\.0\\.0\\.1:\\d+");
private static final Pattern TS = Pattern.compile("\\d{4}-\\d\\d-\\d\\dT[\\d:.]+Z");
private Transcript() {
}
/** Ports, ids and timestamps change on every run, so they are replaced by stable placeholders. */
static String scrub(String s) {
var ids = new java.util.LinkedHashMap<String, String>();
var m = UUID.matcher(s);
StringBuilder out = new StringBuilder();
while (m.find()) {
m.appendReplacement(out, ids.computeIfAbsent(m.group(), k -> "<id-" + (ids.size() + 1) + ">"));
}
m.appendTail(out);
return TS.matcher(PORT.matcher(out.toString()).replaceAll("127.0.0.1:PORT")).replaceAll("<time>");
}
static void write(String name, String content) {
try {
Files.createDirectories(DIR);
Files.writeString(DIR.resolve(name), scrub(content));
} catch (IOException e) {
throw new IllegalStateException(e);
}
}
}
+4
View File
@@ -13,6 +13,7 @@
<module>embabel</module>
<module>a2a</module>
<module>jlama</module>
<module>onnx-djl</module>
</modules>
<properties>
@@ -22,6 +23,9 @@
<embabel.version>1.5.3</embabel.version>
<a2a.sdk.version>1.0.0.Alpha3</a2a.sdk.version>
<jlama.version>0.8.4</jlama.version>
<djl.version>0.38.0</djl.version>
<onnxruntime.version>1.31.0</onnxruntime.version>
<spring-boot.version>4.1.1</spring-boot.version>
<junit.version>6.1.3</junit.version>
<assertj.version>3.27.6</assertj.version>
<slf4j.version>2.0.17</slf4j.version>