diff --git a/bindings/python/README.md b/bindings/python/README.md index 094bb9cb7..c24be3453 100644 --- a/bindings/python/README.md +++ b/bindings/python/README.md @@ -43,6 +43,17 @@ and the one-shot `transcribe()` helper. result = session.run(pcm, pnc="off", itn="on") ``` +### Prompting + +`vocabulary` (custom terms), `prompt` (context, or the instruction under +`task="instruct"`) and `prefix` (text the model continues from) take effect +where `model.supports()` reports `"vocabulary"`, `"context_prompt"`, +`"instruct"` or `"transcript_prefix"`. + +```python +result = session.run(pcm, vocabulary=["Kubernetes", "gRPC"]) +``` + Streaming models expose incremental transcription with committed/tentative text views — see `examples/stream_wav.py`: diff --git a/bindings/python/src/transcribe_cpp/__init__.py b/bindings/python/src/transcribe_cpp/__init__.py index 3187ee245..5e655b550 100644 --- a/bindings/python/src/transcribe_cpp/__init__.py +++ b/bindings/python/src/transcribe_cpp/__init__.py @@ -52,7 +52,7 @@ # String-enum types, exported so callers (and type checkers) can name them. Backend = Literal["auto", "cpu", "metal", "vulkan", "cpu_accel", "cuda", "rocm"] KVType = Literal["auto", "f32", "f16"] -Task = Literal["transcribe", "translate"] +Task = Literal["transcribe", "translate", "instruct"] Timestamps = Literal["none", "auto", "segment", "word", "token"] Pnc = Literal["default", "off", "on"] Itn = Literal["default", "off", "on"] @@ -62,6 +62,7 @@ Feature = Literal[ "initial_prompt", "temperature_fallback", "long_form", "cancellation", "pnc", "itn", "diarization", + "vocabulary", "context_prompt", "instruct", "transcript_prefix", ] __all__ = [ @@ -203,6 +204,7 @@ _TASKS = { "transcribe": _generated.TRANSCRIBE_TASK_TRANSCRIBE, "translate": _generated.TRANSCRIBE_TASK_TRANSLATE, + "instruct": _generated.TRANSCRIBE_TASK_INSTRUCT, } _TIMESTAMPS = { "none": _generated.TRANSCRIBE_TIMESTAMPS_NONE, @@ -250,6 +252,10 @@ "pnc": _generated.TRANSCRIBE_FEATURE_PNC, "itn": _generated.TRANSCRIBE_FEATURE_ITN, "diarization": _generated.TRANSCRIBE_FEATURE_DIARIZATION, + "vocabulary": _generated.TRANSCRIBE_FEATURE_VOCABULARY, + "context_prompt": _generated.TRANSCRIBE_FEATURE_CONTEXT_PROMPT, + "instruct": _generated.TRANSCRIBE_FEATURE_INSTRUCT, + "transcript_prefix": _generated.TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX, } @@ -653,9 +659,17 @@ def _stream_update_from(u) -> StreamUpdate: ) +def _cstr(value: str, name: str) -> bytes: + """UTF-8 bytes for a C string option; a NUL would silently cut it in C.""" + if "\x00" in value: + raise InvalidArgument(f"{name} contains a NUL character") + return value.encode("utf-8") + + def _build_run_params(task, language, target_language, timestamps, keep_special_tags, spec_k_drafts, diarize="default", - pnc="default", itn="default"): + pnc="default", itn="default", vocabulary=None, + prompt=None, prefix=None): if not isinstance(spec_k_drafts, int) or spec_k_drafts < -1: raise InvalidArgument( f"spec_k_drafts must be -1 (family default), 0 (disabled), or a " @@ -668,10 +682,27 @@ def _build_run_params(task, language, target_language, timestamps, params.pnc = _enum(_PNC, pnc, "pnc") params.itn = _enum(_ITN, itn, "itn") params.diarize = _enum(_DIARIZE, diarize, "diarize") - params.language = language.encode("utf-8") if language else None - params.target_language = target_language.encode("utf-8") if target_language else None + params.language = _cstr(language, "language") if language else None + params.target_language = _cstr(target_language, "target_language") if target_language else None params.keep_special_tags = keep_special_tags params.spec_k_drafts = spec_k_drafts + if vocabulary is not None: + if isinstance(vocabulary, (str, bytes)): + raise InvalidArgument("vocabulary must be a sequence of terms, not a single string") + if isinstance(vocabulary, (set, frozenset, dict)): + raise InvalidArgument("vocabulary is in priority order; pass a list, not a set or dict") + vocabulary = list(vocabulary) # an iterator would be used up by the check below + if not all(isinstance(t, str) for t in vocabulary): + raise InvalidArgument("vocabulary terms must be strings") + terms = [_cstr(t, "vocabulary") for t in vocabulary] + if terms: + arr = (ctypes.c_char_p * len(terms))(*terms) + params.vocabulary = ctypes.cast(arr, ctypes.POINTER(ctypes.c_char_p)) + params.n_vocabulary = len(terms) + # The C struct holds raw pointers; keep the buffers alive with it. + params._prompting_keepalive = (arr, terms) + params.prompt = _cstr(prompt, "prompt") if prompt else None + params.prefix = _cstr(prefix, "prefix") if prefix else None return params @@ -740,7 +771,7 @@ def __init__(self, *, initial_prompt: str | None = None, def _apply(self, ext) -> None: if self.initial_prompt is not None: - ext.initial_prompt = self.initial_prompt.encode("utf-8") + ext.initial_prompt = _cstr(self.initial_prompt, "initial_prompt") if self.condition_on_prev_tokens is not None: ext.condition_on_prev_tokens = self.condition_on_prev_tokens if self.temperature is not None: @@ -981,8 +1012,7 @@ def capabilities(self) -> Capabilities: ) def supports(self, feature: Feature) -> bool: - """Whether the model exposes a behavioral feature (initial prompt, - temperature fallback, long-form, cancellation, pnc, itn).""" + """Whether the model exposes a behavioral feature (see ``Feature``).""" return bool(_lib.transcribe_model_supports( self._h, _enum(_FEATURES, feature, "feature"))) @@ -1103,7 +1133,10 @@ def run(self, pcm: PCMLike, *, task: Task = "transcribe", diarize: Diarize = "default", keep_special_tags: bool = False, spec_k_drafts: int = -1, - family: FamilyExtension | None = None) -> Result: + family: FamilyExtension | None = None, + vocabulary: Sequence[str] | None = None, + prompt: str | None = None, + prefix: str | None = None) -> Result: """Transcribe 16 kHz mono float32 PCM and return a materialized Result. ``pnc`` controls punctuation/capitalization and ``itn`` controls @@ -1113,6 +1146,14 @@ def run(self, pcm: PCMLike, *, task: Task = "transcribe", ``spec_k_drafts`` tunes speculative decoding on models whose capabilities advertise ``supports_spec_decode`` (-1 = family default, 0 = disabled, >0 = draft length; silently ignored elsewhere). + ``vocabulary`` (terms, priority order), ``prompt`` and ``prefix`` are + the generic prompting inputs; probe ``model.supports()`` for + ``"vocabulary"``, ``"context_prompt"``, ``"instruct"`` and + ``"transcript_prefix"``. With ``task="instruct"`` the ``prompt`` is the + required instruction and the output is free text. Unsupported + vocabulary/context is ignored with a warning; an unsupported prefix + or instruct task raises. With ``prefix``, ``text`` holds only the + continuation and ``raw_text`` leads with the prefix. On ``Aborted`` (via :meth:`cancel`) and ``OutputTruncated`` (including its ``OutputRepetition`` subclass) the partial transcript is preserved @@ -1120,7 +1161,8 @@ def run(self, pcm: PCMLike, *, task: Task = "transcribe", self._cancel.clear() array, n_samples = _pcm_to_carray(pcm) params = _build_run_params(task, language, target_language, timestamps, - keep_special_tags, spec_k_drafts, diarize, pnc, itn) + keep_special_tags, spec_k_drafts, diarize, pnc, itn, + vocabulary, prompt, prefix) ext = self._resolve_family(family, "run") if family is not None else None if ext is not None: params.family = ctypes.cast( @@ -1146,7 +1188,9 @@ def run_batch(self, pcms: Sequence[PCMLike], *, task: Task = "transcribe", keep_special_tags: bool = False, spec_k_drafts: int = -1, family: FamilyExtension | None = None, - return_exceptions: bool = False) -> list[Result | TranscribeError]: + return_exceptions: bool = False, + vocabulary: Sequence[str] | None = None, + prompt: str | None = None) -> list[Result | TranscribeError]: """Transcribe several utterances in one dispatch — one Result each. Families with a batched compute path process every utterance in a single @@ -1163,7 +1207,10 @@ def run_batch(self, pcms: Sequence[PCMLike], *, task: Task = "transcribe", view (``Result`` or ``TranscribeError`` each) so completed work is never discarded. With ``return_exceptions=True`` no exception is raised for utterance failures and that mixed list is returned - directly (the ``asyncio.gather`` convention).""" + directly (the ``asyncio.gather`` convention). + + ``vocabulary`` / ``prompt`` apply to every utterance (see :meth:`run`); + a transcript prefix is per-utterance and so is not accepted here.""" self._cancel.clear() pcms = list(pcms) if not pcms: @@ -1179,7 +1226,8 @@ def run_batch(self, pcms: Sequence[PCMLike], *, task: Task = "transcribe", counts[k] = n params = _build_run_params(task, language, target_language, timestamps, - keep_special_tags, spec_k_drafts, diarize, pnc, itn) + keep_special_tags, spec_k_drafts, diarize, pnc, itn, + vocabulary, prompt) ext = self._resolve_family(family, "run") if family is not None else None if ext is not None: params.family = ctypes.cast( @@ -1234,19 +1282,24 @@ def stream(self, *, task: Task = "transcribe", language: str | None = None, diarize: Diarize = "default", keep_special_tags: bool = False, commit_policy: CommitPolicy = "auto", stable_prefix_agreement_n: int = 0, - family: FamilyExtension | None = None) -> Stream: + family: FamilyExtension | None = None, + vocabulary: Sequence[str] | None = None, + prompt: str | None = None) -> Stream: """Begin streaming on this session and return a Stream to feed audio to. Requires a model whose capabilities advertise ``supports_streaming``; otherwise raises NotImplementedByModel. ``family`` is an optional family-specific stream extension (e.g. MoonshineStreamingOptions). The session is single-threaded and runs at most one stream at a time. Use - the Stream as a context manager so it is reset when you are done.""" + the Stream as a context manager so it is reset when you are done. + ``vocabulary`` and ``prompt`` are as in :meth:`run`; ``task="instruct"`` + raises.""" self._cancel.clear() # spec_k_drafts is an offline-decode knob; streaming always uses the # family default (-1). run_params = _build_run_params(task, language, target_language, timestamps, - keep_special_tags, -1, diarize, pnc, itn) + keep_special_tags, -1, diarize, pnc, itn, + vocabulary, prompt) sp = _StreamParams() _lib.transcribe_stream_params_init(_byref(sp)) sp.commit_policy = _enum(_COMMIT_POLICIES, commit_policy, "commit_policy") @@ -1499,6 +1552,9 @@ def transcribe( keep_special_tags: bool = False, spec_k_drafts: int = -1, family: FamilyExtension | None = None, + vocabulary: Sequence[str] | None = None, + prompt: str | None = None, + prefix: str | None = None, ) -> Result: """Transcribe *pcm* in one call and return a materialized Result. @@ -1507,13 +1563,15 @@ def transcribe( many clips keep a Model and call ``model.session().run(...)`` yourself; this helper is for the one-shot case. ``backend`` / ``device`` apply only when *model* is a path — they are ignored when an already-loaded Model is passed. - ``family`` / ``spec_k_drafts`` pass through to :meth:`Session.run`. + ``family`` / ``spec_k_drafts`` and the prompting inputs (``vocabulary``, + ``prompt``, ``prefix``) pass through to :meth:`Session.run`. """ session_opts = dict(n_threads=n_threads, kv_type=kv_type, n_ctx=n_ctx) run_opts = dict(task=task, language=language, target_language=target_language, timestamps=timestamps, pnc=pnc, itn=itn, diarize=diarize, keep_special_tags=keep_special_tags, - spec_k_drafts=spec_k_drafts, family=family) + spec_k_drafts=spec_k_drafts, family=family, + vocabulary=vocabulary, prompt=prompt, prefix=prefix) if isinstance(model, Model): with model.session(**session_opts) as session: diff --git a/bindings/python/src/transcribe_cpp/_generated.py b/bindings/python/src/transcribe_cpp/_generated.py index 9c7d47156..d1122d2ae 100644 --- a/bindings/python/src/transcribe_cpp/_generated.py +++ b/bindings/python/src/transcribe_cpp/_generated.py @@ -13,7 +13,7 @@ # Stable digest of the ABI surface below (structs, enums, macros, layout, # prototypes). A native provider package echoes this back so the API # package can reject an ABI-mismatched provider before dlopen. -PUBLIC_HEADER_HASH = "9866413f80138057" +PUBLIC_HEADER_HASH = "59b9a92b47074666" # === enum constants === TRANSCRIBE_OK = 0 @@ -59,6 +59,7 @@ TRANSCRIBE_LOG_LEVEL_CONT = 5 TRANSCRIBE_TASK_TRANSCRIBE = 0 TRANSCRIBE_TASK_TRANSLATE = 1 +TRANSCRIBE_TASK_INSTRUCT = 2 TRANSCRIBE_TIMESTAMPS_NONE = 0 TRANSCRIBE_TIMESTAMPS_AUTO = 1 TRANSCRIBE_TIMESTAMPS_SEGMENT = 2 @@ -96,6 +97,10 @@ TRANSCRIBE_FEATURE_PNC = 4 TRANSCRIBE_FEATURE_ITN = 5 TRANSCRIBE_FEATURE_DIARIZATION = 6 +TRANSCRIBE_FEATURE_VOCABULARY = 7 +TRANSCRIBE_FEATURE_CONTEXT_PROMPT = 8 +TRANSCRIBE_FEATURE_INSTRUCT = 9 +TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX = 10 TRANSCRIBE_STREAM_IDLE = 0 TRANSCRIBE_STREAM_ACTIVE = 1 TRANSCRIBE_STREAM_FINISHED = 2 @@ -168,7 +173,7 @@ class transcribe_whisper_chunk_trace(_c.Structure): transcribe_device_info._fields_ = [("struct_size", _c.c_uint64), ("name", _c.c_char_p), ("description", _c.c_char_p), ("kind", _c.c_char_p), ("device_id", _c.c_char_p), ("memory_total", _c.c_uint64), ("memory_free", _c.c_uint64), ("device_type", _c.c_int)] transcribe_model_load_params._fields_ = [("struct_size", _c.c_uint64), ("backend", _c.c_int), ("device", _c.c_void_p)] transcribe_session_params._fields_ = [("struct_size", _c.c_uint64), ("n_threads", _c.c_int), ("kv_type", _c.c_int), ("n_ctx", _c.c_int32)] -transcribe_run_params._fields_ = [("struct_size", _c.c_uint64), ("task", _c.c_int), ("timestamps", _c.c_int), ("pnc", _c.c_int), ("itn", _c.c_int), ("diarize", _c.c_int), ("language", _c.c_char_p), ("target_language", _c.c_char_p), ("keep_special_tags", _c.c_bool), ("family", _c.POINTER(transcribe_ext)), ("spec_k_drafts", _c.c_int32)] +transcribe_run_params._fields_ = [("struct_size", _c.c_uint64), ("task", _c.c_int), ("timestamps", _c.c_int), ("pnc", _c.c_int), ("itn", _c.c_int), ("diarize", _c.c_int), ("language", _c.c_char_p), ("target_language", _c.c_char_p), ("keep_special_tags", _c.c_bool), ("family", _c.POINTER(transcribe_ext)), ("spec_k_drafts", _c.c_int32), ("vocabulary", _c.POINTER(_c.c_char_p)), ("n_vocabulary", _c.c_int32), ("prompt", _c.c_char_p), ("prefix", _c.c_char_p)] transcribe_capabilities._fields_ = [("struct_size", _c.c_uint64), ("native_sample_rate", _c.c_int32), ("n_languages", _c.c_int), ("languages", _c.POINTER(_c.c_char_p)), ("max_timestamp_kind", _c.c_int), ("supports_language_detect", _c.c_bool), ("supports_translate", _c.c_bool), ("supports_streaming", _c.c_bool), ("supports_spec_decode", _c.c_bool), ("max_audio_ms", _c.c_int64), ("n_translate_target_languages", _c.c_int), ("translate_target_languages", _c.POINTER(_c.c_char_p))] transcribe_session_limits._fields_ = [("struct_size", _c.c_uint64), ("effective_n_ctx", _c.c_int32), ("effective_max_audio_ms", _c.c_int64), ("max_kv_bytes", _c.c_int64)] transcribe_stream_params._fields_ = [("struct_size", _c.c_uint64), ("family", _c.POINTER(transcribe_ext)), ("commit_policy", _c.c_int), ("stable_prefix_agreement_n", _c.c_uint32)] @@ -213,7 +218,7 @@ class transcribe_whisper_chunk_trace(_c.Structure): 'transcribe_device_info': {'size': 64, 'align': 8, 'offsets': {'struct_size': 0, 'name': 8, 'description': 16, 'kind': 24, 'device_id': 32, 'memory_total': 40, 'memory_free': 48, 'device_type': 56}}, 'transcribe_model_load_params': {'size': 24, 'align': 8, 'offsets': {'struct_size': 0, 'backend': 8, 'device': 16}}, 'transcribe_session_params': {'size': 24, 'align': 8, 'offsets': {'struct_size': 0, 'n_threads': 8, 'kv_type': 12, 'n_ctx': 16}}, - 'transcribe_run_params': {'size': 72, 'align': 8, 'offsets': {'struct_size': 0, 'task': 8, 'timestamps': 12, 'pnc': 16, 'itn': 20, 'diarize': 24, 'language': 32, 'target_language': 40, 'keep_special_tags': 48, 'family': 56, 'spec_k_drafts': 64}}, + 'transcribe_run_params': {'size': 104, 'align': 8, 'offsets': {'struct_size': 0, 'task': 8, 'timestamps': 12, 'pnc': 16, 'itn': 20, 'diarize': 24, 'language': 32, 'target_language': 40, 'keep_special_tags': 48, 'family': 56, 'spec_k_drafts': 64, 'vocabulary': 72, 'n_vocabulary': 80, 'prompt': 88, 'prefix': 96}}, 'transcribe_capabilities': {'size': 56, 'align': 8, 'offsets': {'struct_size': 0, 'native_sample_rate': 8, 'n_languages': 12, 'languages': 16, 'max_timestamp_kind': 24, 'supports_language_detect': 28, 'supports_translate': 29, 'supports_streaming': 30, 'supports_spec_decode': 31, 'max_audio_ms': 32, 'n_translate_target_languages': 40, 'translate_target_languages': 48}}, 'transcribe_session_limits': {'size': 32, 'align': 8, 'offsets': {'struct_size': 0, 'effective_n_ctx': 8, 'effective_max_audio_ms': 16, 'max_kv_bytes': 24}}, 'transcribe_stream_params': {'size': 24, 'align': 8, 'offsets': {'struct_size': 0, 'family': 8, 'commit_policy': 16, 'stable_prefix_agreement_n': 20}}, diff --git a/bindings/python/tests/test_prompting.py b/bindings/python/tests/test_prompting.py new file mode 100644 index 000000000..ca08cbc09 --- /dev/null +++ b/bindings/python/tests/test_prompting.py @@ -0,0 +1,116 @@ +"""Generic prompting inputs: run-params marshalling and model-gated behavior.""" + +import ctypes + +import pytest + +import transcribe_cpp as t +from transcribe_cpp import _generated + + +def test_run_params_carry_prompting_fields(): + params = t._build_run_params("instruct", None, None, "none", False, -1, + vocabulary=["GGUF", "ggml"], prompt="Summarize.", + prefix="And so") + assert params.task == _generated.TRANSCRIBE_TASK_INSTRUCT + assert params.n_vocabulary == 2 + assert [params.vocabulary[i] for i in range(2)] == [b"GGUF", b"ggml"] + assert params.prompt == b"Summarize." + assert params.prefix == b"And so" + + +def test_run_params_default_to_no_prompting(): + params = t._build_run_params("transcribe", None, None, "auto", False, -1) + assert params.n_vocabulary == 0 + assert not params.vocabulary + assert params.prompt is None and params.prefix is None + + +def test_vocabulary_rejects_single_string(): + with pytest.raises(t.InvalidArgument): + t._build_run_params("transcribe", None, None, "auto", False, -1, vocabulary="GGUF") + + +def test_nul_in_string_option_is_rejected(): + for kwargs in ({"prompt": "a\x00b"}, {"prefix": "a\x00b"}, {"vocabulary": ["a\x00b"]}): + with pytest.raises(t.InvalidArgument): + t._build_run_params("transcribe", None, None, "auto", False, -1, **kwargs) + + +def test_vocabulary_accepts_iterators_rejects_sets(): + params = t._build_run_params("transcribe", None, None, "auto", False, -1, + vocabulary=(w for w in ["GGUF", "ggml"])) + assert params.n_vocabulary == 2 + with pytest.raises(t.InvalidArgument): + t._build_run_params("transcribe", None, None, "auto", False, -1, vocabulary={"GGUF"}) + + +def test_prompting_features_probe(model_path): + with t.Model(model_path, backend="cpu") as model: + for feature in ("vocabulary", "context_prompt", "instruct", "transcript_prefix"): + assert isinstance(model.supports(feature), bool) + + +def test_unsupported_prefix_raises(streaming_model_path, audio_pcm): + with t.Model(streaming_model_path, backend="cpu") as model, model.session() as session: + if model.supports("transcript_prefix"): + pytest.skip("model supports a transcript prefix") + with pytest.raises(t.InvalidArgument): + session.run(audio_pcm, prefix="And so") + + +def test_stream_prompting(streaming_model_path): + with t.Model(streaming_model_path, backend="cpu") as model, model.session() as session: + with pytest.raises(t.UnsupportedRequest): + session.stream(task="instruct", prompt="Summarize.") + with session.stream(vocabulary=["Kennedy", "Americans"], prompt="A speech."): + pass + + +def _whisper(model): + return model.arch == "whisper" + + +def test_whisper_prefix_contract(model_path, audio_pcm): + """The prefix is forced decoder text: text holds only the continuation, + raw_text leads with the prefix, and nothing is duplicated.""" + prefix = "And so my fellow Americans," + with t.Model(model_path, backend="cpu") as model, model.session() as session: + if not _whisper(model): + pytest.skip("whisper-specific rendering") + assert model.supports("transcript_prefix") + res = session.run(audio_pcm, prefix=prefix) + assert res.raw_text.strip().startswith(prefix) + assert not res.text.lower().startswith("and so") + assert "ask not" in res.text.lower() + # Whisper's timestamp rules do not compose with a prefix. + with t.Model(model_path, backend="cpu") as model, model.session() as session: + with pytest.raises(t.InvalidArgument): + session.run(audio_pcm, prefix=prefix, timestamps="segment") + + +def test_whisper_vocabulary_and_context(model_path, audio_pcm): + with t.Model(model_path, backend="cpu") as model, model.session() as session: + if not _whisper(model): + pytest.skip("whisper-specific rendering") + assert model.supports("vocabulary") and model.supports("context_prompt") + res = session.run(audio_pcm, vocabulary=["Kennedy", "Americans"], + prompt="An inaugural address.") + assert "country" in res.text.lower() + batch = session.run_batch([audio_pcm, audio_pcm], vocabulary=["Kennedy", "Americans"]) + assert len(batch) == 2 and all("country" in r.text.lower() for r in batch) + # The whisper extension's prompt and the generic fields share one slot. + with pytest.raises(t.InvalidArgument): + session.run(audio_pcm, vocabulary=["Kennedy"], + family=t.WhisperRunOptions(initial_prompt="x")) + + +def test_control_token_literal_rejected(model_path, audio_pcm): + """A control-token literal is rejected and leaves the previous result.""" + with t.Model(model_path, backend="cpu") as model, model.session() as session: + if not _whisper(model): + pytest.skip("whisper-specific rendering") + first = session.run(audio_pcm) + with pytest.raises(t.InvalidArgument): + session.run(audio_pcm, prompt="hello <|endoftext|>") + assert session._materialize().text == first.text diff --git a/bindings/rust/sys/src/transcribe_sys.rs b/bindings/rust/sys/src/transcribe_sys.rs index fc6483640..cbccfab04 100644 --- a/bindings/rust/sys/src/transcribe_sys.rs +++ b/bindings/rust/sys/src/transcribe_sys.rs @@ -1,11 +1,11 @@ // @generated by `cargo xtask bindgen` from include/transcribe/extensions.h // DO NOT EDIT BY HAND. Regenerate: `cargo xtask bindgen`. -// Pinned to include/transcribe.abihash = 9866413f80138057 +// Pinned to include/transcribe.abihash = 59b9a92b47074666 /// The public-ABI digest these bindings were generated against /// (sha256/16 over the normalized FFI surface). The load-time version /// gate and the CI drift check both anchor on this value. -pub const PUBLIC_HEADER_HASH: &str = "9866413f80138057"; +pub const PUBLIC_HEADER_HASH: &str = "59b9a92b47074666"; /* automatically generated by rust-bindgen 0.72.1 */ @@ -100,6 +100,7 @@ unsafe extern "C" { impl transcribe_task { pub const TRANSCRIBE_TASK_TRANSCRIBE: transcribe_task = transcribe_task(0); pub const TRANSCRIBE_TASK_TRANSLATE: transcribe_task = transcribe_task(1); + pub const TRANSCRIBE_TASK_INSTRUCT: transcribe_task = transcribe_task(2); } #[repr(transparent)] #[derive(Debug, Copy, Clone, Hash, PartialEq, Eq)] @@ -347,10 +348,14 @@ pub struct transcribe_run_params { pub keep_special_tags: bool, pub family: *const transcribe_ext, pub spec_k_drafts: i32, + pub vocabulary: *const *const ::std::os::raw::c_char, + pub n_vocabulary: i32, + pub prompt: *const ::std::os::raw::c_char, + pub prefix: *const ::std::os::raw::c_char, } #[allow(clippy::unnecessary_operation, clippy::identity_op)] const _: () = { - ["Size of transcribe_run_params"][::std::mem::size_of::() - 72usize]; + ["Size of transcribe_run_params"][::std::mem::size_of::() - 104usize]; ["Alignment of transcribe_run_params"] [::std::mem::align_of::() - 8usize]; ["Offset of field: transcribe_run_params::struct_size"] @@ -375,6 +380,14 @@ const _: () = { [::std::mem::offset_of!(transcribe_run_params, family) - 56usize]; ["Offset of field: transcribe_run_params::spec_k_drafts"] [::std::mem::offset_of!(transcribe_run_params, spec_k_drafts) - 64usize]; + ["Offset of field: transcribe_run_params::vocabulary"] + [::std::mem::offset_of!(transcribe_run_params, vocabulary) - 72usize]; + ["Offset of field: transcribe_run_params::n_vocabulary"] + [::std::mem::offset_of!(transcribe_run_params, n_vocabulary) - 80usize]; + ["Offset of field: transcribe_run_params::prompt"] + [::std::mem::offset_of!(transcribe_run_params, prompt) - 88usize]; + ["Offset of field: transcribe_run_params::prefix"] + [::std::mem::offset_of!(transcribe_run_params, prefix) - 96usize]; }; unsafe extern "C" { pub fn transcribe_run_params_init(params: *mut transcribe_run_params); @@ -442,6 +455,10 @@ impl transcribe_feature { pub const TRANSCRIBE_FEATURE_PNC: transcribe_feature = transcribe_feature(4); pub const TRANSCRIBE_FEATURE_ITN: transcribe_feature = transcribe_feature(5); pub const TRANSCRIBE_FEATURE_DIARIZATION: transcribe_feature = transcribe_feature(6); + pub const TRANSCRIBE_FEATURE_VOCABULARY: transcribe_feature = transcribe_feature(7); + pub const TRANSCRIBE_FEATURE_CONTEXT_PROMPT: transcribe_feature = transcribe_feature(8); + pub const TRANSCRIBE_FEATURE_INSTRUCT: transcribe_feature = transcribe_feature(9); + pub const TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX: transcribe_feature = transcribe_feature(10); } #[repr(transparent)] #[derive(Debug, Copy, Clone, Hash, PartialEq, Eq)] diff --git a/bindings/rust/transcribe-cpp/README.md b/bindings/rust/transcribe-cpp/README.md index 95f9f969d..cdaa40fb5 100644 --- a/bindings/rust/transcribe-cpp/README.md +++ b/bindings/rust/transcribe-cpp/README.md @@ -50,6 +50,20 @@ let result = session.run(&pcm, &options)?; # Ok::<(), transcribe_cpp::Error>(()) ``` +### Prompting + +`RunOptions::vocabulary` (custom terms), `prompt` (context, or the instruction +under `Task::Instruct`) and `prefix` (text the model continues from) take +effect where `model.supports()` reports `Feature::Vocabulary`, +`ContextPrompt`, `Instruct` or `TranscriptPrefix`. + +```rust +use transcribe_cpp::RunOptions; +let options = RunOptions { vocabulary: vec!["Kubernetes".into()], ..Default::default() }; +let result = session.run(&pcm, &options)?; +# Ok::<(), transcribe_cpp::Error>(()) +``` + Streaming exposes both UI-stable text and a fully materialized structured snapshot: diff --git a/bindings/rust/transcribe-cpp/src/family.rs b/bindings/rust/transcribe-cpp/src/family.rs index 78b7b877d..e17cadd92 100644 --- a/bindings/rust/transcribe-cpp/src/family.rs +++ b/bindings/rust/transcribe-cpp/src/family.rs @@ -19,6 +19,7 @@ use crate::error::Result; /// Whisper run-extension knobs (run slot): initial prompt, temperature /// fallback, and decode thresholds. `None` keeps the family default. +/// `initial_prompt` cannot be combined with `RunOptions::vocabulary` / `prompt`. #[derive(Debug, Clone, Default, PartialEq)] pub struct WhisperRunOptions { pub initial_prompt: Option, diff --git a/bindings/rust/transcribe-cpp/src/session.rs b/bindings/rust/transcribe-cpp/src/session.rs index f15f66e3c..af009db38 100644 --- a/bindings/rust/transcribe-cpp/src/session.rs +++ b/bindings/rust/transcribe-cpp/src/session.rs @@ -43,6 +43,17 @@ pub struct RunOptions { pub spec_k_drafts: i32, /// Optional family-specific run extension (e.g. whisper decode knobs). pub family: Option, + /// Custom terms in priority order, formatted per family + /// (`Feature::Vocabulary`; ignored with a warning elsewhere). + pub vocabulary: Vec, + /// Context text under `Task::Transcribe`/`Translate` + /// (`Feature::ContextPrompt`); the required instruction under + /// `Task::Instruct`. + pub prompt: Option, + /// Transcript text the model continues from (`Feature::TranscriptPrefix`; + /// an error elsewhere, and in batch and streaming runs). `text` holds only + /// the continuation; `raw_text` leads with the prefix. + pub prefix: Option, } impl Default for RunOptions { @@ -58,6 +69,9 @@ impl Default for RunOptions { keep_special_tags: false, spec_k_drafts: -1, family: None, + vocabulary: Vec::new(), + prompt: None, + prefix: None, } } } @@ -150,7 +164,7 @@ impl Session { /// transcript is preserved on the returned [`Error::Aborted`] / /// [`Error::OutputTruncated`] / [`Error::OutputRepetition`]. pub fn run(&mut self, pcm: &[f32], options: &RunOptions) -> Result { - let (params, _lang, _target, _family) = build_run_params(options)?; + let (params, _lang, _target, _family, _prompting) = build_run_params(options)?; let n = clamp_len(pcm.len())?; // The compute path is serialized per model; hold the lock for the native @@ -194,7 +208,7 @@ impl Session { pcms: &[&[f32]], options: &RunOptions, ) -> Result>> { - let (params, _lang, _target, _family) = build_run_params(options)?; + let (params, _lang, _target, _family, _prompting) = build_run_params(options)?; let ptrs: Vec<*const f32> = pcms.iter().map(|p| p.as_ptr()).collect(); let lens: Vec = pcms .iter() @@ -266,12 +280,13 @@ impl Session { /// Begin a streaming run, returning a [`Stream`] that borrows this session /// for its lifetime (so the session can't be used for an offline `run` - /// while a stream is active). `run` supplies task / language / timestamps; + /// while a stream is active). `run` supplies task / language / timestamps / + /// vocabulary / prompt; /// `stream` supplies the commit policy and any stream-slot family extension. /// Dropping the returned `Stream` abandons it and returns the session to /// idle. pub fn stream(&mut self, run: &RunOptions, stream: &StreamOptions) -> Result> { - let (run_params, _lang, _target, _family) = build_run_params(run)?; + let (run_params, _lang, _target, _family, _prompting) = build_run_params(run)?; let (stream_params, _stream_family) = build_stream_params(stream); { // Claim the model's compute lease for the whole stream lifetime: a @@ -427,8 +442,17 @@ type RunParamsBundle = ( Option, Option, Option, + PromptingKeepalive, ); +/// Owns the buffers behind the prompting pointers of a `transcribe_run_params`. +struct PromptingKeepalive { + _terms: Vec, + _term_ptrs: Vec<*const std::os::raw::c_char>, + _prompt: Option, + _prefix: Option, +} + /// Build `transcribe_run_params` from options. The returned keepalives own the /// buffers the params' pointers borrow, so the caller must hold them for the /// duration of the native call. @@ -456,7 +480,30 @@ fn build_run_params(o: &RunOptions) -> Result { .transpose()?; params.family = family.as_ref().map_or(std::ptr::null(), |f| f.ext_ptr()); - Ok((params, lang, target, family)) + let terms = o + .vocabulary + .iter() + .map(|t| CString::new(t.as_str())) + .collect::, _>>()?; + let term_ptrs: Vec<*const std::os::raw::c_char> = terms.iter().map(|c| c.as_ptr()).collect(); + let prompt = o.prompt.as_deref().map(CString::new).transpose()?; + let prefix = o.prefix.as_deref().map(CString::new).transpose()?; + params.vocabulary = if term_ptrs.is_empty() { + std::ptr::null() + } else { + term_ptrs.as_ptr() + }; + params.n_vocabulary = clamp_len(term_ptrs.len())?; + params.prompt = prompt.as_ref().map_or(std::ptr::null(), |c| c.as_ptr()); + params.prefix = prefix.as_ref().map_or(std::ptr::null(), |c| c.as_ptr()); + let prompting = PromptingKeepalive { + _terms: terms, + _term_ptrs: term_ptrs, + _prompt: prompt, + _prefix: prefix, + }; + + Ok((params, lang, target, family, prompting)) } /// PCM/utterance lengths cross the ABI as `int`; reject anything that overflows. diff --git a/bindings/rust/transcribe-cpp/src/types.rs b/bindings/rust/transcribe-cpp/src/types.rs index 02089439d..df7bf1227 100644 --- a/bindings/rust/transcribe-cpp/src/types.rs +++ b/bindings/rust/transcribe-cpp/src/types.rs @@ -9,12 +9,17 @@ use transcribe_cpp_sys as sys; /// The task a run performs. #[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] +#[non_exhaustive] pub enum Task { /// Transcribe speech in its source language. #[default] Transcribe, /// Translate speech into the target language (model must support it). Translate, + /// `RunOptions::prompt` replaces the task instruction; the output is free + /// text in `text` / `raw_text` (model must support `Feature::Instruct`; + /// offline only). + Instruct, } impl Task { @@ -22,6 +27,7 @@ impl Task { match self { Task::Transcribe => sys::transcribe_task::TRANSCRIBE_TASK_TRANSCRIBE, Task::Translate => sys::transcribe_task::TRANSCRIBE_TASK_TRANSLATE, + Task::Instruct => sys::transcribe_task::TRANSCRIBE_TASK_INSTRUCT, } } } @@ -196,8 +202,9 @@ impl Backend { /// A yes/no model capability probe (`transcribe_model_supports`). #[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[non_exhaustive] pub enum Feature { - /// Accepts a free-text/token decode prompt (whisper today). + /// The whisper run extension's initial prompt / prompt tokens. InitialPrompt, /// Runs a multi-tier temperature fallback loop (whisper today). TemperatureFallback, @@ -211,6 +218,14 @@ pub enum Feature { Itn, /// Produces structured speaker attribution. Diarization, + /// `RunOptions::vocabulary` is formatted for this model. + Vocabulary, + /// `RunOptions::prompt` reaches a transcription-conditioning slot. + ContextPrompt, + /// Supports `Task::Instruct`. + Instruct, + /// Honors `RunOptions::prefix` as forced decoder text. + TranscriptPrefix, } impl Feature { @@ -224,6 +239,10 @@ impl Feature { Feature::Pnc => F::TRANSCRIBE_FEATURE_PNC, Feature::Itn => F::TRANSCRIBE_FEATURE_ITN, Feature::Diarization => F::TRANSCRIBE_FEATURE_DIARIZATION, + Feature::Vocabulary => F::TRANSCRIBE_FEATURE_VOCABULARY, + Feature::ContextPrompt => F::TRANSCRIBE_FEATURE_CONTEXT_PROMPT, + Feature::Instruct => F::TRANSCRIBE_FEATURE_INSTRUCT, + Feature::TranscriptPrefix => F::TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX, } } } diff --git a/bindings/rust/transcribe-cpp/tests/prompting.rs b/bindings/rust/transcribe-cpp/tests/prompting.rs new file mode 100644 index 000000000..074c15dcb --- /dev/null +++ b/bindings/rust/transcribe-cpp/tests/prompting.rs @@ -0,0 +1,110 @@ +//! Generic prompting (vocabulary / prompt / prefix) against the whisper and +//! streaming canaries. Mirrors bindings/python/tests/test_prompting.py. + +mod common; + +use transcribe_cpp::{Error, Feature, Model, RunOptions, StreamOptions, Task, TimestampKind}; + +fn terms() -> Vec { + ["Kennedy", "Americans", "GGUF"].map(String::from).to_vec() +} + +#[test] +fn nul_in_prompting_text_is_an_error() { + let Some((model_path, pcm)) = common::smoke_fixtures("nul_in_prompting_text_is_an_error") + else { + return; + }; + let mut session = Model::load(&model_path).unwrap().session().unwrap(); + let options = RunOptions { + vocabulary: vec!["a\0b".into()], + ..Default::default() + }; + assert!(matches!(session.run(&pcm, &options), Err(Error::Nul(_)))); +} + +#[test] +fn vocabulary_and_prompt_reach_run_and_batch() { + let Some((model_path, pcm)) = + common::smoke_fixtures("vocabulary_and_prompt_reach_run_and_batch") + else { + return; + }; + let model = Model::load(&model_path).unwrap(); + let mut session = model.session().unwrap(); + let options = RunOptions { + vocabulary: terms(), + prompt: Some("A speech.".into()), + ..Default::default() + }; + let result = session.run(&pcm, &options).unwrap(); + assert!(result.text.to_lowercase().contains("country")); + let batch = session.run_batch(&[&pcm, &pcm], &options).unwrap(); + assert!(batch.iter().all(|r| r.is_ok())); + // A control-token literal is rejected; a prefix cannot apply to a batch. + let control = RunOptions { + prompt: Some("hi <|endoftext|>".into()), + ..Default::default() + }; + assert!(matches!( + session.run(&pcm, &control), + Err(Error::InvalidArgument(_)) + )); + let prefix = RunOptions { + prefix: Some("And so".into()), + ..Default::default() + }; + assert!(matches!( + session.run_batch(&[&pcm], &prefix), + Err(Error::InvalidArgument(_)) + )); +} + +#[test] +fn prefix_text_is_the_continuation() { + let Some((model_path, pcm)) = common::smoke_fixtures("prefix_text_is_the_continuation") else { + return; + }; + let model = Model::load(&model_path).unwrap(); + if !model.supports(Feature::TranscriptPrefix) { + eprintln!("skip: model does not support a transcript prefix"); + return; + } + let mut session = model.session().unwrap(); + let options = RunOptions { + prefix: Some("And so my fellow Americans".into()), + timestamps: TimestampKind::None, + ..Default::default() + }; + let result = session.run(&pcm, &options).unwrap(); + assert!(result + .raw_text + .trim_start() + .starts_with("And so my fellow Americans")); + assert!(!result.text.contains("fellow Americans")); + assert!(result.text.to_lowercase().contains("ask not")); +} + +#[test] +fn stream_accepts_vocabulary_rejects_instruct() { + let Some(model_path) = common::smoke_streaming_model() else { + eprintln!("skip: streaming canary unavailable"); + return; + }; + let mut session = Model::load(&model_path).unwrap().session().unwrap(); + let instruct = RunOptions { + task: Task::Instruct, + prompt: Some("Summarize.".into()), + ..Default::default() + }; + assert!(matches!( + session.stream(&instruct, &StreamOptions::default()), + Err(Error::Unsupported(_)) + )); + let options = RunOptions { + vocabulary: terms(), + prompt: Some("A speech.".into()), + ..Default::default() + }; + assert!(session.stream(&options, &StreamOptions::default()).is_ok()); +} diff --git a/bindings/swift/README.md b/bindings/swift/README.md index fd34d7b60..e647712a4 100644 --- a/bindings/swift/README.md +++ b/bindings/swift/README.md @@ -72,6 +72,17 @@ let options = RunOptions(pnc: .off, itn: .on) let transcript = try session.run(pcm, options: options) ``` +### Prompting + +`vocabulary` (custom terms), `prompt` (context, or the instruction under +`.instruct`) and `prefix` (text the model continues from) take effect where +`model.supports()` reports `.vocabulary`, `.contextPrompt`, `.instruct` or +`.transcriptPrefix`. + +```swift +let transcript = try session.run(pcm, options: RunOptions(vocabulary: ["Kubernetes", "gRPC"])) +``` + Streaming models expose committed/tentative text for UI display: ```swift diff --git a/bindings/swift/Sources/TranscribeCpp/ABIHash.swift b/bindings/swift/Sources/TranscribeCpp/ABIHash.swift index 535db4165..f55cc4c77 100644 --- a/bindings/swift/Sources/TranscribeCpp/ABIHash.swift +++ b/bindings/swift/Sources/TranscribeCpp/ABIHash.swift @@ -13,7 +13,7 @@ import CTranscribe extension Transcribe { /// sha256/16 of the normalized public FFI surface, pinned to the value in /// include/transcribe.abihash at the time this binding was last reviewed. - public static let pinnedHeaderHash = "9866413f80138057" + public static let pinnedHeaderHash = "59b9a92b47074666" /// The public-ABI digest this binding was reviewed against (16 hex chars). public static func headerHash() -> String { pinnedHeaderHash } diff --git a/bindings/swift/Sources/TranscribeCpp/Options.swift b/bindings/swift/Sources/TranscribeCpp/Options.swift index 10ff81f30..9477942fb 100644 --- a/bindings/swift/Sources/TranscribeCpp/Options.swift +++ b/bindings/swift/Sources/TranscribeCpp/Options.swift @@ -1,14 +1,21 @@ import CTranscribe +import Foundation // MARK: - Enums -/// The run mode: plain transcription or speech translation. Named -/// `TranscriptionTask` (not `Task`) so it does not shadow Swift's +/// The run mode: plain transcription, speech translation, or `instruct` +/// (`RunOptions.prompt` replaces the task instruction; free-text output; +/// offline only). +/// Named `TranscriptionTask` (not `Task`) so it does not shadow Swift's /// `_Concurrency.Task` in files that `import TranscribeCpp`. public enum TranscriptionTask: Sendable { - case transcribe, translate + case transcribe, translate, instruct var cValue: transcribe_task { - self == .transcribe ? TRANSCRIBE_TASK_TRANSCRIBE : TRANSCRIBE_TASK_TRANSLATE + switch self { + case .transcribe: return TRANSCRIBE_TASK_TRANSCRIBE + case .translate: return TRANSCRIBE_TASK_TRANSLATE + case .instruct: return TRANSCRIBE_TASK_INSTRUCT + } } } @@ -82,6 +89,7 @@ public enum Diarize: Sendable { public enum Feature: Sendable { case initialPrompt, temperatureFallback, longForm, cancellation, pnc, itn, diarization + case vocabulary, contextPrompt, instruct, transcriptPrefix var cValue: transcribe_feature { switch self { case .initialPrompt: return TRANSCRIBE_FEATURE_INITIAL_PROMPT @@ -91,6 +99,10 @@ public enum Feature: Sendable { case .pnc: return TRANSCRIBE_FEATURE_PNC case .itn: return TRANSCRIBE_FEATURE_ITN case .diarization: return TRANSCRIBE_FEATURE_DIARIZATION + case .vocabulary: return TRANSCRIBE_FEATURE_VOCABULARY + case .contextPrompt: return TRANSCRIBE_FEATURE_CONTEXT_PROMPT + case .instruct: return TRANSCRIBE_FEATURE_INSTRUCT + case .transcriptPrefix: return TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX } } } @@ -136,6 +148,16 @@ public struct RunOptions: Sendable { public var specKDrafts: Int32 /// Family-specific run extension (whisper run options); M3. public var family: RunExtension? + /// Custom terms in priority order, formatted per family + /// (`Feature.vocabulary`; ignored with a warning elsewhere). + public var vocabulary: [String] + /// Context text under transcribe/translate (`Feature.contextPrompt`); + /// the required instruction under `.instruct`. + public var prompt: String? + /// Transcript text the model continues from (`Feature.transcriptPrefix`; + /// an error elsewhere, and in batch and streaming runs). `text` holds only + /// the continuation; `rawText` leads with the prefix. + public var prefix: String? public init( task: TranscriptionTask = .transcribe, @@ -150,7 +172,10 @@ public struct RunOptions: Sendable { targetLanguage: String? = nil, keepSpecialTags: Bool = false, specKDrafts: Int32 = -1, - family: RunExtension? = nil + family: RunExtension? = nil, + vocabulary: [String] = [], + prompt: String? = nil, + prefix: String? = nil ) { self.task = task self.timestamps = timestamps @@ -162,11 +187,25 @@ public struct RunOptions: Sendable { self.keepSpecialTags = keepSpecialTags self.specKDrafts = specKDrafts self.family = family + self.vocabulary = vocabulary + self.prompt = prompt + self.prefix = prefix + } + + /// Throws `.invalidArgument` if a string option contains a NUL character, + /// where C would silently cut it. + func checkCStrings() throws { + var strings = [language, targetLanguage, prompt, prefix].compactMap { $0 } + vocabulary + if case .whisper(let o)? = family, let p = o.initialPrompt { strings.append(p) } + if strings.contains(where: { $0.contains("\0") }) { + throw TranscribeError.invalidArgument("a string option contains a NUL character") + } } /// Materialize a `transcribe_run_params` and run `body` with a pointer to - /// it. The `language` / `target_language` C strings are kept alive for the - /// duration of `body` (the C side copies them before returning). + /// it. The C strings (language, target language, vocabulary, prompt, + /// prefix) are kept alive for the duration of `body` (the C side copies + /// them before returning). func withCParams(_ body: (UnsafePointer) throws -> R) rethrows -> R { var params = transcribe_run_params() transcribe_run_params_init(¶ms) @@ -183,13 +222,35 @@ public struct RunOptions: Sendable { params.target_language = tgt return try withRunExtension(family) { ext in params.family = ext - return try withUnsafePointer(to: ¶ms) { try body($0) } + return try withCStringArray(vocabulary) { terms, n in + params.vocabulary = terms + params.n_vocabulary = n + return try withOptionalCString(prompt) { p in + params.prompt = p + return try withOptionalCString(prefix) { x in + params.prefix = x + return try withUnsafePointer(to: ¶ms) { try body($0) } + } + } + } } } } } } +/// Run `body` with a C array of NUL-terminated copies of `strings` (NULL when +/// empty), freed when `body` returns. +func withCStringArray( + _ strings: [String], _ body: (UnsafePointer?>?, Int32) throws -> R +) rethrows -> R { + if strings.isEmpty { return try body(nil, 0) } + let copies: [UnsafeMutablePointer?] = strings.map { strdup($0) } + defer { copies.forEach { free($0) } } + let ptrs: [UnsafePointer?] = copies.map { UnsafePointer($0) } + return try ptrs.withUnsafeBufferPointer { try body($0.baseAddress, Int32(strings.count)) } +} + func withOptionalCString( _ string: String?, _ body: (UnsafePointer?) throws -> R ) rethrows -> R { diff --git a/bindings/swift/Sources/TranscribeCpp/Session.swift b/bindings/swift/Sources/TranscribeCpp/Session.swift index a5a1e5994..9357899d3 100644 --- a/bindings/swift/Sources/TranscribeCpp/Session.swift +++ b/bindings/swift/Sources/TranscribeCpp/Session.swift @@ -50,6 +50,7 @@ public final class Session { /// Transcribe one utterance. `pcm` is mono float32 at 16 kHz in [-1, 1]. public func run(_ pcm: [Float], options: RunOptions = .init()) throws -> Transcript { + try options.checkCStrings() model.runLock.lock() defer { model.runLock.unlock() } if model.streamActive { @@ -70,6 +71,7 @@ public final class Session { public func runBatch( _ inputs: [[Float]], options: RunOptions = .init() ) throws -> [Result] { + try options.checkCStrings() model.runLock.lock() defer { model.runLock.unlock() } if model.streamActive { diff --git a/bindings/swift/Sources/TranscribeCpp/Streaming.swift b/bindings/swift/Sources/TranscribeCpp/Streaming.swift index 455f2a3bf..70b0c0dfd 100644 --- a/bindings/swift/Sources/TranscribeCpp/Streaming.swift +++ b/bindings/swift/Sources/TranscribeCpp/Streaming.swift @@ -191,6 +191,7 @@ extension Session { public func stream( _ runOptions: RunOptions = .init(), _ streamOptions: StreamOptions = .init() ) throws -> Stream { + try runOptions.checkCStrings() model.runLock.lock() defer { model.runLock.unlock() } if model.streamActive { diff --git a/bindings/swift/Tests/TranscribeCppTests/PromptingTests.swift b/bindings/swift/Tests/TranscribeCppTests/PromptingTests.swift new file mode 100644 index 000000000..ff4abde33 --- /dev/null +++ b/bindings/swift/Tests/TranscribeCppTests/PromptingTests.swift @@ -0,0 +1,67 @@ +import CTranscribe +import XCTest + +@testable import TranscribeCpp + +/// Generic prompting (vocabulary / prompt / prefix). The marshalling tests run +/// without a model; the rest use the whisper and streaming canaries. Mirrors +/// bindings/python/tests/test_prompting.py. +final class PromptingTests: XCTestCase { + func testEnumValues() { + XCTAssertEqual(TranscriptionTask.instruct.cValue.rawValue, 2) + XCTAssertEqual(Feature.vocabulary.cValue.rawValue, 7) + XCTAssertEqual(Feature.contextPrompt.cValue.rawValue, 8) + XCTAssertEqual(Feature.instruct.cValue.rawValue, 9) + XCTAssertEqual(Feature.transcriptPrefix.cValue.rawValue, 10) + } + + func testRunParamsCarryPromptingFields() { + let options = RunOptions( + task: .instruct, vocabulary: ["GGUF", "ggml"], prompt: "Summarize.", prefix: "And so") + options.withCParams { p in + XCTAssertEqual(p.pointee.n_vocabulary, 2) + XCTAssertEqual(String(cString: p.pointee.vocabulary![0]!), "GGUF") + XCTAssertEqual(String(cString: p.pointee.vocabulary![1]!), "ggml") + XCTAssertEqual(String(cString: p.pointee.prompt!), "Summarize.") + XCTAssertEqual(String(cString: p.pointee.prefix!), "And so") + } + XCTAssertThrowsError(try RunOptions(prompt: "a\0b").checkCStrings()) + XCTAssertThrowsError(try RunOptions(vocabulary: ["a\0b"]).checkCStrings()) + RunOptions().withCParams { p in + XCTAssertNil(p.pointee.vocabulary) + XCTAssertEqual(p.pointee.n_vocabulary, 0) + XCTAssertNil(p.pointee.prompt) + XCTAssertNil(p.pointee.prefix) + } + } + + func testVocabularyAndPromptReachRunAndBatch() throws { + let (path, pcm) = try Fixtures.modelAndAudio() + let session = try Model(path: path).session() + let options = RunOptions(vocabulary: ["Kennedy", "Americans", "GGUF"], prompt: "A speech.") + XCTAssertTrue(try session.run(pcm, options: options).text.lowercased().contains("country")) + let batch = try session.runBatch([pcm, pcm], options: options) + XCTAssertEqual(batch.count, 2) + for item in batch { XCTAssertNoThrow(try item.get()) } + XCTAssertThrowsError(try session.run(pcm, options: RunOptions(prompt: "hi <|endoftext|>"))) + XCTAssertThrowsError(try session.runBatch([pcm], options: RunOptions(prefix: "And so"))) + } + + func testPrefixTextIsTheContinuation() throws { + let (path, pcm) = try Fixtures.modelAndAudio() + let model = try Model(path: path) + guard model.supports(.transcriptPrefix) else { throw XCTSkip("no transcript prefix") } + let prefix = "And so my fellow Americans" + let t = try model.session().run(pcm, options: RunOptions(timestamps: .none, prefix: prefix)) + XCTAssertTrue(t.rawText.trimmingCharacters(in: .whitespaces).hasPrefix(prefix), t.rawText) + XCTAssertFalse(t.text.contains("fellow Americans"), t.text) + XCTAssertTrue(t.text.lowercased().contains("ask not"), t.text) + } + + func testStreamAcceptsVocabularyRejectsInstruct() throws { + guard let path = Fixtures.streamingModelPath() else { throw XCTSkip("no streaming canary") } + let session = try Model(path: path).session() + XCTAssertThrowsError(try session.stream(RunOptions(task: .instruct, prompt: "Summarize."))) + XCTAssertNoThrow(try session.stream(RunOptions(vocabulary: ["Kennedy"], prompt: "A speech."))) + } +} diff --git a/bindings/typescript/README.md b/bindings/typescript/README.md index 416a45ee2..5121e7999 100644 --- a/bindings/typescript/README.md +++ b/bindings/typescript/README.md @@ -46,6 +46,17 @@ and streams. const result = await model.transcribe(pcm, { pnc: "off", itn: "on" }); ``` +### Prompting + +`vocabulary` (custom terms), `prompt` (context, or the instruction under +`task: "instruct"`) and `prefix` (text the model continues from) take effect +where `model.supports()` reports `"vocabulary"`, `"context_prompt"`, +`"instruct"` or `"transcript_prefix"`. + +```ts +const result = await model.transcribe(pcm, { vocabulary: ["Kubernetes", "gRPC"] }); +``` + ### Streaming ```ts diff --git a/bindings/typescript/src/_generated.ts b/bindings/typescript/src/_generated.ts index 8864ed420..7c4db3ce9 100644 --- a/bindings/typescript/src/_generated.ts +++ b/bindings/typescript/src/_generated.ts @@ -11,7 +11,7 @@ // Stable digest of the ABI surface (structs, enums, macros, layout, // prototypes), computed by the Python oracle and pinned here so a header // ABI change turns this binding's drift check red for conscious review. -export const PUBLIC_HEADER_HASH = "9866413f80138057"; +export const PUBLIC_HEADER_HASH = "59b9a92b47074666"; // === enum constants === export const TRANSCRIBE_OK = 0; @@ -57,6 +57,7 @@ export const TRANSCRIBE_LOG_LEVEL_DEBUG = 4; export const TRANSCRIBE_LOG_LEVEL_CONT = 5; export const TRANSCRIBE_TASK_TRANSCRIBE = 0; export const TRANSCRIBE_TASK_TRANSLATE = 1; +export const TRANSCRIBE_TASK_INSTRUCT = 2; export const TRANSCRIBE_TIMESTAMPS_NONE = 0; export const TRANSCRIBE_TIMESTAMPS_AUTO = 1; export const TRANSCRIBE_TIMESTAMPS_SEGMENT = 2; @@ -94,6 +95,10 @@ export const TRANSCRIBE_FEATURE_CANCELLATION = 3; export const TRANSCRIBE_FEATURE_PNC = 4; export const TRANSCRIBE_FEATURE_ITN = 5; export const TRANSCRIBE_FEATURE_DIARIZATION = 6; +export const TRANSCRIBE_FEATURE_VOCABULARY = 7; +export const TRANSCRIBE_FEATURE_CONTEXT_PROMPT = 8; +export const TRANSCRIBE_FEATURE_INSTRUCT = 9; +export const TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX = 10; export const TRANSCRIBE_STREAM_IDLE = 0; export const TRANSCRIBE_STREAM_ACTIVE = 1; export const TRANSCRIBE_STREAM_FINISHED = 2; @@ -122,7 +127,7 @@ export const STRUCT_LAYOUT: Record = { 'transcribe_device_info': { size: 64, align: 8, offsets: {'struct_size': 0, 'name': 8, 'description': 16, 'kind': 24, 'device_id': 32, 'memory_total': 40, 'memory_free': 48, 'device_type': 56} }, 'transcribe_model_load_params': { size: 24, align: 8, offsets: {'struct_size': 0, 'backend': 8, 'device': 16} }, 'transcribe_session_params': { size: 24, align: 8, offsets: {'struct_size': 0, 'n_threads': 8, 'kv_type': 12, 'n_ctx': 16} }, - 'transcribe_run_params': { size: 72, align: 8, offsets: {'struct_size': 0, 'task': 8, 'timestamps': 12, 'pnc': 16, 'itn': 20, 'diarize': 24, 'language': 32, 'target_language': 40, 'keep_special_tags': 48, 'family': 56, 'spec_k_drafts': 64} }, + 'transcribe_run_params': { size: 104, align: 8, offsets: {'struct_size': 0, 'task': 8, 'timestamps': 12, 'pnc': 16, 'itn': 20, 'diarize': 24, 'language': 32, 'target_language': 40, 'keep_special_tags': 48, 'family': 56, 'spec_k_drafts': 64, 'vocabulary': 72, 'n_vocabulary': 80, 'prompt': 88, 'prefix': 96} }, 'transcribe_capabilities': { size: 56, align: 8, offsets: {'struct_size': 0, 'native_sample_rate': 8, 'n_languages': 12, 'languages': 16, 'max_timestamp_kind': 24, 'supports_language_detect': 28, 'supports_translate': 29, 'supports_streaming': 30, 'supports_spec_decode': 31, 'max_audio_ms': 32, 'n_translate_target_languages': 40, 'translate_target_languages': 48} }, 'transcribe_session_limits': { size: 32, align: 8, offsets: {'struct_size': 0, 'effective_n_ctx': 8, 'effective_max_audio_ms': 16, 'max_kv_bytes': 24} }, 'transcribe_stream_params': { size: 24, align: 8, offsets: {'struct_size': 0, 'family': 8, 'commit_policy': 16, 'stable_prefix_agreement_n': 20} }, @@ -167,7 +172,7 @@ export function defineTypes(koffi: any): Record { T['transcribe_device_info'] = koffi.struct({ struct_size: 'uint64_t', name: 'char *', description: 'char *', kind: 'char *', device_id: 'char *', memory_total: 'uint64_t', memory_free: 'uint64_t', device_type: 'int' }); T['transcribe_model_load_params'] = koffi.struct({ struct_size: 'uint64_t', backend: 'int', device: 'void *' }); T['transcribe_session_params'] = koffi.struct({ struct_size: 'uint64_t', n_threads: 'int', kv_type: 'int', n_ctx: 'int32_t' }); - T['transcribe_run_params'] = koffi.struct({ struct_size: 'uint64_t', task: 'int', timestamps: 'int', pnc: 'int', itn: 'int', diarize: 'int', language: 'char *', target_language: 'char *', keep_special_tags: 'bool', family: 'void *', spec_k_drafts: 'int32_t' }); + T['transcribe_run_params'] = koffi.struct({ struct_size: 'uint64_t', task: 'int', timestamps: 'int', pnc: 'int', itn: 'int', diarize: 'int', language: 'char *', target_language: 'char *', keep_special_tags: 'bool', family: 'void *', spec_k_drafts: 'int32_t', vocabulary: 'void *', n_vocabulary: 'int32_t', prompt: 'char *', prefix: 'char *' }); T['transcribe_capabilities'] = koffi.struct({ struct_size: 'uint64_t', native_sample_rate: 'int32_t', n_languages: 'int', languages: 'void *', max_timestamp_kind: 'int', supports_language_detect: 'bool', supports_translate: 'bool', supports_streaming: 'bool', supports_spec_decode: 'bool', max_audio_ms: 'int64_t', n_translate_target_languages: 'int', translate_target_languages: 'void *' }); T['transcribe_session_limits'] = koffi.struct({ struct_size: 'uint64_t', effective_n_ctx: 'int32_t', effective_max_audio_ms: 'int64_t', max_kv_bytes: 'int64_t' }); T['transcribe_stream_params'] = koffi.struct({ struct_size: 'uint64_t', family: 'void *', commit_policy: 'int', stable_prefix_agreement_n: 'uint32_t' }); diff --git a/bindings/typescript/src/index.ts b/bindings/typescript/src/index.ts index 27f416808..c2e07394a 100644 --- a/bindings/typescript/src/index.ts +++ b/bindings/typescript/src/index.ts @@ -82,6 +82,7 @@ const KV_TYPES: Record = { const TASKS = { transcribe: g.TRANSCRIBE_TASK_TRANSCRIBE, translate: g.TRANSCRIBE_TASK_TRANSLATE, + instruct: g.TRANSCRIBE_TASK_INSTRUCT, }; const TIMESTAMPS: Record = { none: g.TRANSCRIBE_TIMESTAMPS_NONE, @@ -116,6 +117,10 @@ const FEATURES: Record = { pnc: g.TRANSCRIBE_FEATURE_PNC, itn: g.TRANSCRIBE_FEATURE_ITN, diarization: g.TRANSCRIBE_FEATURE_DIARIZATION, + vocabulary: g.TRANSCRIBE_FEATURE_VOCABULARY, + context_prompt: g.TRANSCRIBE_FEATURE_CONTEXT_PROMPT, + instruct: g.TRANSCRIBE_FEATURE_INSTRUCT, + transcript_prefix: g.TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX, }; // ---- helpers --------------------------------------------------------------- @@ -563,7 +568,8 @@ const FAMILY: Record = { type: "transcribe_whisper_run_ext", init: "whisperRunExtInit", map: (o) => ({ - initial_prompt: o.initialPrompt, + initial_prompt: + o.initialPrompt === undefined ? undefined : cstr(o.initialPrompt, "initialPrompt"), condition_on_prev_tokens: o.conditionOnPrevTokens, temperature: o.temperature, temperature_inc: o.temperatureInc, @@ -663,6 +669,22 @@ function buildFamily( return buf; } +/** A string option as passed to C, which would silently cut it at a NUL. */ +function cstr(value: string, name: string): string { + if (value.includes("\0")) + throw new InvalidArgument(`${name} contains a NUL character`); + return value; +} + +/** Free what #buildRunParams allocated; call once the native call returns. */ +function freeRunParams(n: Native, p: any): void { + if (p.vocabulary) { + n.koffi.free(p.vocabulary); + p.vocabulary = null; + p.n_vocabulary = 0; + } +} + function toStreamUpdate(u: any): StreamUpdate { return { resultChanged: u.result_changed, @@ -832,7 +854,7 @@ export class Session { aborted: F.wasAborted(h), truncated: F.wasTruncated(h), }; - }); + }).finally(() => freeRunParams(n, p)); } #buildRunParams(opts: TranscribeOptions): any { @@ -847,14 +869,30 @@ export class Session { p.pnc = lookup(PNC, opts.pnc ?? "default", "pnc"); p.itn = lookup(ITN, opts.itn ?? "default", "itn"); p.diarize = lookup(DIARIZE, opts.diarize ?? "default", "diarize"); - if (opts.language !== undefined) p.language = opts.language; + if (opts.language !== undefined) p.language = cstr(opts.language, "language"); if (opts.targetLanguage !== undefined) - p.target_language = opts.targetLanguage; + p.target_language = cstr(opts.targetLanguage, "targetLanguage"); if (opts.keepSpecialTags !== undefined) p.keep_special_tags = opts.keepSpecialTags; if (opts.specKDrafts !== undefined) p.spec_k_drafts = opts.specKDrafts; if (opts.family) p.family = buildFamily(n, this.#model.handle, opts.family, "run"); + if (opts.vocabulary !== undefined) { + const terms = opts.vocabulary; + if (!Array.isArray(terms) || !terms.every((t) => typeof t === "string")) + throw new InvalidArgument("vocabulary must be an array of strings"); + if (terms.length > 0) { + // Freed by freeRunParams after the call; the library copies the terms. + terms.forEach((t) => cstr(t, "vocabulary")); + const type = n.koffi.array("char *", terms.length); + const arr = n.koffi.alloc(type, 1); + n.koffi.encode(arr, type, terms); + p.vocabulary = arr; + p.n_vocabulary = terms.length; + } + } + if (opts.prompt !== undefined) p.prompt = cstr(opts.prompt, "prompt"); + if (opts.prefix !== undefined) p.prefix = cstr(opts.prefix, "prefix"); return p; } @@ -937,7 +975,7 @@ export class Session { } } return out; - }); + }).finally(() => freeRunParams(n, p)); } /** Begin a streaming session. The returned Stream owns the begin params. */ @@ -955,6 +993,8 @@ export class Session { diarize: opts.diarize, keepSpecialTags: opts.keepSpecialTags, specKDrafts: -1, + vocabulary: opts.vocabulary, + prompt: opts.prompt, }); const sp: any = {}; F.streamParamsInit(sp); @@ -984,7 +1024,7 @@ export class Session { if (!control) throw new TranscribeError("session control is missing"); control.replaceCurrentStream(stream); return stream; - }); + }).finally(() => freeRunParams(n, rp)); // begin copied the prompting strings } /** diff --git a/bindings/typescript/src/types.ts b/bindings/typescript/src/types.ts index 0a439518f..d986a5c8f 100644 --- a/bindings/typescript/src/types.ts +++ b/bindings/typescript/src/types.ts @@ -4,7 +4,9 @@ import type { TranscribeError } from "./errors.js"; export type Backend = "auto" | "cpu" | "cpu_accel" | "cuda" | "rocm" | "vulkan" | "metal"; export type KvType = "auto" | "f32" | "f16"; -export type Task = "transcribe" | "translate"; +/** "instruct": `prompt` replaces the task instruction and the output is free + * text ("instruct" feature; offline only). */ +export type Task = "transcribe" | "translate" | "instruct"; export type TimestampKind = "none" | "auto" | "segment" | "word" | "token"; export type Pnc = "default" | "off" | "on"; export type Itn = "default" | "off" | "on"; @@ -16,7 +18,11 @@ export type Feature = | "cancellation" | "pnc" | "itn" - | "diarization"; + | "diarization" + | "vocabulary" + | "context_prompt" + | "instruct" + | "transcript_prefix"; /** Mono float32 PCM at the model's native sample rate (16 kHz for v1). */ export type PcmLike = Float32Array | number[] | ArrayBuffer | Buffer; @@ -168,6 +174,16 @@ export interface TranscribeOptions { signal?: AbortSignal; /** A run-slot family extension (e.g. whisper). */ family?: FamilyExtension; + /** Custom terms in priority order, formatted per family ("vocabulary" + * feature; ignored with a warning elsewhere). */ + vocabulary?: readonly string[]; + /** Context text under transcribe/translate ("context_prompt" feature); the + * required instruction under task "instruct". */ + prompt?: string; + /** Transcript text the model continues from ("transcript_prefix" feature; + * an error elsewhere, and in runBatch). `text` holds only the continuation; + * `rawText` leads with the prefix. */ + prefix?: string; } /** One result of a batch run: success carries the transcript, failure the error. @@ -214,6 +230,10 @@ export interface StreamOptions { keepSpecialTags?: boolean; commitPolicy?: CommitPolicy; stablePrefixAgreementN?: number; + /** Custom terms in priority order (see TranscribeOptions.vocabulary). */ + vocabulary?: readonly string[]; + /** Context text (see TranscribeOptions.prompt). */ + prompt?: string; /** A stream-slot family extension (moonshine, parakeet, voxtral). */ family?: FamilyExtension; } diff --git a/bindings/typescript/test/prompting.test.mjs b/bindings/typescript/test/prompting.test.mjs new file mode 100644 index 000000000..fdac27bca --- /dev/null +++ b/bindings/typescript/test/prompting.test.mjs @@ -0,0 +1,64 @@ +// Generic prompting (vocabulary / prompt / prefix) against the whisper and +// streaming canaries. Mirrors bindings/python/tests/test_prompting.py. + +import assert from "node:assert/strict"; +import { modelTest, MODEL, STREAMING_MODEL, jfk, feedChunks } from "./common.mjs"; +import { TranscribeModel, InvalidArgument, UnsupportedRequest } from "../dist/index.js"; + +const TERMS = ["Kennedy", "Americans", "GGUF"]; + +async function withSession(path, fn) { + const m = await TranscribeModel.load(path); + try { + const s = m.createSession(); + try { + await fn(m, s); + } finally { + s.dispose(); + } + } finally { + m.dispose(); + } +} + +modelTest("vocabulary and prompt reach run and runBatch", MODEL, async () => { + await withSession(MODEL, async (m, s) => { + for (const f of ["vocabulary", "context_prompt", "instruct", "transcript_prefix"]) { + assert.equal(typeof m.supports(f), "boolean"); + } + const r = await s.run(jfk(), { vocabulary: TERMS, prompt: "A speech." }); + assert.match(r.text, /ask not/i); + const items = await s.runBatch([jfk(), jfk()], { vocabulary: TERMS }); + assert.ok(items.every((i) => i.ok)); + }); +}); + +modelTest("prefix: text is the continuation, rawText leads with it", MODEL, async () => { + await withSession(MODEL, async (m, s) => { + if (!m.supports("transcript_prefix")) return; + const r = await s.run(jfk(), { prefix: "And so my fellow Americans", timestamps: "none" }); + assert.match(r.rawText, /^\s*And so my fellow Americans/); + assert.doesNotMatch(r.text, /fellow Americans/); + assert.match(r.text, /ask not/i); + }); +}); + +modelTest("prompting input errors", MODEL, async () => { + await withSession(MODEL, async (_m, s) => { + await assert.rejects(() => s.run(jfk(), { vocabulary: "Kennedy" }), InvalidArgument); + await assert.rejects(() => s.run(jfk(), { prompt: "a\0b" }), InvalidArgument); + await assert.rejects(() => s.run(jfk(), { prompt: "hi <|endoftext|>" }), InvalidArgument); + await assert.rejects(() => s.runBatch([jfk()], { prefix: "And so" }), InvalidArgument); + }); +}); + +modelTest("stream accepts vocabulary, rejects instruct", STREAMING_MODEL, async () => { + await withSession(STREAMING_MODEL, async (_m, s) => { + await assert.rejects(() => s.stream({ task: "instruct", prompt: "Summarize." }), UnsupportedRequest); + const stream = await s.stream({ vocabulary: TERMS, prompt: "A speech." }); + await feedChunks(stream, jfk()); + const fin = await stream.finalize(); + assert.equal(fin.isFinal, true); + stream.reset(); + }); +}); diff --git a/docs/models/canary-180m-flash.md b/docs/models/canary-180m-flash.md index 1bc1888d2..30ff3d538 100644 --- a/docs/models/canary-180m-flash.md +++ b/docs/models/canary-180m-flash.md @@ -25,6 +25,8 @@ Not a streaming model. Word and segment timestamps are upstream-experimental and not exposed in the v1 port (deferred — would require porting the `_timestamps_asr_model` CTC aligner from the `.nemo` archive). +**Prompting:** transcript prefix (`--prefix`) via NeMo's `user_prefix` turn. The decoder-context slot is not exposed: any context text made this model stop early in testing. + See NVIDIA's [model card](https://huggingface.co/nvidia/canary-180m-flash) for training data, intended use, and upstream evaluation methodology. diff --git a/docs/models/canary-1b-flash.md b/docs/models/canary-1b-flash.md index d45ad98af..b443dcde7 100644 --- a/docs/models/canary-1b-flash.md +++ b/docs/models/canary-1b-flash.md @@ -21,6 +21,8 @@ Offline multilingual speech-to-text and translation. The model takes a - **Translation** between English and German, Spanish, or French (both directions). +**Prompting:** transcript prefix (`--prefix`) via NeMo's `user_prefix` turn. The decoder-context slot is not exposed: any context text made this model stop early in testing. + See NVIDIA's [model card](https://huggingface.co/nvidia/canary-1b-flash) for training data, intended use, and upstream evaluation methodology. diff --git a/docs/models/canary-1b-v2.md b/docs/models/canary-1b-v2.md index 6d1c9f708..4ed8bd40f 100644 --- a/docs/models/canary-1b-v2.md +++ b/docs/models/canary-1b-v2.md @@ -33,6 +33,8 @@ variants (180m-flash, 1b-flash) cover only English/German/Spanish/French. Not a streaming model. Word and segment timestamps from the upstream model are not exposed in the v1 port. +**Prompting:** transcript prefix (`--prefix`) via NeMo's `user_prefix` turn. The decoder-context slot is not exposed: context text made the output loop and over-insert in testing. + See NVIDIA's [model card](https://huggingface.co/nvidia/canary-1b-v2) for training data, intended use, and upstream evaluation methodology. diff --git a/docs/models/fun-asr-mlt-nano-2512.md b/docs/models/fun-asr-mlt-nano-2512.md index b90b594a4..ea36dcd2f 100644 --- a/docs/models/fun-asr-mlt-nano-2512.md +++ b/docs/models/fun-asr-mlt-nano-2512.md @@ -44,6 +44,8 @@ For Mandarin-only / dialect-heavy use, the regular [Fun-ASR-Nano](fun-asr-nano-2512.md) was trained on a much larger zh/en/ja corpus and may give better Chinese accuracy. +**Prompting:** vocabulary (`--vocabulary`) as the upstream hotword list (`热词列表:[…]`). + See FunAudioLLM's [model card](https://huggingface.co/FunAudioLLM/Fun-ASR-MLT-Nano-2512) for training data, intended use, and upstream evaluation methodology. diff --git a/docs/models/fun-asr-nano-2512.md b/docs/models/fun-asr-nano-2512.md index 43d335dde..4a060adbc 100644 --- a/docs/models/fun-asr-nano-2512.md +++ b/docs/models/fun-asr-nano-2512.md @@ -31,6 +31,8 @@ is supported by the model. Pass `--itn` on the CLI, or set For multilingual coverage beyond zh/en/ja, see the sibling [Fun-ASR-MLT-Nano](fun-asr-mlt-nano-2512.md) (31 languages). +**Prompting:** vocabulary (`--vocabulary`) as the upstream hotword list (`热词列表:[…]`). + See FunAudioLLM's [model card](https://huggingface.co/FunAudioLLM/Fun-ASR-Nano-2512) for training data, intended use, and upstream evaluation methodology. diff --git a/docs/models/granite-4.0-1b-speech.md b/docs/models/granite-4.0-1b-speech.md index a4f267b7e..f4a1b5e37 100644 --- a/docs/models/granite-4.0-1b-speech.md +++ b/docs/models/granite-4.0-1b-speech.md @@ -25,6 +25,8 @@ English-to-Mandarin. Always via English — there is no direct fr↔de, fr↔es, etc. Pass the target language as a BCP-47 code via `--translate --target-language `; the source language is inferred from the audio. +**Prompting:** vocabulary (`--vocabulary`) as IBM's `Keywords:` list biasing on transcription; it is ignored under translation, where keywords make this model drop the translation. Needs a GGUF converted with the `stt.capability.*` prompting keys; older GGUFs report no prompting support. + See IBM's [model card](https://huggingface.co/ibm-granite/granite-4.0-1b-speech) for training data, intended use, and upstream evaluation methodology. diff --git a/docs/models/granite-speech-4.1-2b-plus.md b/docs/models/granite-speech-4.1-2b-plus.md index 0e51855d6..0b559bf4b 100644 --- a/docs/models/granite-speech-4.1-2b-plus.md +++ b/docs/models/granite-speech-4.1-2b-plus.md @@ -33,6 +33,8 @@ This variant is transcription-only. Unlike the base [`granite-speech-4.1-2b`](granite-speech-4.1-2b.md), it does not perform speech translation. +**Prompting:** vocabulary (`--vocabulary`, `Keywords:` list biasing) and transcript prefix (`--prefix`), in plain transcription mode only. With word timestamps or speaker attribution the vocabulary is ignored with a warning and a prefix is rejected. Needs a GGUF converted with the `stt.capability.*` prompting keys; older GGUFs report no prompting support. + See IBM's [model card](https://huggingface.co/ibm-granite/granite-speech-4.1-2b-plus) for training data, intended use, and upstream evaluation methodology. diff --git a/docs/models/granite-speech-4.1-2b.md b/docs/models/granite-speech-4.1-2b.md index 556b62695..0f91025f8 100644 --- a/docs/models/granite-speech-4.1-2b.md +++ b/docs/models/granite-speech-4.1-2b.md @@ -26,6 +26,8 @@ English-to-Mandarin. Always via English — there is no direct fr↔de, fr↔es, etc. Pass the target language as a BCP-47 code via `--translate --target-language `; the source language is inferred from the audio. +**Prompting:** vocabulary (`--vocabulary`) as IBM's `Keywords:` list biasing, on transcription and translation. Needs a GGUF converted with the `stt.capability.*` prompting keys; older GGUFs report no prompting support. + See IBM's [model card](https://huggingface.co/ibm-granite/granite-speech-4.1-2b) for training data, intended use, and upstream evaluation methodology. diff --git a/docs/models/moss-transcribe-diarize.md b/docs/models/moss-transcribe-diarize.md index 856e80d27..9f4feb3ad 100644 --- a/docs/models/moss-transcribe-diarize.md +++ b/docs/models/moss-transcribe-diarize.md @@ -21,6 +21,8 @@ emergent text markers into clean `full_text`, segment rows, and—when `diarize=ON`—speaker IDs and speaker-turn rows. Built for long-form, multi-speaker audio. No translation; not a streaming model. +**Prompting:** vocabulary (`--vocabulary`) as the upstream `热词提示:` hotword hint. It needs a GGUF converted with the instruction split; older GGUFs report no vocabulary support. + See OpenMOSS's [model card](https://huggingface.co/OpenMOSS-Team/MOSS-Transcribe-Diarize) for training data, intended use, and upstream evaluation. All of OpenMOSS's published metrics are Chinese multi-speaker diarization CER/cpCER; LibriSpeech diff --git a/docs/models/parakeet.md b/docs/models/parakeet.md index 1cf84d449..23588c03a 100644 --- a/docs/models/parakeet.md +++ b/docs/models/parakeet.md @@ -120,7 +120,8 @@ CLI surface. Other Parakeet variants run offline only. What's not supported (consistent across the family): translation, VAD, speaker diarization. Language coverage is English-only except `parakeet-tdt-0.6b-v3` and `parakeet-primeline` (25 European languages, -no auto-detect — language hint required). Note that the v3 lineage, +auto-detected; the language hint is accepted but does not change the +output). Note that the v3 lineage, including `parakeet-primeline`, writes German `ss` where standard orthography uses `ß`; see [parakeet-primeline.md](parakeet-primeline.md#orthography-ß-vs-ss) for diff --git a/docs/models/qwen3-asr-0.6b.md b/docs/models/qwen3-asr-0.6b.md index 93e2c4e36..9d8731380 100644 --- a/docs/models/qwen3-asr-0.6b.md +++ b/docs/models/qwen3-asr-0.6b.md @@ -18,6 +18,8 @@ covered including English, Chinese, Japanese, Korean, German, French, Spanish, Arabic, Russian, Hindi, and Vietnamese. `transcribe-cli` reads a 16 kHz mono WAV and returns the transcript text. +**Prompting:** vocabulary (`--vocabulary`, space-joined) and context prompt (`--prompt`), both in the system message. This size echoes the dictionary into the output more often than the 1.7B. + See the [Qwen3-ASR model card](https://huggingface.co/Qwen/Qwen3-ASR-0.6B) for training data, intended use, and upstream evaluation methodology. diff --git a/docs/models/qwen3-asr-1.7b.md b/docs/models/qwen3-asr-1.7b.md index c092028ce..a01a5d821 100644 --- a/docs/models/qwen3-asr-1.7b.md +++ b/docs/models/qwen3-asr-1.7b.md @@ -18,6 +18,8 @@ Same contract as the 0.6B: offline multilingual STT, 30-language auto-detect, 16 kHz mono WAV in → transcript text out. Targets the same use cases as Qwen3-ASR-0.6B with more parameters for accuracy headroom. +**Prompting:** vocabulary (`--vocabulary`, space-joined) and context prompt (`--prompt`), both in the system message. + See the [Qwen3-ASR-1.7B model card](https://huggingface.co/Qwen/Qwen3-ASR-1.7B) for training data and upstream evaluation. diff --git a/docs/models/voxtral-mini-3b-2507.md b/docs/models/voxtral-mini-3b-2507.md index ce1f03299..567e78938 100644 --- a/docs/models/voxtral-mini-3b-2507.md +++ b/docs/models/voxtral-mini-3b-2507.md @@ -25,6 +25,8 @@ mono WAV and produces a transcript via greedy decoding. mistral-common instruct template ("Translate this to {Language}.") to translate non-English speech into the target language's text. +**Prompting:** `--task instruct --prompt ""` replaces the transcription request with a free-text instruction (summaries, questions, reformatting); the output is free text. + See Mistral's [model card](https://huggingface.co/mistralai/Voxtral-Mini-3B-2507) for training data, intended use, and upstream evaluation. diff --git a/docs/models/voxtral-small-24b-2507.md b/docs/models/voxtral-small-24b-2507.md index 63c5e8b08..e6e3614c6 100644 --- a/docs/models/voxtral-small-24b-2507.md +++ b/docs/models/voxtral-small-24b-2507.md @@ -23,6 +23,8 @@ mono WAV and produces a transcript via greedy decoding. mistral-common instruct template ("Translate this to {Language}.") to translate non-English speech into the target language's text. +**Prompting:** `--task instruct --prompt ""` replaces the transcription request with a free-text instruction (summaries, questions, reformatting); the output is free text. + See Mistral's [model card](https://huggingface.co/mistralai/Voxtral-Small-24B-2507) for training data, intended use, and upstream evaluation. diff --git a/docs/models/whisper-base.en.md b/docs/models/whisper-base.en.md index 1173ac0aa..2007031c6 100644 --- a/docs/models/whisper-base.en.md +++ b/docs/models/whisper-base.en.md @@ -10,6 +10,8 @@ OpenAI Whisper base.en — converted to GGUF for transcribe.cpp. English-only; f Offline English speech-to-text. The model takes a 16 kHz mono WAV and returns a transcript. English-only checkpoints are typically faster and slightly more accurate than the multilingual model at the same parameter count, but they cannot transcribe other languages and cannot translate. Long audio is handled via 30-second chunked decoding. +**Prompting:** vocabulary (`--vocabulary`, rendered as `Glossary: …`), context prompt (`--prompt`) and transcript prefix (`--prefix`, first 30 s window only) through the generic `transcribe_run_params` fields; vocabulary and context share Whisper's 223-token prompt budget and cannot be combined with the Whisper extension's `initial_prompt` / `prompt_tokens`. A prefix does not compose with segment timestamps: `AUTO` falls back to none and an explicit segment request is an error. + See the [upstream model card](https://huggingface.co/openai/whisper-base.en) for training data, intended use, and the original evaluation methodology. diff --git a/docs/models/whisper-base.md b/docs/models/whisper-base.md index aed1660b9..bd3b964fe 100644 --- a/docs/models/whisper-base.md +++ b/docs/models/whisper-base.md @@ -10,6 +10,8 @@ OpenAI Whisper base — converted to GGUF for transcribe.cpp. Multilingual trans Offline multilingual speech-to-text and any-language → English speech translation. The model auto-detects the audio's language (99 languages covered) and emits a transcript in that language; passing `language=""` and `task="translate"` to the underlying `whisper_full_params` produces an English translation instead. `transcribe-cli` reads a 16 kHz mono WAV and returns the transcript text. Long audio is handled via 30-second chunked decoding. +**Prompting:** vocabulary (`--vocabulary`, rendered as `Glossary: …`), context prompt (`--prompt`) and transcript prefix (`--prefix`, first 30 s window only) through the generic `transcribe_run_params` fields; vocabulary and context share Whisper's 223-token prompt budget and cannot be combined with the Whisper extension's `initial_prompt` / `prompt_tokens`. A prefix does not compose with segment timestamps: `AUTO` falls back to none and an explicit segment request is an error. + See the [upstream model card](https://huggingface.co/openai/whisper-base) for training data, intended use, and the original evaluation methodology. diff --git a/docs/models/whisper-large-v2.md b/docs/models/whisper-large-v2.md index 2c1ca6776..898d4f831 100644 --- a/docs/models/whisper-large-v2.md +++ b/docs/models/whisper-large-v2.md @@ -10,6 +10,8 @@ OpenAI Whisper large-v2 — converted to GGUF for transcribe.cpp. Multilingual t Offline multilingual speech-to-text and any-language → English speech translation. The model auto-detects the audio's language (99 languages covered) and emits a transcript in that language; passing `language=""` and `task="translate"` to the underlying `whisper_full_params` produces an English translation instead. `transcribe-cli` reads a 16 kHz mono WAV and returns the transcript text. Long audio is handled via 30-second chunked decoding. +**Prompting:** vocabulary (`--vocabulary`, rendered as `Glossary: …`), context prompt (`--prompt`) and transcript prefix (`--prefix`, first 30 s window only) through the generic `transcribe_run_params` fields; vocabulary and context share Whisper's 223-token prompt budget and cannot be combined with the Whisper extension's `initial_prompt` / `prompt_tokens`. A prefix does not compose with segment timestamps: `AUTO` falls back to none and an explicit segment request is an error. + See the [upstream model card](https://huggingface.co/openai/whisper-large-v2) for training data, intended use, and the original evaluation methodology. diff --git a/docs/models/whisper-large-v3-turbo.md b/docs/models/whisper-large-v3-turbo.md index 266d20295..c5f3fa143 100644 --- a/docs/models/whisper-large-v3-turbo.md +++ b/docs/models/whisper-large-v3-turbo.md @@ -10,6 +10,8 @@ OpenAI Whisper large-v3-turbo — converted to GGUF for transcribe.cpp. Multilin Offline multilingual speech-to-text and any-language → English speech translation. The model auto-detects the audio's language (100 languages covered) and emits a transcript in that language; passing `language=""` and `task="translate"` to the underlying `whisper_full_params` produces an English translation instead. `transcribe-cli` reads a 16 kHz mono WAV and returns the transcript text. Long audio is handled via 30-second chunked decoding. v3 family adds Cantonese (yue) on top of v2's 99 languages and switches to a 128-bin mel input. +**Prompting:** vocabulary (`--vocabulary`, rendered as `Glossary: …`), context prompt (`--prompt`) and transcript prefix (`--prefix`, first 30 s window only) through the generic `transcribe_run_params` fields; vocabulary and context share Whisper's 223-token prompt budget and cannot be combined with the Whisper extension's `initial_prompt` / `prompt_tokens`. A prefix does not compose with segment timestamps: `AUTO` falls back to none and an explicit segment request is an error. + See the [upstream model card](https://huggingface.co/openai/whisper-large-v3-turbo) for training data, intended use, and the original evaluation methodology. diff --git a/docs/models/whisper-large-v3.md b/docs/models/whisper-large-v3.md index 538cc43ae..fd4464ba3 100644 --- a/docs/models/whisper-large-v3.md +++ b/docs/models/whisper-large-v3.md @@ -10,6 +10,8 @@ OpenAI Whisper large-v3 — converted to GGUF for transcribe.cpp. Multilingual t Offline multilingual speech-to-text and any-language → English speech translation. The model auto-detects the audio's language (100 languages covered) and emits a transcript in that language; passing `language=""` and `task="translate"` to the underlying `whisper_full_params` produces an English translation instead. `transcribe-cli` reads a 16 kHz mono WAV and returns the transcript text. Long audio is handled via 30-second chunked decoding. v3 family adds Cantonese (yue) on top of v2's 99 languages and switches to a 128-bin mel input. +**Prompting:** vocabulary (`--vocabulary`, rendered as `Glossary: …`), context prompt (`--prompt`) and transcript prefix (`--prefix`, first 30 s window only) through the generic `transcribe_run_params` fields; vocabulary and context share Whisper's 223-token prompt budget and cannot be combined with the Whisper extension's `initial_prompt` / `prompt_tokens`. A prefix does not compose with segment timestamps: `AUTO` falls back to none and an explicit segment request is an error. + See the [upstream model card](https://huggingface.co/openai/whisper-large-v3) for training data, intended use, and the original evaluation methodology. diff --git a/docs/models/whisper-large.md b/docs/models/whisper-large.md index ba75f67b7..339162e6d 100644 --- a/docs/models/whisper-large.md +++ b/docs/models/whisper-large.md @@ -10,6 +10,8 @@ OpenAI Whisper large — converted to GGUF for transcribe.cpp. Multilingual tran Offline multilingual speech-to-text and any-language → English speech translation. The model auto-detects the audio's language (99 languages covered) and emits a transcript in that language; passing `language=""` and `task="translate"` to the underlying `whisper_full_params` produces an English translation instead. `transcribe-cli` reads a 16 kHz mono WAV and returns the transcript text. Long audio is handled via 30-second chunked decoding. +**Prompting:** vocabulary (`--vocabulary`, rendered as `Glossary: …`), context prompt (`--prompt`) and transcript prefix (`--prefix`, first 30 s window only) through the generic `transcribe_run_params` fields; vocabulary and context share Whisper's 223-token prompt budget and cannot be combined with the Whisper extension's `initial_prompt` / `prompt_tokens`. A prefix does not compose with segment timestamps: `AUTO` falls back to none and an explicit segment request is an error. + See the [upstream model card](https://huggingface.co/openai/whisper-large) for training data, intended use, and the original evaluation methodology. diff --git a/docs/models/whisper-medium.en.md b/docs/models/whisper-medium.en.md index b9b78d06a..72dcea2f0 100644 --- a/docs/models/whisper-medium.en.md +++ b/docs/models/whisper-medium.en.md @@ -10,6 +10,8 @@ OpenAI Whisper medium.en — converted to GGUF for transcribe.cpp. English-only; Offline English speech-to-text. The model takes a 16 kHz mono WAV and returns a transcript. English-only checkpoints are typically faster and slightly more accurate than the multilingual model at the same parameter count, but they cannot transcribe other languages and cannot translate. Long audio is handled via 30-second chunked decoding. +**Prompting:** vocabulary (`--vocabulary`, rendered as `Glossary: …`), context prompt (`--prompt`) and transcript prefix (`--prefix`, first 30 s window only) through the generic `transcribe_run_params` fields; vocabulary and context share Whisper's 223-token prompt budget and cannot be combined with the Whisper extension's `initial_prompt` / `prompt_tokens`. A prefix does not compose with segment timestamps: `AUTO` falls back to none and an explicit segment request is an error. + See the [upstream model card](https://huggingface.co/openai/whisper-medium.en) for training data, intended use, and the original evaluation methodology. diff --git a/docs/models/whisper-medium.md b/docs/models/whisper-medium.md index 878b2d441..fcc92a470 100644 --- a/docs/models/whisper-medium.md +++ b/docs/models/whisper-medium.md @@ -10,6 +10,8 @@ OpenAI Whisper medium — converted to GGUF for transcribe.cpp. Multilingual tra Offline multilingual speech-to-text and any-language → English speech translation. The model auto-detects the audio's language (99 languages covered) and emits a transcript in that language; passing `language=""` and `task="translate"` to the underlying `whisper_full_params` produces an English translation instead. `transcribe-cli` reads a 16 kHz mono WAV and returns the transcript text. Long audio is handled via 30-second chunked decoding. +**Prompting:** vocabulary (`--vocabulary`, rendered as `Glossary: …`), context prompt (`--prompt`) and transcript prefix (`--prefix`, first 30 s window only) through the generic `transcribe_run_params` fields; vocabulary and context share Whisper's 223-token prompt budget and cannot be combined with the Whisper extension's `initial_prompt` / `prompt_tokens`. A prefix does not compose with segment timestamps: `AUTO` falls back to none and an explicit segment request is an error. + See the [upstream model card](https://huggingface.co/openai/whisper-medium) for training data, intended use, and the original evaluation methodology. diff --git a/docs/models/whisper-small.en.md b/docs/models/whisper-small.en.md index 0c9b9ef32..9803ee84d 100644 --- a/docs/models/whisper-small.en.md +++ b/docs/models/whisper-small.en.md @@ -10,6 +10,8 @@ OpenAI Whisper small.en — converted to GGUF for transcribe.cpp. English-only; Offline English speech-to-text. The model takes a 16 kHz mono WAV and returns a transcript. English-only checkpoints are typically faster and slightly more accurate than the multilingual model at the same parameter count, but they cannot transcribe other languages and cannot translate. Long audio is handled via 30-second chunked decoding. +**Prompting:** vocabulary (`--vocabulary`, rendered as `Glossary: …`), context prompt (`--prompt`) and transcript prefix (`--prefix`, first 30 s window only) through the generic `transcribe_run_params` fields; vocabulary and context share Whisper's 223-token prompt budget and cannot be combined with the Whisper extension's `initial_prompt` / `prompt_tokens`. A prefix does not compose with segment timestamps: `AUTO` falls back to none and an explicit segment request is an error. + See the [upstream model card](https://huggingface.co/openai/whisper-small.en) for training data, intended use, and the original evaluation methodology. diff --git a/docs/models/whisper-small.md b/docs/models/whisper-small.md index adaff7289..f559b3de6 100644 --- a/docs/models/whisper-small.md +++ b/docs/models/whisper-small.md @@ -10,6 +10,8 @@ OpenAI Whisper small — converted to GGUF for transcribe.cpp. Multilingual tran Offline multilingual speech-to-text and any-language → English speech translation. The model auto-detects the audio's language (99 languages covered) and emits a transcript in that language; passing `language=""` and `task="translate"` to the underlying `whisper_full_params` produces an English translation instead. `transcribe-cli` reads a 16 kHz mono WAV and returns the transcript text. Long audio is handled via 30-second chunked decoding. +**Prompting:** vocabulary (`--vocabulary`, rendered as `Glossary: …`), context prompt (`--prompt`) and transcript prefix (`--prefix`, first 30 s window only) through the generic `transcribe_run_params` fields; vocabulary and context share Whisper's 223-token prompt budget and cannot be combined with the Whisper extension's `initial_prompt` / `prompt_tokens`. A prefix does not compose with segment timestamps: `AUTO` falls back to none and an explicit segment request is an error. + See the [upstream model card](https://huggingface.co/openai/whisper-small) for training data, intended use, and the original evaluation methodology. diff --git a/docs/models/whisper-tiny.en.md b/docs/models/whisper-tiny.en.md index d6dee6e2f..a8deb0868 100644 --- a/docs/models/whisper-tiny.en.md +++ b/docs/models/whisper-tiny.en.md @@ -10,6 +10,8 @@ OpenAI Whisper tiny.en — converted to GGUF for transcribe.cpp. English-only; f Offline English speech-to-text. The model takes a 16 kHz mono WAV and returns a transcript. English-only checkpoints are typically faster and slightly more accurate than the multilingual model at the same parameter count, but they cannot transcribe other languages and cannot translate. Long audio is handled via 30-second chunked decoding. +**Prompting:** vocabulary (`--vocabulary`, rendered as `Glossary: …`), context prompt (`--prompt`) and transcript prefix (`--prefix`, first 30 s window only) through the generic `transcribe_run_params` fields; vocabulary and context share Whisper's 223-token prompt budget and cannot be combined with the Whisper extension's `initial_prompt` / `prompt_tokens`. A prefix does not compose with segment timestamps: `AUTO` falls back to none and an explicit segment request is an error. + See the [upstream model card](https://huggingface.co/openai/whisper-tiny.en) for training data, intended use, and the original evaluation methodology. diff --git a/docs/models/whisper-tiny.md b/docs/models/whisper-tiny.md index 9b20eeb63..2e4905d73 100644 --- a/docs/models/whisper-tiny.md +++ b/docs/models/whisper-tiny.md @@ -10,6 +10,8 @@ OpenAI Whisper tiny — converted to GGUF for transcribe.cpp. Multilingual trans Offline multilingual speech-to-text and any-language → English speech translation. The model auto-detects the audio's language (99 languages covered) and emits a transcript in that language; passing `language=""` and `task="translate"` to the underlying `whisper_full_params` produces an English translation instead. `transcribe-cli` reads a 16 kHz mono WAV and returns the transcript text. Long audio is handled via 30-second chunked decoding. +**Prompting:** vocabulary (`--vocabulary`, rendered as `Glossary: …`), context prompt (`--prompt`) and transcript prefix (`--prefix`, first 30 s window only) through the generic `transcribe_run_params` fields; vocabulary and context share Whisper's 223-token prompt budget and cannot be combined with the Whisper extension's `initial_prompt` / `prompt_tokens`. A prefix does not compose with segment timestamps: `AUTO` falls back to none and an explicit segment request is an error. + See the [upstream model card](https://huggingface.co/openai/whisper-tiny) for training data, intended use, and the original evaluation methodology. diff --git a/docs/porting/families/qwen3_asr.md b/docs/porting/families/qwen3_asr.md index f3bf8468a..222124793 100644 --- a/docs/porting/families/qwen3_asr.md +++ b/docs/porting/families/qwen3_asr.md @@ -179,9 +179,11 @@ code alone. back to `token_embd.weight` with `TENSOR_DUPLICATED` (same as llama.cpp and the existing Cohere decoder path). - **Prompt template.** The Qwen3 chat template (in - `chat_template.json`) carries language and hotword context fields. - The rendered prompt is embedded into the GGUF as a string KV; at - inference the caller-provided language/context is spliced into it. + `chat_template.json`) has no dedicated hotword field: context is just + the system message, and a language hint is the `language X` + assistant prefix. The rendered prompt is embedded into the GGUF as a + string KV; at inference the caller-provided language/context is + spliced into it. The tokenizer merge table and special-token ids are the durable part; the template is a separate KV. - **Reuse Cohere's mel frontend.** Same underlying Whisper diff --git a/docs/prompting.md b/docs/prompting.md new file mode 100644 index 000000000..d908e5c2d --- /dev/null +++ b/docs/prompting.md @@ -0,0 +1,37 @@ +# Prompting + +Generic prompting fields on `transcribe_run_params` (CLI flag in parentheses). +Probe `transcribe_model_supports()` for the matching feature bit. + +| Field | Feature bit | Effect | +|---|---|---| +| `vocabulary` (`--vocabulary`, `--vocabulary-file`) | `VOCABULARY` | Custom terms, priority order, formatted per family. | +| `prompt` (`--prompt`) | `CONTEXT_PROMPT` | Context text in the model's conditioning slot. | +| `task = INSTRUCT` + `prompt` (`--task instruct`) | `INSTRUCT` | `prompt` replaces the transcription request; output is free text. | +| `prefix` (`--prefix`) | `TRANSCRIPT_PREFIX` | Transcript text the model continues from. Unsupported is an error. | + +| Family | Vocabulary | Context prompt | Instruct | Prefix | +|---|---|---|---|---| +| Whisper | yes | yes | | yes | +| Qwen3-ASR | yes | yes | | | +| Voxtral (2507) | | | yes | | +| Granite 4.0-1b / 4.1-2b | yes | | | | +| Granite 4.1-2b-plus | yes | | | yes | +| Canary 180m-flash / 1b-flash / 1b-v2 | | | | yes | +| Fun-ASR-Nano | yes | | | | +| MOSS-Transcribe-Diarize | yes | | | | + +Per-model formats and restrictions are in each model doc's **Prompting:** +note under [`models/`](models/). + +## Limits + +- Over the prompt budget, `vocabulary` drops terms from the end of the list and + `prompt` keeps its most recent text, both with a WARN. An INSTRUCT prompt + that does not fit is an error. +- INSTRUCT requires `target_language == NULL` and timestamps NONE or AUTO. +- `prefix`: the audio must contain the prefix's speech, and long-form models + apply it to the first window only. +- A feature bit covers plain transcription. Under another task or output mode + a model may ignore `vocabulary` or `prompt` with a WARN; its model doc says + when. diff --git a/examples/cli/main.cpp b/examples/cli/main.cpp index 9eae9ac4f..1867af445 100644 --- a/examples/cli/main.cpp +++ b/examples/cli/main.cpp @@ -12,6 +12,7 @@ #include "wav.h" #include +#include #include #include #include @@ -19,6 +20,7 @@ #include #include #include +#include #include #include @@ -220,12 +222,13 @@ struct cli_args { std::string wav_path; std::string model_path; std::string language; - std::string target_language; // --target-language: target lang for translation - std::string batch_file; // --batch: one wav path per line - int batch_size = 0; // --batch-size: >1 groups utterances into - // transcribe_run_batch calls (offline only). - // 0/1 keeps the per-file serial loop. + std::string target_language; // --target-language: target lang for translation + std::string batch_file; // --batch: one wav path per line + int batch_size = 0; // --batch-size: >1 groups utterances into + // transcribe_run_batch calls (offline only). + // 0/1 keeps the per-file serial loop. bool translate = false; + bool instruct = false; // --task instruct bool quiet = false; bool list_devices = false; // --list-devices: print devices and exit bool batch_jsonl = false; // --batch-jsonl: output JSONL @@ -238,6 +241,11 @@ struct cli_args { int device_index = -1; // --device N: -1 = auto, >=0 = exact registry device transcribe_timestamp_kind timestamps = TRANSCRIBE_TIMESTAMPS_AUTO; + // Generic prompting (transcribe_run_params::vocabulary / prompt / prefix). + std::vector vocabulary; // --vocabulary TERMS / --vocabulary-file PATH + std::string prompt; // --prompt TEXT + std::string prefix; // --prefix TEXT + // Whisper-family knobs. Ignored for non-Whisper models. std::string initial_prompt; // --initial-prompt TEXT bool whisper_set = false; @@ -299,6 +307,22 @@ struct cli_args { int spec_k_drafts = -1; }; +// Point rp's generic prompting fields at args' storage; vocabulary_ptrs +// backs rp.vocabulary and must outlive the run. +void apply_prompting(const cli_args & args, transcribe_run_params & rp, std::vector & vocabulary_ptrs) { + if (args.instruct) { + rp.task = TRANSCRIBE_TASK_INSTRUCT; + } + vocabulary_ptrs.clear(); + for (const std::string & term : args.vocabulary) { + vocabulary_ptrs.push_back(term.c_str()); + } + rp.vocabulary = vocabulary_ptrs.empty() ? nullptr : vocabulary_ptrs.data(); + rp.n_vocabulary = static_cast(vocabulary_ptrs.size()); + rp.prompt = args.prompt.empty() ? nullptr : args.prompt.c_str(); + rp.prefix = args.prefix.empty() ? nullptr : args.prefix.c_str(); +} + void print_usage(const char * argv0) { std::fprintf(stderr, "usage: %s [options] audio.wav\n" @@ -307,6 +331,13 @@ void print_usage(const char * argv0) { " -m, --model PATH GGUF model file\n" " -l, --language ISO BCP-47-ish language hint (e.g. en, de)\n" " -t, --translate set task to TRANSLATE\n" + " --task T transcribe, translate or instruct (instruct: --prompt\n" + " is the instruction; output is free text)\n" + " --vocabulary TERMS comma-separated custom terms, priority order;\n" + " repeatable (models with the vocabulary feature)\n" + " --vocabulary-file P custom terms, one per line\n" + " --prompt TEXT context text, or the instruction for --task instruct\n" + " --prefix TEXT transcript text the model continues from\n" " --target-language ISO target language for translation (e.g. de, es, fr)\n" " -q, --quiet suppress library log output\n" " -r, --repeat N run N times per file (benchmark)\n" @@ -444,6 +475,68 @@ bool parse_args(int argc, char ** argv, cli_args & out) { out.target_language = v; } else if (a == "-t" || a == "--translate") { out.translate = true; + out.instruct = false; + } else if (a == "--task") { + const char * v = take_value(a.c_str()); + if (!v) { + return false; + } + const std::string t = v; + out.translate = t == "translate"; + out.instruct = t == "instruct"; + if (!out.translate && !out.instruct && t != "transcribe") { + std::fprintf(stderr, "error: --task must be transcribe, translate or instruct\n"); + return false; + } + } else if (a == "--vocabulary" || a == "--vocabulary-file") { + const char * v = take_value(a.c_str()); + if (!v) { + return false; + } + std::string text = v; + char sep = ','; + if (a == "--vocabulary-file") { + std::ifstream f(v); + if (!f) { + std::fprintf(stderr, "error: cannot read %s\n", v); + return false; + } + text.assign(std::istreambuf_iterator(f), std::istreambuf_iterator()); + if (text.rfind("\xEF\xBB\xBF", 0) == 0) { + text.erase(0, 3); // UTF-8 byte-order mark + } + sep = '\n'; + } + size_t start = 0; + while (start <= text.size()) { + size_t end = text.find(sep, start); + if (end == std::string::npos) { + end = text.size(); + } + size_t a0 = start, b0 = end; + while (a0 < b0 && std::isspace(static_cast(text[a0]))) { + ++a0; + } + while (b0 > a0 && std::isspace(static_cast(text[b0 - 1]))) { + --b0; + } + if (b0 > a0) { + out.vocabulary.emplace_back(text.substr(a0, b0 - a0)); + } + start = end + 1; + } + } else if (a == "--prompt") { + const char * v = take_value(a.c_str()); + if (!v) { + return false; + } + out.prompt = v; + } else if (a == "--prefix") { + const char * v = take_value(a.c_str()); + if (!v) { + return false; + } + out.prefix = v; } else if (a == "-q" || a == "--quiet") { out.quiet = true; } else if (a == "-r" || a == "--repeat") { @@ -698,6 +791,10 @@ bool parse_args(int argc, char ** argv, cli_args & out) { std::fprintf(stderr, "error: cannot combine positional audio.wav with --batch\n"); return false; } + if (!out.prefix.empty() && !out.batch_file.empty()) { + std::fprintf(stderr, "error: --prefix describes one utterance and cannot be combined with --batch\n"); + return false; + } if (out.stream_chunk_ms > 0 && out.repeat > 1) { std::fprintf(stderr, "error: --stream-chunk-ms cannot be combined with --repeat\n"); return false; @@ -853,6 +950,8 @@ int main(int argc, char ** argv) { if (args.translate) { rp.task = TRANSCRIBE_TASK_TRANSLATE; } + std::vector vocabulary_ptrs; + apply_prompting(args, rp, vocabulary_ptrs); if (!args.language.empty()) { rp.language = args.language.c_str(); } @@ -1289,6 +1388,8 @@ int main(int argc, char ** argv) { if (args.translate) { rp.task = TRANSCRIBE_TASK_TRANSLATE; } + std::vector vocabulary_ptrs; + apply_prompting(args, rp, vocabulary_ptrs); if (!args.language.empty()) { rp.language = args.language.c_str(); } diff --git a/include/transcribe.abihash b/include/transcribe.abihash index f126c9a74..ae00915b6 100644 --- a/include/transcribe.abihash +++ b/include/transcribe.abihash @@ -1 +1 @@ -9866413f80138057 +59b9a92b47074666 diff --git a/include/transcribe.h b/include/transcribe.h index ae05c95a7..caa589592 100644 --- a/include/transcribe.h +++ b/include/transcribe.h @@ -446,9 +446,12 @@ TRANSCRIBE_API void transcribe_log_set(transcribe_log_callback cb, void * userda /* Task / timestamps */ /* ----------------------------------------------------------------------- */ +/* INSTRUCT: transcribe_run_params::prompt is the instruction and the output + * is free text (only full_text / raw_text are guaranteed). Offline only. */ typedef enum { TRANSCRIBE_TASK_TRANSCRIBE = 0, TRANSCRIBE_TASK_TRANSLATE = 1, + TRANSCRIBE_TASK_INSTRUCT = 2, } transcribe_task; /* @@ -1038,8 +1041,9 @@ TRANSCRIBE_API void transcribe_session_params_init(struct transcribe_session_par * caller-declared input rate, at which point TRANSCRIBE_ERR_SAMPLE_RATE * (currently reserved) will become observable. * - * task: TRANSCRIBE or TRANSLATE. The model must declare support - * for translate via its capabilities; otherwise the run + * task: TRANSCRIBE, TRANSLATE or INSTRUCT. The model must declare + * support for translate via its capabilities, and for + * INSTRUCT via TRANSCRIBE_FEATURE_INSTRUCT; otherwise the run * returns TRANSCRIBE_ERR_UNSUPPORTED_TASK. * * timestamps: requested granularity. Default params request AUTO, @@ -1072,8 +1076,9 @@ TRANSCRIBE_API void transcribe_session_params_init(struct transcribe_session_par * * target_language: target language for translation tasks, or NULL. * - * String-pointer lifetime (language / target_language): caller-owned, and - * the library copies what it needs before the API call returns. This holds + * String-pointer lifetime (language / target_language / vocabulary / + * prompt / prefix): caller-owned, and the library copies what it needs + * before the API call returns. This holds * for transcribe_run / transcribe_run_batch (synchronous) AND for * transcribe_stream_begin: the dispatcher copies these strings into * session-owned storage at begin, so the caller may free its params — @@ -1102,6 +1107,29 @@ TRANSCRIBE_API void transcribe_session_params_init(struct transcribe_session_par * ext` as field 0. Use transcribe_model_accepts_ext_kind * to probe whether the loaded model accepts a given kind * before pointing `family` at it. + * + * spec_k_drafts: speculative-decode draft length for offline runs: -1 is + * the model default, 0 disables it, >0 drafts K tokens per + * verify pass. Ignored unless the model reports + * transcribe_capabilities::supports_spec_decode. + * + * Generic prompting (vocabulary, prompt, prefix): NULL / 0 / "" means + * unused. Each field is gated by the TRANSCRIBE_FEATURE_* bit in + * parentheses; limits are in docs/prompting.md. + * + * vocabulary / n_vocabulary: custom terms in priority order, formatted for + * the family (VOCABULARY). Ignored with a WARN when + * unsupported. + * + * prompt: context text under TRANSCRIBE / TRANSLATE (CONTEXT_PROMPT), + * ignored with a WARN when unsupported; the required + * instruction under INSTRUCT. Plain text only: control-token + * literals are rejected. + * + * prefix: transcript text the model continues from + * (TRANSCRIPT_PREFIX). Results hold only the continuation, + * except raw_text. An error when unsupported, and under + * INSTRUCT, batch or streaming. */ struct transcribe_run_params { uint64_t struct_size; @@ -1115,29 +1143,12 @@ struct transcribe_run_params { const char * target_language; bool keep_special_tags; const struct transcribe_ext * family; + int32_t spec_k_drafts; - /* - * spec_k_drafts: n-gram-lookup speculative-decode draft length for the - * offline autoregressive decode step. Family-portable strategy knob; - * the family decides how K maps to its internal verify graph. - * - * Convention: - * -1: family default (each family picks its tuned K). - * 0: spec decoding explicitly disabled — standard 1-token-per-step - * autoregression. Use this for byte-equal reproduction of - * pre-spec behavior or when measuring baseline performance. - * >0: draft K tokens per verify pass. Practical range is 1..8; - * optimal K is hardware-dependent (compute-bound hardware - * prefers small K, bandwidth-bound prefers larger K — see - * docs/models/.md for per-family guidance). - * - * Families gate this via transcribe_capabilities::supports_spec_decode. - * Setting spec_k_drafts != -1 on a family with - * supports_spec_decode == false is silently ignored (the run proceeds - * as ordinary autoregression). Probe the capability bit if you want - * to know whether the field will take effect. - */ - int32_t spec_k_drafts; + const char * const * vocabulary; + int32_t n_vocabulary; + const char * prompt; + const char * prefix; }; TRANSCRIBE_API void transcribe_run_params_init(struct transcribe_run_params * params); @@ -1217,14 +1228,8 @@ struct transcribe_capabilities { bool supports_streaming; /* - * supports_spec_decode: gates transcribe_run_params::spec_k_drafts. - * True means the family's offline (transcribe_run / transcribe_run_batch) - * path implements n-gram-lookup speculative decoding. A non-zero - * spec_k_drafts on a model with supports_spec_decode == false is - * silently ignored — the run proceeds as ordinary autoregression. This - * is a soft gate (no error) because spec is purely a performance - * strategy; callers can probe this bit if they want to know whether - * passing K will actually do anything. + * supports_spec_decode: the offline path honors + * transcribe_run_params::spec_k_drafts; elsewhere it is ignored. */ bool supports_spec_decode; @@ -1325,9 +1330,10 @@ TRANSCRIBE_API transcribe_status transcribe_model_get_capabilities(const struct * * Feature meanings: * - * INITIAL_PROMPT The model accepts a free-text or token - * prompt to bias decoding. Today: whisper - * only; reached via transcribe_whisper_run_ext. + * INITIAL_PROMPT The Whisper run extension's initial_prompt / + * prompt_tokens (transcribe_whisper_run_ext). + * For portable prompting use the generic + * fields and the four bits below. * * TEMPERATURE_FALLBACK The model runs a multi-tier temperature loop * with metric-driven fallback. Today: whisper. @@ -1366,6 +1372,15 @@ TRANSCRIBE_API transcribe_status transcribe_model_get_capabilities(const struct * against a model where this returns false emits * a WARN and proceeds. * + * VOCABULARY transcribe_run_params::vocabulary takes effect. + * + * CONTEXT_PROMPT transcribe_run_params::prompt conditions + * TRANSCRIBE / TRANSLATE. + * + * INSTRUCT TRANSCRIBE_TASK_INSTRUCT is available. + * + * TRANSCRIPT_PREFIX transcribe_run_params::prefix is honored. + * * Returns false on NULL model or unknown feature enum. */ typedef enum { @@ -1376,6 +1391,10 @@ typedef enum { TRANSCRIBE_FEATURE_PNC = 4, TRANSCRIBE_FEATURE_ITN = 5, TRANSCRIBE_FEATURE_DIARIZATION = 6, + TRANSCRIBE_FEATURE_VOCABULARY = 7, + TRANSCRIBE_FEATURE_CONTEXT_PROMPT = 8, + TRANSCRIBE_FEATURE_INSTRUCT = 9, + TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX = 10, } transcribe_feature; TRANSCRIBE_API bool transcribe_model_supports(const struct transcribe_model * model, transcribe_feature feature); diff --git a/scripts/convert-granite.py b/scripts/convert-granite.py index 719aaaa3b..c8e671189 100644 --- a/scripts/convert-granite.py +++ b/scripts/convert-granite.py @@ -56,6 +56,9 @@ stt.variant = e.g. "granite-4.0-1b-speech" stt.capability.translate = bool (true for 1b/2b; false for plus) + stt.capability.word_timestamps / speaker_diarization = bool (plus only) + stt.capability.vocabulary = bool (all three variants) + stt.capability.transcript_prefix = bool (plus only) stt.translation.target_languages = BCP-47 target list when translation is true tokenizer.ggml.model = "gpt2" (BPE with byte-level pre-tokenizer) @@ -704,6 +707,19 @@ def convert(model_dir: Path, out_path: Path, variant: str, repo_id: str | None = variant == "granite-speech-4.1-2b-plus") writer.add_bool("stt.capability.speaker_diarization", variant == "granite-speech-4.1-2b-plus") + # Generic prompting (transcribe_run_params). Keyword-list biasing + # (the vocabulary field) is documented on each of these model cards + # and measured on all three; the transcript prefix (IBM's + # prefix_text) only on -plus. Unknown variants advertise neither. + vocabulary_caps = { + "granite-4.0-1b-speech": True, + "granite-speech-4.1-2b": True, + "granite-speech-4.1-2b-plus": True, + } + writer.add_bool("stt.capability.vocabulary", + bool(vocabulary_caps.get(variant, False))) + writer.add_bool("stt.capability.transcript_prefix", + variant == "granite-speech-4.1-2b-plus") # ---- tokenizer.ggml.* (llama.cpp "gpt2" byte-level BPE) ---- # diff --git a/scripts/convert-moss.py b/scripts/convert-moss.py index bebae8ebd..d05cffc19 100644 --- a/scripts/convert-moss.py +++ b/scripts/convert-moss.py @@ -57,6 +57,11 @@ stt.moss.audio_token_id / audio_tokens_per_second / time_marker_every_seconds / enable_time_marker (audio-span + time-marker construction — see modeling/processing notes) + stt.moss.prompt_prefix_tokens / prompt_suffix_tokens / digit_tokens + baked fixed prompt around the audio span + stt.moss.prompt_instruction / prompt_instruction_head_tokens / + prompt_instruction_tail_tokens (the suffix split around the + instruction text, for runtime hotword prompting) stt.frontend.* Whisper frontend parameters CLI: @@ -396,6 +401,20 @@ def compute_prompt_tokens(model_dir: Path) -> dict: prefix_ids = [int(i) for i in tokenizer.encode(before_audio, add_special_tokens=False)] suffix_ids = [int(i) for i in tokenizer.encode(after_audio, add_special_tokens=False)] + # The suffix again, split around the instruction text, so the runtime can + # re-encode the instruction with an appended hotword list (upstream + # examples/prompts.md: "...语音范围。热词提示:{terms}"). The instruction is + # stored as text because a suffix changes its last BPE pretoken. Checked + # here: head + encode(instruction) + tail must equal the baked suffix. + if after_audio.count(DEFAULT_PROMPT) != 1: + raise ValueError("instruction text not found exactly once after the audio placeholder") + head_text, tail_text = after_audio.split(DEFAULT_PROMPT, maxsplit=1) + instr_head_ids = [int(i) for i in tokenizer.encode(head_text, add_special_tokens=False)] + instr_tail_ids = [int(i) for i in tokenizer.encode(tail_text, add_special_tokens=False)] + instr_ids = [int(i) for i in tokenizer.encode(DEFAULT_PROMPT, add_special_tokens=False)] + if instr_head_ids + instr_ids + instr_tail_ids != suffix_ids: + raise ValueError("instruction split does not reproduce the baked prompt suffix") + digit_ids = [] for d in "0123456789": ids = tokenizer.encode(d, add_special_tokens=False) @@ -405,7 +424,8 @@ def compute_prompt_tokens(model_dir: Path) -> dict: print(f"Prompt tokens: prefix={len(prefix_ids)} suffix={len(suffix_ids)} " f"digits={digit_ids}") - return {"prefix_ids": prefix_ids, "suffix_ids": suffix_ids, "digit_ids": digit_ids} + return {"prefix_ids": prefix_ids, "suffix_ids": suffix_ids, "digit_ids": digit_ids, + "instr_head_ids": instr_head_ids, "instr_tail_ids": instr_tail_ids} def compute_size_label(total_params: int) -> str: @@ -547,6 +567,13 @@ def convert(model_dir: Path, out_path: Path, variant: str, repo_id: str | None = writer.add_array("stt.moss.prompt_prefix_tokens", prompt["prefix_ids"]) writer.add_array("stt.moss.prompt_suffix_tokens", prompt["suffix_ids"]) writer.add_array("stt.moss.digit_tokens", prompt["digit_ids"]) + # The suffix split around the instruction (see compute_prompt_tokens): + # prompt_suffix_tokens == instruction_head + encode(instruction) + + # instruction_tail. The runtime appends the hotword list to the + # instruction; GGUFs without these keys keep the fixed prompt. + writer.add_string("stt.moss.prompt_instruction", DEFAULT_PROMPT) + writer.add_array("stt.moss.prompt_instruction_head_tokens", prompt["instr_head_ids"]) + writer.add_array("stt.moss.prompt_instruction_tail_tokens", prompt["instr_tail_ids"]) # ---- stt.frontend.* (Whisper feature extractor) ---- writer.add_string("stt.frontend.type", "mel") diff --git a/scripts/gen_sentencepiece_bpe_fixture.py b/scripts/gen_sentencepiece_bpe_fixture.py new file mode 100644 index 000000000..3683755e8 --- /dev/null +++ b/scripts/gen_sentencepiece_bpe_fixture.py @@ -0,0 +1,28 @@ +"""Generate SentencePiece BPE id fixtures for tests/sentencepiece_bpe_parity.cpp. + +Encodes a text list with the reference `sentencepiece` model and writes one +JSON line per text: {"text": ..., "ids": [...]}, ids offset into the GGUF +vocabulary (aggregate tokenizers put each language's sub-vocab at an offset). + +Usage: + uv run --project scripts/envs/canary scripts/gen_sentencepiece_bpe_fixture.py \ + --model [--offset N] --texts --out +""" +import argparse +import json + +import sentencepiece as spm + +ap = argparse.ArgumentParser() +ap.add_argument("--model", required=True) +ap.add_argument("--offset", type=int, default=0) +ap.add_argument("--texts", required=True) +ap.add_argument("--out", required=True) +args = ap.parse_args() + +sp = spm.SentencePieceProcessor(model_file=args.model) +with open(args.texts, encoding="utf-8") as f, open(args.out, "w", encoding="utf-8") as out: + for line in f: + text = line.rstrip("\n") + ids = [i + args.offset for i in sp.encode(text)] + out.write(json.dumps({"text": text, "ids": ids}, ensure_ascii=False) + "\n") diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt index e40e29271..5b0931f71 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -19,6 +19,7 @@ add_library(transcribe transcribe-mel.cpp transcribe-model.cpp transcribe-tokenizer.cpp + transcribe-prompting.cpp transcribe-unicode.cpp transcribe-unicode-data.cpp transcribe-debug.cpp diff --git a/src/arch/canary/model.cpp b/src/arch/canary/model.cpp index 4676df1c6..3d8288eaf 100644 --- a/src/arch/canary/model.cpp +++ b/src/arch/canary/model.cpp @@ -22,6 +22,7 @@ #include "transcribe-log.h" #include "transcribe-mel.h" #include "transcribe-meta.h" +#include "transcribe-prompting.h" #include "transcribe-repetition-guard.h" #include "weights.h" @@ -383,6 +384,11 @@ transcribe_status load(Loader & loader, const transcribe_model_load_params * par if (const transcribe_status st = read_canary_hparams(loader.gguf(), m->hparams); st != TRANSCRIBE_OK) { return st; } + // Transcript prefix: canary2's user_prefix (measured clean on all three + // canary2 checkpoints). Needs a sub-vocab range for aggregate tokenizers. + transcribe::set_feature(m.get(), TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX, + m->hparams.prompt_format == "canary2" && + (m->hparams.tokenizer_single_sp || !m->hparams.tok_lang_codes.empty())); // Publish the input-length ceiling now that the encoder positional span // and frontend rate are known (apply_family_invariants ran before the @@ -593,18 +599,52 @@ std::vector build_prompt_canary(const CanaryHParams & hp, return ids; } +// Ids the canary2 template emits besides the transcript prefix (see +// build_prompt_canary2): nine task slots, plus the empty decoder-context +// marker on single-SP tokenizers. +int canary2_base_prompt_tokens(const CanaryHParams & hp) { + return hp.tokenizer_single_sp ? 10 : 9; +} + +// The decoder prompt must leave the generation reserve inside the decoder +// self-KV ceiling. Only a transcript prefix makes it variable-length; an +// unbounded one would overrun the KV cache during the prompt prefill. +transcribe_status check_prompt_fits(int prompt_len, int ceiling) { + if (prompt_len + k_gen_reserve > ceiling) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, + "canary run: transcript prefix is too long — a %d-token prompt leaves no room for output " + "within the %d-token decoder context (need %d)", + prompt_len, ceiling, prompt_len + k_gen_reserve); + return TRANSCRIBE_ERR_INVALID_ARG; + } + return TRANSCRIBE_OK; +} + +// Target language of a run: the explicit target when translating, else the +// source language (default "en"). +const char * canary_target_language(const transcribe_run_params * params) { + const char * lang = (params && params->language) ? params->language : "en"; + if (params != nullptr && params->task == TRANSCRIBE_TASK_TRANSLATE && params->target_language) { + return params->target_language; + } + return lang; +} + std::vector build_prompt_canary2(const CanaryModel & cm, const CanaryHParams & hp, int src_lang_id, int tgt_lang_id, const char * /*task*/, - bool pnc) { + bool pnc, + const std::vector & prefix_ids) { // canary2 prompt template: // <|startofcontext|> [decodercontext] <|startoftranscript|> // <|emo:?|> <|src_lang|> <|tgt_lang|> <|pnc|> <|itn|> <|timestamp|> <|diarize|> - // ASR with empty decoder context realizes 9 tokens. + // [user_prefix] + // ASR with empty decoder context realizes 9 tokens. `prefix_ids` (the + // transcript prefix, NeMo's user_prefix turn) follows the last slot. std::vector ids; - ids.reserve(9); + ids.reserve(static_cast(canary2_base_prompt_tokens(hp)) + prefix_ids.size()); if (hp.startofcontext_id < 0 || hp.startoftranscript_id < 0 || src_lang_id < 0 || tgt_lang_id < 0) { return {}; @@ -654,10 +694,47 @@ std::vector build_prompt_canary2(const CanaryModel & cm, return {}; } ids.push_back(hp.nodiarize_id); + ids.insert(ids.end(), prefix_ids.begin(), prefix_ids.end()); return ids; } +// Transcript prefix ids for canary2: NeMo tokenizes the user_prefix turn on +// its own with the target language's SentencePiece (BPE) tokenizer, so it +// carries the dummy-prefix space. Aggregate tokenizers encode within that language's +// sub-vocab; single-SP (canary-1b-v2) over the whole vocab. +transcribe_status encode_canary2_prefix(const CanaryModel & cm, + const char * prefix, + const char * tgt_lang, + std::vector & out) { + out.clear(); + if (prefix == nullptr || prefix[0] == '\0') { + return TRANSCRIBE_OK; + } + if (const transcribe_status st = transcribe::prompting::check_plain_text(cm.tok, prefix, "prefix"); + st != TRANSCRIBE_OK) { + return st; + } + int lo = 0, hi = -1; + if (!cm.hparams.tokenizer_single_sp) { + lo = hi = -1; + for (size_t i = 0; i < cm.hparams.tok_lang_codes.size(); ++i) { + if (cm.hparams.tok_lang_codes[i] == tgt_lang) { + lo = cm.hparams.tok_lang_offsets[i]; + hi = lo + cm.hparams.tok_lang_sizes[i]; + } + } + if (lo < 0) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "canary run: no tokenizer for prefix language '%s'", tgt_lang); + return TRANSCRIBE_ERR_UNSUPPORTED_LANGUAGE; + } + } + // canary-1b-v2's normalizer keeps extra whitespace; the flash models' + // per-language tokenizers collapse it (their SentencePiece model specs). + return cm.tok.encode_sentencepiece_bpe(prefix, out, lo, hi, + /*remove_extra_whitespaces=*/!cm.hparams.tokenizer_single_sp); +} + int find_language_id(const CanaryHParams & hp, const char * lang) { if (lang == nullptr) { return -1; @@ -871,11 +948,8 @@ transcribe_status run(transcribe_session * session, // Build multitask prompt. const char * lang = (params && params->language) ? params->language : "en"; const bool is_translate = (params != nullptr && params->task == TRANSCRIBE_TASK_TRANSLATE); - const char * tgt_lang = lang; - if (is_translate && params && params->target_language) { - tgt_lang = params->target_language; - } - const char * task = is_translate ? "translate" : "asr"; + const char * tgt_lang = canary_target_language(params); + const char * task = is_translate ? "translate" : "asr"; const int src_id = find_language_id(cm->hparams, lang); const int tgt_id = find_language_id(cm->hparams, tgt_lang); @@ -921,8 +995,14 @@ transcribe_status run(transcribe_session * session, } std::vector prompt_ids; + std::vector prefix_ids; if (cm->hparams.prompt_format == "canary2") { - prompt_ids = build_prompt_canary2(*cm, cm->hparams, src_id, tgt_id, task, pnc); + if (const transcribe_status st = + encode_canary2_prefix(*cm, params != nullptr ? params->prefix : nullptr, tgt_lang, prefix_ids); + st != TRANSCRIBE_OK) { + return st; + } + prompt_ids = build_prompt_canary2(*cm, cm->hparams, src_id, tgt_id, task, pnc, prefix_ids); } else if (cm->hparams.prompt_format == "canary") { prompt_ids = build_prompt_canary(cm->hparams, src_id, tgt_id, task, pnc); } @@ -932,6 +1012,11 @@ transcribe_status run(transcribe_session * session, return TRANSCRIBE_ERR_INVALID_ARG; } const int prompt_len = static_cast(prompt_ids.size()); + // Backstop for run_validate's pre-clear bound: never prefill past the KV. + if (const transcribe_status st = check_prompt_fits(prompt_len, canary_context_ceiling(cc->n_ctx, cm->hparams)); + st != TRANSCRIBE_OK) { + return st; + } // Init KV cache. { @@ -1179,8 +1264,11 @@ transcribe_status run(transcribe_session * session, seg.text = full; cc->segments.push_back(std::move(seg)); - cc->raw_text = - tok.decode(generated_ids.data(), static_cast(generated_ids.size())); // unfiltered decode + // Unfiltered decode, led by the transcript prefix when one was + // forced (full_text / segments hold only the continuation). + std::vector raw_ids = prefix_ids; + raw_ids.insert(raw_ids.end(), generated_ids.begin(), generated_ids.end()); + cc->raw_text = tok.decode(raw_ids.data(), static_cast(raw_ids.size())); cc->full_text = std::move(full); cc->result_kind = TRANSCRIBE_TIMESTAMPS_NONE; cc->has_result = true; @@ -1536,13 +1624,10 @@ transcribe_status run_batch(transcribe_session * session, // Shared multitask prompt (identical across the batch). const char * lang = (params && params->language) ? params->language : "en"; const bool is_translate = (params != nullptr && params->task == TRANSCRIBE_TASK_TRANSLATE); - const char * tgt_lang = lang; - if (is_translate && params && params->target_language) { - tgt_lang = params->target_language; - } - const char * task = is_translate ? "translate" : "asr"; - const int src_id = find_language_id(hp, lang); - const int tgt_id = find_language_id(hp, tgt_lang); + const char * tgt_lang = canary_target_language(params); + const char * task = is_translate ? "translate" : "asr"; + const int src_id = find_language_id(hp, lang); + const int tgt_id = find_language_id(hp, tgt_lang); if (src_id < 0 || tgt_id < 0) { return TRANSCRIBE_ERR_UNSUPPORTED_LANGUAGE; } @@ -1555,7 +1640,7 @@ transcribe_status run_batch(transcribe_session * session, } std::vector prompt_ids; if (hp.prompt_format == "canary2") { - prompt_ids = build_prompt_canary2(*cm, hp, src_id, tgt_id, task, pnc); + prompt_ids = build_prompt_canary2(*cm, hp, src_id, tgt_id, task, pnc, {}); } else if (hp.prompt_format == "canary") { prompt_ids = build_prompt_canary(hp, src_id, tgt_id, task, pnc); } @@ -1838,6 +1923,28 @@ transcribe_status run_batch(transcribe_session * session, return TRANSCRIBE_OK; } +// Pre-clear gate for the transcript prefix (canary2): it must encode in the +// target language's tokenizer and fit the decoder context with the +// generation reserve, checked before the dispatcher clears the previous +// result. +transcribe_status run_validate(const transcribe_session * session, const transcribe_run_params * params) { + if (session == nullptr || session->model == nullptr || params == nullptr || params->prefix == nullptr) { + return TRANSCRIBE_OK; + } + const auto * cm = static_cast(session->model); + if (cm->hparams.prompt_format != "canary2") { + return TRANSCRIBE_OK; + } + std::vector prefix_ids; + if (const transcribe_status st = + encode_canary2_prefix(*cm, params->prefix, canary_target_language(params), prefix_ids); + st != TRANSCRIBE_OK) { + return st; + } + return check_prompt_fits(canary2_base_prompt_tokens(cm->hparams) + static_cast(prefix_ids.size()), + canary_context_ceiling(session->n_ctx, cm->hparams)); +} + } // namespace extern const Arch arch = { @@ -1852,6 +1959,7 @@ extern const Arch arch = { /* .stream_finalize = */ nullptr, /* .stream_reset = */ nullptr, /* .accepts_ext_kind = */ nullptr, + /* .run_validate = */ run_validate, }; } // namespace transcribe::canary diff --git a/src/arch/canary/weights.cpp b/src/arch/canary/weights.cpp index 11d1d2a9d..bb3d18d9c 100644 --- a/src/arch/canary/weights.cpp +++ b/src/arch/canary/weights.cpp @@ -127,6 +127,18 @@ transcribe_status read_canary_hparams(const gguf_context * gguf, CanaryHParams & st != TRANSCRIBE_OK) { return st; } + // Sub-vocab ranges are optional: only the canary2 transcript prefix + // encodes text, and it is not advertised without them. + if (!hp.tokenizer_single_sp && + read_string_array_kv(gguf, "stt.canary.tokenizer.lang_codes", hp.tok_lang_codes) == KvResult::Ok && + (read_int32_array_kv(gguf, "stt.canary.tokenizer.lang_offsets", hp.tok_lang_offsets) != KvResult::Ok || + read_int32_array_kv(gguf, "stt.canary.tokenizer.lang_sizes", hp.tok_lang_sizes) != KvResult::Ok || + hp.tok_lang_offsets.size() != hp.tok_lang_codes.size() || + hp.tok_lang_sizes.size() != hp.tok_lang_codes.size())) { + hp.tok_lang_codes.clear(); + hp.tok_lang_offsets.clear(); + hp.tok_lang_sizes.clear(); + } auto require_special = [&](const char * key, int32_t & out) -> transcribe_status { const auto r = read_token_id_required(gguf, key, out); diff --git a/src/arch/canary/weights.h b/src/arch/canary/weights.h index 99aff8645..979fe4abb 100644 --- a/src/arch/canary/weights.h +++ b/src/arch/canary/weights.h @@ -51,7 +51,13 @@ struct CanaryHParams { // render an empty decoder-context slot as a leading whitespace // marker (`▁`) in the canary2 prompt — adds one token to the prompt // length. Aggregate tokenizers skip the empty slot entirely. - bool tokenizer_single_sp = false; + bool tokenizer_single_sp = false; + // Aggregate tokenizers: per-language sub-vocab id ranges + // (stt.canary.tokenizer.lang_codes / lang_offsets / lang_sizes). Empty + // for single-SP tokenizers. + std::vector tok_lang_codes; + std::vector tok_lang_offsets; + std::vector tok_lang_sizes; // Token IDs (filled from tokenizer at load time). int32_t bos_token_id = -1; diff --git a/src/arch/funasr_nano/capabilities.cpp b/src/arch/funasr_nano/capabilities.cpp index b3c163c77..c47be8008 100644 --- a/src/arch/funasr_nano/capabilities.cpp +++ b/src/arch/funasr_nano/capabilities.cpp @@ -24,6 +24,8 @@ void apply_family_invariants(transcribe_model & model) { // runtime toggle. transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CANCELLATION, true); transcribe::set_feature(&model, TRANSCRIBE_FEATURE_ITN, true); + // Generic vocabulary as the upstream get_prompt hotword list. + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_VOCABULARY, true); } } // namespace transcribe::funasr_nano diff --git a/src/arch/funasr_nano/model.cpp b/src/arch/funasr_nano/model.cpp index 21970c98c..b13fccdbc 100644 --- a/src/arch/funasr_nano/model.cpp +++ b/src/arch/funasr_nano/model.cpp @@ -20,6 +20,7 @@ #include "transcribe-loader.h" #include "transcribe-log.h" #include "transcribe-meta.h" +#include "transcribe-prompting.h" #include "transcribe-repetition-guard.h" #include "weights.h" @@ -83,7 +84,9 @@ constexpr const char k_default_variant[] = "fun-asr-nano-2512"; // --------------------------------------------------------------------------- // Generation reserve: what the input gate keeps free, and the decode-budget floor. -constexpr int k_gen_reserve = 256; +constexpr int k_gen_reserve = 256; +// Representative non-audio prompt overhead (chat affixes); advisory. +constexpr int k_prompt_overhead_tokens = 48; // Effective decoder context ceiling, in tokens: the model's trained maximum, // optionally lowered — never raised — by the caller's session n_ctx knob. @@ -104,8 +107,7 @@ int64_t funasr_nano_max_audio_ms(const FunAsrNanoHParams & hp) { if (hp.dec_max_position_embeddings <= 0 || hp.fe_hop_length <= 0 || hp.fe_sample_rate <= 0 || hp.fe_lfr_n <= 0) { return 0; } - constexpr int k_prompt_overhead = 48; // chat affixes; advisory - const int max_audio_tokens = hp.dec_max_position_embeddings - k_prompt_overhead - k_gen_reserve; + const int max_audio_tokens = hp.dec_max_position_embeddings - k_prompt_overhead_tokens - k_gen_reserve; if (max_audio_tokens <= 0) { return 0; } @@ -140,16 +142,49 @@ transcribe_status resolve_chat_tokens(const transcribe::Tokenizer & tok, ChatTok return TRANSCRIBE_OK; } -// Build the language/itn prompt text that the reference's -// FunASRNano.get_prompt produces. hotwords-empty path only. -std::string build_funasr_prompt_text(const char * lang, bool use_itn) { - std::string out; +// get_prompt's hotword block up to the list: +// 请结合上下文信息,更加准确地完成语音转写任务。如果没有相关信息,我们会留空。 +// \n\n\n**上下文信息:**\n\n\n热词列表:[ +constexpr const char k_hotword_preamble[] = + "\xE8\xAF\xB7\xE7\xBB\x93\xE5\x90\x88\xE4\xB8\x8A\xE4\xB8\x8B\xE6\x96\x87\xE4\xBF\xA1\xE6\x81\xAF" + "\xEF\xBC\x8C\xE6\x9B\xB4\xE5\x8A\xA0\xE5\x87\x86\xE7\xA1\xAE\xE5\x9C\xB0\xE5\xAE\x8C\xE6\x88\x90" + "\xE8\xAF\xAD\xE9\x9F\xB3\xE8\xBD\xAC\xE5\x86\x99\xE4\xBB\xBB\xE5\x8A\xA1\xE3\x80\x82\xE5\xA6\x82" + "\xE6\x9E\x9C\xE6\xB2\xA1\xE6\x9C\x89\xE7\x9B\xB8\xE5\x85\xB3\xE4\xBF\xA1\xE6\x81\xAF\xEF\xBC\x8C" + "\xE6\x88\x91\xE4\xBB\xAC\xE4\xBC\x9A\xE7\x95\x99\xE7\xA9\xBA\xE3\x80\x82" + "\n\n\n**\xE4\xB8\x8A\xE4\xB8\x8B\xE6\x96\x87\xE4\xBF\xA1\xE6\x81\xAF\xEF\xBC\x9A**\n\n\n" + "\xE7\x83\xAD\xE8\xAF\x8D\xE5\x88\x97\xE8\xA1\xA8\xEF\xBC\x9A["; + +// Token room for the hotword block alongside `n_audio` audio tokens: what the +// context window leaves after the generation reserve and the representative +// prompt overhead. run() passes its clip; run_batch() passes 0. +int hotword_budget(int ceiling, int n_audio) { + return ceiling - k_gen_reserve - n_audio - k_prompt_overhead_tokens; +} + +// Generic vocabulary -> the upstream hotword block (preamble, ", "-joined +// terms, "]\n"), fitted to `budget` tokens. fit_terms_and_context also rejects +// control-token literals, which encode_with_chat_specials would otherwise +// honor. `fit.terms_text` is the block; empty without terms. +transcribe_status fit_hotwords(const transcribe::Tokenizer & tok, + const transcribe_run_params * params, + int budget, + const char * who, + transcribe::prompting::FittedPrompt & fit) { + return transcribe::prompting::fit_terms_and_context(tok, transcribe::prompting::terms(params), + { k_hotword_preamble, ", ", "]\n" }, "", budget, who, fit); +} + +// Build the prompt text the reference's FunASRNano.get_prompt produces, +// byte for byte: the hotword block (fit_hotwords; may be empty), then the +// language / itn transcription instruction. +std::string build_funasr_prompt_text(const std::string & hotword_block, const char * lang, bool use_itn) { + std::string out = hotword_block; if (lang != nullptr && lang[0] != '\0') { // 语音转写成 = "transcribe to" / "transcribe into" - out = "\xE8\xAF\xAD\xE9\x9F\xB3\xE8\xBD\xAC\xE5\x86\x99\xE6\x88\x90"; + out += "\xE8\xAF\xAD\xE9\x9F\xB3\xE8\xBD\xAC\xE5\x86\x99\xE6\x88\x90"; out += lang; } else { - out = "\xE8\xAF\xAD\xE9\x9F\xB3\xE8\xBD\xAC\xE5\x86\x99"; + out += "\xE8\xAF\xAD\xE9\x9F\xB3\xE8\xBD\xAC\xE5\x86\x99"; } if (!use_itn) { // ,不进行文本规整 = "; do not apply text normalization" @@ -221,6 +256,7 @@ transcribe_status encode_with_chat_specials(const transcribe::Tokenizer & tok, // <|im_start|>/<|im_end|> within each segment. transcribe_status build_funasr_nano_prompt(const transcribe::Tokenizer & tok, const ChatTokens & ct, + const std::string & hotword_block, const char * language, bool use_itn, int fake_token_len, @@ -229,7 +265,7 @@ transcribe_status build_funasr_nano_prompt(const transcribe::Tokenizer & tok, out_ids.clear(); out_fbank_beg = 0; - const std::string prompt_text = build_funasr_prompt_text(language, use_itn); + const std::string prompt_text = build_funasr_prompt_text(hotword_block, language, use_itn); std::string seg_a = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n" @@ -310,7 +346,7 @@ transcribe_status load(Loader & loader, const transcribe_model_load_params * par const int folds = m->hparams.adaptor_use_low_frame_rate ? 8 : 1; m->limits.has_context_cap = true; m->limits.model_max_ctx = m->hparams.dec_max_position_embeddings; - m->limits.prompt_overhead = 48; + m->limits.prompt_overhead = k_prompt_overhead_tokens; m->limits.gen_reserve = k_gen_reserve; m->limits.ms_per_audio_token = static_cast(folds) * m->hparams.fe_lfr_n * m->hparams.fe_hop_length * 1000.0 / m->hparams.fe_sample_rate; @@ -654,10 +690,20 @@ transcribe_status run(transcribe_session * session, } } + // Hotword list, fitted to the room the context window leaves after the + // audio, the rest of the prompt and the generation reserve. + transcribe::prompting::FittedPrompt hotwords; + if (const transcribe_status st = + fit_hotwords(cm->tok, params, hotword_budget(funasr_nano_context_ceiling(cc->n_ctx, hp), fake_token_len), + "funasr_nano run", hotwords); + st != TRANSCRIBE_OK) { + return st; + } + std::vector prompt_ids; int fbank_beg = 0; - if (const transcribe_status st = - build_funasr_nano_prompt(cm->tok, cm->chat_tokens, lang, use_itn, fake_token_len, prompt_ids, fbank_beg); + if (const transcribe_status st = build_funasr_nano_prompt(cm->tok, cm->chat_tokens, hotwords.terms_text, lang, + use_itn, fake_token_len, prompt_ids, fbank_beg); st != TRANSCRIBE_OK) { return st; } @@ -1108,6 +1154,15 @@ transcribe_status run_batch(transcribe_session * session, const char * lang = (params != nullptr) ? params->language : nullptr; bool use_itn = (params != nullptr && params->itn == TRANSCRIBE_ITN_MODE_ON); + // Shared hotword list (one run_params per batch), fitted as if there were + // no audio. A row whose own budget is smaller than that fit would get + // fewer hotwords from run(), so the batch then goes serial (see + // fit_terms_and_context: otherwise the fits match). + transcribe::prompting::FittedPrompt hotwords; + if (fit_hotwords(cm->tok, params, hotword_budget(ceiling, 0), "funasr_nano run_batch", hotwords) != TRANSCRIBE_OK) { + return run_batch_serial(cc, pcm, n_samples, n, params); + } + // ---- Pass 0: parallel frontend (kaldi-fbank, host-side, thread-safe) ---- std::vector> fbufs(n); std::vector T_lfr(n, 0); @@ -1139,9 +1194,12 @@ transcribe_status run_batch(transcribe_session * session, if (audio_embed_one(cc, cm, fbufs[b], T_lfr[b], audio_hosts[b], T_audio[b], enc_us) != TRANSCRIBE_OK) { continue; } + if (static_cast(hotwords.n_tokens()) > std::max(hotword_budget(ceiling, T_audio[b]), 0)) { + return run_batch_serial(cc, pcm, n_samples, n, params); + } int fbank_beg = 0; - if (build_funasr_nano_prompt(cm->tok, cm->chat_tokens, lang, use_itn, T_audio[b], prompt_ids[b], fbank_beg) != - TRANSCRIBE_OK) { + if (build_funasr_nano_prompt(cm->tok, cm->chat_tokens, hotwords.terms_text, lang, use_itn, T_audio[b], + prompt_ids[b], fbank_beg) != TRANSCRIBE_OK) { continue; } T_prompt[b] = static_cast(prompt_ids[b].size()); diff --git a/src/arch/granite/model.cpp b/src/arch/granite/model.cpp index 3d621fbf6..ea725c851 100644 --- a/src/arch/granite/model.cpp +++ b/src/arch/granite/model.cpp @@ -19,6 +19,7 @@ #include "transcribe-log.h" #include "transcribe-mel.h" #include "transcribe-meta.h" +#include "transcribe-prompting.h" #include "transcribe-repetition-guard.h" #include "weights.h" @@ -83,7 +84,9 @@ constexpr float kBnEps = 1e-5f; // rather than silently aliasing RoPE past the trained range. // Generation reserve: what the input gate keeps free, and the decode-budget floor. -constexpr int k_gen_reserve = 256; +constexpr int k_gen_reserve = 256; +// Representative non-audio prompt overhead (chat affixes); advisory. +constexpr int k_prompt_overhead_tokens = 64; // Effective decoder context ceiling, in tokens: the model's trained maximum, // optionally lowered — never raised — by the caller's session n_ctx knob. @@ -117,9 +120,7 @@ int64_t granite_max_audio_ms(const GraniteHParams & hp) { hp.dec_max_position_embeddings <= 0) { return 0; } - // Representative non-audio prompt overhead (chat affixes); advisory. - constexpr int k_prompt_overhead = 64; - const int max_audio_tokens = hp.dec_max_position_embeddings - k_prompt_overhead - k_gen_reserve; + const int max_audio_tokens = hp.dec_max_position_embeddings - k_prompt_overhead_tokens - k_gen_reserve; if (max_audio_tokens <= 0) { return 0; } @@ -264,6 +265,26 @@ transcribe_status load(Loader & loader, const transcribe_model_load_params * par } transcribe::set_feature(m.get(), TRANSCRIBE_FEATURE_DIARIZATION, diar); } + // Generic prompting, variant-scoped the same way: the converter writes + // stt.capability.vocabulary (keyword-list biasing, build_granite_affixes) + // and stt.capability.transcript_prefix (-plus's prefix_text) for the + // variants whose model cards document them. Absent keys mean unsupported. + { + bool vocabulary = false; + bool prefix = false; + if (const transcribe_status st = + read_optional_bool_kv(loader.gguf(), "stt.capability.vocabulary", "granite", false, vocabulary); + st != TRANSCRIBE_OK) { + return st; + } + if (const transcribe_status st = + read_optional_bool_kv(loader.gguf(), "stt.capability.transcript_prefix", "granite", false, prefix); + st != TRANSCRIBE_OK) { + return st; + } + transcribe::set_feature(m.get(), TRANSCRIBE_FEATURE_VOCABULARY, vocabulary); + transcribe::set_feature(m.get(), TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX, prefix); + } if (const transcribe_status st = read_languages_kv(loader.gguf(), *m); st != TRANSCRIBE_OK) { return st; @@ -305,7 +326,7 @@ transcribe_status load(Loader & loader, const transcribe_model_load_params * par m->hparams.fe_sample_rate > 0 && m->hparams.dec_max_position_embeddings > 0) { m->limits.has_context_cap = true; m->limits.model_max_ctx = m->hparams.dec_max_position_embeddings; - m->limits.prompt_overhead = 64; // match granite_max_audio_ms's k_prompt_overhead + m->limits.prompt_overhead = k_prompt_overhead_tokens; m->limits.gen_reserve = k_gen_reserve; // ms per audio token: granite emits num_queries tokens per // window_size encoder frames; t_enc = mel_frames/2; @@ -508,13 +529,40 @@ static const char * granite_target_language_name(const char * code_or_name) { return nullptr; } +// Whether this run uses IBM's word-timestamps task (-plus only; other variants +// advertise NONE): an explicit WORD request. AUTO does not request it. +static bool granite_word_timestamps(const transcribe_model * m, const transcribe_run_params * params) { + return params != nullptr && m->caps.max_timestamp_kind == TRANSCRIBE_TIMESTAMPS_WORD && + params->task == TRANSCRIBE_TASK_TRANSCRIBE && params->timestamps == TRANSCRIBE_TIMESTAMPS_WORD; +} + +// Predicted transcript length for the decode budget. The word-timestamps task +// follows every word with a "[T:N]" marker (about four more tokens), so it +// gets three times the plain-text prediction; the budget is only a ceiling. +static int granite_predicted_tokens(const GraniteModel * cm, const transcribe_run_params * params, int n_audio) { + const int plain = transcribe::predict_transcript_tokens(n_audio, cm->limits.ms_per_audio_token); + return granite_word_timestamps(cm, params) ? 3 * plain : plain; +} + +// Token room for the vocabulary's keyword list alongside `n_audio` audio +// tokens: what the context window leaves after the generation reserve and the +// representative prompt overhead. run() passes its clip; run_batch() passes 0. +static int granite_keyword_room(int ceiling, int n_audio) { + return ceiling - k_gen_reserve - n_audio - k_prompt_overhead_tokens; +} + // Build the prompt prefix/suffix token-id lists from the shared run params and // model variant (the audio tokens splice in between). Single source of truth -// for run() and run_batch(). +// for run() and run_batch(). `keyword_room` bounds the vocabulary's tokens; +// `n_keyword_tokens` (optional) receives how many the fitted list took. static transcribe_status build_granite_affixes(GraniteModel * cm, const transcribe_run_params * params, + int keyword_room, std::vector & prefix_ids, - std::vector & suffix_ids) { + std::vector & suffix_ids, + size_t * n_keyword_tokens = nullptr) { + const bool is_plus = cm->hparams.variant == "granite-speech-4.1-2b-plus"; + bool asr_mode = true; // plain transcription instruction (no task swap) std::string instruction; if (cm->hparams.variant == "granite-speech-4.1-2b") { instruction = "transcribe the speech with proper punctuation and capitalization."; @@ -535,11 +583,12 @@ static transcribe_status build_granite_affixes(GraniteModel * cm, params->target_language); return TRANSCRIBE_ERR_INVALID_ARG; } - instruction = std::string("can you translate the speech into ") + lang_name + "?"; - } else if (params->timestamps == TRANSCRIBE_TIMESTAMPS_WORD) { - // -plus only (1b/2b advertise NONE, gated out upstream). AUTO does - // NOT request timestamps. IBM's verbatim prompt; the model emits - // per-word "[T:N]" centisecond markers (parsed in run()). + instruction = std::string("translate the speech to ") + lang_name + "."; + asr_mode = false; + } else if (granite_word_timestamps(cm, params)) { + // -plus only (see granite_word_timestamps). IBM's verbatim prompt; + // the model emits per-word "[T:N]" centisecond markers (parsed in + // run()). if (diarize_requested(cm, params)) { // Upstream defines timestamps and speaker attribution as // separate tasks (one instruction each); they do not compose. @@ -551,23 +600,61 @@ static transcribe_status build_granite_affixes(GraniteModel * cm, instruction = " Timestamps: Transcribe the speech. After each word, add a timestamp tag " "showing the end time in centiseconds, e.g. hello [T:45] world [T:82]"; + asr_mode = false; } else if (diarize_requested(cm, params)) { // -plus only (the DIARIZATION feature bit gates this). IBM's // verbatim speaker-attribution instruction; the model emits // "[Speaker N]:" tags before turns (split after decode). instruction = k_saa_instruction; + asr_mode = false; + } + } + + // Generic vocabulary as IBM's keyword-list biasing: " Keywords: {terms}" + // (", ") appended to the instruction. On -plus the documented KWB form of + // the ASR instruction is capitalized ("Can you ..."). Translate + keywords + // is IBM-documented for 4.1 and measured usable there; granite-4.0 drops + // the translation for most inputs when keywords are added, and keywords + // with -plus timestamps / speaker attribution are untested, so those + // combinations ignore the vocabulary. + if (const std::vector terms = transcribe::prompting::terms(params); !terms.empty()) { + const bool translate = params->task == TRANSCRIBE_TASK_TRANSLATE; + if (translate && cm->hparams.variant != "granite-speech-4.1-2b") { + log_msg(TRANSCRIBE_LOG_LEVEL_WARN, + "granite: vocabulary is not supported with translation on this variant; ignoring %zu term(s)", + terms.size()); + } else if (!translate && !asr_mode) { + log_msg(TRANSCRIBE_LOG_LEVEL_WARN, + "granite: vocabulary applies to plain transcription only (not word timestamps or speaker " + "attribution); ignoring %zu term(s)", + terms.size()); + } else { + transcribe::prompting::FittedPrompt fit; + if (const transcribe_status st = transcribe::prompting::fit_terms_and_context( + cm->tok, terms, { " Keywords: ", ", ", "" }, "", keyword_room, "granite", fit); + st != TRANSCRIBE_OK) { + return st; + } + if (fit.n_terms > 0) { + if (is_plus && !translate) { + instruction = " Can you transcribe the speech into a written format?"; + } + instruction += fit.terms_text; + } + if (n_keyword_tokens != nullptr) { + *n_keyword_tokens = fit.n_tokens(); + } } } const bool use_granite4_chat = cm->chat_template.find("<|start_of_role|>") != std::string::npos && cm->chat_tokens.start_of_role >= 0 && cm->chat_tokens.end_of_role >= 0; if (use_granite4_chat) { - const char * system_content = (cm->hparams.variant == "granite-speech-4.1-2b-plus") ? - "Knowledge Cutoff Date: April 2024.\n" - "Today's Date: December 19, 2024.\n" - "You are Granite, developed by IBM. You are a helpful AI assistant" : - "You are a helpful assistant. Please ensure responses " - "are professional, accurate, and safe."; + const char * system_content = is_plus ? "Knowledge Cutoff Date: April 2024.\n" + "Today's Date: December 19, 2024.\n" + "You are Granite, developed by IBM. You are a helpful AI assistant" : + "You are a helpful assistant. Please ensure responses " + "are professional, accurate, and safe."; std::vector text_a, text_b; if (const transcribe_status st = cm->tok.encode("system", text_a); st != TRANSCRIBE_OK) { return st; @@ -609,6 +696,26 @@ static transcribe_status build_granite_affixes(GraniteModel * cm, } suffix_ids.insert(suffix_ids.end(), asst_ids.begin(), asst_ids.end()); suffix_ids.push_back(cm->chat_tokens.end_of_role); + + // Transcript prefix (IBM's prefix_text; the dispatcher passes one only + // when stt.capability.transcript_prefix is set): the assistant turn + // opens with it verbatim. Its composition with the word-timestamp and + // speaker-attribution tasks is untested, so those reject it. + if (params != nullptr && params->prefix != nullptr) { + if (!asr_mode) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, + "granite: a transcript prefix is supported in plain transcription only (not word timestamps " + "or speaker attribution)"); + return TRANSCRIBE_ERR_INVALID_ARG; + } + std::vector transcript_prefix_ids; + if (const transcribe_status st = + transcribe::prompting::encode_plain(cm->tok, params->prefix, transcript_prefix_ids, "prefix"); + st != TRANSCRIBE_OK) { + return st; + } + suffix_ids.insert(suffix_ids.end(), transcript_prefix_ids.begin(), transcript_prefix_ids.end()); + } } else { const std::string prefix_text = "USER: "; const std::string suffix_text = instruction + "\n ASSISTANT:"; @@ -750,7 +857,7 @@ void finalize_granite_result(GraniteModel * cm, int64_t audio_ms, Result & out) { out.raw_text = raw_text; // pre-parse marker text, via transcribe_raw_text - if (params != nullptr && params->timestamps == TRANSCRIBE_TIMESTAMPS_WORD) { + if (granite_word_timestamps(cm, params)) { out.full_text = parse_granite_word_timestamps(raw_text, audio_ms, out.words, out.segments); if (!out.words.empty()) { out.result_kind = TRANSCRIBE_TIMESTAMPS_WORD; @@ -1010,17 +1117,20 @@ transcribe_status run(transcribe_session * ctx_base, // PLUS a Granite-specific system prompt (see use_granite4_chat below). // word-timestamps task (-plus) : IBM's verbatim timestamps prompt; the model emits // per-word "[T:N]" centisecond markers (1b/2b: NONE). - // translate task : "can you translate the speech into ?" + // translate task : "translate the speech to ." (IBM model card) std::vector prefix_ids; std::vector suffix_ids; - if (const transcribe_status st = build_granite_affixes(cm, params, prefix_ids, suffix_ids); st != TRANSCRIBE_OK) { + const int n_audio_tokens = cc->n_audio_tokens; + const int ceiling = granite_context_ceiling(cc->n_ctx, cm->hparams); + if (const transcribe_status st = + build_granite_affixes(cm, params, granite_keyword_room(ceiling, n_audio_tokens), prefix_ids, suffix_ids); + st != TRANSCRIBE_OK) { return st; } - const int n_audio_tokens = cc->n_audio_tokens; - const int prefix_len = static_cast(prefix_ids.size()); - const int suffix_len = static_cast(suffix_ids.size()); - const int T_prompt = prefix_len + n_audio_tokens + suffix_len; + const int prefix_len = static_cast(prefix_ids.size()); + const int suffix_len = static_cast(suffix_ids.size()); + const int T_prompt = prefix_len + n_audio_tokens + suffix_len; // Reference quirk: HF replaces audio_token_id with 0 before the // embed_tokens lookup (those rows are overwritten by the audio scatter @@ -1041,20 +1151,18 @@ transcribe_status run(transcribe_session * ctx_base, // autoregressive decode, instead of growing the cache unboundedly and // aliasing RoPE past the trained range. Reserving the full generation // budget means an accepted clip always has room for a real transcript. - const int ceiling = granite_context_ceiling(cc->n_ctx, cm->hparams); if (T_prompt + k_gen_reserve > ceiling) { transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "granite run: input too long — %d audio + %d prompt tokens leave " "no room for output within the %d-token context (need %d). " - "Shorten the audio (see transcribe_capabilities.max_audio_ms) or " - "split it into segments.", + "Shorten the audio (see transcribe_capabilities.max_audio_ms), " + "split it into segments, or shorten the transcript prefix.", n_audio_tokens, prefix_len + suffix_len, ceiling, T_prompt + k_gen_reserve); return TRANSCRIBE_ERR_INPUT_TOO_LONG; } - const int gen_budget = transcribe::pick_decode_budget( - transcribe::predict_transcript_tokens(n_audio_tokens, cm->limits.ms_per_audio_token), k_gen_reserve, T_prompt, - ceiling); + const int gen_budget = transcribe::pick_decode_budget(granite_predicted_tokens(cm, params, n_audio_tokens), + k_gen_reserve, T_prompt, ceiling); // Size the KV cache dynamically: T_prompt + room for the longest // generation we'll emit, clamped to the context ceiling. Matches the @@ -1298,13 +1406,25 @@ transcribe_status run(transcribe_session * ctx_base, } // Detokenize. - const std::string raw_text = cm->tok.decode(gen_ids.data(), static_cast(gen_ids.size())); + std::string raw_text = cm->tok.decode(gen_ids.data(), static_cast(gen_ids.size())); const int64_t audio_ms = static_cast(n_samples) * 1000 / static_cast(cm->hparams.fe_sample_rate); + // Under a transcript prefix the decode is the continuation: the text + // fields hold it (without the joining space), raw_text leads with the + // prefix. + const bool has_prefix = params != nullptr && params->prefix != nullptr; + if (has_prefix) { + const size_t lead = raw_text.find_first_not_of(' '); + raw_text.erase(0, lead == std::string::npos ? raw_text.size() : lead); + } + cc->has_result = true; // -plus word timestamps / speaker attribution / plain text; shared with // the run_batch capture loop via finalize_granite_result. finalize_granite_result(cm, params, raw_text, audio_ms, *cc); + if (has_prefix) { + cc->raw_text = std::string(params->prefix) + cm->tok.decode(gen_ids.data(), static_cast(gen_ids.size())); + } // Output truncation (decode hit the generation budget / context ceiling // before EOS) is a hard status, not a silent success: surface it so the @@ -1473,7 +1593,20 @@ transcribe_status run_batch_serial(GraniteSession * cc, // and cannot be composed. Validate the mode-dependent timestamp contract before // the dispatcher clears the previous result snapshot. transcribe_status run_validate(const transcribe_session * ctx, const transcribe_run_params * params) { - if (ctx == nullptr || ctx->model == nullptr || params == nullptr || !diarize_requested(ctx->model, params)) { + if (ctx == nullptr || ctx->model == nullptr || params == nullptr) { + return TRANSCRIBE_OK; + } + // A transcript prefix only composes with plain transcription (see + // build_granite_affixes); reject the untested task combinations here, + // before the previous result is cleared. + if (params->prefix != nullptr && + (params->timestamps == TRANSCRIBE_TIMESTAMPS_WORD || diarize_requested(ctx->model, params))) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, + "granite: a transcript prefix is supported in plain transcription only " + "(not word timestamps or speaker attribution)"); + return TRANSCRIBE_ERR_INVALID_ARG; + } + if (!diarize_requested(ctx->model, params)) { return TRANSCRIBE_OK; } if (params->task != TRANSCRIBE_TASK_TRANSCRIBE) { @@ -1506,9 +1639,14 @@ transcribe_status run_batch(transcribe_session * session, transcribe::debug::init(); const auto & hp = cm->hparams; - // Shared prompt affixes (one run_params across the batch). + // Shared prompt affixes (one run_params across the batch), with the + // keyword list fitted as if there were no audio. A row whose own room is + // smaller than that fit would get fewer keywords from run(), so the batch + // then goes serial (see fit_terms_and_context: otherwise the fits match). std::vector prefix_ids, suffix_ids; - if (build_granite_affixes(cm, params, prefix_ids, suffix_ids) != TRANSCRIBE_OK) { + size_t n_keyword_tokens = 0; + if (build_granite_affixes(cm, params, granite_keyword_room(granite_context_ceiling(cc->n_ctx, hp), 0), prefix_ids, + suffix_ids, &n_keyword_tokens) != TRANSCRIBE_OK) { return TRANSCRIBE_ERR_INVALID_ARG; } const int prefix_len = static_cast(prefix_ids.size()); @@ -1567,6 +1705,9 @@ transcribe_status run_batch(transcribe_session * session, if (n_audio[b] <= 0) { continue; } + if (n_keyword_tokens > static_cast(std::max(granite_keyword_room(ceiling, n_audio[b]), 0))) { + return run_batch_serial(cc, pcm, n_samples, n, params); + } prompt_ids[b] = prefix_ids; for (int i = 0; i < n_audio[b]; ++i) { prompt_ids[b].push_back(0); // placeholder @@ -1604,10 +1745,9 @@ transcribe_status run_batch(transcribe_session * session, return TRANSCRIBE_OK; } n_audio_max = std::max(1, n_audio_max); - const int max_new = transcribe::pick_decode_budget( - transcribe::predict_transcript_tokens(n_audio_max, cm->limits.ms_per_audio_token), k_gen_reserve, max_T_prompt, - ceiling); - int max_n_kv = 1024; + const int max_new = transcribe::pick_decode_budget(granite_predicted_tokens(cm, params, n_audio_max), k_gen_reserve, + max_T_prompt, ceiling); + int max_n_kv = 1024; while (max_n_kv < max_T_prompt + max_new) { max_n_kv *= 2; } diff --git a/src/arch/moss/model.cpp b/src/arch/moss/model.cpp index 2cfb8a4f8..e2ead141b 100644 --- a/src/arch/moss/model.cpp +++ b/src/arch/moss/model.cpp @@ -24,6 +24,7 @@ #include "transcribe-log.h" #include "transcribe-mel.h" #include "transcribe-meta.h" +#include "transcribe-prompting.h" #include "transcribe-repetition-guard.h" #include "weights.h" @@ -171,10 +172,11 @@ void build_audio_span(const MossHParams & hp, } } -void build_prompt_tokens(const MossHParams & hp, - int audio_seq_len, - std::vector & out_ids, - std::vector & out_audio_positions) { +void build_prompt_tokens(const MossHParams & hp, + int audio_seq_len, + std::vector & out_ids, + std::vector & out_audio_positions, + const std::vector * suffix) { out_ids.clear(); out_audio_positions.clear(); @@ -189,14 +191,86 @@ void build_prompt_tokens(const MossHParams & hp, out_audio_positions.push_back(prefix_len + off); } - out_ids.insert(out_ids.end(), hp.prompt_suffix_tokens.begin(), hp.prompt_suffix_tokens.end()); + const std::vector & tail = suffix != nullptr ? *suffix : hp.prompt_suffix_tokens; + out_ids.insert(out_ids.end(), tail.begin(), tail.end()); } namespace { +// True when the GGUF carries the instruction split and it reproduces the +// baked suffix with this tokenizer, i.e. the runtime can re-encode the +// instruction with a hotword list appended. +bool moss_supports_hotwords(const MossModel & m) { + const MossHParams & hp = m.hparams; + if (hp.prompt_instruction.empty() || !m.tok.has_encoder()) { + return false; + } + std::vector ids = hp.prompt_instruction_head_tokens; + std::vector instr; + if (m.tok.encode(hp.prompt_instruction, instr) != TRANSCRIBE_OK) { + return false; + } + ids.insert(ids.end(), instr.begin(), instr.end()); + ids.insert(ids.end(), hp.prompt_instruction_tail_tokens.begin(), hp.prompt_instruction_tail_tokens.end()); + return ids == hp.prompt_suffix_tokens; +} + +// Prompt suffix for the run: the baked one, or with the generic vocabulary +// appended to the instruction as upstream's hotword hint +// (examples/prompts.md) in the instruction's language: "热词提示:{terms}" +// after a Chinese instruction, " Hotwords: {terms}" otherwise (", "-joined). +// Terms are fitted to `budget` tokens; `n_hint_tokens` receives how many the +// fitted hint took (0 without terms). +transcribe_status moss_prompt_suffix(const MossModel & m, + const transcribe_run_params * params, + int budget, + std::vector & out, + size_t & n_hint_tokens) { + const MossHParams & hp = m.hparams; + const std::vector terms = transcribe::prompting::terms(params); + out = hp.prompt_suffix_tokens; + n_hint_tokens = 0; + if (terms.empty()) { + return TRANSCRIBE_OK; + } + bool cjk = false; + for (size_t i = 0; i + 2 < hp.prompt_instruction.size() && !cjk; ++i) { + const unsigned char b = static_cast(hp.prompt_instruction[i]); + cjk = b >= 0xE4 && b <= 0xE9; // lead bytes of U+4E00..U+9FFF + } + const std::string lead = + cjk ? "\xE7\x83\xAD\xE8\xAF\x8D\xE6\x8F\x90\xE7\xA4\xBA\xEF\xBC\x9A" /* 热词提示: */ : " Hotwords: "; + transcribe::prompting::FittedPrompt fit; + if (const transcribe_status st = + transcribe::prompting::fit_terms_and_context(m.tok, terms, { lead, ", ", "" }, "", budget, "moss run", fit); + st != TRANSCRIBE_OK) { + return st; + } + n_hint_tokens = fit.n_tokens(); + if (fit.n_terms == 0) { + return TRANSCRIBE_OK; + } + std::vector instr; + if (const transcribe_status st = m.tok.encode(hp.prompt_instruction + fit.terms_text, instr); st != TRANSCRIBE_OK) { + return st; + } + out = hp.prompt_instruction_head_tokens; + out.insert(out.end(), instr.begin(), instr.end()); + out.insert(out.end(), hp.prompt_instruction_tail_tokens.begin(), hp.prompt_instruction_tail_tokens.end()); + return TRANSCRIBE_OK; +} + constexpr const char k_default_variant[] = "moss-transcribe-diarize"; constexpr int k_max_new = 256; +// Token room for the hotword hint: what the context window leaves after the +// generation reserve and a prompt of `base_prompt_len` tokens built with the +// baked suffix. run() passes its clip's prompt; run_batch() passes one +// without audio. +int moss_hint_budget(int ceiling, int base_prompt_len) { + return ceiling - k_max_new - base_prompt_len; +} + int moss_context_ceiling(int32_t n_ctx_knob, const MossHParams & hp) { int ceiling = hp.dec_max_position_embeddings; if (n_ctx_knob > 0 && n_ctx_knob < ceiling) { @@ -233,6 +307,9 @@ transcribe_status load(Loader & loader, const transcribe_model_load_params * par if (const transcribe_status st = read_moss_hparams(loader.gguf(), m->hparams); st != TRANSCRIBE_OK) { return st; } + // Generic vocabulary needs the instruction split, which GGUFs converted + // before it lack; those keep the fixed prompt and do not advertise it. + transcribe::set_feature(m.get(), TRANSCRIBE_FEATURE_VOCABULARY, moss_supports_hotwords(*m)); m->hparams.vocab_size = m->tok.n_tokens(); m->hparams.bos_token_id = m->tok.bos_id(); @@ -748,6 +825,20 @@ transcribe_status run(transcribe_session * session, std::vector prompt_ids; std::vector audio_positions; build_prompt_tokens(cm->hparams, T_enc, prompt_ids, audio_positions); + if (params != nullptr && params->n_vocabulary > 0) { + // The suffix follows the audio, so swapping it leaves audio_positions. + std::vector suffix; + size_t n_hint_tokens = 0; + if (const transcribe_status st = moss_prompt_suffix( + *cm, params, + moss_hint_budget(moss_context_ceiling(cc->n_ctx, cm->hparams), static_cast(prompt_ids.size())), + suffix, n_hint_tokens); + st != TRANSCRIBE_OK) { + return st; + } + prompt_ids.resize(prompt_ids.size() - cm->hparams.prompt_suffix_tokens.size()); + prompt_ids.insert(prompt_ids.end(), suffix.begin(), suffix.end()); + } const int T_prompt = static_cast(prompt_ids.size()); if (static_cast(audio_positions.size()) != T_enc) { log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "moss run: audio_positions(%zu) != T_enc(%d)", audio_positions.size(), @@ -1089,6 +1180,20 @@ transcribe_status run_batch(transcribe_session * session, // length is a pure function of the sample count, so predict it here — no // encoder pass needed — and hand the whole batch to the serial path, // which goes through run() and therefore chunks. + // Shared hotword-extended suffix (one run_params per batch), fitted as if + // there were no audio. A row whose own budget is smaller than that fit + // would get fewer hotwords from run(), so the batch then goes serial (see + // fit_terms_and_context: otherwise the fits match); checked below with + // the predicted prompt length. + const int ceiling = moss_context_ceiling(cc->n_ctx, cm->hparams); + std::vector suffix; + size_t n_hint_tokens = 0; + if (moss_prompt_suffix(*cm, params, + moss_hint_budget(ceiling, static_cast(cm->hparams.prompt_prefix_tokens.size() + + cm->hparams.prompt_suffix_tokens.size())), + suffix, n_hint_tokens) != TRANSCRIBE_OK) { + return run_batch_serial(cc, pcm, n_samples, n, params); + } { const int chunk_size = causal_lm::prefill_chunk_size(); for (int b = 0; b < n; ++b) { @@ -1096,7 +1201,11 @@ transcribe_status run_batch(transcribe_session * session, continue; } std::vector ids, positions; - build_prompt_tokens(cm->hparams, audio_token_length(n_samples[b], cm->hparams), ids, positions); + build_prompt_tokens(cm->hparams, audio_token_length(n_samples[b], cm->hparams), ids, positions, &suffix); + const int base_len = static_cast(ids.size() - suffix.size() + cm->hparams.prompt_suffix_tokens.size()); + if (static_cast(n_hint_tokens) > std::max(moss_hint_budget(ceiling, base_len), 0)) { + return run_batch_serial(cc, pcm, n_samples, n, params); + } if (static_cast(ids.size()) > chunk_size) { log_msg(TRANSCRIBE_LOG_LEVEL_DEBUG, "moss run_batch: utterance %d needs %zu prompt tokens (> %d) — running the batch serially so " @@ -1120,9 +1229,8 @@ transcribe_status run_batch(transcribe_session * session, std::vector fail_status(n, TRANSCRIBE_ERR_INVALID_ARG); int64_t mel_us = 0, enc_us = 0; - const int ceiling = moss_context_ceiling(cc->n_ctx, cm->hparams); - int max_T_prompt = 0; - int max_T_enc = 0; + int max_T_prompt = 0; + int max_T_enc = 0; for (int b = 0; b < n; ++b) { if (cc->poll_abort()) { return TRANSCRIBE_ERR_ABORTED; @@ -1138,7 +1246,7 @@ transcribe_status run_batch(transcribe_session * session, continue; } T_enc[b] = te; - build_prompt_tokens(cm->hparams, te, prompt_ids[b], audio_positions[b]); + build_prompt_tokens(cm->hparams, te, prompt_ids[b], audio_positions[b], &suffix); T_prompt[b] = static_cast(prompt_ids[b].size()); if (T_prompt[b] + k_max_new > ceiling) { fail_status[b] = TRANSCRIBE_ERR_INPUT_TOO_LONG; diff --git a/src/arch/moss/moss.h b/src/arch/moss/moss.h index e09cf3fc5..9eef5aecd 100644 --- a/src/arch/moss/moss.h +++ b/src/arch/moss/moss.h @@ -55,10 +55,12 @@ void build_audio_span(const MossHParams & hp, // out_audio_positions holds the absolute prompt positions of the audio_pad // tokens (in order), so the b-th audio feature is scattered to // input_ids[out_audio_positions[b]]. -void build_prompt_tokens(const MossHParams & hp, - int audio_seq_len, - std::vector & out_ids, - std::vector & out_audio_positions); +// `suffix` overrides hp.prompt_suffix_tokens (the hotword-extended suffix). +void build_prompt_tokens(const MossHParams & hp, + int audio_seq_len, + std::vector & out_ids, + std::vector & out_audio_positions, + const std::vector * suffix = nullptr); struct MossModel final : public transcribe_model { Tokenizer tok; diff --git a/src/arch/moss/weights.cpp b/src/arch/moss/weights.cpp index d02b3ca7c..b05e1b9e9 100644 --- a/src/arch/moss/weights.cpp +++ b/src/arch/moss/weights.cpp @@ -170,6 +170,20 @@ transcribe_status read_moss_hparams(const gguf_context * gguf, MossHParams & hp) if (auto st = read_required_i32_array(gguf, "stt.moss.digit_tokens", hp.digit_tokens); st != TRANSCRIBE_OK) { return st; } + // Optional instruction split (generic vocabulary); absent on GGUFs that + // predate it, which then keep the fixed prompt. + if (auto st = read_optional_string_kv(gguf, "stt.moss.prompt_instruction", kFamilyTag, "", hp.prompt_instruction); + st != TRANSCRIBE_OK) { + return st; + } + if (read_int32_array_kv(gguf, "stt.moss.prompt_instruction_head_tokens", hp.prompt_instruction_head_tokens) != + KvResult::Ok || + read_int32_array_kv(gguf, "stt.moss.prompt_instruction_tail_tokens", hp.prompt_instruction_tail_tokens) != + KvResult::Ok) { + hp.prompt_instruction.clear(); + hp.prompt_instruction_head_tokens.clear(); + hp.prompt_instruction_tail_tokens.clear(); + } if (hp.digit_tokens.size() != 10) { log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "moss: stt.moss.digit_tokens must have 10 entries, got %zu", hp.digit_tokens.size()); diff --git a/src/arch/moss/weights.h b/src/arch/moss/weights.h index fcfd540b5..eaf3318da 100644 --- a/src/arch/moss/weights.h +++ b/src/arch/moss/weights.h @@ -65,6 +65,12 @@ struct MossHParams { std::vector prompt_prefix_tokens; std::vector prompt_suffix_tokens; std::vector digit_tokens; // ids for '0'..'9' + // The suffix split around the instruction text (newer GGUFs; empty on + // older ones): suffix == head + encode(instruction) + tail. Lets the + // runtime append a hotword list to the instruction. + std::string prompt_instruction; + std::vector prompt_instruction_head_tokens; + std::vector prompt_instruction_tail_tokens; // Token ids (resolved from tokenizer KV at load). int32_t bos_token_id = -1; diff --git a/src/arch/qwen3_asr/capabilities.cpp b/src/arch/qwen3_asr/capabilities.cpp index 4737e5883..1dd90c977 100644 --- a/src/arch/qwen3_asr/capabilities.cpp +++ b/src/arch/qwen3_asr/capabilities.cpp @@ -20,6 +20,11 @@ void apply_family_invariants(transcribe_model & model) { // Cancellation is wired at the per-run level. No PNC/ITN toggle; the // Whisper-specific features do not apply here. transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CANCELLATION, true); + // Generic vocabulary and context prompt both go in the system message, + // the model's only context slot (prompting A/B, + // notes/prompting-ab-results.md). + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_VOCABULARY, true); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CONTEXT_PROMPT, true); } } // namespace transcribe::qwen3_asr diff --git a/src/arch/qwen3_asr/model.cpp b/src/arch/qwen3_asr/model.cpp index e8e62204c..a8d6d977d 100644 --- a/src/arch/qwen3_asr/model.cpp +++ b/src/arch/qwen3_asr/model.cpp @@ -18,6 +18,7 @@ #include "transcribe-log.h" #include "transcribe-mel.h" #include "transcribe-meta.h" +#include "transcribe-prompting.h" #include "transcribe-repetition-guard.h" #include "weights.h" @@ -364,12 +365,14 @@ transcribe_status resolve_chat_tokens(const transcribe::Tokenizer & tok, ChatTok // <|im_start|>user\n<|audio_start|><|audio_pad|>*T_enc<|audio_end|><|im_end|>\n // <|im_start|>assistant\n[language {Name}]? // -// System prompt is empty. A non-null `lang_prefix_ids` (resolved via +// The system message carries the generic prompting context (`system_ids`, +// empty by default). A non-null `lang_prefix_ids` (resolved via // encode_language_prefix) is appended after the trailing newline to force an // output language; kept out of here so this stays a pure token-id assembler. void build_prompt_tokens(const QwenAsrHParams & hp, const ChatTokens & ct, int T_enc, + const std::vector & system_ids, const std::vector * lang_prefix_ids, std::vector & out_ids, std::vector & out_audio_positions) { @@ -379,6 +382,7 @@ void build_prompt_tokens(const QwenAsrHParams & hp, out_ids.push_back(ct.im_start); out_ids.push_back(ct.role_system); out_ids.push_back(ct.newline); + out_ids.insert(out_ids.end(), system_ids.begin(), system_ids.end()); out_ids.push_back(ct.im_end); out_ids.push_back(ct.newline); @@ -406,6 +410,45 @@ void build_prompt_tokens(const QwenAsrHParams & hp, } } +// Ids build_prompt_tokens emits besides the system ids, the audio pads and +// the language prefix (the role / newline / audio-boundary tokens above). +constexpr int k_chat_frame_tokens = 15; + +// Length build_prompt_tokens produces with an empty system message, so the +// system context can be budgeted before the prompt is built. +int prompt_tokens_without_system(int T_enc, const std::vector * lang_prefix_ids) { + return k_chat_frame_tokens + T_enc + (lang_prefix_ids != nullptr ? static_cast(lang_prefix_ids->size()) : 0); +} + +// Generic prompting -> system-message ids: vocabulary joined " " (measured +// better than ", ": fewer whole-dictionary dumps into the output), then +// " " + prompt verbatim. `budget` is the room the context window leaves +// after the rest of the prompt and the generation reserve; overflow trims +// the context first, then terms (see fit_terms_and_context). +transcribe_status encode_system_context(const transcribe::Tokenizer & tok, + const transcribe_run_params * params, + int budget, + std::vector & out) { + out.clear(); + if (params == nullptr) { + return TRANSCRIBE_OK; + } + const std::vector terms = transcribe::prompting::terms(params); + std::string ctx = params->prompt != nullptr ? params->prompt : ""; + if (!terms.empty() && !ctx.empty()) { + ctx = " " + ctx; + } + transcribe::prompting::FittedPrompt fit; + if (const transcribe_status st = transcribe::prompting::fit_terms_and_context( + tok, terms, { "", " ", "" }, ctx, std::max(budget, 0), "qwen3_asr run", fit); + st != TRANSCRIBE_OK) { + return st; + } + out = std::move(fit.term_ids); + out.insert(out.end(), fit.ctx_ids.begin(), fit.ctx_ids.end()); + return TRANSCRIBE_OK; +} + } // namespace // below; encode_language_prefix matches the qwen3_asr.h declaration.) @@ -721,7 +764,15 @@ transcribe_status run(transcribe_session * session, // Prompt construction. std::vector prompt_ids; std::vector audio_positions; - build_prompt_tokens(cm->hparams, cm->chat_tokens, T_enc, lang_prefix_ptr, prompt_ids, audio_positions); + const int ceiling = qwen3_context_ceiling(cc->n_ctx, cm->hparams); + std::vector system_ids; + if (const transcribe_status st = encode_system_context( + cm->tok, params, ceiling - k_gen_reserve - prompt_tokens_without_system(T_enc, lang_prefix_ptr), + system_ids); + st != TRANSCRIBE_OK) { + return st; + } + build_prompt_tokens(cm->hparams, cm->chat_tokens, T_enc, system_ids, lang_prefix_ptr, prompt_ids, audio_positions); const int T_prompt = static_cast(prompt_ids.size()); const int prefix_len = audio_positions.empty() ? 0 : static_cast(audio_positions.front()); const int suffix_len = T_prompt - prefix_len - T_enc; @@ -729,7 +780,6 @@ transcribe_status run(transcribe_session * session, // Input-length gate: audio + prompt + generation must fit the decoder // context window. Reject an over-length clip here, before prefill/decode. - const int ceiling = qwen3_context_ceiling(cc->n_ctx, cm->hparams); if (T_prompt + k_gen_reserve > ceiling) { transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "qwen3_asr run: input too long — %d audio + %d prompt tokens " @@ -1523,6 +1573,19 @@ transcribe_status run_batch(transcribe_session * session, lang_prefix_ptr = &lang_prefix_ids; } + // Shared system context (vocabulary / prompt), one run_params per batch, + // fitted as if there were no audio. A row whose own budget is smaller than + // that fit would get a different system context from run(), so the batch + // then goes serial (see fit_terms_and_context: otherwise the fits match). + const int ceiling = qwen3_context_ceiling(cc->n_ctx, cm->hparams); + auto system_budget = [&](int T_enc) { + return std::max(ceiling - k_gen_reserve - prompt_tokens_without_system(T_enc, lang_prefix_ptr), 0); + }; + std::vector system_ids; + if (encode_system_context(cm->tok, params, system_budget(0), system_ids) != TRANSCRIBE_OK) { + return run_batch_serial(cc, pcm, n_samples, n, params); + } + // Pass 1: per-utterance encoder + prefill into KV slabs. std::vector> generated(n); std::vector T_prompt(n, 0); @@ -1550,15 +1613,17 @@ transcribe_status run_batch(transcribe_session * session, int prefix_len = 0; // Per-utterance terminal status for rejected rows. Defaults to INVALID_ARG; // over-length rows below are upgraded to INPUT_TOO_LONG. - const int ceiling = qwen3_context_ceiling(cc->n_ctx, cm->hparams); std::vector fail_status(n, TRANSCRIBE_ERR_INVALID_ARG); std::vector> prompt_ids(n); for (int b = 0; b < n; ++b) { if (!valid[b]) { continue; } + if (static_cast(system_ids.size()) > system_budget(T_enc[b])) { + return run_batch_serial(cc, pcm, n_samples, n, params); + } std::vector ap; - build_prompt_tokens(cm->hparams, cm->chat_tokens, T_enc[b], lang_prefix_ptr, prompt_ids[b], ap); + build_prompt_tokens(cm->hparams, cm->chat_tokens, T_enc[b], system_ids, lang_prefix_ptr, prompt_ids[b], ap); T_prompt[b] = static_cast(prompt_ids[b].size()); prefix_len = ap.empty() ? 0 : static_cast(ap.front()); // Same gate as single-shot run(); the rest of the batch still runs. diff --git a/src/arch/voxtral/capabilities.cpp b/src/arch/voxtral/capabilities.cpp index 58e691081..4de2f81c4 100644 --- a/src/arch/voxtral/capabilities.cpp +++ b/src/arch/voxtral/capabilities.cpp @@ -17,15 +17,18 @@ void apply_family_invariants(transcribe_model & model) { // a task token: translation, Q&A and summarization all go through the // mistral-common instruct template (audio + free-text instruction). // The runtime exposes it via --translate / --target-language (which - // synthesize a translate instruction) and a general free-text prompt. + // synthesize a translate instruction). // The GGUF's stt.capability.translate=true is read into supports_translate // by read_capability_kv at load; default it true here as a fallback. caps.supports_translate = true; - // Per-run cancellation; the free-text prompt is wired via the - // INITIAL_PROMPT feature on the Voxtral run extension. + // Per-run cancellation. No INITIAL_PROMPT: Voxtral has no run extension + // that accepts a prompt. transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CANCELLATION, true); - transcribe::set_feature(&model, TRANSCRIBE_FEATURE_INITIAL_PROMPT, true); + // TRANSCRIBE_TASK_INSTRUCT: the caller's prompt through the same chat + // path as translation. Both 2507 sizes follow free-text instructions + // (prompting A/B, notes/prompting-ab-results.md). + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_INSTRUCT, true); } } // namespace transcribe::voxtral diff --git a/src/arch/voxtral/model.cpp b/src/arch/voxtral/model.cpp index f5f973b0c..f4462ec43 100644 --- a/src/arch/voxtral/model.cpp +++ b/src/arch/voxtral/model.cpp @@ -26,6 +26,7 @@ #include "transcribe-log.h" #include "transcribe-mel.h" #include "transcribe-meta.h" +#include "transcribe-prompting.h" #include "transcribe-repetition-guard.h" #include "voxtral.h" #include "weights.h" @@ -325,6 +326,25 @@ transcribe_status build_transcription_prompt(const VoxtralModel & m, return TRANSCRIBE_OK; } +// Instruct-template text for the run, if any: TRANSLATE synthesizes +// "Translate this to {Language}."; TRANSCRIBE_TASK_INSTRUCT sends the caller's +// prompt verbatim (the mistral-common chat path: audio, then the text, with +// no [TRANSCRIBE] token). Returns false for plain transcription. +bool instruct_instruction(const transcribe_run_params * params, std::string & instruction) { + if (params == nullptr) { + return false; + } + if (params->task == TRANSCRIBE_TASK_TRANSLATE) { + instruction = std::string("Translate this to ") + lang_name_for(params->target_language) + "."; + return true; + } + if (params->task == TRANSCRIBE_TASK_INSTRUCT && params->prompt != nullptr) { + instruction = params->prompt; + return true; + } + return false; +} + // Build the instruct prompt: audio + BPE(instruction) + [/INST]. transcribe_status build_instruct_prompt(const VoxtralModel & m, const std::string & instruction, @@ -342,7 +362,8 @@ transcribe_status build_instruct_prompt(const VoxtralModel & m, out_ids.push_back(m.hparams.audio_token_id); } std::vector instr_ids; - if (const transcribe_status st = m.tok.encode(instruction, instr_ids); st != TRANSCRIBE_OK) { + if (const transcribe_status st = transcribe::prompting::encode_plain(m.tok, instruction, instr_ids, "prompt"); + st != TRANSCRIBE_OK) { log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "voxtral: failed to encode instruction text"); return st; } @@ -588,12 +609,8 @@ transcribe_status run(transcribe_session * session, transcribe::debug::init(); // ----- Prompt mode ----- - const bool translate = (params != nullptr && params->task == TRANSCRIBE_TASK_TRANSLATE); std::string instruction; - if (translate) { - const char * tgt = (params != nullptr) ? params->target_language : nullptr; - instruction = std::string("Translate this to ") + lang_name_for(tgt) + "."; - } + const bool use_instruct_prompt = instruct_instruction(params, instruction); if (!cm->mel.has_value()) { log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "voxtral run: model has no MelFrontend"); @@ -732,7 +749,7 @@ transcribe_status run(transcribe_session * session, // ----- Prompt construction ----- std::vector prompt_ids; int prefix_len = 0, suffix_len = 0; - if (translate) { + if (use_instruct_prompt) { if (const transcribe_status st = build_instruct_prompt(*cm, instruction, n_audio_total, prompt_ids, prefix_len, suffix_len); st != TRANSCRIBE_OK) { @@ -764,8 +781,14 @@ transcribe_status run(transcribe_session * session, n_audio_total, T_prompt - n_audio_total, model_max, T_prompt + k_gen_reserve); return TRANSCRIBE_ERR_INPUT_TOO_LONG; } - const int max_new = transcribe::pick_decode_budget(n_audio_total, k_decode_budget_min, T_prompt, model_max); - const int want_ctx = causal_lm::pick_kv_cache_context(T_prompt + max_new, model_max); + // A transcript's length follows the audio, which sets both the budget and + // the KV size. Free text (TRANSCRIBE_TASK_INSTRUCT) does not: it runs until + // EOS, a repetition stop or the context ceiling, growing the KV cache on + // demand from the same starting size. + const bool free_text = params != nullptr && params->task == TRANSCRIBE_TASK_INSTRUCT; + const int predicted = transcribe::pick_decode_budget(n_audio_total, k_decode_budget_min, T_prompt, model_max); + const int max_new = free_text ? model_max - T_prompt : predicted; + const int want_ctx = causal_lm::pick_kv_cache_context(T_prompt + predicted, model_max); if (cc->kv_cache.n_ctx < want_ctx) { const ggml_type kv_type = (cc->kv_type == TRANSCRIBE_KV_TYPE_F32) ? GGML_TYPE_F32 : GGML_TYPE_F16; cc->kv_cache.free(); @@ -910,18 +933,18 @@ transcribe_status run(transcribe_session * session, int cur_past = T_prompt; int max_n_kv = 1024; - while (max_n_kv < T_prompt + max_new) { + while (max_n_kv < T_prompt + std::min(max_new, predicted)) { max_n_kv *= 2; } - if (max_n_kv > cc->kv_cache.n_ctx) { - max_n_kv = cc->kv_cache.n_ctx; - } + max_n_kv = std::min(max_n_kv, cc->kv_cache.n_ctx); - if (cc->compute_ctx != nullptr) { - ggml_free(cc->compute_ctx); - cc->compute_ctx = nullptr; - } - { + // (Re)build the step graph for an attention width of max_n_kv. + StepBuild sb; + const auto build_step = [&]() -> transcribe_status { + if (cc->compute_ctx != nullptr) { + ggml_free(cc->compute_ctx); + cc->compute_ctx = nullptr; + } ggml_init_params ip{}; ip.mem_size = 16 * 1024 * 1024; ip.no_alloc = true; @@ -930,29 +953,54 @@ transcribe_status run(transcribe_session * session, transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "voxtral run: ggml_init (step) failed — out of memory."); return TRANSCRIBE_ERR_OOM; } + sb = build_step_graph(cc->compute_ctx, cm->weights, cm->hparams, cc->kv_cache, max_n_kv, cc->decoder_use_flash); + if (sb.graph == nullptr || sb.out == nullptr) { + return TRANSCRIBE_ERR_GGUF; + } + ggml_backend_sched_reset(cc->sched); + if (!ggml_backend_sched_alloc_graph(cc->sched, sb.graph)) { + transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, + "voxtral run: step graph allocation failed — out of memory. " + "Lower transcribe_session_params.n_ctx or shorten the audio."); + return TRANSCRIBE_ERR_OOM; + } + set_sched_threads(cc->sched, cc->n_threads); + return TRANSCRIBE_OK; + }; + if (const transcribe_status st = build_step(); st != TRANSCRIBE_OK) { + return st; } - StepBuild sb = - build_step_graph(cc->compute_ctx, cm->weights, cm->hparams, cc->kv_cache, max_n_kv, cc->decoder_use_flash); - if (sb.graph == nullptr || sb.out == nullptr) { - return TRANSCRIBE_ERR_GGUF; - } - ggml_backend_sched_reset(cc->sched); - if (!ggml_backend_sched_alloc_graph(cc->sched, sb.graph)) { - transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, - "voxtral run: step graph allocation failed — out of memory. " - "Lower transcribe_session_params.n_ctx or shorten the audio."); - return TRANSCRIBE_ERR_OOM; - } - set_sched_threads(cc->sched, cc->n_threads); const ggml_fp16_t mz = ggml_fp32_to_fp16(0.0f); const ggml_fp16_t mn = ggml_fp32_to_fp16(-INFINITY); std::vector step_mask(max_n_kv, mn); bool repeating = false; - while (next_tok != eos_id && static_cast(generated_ids.size()) < max_new && cur_past + 1 <= max_n_kv) { + while (next_tok != eos_id && static_cast(generated_ids.size()) < max_new) { if (cc->poll_abort()) { return TRANSCRIBE_ERR_ABORTED; } + if (cur_past + 1 > max_n_kv) { + // Out of attention width: widen (doubling, capped at the ceiling), + // growing the KV cache first when it is the limit. + if (max_n_kv >= model_max) { + break; + } + max_n_kv = std::min(max_n_kv * 2, model_max); + if (cc->kv_cache.n_ctx < max_n_kv && + !causal_lm::kv_grow(cc->kv_cache, cm->plan.primary, + causal_lm::pick_kv_cache_context(max_n_kv, model_max), cm->hparams.dec_n_kv_heads, + cm->hparams.dec_head_dim, cm->hparams.dec_n_layers)) { + transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, + "voxtral run: KV cache growth to %d positions failed — out of memory. " + "Lower transcribe_session_params.n_ctx.", + max_n_kv); + return TRANSCRIBE_ERR_OOM; + } + if (const transcribe_status st = build_step(); st != TRANSCRIBE_OK) { + return st; + } + step_mask.resize(max_n_kv, mn); + } ggml_backend_tensor_set(sb.input_id_in, &next_tok, 0, sizeof(int32_t)); const int32_t pos_val = cur_past; ggml_backend_tensor_set(sb.position_in, &pos_val, 0, sizeof(int32_t)); @@ -1067,8 +1115,11 @@ transcribe_status run_batch(transcribe_session * session, } // The batched encoder + causal_lm batched blocks are flash-only; dump mode - // and n==1 take the established single-shot path for byte-parity. - if (!cc->decoder_use_flash || !cc->encoder_use_flash || transcribe::debug::enabled() || n == 1) { + // and n==1 take the established single-shot path for byte-parity. Free + // text (INSTRUCT) grows its KV cache per utterance, which the packed + // batched cache cannot, so it runs serially too. + if (!cc->decoder_use_flash || !cc->encoder_use_flash || transcribe::debug::enabled() || n == 1 || + (params != nullptr && params->task == TRANSCRIBE_TASK_INSTRUCT)) { return run_batch_serial(cc, pcm, n_samples, n, params); } @@ -1104,13 +1155,9 @@ transcribe_status run_batch(transcribe_session * session, } // ----- Prompt mode (uniform across the batch) ----- - const bool translate = (params != nullptr && params->task == TRANSCRIBE_TASK_TRANSLATE); - std::string instruction; - if (translate) { - const char * tgt = (params != nullptr) ? params->target_language : nullptr; - instruction = std::string("Translate this to ") + lang_name_for(tgt) + "."; - } - const char * lang = (params != nullptr) ? params->language : nullptr; + std::string instruction; + const bool use_instruct_prompt = instruct_instruction(params, instruction); + const char * lang = (params != nullptr) ? params->language : nullptr; // ----- Chunk geometry ----- int samples_per_chunk = hp.fe_n_samples; @@ -1271,14 +1318,14 @@ transcribe_status run_batch(transcribe_session * session, const int ctx_ceiling = voxtral_context_ceiling(cc->n_ctx, hp); std::vector> prompt_ids(n); std::vector T_prompt(n, 0), T_audio(n, 0); - int prefix_len = 0, suffix_len = 0; + int prefix_len = 0; for (int b = 0; b < n; ++b) { if (!valid[b]) { continue; } const int n_audio = n_chunks[b] * audio_per_chunk; int pfx = 0, sfx = 0; - const transcribe_status st = translate ? + const transcribe_status st = use_instruct_prompt ? build_instruct_prompt(*cm, instruction, n_audio, prompt_ids[b], pfx, sfx) : build_transcription_prompt(*cm, lang, n_audio, prompt_ids[b], pfx, sfx); if (st != TRANSCRIBE_OK) { @@ -1298,8 +1345,7 @@ transcribe_status run_batch(transcribe_session * session, over_length[b] = 1; continue; } - prefix_len = pfx; - suffix_len = sfx; // uniform across the batch + prefix_len = pfx; // uniform across the batch T_prompt[b] = t_prompt; T_audio[b] = n_audio; } diff --git a/src/arch/whisper/capabilities.cpp b/src/arch/whisper/capabilities.cpp index 565a9dafe..4390d31f9 100644 --- a/src/arch/whisper/capabilities.cpp +++ b/src/arch/whisper/capabilities.cpp @@ -26,6 +26,13 @@ void apply_family_invariants(transcribe_model & model) { transcribe::set_feature(&model, TRANSCRIBE_FEATURE_TEMPERATURE_FALLBACK, true); transcribe::set_feature(&model, TRANSCRIBE_FEATURE_LONG_FORM, true); transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CANCELLATION, true); + // Generic vocabulary (`Glossary: {terms}`) and context prompt, both in the + // <|startofprev|> slot (prompting A/B, notes/prompting-ab-results.md). + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_VOCABULARY, true); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CONTEXT_PROMPT, true); + // Transcript prefix after the SOT sequence (openai DecodingOptions.prefix), + // first window only. + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX, true); } } // namespace transcribe::whisper diff --git a/src/arch/whisper/model.cpp b/src/arch/whisper/model.cpp index b7c82aeab..66e855ac9 100644 --- a/src/arch/whisper/model.cpp +++ b/src/arch/whisper/model.cpp @@ -19,6 +19,7 @@ #include "transcribe-loader.h" #include "transcribe-log.h" #include "transcribe-meta.h" +#include "transcribe-prompting.h" #include "weights.h" #include "whisper.h" @@ -28,6 +29,7 @@ #include "third_party/miniz/miniz.h" #include +#include #include #include #include @@ -943,6 +945,80 @@ transcribe_status load_mel_from_ref(const char * ref_dir, int n_mels, int n_mel_ namespace { +bool whisper_has_generic_prompt(const transcribe_run_params * params) { + return params != nullptr && (params->n_vocabulary > 0 || transcribe::prompting::has_text(params->prompt)); +} + +// Token cap on the <|startofprev|> slot: the run extension's +// max_prev_context_tokens, else half the decoder window minus one (openai). +int whisper_prev_cap(const transcribe_whisper_run_ext & wp, const WhisperHParams & hp) { + return wp.max_prev_context_tokens > 0 ? wp.max_prev_context_tokens : hp.dec_max_target_positions / 2 - 1; +} + +// Transcript prefix ids (openai DecodingOptions.prefix): " " + strip(prefix), +// plain text only. Empty when the prefix is absent or all whitespace. Shared +// by whisper_run and whisper_run_validate so the validated length is the one +// that runs. +transcribe_status whisper_prefix_ids(const WhisperModel & cm, const char * prefix, std::vector & out) { + out.clear(); + const std::string text = prefix != nullptr ? transcribe::prompting::strip(prefix) : std::string(); + if (text.empty()) { + return TRANSCRIBE_OK; + } + if (const transcribe_status st = transcribe::prompting::encode_plain(cm.tok, " " + text, out, "prefix"); + st != TRANSCRIBE_OK) { + return st; + } + const int eos_id = cm.tok.eos_id() >= 0 ? cm.tok.eos_id() : 50257; // as whisper_run + for (int32_t id : out) { + if (id >= eos_id) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "whisper run: prefix encodes to special token id %d", id); + return TRANSCRIBE_ERR_INVALID_ARG; + } + } + return TRANSCRIBE_OK; +} + +// Generic prompting (transcribe_run_params::vocabulary / prompt) rendered into +// the <|startofprev|> slot, as text-only ids (the caller prepends the marker). +// Text is `Glossary: {terms}` (", "-joined), then " " + prompt, tokenized in +// HF get_prompt_ids form (" " + text.strip()). The slot holds at most +// `budget` tokens; Whisper natively keeps the LAST tokens, which would drop +// the highest-priority terms, so overflow is resolved here instead: context +// is trimmed first (keeping its most recent tokens), then terms are dropped +// from the end of the list, with one WARN naming what was dropped. Returns +// INVALID_ARG on control-token literals; an empty `out` means nothing to prime. +transcribe_status whisper_generic_prompt_ids(const WhisperModel & cm, + const transcribe_run_params * params, + int budget, + int eos_id, + std::vector & out) { + out.clear(); + std::string ctx = params->prompt != nullptr ? transcribe::prompting::strip(params->prompt) : std::string(); + if (!ctx.empty()) { + ctx = " " + ctx; + } + transcribe::prompting::FittedPrompt fit; + if (const transcribe_status st = transcribe::prompting::fit_terms_and_context( + cm.tok, transcribe::prompting::terms(params), { " Glossary: ", ", ", "" }, ctx, budget, "whisper run", fit); + st != TRANSCRIBE_OK) { + return st; + } + out = std::move(fit.term_ids); + out.insert(out.end(), fit.ctx_ids.begin(), fit.ctx_ids.end()); + for (int32_t id : out) { + if (id >= eos_id) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "whisper run: prompting text encodes to special token id %d", id); + return TRANSCRIBE_ERR_INVALID_ARG; + } + } + return TRANSCRIBE_OK; +} + +} // namespace + +namespace { + // Whisper timestamp logits processor, shared by whisper_run and // whisper_run_batch so per-step masking is identical. Mirrors transformers' // WhisperTimeStampLogitsProcessor. Mutates `logits` in place; no-op when @@ -1341,8 +1417,13 @@ transcribe_status whisper_run(transcribe_session * session, if (requested_timestamps == TRANSCRIBE_TIMESTAMPS_WORD) { return TRANSCRIBE_ERR_UNSUPPORTED_TIMESTAMPS; } - const bool want_segment_timestamps = - requested_timestamps == TRANSCRIBE_TIMESTAMPS_AUTO || requested_timestamps == TRANSCRIBE_TIMESTAMPS_SEGMENT; + // AUTO resolves to NONE under a transcript prefix: the timestamp rules + // restart after the prefix and force an initial timestamp near 0 s while + // the prefix's speech is still playing, which derails the continuation. + // An explicit SEGMENT request with a prefix is rejected in run_validate. + const bool has_prefix = params != nullptr && transcribe::prompting::has_text(params->prefix); + const bool want_segment_timestamps = (requested_timestamps == TRANSCRIBE_TIMESTAMPS_AUTO && !has_prefix) || + requested_timestamps == TRANSCRIBE_TIMESTAMPS_SEGMENT; // Multilingual variants emit <|lang|> + <|task|> in the decoder prefix; // .en variants have just <|sot|> and no translate/transcribe/language @@ -1523,8 +1604,7 @@ transcribe_status whisper_run(transcribe_session * session, // text-side only); or initial_prompt string, tokenized as HF's // get_prompt_ids form ("<|startofprev|> " + strip) with any special token // (id >= eos_id) in the text rejected (tokenization_whisper.py). - const int max_prev_cap = - wp->max_prev_context_tokens > 0 ? wp->max_prev_context_tokens : (cm->hparams.dec_max_target_positions / 2 - 1); + const int max_prev_cap = whisper_prev_cap(*wp, cm->hparams); std::vector prompt_text_ids; if (wp->prompt_tokens != nullptr && wp->n_prompt_tokens > 0) { // The library prepends <|startofprev|>; a leading prev_sot id from the @@ -1602,17 +1682,55 @@ transcribe_status whisper_run(transcribe_session * session, } } } - // Cap prompt tokens to max_prev_cap (left-truncate, keep most-recent). - if (static_cast(prompt_text_ids.size()) > max_prev_cap) { - prompt_text_ids.erase(prompt_text_ids.begin(), prompt_text_ids.end() - max_prev_cap); + // Transcript prefix (openai DecodingOptions.prefix), right after the SOT + // sequence on the first window only. It sits in the prompt, so the + // timestamp rules (which read generated_ids) begin after it and the result + // text holds only the continuation; raw_text leads with it. + std::vector prefix_ids; + if (const transcribe_status st = whisper_prefix_ids(*cm, params != nullptr ? params->prefix : nullptr, prefix_ids); + st != TRANSCRIBE_OK) { + return st; + } + all_raw_ids.insert(all_raw_ids.end(), prefix_ids.begin(), prefix_ids.end()); + + // Prompt ids capped to `cap` tokens: the extension prompt keeps its most + // recent tokens; the generic vocabulary / context prompt, which share the + // slot (whisper_run_validate rejects them together), are fitted. + const std::vector ext_prompt_ids = std::move(prompt_text_ids); + const auto capped_prompt = [&](int cap, std::vector & out) -> transcribe_status { + if (whisper_has_generic_prompt(params)) { + if (prev_sot_id < 0) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, + "whisper run: model has no <|startofprev|> token; prompting unavailable"); + return TRANSCRIBE_ERR_GGUF; + } + return whisper_generic_prompt_ids(*cm, params, cap, eos_id, out); + } + const size_t keep = std::min(ext_prompt_ids.size(), static_cast(std::max(cap, 0))); + out.assign(ext_prompt_ids.end() - static_cast(keep), ext_prompt_ids.end()); + return TRANSCRIBE_OK; + }; + // Later windows get the full slot. The first window shares it with the + // prefix, so its prompt + prefix stay within max_prev_cap and the rest of + // the decoder window is left for the continuation. + if (const transcribe_status st = capped_prompt(max_prev_cap, prompt_text_ids); st != TRANSCRIBE_OK) { + return st; + } + std::vector first_prompt_ids = prompt_text_ids; + if (!prefix_ids.empty()) { + if (const transcribe_status st = + capped_prompt(max_prev_cap - static_cast(prefix_ids.size()), first_prompt_ids); + st != TRANSCRIBE_OK) { + return st; + } } // History stored as segment token slices (not one flat vector) because // skip_ending_double_timestamps applies per-segment. FIRST_SEGMENT puts the // prompt at the head; ALL_SEGMENTS starts empty and re-prepends per chunk. std::vector> prev_history_segments; - if (wp->prompt_condition == TRANSCRIBE_WHISPER_PROMPT_FIRST_SEGMENT && !prompt_text_ids.empty()) { - prev_history_segments.push_back(prompt_text_ids); + if (wp->prompt_condition == TRANSCRIBE_WHISPER_PROMPT_FIRST_SEGMENT && !first_prompt_ids.empty()) { + prev_history_segments.push_back(first_prompt_ids); } // Per-chunk; HF auto-disables when the previous chunk's accepted @@ -1726,11 +1844,12 @@ transcribe_status whisper_run(transcribe_session * session, // else: empty. // We diverge from HF for FIRST_SEGMENT (the default): prime only the // first window, matching whisper.cpp / OpenAI. - std::vector prev_tokens; + const std::vector & window_prompt = is_first_chunk ? first_prompt_ids : prompt_text_ids; + std::vector prev_tokens; if (do_condition_on_prev_tokens && !prev_history_segments.empty() && prev_sot_id >= 0) { prev_tokens.push_back(prev_sot_id); - if (wp->prompt_condition == TRANSCRIBE_WHISPER_PROMPT_ALL_SEGMENTS && !prompt_text_ids.empty()) { - prev_tokens.insert(prev_tokens.end(), prompt_text_ids.begin(), prompt_text_ids.end()); + if (wp->prompt_condition == TRANSCRIBE_WHISPER_PROMPT_ALL_SEGMENTS && !window_prompt.empty()) { + prev_tokens.insert(prev_tokens.end(), window_prompt.begin(), window_prompt.end()); } std::vector hist; for (const auto & seg : prev_history_segments) { @@ -1744,15 +1863,15 @@ transcribe_status whisper_run(transcribe_session * session, } const int cap = std::min(static_cast(hist.size()), max_prev_cap); prev_tokens.insert(prev_tokens.end(), hist.end() - cap, hist.end()); - } else if (!prompt_text_ids.empty() && prev_sot_id >= 0 && + } else if (!window_prompt.empty() && prev_sot_id >= 0 && (is_first_chunk || wp->prompt_condition == TRANSCRIBE_WHISPER_PROMPT_ALL_SEGMENTS)) { // FIRST_SEGMENT primes the initial prompt on the first window prev_tokens.push_back(prev_sot_id); - prev_tokens.insert(prev_tokens.end(), prompt_text_ids.begin(), prompt_text_ids.end()); + prev_tokens.insert(prev_tokens.end(), window_prompt.begin(), window_prompt.end()); } // Prefix for this chunk: - // multilingual: prev_tokens + [SOT, lang, task, notimestamps?] + // multilingual: prev_tokens + [SOT, lang, task, notimestamps?] + prefix? // .en: prev_tokens + [SOT, notimestamps?] // .en vocab has no <|lang|>/<|task|> tokens; emitting them would land // on a garbage id. @@ -1767,6 +1886,9 @@ transcribe_status whisper_run(transcribe_session * session, if (!want_segment_timestamps) { prompt_ids.push_back(cm->hparams.no_timestamps_token_id); } + if (is_first_chunk) { + prompt_ids.insert(prompt_ids.end(), prefix_ids.begin(), prefix_ids.end()); + } const int seq_len = static_cast(prompt_ids.size()); // Position of the SOT token within the prefix. Used to read the @@ -2617,12 +2739,21 @@ transcribe_status whisper_run_batch(transcribe_session * session, const bool prompt_requested = (wp->prompt_tokens != nullptr && wp->n_prompt_tokens > 0) || (wp->initial_prompt != nullptr && wp->initial_prompt[0] != '\0'); std::vector prev_tokens; - if (prompt_requested) { + if (whisper_has_generic_prompt(params)) { + std::vector ptext; + if (prev_sot_id < 0 || + whisper_generic_prompt_ids(*cm, params, whisper_prev_cap(*wp, hp), eos_id, ptext) != TRANSCRIBE_OK) { + return whisper_run_batch_serial(cc, pcm, n_samples, n, params); + } + if (!ptext.empty()) { + prev_tokens.push_back(prev_sot_id); + prev_tokens.insert(prev_tokens.end(), ptext.begin(), ptext.end()); + } + } else if (prompt_requested) { if (prev_sot_id < 0) { return whisper_run_batch_serial(cc, pcm, n_samples, n, params); } - const int max_prev_cap = - wp->max_prev_context_tokens > 0 ? wp->max_prev_context_tokens : (hp.dec_max_target_positions / 2 - 1); + const int max_prev_cap = whisper_prev_cap(*wp, hp); std::vector ptext; if (wp->prompt_tokens != nullptr && wp->n_prompt_tokens > 0) { if (wp->prompt_tokens[0] == prev_sot_id) { @@ -3392,9 +3523,48 @@ static bool whisper_accepts_ext_kind(const transcribe_model * model, transcribe_ // the snapshot is cleared — an accepted gap, since run() is one-shot with no // accumulating transcript to protect. static transcribe_status whisper_run_validate(const transcribe_session * ctx, const transcribe_run_params * params) { - (void) ctx; - return transcribe_ext_check(params != nullptr ? params->family : nullptr, TRANSCRIBE_EXT_KIND_WHISPER_RUN, - sizeof(struct transcribe_whisper_run_ext)); + if (const transcribe_status st = + transcribe_ext_check(params != nullptr ? params->family : nullptr, TRANSCRIBE_EXT_KIND_WHISPER_RUN, + sizeof(struct transcribe_whisper_run_ext)); + st != TRANSCRIBE_OK) { + return st; + } + // Transcript prefix: with explicit SEGMENT timestamps the timestamp rules + // restart after the prefix and force an initial timestamp while the + // prefix's speech is still playing, which in practice ends the decode + // (AUTO resolves to NONE instead; see whisper_run). And like openai, the + // prefix may take at most half the decoder window. + if (params != nullptr && transcribe::prompting::has_text(params->prefix)) { + if (params->timestamps == TRANSCRIBE_TIMESTAMPS_SEGMENT) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, + "whisper run: a transcript prefix does not compose with segment timestamps; use NONE or AUTO"); + return TRANSCRIBE_ERR_INVALID_ARG; + } + const auto * cm = static_cast(ctx->model); + std::vector ids; + if (const transcribe_status st = whisper_prefix_ids(*cm, params->prefix, ids); st != TRANSCRIBE_OK) { + return st; + } + const int limit = cm->hparams.dec_max_target_positions / 2 - 1; + if (static_cast(ids.size()) > limit) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "whisper run: transcript prefix is %zu tokens; the limit is %d", + ids.size(), limit); + return TRANSCRIBE_ERR_INVALID_ARG; + } + } + // The generic prompting fields and the extension's own prompt fill the + // same <|startofprev|> slot; there is no documented way to merge them. + if (whisper_has_generic_prompt(params) && params->family != nullptr) { + const auto * wp = reinterpret_cast(params->family); + if ((wp->prompt_tokens != nullptr && wp->n_prompt_tokens > 0) || + (wp->initial_prompt != nullptr && wp->initial_prompt[0] != '\0')) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, + "whisper run: generic vocabulary/prompt cannot be combined with the whisper extension's " + "initial_prompt/prompt_tokens"); + return TRANSCRIBE_ERR_INVALID_ARG; + } + } + return TRANSCRIBE_OK; } } // namespace diff --git a/src/causal_lm/causal_lm.cpp b/src/causal_lm/causal_lm.cpp index 284aa780a..4db8c0730 100644 --- a/src/causal_lm/causal_lm.cpp +++ b/src/causal_lm/causal_lm.cpp @@ -105,6 +105,33 @@ bool kv_init(KvCache & cache, return true; } +bool kv_grow(KvCache & cache, ggml_backend_t backend, int n_ctx, int n_kv_heads, int head_dim, int n_layer) { + if (cache.self_k == nullptr || cache.n_batch != 1 || n_ctx <= cache.n_ctx) { + return false; + } + KvCache grown; + if (!kv_init(grown, backend, n_ctx, n_kv_heads, head_dim, n_layer, cache.self_k->type)) { + return false; + } + // Layout is layer-major (layer, position, head, dim), so each layer's + // filled rows move to a new offset; copy them layer by layer. + const size_t row = static_cast(n_kv_heads) * head_dim * ggml_type_size(cache.self_k->type); + const size_t keep = static_cast(std::min(cache.n, cache.n_ctx)) * row; + std::vector host(keep); + for (ggml_tensor * const * pair : { &cache.self_k, &cache.self_v }) { + ggml_tensor * dst = pair == &cache.self_k ? grown.self_k : grown.self_v; + for (int l = 0; l < n_layer && keep > 0; ++l) { + ggml_backend_tensor_get(*pair, host.data(), static_cast(l) * cache.n_ctx * row, keep); + ggml_backend_tensor_set(dst, host.data(), static_cast(l) * n_ctx * row, keep); + } + } + grown.n = cache.n; + grown.head = cache.head; + cache.free(); + cache = grown; + return true; +} + bool kv_init_batched(KvCache & cache, ggml_backend_t backend, int n_ctx, diff --git a/src/causal_lm/causal_lm.h b/src/causal_lm/causal_lm.h index b2129d86e..1992f86bf 100644 --- a/src/causal_lm/causal_lm.h +++ b/src/causal_lm/causal_lm.h @@ -101,6 +101,12 @@ bool kv_init_batched(KvCache & cache, int n_batch, ggml_type kv_type); +// Grow a single-utterance cache (n_batch == 1) to n_ctx positions in place: +// a new cache is allocated, rows [0, cache.n) of every layer are copied over, +// and the fill / write head carry across. For decodes whose length is not +// known up front. On failure (allocation) the old cache is left intact. +bool kv_grow(KvCache & cache, ggml_backend_t backend, int n_ctx, int n_kv_heads, int head_dim, int n_layer); + struct BlockOpts { bool use_flash = true; diff --git a/src/transcribe-prompting.cpp b/src/transcribe-prompting.cpp new file mode 100644 index 000000000..312aa5ab9 --- /dev/null +++ b/src/transcribe-prompting.cpp @@ -0,0 +1,168 @@ +// transcribe-prompting.cpp - shared helpers for the generic prompting fields. + +#include "transcribe-prompting.h" + +#include "transcribe-log.h" +#include "transcribe-tokenizer.h" + +#include +#include +#include + +namespace transcribe::prompting { + +std::vector terms(const transcribe_run_params * p) { + std::vector out; + if (p == nullptr || p->vocabulary == nullptr || p->n_vocabulary <= 0) { + return out; + } + out.reserve(static_cast(p->n_vocabulary)); + for (int32_t i = 0; i < p->n_vocabulary; ++i) { + if (has_text(p->vocabulary[i])) { + out.emplace_back(p->vocabulary[i]); + } + } + return out; +} + +std::string strip(const std::string & s) { + size_t a = 0, b = s.size(); + while (a < b && std::isspace(static_cast(s[a]))) { + ++a; + } + while (b > a && std::isspace(static_cast(s[b - 1]))) { + --b; + } + return s.substr(a, b - a); +} + +std::string join(const std::vector & terms, const char * sep) { + std::string out; + for (size_t i = 0; i < terms.size(); ++i) { + if (i != 0) { + out += sep; + } + out += terms[i]; + } + return out; +} + +transcribe_status check_plain_text(const Tokenizer & tok, const std::string & text, const char * what) { + // Candidate literals are "<...>" and "[...]" spans up to a control + // token's plausible length. "<|...|>" pieces are rejected whenever the + // vocab has them (Whisper's rule: those are never plain text); other + // shapes only when the vocab types them CONTROL or they are BOS/EOS. + constexpr size_t k_max_literal = 48; + for (size_t i = 0; i < text.size(); ++i) { + const char open = text[i]; + if (open != '<' && open != '[') { + continue; + } + const char close = open == '<' ? '>' : ']'; + const size_t end = text.find(close, i + 1); + if (end == std::string::npos || end - i + 1 > k_max_literal) { + continue; + } + const std::string piece = text.substr(i, end - i + 1); + const int id = tok.find(piece); + if (id < 0) { + continue; + } + const bool pipe_form = piece.size() >= 4 && piece[1] == '|' && piece[piece.size() - 2] == '|'; + if (pipe_form || tok.is_control(id) || id == tok.bos_id() || id == tok.eos_id()) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, + "%s contains the control token \"%s\" (id %d); control tokens are " + "not accepted in prompting text", + what, piece.c_str(), id); + return TRANSCRIBE_ERR_INVALID_ARG; + } + } + return TRANSCRIBE_OK; +} + +transcribe_status encode_plain(const Tokenizer & tok, + const std::string & text, + std::vector & out_ids, + const char * what) { + if (const transcribe_status st = check_plain_text(tok, text, what); st != TRANSCRIBE_OK) { + return st; + } + out_ids.clear(); + if (text.empty()) { + return TRANSCRIBE_OK; + } + return tok.encode(text, out_ids); +} + +transcribe_status fit_terms_and_context(const Tokenizer & tok, + const std::vector & terms, + const TermsFormat & fmt, + const std::string & ctx, + int budget, + const char * family, + FittedPrompt & out) { + out = FittedPrompt{}; + const size_t cap = static_cast(std::max(budget, 0)); + // The first k terms as rendered text, and its ids. + auto render = [&](size_t k) { + std::string text; + if (k > 0) { + text = fmt.lead; + for (size_t i = 0; i < k; ++i) { + text += (i != 0 ? fmt.sep : std::string()) + terms[i]; + } + text += fmt.trail; + } + return text; + }; + auto encode_terms = [&](size_t k, std::vector & ids) -> transcribe_status { + ids.clear(); + return k == 0 ? TRANSCRIBE_OK : encode_plain(tok, render(k), ids, "vocabulary"); + }; + size_t kept = terms.size(); + if (const transcribe_status st = encode_terms(kept, out.term_ids); st != TRANSCRIBE_OK) { + return st; + } + if (const transcribe_status st = encode_plain(tok, ctx, out.ctx_ids, "prompt"); st != TRANSCRIBE_OK) { + return st; + } + const size_t ctx_in = out.ctx_ids.size(); + if (out.n_tokens() > cap) { + const size_t room = out.term_ids.size() < cap ? cap - out.term_ids.size() : 0; + out.ctx_ids.erase(out.ctx_ids.begin(), out.ctx_ids.end() - std::min(room, out.ctx_ids.size())); + if (out.term_ids.size() > cap) { + // The most terms that fit, by binary search over the count (the + // token count grows with it): 0 terms always fit, all do not. + size_t lo = 0, hi = kept; + std::vector ids; + while (hi - lo > 1) { + const size_t mid = lo + (hi - lo) / 2; + if (const transcribe_status st = encode_terms(mid, ids); st != TRANSCRIBE_OK) { + return st; + } + (ids.size() <= cap ? lo : hi) = mid; + } + kept = lo; + if (const transcribe_status st = encode_terms(kept, out.term_ids); st != TRANSCRIBE_OK) { + return st; + } + } + char terms_note[96] = ""; + if (kept < terms.size()) { + std::snprintf(terms_note, sizeof(terms_note), "dropped %zu of %zu vocabulary terms", terms.size() - kept, + terms.size()); + } + char ctx_note[96] = ""; + if (out.ctx_ids.size() < ctx_in) { + std::snprintf(ctx_note, sizeof(ctx_note), "kept the last %zu of %zu context tokens", out.ctx_ids.size(), + ctx_in); + } + log_msg(TRANSCRIBE_LOG_LEVEL_WARN, "%s: %s%s%s (prompt budget: %zu tokens)", family, terms_note, + (terms_note[0] != '\0' && ctx_note[0] != '\0') ? "; " : "", ctx_note, cap); + } + out.n_terms = kept; + out.terms_text = render(kept); + return TRANSCRIBE_OK; +} + +} // namespace transcribe::prompting diff --git a/src/transcribe-prompting.h b/src/transcribe-prompting.h new file mode 100644 index 000000000..5e8850c0e --- /dev/null +++ b/src/transcribe-prompting.h @@ -0,0 +1,83 @@ +// transcribe-prompting.h - shared helpers for the generic prompting fields +// (transcribe_run_params::vocabulary / prompt / prefix, TASK_INSTRUCT). +// +// INTERNAL. The dispatcher validates the fields and removes the ones a model +// ignores, so a family acts on whatever is set in the params view it +// receives. Families own the rendering and their budget rules. + +#pragma once + +#include "transcribe.h" + +#include +#include +#include + +namespace transcribe { + +class Tokenizer; + +namespace prompting { + +inline bool has_text(const char * s) { + return s != nullptr && s[0] != '\0'; +} + +inline bool is_instruct(const transcribe_run_params * p) { + return p != nullptr && p->task == TRANSCRIBE_TASK_INSTRUCT; +} + +// Non-empty vocabulary terms in caller order. Assumes the dispatcher's +// shape validation (non-negative count, no NULL entries) already passed. +std::vector terms(const transcribe_run_params * p); + +std::string join(const std::vector & terms, const char * sep); + +// `s` without leading / trailing whitespace (std::isspace). +std::string strip(const std::string & s); + +// Returns INVALID_ARG when `text` contains a literal of one of the +// tokenizer's control tokens (e.g. "<|im_end|>", "[INST]"): the upstream +// tokenizers would encode it as a control id, ours never do. `what` names +// the field in the error log. +transcribe_status check_plain_text(const Tokenizer & tok, const std::string & text, const char * what); + +// check_plain_text + encode. +transcribe_status encode_plain(const Tokenizer & tok, + const std::string & text, + std::vector & out_ids, + const char * what); + +// Tokenized vocabulary + context for one text slot, fitted to `budget` +// tokens (negative = 0). Terms render as `lead + join(terms, sep) + trail`; +// the context is encoded as given. Over budget, the context is trimmed +// first, keeping its most recent tokens, then terms are dropped from the end +// of the list, with one WARN. Both are control-token checked. +// +// A fit that comes in under a smaller budget is also the fit for that +// budget; batch paths rely on this to share one fit across rows. +struct TermsFormat { + std::string lead; + std::string sep; + std::string trail; +}; + +struct FittedPrompt { + std::vector term_ids; + std::vector ctx_ids; + size_t n_terms = 0; // terms kept + std::string terms_text; // the kept terms rendered (lead + join + trail); empty if none + + size_t n_tokens() const { return term_ids.size() + ctx_ids.size(); } +}; + +transcribe_status fit_terms_and_context(const Tokenizer & tok, + const std::vector & terms, + const TermsFormat & fmt, + const std::string & ctx, + int budget, + const char * family, + FittedPrompt & out); + +} // namespace prompting +} // namespace transcribe diff --git a/src/transcribe-session.h b/src/transcribe-session.h index 73cc69c79..e3984b7e3 100644 --- a/src/transcribe-session.h +++ b/src/transcribe-session.h @@ -287,8 +287,12 @@ struct transcribe_session { // storage (the public contract lets the caller free its params pointers // the moment begin returns). Stable for the stream's lifetime; only the // next begin mutates them. - std::string stream_language_owned; - std::string stream_target_language_owned; + std::string stream_language_owned; + std::string stream_target_language_owned; + // Generic prompting strings for the stream's run-params view. + std::vector stream_vocabulary_owned; + std::vector stream_vocabulary_ptrs; + std::string stream_prompt_owned; // UI-facing streaming text state. `full_text` above remains the raw // model hypothesis. `stream_committed_text` is the append-only public diff --git a/src/transcribe-tokenizer.cpp b/src/transcribe-tokenizer.cpp index b46c5647c..8d1fc4889 100644 --- a/src/transcribe-tokenizer.cpp +++ b/src/transcribe-tokenizer.cpp @@ -18,10 +18,13 @@ #include "transcribe-unicode.h" #include +#include #include #include +#include #include #include +#include #include namespace transcribe { @@ -466,6 +469,136 @@ transcribe_status encode_tiktoken_raw_bytes(const std::string & } // namespace +transcribe_status Tokenizer::encode_sentencepiece_bpe(const std::string & text, + std::vector & out_ids, + int lo, + int hi, + bool remove_extra_whitespaces) const { + out_ids.clear(); + const int n_vocab = static_cast(tokens_.size()); + if (model_ != "unigram" && model_ != "bpe") { + return TRANSCRIBE_ERR_NOT_IMPLEMENTED; + } + lo = std::max(lo, 0); + hi = (hi < 0 || hi > n_vocab) ? n_vocab : hi; + + // nmt_nfkc, approximated: the common NFKC compatibility mappings + // (ellipsis, no-break / ideographic space, full-width ASCII, f-ligatures) + // and whitespace handling -- tabs and newlines are spaces; optionally + // collapse runs and trim; then the dummy prefix, and every space becomes + // U+2581. Other NFKC mappings are not applied. + static const std::string k_space = "\xe2\x96\x81"; + std::string ws; + for (size_t i = 0; i < text.size();) { + const unsigned char c0 = static_cast(text[i]); + if (c0 < 0x80) { + ws += (c0 == '\t' || c0 == '\n' || c0 == '\r') ? ' ' : static_cast(c0); + ++i; + continue; + } + size_t len = 1; + while (i + len < text.size() && (static_cast(text[i + len]) & 0xC0) == 0x80) { + ++len; + } + uint32_t cp = len == 2 ? (c0 & 0x1Fu) : len == 3 ? (c0 & 0x0Fu) : (c0 & 0x07u); + for (size_t k = 1; k < len; ++k) { + cp = (cp << 6) | (static_cast(text[i + k]) & 0x3Fu); + } + if (cp == 0x2026) { + ws += "..."; + } else if (cp == 0x00A0 || cp == 0x3000) { + ws += ' '; + } else if (cp >= 0xFF01 && cp <= 0xFF5E) { + ws += static_cast(cp - 0xFF01 + 0x21); + } else if (cp >= 0xFB00 && cp <= 0xFB04) { + static const char * const k_lig[] = { "ff", "fi", "fl", "ffi", "ffl" }; + ws += k_lig[cp - 0xFB00]; + } else { + ws.append(text, i, len); + } + i += len; + } + if (remove_extra_whitespaces) { + std::string collapsed; + for (const char c : ws) { + if (c == ' ' && (collapsed.empty() || collapsed.back() == ' ')) { + continue; + } + collapsed += c; + } + while (!collapsed.empty() && collapsed.back() == ' ') { + collapsed.pop_back(); + } + ws = collapsed; + } + if (ws.empty()) { + return TRANSCRIBE_OK; + } + std::vector symbols{ k_space }; + for (size_t i = 0; i < ws.size();) { + size_t j = i + 1; + while (j < ws.size() && (static_cast(ws[j]) & 0xC0) == 0x80) { + ++j; + } + symbols.push_back(ws[i] == ' ' ? k_space : ws.substr(i, j - i)); + i = j; + } + + // Mergeable pieces in range; a lower id is a higher merge priority (a + // SentencePiece BPE score is minus the merge rank, and pieces are stored + // in rank order). Control / unknown pieces never match text. + constexpr int32_t k_type_unknown = 2; + constexpr int32_t k_type_control = 3; + std::unordered_map pieces; + int32_t unk = -1; + for (int id = lo; id < hi; ++id) { + const std::string & piece = tokens_[static_cast(id)]; + const int32_t type = token_type_.empty() ? 1 : token_type_[static_cast(id)]; + if (type == k_type_unknown || piece == "") { + unk = unk < 0 ? id : unk; + continue; + } + if (type != k_type_control && !piece.empty()) { + pieces.emplace(piece, id); + } + } + if (unk < 0) { + unk = unk_id_; + } + if (unk < 0) { + return TRANSCRIBE_ERR_GGUF; + } + + // Greedy merges: the highest-priority adjacent pair, leftmost on ties. + for (;;) { + size_t best_at = symbols.size(); + int32_t best_id = std::numeric_limits::max(); + for (size_t i = 0; i + 1 < symbols.size(); ++i) { + const auto it = pieces.find(symbols[i] + symbols[i + 1]); + if (it != pieces.end() && it->second < best_id) { + best_id = it->second; + best_at = i; + } + } + if (best_at == symbols.size()) { + break; + } + symbols[best_at] += symbols[best_at + 1]; + symbols.erase(symbols.begin() + static_cast(best_at) + 1); + } + + // As SentencePiece: a run of unknown symbols is one unknown piece. + for (const std::string & sym : symbols) { + const auto it = pieces.find(sym); + const int32_t id = it != pieces.end() ? it->second : unk; + if (id == unk && !out_ids.empty() && out_ids.back() == unk) { + continue; + } + out_ids.push_back(id); + } + return TRANSCRIBE_OK; +} + transcribe_status Tokenizer::encode(const std::string & text, std::vector & out_ids) const { out_ids.clear(); diff --git a/src/transcribe-tokenizer.h b/src/transcribe-tokenizer.h index fb35823da..0b5f35970 100644 --- a/src/transcribe-tokenizer.h +++ b/src/transcribe-tokenizer.h @@ -183,6 +183,18 @@ class Tokenizer { // merges (the encoder needs them). transcribe_status encode(const std::string & text, std::vector & out_ids) const; + // SentencePiece BPE encode over the pieces with ids in [lo, hi) (hi < 0 = + // whole vocab; canary passes one language's sub-vocab range). Piece + // order is the merge rank, so scores are not needed. Input normalization + // approximates nmt_nfkc, not full NFKC. A run of uncoverable characters + // becomes one unknown piece. Accepts GGUF model "unigram" or "bpe" (the + // canary converter labels these BPE models "unigram"). + transcribe_status encode_sentencepiece_bpe(const std::string & text, + std::vector & out_ids, + int lo, + int hi, + bool remove_extra_whitespaces) const; + // Identification + special token ids. -1 if the corresponding key // was absent from the GGUF. const std::string & model_type() const { return model_; } diff --git a/src/transcribe.cpp b/src/transcribe.cpp index 2acf9a606..852702d60 100644 --- a/src/transcribe.cpp +++ b/src/transcribe.cpp @@ -26,6 +26,7 @@ #include "transcribe-log.h" #include "transcribe-model.h" #include "transcribe-path.h" +#include "transcribe-prompting.h" #include "transcribe-session.h" #include "transcribe-tokenizer.h" #include "transcribe/whisper.h" @@ -300,6 +301,8 @@ int timestamp_rank(transcribe_timestamp_kind k) { // rejection differs between the two (run mirrors supports_translate, // streaming-begin rejects unconditionally in v1), so each caller // applies its own translate check before reaching this helper. +transcribe_status validate_prompting(const transcribe_model * model, const transcribe_run_params * params); + transcribe_status validate_run_params_common(const transcribe_session * session, const transcribe_run_params * params) { // Raw-validate every enum field before its first enum-typed load (see // enum_field_raw). Once a field passes here, downstream typed reads — @@ -307,6 +310,7 @@ transcribe_status validate_run_params_common(const transcribe_session * session, switch (enum_field_raw(¶ms->task)) { case TRANSCRIBE_TASK_TRANSCRIBE: case TRANSCRIBE_TASK_TRANSLATE: + case TRANSCRIBE_TASK_INSTRUCT: break; default: return TRANSCRIBE_ERR_INVALID_ARG; @@ -396,9 +400,130 @@ transcribe_status validate_run_params_common(const transcribe_session * session, !session->model->allows_translation_pair(params->language, params->target_language)) { return TRANSCRIBE_ERR_UNSUPPORTED_LANGUAGE; } + return validate_prompting(session->model, params); +} + +// Shape and hard-gate checks for the generic prompting fields, on a +// normalized view. Soft inputs a model ignores are removed later by +// prepare_prompting. +transcribe_status validate_prompting(const transcribe_model * model, const transcribe_run_params * params) { + auto reject = [](transcribe_status st, const char * why) { + transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "transcribe_run: %s", why); + return st; + }; + if (params->n_vocabulary < 0 || (params->n_vocabulary > 0 && params->vocabulary == nullptr)) { + return reject(TRANSCRIBE_ERR_INVALID_ARG, "vocabulary is NULL or n_vocabulary is negative"); + } + for (int32_t i = 0; i < params->n_vocabulary; ++i) { + if (params->vocabulary[i] == nullptr) { + return reject(TRANSCRIBE_ERR_INVALID_ARG, "vocabulary has a NULL entry"); + } + } + const bool has_prefix = transcribe::prompting::has_text(params->prefix); + if (params->task == TRANSCRIBE_TASK_INSTRUCT) { + if (!transcribe::has_feature(model, TRANSCRIBE_FEATURE_INSTRUCT)) { + return reject(TRANSCRIBE_ERR_UNSUPPORTED_TASK, + "this model does not support TRANSCRIBE_TASK_INSTRUCT (TRANSCRIBE_FEATURE_INSTRUCT)"); + } + // Output is free text: no target language and no alignment. A prefix + // as answer prefill is untested on every INSTRUCT family. + if (!transcribe::prompting::has_text(params->prompt)) { + return reject(TRANSCRIBE_ERR_INVALID_ARG, "TRANSCRIBE_TASK_INSTRUCT requires a non-empty prompt"); + } + if (params->target_language != nullptr) { + return reject(TRANSCRIBE_ERR_INVALID_ARG, "TRANSCRIBE_TASK_INSTRUCT does not take a target_language"); + } + if (params->timestamps != TRANSCRIBE_TIMESTAMPS_NONE && params->timestamps != TRANSCRIBE_TIMESTAMPS_AUTO) { + return reject(TRANSCRIBE_ERR_INVALID_ARG, "TRANSCRIBE_TASK_INSTRUCT supports timestamps NONE or AUTO only"); + } + if (has_prefix) { + return reject(TRANSCRIBE_ERR_INVALID_ARG, + "a transcript prefix is not supported with TRANSCRIBE_TASK_INSTRUCT"); + } + } + // Ignoring a prefix would make the output repeat the prefix's words and + // silently break callers that stitch text together, so it is a hard gate. + if (has_prefix && !transcribe::has_feature(model, TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX)) { + return reject(TRANSCRIBE_ERR_INVALID_ARG, + "this model does not support a transcript prefix (TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX)"); + } + return TRANSCRIBE_OK; +} + +// Rejects control-token literals in the prompting text before the result +// snapshot is cleared. Runs after strip_ignored_prompting, so ignored inputs +// are not rejected. +transcribe_status check_prompting_text(const transcribe_model * model, const transcribe_run_params * params) { + const transcribe::Tokenizer * tok = model->tokenizer(); + if (tok == nullptr) { + return TRANSCRIBE_OK; + } + for (int32_t i = 0; i < params->n_vocabulary; ++i) { + if (const transcribe_status st = + transcribe::prompting::check_plain_text(*tok, params->vocabulary[i], "vocabulary"); + st != TRANSCRIBE_OK) { + return st; + } + } + if (params->prompt != nullptr) { + if (const transcribe_status st = transcribe::prompting::check_plain_text(*tok, params->prompt, "prompt"); + st != TRANSCRIBE_OK) { + return st; + } + } + if (params->prefix != nullptr) { + return transcribe::prompting::check_plain_text(*tok, params->prefix, "prefix"); + } return TRANSCRIBE_OK; } +// Full-size copy of a caller's run params: defaults first, then only the +// prefix the caller's struct_size covers, so every trailing field is +// readable (NULL/0 for an older caller). struct_size is preserved so +// has_field() gating still sees the caller's true layout. Idempotent. +void normalize_run_params(const transcribe_run_params * in, transcribe_run_params * out) { + transcribe_run_params_init(out); + std::memcpy(out, in, static_cast(std::min(in->struct_size, sizeof(*out)))); +} + +// Warn about, then remove, the soft prompting inputs this model ignores, so +// a family only ever sees inputs it should act on. Idempotent: the batch +// serial fallback re-enters run_one_inner per utterance. +void strip_ignored_prompting(const transcribe_model * model, transcribe_run_params * params) { + const char * arch_name = (model->arch != nullptr && model->arch->name != nullptr) ? model->arch->name : "(unknown)"; + const bool instruct = params->task == TRANSCRIBE_TASK_INSTRUCT; + if (params->n_vocabulary > 0 && !transcribe::has_feature(model, TRANSCRIBE_FEATURE_VOCABULARY)) { + transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_WARN, + "transcribe_run: model '%s' does not support vocabulary; ignoring %d term(s). Use " + "transcribe_model_supports(model, TRANSCRIBE_FEATURE_VOCABULARY) to pre-check.", + arch_name, params->n_vocabulary); + params->vocabulary = nullptr; + params->n_vocabulary = 0; + } + if (!instruct && transcribe::prompting::has_text(params->prompt) && + !transcribe::has_feature(model, TRANSCRIBE_FEATURE_CONTEXT_PROMPT)) { + transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_WARN, + "transcribe_run: model '%s' has no context-prompt slot; ignoring prompt. Use " + "transcribe_model_supports(model, TRANSCRIBE_FEATURE_CONTEXT_PROMPT) to pre-check, or " + "TRANSCRIBE_TASK_INSTRUCT on models with TRANSCRIBE_FEATURE_INSTRUCT.", + arch_name); + params->prompt = nullptr; + } + if (!transcribe::prompting::has_text(params->prompt)) { + params->prompt = nullptr; + } + if (!transcribe::prompting::has_text(params->prefix)) { + params->prefix = nullptr; + } +} + +// The pre-clear prompting step every entry point shares, on a validated +// normalized view. +transcribe_status prepare_prompting(const transcribe_model * model, transcribe_run_params * params) { + strip_ignored_prompting(model, params); + return check_prompting_text(model, params); +} + } // namespace // Logging @@ -1738,6 +1863,11 @@ static transcribe_status transcribe_stream_begin_impl(struct transcribe_session if (const auto st = check_input_struct_size(run_params->struct_size, k_min_run_params_size); st != TRANSCRIBE_OK) { return st; } + // Full-size view (see normalize_run_params); the family hook gets a + // further copy whose strings the library owns, built below. + struct transcribe_run_params run_params_view; + normalize_run_params(run_params, &run_params_view); + run_params = &run_params_view; if (const auto st = check_input_struct_size(stream_params->struct_size, k_min_stream_params_size); st != TRANSCRIBE_OK) { return st; @@ -1791,9 +1921,17 @@ static transcribe_status transcribe_stream_begin_impl(struct transcribe_session // family that supports streaming translate would loosen this in // its stream_begin hook, but the central dispatcher refuses // upfront so partially-wired callers fail fast. - if (run_params->task == TRANSCRIBE_TASK_TRANSLATE) { + if (run_params->task == TRANSCRIBE_TASK_TRANSLATE || run_params->task == TRANSCRIBE_TASK_INSTRUCT) { return TRANSCRIBE_ERR_UNSUPPORTED_TASK; } + // A prefix is forced decoder text for one utterance's opening; a stream + // has no fixed opening to force it onto. + if (transcribe::prompting::has_text(run_params->prefix)) { + transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, + "transcribe_stream_begin: a transcript prefix is not accepted " + "for streaming"); + return TRANSCRIBE_ERR_INVALID_ARG; + } if (stream_params->family != nullptr) { if (stream_params->family->size < sizeof(struct transcribe_ext)) { @@ -1810,6 +1948,9 @@ static transcribe_status transcribe_stream_begin_impl(struct transcribe_session // clear_result so the pre-hook "snapshot preserved on rejection" // contract is undisturbed. warn_unsupported_advisory(session->model, run_params); + if (const transcribe_status st = prepare_prompting(session->model, &run_params_view); st != TRANSCRIBE_OK) { + return st; + } // Optional family preflight: validates extension field values // (e.g. parakeet's (L, C, R) menu) without mutating state. On @@ -1845,21 +1986,27 @@ static transcribe_status transcribe_stream_begin_impl(struct transcribe_session // wants a run-slot ext at stream begin must plumb it deliberately. session->stream_language_owned = run_params->language != nullptr ? run_params->language : ""; session->stream_target_language_owned = run_params->target_language != nullptr ? run_params->target_language : ""; - // PREFIX copy, not struct assignment: the size gate above admits any - // struct_size >= k_min_run_params_size, so a conforming caller's - // allocation may be SHORTER than sizeof (fields past `family`, e.g. - // spec_k_drafts, absent). Init first so bytes past the caller's - // prefix hold their documented defaults, then copy only what the - // caller owns. The caller's struct_size is preserved by the copy, so - // downstream has_field() gating still sees the caller's true layout. - struct transcribe_run_params run_params_owned; - transcribe_run_params_init(&run_params_owned); - std::memcpy(&run_params_owned, run_params, - static_cast(std::min(run_params->struct_size, sizeof(run_params_owned)))); + // Copied from the normalized view (never the caller's struct, whose + // allocation may end before the trailing fields); the view keeps the + // caller's struct_size, so has_field() gating still sees its layout. + struct transcribe_run_params run_params_owned = run_params_view; run_params_owned.language = run_params->language != nullptr ? session->stream_language_owned.c_str() : nullptr; run_params_owned.target_language = run_params->target_language != nullptr ? session->stream_target_language_owned.c_str() : nullptr; - run_params_owned.family = nullptr; + run_params_owned.family = nullptr; + // Generic prompting strings, same ownership rule. Only non-empty terms + // are kept, so the view's count matches the owned array. + session->stream_vocabulary_owned = transcribe::prompting::terms(run_params); + session->stream_vocabulary_ptrs.clear(); + for (const std::string & term : session->stream_vocabulary_owned) { + session->stream_vocabulary_ptrs.push_back(term.c_str()); + } + session->stream_prompt_owned = run_params->prompt != nullptr ? run_params->prompt : ""; + run_params_owned.vocabulary = + session->stream_vocabulary_ptrs.empty() ? nullptr : session->stream_vocabulary_ptrs.data(); + run_params_owned.n_vocabulary = static_cast(session->stream_vocabulary_ptrs.size()); + run_params_owned.prompt = run_params->prompt != nullptr ? session->stream_prompt_owned.c_str() : nullptr; + run_params_owned.prefix = nullptr; const transcribe_status st = session->model->arch->stream_begin(session, &run_params_owned, stream_params); if (st != TRANSCRIBE_OK) { @@ -2130,6 +2277,12 @@ static transcribe_status run_one_inner(struct transcribe_session * sess if (const auto st = check_input_struct_size(params->struct_size, k_min_run_params_size); st != TRANSCRIBE_OK) { return st; } + // Everything downstream reads this full-size view, never the caller's + // struct: an older caller's allocation may end before the prompting + // fields. Strings stay caller-owned; the call is synchronous. + struct transcribe_run_params params_view; + normalize_run_params(params, ¶ms_view); + params = ¶ms_view; // A run cannot replace an active stream's results — that would // strand the in-flight stream's per-family state. Caller must // finalize or reset first. FINISHED and FAILED both fall through; @@ -2179,6 +2332,9 @@ static transcribe_status run_one_inner(struct transcribe_session * sess if (params->task == TRANSCRIBE_TASK_TRANSLATE && !session->model->caps.supports_translate) { return TRANSCRIBE_ERR_UNSUPPORTED_TASK; } + if (const transcribe_status st = prepare_prompting(session->model, ¶ms_view); st != TRANSCRIBE_OK) { + return st; + } // Family run-ext validation (the _RUN analogue of stream_validate), // the final pre-clear gate. Runs AFTER the run-param checks above, @@ -2320,9 +2476,20 @@ static transcribe_status transcribe_run_batch_impl(struct transcribe_session * if (const auto st = check_input_struct_size(params->struct_size, k_min_run_params_size); st != TRANSCRIBE_OK) { return st; } + struct transcribe_run_params params_view; + normalize_run_params(params, ¶ms_view); + params = ¶ms_view; if (session->stream_state == TRANSCRIBE_STREAM_ACTIVE) { return TRANSCRIBE_ERR_INVALID_ARG; } + // One shared params across different audio: a transcript prefix can + // only describe one of them. + if (transcribe::prompting::has_text(params->prefix)) { + transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, + "transcribe_run_batch: a transcript prefix is per-utterance " + "and is not accepted in a batch"); + return TRANSCRIBE_ERR_INVALID_ARG; + } // Shared-param validation, ONCE, mirroring transcribe_run's pre-clear // gates (ext shape/kind, pnc/itn advisory, enum range, timestamp @@ -2344,6 +2511,9 @@ static transcribe_status transcribe_run_batch_impl(struct transcribe_session * if (params->task == TRANSCRIBE_TASK_TRANSLATE && !session->model->caps.supports_translate) { return TRANSCRIBE_ERR_UNSUPPORTED_TASK; } + if (const transcribe_status st = prepare_prompting(session->model, ¶ms_view); st != TRANSCRIBE_OK) { + return st; + } if (session->model->arch != nullptr && session->model->arch->run_validate != nullptr) { if (const transcribe_status st = session->model->arch->run_validate(session, params); st != TRANSCRIBE_OK) { return st; diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index db3d64be3..006448b37 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -1245,6 +1245,23 @@ if(TRANSCRIBE_BUILD_REAL_MODEL_TESTS) set_tests_properties(transcribe_qwen3_asr_bpe_parity PROPERTIES SKIP_RETURN_CODE 77) + # SentencePiece BPE encoder parity (canary2 transcript prefix), + # gated on TRANSCRIBE_CANARY_1B_V2_GGUF / TRANSCRIBE_CANARY_FLASH_GGUF. + # Expected ids come from the reference sentencepiece library via + # scripts/gen_sentencepiece_bpe_fixture.py. + add_executable(transcribe_sentencepiece_bpe_parity + sentencepiece_bpe_parity.cpp) + target_link_libraries(transcribe_sentencepiece_bpe_parity PRIVATE transcribe ggml) + target_include_directories(transcribe_sentencepiece_bpe_parity PRIVATE + ${CMAKE_SOURCE_DIR}/src) + target_compile_definitions(transcribe_sentencepiece_bpe_parity PRIVATE + "TRANSCRIBE_TEST_FIXTURES_DIR=\"${CMAKE_CURRENT_SOURCE_DIR}/fixtures\"") + transcribe_apply_warnings(transcribe_sentencepiece_bpe_parity) + add_test(NAME transcribe_sentencepiece_bpe_parity + COMMAND transcribe_sentencepiece_bpe_parity) + set_tests_properties(transcribe_sentencepiece_bpe_parity PROPERTIES + SKIP_RETURN_CODE 77) + add_executable(transcribe_qwen3_asr_e2e_smoke qwen3_asr_e2e_smoke.cpp) diff --git a/tests/fixtures/sentencepiece_bpe/canary-1b-v2.jsonl b/tests/fixtures/sentencepiece_bpe/canary-1b-v2.jsonl new file mode 100644 index 000000000..39eb3a2f2 --- /dev/null +++ b/tests/fixtures/sentencepiece_bpe/canary-1b-v2.jsonl @@ -0,0 +1,23 @@ +{"text": "Our annual contract value increased by eight percentage points", "ids": [1342, 1222, 1274, 2179, 1185, 7003, 1937, 1826, 2243, 10360, 8491, 1517, 11002, 13810, 2763, 1448, 1244, 11792]} +{"text": "And so, my fellow Americans, ask not what your country can do for you", "ids": [2572, 1402, 16067, 1904, 15893, 1458, 4092, 1252, 1444, 16067, 1473, 16066, 2019, 3360, 3679, 11618, 2585, 1343, 1378, 1789]} +{"text": "The quick brown fox jumps over the lazy dog.", "ids": [1839, 1259, 2172, 4033, 2210, 2143, 16138, 10077, 2392, 2255, 1289, 1273, 1470, 5391, 16073]} +{"text": "We expect revenue of $4.2 billion in Q3 2021, up 17% year-over-year.", "ids": [2505, 1593, 6409, 12383, 2243, 1381, 16053, 16372, 1127, 16073, 1125, 10082, 1177, 1220, 2091, 1126, 16053, 1125, 1123, 1125, 1124, 16067, 2567, 16053, 1124, 1130, 16320, 7727, 16107, 3902, 16107, 6802, 1176, 16073]} +{"text": "I'm sure it's fine; they'd've said so if it weren't.", "ids": [1262, 16122, 16065, 11471, 1612, 16122, 16060, 10041, 0, 2761, 16122, 16063, 16122, 1395, 4412, 1402, 3499, 1612, 1213, 1985, 16122, 16058, 16073]} +{"text": " leading and trailing spaces ", "ids": [16053, 16053, 14425, 1260, 1392, 1880, 1194, 1260, 1370, 1418, 1172, 16053, 16053, 16053]} +{"text": "double spaces inside this line", "ids": [3980, 2655, 16053, 1370, 1418, 1172, 16053, 16053, 3477, 1721, 16053, 16053, 16053, 2259, 10459]} +{"text": "COVID-19 hit EMEA and APAC; EBITDA margins fell to 12.5%.", "ids": [11267, 16160, 11876, 16107, 1124, 1132, 6740, 1254, 16125, 16127, 16112, 1392, 1235, 16136, 16112, 16148, 0, 1254, 16152, 10095, 12969, 1989, 16072, 2248, 15893, 1237, 16053, 1124, 1125, 16073, 1128, 16320, 16073]} +{"text": "Hello, world! How are you? Fine... thanks.", "ids": [1305, 3832, 16067, 9151, 16210, 8668, 2139, 1789, 16145, 1403, 1412, 16073, 16073, 16073, 5687, 1452, 16073]} +{"text": "naïve café résumé coöperate", "ids": [1248, 16316, 1395, 1186, 2289, 16096, 6203, 1276, 16096, 1690, 16131, 1728, 1459]} +{"text": "Straße und Fußgänger", "ids": [12790, 16244, 16054, 1626, 13033, 16244, 16072, 1382, 2101]} +{"text": "Temperatures of 30°C — roughly 86°F — are common.", "ids": [1240, 1200, 1728, 1184, 4565, 1381, 16053, 1126, 1123, 0, 16148, 16053, 0, 1483, 5232, 1617, 16053, 1131, 1129, 0, 16183, 16053, 0, 2139, 1375, 3540, 16073]} +{"text": "An em-dash—without spaces—and an en-dash – with them.", "ids": [2199, 1633, 16107, 5892, 16071, 0, 16085, 1189, 7409, 1370, 1418, 1172, 0, 1310, 1274, 1258, 16107, 5892, 16071, 16053, 0, 2151, 3846, 16073]} +{"text": "Numbers: 0 1 2 3 4 5 6 7 8 9 10 100 1,000 1.5 3/4", "ids": [1277, 8489, 1385, 16223, 16053, 1123, 16053, 1124, 16053, 1125, 16053, 1126, 16053, 1127, 16053, 1128, 16053, 1129, 16053, 1130, 16053, 1131, 16053, 1132, 16053, 1124, 1123, 16053, 1124, 1123, 1123, 16053, 1124, 16067, 1123, 1123, 1123, 16053, 1124, 16073, 1128, 16053, 1126, 16336, 1127]} +{"text": "e-mail: someone@example.com, URL https://example.com/path?q=1", "ids": [1216, 16107, 1371, 1194, 16223, 4194, 1911, 0, 6308, 1205, 3775, 16073, 5231, 16067, 1413, 16179, 16157, 1191, 2181, 2392, 16223, 16336, 16336, 6308, 1205, 3775, 16073, 5231, 16336, 4408, 16071, 16145, 16109, 0, 1124]} +{"text": "\"Quoted text\" and 'single quotes' and (parentheses) [brackets] {braces}", "ids": [16053, 0, 16237, 6289, 2960, 7965, 0, 1392, 8348, 16060, 1260, 1226, 1259, 1265, 1172, 16122, 1392, 16053, 0, 2475, 1227, 13482, 1172, 0, 16053, 0, 3697, 1600, 3305, 0, 16053, 0, 3697, 1970, 0]} +{"text": "Mixed CASE WoRdS and CamelCaseIdentifiers like TeamViewer or iPhone", "ids": [1245, 3781, 1212, 1303, 9375, 16127, 6983, 16179, 16063, 16113, 1392, 8430, 1181, 16148, 2330, 16132, 1598, 1210, 3072, 1385, 3024, 2680, 1205, 16160, 1195, 3347, 1684, 14441, 16000]} +{"text": "A", "ids": [1235]} +{"text": "a b c d e f g", "ids": [1168, 1192, 1186, 1165, 1216, 1197, 1206]} +{"text": "supercalifragilisticexpialidocious antidisestablishmentarianism", "ids": [3922, 16070, 1457, 16083, 4926, 1194, 2393, 1339, 16138, 2197, 1185, 2408, 1296, 3636, 1274, 3854, 1178, 2272, 1722, 3122, 2023, 2203, 1173, 5406]} +{"text": "日本語のテキスト", "ids": [16053, 0]} +{"text": "tab\tseparated\twords", "ids": [6687, 10068, 4162, 3027, 16060]} +{"text": "Full-width ABC and “curly” quotes… with fine ligatures and a no-break space", "ids": [1403, 2795, 16107, 16085, 1218, 4165, 1235, 16152, 16148, 1392, 16053, 0, 6228, 1617, 0, 1259, 1265, 1172, 16073, 16073, 16073, 2151, 10041, 4484, 1184, 4565, 1392, 1168, 1391, 16107, 2102, 1229, 1370, 4744]} diff --git a/tests/fixtures/sentencepiece_bpe/canary-flash-en.jsonl b/tests/fixtures/sentencepiece_bpe/canary-flash-en.jsonl new file mode 100644 index 000000000..1cd32c44d --- /dev/null +++ b/tests/fixtures/sentencepiece_bpe/canary-flash-en.jsonl @@ -0,0 +1,23 @@ +{"text": "Our annual contract value increased by eight percentage points", "ids": [1452, 1206, 1211, 2025, 1990, 1236, 1626, 1323, 1188, 1192, 1319, 1201, 2035, 1167, 1178, 1246, 1691, 1175, 1553, 1403, 2035, 1182, 1642, 1415, 1159, 1490]} +{"text": "And so, my fellow Americans, ask not what your country can do for you", "ids": [1988, 1268, 2043, 1687, 1784, 1214, 1335, 1251, 1523, 1234, 1309, 2043, 1330, 2049, 1533, 1326, 1185, 1205, 1237, 1493, 1168, 1298, 2045, 1527, 1421, 1340, 1324]} +{"text": "The quick brown fox jumps over the lazy dog.", "ids": [1753, 1187, 1726, 1191, 1209, 2044, 2025, 1638, 2054, 1222, 1261, 1768, 1200, 1369, 1203, 1204, 2050, 2045, 1421, 2039, 2048]} +{"text": "We expect revenue of $4.2 billion in Q3 2021, up 17% year-over-year.", "ids": [1719, 1362, 1390, 1323, 1231, 1582, 1319, 1253, 2023, 2148, 2128, 2048, 2113, 1191, 1379, 1226, 1201, 2023, 2095, 2126, 2023, 2113, 2063, 2113, 2108, 2043, 1752, 2023, 2108, 2130, 2135, 1644, 1174, 2071, 2031, 1369, 2071, 2045, 2024, 1174, 2048]} +{"text": "I'm sure it's fine; they'd've said so if it weren't.", "ids": [1258, 2051, 2037, 1281, 1167, 1320, 2051, 2027, 1190, 1360, 2115, 1203, 2045, 2051, 2034, 2051, 1220, 1317, 1244, 1268, 1704, 1320, 1537, 1153, 2051, 2029, 2048]} +{"text": " leading and trailing spaces ", "ids": [1240, 1212, 1219, 1241, 1440, 1228, 1219, 1610, 2026, 1443]} +{"text": "double spaces inside this line", "ids": [1154, 1173, 1617, 1610, 2026, 1443, 1949, 1596, 1183, 1170, 1163, 1360]} +{"text": "COVID-19 hit EMEA and APAC; EBITDA margins fell to 12.5%.", "ids": [1315, 2087, 2084, 2060, 2067, 2071, 2108, 2120, 1186, 1177, 1254, 2061, 2057, 2058, 1241, 1251, 2069, 2058, 2072, 2115, 1254, 2070, 2060, 2068, 2067, 2058, 1571, 2039, 1530, 1784, 1214, 1221, 2023, 2108, 2113, 2048, 2125, 2135, 2048]} +{"text": "Hello, world! How are you? Fine... thanks.", "ids": [2014, 1214, 2031, 2043, 1724, 1457, 2117, 1345, 1335, 1480, 1324, 2086, 1332, 1360, 2048, 2048, 2048, 1183, 1792, 2027, 2048]} +{"text": "naïve café résumé coöperate", "ids": [1181, 2026, 2111, 1220, 1164, 2026, 2041, 2047, 1243, 1380, 1261, 2047, 1493, 2078, 1478, 1468]} +{"text": "Straße und Fußgänger", "ids": [1215, 1626, 2088, 2024, 1248, 1332, 2032, 2088, 2039, 1629, 1842]} +{"text": "Temperatures of 30°C — roughly 86°F — are common.", "ids": [1294, 1232, 1478, 1185, 2006, 1253, 2023, 2126, 2063, 2154, 2072, 2023, 2136, 1243, 1891, 1348, 2023, 2129, 2132, 2154, 2076, 2023, 2136, 1480, 1238, 2037, 1161, 2048]} +{"text": "An em-dash—without spaces—and an en-dash – with them.", "ids": [1529, 1175, 2037, 2071, 2034, 1178, 2036, 2136, 2044, 1437, 1367, 1610, 2026, 1443, 2136, 1271, 1211, 1210, 2071, 2034, 1178, 2036, 2023, 2140, 1172, 1437, 1203, 2037, 2048]} +{"text": "Numbers: 0 1 2 3 4 5 6 7 8 9 10 100 1,000 1.5 3/4", "ids": [1361, 1261, 2040, 1257, 2107, 2023, 2063, 2023, 2108, 2023, 2113, 2023, 2126, 2023, 2128, 2023, 2125, 2023, 2132, 2023, 2130, 2023, 2129, 2023, 2120, 2023, 2108, 2063, 2023, 2108, 1278, 2023, 2108, 2043, 1278, 2063, 2023, 2108, 2048, 2125, 2023, 2126, 2134, 2128]} +{"text": "e-mail: someone@example.com, URL https://example.com/path?q=1", "ids": [1175, 2071, 1484, 1228, 2107, 1536, 2024, 1811, 2167, 2024, 2054, 1223, 1495, 2048, 2035, 1197, 2043, 1414, 2080, 2065, 1186, 2029, 2029, 1768, 2107, 2134, 2134, 2024, 2054, 1223, 1495, 2048, 2035, 1197, 2134, 2038, 1185, 2036, 2086, 2046, 1152, 2108]} +{"text": "\"Quoted text\" and 'single quotes' and (parentheses) [brackets] {braces}", "ids": [2023, 2102, 1819, 2031, 1526, 1540, 2054, 2029, 2102, 1241, 1649, 2027, 1219, 1195, 1187, 2031, 1598, 2051, 1241, 2023, 2122, 1640, 1182, 2036, 1156, 1156, 2121, 2023, 2155, 2040, 1218, 1388, 1229, 2027, 2157, 2023, 1152, 2040, 1218, 1443, 1152]} +{"text": "Mixed CASE WoRdS and CamelCaseIdentifiers like TeamViewer or iPhone", "ids": [1267, 1946, 1246, 1315, 2058, 2053, 2057, 1308, 2031, 2080, 2034, 2053, 1241, 1315, 1223, 1235, 2072, 1178, 2024, 2060, 2034, 1182, 1420, 2028, 1257, 1521, 1331, 1928, 1223, 2084, 1180, 1909, 1449, 1199, 2069, 2036, 1811]} +{"text": "A", "ids": [1251]} +{"text": "a b c d e f g", "ids": [1157, 1191, 1164, 1154, 1175, 1190, 1198]} +{"text": "supercalifragilisticexpialidocious antidisestablishmentarianism", "ids": [1281, 1478, 2035, 1192, 1420, 1218, 2039, 1228, 1170, 1208, 1273, 2054, 2038, 1851, 1581, 1207, 1230, 1211, 1208, 2034, 1170, 1280, 1227, 2033, 1170, 2036, 1290, 1174, 2028, 1166, 1170, 2037]} +{"text": "日本語のテキスト", "ids": [2023, 1152]} +{"text": "tab\tseparated\twords", "ids": [1160, 1227, 1217, 1640, 1185, 1246, 1172, 1568, 2027]} +{"text": "Full-width ABC and “curly” quotes… with fine ligatures and a no-break space", "ids": [1332, 2032, 1214, 2071, 2044, 1244, 1410, 1251, 2070, 2072, 1241, 2023, 2131, 2035, 1206, 1348, 1152, 1187, 2031, 1598, 2048, 2048, 2048, 1172, 1437, 1190, 1360, 1163, 1225, 1185, 2006, 1241, 1157, 1260, 2071, 1496, 1746, 1610, 1986]} diff --git a/tests/run_dispatch_unit.cpp b/tests/run_dispatch_unit.cpp index e3bf21508..4ee967f1c 100644 --- a/tests/run_dispatch_unit.cpp +++ b/tests/run_dispatch_unit.cpp @@ -6,6 +6,7 @@ #include "transcribe-session.h" #include "transcribe.h" +#include #include #include #include @@ -616,6 +617,178 @@ void test_batch_serial_truncation_is_per_utterance() { check_truncated_then_clean(dispatcher_arch); } +// --------------------------------------------------------------------------- +// Generic prompting fields: validation, warn-and-strip, and the normalized +// full-size view families receive. +// --------------------------------------------------------------------------- + +transcribe_run_params g_seen_params; +int g_prompt_runs = 0; + +transcribe_status capture_run(transcribe_session * session, + const float * pcm, + int n_samples, + const transcribe_run_params * params) { + (void) pcm; + (void) n_samples; + g_seen_params = *params; + ++g_prompt_runs; + session->full_text = "fresh result"; + session->has_result = true; + return TRANSCRIBE_OK; +} + +const transcribe::Arch & capture_arch() { + static const transcribe::Arch arch = { + "fake-prompt", nullptr, nullptr, capture_run, nullptr, nullptr, + nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, + }; + return arch; +} + +transcribe_status prompt_run(transcribe_model & model, const transcribe_run_params & params) { + transcribe_session session; + session.model = &model; + session.full_text = "previous result"; + session.has_result = true; + float pcm = 0.0f; + const transcribe_status st = transcribe_run(&session, &pcm, 1, ¶ms); + if (st != TRANSCRIBE_OK) { + // Every dispatcher-level prompting rejection is pre-clear (family + // text checks go through check_prompting_text / run_validate). + CHECK(session.has_result); + CHECK(session.full_text == "previous result"); + } + return st; +} + +void test_prompting_validation() { + transcribe_model model; + model.arch = &capture_arch(); + const char * terms[] = { "GGUF", "", "ggml" }; + + transcribe_run_params p; + transcribe_run_params_init(&p); + CHECK(p.vocabulary == nullptr && p.n_vocabulary == 0 && p.prompt == nullptr && p.prefix == nullptr); + + // Vocabulary shape. + p.n_vocabulary = -1; + CHECK(prompt_run(model, p) == TRANSCRIBE_ERR_INVALID_ARG); + p.n_vocabulary = 2; + CHECK(prompt_run(model, p) == TRANSCRIBE_ERR_INVALID_ARG); + const char * with_null[] = { "a", nullptr }; + p.vocabulary = with_null; + CHECK(prompt_run(model, p) == TRANSCRIBE_ERR_INVALID_ARG); + + // INSTRUCT without the bit, then its argument rules with it. + transcribe_run_params_init(&p); + p.task = TRANSCRIBE_TASK_INSTRUCT; + p.prompt = "Summarize."; + CHECK(prompt_run(model, p) == TRANSCRIBE_ERR_UNSUPPORTED_TASK); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_INSTRUCT, true); + CHECK(prompt_run(model, p) == TRANSCRIBE_OK); + CHECK(g_seen_params.task == TRANSCRIBE_TASK_INSTRUCT); + CHECK(std::strcmp(g_seen_params.prompt, "Summarize.") == 0); + p.prompt = ""; + CHECK(prompt_run(model, p) == TRANSCRIBE_ERR_INVALID_ARG); + p.prompt = nullptr; + CHECK(prompt_run(model, p) == TRANSCRIBE_ERR_INVALID_ARG); + p.prompt = "Summarize."; + p.target_language = "fr"; + CHECK(prompt_run(model, p) == TRANSCRIBE_ERR_INVALID_ARG); + p.target_language = nullptr; + p.timestamps = TRANSCRIBE_TIMESTAMPS_NONE; + CHECK(prompt_run(model, p) == TRANSCRIBE_OK); + + // Prefix: hard gate, and never under INSTRUCT. + p.prefix = "Good morning"; + CHECK(prompt_run(model, p) == TRANSCRIBE_ERR_INVALID_ARG); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX, true); + CHECK(prompt_run(model, p) == TRANSCRIBE_ERR_INVALID_ARG); + p.task = TRANSCRIBE_TASK_TRANSCRIBE; + CHECK(prompt_run(model, p) == TRANSCRIBE_OK); + CHECK(std::strcmp(g_seen_params.prefix, "Good morning") == 0); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX, false); + p.prefix = ""; // empty == absent + CHECK(prompt_run(model, p) == TRANSCRIBE_OK); + CHECK(g_seen_params.prefix == nullptr); + + // Soft inputs without their bits: warn, run, and the family sees none. + transcribe_run_params_init(&p); + p.vocabulary = terms; + p.n_vocabulary = 3; + p.prompt = "Earnings call."; + CHECK(prompt_run(model, p) == TRANSCRIBE_OK); + CHECK(g_seen_params.vocabulary == nullptr && g_seen_params.n_vocabulary == 0); + CHECK(g_seen_params.prompt == nullptr); + + // With the bits, they pass through untouched. + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_VOCABULARY, true); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CONTEXT_PROMPT, true); + CHECK(prompt_run(model, p) == TRANSCRIBE_OK); + CHECK(g_seen_params.vocabulary == terms && g_seen_params.n_vocabulary == 3); + CHECK(std::strcmp(g_seen_params.prompt, "Earnings call.") == 0); + + // Vocabulary under INSTRUCT needs both V and I. + p.task = TRANSCRIBE_TASK_INSTRUCT; + p.prompt = "Summarize."; + CHECK(prompt_run(model, p) == TRANSCRIBE_OK); + CHECK(g_seen_params.n_vocabulary == 3); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_VOCABULARY, false); + CHECK(prompt_run(model, p) == TRANSCRIBE_OK); + CHECK(g_seen_params.n_vocabulary == 0); +} + +// A caller compiled before the prompting fields existed passes a struct that +// ends at spec_k_drafts. Whatever lies past it must be read as defaults. +void test_prompting_short_struct_reads_defaults() { + transcribe_model model; + model.arch = &capture_arch(); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_VOCABULARY, true); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX, true); + + transcribe_run_params p; + std::memset(&p, 0xA5, sizeof(p)); + transcribe_run_params base; + transcribe_run_params_init(&base); + const size_t old_size = offsetof(transcribe_run_params, spec_k_drafts) + sizeof(base.spec_k_drafts); + std::memcpy(&p, &base, old_size); + p.struct_size = old_size; + + CHECK(prompt_run(model, p) == TRANSCRIBE_OK); + CHECK(g_seen_params.vocabulary == nullptr && g_seen_params.n_vocabulary == 0); + CHECK(g_seen_params.prompt == nullptr && g_seen_params.prefix == nullptr); + CHECK(g_seen_params.struct_size == old_size); +} + +void test_prompting_batch_rejects_prefix() { + transcribe_model model; + model.arch = &capture_arch(); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX, true); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_VOCABULARY, true); + + transcribe_session session; + session.model = &model; + + transcribe_run_params p; + transcribe_run_params_init(&p); + p.prefix = "Good morning"; + float a = 0.0f; + const float * pcm[2] = { &a, &a }; + const int ns[2] = { 1, 1 }; + CHECK(transcribe_run_batch(&session, pcm, ns, 2, &p) == TRANSCRIBE_ERR_INVALID_ARG); + + // Vocabulary is fine in a batch and reaches every utterance. + const char * terms[] = { "GGUF" }; + p.prefix = nullptr; + p.vocabulary = terms; + p.n_vocabulary = 1; + g_prompt_runs = 0; + CHECK(transcribe_run_batch(&session, pcm, ns, 2, &p) == TRANSCRIBE_OK); + CHECK(g_prompt_runs == 2); + CHECK(g_seen_params.n_vocabulary == 1); +} + } // namespace int main() { @@ -628,5 +801,8 @@ int main() { test_batch_abort_pads_missing_to_n(); test_batch_fastpath_abort_pads_missing_to_n(); test_raw_text_single_batch_and_alias(); + test_prompting_validation(); + test_prompting_short_struct_reads_defaults(); + test_prompting_batch_rejects_prefix(); return g_failures == 0 ? EXIT_SUCCESS : EXIT_FAILURE; } diff --git a/tests/sentencepiece_bpe_parity.cpp b/tests/sentencepiece_bpe_parity.cpp new file mode 100644 index 000000000..13cfcc1c4 --- /dev/null +++ b/tests/sentencepiece_bpe_parity.cpp @@ -0,0 +1,165 @@ +// sentencepiece_bpe_parity.cpp - Tokenizer::encode_sentencepiece_bpe +// against ids from the reference `sentencepiece` library. +// +// Fixtures: tests/fixtures/sentencepiece_bpe/.jsonl, one +// {"text": ..., "ids": [...]} per line, generated by +// scripts/gen_sentencepiece_bpe_fixture.py. Only the GGUF's tokenizer KVs +// are read (no weights). Gated per model: +// TRANSCRIBE_CANARY_1B_V2_GGUF -> canary-1b-v2.jsonl (single SentencePiece) +// TRANSCRIBE_CANARY_FLASH_GGUF -> canary-flash-en.jsonl (English sub-vocab) +// Ad hoc: sentencepiece_bpe_parity [lo hi rmws]. + +#include "gguf.h" +#include "transcribe-tokenizer.h" + +#include +#include +#include +#include +#include + +namespace { + +struct Case { + std::string text; + std::vector ids; +}; + +// Minimal parser for the fixture's own JSON lines: a "text" string (with +// \\, \", \t, \n, \uXXXX escapes) and an "ids" integer array. +bool parse_line(const std::string & line, Case & out) { + size_t p = line.find("\"text\": \""); + if (p == std::string::npos) { + return false; + } + p += 9; + out.text.clear(); + while (p < line.size() && line[p] != '"') { + char c = line[p++]; + if (c != '\\') { + out.text += c; + continue; + } + const char e = line[p++]; + if (e == 'n') { + out.text += '\n'; + } else if (e == 't') { + out.text += '\t'; + } else if (e == 'u') { + const unsigned cp = static_cast(std::strtoul(line.substr(p, 4).c_str(), nullptr, 16)); + p += 4; + if (cp < 0x80) { + out.text += static_cast(cp); + } else if (cp < 0x800) { + out.text += static_cast(0xC0 | (cp >> 6)); + out.text += static_cast(0x80 | (cp & 0x3F)); + } else { + out.text += static_cast(0xE0 | (cp >> 12)); + out.text += static_cast(0x80 | ((cp >> 6) & 0x3F)); + out.text += static_cast(0x80 | (cp & 0x3F)); + } + } else { + out.text += e; + } + } + p = line.find("\"ids\": [", p); + if (p == std::string::npos) { + return false; + } + p += 8; + out.ids.clear(); + while (p < line.size() && line[p] != ']') { + char * end = nullptr; + out.ids.push_back(static_cast(std::strtol(line.c_str() + p, &end, 10))); + p = static_cast(end - line.c_str()); + while (p < line.size() && (line[p] == ',' || line[p] == ' ')) { + ++p; + } + } + return true; +} + +// Returns the number of mismatches, or -1 when the inputs cannot be loaded. +int run(const char * gguf_path, const std::string & fixture, int lo, int hi, bool rmws) { + gguf_init_params gp{}; + gp.no_alloc = true; + gp.ctx = nullptr; + gguf_context * g = gguf_init_from_file(gguf_path, gp); + if (g == nullptr) { + std::fprintf(stderr, "cannot read %s\n", gguf_path); + return -1; + } + transcribe::Tokenizer tok; + const bool loaded = tok.load(g) == TRANSCRIBE_OK; + gguf_free(g); + std::ifstream f(fixture); + if (!loaded || !f) { + std::fprintf(stderr, "cannot load tokenizer or fixture %s\n", fixture.c_str()); + return -1; + } + int failures = 0, n = 0; + std::string line; + while (std::getline(f, line)) { + Case c; + if (!parse_line(line, c)) { + continue; + } + ++n; + std::vector got; + if (tok.encode_sentencepiece_bpe(c.text, got, lo, hi, rmws) != TRANSCRIBE_OK || got != c.ids) { + ++failures; + std::fprintf(stderr, "FAIL \"%s\"\n expected:", c.text.c_str()); + for (int32_t id : c.ids) { + std::fprintf(stderr, " %d", id); + } + std::fprintf(stderr, "\n actual: "); + for (int32_t id : got) { + std::fprintf(stderr, " %d", id); + } + std::fprintf(stderr, "\n"); + } + } + std::fprintf(stderr, "%s: %d/%d match\n", fixture.c_str(), n - failures, n); + return failures; +} + +} // namespace + +int main(int argc, char ** argv) { + if (argc >= 3) { + const int lo = argc >= 5 ? std::atoi(argv[3]) : 0; + const int hi = argc >= 5 ? std::atoi(argv[4]) : -1; + const bool rmws = argc >= 6 && std::atoi(argv[5]) != 0; + return run(argv[1], argv[2], lo, hi, rmws) == 0 ? EXIT_SUCCESS : EXIT_FAILURE; + } + const std::string dir = std::string(TRANSCRIBE_TEST_FIXTURES_DIR) + "/sentencepiece_bpe/"; + + struct Gate { + const char * env; + const char * fixture; + int lo, hi; + bool rmws; // the SentencePiece model's remove_extra_whitespaces + }; + + // canary-1b/180m-flash: the English sub-vocab is lang_offsets[1] = 1152, + // 1024 pieces (stt.canary.tokenizer.lang_*). + const Gate gates[] = { + { "TRANSCRIBE_CANARY_1B_V2_GGUF", "canary-1b-v2.jsonl", 0, -1, false }, + { "TRANSCRIBE_CANARY_FLASH_GGUF", "canary-flash-en.jsonl", 1152, 2176, true }, + }; + int ran = 0, failures = 0; + for (const Gate & gate : gates) { + const char * path = std::getenv(gate.env); + if (path == nullptr || path[0] == '\0') { + continue; + } + ++ran; + const int r = run(path, dir + gate.fixture, gate.lo, gate.hi, gate.rmws); + failures += r < 0 ? 1 : r; + } + if (ran == 0) { + std::fprintf(stderr, "sentencepiece_bpe_parity: no model env set; skipping\n"); + return 77; + } + return failures == 0 ? EXIT_SUCCESS : EXIT_FAILURE; +} diff --git a/tests/stream_dispatch_unit.cpp b/tests/stream_dispatch_unit.cpp index fa0df5e80..2725321b4 100644 --- a/tests/stream_dispatch_unit.cpp +++ b/tests/stream_dispatch_unit.cpp @@ -1546,8 +1546,72 @@ void test_begin_accepts_min_prefix_run_params() { } // namespace +transcribe_run_params g_stream_seen; + +transcribe_status capture_stream_begin(transcribe_session * session, + const transcribe_run_params * run_params, + const transcribe_stream_params * stream_params) { + (void) session; + (void) stream_params; + g_stream_seen = *run_params; + ++g_begin_calls; + return TRANSCRIBE_OK; +} + +// Streaming rejects a prefix and INSTRUCT before the hook; vocabulary and a +// context prompt reach the hook through library-owned copies. +void test_begin_prompting() { + const transcribe::Arch arch = { + "fake-stream-prompt", nullptr, nullptr, nullptr, nullptr, nullptr, capture_stream_begin, fake_stream_feed, + fake_stream_finalize, nullptr, nullptr, nullptr, + }; + transcribe_model model; + model.arch = &arch; + model.caps.supports_streaming = true; + model.caps.max_timestamp_kind = TRANSCRIBE_TIMESTAMPS_NONE; + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_VOCABULARY, true); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_CONTEXT_PROMPT, true); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_INSTRUCT, true); + transcribe::set_feature(&model, TRANSCRIBE_FEATURE_TRANSCRIPT_PREFIX, true); + + transcribe_session session; + session.model = &model; + transcribe_stream_params sp; + transcribe_stream_params_init(&sp); + + transcribe_run_params rp; + transcribe_run_params_init(&rp); + rp.prefix = "Good morning"; + g_begin_calls = 0; + CHECK(transcribe_stream_begin(&session, &rp, &sp) == TRANSCRIBE_ERR_INVALID_ARG); + rp.prefix = nullptr; + rp.task = TRANSCRIBE_TASK_INSTRUCT; + rp.prompt = "Summarize."; + CHECK(transcribe_stream_begin(&session, &rp, &sp) == TRANSCRIBE_ERR_UNSUPPORTED_TASK); + CHECK(g_begin_calls == 0); + + std::string t0 = "GGUF", t1 = "", t2 = "ggml", ctx = "Earnings call."; + const char * terms[] = { t0.c_str(), t1.c_str(), t2.c_str() }; + rp.task = TRANSCRIBE_TASK_TRANSCRIBE; + rp.vocabulary = terms; + rp.n_vocabulary = 3; + rp.prompt = ctx.c_str(); + CHECK(transcribe_stream_begin(&session, &rp, &sp) == TRANSCRIBE_OK); + CHECK(g_begin_calls == 1); + CHECK(g_stream_seen.n_vocabulary == 2); // empty term dropped + CHECK(g_stream_seen.vocabulary != terms); + CHECK(g_stream_seen.prompt != ctx.c_str()); + t0.assign("XXXX"); + ctx.assign("clobbered"); + CHECK(std::strcmp(g_stream_seen.vocabulary[0], "GGUF") == 0); + CHECK(std::strcmp(g_stream_seen.vocabulary[1], "ggml") == 0); + CHECK(std::strcmp(g_stream_seen.prompt, "Earnings call.") == 0); + transcribe_stream_reset(&session); +} + int main() { test_accessors_on_null_ctx(); + test_begin_prompting(); test_accessors_on_idle_ctx(); test_default_params(); test_begin_null_args();