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
+5
View File
@@ -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)
+12
View File
@@ -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,