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:
@@ -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.
|
||||
|
||||
@@ -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]
|
||||
@@ -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.
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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>
|
||||
|
||||
Reference in New Issue
Block a user