diff --git a/PROJECT_KNOWLEDGE.md b/PROJECT_KNOWLEDGE.md index d809f8f..3ef2413 100644 --- a/PROJECT_KNOWLEDGE.md +++ b/PROJECT_KNOWLEDGE.md @@ -295,3 +295,7 @@ The next recommended evaluation step is to use the consolidated V0 result as input for the unchanged Working Protocol renderer and compare that output against the Working Protocol Synthesizer V0 baseline and the human reference protocol. + +Protocol calls accept an optional positive `protocol_num_thread` setting through +both initial processing and regeneration. With `None`, Ollama receives no +`num_thread` override and selects its own thread configuration. diff --git a/src/meeting_lab/llm/ollama.py b/src/meeting_lab/llm/ollama.py index 6c860d4..9c2ef0a 100644 --- a/src/meeting_lab/llm/ollama.py +++ b/src/meeting_lab/llm/ollama.py @@ -62,6 +62,7 @@ def generate_once( timeout: int, num_ctx: int, num_predict: int, + num_thread: int | None = None, ) -> OllamaGeneration: payload = { "model": model, @@ -74,6 +75,10 @@ def generate_once( "num_predict": num_predict, }, } + if num_thread is not None: + if type(num_thread) is not int or num_thread <= 0: + raise ValueError("Protocol thread count must be a positive integer.") + payload["options"]["num_thread"] = num_thread started = time.perf_counter() try: response = requests.post(generate_url(endpoint), json=payload, timeout=timeout) diff --git a/src/meeting_lab/orchestration/mvp.py b/src/meeting_lab/orchestration/mvp.py index 3356fee..e036a6f 100644 --- a/src/meeting_lab/orchestration/mvp.py +++ b/src/meeting_lab/orchestration/mvp.py @@ -56,6 +56,7 @@ class MvpMeetingConfig: threads: str | int = "auto" model: str = DEFAULT_MODEL ollama_endpoint: str = DEFAULT_ENDPOINT + protocol_num_thread: int | None = None protocol_num_ctx: int = DEFAULT_NUM_CTX protocol_safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET diarization: str = "off" @@ -77,6 +78,7 @@ def regenerate_mvp_protocol( meeting_context: ContextInput, model: str = DEFAULT_MODEL, ollama_endpoint: str = DEFAULT_ENDPOINT, + protocol_num_thread: int | None = None, protocol_num_ctx: int = DEFAULT_NUM_CTX, protocol_safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET, progress_sink: ProgressSink | None = None, @@ -87,6 +89,10 @@ def regenerate_mvp_protocol( context = _effective_context(meeting_context) if context is None: raise ValueError("Meeting Context is required for protocol regeneration.") + if protocol_num_thread is not None and ( + type(protocol_num_thread) is not int or protocol_num_thread <= 0 + ): + raise ValueError("Protocol Ollama thread count must be a positive integer.") if protocol_num_ctx <= 0: raise ValueError("Protocol Ollama context size must be positive.") if protocol_safe_input_token_budget <= 0: @@ -112,6 +118,7 @@ def regenerate_mvp_protocol( model=model, endpoint=ollama_endpoint, num_ctx=protocol_num_ctx, + num_thread=protocol_num_thread, safe_input_token_budget=protocol_safe_input_token_budget, ) protocol_path = _persist_protocol(run_dir, result) @@ -191,6 +198,10 @@ def _validate_inputs( raise ValueError("A diarization container image is required.") if config.protocol_safe_input_token_budget <= 0: raise ValueError("Protocol safe input token budget must be positive.") + if config.protocol_num_thread is not None and ( + type(config.protocol_num_thread) is not int or config.protocol_num_thread <= 0 + ): + raise ValueError("Protocol Ollama thread count must be a positive integer.") if config.protocol_num_ctx <= 0: raise ValueError("Protocol Ollama context size must be positive.") @@ -426,6 +437,7 @@ def run_mvp_meeting( model=config.model, endpoint=config.ollama_endpoint, num_ctx=config.protocol_num_ctx, + num_thread=config.protocol_num_thread, safe_input_token_budget=config.protocol_safe_input_token_budget, ) stage_runtimes["protocol"] = round(time.perf_counter() - stage_started, 3) diff --git a/src/meeting_lab/protocol/generate_direct_protocol.py b/src/meeting_lab/protocol/generate_direct_protocol.py index 7095fb2..7cf8014 100644 --- a/src/meeting_lab/protocol/generate_direct_protocol.py +++ b/src/meeting_lab/protocol/generate_direct_protocol.py @@ -170,6 +170,7 @@ def generate_direct_protocol( timeout: int = DEFAULT_TIMEOUT, num_ctx: int = DEFAULT_NUM_CTX, num_predict: int = DEFAULT_NUM_PREDICT, + num_thread: int | None = None, safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET, model_check: Callable[[str, str, int], dict[str, Any]] = require_model, generation_call: Callable[..., OllamaGeneration] = generate_once, @@ -198,6 +199,7 @@ def generate_direct_protocol( timeout=timeout, num_ctx=num_ctx, num_predict=num_predict, + num_thread=num_thread, ) data = generation.raw_response runtime_metadata = { @@ -216,6 +218,7 @@ def generate_direct_protocol( "think": False, "num_ctx": num_ctx, "num_predict": num_predict, + "num_thread": num_thread, "selected_transcript_representation": selected.representation, "estimated_input_tokens": selected.estimated_input_tokens, "safe_input_token_budget": selected.safe_input_token_budget, diff --git a/tests/test_direct_protocol.py b/tests/test_direct_protocol.py index fa91d93..850ba64 100644 --- a/tests/test_direct_protocol.py +++ b/tests/test_direct_protocol.py @@ -198,6 +198,23 @@ class OllamaTests(unittest.TestCase): self.assertFalse(payload["stream"]) self.assertEqual(result.raw_response, raw) + def test_request_omits_num_thread_when_backend_selection_is_requested(self) -> None: + response = Mock() + response.raise_for_status.return_value = None + response.json.return_value = {"response": "# Meeting Protocol"} + with patch.object(ollama.requests, "post", return_value=response) as post: + ollama.generate_once( + "http://localhost:11434", + "qwen3.8:27b", + "prompt", + timeout=30, + num_ctx=32768, + num_predict=8192, + ) + + payload = post.call_args.kwargs["json"] + self.assertNotIn("num_thread", payload["options"]) + def test_malformed_response_failure_without_retry(self) -> None: response = Mock() response.raise_for_status.return_value = None diff --git a/tests/test_protocol_threads.py b/tests/test_protocol_threads.py new file mode 100644 index 0000000..4e29c1e --- /dev/null +++ b/tests/test_protocol_threads.py @@ -0,0 +1,58 @@ +"""Protocol thread configuration through the public pipeline boundaries.""" + +import tempfile +import unittest +from dataclasses import replace +from pathlib import Path +from unittest.mock import Mock, patch + +from src.meeting_lab.llm import ollama +from src.meeting_lab.orchestration import mvp +from tests import test_mvp_api as fixtures + + +class ProtocolThreadTests(unittest.TestCase): + def test_initial_and_regeneration_forward_default_and_explicit_threads(self): + for threads in (None, 4, 10, 16): + with self.subTest(threads=threads), tempfile.TemporaryDirectory() as directory: + config = replace( + fixtures.MvpApiTests().config(Path(directory)), protocol_num_thread=threads + ) + with ( + patch.object(mvp, "prepare_audio", side_effect=fixtures.fake_prepare), + patch.object(mvp, "transcribe_audio", side_effect=fixtures.fake_transcribe), + patch.object( + mvp, "generate_direct_protocol", side_effect=fixtures.fake_protocol + ) as generate, + ): + result = mvp.run_mvp_meeting(config, meeting_context=fixtures.context_data()) + self.assertEqual(result.exit_code, 0) + self.assertEqual(generate.call_args.kwargs["num_thread"], threads) + mvp.regenerate_mvp_protocol( + result.run_dir, + meeting_context=fixtures.context_data(), + protocol_num_thread=threads, + ) + self.assertEqual(generate.call_args.kwargs["num_thread"], threads) + + def test_explicit_payload_and_invalid_values(self): + response = Mock() + response.json.return_value = {"response": "protocol"} + with patch.object(ollama.requests, "post", return_value=response) as post: + ollama.generate_once( + "url", "model", "prompt", timeout=1, num_ctx=32, num_predict=8, num_thread=10 + ) + self.assertEqual(post.call_args.kwargs["json"]["options"]["num_thread"], 10) + post.reset_mock() + for value in (0, -1, True, 1.5): + with self.assertRaises(ValueError): + ollama.generate_once( + "url", + "model", + "prompt", + timeout=1, + num_ctx=32, + num_predict=8, + num_thread=value, + ) + post.assert_not_called()