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
|
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.
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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