Skip to content

feat(provider): add Onnx in-process embedding provider - #454

Open
akhil838 wants to merge 3 commits into
kestra-io:mainfrom
akhil838:feat/onnx-embedding-provider
Open

akhil838 wants to merge 3 commits into
kestra-io:mainfrom
akhil838:feat/onnx-embedding-provider

Conversation

@akhil838

@akhil838 akhil838 commented Oct 2, 2026 •

Copy link
Copy Markdown

What changes are being made and why?

closes #354

Adds an Onnx provider that runs a sentence-embedding model inside the worker with ONNX Runtime, so you can do RAG without an API key or a model server. It only does embeddings, so it works anywhere an embedding provider is used (IngestDocument, Search, rag.ChatCompletion, EmbeddingStoreRetriever). chatModel and imageModel throw UnsupportedOperationException.

  • modelUri and tokenizerUri point to the .onnx file and its tokenizer.json. They are read with URIFetcher, so nsfile://, kestra:// and allowed file:// all work. No model is bundled in the plugin.
  • poolingMode defaults to MEAN (all-MiniLM, E5). BGE models need CLS.
  • Loaded models are cached per worker, keyed on the sha256 of the model + tokenizer + pooling mode. langchain4j never closes the ONNX session, and when I reloaded the model on every run, memory grew by about 115 MB each time.
  • On the JAR size concern from the issue: adding langchain4j-embeddings takes the shaded JAR from 233 MB to 348 MB, but 64 MB of that is onnxruntime debug symbols (.pdb / .dSYM). I excluded them, so the JAR ends up at 284 MB (+51 MB) with all the native libs still there.

How the changes have been QAed?

OnnxTest uses the all-MiniLM files that are already a test dependency, so it needs no network or container:

./gradlew test --tests 'io.kestra.plugin.ai.provider.OnnxTest'

I also tested it manually on a local Kestra with the plugin built from this branch. This flow downloads all-MiniLM, ingests three short documents and searches them; search returns the Kestra one:

id: onnx_rag
namespace: company.ai

tasks:
  - id: model
    type: io.kestra.plugin.core.http.Download
    uri: https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2/resolve/main/onnx/model.onnx

  - id: tokenizer
    type: io.kestra.plugin.core.http.Download
    uri: https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2/resolve/main/tokenizer.json

  - id: ingest
    type: io.kestra.plugin.ai.rag.IngestDocument
    provider:
      type: io.kestra.plugin.ai.provider.Onnx
      modelName: all-MiniLM-L6-v2
      modelUri: "{{ outputs.model.uri }}"
      tokenizerUri: "{{ outputs.tokenizer.uri }}"
    embeddings:
      type: io.kestra.plugin.ai.embeddings.KestraKVStore
    drop: true
    fromDocuments:
      - content: Kestra is an open-source orchestration platform for data and AI workflows.
      - content: Bananas are a good source of potassium.
      - content: PostgreSQL is a relational database.

  - id: search
    type: io.kestra.plugin.ai.rag.Search
    provider:
      type: io.kestra.plugin.ai.provider.Onnx
      modelName: all-MiniLM-L6-v2
      modelUri: "{{ outputs.model.uri }}"
      tokenizerUri: "{{ outputs.tokenizer.uri }}"
    embeddings:
      type: io.kestra.plugin.ai.embeddings.KestraKVStore
    query: Which tool orchestrates workflows?
    maxResults: 1
    minScore: 0.3
    fetchType: FETCH

I also checked loading the model from namespace files in another namespace (nsfile://company.models/...) and ingesting a file uploaded through a FILE input.

I didn't run the full test suite, since most of it needs containers or API keys. spotlessJavaCheck already fails on main, so I only formatted the two new Java files.


Contributor Checklist ✅

Adds an Onnx provider that computes embeddings inside the worker with ONNX Runtime (CPU), from a BERT-style .onnx model and its Hugging Face tokenizer.json supplied as nsfile://, kestra:// or allowed file:// URIs. Embeddings only.

Loaded models are cached per worker by sha256 of model + tokenizer + pooling mode: langchain4j never closes the ONNX Runtime session, so reloading on every task run leaked ~115 MB of native memory per run.

shadowJar excludes onnxruntime debug symbols (*.pdb, *.dSYM), cutting the added size from ~115 MB to ~51 MB.

close kestra-io#354
@kestrabot kestrabot Bot added this to Pull Requests Oct 2, 2026
@akhil838
akhil838 marked this pull request as ready for review October 2, 2026 15:08
@akhil838

akhil838 commented Oct 3, 2026 •

Copy link
Copy Markdown
Author

Hi @fdelbrayelle requesting your review here,

image image

also wanted to discuss about things that may pop up in future.

  1. CPU resource limiting for ONNX Model by using a shared resource pool. would be helpfull while running large number of workflows at once
  2. Model Lifecycle. (ability to check models that are already downloaded and delete them from the APP/ UI to recover disk space )

@fdelbrayelle
fdelbrayelle requested review from a team and jymaire October 4, 2026 11:25
@fdelbrayelle fdelbrayelle added area/plugin Plugin-related issue or feature request kind/external Pull requests raised by community contributors labels Oct 4, 2026

@jymaire jymaire left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks a lot for this contribution @akhil838, and for the detailed write-up on the JAR size!

I went through it against the plan agreed in #354, and the scope matches: embeddings only, no bundled model, MEAN/CLS pooling, and it plugs into the existing RAG tasks without any change to them. It also follows the conventions of the other providers (UnsupportedOperationException for unsupported model types, internalStorageURI properties read through URIFetcher, icon, and the doc entry).

What I checked locally on 159574e:

  • ./gradlew test --tests 'io.kestra.plugin.ai.provider.OnnxTest': 2/2 passed.
  • shadowJar: 270.4 MiB. The onnxruntime .pdb/.dSYM debug symbols are gone, and the onnxruntime and DJL tokenizers native libraries are still there for linux-x64, linux-aarch64, osx-aarch64, osx-x64 and win-x64. So the +51 MB figure from the issue holds.

I have one point, about the model cache. See the inline comment.

)
public class Onnx extends ModelProvider {
// ONNX Runtime sessions hold native memory that is never released, so each distinct model is loaded once per worker.
private static final Map<String, EmbeddingModel> MODELS = new ConcurrentHashMap<>();

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

[MEDIUM] This cache is static and never evicts anything. Each distinct combination of model bytes, tokenizer bytes and poolingMode keeps its ONNX Runtime session in native memory until the worker stops. In onnxruntime 1.20.0, OrtSession only frees that memory through an explicit close() (it has no cleaner or finalizer), and OnnxBertBiEncoder never calls it. So even dropping an entry from the map would not free anything.

Workers are often shared across namespaces (and tenants in EE), so any flow can keep adding models (a different model, a re-exported file, a different pooling mode) and grow the worker's native memory until it gets OOM-killed. That would take down every other execution on that worker.

Could you put a limit on how much this cache can hold? For example, a small LRU where the plugin creates and owns the OrtSession, so it can close the session when the entry is evicted. langchain4j has a public OnnxBertBiEncoder(OrtEnvironment, OrtSession, InputStream, PoolingMode) constructor, plus AbstractInProcessEmbeddingModel, for exactly this. You would need to make sure a session is never closed while another task is still using it. If you'd rather keep the current approach, it would be good to agree with the maintainers on the trade-off and say clearly in the docs that every distinct model stays in worker memory until restart.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

hey @jymaire thanks for the review,
i pushed the LRU caching to unload the model with 0 references/uses automatically.

…sions

The provider now owns each OrtSession and keeps at most max-loaded-models (default 2) loaded per worker, unloading the least recently used idle model. A model is never unloaded while an embedding call uses it, and close() waits for segments still running.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

area/plugin Plugin-related issue or feature request kind/external Pull requests raised by community contributors

Projects

Status: To review

Development

Successfully merging this pull request may close these issues.

Kestra AI Plugin for ONNX embedder

3 participants