From 06b25c0c7f0a20554f818006e8e42c5349c0a6f5 Mon Sep 17 00:00:00 2001 From: Johnny Wilson Dougherty <192861341+JohnnyWilson-Portfolio@users.noreply.github.com> Date: Sun, 13 Sep 2026 22:36:18 +0530 Subject: [PATCH] rag: set embedding server batch and ubatch size to match context size When running `ramalama rag`, `doc2rag` chunks documents using `cl100k_base` with a 400 token limit. However, the embedding model (`embeddinggemma-300m-GGUF`) tokenizes using SentencePiece, which expands tokens on code, URLs, and structured syntax to 500-700+ tokens. `llama-server` defaults `--ubatch-size` to 512, which causes requests with chunks >512 tokens to fail with HTTP 500 (`input (N tokens) is too large to process. increase the physical batch size (current batch size: 512)`). This commit: - Defaults `embed_batch_size` to 2048 (or user-supplied `embed_ctx_size`). - Passes `--batch-size` and `--ubatch-size` matching `embed_batch_size` to the internal embedding server. - Sets `ctx_size=embed_batch_size` to maintain `ctx_size >= batch_size >= ubatch_size`. - Adds unit tests in `test/unit/test_rag_unit.py` verifying both default and custom batch size propagation and command generation. Fixes: #2836 Signed-off-by: Johnny Wilson Dougherty <192861341+JohnnyWilson-Portfolio@users.noreply.github.com> --- .../plugins/runtimes/inference/rag/handler.py | 12 +- test/unit/test_rag_unit.py | 125 +++++++++++++++++- 2 files changed, 133 insertions(+), 4 deletions(-) diff --git a/ramalama/plugins/runtimes/inference/rag/handler.py b/ramalama/plugins/runtimes/inference/rag/handler.py index ad6ef961e..ca5ee9bd3 100644 --- a/ramalama/plugins/runtimes/inference/rag/handler.py +++ b/ramalama/plugins/runtimes/inference/rag/handler.py @@ -67,7 +67,7 @@ def rag_handler(plugin: RuntimePlugin, args: argparse.Namespace) -> None: try: # Build serve args for the VLM and embedding servers vlm_ctx_size = getattr(args, "ctx_size", 8192) - embed_ctx_size = getattr(args, "embed_ctx_size", None) + embed_batch_size = getattr(args, "embed_ctx_size", 0) or 2048 docling_serve_args = _build_serve_args( args, docling_model, docling_port, runtime_args=["--special"], ctx_size=vlm_ctx_size ) @@ -75,8 +75,14 @@ def rag_handler(plugin: RuntimePlugin, args: argparse.Namespace) -> None: args, embedding_model, embed_port, - runtime_args=["--embedding"], - ctx_size=embed_ctx_size, + runtime_args=[ + "--embedding", + "--batch-size", + str(embed_batch_size), + "--ubatch-size", + str(embed_batch_size), + ], + ctx_size=embed_batch_size, cache_reuse=0, ) all_serve_args = [docling_serve_args, embed_serve_args] diff --git a/test/unit/test_rag_unit.py b/test/unit/test_rag_unit.py index 0ec960938..3c5ac80af 100644 --- a/test/unit/test_rag_unit.py +++ b/test/unit/test_rag_unit.py @@ -1,9 +1,11 @@ from argparse import Namespace from pathlib import Path -from unittest.mock import MagicMock +from unittest.mock import MagicMock, patch import pytest +from ramalama.plugins.runtimes.inference.llama_cpp import LlamaCppPlugin +from ramalama.plugins.runtimes.inference.rag.handler import rag_handler from ramalama.rag import RagSource, RagTransport @@ -87,3 +89,124 @@ def test_bare_name_syncs_args_rag(self, tmp_path: Path, force_oci_image: None) - RagTransport(imodel=MagicMock(), cmd=[], args=args) assert args.rag == "localhost/myrag:latest" + + +class TestRagEmbeddingServeArgs: + """Regression tests for issue #2836: ensure embedding server uses + appropriate context and batch sizes (--batch-size, --ubatch-size) + so chunks up to 2048 tokens (or user-specified embed_ctx_size) are accepted.""" + + @patch("ramalama.plugins.loader.get_runtime", return_value=LlamaCppPlugin()) + @patch("ramalama.rag.Rag") + @patch("ramalama.plugins.runtimes.inference.rag.handler.New") + @patch("ramalama.plugins.runtimes.inference.rag.handler.compute_serving_port", side_effect=[8080, 8081]) + @patch("ramalama.plugins.runtimes.inference.rag.handler.set_accel_env_vars") + def test_default_embed_batch_size_2048( + self, + mock_accel: MagicMock, + mock_port: MagicMock, + mock_new: MagicMock, + mock_rag: MagicMock, + mock_runtime: MagicMock, + ) -> None: + mock_transport = MagicMock() + mock_proc = MagicMock() + mock_proc.poll.return_value = 0 + mock_transport.serve_nonblocking.return_value = mock_proc + mock_new.return_value = mock_transport + + args = Namespace( + container=True, + engine="podman", + store="/tmp/store", + dryrun=True, + debug=False, + image="rag-image", + rag_image="rag-image", + DOCUMENTS=["doc.pdf"], + DESTINATION="myrag", + embed_ctx_size=0, + runtime="llama.cpp", + subcommand="serve", + ) + plugin = MagicMock() + rag_handler(plugin, args) + + assert mock_new.call_count >= 2 + embed_call = mock_new.call_args_list[1] + _, embed_serve_args = embed_call[0] + + assert embed_serve_args.ctx_size == 2048 + assert "--embedding" in embed_serve_args.runtime_args + assert "--batch-size" in embed_serve_args.runtime_args + assert "--ubatch-size" in embed_serve_args.runtime_args + + idx_batch = embed_serve_args.runtime_args.index("--batch-size") + assert embed_serve_args.runtime_args[idx_batch + 1] == "2048" + + idx_ubatch = embed_serve_args.runtime_args.index("--ubatch-size") + assert embed_serve_args.runtime_args[idx_ubatch + 1] == "2048" + + embed_call_serve = mock_transport.serve_nonblocking.call_args_list[1] + _, called_cmd = embed_call_serve[0] + assert "--ctx-size" in called_cmd + assert called_cmd[called_cmd.index("--ctx-size") + 1] == "2048" + assert "--batch-size" in called_cmd + assert called_cmd[called_cmd.index("--batch-size") + 1] == "2048" + assert "--ubatch-size" in called_cmd + assert called_cmd[called_cmd.index("--ubatch-size") + 1] == "2048" + + @patch("ramalama.plugins.loader.get_runtime", return_value=LlamaCppPlugin()) + @patch("ramalama.rag.Rag") + @patch("ramalama.plugins.runtimes.inference.rag.handler.New") + @patch("ramalama.plugins.runtimes.inference.rag.handler.compute_serving_port", side_effect=[8080, 8081]) + @patch("ramalama.plugins.runtimes.inference.rag.handler.set_accel_env_vars") + def test_custom_embed_ctx_size_propagates( + self, + mock_accel: MagicMock, + mock_port: MagicMock, + mock_new: MagicMock, + mock_rag: MagicMock, + mock_runtime: MagicMock, + ) -> None: + mock_transport = MagicMock() + mock_proc = MagicMock() + mock_proc.poll.return_value = 0 + mock_transport.serve_nonblocking.return_value = mock_proc + mock_new.return_value = mock_transport + + args = Namespace( + container=True, + engine="podman", + store="/tmp/store", + dryrun=True, + debug=False, + image="rag-image", + rag_image="rag-image", + DOCUMENTS=["doc.pdf"], + DESTINATION="myrag", + embed_ctx_size=4096, + runtime="llama.cpp", + subcommand="serve", + ) + plugin = MagicMock() + rag_handler(plugin, args) + + assert mock_new.call_count >= 2 + embed_call = mock_new.call_args_list[1] + _, embed_serve_args = embed_call[0] + + assert embed_serve_args.ctx_size == 4096 + idx_batch = embed_serve_args.runtime_args.index("--batch-size") + assert embed_serve_args.runtime_args[idx_batch + 1] == "4096" + idx_ubatch = embed_serve_args.runtime_args.index("--ubatch-size") + assert embed_serve_args.runtime_args[idx_ubatch + 1] == "4096" + + embed_call_serve = mock_transport.serve_nonblocking.call_args_list[1] + _, called_cmd = embed_call_serve[0] + assert "--ctx-size" in called_cmd + assert called_cmd[called_cmd.index("--ctx-size") + 1] == "4096" + assert "--batch-size" in called_cmd + assert called_cmd[called_cmd.index("--batch-size") + 1] == "4096" + assert "--ubatch-size" in called_cmd + assert called_cmd[called_cmd.index("--ubatch-size") + 1] == "4096"