Support protocol thread configuration

This commit is contained in:
2026-09-12 11:40:28 +02:00
parent 21082e66b3
commit 9177a2660b
6 changed files with 99 additions and 0 deletions
+4
View File
@@ -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 input for the unchanged Working Protocol renderer and compare that output
against the Working Protocol Synthesizer V0 baseline and the human reference against the Working Protocol Synthesizer V0 baseline and the human reference
protocol. 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.
+5
View File
@@ -62,6 +62,7 @@ def generate_once(
timeout: int, timeout: int,
num_ctx: int, num_ctx: int,
num_predict: int, num_predict: int,
num_thread: int | None = None,
) -> OllamaGeneration: ) -> OllamaGeneration:
payload = { payload = {
"model": model, "model": model,
@@ -74,6 +75,10 @@ def generate_once(
"num_predict": num_predict, "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() started = time.perf_counter()
try: try:
response = requests.post(generate_url(endpoint), json=payload, timeout=timeout) response = requests.post(generate_url(endpoint), json=payload, timeout=timeout)
+12
View File
@@ -56,6 +56,7 @@ class MvpMeetingConfig:
threads: str | int = "auto" threads: str | int = "auto"
model: str = DEFAULT_MODEL model: str = DEFAULT_MODEL
ollama_endpoint: str = DEFAULT_ENDPOINT ollama_endpoint: str = DEFAULT_ENDPOINT
protocol_num_thread: int | None = None
protocol_num_ctx: int = DEFAULT_NUM_CTX protocol_num_ctx: int = DEFAULT_NUM_CTX
protocol_safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET protocol_safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET
diarization: str = "off" diarization: str = "off"
@@ -77,6 +78,7 @@ def regenerate_mvp_protocol(
meeting_context: ContextInput, meeting_context: ContextInput,
model: str = DEFAULT_MODEL, model: str = DEFAULT_MODEL,
ollama_endpoint: str = DEFAULT_ENDPOINT, ollama_endpoint: str = DEFAULT_ENDPOINT,
protocol_num_thread: int | None = None,
protocol_num_ctx: int = DEFAULT_NUM_CTX, protocol_num_ctx: int = DEFAULT_NUM_CTX,
protocol_safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET, protocol_safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET,
progress_sink: ProgressSink | None = None, progress_sink: ProgressSink | None = None,
@@ -87,6 +89,10 @@ def regenerate_mvp_protocol(
context = _effective_context(meeting_context) context = _effective_context(meeting_context)
if context is None: if context is None:
raise ValueError("Meeting Context is required for protocol regeneration.") 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: if protocol_num_ctx <= 0:
raise ValueError("Protocol Ollama context size must be positive.") raise ValueError("Protocol Ollama context size must be positive.")
if protocol_safe_input_token_budget <= 0: if protocol_safe_input_token_budget <= 0:
@@ -112,6 +118,7 @@ def regenerate_mvp_protocol(
model=model, model=model,
endpoint=ollama_endpoint, endpoint=ollama_endpoint,
num_ctx=protocol_num_ctx, num_ctx=protocol_num_ctx,
num_thread=protocol_num_thread,
safe_input_token_budget=protocol_safe_input_token_budget, safe_input_token_budget=protocol_safe_input_token_budget,
) )
protocol_path = _persist_protocol(run_dir, result) protocol_path = _persist_protocol(run_dir, result)
@@ -191,6 +198,10 @@ def _validate_inputs(
raise ValueError("A diarization container image is required.") raise ValueError("A diarization container image is required.")
if config.protocol_safe_input_token_budget <= 0: if config.protocol_safe_input_token_budget <= 0:
raise ValueError("Protocol safe input token budget must be positive.") 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: if config.protocol_num_ctx <= 0:
raise ValueError("Protocol Ollama context size must be positive.") raise ValueError("Protocol Ollama context size must be positive.")
@@ -426,6 +437,7 @@ def run_mvp_meeting(
model=config.model, model=config.model,
endpoint=config.ollama_endpoint, endpoint=config.ollama_endpoint,
num_ctx=config.protocol_num_ctx, num_ctx=config.protocol_num_ctx,
num_thread=config.protocol_num_thread,
safe_input_token_budget=config.protocol_safe_input_token_budget, safe_input_token_budget=config.protocol_safe_input_token_budget,
) )
stage_runtimes["protocol"] = round(time.perf_counter() - stage_started, 3) stage_runtimes["protocol"] = round(time.perf_counter() - stage_started, 3)
@@ -170,6 +170,7 @@ def generate_direct_protocol(
timeout: int = DEFAULT_TIMEOUT, timeout: int = DEFAULT_TIMEOUT,
num_ctx: int = DEFAULT_NUM_CTX, num_ctx: int = DEFAULT_NUM_CTX,
num_predict: int = DEFAULT_NUM_PREDICT, num_predict: int = DEFAULT_NUM_PREDICT,
num_thread: int | None = None,
safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET, safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET,
model_check: Callable[[str, str, int], dict[str, Any]] = require_model, model_check: Callable[[str, str, int], dict[str, Any]] = require_model,
generation_call: Callable[..., OllamaGeneration] = generate_once, generation_call: Callable[..., OllamaGeneration] = generate_once,
@@ -198,6 +199,7 @@ def generate_direct_protocol(
timeout=timeout, timeout=timeout,
num_ctx=num_ctx, num_ctx=num_ctx,
num_predict=num_predict, num_predict=num_predict,
num_thread=num_thread,
) )
data = generation.raw_response data = generation.raw_response
runtime_metadata = { runtime_metadata = {
@@ -216,6 +218,7 @@ def generate_direct_protocol(
"think": False, "think": False,
"num_ctx": num_ctx, "num_ctx": num_ctx,
"num_predict": num_predict, "num_predict": num_predict,
"num_thread": num_thread,
"selected_transcript_representation": selected.representation, "selected_transcript_representation": selected.representation,
"estimated_input_tokens": selected.estimated_input_tokens, "estimated_input_tokens": selected.estimated_input_tokens,
"safe_input_token_budget": selected.safe_input_token_budget, "safe_input_token_budget": selected.safe_input_token_budget,
+17
View File
@@ -198,6 +198,23 @@ class OllamaTests(unittest.TestCase):
self.assertFalse(payload["stream"]) self.assertFalse(payload["stream"])
self.assertEqual(result.raw_response, raw) 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: def test_malformed_response_failure_without_retry(self) -> None:
response = Mock() response = Mock()
response.raise_for_status.return_value = None response.raise_for_status.return_value = None
+58
View File
@@ -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()