Support protocol thread configuration
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user