feat: configure 32k context for protocol generation
This commit is contained in:
+3
-3
@@ -48,9 +48,9 @@ derived representation sent to prompt construction.
|
|||||||
|
|
||||||
Before contacting Ollama, Meeting Lab conservatively estimates prompt tokens
|
Before contacting Ollama, Meeting Lab conservatively estimates prompt tokens
|
||||||
from UTF-8 byte count without adding a model tokenizer dependency. The default safe
|
from UTF-8 byte count without adding a model tokenizer dependency. The default safe
|
||||||
budget is 16,200 estimated tokens, below the observed 16,386-token effective
|
budget is 29,000 estimated tokens within the explicitly configured 32,768-token
|
||||||
boundary even when a larger `num_ctx` was requested. The estimate is calibrated
|
Ollama context. The estimate is calibrated against the currently validated
|
||||||
against the currently validated German BPD input and is configurable through
|
German BPD input and is configurable through
|
||||||
`MvpMeetingConfig.protocol_safe_input_token_budget` or
|
`MvpMeetingConfig.protocol_safe_input_token_budget` or
|
||||||
`--protocol-safe-input-token-budget`.
|
`--protocol-safe-input-token-budget`.
|
||||||
|
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ from src.meeting_lab.models.meeting_context import (
|
|||||||
from src.meeting_lab.progress import ProgressEvent, ProgressSink, ProgressStatus
|
from src.meeting_lab.progress import ProgressEvent, ProgressSink, ProgressStatus
|
||||||
from src.meeting_lab.protocol.generate_direct_protocol import (
|
from src.meeting_lab.protocol.generate_direct_protocol import (
|
||||||
DEFAULT_MODEL,
|
DEFAULT_MODEL,
|
||||||
|
DEFAULT_NUM_CTX,
|
||||||
DEFAULT_SAFE_INPUT_TOKEN_BUDGET,
|
DEFAULT_SAFE_INPUT_TOKEN_BUDGET,
|
||||||
DirectProtocolResult,
|
DirectProtocolResult,
|
||||||
generate_direct_protocol,
|
generate_direct_protocol,
|
||||||
@@ -55,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_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"
|
||||||
diarization_runtime: str = "native"
|
diarization_runtime: str = "native"
|
||||||
@@ -131,6 +133,8 @@ 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_ctx <= 0:
|
||||||
|
raise ValueError("Protocol Ollama context size must be positive.")
|
||||||
|
|
||||||
|
|
||||||
def _emit(
|
def _emit(
|
||||||
@@ -363,6 +367,7 @@ def run_mvp_meeting(
|
|||||||
preserved_context,
|
preserved_context,
|
||||||
model=config.model,
|
model=config.model,
|
||||||
endpoint=config.ollama_endpoint,
|
endpoint=config.ollama_endpoint,
|
||||||
|
num_ctx=config.protocol_num_ctx,
|
||||||
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)
|
||||||
|
|||||||
@@ -33,7 +33,7 @@ DEFAULT_MODEL = "qwen3.6:35B-A3B"
|
|||||||
DEFAULT_NUM_CTX = 32768
|
DEFAULT_NUM_CTX = 32768
|
||||||
DEFAULT_NUM_PREDICT = 8192
|
DEFAULT_NUM_PREDICT = 8192
|
||||||
DEFAULT_TIMEOUT = 1800
|
DEFAULT_TIMEOUT = 1800
|
||||||
DEFAULT_SAFE_INPUT_TOKEN_BUDGET = 16_200
|
DEFAULT_SAFE_INPUT_TOKEN_BUDGET = 29_000
|
||||||
ESTIMATED_UTF8_BYTES_PER_TOKEN = 4.4
|
ESTIMATED_UTF8_BYTES_PER_TOKEN = 4.4
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -116,6 +116,7 @@ class GeneratorTests(unittest.TestCase):
|
|||||||
self.assertEqual(check.call_count, 1)
|
self.assertEqual(check.call_count, 1)
|
||||||
self.assertEqual(call.call_count, 1)
|
self.assertEqual(call.call_count, 1)
|
||||||
self.assertEqual(call.call_args.args[1], "qwen3.6:35B-A3B")
|
self.assertEqual(call.call_args.args[1], "qwen3.6:35B-A3B")
|
||||||
|
self.assertEqual(call.call_args.kwargs["num_ctx"], 32768)
|
||||||
self.assertEqual(result.runtime_metadata["request_count"], 1)
|
self.assertEqual(result.runtime_metadata["request_count"], 1)
|
||||||
self.assertEqual(result.runtime_metadata["prompt_token_count"], 123)
|
self.assertEqual(result.runtime_metadata["prompt_token_count"], 123)
|
||||||
self.assertFalse(result.runtime_metadata["think"])
|
self.assertFalse(result.runtime_metadata["think"])
|
||||||
@@ -181,7 +182,7 @@ class OllamaTests(unittest.TestCase):
|
|||||||
with patch.object(ollama.requests, "post", return_value=response) as post:
|
with patch.object(ollama.requests, "post", return_value=response) as post:
|
||||||
result = ollama.generate_once(
|
result = ollama.generate_once(
|
||||||
"http://localhost:11434",
|
"http://localhost:11434",
|
||||||
"chosen:model",
|
"qwen3.8:27b",
|
||||||
"prompt",
|
"prompt",
|
||||||
timeout=30,
|
timeout=30,
|
||||||
num_ctx=32768,
|
num_ctx=32768,
|
||||||
@@ -190,8 +191,9 @@ class OllamaTests(unittest.TestCase):
|
|||||||
|
|
||||||
self.assertEqual(post.call_count, 1)
|
self.assertEqual(post.call_count, 1)
|
||||||
payload = post.call_args.kwargs["json"]
|
payload = post.call_args.kwargs["json"]
|
||||||
self.assertEqual(payload["model"], "chosen:model")
|
self.assertEqual(payload["model"], "qwen3.8:27b")
|
||||||
self.assertEqual(payload["options"]["temperature"], 0.0)
|
self.assertEqual(payload["options"]["temperature"], 0.0)
|
||||||
|
self.assertEqual(payload["options"]["num_ctx"], 32768)
|
||||||
self.assertFalse(payload["think"])
|
self.assertFalse(payload["think"])
|
||||||
self.assertFalse(payload["stream"])
|
self.assertFalse(payload["stream"])
|
||||||
self.assertEqual(result.raw_response, raw)
|
self.assertEqual(result.raw_response, raw)
|
||||||
|
|||||||
@@ -115,7 +115,7 @@ class MvpApiTests(unittest.TestCase):
|
|||||||
patch.object(mvp_api, "prepare_audio", side_effect=fake_prepare),
|
patch.object(mvp_api, "prepare_audio", side_effect=fake_prepare),
|
||||||
patch.object(
|
patch.object(
|
||||||
mvp_api, "generate_direct_protocol", side_effect=fake_protocol
|
mvp_api, "generate_direct_protocol", side_effect=fake_protocol
|
||||||
),
|
) as protocol_generator,
|
||||||
patch.object(subprocess, "run") as subprocess_run,
|
patch.object(subprocess, "run") as subprocess_run,
|
||||||
):
|
):
|
||||||
result = mvp_api.run_mvp_meeting(
|
result = mvp_api.run_mvp_meeting(
|
||||||
@@ -180,7 +180,8 @@ class MvpApiTests(unittest.TestCase):
|
|||||||
self.assertEqual(delegated.whisper_executable, "whisper-cli")
|
self.assertEqual(delegated.whisper_executable, "whisper-cli")
|
||||||
self.assertEqual(delegated.ffmpeg_executable, "ffmpeg")
|
self.assertEqual(delegated.ffmpeg_executable, "ffmpeg")
|
||||||
self.assertTrue(delegated.audio_normalization)
|
self.assertTrue(delegated.audio_normalization)
|
||||||
self.assertEqual(delegated.protocol_safe_input_token_budget, 16_200)
|
self.assertEqual(delegated.protocol_num_ctx, 32_768)
|
||||||
|
self.assertEqual(delegated.protocol_safe_input_token_budget, 29_000)
|
||||||
self.assertEqual(api.call_args.kwargs["meeting_context"], context_data())
|
self.assertEqual(api.call_args.kwargs["meeting_context"], context_data())
|
||||||
|
|
||||||
def test_cli_explicit_audio_normalization_values_are_propagated(self):
|
def test_cli_explicit_audio_normalization_values_are_propagated(self):
|
||||||
|
|||||||
Reference in New Issue
Block a user