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"