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