diff --git a/.claude/skills/porting-1-intake/SKILL.md b/.claude/skills/porting-1-intake/SKILL.md index 0da7b7a89..ebfa1198d 100644 --- a/.claude/skills/porting-1-intake/SKILL.md +++ b/.claude/skills/porting-1-intake/SKILL.md @@ -53,13 +53,20 @@ field semantics. 1. **Reference framework choice.**: The reference implementation should always be the source of truth. If there is not one, this is a red flag and must immediately be told to the human. -2. **Architecture pattern.** One of `encoder-transducer`, `encoder-decoder`, `audio-llm`, `encoder-ctc`. The script's `config.architecture_candidates` is a heuristic starting point. If it doesn't fit, propose a new pattern and have the user accept it based on your research of the architecture. +2. **Architecture pattern.** One of `encoder-transducer`, `encoder-decoder`, `audio-llm`, `encoder-ctc` (`encoder-diarizer` / `encoder-classifier` for the DIARIZE / LANGID roles). The script's `config.architecture_candidates` is a heuristic starting point. If it doesn't fit, propose a new pattern and have the user accept it based on your research of the architecture. **Role.** ASR (the product is a transcript) or DIARIZE (the product is who spoke when; see `docs/roles.md`). A DIARIZE port implements `DiarizeOps` (`src/transcribe-diarize.h`) instead of the ASR hooks, sets `"role": "diarize"` in its catalog record, and is accepted on DER - (`scripts/diar/`), not WER. + (`scripts/diar/`), not WER. A LANGID port (the product is a language + decision; `docs/langid.md`) implements `LangidOps` + (`src/transcribe-langid.h`) and only returns logits over its label table, + sets `"role": "langid"` with `metric: "accuracy"` / `acc_pct` rows, and + is accepted on top-1 decision parity with the reference over FLEURS + (`scripts/langid/`), not WER. Its golden cases set `"language": null` + and carry the expected label as `expected_language`, so validate.py never + hands the answer to the model. 3. **Acceptance dataset.** Default: LibriSpeech test-clean. Capture any publisher-reported score in `upstream_benchmarks` when available for diff --git a/.claude/skills/porting-7-wer/SKILL.md b/.claude/skills/porting-7-wer/SKILL.md index 4005e5e55..9f7312c86 100644 --- a/.claude/skills/porting-7-wer/SKILL.md +++ b/.claude/skills/porting-7-wer/SKILL.md @@ -42,7 +42,7 @@ WER progress: ### Step 1: Acceptance manifest (execute or ask-point) -Read the acceptance dataset and metric from `upstream_benchmarks[0].{dataset, metric}`. Slugify: lowercase, spaces → hyphens. `"LibriSpeech test-clean"` → `librispeech-test-clean`. Resolve to `$MANIFEST`; every later step uses `$MANIFEST`, never a reconstructed path. Confirm `metric` is `wer` or `cer` (`der` for a DIARIZE-role model: score with `scripts/diar/` instead; anything else is out of scope). Do not use any publisher-reported score for pass/fail; the measured Oracle reference score is the gate target. +Read the acceptance dataset and metric from `upstream_benchmarks[0].{dataset, metric}`. Slugify: lowercase, spaces → hyphens. `"LibriSpeech test-clean"` → `librispeech-test-clean`. Resolve to `$MANIFEST`; every later step uses `$MANIFEST`, never a reconstructed path. Confirm `metric` is `wer` or `cer` (`der` for a DIARIZE-role model: score with `scripts/diar/` instead; `accuracy` for a LANGID-role model: run `scripts/langid/run.py` for the reference and every shipped quant, then score each with `scripts/langid/score.py --ref `, which gates the ref dtype on top-1 agreement (`--report-only` for the quants); anything else is out of scope). Do not use any publisher-reported score for pass/fail; the measured Oracle reference score is the gate target. If the intake's dataset is not covered by `scripts/wer/ingest.py`, extend that script before running this step. diff --git a/.gitignore b/.gitignore index 3c60cc0cd..7a3c3366a 100644 --- a/.gitignore +++ b/.gitignore @@ -59,6 +59,8 @@ __pycache__/ # WER evaluation data + generated working reports. /samples/wer/ +# FLEURS language ID corpus (scripts/langid/ingest.py). +/samples/langid/ # Diarization eval corpora (AMI audio + fetched RTTMs); regenerable via # scripts/diar/ingest_ami.py + fetch_ami_forced_alignment.py. The committed # oracle case lives at samples/sortformer-2spk-mix.wav (outside samples/diar/). diff --git a/README.md b/README.md index 1393cc9f2..8a483416e 100644 --- a/README.md +++ b/README.md @@ -39,6 +39,14 @@ C/C++ speech-to-text inference library. Runs diverse STT model families via [GGU | Streaming Sortformer Diarizer 4spk v2.1 | `diar_streaming_sortformer_4spk-v2.1` | diarize, streaming | [docs/models/diar_streaming_sortformer_4spk-v2.1.md](docs/models/diar_streaming_sortformer_4spk-v2.1.md) | +**Language ID models** (no transcription; verified by top-1 decision parity on FLEURS; see [`docs/langid.md`](docs/langid.md)): + + +| Family | Variants | Available capabilities | Docs | +| --- | --- | --- | --- | +| VoxLingua107 ECAPA-TDNN (language ID) | `lang-id-voxlingua107-ecapa` | language ID (107 languages) | [docs/models/lang-id-voxlingua107-ecapa.md](docs/models/lang-id-voxlingua107-ecapa.md) | + + Per-variant model cards live under [`docs/models/`](docs/models/). ## Model catalog diff --git a/bindings/python/README.md b/bindings/python/README.md index 9422adf99..bb70a14e8 100644 --- a/bindings/python/README.md +++ b/bindings/python/README.md @@ -88,6 +88,25 @@ with model.diarize_session() as diarizer: print(turn.speaker_id, turn.t0_ms, turn.t1_ms) ``` +### Language ID + +Models whose `model.roles` include `Role.LANGID` (VoxLingua107 ECAPA-TDNN) +identify the spoken language through a langid session. `run()` returns a +`LangIdResult` with candidates ranked by `p`. Codes are the model's own labels +(`"iw"`, `"jw"`; `model.langid_label_index("he")` resolves aliases), so match +`result.code` against the ASR model's `capabilities.languages` yourself. +`allowed=None` scores every label; `allowed_mass` is the unrestricted +probability the allowed set captured. An empty `allowed` list raises +`InvalidArgument`; clips under `model.langid_info.min_audio_ms` raise +`InputTooShort`. Locking, `Busy`, `cancel()` and `close()` work as on +`Session`. + +```python +with model.langid_session() as lid: + result = lid.run(pcm, allowed=["en", "de", "fr"], top_k=3) + print(result.code, result.candidates[0].p, result.allowed_mass) +``` + ## Backends `Model(backend=...)` applies a backend policy (`"auto"` uses the best diff --git a/bindings/python/src/transcribe_cpp/__init__.py b/bindings/python/src/transcribe_cpp/__init__.py index a77f20225..3b0cb46f4 100644 --- a/bindings/python/src/transcribe_cpp/__init__.py +++ b/bindings/python/src/transcribe_cpp/__init__.py @@ -39,6 +39,7 @@ BackendError, Busy, InputTooLong, + InputTooShort, InvalidArgument, ModelFileNotFound, ModelLoadError, @@ -78,6 +79,10 @@ "Session", "DiarizeSession", "DiarizeInfo", + "LangIdSession", + "LangIdInfo", + "LangIdResult", + "LangIdCandidate", "Role", "Result", "Segment", @@ -120,6 +125,7 @@ "Aborted", "Busy", "InputTooLong", + "InputTooShort", "OutputTruncated", "OutputRepetition", "native_version", @@ -190,6 +196,8 @@ _Timings = _generated.transcribe_timings _Segment = _generated.transcribe_segment _SpeakerSegment = _generated.transcribe_speaker_segment +_LangIdCandidate = _generated.transcribe_langid_candidate +_LangIdResult = _generated.transcribe_langid_result _Word = _generated.transcribe_word _Token = _generated.transcribe_token _StreamParams = _generated.transcribe_stream_params @@ -558,10 +566,12 @@ class Capabilities: class Role(enum.Enum): """What a model serves (``Model.roles``): ASR is ``Model.session()``, - DIARIZE is ``Model.diarize_session()``.""" + DIARIZE is ``Model.diarize_session()``, LANGID is + ``Model.langid_session()``.""" ASR = _generated.TRANSCRIBE_ROLE_ASR DIARIZE = _generated.TRANSCRIBE_ROLE_DIARIZE + LANGID = _generated.TRANSCRIBE_ROLE_LANGID @dataclass(frozen=True) @@ -570,6 +580,43 @@ class DiarizeInfo: max_speakers: int # speaker_id is in [1, max_speakers] +@dataclass(frozen=True) +class LangIdInfo: + sample_rate: int + n_labels: int # label indices are [0, n_labels) + min_audio_ms: int # shorter scored audio raises InputTooShort + + +@dataclass(frozen=True) +class LangIdCandidate: + """One ranked label. ``code`` is the model's own label ("en", "iw"); + ``p`` is renormalized over the allowed set.""" + + index: int + code: str + name: str + p: float + logit: float + + +@dataclass(frozen=True) +class LangIdResult: + """``candidates`` are ranked by ``p`` (ties keep label order). + ``allowed_mass`` is the unrestricted probability inside the allowed set + (1.0 when unrestricted); a low value means the speech is probably outside + it. ``audio_ms`` is what was scored, after the crop.""" + + candidates: tuple[LangIdCandidate, ...] + n_allowed: int + allowed_mass: float + audio_ms: int + + @property + def code(self) -> str | None: + """The top candidate's code, or None when there are no candidates.""" + return self.candidates[0].code if self.candidates else None + + @dataclass(frozen=True) class SessionLimits: """Effective per-session limits (model bound narrowed by session params). @@ -1227,6 +1274,39 @@ def diarize_session(self, *, n_threads: int = 0) -> "DiarizeSession": """Raises :class:`UnsupportedRole` without the DIARIZE role.""" return DiarizeSession(self, n_threads=n_threads) + @property + def langid_info(self) -> LangIdInfo: + """Raises :class:`UnsupportedRole` without the LANGID role.""" + info = _generated.transcribe_langid_info() + _lib.transcribe_langid_info_init(_byref(info)) + _check(_lib.transcribe_langid_get_info(self._h, _byref(info)), + "reading langid info") + return LangIdInfo(sample_rate=info.sample_rate, n_labels=info.n_labels, + min_audio_ms=info.min_audio_ms) + + @property + def langid_labels(self) -> tuple[tuple[str, str], ...]: + """``(code, name)`` per label index. Raises :class:`UnsupportedRole` + without the LANGID role.""" + n = self.langid_info.n_labels + return tuple((_decode(_lib.transcribe_langid_label_code(self._h, i)), + _decode(_lib.transcribe_langid_label_name(self._h, i))) + for i in range(n)) + + def langid_label_index(self, code: str) -> int | None: + """Label index of a code or alias ("he" and "iw" name the same + label), or None. None as well on a model without the LANGID role. + Raises :class:`InvalidArgument` if ``code`` contains a NUL.""" + i = _lib.transcribe_langid_label_index(self._h, _cstr(code, "code")) + return i if i >= 0 else None + + def langid_session(self, *, n_threads: int = 0, + max_audio_ms: int = 0) -> "LangIdSession": + """``max_audio_ms``: longer input is scored on its last + ``max_audio_ms`` (0 = 30000). Raises :class:`UnsupportedRole` + without the LANGID role.""" + return LangIdSession(self, n_threads=n_threads, max_audio_ms=max_audio_ms) + def close(self) -> None: """Free the model. Any session still open on it is closed first — the C contract forbids freeing a model before its sessions, so this @@ -1888,6 +1968,91 @@ def timings(self) -> Timings: return _timings_from(tm) +class LangIdSession(_SessionBase): + """A language ID context on a model with the LANGID role. Compute + locking and ``Busy`` rules: see ``Model``.""" + + _free_fn = "transcribe_langid_session_free" + + def __init__(self, model: Model, *, n_threads: int = 0, max_audio_ms: int = 0): + self._model = model # keep the model alive for the session's lifetime + params = _generated.transcribe_langid_session_params() + _lib.transcribe_langid_session_params_init(_byref(params)) + params.n_threads = n_threads + params.max_audio_ms = max_audio_ms + + handle = ctypes.c_void_p() + _check(_lib.transcribe_langid_session_init(model._h, _byref(params), _byref(handle)), + "opening langid session") + if not handle.value: + raise TranscribeError("langid session init returned a null handle") + self._handle = handle + self._arm_abort(_lib.transcribe_langid_set_abort_callback) + + def run(self, pcm: PCMLike, *, allowed: "Sequence[str] | None" = None, + top_k: int = 0) -> LangIdResult: + """Identify the language of one clip (16 kHz mono float32 PCM). + Longer input than the session's ``max_audio_ms`` is scored on its + tail. ``allowed`` restricts the decision to those codes or aliases; + None means every label, and an empty list is rejected. ``top_k`` + keeps the best candidates (0 = every allowed label). + + Raises :class:`InputTooShort` below ``LangIdInfo.min_audio_ms``, + :class:`UnsupportedRequest` for an unknown code, :class:`Aborted` + after :meth:`cancel`, and :class:`Busy` if a stream is active on this + model.""" + self._cancel.clear() # before the lock wait, as in Session.run() + array, n_samples = _pcm_to_carray(pcm) + params = _generated.transcribe_langid_params() + _lib.transcribe_langid_params_init(_byref(params)) + params.top_k = top_k + if allowed is not None: + if isinstance(allowed, (str, bytes)): + raise InvalidArgument("allowed must be a sequence of codes, not a string") + codes = list(allowed) + # An empty list must not reach native as NULL, which means "all". + if not codes: + raise InvalidArgument("allowed is empty; pass None for every label") + if not all(isinstance(c, str) for c in codes): + raise InvalidArgument("allowed entries must be str") + encoded = [_cstr(c, "allowed") for c in codes] + arr = (ctypes.c_char_p * len(encoded))(*encoded) + params.allowed = ctypes.cast(arr, type(params.allowed)) + params.n_allowed = len(encoded) + params._allowed_keepalive = (encoded, arr) + with self._model._exclusive( + "langid_run", busy="a stream is active on this model; " + "finish or drop it before langid run()"): + h = self._h # captured under the lock; close() defers its free + _check(_lib.transcribe_langid_run(h, array, n_samples, _byref(params)), + "transcribe_langid_run") + res = _LangIdResult() + _lib.transcribe_langid_result_init(_byref(res)) + _check(_lib.transcribe_langid_get_result(h, _byref(res)), + "transcribe_langid_get_result") + rows = [] + for i in range(res.n_candidates): + c = _LangIdCandidate() + _lib.transcribe_langid_candidate_init(_byref(c)) + _check(_lib.transcribe_langid_get_candidate(h, i, _byref(c)), + "transcribe_langid_get_candidate") + rows.append(LangIdCandidate( + index=c.index, code=_decode(c.code), name=_decode(c.name), + p=c.p, logit=c.logit)) + return LangIdResult(candidates=tuple(rows), n_allowed=res.n_allowed, + allowed_mass=res.allowed_mass, audio_ms=res.audio_ms) + + @property + def timings(self) -> Timings: + """Load time plus the last run's mel / encode time. Not locked, like + ``Session.limits``.""" + tm = _Timings() + _lib.transcribe_timings_init(_byref(tm)) + _check(_lib.transcribe_langid_get_timings(self._h, _byref(tm)), + "transcribe_langid_get_timings") + return _timings_from(tm) + + def transcribe( model: Model | str | os.PathLike, pcm: PCMLike, diff --git a/bindings/python/src/transcribe_cpp/_generated.py b/bindings/python/src/transcribe_cpp/_generated.py index f908877ae..2354b16a5 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 = "bd3273dabb25a1fe" +PUBLIC_HEADER_HASH = "d544d70a2b5acf50" # === enum constants === TRANSCRIBE_OK = 0 @@ -37,6 +37,7 @@ TRANSCRIBE_ERR_OUTPUT_TRUNCATED = 18 TRANSCRIBE_ERR_OUTPUT_REPETITION = 19 TRANSCRIBE_ERR_UNSUPPORTED_ROLE = 20 +TRANSCRIBE_ERR_INPUT_TOO_SHORT = 21 TRANSCRIBE_ABI_MODEL_LOAD_PARAMS = 0 TRANSCRIBE_ABI_SESSION_PARAMS = 1 TRANSCRIBE_ABI_RUN_PARAMS = 2 @@ -56,6 +57,11 @@ TRANSCRIBE_ABI_DIARIZE_INFO = 16 TRANSCRIBE_ABI_DIARIZE_SESSION_PARAMS = 17 TRANSCRIBE_ABI_DIARIZE_PARAMS = 18 +TRANSCRIBE_ABI_LANGID_INFO = 19 +TRANSCRIBE_ABI_LANGID_SESSION_PARAMS = 20 +TRANSCRIBE_ABI_LANGID_PARAMS = 21 +TRANSCRIBE_ABI_LANGID_RESULT = 22 +TRANSCRIBE_ABI_LANGID_CANDIDATE = 23 TRANSCRIBE_LOG_LEVEL_NONE = 0 TRANSCRIBE_LOG_LEVEL_INFO = 1 TRANSCRIBE_LOG_LEVEL_WARN = 2 @@ -98,6 +104,7 @@ TRANSCRIBE_DEVICE_TYPE_ACCEL = 3 TRANSCRIBE_ROLE_ASR = 1 TRANSCRIBE_ROLE_DIARIZE = 2 +TRANSCRIBE_ROLE_LANGID = 4 TRANSCRIBE_FEATURE_INITIAL_PROMPT = 0 TRANSCRIBE_FEATURE_TEMPERATURE_FALLBACK = 1 TRANSCRIBE_FEATURE_LONG_FORM = 2 @@ -177,6 +184,16 @@ class transcribe_diarize_session_params(_c.Structure): pass class transcribe_diarize_params(_c.Structure): pass +class transcribe_langid_info(_c.Structure): + pass +class transcribe_langid_session_params(_c.Structure): + pass +class transcribe_langid_params(_c.Structure): + pass +class transcribe_langid_result(_c.Structure): + pass +class transcribe_langid_candidate(_c.Structure): + pass class transcribe_moonshine_streaming_stream_ext(_c.Structure): pass class transcribe_parakeet_stream_ext(_c.Structure): @@ -211,6 +228,11 @@ class transcribe_whisper_chunk_trace(_c.Structure): transcribe_diarize_info._fields_ = [("struct_size", _c.c_uint64), ("sample_rate", _c.c_int32), ("max_speakers", _c.c_int32)] transcribe_diarize_session_params._fields_ = [("struct_size", _c.c_uint64), ("n_threads", _c.c_int32)] transcribe_diarize_params._fields_ = [("struct_size", _c.c_uint64), ("family", _c.POINTER(transcribe_ext))] +transcribe_langid_info._fields_ = [("struct_size", _c.c_uint64), ("sample_rate", _c.c_int32), ("n_labels", _c.c_int32), ("min_audio_ms", _c.c_int32)] +transcribe_langid_session_params._fields_ = [("struct_size", _c.c_uint64), ("n_threads", _c.c_int32), ("max_audio_ms", _c.c_int32)] +transcribe_langid_params._fields_ = [("struct_size", _c.c_uint64), ("allowed", _c.POINTER(_c.c_char_p)), ("n_allowed", _c.c_int32), ("top_k", _c.c_int32)] +transcribe_langid_result._fields_ = [("struct_size", _c.c_uint64), ("n_candidates", _c.c_int32), ("n_allowed", _c.c_int32), ("allowed_mass", _c.c_float), ("audio_ms", _c.c_int64)] +transcribe_langid_candidate._fields_ = [("struct_size", _c.c_uint64), ("index", _c.c_int32), ("code", _c.c_char_p), ("name", _c.c_char_p), ("p", _c.c_float), ("logit", _c.c_float)] transcribe_moonshine_streaming_stream_ext._fields_ = [("ext", transcribe_ext), ("min_decode_interval_ms", _c.c_int32)] transcribe_parakeet_stream_ext._fields_ = [("ext", transcribe_ext), ("att_context_right", _c.c_int32)] transcribe_parakeet_buffered_stream_ext._fields_ = [("ext", transcribe_ext), ("left_ms", _c.c_int32), ("chunk_ms", _c.c_int32), ("right_ms", _c.c_int32)] @@ -241,6 +263,11 @@ class transcribe_whisper_chunk_trace(_c.Structure): 'transcribe_diarize_info': 16, 'transcribe_diarize_session_params': 17, 'transcribe_diarize_params': 18, + 'transcribe_langid_info': 19, + 'transcribe_langid_session_params': 20, + 'transcribe_langid_params': 21, + 'transcribe_langid_result': 22, + 'transcribe_langid_candidate': 23, } # C-compiler layout captured at generation (for offset self-check). @@ -264,6 +291,11 @@ class transcribe_whisper_chunk_trace(_c.Structure): 'transcribe_diarize_info': {'size': 16, 'align': 8, 'offsets': {'struct_size': 0, 'sample_rate': 8, 'max_speakers': 12}}, 'transcribe_diarize_session_params': {'size': 16, 'align': 8, 'offsets': {'struct_size': 0, 'n_threads': 8}}, 'transcribe_diarize_params': {'size': 16, 'align': 8, 'offsets': {'struct_size': 0, 'family': 8}}, + 'transcribe_langid_info': {'size': 24, 'align': 8, 'offsets': {'struct_size': 0, 'sample_rate': 8, 'n_labels': 12, 'min_audio_ms': 16}}, + 'transcribe_langid_session_params': {'size': 16, 'align': 8, 'offsets': {'struct_size': 0, 'n_threads': 8, 'max_audio_ms': 12}}, + 'transcribe_langid_params': {'size': 24, 'align': 8, 'offsets': {'struct_size': 0, 'allowed': 8, 'n_allowed': 16, 'top_k': 20}}, + 'transcribe_langid_result': {'size': 32, 'align': 8, 'offsets': {'struct_size': 0, 'n_candidates': 8, 'n_allowed': 12, 'allowed_mass': 16, 'audio_ms': 24}}, + 'transcribe_langid_candidate': {'size': 40, 'align': 8, 'offsets': {'struct_size': 0, 'index': 8, 'code': 16, 'name': 24, 'p': 32, 'logit': 36}}, 'transcribe_moonshine_streaming_stream_ext': {'size': 24, 'align': 8, 'offsets': {'ext': 0, 'min_decode_interval_ms': 16}}, 'transcribe_parakeet_stream_ext': {'size': 24, 'align': 8, 'offsets': {'ext': 0, 'att_context_right': 16}}, 'transcribe_parakeet_buffered_stream_ext': {'size': 32, 'align': 8, 'offsets': {'ext': 0, 'left_ms': 16, 'chunk_ms': 20, 'right_ms': 24}}, @@ -378,6 +410,38 @@ def configure(lib): lib.transcribe_init_backends_default.argtypes = [] lib.transcribe_init_backends_ex.restype = _c.c_int lib.transcribe_init_backends_ex.argtypes = [_c.POINTER(transcribe_backend_init_params)] + lib.transcribe_langid_candidate_init.restype = None + lib.transcribe_langid_candidate_init.argtypes = [_c.POINTER(transcribe_langid_candidate)] + lib.transcribe_langid_get_candidate.restype = _c.c_int + lib.transcribe_langid_get_candidate.argtypes = [_c.c_void_p, _c.c_int, _c.POINTER(transcribe_langid_candidate)] + lib.transcribe_langid_get_info.restype = _c.c_int + lib.transcribe_langid_get_info.argtypes = [_c.c_void_p, _c.POINTER(transcribe_langid_info)] + lib.transcribe_langid_get_result.restype = _c.c_int + lib.transcribe_langid_get_result.argtypes = [_c.c_void_p, _c.POINTER(transcribe_langid_result)] + lib.transcribe_langid_get_timings.restype = _c.c_int + lib.transcribe_langid_get_timings.argtypes = [_c.c_void_p, _c.POINTER(transcribe_timings)] + lib.transcribe_langid_info_init.restype = None + lib.transcribe_langid_info_init.argtypes = [_c.POINTER(transcribe_langid_info)] + lib.transcribe_langid_label_code.restype = _c.c_char_p + lib.transcribe_langid_label_code.argtypes = [_c.c_void_p, _c.c_int32] + lib.transcribe_langid_label_index.restype = _c.c_int32 + lib.transcribe_langid_label_index.argtypes = [_c.c_void_p, _c.c_char_p] + lib.transcribe_langid_label_name.restype = _c.c_char_p + lib.transcribe_langid_label_name.argtypes = [_c.c_void_p, _c.c_int32] + lib.transcribe_langid_params_init.restype = None + lib.transcribe_langid_params_init.argtypes = [_c.POINTER(transcribe_langid_params)] + lib.transcribe_langid_result_init.restype = None + lib.transcribe_langid_result_init.argtypes = [_c.POINTER(transcribe_langid_result)] + lib.transcribe_langid_run.restype = _c.c_int + lib.transcribe_langid_run.argtypes = [_c.c_void_p, _c.POINTER(_c.c_float), _c.c_int, _c.POINTER(transcribe_langid_params)] + lib.transcribe_langid_session_free.restype = None + lib.transcribe_langid_session_free.argtypes = [_c.c_void_p] + lib.transcribe_langid_session_init.restype = _c.c_int + lib.transcribe_langid_session_init.argtypes = [_c.c_void_p, _c.POINTER(transcribe_langid_session_params), _c.POINTER(_c.c_void_p)] + lib.transcribe_langid_session_params_init.restype = None + lib.transcribe_langid_session_params_init.argtypes = [_c.POINTER(transcribe_langid_session_params)] + lib.transcribe_langid_set_abort_callback.restype = None + lib.transcribe_langid_set_abort_callback.argtypes = [_c.c_void_p, _c.CFUNCTYPE(_c.c_bool, _c.c_void_p), _c.c_void_p] lib.transcribe_log_set.restype = None lib.transcribe_log_set.argtypes = [_c.CFUNCTYPE(None, _c.c_int, _c.c_char_p, _c.c_void_p), _c.c_void_p] lib.transcribe_model_accepts_ext_kind.restype = _c.c_bool diff --git a/bindings/python/src/transcribe_cpp/errors.py b/bindings/python/src/transcribe_cpp/errors.py index 236e94c67..b82fea02d 100644 --- a/bindings/python/src/transcribe_cpp/errors.py +++ b/bindings/python/src/transcribe_cpp/errors.py @@ -19,6 +19,7 @@ TRANSCRIBE_ERR_FILE_NOT_FOUND as ERR_FILE_NOT_FOUND, TRANSCRIBE_ERR_GGUF as ERR_GGUF, TRANSCRIBE_ERR_INPUT_TOO_LONG as ERR_INPUT_TOO_LONG, + TRANSCRIBE_ERR_INPUT_TOO_SHORT as ERR_INPUT_TOO_SHORT, TRANSCRIBE_ERR_INVALID_ARG as ERR_INVALID_ARG, TRANSCRIBE_ERR_NOT_IMPLEMENTED as ERR_NOT_IMPLEMENTED, TRANSCRIBE_ERR_OOM as ERR_OOM, @@ -99,6 +100,11 @@ class InputTooLong(TranscribeError): pass +class InputTooShort(TranscribeError): + """The audio is shorter than the role's minimum (e.g. language ID scores + at least ``LangIdInfo.min_audio_ms``). Nothing was computed.""" + + class Busy(TranscribeError): """A stream is active on this model, so the call was refused instead of started (see ``Model``). Finalize or reset the stream first, or use one @@ -160,6 +166,7 @@ class OutputRepetition(OutputTruncated): ERR_OUTPUT_TRUNCATED: OutputTruncated, ERR_OUTPUT_REPETITION: OutputRepetition, ERR_UNSUPPORTED_ROLE: UnsupportedRole, + ERR_INPUT_TOO_SHORT: InputTooShort, } diff --git a/bindings/python/tests/conftest.py b/bindings/python/tests/conftest.py index 392ef2935..a1eb6a5b6 100644 --- a/bindings/python/tests/conftest.py +++ b/bindings/python/tests/conftest.py @@ -64,6 +64,9 @@ / "models/diar_streaming_sortformer_4spk-v2.1" / "diar_streaming_sortformer_4spk-v2.1-F32.gguf" ) +LANGID_MODEL = ( + REPO / "models/lang-id-voxlingua107-ecapa/lang-id-voxlingua107-ecapa-Q8_0.gguf" +) PNC_MODEL = REPO / "models/canary-180m-flash/canary-180m-flash-Q8_0.gguf" ITN_MODEL = REPO / "models/SenseVoiceSmall/SenseVoiceSmall-Q8_0.gguf" @@ -172,6 +175,28 @@ def sortformer_model_path() -> Path: return _family_model("TRANSCRIBE_SMOKE_SORTFORMER_MODEL", SORTFORMER_MODEL) +@pytest.fixture(scope="session") +def langid_model_path() -> Path: + """VoxLingua107 ECAPA-TDNN (serves only the LANGID role).""" + return _family_model("TRANSCRIBE_SMOKE_LANGID_MODEL", LANGID_MODEL) + + +@pytest.fixture(scope="session") +def langid_toy_model_path(tmp_path_factory) -> Path: + """The toy ecapa_tdnn GGUF from tests/fixtures/make_gguf_fixtures.py + (5 labels aa..ee, alias xx=aa, random weights). The generator is + dependency-free, so this always runs; results are structural only.""" + import importlib.util + + gen = REPO / "tests/fixtures/make_gguf_fixtures.py" + spec = importlib.util.spec_from_file_location("make_gguf_fixtures", gen) + mod = importlib.util.module_from_spec(spec) + spec.loader.exec_module(mod) + path = tmp_path_factory.mktemp("langid") / "ecapa_tdnn_toy.gguf" + path.write_bytes(mod._ecapa_tdnn_gguf(mod.ECAPA_LABEL_CODES, mod.ECAPA_LABEL_NAMES, ["xx=aa"])) + return path + + @pytest.fixture(scope="session") def pnc_model_path() -> Path: """Canary model whose generic PNC run parameter changes the prompt.""" diff --git a/bindings/python/tests/test_compute_lock.py b/bindings/python/tests/test_compute_lock.py index ef65ddf8a..dc910afcb 100644 --- a/bindings/python/tests/test_compute_lock.py +++ b/bindings/python/tests/test_compute_lock.py @@ -55,6 +55,8 @@ def fake(monkeypatch): lambda h: frees.append(("model", h.value))) monkeypatch.setattr(t._lib, "transcribe_diarize_session_free", lambda h: frees.append(("diarize", h.value))) + monkeypatch.setattr(t._lib, "transcribe_langid_session_free", + lambda h: frees.append(("langid", h.value))) resets: list = [] # session handles transcribe_stream_reset was given monkeypatch.setattr(t._lib, "transcribe_stream_reset", lambda h: resets.append(h.value)) @@ -183,6 +185,7 @@ def test_every_compute_site_holds_the_lock(fake, monkeypatch): m = fake.model() s = fake.session(m) d = fake.session(m, cls=t.DiarizeSession) + lid = fake.session(m, cls=t.LangIdSession) seen: list = [] def holds(name, ret=0): @@ -196,7 +199,8 @@ def probe(*args): "transcribe_stream_begin", "transcribe_stream_feed", "transcribe_stream_finalize", "transcribe_stream_reset", "transcribe_batch_status", "transcribe_diarize_run", - "transcribe_diarize_n_segments"): + "transcribe_diarize_n_segments", "transcribe_langid_run", + "transcribe_langid_get_result"): monkeypatch.setattr(t._lib, name, holds(name)) monkeypatch.setattr(t._lib, "transcribe_batch_n_results", holds("transcribe_batch_n_results", ret=1)) @@ -215,6 +219,7 @@ def materialize(self, h, utt=None): stream.finalize() stream.reset() assert d.run(PCM) == [] + assert lid.run(PCM).candidates == () names = [n for n, _ in seen] for site in ("transcribe_run", "transcribe_run_batch", @@ -222,6 +227,7 @@ def materialize(self, h, utt=None): "transcribe_stream_begin", "transcribe_stream_feed", "transcribe_stream_finalize", "transcribe_stream_reset", "transcribe_diarize_run", "transcribe_diarize_n_segments", + "transcribe_langid_run", "transcribe_langid_get_result", "copy-out"): assert site in names, f"{site} never reached" assert all(ok for _, ok in seen), [n for n, ok in seen if not ok] @@ -571,3 +577,45 @@ def test_diarize_run_busy_while_stream_active(fake, native, monkeypatch): assert calls == [] and not m._compute_lock.locked() stream.finalize() assert d.run(PCM) == [] and calls == ["run"] + + +# --- LangIdSession -------------------------------------------------------------- + +LANGID_BUSY = ("a stream is active on this model; " + "finish or drop it before langid run()") + + +def test_langid_run_busy_while_stream_active(fake, native, monkeypatch): + m = fake.model() + s, lid = fake.session(m), fake.session(m, cls=t.LangIdSession) + calls: list = [] + monkeypatch.setattr(t._lib, "transcribe_langid_run", + lambda *a: calls.append("run") or 0) + monkeypatch.setattr(t._lib, "transcribe_langid_get_result", lambda h, out: 0) + stream = s.stream() + with pytest.raises(t.Busy) as ei: + lid.run(PCM) + assert str(ei.value) == LANGID_BUSY + assert calls == [] and not m._compute_lock.locked() + stream.finalize() + assert lid.run(PCM).candidates == () and calls == ["run"] + + +def test_langid_allowed_kept_alive_and_empty_rejected(fake, monkeypatch): + m = fake.model() + lid = fake.session(m, cls=t.LangIdSession) + seen: list = [] + + def run(h, pcm, n, params): + p = params._obj + seen.append([p.allowed[i].decode() for i in range(p.n_allowed)]) + return 0 + + monkeypatch.setattr(t._lib, "transcribe_langid_run", run) + monkeypatch.setattr(t._lib, "transcribe_langid_get_result", lambda h, out: 0) + lid.run(PCM, allowed=["en", "de"]) + lid.run(PCM, allowed=None) + assert seen == [["en", "de"], []] + with pytest.raises(t.InvalidArgument): + lid.run(PCM, allowed=[]) + assert len(seen) == 2 # rejected before the native call diff --git a/bindings/python/tests/test_errors.py b/bindings/python/tests/test_errors.py index 1c7b64d71..92954b9ad 100644 --- a/bindings/python/tests/test_errors.py +++ b/bindings/python/tests/test_errors.py @@ -43,6 +43,7 @@ def test_every_status_maps_to_documented_subclass(): errors.ERR_OUTPUT_TRUNCATED: t.OutputTruncated, errors.ERR_OUTPUT_REPETITION: t.OutputRepetition, errors.ERR_UNSUPPORTED_ROLE: t.UnsupportedRole, + errors.ERR_INPUT_TOO_SHORT: t.InputTooShort, } # The mapping table covers every non-OK status the header defines, and # nothing else (a new C status must be mapped deliberately, not by diff --git a/bindings/python/tests/test_langid.py b/bindings/python/tests/test_langid.py new file mode 100644 index 000000000..6888e9367 --- /dev/null +++ b/bindings/python/tests/test_langid.py @@ -0,0 +1,147 @@ +"""LANGID role: roles, info, label table, LangIdSession.run and its allowed / +top_k contract. Contract tests run on the toy ecapa_tdnn fixture (always +available); accuracy tests need the real VoxLingua107 GGUF. The locking / +Busy / close / cancel rules are pinned model-free in test_compute_lock.py.""" + +from __future__ import annotations + +import array +import math + +import pytest + +import transcribe_cpp as t +from conftest import SAMPLES, load_wav + + +def _noise(n: int, seed: int = 7) -> array.array: + out = array.array("f") + s = seed | 1 + for _ in range(n): + s = (s * 1664525 + 1013904223) & 0xFFFFFFFF + out.append(((s >> 8) & 0xFFFFFF) / 16777216.0 - 0.5) + return out + + +@pytest.fixture(scope="module") +def noise_1s(): + return _noise(16000) + + +def test_toy_roles_info_labels(langid_toy_model_path): + with t.Model(langid_toy_model_path, backend="cpu") as model: + assert model.roles == {t.Role.LANGID} + assert model.langid_info == t.LangIdInfo(sample_rate=16000, n_labels=5, min_audio_ms=500) + assert model.langid_labels[2] == ("cc", "Charlie") + assert model.langid_label_index("xx") == 0 + assert model.langid_label_index("zz") is None + with pytest.raises(t.UnsupportedRole): + model.capabilities + with pytest.raises(t.UnsupportedRole): + model.session() + with pytest.raises(t.UnsupportedRole): + model.diarize_session() + + +def test_asr_only_model_rejects_langid(model_path): + with t.Model(model_path) as model: + with pytest.raises(t.UnsupportedRole): + model.langid_info + with pytest.raises(t.UnsupportedRole): + model.langid_session() + assert model.langid_label_index("en") is None + + +def test_toy_run_contract(langid_toy_model_path, noise_1s): + with t.Model(langid_toy_model_path, backend="cpu") as model: + with model.langid_session(n_threads=1) as lid: + r = lid.run(noise_1s) + assert len(r.candidates) == 5 and r.n_allowed == 5 + assert r.allowed_mass == 1.0 and r.audio_ms == 1000 + assert r.code == r.candidates[0].code + assert math.isclose(sum(c.p for c in r.candidates), 1.0, rel_tol=1e-5) + ps = [c.p for c in r.candidates] + assert ps == sorted(ps, reverse=True) + + restricted = lid.run(noise_1s, allowed=["bb", "dd"]) + assert {c.code for c in restricted.candidates} == {"bb", "dd"} + assert restricted.n_allowed == 2 and 0.0 < restricted.allowed_mass < 1.0 + + alias = lid.run(noise_1s, allowed=("xx",)) + assert [c.code for c in alias.candidates] == ["aa"] and alias.candidates[0].p == 1.0 + + top = lid.run(noise_1s, top_k=2) + assert len(top.candidates) == 2 and top.n_allowed == 5 + + assert lid.timings.encode_ms > 0.0 + + +def test_toy_allowed_rejections(langid_toy_model_path, noise_1s): + with t.Model(langid_toy_model_path, backend="cpu") as model: + with model.langid_session() as lid: + # An empty list must not silently mean "all labels". + with pytest.raises(t.InvalidArgument): + lid.run(noise_1s, allowed=[]) + with pytest.raises(t.InvalidArgument): + lid.run(noise_1s, allowed="en") + with pytest.raises(t.InvalidArgument): + lid.run(noise_1s, allowed=["aa", None]) + with pytest.raises(t.UnsupportedRequest): + lid.run(noise_1s, allowed=["zz"]) + with pytest.raises(t.InvalidArgument): + lid.run(noise_1s, top_k=-1) + lid.run(noise_1s, allowed=None) # None is every label + + +def test_toy_interior_nul_is_rejected(langid_toy_model_path, noise_1s): + # C would cut "bb\0zz" to "bb"; the code must not silently narrow. + with t.Model(langid_toy_model_path, backend="cpu") as model: + with pytest.raises(t.InvalidArgument): + model.langid_label_index("bb\0zz") + with model.langid_session() as lid: + with pytest.raises(t.InvalidArgument): + lid.run(noise_1s, allowed=["bb\0zz"]) + + +def test_toy_input_rules(langid_toy_model_path, noise_1s): + with t.Model(langid_toy_model_path, backend="cpu") as model: + with model.langid_session() as lid: + with pytest.raises(t.InputTooShort): + lid.run(noise_1s[:6400]) # 400 ms + assert lid.run(noise_1s[:8000]).audio_ms == 500 + assert lid.run(_noise(16000 * 31)).audio_ms == 30000 + bad = array.array("f", noise_1s) + bad[10] = float("nan") + with pytest.raises(t.InvalidArgument): + lid.run(bad) + with pytest.raises(t.InvalidArgument): + model.langid_session(max_audio_ms=400) + with model.langid_session(max_audio_ms=500) as short: + assert short.run(noise_1s).audio_ms == 500 + + +def test_toy_cancel(langid_toy_model_path, noise_1s): + with t.Model(langid_toy_model_path, backend="cpu") as model: + with model.langid_session() as lid: + lid.cancel() + # run() clears a stale cancel before it starts, so it completes. + assert lid.run(noise_1s).candidates + + +@pytest.mark.parametrize("code", ["en", "de", "fr", "es", "ja", "zh", "ru", "id"]) +def test_real_fleurs_top1(langid_model_path, code): + pcm = load_wav(SAMPLES / f"fleurs-{code}.wav") + with t.Model(langid_model_path, backend="cpu") as model: + assert model.langid_info.n_labels == 107 + assert model.langid_label_index("he") == model.langid_label_index("iw") + with model.langid_session() as lid: + r = lid.run(pcm, top_k=3) + assert r.code == code and r.candidates[0].p >= 0.5 + + +def test_real_allowed_mass_signals_out_of_set(langid_model_path): + pcm = load_wav(SAMPLES / "fleurs-ja.wav") + with t.Model(langid_model_path, backend="cpu") as model: + with model.langid_session() as lid: + r = lid.run(pcm, allowed=["en", "de"]) + assert r.code in ("en", "de") and r.allowed_mass < 0.01 diff --git a/bindings/rust/sys/src/transcribe_sys.rs b/bindings/rust/sys/src/transcribe_sys.rs index 2c08ac549..778e5b4b5 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 = bd3273dabb25a1fe +// Pinned to include/transcribe.abihash = d544d70a2b5acf50 /// 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 = "bd3273dabb25a1fe"; +pub const PUBLIC_HEADER_HASH: &str = "d544d70a2b5acf50"; /* automatically generated by rust-bindgen 0.72.1 */ @@ -44,6 +44,7 @@ impl transcribe_status { pub const TRANSCRIBE_ERR_OUTPUT_TRUNCATED: transcribe_status = transcribe_status(18); pub const TRANSCRIBE_ERR_OUTPUT_REPETITION: transcribe_status = transcribe_status(19); pub const TRANSCRIBE_ERR_UNSUPPORTED_ROLE: transcribe_status = transcribe_status(20); + pub const TRANSCRIBE_ERR_INPUT_TOO_SHORT: transcribe_status = transcribe_status(21); } #[repr(transparent)] #[derive(Debug, Copy, Clone, Hash, PartialEq, Eq)] @@ -79,6 +80,12 @@ impl transcribe_abi_struct { pub const TRANSCRIBE_ABI_DIARIZE_SESSION_PARAMS: transcribe_abi_struct = transcribe_abi_struct(17); pub const TRANSCRIBE_ABI_DIARIZE_PARAMS: transcribe_abi_struct = transcribe_abi_struct(18); + pub const TRANSCRIBE_ABI_LANGID_INFO: transcribe_abi_struct = transcribe_abi_struct(19); + pub const TRANSCRIBE_ABI_LANGID_SESSION_PARAMS: transcribe_abi_struct = + transcribe_abi_struct(20); + pub const TRANSCRIBE_ABI_LANGID_PARAMS: transcribe_abi_struct = transcribe_abi_struct(21); + pub const TRANSCRIBE_ABI_LANGID_RESULT: transcribe_abi_struct = transcribe_abi_struct(22); + pub const TRANSCRIBE_ABI_LANGID_CANDIDATE: transcribe_abi_struct = transcribe_abi_struct(23); } #[repr(transparent)] #[derive(Debug, Copy, Clone, Hash, PartialEq, Eq)] @@ -489,6 +496,7 @@ unsafe extern "C" { impl transcribe_role { pub const TRANSCRIBE_ROLE_ASR: transcribe_role = transcribe_role(1); pub const TRANSCRIBE_ROLE_DIARIZE: transcribe_role = transcribe_role(2); + pub const TRANSCRIBE_ROLE_LANGID: transcribe_role = transcribe_role(4); } #[repr(transparent)] #[derive(Debug, Copy, Clone, Hash, PartialEq, Eq)] @@ -1281,6 +1289,214 @@ unsafe extern "C" { } #[repr(C)] #[derive(Debug, Copy, Clone)] +pub struct transcribe_langid_session { + _unused: [u8; 0], +} +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct transcribe_langid_info { + pub struct_size: u64, + pub sample_rate: i32, + pub n_labels: i32, + pub min_audio_ms: i32, +} +#[allow(clippy::unnecessary_operation, clippy::identity_op)] +const _: () = { + ["Size of transcribe_langid_info"][::std::mem::size_of::() - 24usize]; + ["Alignment of transcribe_langid_info"] + [::std::mem::align_of::() - 8usize]; + ["Offset of field: transcribe_langid_info::struct_size"] + [::std::mem::offset_of!(transcribe_langid_info, struct_size) - 0usize]; + ["Offset of field: transcribe_langid_info::sample_rate"] + [::std::mem::offset_of!(transcribe_langid_info, sample_rate) - 8usize]; + ["Offset of field: transcribe_langid_info::n_labels"] + [::std::mem::offset_of!(transcribe_langid_info, n_labels) - 12usize]; + ["Offset of field: transcribe_langid_info::min_audio_ms"] + [::std::mem::offset_of!(transcribe_langid_info, min_audio_ms) - 16usize]; +}; +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct transcribe_langid_session_params { + pub struct_size: u64, + pub n_threads: i32, + pub max_audio_ms: i32, +} +#[allow(clippy::unnecessary_operation, clippy::identity_op)] +const _: () = { + ["Size of transcribe_langid_session_params"] + [::std::mem::size_of::() - 16usize]; + ["Alignment of transcribe_langid_session_params"] + [::std::mem::align_of::() - 8usize]; + ["Offset of field: transcribe_langid_session_params::struct_size"] + [::std::mem::offset_of!(transcribe_langid_session_params, struct_size) - 0usize]; + ["Offset of field: transcribe_langid_session_params::n_threads"] + [::std::mem::offset_of!(transcribe_langid_session_params, n_threads) - 8usize]; + ["Offset of field: transcribe_langid_session_params::max_audio_ms"] + [::std::mem::offset_of!(transcribe_langid_session_params, max_audio_ms) - 12usize]; +}; +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct transcribe_langid_params { + pub struct_size: u64, + pub allowed: *const *const ::std::os::raw::c_char, + pub n_allowed: i32, + pub top_k: i32, +} +#[allow(clippy::unnecessary_operation, clippy::identity_op)] +const _: () = { + ["Size of transcribe_langid_params"] + [::std::mem::size_of::() - 24usize]; + ["Alignment of transcribe_langid_params"] + [::std::mem::align_of::() - 8usize]; + ["Offset of field: transcribe_langid_params::struct_size"] + [::std::mem::offset_of!(transcribe_langid_params, struct_size) - 0usize]; + ["Offset of field: transcribe_langid_params::allowed"] + [::std::mem::offset_of!(transcribe_langid_params, allowed) - 8usize]; + ["Offset of field: transcribe_langid_params::n_allowed"] + [::std::mem::offset_of!(transcribe_langid_params, n_allowed) - 16usize]; + ["Offset of field: transcribe_langid_params::top_k"] + [::std::mem::offset_of!(transcribe_langid_params, top_k) - 20usize]; +}; +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct transcribe_langid_result { + pub struct_size: u64, + pub n_candidates: i32, + pub n_allowed: i32, + pub allowed_mass: f32, + pub audio_ms: i64, +} +#[allow(clippy::unnecessary_operation, clippy::identity_op)] +const _: () = { + ["Size of transcribe_langid_result"] + [::std::mem::size_of::() - 32usize]; + ["Alignment of transcribe_langid_result"] + [::std::mem::align_of::() - 8usize]; + ["Offset of field: transcribe_langid_result::struct_size"] + [::std::mem::offset_of!(transcribe_langid_result, struct_size) - 0usize]; + ["Offset of field: transcribe_langid_result::n_candidates"] + [::std::mem::offset_of!(transcribe_langid_result, n_candidates) - 8usize]; + ["Offset of field: transcribe_langid_result::n_allowed"] + [::std::mem::offset_of!(transcribe_langid_result, n_allowed) - 12usize]; + ["Offset of field: transcribe_langid_result::allowed_mass"] + [::std::mem::offset_of!(transcribe_langid_result, allowed_mass) - 16usize]; + ["Offset of field: transcribe_langid_result::audio_ms"] + [::std::mem::offset_of!(transcribe_langid_result, audio_ms) - 24usize]; +}; +#[repr(C)] +#[derive(Debug, Copy, Clone)] +pub struct transcribe_langid_candidate { + pub struct_size: u64, + pub index: i32, + pub code: *const ::std::os::raw::c_char, + pub name: *const ::std::os::raw::c_char, + pub p: f32, + pub logit: f32, +} +#[allow(clippy::unnecessary_operation, clippy::identity_op)] +const _: () = { + ["Size of transcribe_langid_candidate"] + [::std::mem::size_of::() - 40usize]; + ["Alignment of transcribe_langid_candidate"] + [::std::mem::align_of::() - 8usize]; + ["Offset of field: transcribe_langid_candidate::struct_size"] + [::std::mem::offset_of!(transcribe_langid_candidate, struct_size) - 0usize]; + ["Offset of field: transcribe_langid_candidate::index"] + [::std::mem::offset_of!(transcribe_langid_candidate, index) - 8usize]; + ["Offset of field: transcribe_langid_candidate::code"] + [::std::mem::offset_of!(transcribe_langid_candidate, code) - 16usize]; + ["Offset of field: transcribe_langid_candidate::name"] + [::std::mem::offset_of!(transcribe_langid_candidate, name) - 24usize]; + ["Offset of field: transcribe_langid_candidate::p"] + [::std::mem::offset_of!(transcribe_langid_candidate, p) - 32usize]; + ["Offset of field: transcribe_langid_candidate::logit"] + [::std::mem::offset_of!(transcribe_langid_candidate, logit) - 36usize]; +}; +unsafe extern "C" { + pub fn transcribe_langid_info_init(out: *mut transcribe_langid_info); +} +unsafe extern "C" { + pub fn transcribe_langid_session_params_init(params: *mut transcribe_langid_session_params); +} +unsafe extern "C" { + pub fn transcribe_langid_params_init(params: *mut transcribe_langid_params); +} +unsafe extern "C" { + pub fn transcribe_langid_result_init(out: *mut transcribe_langid_result); +} +unsafe extern "C" { + pub fn transcribe_langid_candidate_init(out: *mut transcribe_langid_candidate); +} +unsafe extern "C" { + pub fn transcribe_langid_get_info( + model: *const transcribe_model, + out: *mut transcribe_langid_info, + ) -> transcribe_status; +} +unsafe extern "C" { + pub fn transcribe_langid_label_code( + model: *const transcribe_model, + i: i32, + ) -> *const ::std::os::raw::c_char; +} +unsafe extern "C" { + pub fn transcribe_langid_label_name( + model: *const transcribe_model, + i: i32, + ) -> *const ::std::os::raw::c_char; +} +unsafe extern "C" { + pub fn transcribe_langid_label_index( + model: *const transcribe_model, + code_or_alias: *const ::std::os::raw::c_char, + ) -> i32; +} +unsafe extern "C" { + pub fn transcribe_langid_session_init( + model: *mut transcribe_model, + params: *const transcribe_langid_session_params, + out_session: *mut *mut transcribe_langid_session, + ) -> transcribe_status; +} +unsafe extern "C" { + pub fn transcribe_langid_session_free(session: *mut transcribe_langid_session); +} +unsafe extern "C" { + pub fn transcribe_langid_set_abort_callback( + session: *mut transcribe_langid_session, + cb: transcribe_abort_callback, + user_data: *mut ::std::os::raw::c_void, + ); +} +unsafe extern "C" { + pub fn transcribe_langid_run( + session: *mut transcribe_langid_session, + pcm: *const f32, + n_samples: ::std::os::raw::c_int, + params: *const transcribe_langid_params, + ) -> transcribe_status; +} +unsafe extern "C" { + pub fn transcribe_langid_get_result( + session: *const transcribe_langid_session, + out: *mut transcribe_langid_result, + ) -> transcribe_status; +} +unsafe extern "C" { + pub fn transcribe_langid_get_candidate( + session: *const transcribe_langid_session, + i: ::std::os::raw::c_int, + out: *mut transcribe_langid_candidate, + ) -> transcribe_status; +} +unsafe extern "C" { + pub fn transcribe_langid_get_timings( + session: *const transcribe_langid_session, + out: *mut transcribe_timings, + ) -> transcribe_status; +} +#[repr(C)] +#[derive(Debug, Copy, Clone)] pub struct transcribe_moonshine_streaming_stream_ext { pub ext: transcribe_ext, pub min_decode_interval_ms: i32, diff --git a/bindings/rust/transcribe-cpp/README.md b/bindings/rust/transcribe-cpp/README.md index dd12b78c7..bfb8fff18 100644 --- a/bindings/rust/transcribe-cpp/README.md +++ b/bindings/rust/transcribe-cpp/README.md @@ -93,6 +93,25 @@ for turn in diarize.run(&pcm, &DiarizeOptions::default())? { # Ok::<(), transcribe_cpp::Error>(()) ``` +### Language ID + +A model whose `roles()` contain `Role::LangId` (VoxLingua107 ECAPA-TDNN) +opens a `LangIdSession`. Codes are the model's own labels; match them against +an ASR model's `capabilities().languages` yourself. `allowed: None` scores +every label, `Some(vec![])` is `Error::InvalidArgument`, and clips under +`langid_info()?.min_audio_ms` are `Error::InputTooShort`. A low +`allowed_mass` means the speech is probably outside the allowed set. + +```rust +use transcribe_cpp::{LangIdOptions, Model}; +let model = Model::load("lang-id-voxlingua107-ecapa-Q8_0.gguf")?; +let mut lid = model.langid_session()?; +let opts = LangIdOptions { allowed: Some(vec!["en".into(), "de".into()]), top_k: 3 }; +let result = lid.run(&pcm, &opts)?; +println!("{:?} (mass {})", result.code(), result.allowed_mass); +# Ok::<(), transcribe_cpp::Error>(()) +``` + Runnable examples: ```sh diff --git a/bindings/rust/transcribe-cpp/src/error.rs b/bindings/rust/transcribe-cpp/src/error.rs index 40577cc57..8ca913731 100644 --- a/bindings/rust/transcribe-cpp/src/error.rs +++ b/bindings/rust/transcribe-cpp/src/error.rs @@ -78,6 +78,10 @@ pub enum Error { /// `TRANSCRIBE_ERR_UNSUPPORTED_ROLE` — see [`Model::roles`](crate::Model::roles). #[error("unsupported role: {0}")] UnsupportedRole(String), + /// `TRANSCRIBE_ERR_INPUT_TOO_SHORT` — the audio is shorter than the role's + /// minimum (e.g. [`LangIdInfo::min_audio_ms`](crate::LangIdInfo)). + #[error("input too short: {0}")] + InputTooShort(String), /// The loaded library's base version disagrees with the headers this crate /// was generated against (the pre-1.0 version lock). Raised on first use. #[error("native library version mismatch: {0}")] @@ -127,6 +131,7 @@ pub enum ErrorKind { VersionMismatch, Nul, Busy, + InputTooShort, #[cfg_attr(feature = "serde", serde(other))] Other, } @@ -231,6 +236,7 @@ impl Error { Error::OutputTruncated { .. } => ErrorKind::OutputTruncated, Error::OutputRepetition { .. } => ErrorKind::OutputRepetition, Error::UnsupportedRole(_) => ErrorKind::UnsupportedRole, + Error::InputTooShort(_) => ErrorKind::InputTooShort, Error::VersionMismatch(_) => ErrorKind::VersionMismatch, Error::Nul(_) => ErrorKind::Nul, Error::Busy(_) => ErrorKind::Busy, @@ -256,6 +262,7 @@ impl Error { Error::OutputTruncated { .. } => S::TRANSCRIBE_ERR_OUTPUT_TRUNCATED, Error::OutputRepetition { .. } => S::TRANSCRIBE_ERR_OUTPUT_REPETITION, Error::UnsupportedRole(_) => S::TRANSCRIBE_ERR_UNSUPPORTED_ROLE, + Error::InputTooShort(_) => S::TRANSCRIBE_ERR_INPUT_TOO_SHORT, _ => S::TRANSCRIBE_OK, }; s.0 as i32 @@ -330,6 +337,7 @@ pub(crate) fn error_for_status(status: sys::transcribe_status, context: &str) -> partial: None, }, S::TRANSCRIBE_ERR_UNSUPPORTED_ROLE => Error::UnsupportedRole(msg), + S::TRANSCRIBE_ERR_INPUT_TOO_SHORT => Error::InputTooShort(msg), _ => Error::Other(msg), } } diff --git a/bindings/rust/transcribe-cpp/src/langid.rs b/bindings/rust/transcribe-cpp/src/langid.rs new file mode 100644 index 000000000..b9d513229 --- /dev/null +++ b/bindings/rust/transcribe-cpp/src/langid.rs @@ -0,0 +1,303 @@ +//! [`LangIdSession`] — the LANGID role (which language is spoken), from +//! [`Model::langid_session`]. Threading, lifetime, and compute-lock rules are +//! those of [`Session`](crate::Session). + +use std::ffi::CString; +use std::os::raw::{c_char, c_void}; +use std::sync::atomic::AtomicBool; +use std::sync::Arc; + +use transcribe_cpp_sys as sys; + +use crate::cancel::{abort_trampoline, CancelToken}; +use crate::error::{check, Error, Result}; +use crate::model::{Model, ModelInner}; +use crate::result::{owned_str, Timings}; +use crate::session::clamp_len; + +/// Static facts about a language ID model ([`Model::langid_info`]). +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[non_exhaustive] +#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))] +pub struct LangIdInfo { + /// Input PCM rate (16000). + pub sample_rate: i32, + /// Label indices are `[0, n_labels)`. + pub n_labels: i32, + /// Shorter scored audio is [`Error::InputTooShort`]. + pub min_audio_ms: i32, +} + +/// Options for creating a language ID session. +#[derive(Debug, Clone, Default, PartialEq, Eq)] +#[cfg_attr( + feature = "serde", + derive(serde::Serialize, serde::Deserialize), + serde(default) +)] +pub struct LangIdSessionOptions { + /// CPU threads for ops that run on CPU; 0 = library default. + pub n_threads: i32, + /// Longer input is scored on its last `max_audio_ms`; 0 = 30000. Below + /// [`LangIdInfo::min_audio_ms`] is [`Error::InvalidArgument`]. + pub max_audio_ms: i32, +} + +/// Per-run language ID parameters. +#[derive(Debug, Clone, Default, PartialEq, Eq)] +#[cfg_attr( + feature = "serde", + derive(serde::Serialize, serde::Deserialize), + serde(default) +)] +pub struct LangIdOptions { + /// Restrict the decision to these codes or aliases. `None` means every + /// label; `Some` of an empty list is [`Error::InvalidArgument`]. An + /// unknown code is [`Error::Unsupported`]. + pub allowed: Option>, + /// Keep the best `top_k` candidates; 0 = every allowed label. + pub top_k: i32, +} + +/// One ranked label. `code` is the model's own label (`"en"`, `"iw"`). +#[derive(Debug, Clone, Default, PartialEq)] +#[non_exhaustive] +#[cfg_attr( + feature = "serde", + derive(serde::Serialize, serde::Deserialize), + serde(default) +)] +pub struct LangIdCandidate { + /// Label index. + pub index: i32, + pub code: String, + pub name: String, + /// Softmax renormalized over the allowed set. + pub p: f32, + pub logit: f32, +} + +/// The result of one [`LangIdSession::run`]. +#[derive(Debug, Clone, Default, PartialEq)] +#[non_exhaustive] +#[cfg_attr( + feature = "serde", + derive(serde::Serialize, serde::Deserialize), + serde(default) +)] +pub struct LangIdResult { + /// Ranked by `p`, descending; ties keep label order. + pub candidates: Vec, + /// Labels in the allowed set (before `top_k`). + pub n_allowed: i32, + /// Unrestricted probability inside the allowed set (1.0 when + /// unrestricted); a low value means the speech is probably outside it. + pub allowed_mass: f32, + /// Audio actually scored, after the crop. + pub audio_ms: i64, +} + +impl LangIdResult { + /// The top candidate's code, if any. + pub fn code(&self) -> Option<&str> { + self.candidates.first().map(|c| c.code.as_str()) + } +} + +/// A language ID session. +pub struct LangIdSession { + ptr: *mut sys::transcribe_langid_session, + // Keeps the native model alive and carries the per-model compute lock. + model: Arc, + // Keeps the abort callback's userdata alive while installed. + cancel: Option>, +} + +impl std::fmt::Debug for LangIdSession { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("LangIdSession").finish_non_exhaustive() + } +} + +// SAFETY: as for `Session` — `&mut self` on the mutating calls keeps use to +// one thread at a time; deliberately NOT Sync. +unsafe impl Send for LangIdSession {} + +impl Drop for LangIdSession { + fn drop(&mut self) { + unsafe { sys::transcribe_langid_session_free(self.ptr) }; + } +} + +impl LangIdSession { + pub(crate) fn new(model: &Model, options: &LangIdSessionOptions) -> Result { + let mut params: sys::transcribe_langid_session_params = unsafe { std::mem::zeroed() }; + unsafe { sys::transcribe_langid_session_params_init(&mut params) }; + params.n_threads = options.n_threads; + params.max_audio_ms = options.max_audio_ms; + + let mut out: *mut sys::transcribe_langid_session = std::ptr::null_mut(); + let status = + unsafe { sys::transcribe_langid_session_init(model.inner.ptr, ¶ms, &mut out) }; + check(status, "langid session init")?; + debug_assert!(!out.is_null()); + + Ok(LangIdSession { + ptr: out, + model: Arc::clone(&model.inner), + cancel: None, + }) + } + + /// Install a [`CancelToken`] so an in-flight run can be aborted from + /// another thread (the run then returns [`Error::Aborted`]). + /// Replaces any previously installed token. + pub fn set_cancel_token(&mut self, token: &CancelToken) { + let flag = Arc::clone(&token.flag); + let userdata = Arc::as_ptr(&flag) as *mut c_void; + unsafe { + sys::transcribe_langid_set_abort_callback(self.ptr, Some(abort_trampoline), userdata) + }; + self.cancel = Some(flag); + } + + /// Remove any installed cancel token. + pub fn clear_cancel_token(&mut self) { + unsafe { sys::transcribe_langid_set_abort_callback(self.ptr, None, std::ptr::null_mut()) }; + self.cancel = None; + } + + /// Identify the language of one clip of 16 kHz mono float32 PCM. Input + /// longer than the session's `max_audio_ms` is scored on its tail. + pub fn run(&mut self, pcm: &[f32], options: &LangIdOptions) -> Result { + let n = clamp_len(pcm.len())?; + // The CStrings and the pointer array stay alive until the end of this + // call, past the native run. + let codes: Option> = match &options.allowed { + None => None, + Some(list) if list.is_empty() => { + // NULL would mean "every label", the opposite of an empty list. + return Err(Error::InvalidArgument( + "allowed is empty; pass None for every label".into(), + )); + } + Some(list) => Some( + list.iter() + .map(|c| CString::new(c.as_str())) + .collect::>()?, + ), + }; + let ptrs: Option> = codes + .as_ref() + .map(|cs| cs.iter().map(|c| c.as_ptr()).collect()); + let mut params: sys::transcribe_langid_params = unsafe { std::mem::zeroed() }; + unsafe { sys::transcribe_langid_params_init(&mut params) }; + if let Some(p) = &ptrs { + params.allowed = p.as_ptr(); + params.n_allowed = i32::try_from(p.len()) + .map_err(|_| Error::InvalidArgument("allowed list too long".into()))?; + } + params.top_k = options.top_k; + + // Results are copied out under the compute lock, like every other + // binding, so a concurrent run on the same model cannot interleave. + let ptr = self.ptr; + self.model.with_compute( + Some("a stream is active on this model; finish or drop it before langid run()"), + |_| -> Result { + check( + unsafe { sys::transcribe_langid_run(ptr, pcm.as_ptr(), n, ¶ms) }, + "langid run", + )?; + let mut res: sys::transcribe_langid_result = unsafe { std::mem::zeroed() }; + unsafe { sys::transcribe_langid_result_init(&mut res) }; + check( + unsafe { sys::transcribe_langid_get_result(ptr, &mut res) }, + "langid result", + )?; + let candidates = (0..res.n_candidates) + .map(|i| { + let mut raw: sys::transcribe_langid_candidate = + unsafe { std::mem::zeroed() }; + unsafe { sys::transcribe_langid_candidate_init(&mut raw) }; + let _ = unsafe { sys::transcribe_langid_get_candidate(ptr, i, &mut raw) }; + LangIdCandidate { + index: raw.index, + code: owned_str(raw.code), + name: owned_str(raw.name), + p: raw.p, + logit: raw.logit, + } + }) + .collect(); + Ok(LangIdResult { + candidates, + n_allowed: res.n_allowed, + allowed_mass: res.allowed_mass, + audio_ms: res.audio_ms, + }) + }, + )? + } + + /// Model load time plus the last run's stage timings (`decode_ms` is 0). + pub fn timings(&self) -> Timings { + let mut raw: sys::transcribe_timings = unsafe { std::mem::zeroed() }; + unsafe { sys::transcribe_timings_init(&mut raw) }; + let _ = unsafe { sys::transcribe_langid_get_timings(self.ptr, &mut raw) }; + Timings::from_raw(&raw) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::Mutex; + + fn null_session(streaming: bool) -> LangIdSession { + // Null native handles (both frees are NULL no-ops): every check below + // fires before any native call. + LangIdSession { + ptr: std::ptr::null_mut(), + model: Arc::new(ModelInner { + ptr: std::ptr::null_mut(), + compute_lock: Mutex::new(streaming), + }), + cancel: None, + } + } + + #[test] + fn run_is_busy_while_a_stream_holds_the_lease() { + let err = null_session(true) + .run(&[0.0; 160], &LangIdOptions::default()) + .unwrap_err(); + let Error::Busy(msg) = err else { + panic!("expected Busy, got {err:?}"); + }; + assert_eq!( + msg, + "a stream is active on this model; finish or drop it before langid run()" + ); + } + + #[test] + fn empty_allowed_is_rejected_not_all() { + let opts = LangIdOptions { + allowed: Some(vec![]), + top_k: 0, + }; + let err = null_session(false).run(&[0.0; 160], &opts).unwrap_err(); + assert!(matches!(err, Error::InvalidArgument(_)), "{err:?}"); + } + + #[test] + fn interior_nul_in_allowed_is_rejected() { + let opts = LangIdOptions { + allowed: Some(vec!["e\0n".into()]), + top_k: 0, + }; + let err = null_session(false).run(&[0.0; 160], &opts).unwrap_err(); + assert!(matches!(err, Error::Nul(_)), "{err:?}"); + } +} diff --git a/bindings/rust/transcribe-cpp/src/lib.rs b/bindings/rust/transcribe-cpp/src/lib.rs index 2f2913757..403924964 100644 --- a/bindings/rust/transcribe-cpp/src/lib.rs +++ b/bindings/rust/transcribe-cpp/src/lib.rs @@ -51,6 +51,7 @@ mod cancel; mod diarize; mod error; mod family; +mod langid; mod logging; mod model; mod result; @@ -71,6 +72,9 @@ pub use family::{ ParakeetStreamOptions, RunExtension, SortformerDiarizeOptions, SortformerPreset, StreamExtension, VoxtralRealtimeStreamOptions, WhisperRunOptions, }; +pub use langid::{ + LangIdCandidate, LangIdInfo, LangIdOptions, LangIdResult, LangIdSession, LangIdSessionOptions, +}; pub use logging::{disable_logging, init_logging}; pub use model::{Capabilities, Model, ModelOptions, SessionLimits, SessionOptions}; pub use result::{Segment, SpeakerSegment, Timings, Token, Transcript, Word}; diff --git a/bindings/rust/transcribe-cpp/src/model.rs b/bindings/rust/transcribe-cpp/src/model.rs index 615f28506..fd902900b 100644 --- a/bindings/rust/transcribe-cpp/src/model.rs +++ b/bindings/rust/transcribe-cpp/src/model.rs @@ -25,6 +25,7 @@ use transcribe_cpp_sys as sys; use crate::backend::Device; use crate::diarize::{DiarizeInfo, DiarizeSession, DiarizeSessionOptions}; use crate::error::{check, Error, Result}; +use crate::langid::{LangIdInfo, LangIdSession, LangIdSessionOptions}; use crate::result::owned_str; use crate::session::Session; use crate::types::{Backend, ExtSlot, Feature, Roles, TimestampKind}; @@ -200,6 +201,55 @@ impl Model { }) } + /// Open a language ID session with default options. Errors with + /// [`Error::UnsupportedRole`] unless the model serves [`Role::LangId`](crate::Role). + pub fn langid_session(&self) -> Result { + self.langid_session_with(&LangIdSessionOptions::default()) + } + + /// Open a language ID session with explicit options. + pub fn langid_session_with(&self, options: &LangIdSessionOptions) -> Result { + LangIdSession::new(self, options) + } + + /// Static facts about a language ID model. Errors with + /// [`Error::UnsupportedRole`] unless the model serves [`Role::LangId`](crate::Role). + pub fn langid_info(&self) -> Result { + let mut raw: sys::transcribe_langid_info = unsafe { std::mem::zeroed() }; + unsafe { sys::transcribe_langid_info_init(&mut raw) }; + check( + unsafe { sys::transcribe_langid_get_info(self.inner.ptr, &mut raw) }, + "langid info", + )?; + Ok(LangIdInfo { + sample_rate: raw.sample_rate, + n_labels: raw.n_labels, + min_audio_ms: raw.min_audio_ms, + }) + } + + /// `(code, name)` per label index. Errors with [`Error::UnsupportedRole`] + /// unless the model serves [`Role::LangId`](crate::Role). + pub fn langid_labels(&self) -> Result> { + let n = self.langid_info()?.n_labels; + Ok((0..n) + .map(|i| unsafe { + ( + owned_str(sys::transcribe_langid_label_code(self.inner.ptr, i)), + owned_str(sys::transcribe_langid_label_name(self.inner.ptr, i)), + ) + }) + .collect()) + } + + /// Label index of a code or alias (`"he"` and `"iw"` name the same label); + /// `None` when unknown or the model does not serve the LANGID role. + pub fn langid_label_index(&self, code: &str) -> Option { + let c = CString::new(code).ok()?; + let i = unsafe { sys::transcribe_langid_label_index(self.inner.ptr, c.as_ptr()) }; + (i >= 0).then_some(i) + } + /// The roles (kinds of work) this model serves. pub fn roles(&self) -> Roles { Roles(unsafe { sys::transcribe_model_roles(self.inner.ptr) }) diff --git a/bindings/rust/transcribe-cpp/src/types.rs b/bindings/rust/transcribe-cpp/src/types.rs index eb1554eda..6210863dd 100644 --- a/bindings/rust/transcribe-cpp/src/types.rs +++ b/bindings/rust/transcribe-cpp/src/types.rs @@ -324,6 +324,11 @@ pub enum AbiStruct { DiarizeInfo, DiarizeSessionParams, DiarizeParams, + LangIdInfo, + LangIdSessionParams, + LangIdParams, + LangIdResult, + LangIdCandidate, } impl AbiStruct { @@ -349,6 +354,11 @@ impl AbiStruct { AbiStruct::DiarizeInfo => A::TRANSCRIBE_ABI_DIARIZE_INFO, AbiStruct::DiarizeSessionParams => A::TRANSCRIBE_ABI_DIARIZE_SESSION_PARAMS, AbiStruct::DiarizeParams => A::TRANSCRIBE_ABI_DIARIZE_PARAMS, + AbiStruct::LangIdInfo => A::TRANSCRIBE_ABI_LANGID_INFO, + AbiStruct::LangIdSessionParams => A::TRANSCRIBE_ABI_LANGID_SESSION_PARAMS, + AbiStruct::LangIdParams => A::TRANSCRIBE_ABI_LANGID_PARAMS, + AbiStruct::LangIdResult => A::TRANSCRIBE_ABI_LANGID_RESULT, + AbiStruct::LangIdCandidate => A::TRANSCRIBE_ABI_LANGID_CANDIDATE, } } } @@ -386,6 +396,8 @@ pub enum Role { Asr, /// Speaker diarization: [`DiarizeSession`](crate::DiarizeSession). Diarize, + /// Language identification: [`LangIdSession`](crate::LangIdSession). + LangId, } /// The set of [`Role`]s a model serves ([`Model::roles`](crate::Model::roles)). @@ -399,6 +411,7 @@ impl Roles { let bit = match role { Role::Asr => sys::transcribe_role::TRANSCRIBE_ROLE_ASR, Role::Diarize => sys::transcribe_role::TRANSCRIBE_ROLE_DIARIZE, + Role::LangId => sys::transcribe_role::TRANSCRIBE_ROLE_LANGID, }; self.0 & bit.0 != 0 } diff --git a/bindings/rust/transcribe-cpp/tests/common/mod.rs b/bindings/rust/transcribe-cpp/tests/common/mod.rs index 6e66a8a12..4ea467d4c 100644 --- a/bindings/rust/transcribe-cpp/tests/common/mod.rs +++ b/bindings/rust/transcribe-cpp/tests/common/mod.rs @@ -128,6 +128,34 @@ pub fn smoke_sortformer_fixtures(test: &str) -> Option<(PathBuf, Vec)> { } } +/// VoxLingua107 ECAPA-TDNN (LANGID role only), or `None` (with a skip note). +pub fn smoke_langid_model(test: &str) -> Option { + let model = family_model( + "TRANSCRIBE_SMOKE_LANGID_MODEL", + "models/lang-id-voxlingua107-ecapa/lang-id-voxlingua107-ecapa-Q8_0.gguf", + ); + if model.is_none() { + eprintln!("skip {test}: langid model absent (set TRANSCRIBE_SMOKE_LANGID_MODEL)"); + } + model +} + +/// The toy ecapa_tdnn GGUF the C++ build generates under tests/fixtures/ +/// (5 labels aa..ee, alias xx=aa, random weights), or `None` (with a skip +/// note) before the C++ test fixtures have been built. +pub fn langid_toy_model(test: &str) -> Option { + ensure_backends(); + let path = repo_root().join("tests/fixtures/arch_ecapa_tdnn_minimal.gguf"); + if !path.is_file() { + eprintln!( + "skip {test}: {} absent (build the C++ `fixtures` target)", + path.display() + ); + return None; + } + Some(path) +} + /// Both fixtures together; prints a skip note and returns `None` if either is /// missing (so the caller can `return` early — the Rust equivalent of skip). pub fn smoke_fixtures(test: &str) -> Option<(PathBuf, Vec)> { @@ -142,7 +170,7 @@ pub fn smoke_fixtures(test: &str) -> Option<(PathBuf, Vec)> { } } -fn load_wav(path: &std::path::Path) -> Vec { +pub fn load_wav(path: &std::path::Path) -> Vec { let mut reader = hound::WavReader::open(path).expect("open wav"); let spec = reader.spec(); assert_eq!(spec.channels, 1, "{path:?} must be mono"); diff --git a/bindings/rust/transcribe-cpp/tests/langid.rs b/bindings/rust/transcribe-cpp/tests/langid.rs new file mode 100644 index 000000000..e474a8835 --- /dev/null +++ b/bindings/rust/transcribe-cpp/tests/langid.rs @@ -0,0 +1,156 @@ +//! LANGID role: contract checks on the toy ecapa_tdnn fixture +//! (tests/fixtures/arch_ecapa_tdnn_minimal.gguf, built by the C++ test +//! fixtures) and top-1 on the real VoxLingua107 model +//! (TRANSCRIBE_SMOKE_LANGID_MODEL). Busy / empty-allowed are covered in +//! `src/langid.rs`. + +mod common; + +use transcribe_cpp::{ + Backend, CancelToken, Error, LangIdOptions, LangIdSessionOptions, Model, ModelOptions, Role, +}; + +fn noise(n: usize, seed: u32) -> Vec { + let mut s = seed | 1; + (0..n) + .map(|_| { + s = s.wrapping_mul(1664525).wrapping_add(1013904223); + ((s >> 8) & 0xFF_FFFF) as f32 / 16_777_216.0 - 0.5 + }) + .collect() +} + +fn load_cpu(path: &std::path::Path) -> Model { + let opts = ModelOptions { + backend: Backend::Cpu, + ..Default::default() + }; + Model::load_with(path, &opts).unwrap() +} + +#[test] +fn toy_roles_info_labels() { + let Some(path) = common::langid_toy_model("toy_roles_info_labels") else { + return; + }; + let model = load_cpu(&path); + let roles = model.roles(); + assert!(roles.contains(Role::LangId) && !roles.contains(Role::Asr)); + let info = model.langid_info().unwrap(); + assert_eq!( + (info.sample_rate, info.n_labels, info.min_audio_ms), + (16000, 5, 500) + ); + assert_eq!( + model.langid_labels().unwrap()[2], + ("cc".into(), "Charlie".into()) + ); + assert_eq!(model.langid_label_index("xx"), Some(0)); + assert_eq!(model.langid_label_index("zz"), None); + assert!(matches!( + model.capabilities(), + Err(Error::UnsupportedRole(_)) + )); + assert!(matches!(model.session(), Err(Error::UnsupportedRole(_)))); +} + +#[test] +fn toy_run_contract() { + let Some(path) = common::langid_toy_model("toy_run_contract") else { + return; + }; + let model = load_cpu(&path); + let mut lid = model.langid_session().unwrap(); + let pcm = noise(16000, 7); + + let r = lid.run(&pcm, &LangIdOptions::default()).unwrap(); + assert_eq!((r.candidates.len(), r.n_allowed, r.audio_ms), (5, 5, 1000)); + assert_eq!(r.allowed_mass, 1.0); + assert_eq!(r.code(), Some(r.candidates[0].code.as_str())); + let sum: f32 = r.candidates.iter().map(|c| c.p).sum(); + assert!((sum - 1.0).abs() < 1e-5); + assert!(r.candidates.windows(2).all(|w| w[0].p >= w[1].p)); + + let opts = LangIdOptions { + allowed: Some(vec!["bb".into(), "dd".into()]), + top_k: 0, + }; + let r = lid.run(&pcm, &opts).unwrap(); + assert_eq!(r.n_allowed, 2); + assert!(r + .candidates + .iter() + .all(|c| c.code == "bb" || c.code == "dd")); + assert!(r.allowed_mass > 0.0 && r.allowed_mass < 1.0); + + let opts = LangIdOptions { + allowed: Some(vec!["zz".into()]), + top_k: 0, + }; + assert!(matches!(lid.run(&pcm, &opts), Err(Error::Unsupported(_)))); + assert!(matches!( + lid.run(&pcm[..6400], &LangIdOptions::default()), + Err(Error::InputTooShort(_)) + )); + let top = LangIdOptions { + allowed: None, + top_k: 2, + }; + assert_eq!(lid.run(&pcm, &top).unwrap().candidates.len(), 2); + assert!(lid.timings().encode_ms > 0.0); + + let bad = LangIdSessionOptions { + n_threads: 0, + max_audio_ms: 400, + }; + assert!(matches!( + model.langid_session_with(&bad), + Err(Error::InvalidArgument(_)) + )); +} + +#[test] +fn toy_cancel_aborts() { + let Some(path) = common::langid_toy_model("toy_cancel_aborts") else { + return; + }; + let model = load_cpu(&path); + let mut lid = model.langid_session().unwrap(); + let token = CancelToken::new(); + lid.set_cancel_token(&token); + token.cancel(); + assert!(matches!( + lid.run(&noise(16000, 3), &LangIdOptions::default()), + Err(Error::Aborted { .. }) + )); + lid.clear_cancel_token(); + assert!(lid.run(&noise(16000, 3), &LangIdOptions::default()).is_ok()); +} + +#[test] +fn real_fleurs_top1() { + let Some(path) = common::smoke_langid_model("real_fleurs_top1") else { + return; + }; + let model = load_cpu(&path); + assert_eq!(model.langid_info().unwrap().n_labels, 107); + assert_eq!( + model.langid_label_index("he"), + model.langid_label_index("iw") + ); + let mut lid = model.langid_session().unwrap(); + for code in ["en", "de", "fr", "es", "ja", "zh", "ru", "id"] { + let pcm = common::load_wav(&common::repo_root().join(format!("samples/fleurs-{code}.wav"))); + let r = lid + .run( + &pcm, + &LangIdOptions { + allowed: None, + top_k: 3, + }, + ) + .unwrap(); + assert_eq!(r.code(), Some(code), "{code}: {:?}", r.candidates); + assert!(r.candidates[0].p >= 0.5); + } +} diff --git a/bindings/rust/transcribe-cpp/tests/no_model.rs b/bindings/rust/transcribe-cpp/tests/no_model.rs index ed7df746d..5d3985554 100644 --- a/bindings/rust/transcribe-cpp/tests/no_model.rs +++ b/bindings/rust/transcribe-cpp/tests/no_model.rs @@ -38,6 +38,11 @@ fn abi_struct_sizes_are_live() { AbiStruct::DiarizeInfo, AbiStruct::DiarizeSessionParams, AbiStruct::DiarizeParams, + AbiStruct::LangIdInfo, + AbiStruct::LangIdSessionParams, + AbiStruct::LangIdParams, + AbiStruct::LangIdResult, + AbiStruct::LangIdCandidate, ] { assert!(abi_struct_size(which) > 0, "{which:?} reported size 0"); } @@ -142,5 +147,6 @@ fn handles_are_send_sync() { assert_send_sync::(); assert_send::(); assert_send::(); + assert_send::(); // Sessions are intentionally NOT Sync (single-threaded use). } diff --git a/bindings/rust/transcribe-cpp/tests/serde.rs b/bindings/rust/transcribe-cpp/tests/serde.rs index a9592bc6e..5287f5475 100644 --- a/bindings/rust/transcribe-cpp/tests/serde.rs +++ b/bindings/rust/transcribe-cpp/tests/serde.rs @@ -5,12 +5,13 @@ use transcribe_cpp::{ Backend, Capabilities, CommitPolicy, DeviceType, Diarize, DiarizeExtension, DiarizeInfo, - DiarizeOptions, DiarizeSessionOptions, ExtSlot, Feature, Itn, KvType, - MoonshineStreamingOptions, ParakeetBufferedStreamOptions, ParakeetStreamOptions, Pnc, Role, - Roles, RunExtension, RunOptions, Segment, SessionLimits, SessionOptions, - SortformerDiarizeOptions, SortformerPreset, SpeakerSegment, StreamExtension, StreamOptions, - StreamState, StreamText, StreamUpdate, Task, TimestampKind, Timings, Token, Transcript, - VoxtralRealtimeStreamOptions, WhisperRunOptions, Word, + DiarizeOptions, DiarizeSessionOptions, ExtSlot, Feature, Itn, KvType, LangIdCandidate, + LangIdInfo, LangIdOptions, LangIdResult, LangIdSessionOptions, MoonshineStreamingOptions, + ParakeetBufferedStreamOptions, ParakeetStreamOptions, Pnc, Role, Roles, RunExtension, + RunOptions, Segment, SessionLimits, SessionOptions, SortformerDiarizeOptions, SortformerPreset, + SpeakerSegment, StreamExtension, StreamOptions, StreamState, StreamText, StreamUpdate, Task, + TimestampKind, Timings, Token, Transcript, VoxtralRealtimeStreamOptions, WhisperRunOptions, + Word, }; fn assert_serde() {} @@ -34,8 +35,11 @@ fn plain_data_types_are_serializable() { assert_serde::(); assert_serde::(); assert_serde::(); + assert_serde::(); + assert_serde::(); // Results. assert_serde::(); + assert_serde::(); assert_serde::(); assert_serde::(); assert_serde::(); @@ -46,6 +50,8 @@ fn plain_data_types_are_serializable() { assert_serde::(); assert_serde::(); assert_serde::(); + assert_serde::(); + assert_serde::(); // Enums. assert_serde::(); assert_serde::(); @@ -135,6 +141,13 @@ fn missing_fields_take_defaults() { assert!(token.p.is_nan(), "missing p decoded as {}", token.p); let speaker: SpeakerSegment = serde_json::from_str(r#"{"speaker_id":2}"#).unwrap(); assert!(speaker.p.is_nan(), "missing p decoded as {}", speaker.p); + + let langid: LangIdResult = serde_json::from_str(r#"{"n_allowed":3}"#).unwrap(); + assert_eq!(langid.n_allowed, 3); + assert!(langid.candidates.is_empty()); + let candidate: LangIdCandidate = serde_json::from_str(r#"{"code":"en"}"#).unwrap(); + assert_eq!(candidate.code, "en"); + assert_eq!(candidate.p, 0.0); } mod errors { @@ -218,6 +231,7 @@ mod errors { ErrorKind::VersionMismatch, ErrorKind::Nul, ErrorKind::Busy, + ErrorKind::InputTooShort, ErrorKind::Other, ]; for kind in kinds { diff --git a/bindings/swift/README.md b/bindings/swift/README.md index 8a5bed710..328228a56 100644 --- a/bindings/swift/README.md +++ b/bindings/swift/README.md @@ -115,6 +115,20 @@ for t in turns { print(t.speakerId, t.t0Ms, t.t1Ms) } `capabilities` and `session()` throw `.unsupportedRole` on a model without `.asr`. +## Language ID + +Language ID models (`.langId`, e.g. VoxLingua107 ECAPA-TDNN) rank the model's +own label codes from a `LangIdSession`; match `result.code` against an ASR +model's `capabilities.languages` yourself. `allowed: nil` scores every label, +an empty array throws `.invalidArgument`, and clips under +`langIdInfo.minAudioMs` throw `.inputTooShort`. + +```swift +let model = try Model(path: "lang-id-voxlingua107-ecapa-Q8_0.gguf") +let result = try model.langIdSession().run(pcm, options: LangIdOptions(allowed: ["en", "de"], topK: 3)) +print(result.code ?? "-", result.allowedMass) +``` + ## Backends Backends are compiled into the xcframework per Apple slice: diff --git a/bindings/swift/Sources/TranscribeCpp/ABIHash.swift b/bindings/swift/Sources/TranscribeCpp/ABIHash.swift index 5cae2c4dd..c5f9e2e00 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 = "bd3273dabb25a1fe" + public static let pinnedHeaderHash = "d544d70a2b5acf50" /// 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/Backend.swift b/bindings/swift/Sources/TranscribeCpp/Backend.swift index 9eab54f63..837d037ff 100644 --- a/bindings/swift/Sources/TranscribeCpp/Backend.swift +++ b/bindings/swift/Sources/TranscribeCpp/Backend.swift @@ -113,6 +113,11 @@ public enum AbiStruct: Sendable { case diarizeInfo case diarizeSessionParams case diarizeParams + case langIdInfo + case langIdSessionParams + case langIdParams + case langIdResult + case langIdCandidate var cValue: transcribe_abi_struct { switch self { @@ -133,6 +138,11 @@ public enum AbiStruct: Sendable { case .diarizeInfo: return TRANSCRIBE_ABI_DIARIZE_INFO case .diarizeSessionParams: return TRANSCRIBE_ABI_DIARIZE_SESSION_PARAMS case .diarizeParams: return TRANSCRIBE_ABI_DIARIZE_PARAMS + case .langIdInfo: return TRANSCRIBE_ABI_LANGID_INFO + case .langIdSessionParams: return TRANSCRIBE_ABI_LANGID_SESSION_PARAMS + case .langIdParams: return TRANSCRIBE_ABI_LANGID_PARAMS + case .langIdResult: return TRANSCRIBE_ABI_LANGID_RESULT + case .langIdCandidate: return TRANSCRIBE_ABI_LANGID_CANDIDATE } } } diff --git a/bindings/swift/Sources/TranscribeCpp/Cancellation.swift b/bindings/swift/Sources/TranscribeCpp/Cancellation.swift index fe5221932..7142c3685 100644 --- a/bindings/swift/Sources/TranscribeCpp/Cancellation.swift +++ b/bindings/swift/Sources/TranscribeCpp/Cancellation.swift @@ -39,6 +39,21 @@ extension Session { } } +extension LangIdSession { + /// Install a cancellation token; a cancelled run throws `.aborted`. + public func setCancellationToken(_ token: CancellationToken) { + cancelToken = token + let context = Unmanaged.passUnretained(token).toOpaque() + transcribe_langid_set_abort_callback(ptr, abortTrampoline, context) + } + + /// Remove any installed cancellation token. + public func clearCancellationToken() { + transcribe_langid_set_abort_callback(ptr, nil, nil) + cancelToken = nil + } +} + extension DiarizeSession { /// Install a cancellation token; a cancelled run throws `.aborted`. public func setCancellationToken(_ token: CancellationToken) { diff --git a/bindings/swift/Sources/TranscribeCpp/LangId.swift b/bindings/swift/Sources/TranscribeCpp/LangId.swift new file mode 100644 index 000000000..ea30b9cb5 --- /dev/null +++ b/bindings/swift/Sources/TranscribeCpp/LangId.swift @@ -0,0 +1,190 @@ +import CTranscribe +import Foundation + +/// Static facts about a language ID model. +public struct LangIdInfo: Sendable, Equatable { + /// Input PCM rate (16000). + public let sampleRate: Int32 + /// Label indices are `0.. Int32? { + if code.contains("\0") { return nil } + let i = transcribe_langid_label_index(ptr, code) + return i >= 0 ? i : nil + } + + /// Create a LANGID session (`threads` 0 = library default). Longer input + /// than `maxAudioMs` is scored on its tail (0 = 30000). Throws + /// `.unsupportedRole` when `roles` lacks `.langId`. + public func langIdSession(threads: Int32 = 0, maxAudioMs: Int32 = 0) throws -> LangIdSession { + var params = transcribe_langid_session_params() + transcribe_langid_session_params_init(¶ms) + params.n_threads = threads + params.max_audio_ms = maxAudioMs + var out: OpaquePointer? + try TranscribeError.check( + transcribe_langid_session_init(ptr, ¶ms, &out), context: "creating langid session") + guard let out else { + throw TranscribeError.other(status: 0, message: "null langid session handle") + } + return LangIdSession(model: self, ptr: out) + } +} + +/// A LANGID-role session: which language is spoken. Same threading and +/// lifetime contract as `Session` (single-threaded; holds its `Model` alive), +/// and its runs share the model's compute lock and stream lease with every +/// `Session`. +public final class LangIdSession { + let model: Model + let ptr: OpaquePointer + /// Strong ref to the installed cancellation token (see `Session`). + var cancelToken: CancellationToken? + + init(model: Model, ptr: OpaquePointer) { + self.model = model + self.ptr = ptr + } + + deinit { transcribe_langid_session_free(ptr) } + + /// Identify the language of one clip (16 kHz mono float32). + public func run(_ pcm: [Float], options: LangIdOptions = .init()) throws -> LangIdResult { + if let allowed = options.allowed, allowed.isEmpty { + // NULL would mean "every label", the opposite of an empty list. + throw TranscribeError.invalidArgument("allowed is empty; pass nil for every label") + } + if let allowed = options.allowed, allowed.contains(where: { $0.contains("\0") }) { + throw TranscribeError.invalidArgument("allowed contains a NUL character") + } + // strdup'd copies stay alive until the native call returns. + let owned: [UnsafeMutablePointer?] = options.allowed?.map { strdup($0) } ?? [] + defer { owned.forEach { free($0) } } + let codes: [UnsafePointer?] = owned.map { UnsafePointer($0) } + + return try model.withCompute( + busyIfStreaming: "a stream is active on this model; finish or drop it before langid run()" + ) { + let status = codes.withUnsafeBufferPointer { codeBuf -> transcribe_status in + var params = transcribe_langid_params() + transcribe_langid_params_init(¶ms) + if options.allowed != nil { + params.allowed = codeBuf.baseAddress + params.n_allowed = Int32(codeBuf.count) + } + params.top_k = options.topK + return pcm.withUnsafeBufferPointer { + transcribe_langid_run(ptr, $0.baseAddress, Int32($0.count), ¶ms) + } + } + try TranscribeError.check(status, context: "langid_run") + var res = transcribe_langid_result(); transcribe_langid_result_init(&res) + try TranscribeError.check(transcribe_langid_get_result(ptr, &res), context: "langid_get_result") + let candidates = (0.. LangIdCandidate in + var c = transcribe_langid_candidate(); transcribe_langid_candidate_init(&c) + _ = transcribe_langid_get_candidate(ptr, i, &c) + return LangIdCandidate( + index: c.index, code: c.code.map { String(cString: $0) } ?? "", + name: c.name.map { String(cString: $0) } ?? "", + p: c.p, logit: c.logit) + } + return LangIdResult(candidates: candidates, nAllowed: res.n_allowed, + allowedMass: res.allowed_mass, audioMs: res.audio_ms) + } + } + + /// `run` hopped off the caller's thread/actor onto a background queue, with + /// Swift task cancellation bridged to the native abort. Same contract as + /// `Session.run(_:options:) async`. + public func run(_ pcm: [Float], options: LangIdOptions = .init()) async throws -> LangIdResult { + nonisolated(unsafe) let this = self + let bridged = (cancelToken == nil) ? CancellationToken() : nil + if let bridged { setCancellationToken(bridged) } + defer { if bridged != nil { clearCancellationToken() } } + return try await withTaskCancellationHandler { + try await withCheckedThrowingContinuation { + (cont: CheckedContinuation) in + DispatchQueue.global().async { + cont.resume(with: Result { try this.run(pcm, options: options) }) + } + } + } onCancel: { + bridged?.cancel() + } + } + + /// Timings from the most recent run (`decodeMs` is 0). + public var timings: Timings { + var t = transcribe_timings(); transcribe_timings_init(&t) + _ = transcribe_langid_get_timings(ptr, &t) + return Timings(t) + } +} diff --git a/bindings/swift/Sources/TranscribeCpp/Model.swift b/bindings/swift/Sources/TranscribeCpp/Model.swift index 34ef8dc7e..c3a32d870 100644 --- a/bindings/swift/Sources/TranscribeCpp/Model.swift +++ b/bindings/swift/Sources/TranscribeCpp/Model.swift @@ -2,12 +2,13 @@ import CTranscribe import Foundation /// The kinds of work a model serves (`transcribe_model_roles`). ASR is -/// `Session`; DIARIZE is `DiarizeSession`. +/// `Session`; DIARIZE is `DiarizeSession`; LANGID is `LangIdSession`. public struct Roles: OptionSet, Sendable { public let rawValue: UInt32 public init(rawValue: UInt32) { self.rawValue = rawValue } public static let asr = Roles(rawValue: TRANSCRIBE_ROLE_ASR.rawValue) public static let diarize = Roles(rawValue: TRANSCRIBE_ROLE_DIARIZE.rawValue) + public static let langId = Roles(rawValue: TRANSCRIBE_ROLE_LANGID.rawValue) } /// A loaded model. Safe to share across threads (`@unchecked Sendable`): the C diff --git a/bindings/swift/Sources/TranscribeCpp/TranscribeError.swift b/bindings/swift/Sources/TranscribeCpp/TranscribeError.swift index 4e0db09c9..ea6798609 100644 --- a/bindings/swift/Sources/TranscribeCpp/TranscribeError.swift +++ b/bindings/swift/Sources/TranscribeCpp/TranscribeError.swift @@ -30,6 +30,9 @@ public enum TranscribeError: Error { case outputRepetition(message: String, partial: Transcript?) /// The model's `roles` lack the one the call needs (`TRANSCRIBE_ERR_UNSUPPORTED_ROLE`). case unsupportedRole(String) + /// The audio is shorter than the role's minimum (`TRANSCRIBE_ERR_INPUT_TOO_SHORT`, + /// e.g. `LangIdInfo.minAudioMs`). + case inputTooShort(String) case versionMismatch(String) case busy(String) case other(status: Int32, message: String) @@ -73,6 +76,8 @@ public enum TranscribeError: Error { return .outputRepetition(message: message, partial: nil) case TRANSCRIBE_ERR_UNSUPPORTED_ROLE: return .unsupportedRole(message) + case TRANSCRIBE_ERR_INPUT_TOO_SHORT: + return .inputTooShort(message) default: return .other(status: raw, message: message) } diff --git a/bindings/swift/Tests/TranscribeCppTests/LangIdTests.swift b/bindings/swift/Tests/TranscribeCppTests/LangIdTests.swift new file mode 100644 index 000000000..e6f48036a --- /dev/null +++ b/bindings/swift/Tests/TranscribeCppTests/LangIdTests.swift @@ -0,0 +1,101 @@ +import XCTest + +@testable import TranscribeCpp + +/// LANGID role (`LangIdSession`): contract checks on the toy ecapa_tdnn +/// fixture, top-1 on the real VoxLingua107 model. +final class LangIdTests: XCTestCase { + private func noise(_ n: Int, seed: UInt32 = 7) -> [Float] { + var s = seed | 1 + return (0..> 8) & 0xFFFFFF) / 16777216.0 - 0.5 + } + } + + private func cpuModel(_ path: String) throws -> Model { + try Model(path: path, options: ModelOptions(backend: .cpu)) + } + + func testToyRolesInfoLabels() throws { + let model = try cpuModel(try Fixtures.langIdToyModelPath()) + XCTAssertEqual(model.roles, .langId) + XCTAssertEqual(try model.langIdInfo, LangIdInfo(sampleRate: 16000, nLabels: 5, minAudioMs: 500)) + let labels = try model.langIdLabels + XCTAssertEqual(labels[2].code, "cc") + XCTAssertEqual(labels[2].name, "Charlie") + XCTAssertEqual(model.langIdLabelIndex("xx"), 0) + XCTAssertNil(model.langIdLabelIndex("zz")) + XCTAssertThrowsError(try model.session()) { error in + guard case TranscribeError.unsupportedRole = error else { return XCTFail("\(error)") } + } + } + + func testToyRunContract() throws { + let model = try cpuModel(try Fixtures.langIdToyModelPath()) + let lid = try model.langIdSession(threads: 1) + let pcm = noise(16000) + + let r = try lid.run(pcm) + XCTAssertEqual(r.candidates.count, 5) + XCTAssertEqual(r.nAllowed, 5) + XCTAssertEqual(r.allowedMass, 1.0) + XCTAssertEqual(r.audioMs, 1000) + XCTAssertEqual(r.code, r.candidates[0].code) + XCTAssertEqual(r.candidates.map(\.p).reduce(0, +), 1.0, accuracy: 1e-5) + + let restricted = try lid.run(pcm, options: LangIdOptions(allowed: ["bb", "dd"])) + XCTAssertEqual(Set(restricted.candidates.map(\.code)), ["bb", "dd"]) + XCTAssertLessThan(restricted.allowedMass, 1.0) + + XCTAssertEqual(try lid.run(pcm, options: LangIdOptions(topK: 2)).candidates.count, 2) + + XCTAssertThrowsError(try lid.run(pcm, options: LangIdOptions(allowed: []))) { error in + guard case TranscribeError.invalidArgument = error else { return XCTFail("\(error)") } + } + XCTAssertThrowsError(try lid.run(pcm, options: LangIdOptions(allowed: ["zz"]))) { error in + guard case TranscribeError.unsupported = error else { return XCTFail("\(error)") } + } + XCTAssertThrowsError(try lid.run(Array(pcm[..<6400]))) { error in + guard case TranscribeError.inputTooShort = error else { return XCTFail("\(error)") } + } + XCTAssertGreaterThan(lid.timings.encodeMs, 0) + } + + func testToyInteriorNulIsRejected() throws { + // C would cut "bb\0zz" to "bb"; the code must not silently narrow. + let model = try cpuModel(try Fixtures.langIdToyModelPath()) + XCTAssertNil(model.langIdLabelIndex("bb\0zz")) + let lid = try model.langIdSession() + XCTAssertThrowsError(try lid.run(noise(16000), options: LangIdOptions(allowed: ["bb\0zz"]))) { error in + guard case TranscribeError.invalidArgument = error else { return XCTFail("\(error)") } + } + } + + func testToyCancelAborts() throws { + let model = try cpuModel(try Fixtures.langIdToyModelPath()) + let lid = try model.langIdSession() + let token = CancellationToken() + lid.setCancellationToken(token) + token.cancel() + XCTAssertThrowsError(try lid.run(noise(16000))) { error in + guard case TranscribeError.aborted = error else { return XCTFail("\(error)") } + } + lid.clearCancellationToken() + XCTAssertNoThrow(try lid.run(noise(16000))) + } + + func testRealFleursTop1() throws { + let model = try cpuModel(try Fixtures.langIdModelPath()) + XCTAssertEqual(try model.langIdInfo.nLabels, 107) + XCTAssertEqual(model.langIdLabelIndex("he"), model.langIdLabelIndex("iw")) + let lid = try model.langIdSession() + for code in ["en", "de", "fr", "es", "ja", "zh", "ru", "id"] { + let pcm = try Fixtures.loadWav( + Fixtures.repoRoot().appendingPathComponent("samples/fleurs-\(code).wav").path) + let r = try lid.run(pcm, options: LangIdOptions(topK: 3)) + XCTAssertEqual(r.code, code) + XCTAssertGreaterThanOrEqual(r.candidates[0].p, 0.5) + } + } +} diff --git a/bindings/swift/Tests/TranscribeCppTests/NoModelTests.swift b/bindings/swift/Tests/TranscribeCppTests/NoModelTests.swift index 77d5d0f5a..dc4fc5973 100644 --- a/bindings/swift/Tests/TranscribeCppTests/NoModelTests.swift +++ b/bindings/swift/Tests/TranscribeCppTests/NoModelTests.swift @@ -21,7 +21,7 @@ final class NoModelTests: XCTestCase { func testAbiStructSizesAreLive() { // A real layout is non-zero; a garbage/empty one would be 0. - for s in [AbiStruct.runParams, .capabilities, .segment, .sessionLimits] { + for s in [AbiStruct.runParams, .capabilities, .segment, .sessionLimits, .langIdResult, .langIdCandidate] { XCTAssertGreaterThan(Transcribe.abiStructSize(s), 0, "\(s)") } } @@ -60,15 +60,15 @@ final class NoModelTests: XCTestCase { 10: "unsupported", 11: "unsupported", 12: "unsupported", 13: "aborted", 14: "badStructSize", 15: "unsupported", 16: "unsupported", 17: "inputTooLong", 18: "outputTruncated", - 19: "outputRepetition", 20: "unsupportedRole", + 19: "outputRepetition", 20: "unsupportedRole", 21: "inputTooShort", ] - for raw in 1...20 { + for raw in 1...21 { let status = transcribe_status(rawValue: UInt32(raw)) XCTAssertNotEqual(Transcribe.statusString(Int32(raw)), "unknown status", "status \(raw)") let name = String(describing: TranscribeError.make(status)).prefix { $0 != "(" } XCTAssertEqual(String(name), expected[Int32(raw)], "status \(raw)") } - XCTAssertEqual(Transcribe.statusString(21), "unknown status", + XCTAssertEqual(Transcribe.statusString(22), "unknown status", "a new status was appended; map it in TranscribeError.make") } diff --git a/bindings/swift/Tests/TranscribeCppTests/TestSupport.swift b/bindings/swift/Tests/TranscribeCppTests/TestSupport.swift index 5d495ecaa..9875f77e4 100644 --- a/bindings/swift/Tests/TranscribeCppTests/TestSupport.swift +++ b/bindings/swift/Tests/TranscribeCppTests/TestSupport.swift @@ -96,6 +96,25 @@ enum Fixtures { return (model, try loadWav(audio)) } + /// VoxLingua107 ECAPA-TDNN (LANGID only), or `XCTSkip` when absent. + static func langIdModelPath() throws -> String { + guard let model = familyModel( + "TRANSCRIBE_SMOKE_LANGID_MODEL", + "models/lang-id-voxlingua107-ecapa/lang-id-voxlingua107-ecapa-Q8_0.gguf") + else { throw XCTSkip("no language ID model (set TRANSCRIBE_SMOKE_LANGID_MODEL)") } + return model + } + + /// The toy ecapa_tdnn GGUF the C++ build generates under tests/fixtures/, + /// or `XCTSkip` before the C++ test fixtures have been built. + static func langIdToyModelPath() throws -> String { + let path = repoRoot().appendingPathComponent("tests/fixtures/arch_ecapa_tdnn_minimal.gguf").path + guard FileManager.default.fileExists(atPath: path) else { + throw XCTSkip("no toy ecapa_tdnn fixture (build the C++ `fixtures` target)") + } + return path + } + /// The model path + decoded PCM, or `XCTSkip` when either is absent. static func modelAndAudio() throws -> (model: String, pcm: [Float]) { guard let model = modelPath() else { diff --git a/bindings/typescript/README.md b/bindings/typescript/README.md index da016a267..a3a714fbe 100644 --- a/bindings/typescript/README.md +++ b/bindings/typescript/README.md @@ -127,9 +127,22 @@ for (const t of turns) console.log(t.speakerId, t.t0Ms, t.t1Ms); `diarizer.timings` reports the last run. Diarize runs wait on the same model-wide lock as other compute calls (see below). +### Language ID (LANGID role) + +A `"langid"` model (VoxLingua107 ECAPA-TDNN) ranks its own label codes; match +`result.code` against an ASR model's `capabilities.languages` yourself. +Omitting `allowed` scores every label, `[]` throws `InvalidArgument`, and +clips under `model.langidInfo.minAudioMs` throw `InputTooShort`. + +```ts +using lid = model.createLangIdSession(); +const result = await lid.run(pcm, { allowed: ["en", "de", "fr"], topK: 3 }); +console.log(result.code, result.candidates[0].p, result.allowedMass); +``` + ### Resource management -`TranscribeModel`, `Session`, `DiarizeSession`, and `Stream` all implement +`TranscribeModel`, `Session`, `DiarizeSession`, `LangIdSession`, and `Stream` all implement `Symbol.dispose`, so `using` works (TypeScript 5.2+ / Node 22+): ```ts diff --git a/bindings/typescript/src/_generated.ts b/bindings/typescript/src/_generated.ts index 5e7ea6f31..dae860c8c 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 = "bd3273dabb25a1fe"; +export const PUBLIC_HEADER_HASH = "d544d70a2b5acf50"; // === enum constants === export const TRANSCRIBE_OK = 0; @@ -35,6 +35,7 @@ export const TRANSCRIBE_ERR_INPUT_TOO_LONG = 17; export const TRANSCRIBE_ERR_OUTPUT_TRUNCATED = 18; export const TRANSCRIBE_ERR_OUTPUT_REPETITION = 19; export const TRANSCRIBE_ERR_UNSUPPORTED_ROLE = 20; +export const TRANSCRIBE_ERR_INPUT_TOO_SHORT = 21; export const TRANSCRIBE_ABI_MODEL_LOAD_PARAMS = 0; export const TRANSCRIBE_ABI_SESSION_PARAMS = 1; export const TRANSCRIBE_ABI_RUN_PARAMS = 2; @@ -54,6 +55,11 @@ export const TRANSCRIBE_ABI_BACKEND_INIT_PARAMS = 15; export const TRANSCRIBE_ABI_DIARIZE_INFO = 16; export const TRANSCRIBE_ABI_DIARIZE_SESSION_PARAMS = 17; export const TRANSCRIBE_ABI_DIARIZE_PARAMS = 18; +export const TRANSCRIBE_ABI_LANGID_INFO = 19; +export const TRANSCRIBE_ABI_LANGID_SESSION_PARAMS = 20; +export const TRANSCRIBE_ABI_LANGID_PARAMS = 21; +export const TRANSCRIBE_ABI_LANGID_RESULT = 22; +export const TRANSCRIBE_ABI_LANGID_CANDIDATE = 23; export const TRANSCRIBE_LOG_LEVEL_NONE = 0; export const TRANSCRIBE_LOG_LEVEL_INFO = 1; export const TRANSCRIBE_LOG_LEVEL_WARN = 2; @@ -96,6 +102,7 @@ export const TRANSCRIBE_DEVICE_TYPE_IGPU = 2; export const TRANSCRIBE_DEVICE_TYPE_ACCEL = 3; export const TRANSCRIBE_ROLE_ASR = 1; export const TRANSCRIBE_ROLE_DIARIZE = 2; +export const TRANSCRIBE_ROLE_LANGID = 4; export const TRANSCRIBE_FEATURE_INITIAL_PROMPT = 0; export const TRANSCRIBE_FEATURE_TEMPERATURE_FALLBACK = 1; export const TRANSCRIBE_FEATURE_LONG_FORM = 2; @@ -157,6 +164,11 @@ export const STRUCT_LAYOUT: Record = { 'transcribe_diarize_info': { size: 16, align: 8, offsets: {'struct_size': 0, 'sample_rate': 8, 'max_speakers': 12} }, 'transcribe_diarize_session_params': { size: 16, align: 8, offsets: {'struct_size': 0, 'n_threads': 8} }, 'transcribe_diarize_params': { size: 16, align: 8, offsets: {'struct_size': 0, 'family': 8} }, + 'transcribe_langid_info': { size: 24, align: 8, offsets: {'struct_size': 0, 'sample_rate': 8, 'n_labels': 12, 'min_audio_ms': 16} }, + 'transcribe_langid_session_params': { size: 16, align: 8, offsets: {'struct_size': 0, 'n_threads': 8, 'max_audio_ms': 12} }, + 'transcribe_langid_params': { size: 24, align: 8, offsets: {'struct_size': 0, 'allowed': 8, 'n_allowed': 16, 'top_k': 20} }, + 'transcribe_langid_result': { size: 32, align: 8, offsets: {'struct_size': 0, 'n_candidates': 8, 'n_allowed': 12, 'allowed_mass': 16, 'audio_ms': 24} }, + 'transcribe_langid_candidate': { size: 40, align: 8, offsets: {'struct_size': 0, 'index': 8, 'code': 16, 'name': 24, 'p': 32, 'logit': 36} }, 'transcribe_moonshine_streaming_stream_ext': { size: 24, align: 8, offsets: {'ext': 0, 'min_decode_interval_ms': 16} }, 'transcribe_parakeet_stream_ext': { size: 24, align: 8, offsets: {'ext': 0, 'att_context_right': 16} }, 'transcribe_parakeet_buffered_stream_ext': { size: 32, align: 8, offsets: {'ext': 0, 'left_ms': 16, 'chunk_ms': 20, 'right_ms': 24} }, @@ -186,6 +198,11 @@ export const ABI_STRUCT_IDS: Record = { 'transcribe_diarize_info': 16, 'transcribe_diarize_session_params': 17, 'transcribe_diarize_params': 18, + 'transcribe_langid_info': 19, + 'transcribe_langid_session_params': 20, + 'transcribe_langid_params': 21, + 'transcribe_langid_result': 22, + 'transcribe_langid_candidate': 23, }; // Build koffi struct types; returns a name -> koffi.IKoffiCType map. @@ -210,6 +227,11 @@ export function defineTypes(koffi: any): Record { T['transcribe_diarize_info'] = koffi.struct({ struct_size: 'uint64_t', sample_rate: 'int32_t', max_speakers: 'int32_t' }); T['transcribe_diarize_session_params'] = koffi.struct({ struct_size: 'uint64_t', n_threads: 'int32_t' }); T['transcribe_diarize_params'] = koffi.struct({ struct_size: 'uint64_t', family: 'void *' }); + T['transcribe_langid_info'] = koffi.struct({ struct_size: 'uint64_t', sample_rate: 'int32_t', n_labels: 'int32_t', min_audio_ms: 'int32_t' }); + T['transcribe_langid_session_params'] = koffi.struct({ struct_size: 'uint64_t', n_threads: 'int32_t', max_audio_ms: 'int32_t' }); + T['transcribe_langid_params'] = koffi.struct({ struct_size: 'uint64_t', allowed: 'void *', n_allowed: 'int32_t', top_k: 'int32_t' }); + T['transcribe_langid_result'] = koffi.struct({ struct_size: 'uint64_t', n_candidates: 'int32_t', n_allowed: 'int32_t', allowed_mass: 'float', audio_ms: 'int64_t' }); + T['transcribe_langid_candidate'] = koffi.struct({ struct_size: 'uint64_t', index: 'int32_t', code: 'char *', name: 'char *', p: 'float', logit: 'float' }); T['transcribe_moonshine_streaming_stream_ext'] = koffi.struct({ ext: T['transcribe_ext'], min_decode_interval_ms: 'int32_t' }); T['transcribe_parakeet_stream_ext'] = koffi.struct({ ext: T['transcribe_ext'], att_context_right: 'int32_t' }); T['transcribe_parakeet_buffered_stream_ext'] = koffi.struct({ ext: T['transcribe_ext'], left_ms: 'int32_t', chunk_ms: 'int32_t', right_ms: 'int32_t' }); @@ -273,6 +295,22 @@ export const FUNCTION_SIGNATURES: Record = { 'transcribe_init_backends': { ret: 'transcribe_status', args: ['const char *'] }, 'transcribe_init_backends_default': { ret: 'transcribe_status', args: [] }, 'transcribe_init_backends_ex': { ret: 'transcribe_status', args: ['const struct transcribe_backend_init_params *'] }, + 'transcribe_langid_candidate_init': { ret: 'void', args: ['struct transcribe_langid_candidate *'] }, + 'transcribe_langid_get_candidate': { ret: 'transcribe_status', args: ['const struct transcribe_langid_session *', 'int', 'struct transcribe_langid_candidate *'] }, + 'transcribe_langid_get_info': { ret: 'transcribe_status', args: ['const struct transcribe_model *', 'struct transcribe_langid_info *'] }, + 'transcribe_langid_get_result': { ret: 'transcribe_status', args: ['const struct transcribe_langid_session *', 'struct transcribe_langid_result *'] }, + 'transcribe_langid_get_timings': { ret: 'transcribe_status', args: ['const struct transcribe_langid_session *', 'struct transcribe_timings *'] }, + 'transcribe_langid_info_init': { ret: 'void', args: ['struct transcribe_langid_info *'] }, + 'transcribe_langid_label_code': { ret: 'const char *', args: ['const struct transcribe_model *', 'int32_t'] }, + 'transcribe_langid_label_index': { ret: 'int32_t', args: ['const struct transcribe_model *', 'const char *'] }, + 'transcribe_langid_label_name': { ret: 'const char *', args: ['const struct transcribe_model *', 'int32_t'] }, + 'transcribe_langid_params_init': { ret: 'void', args: ['struct transcribe_langid_params *'] }, + 'transcribe_langid_result_init': { ret: 'void', args: ['struct transcribe_langid_result *'] }, + 'transcribe_langid_run': { ret: 'transcribe_status', args: ['struct transcribe_langid_session *', 'const float *', 'int', 'const struct transcribe_langid_params *'] }, + 'transcribe_langid_session_free': { ret: 'void', args: ['struct transcribe_langid_session *'] }, + 'transcribe_langid_session_init': { ret: 'transcribe_status', args: ['struct transcribe_model *', 'const struct transcribe_langid_session_params *', 'struct transcribe_langid_session **'] }, + 'transcribe_langid_session_params_init': { ret: 'void', args: ['struct transcribe_langid_session_params *'] }, + 'transcribe_langid_set_abort_callback': { ret: 'void', args: ['struct transcribe_langid_session *', 'transcribe_abort_callback', 'void *'] }, 'transcribe_log_set': { ret: 'void', args: ['transcribe_log_callback', 'void *'] }, 'transcribe_model_accepts_ext_kind': { ret: '_Bool', args: ['const struct transcribe_model *', 'transcribe_ext_slot', 'uint32_t'] }, 'transcribe_model_arch_string': { ret: 'const char *', args: ['const struct transcribe_model *'] }, diff --git a/bindings/typescript/src/errors.ts b/bindings/typescript/src/errors.ts index 1eec96790..1bb4d7f40 100644 --- a/bindings/typescript/src/errors.ts +++ b/bindings/typescript/src/errors.ts @@ -35,6 +35,8 @@ export class BackendError extends TranscribeError {} export class UnsupportedRequest extends TranscribeError {} export class AbiError extends TranscribeError {} export class InputTooLong extends TranscribeError {} +/** Raised when the audio is shorter than the role's minimum (e.g. language ID's minAudioMs). */ +export class InputTooShort extends TranscribeError {} export class VersionMismatch extends TranscribeError {} /** @@ -43,7 +45,7 @@ export class VersionMismatch extends TranscribeError {} */ export class BackendInitializing extends TranscribeError {} -/** Raised when the model does not serve the role (ASR, diarize) the call needs. */ +/** Raised when the model does not serve the role (ASR, diarize, langid) the call needs. */ export class UnsupportedRole extends TranscribeError {} /** Raised when a run is cancelled; carries any partial transcript in `partialResult`. */ @@ -80,6 +82,7 @@ const STATUS_TO_EXC: Record TranscribeErr [g.TRANSCRIBE_ERR_OUTPUT_TRUNCATED]: OutputTruncated, [g.TRANSCRIBE_ERR_OUTPUT_REPETITION]: OutputRepetition, [g.TRANSCRIBE_ERR_UNSUPPORTED_ROLE]: UnsupportedRole, + [g.TRANSCRIBE_ERR_INPUT_TOO_SHORT]: InputTooShort, }; /** Build (do not throw) the mapped exception for a status. */ diff --git a/bindings/typescript/src/ffi.ts b/bindings/typescript/src/ffi.ts index c6d3c4bf6..7ae9bc9dd 100644 --- a/bindings/typescript/src/ffi.ts +++ b/bindings/typescript/src/ffi.ts @@ -220,6 +220,60 @@ export function bindLibrary(libraryPath: string): Bound { iop(T.transcribe_timings), ]), + // langid role + langidInfoInit: lib.func("transcribe_langid_info_init", "void", [ + outp(T.transcribe_langid_info), + ]), + langidGetInfo: lib.func("transcribe_langid_get_info", "int", [ + "void *", + iop(T.transcribe_langid_info), + ]), + langidLabelCode: lib.func("transcribe_langid_label_code", "const char *", ["void *", "int32_t"]), + langidLabelName: lib.func("transcribe_langid_label_name", "const char *", ["void *", "int32_t"]), + langidLabelIndex: lib.func("transcribe_langid_label_index", "int32_t", ["void *", "const char *"]), + langidSessionParamsInit: lib.func("transcribe_langid_session_params_init", "void", [ + outp(T.transcribe_langid_session_params), + ]), + langidSessionInit: lib.func("transcribe_langid_session_init", "int", [ + "void *", + inp(T.transcribe_langid_session_params), + handleOut, + ]), + langidSessionFree: lib.func("transcribe_langid_session_free", "void", ["void *"]), + langidSetAbortCallback: lib.func("transcribe_langid_set_abort_callback", "void", [ + "void *", + "void *", + "void *", + ]), + langidParamsInit: lib.func("transcribe_langid_params_init", "void", [ + outp(T.transcribe_langid_params), + ]), + langidRun: lib.func("transcribe_langid_run", "int", [ + "void *", + inp("float"), + "int", + inp(T.transcribe_langid_params), + ]), + langidResultInit: lib.func("transcribe_langid_result_init", "void", [ + outp(T.transcribe_langid_result), + ]), + langidGetResult: lib.func("transcribe_langid_get_result", "int", [ + "void *", + iop(T.transcribe_langid_result), + ]), + langidCandidateInit: lib.func("transcribe_langid_candidate_init", "void", [ + outp(T.transcribe_langid_candidate), + ]), + langidGetCandidate: lib.func("transcribe_langid_get_candidate", "int", [ + "void *", + "int", + iop(T.transcribe_langid_candidate), + ]), + langidGetTimings: lib.func("transcribe_langid_get_timings", "int", [ + "void *", + iop(T.transcribe_timings), + ]), + // batch (offline) runBatch: lib.func("transcribe_run_batch", "int", [ "void *", diff --git a/bindings/typescript/src/index.ts b/bindings/typescript/src/index.ts index f37a605c6..b831f1bfb 100644 --- a/bindings/typescript/src/index.ts +++ b/bindings/typescript/src/index.ts @@ -48,6 +48,11 @@ import type { Feature, Itn, KvType, + LangIdCandidate, + LangIdInfo, + LangIdOptions, + LangIdResult, + LangIdSessionOptions, ModelOptions, PcmLike, Pnc, @@ -135,6 +140,7 @@ const FEATURES: Record = { const ROLES: Record = { asr: g.TRANSCRIBE_ROLE_ASR, diarize: g.TRANSCRIBE_ROLE_DIARIZE, + langid: g.TRANSCRIBE_ROLE_LANGID, }; // ---- helpers --------------------------------------------------------------- @@ -1463,13 +1469,117 @@ export class DiarizeSession { } } +// ---- LangIdSession --------------------------------------------------------- + +/** A LANGID-role session: which language is spoken. Same compute rules as Session. */ +export class LangIdSession { + #n: Native; + #core: SessionCore; + #model: TranscribeModel; // keep the model alive while this session lives + #untrack: (self: LangIdSession) => void; + + /** @internal */ + constructor( + n: Native, + model: TranscribeModel, + handle: any, + lock: Mutex, + untrack: (self: LangIdSession) => void, + ) { + this.#n = n; + this.#model = model; + this.#core = new SessionCore(n, handle, lock, n.F.langidSetAbortCallback); + this.#untrack = untrack; + } + + /** + * Identify the language of one clip; input longer than the session's + * maxAudioMs is scored on its tail. The input PCM is borrowed, not copied + * (see Session.run). + */ + async run(pcm: PcmLike, opts: LangIdOptions = {}): Promise { + const n = this.#n; + const F = n.F; + const h = this.#core.handle; + const samples = toFloat32(pcm); + const p: any = {}; + F.langidParamsInit(p); + if (opts.topK !== undefined) p.top_k = opts.topK; + if (opts.allowed !== undefined && opts.allowed !== null) { + const codes = opts.allowed; + if (!Array.isArray(codes) || !codes.every((c) => typeof c === "string")) + throw new InvalidArgument("allowed must be an array of strings"); + // NULL would mean "every label", the opposite of an empty list. + if (codes.length === 0) + throw new InvalidArgument("allowed is empty; omit it for every label"); + codes.forEach((c) => cstr(c, "allowed")); + // Freed after the call returns (including async worker calls). + const type = n.koffi.array("char *", codes.length); + const arr = n.koffi.alloc(type, 1); + n.koffi.encode(arr, type, codes); + p.allowed = arr; + p.n_allowed = codes.length; + } + + return this.#core.exclusive("langid", async (call) => { + const status = await call("run()", opts.signal, F.langidRun, h, samples, samples.length, p); + check(n, status, "transcribe_langid_run"); + const res: any = {}; + F.langidResultInit(res); + check(n, F.langidGetResult(h, res), "transcribe_langid_get_result"); + const candidates: LangIdCandidate[] = []; + for (let i = 0; i < res.n_candidates; i++) { + const c: any = {}; + F.langidCandidateInit(c); + check(n, F.langidGetCandidate(h, i, c), "transcribe_langid_get_candidate"); + candidates.push({ + index: c.index, + code: c.code ?? "", + name: c.name ?? "", + p: c.p, + logit: c.logit, + }); + } + return { + candidates, + code: candidates.length > 0 ? candidates[0].code : null, + nAllowed: res.n_allowed, + allowedMass: res.allowed_mass, + audioMs: Number(res.audio_ms), + }; + }).finally(() => { + if (p.allowed) { + n.koffi.free(p.allowed); + p.allowed = null; + } + }); + } + + /** load_ms plus the last run's mel / encode time. */ + get timings(): Timings { + this.#core.assertNotComputing("session timings"); + const h = this.#core.handle; + return readTimings(this.#n, (o) => this.#n.F.langidGetTimings(h, o)); + } + + dispose(): void { + if (this.#core.disposed) return; + this.#untrack(this); + this.#core.dispose(this.#n.F.langidSessionFree); + } + + [Symbol.dispose](): void { + this.dispose(); + } +} + // ---- Model ----------------------------------------------------------------- export class TranscribeModel { #n: Native; #h: any; #disposed = false; - #sessions = new Set(); + #sessions = new Set(); #lock = new Mutex(); // serializes compute across all sessions of this model private constructor(n: Native, handle: any) { @@ -1584,6 +1694,49 @@ export class TranscribeModel { return session; } + /** Static facts of a "langid" model; UnsupportedRole otherwise. */ + get langidInfo(): LangIdInfo { + const n = this.#n; + const info: any = {}; + n.F.langidInfoInit(info); + check(n, n.F.langidGetInfo(this.handle, info), "reading langid info"); + return { sampleRate: info.sample_rate, nLabels: info.n_labels, minAudioMs: info.min_audio_ms }; + } + + /** [code, name] per label index of a "langid" model; UnsupportedRole otherwise. */ + get langidLabels(): Array<[string, string]> { + const n = this.langidInfo.nLabels; + const F = this.#n.F; + const out: Array<[string, string]> = []; + for (let i = 0; i < n; i++) + out.push([F.langidLabelCode(this.handle, i) ?? "", F.langidLabelName(this.handle, i) ?? ""]); + return out; + } + + /** Label index of a code or alias ("he" and "iw" name the same label), or null. */ + langidLabelIndex(code: string): number | null { + const i = this.#n.F.langidLabelIndex(this.handle, cstr(code, "code")); + return i >= 0 ? i : null; + } + + /** Open a LANGID-role session; UnsupportedRole on a model without "langid". */ + createLangIdSession(opts: LangIdSessionOptions = {}): LangIdSession { + const n = this.#n; + const p: any = {}; + n.F.langidSessionParamsInit(p); + if (opts.nThreads !== undefined) p.n_threads = opts.nThreads; + if (opts.maxAudioMs !== undefined) p.max_audio_ms = opts.maxAudioMs; + const out: any[] = [null]; + check(n, n.F.langidSessionInit(this.handle, p, out), "opening langid session"); + if (!out[0]) + throw new TranscribeError("langid session init returned a null handle"); + const session = new LangIdSession(n, this, out[0], this.#lock, (s) => + this.#sessions.delete(s), + ); + this.#sessions.add(session); + return session; + } + get capabilities(): Capabilities { const n = this.#n; const c: any = {}; diff --git a/bindings/typescript/src/types.ts b/bindings/typescript/src/types.ts index 1ddd6bfb1..1a9368b09 100644 --- a/bindings/typescript/src/types.ts +++ b/bindings/typescript/src/types.ts @@ -291,8 +291,8 @@ export type FamilyExtension = // ---- roles ----------------------------------------------------------------- -/** What a model serves: "asr" (transcription) and/or "diarize" (speaker turns). */ -export type Role = "asr" | "diarize"; +/** What a model serves: "asr" (transcription), "diarize" (speaker turns), "langid" (language). */ +export type Role = "asr" | "diarize" | "langid"; export interface DiarizeInfo { /** Input PCM rate. */ @@ -312,3 +312,55 @@ export interface DiarizeOptions { /** A diarize_run-slot family extension (e.g. sortformer_diarize). */ family?: FamilyExtension; } + +export interface LangIdInfo { + /** Input PCM rate. */ + sampleRate: number; + /** Label indices are [0, nLabels). */ + nLabels: number; + /** Shorter scored audio throws InputTooShort. */ + minAudioMs: number; +} + +export interface LangIdSessionOptions { + /** CPU threads for CPU-side ops; 0 = library default. */ + nThreads?: number; + /** Longer input is scored on its last maxAudioMs; 0 = 30000. */ + maxAudioMs?: number; +} + +export interface LangIdOptions { + /** Cancel the run cooperatively. */ + signal?: AbortSignal; + /** + * Restrict the decision to these codes or aliases. Omitted (or undefined / + * null) means every label; an empty array throws InvalidArgument; an + * unknown code throws UnsupportedRequest. + */ + allowed?: readonly string[] | null; + /** Keep the best topK candidates; 0 = every allowed label. */ + topK?: number; +} + +/** One ranked label; `code` is the model's own label ("en", "iw"). */ +export interface LangIdCandidate { + index: number; + code: string; + name: string; + /** Softmax renormalized over the allowed set. */ + p: number; + logit: number; +} + +export interface LangIdResult { + /** Ranked by p, descending; ties keep label order. */ + candidates: LangIdCandidate[]; + /** The top candidate's code, or null with no candidates. */ + code: string | null; + /** Labels in the allowed set (before topK). */ + nAllowed: number; + /** Unrestricted probability inside the allowed set (1 when unrestricted). */ + allowedMass: number; + /** Audio actually scored, after the crop. */ + audioMs: number; +} diff --git a/bindings/typescript/test/common.mjs b/bindings/typescript/test/common.mjs index de5e68e0e..c14942fa5 100644 --- a/bindings/typescript/test/common.mjs +++ b/bindings/typescript/test/common.mjs @@ -26,6 +26,14 @@ export const VOXTRAL_MODEL = process.env.TRANSCRIBE_SMOKE_VOXTRAL_MODEL || ""; export const SORTFORMER_MODEL = process.env.TRANSCRIBE_SMOKE_SORTFORMER_MODEL || ""; export const SORTFORMER_AUDIO = path.resolve(HERE, "../../../samples/sortformer-2spk-mix.wav"); +// VoxLingua107 ECAPA-TDNN (LANGID role only); the toy fixture is generated by +// the C++ test build (tests/fixtures), so it skips until that has been built. +export const LANGID_MODEL = + process.env.TRANSCRIBE_SMOKE_LANGID_MODEL || + path.resolve(HERE, "../../../models/lang-id-voxlingua107-ecapa/lang-id-voxlingua107-ecapa-Q8_0.gguf"); +export const LANGID_TOY_MODEL = path.resolve(HERE, "../../../tests/fixtures/arch_ecapa_tdnn_minimal.gguf"); +export const SAMPLES = path.resolve(HERE, "../../../samples"); + // jfk.wav ships in-repo; fetch-canary exports only the model paths. export const AUDIO = process.env.TRANSCRIBE_SMOKE_AUDIO || path.resolve(HERE, "../../../samples/jfk.wav"); diff --git a/bindings/typescript/test/errors.test.mjs b/bindings/typescript/test/errors.test.mjs index 0c0d13a2c..3dafd907e 100644 --- a/bindings/typescript/test/errors.test.mjs +++ b/bindings/typescript/test/errors.test.mjs @@ -20,6 +20,7 @@ import { InputTooLong, OutputTruncated, OutputRepetition, + InputTooShort, } from "../dist/index.js"; const EXPECTED = { @@ -43,6 +44,7 @@ const EXPECTED = { TRANSCRIBE_ERR_OUTPUT_TRUNCATED: OutputTruncated, TRANSCRIBE_ERR_OUTPUT_REPETITION: OutputRepetition, TRANSCRIBE_ERR_UNSUPPORTED_ROLE: UnsupportedRole, + TRANSCRIBE_ERR_INPUT_TOO_SHORT: InputTooShort, }; test("every transcribe_status maps to its documented error class", () => { diff --git a/bindings/typescript/test/langid.test.mjs b/bindings/typescript/test/langid.test.mjs new file mode 100644 index 000000000..69e0bbdff --- /dev/null +++ b/bindings/typescript/test/langid.test.mjs @@ -0,0 +1,104 @@ +// LANGID role: model.roles, langidInfo / labels and LangIdSession wiring. +// Contract checks run on the toy ecapa_tdnn fixture; top-1 on the real +// VoxLingua107 model. + +import assert from "node:assert/strict"; +import * as path from "node:path"; +import { modelTest, LANGID_MODEL, LANGID_TOY_MODEL, SAMPLES, readWav } from "./common.mjs"; +import { + TranscribeModel, + InvalidArgument, + InputTooShort, + UnsupportedRequest, + UnsupportedRole, + Aborted, +} from "../dist/index.js"; + +function noise(n, seed = 7) { + const out = new Float32Array(n); + let s = (seed | 1) >>> 0; + for (let i = 0; i < n; i++) { + s = (Math.imul(s, 1664525) + 1013904223) >>> 0; + out[i] = ((s >>> 8) & 0xffffff) / 16777216 - 0.5; + } + return out; +} + +modelTest("toy: roles, info, labels; ASR calls are refused", LANGID_TOY_MODEL, async () => { + const m = await TranscribeModel.load(LANGID_TOY_MODEL, { backend: "cpu" }); + try { + assert.deepEqual(m.roles, ["langid"]); + assert.deepEqual(m.langidInfo, { sampleRate: 16000, nLabels: 5, minAudioMs: 500 }); + assert.deepEqual(m.langidLabels[2], ["cc", "Charlie"]); + assert.equal(m.langidLabelIndex("xx"), 0); + assert.equal(m.langidLabelIndex("zz"), null); + assert.throws(() => m.capabilities, UnsupportedRole); + assert.throws(() => m.createSession(), UnsupportedRole); + } finally { + m.dispose(); + } +}); + +modelTest("toy: run contract (allowed, topK, crop, minimum)", LANGID_TOY_MODEL, async () => { + const m = await TranscribeModel.load(LANGID_TOY_MODEL, { backend: "cpu" }); + const lid = m.createLangIdSession({ nThreads: 1 }); + try { + const pcm = noise(16000); + const r = await lid.run(pcm); + assert.equal(r.candidates.length, 5); + assert.equal(r.nAllowed, 5); + assert.equal(r.allowedMass, 1); + assert.equal(r.audioMs, 1000); + assert.equal(r.code, r.candidates[0].code); + const sum = r.candidates.reduce((a, c) => a + c.p, 0); + assert.ok(Math.abs(sum - 1) < 1e-5); + + const restricted = await lid.run(pcm, { allowed: ["bb", "dd"] }); + assert.deepEqual(new Set(restricted.candidates.map((c) => c.code)), new Set(["bb", "dd"])); + assert.ok(restricted.allowedMass < 1); + + assert.equal((await lid.run(pcm, { topK: 2 })).candidates.length, 2); + assert.equal((await lid.run(pcm, { allowed: null })).nAllowed, 5); + // An empty list must not silently mean "every label". + await assert.rejects(lid.run(pcm, { allowed: [] }), InvalidArgument); + await assert.rejects(lid.run(pcm, { allowed: ["zz"] }), UnsupportedRequest); + await assert.rejects(lid.run(pcm.subarray(0, 6400)), InputTooShort); + assert.equal((await lid.run(noise(16000 * 31))).audioMs, 30000); + assert.ok(lid.timings.encodeMs > 0); + assert.throws(() => m.createLangIdSession({ maxAudioMs: 400 }), InvalidArgument); + } finally { + lid.dispose(); + m.dispose(); + } +}); + +modelTest("toy: an aborted signal cancels the run", LANGID_TOY_MODEL, async () => { + const m = await TranscribeModel.load(LANGID_TOY_MODEL, { backend: "cpu" }); + const lid = m.createLangIdSession(); + try { + const ac = new AbortController(); + ac.abort(); + await assert.rejects(lid.run(noise(16000), { signal: ac.signal }), Aborted); + assert.ok((await lid.run(noise(16000))).candidates.length > 0); + } finally { + lid.dispose(); + m.dispose(); + } +}); + +modelTest("real model: top-1 on the FLEURS clips", LANGID_MODEL, async () => { + const m = await TranscribeModel.load(LANGID_MODEL, { backend: "cpu" }); + const lid = m.createLangIdSession(); + try { + assert.equal(m.langidInfo.nLabels, 107); + assert.equal(m.langidLabelIndex("he"), m.langidLabelIndex("iw")); + for (const code of ["en", "de", "fr", "es", "ja", "zh", "ru", "id"]) { + const r = await lid.run(readWav(path.join(SAMPLES, `fleurs-${code}.wav`)), { topK: 3 }); + assert.equal(r.code, code); + assert.ok(r.candidates[0].p >= 0.5); + } + } finally { + lid.dispose(); + m.dispose(); + } +}); diff --git a/catalog/_benchmark_profiles.json b/catalog/_benchmark_profiles.json index ae2fd44b7..19731cc7d 100644 --- a/catalog/_benchmark_profiles.json +++ b/catalog/_benchmark_profiles.json @@ -1,5 +1,8 @@ { "default": "asr-publication-v2", + "roles": { + "langid": "langid-publication-v1" + }, "profiles": { "asr-publication-v2": { "description": "The complete benchmark set published for transcription models. v2: accuracy on L40S; batch_size is the recommended setting for new runs, and a cell is satisfied at any batch size (the row records the one used).", @@ -67,6 +70,68 @@ "reason": "A variant that supports exactly one non-English language is benched on that language at the same two clip lengths as jfk/dots. English jfk/dots decode out of distribution on a single-language fine-tune and can loop until the position cap, and whether that happens differs between CPU and GPU, so the figure would not be comparable." } } + }, + "langid-publication-v1": { + "role": "langid", + "description": "The benchmark set published for language ID (LANGID role) models. Accuracy: open-set top-1 over every label, the macro mean over the FLEURS test languages of scripts/langid/ingest.py, on the first crop_s seconds of each clip without silence trimming, C++ on CPU, every shipped quant; scored by scripts/langid/score.py from one scripts/langid/run.py sweep (all crops), which it also compares against the reference sweep (--ref) for the row's agreement. Speed: scripts/langid/bench.py through the Python binding; a sample -s is the first N seconds of samples/.wav.", + "accuracy": [ + { + "dataset": "fleurs", + "split": "test", + "languages": "pooled", + "metric": "accuracy", + "quants": "all-downloads", + "batch_size": 1, + "timestamps": "none", + "backend": "cpu", + "crop_s": "5", + "pooled_languages": [ + "en", + "zh", + "de", + "fr", + "es", + "pt", + "ja", + "ko", + "ru", + "cs", + "sk", + "no", + "da", + "id", + "ms" + ] + } + ], + "speed": { + "quants": "all-downloads", + "samples": [ + "ru-long-10s", + "ru-long-30s" + ], + "iterations": 20, + "warmup": 3, + "targets": [ + { + "machine": "m4-max", + "display": "Apple M4 Max", + "backends": [ + "cpu", + "metal" + ] + }, + { + "machine": "ryzen-4750u", + "display": "AMD Ryzen 7 PRO 4750U (Radeon RADV RENOIR)", + "backends": [ + "cpu", + "vulkan" + ], + "cooldown_tctl_c": 55.0 + } + ] + } } } } diff --git a/catalog/_schema.json b/catalog/_schema.json index 219e539b4..4334ddc2c 100644 --- a/catalog/_schema.json +++ b/catalog/_schema.json @@ -25,7 +25,7 @@ "description": "Loader family (the GGUF's general.architecture)" }, "role": { - "enum": ["asr", "diarize"], + "enum": ["asr", "diarize", "langid"], "description": "The role the model is published for (transcribe_model_roles). Absent means asr." }, "display_name": { @@ -286,7 +286,7 @@ "type": "object", "additionalProperties": false, "required": [ - "dataset", "split", "language", "quant", "metric", "err_pct", + "dataset", "split", "language", "quant", "metric", "ci95", "n_utts", "batch_size", "timestamps", "engine_sha" ], "properties": { @@ -320,7 +320,14 @@ }, "err_pct": { "type": "number", - "minimum": 0 + "minimum": 0, + "description": "Error rate in percent (wer / cer / der / cpwer rows). Accuracy rows carry acc_pct instead." + }, + "acc_pct": { + "type": "number", + "minimum": 0, + "maximum": 100, + "description": "Top-1 accuracy in percent; only on metric=accuracy rows (language ID), in place of err_pct." }, "ci95": { "type": "array", @@ -391,8 +398,26 @@ "mode": { "type": "string", "description": "Decoding mode for models that publish one metric under several modes (multitalker \"kernel\" / \"masked\" cpWER). Rows with a mode are a separate result set." + }, + "agreement": { + "type": "object", + "additionalProperties": false, + "required": ["n_agree", "n"], + "properties": { + "n_agree": {"type": "integer", "minimum": 0}, + "n": {"type": "integer", "minimum": 1}, + "max_abs_logit_delta": {"type": ["number", "null"], "minimum": 0} + }, + "description": "Language ID only: top-1 decisions of the run this row was scored from that match the reference implementation on the same audio (scripts/langid/score.py --ref, the port's ship gate), over every crop of the sweep rather than only the scored one, and the largest logit difference seen." } - } + }, + "allOf": [ + { + "if": {"properties": {"metric": {"const": "accuracy"}}}, + "then": {"required": ["acc_pct"], "not": {"required": ["err_pct"]}}, + "else": {"required": ["err_pct"], "not": {"required": ["acc_pct"]}} + } + ] } }, "headline_benchmark": { diff --git a/catalog/lang-id-voxlingua107-ecapa.json b/catalog/lang-id-voxlingua107-ecapa.json new file mode 100644 index 000000000..7357a197a --- /dev/null +++ b/catalog/lang-id-voxlingua107-ecapa.json @@ -0,0 +1,83 @@ +{ + "schema": "transcribe-catalog-v1", + "variant": "lang-id-voxlingua107-ecapa", + "family": "ecapa_tdnn", + "role": "langid", + "display_name": "VoxLingua107 ECAPA-TDNN", + "params": 21244679, + "license": { + "spdx": "apache-2.0", + "display": "Apache-2.0" + }, + "upstream_repo": "speechbrain/lang-id-voxlingua107-ecapa", + "upstream_commit": "0253049", + "published_repo": "handy-computer/lang-id-voxlingua107-ecapa-gguf", + "docs_page": "lang-id-voxlingua107-ecapa.md", + "languages": [ + "ab", "af", "am", "ar", "as", "az", "ba", "be", "bg", "bn", "bo", "br", + "bs", "ca", "ceb", "cs", "cy", "da", "de", "el", "en", "eo", "es", "et", + "eu", "fa", "fi", "fo", "fr", "gl", "gn", "gu", "gv", "ha", "haw", "hi", + "hr", "ht", "hu", "hy", "ia", "id", "is", "it", "iw", "ja", "jw", "ka", + "kk", "km", "kn", "ko", "la", "lb", "ln", "lo", "lt", "lv", "mg", "mi", + "mk", "ml", "mn", "mr", "ms", "mt", "my", "ne", "nl", "nn", "no", "oc", + "pa", "pl", "ps", "pt", "ro", "ru", "sa", "sco", "sd", "si", "sk", "sl", + "sn", "so", "sq", "sr", "su", "sv", "sw", "ta", "te", "tg", "th", "tk", + "tl", "tr", "tt", "uk", "ur", "uz", "vi", "war", "yi", "yo", "zh" + ], + "language_tag_form": "bare-bcp47", + "long_form_strategy": "hard-cap", + "max_audio_s": 30, + "capabilities": { + "transcribe": {"supported":false}, + "translate": {"supported":false}, + "lang_detect": {"supported":false}, + "timestamps": {"supported":false}, + "streaming": {"supported":false}, + "diarize": {"supported":false}, + "batching": {"supported":false} + }, + "downloads": [ + {"quant":"F32","filename":"lang-id-voxlingua107-ecapa-F32.gguf","size_bytes":84993184}, + {"quant":"F16","filename":"lang-id-voxlingua107-ecapa-F16.gguf","size_bytes":45299872}, + {"quant":"Q8_0","filename":"lang-id-voxlingua107-ecapa-Q8_0.gguf","size_bytes":26693632} + ], + "accuracy_benchmarks": [ + {"dataset":"fleurs","split":"test","language":"mul","backend":"cpu","quant":"F32","metric":"accuracy","acc_pct":85.23,"ci95":[84.07,86.47],"n_utts":3000,"batch_size":1,"timestamps":"none","engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","agreement":{"n_agree":12000,"n":12000,"max_abs_logit_delta":7.3e-05}}, + {"dataset":"fleurs","split":"test","language":"mul","backend":"cpu","quant":"F16","metric":"accuracy","acc_pct":85.17,"ci95":[84.0,86.43],"n_utts":3000,"batch_size":1,"timestamps":"none","engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","agreement":{"n_agree":11986,"n":12000,"max_abs_logit_delta":0.095}}, + {"dataset":"fleurs","split":"test","language":"mul","backend":"cpu","quant":"Q8_0","metric":"accuracy","acc_pct":86.3,"ci95":[85.17,87.5],"n_utts":3000,"batch_size":1,"timestamps":"none","engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","agreement":{"n_agree":11705,"n":12000,"max_abs_logit_delta":2.9}} + ], + "headline_benchmark": { + "dataset": "fleurs", + "split": "test", + "language": "mul", + "metric": "accuracy", + "batch_size": 1, + "timestamps": "none" + }, + "speed_benchmarks": [ + {"machine":"m4-max","backend":"cpu","quant":"F16","sample":"ru-long-10s","sample_duration_s":10.0,"total_ms":78.9,"xrt_compute":126.69,"wall_ms":80.0,"xrt_wall":125.03,"load_ms":10.6,"mel_ms":1.1,"encode_ms":77.8,"decode_ms":null,"engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"m4-max","backend":"cpu","quant":"F16","sample":"ru-long-30s","sample_duration_s":30.0,"total_ms":241.8,"xrt_compute":124.06,"wall_ms":243.9,"xrt_wall":123.02,"load_ms":10.6,"mel_ms":3.1,"encode_ms":238.7,"decode_ms":null,"engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"m4-max","backend":"cpu","quant":"F32","sample":"ru-long-10s","sample_duration_s":10.0,"total_ms":123.8,"xrt_compute":80.77,"wall_ms":124.8,"xrt_wall":80.14,"load_ms":21.3,"mel_ms":1.0,"encode_ms":122.8,"decode_ms":null,"engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"m4-max","backend":"cpu","quant":"F32","sample":"ru-long-30s","sample_duration_s":30.0,"total_ms":382.3,"xrt_compute":78.48,"wall_ms":384.3,"xrt_wall":78.07,"load_ms":21.3,"mel_ms":3.0,"encode_ms":379.3,"decode_ms":null,"engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"m4-max","backend":"cpu","quant":"Q8_0","sample":"ru-long-10s","sample_duration_s":10.0,"total_ms":80.4,"xrt_compute":124.38,"wall_ms":81.5,"xrt_wall":122.77,"load_ms":10.5,"mel_ms":1.1,"encode_ms":79.3,"decode_ms":null,"engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"m4-max","backend":"cpu","quant":"Q8_0","sample":"ru-long-30s","sample_duration_s":30.0,"total_ms":247.1,"xrt_compute":121.43,"wall_ms":249.1,"xrt_wall":120.42,"load_ms":10.5,"mel_ms":3.2,"encode_ms":243.9,"decode_ms":null,"engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"m4-max","backend":"metal","quant":"F16","sample":"ru-long-10s","sample_duration_s":10.0,"total_ms":8.3,"xrt_compute":1208.9,"wall_ms":10.8,"xrt_wall":926.08,"load_ms":11.8,"mel_ms":1.6,"encode_ms":6.6,"decode_ms":null,"engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"m4-max","backend":"metal","quant":"F16","sample":"ru-long-30s","sample_duration_s":30.0,"total_ms":20.4,"xrt_compute":1468.16,"wall_ms":26.4,"xrt_wall":1136.04,"load_ms":11.8,"mel_ms":3.8,"encode_ms":16.6,"decode_ms":null,"engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"m4-max","backend":"metal","quant":"F32","sample":"ru-long-10s","sample_duration_s":10.0,"total_ms":8.2,"xrt_compute":1217.01,"wall_ms":10.8,"xrt_wall":926.9,"load_ms":20.3,"mel_ms":1.3,"encode_ms":6.9,"decode_ms":null,"engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"m4-max","backend":"metal","quant":"F32","sample":"ru-long-30s","sample_duration_s":30.0,"total_ms":21.3,"xrt_compute":1406.64,"wall_ms":27.2,"xrt_wall":1104.59,"load_ms":20.3,"mel_ms":4.1,"encode_ms":17.2,"decode_ms":null,"engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"m4-max","backend":"metal","quant":"Q8_0","sample":"ru-long-10s","sample_duration_s":10.0,"total_ms":8.2,"xrt_compute":1212.15,"wall_ms":11.0,"xrt_wall":912.33,"load_ms":11.4,"mel_ms":1.7,"encode_ms":6.6,"decode_ms":null,"engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"m4-max","backend":"metal","quant":"Q8_0","sample":"ru-long-30s","sample_duration_s":30.0,"total_ms":21.0,"xrt_compute":1426.25,"wall_ms":27.0,"xrt_wall":1112.49,"load_ms":11.4,"mel_ms":4.2,"encode_ms":16.8,"decode_ms":null,"engine_sha":"2fdb5c95","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"ryzen-4750u","backend":"cpu","quant":"F16","sample":"ru-long-10s","sample_duration_s":10.0,"total_ms":211.4,"xrt_compute":47.3,"wall_ms":218.9,"xrt_wall":45.68,"load_ms":30.2,"mel_ms":7.3,"encode_ms":204.2,"decode_ms":null,"engine_sha":"8d00eb4a","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"ryzen-4750u","backend":"cpu","quant":"F16","sample":"ru-long-30s","sample_duration_s":30.0,"total_ms":818.3,"xrt_compute":36.66,"wall_ms":840.9,"xrt_wall":35.68,"load_ms":30.2,"mel_ms":22.0,"encode_ms":796.3,"decode_ms":null,"engine_sha":"8d00eb4a","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"ryzen-4750u","backend":"cpu","quant":"F32","sample":"ru-long-10s","sample_duration_s":10.0,"total_ms":191.0,"xrt_compute":52.37,"wall_ms":198.3,"xrt_wall":50.43,"load_ms":68.7,"mel_ms":7.0,"encode_ms":183.9,"decode_ms":null,"engine_sha":"8d00eb4a","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"ryzen-4750u","backend":"cpu","quant":"F32","sample":"ru-long-30s","sample_duration_s":30.0,"total_ms":730.1,"xrt_compute":41.09,"wall_ms":751.1,"xrt_wall":39.94,"load_ms":68.7,"mel_ms":22.5,"encode_ms":707.6,"decode_ms":null,"engine_sha":"8d00eb4a","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"ryzen-4750u","backend":"cpu","quant":"Q8_0","sample":"ru-long-10s","sample_duration_s":10.0,"total_ms":213.2,"xrt_compute":46.9,"wall_ms":221.2,"xrt_wall":45.22,"load_ms":58.5,"mel_ms":7.1,"encode_ms":206.1,"decode_ms":null,"engine_sha":"8d00eb4a","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"ryzen-4750u","backend":"cpu","quant":"Q8_0","sample":"ru-long-30s","sample_duration_s":30.0,"total_ms":801.7,"xrt_compute":37.42,"wall_ms":823.9,"xrt_wall":36.41,"load_ms":58.5,"mel_ms":20.9,"encode_ms":780.8,"decode_ms":null,"engine_sha":"8d00eb4a","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"ryzen-4750u","backend":"vulkan","quant":"F16","sample":"ru-long-10s","sample_duration_s":10.0,"total_ms":102.6,"xrt_compute":97.47,"wall_ms":107.4,"xrt_wall":93.11,"load_ms":40.7,"mel_ms":15.4,"encode_ms":87.2,"decode_ms":null,"engine_sha":"8d00eb4a","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"ryzen-4750u","backend":"vulkan","quant":"F16","sample":"ru-long-30s","sample_duration_s":30.0,"total_ms":283.6,"xrt_compute":105.79,"wall_ms":288.3,"xrt_wall":104.04,"load_ms":40.7,"mel_ms":30.0,"encode_ms":253.6,"decode_ms":null,"engine_sha":"8d00eb4a","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"ryzen-4750u","backend":"vulkan","quant":"F32","sample":"ru-long-10s","sample_duration_s":10.0,"total_ms":107.5,"xrt_compute":93.05,"wall_ms":112.4,"xrt_wall":88.95,"load_ms":48.3,"mel_ms":15.1,"encode_ms":92.4,"decode_ms":null,"engine_sha":"8d00eb4a","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"ryzen-4750u","backend":"vulkan","quant":"F32","sample":"ru-long-30s","sample_duration_s":30.0,"total_ms":301.9,"xrt_compute":99.38,"wall_ms":306.7,"xrt_wall":97.83,"load_ms":48.3,"mel_ms":30.9,"encode_ms":271.0,"decode_ms":null,"engine_sha":"8d00eb4a","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"ryzen-4750u","backend":"vulkan","quant":"Q8_0","sample":"ru-long-10s","sample_duration_s":10.0,"total_ms":102.5,"xrt_compute":97.56,"wall_ms":107.2,"xrt_wall":93.25,"load_ms":63.5,"mel_ms":15.4,"encode_ms":87.1,"decode_ms":null,"engine_sha":"8d00eb4a","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null}, + {"machine":"ryzen-4750u","backend":"vulkan","quant":"Q8_0","sample":"ru-long-30s","sample_duration_s":30.0,"total_ms":283.4,"xrt_compute":105.85,"wall_ms":288.3,"xrt_wall":104.06,"load_ms":63.5,"mel_ms":29.7,"encode_ms":253.8,"decode_ms":null,"engine_sha":"8d00eb4a","publication_profile":"langid-publication-v1","measured_on":"2026-10-07","thermal_gated":null} + ] +} diff --git a/docs/bindings.md b/docs/bindings.md index 9440f3621..87431da47 100644 --- a/docs/bindings.md +++ b/docs/bindings.md @@ -110,16 +110,17 @@ several sessions of one model each hold an active stream and interleave their feeds; a binding allows one active stream per model. Starting a stream takes the model's stream lease, and ending it (finalize, reset, a feed that leaves the stream FAILED, or dropping / closing the stream) releases it. -While the lease is held, `run`, `run_batch`, a new stream and a diarize run -on any session of that model raise `Busy` instead of waiting. A feed +While the lease is held, `run`, `run_batch`, a new stream, a diarize run and +a langid run on any session of that model raise `Busy` instead of waiting. A feed rejected before the native call (e.g. NaN input) keeps the lease. ## Roles and the DIARIZE session Each binding exposes the role mask as `Model.roles` (Python `frozenset[Role]`, -TypeScript `readonly ('asr' | 'diarize')[]`, Rust `Roles`, Swift `Roles` -option set). Status 20 surfaces as `UnsupportedRole` (`.unsupportedRole` in -Swift), including from capabilities on a model without ASR. +TypeScript `readonly ('asr' | 'diarize' | 'langid')[]`, Rust `Roles`, Swift +`Roles` option set). Status 20 surfaces as `UnsupportedRole` +(`.unsupportedRole` in Swift), including from capabilities on a model without +ASR. Status 21 surfaces as `InputTooShort` (`.inputTooShort` in Swift). A DIARIZE model (`docs/roles.md`) gets its own session type: `DiarizeSession` from `model.diarize_session()` (Python, Rust), `model.diarizeSession()` @@ -134,6 +135,27 @@ model-wide compute lock, waits behind other compute on the model, raises keeps its model alive, honours cancellation, and defers native frees that race an in-flight call. +## The LANGID session + +A LANGID model gets `LangIdSession` from `model.langid_session()` (Python, +Rust), `model.langIdSession()` (Swift) or `model.createLangIdSession()` +(TypeScript), with `max_audio_ms` / `maxAudioMs` as a session option, plus +`langid_info` / `langIdInfo` / `langidInfo` (sample rate, label count, +minimum audio) and the label table (codes, names, and alias lookup). +`run(pcm, allowed=…, top_k=…)` returns a copied-out result: candidates ranked +by `p` (index, code, name, `p`, `logit`), `n_allowed`, +`allowed_mass` and the scored `audio_ms`. + +`allowed` is the first caller-owned `const char * const *` input. Every +binding keeps the array and each encoded string alive until the native call +returns (TypeScript frees them only after its async worker call settles). +Omitted / `None` / `nil` / `null` passes NULL, meaning every label; an empty +list raises `InvalidArgument`. + +A langid run follows the same execution rules as an ASR or diarize run: +model-wide compute lock, `Busy` under a stream lease, results copied out +under the lock, cancellation, and deferred frees. + ## Raw text Every first-class binding exposes `raw_text` / `rawText` on the materialized diff --git a/docs/environment-variables.md b/docs/environment-variables.md index 053b0f7e6..4e1ae2237 100644 --- a/docs/environment-variables.md +++ b/docs/environment-variables.md @@ -87,6 +87,7 @@ its var is unset. Convention: `TRANSCRIBE__GGUF`. | `TRANSCRIBE_GIGAAM_GGUF` | `gigaam_workspace_release_smoke` | | `TRANSCRIBE_MULTITALKER_BUNDLE_GGUF` | `parakeet_multitalker_e2e_smoke` | | `TRANSCRIBE_SORTFORMER_GGUF` | `sortformer_diarize_unit`, `cli_diarize_smoke` | +| `TRANSCRIBE_ECAPA_TDNN_GGUF` | `ecapa_tdnn_real_smoke`, `cli_langid_smoke` (its real-model check) | | `TRANSCRIBE_COHERE_GGUF` | `cohere_real_smoke`, `cohere_e2e_smoke` | | `TRANSCRIBE_GRANITE5_CTC_GGUF` | `granite5_ctc_real_smoke`, `granite5_ctc_e2e_smoke` | | `TRANSCRIBE_WHISPER_GGUF` | `whisper_e2e_smoke`, `whisper_tokenize_parity` | diff --git a/docs/langid.md b/docs/langid.md new file mode 100644 index 000000000..20a9d8035 --- /dev/null +++ b/docs/langid.md @@ -0,0 +1,66 @@ +# Language ID + +A LANGID model (`transcribe_model_roles()` has `TRANSCRIBE_ROLE_LANGID`) +identifies the spoken language of a clip and returns its labels ranked by +probability. API: `include/transcribe/langid.h`. The shipped model is +[VoxLingua107 ECAPA-TDNN](models/lang-id-voxlingua107-ecapa.md), 107 labels. + +```c +#include "transcribe/langid.h" + +struct transcribe_model * m = NULL; +transcribe_model_load_file("lang-id-voxlingua107-ecapa-Q8_0.gguf", NULL, &m); + +struct transcribe_langid_session * lid = NULL; +transcribe_langid_session_init(m, NULL, &lid); + +const char * allowed[] = { "en", "de", "fr" }; +struct transcribe_langid_params lp; +transcribe_langid_params_init(&lp); +lp.allowed = allowed; +lp.n_allowed = 3; +transcribe_langid_run(lid, pcm, n_samples, &lp); /* 16 kHz mono float32 */ + +struct transcribe_langid_result r; +transcribe_langid_result_init(&r); +transcribe_langid_get_result(lid, &r); +struct transcribe_langid_candidate c; +transcribe_langid_candidate_init(&c); +transcribe_langid_get_candidate(lid, 0, &c); /* c.code, c.p, r.allowed_mass */ + +transcribe_langid_session_free(lid); +transcribe_model_free(m); +``` + +From the CLI: `transcribe-cli -m lang-id-voxlingua107-ecapa-Q8_0.gguf --allow en,de,fr --top 3 clip.wav`. + +## Parameters and results + +- **`allowed`** restricts the decision to the listed labels: candidates are + ranked by `p`, the softmax over the allowed set. The network still scores + every label. NULL with + `n_allowed == 0` means every label. A non-NULL list with `n_allowed <= 0`, + a NULL list with `n_allowed > 0`, or a NULL element is + `TRANSCRIBE_ERR_INVALID_ARG`; an unknown code is + `TRANSCRIBE_ERR_UNSUPPORTED_LANGUAGE`; duplicates count once. +- **`allowed_mass`** (`transcribe_langid_result`) is the share of the softmax + over every label that falls in the allowed set, 1 when unrestricted. A low + value means the unrestricted model puts most of its probability outside the + allowed set. +- **`top_k`** limits the returned candidates (0 = every allowed label). It + does not change `p` or `allowed_mass`. +- **Length.** Input longer than the session's `max_audio_ms` (default 30000) + is scored on its last `max_audio_ms`; `audio_ms` reports what was scored. + Scored audio shorter than `transcribe_langid_info::min_audio_ms` (500 ms) + is `TRANSCRIBE_ERR_INPUT_TOO_SHORT`. +- **Labels** are the model's own codes. VoxLingua107 uses `iw` (Hebrew), + `jw` (Javanese), `tl` (Filipino) and `no` (Norwegian Bokmal); the GGUF + stores `he`, `jv`, `fil` and `nb` as aliases, accepted in `allowed` and by + `transcribe_langid_label_index`. Results always report the model's code. + Codes are not canonicalized across models: match `code` against an ASR + model's `transcribe_capabilities::languages` yourself. + +The model is a closed-set classifier with no "unknown" or "no speech" label. + +Sessions, threading and model lifetime follow the rules shared by every role +(`docs/roles.md`). diff --git a/docs/migrating-to-0.4.md b/docs/migrating-to-0.4.md index 42c8ed57f..2f033d1d8 100644 --- a/docs/migrating-to-0.4.md +++ b/docs/migrating-to-0.4.md @@ -25,6 +25,12 @@ Segments are byte-identical to 0.3 for the same audio, preset and backend. Speaker-attributing ASR models (granite, moss, multitalker parakeet) are unchanged; see `docs/roles.md`. +## New: the LANGID role + +Language identification is a new role (`include/transcribe/langid.h`, +`docs/langid.md`) with a new status, `TRANSCRIBE_ERR_INPUT_TOO_SHORT` (21). +It is new API, so nothing migrates. + ## Capabilities are ASR-only `transcribe_model_get_capabilities` returns `TRANSCRIBE_ERR_UNSUPPORTED_ROLE` @@ -46,8 +52,9 @@ from the role's own query (`transcribe_diarize_get_info`). ## Language bindings -- New `Model.roles`, `DiarizeSession` and an `UnsupportedRole` error for - status 20 (`.unsupportedRole` in Swift); see `docs/bindings.md`. The +- New `Model.roles`, `DiarizeSession`, `LangIdSession`, an `UnsupportedRole` + error for status 20 (`.unsupportedRole` in Swift) and an `InputTooShort` + error for status 21; see `docs/bindings.md`. The Sortformer ASR-path extension (`SortformerStreamOptions` and equivalents) is removed. Sortformer runs through a diarize session instead: @@ -68,7 +75,7 @@ from the role's own query (`transcribe_diarize_get_info`). (e.g. NaN) keeps the lease and the stream stays usable. - **Rust:** `Model::capabilities()` returns `Result` (`Err(Error::UnsupportedRole)` on a model without ASR). `ExtSlot` gains - `DiarizeRun` and `AbiStruct` gains the three diarize structs; both are now + `DiarizeRun` and `AbiStruct` gains the diarize and langid structs; both are now `#[non_exhaustive]`, so later roles add variants without another break. - **Swift:** `Model.capabilities` is `get throws`. - **Python:** calls on one model now serialize (they used to race), and a diff --git a/docs/models/lang-id-voxlingua107-ecapa.md b/docs/models/lang-id-voxlingua107-ecapa.md new file mode 100644 index 000000000..0ec1cd9b7 --- /dev/null +++ b/docs/models/lang-id-voxlingua107-ecapa.md @@ -0,0 +1,206 @@ +# VoxLingua107 ECAPA-TDNN (language ID) + + +Upstream: [`speechbrain/lang-id-voxlingua107-ecapa`](https://huggingface.co/speechbrain/lang-id-voxlingua107-ecapa) at [`0253049`](https://huggingface.co/speechbrain/lang-id-voxlingua107-ecapa/commit/0253049). + +Spoken language identification over 107 languages: SpeechBrain's +ECAPA-TDNN trained on VoxLingua107. NOT a transcription model: a run +returns the model's language labels ranked by probability, optionally +restricted to a caller-chosen set. Takes 16 kHz mono WAV; scores up to +the last 30 s of a clip. + + +## What it's for + +Spoken language identification over 107 languages. This is **not a +transcription model**. It serves the LANGID role +(`include/transcribe/langid.h`, [`docs/langid.md`](../langid.md)): a run +returns the model's language labels ranked by probability, optionally +restricted to a caller-chosen set, plus `allowed_mass`, the unrestricted +probability that set captured. + +See SpeechBrain's [model card](https://huggingface.co/speechbrain/lang-id-voxlingua107-ecapa) +for training data and the label list. + + +Licensed Apache-2.0. Ported from upstream commit [`0253049`](https://huggingface.co/speechbrain/lang-id-voxlingua107-ecapa/commit/0253049), pinned 2026-10-05. Validated against the SpeechBrain reference at transcribe.cpp commit [`aa4d0f47`](https://github.com/handy-computer/transcribe.cpp/tree/aa4d0f47) on 2026-10-08. + + +## Download + + +| Quantization | Download | Size | Top-1 accuracy (FLEURS multilingual) | +| --- | --- | ---: | ---: | +| F32 | [lang-id-voxlingua107-ecapa-F32.gguf](https://huggingface.co/handy-computer/lang-id-voxlingua107-ecapa-gguf/resolve/main/lang-id-voxlingua107-ecapa-F32.gguf) | 85 MB | 85.23% | +| F16 | [lang-id-voxlingua107-ecapa-F16.gguf](https://huggingface.co/handy-computer/lang-id-voxlingua107-ecapa-gguf/resolve/main/lang-id-voxlingua107-ecapa-F16.gguf) | 45 MB | 85.17% | +| Q8_0 | [lang-id-voxlingua107-ecapa-Q8_0.gguf](https://huggingface.co/handy-computer/lang-id-voxlingua107-ecapa-gguf/resolve/main/lang-id-voxlingua107-ecapa-Q8_0.gguf) | 27 MB | 86.30% | + + + +Top-1 accuracy on FLEURS multilingual (3,000 utterances), scored on cpu. Measured at transcribe.cpp `2fdb5c95` on 2026-10-07. + + + +Open-set top-1 accuracy over all 107 labels, the mean over 15 FLEURS +languages, on the first 5 s of each clip without silence trimming. +Agreement with the SpeechBrain reference, per GGUF, is on the +transcribe.cpp model page. + + +## Accuracy + + +| GGUF | Top-1 accuracy (95% CI) | Top-1 agreement with SpeechBrain | Max abs logit difference | +| --- | ---: | ---: | ---: | +| F32 | 85.23% (84.07-86.47) | 12000 / 12000 | 7.3e-05 | +| F16 | 85.17% (84.00-86.43) | 11986 / 12000 | 0.095 | +| Q8_0 | 86.30% (85.17-87.50) | 11705 / 12000 | 2.9 | + +Measured at transcribe.cpp `2fdb5c95` on 2026-10-07. + + +Agreement counts the scored sweep's top-1 decisions that match the +SpeechBrain reference on the same audio, over the 3 / 5 / 10 s / full crops. + +## Quick Start + +```bash +cmake -B build +cmake --build build + +build/bin/transcribe-cli \ + -m models/lang-id-voxlingua107-ecapa/lang-id-voxlingua107-ecapa-Q8_0.gguf \ + --allow en,de,fr --top 3 audio.wav +# language: de index=18 p=0.999998 +# ... +``` + +If your audio is not already 16 kHz mono WAV, convert it first: + +```bash +ffmpeg -i input.mp3 -ar 16000 -ac 1 output.wav +``` + +From the C API, use the LANGID role (`include/transcribe/langid.h`): +`transcribe_langid_session_init`, `transcribe_langid_run`, then +`transcribe_langid_get_result` / `transcribe_langid_get_candidate`. Labels are +the model's own codes (`iw`, `jw`, `tl`, `no`; `he`, `jv`, `fil`, `nb` are +accepted as aliases). + +## Performance + +### Apple M4 Max + + +Compute latency (mel + encode), speedup over realtime in parentheses; profile `langid-publication-v1`: mean over 20 iterations after 3 warmup. + +| Backend | Sample | F32 | F16 | Q8_0 | +| ------- | ------------------- | -----------------: | -----------------: | -----------------: | +| Metal | ru-long-10s (10.0s) | 8.2 ms (1217.01×) | 8.3 ms (1208.90×) | 8.2 ms (1212.15×) | +| Metal | ru-long-30s (30.0s) | 21.3 ms (1406.64×) | 20.4 ms (1468.16×) | 21.0 ms (1426.25×) | +| CPU | ru-long-10s (10.0s) | 123.8 ms (80.77×) | 78.9 ms (126.69×) | 80.4 ms (124.38×) | +| CPU | ru-long-30s (30.0s) | 382.3 ms (78.48×) | 241.8 ms (124.06×) | 247.1 ms (121.43×) | + +Apple M4 Max: transcribe.cpp `2fdb5c95` on 2026-10-07. + + +### AMD Ryzen 7 PRO 4750U + + +Compute latency (mel + encode), speedup over realtime in parentheses; profile `langid-publication-v1`: mean over 20 iterations after 3 warmup. + +| Backend | Sample | F32 | F16 | Q8_0 | +| ------- | ------------------- | ----------------: | -----------------: | -----------------: | +| Vulkan | ru-long-10s (10.0s) | 107.5 ms (93.05×) | 102.6 ms (97.47×) | 102.5 ms (97.56×) | +| Vulkan | ru-long-30s (30.0s) | 301.9 ms (99.38×) | 283.6 ms (105.79×) | 283.4 ms (105.85×) | +| CPU | ru-long-10s (10.0s) | 191.0 ms (52.37×) | 211.4 ms (47.30×) | 213.2 ms (46.90×) | +| CPU | ru-long-30s (30.0s) | 730.1 ms (41.09×) | 818.3 ms (36.66×) | 801.7 ms (37.42×) | + +AMD Ryzen 7 PRO 4750U (Radeon RADV RENOIR): transcribe.cpp `8d00eb4a` on 2026-10-07. + + +Benchmark reproduction (`tools/transcribe-bench` is ASR-only): + +```bash +uv run --project scripts/envs/ecapa_tdnn scripts/langid/bench.py --profile \ + --library build-shared/src/libtranscribe.dylib +``` + +## Numerical Validation + +transcribe.cpp is validated tensor-by-tensor against SpeechBrain on eight +FLEURS clips (`samples/fleurs-*.wav`): every stage tensor gated in +`tests/tolerances/ecapa_tdnn.json` (front end, every encoder block, pooling, +embedding, logits) falls within tolerance, and the top-1 label matches the +reference on every clip. + +| Field | Value | +| --- | --- | +| Reference | SpeechBrain 1.1.1, `speechbrain/lang-id-voxlingua107-ecapa` | +| Dump script | `scripts/dump_reference_ecapa_tdnn_speechbrain.py` | +| Manifest | `tests/golden/ecapa_tdnn/lang-id-voxlingua107-ecapa.manifest.json` | +| Command | `uv run scripts/validate.py all --family ecapa_tdnn` | +| Dataset gate | `scripts/langid/score.py --ref` (FLEURS, every crop) | + +## Known Limitations + +- **Closed set.** There is no "unknown" or "no speech" label: silence, music + and out-of-set languages still get a top candidate. +- **No Cantonese label**: Cantonese is classified as `zh`. +- **Minimum 500 ms** of scored audio (`TRANSCRIBE_ERR_INPUT_TOO_SHORT`). +- **Scores at most the last 30 s** by default + (`transcribe_langid_session_params::max_audio_ms`). +- **Q8_0 weights are widened to F16 at load**, so Q8_0 computes and uses + memory like F16. + +## Reproduction + +### Convert + +Loads the SpeechBrain checkpoint through SpeechBrain from the Hugging Face +cache and writes F32 only. + +```bash +uv run --project scripts/envs/ecapa_tdnn \ + scripts/convert-ecapa_tdnn.py speechbrain/lang-id-voxlingua107-ecapa +``` + +### Quantize + +```bash +build/bin/transcribe-quantize \ + models/lang-id-voxlingua107-ecapa/lang-id-voxlingua107-ecapa-F32.gguf \ + models/lang-id-voxlingua107-ecapa/lang-id-voxlingua107-ecapa-Q8_0.gguf \ + --quant Q8_0 +# repeat with F16 +``` + +### Validate + +```bash +uv run scripts/validate.py all --family ecapa_tdnn +``` + +### Accuracy acceptance + +One reference sweep, then one C++ sweep per shipped GGUF (shown for F32; +repeat with F16 and Q8_0, passing `--report-only` to `score.py`, which gates +F32 only). `score.py --json` also writes the `.agreement.json` beside the score. + +```bash +uv run --project scripts/envs/ecapa_tdnn scripts/langid/ingest.py fleurs --lang all +uv run --project scripts/envs/ecapa_tdnn scripts/langid/run.py --engine speechbrain \ + --model speechbrain/lang-id-voxlingua107-ecapa \ + --manifest samples/langid/fleurs-*.manifest.jsonl --crops 3,5,10,full \ + --out reports/langid/ref-speechbrain-untrimmed.jsonl +uv run --project scripts/envs/ecapa_tdnn scripts/langid/run.py --engine cpp \ + --model models/lang-id-voxlingua107-ecapa/lang-id-voxlingua107-ecapa-F32.gguf \ + --library build-shared/src/libtranscribe.dylib \ + --manifest samples/langid/fleurs-*.manifest.jsonl --crops 3,5,10,full \ + --out reports/langid/cpp-f32-untrimmed.jsonl +uv run scripts/langid/score.py reports/langid/cpp-f32-untrimmed.jsonl \ + --ref reports/langid/ref-speechbrain-untrimmed.jsonl \ + --json reports/langid/lang-id-voxlingua107-ecapa-F32.fleurs-mul.score.json +uv run scripts/catalog/ingest_accuracy.py --models lang-id-voxlingua107-ecapa +uv run scripts/catalog/render.py +``` diff --git a/docs/porting/families/_intake-schema.json b/docs/porting/families/_intake-schema.json index 7d7d33ae7..6e80cb879 100644 --- a/docs/porting/families/_intake-schema.json +++ b/docs/porting/families/_intake-schema.json @@ -85,7 +85,7 @@ "type": "array", "items": { "type": "string", - "enum": ["encoder-transducer", "encoder-decoder", "audio-llm", "encoder-ctc", "encoder-diarizer"] + "enum": ["encoder-transducer", "encoder-decoder", "audio-llm", "encoder-ctc", "encoder-diarizer", "encoder-classifier"] }, "description": "Heuristic matches against known patterns. Human selects one in architecture_pattern." }, @@ -147,12 +147,12 @@ "fft_size": {"type": ["integer", "null"]}, "window": { "type": ["string", "null"], - "enum": ["hann_periodic", "hann_symmetric", "hamming", "blackman", null], + "enum": ["hann_periodic", "hann_symmetric", "hamming", "hamming_periodic", "blackman", null], "description": "Window type. 'hann_periodic' and 'hann_symmetric' are a classic mismatch source." }, "normalization": { "type": ["string", "null"], - "enum": ["per_feature", "global", "per_utterance", "none", null] + "enum": ["per_feature", "global", "per_utterance", "sentence_mean", "none", null] }, "preemphasis": { "type": ["number", "null"], @@ -257,7 +257,7 @@ }, "metric": { "type": "string", - "enum": ["wer", "cer", "bleu", "other"] + "enum": ["wer", "cer", "bleu", "accuracy", "other"] }, "score": { "type": ["number", "null"], @@ -288,7 +288,7 @@ }, "architecture_pattern": { "type": ["string", "null"], - "enum": ["encoder-transducer", "encoder-decoder", "audio-llm", "encoder-ctc", "encoder-diarizer", null] + "enum": ["encoder-transducer", "encoder-decoder", "audio-llm", "encoder-ctc", "encoder-diarizer", "encoder-classifier", null] }, "known_risks": { "type": "array", diff --git a/docs/porting/families/ecapa_tdnn.md b/docs/porting/families/ecapa_tdnn.md new file mode 100644 index 000000000..12ea2266c --- /dev/null +++ b/docs/porting/families/ecapa_tdnn.md @@ -0,0 +1,115 @@ +# ECAPA-TDNN (language ID) + +Status: supported (LANGID role). Ported from `handy-computer/langid.cpp` +(`96af3a5`). + +ECAPA-TDNN is SpeechBrain's speaker / language embedder (architecture pattern +`encoder-classifier`): SpeechBrain Fbank front end, a TDNN block, three +SERes2Net blocks, multi-layer feature aggregation, attentive statistics +pooling, a 256-d embedding and a two-layer classifier over the VoxLingua107 +labels. It is not a transcription model. The family only produces logits; the +LANGID role dispatcher (`src/transcribe-langid.cpp`) owns the crop to the +last `max_audio_ms`, the allowed set, softmax, `allowed_mass`, ranking and +top-k (`docs/langid.md`). + +Acceptance: tensor parity on eight FLEURS clips (`validate.py`), and top-1 +decision parity with SpeechBrain on FLEURS (15 languages x 200 utterances x +3 / 5 / 10 s / full crops, C++ F32 on CPU). Shipped matrix: F32 + F16 + Q8_0. +Each GGUF's accuracy and agreement live in the catalog and are rendered on +[the model page](../../models/lang-id-voxlingua107-ecapa.md#accuracy). The +loader widens Q8_0 weights to F16 (`widen_q8_0_weights` in +`src/arch/ecapa_tdnn/model.cpp`). + +## Identity + +- Family key: `ecapa_tdnn` +- Upstream architecture: `speechbrain.lobes.models.ECAPA_TDNN.ECAPA_TDNN` + `speechbrain.lobes.models.Xvector.Classifier` +- Hugging Face repo: `speechbrain/lang-id-voxlingua107-ecapa` +- Hugging Face revision: `0253049ae131d6a4be1c4f0d8b0ff483a0f8c8e9` +- License: Apache-2.0 +- Variants: `lang-id-voxlingua107-ecapa` (107 labels; legacy codes `iw`, `jw`, `tl`, `no` with aliases `he`, `jv`, `fil`, `nb`) + +## References + +- Canonical reference: SpeechBrain 1.1.1 `EncoderClassifier` (the only + implementation; `hyperparams.yaml` instantiates SpeechBrain classes by name), + torch 2.13.0, CPU, one thread. +- Instrumented reference: `scripts/dump_reference_ecapa_tdnn_speechbrain.py` + (forward hooks on one `classify_batch`; asserts every hook fires once). + +## GGUF contract + +Keys: `stt.variant`, `stt.frontend.*` (`window = hamming_periodic`, +`pad_mode = constant`, `log_clamp_min = 1e-10`, `top_db = 80`, +`normalize = sentence_mean`, plus the usual sizes), `stt.ecapa_tdnn.*` +(channels, kernel sizes, dilations, res2net scale, SE / attention / +embedding / classifier widths, `asp_eps`, leaky slope) and +`stt.langid.labels.{codes,names,aliases}`. Integer scalars are uint32. + +The converter applies four exact rewrites: BatchNorm to a `scale` / `shift` +pair (TDNNBlock is conv -> ReLU -> BN, so it cannot fold backward); the +three activation-free BNs folded forward into `fc`, `cls.l1`, `cls.out`; and +the ASP and MFA input weights split per concatenated operand. Tensor names +follow the `tools/transcribe-quantize` rules (`.bias` / `.bn.` F32, +`.conv.weight` for k>1 tap-major kernels F32 / F16, every other `.weight` +quantizable, `frontend.mel_filterbank` F32), so the quant policy has no +ECAPA entries. +The full catalogue is in `src/arch/ecapa_tdnn/weights.h`. + +The front end is the shared `transcribe-mel` (`src/transcribe-mel.cpp`) +with `window_type = "hamming_periodic"` and `normalize = "sentence_mean"` +(`10*log10`, 80 dB top-db floor, per-bin mean subtraction, no frame drop). +The filterbank is SpeechBrain's own, captured by the converter and stored in +the GGUF; it is never rebuilt. + +## Commands + +Reference dumps and tensor parity: + +```bash +uv run scripts/validate.py all --family ecapa_tdnn +``` + +Conversion (F32 only; quantize afterwards): + +```bash +uv run --project scripts/envs/ecapa_tdnn scripts/convert-ecapa_tdnn.py speechbrain/lang-id-voxlingua107-ecapa +build/bin/transcribe-quantize models/lang-id-voxlingua107-ecapa/lang-id-voxlingua107-ecapa-F32.gguf \ + models/lang-id-voxlingua107-ecapa/lang-id-voxlingua107-ecapa-Q8_0.gguf --quant Q8_0 +``` + +Accuracy and decision parity (FLEURS from the HF cache): + +```bash +uv run --project scripts/envs/ecapa_tdnn scripts/langid/ingest.py fleurs --lang all +uv run --project scripts/envs/ecapa_tdnn scripts/langid/run.py --engine speechbrain ... --out reports/langid/ref-speechbrain-untrimmed.jsonl +uv run --project scripts/envs/ecapa_tdnn scripts/langid/run.py --engine cpp --library build-shared/src/libtranscribe.dylib ... --out reports/langid/cpp-f32-untrimmed.jsonl +uv run scripts/langid/score.py reports/langid/cpp-f32-untrimmed.jsonl --ref reports/langid/ref-speechbrain-untrimmed.jsonl \ + --json reports/langid/lang-id-voxlingua107-ecapa-F32.fleurs-mul.score.json +uv run scripts/catalog/ingest_accuracy.py --models lang-id-voxlingua107-ecapa +``` + +The catalog rows follow the `langid-publication-v1` profile +(`catalog/_benchmark_profiles.json`): the open-set mean on the 5 s untrimmed +crop of each shipped GGUF's sweep, with that sweep's agreement. + +`score.py --ref` fails on any top-1 disagreement with the reference sweep or +any row the two sweeps do not share, and refuses sweeps whose labels, recipe, +manifest hashes or checkpoint revision differ. With `--json` it writes the +score and, beside it, the agreement. Only F32 is gated: score the F16 and Q8_0 +sweeps with `--report-only`, which still writes their agreement. + +Benchmarks: `scripts/langid/bench.py --profile` (`tools/transcribe-bench` is +ASR-only) writes bench-driver reports under `reports/perf/`, which +`scripts/catalog/ingest_perf.py` folds into the catalog. + +## Capability Validation + +| Capability | Mode | Command / test | Expected observable | Target | Status | +|------------|------|----------------|---------------------|--------|--------| +| Language ID | open set | `build/bin/transcribe-cli -m models/lang-id-voxlingua107-ecapa/lang-id-voxlingua107-ecapa-F32.gguf samples/fleurs-ja.wav` | `language: ja` | MUST PASS | PASS | +| Language ID | allowed set | `... --allow en,de samples/fleurs-de.wav` | `language: de`, 2 candidates | MUST PASS | PASS | +| Crop | 30 s window | `transcribe_ecapa_tdnn_real_smoke` (`ru-long.wav`) | `audio_ms == 30000` | MUST PASS | PASS | +| Minimum length | 500 ms | `transcribe_ecapa_tdnn_smoke`, `transcribe_langid_dispatch_unit` | `INPUT_TOO_SHORT` below 500 ms | MUST PASS | PASS | +| Transcribe / translate / timestamps / streaming | n/a | `transcribe_session_init` | `UNSUPPORTED_ROLE` | OUT OF SCOPE — not an ASR model | SKIP — not exposed by runtime | +| Batch (offline) | n/a | `transcribe-cli --batch` | refused: ASR-only | OUT OF SCOPE — no batch entry points for new roles in v1 | ACCEPTED GAP — one clip per call | diff --git a/docs/roles.md b/docs/roles.md index 7330de12a..e6f3109e3 100644 --- a/docs/roles.md +++ b/docs/roles.md @@ -9,6 +9,7 @@ A loaded model serves one or more **roles**, the kinds of work it can do. |---|---|---|---| | `TRANSCRIBE_ROLE_ASR` | `include/transcribe.h` | `transcribe_session` | text, timestamps, speaker-attributed segments | | `TRANSCRIBE_ROLE_DIARIZE` | `include/transcribe/diarize.h` | `transcribe_diarize_session` | who spoke when | +| `TRANSCRIBE_ROLE_LANGID` | `include/transcribe/langid.h` | `transcribe_langid_session` | which language (`docs/langid.md`) | Every role follows the same rules: diff --git a/examples/cli/CMakeLists.txt b/examples/cli/CMakeLists.txt index 07a681522..43eff1e6f 100644 --- a/examples/cli/CMakeLists.txt +++ b/examples/cli/CMakeLists.txt @@ -8,6 +8,7 @@ add_executable(transcribe-cli main.cpp asr.cpp diarize.cpp + langid.cpp ) target_link_libraries(transcribe-cli diff --git a/examples/cli/cli.h b/examples/cli/cli.h index fbe4b5b50..e4c24ec50 100644 --- a/examples/cli/cli.h +++ b/examples/cli/cli.h @@ -1,7 +1,7 @@ // cli.h - shared declarations for the transcribe-cli example. // // main.cpp parses arguments, owns process setup (log sink, output file), and -// routes each model to its role's driver (asr.cpp, diarize.cpp). Shared +// routes each model to its role's driver (asr.cpp, diarize.cpp, langid.cpp). Shared // helpers live in namespace transcribe_cli next to the WAV loader in // examples/common. @@ -103,6 +103,14 @@ struct cli_args { // explicit K. Silently ignored by families without // supports_spec_decode. Set by --spec-k-drafts N. int spec_k_drafts = -1; + + // LANGID role (language ID models). --allow restricts the decision to + // these labels (codes or aliases); --top keeps the best N candidates + // (0 = every allowed label); --max-audio-ms scores the last N ms + // (0 = library default, 30000). + std::vector langid_allow; + int langid_top_k = 0; + int langid_max_audio_ms = 0; }; // -o/--output: write `text` (newline-terminated) to `output` if non-null. @@ -129,4 +137,10 @@ int run_diarize_file(const cli_args & args, double duration_s, std::ofstream * output); +// With -o, writes the candidate lines. +int run_langid_file(const cli_args & args, + transcribe_model * model, + const std::vector & pcm, + std::ofstream * output); + } // namespace transcribe_cli diff --git a/examples/cli/langid.cpp b/examples/cli/langid.cpp new file mode 100644 index 000000000..566eedef9 --- /dev/null +++ b/examples/cli/langid.cpp @@ -0,0 +1,101 @@ +// langid.cpp - transcribe-cli LANGID driver: ranked language candidates for +// one file through include/transcribe/langid.h. +// +// Output lines are stable (scripts/validate.py parses `language:`): +// language: index= p=

+// candidate: index= p=

logit= + +#include "transcribe/langid.h" + +#include "cli.h" + +#include +#include +#include +#include +#include + +int transcribe_cli::run_langid_file(const cli_args & args, + transcribe_model * model, + const std::vector & pcm, + std::ofstream * output) { + if (args.stream_chunk_ms > 0) { + std::fprintf(stderr, "stream: the language ID path has no streaming entry point; drop --stream-chunk-ms\n"); + transcribe_model_free(model); + return EXIT_FAILURE; + } + + transcribe_langid_session_params sp; + transcribe_langid_session_params_init(&sp); + sp.n_threads = args.n_threads; + sp.max_audio_ms = args.langid_max_audio_ms; + transcribe_langid_session * session = nullptr; + transcribe_status st = transcribe_langid_session_init(model, &sp, &session); + if (st != TRANSCRIBE_OK) { + std::fprintf(stderr, "langid session init: %s\n", transcribe_status_string(st)); + transcribe_model_free(model); + return EXIT_FAILURE; + } + + std::vector allowed; + for (const std::string & code : args.langid_allow) { + allowed.push_back(code.c_str()); + } + transcribe_langid_params lp; + transcribe_langid_params_init(&lp); + lp.allowed = allowed.empty() ? nullptr : allowed.data(); + lp.n_allowed = static_cast(allowed.size()); + lp.top_k = args.langid_top_k; + + // --repeat N runs transcribe_langid_run() N times for steady-state perf + // measurements. + for (int r = 0; r < args.repeat; ++r) { + st = transcribe_langid_run(session, pcm.data(), static_cast(pcm.size()), &lp); + if (st != TRANSCRIBE_OK) { + break; + } + } + std::printf("run: %s\n", transcribe_status_string(st)); + bool output_ok = true; + double scored_s = 0.0; // only the last max_audio_ms is scored, not the whole file + if (st == TRANSCRIBE_OK) { + transcribe_langid_result res; + transcribe_langid_result_init(&res); + transcribe_langid_get_result(session, &res); + scored_s = static_cast(res.audio_ms) / 1000.0; + + std::string lines; + for (int i = 0; i < res.n_candidates; ++i) { + transcribe_langid_candidate c; + transcribe_langid_candidate_init(&c); + transcribe_langid_get_candidate(session, i, &c); + if (i == 0) { + std::printf("language: %s index=%d p=%.6f\n", c.code, c.index, static_cast(c.p)); + std::printf(" name: %s\n", c.name); + std::printf(" allowed_mass: %.6f (%d allowed)\n", static_cast(res.allowed_mass), + res.n_allowed); + std::printf(" audio_ms: %lld\n", static_cast(res.audio_ms)); + } + char line[160]; + std::snprintf(line, sizeof(line), "candidate: %d %s index=%d p=%.6f logit=%.4f\n", i + 1, c.code, c.index, + static_cast(c.p), static_cast(c.logit)); + std::printf(" %s", line); + lines += line; + } + output_ok = write_output_file(output, args.output_path, lines.c_str()); + } + + transcribe_timings tm; + transcribe_timings_init(&tm); + transcribe_langid_get_timings(session, &tm); + const double total_ms = tm.mel_ms + tm.encode_ms; + if (total_ms > 0.0 && scored_s > 0.0) { + std::printf(" realtime: %.0fx (%.1f ms for %.1f s; mel %.1f ms, encode %.1f ms)\n", + scored_s * 1000.0 / total_ms, total_ms, scored_s, static_cast(tm.mel_ms), + static_cast(tm.encode_ms)); + } + + transcribe_langid_session_free(session); + transcribe_model_free(model); + return st == TRANSCRIBE_OK && output_ok ? EXIT_SUCCESS : EXIT_FAILURE; +} diff --git a/examples/cli/main.cpp b/examples/cli/main.cpp index 0d5f4c0b9..95f0bb418 100644 --- a/examples/cli/main.cpp +++ b/examples/cli/main.cpp @@ -83,6 +83,9 @@ void print_usage(const char * argv0) { " --diarize (moss/granite-plus) speaker attribution: segments carry\n" " speaker ids; granite-plus requests its speaker task\n" " --no-diarize disable speaker attribution (the library default)\n" + " --allow CODES (language ID) comma-separated labels to choose from\n" + " --top N (language ID) print the best N candidates (0 = all)\n" + " --max-audio-ms N (language ID) score the last N ms (0 = 30000)\n" " --raw-tokens keep <|...|> control tokens in output text\n" " --stream-chunk-ms N single-file: drive the streaming API by feeding\n" " N-ms PCM slices; requires model to advertise\n" @@ -418,6 +421,41 @@ bool parse_args(int argc, char ** argv, cli_args & out) { } else if (a == "--no-pnc") { out.canary_pnc = false; out.canary_pnc_set = true; + } else if (a == "--allow") { + const char * v = take_value(a.c_str()); + if (!v) { + return false; + } + const std::string list = v; + size_t pos = 0; + while (pos <= list.size()) { + const size_t comma = list.find(',', pos); + const size_t end = comma == std::string::npos ? list.size() : comma; + if (end > pos) { + out.langid_allow.push_back(list.substr(pos, end - pos)); + } + pos = end + 1; + } + if (out.langid_allow.empty()) { + std::fprintf(stderr, "error: --allow needs at least one label\n"); + return false; + } + } else if (a == "--top") { + const char * v = take_value(a.c_str()); + if (!v) { + return false; + } + out.langid_top_k = std::atoi(v); + if (out.langid_top_k < 0) { + std::fprintf(stderr, "error: --top must be >= 0\n"); + return false; + } + } else if (a == "--max-audio-ms") { + const char * v = take_value(a.c_str()); + if (!v) { + return false; + } + out.langid_max_audio_ms = std::atoi(v); } else if (a == "--diarize") { out.diarize = true; out.diarize_set = true; @@ -583,9 +621,13 @@ int run_file(const cli_args & args, std::ofstream * output) { std::printf(" license: %s\n", lic); } - if ((transcribe_model_roles(model) & TRANSCRIBE_ROLE_ASR) != 0) { + const uint32_t roles = transcribe_model_roles(model); + if ((roles & TRANSCRIBE_ROLE_ASR) != 0) { return transcribe_cli::run_asr_file(args, model, pcm, duration_s, output); } + if ((roles & TRANSCRIBE_ROLE_LANGID) != 0) { + return transcribe_cli::run_langid_file(args, model, pcm, output); + } return transcribe_cli::run_diarize_file(args, model, pcm, duration_s, output); } diff --git a/include/transcribe.abihash b/include/transcribe.abihash index 8fce33b22..5e88d1778 100644 --- a/include/transcribe.abihash +++ b/include/transcribe.abihash @@ -1 +1 @@ -bd3273dabb25a1fe +d544d70a2b5acf50 diff --git a/include/transcribe.h b/include/transcribe.h index c36df7b75..ec016f3c4 100644 --- a/include/transcribe.h +++ b/include/transcribe.h @@ -315,6 +315,9 @@ typedef enum { TRANSCRIBE_ERR_OUTPUT_REPETITION = 19, /* The model does not serve the role this call needs; see transcribe_model_roles(). */ TRANSCRIBE_ERR_UNSUPPORTED_ROLE = 20, + /* The audio is shorter than the role's minimum (e.g. language ID's + * transcribe_langid_info::min_audio_ms). Nothing was computed. */ + TRANSCRIBE_ERR_INPUT_TOO_SHORT = 21, } transcribe_status; /* @@ -391,6 +394,12 @@ typedef enum { TRANSCRIBE_ABI_DIARIZE_INFO = 16, TRANSCRIBE_ABI_DIARIZE_SESSION_PARAMS = 17, TRANSCRIBE_ABI_DIARIZE_PARAMS = 18, + /* include/transcribe/langid.h */ + TRANSCRIBE_ABI_LANGID_INFO = 19, + TRANSCRIBE_ABI_LANGID_SESSION_PARAMS = 20, + TRANSCRIBE_ABI_LANGID_PARAMS = 21, + TRANSCRIBE_ABI_LANGID_RESULT = 22, + TRANSCRIBE_ABI_LANGID_CANDIDATE = 23, } transcribe_abi_struct; /* sizeof / alignof of the selected public struct, or 0 for an unknown id. @@ -1362,6 +1371,7 @@ TRANSCRIBE_API void transcribe_capabilities_init(struct transcribe_capabilities typedef enum { TRANSCRIBE_ROLE_ASR = 1u << 0, TRANSCRIBE_ROLE_DIARIZE = 1u << 1, + TRANSCRIBE_ROLE_LANGID = 1u << 2, } transcribe_role; /* diff --git a/include/transcribe/extensions.h b/include/transcribe/extensions.h index 16c5684ad..f64cd75d6 100644 --- a/include/transcribe/extensions.h +++ b/include/transcribe/extensions.h @@ -15,6 +15,7 @@ #include "transcribe.h" #include "transcribe/diarize.h" +#include "transcribe/langid.h" #include "transcribe/moonshine_streaming.h" #include "transcribe/parakeet.h" #include "transcribe/sortformer.h" diff --git a/include/transcribe/langid.h b/include/transcribe/langid.h new file mode 100644 index 000000000..99dfa5365 --- /dev/null +++ b/include/transcribe/langid.h @@ -0,0 +1,142 @@ +/* + * include/transcribe/langid.h - LANGID role: which language is spoken. + * + * Includes transcribe.h; safe to include in C or C++ TUs. For models whose + * product is a language decision (transcribe_model_roles() has + * TRANSCRIBE_ROLE_LANGID). + * + * Open a transcribe_langid_session on a loaded model, run it on a clip, + * then read the ranked candidates. Labels are the model's own codes ("en", + * "iw", "jw"); there is no canonicalization, so match a code against an ASR + * model's transcribe_capabilities::languages yourself. Usage guide: + * docs/langid.md. Threading and lifetime rules: docs/roles.md. + */ + +#ifndef TRANSCRIBE_LANGID_H +#define TRANSCRIBE_LANGID_H + +#include "transcribe.h" + +#ifdef __cplusplus +extern "C" { +#endif + +struct transcribe_langid_session; + +/* Static facts about a language ID model. */ +struct transcribe_langid_info { + uint64_t struct_size; + int32_t sample_rate; /* input PCM rate (16000) */ + int32_t n_labels; /* label indices are [0, n_labels) */ + int32_t min_audio_ms; /* shorter scored audio is TRANSCRIBE_ERR_INPUT_TOO_SHORT */ +}; + +struct transcribe_langid_session_params { + uint64_t struct_size; + int32_t n_threads; /* 0 = library default */ + /* Longer input is scored on its LAST max_audio_ms. 0 = 30000. A value + * below transcribe_langid_info::min_audio_ms is INVALID_ARG. */ + int32_t max_audio_ms; +}; + +struct transcribe_langid_params { + uint64_t struct_size; + /* Restrict the decision to these labels (codes or aliases, borrowed for + * the call). NULL with n_allowed == 0 means every label; that is the + * only "all" form. A non-NULL list with n_allowed <= 0, a NULL list + * with n_allowed > 0, or a NULL element is INVALID_ARG. An unknown code + * is TRANSCRIBE_ERR_UNSUPPORTED_LANGUAGE. Duplicates count once. */ + const char * const * allowed; + int32_t n_allowed; + int32_t top_k; /* keep the best top_k candidates; 0 = every allowed label */ +}; + +/* Summary of the last run. */ +struct transcribe_langid_result { + uint64_t struct_size; + int32_t n_candidates; /* rows readable via transcribe_langid_get_candidate */ + int32_t n_allowed; /* labels in the allowed set (before top_k) */ + /* Share of the softmax over every label that falls in the allowed set; + * 1 when unrestricted. A low value means the speech is probably outside + * the allowed set. */ + float allowed_mass; + int64_t audio_ms; /* audio actually scored, after the crop */ +}; + +/* One ranked label. code and name point at model-owned storage, valid until + * transcribe_model_free. */ +struct transcribe_langid_candidate { + uint64_t struct_size; + int32_t index; /* label index */ + const char * code; + const char * name; + float p; /* softmax over the allowed set */ + float logit; +}; + +TRANSCRIBE_API void transcribe_langid_info_init(struct transcribe_langid_info * out); +TRANSCRIBE_API void transcribe_langid_session_params_init(struct transcribe_langid_session_params * params); +TRANSCRIBE_API void transcribe_langid_params_init(struct transcribe_langid_params * params); +TRANSCRIBE_API void transcribe_langid_result_init(struct transcribe_langid_result * out); +TRANSCRIBE_API void transcribe_langid_candidate_init(struct transcribe_langid_candidate * out); + +/* UNSUPPORTED_ROLE when the model does not serve TRANSCRIBE_ROLE_LANGID. */ +TRANSCRIBE_API transcribe_status transcribe_langid_get_info(const struct transcribe_model * model, + struct transcribe_langid_info * out); + +/* The label table. Code and name of label i, or NULL when i is out of range + * or the model does not serve TRANSCRIBE_ROLE_LANGID; model-owned. */ +TRANSCRIBE_API const char * transcribe_langid_label_code(const struct transcribe_model * model, int32_t i); +TRANSCRIBE_API const char * transcribe_langid_label_name(const struct transcribe_model * model, int32_t i); +/* Index of a code or alias ("he" and "iw" name the same label), or -1. */ +TRANSCRIBE_API int32_t transcribe_langid_label_index(const struct transcribe_model * model, const char * code_or_alias); + +/* params may be NULL for defaults. UNSUPPORTED_ROLE when the model does not + * serve TRANSCRIBE_ROLE_LANGID. On failure *out_session is NULL. */ +TRANSCRIBE_API transcribe_status transcribe_langid_session_init(struct transcribe_model * model, + const struct transcribe_langid_session_params * params, + struct transcribe_langid_session ** out_session); + +/* NULL is a no-op. */ +TRANSCRIBE_API void transcribe_langid_session_free(struct transcribe_langid_session * session); + +/* Polled during a run; returning true stops it with TRANSCRIBE_ERR_ABORTED. */ +TRANSCRIBE_API void transcribe_langid_set_abort_callback(struct transcribe_langid_session * session, + transcribe_abort_callback cb, + void * user_data); + +/* + * Identify the language of one clip: 16 kHz mono float32 PCM, every sample + * finite. params may be NULL for defaults. Input longer than the session's + * max_audio_ms is scored on its last max_audio_ms. Replaces the previous + * result. Malformed input (NULL pointers, n_samples <= 0, NaN / Inf, a bad + * allowed list, an unknown code, scored audio shorter than min_audio_ms) + * returns an error before the previous result is touched. + */ +TRANSCRIBE_API transcribe_status transcribe_langid_run(struct transcribe_langid_session * session, + const float * pcm, + int n_samples, + const struct transcribe_langid_params * params); + +/* + * The last run's summary. Zeroed before any run and after a run that passed + * input validation but then failed (aborted, backend error). + */ +TRANSCRIBE_API transcribe_status transcribe_langid_get_result(const struct transcribe_langid_session * session, + struct transcribe_langid_result * out); + +/* Candidate i of the last run, ranked by p (descending; ties keep label + * order). An out-of-range index returns OK with a zeroed row (code NULL). */ +TRANSCRIBE_API transcribe_status transcribe_langid_get_candidate(const struct transcribe_langid_session * session, + int i, + struct transcribe_langid_candidate * out); + +/* load_ms plus the last run's stage times (mel_ms, encode_ms). */ +TRANSCRIBE_API transcribe_status transcribe_langid_get_timings(const struct transcribe_langid_session * session, + struct transcribe_timings * out); + +#ifdef __cplusplus +} +#endif + +#endif /* TRANSCRIBE_LANGID_H */ diff --git a/reports/porting/ecapa_tdnn/lang-id-voxlingua107-ecapa/intake.json b/reports/porting/ecapa_tdnn/lang-id-voxlingua107-ecapa/intake.json new file mode 100644 index 000000000..2ce5a1a11 --- /dev/null +++ b/reports/porting/ecapa_tdnn/lang-id-voxlingua107-ecapa/intake.json @@ -0,0 +1,154 @@ +{ + "schema_version": "transcribe-intake-v1", + "family": "ecapa_tdnn", + "hf_repo": "speechbrain/lang-id-voxlingua107-ecapa", + "hf_revision": "0253049ae131d6a4be1c4f0d8b0ff483a0f8c8e9", + "sources": { + "config": { + "kind": "hf_file", + "path": "hyperparams.yaml", + "status": "found", + "detail": "SpeechBrain hyperparams.yaml (no HF config.json). Instantiates lobes.features.Fbank, lobes.models.ECAPA_TDNN.ECAPA_TDNN and lobes.models.Xvector.Classifier by class name; only channels / kernel_sizes / dilations / attention_channels / lin_neurons and four Fbank fields are explicit, the rest are speechbrain 1.1.1 defaults (asserted against the live modules by scripts/convert-ecapa_tdnn.py)." + }, + "preprocessor": { + "kind": "hf_file", + "path": "hyperparams.yaml::compute_features", + "status": "found", + "detail": "Fbank(n_mels=60, left_frames=0, right_frames=0, deltas=False) + InputNormalization(norm_type=sentence, std_norm=False)." + }, + "tokenizer_config": { + "kind": "hf_file", + "path": "tokenizer_config.json", + "status": "missing" + }, + "tokenizer_json": { + "kind": "hf_file", + "path": "tokenizer.json", + "status": "missing" + }, + "generation_config": { + "kind": "hf_file", + "path": "generation_config.json", + "status": "missing" + }, + "safetensors_metadata": { + "kind": "hf_api", + "path": "HfApi.get_safetensors_metadata", + "status": "missing", + "detail": "header-only floating dtype distribution; no tensor payloads downloaded" + }, + "reference_modeling_code": { + "kind": "reference_code", + "path": "speechbrain==1.1.1 speechbrain.inference.classifiers.EncoderClassifier", + "status": "found", + "detail": "The only published implementation; ported from handy-computer/langid.cpp@96af3a5 where the C++ was first validated against it." + } + }, + "variants": [ + { + "name": "lang-id-voxlingua107-ecapa", + "memory_gb": 0.1, + "files": [ + "hyperparams.yaml", + "embedding_model.ckpt", + "classifier.ckpt", + "label_encoder.txt" + ] + } + ], + "config": { + "architecture_candidates": [ + "encoder-classifier" + ], + "key_fields": { + "channels": [ + 1024, + 1024, + 1024, + 1024, + 3072 + ], + "kernel_sizes": [ + 5, + 3, + 3, + 3, + 1 + ], + "dilations": [ + 1, + 2, + 3, + 4, + 1 + ], + "res2net_scale": 8, + "se_channels": 128, + "attention_channels": 128, + "lin_neurons": 256, + "classifier_hidden": 512, + "n_labels": 107 + } + }, + "dtype": { + "expected": "float32", + "source": "weights_header", + "evidence": "embedding_model.ckpt / classifier.ckpt are fp32 PyTorch state dicts (21,058,432 + 188,011 params); the converter asserts every parameter is torch.float32.", + "details": { + "config_declared": null, + "header_distribution": { + "float32": 21246443 + } + } + }, + "frontend": { + "sample_rate": 16000, + "n_mels": 60, + "hop_length": 160, + "fft_size": 400, + "window": "hamming_periodic", + "normalization": "sentence_mean", + "preemphasis": null, + "dither": null, + "center": true, + "padding_mode": "constant", + "mel_filterbank_norm": null + }, + "tokenizer": { + "type": "other", + "vocab_size": 0, + "special_tokens": {}, + "has_language_tokens": false, + "vocab_sha256": null + }, + "capabilities": { + "languages": [], + "language_detection": false, + "translation": false, + "timestamps": [], + "streaming": false, + "speaker_diarization": false + }, + "upstream_benchmarks": [ + { + "dataset": "VoxLingua107 dev", + "language": null, + "metric": "accuracy", + "score": 93.3, + "score_unit": "percent", + "source": "model card (6.7% error rate on the VoxLingua107 development set)", + "notes": "Language ID role: top-1 accuracy over 107 labels. Not a WER figure." + } + ], + "reference_framework": "author_repo_speechbrain", + "reference_rationale": "SpeechBrain is the only implementation of this checkpoint: hyperparams.yaml instantiates SpeechBrain classes by name and there is no Transformers shim. Pinned to speechbrain 1.1.1 / torch 2.13.0, the versions langid.cpp's tolerances were measured against.", + "architecture_pattern": "encoder-classifier", + "known_risks": [ + "Role is LANGID, not ASR: acceptance is FLEURS decision parity vs SpeechBrain (12000 decisions) plus top-1 accuracy, not WER on LibriSpeech.", + "capabilities.language_detection / stt.capability.lang_detect mean ASR auto-detect and stay false; language ID is the model's role (transcribe_model_roles), not an ASR capability.", + "SpeechBrain front end (periodic Hamming, 10*log10, 80 dB top-db, per-bin mean) runs on the shared transcribe-mel (src/transcribe-mel.cpp) with window_type hamming_periodic and normalize sentence_mean.", + "Labels use VoxLingua107's legacy codes (iw, jw, tl, no); modern codes are stored as aliases. Codes are not canonicalized against ASR models' language lists.", + "Closed-set classifier: out-of-set speech still gets a confident label; callers use allowed_mass / restrict the allowed set." + ], + "intake_gaps": [] +} diff --git a/samples/README.md b/samples/README.md index d3fd433f2..80f271e1f 100644 --- a/samples/README.md +++ b/samples/README.md @@ -40,6 +40,31 @@ and is recorded here so each clip can be traced back. Regenerate the source pool, not the clips themselves, with `uv run scripts/wer/ingest.py fleurs `. +## Language ID clips (FLEURS) + +Fixtures for the ecapa_tdnn (LANGID) family, carried over byte-for-byte from +langid.cpp (`handy-computer/langid.cpp@96af3a5`, `scripts/make_samples_fleurs.py`). +Rule for the eight +`fleurs-` clips: the first test-split utterance of 4-12 s, in parquet +order, that the SpeechBrain reference (`speechbrain/lang-id-voxlingua107-ecapa` +@ `0253049a`) classifies as the expected VoxLingua107 code. They back +`tests/golden/ecapa_tdnn/`, `transcribe_ecapa_tdnn_real_smoke` and +`transcribe_cli_langid_smoke`. + +Source: [google/fleurs](https://huggingface.co/datasets/google/fleurs), test +split, licensed **CC-BY-4.0**. + +| file | duration | config | FLEURS id | +| --- | ---: | --- | --- | +| `fleurs-en.wav` | 10.56 s | `en_us` | 1904 | +| `fleurs-de.wav` | 11.16 s | `de_de` | 1738 | +| `fleurs-fr.wav` | 10.20 s | `fr_fr` | 1987 | +| `fleurs-es.wav` | 9.72 s | `es_419` | 1764 | +| `fleurs-ja.wav` | 10.44 s | `ja_jp` | 1828 | +| `fleurs-zh.wav` | 10.38 s | `cmn_hans_cn` | 1906 | +| `fleurs-ru.wav` | 7.92 s | `ru_ru` | 1917 | +| `fleurs-id.wav` | 8.58 s | `id_id` | 1909 | + ## Everything else These predate this file and arrived inside unrelated commits, so their source diff --git a/samples/fleurs-de.wav b/samples/fleurs-de.wav new file mode 100644 index 000000000..d07c69e8d Binary files /dev/null and b/samples/fleurs-de.wav differ diff --git a/samples/fleurs-en.wav b/samples/fleurs-en.wav new file mode 100644 index 000000000..74121e55c Binary files /dev/null and b/samples/fleurs-en.wav differ diff --git a/samples/fleurs-es.wav b/samples/fleurs-es.wav new file mode 100644 index 000000000..ba264d20d Binary files /dev/null and b/samples/fleurs-es.wav differ diff --git a/samples/fleurs-fr.wav b/samples/fleurs-fr.wav new file mode 100644 index 000000000..9adcd8bf3 Binary files /dev/null and b/samples/fleurs-fr.wav differ diff --git a/samples/fleurs-id.wav b/samples/fleurs-id.wav new file mode 100644 index 000000000..19415ce75 Binary files /dev/null and b/samples/fleurs-id.wav differ diff --git a/samples/fleurs-ja.wav b/samples/fleurs-ja.wav new file mode 100644 index 000000000..d793a00ee Binary files /dev/null and b/samples/fleurs-ja.wav differ diff --git a/samples/fleurs-ru.wav b/samples/fleurs-ru.wav new file mode 100644 index 000000000..35c32bc09 Binary files /dev/null and b/samples/fleurs-ru.wav differ diff --git a/samples/fleurs-zh.wav b/samples/fleurs-zh.wav new file mode 100644 index 000000000..8f0660a5f Binary files /dev/null and b/samples/fleurs-zh.wav differ diff --git a/scripts/catalog/check.py b/scripts/catalog/check.py index a0facf246..5c72c9470 100755 --- a/scripts/catalog/check.py +++ b/scripts/catalog/check.py @@ -116,18 +116,22 @@ def pairing_pass(records: dict, selected: bool = False) -> int: def publication_pass(records: dict, profile_id: str | None, enforce: bool) -> int: - """Check publication matrices, including explicit legacy accuracy rows.""" + """Check publication matrices, including explicit legacy accuracy rows. + + Each record is held to its role's profile (language ID has its own).""" try: - resolved_id, profile = profiles.load_profile(profile_id) + profiles.load_profile(profile_id) except (OSError, ValueError, json.JSONDecodeError) as exc: print(f"publication FAIL: {exc}") return 1 problems = 0 - totals = collections.Counter() + totals: dict[str, collections.Counter] = collections.defaultdict(collections.Counter) for name, record in records.items(): + resolved_id, profile = profiles.profile_for(record, profile_id) # An ASR publication profile does not apply to standalone diarizers. - if not record.get("capabilities", {}).get("transcribe", {}).get("supported"): + if not profile.get("role") and \ + not record.get("capabilities", {}).get("transcribe", {}).get("supported"): continue accuracy_raw = profiles.expected_accuracy(record, profile) speed_raw = profiles.expected_speed(record, profile) @@ -177,7 +181,7 @@ def publication_pass(records: dict, profile_id: str | None, enforce: bool) -> in # Preserve pre-profile rows that were already published, # even when their quant was not selected by today's # publication matrix. They are archive data, not drift. - totals["accuracy_archived"] += 1 + totals[resolved_id]["accuracy_archived"] += 1 else: accuracy_extra.add(key) continue @@ -193,9 +197,9 @@ def publication_pass(records: dict, profile_id: str | None, enforce: bool) -> in per_model["accuracy_extra"] = len(accuracy_extra) per_model["accuracy_duplicate"] = sum( count - 1 for count in accuracy_counts.values() if count > 1) - totals["accuracy_required"] += len(expected_keys) + totals[resolved_id]["accuracy_required"] += len(expected_keys) for suffix in ("missing", "invalid", "extra", "duplicate"): - totals[f"accuracy_{suffix}"] += per_model[f"accuracy_{suffix}"] + totals[resolved_id][f"accuracy_{suffix}"] += per_model[f"accuracy_{suffix}"] # Both published samples are required for each quant selected by the # profile on every machine/backend target. @@ -215,9 +219,9 @@ def publication_pass(records: dict, profile_id: str | None, enforce: bool) -> in per_model["speed_extra"] = len(speed_actual - speed_expected) per_model["speed_duplicate"] = sum( count - 1 for count in speed_counts.values() if count > 1) - totals["speed_required"] += len(speed_expected) + totals[resolved_id]["speed_required"] += len(speed_expected) for suffix in ("missing", "invalid", "extra", "duplicate"): - totals[f"speed_{suffix}"] += per_model[f"speed_{suffix}"] + totals[resolved_id][f"speed_{suffix}"] += per_model[f"speed_{suffix}"] count = sum(per_model.values()) if count: @@ -225,11 +229,12 @@ def publication_pass(records: dict, profile_id: str | None, enforce: bool) -> in details = ", ".join(f"{key}={value}" for key, value in per_model.items() if value) print(f" {'FAIL' if enforce else 'TODO'} {name}: {details}") - print(f"publication {resolved_id}: accuracy {totals['accuracy_required']} required, " - f"{totals['accuracy_missing']} missing, {totals['accuracy_invalid']} invalid, " - f"{totals['accuracy_extra']} extra, {totals['accuracy_archived']} archived legacy; speed " - f"{totals['speed_required']} required, {totals['speed_missing']} missing, " - f"{totals['speed_invalid']} invalid, {totals['speed_extra']} extra") + for resolved_id, total in sorted(totals.items()): + print(f"publication {resolved_id}: accuracy {total['accuracy_required']} required, " + f"{total['accuracy_missing']} missing, {total['accuracy_invalid']} invalid, " + f"{total['accuracy_extra']} extra, {total['accuracy_archived']} archived legacy; speed " + f"{total['speed_required']} required, {total['speed_missing']} missing, " + f"{total['speed_invalid']} invalid, {total['speed_extra']} extra") if problems and not enforce: print(" audit only; pass --publication-profile to enforce this gate") return problems if enforce else 0 diff --git a/scripts/catalog/common.py b/scripts/catalog/common.py index f1bba3ad4..dc0d5d057 100644 --- a/scripts/catalog/common.py +++ b/scripts/catalog/common.py @@ -108,10 +108,18 @@ def dataset_spec(row: dict) -> str: def dataset_label(dataset: str, split: str, language: str) -> str: if dataset == "fleurs": - return f"FLEURS {language}" + # "mul": a multilingual pool (language ID scores several FLEURS + # languages as one result set). + return "FLEURS multilingual" if language == "mul" else f"FLEURS {language}" return DATASET_LABELS.get((dataset, split), f"{dataset} {split}") +def metric_label(metric: str) -> str: + """A metric as a column header prints it: WER / CER / DER / CPWER, or + "Top-1 accuracy" for language ID.""" + return "Top-1 accuracy" if metric == "accuracy" else metric.upper() + + def headline(record: dict) -> dict | None: """The benchmark row-set a variant publishes in its download table. @@ -156,20 +164,24 @@ def headline_recipe(record: dict) -> str: sample = measured[0] n_utts = max(row["n_utts"] for row in rows) unit = "meetings" if target["metric"] in ("der", "cpwer") else "utterances" - parts = [f"{target['metric'].upper()} on the full {headline_label(record)} split " - f"({n_utts:,} {unit})"] + # Language ID scores a pooled, cropped subset (the row's notes say which), + # and has no batch size or timestamps. + lid = target["metric"] == "accuracy" + scope = "" if lid else "the full " + parts = [f"{metric_label(target['metric'])} on {scope}{headline_label(record)} " + f"{'' if lid else 'split '}({n_utts:,} {unit})"] batch_sizes = sorted({row["batch_size"] for row in rows - if row.get("batch_size") is not None}) + if row.get("batch_size") is not None and not lid}) if len(batch_sizes) == 1: parts.append(f"batch size {batch_sizes[0]}") elif batch_sizes: parts.append("batch sizes " + " and ".join(str(size) for size in batch_sizes)) - if sample.get("timestamps"): + if sample.get("timestamps") and not lid: parts.append(f"timestamps {sample['timestamps']}") if sample.get("language_hint"): parts.append(f"language hint `{sample['language_hint']}`") if sample.get("backend"): - parts.append(f"decoded on {sample['backend']}") + parts.append(f"{'scored' if lid else 'decoded'} on {sample['backend']}") text = ", ".join(parts) + "." shas = sorted({(row["engine_sha"], row.get("measured_on") or "") for row in rows if row.get("engine_sha")}) @@ -181,11 +193,18 @@ def headline_recipe(record: dict) -> str: return text +def row_pct(row: dict) -> float: + """A row's percentage: the error rate, or the accuracy on metric=accuracy + (language ID) rows, which carry acc_pct instead of err_pct.""" + return row["acc_pct"] if row.get("metric") == "accuracy" else row["err_pct"] + + def fmt_err(row: dict | None, dp: int = 2) -> str: - """An error rate as a card prints it. `-` when the cell was not measured.""" + """An error rate (or, for accuracy rows, the accuracy) as a card prints + it. `-` when the cell was not measured.""" if row is None: return "-" - return f"{row['err_pct']:.{dp}f}%" + return f"{row_pct(row):.{dp}f}%" # -------------------------------------------------------------------------- @@ -243,6 +262,8 @@ def capabilities_summary(record: dict) -> str: """The extras beyond plain transcription, as a short comma list.""" caps = record.get("capabilities", {}) out = [] + if record.get("role") == "langid": + out.append(f"language ID ({len(record.get('languages', []))} languages)") for name, label in (("translate", "translate"), ("streaming", "streaming"), ("diarize", "diarize")): if caps.get(name, {}).get("supported"): diff --git a/scripts/catalog/db.py b/scripts/catalog/db.py index fbd39628b..532de458c 100755 --- a/scripts/catalog/db.py +++ b/scripts/catalog/db.py @@ -101,7 +101,8 @@ metric TEXT NOT NULL, language_hint TEXT, backend TEXT, - err_pct REAL NOT NULL CHECK(err_pct >= 0), + err_pct REAL CHECK(err_pct >= 0), + acc_pct REAL CHECK(acc_pct BETWEEN 0 AND 100), ci_lo REAL, ci_hi REAL, n_utts INTEGER NOT NULL CHECK(n_utts > 0), @@ -117,7 +118,10 @@ utts_over_50pct INTEGER, publication_profile TEXT, scoring TEXT, - mode TEXT + mode TEXT, + agreement_n_agree INTEGER, + agreement_n INTEGER, + agreement_max_abs_logit_delta REAL ); CREATE UNIQUE INDEX accuracy_identity ON accuracy( dataset_id, variant, quant, metric, @@ -153,7 +157,7 @@ -- The per-quant column a model card and its doc print. CREATE VIEW headline AS SELECT a.variant, d.dataset, d.split, d.language, - a.quant, a.metric, a.err_pct, a.ci_lo, a.ci_hi, a.n_utts + a.quant, a.metric, a.err_pct, a.acc_pct, a.ci_lo, a.ci_hi, a.n_utts FROM accuracy a JOIN models m ON m.variant = a.variant JOIN datasets d ON d.dataset_id = a.dataset_id @@ -237,9 +241,10 @@ def build(records: dict[str, dict], out: pathlib.Path) -> dict[str, int]: (variant, item["quant"], item["filename"], item["size_bytes"]) for item in record.get("downloads", [])]) con.executemany( - "INSERT INTO accuracy VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)", [ + "INSERT INTO accuracy VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)", [ (dataset_id(row), variant, row["quant"], row["metric"], - row.get("language_hint"), row.get("backend"), row["err_pct"], + row.get("language_hint"), row.get("backend"), row.get("err_pct"), + row.get("acc_pct"), (row.get("ci95") or [None, None])[0], (row.get("ci95") or [None, None])[1], row["n_utts"], row.get("batch_size"), row.get("timestamps"), row.get("engine_sha"), @@ -248,7 +253,10 @@ def build(records: dict[str, dict], out: pathlib.Path) -> dict[str, int]: (row.get("errors") or {}).get("del"), (row.get("errors") or {}).get("ins"), row.get("empty_hyp"), row.get("utts_over_50pct"), row.get("publication_profile"), - row.get("scoring"), row.get("mode")) + row.get("scoring"), row.get("mode"), + (row.get("agreement") or {}).get("n_agree"), + (row.get("agreement") or {}).get("n"), + (row.get("agreement") or {}).get("max_abs_logit_delta")) for row in record.get("accuracy_benchmarks", [])]) con.executemany( "INSERT INTO speed VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)", [ diff --git a/scripts/catalog/ingest_accuracy.py b/scripts/catalog/ingest_accuracy.py index 85a85691c..8b78593f7 100755 --- a/scripts/catalog/ingest_accuracy.py +++ b/scripts/catalog/ingest_accuracy.py @@ -10,6 +10,13 @@ for f in reports/wer/*.jsonl; do uv run scripts/wer/score.py "$f"; done uv run scripts/catalog/ingest_accuracy.py --dry-run uv run scripts/catalog/ingest_accuracy.py + +Language ID records follow their own profile (catalog/_benchmark_profiles.json +`roles`). Their scores come from scripts/langid/score.py --ref --json under +reports/langid/, named after the GGUF like a WER score +(`.fleurs-mul.score.json`, with its `.agreement.json`); the profile +names the crop, and the importer checks the sweep's own recipe against it +instead of a stamp. """ from __future__ import annotations @@ -23,6 +30,8 @@ import profiles # noqa: E402 REPORTS = common.REPO / "reports" / "wer" +LANGID_REPORTS = common.REPO / "reports" / "langid" +LANGID_SCORE = "transcribe-langid-score-v1" def score_path(record: dict, cell: dict, reports: pathlib.Path) -> pathlib.Path: @@ -67,17 +76,73 @@ def row_from_score(cell: dict, score: dict, profile_id: str) -> dict: } +def langid_row(record: dict, cell: dict, score: dict, agreement: dict | None, + profile_id: str) -> tuple[dict | None, list[str]]: + """The catalog row for one language ID cell, or why the score cannot be + it. The sweep's header travels in the score, so the recipe is checked + here rather than trusted from a stamp.""" + filename = next(item["filename"] for item in record["downloads"] + if item["quant"] == cell["quant"]) + crop = (score.get("crops") or {}).get(cell["crop_s"]) + reasons = [] + if score.get("schema") != LANGID_SCORE: + reasons.append(f"schema={score.get('schema')!r}") + if score.get("engine") != "cpp": + reasons.append(f"engine={score.get('engine')!r}") + if pathlib.PurePath(str(score.get("model", ""))).name != filename: + reasons.append(f"model={score.get('model')!r}") + for field in ("dataset", "split", "language", "backend"): + if score.get(field) != cell.get(field): + reasons.append(f"{field}={score.get(field)!r}") + if crop is None: + reasons.append(f"no {cell['crop_s']} s crop") + if not score.get("engine_sha"): + reasons.append("engine_sha is empty") + if sorted(score.get("languages") or []) != sorted(cell.get("pooled_languages") or []): + reasons.append(f"languages={score.get('languages')!r} are not the profile's pooled_languages") + if agreement is None: + reasons.append("no agreement (write it with score.py --ref --json)") + elif agreement.get("run") != score.get("run"): + reasons.append(f"agreement is for {agreement.get('run')!r}, not {score.get('run')!r}") + elif not agreement.get("same_rows", False): + reasons.append("agreement run does not cover the reference's rows") + if reasons: + return None, reasons + row = { + "dataset": cell["dataset"], + "split": cell["split"], + "language": cell["language"], + "backend": cell["backend"], + "quant": cell["quant"], + "metric": cell["metric"], + "acc_pct": crop["acc_pct"], + "ci95": crop["ci95"], + "n_utts": crop["n"], + "batch_size": cell["batch_size"], + "timestamps": cell["timestamps"], + "engine_sha": score["engine_sha"], + "publication_profile": profile_id, + "measured_on": (score.get("created") or "")[:10] or None, + } + row["agreement"] = { + "n_agree": agreement["n_agree"], + "n": agreement["n"], + "max_abs_logit_delta": float(f"{agreement['max_abs_logit_delta']:.2g}"), + } + return row, [] + + def main() -> int: parser = argparse.ArgumentParser() - parser.add_argument("--reports", default=str(REPORTS)) + parser.add_argument("--reports", default=None, + help="score directory (default: reports/wer, and " + "reports/langid for language ID records)") parser.add_argument("--profile", default=None) parser.add_argument("--models", default="", help="comma-separated variants (default: all)") parser.add_argument("--dry-run", action="store_true") args = parser.parse_args() - profile_id, profile = profiles.load_profile(args.profile) - reports = pathlib.Path(args.reports) selected = {item.strip() for item in args.models.split(",") if item.strip()} records = common.load_records() unknown = selected - records.keys() @@ -86,41 +151,55 @@ def main() -> int: return 2 added = replaced = rejected = missing = unstamped = 0 + used: set[str] = set() for variant, record in records.items(): if selected and variant not in selected: continue + profile_id, profile = profiles.profile_for(record, args.profile) + langid = profile.get("role") == "langid" + reports = pathlib.Path(args.reports) if args.reports else ( + LANGID_REPORTS if langid else REPORTS) path = common.CATALOG_DIR / f"{variant}.json" rows = record.get("accuracy_benchmarks", []) changed = False expected = profiles.apply_exceptions( record, "accuracy", profiles.expected_accuracy(record, profile)) + if expected: + used.add(profile_id) for cell in expected: source_path = score_path(record, cell, reports) if not source_path.exists(): missing += 1 continue score = json.loads(source_path.read_text()) - recipe = score.get("recipe") or {} - covered = any(profiles.profile_key(row) == profiles.profile_key(cell) - or (row.get("measurement_provenance") == "legacy-published" - and profiles.accuracy_core_key(row) == profiles.accuracy_core_key(cell)) - for row in rows) - if not recipe.get("publication_profile") and covered: - # A score from before profile stamping, for a cell the catalog - # already publishes: superseded history, not a problem. - unstamped += 1 - continue - reasons = [] - if recipe.get("publication_profile") != profile_id: - reasons.append(f"profile={recipe.get('publication_profile')!r}") - if score.get("metric") != cell["metric"]: - reasons.append(f"metric={score.get('metric')!r}") - if score.get("timestamps") != cell["timestamps"]: - reasons.append(f"timestamps={score.get('timestamps')!r}") - if recipe.get("backend") != cell["backend"]: - reasons.append(f"backend={recipe.get('backend')!r}") - if not score.get("engine_sha"): - reasons.append("engine_sha is empty") + if langid: + agreement_path = source_path.with_name( + source_path.name.replace(".score.json", ".agreement.json")) + agreement = (json.loads(agreement_path.read_text()) + if agreement_path.exists() else None) + new_row, reasons = langid_row(record, cell, score, agreement, profile_id) + else: + new_row, reasons = None, [] + recipe = score.get("recipe") or {} + covered = any(profiles.profile_key(row) == profiles.profile_key(cell) + or (row.get("measurement_provenance") == "legacy-published" + and profiles.accuracy_core_key(row) == profiles.accuracy_core_key(cell)) + for row in rows) + if not recipe.get("publication_profile") and covered: + # A score from before profile stamping, for a cell the catalog + # already publishes: superseded history, not a problem. + unstamped += 1 + continue + if recipe.get("publication_profile") != profile_id: + reasons.append(f"profile={recipe.get('publication_profile')!r}") + if score.get("metric") != cell["metric"]: + reasons.append(f"metric={score.get('metric')!r}") + if score.get("timestamps") != cell["timestamps"]: + reasons.append(f"timestamps={score.get('timestamps')!r}") + if recipe.get("backend") != cell["backend"]: + reasons.append(f"backend={recipe.get('backend')!r}") + if not score.get("engine_sha"): + reasons.append("engine_sha is empty") if reasons: rejected += 1 print(f" reject {source_path.name}: {', '.join(reasons)}") @@ -138,7 +217,7 @@ def main() -> int: indices = [index for index, row in enumerate(rows) if row.get("measurement_provenance") == "legacy-published" and profiles.accuracy_core_key(row) == core] - new_row = row_from_score(cell, score, profile_id) + new_row = new_row or row_from_score(cell, score, profile_id) if indices: first = indices[0] if rows[first] == new_row and len(indices) == 1: @@ -154,7 +233,7 @@ def main() -> int: if changed and not args.dry_run: common.write_record(path, record) - print(f"profile {profile_id}: {added} added, {replaced} replaced, " + print(f"profile {', '.join(sorted(used)) or '-'}: {added} added, {replaced} replaced, " f"{rejected} rejected, {unstamped} unstamped score(s) for already published " f"cells skipped, {missing} score file(s) absent") if args.dry_run: diff --git a/scripts/catalog/ingest_perf.py b/scripts/catalog/ingest_perf.py index 1c4c62370..6aab0ca35 100755 --- a/scripts/catalog/ingest_perf.py +++ b/scripts/catalog/ingest_perf.py @@ -5,7 +5,9 @@ """Fold bench driver reports into the catalog's speed_benchmarks rows. `scripts/bench/run.py` writes one report per (variant, backend) under -reports/perf//, and reports/ is gitignored -- so the latency +reports/perf// (`scripts/langid/bench.py` writes the same shape +for language ID, which transcribe-bench cannot run), and reports/ is +gitignored -- so the latency breakdown only exists on the machine that measured it. This is the hop that moves it into the catalog, where it is durable and publishable. @@ -112,7 +114,9 @@ def mean(field: str): # the canonical backend; the driver records that at the top level. "backend": (report.get("backend") or run.get("backend", "")).lower(), "quant": quant, - "sample": pathlib.PurePosixPath(run.get("sample_path", "")).stem, + # The driver names a cell after its clip; scripts/langid/bench.py + # scores several lengths of one clip and names each. + "sample": run.get("sample") or pathlib.PurePosixPath(run.get("sample_path", "")).stem, "sample_duration_s": duration, "total_ms": round(total, 1), # xrt is recomputed from the unrounded mean rather than carried @@ -214,7 +218,6 @@ def main() -> int: for note in notes: print(f" note: {note}") - profile_id, profile = profiles.load_profile() filled = updated = added = matched = 0 drift, refused, unmatched = [], [], [] for variant, record in common.load_records().items(): @@ -261,6 +264,7 @@ def main() -> int: # Profile runs can create rows; the old importer could only refresh # placeholders, which made a newly required quant impossible to ingest # without first hand-authoring empty catalog cells. + profile_id, profile = profiles.profile_for(record) expected = profiles.apply_exceptions( record, "speed", profiles.expected_speed(record, profile)) expected_keys = { diff --git a/scripts/catalog/profiles.py b/scripts/catalog/profiles.py index 6d926da0a..7fdd7b56d 100644 --- a/scripts/catalog/profiles.py +++ b/scripts/catalog/profiles.py @@ -42,6 +42,10 @@ def load_profiles(path: pathlib.Path = PROFILE_PATH) -> dict: default = data.get("default") if default not in data["profiles"]: raise ValueError(f"{path}: default profile {default!r} is not defined") + for role, profile_id in (data.get("roles") or {}).items(): + if data["profiles"].get(profile_id, {}).get("role") != role: + raise ValueError(f"{path}: roles.{role} must name a defined profile " + f"whose role is {role!r}") return data @@ -57,6 +61,30 @@ def load_profile(profile_id: str | None = None) -> tuple[str, dict]: ) from exc +def profile_for(record: dict, profile_id: str | None = None) -> tuple[str, dict]: + """The profile that governs one record: `profile_id` when it is written + for the record's role, else the profile its role names in `roles`, else + the default. ASR and diarize records share the default; language ID has + a profile of its own, since a top-1 accuracy over a pooled set and a + classifier's latency are not cells of the ASR matrix.""" + data = load_profiles() + if profile_id: + named = load_profile(profile_id)[1] + if governs(named, record): + return profile_id, named + chosen = data.get("roles", {}).get(record.get("role", "asr"), data["default"]) + return chosen, data["profiles"][chosen] + + +def governs(profile: dict, record: dict) -> bool: + """A profile with a `role` governs only records of that role; a profile + without one governs every record no role-specific profile claims.""" + role = record.get("role", "asr") + if profile.get("role"): + return profile["role"] == role + return role not in load_profiles().get("roles", {}) + + def canonical_machine(slug: str) -> str: return MACHINE_ALIASES.get(slug, slug) @@ -117,7 +145,10 @@ def _quants(spec: str | list[str], record: dict) -> list[str]: def expected_accuracy(record: dict, profile: dict) -> list[dict]: """Expand a profile into publication accuracy cells for one model.""" cells: list[dict] = [] - if not (record.get("capabilities", {}).get("transcribe", {}).get("supported")): + if not governs(profile, record): + return cells + if not profile.get("role") and not ( + record.get("capabilities", {}).get("transcribe", {}).get("supported")): return cells for suite in profile.get("accuracy", []): selector = suite["languages"] @@ -125,6 +156,9 @@ def expected_accuracy(record: dict, profile: dict) -> list[dict]: languages = ["en"] if "en" in fleurs_languages(record) else [] elif selector == "supported-intersect-fleurs": languages = fleurs_languages(record) + elif selector == "pooled": + # One result over every evaluated language at once (language ID). + languages = ["mul"] else: raise ValueError(f"unknown language selector {selector!r}") for language in languages: @@ -133,14 +167,21 @@ def expected_accuracy(record: dict, profile: dict) -> list[dict]: "dataset": suite["dataset"], "split": suite["split"], "language": language, - "runtime_language": runtime_language(record, language), + "runtime_language": (None if selector == "pooled" + else runtime_language(record, language)), "quant": quant, - "metric": "cer" if language in CER_LANGUAGES else "wer", + "metric": suite.get("metric") or ( + "cer" if language in CER_LANGUAGES else "wer"), "batch_size": suite["batch_size"], "sort_by_length": suite.get("sort_by_length", False), "timestamps": suite["timestamps"], "gpu": suite.get("gpu"), "backend": suite.get("backend"), + # Language ID's crop: which slice of a run.py sweep the + # published number is, and which languages its pooled + # mean is over. + **{key: suite[key] for key in ("crop_s", "pooled_languages") + if key in suite}, }) return cells @@ -174,6 +215,8 @@ def all_speed_samples(profile: dict, records: dict[str, dict]) -> list[str]: def expected_speed(record: dict, profile: dict) -> list[dict]: """Expand the exact publication speed matrix for one model.""" + if not governs(profile, record): + return [] spec = profile["speed"] samples = speed_samples(record, profile) cells: list[dict] = [] diff --git a/scripts/catalog/render.py b/scripts/catalog/render.py index 76e1bf0fb..8b4e3ce0c 100755 --- a/scripts/catalog/render.py +++ b/scripts/catalog/render.py @@ -16,8 +16,9 @@ Blocks: `downloads`, `perf machine=`, `accuracy` (one table per -dataset split beyond the headline), `recipe` (the mechanical WER sentence -from the headline rows), `pin` (licence, upstream and validation pins), +dataset split beyond the headline), `agreement` (language ID: headline +accuracy and reference agreement per GGUF), `recipe` (the mechanical WER +sentence from the headline rows), `pin` (licence, upstream and validation pins), `intro` (upstream link plus the card spec's `summary`), `prose field=wer.notes` (any `|` text field of the spec, dotted path), `family variants=a,b,c` (a roll-up row per variant, for family pages), and `family-index` (the root README's supported-models table, @@ -84,7 +85,7 @@ def block_downloads(record: dict, attrs: dict[str, str]) -> list[str]: aligns = ["l", "l", "r"] if want_metric: label = attrs.get("label") or common.headline_label(record) - metric = attrs.get("metric_name") or target["metric"].upper() + metric = attrs.get("metric_name") or common.metric_label(target["metric"]) header.append(f"{metric} ({label})") aligns.append("r") @@ -157,13 +158,15 @@ def ordered(index: int, override: str | None, rank) -> list[str]: table = common.render_table(["Backend", "Sample"] + quants, ["l", "l"] + ["r"] * len(quants), body, rule_fill=True, max_pad=20) - return perf_methodology(rows) + [""] + table + [""] + perf_provenance(machine, rows) + return (perf_methodology(record, rows) + [""] + table + [""] + + perf_provenance(record, machine, rows)) -def perf_methodology(rows: dict) -> list[str]: +def perf_methodology(record: dict, rows: dict) -> list[str]: """What a cell is. Iterations and warmup are claimed only for rows that name the profile they were measured under.""" - line = "Compute latency (mel + encode + decode), speedup over realtime in parentheses" + stages = "mel + encode" if record.get("role") == "langid" else "mel + encode + decode" + line = f"Compute latency ({stages}), speedup over realtime in parentheses" ids = sorted({row["publication_profile"] for row in rows.values() if row.get("publication_profile")}) claims = [] @@ -176,9 +179,9 @@ def perf_methodology(rows: dict) -> list[str]: return [line + "."] -def perf_provenance(machine: str, rows: dict) -> list[str]: +def perf_provenance(record: dict, machine: str, rows: dict) -> list[str]: """Where the numbers came from: machine, engine commit, date.""" - _, profile = profiles.load_profile() + _, profile = profiles.profile_for(record) display = profiles.machine_display(profile, machine) builds: dict[tuple, int] = {} for row in rows.values(): @@ -331,6 +334,37 @@ def block_accuracy(record: dict, attrs: dict[str, str]) -> list[str]: return out +def block_agreement(record: dict, attrs: dict[str, str]) -> list[str]: + """Language ID: per shipped GGUF, the headline accuracy with its interval + and the scored run's top-1 agreement with the reference (the ship gate), + then which build measured it.""" + rows = common.headline_rows(record) + reference = (spec_for(record).get("validation") or {}).get("reference", "reference") + body = [] + for item in record.get("downloads", []): + row = rows.get(item["quant"]) + if row is None: + continue + agreement = row.get("agreement") or {} + lo, hi = row["ci95"] + delta = agreement.get("max_abs_logit_delta") + body.append([item["quant"], + common.fmt_err(row) + ("" if lo is None else f" ({lo:.2f}-{hi:.2f})"), + f"{agreement['n_agree']} / {agreement['n']}" if agreement else "-", + "-" if delta is None else f"{delta:.2g}"]) + if not any(cells[2] != "-" for cells in body): + raise RenderError("no headline row carries an agreement") + table = common.render_table( + ["GGUF", f"{common.metric_label(common.headline(record)['metric'])} (95% CI)", + f"Top-1 agreement with {reference}", "Max abs logit difference"], + ["l", "r", "r", "r"], body) + builds = sorted({(row["engine_sha"], row.get("measured_on") or "") + for row in rows.values() if row.get("engine_sha")}) + line = "Measured at " + "; ".join( + f"transcribe.cpp `{sha}`" + (f" on {date}" if date else "") for sha, date in builds) + "." + return table + ["", line] if builds else table + + def block_family(records: dict[str, dict], attrs: dict[str, str]) -> list[str]: """A family roll-up: one row per variant, headline number at one quant.""" names = [name for name in attrs.get("variants", "").split(",") if name] @@ -352,7 +386,7 @@ def block_family(records: dict[str, dict], attrs: dict[str, str]) -> list[str]: body.append([ f"`{name}`", common.fmt_params(record["params"]), common.languages_summary(record), common.fmt_size(download["size_bytes"]) if download else "-", - (f"{common.headline_label(record)} ({headline['metric'].upper()})" + (f"{common.headline_label(record)} ({common.metric_label(headline['metric'])})" if headline else "-"), common.fmt_err(common.headline_rows(record).get(quant)), common.capabilities_summary(record), link]) @@ -402,7 +436,7 @@ def block_family_index(records: dict[str, dict], attrs: dict[str, str]) -> list[ BLOCKS = {"downloads": block_downloads, "perf": block_perf, "intro": block_intro, "prose": block_prose, "accuracy": block_accuracy, - "recipe": block_recipe, "pin": block_pin} + "recipe": block_recipe, "pin": block_pin, "agreement": block_agreement} # -------------------------------------------------------------------------- diff --git a/scripts/catalog/test_db_mapping.py b/scripts/catalog/test_db_mapping.py index 44a23f23a..04ddf46eb 100644 --- a/scripts/catalog/test_db_mapping.py +++ b/scripts/catalog/test_db_mapping.py @@ -16,7 +16,8 @@ # Row properties that are flattened or renamed rather than stored one to one. ACCURACY_MAPPED = {"dataset": "dataset_id", "split": "dataset_id", "language": "dataset_id", - "ci95": "ci_lo/ci_hi", "errors": "substitutions/deletions/insertions"} + "ci95": "ci_lo/ci_hi", "errors": "substitutions/deletions/insertions", + "agreement": "agreement_n_agree/agreement_n/agreement_max_abs_logit_delta"} SPEED_MAPPED = {} diff --git a/scripts/convert-ecapa_tdnn.py b/scripts/convert-ecapa_tdnn.py new file mode 100644 index 000000000..679fef953 --- /dev/null +++ b/scripts/convert-ecapa_tdnn.py @@ -0,0 +1,437 @@ +#!/usr/bin/env python3 +""" +convert-ecapa_tdnn.py - convert SpeechBrain's VoxLingua107 ECAPA-TDNN +language-ID checkpoint into the F32 reference GGUF of the ecapa_tdnn family. +Loads the model through SpeechBrain (the only implementation), asserts the +hyperparameters hard-coded below against the live modules, and writes the +tensor layout of src/arch/ecapa_tdnn/weights.h. Quantize afterwards with +tools/transcribe-quantize. See docs/porting/families/ecapa_tdnn.md. + + uv run --project scripts/envs/ecapa_tdnn \ + scripts/convert-ecapa_tdnn.py speechbrain/lang-id-voxlingua107-ecapa +""" + +from __future__ import annotations + +import argparse +import os +import sys +from pathlib import Path + +# XET-backed transfers occasionally stall on this checkpoint's small .ckpt +# files. Set before huggingface_hub is imported (via speechbrain). +os.environ.setdefault("HF_HUB_DISABLE_XET", "1") + +import numpy as np +import torch +from gguf import GGMLQuantizationType, LlamaFileType + +REPO_ROOT = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(REPO_ROOT / "scripts")) + +from lib.gguf_common import ( # noqa: E402 + add_general_identity, + canonicalize_normalize, + encode_for_gguf, + gguf_name, + gguf_writer, + slug_from_repo_id, +) +from lib.hf_source import looks_like_repo_id, resolve_model_dir # noqa: E402 + +DEFAULT_REPO = "speechbrain/lang-id-voxlingua107-ecapa" +DEFAULT_REVISION = "0253049ae131d6a4be1c4f0d8b0ff483a0f8c8e9" +DEFAULT_VARIANT = "lang-id-voxlingua107-ecapa" +ARCH = "ecapa_tdnn" + +# Written into the GGUF and asserted against the live modules below. +SAMPLE_RATE = 16000 +N_FFT = 400 +WIN = 400 +HOP = 160 +N_STFT = N_FFT // 2 + 1 # 201 +N_MELS = 60 +LOG_FLOOR = 1e-10 +TOP_DB = 80.0 +CHANNELS = [1024, 1024, 1024, 1024, 3072] +KERNEL_SIZES = [5, 3, 3, 3, 1] +DILATIONS = [1, 2, 3, 4, 1] +RES2NET_SCALE = 8 +SE_CHANNELS = 128 +ATTENTION_CHANNELS = 128 +ASP_EPS = 1e-12 +EMB_DIM = 256 +HIDDEN = 512 +N_LABELS = 107 +LEAKY_SLOPE = 0.01 +BN_EPS = 1e-5 +N_PARAMS_EMBEDDING = 21_058_432 +N_PARAMS_CLASSIFIER = 188_011 + +# Modern ISO codes users type -> the legacy codes VoxLingua107 uses. +LABEL_ALIASES = ["he=iw", "jv=jw", "fil=tl", "nb=no"] + + +def fail(msg: str) -> None: + print(f"error: {msg}", file=sys.stderr) + raise SystemExit(2) + + +def load_reference(src: Path): + """SpeechBrain EncoderClassifier in eval mode. + + `pretrained_path` must be overridden to the local snapshot, otherwise + SpeechBrain re-fetches the .ckpt files from the HF `main` branch and the + pinned revision is ignored. `savedir` must differ from `src` (SpeechBrain + symlinks the loadables into it). + """ + from speechbrain.inference.classifiers import EncoderClassifier + + savedir = REPO_ROOT / "build" / "speechbrain" / ARCH + savedir.mkdir(parents=True, exist_ok=True) + clf = EncoderClassifier.from_hparams( + source=str(src), savedir=str(savedir), + overrides={"pretrained_path": str(src)}, run_opts={"device": "cpu"}) + clf.mods.eval() + return clf + + +def bn_of(module) -> torch.nn.BatchNorm1d: + """Walk SpeechBrain's BatchNorm1d wrapper(s) to the torch layer (the + wrapper depth differs between the TDNN blocks and `asp_bn`).""" + m = module + for _ in range(4): + if isinstance(m, torch.nn.BatchNorm1d): + return m + m = getattr(m, "norm", None) + fail(f"no torch BatchNorm1d under {type(module).__name__}") + + +def assert_architecture(clf) -> dict[str, list[str]]: + """Check every hyperparameter the GGUF hard-codes that tensor shapes (and + hence the C++ loader) cannot see. Returns the parsed label table.""" + em, cl = clf.mods.embedding_model, clf.mods.classifier + fe, mvn = clf.mods.compute_features, clf.mods.mean_var_norm + st, fb = fe.compute_STFT, fe.compute_fbanks + asp = em.asp + + def name(m) -> str: + return type(m).__name__ + + checks = [ # (label, actual, expected) + ("param dtypes", {p.dtype for p in clf.mods.parameters()}, {torch.float32}), + ("stft.n_fft", st.n_fft, N_FFT), + ("stft.win_length", st.win_length, WIN), + ("stft.hop_length", st.hop_length, HOP), + ("stft.sample_rate", st.sample_rate, SAMPLE_RATE), + ("stft.center", st.center, True), + ("stft.pad_mode", st.pad_mode, "constant"), + ("stft.onesided", st.onesided, True), + ("stft.normalized_stft", st.normalized_stft, False), + ("stft.window == hamming_window(400)", + torch.equal(st.window.to(torch.float32).cpu(), torch.hamming_window(WIN)), True), + ("fbank.n_mels", fb.n_mels, N_MELS), + ("fbank.filter_shape", fb.filter_shape, "triangular"), + ("fbank.f_min", fb.f_min, 0.0), + ("fbank.f_max", fb.f_max, 8000.0), + ("fbank.n_stft", fb.n_stft, N_STFT), + ("fbank.power_spectrogram", fb.power_spectrogram, 2), + ("fbank.multiplier", fb.multiplier, 10), + ("fbank.amin", fb.amin, LOG_FLOOR), + ("fbank.ref_value", fb.ref_value, 1.0), + ("fbank.db_multiplier", fb.db_multiplier, 0.0), + ("fbank.top_db", fb.top_db, TOP_DB), + ("fbank.log_mel", fb.log_mel, True), + ("fbank.freeze", fb.freeze, True), + ("fbank.param_rand_factor", fb.param_rand_factor, 0.0), + ("fbank.deltas", fe.deltas, False), + ("fbank.context", fe.context, False), + ("mvn.norm_type", mvn.norm_type, "sentence"), + ("mvn.std_norm", mvn.std_norm, False), + ("mvn.length_dim", mvn.length_dim, 1), + ("mvn.avoid_padding_norm", mvn.avoid_padding_norm, False), + ("len(blocks)", len(em.blocks), 4), + ("asp.global_context", asp.global_context, True), + ("asp.eps", asp.eps, ASP_EPS), + ("asp.tanh", name(asp.tanh), "Tanh"), + ("asp_bn.eps", bn_of(em.asp_bn).eps, BN_EPS), + ("cls.act", name(cl.act), "LeakyReLU"), + ("cls.act slope", cl.act.negative_slope, LEAKY_SLOPE), + ("cls.bn0.eps", bn_of(cl.norm).eps, BN_EPS), + ("len(cls.DNN)", len(cl.DNN), 1), + ("cls.hidden act", name(cl.DNN.block_0.act), "LeakyReLU"), + ("cls.hidden slope", cl.DNN.block_0.act.negative_slope, LEAKY_SLOPE), + ("cls.bn1.eps", bn_of(cl.DNN.block_0.norm).eps, BN_EPS), + ("embedding params", sum(p.numel() for p in em.parameters()), N_PARAMS_EMBEDDING), + ("classifier params", sum(p.numel() for p in cl.parameters()), N_PARAMS_CLASSIFIER), + ] + + def conv(label, c, k, d=1): + # Every conv: kernel k, dilation d, reflect 'same' pad, bias, groups 1. + return [ + (f"{label}.kernel", c.conv.weight.shape[2], k), + (f"{label}.dilation", c.dilation, d), + (f"{label}.padding", (c.padding, c.padding_mode), ("same", "reflect")), + (f"{label}.groups", c.conv.groups, 1), + (f"{label}.has_bias", c.conv.bias is not None, True), + ] + + def tdnn(label, blk, k, d): + # conv -> ReLU -> BN -> Dropout1d(p=0) (identity in eval). + return conv(f"{label}.conv", blk.conv, k, d) + [ + (f"{label}.activation", name(blk.activation), "ReLU"), + (f"{label}.bn.eps", bn_of(blk.norm).eps, BN_EPS), + (f"{label}.dropout.p", blk.dropout.p, 0.0), + ] + + checks += tdnn("blocks.0", em.blocks[0], KERNEL_SIZES[0], DILATIONS[0]) + for i in (1, 2, 3): + b, p = em.blocks[i], f"blocks.{i}" + r2 = b.res2net_block + checks += [(f"{p}.shortcut", b.shortcut, None), + (f"{p}.res2net.scale", r2.scale, RES2NET_SCALE), + (f"{p}.res2net blocks", len(r2.blocks), RES2NET_SCALE - 1), + (f"{p}.se.relu", name(b.se_block.relu), "ReLU"), + (f"{p}.se.sigmoid", name(b.se_block.sigmoid), "Sigmoid")] + checks += tdnn(f"{p}.tdnn1", b.tdnn1, 1, 1) + tdnn(f"{p}.tdnn2", b.tdnn2, 1, 1) + for j, sub in enumerate(r2.blocks): + checks += tdnn(f"{p}.res2net.{j}", sub, KERNEL_SIZES[i], DILATIONS[i]) + checks += conv(f"{p}.se.conv1", b.se_block.conv1, 1) + checks += conv(f"{p}.se.conv2", b.se_block.conv2, 1) + checks += tdnn("mfa", em.mfa, KERNEL_SIZES[4], DILATIONS[4]) + checks += tdnn("asp.tdnn", asp.tdnn, 1, 1) + checks += conv("asp.conv", asp.conv, 1) + conv("fc", em.fc, 1) + + ind2lab = clf.hparams.label_encoder.ind2lab + raw = [str(ind2lab[i]) for i in range(len(ind2lab))] + pairs = [r.split(": ", 1) for r in raw] + codes = [c.strip() for c, *_ in pairs] + checks += [("label count", len(raw), N_LABELS), + ("labels in ': ' form", all(len(x) == 2 for x in pairs), True), + ("distinct label codes", len(set(codes)), N_LABELS)] + for alias in LABEL_ALIASES: + a, c = alias.split("=", 1) + checks += [(f"alias target {c!r} is a label", c in codes, True), + (f"alias {a!r} is not a label", a in codes, False)] + + for label, got, want in checks: + if got != want: + fail(f"architecture mismatch: {label}: got {got!r}, expected {want!r} " + "(only the pinned SpeechBrain VoxLingua107 ECAPA-TDNN is supported)") + print(f"Architecture verified: {len(checks)} checks passed.") + return {"codes": codes, "names": [x[1].strip() for x in pairs]} + + +def capture_fbank_matrix(clf) -> np.ndarray: + """The exact [201, 60] fbank matrix SpeechBrain's forward builds. + + Filterbank has no stored matrix; it is built inside forward() by + `_create_fbank_matrix`. Wrap that, run 0.1 s of silence, keep the result + (bit-identical to the reference; a float64 rebuild is off by ~1e-5). + """ + fbanks = clf.mods.compute_features.compute_fbanks + original = fbanks._create_fbank_matrix + captured = [] + + def capture(*a): + captured.append(original(*a)) + return captured[-1] + + fbanks._create_fbank_matrix = capture + try: + with torch.inference_mode(): + clf.mods.compute_features(torch.zeros(1, SAMPLE_RATE // 10)) + finally: + fbanks._create_fbank_matrix = original + if len(captured) != 1 or tuple(captured[0].shape) != (N_STFT, N_MELS): + fail(f"fbank capture: {len(captured)} calls, expected one [{N_STFT}, {N_MELS}] matrix") + return captured[0].detach().to(dtype=torch.float32, device="cpu").numpy() + + +def f64(t: torch.Tensor) -> np.ndarray: + return t.detach().to(torch.float64).cpu().numpy() + + +def bn_affine(module) -> tuple[np.ndarray, np.ndarray]: + """BatchNorm -> (scale, shift), `y = x * scale + shift`, in float64.""" + bn = bn_of(module) + scale = f64(bn.weight) / np.sqrt(f64(bn.running_var) + float(bn.eps)) + return scale, f64(bn.bias) - f64(bn.running_mean) * scale + + +def fold_bn_into_linear(module, w: np.ndarray, b: np.ndarray): + """`W(x * s + t) + b = (W diag(s)) x + (W t + b)`: fold a BN forward into + the linear map that directly consumes it.""" + s, t = bn_affine(module) + return w * s[None, :], w @ t + b + + +def convert(clf, out_path: Path, *, variant: str, repo_id: str, revision: str, + filters_201x60: np.ndarray, labels: dict[str, list[str]]) -> None: + em, cl = clf.mods.embedding_model, clf.mods.classifier + + out_path.parent.mkdir(parents=True, exist_ok=True) + writer = gguf_writer(str(out_path), ARCH) + + add_general_identity( + writer, + name=repo_id, + basename=slug_from_repo_id(repo_id), + size_label="21M", + file_type=int(LlamaFileType.ALL_F32), + license="apache-2.0", + license_name="Apache License 2.0", + license_link="https://www.apache.org/licenses/LICENSE-2.0", + author="SpeechBrain", + organization="speechbrain", + repo_url=f"https://huggingface.co/{repo_id}", + source_url=f"https://huggingface.co/{repo_id}", + languages=labels["codes"], + tags=["language-identification", "ecapa-tdnn", "voxlingua107", + "speechbrain"], + description=( + "SpeechBrain ECAPA-TDNN spoken-language identification trained " + "on VoxLingua107 (107 languages). Converted from " + f"{repo_id} for transcribe.cpp." + ), + ) + writer.add_string("general.source.commit", revision) + + writer.add_string("stt.variant", variant) + writer.add_string("stt.frontend.type", "speechbrain_fbank") + writer.add_uint32("stt.frontend.sample_rate", SAMPLE_RATE) + writer.add_uint32("stt.frontend.n_fft", N_FFT) + writer.add_uint32("stt.frontend.hop_length", HOP) + writer.add_uint32("stt.frontend.win_length", WIN) + writer.add_uint32("stt.frontend.num_mels", N_MELS) + writer.add_string("stt.frontend.window", "hamming_periodic") + writer.add_string("stt.frontend.pad_mode", "constant") + writer.add_float32("stt.frontend.log_clamp_min", LOG_FLOOR) + writer.add_float32("stt.frontend.top_db", TOP_DB) + writer.add_string("stt.frontend.normalize", canonicalize_normalize("sentence_mean")) + + writer.add_array("stt.ecapa_tdnn.channels", CHANNELS) + writer.add_array("stt.ecapa_tdnn.kernel_sizes", KERNEL_SIZES) + writer.add_array("stt.ecapa_tdnn.dilations", DILATIONS) + writer.add_uint32("stt.ecapa_tdnn.res2net_scale", RES2NET_SCALE) + writer.add_uint32("stt.ecapa_tdnn.se_channels", SE_CHANNELS) + writer.add_uint32("stt.ecapa_tdnn.attention_channels", ATTENTION_CHANNELS) + writer.add_float32("stt.ecapa_tdnn.asp_eps", ASP_EPS) + writer.add_uint32("stt.ecapa_tdnn.embedding_dim", EMB_DIM) + writer.add_uint32("stt.ecapa_tdnn.classifier_hidden", HIDDEN) + writer.add_float32("stt.ecapa_tdnn.classifier_leaky_slope", LEAKY_SLOPE) + + writer.add_array("stt.langid.labels.codes", labels["codes"]) + writer.add_array("stt.langid.labels.names", labels["names"]) + writer.add_array("stt.langid.labels.aliases", LABEL_ALIASES) + + def add(name: str, arr: np.ndarray) -> None: + a = np.ascontiguousarray(arr.astype(np.float32)) + encoded, raw_dtype = encode_for_gguf(a, GGMLQuantizationType.F32) + writer.add_tensor(name, encoded, raw_dtype=raw_dtype) + + def add_tdnn(p: str, blk, split: tuple[str, ...] = ()) -> None: + """conv weight + bias, then the BN as a scale/shift pair (TDNNBlock is + conv -> ReLU -> BN, so the BN cannot fold into the conv). + + k>1 kernels are `

.conv.weight`, tap-major numpy [K, OC, IC] (ggml + ne [IC, OC, K]); 1x1 kernels are `

.weight` [OC, IC]. `split` cuts a + 1x1 weight's input columns into one `

..weight` per member of + the concatenation SpeechBrain feeds it. + """ + w = f64(blk.conv.conv.weight) + c = f"{p}.conv" if w.shape[2] > 1 else p + if w.shape[2] > 1: + add(f"{c}.weight", np.transpose(w, (2, 0, 1))) + elif split: + n = w.shape[1] // len(split) + for k, part in enumerate(split): + add(f"{p}.{part}.weight", w[:, k * n:(k + 1) * n, 0]) + else: + add(f"{c}.weight", w[:, :, 0]) + add(f"{c}.bias", f64(blk.conv.conv.bias)) + s, t = bn_affine(blk.norm) + add(f"{p}.bn.scale", s) + add(f"{p}.bn.shift", t) + + # Mel-major: numpy [60, 201] -> ggml ne [201, 60]. + add("frontend.mel_filterbank", filters_201x60.T) + add_tdnn("blk.0", em.blocks[0]) + for i in (1, 2, 3): + b, p = em.blocks[i], f"blk.{i}" + add_tdnn(f"{p}.tdnn1", b.tdnn1) + # res2.{j} convolves chunk j+1 (plus the previous sub-block's output). + for j, sub in enumerate(b.res2net_block.blocks): + add_tdnn(f"{p}.res2.{j}", sub) + add_tdnn(f"{p}.tdnn2", b.tdnn2) + for c, m in (("c1", b.se_block.conv1), ("c2", b.se_block.conv2)): + add(f"{p}.se.{c}.weight", f64(m.conv.weight)[:, :, 0]) + add(f"{p}.se.{c}.bias", f64(m.conv.bias)) + # mfa consumes cat(blocks.1..3); asp.tdnn consumes cat([x, mean, std]). + add_tdnn("mfa", em.mfa, ("w1", "w2", "w3")) + add_tdnn("asp.tdnn", em.asp.tdnn, ("x", "mean", "std")) + add("asp.attn.weight", f64(em.asp.conv.conv.weight)[:, :, 0]) + add("asp.attn.bias", f64(em.asp.conv.conv.bias)) + + # BNs with no activation before their linear map fold forward into it: + # asp_bn -> fc, cls.bn0 -> cls.l1, cls.bn1 -> cls.out. The classifier's + # LeakyReLUs sit before each BN, so they survive in the graph. + for name, bn, lin in ( + ("fc", em.asp_bn, em.fc.conv), + ("cls.l1", cl.norm, cl.DNN.block_0.linear.w), + ("cls.out", cl.DNN.block_0.norm, cl.out.w), + ): + w = f64(lin.weight) + w, b = fold_bn_into_linear(bn, w.reshape(w.shape[0], w.shape[1]), f64(lin.bias)) + add(f"{name}.weight", w) + add(f"{name}.bias", b) + + writer.write_header_to_file() + writer.write_kv_data_to_file() + writer.write_tensors_to_file() + writer.close() + print(f"Wrote {out_path} ({out_path.stat().st_size:,} bytes)") + + +def main(argv: list[str] | None = None) -> int: + p = argparse.ArgumentParser( + description="Convert SpeechBrain's VoxLingua107 ECAPA-TDNN checkpoint " + "to the F32 reference GGUF.") + p.add_argument("model", nargs="?", default=DEFAULT_REPO, + help=f"HF repo id or local checkpoint dir " + f"(default: {DEFAULT_REPO})") + p.add_argument("out_path", type=Path, nargs="?", + help="Output .gguf path (derived from the repo id when " + "omitted)") + p.add_argument("--revision", default=DEFAULT_REVISION, + help="HF revision (commit SHA) to pin the download to") + p.add_argument("--repo-id", default=None, + help="HF repo id used for the output slug and metadata " + "when converting from a local path") + args = p.parse_args(argv) + + repo_id = args.repo_id or (args.model if looks_like_repo_id(args.model) + else None) + if repo_id is None: + fail("cannot infer the source repo id from a local path; pass --repo-id") + + src = resolve_model_dir(args.model, args.revision) + needed = ["hyperparams.yaml", "embedding_model.ckpt", "classifier.ckpt", + "label_encoder.txt"] + missing = [f for f in needed if not (src / f).exists()] + if missing: + fail(f"{src} is missing {missing}") + slug = slug_from_repo_id(repo_id) + out_path = args.out_path or (REPO_ROOT / "models" / slug / + gguf_name(slug, "F32")) + + clf = load_reference(src) + labels = assert_architecture(clf) + filters = capture_fbank_matrix(clf) + convert(clf, out_path, variant=DEFAULT_VARIANT, repo_id=repo_id, + revision=args.revision, filters_201x60=filters, labels=labels) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/scripts/dump_reference_ecapa_tdnn_speechbrain.py b/scripts/dump_reference_ecapa_tdnn_speechbrain.py new file mode 100644 index 000000000..8acb2c260 --- /dev/null +++ b/scripts/dump_reference_ecapa_tdnn_speechbrain.py @@ -0,0 +1,228 @@ +#!/usr/bin/env python3 +""" +dump_reference_ecapa_tdnn_speechbrain.py - generate SpeechBrain VoxLingua107 +ECAPA-TDNN reference tensors (SpeechBrain is the only implementation). + + uv run --project scripts/envs/ecapa_tdnn \ + scripts/dump_reference_ecapa_tdnn_speechbrain.py encoder \ + --audio samples/fleurs-en.wav \ + --out build/validate/ecapa_tdnn/lang-id-voxlingua107-ecapa/fleurs-en/ref + +The single `encoder` stage hooks one `classify_batch` call and writes +`.f32` / `.json` pairs (scripts/lib/ref_dump.py) for the dump +points listed in cmd_encoder, plus prediction.json (top-1 label). Every hook must +fire exactly once. Activations are written time-major `[T, C]` (PyTorch's +`[1, C, T]` transposed) to match the C++ dumper (src/arch/ecapa_tdnn/graph.h). +""" + +from __future__ import annotations + +import argparse +import json +import math +import sys +from pathlib import Path + +import numpy as np +import torch + +REPO_ROOT = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(REPO_ROOT / "scripts")) + +from lib.hf_source import resolve_model_dir # noqa: E402 +from lib.ref_dump import write_tensor # noqa: E402 + +DEFAULT_REPO = "speechbrain/lang-id-voxlingua107-ecapa" +DEFAULT_REVISION = "0253049ae131d6a4be1c4f0d8b0ff483a0f8c8e9" + +SAMPLE_RATE = 16000 +HOP = 160 +N_MELS = 60 +N_LABELS = 107 + + +def configure_torch(args: argparse.Namespace) -> None: + torch.manual_seed(0) + if args.torch_threads > 0: + torch.set_num_threads(args.torch_threads) + try: + torch.set_num_interop_threads(1) + except RuntimeError: + pass # already set earlier in this process + torch.use_deterministic_algorithms(True, warn_only=True) + + +def load_reference(args: argparse.Namespace): + """Return `(EncoderClassifier, checkpoint_dir)`, in eval mode. + + `pretrained_path` must be overridden to the local snapshot, otherwise + SpeechBrain re-fetches the .ckpt files from the HF `main` branch and the + pinned revision is ignored. `savedir` must differ from the snapshot + (SpeechBrain symlinks the loadables into it). + """ + from speechbrain.inference.classifiers import EncoderClassifier + + src = resolve_model_dir(args.model, args.revision) + savedir = REPO_ROOT / "build" / "speechbrain" / src.parent.name + savedir.mkdir(parents=True, exist_ok=True) + clf = EncoderClassifier.from_hparams( + source=str(src), savedir=str(savedir), + overrides={"pretrained_path": str(src)}, run_opts={"device": args.device}) + clf.mods.eval() + return clf, src + + +def split_label(label: str) -> tuple[str, str]: + """`'zh: Chinese'` -> `('zh', 'Chinese')`.""" + code, _, name = label.partition(": ") + return code.strip(), name.strip() + + +def to_np(t: torch.Tensor, shape: tuple[int, ...], *, ct: bool = False) -> np.ndarray: + """Hooked tensor -> contiguous f32 numpy of exactly `shape`. + + 1-D shapes flatten; otherwise the batch dim is dropped and `ct=True` + transposes `[C, T]` to time-major `[T, C]`. The shape is asserted so a + tensor in the wrong layout cannot pass as its transpose. + """ + a = t.detach().to(dtype=torch.float32, device="cpu") + if len(shape) == 1: + a = a.reshape(-1) + else: + a = a[0].transpose(0, 1) if ct else a[0] + if tuple(a.shape) != shape: + raise SystemExit(f"error: hooked tensor is {tuple(a.shape)}, expected {shape}") + return np.ascontiguousarray(a.numpy(), dtype=np.float32) + + +def cmd_encoder(args: argparse.Namespace) -> int: + import soundfile as sf + import speechbrain + + configure_torch(args) + clf, src_dir = load_reference(args) + + audio_path = Path(args.audio).expanduser().resolve() + out_dir = Path(args.out).expanduser().resolve() + out_dir.mkdir(parents=True, exist_ok=True) + + pcm, sr = sf.read(str(audio_path), dtype="float32", always_2d=False) + if pcm.ndim > 1: + pcm = pcm.mean(axis=1) + if int(sr) != SAMPLE_RATE: + raise SystemExit(f"error: {audio_path} is {sr} Hz; resample to {SAMPLE_RATE} Hz") + pcm = np.ascontiguousarray(pcm, dtype=np.float32) + n_samples = int(pcm.size) + expected_T = n_samples // HOP + 1 + + device = torch.device(args.device) + wav = torch.from_numpy(pcm).to(device).unsqueeze(0) # [1, N] + # wav_lens is RELATIVE to the padded batch max (1.0 for one unpadded + # clip); a sample count would mask the wrong span, silently. + wav_lens = torch.ones(1, device=device) + + em, cl = clf.mods.embedding_model, clf.mods.classifier + # (dump name, module, stage, channels, layout). Layouts: "tc" is already + # [1, T, C]; "ct" is [1, C, T] and gets transposed; "vec" is flattened. + points = [ + ("fe.mel", clf.mods.mean_var_norm, "frontend", N_MELS, "tc"), + ("enc.blk.0.out", em.blocks[0], "encoder", 1024, "ct"), + ("enc.blk.1.tdnn1.out", em.blocks[1].tdnn1, "encoder", 1024, "ct"), + ("enc.blk.1.res2.out", em.blocks[1].res2net_block, "encoder", 1024, "ct"), + ("enc.blk.1.se.out", em.blocks[1].se_block, "encoder", 1024, "ct"), # pre-residual + ("enc.blk.1.out", em.blocks[1], "encoder", 1024, "ct"), + ("enc.blk.2.out", em.blocks[2], "encoder", 1024, "ct"), + ("enc.blk.3.out", em.blocks[3], "encoder", 1024, "ct"), + ("enc.mfa.out", em.mfa, "encoder", 3072, "ct"), + ("enc.asp.attn_logits", em.asp.conv, "encoder", 3072, "ct"), # pre-softmax + ("enc.asp.out", em.asp, "encoder", 6144, "vec"), + ("enc.emb", em.fc, "encoder", 256, "vec"), + ("cls.hidden", cl.DNN.block_0.act, "classifier", 512, "vec"), # pre-BN1 + ("cls.logits_raw", cl.out, "classifier", N_LABELS, "vec"), + ] + + captured: dict[str, list] = {name: [] for name, *_ in points} + + def grab(name): + def hook(_mod, _inp, out): + captured[name].append(out[0] if isinstance(out, tuple) else out) + return hook + + handles = [m.register_forward_hook(grab(name)) for name, m, *_ in points] + try: + with torch.inference_mode(): + out_prob, score, index, text_lab = clf.classify_batch(wav, wav_lens) + finally: + for h in handles: + h.remove() + + # Zero firings = module not on the path; two = the model ran twice. + bad = {n: len(v) for n, v in captured.items() if len(v) != 1} + if bad: + raise SystemExit(f"error: unexpected hook call counts: {bad}") + out = {n: v[0] for n, v in captured.items()} + if not torch.equal(torch.log_softmax(out["cls.logits_raw"].reshape(-1), -1), + out_prob.reshape(-1)): + raise SystemExit("error: log_softmax(cls.logits_raw) != classify_batch out_prob") + + T = int(out["fe.mel"].shape[1]) + if T != expected_T: + raise SystemExit(f"error: fe.mel has T={T}, but the C++ front end assumes " + f"floor({n_samples}/{HOP}) + 1 = {expected_T}") + + source = { + "kind": "ecapa-tdnn-speechbrain", + "framework": "speechbrain", + "framework_version": speechbrain.__version__, + "torch_version": torch.__version__, + "model": args.model, + "revision": args.revision, + "checkpoint_dir": str(src_dir), + "device": args.device, + "torch_threads": args.torch_threads, + "model_dtype": "f32", + "audio": audio_path.name, + "n_samples": n_samples, + "sample_rate": SAMPLE_RATE, + "n_frames": T, + "n_mels": N_MELS, + } + for name, _, stage, c, layout in points: + a = to_np(out[name], (c,) if layout == "vec" else (T, c), ct=layout == "ct") + print(f" {name}: shape={a.shape} min={a.min():.4e} max={a.max():.4e} " + f"mean={a.mean():.6e}") + write_tensor(name, a, stage=stage, source=source, out_dir=out_dir) + + label_index = int(index.item()) + label = str(clf.hparams.label_encoder.ind2lab[label_index]) + if str(text_lab[0]) != label: + raise SystemExit(f"error: decoded label {text_lab[0]!r} != ind2lab[{label_index}] {label!r}") + code, name = split_label(label) + prediction = {"label_index": label_index, "code": code} + (out_dir / "prediction.json").write_text(json.dumps(prediction, indent=2) + "\n") + print(f"prediction: {code} ({name}) prob={math.exp(float(score.item())):.4f}") + return 0 + + +def main(argv: list[str] | None = None) -> int: + p = argparse.ArgumentParser( + description="SpeechBrain VoxLingua107 ECAPA-TDNN reference dumper.") + sub = p.add_subparsers(dest="cmd", required=True) + ep = sub.add_parser("encoder", help="Run the full forward pass; dump front " + "end, encoder, classifier and prediction.json") + ep.add_argument("--model", default=DEFAULT_REPO, + help=f"HF repo id or local checkpoint dir (default: {DEFAULT_REPO})") + ep.add_argument("--revision", default=DEFAULT_REVISION, + help="HF revision (commit SHA); ignored for a local --model") + ep.add_argument("--audio", required=True, help="16 kHz mono WAV path") + ep.add_argument("--out", required=True, help="Output directory for dumps") + ep.add_argument("--device", default="cpu", help="torch device (default: cpu)") + ep.add_argument("--torch-threads", type=int, default=1, + help="torch.set_num_threads (0 = unchanged)") + ep.set_defaults(func=cmd_encoder) + args = p.parse_args(argv) + return args.func(args) + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/scripts/envs/ecapa_tdnn/pyproject.toml b/scripts/envs/ecapa_tdnn/pyproject.toml new file mode 100644 index 000000000..66e98bd66 --- /dev/null +++ b/scripts/envs/ecapa_tdnn/pyproject.toml @@ -0,0 +1,29 @@ +[project] +name = "transcribe-ecapa_tdnn-env" +version = "0.1.0" +description = "Python env for ECAPA-TDNN VoxLingua107 reference dumps and conversion (SpeechBrain)" +requires-python = ">=3.11,<3.13" +dependencies = [ + # Exact pin. The VoxLingua107 checkpoint ships as a SpeechBrain + # hyperparams.yaml that instantiates SpeechBrain classes by name, so the + # module tree the reference dumper hooks (`embedding_model.blocks[1].se_block`, + # `classifier.DNN.block_0.act`, `compute_features.compute_fbanks`) is + # version-specific, and so is the arithmetic inside them. + # tests/tolerances/ecapa_tdnn.json was measured against THIS version + # (recorded as reference.revision in + # tests/golden/ecapa_tdnn/lang-id-voxlingua107-ecapa.manifest.json); a silent bump can + # move both the hook points and the numbers. + "speechbrain==1.1.1", + # CPU torch. The parity regime is fp32 on CPU with one thread; no + # CUDA/Metal build is needed or wanted here. + "torch==2.13.0", + "torchaudio==2.11.0", + "numpy>=1.26", + "soundfile>=0.12", + "huggingface-hub>=0.30", + # Used by the converter (T3), not by the dumper. + "gguf>=0.10", + # FLEURS lives in the local HF hub cache as parquet; the fixture builder + # reads it with pyarrow rather than pulling in `datasets`. + "pyarrow>=15", +] diff --git a/scripts/envs/ecapa_tdnn/uv.lock b/scripts/envs/ecapa_tdnn/uv.lock new file mode 100644 index 000000000..d7db588e2 --- /dev/null +++ b/scripts/envs/ecapa_tdnn/uv.lock @@ -0,0 +1,1031 @@ +version = 1 +requires-python = ">=3.11, <3.13" +resolution-markers = [ + "python_full_version >= '3.12' and sys_platform == 'win32'", + "python_full_version >= '3.12' and sys_platform == 'emscripten'", + "python_full_version >= '3.12' and sys_platform != 'emscripten' and sys_platform != 'win32'", + "python_full_version < '3.12' and sys_platform == 'win32'", + "python_full_version < '3.12' and sys_platform == 'emscripten'", + "python_full_version < '3.12' and sys_platform != 'emscripten' and sys_platform != 'win32'", +] + +[[package]] +name = "anyio" +version = "4.14.2" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "idna" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/61/cc/a381afa6efea9f496eff839d4a6a1aed3bfafc7b3ab4b0d1b243a12573dd/anyio-4.14.2.tar.gz", hash = "sha256:cfa139f3ed1a23ee8f88a145ddb5ac7605b8bbfd8592baacd7ce3d8bb4313c7f", size = 260176 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/da/35/f2287558c17e29fafc8ef3daf819bb9834061cfa43bff8014f7df7f63bdc/anyio-4.14.2-py3-none-any.whl", hash = "sha256:9f505dda5ac9f0c8309b5e8bd445a8c2bf7246f3ce950121e45ea15bc41d1494", size = 125813 }, +] + +[[package]] +name = "certifi" +version = "2026.7.22" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/a3/c2/24167ea9858356b47a87a50d39908bfdb72ceeefe0041586e704e5376b3a/certifi-2026.7.22.tar.gz", hash = "sha256:741e2c3b351ddf169a738da9f2c048608ff7f2c5cc02f1ebc6b118bb090d5d55", size = 138112 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/0b/a7/71ac2cff56fec219ed242bb11b8efb69fcc4bec75db06fb7bfe35de520e6/certifi-2026.7.22-py3-none-any.whl", hash = "sha256:62f22742b58a1a33014a2b6b706588a8d7e2a88ae7bd1a6ebe8c992928483775", size = 136983 }, +] + +[[package]] +name = "cffi" +version = "2.1.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pycparser", marker = "implementation_name != 'PyPy'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/9e/ef/008a1939e372c06329a3fce4279c02f328488f3526744906eeec3da7ad5f/cffi-2.1.1.tar.gz", hash = "sha256:dd31f52ea1086513bb9df30f8fcee9b8918323ae067a3d5b78bc826a000712be", size = 530807 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/70/d2/16d99a0c4948febc0ebd133a13b2f688ff7f8cb04da971e1128872ce0c03/cffi-2.1.1-cp311-cp311-macosx_10_15_x86_64.whl", hash = "sha256:c8d2c9fd1f2d16f780d15127abb050d13d1a76c03a4bd87d7e4980e45e511e12", size = 183838 }, + { url = "https://files.pythonhosted.org/packages/cd/95/31b535a9f0220ae9f357de4a08d57ce89cb417653c2fd9f075f50822a388/cffi-2.1.1-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:398aff33cee2767e3e781d2554c54bd0dff386bb437581e0d8011fde1a942ec1", size = 184168 }, + { url = "https://files.pythonhosted.org/packages/ad/5a/4707a0dc1f203f5dde5a907b0d4e3c25d71120241048bd5bc6f1bb9d4e71/cffi-2.1.1-cp311-cp311-manylinux1_i686.manylinux2014_i686.manylinux_2_17_i686.manylinux_2_5_i686.whl", hash = "sha256:154852545011f779917b11c78db2358d095da62a9a172b78ad0a583ee5adc0d0", size = 211805 }, + { url = "https://files.pythonhosted.org/packages/ad/66/c19feabb28485b6e0bbaaafa90837a1ef5d302e90f2178bd33f17a49879b/cffi-2.1.1-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:3311ed60d36f83378794e1009ac6258bafbf81f7888b4caa7b35a521e3f95813", size = 218716 }, + { url = "https://files.pythonhosted.org/packages/a7/92/500760486c8baab49a7a8a58ba7fc3355ec3974b454b8a09e528efde9e1d/cffi-2.1.1-cp311-cp311-manylinux2014_ppc64le.manylinux_2_17_ppc64le.whl", hash = "sha256:6e192623c49c94421616a5778fba35cf0d5a8d000650c1967ef4448ee5cdd990", size = 205569 }, + { url = "https://files.pythonhosted.org/packages/a5/a7/a67c733254d6e7373f7822f8082d8d6beade791e0cf12a7611f376fa61c7/cffi-2.1.1-cp311-cp311-manylinux2014_s390x.manylinux_2_17_s390x.whl", hash = "sha256:a6e721d4b0e45d5b65e87534470e67b18dcd092c83f68fba09f152b9cbc061af", size = 204907 }, + { url = "https://files.pythonhosted.org/packages/f7/a4/4399daaf8f7dfee9d7c3327fdb0426ee041cc63edc358b93911ceb2bfc7a/cffi-2.1.1-cp311-cp311-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:34e261f78cb6ceaaa36f42f2613f4380d94d9c759a9c73c769ee6e0247364632", size = 217807 }, + { url = "https://files.pythonhosted.org/packages/28/f7/dabe6da2466ecbd82dc62e7342dc6b1065dad990c06f00f0ede9ebf2a0ed/cffi-2.1.1-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:7225e4514edb64eb6740324353e0da0711954fd8d7da4576755b1c6e09b697cd", size = 221252 }, + { url = "https://files.pythonhosted.org/packages/ce/87/616202d8e51342c07d2534c510111c4cc37201775ce8f60802c9335d1edd/cffi-2.1.1-cp311-cp311-musllinux_1_2_i686.whl", hash = "sha256:df913725b79db7bcf03448f36b7bf8815363417d5b58deecf9305e3e30f0f21a", size = 214214 }, + { url = "https://files.pythonhosted.org/packages/b4/c6/ab025d75d2c26c19b087c0124e75ee31cb65032f4fe345d356d8c507ab97/cffi-2.1.1-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:f5cfbc5fe74540d335175b656c725d74d90e3730c626d92575eea35029d9afaa", size = 219408 }, + { url = "https://files.pythonhosted.org/packages/db/e2/7e8109f65445bdc673a7b54f02c677de462db75674220fd1335efc8eb598/cffi-2.1.1-cp311-cp311-win32.whl", hash = "sha256:f8ec5e643a9a937f64e1999eb9f75d072263751912dc5cd06d3c85f8f44be7c3", size = 174470 }, + { url = "https://files.pythonhosted.org/packages/73/c0/77ba02423c2f7d7091143c45cd49e0e6575c4c1967394bb542bd923a9b74/cffi-2.1.1-cp311-cp311-win_amd64.whl", hash = "sha256:42f6930c31dc7f50732c9ae793c2786c7b6b044195967bbdde40bb9be81c4cc0", size = 185096 }, + { url = "https://files.pythonhosted.org/packages/7c/47/9f1f85f9672ceda4984dc6c4f8824e8558992a2972c3d3c81fb8eb28d4ba/cffi-2.1.1-cp311-cp311-win_arm64.whl", hash = "sha256:c7659f22557c5a0bc4855cd635f55edec690cc008a40768527762cb9fb263455", size = 179941 }, + { url = "https://files.pythonhosted.org/packages/10/69/43965eccfdead3b9220015fd1320e117be8c6ed01a62ffab76eeb752f5d5/cffi-2.1.1-cp312-cp312-macosx_10_15_x86_64.whl", hash = "sha256:c8c69575568085ba0b1b10c0249d779a214aea6f6522e949a0fc9fb0fcb449d0", size = 184821 }, + { url = "https://files.pythonhosted.org/packages/54/7d/16e5a096677b5e313ca80cd5e5170efa3ea44624a82bb111925522da64b1/cffi-2.1.1-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:f81b3b8f3d4e343550fa4baa0e479bba9f2d29ce9c2e9b51d1ce1718d7442fcf", size = 184719 }, + { url = "https://files.pythonhosted.org/packages/56/e6/8941622732edec876dd17d0453dce07317ae96db34f2ec1436c9d3785986/cffi-2.1.1-cp312-cp312-manylinux1_i686.manylinux2014_i686.manylinux_2_17_i686.manylinux_2_5_i686.whl", hash = "sha256:811bd1e21d32de12efca32393a0ab3f5133b54fce9bd44b8bd77ab07da14bf6a", size = 214799 }, + { url = "https://files.pythonhosted.org/packages/44/de/f98430906df1545ffde0d543dd124a7a439bc2cd32b36b9c53f805df7333/cffi-2.1.1-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:68e62fe11f30d5ca8289242866f0a5291402d8529ca2178ab8afc5c9694ae890", size = 222389 }, + { url = "https://files.pythonhosted.org/packages/6a/5b/717f1526b9957b34456313c31645c5b82b8fb5c3fe9e4752999be7128bfc/cffi-2.1.1-cp312-cp312-manylinux2014_ppc64le.manylinux_2_17_ppc64le.whl", hash = "sha256:4a7c934f7360e8cd64fe9efadcbd10c7c6364f531e432b9a4bf5ccbc9e0e8b50", size = 210249 }, + { url = "https://files.pythonhosted.org/packages/64/b3/f8aa4f3e34986c7e4ec45072d1b1b9dd295b6b18007b45518d79726dd725/cffi-2.1.1-cp312-cp312-manylinux2014_s390x.manylinux_2_17_s390x.whl", hash = "sha256:3143d81e29e1e20a9ce10901ec369012947876596f75a222235965f2b7ae832e", size = 208775 }, + { url = "https://files.pythonhosted.org/packages/b1/db/dceb9dd5b231e1da801793f8acc9f3c52a7e1afe40bb1aae37e02b0faad5/cffi-2.1.1-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:c1453022f490d2459a11819d83ad1d586e9ff65a12ac3e705ffebd46d3685dcf", size = 221822 }, + { url = "https://files.pythonhosted.org/packages/a0/d2/6cd24ae3be000a634109c247d1475d62e5616d0dc78c82770942ec384248/cffi-2.1.1-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:208f941bb9d18e768138677f0a6d2ce01f590df56043dda1df1535ac57c88517", size = 225232 }, + { url = "https://files.pythonhosted.org/packages/cb/52/3fa190537004dd7f0ab860a6dc7c0175b8667f68d1e618a46f5498d30250/cffi-2.1.1-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:210019b6c7cf07f081b4c54635c8cf744377001350e29cc0f81c4377b4797735", size = 223597 }, + { url = "https://files.pythonhosted.org/packages/80/fb/0bb75b7039588c074b37ae99f40d9bfddf990ecb2fbc346ebccd2e56b9be/cffi-2.1.1-cp312-cp312-win32.whl", hash = "sha256:046bfc24911b37851ee1b51aab8bffe713d89c68c6a057b09484ce9fd5f69b4e", size = 175292 }, + { url = "https://files.pythonhosted.org/packages/d9/79/615cc094e2fb508cade7de88d3b4f6c4ec2bab695c97bce9153dc65aadf5/cffi-2.1.1-cp312-cp312-win_amd64.whl", hash = "sha256:f53e442b08449d42821fa4a4fba000095af9f62742a500f978a9f557ec44339a", size = 185919 }, + { url = "https://files.pythonhosted.org/packages/70/c6/d0ea84713fe46b243a436a18fcd47d639732747e21635c8a27191b06dc30/cffi-2.1.1-cp312-cp312-win_arm64.whl", hash = "sha256:7bde5e4cc5c10140859842b9d383af292b22639a4dffb725314baf45968cef80", size = 180093 }, +] + +[[package]] +name = "charset-normalizer" +version = "3.5.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/e5/3f/143b048436775b0f76ac3eec145c019e8173ccc2885c8f20319b996d5e83/charset_normalizer-3.5.1.tar.gz", hash = "sha256:6117b84ea48435e5356dc737f5121485c30920ba43375fa7b434fd753df0eac3", size = 171764 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/6a/b6/034f6802e9c3f6418966cfabb7db8c9252cc2429c5098f41cc43af804149/charset_normalizer-3.5.1-cp311-cp311-macosx_10_9_universal2.whl", hash = "sha256:eda059b6bc8bc0812d626fd91a7ce01bf583df0a61296eff390fd94141a34e30", size = 363585 }, + { url = "https://files.pythonhosted.org/packages/d5/fa/6a7e2a7c4b5451912b8c417732df79574354443592a88d616de03da66ae5/charset_normalizer-3.5.1-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:aa2bb0b37202dca27175591f761108b5d34096ade1191ffe4808bdf6b1571488", size = 251189 }, + { url = "https://files.pythonhosted.org/packages/a4/c8/ab42b07cfd82e919f427fcfaa7c41abae8242833ad1aad66d42bae40b669/charset_normalizer-3.5.1-cp311-cp311-manylinux2014_armv7l.manylinux_2_17_armv7l.manylinux_2_31_armv7l.whl", hash = "sha256:0b2b1b3fa5670c127b246df1d0c059defd41f689a868a3b9d79df9b1cac42d22", size = 239724 }, + { url = "https://files.pythonhosted.org/packages/e7/80/b9348b5d3041209f98b4cdad7655766369233f1d533f4f4f7558e9717bec/charset_normalizer-3.5.1-cp311-cp311-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:6e5e4d73d588ca5ed09df1b7dcd1b203d1df3c542e3f50d126c947d432b10731", size = 280078 }, + { url = "https://files.pythonhosted.org/packages/82/38/083a24028304bc85bb9e376fed801178423dcbb67495f73b6ea0624e1894/charset_normalizer-3.5.1-cp311-cp311-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:b54e7e13267d49ffbfe68e25b3cbd774dab38fa37238f71265e91b36146eb21c", size = 276650 }, + { url = "https://files.pythonhosted.org/packages/0d/35/731ac04aa0a097fc1c97f0994c375bdb230c6c96619db794208fe664e9ce/charset_normalizer-3.5.1-cp311-cp311-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:c7b742bf31c88566b4bb6335a7f393bb322e580b6bb98df7bd0c25e6e3519ce8", size = 262325 }, + { url = "https://files.pythonhosted.org/packages/f5/28/c2028e7021fb89c6e56868ed0e387b8e9aa811abdd2ab3208d6578d2c930/charset_normalizer-3.5.1-cp311-cp311-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:6ba32c4d2abf1d2fe7cf27d280f4cca5664233b0f885549c7761719eb977f486", size = 261140 }, + { url = "https://files.pythonhosted.org/packages/28/f0/0c0ceec6d98b7daa62e361e418135d59685811d79ba11529aad5cdf15e84/charset_normalizer-3.5.1-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:0722590aabf9dc6a6c0343d523c05458fa2b5047dbe6302fd526bb570600753f", size = 252791 }, + { url = "https://files.pythonhosted.org/packages/f0/3e/48f4cd187b1c33189d86039e9cbe4f92c05454175504b44ff81806d4d1bf/charset_normalizer-3.5.1-cp311-cp311-musllinux_1_2_armv7l.whl", hash = "sha256:aa1099b956fb795e686d073568f6dc002a0bb89765ea6d5b055dd7d9bf1b116c", size = 240730 }, + { url = "https://files.pythonhosted.org/packages/42/85/f9e22af69af67c54cce42be9455d9c81294f918b4ccc454db01f66efcac2/charset_normalizer-3.5.1-cp311-cp311-musllinux_1_2_ppc64le.whl", hash = "sha256:bd6c173f04743d483881bffa1478d5a4624475b8cd1d2194956a75548e191c18", size = 280791 }, + { url = "https://files.pythonhosted.org/packages/fd/4c/9044135f42127630b6fa742feb51256353f6ab87a78f2fdd1de3de955a7f/charset_normalizer-3.5.1-cp311-cp311-musllinux_1_2_riscv64.whl", hash = "sha256:f298e218441525d3794428b4c8b8fb8662c6d3ea79925d4807ee6b9a96a3bca5", size = 259598 }, + { url = "https://files.pythonhosted.org/packages/ba/ed/1dd7cfebb4e75812934c49ca3b79757d11948053f7937ab7070c151f3c55/charset_normalizer-3.5.1-cp311-cp311-musllinux_1_2_s390x.whl", hash = "sha256:6e2912d4babbc65196ac13c2f53468dc57fb8b9c25ef913e8c59ddf7c6dc0e1b", size = 278217 }, + { url = "https://files.pythonhosted.org/packages/bf/eb/239c84503cc9e3ba6eb34686a24bc66e84f3924efdd7e38e751a19f6bc10/charset_normalizer-3.5.1-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:3d27167433c0d5f18dc850f07d0b3816221984fecdc405d6c157a6f0b8f8e9e6", size = 263417 }, + { url = "https://files.pythonhosted.org/packages/37/ab/4e4510e1e288478e2c8333131d1c1382382ba8cd2165053c79e39d1da961/charset_normalizer-3.5.1-cp311-cp311-win32.whl", hash = "sha256:ac00177c4831ffa650f8609e4bdddd5fe09c03b1c0c47acece7e6ea20421598b", size = 181774 }, + { url = "https://files.pythonhosted.org/packages/e3/57/32f0ccea59e8612057c61d6fd22ef2cb63cca93c9fe594094919696ac170/charset_normalizer-3.5.1-cp311-cp311-win_amd64.whl", hash = "sha256:f9b1e28d0e8dbfa858abdba91d6b547beaf2df1a59bec6da6faae7b96a4991a9", size = 206653 }, + { url = "https://files.pythonhosted.org/packages/17/d4/b65c433fc521e58b5f54293982a5e51c05cb5f2dd3f1c7a6acb65b75324e/charset_normalizer-3.5.1-cp311-cp311-win_arm64.whl", hash = "sha256:ae31a1a1db2ee6cc2942fccaf695c934bc7f3db9f2133a3fef1f367cf1a4ab10", size = 185630 }, + { url = "https://files.pythonhosted.org/packages/30/27/78873dc8b6a56357517b74b6bb9568b80450e7bb4f6ef7e3fa9d22aa0bd7/charset_normalizer-3.5.1-cp312-cp312-macosx_10_13_universal2.whl", hash = "sha256:5b6d1386bf0096d26d3a863dc0a487a5b4eb9aa93cf5ba69683d29dde6b9d60f", size = 344456 }, + { url = "https://files.pythonhosted.org/packages/9a/4c/be49ada26b1f0232d57aa89bbebf997a5cc2332a5616b6eca26ff680044d/charset_normalizer-3.5.1-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:4582c27e8c889d64811987b5967fbd3ae0c823fe1fd933b543d55ac20bb475fa", size = 238530 }, + { url = "https://files.pythonhosted.org/packages/76/84/6f1290fa07ae6978d3960caa3eb1b8019bf9284ab7c2297b00c099ef4250/charset_normalizer-3.5.1-cp312-cp312-manylinux2014_armv7l.manylinux_2_17_armv7l.manylinux_2_31_armv7l.whl", hash = "sha256:1d1c7a53a6c2103925cdd6d7229f8c567379f211c869793df679f2e9f738c369", size = 230200 }, + { url = "https://files.pythonhosted.org/packages/e7/a0/47b18adeed31c8f16ba9700f32c1b18594cfa09f47eb672a488c273c22bf/charset_normalizer-3.5.1-cp312-cp312-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:e6621fb2a4988d6e53eedc455e5903e2679f3967b8acb3d639f1b63c14a2e893", size = 262222 }, + { url = "https://files.pythonhosted.org/packages/38/fe/341861ac118dae06f3ec0eb487488af52128f2ef2faf0b11003944d22259/charset_normalizer-3.5.1-cp312-cp312-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:7c0c10730342b0c9b35dd1d619beb8214e520bd96a1f870f452680b238aab3e0", size = 258951 }, + { url = "https://files.pythonhosted.org/packages/6f/89/bb5108dc6c3651dca963f2b0a3ba19bbcb370c94e1b6d3e0e844a58e6dca/charset_normalizer-3.5.1-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:b9af956078716df40d985fb0dfeb2c2120c5ca92ba4ff4b388acfd01cdc14d08", size = 248801 }, + { url = "https://files.pythonhosted.org/packages/b1/ba/ef83ae3aca816393decfa3530976f38a79812d707b80b580ac33b83f9877/charset_normalizer-3.5.1-cp312-cp312-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:f9f8405c2c758532c74fed975dbee57be1f31a6e865c031870c79a6ed3212ada", size = 244070 }, + { url = "https://files.pythonhosted.org/packages/f6/0b/c5292a2462d69b7378ea89793bbb5b2b6fcf6f7dd6d1667f9619094ad553/charset_normalizer-3.5.1-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:96fef3e886d6a9874b14f27fc193fbdc69d5d8035783d86aa4e1cea594e695f9", size = 240110 }, + { url = "https://files.pythonhosted.org/packages/46/22/111e5be3b740d5c2a5bfcedb3d237b6591e5c2e82ae9d6ffcb121fe0909c/charset_normalizer-3.5.1-cp312-cp312-musllinux_1_2_armv7l.whl", hash = "sha256:5d8531a6569d025f68e2321e7638fb7978f23db58e5f69f56913837aae03816e", size = 232836 }, + { url = "https://files.pythonhosted.org/packages/f9/d2/d2aad6fe0dbb44b194bf3becb60f5a0ac48446ade999a47fe7bb41eb09a7/charset_normalizer-3.5.1-cp312-cp312-musllinux_1_2_ppc64le.whl", hash = "sha256:aae2ee51122d3ae968a3837d97dc24a0aeebb0dea23694422cd172bd30017cd6", size = 262712 }, + { url = "https://files.pythonhosted.org/packages/35/5a/337e4663a5eae6de99db940ee8066d4145caafb61327db62deda15313cce/charset_normalizer-3.5.1-cp312-cp312-musllinux_1_2_riscv64.whl", hash = "sha256:7235dc28fc6dd9d832ac7c7bce95367dedb85929f17368a0c2bee1e080b9acbf", size = 242977 }, + { url = "https://files.pythonhosted.org/packages/ca/85/f82f8a92e31c7519410e2e1afdc630f28ec47490ce2c09a11c1a43cbb459/charset_normalizer-3.5.1-cp312-cp312-musllinux_1_2_s390x.whl", hash = "sha256:4abdc5f9ad448c1ecbfae2974b820535d6bc6e7eef63babbab3d81cf46968c71", size = 260207 }, + { url = "https://files.pythonhosted.org/packages/b7/52/643d11ffd60e9ac2fd1fb87e167a19285b9eefeff4a40e63c87cbfbeab36/charset_normalizer-3.5.1-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:ba501e667c17d8411f98e67a022d9604ef179aff0e459b7e292c796837c13573", size = 250562 }, + { url = "https://files.pythonhosted.org/packages/62/16/46556278c2168d12df9da7fede5dc6fc70e60301b26a82bbeec238c9cfe3/charset_normalizer-3.5.1-cp312-cp312-win32.whl", hash = "sha256:cfa1c0cc3a8f9f53f1243a5a99ac36fd003880199383b37672e86ddda9cb07e2", size = 178507 }, + { url = "https://files.pythonhosted.org/packages/9d/7a/4c6c298171e6b3e745633180ff59350fc0ca0db1ffd28df1e369e0579f71/charset_normalizer-3.5.1-cp312-cp312-win_amd64.whl", hash = "sha256:3617ac3cfd8b9888f145ad89dd6e692285834b0201c6074a5eeaad3fd4d668c2", size = 200551 }, + { url = "https://files.pythonhosted.org/packages/cd/d7/eb95a042f0dd22e304b0b6472b154f3546a1a039a9ee89ccb2a7f61591fc/charset_normalizer-3.5.1-cp312-cp312-win_arm64.whl", hash = "sha256:88e85ab89cb822c1e635f51d6d32e488f94e002e70e2f492bdb8b945543f345a", size = 180700 }, + { url = "https://files.pythonhosted.org/packages/5b/97/fb4e82231aba271ffd775a1b4993b0defc4e3059f286ae41d9433409fe85/charset_normalizer-3.5.1-cp37-abi3-macosx_10_9_universal2.whl", hash = "sha256:41876ee62a3dddf48ff1121ad8f0798032aa03f2fd35f21f34a4cab14f18d8d2", size = 331467 }, + { url = "https://files.pythonhosted.org/packages/9f/2f/fe3f187327aac18e2d54e9d2b08e15d27bf9b642d9e51c219f130fc34d1a/charset_normalizer-3.5.1-cp37-abi3-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:a6dac12ff6b846103483683f60c5f8fee205121adc58ffd87e90a90a3af69e99", size = 253057 }, + { url = "https://files.pythonhosted.org/packages/d7/c7/9e48cee5c161fe24da823b61bf381921d77cb994a0a4de148e95018c1984/charset_normalizer-3.5.1-cp37-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:cee5dd7c6fb5dd52a0fe2a740f9bc6e3593f5f8b1788bde49de02086f30182b2", size = 240930 }, + { url = "https://files.pythonhosted.org/packages/49/e0/716601f3cc69be7b198951150c75ead1ece33c3c8036ff6ffa46029659a0/charset_normalizer-3.5.1-cp37-abi3-manylinux2014_armv7l.manylinux_2_17_armv7l.manylinux_2_31_armv7l.whl", hash = "sha256:343fb4f2821043bd87095f7b08a1a181febc8e36ac64212143bbfd0a0e1bc235", size = 230822 }, + { url = "https://files.pythonhosted.org/packages/d3/05/71bfc5caa0abcc45aea1f6a4d50ac68e59605ddc7666fe8494f4cd229665/charset_normalizer-3.5.1-cp37-abi3-manylinux2014_ppc64le.manylinux_2_17_ppc64le.manylinux_2_28_ppc64le.whl", hash = "sha256:ae4a097991662cd4fff0ddc74e0fe7874f82e00042fa0ea00855645ed0c79598", size = 260037 }, + { url = "https://files.pythonhosted.org/packages/c3/92/de7e32ed05341e7a9c4c877c318418197b7f2d66a3b68d561bf2ac57ca3e/charset_normalizer-3.5.1-cp37-abi3-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:4b599739b93b2cbeded49645ae3c8d1405c29ddfbceac1545c87a3f9580a9e96", size = 255097 }, + { url = "https://files.pythonhosted.org/packages/f5/7b/ade0a122600319dfa0b1000ab0f9731c94a817904cf3c5de408c73a4ede7/charset_normalizer-3.5.1-cp37-abi3-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:b39b69b347e5e47a3b5b8cfc005c68c1ba347474e3960236c4944a8ecd174962", size = 250166 }, + { url = "https://files.pythonhosted.org/packages/75/9c/019fbb9f4834491a160951349b1a3714439376f66e5f7cf18b4f18f0c7aa/charset_normalizer-3.5.1-cp37-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:a2028475ba855475b8b4d3cfeb4994269c967aea8b9892dfba907f4263a863a3", size = 241821 }, + { url = "https://files.pythonhosted.org/packages/2b/b8/11d4840bfc99330cc7fbcc2681ee5a044553a6e77655508d8f9b2bff7b34/charset_normalizer-3.5.1-cp37-abi3-musllinux_1_2_armv7l.whl", hash = "sha256:36047af20e17097c3bb9476c2b7655f2f7aa51322c0ba58c07695bedf755a950", size = 232529 }, + { url = "https://files.pythonhosted.org/packages/18/96/2b3a21492d9f65171ac75d872f5018260013d00bfa0ff70ec9f179148cbd/charset_normalizer-3.5.1-cp37-abi3-musllinux_1_2_ppc64le.whl", hash = "sha256:4c4fb141a727957c93edfe5c32a26ceb6b5f6461d67146e2d39f51e16170bea8", size = 260348 }, + { url = "https://files.pythonhosted.org/packages/d6/aa/a69a2028e8bd052476c245460ab19d7de595de084dd968f2d75cd50c3e25/charset_normalizer-3.5.1-cp37-abi3-musllinux_1_2_riscv64.whl", hash = "sha256:2f293479cce755c75f1697e87c409b7ae4c555c7dfecb6e988ad13abba943031", size = 247234 }, + { url = "https://files.pythonhosted.org/packages/35/8a/3d130aeabcaf3d2466af76b7b141c08d9e89c9016ab4b7cdd0f7dc2d1c62/charset_normalizer-3.5.1-cp37-abi3-musllinux_1_2_s390x.whl", hash = "sha256:3588e376b3ea2eea84976f67273d679f229e24c66dce7b82ae45aef04ff6e072", size = 256917 }, + { url = "https://files.pythonhosted.org/packages/80/c2/a7379b840292d0c1ab9fbd17d1f3967aa81794dc95bc74be8999d7fedcf7/charset_normalizer-3.5.1-cp37-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:e199fb99720074809a7720f1c0b4d919eea8b87e88713e0f8f602f7bef543d9d", size = 254846 }, + { url = "https://files.pythonhosted.org/packages/01/65/d43b714731bb2f40d4053dfa00ecfc1c5a301f8e3316c5db3a09af59fe94/charset_normalizer-3.5.1-cp37-abi3-win32.whl", hash = "sha256:dd732602a7009217f658d5863d12d79d373a4de0eebc111094bcdd3bb8e0a6cc", size = 174216 }, + { url = "https://files.pythonhosted.org/packages/35/4f/b911ed898b26a09789eba9c9200c999aff6c61b4bafaf4838e56d1a1e1a3/charset_normalizer-3.5.1-cp37-abi3-win_amd64.whl", hash = "sha256:70055ff39b97c99e7ae40ea3e393fb62aa2e44dbd9b29f8d14f42fb0025c3959", size = 199764 }, + { url = "https://files.pythonhosted.org/packages/f0/a7/920baf467bfd9bf689f3b318340f37aee4572a71f162bd8db51da55ba4fa/charset_normalizer-3.5.1-cp37-abi3-win_arm64.whl", hash = "sha256:87e4f41d375c0b9be2fb5251aee4b8a689169e134535aed81bf085c3b647451e", size = 287318 }, + { url = "https://files.pythonhosted.org/packages/cc/61/d01fc49b8dea277640b55a9e15960dbca9fdc8c9fde18e572d39c59f4019/charset_normalizer-3.5.1-py3-none-any.whl", hash = "sha256:6df0ec430f9a831772c23ca5a224cba36517a58a84bb32c32bb59a9fa67c47f6", size = 68658 }, +] + +[[package]] +name = "click" +version = "8.5.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/c7/0e/7fa0ef50764b67090eca4114772a2abf8b6148198475e54c660b97caeee6/click-8.5.0.tar.gz", hash = "sha256:ba0d2089de75ea0310e2dde03160e6ca10009947fb95a182f9b54021bb272e34", size = 382235 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/58/50/6c0d534c5f134586a8e1ba4e330569e32f057e33372ae556463212fb4cd3/click-8.5.0-py3-none-any.whl", hash = "sha256:255bc9599cf7748b4b1a446ccc735421bd08a2ae529a8b88597d3de5664ee360", size = 125251 }, +] + +[[package]] +name = "cloudpickle" +version = "3.1.2" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/27/fb/576f067976d320f5f0114a8d9fa1215425441bb35627b1993e5afd8111e5/cloudpickle-3.1.2.tar.gz", hash = "sha256:7fda9eb655c9c230dab534f1983763de5835249750e85fbcef43aaa30a9a2414", size = 22330 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/88/39/799be3f2f0f38cc727ee3b4f1445fe6d5e4133064ec2e4115069418a5bb6/cloudpickle-3.1.2-py3-none-any.whl", hash = "sha256:9acb47f6afd73f60dc1df93bb801b472f05ff42fa6c84167d25cb206be1fbf4a", size = 22228 }, +] + +[[package]] +name = "colorama" +version = "0.4.6" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/d8/53/6f443c9a4a8358a93a6792e2acffb9d9d5cb0a5cfd8802644b7b1c9a02e4/colorama-0.4.6.tar.gz", hash = "sha256:08695f5cb7ed6e0531a20572697297273c47b8cae5a63ffc6d6ed5c201be6e44", size = 27697 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/d1/d6/3965ed04c63042e047cb6a3e6ed1a63a35087b6a609aa3a15ed8ac56c221/colorama-0.4.6-py2.py3-none-any.whl", hash = "sha256:4f1d9991f5acc0ca119f9d443620b77f9d6b33703e51011c16baf57afb285fc6", size = 25335 }, +] + +[[package]] +name = "cuda-bindings" +version = "13.3.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "cuda-pathfinder", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, +] +wheels = [ + { url = "https://files.pythonhosted.org/packages/51/6b/457ca12dad3ee9bfcc9a545cfd6b64b359ba49de40f776f6e028e678f262/cuda_bindings-13.3.1-cp311-cp311-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:c5879712accf6e14bb01aa5e67440eb84998b8d104b509cc7a6dc0b8f656a474", size = 6053539 }, + { url = "https://files.pythonhosted.org/packages/95/7a/c5e3c34a409b148f5c0f5a4ea374158f95d488862c1dffedf9aa5c639df9/cuda_bindings-13.3.1-cp311-cp311-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:04436a9364059c84b8f9636f359eccda1cf814341f5b670c71d80d2f79dbc708", size = 6674166 }, + { url = "https://files.pythonhosted.org/packages/ce/67/5e7dba1ba576dd73da5dee894ca076ca5e959450dfff66d6d510a255d1f7/cuda_bindings-13.3.1-cp312-cp312-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:c7855c4868aabc0cfae28abbe83d56734bdfbd08f08fc234ac1912a12858bf49", size = 6025351 }, + { url = "https://files.pythonhosted.org/packages/39/2a/6d2e9047d1fb243dbaa364b01e0297534b9ed7fd27dba1c9f361519cf69b/cuda_bindings-13.3.1-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:e32d08f71ebcdf00f0f41eab2eb37e8da94c8ed411cc9f7f7a019ce6b34abe3a", size = 6657965 }, +] + +[[package]] +name = "cuda-pathfinder" +version = "1.8.0" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/a1/b1/ef21259ec74fe0b265ed201379de1d0ef7c14178313ee03705952f1b7093/cuda_pathfinder-1.8.0-py3-none-any.whl", hash = "sha256:c44e574dc997fae2814721d1ae97d0fd6db76db82decbe9b753bf75de53f515e", size = 62539 }, +] + +[[package]] +name = "cuda-toolkit" +version = "13.0.3.0" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/d1/c7/a79086a62c98befcdb8349656c6f114e2db3b8b2422f6e25c97a7f2a9a3c/cuda_toolkit-13.0.3.0-py2.py3-none-any.whl", hash = "sha256:d693caaa261214ddd7dbb60d68e71cbed884e68c2be7509778f3051da0b91c3f", size = 2512 }, +] + +[package.optional-dependencies] +cublas = [ + { name = "nvidia-cublas", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, + { name = "nvidia-cuda-nvrtc", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, +] +cudart = [ + { name = "nvidia-cuda-runtime", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, +] +cufft = [ + { name = "nvidia-cufft", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, + { name = "nvidia-nvjitlink", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, +] +cufile = [ + { name = "nvidia-cufile", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, +] +cupti = [ + { name = "nvidia-cuda-cupti", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, +] +curand = [ + { name = "nvidia-curand", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, +] +cusolver = [ + { name = "nvidia-cublas", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, + { name = "nvidia-cusolver", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, + { name = "nvidia-cusparse", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, + { name = "nvidia-nvjitlink", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, +] +cusparse = [ + { name = "nvidia-cusparse", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, + { name = "nvidia-nvjitlink", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, +] +nvjitlink = [ + { name = "nvidia-nvjitlink", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, +] +nvrtc = [ + { name = "nvidia-cuda-nvrtc", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, +] +nvtx = [ + { name = "nvidia-nvtx", marker = "(platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, +] + +[[package]] +name = "filelock" +version = "3.32.5" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/0a/a0/50c2c0ce5e74d7721bbb1b19a26ebd339aac5878553a6e35308c2f31f935/filelock-3.32.5.tar.gz", hash = "sha256:f6a6a28f743f9b95ce19db5abe0f376f75eb56517dff21e1a4751e2657d3e83d", size = 222838 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/36/d2/b70a31e13d04456d28493f31d2aa087e99eeb2767ef0293b2625727ccb8c/filelock-3.32.5-py3-none-any.whl", hash = "sha256:142cd9fa77a872c5e78c62329a0d15278fadc686eb89e760017968961a4fd6b2", size = 100003 }, +] + +[[package]] +name = "fsspec" +version = "2026.7.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/00/78/f34251dadb8f3921264a1d9b8946f5e542014ee2614b285261b4e40e6775/fsspec-2026.7.0.tar.gz", hash = "sha256:c803c40f4cf860b49dea58ee3e1c33cb9c790520e233537e1340049f89b82a88", size = 317040 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/fd/3c/6a2bf344106328fd04963664a60b9bb6496fc25df8e962fcdc1367285fb9/fsspec-2026.7.0-py3-none-any.whl", hash = "sha256:b57ddbafedfaef7018c1ecab32aa200a9d7ca26b77965f64e48b70061249d279", size = 206583 }, +] + +[[package]] +name = "gguf" +version = "0.19.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, + { name = "numpy", version = "2.5.2", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "pyyaml" }, + { name = "requests" }, + { name = "tqdm" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/48/ae/17f1308ae45cd7b08ebb521747d5b23f4efc4d172038a4e228dd5106c3ff/gguf-0.19.0.tar.gz", hash = "sha256:dbadcd6cc7ccd44256f2229fe7c2dff5e8aa5cf0612ab987fd2b1a57e428923f", size = 111220 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/b3/bb/d71d6da82763528c2c2ed6b59a9d6142c6595545a4c448e2085d155e88c2/gguf-0.19.0-py3-none-any.whl", hash = "sha256:70bcd10edfe697fb2dad6e40af2234b9d8ece9a41a99761405121ebda1c3c1cd", size = 118475 }, +] + +[[package]] +name = "h11" +version = "0.16.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/01/ee/02a2c011bdab74c6fb3c75474d40b3052059d95df7e73351460c8588d963/h11-0.16.0.tar.gz", hash = "sha256:4e35b956cf45792e4caa5885e69fba00bdbc6ffafbfa020300e549b208ee5ff1", size = 101250 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/04/4b/29cac41a4d98d144bf5f6d33995617b185d14b22401f75ca86f384e87ff1/h11-0.16.0-py3-none-any.whl", hash = "sha256:63cf8bbe7522de3bf65932fda1d9c2772064ffb3dae62d55932da54b31cb6c86", size = 37515 }, +] + +[[package]] +name = "hf-xet" +version = "1.6.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/1b/ab/522a2ab67f27971a9d48ca666d4fca85ef7d5282d142e31fd087e27b1bbe/hf_xet-1.6.0.tar.gz", hash = "sha256:2e58454a340b3556dfa4972d5451aff4fba8dd42a236600ba1a1d2b1514f0fef", size = 920527 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/a2/50/7afa2c9c787405864fc47a0d1bbc02c62e9101947ed43c1f43899fc7d91d/hf_xet-1.6.0-cp38-abi3-macosx_10_12_x86_64.whl", hash = "sha256:633dc0cd71d32da58ab8c03ad38e2fac452c15c2b0a2866ebf6ededfe0a5061d", size = 4071729 }, + { url = "https://files.pythonhosted.org/packages/4b/69/55b8dcf636142ae660fec1869fcac14c4da2e8412e14d6eee1523be77e9f/hf_xet-1.6.0-cp38-abi3-macosx_11_0_arm64.whl", hash = "sha256:f0906082d9932ae0c0057fa194041c22b4e2cdb46b2592ef3b91f020d62a081a", size = 3876287 }, + { url = "https://files.pythonhosted.org/packages/67/4e/a28359bf1c1ecf11eba22123168c138698f7cb576ac678f5a2e16cd5da08/hf_xet-1.6.0-cp38-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:d62671bb130879cef0ee4c9ebe47a14af6c66ec53e6d84dc15936e5ffdfac82f", size = 4464663 }, + { url = "https://files.pythonhosted.org/packages/9a/69/1f0cbc2fb22ae6082d094f743d1b8945a3f36f6089cb95f42b7ee348cda7/hf_xet-1.6.0-cp38-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:0e6e21fa3cdfcdcd76748564bf593870a5e013f47d97cf10aed63aa222cff5b7", size = 4262538 }, + { url = "https://files.pythonhosted.org/packages/d1/3a/4f4f2301ade26e404462d3336fa11f7958d914cabbabdd6e03c3c5d5658c/hf_xet-1.6.0-cp38-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:4fc74352a17015bd0ee90038bc9efe38db894cde45f268b6712b04fce8cd0acb", size = 4460520 }, + { url = "https://files.pythonhosted.org/packages/ab/5f/311725e2a905534dfee2dcb5b08414f249147f1f12252bfc2bd24caa075c/hf_xet-1.6.0-cp38-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:8fb4f71cba6129110c3374a33f919001ff130488fc23553698e34cc1c2a1198c", size = 4675937 }, + { url = "https://files.pythonhosted.org/packages/98/b7/8c59a66d15205024662f1d66968136f13893f96df1ddc5087e2e281fc95f/hf_xet-1.6.0-cp38-abi3-win_amd64.whl", hash = "sha256:fb4fadde1b2b70bf4c0c14a6dccbe7194b1c28947fefd5bbe3fed9d940676c3b", size = 4033128 }, + { url = "https://files.pythonhosted.org/packages/73/63/ca511b6f802f28cf3489b280fe77475bcca8de85e81a6299d7916b5b5555/hf_xet-1.6.0-cp38-abi3-win_arm64.whl", hash = "sha256:3dc3e35441ba395006af5aaacc40ef2e603c51ef46c3530b9156185f00935ea3", size = 3859359 }, +] + +[[package]] +name = "httpcore" +version = "1.0.9" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "certifi" }, + { name = "h11" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/06/94/82699a10bca87a5556c9c59b5963f2d039dbd239f25bc2a63907a05a14cb/httpcore-1.0.9.tar.gz", hash = "sha256:6e34463af53fd2ab5d807f399a9b45ea31c3dfa2276f15a2c3f00afff6e176e8", size = 85484 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/7e/f5/f66802a942d491edb555dd61e3a9961140fd64c90bce1eafd741609d334d/httpcore-1.0.9-py3-none-any.whl", hash = "sha256:2d400746a40668fc9dec9810239072b40b4484b640a8c38fd654a024c7a1bf55", size = 78784 }, +] + +[[package]] +name = "httpx" +version = "0.28.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "anyio" }, + { name = "certifi" }, + { name = "httpcore" }, + { name = "idna" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/b1/df/48c586a5fe32a0f01324ee087459e112ebb7224f646c0b5023f5e79e9956/httpx-0.28.1.tar.gz", hash = "sha256:75e98c5f16b0f35b567856f597f06ff2270a374470a5c2392242528e3e3e42fc", size = 141406 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/2a/39/e50c7c3a983047577ee07d2a9e53faf5a69493943ec3f6a384bdc792deb2/httpx-0.28.1-py3-none-any.whl", hash = "sha256:d909fcccc110f8c7faf814ca82a9a4d816bc5a6dbfea25d6591d6985b8ba59ad", size = 73517 }, +] + +[[package]] +name = "huggingface-hub" +version = "1.29.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "click" }, + { name = "filelock" }, + { name = "fsspec" }, + { name = "hf-xet", marker = "platform_machine == 'AMD64' or platform_machine == 'aarch64' or platform_machine == 'amd64' or platform_machine == 'arm64' or platform_machine == 'x86_64'" }, + { name = "httpx" }, + { name = "packaging" }, + { name = "pyyaml" }, + { name = "tqdm" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/64/35/42316e8f6908b6d21bc8df017cc6efba94fb5edbf99b64e28dd142325e20/huggingface_hub-1.29.0.tar.gz", hash = "sha256:6ebb385a581435325cf6d5c5b233d5d4bc91175834d99fd65dae14379b36e9ad", size = 963121 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/4e/a5/47c2ea9b228ccbcba8467e9a64823146e8ebbad29855e591d8f5eedcc9c7/huggingface_hub-1.29.0-py3-none-any.whl", hash = "sha256:b00f7782afc14db4bc6572763810a635bdfbab8623d957bfb553bd18e03852cd", size = 795768 }, +] + +[[package]] +name = "hyperpyyaml" +version = "1.2.3" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pyyaml" }, + { name = "ruamel-yaml" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/ce/cf/000268ec653f354d464e6101505672365e8cecdcb216034a4bb0dbf9dca0/hyperpyyaml-1.2.3.tar.gz", hash = "sha256:4b135800ae13b846ae90aeb1c25e65e4d7a0138bd4935d3222734828d9f218f6", size = 17355 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/7d/a5/98e20cc4365e293dfc1382e2f8dbf44b1cb14d75658ab5147cb0d6d8718e/hyperpyyaml-1.2.3-py3-none-any.whl", hash = "sha256:0088c8ce97dc7c7d3460112b6ccc92dca46fade4d0320bc39f58b1fdda60f7f1", size = 16456 }, +] + +[[package]] +name = "idna" +version = "3.19" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/5f/f7/abb373e5757eaec4b922b92f97ec8d6d7e057cf06778247604fbc4e7c3f3/idna-3.19.tar.gz", hash = "sha256:5e0811a4383b21dc5838069f801c4fb62113b7447663d2530d2bd6e77b49bf15", size = 215237 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/57/b0/0e52c878c53f245edd3a11020f20979b3f490f245af532c7cae3027754b5/idna-3.19-py3-none-any.whl", hash = "sha256:815e7be7a7806d54abb586dc943addc79e8b2ee16915059658cbeff4b1b43bf4", size = 68550 }, +] + +[[package]] +name = "jinja2" +version = "3.1.6" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "markupsafe" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/df/bf/f7da0350254c0ed7c72f3e33cef02e048281fec7ecec5f032d4aac52226b/jinja2-3.1.6.tar.gz", hash = "sha256:0137fb05990d35f1275a587e9aee6d56da821fc83491a0fb838183be43f66d6d", size = 245115 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/62/a1/3d680cbfd5f4b8f15abc1d571870c5fc3e594bb582bc3b64ea099db13e56/jinja2-3.1.6-py3-none-any.whl", hash = "sha256:85ece4451f492d0c13c5dd7c13a64681a86afae63a5f347908daf103ce6d2f67", size = 134899 }, +] + +[[package]] +name = "joblib" +version = "1.6.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "cloudpickle" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/d5/1d/537ab090f302b838943a1b56497dd53059b9a9b46a074936470173a2e207/joblib-1.6.0.tar.gz", hash = "sha256:2ccc96785b12046c08fd6d55839c12857831b54a3c1673ffadd2f04bfc4eda03", size = 327903 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/18/53/84099323c2ec4be98d935f63c033ac4151ee83836ca1050ede3b3aadf155/joblib-1.6.0-py3-none-any.whl", hash = "sha256:3dbbf9f6e4b592a2357b854608e980fe6390d131d7a82f011a377ef2ebef7aba", size = 306115 }, +] + +[[package]] +name = "markupsafe" +version = "3.0.3" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/7e/99/7690b6d4034fffd95959cbe0c02de8deb3098cc577c67bb6a24fe5d7caa7/markupsafe-3.0.3.tar.gz", hash = "sha256:722695808f4b6457b320fdc131280796bdceb04ab50fe1795cd540799ebe1698", size = 80313 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/08/db/fefacb2136439fc8dd20e797950e749aa1f4997ed584c62cfb8ef7c2be0e/markupsafe-3.0.3-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:1cc7ea17a6824959616c525620e387f6dd30fec8cb44f649e31712db02123dad", size = 11631 }, + { url = "https://files.pythonhosted.org/packages/e1/2e/5898933336b61975ce9dc04decbc0a7f2fee78c30353c5efba7f2d6ff27a/markupsafe-3.0.3-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:4bd4cd07944443f5a265608cc6aab442e4f74dff8088b0dfc8238647b8f6ae9a", size = 12058 }, + { url = "https://files.pythonhosted.org/packages/1d/09/adf2df3699d87d1d8184038df46a9c80d78c0148492323f4693df54e17bb/markupsafe-3.0.3-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:6b5420a1d9450023228968e7e6a9ce57f65d148ab56d2313fcd589eee96a7a50", size = 24287 }, + { url = "https://files.pythonhosted.org/packages/30/ac/0273f6fcb5f42e314c6d8cd99effae6a5354604d461b8d392b5ec9530a54/markupsafe-3.0.3-cp311-cp311-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:0bf2a864d67e76e5c9a34dc26ec616a66b9888e25e7b9460e1c76d3293bd9dbf", size = 22940 }, + { url = "https://files.pythonhosted.org/packages/19/ae/31c1be199ef767124c042c6c3e904da327a2f7f0cd63a0337e1eca2967a8/markupsafe-3.0.3-cp311-cp311-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:bc51efed119bc9cfdf792cdeaa4d67e8f6fcccab66ed4bfdd6bde3e59bfcbb2f", size = 21887 }, + { url = "https://files.pythonhosted.org/packages/b2/76/7edcab99d5349a4532a459e1fe64f0b0467a3365056ae550d3bcf3f79e1e/markupsafe-3.0.3-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:068f375c472b3e7acbe2d5318dea141359e6900156b5b2ba06a30b169086b91a", size = 23692 }, + { url = "https://files.pythonhosted.org/packages/a4/28/6e74cdd26d7514849143d69f0bf2399f929c37dc2b31e6829fd2045b2765/markupsafe-3.0.3-cp311-cp311-musllinux_1_2_riscv64.whl", hash = "sha256:7be7b61bb172e1ed687f1754f8e7484f1c8019780f6f6b0786e76bb01c2ae115", size = 21471 }, + { url = "https://files.pythonhosted.org/packages/62/7e/a145f36a5c2945673e590850a6f8014318d5577ed7e5920a4b3448e0865d/markupsafe-3.0.3-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:f9e130248f4462aaa8e2552d547f36ddadbeaa573879158d721bbd33dfe4743a", size = 22923 }, + { url = "https://files.pythonhosted.org/packages/0f/62/d9c46a7f5c9adbeeeda52f5b8d802e1094e9717705a645efc71b0913a0a8/markupsafe-3.0.3-cp311-cp311-win32.whl", hash = "sha256:0db14f5dafddbb6d9208827849fad01f1a2609380add406671a26386cdf15a19", size = 14572 }, + { url = "https://files.pythonhosted.org/packages/83/8a/4414c03d3f891739326e1783338e48fb49781cc915b2e0ee052aa490d586/markupsafe-3.0.3-cp311-cp311-win_amd64.whl", hash = "sha256:de8a88e63464af587c950061a5e6a67d3632e36df62b986892331d4620a35c01", size = 15077 }, + { url = "https://files.pythonhosted.org/packages/35/73/893072b42e6862f319b5207adc9ae06070f095b358655f077f69a35601f0/markupsafe-3.0.3-cp311-cp311-win_arm64.whl", hash = "sha256:3b562dd9e9ea93f13d53989d23a7e775fdfd1066c33494ff43f5418bc8c58a5c", size = 13876 }, + { url = "https://files.pythonhosted.org/packages/5a/72/147da192e38635ada20e0a2e1a51cf8823d2119ce8883f7053879c2199b5/markupsafe-3.0.3-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:d53197da72cc091b024dd97249dfc7794d6a56530370992a5e1a08983ad9230e", size = 11615 }, + { url = "https://files.pythonhosted.org/packages/9a/81/7e4e08678a1f98521201c3079f77db69fb552acd56067661f8c2f534a718/markupsafe-3.0.3-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:1872df69a4de6aead3491198eaf13810b565bdbeec3ae2dc8780f14458ec73ce", size = 12020 }, + { url = "https://files.pythonhosted.org/packages/1e/2c/799f4742efc39633a1b54a92eec4082e4f815314869865d876824c257c1e/markupsafe-3.0.3-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:3a7e8ae81ae39e62a41ec302f972ba6ae23a5c5396c8e60113e9066ef893da0d", size = 24332 }, + { url = "https://files.pythonhosted.org/packages/3c/2e/8d0c2ab90a8c1d9a24f0399058ab8519a3279d1bd4289511d74e909f060e/markupsafe-3.0.3-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:d6dd0be5b5b189d31db7cda48b91d7e0a9795f31430b7f271219ab30f1d3ac9d", size = 22947 }, + { url = "https://files.pythonhosted.org/packages/2c/54/887f3092a85238093a0b2154bd629c89444f395618842e8b0c41783898ea/markupsafe-3.0.3-cp312-cp312-manylinux_2_31_riscv64.manylinux_2_39_riscv64.whl", hash = "sha256:94c6f0bb423f739146aec64595853541634bde58b2135f27f61c1ffd1cd4d16a", size = 21962 }, + { url = "https://files.pythonhosted.org/packages/c9/2f/336b8c7b6f4a4d95e91119dc8521402461b74a485558d8f238a68312f11c/markupsafe-3.0.3-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:be8813b57049a7dc738189df53d69395eba14fb99345e0a5994914a3864c8a4b", size = 23760 }, + { url = "https://files.pythonhosted.org/packages/32/43/67935f2b7e4982ffb50a4d169b724d74b62a3964bc1a9a527f5ac4f1ee2b/markupsafe-3.0.3-cp312-cp312-musllinux_1_2_riscv64.whl", hash = "sha256:83891d0e9fb81a825d9a6d61e3f07550ca70a076484292a70fde82c4b807286f", size = 21529 }, + { url = "https://files.pythonhosted.org/packages/89/e0/4486f11e51bbba8b0c041098859e869e304d1c261e59244baa3d295d47b7/markupsafe-3.0.3-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:77f0643abe7495da77fb436f50f8dab76dbc6e5fd25d39589a0f1fe6548bfa2b", size = 23015 }, + { url = "https://files.pythonhosted.org/packages/2f/e1/78ee7a023dac597a5825441ebd17170785a9dab23de95d2c7508ade94e0e/markupsafe-3.0.3-cp312-cp312-win32.whl", hash = "sha256:d88b440e37a16e651bda4c7c2b930eb586fd15ca7406cb39e211fcff3bf3017d", size = 14540 }, + { url = "https://files.pythonhosted.org/packages/aa/5b/bec5aa9bbbb2c946ca2733ef9c4ca91c91b6a24580193e891b5f7dbe8e1e/markupsafe-3.0.3-cp312-cp312-win_amd64.whl", hash = "sha256:26a5784ded40c9e318cfc2bdb30fe164bdb8665ded9cd64d500a34fb42067b1c", size = 15105 }, + { url = "https://files.pythonhosted.org/packages/e5/f1/216fc1bbfd74011693a4fd837e7026152e89c4bcf3e77b6692fba9923123/markupsafe-3.0.3-cp312-cp312-win_arm64.whl", hash = "sha256:35add3b638a5d900e807944a078b51922212fb3dedb01633a8defc4b01a3c85f", size = 13906 }, +] + +[[package]] +name = "mpmath" +version = "1.3.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/e0/47/dd32fa426cc72114383ac549964eecb20ecfd886d1e5ccf5340b55b02f57/mpmath-1.3.0.tar.gz", hash = "sha256:7a28eb2a9774d00c7bc92411c19a89209d5da7c4c9a9e227be8330a23a25b91f", size = 508106 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/43/e3/7d92a15f894aa0c9c4b49b8ee9ac9850d6e63b03c9c32c0367a13ae62209/mpmath-1.3.0-py3-none-any.whl", hash = "sha256:a0b2b9fe80bbcd81a6647ff13108738cfb482d481d826cc0e02f5b35e5c88d2c", size = 536198 }, +] + +[[package]] +name = "networkx" +version = "3.6.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/6a/51/63fe664f3908c97be9d2e4f1158eb633317598cfa6e1fc14af5383f17512/networkx-3.6.1.tar.gz", hash = "sha256:26b7c357accc0c8cde558ad486283728b65b6a95d85ee1cd66bafab4c8168509", size = 2517025 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/9e/c9/b2622292ea83fbb4ec318f5b9ab867d0a28ab43c5717bb85b0a5f6b3b0a4/networkx-3.6.1-py3-none-any.whl", hash = "sha256:d47fbf302e7d9cbbb9e2555a0d267983d2aa476bac30e90dfbe5669bd57f3762", size = 2068504 }, +] + +[[package]] +name = "numpy" +version = "2.4.6" +source = { registry = "https://pypi.org/simple" } +resolution-markers = [ + "python_full_version < '3.12' and sys_platform == 'win32'", + "python_full_version < '3.12' and sys_platform == 'emscripten'", + "python_full_version < '3.12' and sys_platform != 'emscripten' and sys_platform != 'win32'", +] +sdist = { url = "https://files.pythonhosted.org/packages/d0/ad/fed0499ce6a338d2a03ebae59cd15093910c8875328855781952abf6c2fe/numpy-2.4.6.tar.gz", hash = "sha256:f3a3570c4a2a16746ac2c31a7c7c7b0c186b95ce902e33db6f28094ed7387dda", size = 20735807 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/b3/49/ec46835a70be8fa6446c495126ac84fdb28cb2558e1620ffb87a10c8b64c/numpy-2.4.6-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:0280e0356c0829a18d9de1cb7eee50ec22ca639878d7240307ca0943d73cd2c4", size = 16969194 }, + { url = "https://files.pythonhosted.org/packages/0e/0d/f5957185c0ee2f3e12f78715aa9e3b353fd83633316c8532b38faa37e3f6/numpy-2.4.6-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:110f8b71aacb688ec69062bb7f6938a0f8acb01b7c1c4beb453c65b6d234584d", size = 14964111 }, + { url = "https://files.pythonhosted.org/packages/ad/40/40a40ee0ddf7ceb782c49af278894b686e586d65d8c1889c8b5da01a3d7d/numpy-2.4.6-cp311-cp311-macosx_14_0_arm64.whl", hash = "sha256:4cfe66903cc32a9921a6733d96b19bb6abf310397581bbad89c228f5abaf0ee8", size = 5469159 }, + { url = "https://files.pythonhosted.org/packages/63/13/f9a8046535cb21deae82f8d03de9617e08882d274fad2539630761888228/numpy-2.4.6-cp311-cp311-macosx_14_0_x86_64.whl", hash = "sha256:8155154c7c691289fe18f510b5d4657c68c67989f293f0535a91360392ff6538", size = 6798936 }, + { url = "https://files.pythonhosted.org/packages/33/a8/6fa8c1a345a8c85dbb21932c447bee07c30a2c2a3f31e369c0a84b300147/numpy-2.4.6-cp311-cp311-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:0ab0a9c4ffb1a6d95ef519fe4247dba8eb6b18ad93999f76b7f657039acabd47", size = 15966692 }, + { url = "https://files.pythonhosted.org/packages/02/03/74fe2a4cb3817d94d86402f2506554130a2f01414e299b5a843e5a8a957f/numpy-2.4.6-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:89cd468399cfd2504718f0ba50e410dca55a170b61a02ad92bb18c8a65186e93", size = 16918164 }, + { url = "https://files.pythonhosted.org/packages/c5/80/3615be3313f7e7696609bc194b9f0101da809df79e859bdb84e0cd043f46/numpy-2.4.6-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:c2d37ab77531417474168eb79d6d80b14f821a966818505d03013d0833edb7a8", size = 17322877 }, + { url = "https://files.pythonhosted.org/packages/ca/ac/a691e0fe2675e370d0e08ff905adc49a1c8830e8cae03efe4477e92cd55d/numpy-2.4.6-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:f407cb6b8e9d6d8c626bc73c945db1706035af8fd632295547bf1c9e46d092d6", size = 18651487 }, + { url = "https://files.pythonhosted.org/packages/15/a7/9bc1cd626d7bf6869bfedf27b91b6ab5dd607758bf8e959d6fa80c6a59cb/numpy-2.4.6-cp311-cp311-win32.whl", hash = "sha256:ddea102b48f9e339f3948bf22040944184627a30fdf7f858667673b9c5f033c8", size = 6233945 }, + { url = "https://files.pythonhosted.org/packages/c5/31/7fc6239c12bce7e931463251cca4426c465e1876ba3cc785402ef4dd8f4e/numpy-2.4.6-cp311-cp311-win_amd64.whl", hash = "sha256:1e254a00cdf42b1e4d5b3d68d33af63268d41340d8885df2ab6470f2e1500147", size = 12608406 }, + { url = "https://files.pythonhosted.org/packages/27/83/140f85a466595a16382996a1bf06b2b54bcd597488921b0c9daaeeda72af/numpy-2.4.6-cp311-cp311-win_arm64.whl", hash = "sha256:ed9749eef4cbd126da3dc1d6bcb3a57f5eb7ac6a6484146bdbf743f552dfc577", size = 10479528 }, + { url = "https://files.pythonhosted.org/packages/95/2a/3d7b5ac8aac24feaf9ad7ed58f45b0bbc06d37e4338ae84c9f2298b570f9/numpy-2.4.6-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:001fbb8e08d942dd57599e781f2472269ee7f2755fae407b4f67b2f0b17da3f1", size = 16689119 }, + { url = "https://files.pythonhosted.org/packages/ea/12/92c4c131527599e8288d6918e888d88726f84d805d784b771f32408aeaef/numpy-2.4.6-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:ebfb099f8dcf083deef3ac1ca4c1503f387cf76296fcb3816b66f5ecb5f54fdb", size = 14699246 }, + { url = "https://files.pythonhosted.org/packages/ad/fe/c0a6b7b2ca128a8fb228575147073b660656734b8ebe4d76c8fd748dcc79/numpy-2.4.6-cp312-cp312-macosx_14_0_arm64.whl", hash = "sha256:3213d622a0283a39a93d188f3cf72b26862df52fbb4ca3697f51705016523d41", size = 5204410 }, + { url = "https://files.pythonhosted.org/packages/f3/d4/9770d14ba719432bb90a421bfd443872ed0f70f7264b64bec12ea363d5fd/numpy-2.4.6-cp312-cp312-macosx_14_0_x86_64.whl", hash = "sha256:357cc07a6d7b0b182ff02249616a03742827ebb1277546b5c7cd7f7620a45698", size = 6551240 }, + { url = "https://files.pythonhosted.org/packages/c9/c6/50a46a6205feba2343f1d6d17438107c5dc491ed1c736e6ea68689fd906b/numpy-2.4.6-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:5f9fb9157b4ce2971008323afe46053787b526ef624fea915b261468a8421a0f", size = 15671012 }, + { url = "https://files.pythonhosted.org/packages/99/60/14115e6364fa676c5397c2ad3004e527e9aa487abf5d0706ec81bbd08529/numpy-2.4.6-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:90f9849678c75fe7afa2d348ac842c168b0a4d3d61919687216dfc547976d853", size = 16645538 }, + { url = "https://files.pythonhosted.org/packages/ae/c5/693cbe59e57db94d2231fa519ca3978dc9e19da5a8f088588f5c6e947ff2/numpy-2.4.6-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:c1a2af6c6ef86344a6b0db6b97834208bf598db514f2b155042439b62605601a", size = 17020706 }, + { url = "https://files.pythonhosted.org/packages/ef/fc/85b7c4eff9b4966ade25c2273cf7e7012e92366c032058653934b37de044/numpy-2.4.6-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:e5805d5a22fd19c8ccff10a9561f9df94436b0545619ea579db2d3c35294bce2", size = 18368541 }, + { url = "https://files.pythonhosted.org/packages/f6/81/e1b27545deedce7f4a0b348618c6b62d74e36a4dc9ccd42f3eb2f85eee32/numpy-2.4.6-cp312-cp312-win32.whl", hash = "sha256:e3eeb0aabd6bd5ce64faae67e9935203a6991b4bc2a485a767fbafb2c5125f45", size = 5962825 }, + { url = "https://files.pythonhosted.org/packages/ab/ca/feab00bd44aa5fe1ad2c18f08b4d3bb92e26484b0b1d1443897809ed528c/numpy-2.4.6-cp312-cp312-win_amd64.whl", hash = "sha256:d8e8286dd7cea7895157318d1b91cdacac64c479f3cbc8dce548331728484751", size = 12321687 }, + { url = "https://files.pythonhosted.org/packages/63/cf/5a6d34850a39d1093558564f77ee8e8e0bee5061151b8f05a55711001ec7/numpy-2.4.6-cp312-cp312-win_arm64.whl", hash = "sha256:4081eb135ac24158bd51cdfbef16f1c64df7063b1143f24731387137c092bec8", size = 10221482 }, + { url = "https://files.pythonhosted.org/packages/de/12/b422cc84439adc0d00de605bf4a308890ae5c26f2c71fbd73e5d08fbb0dd/numpy-2.4.6-pp311-pypy311_pp73-macosx_10_15_x86_64.whl", hash = "sha256:55cced7c52e981362f708ad635198e97a752dfba412cc03c23bbf3bd8d5cd662", size = 16847511 }, + { url = "https://files.pythonhosted.org/packages/44/53/f481bef68011740f8849418d82db07230e825013f31f4eef5ba5b805316a/numpy-2.4.6-pp311-pypy311_pp73-macosx_11_0_arm64.whl", hash = "sha256:d6da64deb6b8ed903e7560180a92f2d804ee1ba5eeb849ac2748b8c1aba1f6d7", size = 14889064 }, + { url = "https://files.pythonhosted.org/packages/7f/57/42ed575c10ced8af951d426bc4e1f8aff16fd851db33f067036215a7f860/numpy-2.4.6-pp311-pypy311_pp73-macosx_14_0_arm64.whl", hash = "sha256:68a5124b13fa6cc2086764a20005d30bc0548146f7f5322f02fce212ca14317f", size = 5394157 }, + { url = "https://files.pythonhosted.org/packages/6a/ef/f66cc724fcc36c1e364c67f51ae9146090b8b584f27d58b97fdae3edd737/numpy-2.4.6-pp311-pypy311_pp73-macosx_14_0_x86_64.whl", hash = "sha256:948424b06129ce883307e8cff868c31396d8dc7630a59c61d70d98dbe70f222c", size = 6708728 }, + { url = "https://files.pythonhosted.org/packages/1a/9c/c531f2293b91265d8b48e9b329f54fdd7ffae73cb4134ea10cca4237e9cc/numpy-2.4.6-pp311-pypy311_pp73-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:5dbbdb29840ca3d91ee0fece42fc29278886d908280bfec0a5846c6f901a3eb0", size = 15798374 }, + { url = "https://files.pythonhosted.org/packages/1a/b0/413077f6b1153ed3cba361401c6783bbad6114804a000cc22eb71c13e190/numpy-2.4.6-pp311-pypy311_pp73-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:8ad03c0965fb3c692200e74d458ca28c1dbb4ce96f9a479a8aa041ad5fabca02", size = 16747286 }, + { url = "https://files.pythonhosted.org/packages/15/ce/e5ec180bc41812edcd8daeb8639d205622c0e8c02259d8ab25a0201b3c2a/numpy-2.4.6-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:2803abfebfc990042cd494d8ce2d5f82e9d847af6d35ec486923aa19dbad5e73", size = 12504263 }, +] + +[[package]] +name = "numpy" +version = "2.5.2" +source = { registry = "https://pypi.org/simple" } +resolution-markers = [ + "python_full_version >= '3.12' and sys_platform == 'win32'", + "python_full_version >= '3.12' and sys_platform == 'emscripten'", + "python_full_version >= '3.12' and sys_platform != 'emscripten' and sys_platform != 'win32'", +] +sdist = { url = "https://files.pythonhosted.org/packages/9a/80/db0b4559e57ec36362bedbb05530a87fafbcb6067708c946967a41d449e7/numpy-2.5.2.tar.gz", hash = "sha256:d482d171c406ae88c5b19cad3b6a1c4c5209f886ab74bc44c2c865c23f52d860", size = 20773161 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/69/72/dccb0aaf40972777283303919f613964227266d0c13adebb79ac124f1c3e/numpy-2.5.2-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:14e373cfc6387177e8409dac3c7159be8eb05cd77096cd7c950268b86f62831c", size = 16891693 }, + { url = "https://files.pythonhosted.org/packages/60/2e/b5aee50a1f74ac815cf8331812cb8251e29024025de462e0c047641c614c/numpy-2.5.2-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:4bbd96c833ecc8cc069ce518078fc8c60cb9cbfb0fea5b7a803ad65035596d03", size = 11903109 }, + { url = "https://files.pythonhosted.org/packages/f3/f4/29e78102a80601cf034d4e9767022cffeca2c3b4c926e1754572ca95593d/numpy-2.5.2-cp312-cp312-macosx_14_0_arm64.whl", hash = "sha256:6e8172ddfcf5cf74b811d372b570b83c60bd2de87a6fbfbebdadb4a9bd9c6cbb", size = 5350202 }, + { url = "https://files.pythonhosted.org/packages/11/4b/dcd3b7eadaf4035d2c7a4289d232523a6964f602598ef7674e4bd7291f93/numpy-2.5.2-cp312-cp312-macosx_14_0_x86_64.whl", hash = "sha256:65f188481f1669e26f62b701e8205d19e460fa4a9b52a1414ba382330e4a3414", size = 6687736 }, + { url = "https://files.pythonhosted.org/packages/e5/21/4947e0e9d6c9fc2e2ff15b8949049ee44f63adb9cacc729ab8793f97e712/numpy-2.5.2-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:8ee9c4eeb8454b3660a8b53493563c3e121c2fc94fbd72b848ef814ed7b676a9", size = 15612696 }, + { url = "https://files.pythonhosted.org/packages/3a/5f/62d28cf019460c7f1394105b4d49d9911a9c444cb77ab0bd95a204c5a6de/numpy-2.5.2-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:3cdec01fa790a186d430433fdd4d4ffb70eed6f0eeb4bf05c8dbe2dce0a9bcb8", size = 16722264 }, + { url = "https://files.pythonhosted.org/packages/14/25/3f0be4c1b9fdf5dd5e708a6806978564d7c46a055c000496309ff2a2f8af/numpy-2.5.2-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:7999d4ddb0c4025018373fd787510d46e04c769467af22869707b3c1cfd459ab", size = 16974396 }, + { url = "https://files.pythonhosted.org/packages/22/72/6262cbdeeb45da9d971e40715f579d791603ba8ec0b5e2db1ac55454421d/numpy-2.5.2-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:c1f017dc0875c9209d219f97feceb7d54c2661bb243deb4114478e1295808af7", size = 18476044 }, + { url = "https://files.pythonhosted.org/packages/36/33/29208b8b075bde62d26a81d14b358c42b0f69b6cabd98d4ff97f37f22b05/numpy-2.5.2-cp312-cp312-win32.whl", hash = "sha256:d6a48072864e3324e194a8fbb3c657bcc5b5c869dbc64c9537b1d5c862572c0a", size = 6072817 }, + { url = "https://files.pythonhosted.org/packages/7f/b9/87fea2769fe1c47c1b5b01d8310772c9d1a85d485de7cf386ef7a3332b02/numpy-2.5.2-cp312-cp312-win_amd64.whl", hash = "sha256:28ac63476ec7651484215ee7fa15a1f78b57c14621f01e392afe17b9a1390ce4", size = 12464674 }, + { url = "https://files.pythonhosted.org/packages/14/52/032b97e00461ab0809bbe4c588b035620e5a14b8cdee47ecddefc7b17d33/numpy-2.5.2-cp312-cp312-win_arm64.whl", hash = "sha256:27650bb0e7140fa3d37b9923b4803645e0b125d190f326eecfd3f4dad8e8ade1", size = 10397131 }, +] + +[[package]] +name = "nvidia-cublas" +version = "13.1.1.3" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "nvidia-cuda-nvrtc", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, +] +wheels = [ + { url = "https://files.pythonhosted.org/packages/a7/a1/0bd24ee8c8d03adac032fd2909426a00c88f8c57961b1277ded97f91119f/nvidia_cublas-13.1.1.3-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:b7a210458267ac818974c53038fbec2e969d5c99f305ab15c72522fa9f001dd5", size = 542848918 }, + { url = "https://files.pythonhosted.org/packages/3b/cd/154ca20c38269e05eff77c1464e6c1da89f50a6390b565e9d82e06bc11e1/nvidia_cublas-13.1.1.3-py3-none-manylinux_2_27_x86_64.whl", hash = "sha256:37936a16db8fe4ac1f065c2139360608a543a09275cb1a1af612e08cfa065436", size = 423138758 }, +] + +[[package]] +name = "nvidia-cuda-cupti" +version = "13.0.85" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/2a/2a/80353b103fc20ce05ef51e928daed4b6015db4aaa9162ed0997090fe2250/nvidia_cuda_cupti-13.0.85-py3-none-manylinux_2_25_aarch64.whl", hash = "sha256:796bd679890ee55fb14a94629b698b6db54bcfd833d391d5e94017dd9d7d3151", size = 10310827 }, + { url = "https://files.pythonhosted.org/packages/33/6d/737d164b4837a9bbd202f5ae3078975f0525a55730fe871d8ed4e3b952b0/nvidia_cuda_cupti-13.0.85-py3-none-manylinux_2_25_x86_64.whl", hash = "sha256:4eb01c08e859bf924d222250d2e8f8b8ff6d3db4721288cf35d14252a4d933c8", size = 10715597 }, +] + +[[package]] +name = "nvidia-cuda-nvrtc" +version = "13.0.88" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/c3/68/483a78f5e8f31b08fb1bb671559968c0ca3a065ac7acabfc7cee55214fd6/nvidia_cuda_nvrtc-13.0.88-py3-none-manylinux2010_x86_64.manylinux_2_12_x86_64.whl", hash = "sha256:ad9b6d2ead2435f11cbb6868809d2adeeee302e9bb94bcf0539c7a40d80e8575", size = 90215200 }, + { url = "https://files.pythonhosted.org/packages/b7/dc/6bb80850e0b7edd6588d560758f17e0550893a1feaf436807d64d2da040f/nvidia_cuda_nvrtc-13.0.88-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:d27f20a0ca67a4bb34268a5e951033496c5b74870b868bacd046b1b8e0c3267b", size = 43015449 }, +] + +[[package]] +name = "nvidia-cuda-runtime" +version = "13.0.96" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/87/4f/17d7b9b8e285199c58ce28e31b5c5bbaa4d8271af06a89b6405258245de2/nvidia_cuda_runtime-13.0.96-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:ef9bcbe90493a2b9d810e43d249adb3d02e98dd30200d86607d8d02687c43f55", size = 2261060 }, + { url = "https://files.pythonhosted.org/packages/2e/24/d1558f3b68b1d26e706813b1d10aa1d785e4698c425af8db8edc3dced472/nvidia_cuda_runtime-13.0.96-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:7f82250d7782aa23b6cfe765ecc7db554bd3c2870c43f3d1821f1d18aebf0548", size = 2243632 }, +] + +[[package]] +name = "nvidia-cudnn-cu13" +version = "9.20.0.48" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "nvidia-cublas", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, +] +wheels = [ + { url = "https://files.pythonhosted.org/packages/56/c5/83384d846b2fd17c44bd499b36c75a45ed4f095fbbb2252294e89cea5c5c/nvidia_cudnn_cu13-9.20.0.48-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:e31454ae00094b0c55319d9d15b6fa2fc50a9e1c0f5c8c80fb75258234e731e1", size = 444574296 }, + { url = "https://files.pythonhosted.org/packages/6e/5e/edb9c0ae051602c3ccaffe424256463636d639e27d7f302dde9975ef9e7a/nvidia_cudnn_cu13-9.20.0.48-py3-none-manylinux_2_27_x86_64.whl", hash = "sha256:0c45dd8eeb50b603f07995b1b300c62ffe6a1980482b82b3bcf94a4ca9d49304", size = 366173588 }, +] + +[[package]] +name = "nvidia-cufft" +version = "12.0.0.61" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "nvidia-nvjitlink", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, +] +wheels = [ + { url = "https://files.pythonhosted.org/packages/8b/ae/f417a75c0259e85c1d2f83ca4e960289a5f814ed0cea74d18c353d3e989d/nvidia_cufft-12.0.0.61-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:2708c852ef8cd89d1d2068bdbece0aa188813a0c934db3779b9b1faa8442e5f5", size = 214053554 }, + { url = "https://files.pythonhosted.org/packages/a8/2f/7b57e29836ea8714f81e9898409196f47d772d5ddedddf1592eadb8ab743/nvidia_cufft-12.0.0.61-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:6c44f692dce8fd5ffd3e3df134b6cdb9c2f72d99cf40b62c32dde45eea9ddad3", size = 214085489 }, +] + +[[package]] +name = "nvidia-cufile" +version = "1.15.1.6" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/3f/70/4f193de89a48b71714e74602ee14d04e4019ad36a5a9f20c425776e72cd6/nvidia_cufile-1.15.1.6-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:08a3ecefae5a01c7f5117351c64f17c7c62efa5fffdbe24fc7d298da19cd0b44", size = 1223672 }, + { url = "https://files.pythonhosted.org/packages/ab/73/cc4a14c9813a8a0d509417cf5f4bdaba76e924d58beb9864f5a7baceefbf/nvidia_cufile-1.15.1.6-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:bdc0deedc61f548bddf7733bdc216456c2fdb101d020e1ab4b88d232d5e2f6d1", size = 1136992 }, +] + +[[package]] +name = "nvidia-curand" +version = "10.4.0.35" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/1e/72/7c2ae24fb6b63a32e6ae5d241cc65263ea18d08802aaae087d9f013335a2/nvidia_curand-10.4.0.35-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:133df5a7509c3e292aaa2b477afd0194f06ce4ea24d714d616ff36439cee349a", size = 61962106 }, + { url = "https://files.pythonhosted.org/packages/a5/9f/be0a41ca4a4917abf5cb9ae0daff1a6060cc5de950aec0396de9f3b52bc5/nvidia_curand-10.4.0.35-py3-none-manylinux_2_27_x86_64.whl", hash = "sha256:1aee33a5da6e1db083fe2b90082def8915f30f3248d5896bcec36a579d941bfc", size = 59544258 }, +] + +[[package]] +name = "nvidia-cusolver" +version = "12.0.4.66" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "nvidia-cublas", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, + { name = "nvidia-cusparse", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, + { name = "nvidia-nvjitlink", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, +] +wheels = [ + { url = "https://files.pythonhosted.org/packages/c8/c3/b30c9e935fc01e3da443ec0116ed1b2a009bb867f5324d3f2d7e533e776b/nvidia_cusolver-12.0.4.66-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:02c2457eaa9e39de20f880f4bd8820e6a1cfb9f9a34f820eb12a155aa5bc92d2", size = 223467760 }, + { url = "https://files.pythonhosted.org/packages/5f/67/cba3777620cdacb99102da4042883709c41c709f4b6323c10781a9c3aa34/nvidia_cusolver-12.0.4.66-py3-none-manylinux_2_27_x86_64.whl", hash = "sha256:0a759da5dea5c0ea10fd307de75cdeb59e7ea4fcb8add0924859b944babf1112", size = 200941980 }, +] + +[[package]] +name = "nvidia-cusparse" +version = "12.6.3.3" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "nvidia-nvjitlink", marker = "sys_platform != 'emscripten' and sys_platform != 'win32'" }, +] +wheels = [ + { url = "https://files.pythonhosted.org/packages/f8/94/5c26f33738ae35276672f12615a64bd008ed5be6d1ebcb23579285d960a9/nvidia_cusparse-12.6.3.3-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:80bcc4662f23f1054ee334a15c72b8940402975e0eab63178fc7e670aa59472c", size = 162155568 }, + { url = "https://files.pythonhosted.org/packages/fa/18/623c77619c31d62efd55302939756966f3ecc8d724a14dab2b75f1508850/nvidia_cusparse-12.6.3.3-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:2b3c89c88d01ee0e477cb7f82ef60a11a4bcd57b6b87c33f789350b59759360b", size = 145942937 }, +] + +[[package]] +name = "nvidia-cusparselt-cu13" +version = "0.8.1" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/46/e1/cdc1797eadf82d3a9a575a19b33fdc871a97edbec42c00b5b5e914f4aff4/nvidia_cusparselt_cu13-0.8.1-py3-none-manylinux2014_aarch64.whl", hash = "sha256:4dca476c50bf4780d46cd0bfbd82e2bc10a08e4fef7950917ce8d7578d22a23f", size = 221051344 }, + { url = "https://files.pythonhosted.org/packages/34/7d/2661f2fb3ac4302f3a246f5fc030213ac60c1fe0bce84f9783dbd831dbb7/nvidia_cusparselt_cu13-0.8.1-py3-none-manylinux2014_x86_64.whl", hash = "sha256:786ce87568c303fadb5afcc7102d454cd3040d75f6f8626f5db460d1871f4dd0", size = 170148586 }, +] + +[[package]] +name = "nvidia-nccl-cu13" +version = "2.29.7" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/72/0d/daf50d44177ee0cbc7ff0a0c91eb5ff676c82be42f9a970bc7597f440c3a/nvidia_nccl_cu13-2.29.7-py3-none-manylinux_2_18_aarch64.whl", hash = "sha256:674a12383e3c38a1bcccae7d4f3633b37852230b6047883cb2f4c2d1b36d9bf5", size = 206014712 }, + { url = "https://files.pythonhosted.org/packages/67/f4/58e4e91b6919367c7aafb8e36fce9aad1a3047e536bf7e2fd560927d3a4c/nvidia_nccl_cu13-2.29.7-py3-none-manylinux_2_18_x86_64.whl", hash = "sha256:edd81538446786ec3b73972543e53bb43bcaf0bfc8ef76cb679fcc390ffe136d", size = 205976000 }, +] + +[[package]] +name = "nvidia-nvjitlink" +version = "13.3.33" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/f0/ee/580ca6f29dcab0221db8706badca1bbbb084f1975c4d4e83329c3a7e31f0/nvidia_nvjitlink-13.3.33-py3-none-manylinux2010_x86_64.manylinux_2_12_x86_64.whl", hash = "sha256:26a6de7fb4c8fdaa7703d3dad720d6d427ddfea5c48a528fd97c11733ad830e5", size = 40742423 }, + { url = "https://files.pythonhosted.org/packages/69/30/45414e35ff2eee7db3da037e5707037ccf9d2b5218ffbdb055ea4d5aa98a/nvidia_nvjitlink-13.3.33-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:ce48b37dfeb3cb1eae4cf85adacb47d7a6539ea2272870c9a3628ce275c2037e", size = 39168635 }, +] + +[[package]] +name = "nvidia-nvshmem-cu13" +version = "3.4.5" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/dc/0f/05cc9c720236dcd2db9c1ab97fff629e96821be2e63103569da0c9b72f19/nvidia_nvshmem_cu13-3.4.5-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:6dc2a197f38e5d0376ad52cd1a2a3617d3cdc150fd5966f4aee9bcebb1d68fe9", size = 60215947 }, + { url = "https://files.pythonhosted.org/packages/3c/35/a9bf80a609e74e3b000fef598933235c908fcefcef9026042b8e6dfde2a9/nvidia_nvshmem_cu13-3.4.5-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:290f0a2ee94c9f3687a02502f3b9299a9f9fe826e6d0287ee18482e78d495b80", size = 60412546 }, +] + +[[package]] +name = "nvidia-nvtx" +version = "13.0.85" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/c2/f3/d86c845465a2723ad7e1e5c36dcd75ddb82898b3f53be47ebd429fb2fa5d/nvidia_nvtx-13.0.85-py3-none-manylinux1_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:4936d1d6780fbe68db454f5e72a42ff64d1fd6397df9f363ae786930fd5c1cd4", size = 148047 }, + { url = "https://files.pythonhosted.org/packages/a8/64/3708a90d1ebe202ffdeb7185f878a3c84d15c2b2c31858da2ce0583e2def/nvidia_nvtx-13.0.85-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:cb7780edb6b14107373c835bf8b72e7a178bac7367e23da7acb108f973f157a6", size = 148878 }, +] + +[[package]] +name = "packaging" +version = "26.3" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/7d/fa/3944b40b07da9ce895c0e6303a5ab7d53da063554f534556b134a54d6093/packaging-26.3.tar.gz", hash = "sha256:94edc256424af38762eb31306eed28beb9f0efc50a8837492c9d6fd6004aed79", size = 313412 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/63/34/ba1c580383c9eada3711951fef0795c80b829a078d72188184bcab9dd527/packaging-26.3-py3-none-any.whl", hash = "sha256:d7193f7c8e4e93f444fde0262bf90af30e16fa0ad0ad44cb553c87339b23cd1c", size = 129956 }, +] + +[[package]] +name = "pyarrow" +version = "25.0.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/3d/e3/27f57f80141379d60defe6703eb50a707325706f07fedfd1312c7a751995/pyarrow-25.0.1.tar.gz", hash = "sha256:9150a83248bfed9813ea3c3af74c3856c1984d444aa28e58bf7733b9750ddf6a", size = 1201653 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/ee/8b/0d23b47702fcfe8b3618d5292035099675c5a1c48258932350c08020f7b5/pyarrow-25.0.1-cp311-cp311-macosx_12_0_arm64.whl", hash = "sha256:51093dd9e10325fbdb3c10a2ae7c4806e5c822d94e74ae4938b26524a3323fee", size = 35946180 }, + { url = "https://files.pythonhosted.org/packages/d8/17/707d17a5476c55a9541fde0db8213ac30979a792864d72415f176ba50c45/pyarrow-25.0.1-cp311-cp311-macosx_12_0_x86_64.whl", hash = "sha256:eb6203482ff3746a5632303a7279ae0b5a304c46985b49ed1378cb350ea6728d", size = 37644787 }, + { url = "https://files.pythonhosted.org/packages/c1/b2/cdc98ecf1a6408280bc3a6a07054cdd99a3f4670acc0545d383ce113e87d/pyarrow-25.0.1-cp311-cp311-manylinux_2_28_aarch64.whl", hash = "sha256:880523be3d29efcf83d3998835d206118ccf35e3871dbd2fb60408cf6b007a80", size = 46834633 }, + { url = "https://files.pythonhosted.org/packages/c8/6e/d3fafc41f378b2c65be43b827798c0fae42049a641c8526633ed3eb573e2/pyarrow-25.0.1-cp311-cp311-manylinux_2_28_x86_64.whl", hash = "sha256:25f8720bf6387d5dc2ebd2622112de630760419e4b66134405dd24110d15f37e", size = 50065507 }, + { url = "https://files.pythonhosted.org/packages/d5/12/8d0698954b8c3001844a898e0a6900bebe83d7ee40c11195174c5122f324/pyarrow-25.0.1-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:4facd65742a024a4a366328a1d2292062d72d6e023c1b7dda8d4c37544933a25", size = 49955690 }, + { url = "https://files.pythonhosted.org/packages/d3/0b/1ecb936ac6409e90a34d58eea1c7cec09a9ae6d2141b9e49ad01a2b1ea47/pyarrow-25.0.1-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:aa0559502e1cd6254d6814614085dd9c5a3dd0419362978a936a3f68a9e5c3df", size = 53128198 }, + { url = "https://files.pythonhosted.org/packages/8e/1c/5236033550633c9b7377b2a53660b2bbb06cb06dc09c4356332d67643ca1/pyarrow-25.0.1-cp311-cp311-win_amd64.whl", hash = "sha256:62cd0d785b8aa6675ee355f9fc02252a340f4441257c42674937826fd7594325", size = 27857263 }, + { url = "https://files.pythonhosted.org/packages/a6/e2/9ab15b88cbfac28e16419ce5439ec29234c5172cb8259301b4ba639bdec0/pyarrow-25.0.1-cp312-cp312-macosx_12_0_arm64.whl", hash = "sha256:df961f2e7ae9cf496459259d798652c70625f6c080650d6952f8c04053c58ee9", size = 35861559 }, + { url = "https://files.pythonhosted.org/packages/58/79/a0036dbe1eabe1f73127427342f1d99982584c4a2cde2651d6c93499c6f6/pyarrow-25.0.1-cp312-cp312-macosx_12_0_x86_64.whl", hash = "sha256:cc4aa407fde9fc660be3939e49ea31f50f3e9fec17c0ec63159f7711edd3efc9", size = 37628383 }, + { url = "https://files.pythonhosted.org/packages/13/49/d93a57d375f4bf0cf82913dd6bb54acafde83dd993be2282c81ac5616cad/pyarrow-25.0.1-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:4340f0ba6c1d2e13f21658de1d7c662ca2545018568d0030a1e9afca159d87e3", size = 46820190 }, + { url = "https://files.pythonhosted.org/packages/60/c9/711ca85d79f1ec98f29a5eae2b051e25b4ecec5de3e3c0e2d5c5dcb15664/pyarrow-25.0.1-cp312-cp312-manylinux_2_28_x86_64.whl", hash = "sha256:5389cdf79447ed1515c9e31620e6e1e2302249564d603f2ad727d4f6d313e4c3", size = 50102437 }, + { url = "https://files.pythonhosted.org/packages/80/53/8fb8359ff17cfb6263a1cf3ebf7caec9fe197de118719e84fcb1d0618026/pyarrow-25.0.1-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:d51592cb7561e87877c506113e7adbf1342ab579e6c21f0ef44b8ba41cb74c80", size = 49942424 }, + { url = "https://files.pythonhosted.org/packages/e8/83/4e5ae02a9341571b18a6fca380ac7a58ce6ddae7ab3c060208c0a1e79f02/pyarrow-25.0.1-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:6109c94d8b9f3b17a041daca16cacb2f651ad8f1ef70a4232c2c0f37a23da2a8", size = 53144206 }, + { url = "https://files.pythonhosted.org/packages/65/ee/197cbf47e49f83e6ebeb946a5259a48a638dea27ac774db42fe78022179d/pyarrow-25.0.1-cp312-cp312-win_amd64.whl", hash = "sha256:8858d7bfc22e3f51529aeaa4077225029724623e4595dc9eff8c793935c34140", size = 27953934 }, +] + +[[package]] +name = "pycparser" +version = "3.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/1b/7d/92392ff7815c21062bea51aa7b87d45576f649f16458d78b7cf94b9ab2e6/pycparser-3.0.tar.gz", hash = "sha256:600f49d217304a5902ac3c37e1281c9fe94e4d0489de643a9504c5cdfdfc6b29", size = 103492 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/0c/c3/44f3fbbfa403ea2a7c779186dc20772604442dde72947e7d01069cbe98e3/pycparser-3.0-py3-none-any.whl", hash = "sha256:b727414169a36b7d524c1c3e31839a521725078d7b2ff038656844266160a992", size = 48172 }, +] + +[[package]] +name = "pyyaml" +version = "6.0.3" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/05/8e/961c0007c59b8dd7729d542c61a4d537767a59645b82a0b521206e1e25c2/pyyaml-6.0.3.tar.gz", hash = "sha256:d76623373421df22fb4cf8817020cbb7ef15c725b9d5e45f17e189bfc384190f", size = 130960 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/6d/16/a95b6757765b7b031c9374925bb718d55e0a9ba8a1b6a12d25962ea44347/pyyaml-6.0.3-cp311-cp311-macosx_10_13_x86_64.whl", hash = "sha256:44edc647873928551a01e7a563d7452ccdebee747728c1080d881d68af7b997e", size = 185826 }, + { url = "https://files.pythonhosted.org/packages/16/19/13de8e4377ed53079ee996e1ab0a9c33ec2faf808a4647b7b4c0d46dd239/pyyaml-6.0.3-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:652cb6edd41e718550aad172851962662ff2681490a8a711af6a4d288dd96824", size = 175577 }, + { url = "https://files.pythonhosted.org/packages/0c/62/d2eb46264d4b157dae1275b573017abec435397aa59cbcdab6fc978a8af4/pyyaml-6.0.3-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:10892704fc220243f5305762e276552a0395f7beb4dbf9b14ec8fd43b57f126c", size = 775556 }, + { url = "https://files.pythonhosted.org/packages/10/cb/16c3f2cf3266edd25aaa00d6c4350381c8b012ed6f5276675b9eba8d9ff4/pyyaml-6.0.3-cp311-cp311-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:850774a7879607d3a6f50d36d04f00ee69e7fc816450e5f7e58d7f17f1ae5c00", size = 882114 }, + { url = "https://files.pythonhosted.org/packages/71/60/917329f640924b18ff085ab889a11c763e0b573da888e8404ff486657602/pyyaml-6.0.3-cp311-cp311-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:b8bb0864c5a28024fac8a632c443c87c5aa6f215c0b126c449ae1a150412f31d", size = 806638 }, + { url = "https://files.pythonhosted.org/packages/dd/6f/529b0f316a9fd167281a6c3826b5583e6192dba792dd55e3203d3f8e655a/pyyaml-6.0.3-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:1d37d57ad971609cf3c53ba6a7e365e40660e3be0e5175fa9f2365a379d6095a", size = 767463 }, + { url = "https://files.pythonhosted.org/packages/f2/6a/b627b4e0c1dd03718543519ffb2f1deea4a1e6d42fbab8021936a4d22589/pyyaml-6.0.3-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:37503bfbfc9d2c40b344d06b2199cf0e96e97957ab1c1b546fd4f87e53e5d3e4", size = 794986 }, + { url = "https://files.pythonhosted.org/packages/45/91/47a6e1c42d9ee337c4839208f30d9f09caa9f720ec7582917b264defc875/pyyaml-6.0.3-cp311-cp311-win32.whl", hash = "sha256:8098f252adfa6c80ab48096053f512f2321f0b998f98150cea9bd23d83e1467b", size = 142543 }, + { url = "https://files.pythonhosted.org/packages/da/e3/ea007450a105ae919a72393cb06f122f288ef60bba2dc64b26e2646fa315/pyyaml-6.0.3-cp311-cp311-win_amd64.whl", hash = "sha256:9f3bfb4965eb874431221a3ff3fdcddc7e74e3b07799e0e84ca4a0f867d449bf", size = 158763 }, + { url = "https://files.pythonhosted.org/packages/d1/33/422b98d2195232ca1826284a76852ad5a86fe23e31b009c9886b2d0fb8b2/pyyaml-6.0.3-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:7f047e29dcae44602496db43be01ad42fc6f1cc0d8cd6c83d342306c32270196", size = 182063 }, + { url = "https://files.pythonhosted.org/packages/89/a0/6cf41a19a1f2f3feab0e9c0b74134aa2ce6849093d5517a0c550fe37a648/pyyaml-6.0.3-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:fc09d0aa354569bc501d4e787133afc08552722d3ab34836a80547331bb5d4a0", size = 173973 }, + { url = "https://files.pythonhosted.org/packages/ed/23/7a778b6bd0b9a8039df8b1b1d80e2e2ad78aa04171592c8a5c43a56a6af4/pyyaml-6.0.3-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:9149cad251584d5fb4981be1ecde53a1ca46c891a79788c0df828d2f166bda28", size = 775116 }, + { url = "https://files.pythonhosted.org/packages/65/30/d7353c338e12baef4ecc1b09e877c1970bd3382789c159b4f89d6a70dc09/pyyaml-6.0.3-cp312-cp312-manylinux2014_s390x.manylinux_2_17_s390x.manylinux_2_28_s390x.whl", hash = "sha256:5fdec68f91a0c6739b380c83b951e2c72ac0197ace422360e6d5a959d8d97b2c", size = 844011 }, + { url = "https://files.pythonhosted.org/packages/8b/9d/b3589d3877982d4f2329302ef98a8026e7f4443c765c46cfecc8858c6b4b/pyyaml-6.0.3-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:ba1cc08a7ccde2d2ec775841541641e4548226580ab850948cbfda66a1befcdc", size = 807870 }, + { url = "https://files.pythonhosted.org/packages/05/c0/b3be26a015601b822b97d9149ff8cb5ead58c66f981e04fedf4e762f4bd4/pyyaml-6.0.3-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:8dc52c23056b9ddd46818a57b78404882310fb473d63f17b07d5c40421e47f8e", size = 761089 }, + { url = "https://files.pythonhosted.org/packages/be/8e/98435a21d1d4b46590d5459a22d88128103f8da4c2d4cb8f14f2a96504e1/pyyaml-6.0.3-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:41715c910c881bc081f1e8872880d3c650acf13dfa8214bad49ed4cede7c34ea", size = 790181 }, + { url = "https://files.pythonhosted.org/packages/74/93/7baea19427dcfbe1e5a372d81473250b379f04b1bd3c4c5ff825e2327202/pyyaml-6.0.3-cp312-cp312-win32.whl", hash = "sha256:96b533f0e99f6579b3d4d4995707cf36df9100d67e0c8303a0c55b27b5f99bc5", size = 137658 }, + { url = "https://files.pythonhosted.org/packages/86/bf/899e81e4cce32febab4fb42bb97dcdf66bc135272882d1987881a4b519e9/pyyaml-6.0.3-cp312-cp312-win_amd64.whl", hash = "sha256:5fcd34e47f6e0b794d17de1b4ff496c00986e1c83f7ab2fb8fcfe9616ff7477b", size = 154003 }, + { url = "https://files.pythonhosted.org/packages/1a/08/67bd04656199bbb51dbed1439b7f27601dfb576fb864099c7ef0c3e55531/pyyaml-6.0.3-cp312-cp312-win_arm64.whl", hash = "sha256:64386e5e707d03a7e172c0701abfb7e10f0fb753ee1d773128192742712a98fd", size = 140344 }, +] + +[[package]] +name = "requests" +version = "2.34.2" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "certifi" }, + { name = "charset-normalizer" }, + { name = "idna" }, + { name = "urllib3" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/ac/c3/e2a2b89f2d3e2179abd6d00ebd70bff6273f37fb3e0cc209f48b39d00cbf/requests-2.34.2.tar.gz", hash = "sha256:f288924cae4e29463698d6d60bc6a4da69c89185ad1e0bcc4104f584e960b9ed", size = 142856 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/a0/f4/c67b0b3f1b9245e8d266f0f112c500d50e5b4e83cb6f3b71b6528104182a/requests-2.34.2-py3-none-any.whl", hash = "sha256:2a0d60c172f83ac6ab31e4554906c0f3b3588d37b5cb939b1c061f4907e278e0", size = 73075 }, +] + +[[package]] +name = "ruamel-yaml" +version = "0.18.17" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "ruamel-yaml-clib", marker = "platform_python_implementation == 'CPython'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/3a/2b/7a1f1ebcd6b3f14febdc003e658778d81e76b40df2267904ee6b13f0c5c6/ruamel_yaml-0.18.17.tar.gz", hash = "sha256:9091cd6e2d93a3a4b157ddb8fabf348c3de7f1fb1381346d985b6b247dcd8d3c", size = 149602 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/af/fe/b6045c782f1fd1ae317d2a6ca1884857ce5c20f59befe6ab25a8603c43a7/ruamel_yaml-0.18.17-py3-none-any.whl", hash = "sha256:9c8ba9eb3e793efdf924b60d521820869d5bf0cb9c6f1b82d82de8295e290b9d", size = 121594 }, +] + +[[package]] +name = "ruamel-yaml-clib" +version = "0.2.15" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/ea/97/60fda20e2fb54b83a61ae14648b0817c8f5d84a3821e40bfbdae1437026a/ruamel_yaml_clib-0.2.15.tar.gz", hash = "sha256:46e4cc8c43ef6a94885f72512094e482114a8a706d3c555a34ed4b0d20200600", size = 225794 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/2c/80/8ce7b9af532aa94dd83360f01ce4716264db73de6bc8efd22c32341f6658/ruamel_yaml_clib-0.2.15-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:c583229f336682b7212a43d2fa32c30e643d3076178fb9f7a6a14dde85a2d8bd", size = 147998 }, + { url = "https://files.pythonhosted.org/packages/53/09/de9d3f6b6701ced5f276d082ad0f980edf08ca67114523d1b9264cd5e2e0/ruamel_yaml_clib-0.2.15-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:56ea19c157ed8c74b6be51b5fa1c3aff6e289a041575f0556f66e5fb848bb137", size = 132743 }, + { url = "https://files.pythonhosted.org/packages/0e/f7/73a9b517571e214fe5c246698ff3ed232f1ef863c8ae1667486625ec688a/ruamel_yaml_clib-0.2.15-cp311-cp311-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:5fea0932358e18293407feb921d4f4457db837b67ec1837f87074667449f9401", size = 731459 }, + { url = "https://files.pythonhosted.org/packages/9b/a2/0dc0013169800f1c331a6f55b1282c1f4492a6d32660a0cf7b89e6684919/ruamel_yaml_clib-0.2.15-cp311-cp311-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:ef71831bd61fbdb7aa0399d5c4da06bea37107ab5c79ff884cc07f2450910262", size = 749289 }, + { url = "https://files.pythonhosted.org/packages/aa/ed/3fb20a1a96b8dc645d88c4072df481fe06e0289e4d528ebbdcc044ebc8b3/ruamel_yaml_clib-0.2.15-cp311-cp311-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:617d35dc765715fa86f8c3ccdae1e4229055832c452d4ec20856136acc75053f", size = 777630 }, + { url = "https://files.pythonhosted.org/packages/60/50/6842f4628bc98b7aa4733ab2378346e1441e150935ad3b9f3c3c429d9408/ruamel_yaml_clib-0.2.15-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:1b45498cc81a4724a2d42273d6cfc243c0547ad7c6b87b4f774cb7bcc131c98d", size = 744368 }, + { url = "https://files.pythonhosted.org/packages/d3/b0/128ae8e19a7d794c2e36130a72b3bb650ce1dd13fb7def6cf10656437dcf/ruamel_yaml_clib-0.2.15-cp311-cp311-musllinux_1_2_i686.whl", hash = "sha256:def5663361f6771b18646620fca12968aae730132e104688766cf8a3b1d65922", size = 745233 }, + { url = "https://files.pythonhosted.org/packages/75/05/91130633602d6ba7ce3e07f8fc865b40d2a09efd4751c740df89eed5caf9/ruamel_yaml_clib-0.2.15-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:014181cdec565c8745b7cbc4de3bf2cc8ced05183d986e6d1200168e5bb59490", size = 770963 }, + { url = "https://files.pythonhosted.org/packages/fd/4b/fd4542e7f33d7d1bc64cc9ac9ba574ce8cf145569d21f5f20133336cdc8c/ruamel_yaml_clib-0.2.15-cp311-cp311-win32.whl", hash = "sha256:d290eda8f6ada19e1771b54e5706b8f9807e6bb08e873900d5ba114ced13e02c", size = 102640 }, + { url = "https://files.pythonhosted.org/packages/bb/eb/00ff6032c19c7537371e3119287999570867a0eafb0154fccc80e74bf57a/ruamel_yaml_clib-0.2.15-cp311-cp311-win_amd64.whl", hash = "sha256:bdc06ad71173b915167702f55d0f3f027fc61abd975bd308a0968c02db4a4c3e", size = 121996 }, + { url = "https://files.pythonhosted.org/packages/72/4b/5fde11a0722d676e469d3d6f78c6a17591b9c7e0072ca359801c4bd17eee/ruamel_yaml_clib-0.2.15-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:cb15a2e2a90c8475df45c0949793af1ff413acfb0a716b8b94e488ea95ce7cff", size = 149088 }, + { url = "https://files.pythonhosted.org/packages/85/82/4d08ac65ecf0ef3b046421985e66301a242804eb9a62c93ca3437dc94ee0/ruamel_yaml_clib-0.2.15-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:64da03cbe93c1e91af133f5bec37fd24d0d4ba2418eaf970d7166b0a26a148a2", size = 134553 }, + { url = "https://files.pythonhosted.org/packages/b9/cb/22366d68b280e281a932403b76da7a988108287adff2bfa5ce881200107a/ruamel_yaml_clib-0.2.15-cp312-cp312-manylinux1_i686.manylinux_2_28_i686.manylinux_2_5_i686.whl", hash = "sha256:f6d3655e95a80325b84c4e14c080b2470fe4f33b6846f288379ce36154993fb1", size = 737468 }, + { url = "https://files.pythonhosted.org/packages/71/73/81230babf8c9e33770d43ed9056f603f6f5f9665aea4177a2c30ae48e3f3/ruamel_yaml_clib-0.2.15-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:71845d377c7a47afc6592aacfea738cc8a7e876d586dfba814501d8c53c1ba60", size = 753349 }, + { url = "https://files.pythonhosted.org/packages/61/62/150c841f24cda9e30f588ef396ed83f64cfdc13b92d2f925bb96df337ba9/ruamel_yaml_clib-0.2.15-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:11e5499db1ccbc7f4b41f0565e4f799d863ea720e01d3e99fa0b7b5fcd7802c9", size = 788211 }, + { url = "https://files.pythonhosted.org/packages/30/93/e79bd9cbecc3267499d9ead919bd61f7ddf55d793fb5ef2b1d7d92444f35/ruamel_yaml_clib-0.2.15-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:4b293a37dc97e2b1e8a1aec62792d1e52027087c8eea4fc7b5abd2bdafdd6642", size = 743203 }, + { url = "https://files.pythonhosted.org/packages/8d/06/1eb640065c3a27ce92d76157f8efddb184bd484ed2639b712396a20d6dce/ruamel_yaml_clib-0.2.15-cp312-cp312-musllinux_1_2_i686.whl", hash = "sha256:512571ad41bba04eac7268fe33f7f4742210ca26a81fe0c75357fa682636c690", size = 747292 }, + { url = "https://files.pythonhosted.org/packages/a5/21/ee353e882350beab65fcc47a91b6bdc512cace4358ee327af2962892ff16/ruamel_yaml_clib-0.2.15-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:e5e9f630c73a490b758bf14d859a39f375e6999aea5ddd2e2e9da89b9953486a", size = 771624 }, + { url = "https://files.pythonhosted.org/packages/57/34/cc1b94057aa867c963ecf9ea92ac59198ec2ee3a8d22a126af0b4d4be712/ruamel_yaml_clib-0.2.15-cp312-cp312-win32.whl", hash = "sha256:f4421ab780c37210a07d138e56dd4b51f8642187cdfb433eb687fe8c11de0144", size = 100342 }, + { url = "https://files.pythonhosted.org/packages/b3/e5/8925a4208f131b218f9a7e459c0d6fcac8324ae35da269cb437894576366/ruamel_yaml_clib-0.2.15-cp312-cp312-win_amd64.whl", hash = "sha256:2b216904750889133d9222b7b873c199d48ecbb12912aca78970f84a5aa1a4bc", size = 119013 }, +] + +[[package]] +name = "scipy" +version = "1.17.1" +source = { registry = "https://pypi.org/simple" } +resolution-markers = [ + "python_full_version < '3.12' and sys_platform == 'win32'", + "python_full_version < '3.12' and sys_platform == 'emscripten'", + "python_full_version < '3.12' and sys_platform != 'emscripten' and sys_platform != 'win32'", +] +dependencies = [ + { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/7a/97/5a3609c4f8d58b039179648e62dd220f89864f56f7357f5d4f45c29eb2cc/scipy-1.17.1.tar.gz", hash = "sha256:95d8e012d8cb8816c226aef832200b1d45109ed4464303e997c5b13122b297c0", size = 30573822 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/df/75/b4ce781849931fef6fd529afa6b63711d5a733065722d0c3e2724af9e40a/scipy-1.17.1-cp311-cp311-macosx_10_14_x86_64.whl", hash = "sha256:1f95b894f13729334fb990162e911c9e5dc1ab390c58aa6cbecb389c5b5e28ec", size = 31613675 }, + { url = "https://files.pythonhosted.org/packages/f7/58/bccc2861b305abdd1b8663d6130c0b3d7cc22e8d86663edbc8401bfd40d4/scipy-1.17.1-cp311-cp311-macosx_12_0_arm64.whl", hash = "sha256:e18f12c6b0bc5a592ed23d3f7b891f68fd7f8241d69b7883769eb5d5dfb52696", size = 28162057 }, + { url = "https://files.pythonhosted.org/packages/6d/ee/18146b7757ed4976276b9c9819108adbc73c5aad636e5353e20746b73069/scipy-1.17.1-cp311-cp311-macosx_14_0_arm64.whl", hash = "sha256:a3472cfbca0a54177d0faa68f697d8ba4c80bbdc19908c3465556d9f7efce9ee", size = 20334032 }, + { url = "https://files.pythonhosted.org/packages/ec/e6/cef1cf3557f0c54954198554a10016b6a03b2ec9e22a4e1df734936bd99c/scipy-1.17.1-cp311-cp311-macosx_14_0_x86_64.whl", hash = "sha256:766e0dc5a616d026a3a1cffa379af959671729083882f50307e18175797b3dfd", size = 22709533 }, + { url = "https://files.pythonhosted.org/packages/4d/60/8804678875fc59362b0fb759ab3ecce1f09c10a735680318ac30da8cd76b/scipy-1.17.1-cp311-cp311-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:744b2bf3640d907b79f3fd7874efe432d1cf171ee721243e350f55234b4cec4c", size = 33062057 }, + { url = "https://files.pythonhosted.org/packages/09/7d/af933f0f6e0767995b4e2d705a0665e454d1c19402aa7e895de3951ebb04/scipy-1.17.1-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:43af8d1f3bea642559019edfe64e9b11192a8978efbd1539d7bc2aaa23d92de4", size = 35349300 }, + { url = "https://files.pythonhosted.org/packages/b4/3d/7ccbbdcbb54c8fdc20d3b6930137c782a163fa626f0aef920349873421ba/scipy-1.17.1-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:cd96a1898c0a47be4520327e01f874acfd61fb48a9420f8aa9f6483412ffa444", size = 35127333 }, + { url = "https://files.pythonhosted.org/packages/e8/19/f926cb11c42b15ba08e3a71e376d816ac08614f769b4f47e06c3580c836a/scipy-1.17.1-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:4eb6c25dd62ee8d5edf68a8e1c171dd71c292fdae95d8aeb3dd7d7de4c364082", size = 37741314 }, + { url = "https://files.pythonhosted.org/packages/95/da/0d1df507cf574b3f224ccc3d45244c9a1d732c81dcb26b1e8a766ae271a8/scipy-1.17.1-cp311-cp311-win_amd64.whl", hash = "sha256:d30e57c72013c2a4fe441c2fcb8e77b14e152ad48b5464858e07e2ad9fbfceff", size = 36607512 }, + { url = "https://files.pythonhosted.org/packages/68/7f/bdd79ceaad24b671543ffe0ef61ed8e659440eb683b66f033454dcee90eb/scipy-1.17.1-cp311-cp311-win_arm64.whl", hash = "sha256:9ecb4efb1cd6e8c4afea0daa91a87fbddbce1b99d2895d151596716c0b2e859d", size = 24599248 }, + { url = "https://files.pythonhosted.org/packages/35/48/b992b488d6f299dbe3f11a20b24d3dda3d46f1a635ede1c46b5b17a7b163/scipy-1.17.1-cp312-cp312-macosx_10_14_x86_64.whl", hash = "sha256:35c3a56d2ef83efc372eaec584314bd0ef2e2f0d2adb21c55e6ad5b344c0dcb8", size = 31610954 }, + { url = "https://files.pythonhosted.org/packages/b2/02/cf107b01494c19dc100f1d0b7ac3cc08666e96ba2d64db7626066cee895e/scipy-1.17.1-cp312-cp312-macosx_12_0_arm64.whl", hash = "sha256:fcb310ddb270a06114bb64bbe53c94926b943f5b7f0842194d585c65eb4edd76", size = 28172662 }, + { url = "https://files.pythonhosted.org/packages/cf/a9/599c28631bad314d219cf9ffd40e985b24d603fc8a2f4ccc5ae8419a535b/scipy-1.17.1-cp312-cp312-macosx_14_0_arm64.whl", hash = "sha256:cc90d2e9c7e5c7f1a482c9875007c095c3194b1cfedca3c2f3291cdc2bc7c086", size = 20344366 }, + { url = "https://files.pythonhosted.org/packages/35/f5/906eda513271c8deb5af284e5ef0206d17a96239af79f9fa0aebfe0e36b4/scipy-1.17.1-cp312-cp312-macosx_14_0_x86_64.whl", hash = "sha256:c80be5ede8f3f8eded4eff73cc99a25c388ce98e555b17d31da05287015ffa5b", size = 22704017 }, + { url = "https://files.pythonhosted.org/packages/da/34/16f10e3042d2f1d6b66e0428308ab52224b6a23049cb2f5c1756f713815f/scipy-1.17.1-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:e19ebea31758fac5893a2ac360fedd00116cbb7628e650842a6691ba7ca28a21", size = 32927842 }, + { url = "https://files.pythonhosted.org/packages/01/8e/1e35281b8ab6d5d72ebe9911edcdffa3f36b04ed9d51dec6dd140396e220/scipy-1.17.1-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:02ae3b274fde71c5e92ac4d54bc06c42d80e399fec704383dcd99b301df37458", size = 35235890 }, + { url = "https://files.pythonhosted.org/packages/c5/5c/9d7f4c88bea6e0d5a4f1bc0506a53a00e9fcb198de372bfe4d3652cef482/scipy-1.17.1-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:8a604bae87c6195d8b1045eddece0514d041604b14f2727bbc2b3020172045eb", size = 35003557 }, + { url = "https://files.pythonhosted.org/packages/65/94/7698add8f276dbab7a9de9fb6b0e02fc13ee61d51c7c3f85ac28b65e1239/scipy-1.17.1-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:f590cd684941912d10becc07325a3eeb77886fe981415660d9265c4c418d0bea", size = 37625856 }, + { url = "https://files.pythonhosted.org/packages/a2/84/dc08d77fbf3d87d3ee27f6a0c6dcce1de5829a64f2eae85a0ecc1f0daa73/scipy-1.17.1-cp312-cp312-win_amd64.whl", hash = "sha256:41b71f4a3a4cab9d366cd9065b288efc4d4f3c0b37a91a8e0947fb5bd7f31d87", size = 36549682 }, + { url = "https://files.pythonhosted.org/packages/bc/98/fe9ae9ffb3b54b62559f52dedaebe204b408db8109a8c66fdd04869e6424/scipy-1.17.1-cp312-cp312-win_arm64.whl", hash = "sha256:f4115102802df98b2b0db3cce5cb9b92572633a1197c77b7553e5203f284a5b3", size = 24547340 }, +] + +[[package]] +name = "scipy" +version = "1.18.1" +source = { registry = "https://pypi.org/simple" } +resolution-markers = [ + "python_full_version >= '3.12' and sys_platform == 'win32'", + "python_full_version >= '3.12' and sys_platform == 'emscripten'", + "python_full_version >= '3.12' and sys_platform != 'emscripten' and sys_platform != 'win32'", +] +dependencies = [ + { name = "numpy", version = "2.5.2", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/7e/74/66de6258867beb2ef08f35f9f2ac017a52cacd5081714d239ff1a442d458/scipy-1.18.1.tar.gz", hash = "sha256:52c4b7422442aba924d03ad4019852b08a92e64ea187b933135687bfe2747307", size = 30781235 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/18/f7/240c110c08693826b4513a52f5717d62ec7c7af72f2920821247c03b17b3/scipy-1.18.1-cp312-cp312-macosx_10_15_x86_64.whl", hash = "sha256:457fd7a2a8edeb044ab6ffbc0aa03ff6cd18491356e5e0c834d76ce621b916d1", size = 31111061 }, + { url = "https://files.pythonhosted.org/packages/05/4a/78c6285577c375e7cf27277ea8ee6961224327f1e1a0c44af5f17f23635c/scipy-1.18.1-cp312-cp312-macosx_12_0_arm64.whl", hash = "sha256:e708533e8b2ae2497d65346538a7dcc92814410b25b81432eac66de0f2af8265", size = 28733332 }, + { url = "https://files.pythonhosted.org/packages/a5/f6/a5b82f8abbe14d134691b8b903696f701d25a081353a29dc655c364d9e62/scipy-1.18.1-cp312-cp312-macosx_14_0_arm64.whl", hash = "sha256:7bbf207c4453ce1ad2e00b17313852b33310b83090c2311bdaf97f93c0380d12", size = 20475078 }, + { url = "https://files.pythonhosted.org/packages/23/22/0858a0bbd6b3e825ceb8cd9baf9eaf3b2f2b1d77727eb6be40500bcdc92f/scipy-1.18.1-cp312-cp312-macosx_14_0_x86_64.whl", hash = "sha256:78c0665edead396b1abb4897c41a5c1d9bf090c8a637a4c20a61678e0a264e66", size = 23108904 }, + { url = "https://files.pythonhosted.org/packages/75/9a/2e71719f31eaefe0e3a1706c4a1ded94e664bfd95ffca2b219a671faee01/scipy-1.18.1-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:3c085faa2cfa879c5141df483f836f4d691045a078224a670fa570fa01612d89", size = 34025113 }, + { url = "https://files.pythonhosted.org/packages/df/64/ff35eb9e54894cf471ff4716abd3c81eb0a0626869217ce3e6ba4ccf17d7/scipy-1.18.1-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:f55fa87b6c612ecd6b058f167c53231b1d14e412efe361d3d6e38b3631c73218", size = 35344199 }, + { url = "https://files.pythonhosted.org/packages/d3/af/c5538be1792f7034c12c7db6ee67cace58253c7b87b122d68253eaf5de89/scipy-1.18.1-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:c35d74ce0e193ff740c2f2be2ac913ddc232fe6c1ff40b26cfecb9c670c63314", size = 35639587 }, + { url = "https://files.pythonhosted.org/packages/91/4c/075e4f66471bac101141ac739e9e135549be1bae584571bd03a530c056e1/scipy-1.18.1-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:d2924a03db38dc2e848bca2fe9f077dafb891480b91a00a0963a8cf86dfc31c1", size = 37480330 }, + { url = "https://files.pythonhosted.org/packages/39/e7/979fd14e75008623df31ba70d6bb144700f68feadcea042021c06a05bf82/scipy-1.18.1-cp312-cp312-win_amd64.whl", hash = "sha256:5e4d44984abc0020154ea81b247adeddcc3ac5527b975ff798bd1ba0adc513c2", size = 36658278 }, + { url = "https://files.pythonhosted.org/packages/c7/0b/e1525354ff9d7d5feb6d1b31af6d14072e5c91e9607b421fa1ec889660b3/scipy-1.18.1-cp312-cp312-win_arm64.whl", hash = "sha256:d65d448389b8436493abcf629cc94ad0cf32aecaf06e1acca1de53cc795f2f12", size = 24400588 }, +] + +[[package]] +name = "sentencepiece" +version = "0.2.2" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/cc/33/ea3cb3839607eb175da835244a798f797f478c5ddf0e8ecdf57ea85a4c70/sentencepiece-0.2.2.tar.gz", hash = "sha256:3d2b5e824b5622038dc7b490897efe05ebbbb9e7350fc142f3ecc8789ef9bdf6", size = 8218435 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/20/31/f23a2efaa0210b883574001b88fa64e499f798f0848a0b610fb9b384d162/sentencepiece-0.2.2-cp311-cp311-macosx_10_9_universal2.whl", hash = "sha256:69e9dc8078e128286ed3b975e37c837ba96e215a50c3ef9f3f8b7ab9e5a832a0", size = 2184255 }, + { url = "https://files.pythonhosted.org/packages/96/f2/1ee0ccb772d71e822f625d6cb5f0ea825835e877f28a9ef299a1291df19e/sentencepiece-0.2.2-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:6dd76f3e5c8b2eb8a3a3efee787bbf5b9a66e52a048fe09cab85eca33fec6790", size = 1438545 }, + { url = "https://files.pythonhosted.org/packages/2a/92/3a6ea4a2c6dd9e7062698a5a33534ca0e20844883338ae9c6b9c122c1a9f/sentencepiece-0.2.2-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:443ac618c7a2a1377cf5c82581fbb849591d14e656d5e5a3e4682d4e36a34e4e", size = 1346997 }, + { url = "https://files.pythonhosted.org/packages/f3/3a/7839048997c7bc0c34c57526f539f835e20c7a57dc2a99f99579b11cdbef/sentencepiece-0.2.2-cp311-cp311-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:0e2aae42960392d6dcb9a72d8e1e65a97294c965071b43c7b3429a42f350250e", size = 1324282 }, + { url = "https://files.pythonhosted.org/packages/06/5f/9117bf854aef817ad0d0ee9310eed0308a7e529e7eaf2e80ad9cd281ef82/sentencepiece-0.2.2-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:1416b92f2f010333786fe6306ed2631121d5ea492219b0841e967b6765e64107", size = 1394242 }, + { url = "https://files.pythonhosted.org/packages/ab/62/9e2569867e3dcff7ad6d89642a9615b9801b5cd698abe7df3b490361f66e/sentencepiece-0.2.2-cp311-cp311-win_amd64.whl", hash = "sha256:70d4ca6f4d06df7f0ccab6fe4f49c8a712c8c8b6847b4f0af9a0e1dbb0e0337e", size = 1246268 }, + { url = "https://files.pythonhosted.org/packages/96/c9/5d781d4ef1124564a45c98b9ff25d531c10cdf568ec6314a2d1946f9251c/sentencepiece-0.2.2-cp311-cp311-win_arm64.whl", hash = "sha256:252908153eeec06c3ca3a32077e64a49d572e3d89881475b4e0f02d99d9fcc7c", size = 1190702 }, + { url = "https://files.pythonhosted.org/packages/b8/13/7a562289c8d5b49ebdf3f9c1e8ab67cf14a8743b1d90c8f406bfdec36b72/sentencepiece-0.2.2-cp312-cp312-macosx_10_13_universal2.whl", hash = "sha256:1edb10e520e4bddf74d85b0f5ae74cc2d60c2b448885080bfb618bc2b3a49f6b", size = 2188384 }, + { url = "https://files.pythonhosted.org/packages/85/d1/912f14fd5eae168aba726ffb6a9a2dc1c71fe7676c53da6f5c442b886d4a/sentencepiece-0.2.2-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:f7c06c751c19d923435a54bff4f7e66e728fad160e8da28254f133abc9725820", size = 1441553 }, + { url = "https://files.pythonhosted.org/packages/bd/44/caa9cab5f261a019e2808bc5046152775dc57352ba9cbae7525e9e7a1ed4/sentencepiece-0.2.2-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:38111ed1f79268f399c505028023d5eaaf0ab4e5eafceb709468b0d3323e7838", size = 1347176 }, + { url = "https://files.pythonhosted.org/packages/19/90/cd798935668cff71d309d8ff10385844ecf216b1fe454f1993ed8bf2cb91/sentencepiece-0.2.2-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:cbce24284f51f71d10a42b7b9c964dcb9048b28f1c8e5db40bcbcb6f428cba6a", size = 1325200 }, + { url = "https://files.pythonhosted.org/packages/b6/2d/37e3da037318a70066ded0d51bc2a7f35491ae6338dd993d5eb1503fc3b5/sentencepiece-0.2.2-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:c8a168b040bc61681293f79a949b5d911c8e25086f4260285b8d97ab5f1195da", size = 1397736 }, + { url = "https://files.pythonhosted.org/packages/8d/11/753fca2e6b109be3ab7867abf357dfe48677fe726ae5a5363d0b54ca9450/sentencepiece-0.2.2-cp312-cp312-win_amd64.whl", hash = "sha256:7c6e7bf684dc12145bfa685d3060beaea55139134ba848289bee514ed42e7383", size = 1248030 }, + { url = "https://files.pythonhosted.org/packages/e2/0a/70efbe861ca182d7d4b6e1a20f58e043400848fa9f2915229f082e221648/sentencepiece-0.2.2-cp312-cp312-win_arm64.whl", hash = "sha256:76ff5814db72e7462dece042d7593cdf102b8ec82c2b1cc201a2add34ee3050d", size = 1187325 }, +] + +[[package]] +name = "setuptools" +version = "84.0.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/6d/44/f5da03a8ef95d369145c5bb53050e7877c9f3d312e128605fd9504829143/setuptools-84.0.0.tar.gz", hash = "sha256:f4695c21257f0d9b537ec2692c941d02ee143b7cc1276941349a546573b2ef73", size = 1168449 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/95/9c/c510029fc6ef33a6275cd2c5d3cecd6613dfd6aa401d57c54f1c18852ccf/setuptools-84.0.0-py3-none-any.whl", hash = "sha256:51a52592b3b99e102b609654876bd65f19f999935166d1352678931132b0c670", size = 818216 }, +] + +[[package]] +name = "soundfile" +version = "0.14.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "cffi" }, + { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, + { name = "numpy", version = "2.5.2", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/d2/db/949331952a6fb1c5b12e9de80fd08747966c2039d1a61db4764fbd3981c2/soundfile-0.14.0.tar.gz", hash = "sha256:ba1c1a2d618bca5c406647c83b89f07cc8810fa506a50622a6993ba130c1de11", size = 47842 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/b1/d1/5e338af9ca6ed0786cd5bb03f6d60de1c325728c1189014f3b59aae7403c/soundfile-0.14.0-py2.py3-none-any.whl", hash = "sha256:8ba81ae3a89fd5ab3bef8a8eb481fbbe794e806309675a89b4df48b8d31908a8", size = 26799 }, + { url = "https://files.pythonhosted.org/packages/7e/72/c6b21e58d3113596e7e8de0a08d6f1d95173492cfbca0a4db14148cbba2a/soundfile-0.14.0-py2.py3-none-macosx_10_9_x86_64.whl", hash = "sha256:19be05428da76ed61a4cad29b8e4bcf43a3e5c100089d2ec81dc961eed1b0dd4", size = 1144568 }, + { url = "https://files.pythonhosted.org/packages/63/7a/dfdd6f8c748988427119f75eb860a3cedd858d1aea1fe28f39ad8559ef22/soundfile-0.14.0-py2.py3-none-macosx_11_0_arm64.whl", hash = "sha256:d828d35a059626da52f1415b5faee610aeab393319cb3fc4a9aef47b619fc14c", size = 1103726 }, + { url = "https://files.pythonhosted.org/packages/4a/f8/fc39fad6f879633461d27394cd1ddaf1f769ffa0597dca35872f51b16461/soundfile-0.14.0-py2.py3-none-manylinux_2_28_aarch64.whl", hash = "sha256:e85724a90bc99a6e8062c0b4ddf725f53b2a3b70afd4da875e9d2cfc4e92f377", size = 1238050 }, + { url = "https://files.pythonhosted.org/packages/7b/a2/70fd4432b924684c372df8b0a45708c36c057ef3596c9eb53e0a806b980b/soundfile-0.14.0-py2.py3-none-manylinux_2_28_x86_64.whl", hash = "sha256:1e38bac1853412871318e82a1ba69a8be677619b56025bbfcccdb41b6cafe82d", size = 1315963 }, + { url = "https://files.pythonhosted.org/packages/d9/34/c9e80783d83eab739a9531fdee03675d53e0bf1b2ccb4bb3af5844675046/soundfile-0.14.0-py2.py3-none-win32.whl", hash = "sha256:0a6ae43c50c71b4e020cc55382925cb89451c1ed1a0c3d0f5d802da269226849", size = 902199 }, + { url = "https://files.pythonhosted.org/packages/ed/97/b39c18ac1df45e755ca22b8b00e872929da5d107998a207a5e4ac831bfda/soundfile-0.14.0-py2.py3-none-win_amd64.whl", hash = "sha256:299491d3499460fb1b74bb4bd78b57ffc2d243a5fafa7b6ec1b264875c78453e", size = 1021480 }, + { url = "https://files.pythonhosted.org/packages/f4/83/55c65e61cf457805ce2ec157c1c6ae17715d0851aa2374422de0538838ca/soundfile-0.14.0-py2.py3-none-win_arm64.whl", hash = "sha256:e090704718e124e7c844695236f1fce8d18a5e761eaf7c82dfcd124620805f98", size = 888858 }, +] + +[[package]] +name = "speechbrain" +version = "1.1.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "huggingface-hub" }, + { name = "hyperpyyaml" }, + { name = "joblib" }, + { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, + { name = "numpy", version = "2.5.2", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "packaging" }, + { name = "requests" }, + { name = "scipy", version = "1.17.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, + { name = "scipy", version = "1.18.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "sentencepiece" }, + { name = "soundfile" }, + { name = "torch" }, + { name = "torchaudio" }, + { name = "tqdm" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/51/7f/4c136fb35f09d9bc93601c0b35a41980d2ab2af9033c1f711926b67635e1/speechbrain-1.1.1.tar.gz", hash = "sha256:3b69d9341662478b3c564837613f7a87bfcfd40d88cd54b4d19d2cca31e3dd3d", size = 1748087 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/08/bd/cb7befa5c2a6bc97b2d260a3481a7b2360c47db0f56ea2d1d2e2816bc3d5/speechbrain-1.1.1-py3-none-any.whl", hash = "sha256:de4f78d3564d40443e11e01648b00b4c47a6f942558c7ab08b5c350264ffcd7c", size = 2313069 }, +] + +[[package]] +name = "sympy" +version = "1.14.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "mpmath" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/83/d3/803453b36afefb7c2bb238361cd4ae6125a569b4db67cd9e79846ba2d68c/sympy-1.14.0.tar.gz", hash = "sha256:d3d3fe8df1e5a0b42f0e7bdf50541697dbe7d23746e894990c030e2b05e72517", size = 7793921 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/a2/09/77d55d46fd61b4a135c444fc97158ef34a095e5681d0a6c10b75bf356191/sympy-1.14.0-py3-none-any.whl", hash = "sha256:e091cc3e99d2141a0ba2847328f5479b05d94a6635cb96148ccb3f34671bd8f5", size = 6299353 }, +] + +[[package]] +name = "torch" +version = "2.13.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "cuda-bindings", marker = "sys_platform == 'linux'" }, + { name = "cuda-toolkit", extra = ["cublas", "cudart", "cufft", "cufile", "cupti", "curand", "cusolver", "cusparse", "nvjitlink", "nvrtc", "nvtx"], marker = "sys_platform == 'linux'" }, + { name = "filelock" }, + { name = "fsspec" }, + { name = "jinja2" }, + { name = "networkx" }, + { name = "nvidia-cudnn-cu13", marker = "sys_platform == 'linux'" }, + { name = "nvidia-cusparselt-cu13", marker = "sys_platform == 'linux'" }, + { name = "nvidia-nccl-cu13", marker = "sys_platform == 'linux'" }, + { name = "nvidia-nvshmem-cu13", marker = "sys_platform == 'linux'" }, + { name = "setuptools" }, + { name = "sympy" }, + { name = "triton", marker = "sys_platform == 'linux'" }, + { name = "typing-extensions" }, +] +wheels = [ + { url = "https://files.pythonhosted.org/packages/5b/fe/cba54dc58523434919b66f13a667e36e436deddd77ca519e96553617d4ec/torch-2.13.0-cp311-cp311-macosx_14_0_arm64.whl", hash = "sha256:e76f9bcecc52b8ff711239a2f7547d5353df95878ab232f0773c1d95928b92f8", size = 111187938 }, + { url = "https://files.pythonhosted.org/packages/c2/59/1e3160e18e12aa3038390efab3ce02b36a9d4d6a527ecdd8520dca2e68d8/torch-2.13.0-cp311-cp311-manylinux_2_28_aarch64.whl", hash = "sha256:092790c696a760c729fd5722835f50b9d81fd7c8f141571f3f3cf4081a8f664c", size = 427199369 }, + { url = "https://files.pythonhosted.org/packages/01/79/1f2d34ad7034ee1c7ffc1cf8bf0f8213af2a81df6ecdb3997ecec107c09d/torch-2.13.0-cp311-cp311-manylinux_2_28_x86_64.whl", hash = "sha256:60fcdcb2f3876e21146cb4524ef06397d727ca9ad5f020818547e25075fe3cb7", size = 526574961 }, + { url = "https://files.pythonhosted.org/packages/6c/fd/0f2ce40f58aefbdb3392f9acce3c8171940943ae2d661f70558bfa73befb/torch-2.13.0-cp311-cp311-win_amd64.whl", hash = "sha256:a0d8b11f16a48d60e2015d8213aa0390744cbebb98e58b62b3514dddc656e330", size = 122015870 }, + { url = "https://files.pythonhosted.org/packages/c4/3a/ed0f4d4d1dcde03bced7aac9a28e800abcdc0cbd06b6775044c9fbd877b7/torch-2.13.0-cp312-cp312-macosx_14_0_arm64.whl", hash = "sha256:2fe228aba290d14b9f31b049be550dbd469c3fd3013d7a19705b30454da97027", size = 111213045 }, + { url = "https://files.pythonhosted.org/packages/df/a9/f6a2a4d763ff1df02e9a64c477029db614295bc9367f4131223791ccc243/torch-2.13.0-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:572df8be8ffb4599c88cbd6a0726f1f854f4da65d2e3c09f0e2c2283333cd6d4", size = 427210998 }, + { url = "https://files.pythonhosted.org/packages/f3/82/fea946351658e6534db52d2cc12bc53087cbf87f9440c5f180f367c1950b/torch-2.13.0-cp312-cp312-manylinux_2_28_x86_64.whl", hash = "sha256:796633c4cdf0fe2cdced72d8f88f22e73dbcfce83132763162f6d4bff13b820b", size = 526605292 }, + { url = "https://files.pythonhosted.org/packages/21/d6/e8f3c6f7e01f626f77259de9860d2a78bc84c40539e28e79b7e98b0bb659/torch-2.13.0-cp312-cp312-win_amd64.whl", hash = "sha256:024c6cc0c1b085f2f91f20a3dc27b0471d021c31ce84b81be3afdc39f791fd9d", size = 122057313 }, +] + +[[package]] +name = "torchaudio" +version = "2.11.0" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/94/77/0eec7f175d88f312296bd5b11c23bd58da37c1021f53da3db4df449ce3ee/torchaudio-2.11.0-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:492dd64645e9d0bb843e94f1d9a4d1e31426262ffc594fafecc1697df9df5eb9", size = 684142 }, + { url = "https://files.pythonhosted.org/packages/b3/f9/6f7ebe071b44592c85269762b55b63ab0a091b5f479f73544738f7564a1e/torchaudio-2.11.0-cp311-cp311-manylinux_2_28_aarch64.whl", hash = "sha256:73dab4841f94d888bc7c2aed7b5547c643edc974306919fe1adfb65d57cccf4b", size = 1626527 }, + { url = "https://files.pythonhosted.org/packages/ac/70/17408e0d154d0c894537a88dcbadc48e8ad3b6e1ef4a1dabda5d40245ee0/torchaudio-2.11.0-cp311-cp311-manylinux_2_28_x86_64.whl", hash = "sha256:1a07ec72fd6f26a588c39b5f029e0130d16bb40bc4221635580bf8fb18fcbc80", size = 1771930 }, + { url = "https://files.pythonhosted.org/packages/c9/75/b6d03fc75b409bdaec597274d1bdd4213db716ed16f6801386b31d59c551/torchaudio-2.11.0-cp311-cp311-win_amd64.whl", hash = "sha256:bb59ba4452bbbe95d75ad3ef18df9824955625f36698ce9a5998a4a9f3c1ba1d", size = 328658 }, + { url = "https://files.pythonhosted.org/packages/f1/b1/77658817acacd01a72b714440c62f419efc4d90170e704e8e7a2c0918988/torchaudio-2.11.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:a1cf1acc883bee9cb906a933572fed6a8a933f86ef34e9ea7d803f72317e8c1b", size = 684226 }, + { url = "https://files.pythonhosted.org/packages/78/28/c7adc053039f286c2aca0038b766cbe3294e66fec6b29a820e95128f9ede/torchaudio-2.11.0-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:bc653defca1c16154398517a1adc98d0fb7f1dd08e58ced217558d213c2c6e29", size = 1626670 }, + { url = "https://files.pythonhosted.org/packages/88/d8/d6d0f896e064aa67377484efef4911cdcc07bce2929474e1417cc0af18c2/torchaudio-2.11.0-cp312-cp312-manylinux_2_28_x86_64.whl", hash = "sha256:6503c0bdb29daf2e6281bb70ea2dfe2c3553b782b619eb5d73bdadd8a3f7cecf", size = 1771992 }, + { url = "https://files.pythonhosted.org/packages/23/a8/941277ecc39f7a0a169d554302a1f1afd87c1d94a8aec828891916cea59a/torchaudio-2.11.0-cp312-cp312-win_amd64.whl", hash = "sha256:478110f981e5d40a8d82221732c57a56c85a1d5895fb8fe646e86ee15eded3bd", size = 328663 }, +] + +[[package]] +name = "tqdm" +version = "4.70.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "colorama", marker = "sys_platform == 'win32'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/21/3b/6c24bec5be5e743ffd99576daa5cc077722fc7d5bbc00bd133fa0c698dc6/tqdm-4.70.0.tar.gz", hash = "sha256:55b0b0dbd97462d06ebee91e4dac24ed4d4702be82b24f07e6c1d27e08cea220", size = 795438 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/f9/1c/01bfd571a64e7f270e6bab5e33777debe0edc56759233ce84f27dec92d14/tqdm-4.70.0-py3-none-any.whl", hash = "sha256:7f585706bfddbdebf89daac705b2dfcc16890130727d3197ca62c732b4310953", size = 80184 }, +] + +[[package]] +name = "transcribe-ecapa-tdnn-env" +version = "0.1.0" +source = { virtual = "." } +dependencies = [ + { name = "gguf" }, + { name = "huggingface-hub" }, + { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, + { name = "numpy", version = "2.5.2", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "pyarrow" }, + { name = "soundfile" }, + { name = "speechbrain" }, + { name = "torch" }, + { name = "torchaudio" }, +] + +[package.metadata] +requires-dist = [ + { name = "gguf", specifier = ">=0.10" }, + { name = "huggingface-hub", specifier = ">=0.30" }, + { name = "numpy", specifier = ">=1.26" }, + { name = "pyarrow", specifier = ">=15" }, + { name = "soundfile", specifier = ">=0.12" }, + { name = "speechbrain", specifier = "==1.1.1" }, + { name = "torch", specifier = "==2.13.0" }, + { name = "torchaudio", specifier = "==2.11.0" }, +] + +[[package]] +name = "triton" +version = "3.7.1" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/7b/f9/19d842d06a08559534fa1eaab6ca551b1bcf40f06620bddec1babaa2772d/triton-3.7.1-cp311-cp311-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:d4a0e1cd4c4a76370ed74a8432a53cea28716827d19e40ffc732233e35ceb3f6", size = 184664887 }, + { url = "https://files.pythonhosted.org/packages/cd/5e/fce69606f7f240297f163e25539906732b199530d486ce67ae319877e821/triton-3.7.1-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:6744957e9fd610a29680ec2346057d0c86948ed3812468670719f391e94b44a5", size = 197701306 }, + { url = "https://files.pythonhosted.org/packages/94/fa/f856e24deb462d5f18bd4b5a746957862ab9b6ee5834bda60605ec348366/triton-3.7.1-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:9497f2e696ee368862a181a90b2dcc03ca978cc4f602abd67c7d81022a6988e1", size = 184692359 }, + { url = "https://files.pythonhosted.org/packages/c4/6f/fb96d15db6f36d6eae4cafb998c2e0353bf59d7c4ea1662d7497f269134a/triton-3.7.1-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:7e40869937a68206ec70d7f25bb7ec6433cb083f9135e1f36dbd318dc449a728", size = 197719725 }, +] + +[[package]] +name = "typing-extensions" +version = "4.16.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/f6/cc/6253133b5bb138fc3306cebfbda2c520f545d36b5be2c7255cc528bb45d6/typing_extensions-4.16.0.tar.gz", hash = "sha256:dc983d19a509c94dba722ee6abd33940f7c05a89e243c47e907eb4db6f1a43e5", size = 113555 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/49/d3/b8441a820a491ddfc024b0b0cf0393375b75ea13866d9c66727e54c2fc80/typing_extensions-4.16.0-py3-none-any.whl", hash = "sha256:481caa481374e813c1b176ada14e97f1f67a4539ce9cfeb3f350d78d6370c2e8", size = 45571 }, +] + +[[package]] +name = "urllib3" +version = "2.7.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/53/0c/06f8b233b8fd13b9e5ee11424ef85419ba0d8ba0b3138bf360be2ff56953/urllib3-2.7.0.tar.gz", hash = "sha256:231e0ec3b63ceb14667c67be60f2f2c40a518cb38b03af60abc813da26505f4c", size = 433602 } +wheels = [ + { url = "https://files.pythonhosted.org/packages/7f/3e/5db95bcf282c52709639744ca2a8b149baccf648e39c8cc87553df9eae0c/urllib3-2.7.0-py3-none-any.whl", hash = "sha256:9fb4c81ebbb1ce9531cce37674bbc6f1360472bc18ca9a553ede278ef7276897", size = 131087 }, +] diff --git a/scripts/hf_cards/generate.py b/scripts/hf_cards/generate.py index 1741ac06e..af35949d0 100755 --- a/scripts/hf_cards/generate.py +++ b/scripts/hf_cards/generate.py @@ -134,7 +134,7 @@ def derive_metric_blocks(record: dict) -> dict[str, dict[str, float]]: chosen[cell] = row blocks: dict[str, dict[str, float]] = {} for (key, quant), row in chosen.items(): - blocks.setdefault(key, {})[quant] = row["err_pct"] + blocks.setdefault(key, {})[quant] = common.row_pct(row) return blocks @@ -233,7 +233,7 @@ def build_context(record: dict, spec: dict) -> dict: "transcribe_docs_url": docs_url(record), } if headline.get("metric"): - ctx["metric"] = headline["metric"].upper() + ctx["metric"] = common.metric_label(headline["metric"]) for key in ("name", "link"): if record["license"].get(key): ctx[f"license_{key}"] = record["license"][key] diff --git a/scripts/hf_cards/lang-id-voxlingua107-ecapa.yaml b/scripts/hf_cards/lang-id-voxlingua107-ecapa.yaml new file mode 100644 index 000000000..73ef4855b --- /dev/null +++ b/scripts/hf_cards/lang-id-voxlingua107-ecapa.yaml @@ -0,0 +1,67 @@ +# Spec for the HF README of handy-computer/lang-id-voxlingua107-ecapa-gguf. +# Prose only; numbers and metadata come from catalog/.json. See README.md. +# +# Language ID family (LANGID role): metric is top-1 accuracy (not WER). + +pin_date: 2026-10-05 + +# Validation pin for the most recent upload. Updated on each release — +# older HF revisions carry whatever value was current at their upload time. +validation: + reference: SpeechBrain + commit: aa4d0f47 + date: 2026-10-08 + +pipeline_tag: audio-classification +tags: + - gguf + - transcribe.cpp + - language-identification + - spoken-language-identification + - ecapa-tdnn + - voxlingua107 + - speechbrain + +summary: | + Spoken language identification over 107 languages: SpeechBrain's + ECAPA-TDNN trained on VoxLingua107. NOT a transcription model: a run + returns the model's language labels ranked by probability, optionally + restricted to a caller-chosen set. Takes 16 kHz mono WAV; scores up to + the last 30 s of a clip. + +usage: | + Build transcribe.cpp from source: + + ```bash + git clone git@github.com:handy-computer/transcribe.cpp.git + cd transcribe.cpp + cmake -B build && cmake --build build + ``` + + Run on a 16 kHz mono WAV. This is a language identifier, not a + transcription model: the CLI prints ranked language candidates. + + ```bash + build/bin/transcribe-cli -m lang-id-voxlingua107-ecapa-Q8_0.gguf --allow en,de,fr --top 3 input.wav + # language: de index=18 p=0.999998 + # ... + ``` + + From the C API, use the LANGID role (`include/transcribe/langid.h`, see + `docs/langid.md`). + + If your audio isn't already 16 kHz mono WAV, convert it first: + + ```bash + ffmpeg -i input.mp3 -ar 16000 -ac 1 output.wav + ``` + + See the [transcribe.cpp model page](https://github.com/handy-computer/transcribe.cpp/blob/main/docs/models/lang-id-voxlingua107-ecapa.md) for performance + numbers, numerical validation, and reproduction steps. + +wer: + notes: | + Open-set top-1 accuracy over all 107 labels, the mean over 15 FLEURS + languages, on the first 5 s of each clip without silence trimming. + Agreement with the SpeechBrain reference, per GGUF, is on the + transcribe.cpp model page. diff --git a/scripts/langid/bench.py b/scripts/langid/bench.py new file mode 100644 index 000000000..1f59191a1 --- /dev/null +++ b/scripts/langid/bench.py @@ -0,0 +1,153 @@ +#!/usr/bin/env python3 +""" +bench.py - this machine's language ID publication speed cells +(catalog/_benchmark_profiles.json) through the Python binding: +tools/transcribe-bench is ASR-only. + + uv run --project scripts/envs/ecapa_tdnn scripts/langid/bench.py --profile \\ + --library build-shared/src/libtranscribe.dylib + uv run scripts/catalog/ingest_perf.py + +The profile names the GGUFs, backends, samples, iterations, warmup and +thermal gate. A sample `-s` is the first N seconds of +samples/.wav; the cell's time is the native mel + encode. Writes one +bench-driver report (scripts/bench/run.py's shape) per backend to +reports/perf//-publication__.json. +""" + +from __future__ import annotations + +import argparse +import importlib.util +import json +import os +import re +import statistics +import sys +import time +from pathlib import Path + +import numpy as np + +REPO_ROOT = Path(__file__).resolve().parents[2] +SAMPLE_RE = re.compile(r"^(?P.+)-(?P\d+(?:\.\d+)?)s$") + + +def bench_driver(): + """scripts/bench/run.py, for its profiles, machine detection and thermal + gate. Loaded by path: this directory has a run.py of its own.""" + spec = importlib.util.spec_from_file_location( + "bench_driver", REPO_ROOT / "scripts" / "bench" / "run.py") + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module # its dataclasses look their module up + spec.loader.exec_module(module) + return module + + +def main(argv: list[str] | None = None) -> int: + p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + p.add_argument("--profile", nargs="?", const="", default="", + help="profile id (default: the variant's role profile)") + p.add_argument("--variant", default="lang-id-voxlingua107-ecapa") + p.add_argument("--library", type=Path, help="a shared libtranscribe (TRANSCRIBE_LIBRARY)") + args = p.parse_args(argv) + + driver = bench_driver() + profiles = driver.benchmark_profiles + machine = driver.detect_machine() + record = driver.catalog_common.load_record(args.variant) + profile_id, profile = profiles.profile_for(record, args.profile or None) + target = profiles.target_for_machine(profile, machine["slug"]) + if target is None: + p.error(f"machine {machine['slug']!r} is not a target in profile {profile_id}") + cells = [cell for cell in profiles.apply_exceptions( + record, "speed", profiles.expected_speed(record, profile)) + if cell["machine"] == target["machine"]] + quants = {cell["quant"] for cell in cells} + ggufs = [REPO_ROOT / "models" / args.variant / item["filename"] + for item in record["downloads"] if item["quant"] in quants] + backends = list(dict.fromkeys(cell["backend"] for cell in cells)) + samples = [SAMPLE_RE.match(name) for name in dict.fromkeys(cell["sample"] for cell in cells)] + if not all(samples) or len({m["clip"] for m in samples}) != 1: + p.error(f"profile samples must be -s on one clip: " + f"{sorted({cell['sample'] for cell in cells})}") + clip = samples[0]["clip"] + durations = [m["seconds"] for m in samples] + warmup, repeat = int(profile["speed"]["warmup"]), int(profile["speed"]["iterations"]) + cooldown_c = float(target.get("cooldown_tctl_c", 0.0)) + + if args.library is not None: + os.environ["TRANSCRIBE_LIBRARY"] = str(args.library.resolve()) + sys.path.insert(0, str(REPO_ROOT / "bindings" / "python" / "src")) + import soundfile as sf + import transcribe_cpp as t + + sample_path = REPO_ROOT / "samples" / f"{clip}.wav" + pcm, sr = sf.read(str(sample_path), dtype="float32") + assert sr == 16000 and pcm.size >= float(max(durations, key=float)) * 16000 + + timestamp = driver.now_utc_iso() + commit = t.native_commit() + name = f"{args.variant}-publication" + for backend in backends: + runs = [] + for gguf in ggufs: + t0 = time.perf_counter() + model = t.Model(str(gguf), backend=backend) + load_ms = (time.perf_counter() - t0) * 1000 + lid = model.langid_session(n_threads=0, max_audio_ms=60000) + for seconds in durations: + d = float(seconds) + driver.cooldown_wait(cooldown_c, 300.0, 10.0) + audio = np.ascontiguousarray(pcm[: int(d * 16000)]) + for _ in range(warmup): + lid.run(audio) + per_iter = [] + for _ in range(repeat): + w0 = time.perf_counter() + lid.run(audio) + wall = (time.perf_counter() - w0) * 1000 + tm = lid.timings + per_iter.append({"mel_ms": tm.mel_ms, "encode_ms": tm.encode_ms, + "total_ms": tm.mel_ms + tm.encode_ms, "wall_ms": wall}) + runs.append({ + "model_path": str(gguf.relative_to(REPO_ROOT)), + "sample": f"{clip}-{seconds}s", + "sample_path": str(sample_path.relative_to(REPO_ROOT)), + "sample_duration_s": d, + "backend": model.backend, + "threads": 0, + "load_ms": round(load_ms, 1), + "per_iter": per_iter, + "summary": {field: {"mean": statistics.fmean(it[field] for it in per_iter)} + for field in ("mel_ms", "encode_ms", "total_ms", "wall_ms")}, + }) + print(f"{gguf.name:42} {backend:6} {d:5.0f}s mean " + f"{runs[-1]['summary']['total_ms']['mean']:8.2f} ms", flush=True) + lid.close() + model.close() + + out = REPO_ROOT / "reports" / "perf" / machine["slug"] / \ + f"{driver.slugify(name)}_{args.variant}_{backend}.json" + out.parent.mkdir(parents=True, exist_ok=True) + out.write_text(json.dumps({ + "schema": "transcribe-bench-driver-v1", + "timestamp": timestamp, + "name": name, + "publication_profile": profile_id, + "machine": machine, + # The build the native library reports, which is what ran. + "git_sha": commit if commit != "unknown" else driver.get_git_sha(REPO_ROOT), + "variant": args.variant, + "backend": backend, + "iters": repeat, + "warmup": warmup, + "tool": "scripts/langid/bench.py", + "runs": runs, + }, indent=2) + "\n") + print(f"wrote {out}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/scripts/langid/ingest.py b/scripts/langid/ingest.py new file mode 100755 index 000000000..088ebf27b --- /dev/null +++ b/scripts/langid/ingest.py @@ -0,0 +1,119 @@ +#!/usr/bin/env python3 +""" +ingest.py - build the language ID corpus from the FLEURS test parquet at the +pinned revision FLEURS_REVISION (read through the Hugging Face hub cache). + + uv run --project scripts/envs/ecapa_tdnn scripts/langid/ingest.py fleurs --lang all + +Output (gitignored): + samples/langid/fleurs-/.wav 16-bit PCM mono 16 kHz + samples/langid/fleurs-.manifest.jsonl {"id","audio","language","duration_s",...} + +The languages are the langid publication profile's `pooled_languages`. Per +language: the first N utterances in parquet order that last at least 1 s. +Ids are `fleurs--` by selection index (FLEURS sentence ids repeat +across speakers); the sentence id is kept as `fleurs_id`. An existing +manifest is left alone unless --force. +""" + +from __future__ import annotations + +import argparse +import io +import json +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(REPO_ROOT / "scripts" / "wer")) +from languages import FLEURS_LANGS # noqa: E402 + +SAMPLE_RATE = 16000 +MIN_DURATION_S = 1.0 +SPLIT = "test" +DATASET = "google/fleurs" +FLEURS_REVISION = "70bb2e84b976b7e960aa89f1c648e09c59f894dd" +LICENCE = "CC-BY-4.0" + + +def profile_languages() -> list[str]: + data = json.loads((REPO_ROOT / "catalog" / "_benchmark_profiles.json").read_text()) + return data["profiles"][data["roles"]["langid"]]["accuracy"][0]["pooled_languages"] + + +def ingest_language(code: str, n: int, force: bool) -> None: + import pyarrow.parquet as pq + import soundfile as sf + from huggingface_hub import hf_hub_download + + config = FLEURS_LANGS[code] + out_dir = REPO_ROOT / "samples" / "langid" / f"fleurs-{code}" + manifest_path = out_dir.with_name(f"fleurs-{code}.manifest.jsonl") + if manifest_path.exists() and not force: + print(f"{code}: {manifest_path.name} exists; skipping (--force to rebuild)") + return + + parquet = hf_hub_download( + repo_id=DATASET, repo_type="dataset", revision=FLEURS_REVISION, + filename=f"parquet-data/{config}/{SPLIT}-00000-of-00001.parquet") + out_dir.mkdir(parents=True, exist_ok=True) + entries: list[dict] = [] + rows = (row for batch in pq.ParquetFile(parquet).iter_batches(batch_size=16) + for row in batch.to_pylist()) + for row_index, row in enumerate(rows): + if len(entries) >= n: + break + audio = row["audio"] + pcm, sr = sf.read(io.BytesIO(audio["bytes"] if isinstance(audio, dict) else audio), + dtype="float32", always_2d=False) + if sr != SAMPLE_RATE: + raise SystemExit(f"error: {code} row {row_index} is {sr} Hz; FLEURS is 16 kHz") + if pcm.ndim > 1: + pcm = pcm.mean(axis=1) + if pcm.size < MIN_DURATION_S * SAMPLE_RATE: + continue + utt_id = f"fleurs-{code}-{len(entries):04d}" + wav_path = out_dir / f"{utt_id}.wav" + sf.write(str(wav_path), pcm, SAMPLE_RATE, subtype="PCM_16") + entries.append({ + "id": utt_id, + "audio": str(wav_path.relative_to(REPO_ROOT)), + "language": code, + "duration_s": round(pcm.size / SAMPLE_RATE, 3), + "fleurs_id": int(row["id"]), + "config": config, + "split": SPLIT, + "row_index": row_index, + "dataset": DATASET, + "licence": LICENCE, + }) + if len(entries) < n: + print(f"warning: {code}: only {len(entries)} of {n} utterances", file=sys.stderr) + manifest_path.write_text("".join(json.dumps(e) + "\n" for e in entries)) + print(f"{code}: {len(entries)} utterances -> {manifest_path.relative_to(REPO_ROOT)}") + + +def main(argv: list[str] | None = None) -> int: + p = argparse.ArgumentParser( + description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + sub = p.add_subparsers(dest="source", required=True) + fp = sub.add_parser("fleurs", help="FLEURS parquet at the pinned revision") + fp.add_argument("--lang", required=True, + help="comma-separated VoxLingua107 codes from the profile, or 'all'") + fp.add_argument("--n", type=int, default=200, help="utterances per language") + fp.add_argument("--force", action="store_true", + help="rebuild even when the manifest already exists") + args = p.parse_args(argv) + + known = profile_languages() + codes = known if args.lang == "all" else args.lang.split(",") + unknown = sorted(set(codes) - set(known)) + if unknown: + p.error(f"unknown language(s) {unknown}; known: {', '.join(known)} or 'all'") + for code in codes: + ingest_language(code, args.n, args.force) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/scripts/langid/run.py b/scripts/langid/run.py new file mode 100755 index 000000000..5d39929e6 --- /dev/null +++ b/scripts/langid/run.py @@ -0,0 +1,212 @@ +#!/usr/bin/env python3 +""" +run.py - score language ID manifests with the SpeechBrain reference or +transcribe.cpp (through the Python binding) into one JSONL sweep: a header +line, then one row per (utterance, crop) with all 107 pre-softmax logits. + + uv run --project scripts/envs/ecapa_tdnn scripts/langid/run.py \\ + --engine speechbrain --model speechbrain/lang-id-voxlingua107-ecapa \\ + --manifest samples/langid/fleurs-*.manifest.jsonl \\ + --out reports/langid/ref-speechbrain-untrimmed.jsonl + + uv run --project scripts/envs/ecapa_tdnn scripts/langid/run.py \\ + --engine cpp --model models/lang-id-voxlingua107-ecapa/lang-id-voxlingua107-ecapa-F32.gguf \\ + --library build-shared/src/libtranscribe.dylib \\ + --manifest samples/langid/fleurs-*.manifest.jsonl \\ + --out reports/langid/cpp-f32-untrimmed.jsonl + +A crop of N is the first N seconds of the clip, `full` the whole clip; a +row's `audio_s` is what was scored. Both engines get the same float samples +(the clips are 16-bit, so a crop is exact). SpeechBrain runs a batch of one: +padding a batch changes its normalisation span. +""" + +from __future__ import annotations + +import argparse +import datetime as _dt +import hashlib +import json +import math +import os +import sys +import time +from pathlib import Path + +import numpy as np + +REPO_ROOT = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(REPO_ROOT / "scripts")) + +SAMPLE_RATE = 16000 + + +class SpeechBrainEngine: + """The pinned checkpoint; logits from a forward hook on the pre-softmax + `mods.classifier.out` (classify_batch only returns log-softmax).""" + + def __init__(self, model: str, threads: int) -> None: + import speechbrain + import torch + + import dump_reference_ecapa_tdnn_speechbrain as dumper + + self.torch = torch + args = argparse.Namespace(model=model, revision=dumper.DEFAULT_REVISION, + device="cpu", torch_threads=threads) + dumper.configure_torch(args) + self.clf, checkpoint_dir = dumper.load_reference(args) + ind2lab = self.clf.hparams.label_encoder.ind2lab + self.labels = [dumper.split_label(str(ind2lab[i]))[0] for i in range(len(ind2lab))] + self._captured: list = [] + self.clf.mods.classifier.out.register_forward_hook( + lambda _m, _i, out: self._captured.append(out)) + self.recipe = { + "revision": dumper.DEFAULT_REVISION, + "checkpoint_dir": str(checkpoint_dir), + "device": "cpu", + "model_dtype": "f32", + "batch_size": 1, + "torch_threads": threads, + "speechbrain": speechbrain.__version__, + "torch": torch.__version__, + } + + def logits(self, pcm: np.ndarray) -> np.ndarray: + torch = self.torch + self._captured.clear() + # wav_lens is relative to the padded batch: 1.0 for one unpadded clip. + with torch.inference_mode(): + self.clf.classify_batch(torch.from_numpy(pcm).unsqueeze(0), torch.ones(1)) + (out,) = self._captured + return out.detach().to(torch.float32).reshape(len(self.labels)).numpy() + + +class CppEngine: + """The GGUF loaded once; each clip is one LangIdSession.run with top_k 0, + so every label's logit comes back.""" + + def __init__(self, model: Path, backend: str, threads: int, max_audio_ms: int, + library: Path | None) -> None: + if library is not None: + os.environ["TRANSCRIBE_LIBRARY"] = str(library.resolve()) + sys.path.insert(0, str(REPO_ROOT / "bindings" / "python" / "src")) + import transcribe_cpp + from gguf import GGUFReader + + self.model = transcribe_cpp.Model(str(model), backend=backend) + self.session = self.model.langid_session(n_threads=threads, max_audio_ms=max_audio_ms) + self.labels = [code for code, _name in self.model.langid_labels] + fields = GGUFReader(str(model)).fields + kv = lambda key: str(fields[key].contents()) if key in fields else None # noqa: E731 + with open(model, "rb") as f: + sha = hashlib.file_digest(f, "sha256").hexdigest() + self.recipe = { + "backend": backend, + "threads": threads, + "max_audio_ms": max_audio_ms, + "gguf_sha256": sha, + "gguf_source_commit": kv("general.source.commit"), + "gguf_source_repo": kv("general.name"), + "native_version": transcribe_cpp.native_version(), + "native_commit": transcribe_cpp.native_commit(), + "backend_bound": self.model.backend, + } + + def logits(self, pcm: np.ndarray) -> np.ndarray: + logits = np.zeros(len(self.labels), dtype=np.float32) + for c in self.session.run(pcm, top_k=0).candidates: + logits[c.index] = c.logit + return logits + + +def read_audio(path: Path) -> np.ndarray: + import soundfile as sf + + pcm, sr = sf.read(str(path), dtype="float32", always_2d=False) + if sr != SAMPLE_RATE or pcm.ndim != 1: + raise SystemExit(f"error: {path} is not 16 kHz mono") + return pcm + + +def main(argv: list[str] | None = None) -> int: + p = argparse.ArgumentParser( + description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + p.add_argument("--engine", required=True, choices=("speechbrain", "cpp")) + p.add_argument("--model", required=True, + help="HF repo id / checkpoint dir (speechbrain) or GGUF path (cpp)") + p.add_argument("--manifest", required=True, nargs="+", type=Path) + p.add_argument("--crops", default="3,5,10,full") + p.add_argument("--out", required=True, type=Path) + p.add_argument("--library", type=Path, + help="cpp only: a shared libtranscribe (TRANSCRIBE_LIBRARY)") + p.add_argument("--backend", default="cpu", help="cpp only") + p.add_argument("--threads", type=int, default=4) + p.add_argument("--max-utts", type=int, default=0, + help="cap utterances per manifest (0 = all); for smoke tests") + args = p.parse_args(argv) + + crops = [c if c == "full" else int(c) for c in args.crops.split(",")] + rows = [] + for m in args.manifest: + lines = [json.loads(l) for l in m.read_text().splitlines() if l.strip()] + rows += lines[: args.max_utts or None] + longest_s = max(r["duration_s"] for r in rows) + + if args.engine == "speechbrain": + engine = SpeechBrainEngine(args.model, args.threads) + else: + # Room for the longest clip, so `full` scores the same audio as SpeechBrain. + max_audio_ms = int(math.ceil(longest_s / 10.0) * 10.0 * 1000) + engine = CppEngine(Path(args.model), args.backend, args.threads, max_audio_ms, + args.library) + + header = { + "type": "header", + "engine": args.engine, + "model": args.model, + "recipe": { + "crops": [str(c) for c in crops], + "manifests": [str(m) for m in args.manifest], + "manifest_sha256": {str(m): hashlib.sha256(m.read_bytes()).hexdigest() + for m in args.manifest}, + "n_utterances": len(rows), + "sample_rate": SAMPLE_RATE, + "logit_kind": "pre_softmax", + "longest_clip_s": longest_s, + **engine.recipe, + }, + "labels": engine.labels, + "created": _dt.datetime.now(_dt.timezone.utc).replace(microsecond=0).isoformat(), + } + + args.out.parent.mkdir(parents=True, exist_ok=True) + t0 = time.time() + with open(args.out, "w") as f: + f.write(json.dumps(header) + "\n") + for crop in crops: + for i, row in enumerate(rows): + pcm = read_audio(REPO_ROOT / row["audio"]) + if crop != "full": + pcm = pcm[: crop * SAMPLE_RATE] + logits = engine.logits(pcm) + z = np.exp(logits.astype(np.float64) - logits.max()) + top = int(np.argmax(z)) + f.write(json.dumps({ + "id": row["id"], + "language": row["language"], + "crop_s": crop, + "audio_s": round(pcm.size / SAMPLE_RATE, 4), + "logits": [round(float(x), 6) for x in logits], + "top1": engine.labels[top], + "top1_prob": round(float(z[top] / z.sum()), 6), + }) + "\n") + if (i + 1) % 200 == 0 or i + 1 == len(rows): + print(f" crop {crop}: {i + 1}/{len(rows)} {time.time() - t0:.0f}s", + file=sys.stderr) + print(f"wrote {args.out} in {time.time() - t0:.0f}s", file=sys.stderr) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/scripts/langid/score.py b/scripts/langid/score.py new file mode 100755 index 000000000..9a0e29ee0 --- /dev/null +++ b/scripts/langid/score.py @@ -0,0 +1,193 @@ +#!/usr/bin/env -S uv run --script +# /// script +# requires-python = ">=3.11" +# dependencies = [ +# "numpy>=1.26", +# ] +# /// +""" +score.py - open-set top-1 accuracy of a scripts/langid/run.py sweep and, with +--ref, its top-1 agreement with the reference sweep (the dataset gate). + + uv run scripts/langid/score.py reports/langid/cpp-f32-untrimmed.jsonl \\ + --ref reports/langid/ref-speechbrain-untrimmed.jsonl \\ + --json reports/langid/lang-id-voxlingua107-ecapa-F32.fleurs-mul.score.json + +Accuracy is the macro mean over languages of strict code equality (`nn` is +not credited for `no`), per crop, with a 95% bootstrap CI that resamples +utterances within each language. + +--json writes the score; with --ref it also writes the agreement next to it +(`.agreement.json`); scripts/catalog/ingest_accuracy.py reads both. + +Agreement joins rows on (id, crop_s). Exit 1 on any top-1 disagreement or +if the sweeps do not cover the same rows; exit 2 if they differ in labels, +recipe or checkpoint revision. Only F32 is gated: --report-only (the quants) +exits 1 on a coverage mismatch only. +""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path + +import numpy as np + +REPO_ROOT = Path(__file__).resolve().parents[2] +N_BOOT = 1000 +BOOT_SEED = 42 +CI = 0.95 +RECIPE_FIELDS = ("crops", "n_utterances", "logit_kind", "sample_rate", "manifest_sha256") + + +def die(msg: str) -> None: + print(f"error: {msg}", file=sys.stderr) + raise SystemExit(2) + + +def rel(path: Path) -> str: + try: + return str(path.resolve().relative_to(REPO_ROOT)) + except ValueError: + return str(path) + + +def load(path: Path) -> tuple[dict, dict[tuple, dict]]: + lines = [json.loads(l) for l in path.read_text().splitlines() if l.strip()] + if not lines or lines[0].get("type") != "header" or len(lines) < 2: + die(f"{path} is not a run.py sweep (header line, then rows)") + rows = {(r["id"], str(r["crop_s"])): r for r in lines[1:]} + if len(rows) != len(lines) - 1: + die(f"{path} has duplicate (id, crop_s) rows") + return lines[0], rows + + +def bootstrap_ci(blocks: list[np.ndarray]) -> tuple[float, float]: + rng = np.random.default_rng(BOOT_SEED) + means = np.empty(N_BOOT, dtype=np.float64) + for b in range(N_BOOT): + acc = 0.0 + for blk in blocks: + acc += blk[rng.integers(0, blk.size, blk.size)].mean() + means[b] = acc / len(blocks) + means.sort() + return (float(means[int((1 - CI) / 2 * N_BOOT)]), + float(means[min(N_BOOT - 1, int((1 + CI) / 2 * N_BOOT))])) + + +def accuracy(header: dict, rows: list[dict]) -> tuple[list[str], dict]: + """Per crop: n, macro-mean accuracy and its CI, in percent.""" + labels = np.array(header["labels"]) + languages = list(dict.fromkeys(r["language"] for r in rows)) + crops = {} + for crop in dict.fromkeys(str(r["crop_s"]) for r in rows): + sel = [r for r in rows if str(r["crop_s"]) == crop] + lang = np.array([r["language"] for r in sel]) + logits = np.array([r["logits"] for r in sel], dtype=np.float64) + correct = (labels[np.argmax(logits, axis=1)] == lang).astype(np.float64) + blocks = [correct[lang == lg] for lg in languages if (lang == lg).any()] + mean = sum(float(blk.mean()) for blk in blocks) / len(blocks) + lo, hi = bootstrap_ci(blocks) + crops[crop] = {"n": len(sel), "acc_pct": round(100.0 * mean, 2), + "ci95": [round(100.0 * lo, 2), round(100.0 * hi, 2)]} + return languages, crops + + +def agreement(ref: Path, header: dict, rows: dict, run: Path) -> tuple[dict, list]: + """The agreement JSON, and the first few disagreeing keys.""" + rh, rrows = load(ref) + if rh["labels"] != header["labels"]: + die("the sweeps disagree on the label set") + differs = [f for f in RECIPE_FIELDS if rh["recipe"].get(f) != header["recipe"].get(f)] + if differs: + die(f"the sweeps' recipes differ in {differs}") + revisions = {h["recipe"].get("revision") or h["recipe"].get("gguf_source_commit") + for h in (rh, header)} + if len(revisions) != 1 or None in revisions: + die(f"the sweeps do not trace back to one checkpoint revision: {revisions}") + keys = sorted(set(rrows) & set(rows)) + if not keys: + die("no shared (id, crop_s) rows") + a = np.array([rrows[k]["logits"] for k in keys], dtype=np.float64) + b = np.array([rows[k]["logits"] for k in keys], dtype=np.float64) + if a.shape != b.shape: + die(f"logit shapes differ: {a.shape} vs {b.shape}") + agree = np.argmax(a, axis=1) == np.argmax(b, axis=1) + return { + "schema": "transcribe-langid-agreement-v1", + "reference_run": rel(ref), + "run": rel(run), + "n": len(keys), + "n_agree": int(agree.sum()), + "max_abs_logit_delta": float(np.abs(a - b).max()), + "disagreements": int((~agree).sum()), + "same_rows": set(rrows) == set(rows), + }, [keys[i] for i in np.flatnonzero(~agree)[:10]] + + +def main(argv: list[str] | None = None) -> int: + p = argparse.ArgumentParser( + description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + p.add_argument("run", type=Path, help="run.py JSONL sweep") + p.add_argument("--ref", type=Path, help="reference sweep to measure agreement against") + p.add_argument("--report-only", action="store_true", + help="record agreement without gating on it (the quants)") + p.add_argument("--json", type=Path, + help="write the score here: .fleurs-mul.score.json " + "under reports/langid/") + args = p.parse_args(argv) + if args.ref and args.json and not args.json.name.endswith(".score.json"): + p.error("--json must end in .score.json (the agreement is written beside it)") + + header, rows = load(args.run) + recipe = header["recipe"] + languages, crops = accuracy(header, list(rows.values())) + for crop, c in crops.items(): + print(f"crop {crop:>4}: {c['acc_pct']:6.2f}% " + f"[{c['ci95'][0]:.2f}, {c['ci95'][1]:.2f}] n={c['n']}") + + status = 0 + agr = None + if args.ref: + agr, disagreeing = agreement(args.ref, header, rows, args.run) + print(f"top-1 agreement with {agr['reference_run']}: {agr['n_agree']}/{agr['n']}, " + f"max |delta logit| {agr['max_abs_logit_delta']:.6g}") + if not agr["same_rows"]: + print("FAIL: the sweeps do not cover the same rows") + status = 1 + elif agr["disagreements"] and not args.report_only: + print(f"FAIL: {agr['disagreements']} disagreements, e.g. {disagreeing}") + status = 1 + + if args.json: + names = [Path(m).name for m in recipe.get("manifests", [])] + native = recipe.get("native_commit") + args.json.parent.mkdir(parents=True, exist_ok=True) + args.json.write_text(json.dumps({ + "schema": "transcribe-langid-score-v1", + "run": rel(args.run), + "engine": header["engine"], + "model": header["model"], + # The build the native library reports, not the checkout scoring it. + "engine_sha": native if native != "unknown" else None, + "backend": recipe.get("backend"), + "dataset": "fleurs" if names and all(n.startswith("fleurs-") for n in names) else None, + "split": "test", + "language": "mul", + "languages": languages, + "created": header.get("created"), + "bootstrap": {"n": N_BOOT, "seed": BOOT_SEED}, + "crops": crops, + }, indent=2) + "\n") + print(f"wrote {args.json}") + if agr: + out = args.json.with_name(args.json.name.replace(".score.json", ".agreement.json")) + out.write_text(json.dumps(agr, indent=2) + "\n") + print(f"wrote {out}") + return status + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/scripts/lib/gguf_common.py b/scripts/lib/gguf_common.py index 1712935ae..f426760bd 100644 --- a/scripts/lib/gguf_common.py +++ b/scripts/lib/gguf_common.py @@ -267,12 +267,16 @@ def encode_for_gguf( "all_features": "global", "global": "global", "per_utterance": "per_utterance", + # SpeechBrain InputNormalization(norm_type="sentence", std_norm=False): + # per-bin mean subtraction only (ecapa_tdnn). + "sentence_mean": "sentence_mean", } def canonicalize_normalize(raw) -> str: """Map a reference-framework normalize value to the canonical enum - in our intake schema (per_feature / global / per_utterance / none). + in our intake schema (per_feature / global / per_utterance / + sentence_mean / none). Unknown values raise — Stage 3 should fail loudly rather than emit a value the C++ loader will not recognise.""" key = raw if raw is None else str(raw) diff --git a/scripts/lib/quant_policy.py b/scripts/lib/quant_policy.py index 472a3eef0..43561c003 100644 --- a/scripts/lib/quant_policy.py +++ b/scripts/lib/quant_policy.py @@ -48,6 +48,11 @@ # either). K tiers also save little (Q8_0 139MB -> Q4_K_M 92MB) and # are slower than Q8_0 on CPU. Ship only the near-reference tiers. "sortformer": ("F16", "Q8_0"), + # ecapa_tdnn (language ID): a 21M-parameter model whose decision is an + # argmax over 107 logits. Q8_0 ships for its download size (the loader + # widens it to F16, see src/arch/ecapa_tdnn/model.cpp); the K tiers would + # save a few MB at most. Ship only the near-reference tiers. + "ecapa_tdnn": ("F16", "Q8_0"), } diff --git a/scripts/lib/test_quant_policy_sync.py b/scripts/lib/test_quant_policy_sync.py index b3d5b9482..e69c0921a 100644 --- a/scripts/lib/test_quant_policy_sync.py +++ b/scripts/lib/test_quant_policy_sync.py @@ -86,6 +86,8 @@ "dec.pos_emb.weight", # whisper decoder pos_emb "frontend.mel_filterbank", # mel frontend buffer "frontend.window", # window frontend buffer + "blk.1.tdnn1.bn.scale", # ecapa_tdnn folded BN affine (.bn.) + "blk.1.res2.0.conv.bias", # ecapa_tdnn conv bias (.bias) ] # Conv bucket: 2D / depthwise / 1x1 pointwise conv kernels. The loader has no @@ -105,6 +107,7 @@ "enc.blocks.3.conv.pointwise2.weight", # conformer 1x1 pointwise "enc.pre_encode.conv.0.weight", # pre-encode subsampling conv "enc.blocks.3.conv.depthwise.weight", # conformer depthwise conv + "blk.1.res2.0.conv.weight", # ecapa_tdnn tap-major k>1 kernel "vad.proj.weight", # parakeet-ultra VAD head 1x1 conv "vad.ctx.weight", # parakeet-ultra VAD head k=5 conv ] @@ -131,6 +134,9 @@ "enc.blocks.3.attn.rel_pos_emb.weight", "enc.blocks.3.attn.kv.weight", # granite5_ctc fused K|V projection "enc.ctc_proj.weight", # granite5_ctc tied CTC head + "blk.1.tdnn1.weight", # ecapa_tdnn 1x1 TDNN (a matmul operand) + "asp.attn.weight", # ecapa_tdnn ASP attention 1x1 (NOT .conv.) + "cls.out.weight", # ecapa_tdnn classifier head ] # KNOWN DRIFT — policy.cpp::classify_tensor places these in the Norm (F32) or diff --git a/scripts/validate.py b/scripts/validate.py index ac87dd4fa..e7e6cfadc 100644 --- a/scripts/validate.py +++ b/scripts/validate.py @@ -327,6 +327,16 @@ def parse_cli_transcript(output: str) -> str | None: return None +def parse_cli_language(output: str) -> dict[str, Any] | None: + """The LANGID CLI's `language: index= p=

` line.""" + for line in output.splitlines(): + if line.startswith("language: "): + parts = line[len("language: "):].split() + fields = dict(p.split("=", 1) for p in parts[1:] if "=" in p) + return {"code": parts[0], "label_index": int(fields["index"]), "p": float(fields["p"])} + return None + + def write_cpp_transcript( out_dir: Path, *, @@ -397,7 +407,7 @@ def cmd_ref(args: argparse.Namespace) -> int: # family's dumper accepts --revision. The dumper itself ignores # --revision when --model resolves to a local directory. hf_revision = (manifest.get("source_model") or {}).get("hf_revision") - if hf_revision and args.family in ("qwen3_asr", "granite_nar", "granite5_ctc"): + if hf_revision and args.family in ("qwen3_asr", "granite_nar", "granite5_ctc", "ecapa_tdnn"): common_args += ["--revision", str(hf_revision)] # Forward any manifest-declared dumper args verbatim. Used today @@ -429,8 +439,14 @@ def cmd_ref(args: argparse.Namespace) -> int: # (--preset) and the C++ side (TRANSCRIBE_SORTFORMER_STREAM_PRESET in # cmd_cpp); unset -> the checkpoint-shipped cfg (single chunk on the # short oracle, i.e. diar.probs == diar.preds_offline). + # + # ecapa_tdnn (language ID) is one forward pass: the `encoder` + # subcommand dumps the front end, encoder, classifier and + # prediction.json from a single classify_batch call. if args.family == "sortformer": stages = ["encoder", "diarize"] + elif args.family == "ecapa_tdnn": + stages = ["encoder"] else: stages = case_stages(case, ["encoder", "decode"]) sf_preset = os.environ.get("VALIDATE_SORTFORMER_PRESET") @@ -571,6 +587,13 @@ def cmd_cpp(args: argparse.Namespace) -> int: transcript = parse_cli_transcript(result.stdout or "") if transcript is None and "speaker segments:" in (result.stdout or ""): continue # a diarizer has no transcript + prediction = parse_cli_language(result.stdout or "") if transcript is None else None + if prediction is not None: + # A language ID model: its behavioural artifact is the top-1 + # label, compared against the reference's prediction.json. + (out_dir / "prediction.json").write_text(json.dumps(prediction, indent=2) + "\n") + print(f" wrote {out_dir / 'prediction.json'}", file=sys.stderr) + continue if transcript is None: raise SystemExit( f"error: cpp dump [{args.family}/{case_name}] did not emit a transcript line" @@ -669,6 +692,23 @@ def cmd_compare(args: argparse.Namespace) -> int: # emits speaker segments, not `text`). Its behavioral artifact is the # diar.probs tensor, gated above; the `diarize` stage's segment lines # are informational only, so skip the text-transcript comparison. + # Language ID: the gate is the top-1 label index, C++ vs reference. + # The case's expected_language is informational (the reference itself + # may miss it); it is reported, never gated. + ref_prediction = ref_dir / "prediction.json" + if ref_prediction.exists(): + ref_pred = json.loads(ref_prediction.read_text()) + cpp_prediction = cpp_dir / "prediction.json" + cpp_pred = (json.loads(cpp_prediction.read_text()) if cpp_prediction.exists() + else {"code": ""}) + match = cpp_pred.get("label_index") == ref_pred["label_index"] + all_passed = all_passed and match + transcript_results.append({"case": case_name, "match": match, "mode": "label", + "reference": ref_pred["code"], "cpp": cpp_pred["code"]}) + expected = case.get("expected_language") if isinstance(case, dict) else None + print(f"\n Prediction: {'ok' if match else 'FAIL'} c++ {cpp_pred['code']!r}, " + f"reference {ref_pred['code']!r}, expected {expected!r}") + ref_transcript = ref_dir / "transcript.json" if ref_transcript.exists() and args.family != "sortformer": transcript_compare = case_transcript_compare(manifest, case) diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt index c20134b77..8e877fa10 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -14,6 +14,7 @@ add_library(transcribe transcribe.cpp transcribe-asr.cpp transcribe-diarize.cpp + transcribe-langid.cpp transcribe-loader.cpp transcribe-arch.cpp transcribe-meta.cpp @@ -67,6 +68,9 @@ add_library(transcribe arch/sortformer/model.cpp arch/sortformer/stream.cpp arch/sortformer/weights.cpp + arch/ecapa_tdnn/graph.cpp + arch/ecapa_tdnn/model.cpp + arch/ecapa_tdnn/weights.cpp arch/voxtral/model.cpp arch/voxtral/capabilities.cpp arch/voxtral/weights.cpp diff --git a/src/arch/ecapa_tdnn/ecapa_tdnn.h b/src/arch/ecapa_tdnn/ecapa_tdnn.h new file mode 100644 index 000000000..f282ad5f0 --- /dev/null +++ b/src/arch/ecapa_tdnn/ecapa_tdnn.h @@ -0,0 +1,64 @@ +// arch/ecapa_tdnn/ecapa_tdnn.h - ECAPA-TDNN family model and session types +// (SpeechBrain ECAPA-TDNN + VoxLingua107 classifier, LANGID role). +// +// INTERNAL to src/arch/ecapa_tdnn/. + +#pragma once + +#include "transcribe-backend.h" +#include "transcribe-langid.h" +#include "transcribe-mel.h" +#include "transcribe-model.h" +#include "weights.h" + +#include +#include +#include + +struct ggml_context; +struct ggml_backend_buffer; +typedef struct ggml_backend_buffer * ggml_backend_buffer_t; + +namespace transcribe { +struct Arch; +} + +namespace transcribe::ecapa_tdnn { + +struct Model final : public transcribe_model { + HParams hparams; + Weights weights; + + // GGUF metadata context: owns the ggml_tensor structs the weight slots + // above point at. Their data lives in backend_buffer. + ggml_context * ctx_meta = nullptr; + ggml_backend_buffer_t backend_buffer = nullptr; + BackendPlan plan; + + // Load-time derived weights (Weights::blk0_w_im2col): their own metadata + // context and backend buffer, since ctx_meta is sized for the file. + ggml_context * ctx_derived = nullptr; + ggml_backend_buffer_t derived_buffer = nullptr; + + // Built at load from stt.langid.labels.*; immutable after. + LangidLabels labels; + + // Built at load from stt.frontend.* and frontend.mel_filterbank; shared + // by every session. + std::unique_ptr mel; + + Model() = default; + ~Model() override; +}; + +struct Session final : public transcribe_langid_session { + // Host scratch, reused across calls. + std::vector mel_raw; // [n_mels * T], mel-major (MelFrontend) + std::vector mel_buf; // [T * n_mels], frame-major + std::vector im2col_buf; // [T * blk0_cols], frame-major + std::vector idx_buf[kNumSeBlocks]; // reflect indices, [T + 2p] +}; + +extern const Arch arch; + +} // namespace transcribe::ecapa_tdnn diff --git a/src/arch/ecapa_tdnn/graph.cpp b/src/arch/ecapa_tdnn/graph.cpp new file mode 100644 index 000000000..b306d7bb8 --- /dev/null +++ b/src/arch/ecapa_tdnn/graph.cpp @@ -0,0 +1,352 @@ +// arch/ecapa_tdnn/graph.cpp - the ECAPA-TDNN forward graph. +// +// SpeechBrain ECAPA_TDNN.forward + Xvector.Classifier, after the converter's +// rewrites: BatchNorm as scale/shift, activation-free BatchNorms folded into +// the next linear map, and the MFA / ASP concats split into weight blocks. +// +// Activations are ne = [C, T]: a 1x1 conv is one mul_mat, reflect padding is +// one get_rows over frames, and time reductions run on a [T, C] transpose. +// +// mel [n_mels, T] +// -> blk.0 TDNNBlock n_mels -> C, k=K0 (host im2col + one matmul) +// -> blk.1..3 SERes2Net C -> C +// -> mfa 1x1 over concat(blk.1..3) [3C, T] +// -> asp attentive statistics pooling [6C] +// -> fc 6C -> emb [emb] +// -> classifier LeakyReLU, emb -> hid, LeakyReLU, hid -> n_labels + +#include "graph.h" + +#include "ecapa_tdnn.h" +#include "ggml.h" +#include "transcribe-debug.h" + +#include +#include +#include +#include + +namespace transcribe::ecapa_tdnn { + +namespace { + +// Name a tensor and return it; NULL-safe. +ggml_tensor * named(ggml_tensor * t, const char * name) { + if (t != nullptr && name != nullptr) { + ggml_set_name(t, name); + } + return t; +} + +// y = W x (+ b): w is ne = [IC, OC], x is ne = [IC, T] or [IC]; b may be null. +ggml_tensor * linear(ggml_context * ctx, ggml_tensor * w, ggml_tensor * x, ggml_tensor * b) { + ggml_tensor * y = ggml_mul_mat(ctx, w, x); + if (w->type == GGML_TYPE_F16) { + ggml_prec_set_acc(y, GGML_PREC_F32); + } + if (b != nullptr) { + y = ggml_add(ctx, y, b); + } + return y; +} + +// BatchNorm stored as an affine map: y = x * scale + shift, both ne = [C]. +ggml_tensor * bn_affine(ggml_context * ctx, ggml_tensor * x, ggml_tensor * scale, ggml_tensor * shift) { + ggml_tensor * y = ggml_mul(ctx, x, scale); + y = ggml_add(ctx, y, shift); + return y; +} + +// Dilated 1-D conv as a sum over taps: xpad is the reflect-padded +// ne = [IC, T + 2p], w the tap-major ne = [IC, OC, K]. +// y[t] = sum_k w_k . xpad[t + k*dilation], ne = [OC, T]. +ggml_tensor * conv_taps(ggml_context * ctx, ggml_tensor * xpad, ggml_tensor * w, ggml_tensor * b, int dilation, int T) { + const int64_t IC = w->ne[0]; + const int64_t K = w->ne[2]; + + ggml_tensor * acc = nullptr; + for (int64_t k = 0; k < K; ++k) { + // Rows [k*d, k*d + T) of xpad, and tap k's [IC, OC] slice of w. + ggml_tensor * v = ggml_view_2d(ctx, xpad, IC, T, xpad->nb[1], static_cast(k * dilation) * xpad->nb[1]); + ggml_tensor * wk = ggml_view_2d(ctx, w, IC, w->ne[1], w->nb[1], static_cast(k) * w->nb[2]); + + ggml_tensor * y = ggml_mul_mat(ctx, wk, v); + if (w->type == GGML_TYPE_F16) { + ggml_prec_set_acc(y, GGML_PREC_F32); + } + acc = (acc == nullptr) ? y : ggml_add(ctx, acc, y); + } + + if (b != nullptr) { + acc = ggml_add(ctx, acc, b); + } + return acc; +} + +// Name a tensor, mark it for dumping (no-op unless TRANSCRIBE_DUMP_DIR is +// set), and stash the pointer. +void mark_dump(ggml_tensor *& slot, ggml_tensor * t, const char * name) { + named(t, name); + debug::mark_tensor_for_dump(t); + if (t->view_src != nullptr) { + // In-place ops return a view; keep its backing tensor alive too. + debug::mark_tensor_for_dump(t->view_src); + } + slot = t; +} + +// One TDNNBlock with k = 1: 1x1 conv -> ReLU -> BatchNorm. +ggml_tensor * tdnn_1x1(ggml_context * ctx, const TdnnLayer & l, ggml_tensor * x) { + ggml_tensor * y = linear(ctx, l.w, x, l.b); + y = ggml_relu(ctx, y); + return bn_affine(ctx, y, l.bn.scale, l.bn.shift); +} + +// One TDNNBlock with k > 1 over an already reflect-padded input: +// dilated conv -> ReLU -> BatchNorm. +ggml_tensor * tdnn_conv(ggml_context * ctx, + ggml_tensor * xpad, + ggml_tensor * w, + ggml_tensor * b, + const BnAffine & bn, + int dilation, + int T) { + ggml_tensor * y = conv_taps(ctx, xpad, w, b, dilation, T); + y = ggml_relu(ctx, y); + return bn_affine(ctx, y, bn.scale, bn.shift); +} + +// Mean over the time axis of x ne = [C, T], returned as ne = [C]. +ggml_tensor * mean_over_time(ggml_context * ctx, ggml_tensor * x) { + ggml_tensor * xT = ggml_cont(ctx, ggml_transpose(ctx, x)); // [T, C] + return ggml_reshape_1d(ctx, ggml_mean(ctx, xT), x->ne[0]); +} + +// SpeechBrain Res2NetBlock with scale S over chunks x0..x(S-1) of C/S: +// y0 = x0, y1 = B0(x1), yi = B(i-1)(xi + y(i-1)), out = concat(y0..y(S-1)) +// Each yi is written over chunk i of h in place with set_rows (h viewed as +// [C/S, S, T]), so there is no concat. +ggml_tensor * res2net(ggml_context * ctx, + const SeRes2NetBlock & blk, + ggml_tensor * h, + ggml_tensor * idx, + ggml_tensor * chunk_ids, + int dilation, + int T, + int chunk) { + const size_t chunk_bytes = static_cast(chunk) * ggml_element_size(h); + + ggml_tensor * out = ggml_reshape_3d(ctx, h, chunk, kRes2NetScale, T); + ggml_tensor * prev = nullptr; // y(i-1) + + for (int i = 1; i < kRes2NetScale; ++i) { + ggml_tensor * ci = ggml_view_2d(ctx, out, chunk, T, out->nb[2], static_cast(i) * chunk_bytes); + ggml_tensor * in = (i == 1) ? ci : ggml_add(ctx, ci, prev); + const Res2Sub & s = blk.res2[i - 1]; + ggml_tensor * yi = tdnn_conv(ctx, ggml_get_rows(ctx, in, idx), s.w, s.b, s.bn, dilation, T); + + out = ggml_set_rows(ctx, out, ggml_reshape_3d(ctx, yi, chunk, 1, T), + ggml_view_1d(ctx, chunk_ids, 1, static_cast(i) * sizeof(int32_t))); + prev = yi; + } + + return ggml_reshape_2d(ctx, out, h->ne[0], T); +} + +// SpeechBrain SEBlock: s = mean_T(x) -> 1x1 -> ReLU -> 1x1 -> sigmoid, then +// x * s broadcast over T. +ggml_tensor * se_block(ggml_context * ctx, const SeBlock & se, ggml_tensor * x) { + ggml_tensor * s = mean_over_time(ctx, x); + s = ggml_relu(ctx, linear(ctx, se.c1_w, s, se.c1_b)); + s = ggml_sigmoid(ctx, linear(ctx, se.c2_w, s, se.c2_b)); + return ggml_mul(ctx, x, s); +} + +// MFA + attentive statistics pooling, stock ops. Returns pooled [2*Cm]. +ggml_tensor * mfa_asp(ggml_context * ctx, + GraphBuild & gb, + const HParams & hp, + const Weights & w, + ggml_tensor * const * blk_out, + int Cm) { + const int64_t T = blk_out[0]->ne[1]; + + // ---- MFA: 1x1 over concat(blk1, blk2, blk3), split into three blocks -- + // Expand all three products first and start the sum from part[2] so the + // two adds stay adjacent (Vulkan fuses them). + ggml_tensor * part[kNumSeBlocks] = {}; + for (int i = 0; i < kNumSeBlocks; ++i) { + part[i] = linear(ctx, w.mfa_w[i], blk_out[i], nullptr); + ggml_build_forward_expand(gb.graph, part[i]); + } + ggml_tensor * m = ggml_add(ctx, ggml_add(ctx, part[2], part[0]), part[1]); + m = ggml_add(ctx, m, w.mfa_b); + m = ggml_relu(ctx, m); + m = bn_affine(ctx, m, w.mfa_bn.scale, w.mfa_bn.shift); + mark_dump(gb.dumps.mfa_out, m, "enc.mfa.out"); + + // ---- attentive statistics pooling ------------------------------------ + // Uniform mean / biased std over T, std clamped at eps (SpeechBrain). + ggml_tensor * mT = ggml_cont(ctx, ggml_transpose(ctx, m)); // [T, Cm] + ggml_tensor * mean = ggml_mean(ctx, mT); // [1, Cm] + ggml_tensor * d = ggml_sub(ctx, mT, mean); // [T, Cm] + ggml_tensor * d2 = ggml_sqr(ctx, d); + ggml_tensor * var = ggml_mean(ctx, d2); + ggml_tensor * sd = ggml_sqrt(ctx, ggml_clamp(ctx, var, hp.asp_eps, FLT_MAX)); + + ggml_tensor * mean1 = ggml_reshape_1d(ctx, mean, Cm); + ggml_tensor * sd1 = ggml_reshape_1d(ctx, sd, Cm); + + // wm@mean + ws@std + b is constant over T: one [att] vector, broadcast. + ggml_tensor * cvec = ggml_add(ctx, linear(ctx, w.asp_wm, mean1, nullptr), linear(ctx, w.asp_ws, sd1, nullptr)); + cvec = ggml_add(ctx, cvec, w.asp_b); + + ggml_tensor * a = ggml_add(ctx, linear(ctx, w.asp_wx, m, nullptr), cvec); // [att, T] + a = ggml_relu(ctx, a); + a = bn_affine(ctx, a, w.asp_bn.scale, w.asp_bn.shift); + a = ggml_tanh(ctx, a); + + // The per-channel bias cancels in the softmax over T; only the dump adds it. + ggml_tensor * al = linear(ctx, w.asp_attn_w, a, nullptr); // [Cm, T] + if (debug::enabled()) { + mark_dump(gb.dumps.asp_attn_logits, ggml_add(ctx, al, w.asp_attn_b), "enc.asp.attn_logits"); + ggml_build_forward_expand(gb.graph, gb.dumps.asp_attn_logits); + } + + ggml_tensor * alT = ggml_cont(ctx, ggml_transpose(ctx, al)); // [T, Cm] + ggml_tensor * attn = ggml_soft_max(ctx, alT); // over ne[0] = T + + // Weighted moments about the uniform mean (sum_t attn = 1): + // e1 = sum_t attn d, e2 = sum_t attn d^2, mu = mean + e1, sigma^2 = e2 - e1^2 + ggml_tensor * attn3 = ggml_reshape_3d(ctx, attn, T, 1, Cm); + ggml_tensor * e1 = ggml_reshape_2d(ctx, ggml_mul_mat(ctx, attn3, ggml_reshape_3d(ctx, d, T, 1, Cm)), 1, Cm); + ggml_tensor * e2 = ggml_reshape_2d(ctx, ggml_mul_mat(ctx, attn3, ggml_reshape_3d(ctx, d2, T, 1, Cm)), 1, Cm); + ggml_tensor * mu = ggml_add(ctx, mean, e1); // [1, Cm] + ggml_tensor * sig = ggml_sub(ctx, e2, ggml_sqr(ctx, e1)); + sig = ggml_sqrt(ctx, ggml_clamp(ctx, sig, hp.asp_eps, FLT_MAX)); + + ggml_tensor * pooled = + ggml_concat(ctx, ggml_reshape_1d(ctx, mu, Cm), ggml_reshape_1d(ctx, sig, Cm), /*dim=*/0); // [2*Cm] + mark_dump(gb.dumps.asp_out, pooled, "enc.asp.out"); + + return pooled; +} + +} // namespace + +void build_blk0_im2col(const HParams & hp, const float * mel, int T, std::vector & out) { + const int n_mels = hp.mel_n_mels; + const int K = hp.kernel_sizes[0]; + const int d = hp.dilations[0]; + const int p = hp.pad(0); + const size_t cols = static_cast(hp.blk0_cols()); + out.assign(static_cast(T) * cols, 0.0f); + for (int t = 0; t < T; ++t) { + float * row = out.data() + static_cast(t) * cols; + for (int k = 0; k < K; ++k) { + int src = t + k * d - p; // torch "reflect", edge not repeated + if (src < 0) { + src = -src; + } else if (src >= T) { + src = 2 * (T - 1) - src; + } + const float * m = mel + static_cast(src) * static_cast(n_mels); + std::copy(m, m + n_mels, row + static_cast(k) * static_cast(n_mels)); + } + } +} + +GraphBuild build_graph(ggml_context * ctx, const Model & model, int T) { + GraphBuild gb{}; + + const HParams & hp = model.hparams; + const Weights & w = model.weights; + + const int Cm = hp.c_mfa(); + const int chunk = hp.c_chunk(); + + gb.graph = ggml_new_graph_custom(ctx, kGraphSize, /*grads=*/false); + + // ---- inputs ---------------------------------------------------------- + gb.blk0_in = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hp.blk0_cols(), T); + named(gb.blk0_in, "fe.mel.im2col"); + ggml_set_input(gb.blk0_in); + + for (int i = 0; i < kNumSeBlocks; ++i) { + char name[32]; + std::snprintf(name, sizeof(name), "reflect.idx.%d", i); + gb.idx[i] = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, T + 2 * hp.pad(i + 1)); + named(gb.idx[i], name); + ggml_set_input(gb.idx[i]); + } + gb.chunk_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, kRes2NetScale); + named(gb.chunk_ids, "res2.chunk_ids"); + ggml_set_input(gb.chunk_ids); + + // ---- stage 0: TDNNBlock n_mels -> C ---------------------------------- + // One matmul over the host im2col, inner dim padded to a multiple of 32 + // for the CPU tinyBLAS tiles. + ggml_tensor * x = linear(ctx, w.blk0_w_im2col, gb.blk0_in, w.blk0_b); + x = ggml_relu(ctx, x); + x = bn_affine(ctx, x, w.blk0_bn.scale, w.blk0_bn.shift); + mark_dump(gb.dumps.blk0_out, x, "enc.blk.0.out"); + + // ---- stages 1..3: SERes2Net ------------------------------------------ + ggml_tensor * blk_out[kNumSeBlocks] = { nullptr, nullptr, nullptr }; + for (int i = 0; i < kNumSeBlocks; ++i) { + const SeRes2NetBlock & blk = w.blocks[i]; + const int d = hp.dilations[static_cast(i + 1)]; + + ggml_tensor * residual = x; + + ggml_tensor * h = tdnn_1x1(ctx, blk.tdnn1, x); + if (i == 0) { + mark_dump(gb.dumps.blk1_tdnn1_out, h, "enc.blk.1.tdnn1.out"); + } + + // Res2Net writes in place; copy so the tdnn1 dump survives. + if (i == 0 && debug::enabled()) { + h = ggml_cont(ctx, h); + } + h = res2net(ctx, blk, h, gb.idx[i], gb.chunk_ids, d, T, chunk); + if (i == 0) { + mark_dump(gb.dumps.blk1_res2_out, h, "enc.blk.1.res2.out"); + } + + h = tdnn_1x1(ctx, blk.tdnn2, h); + h = se_block(ctx, blk.se, h); + if (i == 0) { + mark_dump(gb.dumps.blk1_se_out, h, "enc.blk.1.se.out"); + } + x = ggml_add(ctx, h, residual); + + char name[32]; + std::snprintf(name, sizeof(name), "enc.blk.%d.out", i + 1); + mark_dump(gb.dumps.blk_out[i], x, name); + blk_out[i] = x; + } + + ggml_tensor * pooled = mfa_asp(ctx, gb, hp, w, blk_out, Cm); + + // ---- embedding head (asp_bn folded into fc) --------------------------- + ggml_tensor * emb = linear(ctx, w.fc_w, pooled, w.fc_b); + mark_dump(gb.dumps.emb, emb, "enc.emb"); + + // ---- classifier ------------------------------------------------------- + // LeakyReLU precedes each (folded) BatchNorm. + ggml_tensor * e = ggml_leaky_relu(ctx, emb, hp.leaky_slope, /*inplace=*/false); + ggml_tensor * hcls = + ggml_leaky_relu(ctx, linear(ctx, w.cls_l1_w, e, w.cls_l1_b), hp.leaky_slope, /*inplace=*/false); + mark_dump(gb.dumps.cls_hidden, hcls, "cls.hidden"); + + ggml_tensor * z = linear(ctx, w.cls_out_w, hcls, w.cls_out_b); + mark_dump(gb.dumps.cls_logits, z, "cls.logits_raw"); + gb.logits = z; + + ggml_set_output(gb.logits); + ggml_build_forward_expand(gb.graph, gb.logits); + + return gb; +} + +} // namespace transcribe::ecapa_tdnn diff --git a/src/arch/ecapa_tdnn/graph.h b/src/arch/ecapa_tdnn/graph.h new file mode 100644 index 000000000..e7626c1db --- /dev/null +++ b/src/arch/ecapa_tdnn/graph.h @@ -0,0 +1,70 @@ +// arch/ecapa_tdnn/graph.h - ECAPA-TDNN forward graph builder. +// +// INTERNAL to src/arch/ecapa_tdnn/. Activations are ggml ne = [C, T]. + +#pragma once + +#include "ggml.h" +#include "weights.h" + +#include +#include + +struct ggml_context; +struct ggml_cgraph; +struct ggml_tensor; + +namespace transcribe::ecapa_tdnn { + +struct Model; + +// Graph node capacity (the built graph is ~571 nodes, independent of T). +constexpr size_t kGraphSize = 2048; + +// Named intermediates for the parity dumps (marked as graph outputs when +// TRANSCRIBE_DUMP_DIR is set). Names match +// scripts/dump_reference_ecapa_tdnn_speechbrain.py. +struct Dumps { + ggml_tensor * blk0_out = nullptr; // enc.blk.0.out [C, T] + ggml_tensor * blk1_tdnn1_out = nullptr; // enc.blk.1.tdnn1.out [C, T] + ggml_tensor * blk1_res2_out = nullptr; // enc.blk.1.res2.out [C, T] + ggml_tensor * blk1_se_out = nullptr; // enc.blk.1.se.out [C, T] + ggml_tensor * blk_out[kNumSeBlocks] = { nullptr, nullptr, nullptr }; + // enc.blk.{1,2,3}.out [C, T] + ggml_tensor * mfa_out = nullptr; // enc.mfa.out [Cm, T] + ggml_tensor * asp_attn_logits = nullptr; // enc.asp.attn_logits [Cm, T] + ggml_tensor * asp_out = nullptr; // enc.asp.out [2*Cm] + ggml_tensor * emb = nullptr; // enc.emb [emb] + ggml_tensor * cls_hidden = nullptr; // cls.hidden [hid] + ggml_tensor * cls_logits = nullptr; // cls.logits_raw [n_labels] +}; + +struct GraphBuild { + // Stage-0 input: im2col of the log-mel, ne = [blk0_cols, T] (build_blk0_im2col). + ggml_tensor * blk0_in = nullptr; + + // Reflect-padding gather indices, one per SERes2Net block, I32 + // ne = [T + 2*pad(i+1)]. + ggml_tensor * idx[kNumSeBlocks] = { nullptr, nullptr, nullptr }; + + // 0..kRes2NetScale-1, I32: set_rows targets for the Res2Net chunks. + ggml_tensor * chunk_ids = nullptr; + + // Output. + ggml_tensor * logits = nullptr; // ne = [n_labels] + + Dumps dumps{}; + + ggml_cgraph * graph = nullptr; +}; + +// Build the forward graph for T frames into a fresh no_alloc `ctx`. +// Requires T > hp.pad(i) for every stage. +GraphBuild build_graph(ggml_context * ctx, const Model & model, int T); + +// Fill `out` ([T, hp.blk0_cols()] frame-major, i.e. ggml ne = [cols, T]) +// with the stage-0 im2col of the frame-major log-mel `mel` [T, n_mels]: +// out[t, k*n_mels + m] = mel[reflect(t + k*d0 - p0), m], zeros past K0*n_mels. +void build_blk0_im2col(const HParams & hp, const float * mel, int T, std::vector & out); + +} // namespace transcribe::ecapa_tdnn diff --git a/src/arch/ecapa_tdnn/model.cpp b/src/arch/ecapa_tdnn/model.cpp new file mode 100644 index 000000000..f97b1e684 --- /dev/null +++ b/src/arch/ecapa_tdnn/model.cpp @@ -0,0 +1,555 @@ +// arch/ecapa_tdnn/model.cpp - ECAPA-TDNN family handler (LANGID role): +// load, forward pass, and the LANGID ops table. + +#include "ecapa_tdnn.h" +#include "ggml-alloc.h" +#include "ggml-backend.h" +#include "ggml.h" +#include "gguf.h" +#include "graph.h" +#include "transcribe-arch.h" +#include "transcribe-backend.h" +#include "transcribe-batch-util.h" +#include "transcribe-debug.h" +#include "transcribe-load-common.h" +#include "transcribe-loader.h" +#include "transcribe-log.h" +#include "transcribe-meta.h" +#include "transcribe-path.h" +#include "weights.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace transcribe::ecapa_tdnn { + +namespace { + +constexpr const char * kTag = "ecapa_tdnn"; + +// Default when stt.variant is absent. +constexpr const char k_default_variant[] = "lang-id-voxlingua107-ecapa"; + +constexpr size_t k_compute_ctx_bytes = 4u * 1024u * 1024u; + +// Frees the gguf_context on every exit path of load(). +struct GgufGuard { + gguf_context * ctx = nullptr; + + GgufGuard() = default; + + ~GgufGuard() { + if (ctx != nullptr) { + gguf_free(ctx); + } + } + + GgufGuard(const GgufGuard &) = delete; + GgufGuard & operator=(const GgufGuard &) = delete; +}; + +// Reflect-padding gather indices (torch "reflect", edge not repeated): +// [p, p-1, ..., 1, 0, 1, ..., T-1, T-2, ..., T-1-p]. Requires T > p. +void build_reflect_indices(int T, int p, std::vector & out) { + out.resize(static_cast(T) + 2 * static_cast(p)); + for (int j = 0; j < p; ++j) { + out[static_cast(j)] = p - j; + } + for (int t = 0; t < T; ++t) { + out[static_cast(p + t)] = t; + } + for (int j = 0; j < p; ++j) { + out[static_cast(p + T + j)] = T - 2 - j; + } +} + +// Read stt.langid.labels.* into a label table. Aliases are optional. +transcribe_status read_labels(const gguf_context * gguf, LangidLabels & out) { + std::vector codes; + std::vector names; + std::vector aliases; + if (read_string_array_kv(gguf, "stt.langid.labels.codes", codes) != KvResult::Ok || + read_string_array_kv(gguf, "stt.langid.labels.names", names) != KvResult::Ok) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: stt.langid.labels.codes / names missing or not string arrays", kTag); + return TRANSCRIBE_ERR_GGUF; + } + if (read_string_array_kv(gguf, "stt.langid.labels.aliases", aliases) == KvResult::BadType) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: stt.langid.labels.aliases is not a string array", kTag); + return TRANSCRIBE_ERR_GGUF; + } + return build_langid_labels(std::move(codes), std::move(names), aliases, kTag, out); +} + +// Q8_0 is a download format here, not a compute format: every Q8_0 weight is +// widened to F16 at load. ggml's Q8_0 matmuls also quantize the ACTIVATIONS +// to 8 bits, and on FLEURS that, not the 8-bit weights, is what moves the +// decision; the same weights computed in F16 score at F32's accuracy +// (agreement figures: docs/models/lang-id-voxlingua107-ecapa.md). The cost +// is F16 memory for those weights, and F16 speed. +// +// Retype in place: the tensors come from the no_alloc gguf context and have +// no data or views yet, so only type and strides change. Returns the count. +int widen_q8_0_weights(ggml_context * ctx_meta) { + int n = 0; + for (ggml_tensor * t = ggml_get_first_tensor(ctx_meta); t != nullptr; t = ggml_get_next_tensor(ctx_meta, t)) { + if (t->type != GGML_TYPE_Q8_0) { + continue; + } + t->type = GGML_TYPE_F16; + t->nb[0] = ggml_type_size(GGML_TYPE_F16); + for (int i = 1; i < GGML_MAX_DIMS; ++i) { + t->nb[i] = t->nb[i - 1] * static_cast(t->ne[i - 1]); + } + ++n; + } + return n; +} + +// load_common::stream_tensor_data, plus the Q8_0 -> F16 widening: a tensor +// whose file type is Q8_0 and whose bound type is F16 is dequantized with +// ggml's own Q8_0 reference and rounded to F16 on the way in. Every other +// tensor must be bound with its file type and is copied as-is. +transcribe_status stream_weights(const std::string & path, const gguf_context * gguf, ggml_context * ctx_meta) { + std::ifstream fin(path_from_utf8(path), std::ios::binary); + if (!fin) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: failed to reopen %s for tensor data", kTag, path.c_str()); + return TRANSCRIBE_ERR_GGUF; + } + + const size_t data_offset = gguf_get_data_offset(gguf); + const ggml_to_float_t q8_to_f32 = ggml_get_type_traits(GGML_TYPE_Q8_0)->to_float; + std::vector staging; + std::vector f32; + std::vector f16; + + for (ggml_tensor * t = ggml_get_first_tensor(ctx_meta); t != nullptr; t = ggml_get_next_tensor(ctx_meta, t)) { + const int64_t idx = gguf_find_tensor(gguf, t->name); + if (idx < 0) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: tensor \"%s\" not in gguf data", kTag, t->name); + return TRANSCRIBE_ERR_GGUF; + } + const ggml_type file_type = gguf_get_tensor_type(gguf, idx); + const bool widen = file_type == GGML_TYPE_Q8_0 && t->type == GGML_TYPE_F16; + if (file_type != t->type && !widen) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: tensor \"%s\" is %s in the file but bound as %s", kTag, t->name, + ggml_type_name(file_type), ggml_type_name(t->type)); + return TRANSCRIBE_ERR_GGUF; + } + const size_t nbytes = gguf_get_tensor_size(gguf, idx); + if (!widen && nbytes != ggml_nbytes(t)) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: tensor \"%s\" size mismatch", kTag, t->name); + return TRANSCRIBE_ERR_GGUF; + } + + fin.seekg(static_cast(data_offset) + + static_cast(gguf_get_tensor_offset(gguf, idx))); + if (staging.size() < nbytes) { + staging.resize(nbytes); + } + fin.read(reinterpret_cast(staging.data()), static_cast(nbytes)); + if (!fin) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: short read for tensor \"%s\" (%zu bytes)", kTag, t->name, nbytes); + return TRANSCRIBE_ERR_GGUF; + } + + if (!widen) { + ggml_backend_tensor_set(t, staging.data(), 0, nbytes); + continue; + } + const int64_t n = ggml_nelements(t); + if (nbytes != ggml_row_size(GGML_TYPE_Q8_0, n)) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: Q8_0 tensor \"%s\" size mismatch", kTag, t->name); + return TRANSCRIBE_ERR_GGUF; + } + f32.resize(static_cast(n)); + f16.resize(static_cast(n)); + // Split on Q8_0 block boundaries; each element converts the same way + // on any thread, so the result does not depend on the thread count. + const int64_t qk = ggml_blck_size(GGML_TYPE_Q8_0); + const int64_t n_blocks = n / qk; + const int n_thr = static_cast(std::min(default_n_threads(), std::max(1, n_blocks / 64))); + run_on_threads(n_thr, [&](int tid) { + const int64_t e0 = n_blocks * tid / n_thr * qk; + const int64_t ne = n_blocks * (tid + 1) / n_thr * qk - e0; + q8_to_f32(staging.data() + ggml_row_size(GGML_TYPE_Q8_0, e0), f32.data() + e0, ne); + ggml_fp32_to_fp16_row(f32.data() + e0, f16.data() + e0, ne); + }); + ggml_backend_tensor_set(t, f16.data(), 0, ggml_nbytes(t)); + } + return TRANSCRIBE_OK; +} + +// Build Weights::blk0_w_im2col from blk0_w: the tap-major [n_mels, C, K0] +// kernel regrouped to [K0*n_mels, C] (column k*n_mels + m of output row c is +// tap k, mel m) and zero padded to hp.blk0_cols(). Same type as blk0_w. +transcribe_status build_derived_weights(Model & m) { + const HParams & hp = m.hparams; + ggml_tensor * src = m.weights.blk0_w; + const int64_t n_mels = hp.mel_n_mels; + const int64_t C = hp.c_block(); + const int64_t K = hp.kernel_sizes[0]; + const int64_t cols = hp.blk0_cols(); + + ggml_init_params ip{}; + ip.mem_size = 4 * ggml_tensor_overhead(); + ip.no_alloc = true; + m.ctx_derived = ggml_init(ip); + if (m.ctx_derived == nullptr) { + return TRANSCRIBE_ERR_OOM; + } + ggml_tensor * dst = ggml_new_tensor_2d(m.ctx_derived, src->type, cols, C); + ggml_set_name(dst, "blk.0.conv.weight.im2col"); + m.derived_buffer = ggml_backend_alloc_ctx_tensors(m.ctx_derived, m.plan.primary); + if (m.derived_buffer == nullptr) { + return TRANSCRIBE_ERR_OOM; + } + ggml_backend_buffer_set_usage(m.derived_buffer, GGML_BACKEND_BUFFER_USAGE_WEIGHTS); + + const size_t esz = ggml_type_size(src->type); // F32 / F16, block size 1 + std::vector in(ggml_nbytes(src)); + std::vector out(ggml_nbytes(dst), 0); // zero bits are 0.0 in F32 and F16 + ggml_backend_tensor_get(src, in.data(), 0, in.size()); + for (int64_t k = 0; k < K; ++k) { + for (int64_t c = 0; c < C; ++c) { + const uint8_t * s = in.data() + static_cast(k) * src->nb[2] + static_cast(c) * src->nb[1]; + uint8_t * d = out.data() + static_cast(c) * dst->nb[1] + static_cast(k * n_mels) * esz; + std::memcpy(d, s, static_cast(n_mels) * esz); + } + } + ggml_backend_tensor_set(dst, out.data(), 0, out.size()); + m.weights.blk0_w_im2col = dst; + return TRANSCRIBE_OK; +} + +// --------------------------------------------------------------------------- +// load +// --------------------------------------------------------------------------- + +transcribe_status load(Loader & loader, const transcribe_model_load_params * params, transcribe_model ** out_model) { + const int64_t t_load_start = ggml_time_us(); + + auto m = std::make_unique(); + m->arch = &arch; + m->variant = loader.variant().empty() ? k_default_variant : loader.variant(); + + if (auto st = read_hparams(loader.gguf(), m->hparams); st != TRANSCRIBE_OK) { + return st; + } + const HParams & hp = m->hparams; + + if (auto st = read_labels(loader.gguf(), m->labels); st != TRANSCRIBE_OK) { + return st; + } + m->hparams.n_labels = static_cast(m->labels.codes.size()); + + // ---- front end ------------------------------------------------------ + // The filterbank is ne = [n_freq, n_mels], i.e. mel-major as MelConfig + // expects. + { + const size_t expected = static_cast(hp.mel_n_mels) * static_cast(hp.n_freq()); + + MelConfig cfg; + cfg.sample_rate = hp.sample_rate; + cfg.num_mels = hp.mel_n_mels; + cfg.n_fft = hp.mel_n_fft; + cfg.win_length = hp.mel_win; + cfg.hop_length = hp.mel_hop; + cfg.pre_emphasis = 0.0f; + cfg.window_type = hp.mel_window; + cfg.pad_mode = hp.mel_pad_mode; + cfg.log_clamp_min = hp.mel_log_floor; + cfg.top_db = hp.mel_top_db; + cfg.normalize = hp.mel_normalize; + + const auto rr = load_common::read_f32_tensor_checked(loader.gguf(), loader.path(), "frontend.mel_filterbank", + expected, kTag, cfg.filterbank); + if (rr != load_common::ReadF32Result::Ok) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, + "%s: frontend.mel_filterbank missing or unreadable (expected %zu f32 values)", kTag, expected); + return TRANSCRIBE_ERR_GGUF; + } + + m->mel = std::make_unique(cfg); + } + + // ---- weights --------------------------------------------------------- + GgufGuard guard; + { + gguf_init_params init_params{}; + init_params.no_alloc = true; + init_params.ctx = &m->ctx_meta; + guard.ctx = gguf_init_from_file(loader.path().c_str(), init_params); + if (guard.ctx == nullptr) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: failed to reopen \"%s\" for tensor data", kTag, + loader.path().c_str()); + return TRANSCRIBE_ERR_GGUF; + } + } + + if (auto st = build_weights(m->ctx_meta, hp, m->weights); st != TRANSCRIBE_OK) { + return st; + } + + // After build_weights, which validated each weight's file type. + if (const int n_widened = widen_q8_0_weights(m->ctx_meta); n_widened > 0) { + log_msg(TRANSCRIBE_LOG_LEVEL_INFO, "%s: %d Q8_0 weights widened to F16 at load", kTag, n_widened); + } + + const transcribe_backend_request backend_req = params != nullptr ? params->backend : TRANSCRIBE_BACKEND_AUTO; + if (auto st = load_common::init_backends(backend_req, params != nullptr ? params->device : nullptr, kTag, m->plan); + st != TRANSCRIBE_OK) { + return st; + } + m->backend = ggml_backend_name(m->plan.primary); + m->primary_backend = m->plan.primary; + + m->backend_buffer = ggml_backend_alloc_ctx_tensors(m->ctx_meta, m->plan.primary); + if (m->backend_buffer == nullptr) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: ggml_backend_alloc_ctx_tensors failed", kTag); + return TRANSCRIBE_ERR_OOM; + } + ggml_backend_buffer_set_usage(m->backend_buffer, GGML_BACKEND_BUFFER_USAGE_WEIGHTS); + + if (auto st = stream_weights(loader.path(), guard.ctx, m->ctx_meta); st != TRANSCRIBE_OK) { + return st; + } + if (auto st = build_derived_weights(*m); st != TRANSCRIBE_OK) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: derived weight allocation failed", kTag); + return st; + } + + m->roles = TRANSCRIBE_ROLE_LANGID; + set_feature(m.get(), TRANSCRIBE_FEATURE_CANCELLATION, true); + m->t_load_us = ggml_time_us() - t_load_start; + *out_model = m.release(); + return TRANSCRIBE_OK; +} + +// --------------------------------------------------------------------------- +// forward pass +// --------------------------------------------------------------------------- + +// On success gb_out's tensors stay valid until the session scratch is released. +transcribe_status forward(Session & s, const Model & m, const float * pcm, int n_samples, GraphBuild & gb_out) { + debug::init(); + + const HParams & hp = m.hparams; + const int n_threads = s.n_threads > 0 ? s.n_threads : default_n_threads(); + + // ---- front end (host side) ------------------------------------------- + const int64_t t_mel_start = ggml_time_us(); + int T = 0; + int n_mels = 0; + if (auto st = m.mel->compute(pcm, static_cast(n_samples), s.mel_raw, n_mels, T, n_threads); + st != TRANSCRIBE_OK) { + return st; + } + // MelFrontend emits mel-major [n_mels, T]; transpose to frame-major. + s.mel_buf.resize(static_cast(T) * static_cast(n_mels)); + for (int mi = 0; mi < n_mels; ++mi) { + const float * src = s.mel_raw.data() + static_cast(mi) * static_cast(T); + for (int t = 0; t < T; ++t) { + s.mel_buf[static_cast(t) * static_cast(n_mels) + static_cast(mi)] = src[t]; + } + } + s.t_mel_us = ggml_time_us() - t_mel_start; + + if (s.poll_abort()) { + return TRANSCRIBE_ERR_ABORTED; + } + + if (debug::enabled()) { + const long long shape[2] = { T, hp.mel_n_mels }; + debug::dump_host_f32("fe.mel", s.mel_buf.data(), static_cast(s.mel_buf.size()), shape, 2, + "frontend"); + } + + // Reflect padding needs T > pad. + for (int i = 0; i < kNumStages - 1; ++i) { + if (T <= hp.pad(i)) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s run: T=%d is too short for reflect padding %d", kTag, T, hp.pad(i)); + return TRANSCRIBE_ERR_INPUT_TOO_SHORT; + } + } + + // ---- stage-0 im2col and reflect-padding indices ------------------------ + build_blk0_im2col(hp, s.mel_buf.data(), T, s.im2col_buf); + for (int i = 0; i < kNumSeBlocks; ++i) { + build_reflect_indices(T, hp.pad(i + 1), s.idx_buf[i]); + } + + // ---- graph ------------------------------------------------------------- + if (s.compute_ctx != nullptr) { + ggml_free(s.compute_ctx); + s.compute_ctx = nullptr; + } + { + ggml_init_params init_params{}; + init_params.mem_size = k_compute_ctx_bytes; + init_params.mem_buffer = nullptr; + init_params.no_alloc = true; + s.compute_ctx = ggml_init(init_params); + if (s.compute_ctx == nullptr) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s run: compute context allocation failed", kTag); + return TRANSCRIBE_ERR_OOM; + } + } + + GraphBuild gb = build_graph(s.compute_ctx, m, T); + + if (s.sched == nullptr) { + auto & backends = const_cast &>(m.plan.scheduler_list); + s.sched = ggml_backend_sched_new(backends.data(), nullptr, static_cast(backends.size()), kGraphSize, + /*parallel=*/false, /*op_offload=*/true); + if (s.sched == nullptr) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s run: scheduler allocation failed", kTag); + return TRANSCRIBE_ERR_OOM; + } + } + configure_sched_n_threads(s.sched, n_threads); + + ggml_backend_sched_reset(s.sched); + if (!ggml_backend_sched_alloc_graph(s.sched, gb.graph)) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s run: graph allocation failed", kTag); + return TRANSCRIBE_ERR_OOM; + } + + // ---- inputs ------------------------------------------------------------- + ggml_backend_tensor_set(gb.blk0_in, s.im2col_buf.data(), 0, s.im2col_buf.size() * sizeof(float)); + for (int i = 0; i < kNumSeBlocks; ++i) { + ggml_backend_tensor_set(gb.idx[i], s.idx_buf[i].data(), 0, s.idx_buf[i].size() * sizeof(int32_t)); + } + { + int32_t ids[kRes2NetScale]; + for (int i = 0; i < kRes2NetScale; ++i) { + ids[i] = i; + } + ggml_backend_tensor_set(gb.chunk_ids, ids, 0, sizeof(ids)); + } + + // ---- compute ------------------------------------------------------------- + const int64_t t_enc_start = ggml_time_us(); + if (const ggml_status gs = ggml_backend_sched_graph_compute(s.sched, gb.graph); gs != GGML_STATUS_SUCCESS) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s run: graph compute failed (%d)", kTag, static_cast(gs)); + return TRANSCRIBE_ERR_BACKEND; + } + s.t_encode_us = ggml_time_us() - t_enc_start; + + gb_out = gb; + return TRANSCRIBE_OK; +} + +// Write the contract's stage tensors. No-op unless TRANSCRIBE_DUMP_DIR is set. +void dump_stages(const GraphBuild & gb) { + if (!debug::enabled()) { + return; + } + + auto try_dump = [](const char * name, ggml_tensor * t, const char * stage) { + if (t != nullptr) { + debug::dump_tensor(name, t, stage); + } + }; + + try_dump("enc.blk.0.out", gb.dumps.blk0_out, "encoder"); + try_dump("enc.blk.1.tdnn1.out", gb.dumps.blk1_tdnn1_out, "encoder"); + try_dump("enc.blk.1.res2.out", gb.dumps.blk1_res2_out, "encoder"); + try_dump("enc.blk.1.se.out", gb.dumps.blk1_se_out, "encoder"); + try_dump("enc.blk.1.out", gb.dumps.blk_out[0], "encoder"); + try_dump("enc.blk.2.out", gb.dumps.blk_out[1], "encoder"); + try_dump("enc.blk.3.out", gb.dumps.blk_out[2], "encoder"); + try_dump("enc.mfa.out", gb.dumps.mfa_out, "encoder"); + try_dump("enc.asp.attn_logits", gb.dumps.asp_attn_logits, "encoder"); + try_dump("enc.asp.out", gb.dumps.asp_out, "encoder"); + try_dump("enc.emb", gb.dumps.emb, "encoder"); + try_dump("cls.hidden", gb.dumps.cls_hidden, "classifier"); + try_dump("cls.logits_raw", gb.dumps.cls_logits, "classifier"); +} + +// --------------------------------------------------------------------------- +// LANGID ops +// --------------------------------------------------------------------------- + +const LangidLabels & langid_labels(const transcribe_model * model) { + return static_cast(model)->labels; +} + +transcribe_langid_session * langid_new_session() { + return new Session(); +} + +transcribe_status langid_run(transcribe_langid_session * session, + const float * pcm, + int n_samples, + std::vector & logits) { + auto & s = static_cast(*session); + const auto & m = static_cast(*session->model); + if (s.poll_abort()) { + return TRANSCRIBE_ERR_ABORTED; + } + + GraphBuild gb{}; + if (auto st = forward(s, m, pcm, n_samples, gb); st != TRANSCRIBE_OK) { + return st; + } + + logits.assign(static_cast(m.hparams.n_labels), 0.0f); + ggml_backend_tensor_get(gb.logits, logits.data(), 0, logits.size() * sizeof(float)); + + dump_stages(gb); + return TRANSCRIBE_OK; +} + +const LangidOps k_langid_ops = { langid_labels, langid_new_session, langid_run }; + +} // namespace + +// --------------------------------------------------------------------------- +// Destructor, registry entry +// --------------------------------------------------------------------------- + +Model::~Model() { + if (ctx_meta != nullptr) { + ggml_free(ctx_meta); + } + if (backend_buffer != nullptr) { + safe_buffer_free(backend_buffer); + } + if (ctx_derived != nullptr) { + ggml_free(ctx_derived); + } + if (derived_buffer != nullptr) { + safe_buffer_free(derived_buffer); + } + for (auto it = plan.scheduler_list.rbegin(); it != plan.scheduler_list.rend(); ++it) { + safe_backend_free(*it); + } + plan.scheduler_list.clear(); + plan.primary = nullptr; +} + +const Arch arch = { + /* .name = */ "ecapa_tdnn", + /* .load = */ load, + /* .init_context = */ nullptr, + /* .run = */ nullptr, + /* .run_batch = */ nullptr, + /* .stream_validate = */ nullptr, + /* .stream_begin = */ nullptr, + /* .stream_feed = */ nullptr, + /* .stream_finalize = */ nullptr, + /* .stream_reset = */ nullptr, + /* .accepts_ext_kind = */ nullptr, + /* .run_validate = */ nullptr, + /* .diarize = */ nullptr, + /* .langid = */ &k_langid_ops, +}; + +} // namespace transcribe::ecapa_tdnn diff --git a/src/arch/ecapa_tdnn/weights.cpp b/src/arch/ecapa_tdnn/weights.cpp new file mode 100644 index 000000000..dfcf8e931 --- /dev/null +++ b/src/arch/ecapa_tdnn/weights.cpp @@ -0,0 +1,305 @@ +// arch/ecapa_tdnn/weights.cpp - read_hparams + build_weights. + +#include "weights.h" + +#include "ggml.h" +#include "gguf.h" +#include "transcribe-log.h" +#include "transcribe-meta.h" +#include "transcribe-weights-util.h" + +#include +#include + +namespace transcribe::ecapa_tdnn { + +namespace { + +constexpr const char * kTag = "ecapa_tdnn"; + +// Required int32 array of exactly `n` entries. +transcribe_status read_required_i32_array_kv(const gguf_context * gguf, + const char * key, + size_t n, + std::vector & out) { + if (read_int32_array_kv(gguf, key, out) != KvResult::Ok) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: required KV \"%s\" missing or not an int32 array", kTag, key); + return TRANSCRIBE_ERR_GGUF; + } + if (out.size() != n) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: KV \"%s\" has %zu entries, expected %zu", kTag, key, out.size(), n); + return TRANSCRIBE_ERR_GGUF; + } + return TRANSCRIBE_OK; +} + +} // namespace + +transcribe_status read_hparams(const gguf_context * gguf, HParams & hp) { + if (gguf == nullptr) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + +#define REQ(expr) \ + do { \ + const transcribe_status _st = (expr); \ + if (_st != TRANSCRIBE_OK) { \ + return _st; \ + } \ + } while (0) + + REQ(read_required_u32_kv(gguf, "stt.frontend.sample_rate", kTag, hp.sample_rate)); + REQ(read_required_u32_kv(gguf, "stt.frontend.n_fft", kTag, hp.mel_n_fft)); + REQ(read_required_u32_kv(gguf, "stt.frontend.hop_length", kTag, hp.mel_hop)); + REQ(read_required_u32_kv(gguf, "stt.frontend.win_length", kTag, hp.mel_win)); + REQ(read_required_u32_kv(gguf, "stt.frontend.num_mels", kTag, hp.mel_n_mels)); + REQ(read_required_string_kv(gguf, "stt.frontend.window", kTag, hp.mel_window)); + REQ(read_required_string_kv(gguf, "stt.frontend.pad_mode", kTag, hp.mel_pad_mode)); + REQ(read_required_f32_kv(gguf, "stt.frontend.log_clamp_min", kTag, hp.mel_log_floor)); + REQ(read_required_f32_kv(gguf, "stt.frontend.top_db", kTag, hp.mel_top_db)); + REQ(read_required_string_kv(gguf, "stt.frontend.normalize", kTag, hp.mel_normalize)); + + REQ(read_required_i32_array_kv(gguf, "stt.ecapa_tdnn.channels", kNumStages, hp.channels)); + REQ(read_required_i32_array_kv(gguf, "stt.ecapa_tdnn.kernel_sizes", kNumStages, hp.kernel_sizes)); + REQ(read_required_i32_array_kv(gguf, "stt.ecapa_tdnn.dilations", kNumStages, hp.dilations)); + REQ(read_required_u32_kv(gguf, "stt.ecapa_tdnn.res2net_scale", kTag, hp.res2net_scale)); + REQ(read_required_u32_kv(gguf, "stt.ecapa_tdnn.se_channels", kTag, hp.se_channels)); + REQ(read_required_u32_kv(gguf, "stt.ecapa_tdnn.attention_channels", kTag, hp.attention_channels)); + REQ(read_required_f32_kv(gguf, "stt.ecapa_tdnn.asp_eps", kTag, hp.asp_eps)); + REQ(read_required_u32_kv(gguf, "stt.ecapa_tdnn.embedding_dim", kTag, hp.embedding_dim)); + REQ(read_required_u32_kv(gguf, "stt.ecapa_tdnn.classifier_hidden", kTag, hp.classifier_hidden)); + REQ(read_required_f32_kv(gguf, "stt.ecapa_tdnn.classifier_leaky_slope", kTag, hp.leaky_slope)); + +#undef REQ + + if (hp.sample_rate != 16000) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: stt.frontend.sample_rate is %d; only 16000 is supported", kTag, + hp.sample_rate); + return TRANSCRIBE_ERR_GGUF; + } + if (hp.mel_n_fft <= 0 || hp.mel_hop <= 0 || hp.mel_win <= 0 || hp.mel_n_mels <= 0) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, + "%s: frontend dimensions must be positive (n_fft=%d hop_length=%d win_length=%d num_mels=%d)", kTag, + hp.mel_n_fft, hp.mel_hop, hp.mel_win, hp.mel_n_mels); + return TRANSCRIBE_ERR_GGUF; + } + if (hp.mel_win > hp.mel_n_fft) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: frontend win_length (%d) > n_fft (%d)", kTag, hp.mel_win, + hp.mel_n_fft); + return TRANSCRIBE_ERR_GGUF; + } + // MelFrontend falls back silently on unknown values. + if (hp.mel_window != "hamming_periodic" || hp.mel_pad_mode != "constant" || hp.mel_normalize != "sentence_mean") { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: unsupported front end (window=%s pad_mode=%s normalize=%s)", kTag, + hp.mel_window.c_str(), hp.mel_pad_mode.c_str(), hp.mel_normalize.c_str()); + return TRANSCRIBE_ERR_GGUF; + } + + for (int i = 0; i < kNumStages; ++i) { + const int32_t k = hp.kernel_sizes[static_cast(i)]; + const int32_t d = hp.dilations[static_cast(i)]; + if (k <= 0 || k % 2 == 0 || d <= 0) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: stage %d needs an odd kernel and positive dilation (k=%d d=%d)", + kTag, i, k, d); + return TRANSCRIBE_ERR_GGUF; + } + } + if (hp.kernel_sizes[kNumStages - 1] != 1 || hp.dilations[kNumStages - 1] != 1) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: MFA stage must be k=1 d=1, got k=%d d=%d", kTag, + hp.kernel_sizes[kNumStages - 1], hp.dilations[kNumStages - 1]); + return TRANSCRIBE_ERR_GGUF; + } + // SERes2Net blocks have no shortcut projection: every block width matches. + for (int i = 1; i < kNumStages - 1; ++i) { + if (hp.channels[static_cast(i)] != hp.channels[0]) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: channels[%d] = %d differs from channels[0] = %d", kTag, i, + hp.channels[static_cast(i)], hp.channels[0]); + return TRANSCRIBE_ERR_GGUF; + } + } + if (hp.channels[kNumStages - 1] != (kNumSeBlocks * hp.channels[0])) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: MFA channels %d must equal %d x block channels %d", kTag, + hp.channels[kNumStages - 1], kNumSeBlocks, hp.channels[0]); + return TRANSCRIBE_ERR_GGUF; + } + if (hp.res2net_scale != kRes2NetScale || hp.channels[0] <= 0 || hp.channels[0] % hp.res2net_scale != 0) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: res2net_scale must be %d and divide channels[0] (got %d, %d)", kTag, + kRes2NetScale, hp.res2net_scale, hp.channels[0]); + return TRANSCRIBE_ERR_GGUF; + } + + return TRANSCRIBE_OK; +} + +namespace { + +using transcribe::weights::find_tensor; + +#define GET_F32(slot, name, ...) \ + do { \ + ggml_tensor * _t = find_tensor(ctx_meta, (name), { GGML_TYPE_F32 }, { __VA_ARGS__ }, kTag); \ + if (_t == nullptr) { \ + return TRANSCRIBE_ERR_GGUF; \ + } \ + (slot) = _t; \ + } while (0) + +#define GET_CONV(slot, name, ...) \ + do { \ + ggml_tensor * _t = find_tensor(ctx_meta, (name), { TRANSCRIBE_QUANT_CONV_TYPES }, { __VA_ARGS__ }, kTag); \ + if (_t == nullptr) { \ + return TRANSCRIBE_ERR_GGUF; \ + } \ + (slot) = _t; \ + } while (0) + +#define GET_LIN(slot, name, ...) \ + do { \ + ggml_tensor * _t = find_tensor(ctx_meta, (name), { TRANSCRIBE_QUANT_LINEAR_TYPES }, { __VA_ARGS__ }, kTag); \ + if (_t == nullptr) { \ + return TRANSCRIBE_ERR_GGUF; \ + } \ + (slot) = _t; \ + } while (0) + +std::string cat(const std::string & prefix, const char * suffix) { + return prefix + suffix; +} + +// A k=1 TDNNBlock: [in, out] weight, [out] bias, [out] BN scale/shift. +transcribe_status load_tdnn(ggml_context * ctx_meta, + const std::string & prefix, + int64_t in, + int64_t out, + TdnnLayer & t) { + const std::string n_w = cat(prefix, ".weight"); + const std::string n_b = cat(prefix, ".bias"); + const std::string n_scale = cat(prefix, ".bn.scale"); + const std::string n_shift = cat(prefix, ".bn.shift"); + + GET_LIN(t.w, n_w.c_str(), in, out); + GET_F32(t.b, n_b.c_str(), out); + GET_F32(t.bn.scale, n_scale.c_str(), out); + GET_F32(t.bn.shift, n_shift.c_str(), out); + return TRANSCRIBE_OK; +} + +// One dilated k-tap Res2Net sub-convolution, tap-major ne=[c, c, k]. +transcribe_status load_res2_sub(ggml_context * ctx_meta, + const std::string & prefix, + int64_t c, + int64_t k, + Res2Sub & s) { + const std::string n_w = cat(prefix, ".conv.weight"); + const std::string n_b = cat(prefix, ".conv.bias"); + const std::string n_scale = cat(prefix, ".bn.scale"); + const std::string n_shift = cat(prefix, ".bn.shift"); + + GET_CONV(s.w, n_w.c_str(), c, c, k); + GET_F32(s.b, n_b.c_str(), c); + GET_F32(s.bn.scale, n_scale.c_str(), c); + GET_F32(s.bn.shift, n_shift.c_str(), c); + return TRANSCRIBE_OK; +} + +transcribe_status load_se(ggml_context * ctx_meta, const std::string & prefix, int64_t c, int64_t se, SeBlock & s) { + const std::string n_c1_w = cat(prefix, ".c1.weight"); + const std::string n_c1_b = cat(prefix, ".c1.bias"); + const std::string n_c2_w = cat(prefix, ".c2.weight"); + const std::string n_c2_b = cat(prefix, ".c2.bias"); + + GET_LIN(s.c1_w, n_c1_w.c_str(), c, se); + GET_F32(s.c1_b, n_c1_b.c_str(), se); + GET_LIN(s.c2_w, n_c2_w.c_str(), se, c); + GET_F32(s.c2_b, n_c2_b.c_str(), c); + return TRANSCRIBE_OK; +} + +} // namespace + +transcribe_status build_weights(ggml_context * ctx_meta, const HParams & hp, Weights & w) { + if (ctx_meta == nullptr) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + + const int64_t n_mels = hp.mel_n_mels; + const int64_t n_freq = hp.n_freq(); + const int64_t C = hp.c_block(); + const int64_t Cm = hp.c_mfa(); + const int64_t chunk = hp.c_chunk(); + const int64_t se = hp.se_channels; + const int64_t att = hp.attention_channels; + const int64_t emb = hp.embedding_dim; + const int64_t hid = hp.classifier_hidden; + const int64_t n_lab = hp.n_labels; + const int64_t k_blk0 = hp.kernel_sizes[0]; + + // ---- front end ------------------------------------------------------ + // Mel-major: each filter's n_freq weights contiguous, so ne[0] = n_freq. + GET_F32(w.mel_filters, "frontend.mel_filterbank", n_freq, n_mels); + + // ---- stage 0 -------------------------------------------------------- + GET_CONV(w.blk0_w, "blk.0.conv.weight", n_mels, C, k_blk0); + GET_F32(w.blk0_b, "blk.0.conv.bias", C); + GET_F32(w.blk0_bn.scale, "blk.0.bn.scale", C); + GET_F32(w.blk0_bn.shift, "blk.0.bn.shift", C); + + // ---- stages 1..3 (SERes2Net) ---------------------------------------- + for (int i = 0; i < kNumSeBlocks; ++i) { + SeRes2NetBlock & blk = w.blocks[i]; + const std::string base = "blk." + std::to_string(i + 1); + const int64_t k_res2 = hp.kernel_sizes[static_cast(i + 1)]; + + if (auto st = load_tdnn(ctx_meta, base + ".tdnn1", C, C, blk.tdnn1); st != TRANSCRIBE_OK) { + return st; + } + for (int j = 0; j < kRes2NetSubs; ++j) { + const std::string p = base + ".res2." + std::to_string(j); + if (auto st = load_res2_sub(ctx_meta, p, chunk, k_res2, blk.res2[j]); st != TRANSCRIBE_OK) { + return st; + } + } + if (auto st = load_tdnn(ctx_meta, base + ".tdnn2", C, C, blk.tdnn2); st != TRANSCRIBE_OK) { + return st; + } + if (auto st = load_se(ctx_meta, base + ".se", C, se, blk.se); st != TRANSCRIBE_OK) { + return st; + } + } + + // ---- MFA ------------------------------------------------------------- + GET_LIN(w.mfa_w[0], "mfa.w1.weight", C, Cm); + GET_LIN(w.mfa_w[1], "mfa.w2.weight", C, Cm); + GET_LIN(w.mfa_w[2], "mfa.w3.weight", C, Cm); + GET_F32(w.mfa_b, "mfa.bias", Cm); + GET_F32(w.mfa_bn.scale, "mfa.bn.scale", Cm); + GET_F32(w.mfa_bn.shift, "mfa.bn.shift", Cm); + + // ---- attentive statistics pooling ------------------------------------- + GET_LIN(w.asp_wx, "asp.tdnn.x.weight", Cm, att); + GET_LIN(w.asp_wm, "asp.tdnn.mean.weight", Cm, att); + GET_LIN(w.asp_ws, "asp.tdnn.std.weight", Cm, att); + GET_F32(w.asp_b, "asp.tdnn.bias", att); + GET_F32(w.asp_bn.scale, "asp.tdnn.bn.scale", att); + GET_F32(w.asp_bn.shift, "asp.tdnn.bn.shift", att); + GET_LIN(w.asp_attn_w, "asp.attn.weight", att, Cm); + GET_F32(w.asp_attn_b, "asp.attn.bias", Cm); + + // ---- embedding head --------------------------------------------------- + GET_LIN(w.fc_w, "fc.weight", 2 * Cm, emb); + GET_F32(w.fc_b, "fc.bias", emb); + + // ---- classifier ------------------------------------------------------- + GET_LIN(w.cls_l1_w, "cls.l1.weight", emb, hid); + GET_F32(w.cls_l1_b, "cls.l1.bias", hid); + GET_LIN(w.cls_out_w, "cls.out.weight", hid, n_lab); + GET_F32(w.cls_out_b, "cls.out.bias", n_lab); + + return TRANSCRIBE_OK; +} + +#undef GET_F32 +#undef GET_CONV +#undef GET_LIN + +} // namespace transcribe::ecapa_tdnn diff --git a/src/arch/ecapa_tdnn/weights.h b/src/arch/ecapa_tdnn/weights.h new file mode 100644 index 000000000..3d1a141ee --- /dev/null +++ b/src/arch/ecapa_tdnn/weights.h @@ -0,0 +1,223 @@ +// arch/ecapa_tdnn/weights.h - ECAPA-TDNN hyperparameters, tensor catalogue, +// and the per-instance weight slots. +// +// INTERNAL to src/arch/ecapa_tdnn/. +// +// The names and shapes below are the loader contract with +// scripts/convert-ecapa_tdnn.py. ggml `ne` is fast-to-slow, so a PyTorch +// Linear [OC, IC] lands as ne = [IC, OC] and a PyTorch Conv1d [OC, IC, K] is +// stored TAP-MAJOR as ne = [IC, OC, K] (numpy [K, OC, IC]) so every tap is a +// contiguous [IC, OC] matrix. Names follow the tools/transcribe-quantize +// rules: `.bias` and `.bn.` stay F32, `.conv.weight` (k>1 kernels) stays +// F32 / F16, every other `.weight` is a quantizable matmul operand. +// +// frontend.mel_filterbank F32 ne=[n_freq, n_mels] mel-major +// +// blk.0.conv.weight F32/F16 ne=[n_mels, C, K0] +// blk.0.conv.bias F32 [C] +// blk.0.bn.scale / .shift F32 [C] +// +// blk.{1,2,3}.tdnn1.weight quant ne=[C, C] +// blk.{1,2,3}.tdnn1.bias F32 [C] +// blk.{1,2,3}.tdnn1.bn.scale / .shift [C] +// blk.{1,2,3}.res2.{0..S-2}.conv.weight F32/F16 ne=[C/S, C/S, Ki] +// blk.{1,2,3}.res2.{0..S-2}.conv.bias F32 [C/S] +// blk.{1,2,3}.res2.{0..S-2}.bn.scale / .shift [C/S] +// blk.{1,2,3}.tdnn2.weight quant ne=[C, C] +// blk.{1,2,3}.tdnn2.bias F32 [C] +// blk.{1,2,3}.tdnn2.bn.scale / .shift [C] +// blk.{1,2,3}.se.c1.weight quant ne=[C, se] .bias [se] +// blk.{1,2,3}.se.c2.weight quant ne=[se, C] .bias [C] +// +// mfa.w{1,2,3}.weight quant ne=[C, Cm] (Cm = 3C) +// mfa.bias F32 [Cm] +// mfa.bn.scale / .shift F32 [Cm] +// +// asp.tdnn.{x,mean,std}.weight quant ne=[Cm, att] +// asp.tdnn.bias F32 [att] +// asp.tdnn.bn.scale / .shift F32 [att] +// asp.attn.weight quant ne=[att, Cm] .bias [Cm] +// +// fc.weight quant ne=[2*Cm, emb] .bias [emb] (asp_bn folded in) +// cls.l1.weight quant ne=[emb, hid] .bias [hid] (cls.bn0 folded in) +// cls.out.weight quant ne=[hid, n_lab] .bias [n_lab] (cls.bn1 folded in) +// +// "quant" is TRANSCRIBE_QUANT_LINEAR_TYPES; k>1 kernels are +// TRANSCRIBE_QUANT_CONV_TYPES. Every dimension above is derived from HParams +// so the tiny test fixture and the real checkpoint share one path. + +#pragma once + +#include "transcribe.h" + +#include +#include +#include + +struct gguf_context; +struct ggml_context; +struct ggml_tensor; + +namespace transcribe::ecapa_tdnn { + +// Number of TDNN stages described by stt.ecapa_tdnn.channels: four embedding +// blocks (one plain TDNNBlock + three SERes2Net) plus the MFA stage. +constexpr int kNumStages = 5; + +// SERes2Net blocks, i.e. stages 1..3. +constexpr int kNumSeBlocks = 3; + +// Res2Net scale this implementation supports (fixes the weight-slot count). +constexpr int kRes2NetScale = 8; + +// Sub-convolutions inside one Res2Net block: the first chunk is passed +// through unchanged, so there are scale - 1 of them. +constexpr int kRes2NetSubs = kRes2NetScale - 1; + +// Every stt.* KV the family reads, plus the label count. +struct HParams { + int32_t sample_rate = 0; + + // Front end (stt.frontend.*). + int32_t mel_n_fft = 0; + int32_t mel_hop = 0; + int32_t mel_win = 0; + int32_t mel_n_mels = 0; + std::string mel_window; // "hamming_periodic" + std::string mel_pad_mode; // "constant" + float mel_log_floor = 0.0f; + float mel_top_db = 0.0f; + std::string mel_normalize; // "sentence_mean" + + // Embedding network (stt.ecapa_tdnn.*). + std::vector channels; // kNumStages entries + std::vector kernel_sizes; // kNumStages entries + std::vector dilations; // kNumStages entries + int32_t res2net_scale = 0; + int32_t se_channels = 0; + int32_t attention_channels = 0; + float asp_eps = 0.0f; + int32_t embedding_dim = 0; + + // Classifier (stt.ecapa_tdnn.classifier_*). + int32_t classifier_hidden = 0; + float leaky_slope = 0.0f; + + // Set by load() from the label table. + int32_t n_labels = 0; + + // ---- derived ------------------------------------------------------ + int32_t c_block() const { return channels.empty() ? 0 : channels[0]; } + + int32_t c_mfa() const { return channels.size() < kNumStages ? 0 : channels[kNumStages - 1]; } + + int32_t c_chunk() const { return res2net_scale > 0 ? c_block() / res2net_scale : 0; } + + int32_t n_freq() const { return mel_n_fft / 2 + 1; } + + // Inner width of the stage-0 im2col matmul: K0 * n_mels rounded up to a + // multiple of 32 (every CPU tinyBLAS tile width divides it), zero padded. + int32_t blk0_cols() const { + const int32_t n = kernel_sizes.empty() ? 0 : kernel_sizes[0] * mel_n_mels; + return (n + 31) / 32 * 32; + } + + // Reflect padding on each side of stage `i`: d * (k - 1) / 2, the + // "same"-padding width SpeechBrain's Conv1d uses. + int32_t pad(int i) const { + return dilations[static_cast(i)] * (kernel_sizes[static_cast(i)] - 1) / 2; + } +}; + +// One TDNNBlock's BatchNorm, stored as an affine map (the converter +// cannot fold it into the conv: TDNNBlock is conv -> ReLU -> BN). +struct BnAffine { + ggml_tensor * scale = nullptr; + ggml_tensor * shift = nullptr; +}; + +// A k=1 TDNNBlock: 1x1 conv -> ReLU -> BN. +struct TdnnLayer { + ggml_tensor * w = nullptr; + ggml_tensor * b = nullptr; + BnAffine bn; +}; + +// One of the scale-1 dilated 3-tap convolutions inside a Res2Net block. +struct Res2Sub { + ggml_tensor * w = nullptr; // ne=[c_chunk, c_chunk, K] + ggml_tensor * b = nullptr; // [c_chunk] + BnAffine bn; +}; + +// Squeeze-and-excitation: mean over T -> 1x1 -> ReLU -> 1x1 -> sigmoid -> scale. +struct SeBlock { + ggml_tensor * c1_w = nullptr; // ne=[C, se] + ggml_tensor * c1_b = nullptr; // [se] + ggml_tensor * c2_w = nullptr; // ne=[se, C] + ggml_tensor * c2_b = nullptr; // [C] +}; + +struct SeRes2NetBlock { + TdnnLayer tdnn1; + Res2Sub res2[kRes2NetSubs]; + TdnnLayer tdnn2; + SeBlock se; +}; + +struct Weights { + // Front end. Read back to host memory at load time and handed to the + // MelFrontend; never touched by the graph. + ggml_tensor * mel_filters = nullptr; // ne=[n_freq, n_mels] + + // Stage 0: plain TDNNBlock, n_mels -> C, k = kernel_sizes[0]. + ggml_tensor * blk0_w = nullptr; // ne=[n_mels, C, K0] + // Derived at load (model.cpp, not in the GGUF): blk0_w regrouped as one + // [K0 * n_mels, C] matrix, zero-padded on the inner axis to + // HParams::blk0_cols(), so stage 0 is a single aligned matmul against + // the host-built im2col of the mel. Same type as blk0_w. + ggml_tensor * blk0_w_im2col = nullptr; // ne=[blk0_cols, C] + ggml_tensor * blk0_b = nullptr; // [C] + BnAffine blk0_bn; + + // Stages 1..3. + SeRes2NetBlock blocks[kNumSeBlocks]; + + // Multi-layer feature aggregation. The 3C -> 3C 1x1 conv is split into + // three C -> 3C blocks applied to the three SERes2Net outputs and summed, + // which avoids materialising the concat. + ggml_tensor * mfa_w[kNumSeBlocks] = { nullptr, nullptr, nullptr }; + ggml_tensor * mfa_b = nullptr; // [Cm] + BnAffine mfa_bn; + + // Attentive statistics pooling. asp.tdnn's 3*Cm -> att weight is split + // into the x / mean / std blocks. + ggml_tensor * asp_wx = nullptr; // ne=[Cm, att] + ggml_tensor * asp_wm = nullptr; + ggml_tensor * asp_ws = nullptr; + ggml_tensor * asp_b = nullptr; // [att] + BnAffine asp_bn; + ggml_tensor * asp_attn_w = nullptr; // ne=[att, Cm] + ggml_tensor * asp_attn_b = nullptr; // [Cm] + + // Embedding head (asp_bn folded in). + ggml_tensor * fc_w = nullptr; // ne=[2*Cm, emb] + ggml_tensor * fc_b = nullptr; // [emb] + + // Classifier (cls.bn0 / cls.bn1 folded forward; the LeakyReLUs still run). + ggml_tensor * cls_l1_w = nullptr; // ne=[emb, hid] + ggml_tensor * cls_l1_b = nullptr; // [hid] + ggml_tensor * cls_out_w = nullptr; // ne=[hid, n_labels] + ggml_tensor * cls_out_b = nullptr; // [n_labels] +}; + +// Read every stt.* KV the family needs and reject shapes the graph cannot +// build. Does not set n_labels. +transcribe_status read_hparams(const gguf_context * gguf, HParams & hp); + +// Bind every tensor in the catalogue above to a borrowed pointer in +// `ctx_meta`, validating type and shape against `hp`. On failure the partly +// built `w` is indeterminate and the caller must discard the model. +transcribe_status build_weights(ggml_context * ctx_meta, const HParams & hp, Weights & w); + +} // namespace transcribe::ecapa_tdnn diff --git a/src/transcribe-arch.cpp b/src/transcribe-arch.cpp index 0840408b4..4487cd98a 100644 --- a/src/transcribe-arch.cpp +++ b/src/transcribe-arch.cpp @@ -90,6 +90,10 @@ namespace sortformer { extern const Arch arch; } +namespace ecapa_tdnn { +extern const Arch arch; +} + const Arch * find_arch(const char * name) { if (name == nullptr) { return nullptr; @@ -99,7 +103,7 @@ const Arch * find_arch(const char * name) { ¶keet::arch, &cohere::arch, &canary::arch, &qwen3_asr::arch, &voxtral::arch, &voxtral_realtime::arch, &canary_qwen::arch, &whisper::arch, &moonshine::arch, &moonshine_streaming::arch, &sensevoice::arch, &funasr_nano::arch, &gigaam::arch, &granite::arch, &granite_nar::arch, - &medasr::arch, &moss::arch, &sortformer::arch, &granite5_ctc::arch, + &medasr::arch, &moss::arch, &sortformer::arch, &granite5_ctc::arch, &ecapa_tdnn::arch, }; constexpr size_t k_n = sizeof(k_archs) / sizeof(k_archs[0]); @@ -123,12 +127,13 @@ transcribe_status resolve_roles(transcribe_model * model) { const char * name = arch.name != nullptr ? arch.name : "(unknown)"; const bool has_asr = arch.init_context != nullptr && arch.run != nullptr; const bool has_diarize = arch.diarize != nullptr; + const bool has_langid = arch.langid != nullptr; if (model->roles == 0 && has_asr) { model->roles = TRANSCRIBE_ROLE_ASR; } - const uint32_t known = TRANSCRIBE_ROLE_ASR | TRANSCRIBE_ROLE_DIARIZE; + const uint32_t known = TRANSCRIBE_ROLE_ASR | TRANSCRIBE_ROLE_DIARIZE | TRANSCRIBE_ROLE_LANGID; const char * why = nullptr; if (model->roles == 0) { why = "serves no role"; @@ -138,6 +143,8 @@ transcribe_status resolve_roles(transcribe_model * model) { why = "sets the ASR role without init_context / run hooks"; } else if ((model->roles & TRANSCRIBE_ROLE_DIARIZE) != 0 && !has_diarize) { why = "sets the DIARIZE role without a diarize ops table"; + } else if ((model->roles & TRANSCRIBE_ROLE_LANGID) != 0 && !has_langid) { + why = "sets the LANGID role without a langid ops table"; } if (why != nullptr) { log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "transcribe_model_load_file: arch '%s' %s (roles 0x%x)", name, why, diff --git a/src/transcribe-arch.h b/src/transcribe-arch.h index 6a6e0e0a0..6d06d53c2 100644 --- a/src/transcribe-arch.h +++ b/src/transcribe-arch.h @@ -15,6 +15,7 @@ namespace transcribe { class Loader; struct DiarizeOps; +struct LangidOps; // Per-family trait. Function pointers may be null when an entry point is // not yet implemented; the central dispatch converts null entries into @@ -133,6 +134,7 @@ struct Arch { // Ops tables for the non-ASR roles; nullptr = role not implemented. const DiarizeOps * diarize; + const LangidOps * langid; }; // Look up an architecture by name. Returns nullptr if no registered diff --git a/src/transcribe-langid.cpp b/src/transcribe-langid.cpp new file mode 100644 index 000000000..a90a94192 --- /dev/null +++ b/src/transcribe-langid.cpp @@ -0,0 +1,442 @@ +// transcribe-langid.cpp - LANGID role C ABI (include/transcribe/langid.h): +// label table, allowed set, crop, softmax, ranking and top-k. + +#include "transcribe-langid.h" + +#include "transcribe-abi.h" +#include "transcribe-api-guard.h" +#include "transcribe-arch.h" +#include "transcribe-log.h" +#include "transcribe-model.h" + +#include +#include +#include +#include +#include + +using transcribe::api_guard_status; +using transcribe::api_guard_void; +using transcribe::check_input_struct_size; +using transcribe::check_struct_size; +using transcribe::copy_out_prefix; +using transcribe::LangidCandidateEntry; +using transcribe::LangidLabels; + +namespace { + +constexpr size_t k_min_info_size = TRANSCRIBE_FIELD_END(transcribe_langid_info, min_audio_ms); +constexpr size_t k_min_session_params_size = TRANSCRIBE_FIELD_END(transcribe_langid_session_params, max_audio_ms); +constexpr size_t k_min_params_size = TRANSCRIBE_FIELD_END(transcribe_langid_params, top_k); +constexpr size_t k_min_result_size = TRANSCRIBE_FIELD_END(transcribe_langid_result, audio_ms); +constexpr size_t k_min_candidate_size = TRANSCRIBE_FIELD_END(transcribe_langid_candidate, logit); +constexpr size_t k_min_timings_size = TRANSCRIBE_FIELD_END(transcribe_timings, decode_ms); + +constexpr int64_t k_samples_per_ms = 16; // 16 kHz + +const transcribe::LangidOps * langid_ops(const transcribe_model * model) { + return model != nullptr && (model->roles & TRANSCRIBE_ROLE_LANGID) != 0 ? model->arch->langid : nullptr; +} + +const LangidLabels * langid_labels(const transcribe_model * model) { + const transcribe::LangidOps * ops = langid_ops(model); + return ops != nullptr ? &ops->labels(model) : nullptr; +} + +int32_t label_index(const LangidLabels & labels, const char * code) { + if (code == nullptr) { + return -1; + } + const auto it = labels.index.find(std::string_view(code)); + return it != labels.index.end() ? it->second : -1; +} + +// Resolve params->allowed into a per-label mask. Only NULL with +// n_allowed == 0 means "all". +transcribe_status build_allowed_mask(const LangidLabels & labels, + const transcribe_langid_params * params, + std::vector & mask, + int32_t & n_allowed) { + const size_t n_labels = labels.codes.size(); + if (params->allowed == nullptr) { + if (params->n_allowed != 0) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + mask.assign(n_labels, 1); + n_allowed = static_cast(n_labels); + return TRANSCRIBE_OK; + } + if (params->n_allowed <= 0) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + mask.assign(n_labels, 0); + for (int32_t i = 0; i < params->n_allowed; ++i) { + if (params->allowed[i] == nullptr) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + } + n_allowed = 0; + for (int32_t i = 0; i < params->n_allowed; ++i) { + const int32_t idx = label_index(labels, params->allowed[i]); + if (idx < 0) { + transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_WARN, "transcribe_langid_run: unknown language '%s'", + params->allowed[i]); + return TRANSCRIBE_ERR_UNSUPPORTED_LANGUAGE; + } + if (mask[static_cast(idx)] == 0) { + mask[static_cast(idx)] = 1; + ++n_allowed; + } + } + return TRANSCRIBE_OK; +} + +// Softmax over the entries with mask[i] != 0 (every entry when mask is +// NULL); masked-out entries get 0. Sums in double. +void masked_softmax(const std::vector & logits, const uint8_t * mask, std::vector & out) { + const size_t n = logits.size(); + float mx = -INFINITY; + for (size_t i = 0; i < n; ++i) { + if (mask == nullptr || mask[i] != 0) { + mx = std::max(mx, logits[i]); + } + } + out.assign(n, 0.0f); + double sum = 0.0; + for (size_t i = 0; i < n; ++i) { + if (mask == nullptr || mask[i] != 0) { + const double e = std::exp(static_cast(logits[i]) - static_cast(mx)); + out[i] = static_cast(e); + sum += e; + } + } + for (size_t i = 0; i < n; ++i) { + out[i] = static_cast(static_cast(out[i]) / sum); + } +} + +} // namespace + +transcribe_status transcribe::build_langid_labels(std::vector codes, + std::vector names, + const std::vector & alias_specs, + const char * tag, + LangidLabels & out) { + auto fail = [tag](const char * what, const std::string & item) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: label table: %s '%s'", tag, what, item.c_str()); + return TRANSCRIBE_ERR_GGUF; + }; + if (codes.empty() || codes.size() != names.size()) { + log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "%s: label table: %zu codes vs %zu names", tag, codes.size(), names.size()); + return TRANSCRIBE_ERR_GGUF; + } + LangidLabels staged; + for (size_t i = 0; i < codes.size(); ++i) { + if (codes[i].empty()) { + return fail("empty code at index", std::to_string(i)); + } + if (!staged.index.emplace(codes[i], static_cast(i)).second) { + return fail("duplicate code", codes[i]); + } + } + for (const std::string & spec : alias_specs) { + const size_t eq = spec.find('='); + if (eq == std::string::npos || eq == 0 || eq + 1 == spec.size()) { + return fail("malformed alias (want alias=code)", spec); + } + std::string alias = spec.substr(0, eq); + std::string target = spec.substr(eq + 1); + const auto it = staged.index.find(target); + if (it == staged.index.end() || codes[static_cast(it->second)] != target) { + return fail("alias names an unknown code", spec); + } + if (!staged.index.emplace(std::move(alias), it->second).second) { + return fail("alias collides with a code or another alias", spec); + } + } + staged.codes = std::move(codes); + staged.names = std::move(names); + out = std::move(staged); + return TRANSCRIBE_OK; +} + +extern "C" void transcribe_langid_info_init(struct transcribe_langid_info * p) { + transcribe::init_sized(p); +} + +extern "C" void transcribe_langid_session_params_init(struct transcribe_langid_session_params * p) { + transcribe::init_sized(p); +} + +extern "C" void transcribe_langid_params_init(struct transcribe_langid_params * p) { + transcribe::init_sized(p); +} + +extern "C" void transcribe_langid_result_init(struct transcribe_langid_result * p) { + transcribe::init_sized(p); +} + +extern "C" void transcribe_langid_candidate_init(struct transcribe_langid_candidate * p) { + transcribe::init_sized(p); +} + +static transcribe_status langid_get_info_impl(const transcribe_model * model, transcribe_langid_info * out) { + if (model == nullptr || out == nullptr) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + if (const auto st = check_struct_size(out->struct_size, k_min_info_size); st != TRANSCRIBE_OK) { + return st; + } + const LangidLabels * labels = langid_labels(model); + if (labels == nullptr) { + return TRANSCRIBE_ERR_UNSUPPORTED_ROLE; + } + transcribe_langid_info staged{}; + staged.struct_size = out->struct_size; + staged.sample_rate = 16000; + staged.n_labels = static_cast(labels->codes.size()); + staged.min_audio_ms = transcribe::k_langid_min_audio_ms; + copy_out_prefix(out, &staged, out->struct_size, sizeof(staged)); + return TRANSCRIBE_OK; +} + +static transcribe_status langid_session_init_impl(transcribe_model * model, + const transcribe_langid_session_params * params, + transcribe_langid_session ** out) { + if (out == nullptr) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + *out = nullptr; + if (model == nullptr) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + const transcribe::LangidOps * ops = langid_ops(model); + if (ops == nullptr) { + return TRANSCRIBE_ERR_UNSUPPORTED_ROLE; + } + transcribe_langid_session_params defaults; + transcribe_langid_session_params_init(&defaults); + if (params == nullptr) { + params = &defaults; + } + if (const auto st = check_input_struct_size(params->struct_size, k_min_session_params_size); st != TRANSCRIBE_OK) { + return st; + } + if (params->n_threads < 0) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + const int32_t max_audio_ms = + params->max_audio_ms == 0 ? transcribe::k_langid_default_audio_ms : params->max_audio_ms; + if (max_audio_ms < transcribe::k_langid_min_audio_ms) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + *out = ops->new_session(); + (*out)->model = model; + (*out)->n_threads = params->n_threads; + (*out)->max_audio_ms = max_audio_ms; + return TRANSCRIBE_OK; +} + +static transcribe_status langid_run_impl(transcribe_langid_session * session, + const float * pcm, + int n_samples, + const transcribe_langid_params * params) { + // Everything up to the commit point leaves the previous result intact. + if (session == nullptr || pcm == nullptr || n_samples <= 0 || !transcribe::pcm_is_finite(pcm, n_samples)) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + transcribe_langid_params defaults; + transcribe_langid_params_init(&defaults); + if (params == nullptr) { + params = &defaults; + } + if (const auto st = check_input_struct_size(params->struct_size, k_min_params_size); st != TRANSCRIBE_OK) { + return st; + } + if (params->top_k < 0) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + const transcribe::LangidOps * ops = langid_ops(session->model); + const LangidLabels & labels = ops->labels(session->model); + std::vector mask; + int32_t n_allowed = 0; + if (const auto st = build_allowed_mask(labels, params, mask, n_allowed); st != TRANSCRIBE_OK) { + return st; + } + // Score the last max_audio_ms; the minimum applies to what is scored. + const int64_t max_samples = static_cast(session->max_audio_ms) * k_samples_per_ms; + const int64_t n_used = std::min(n_samples, max_samples); + if (n_used < static_cast(transcribe::k_langid_min_audio_ms) * k_samples_per_ms) { + return TRANSCRIBE_ERR_INPUT_TOO_SHORT; + } + const float * pcm_used = pcm + (n_samples - n_used); + + session->candidates.clear(); + session->n_allowed = 0; + session->allowed_mass = 0.0f; + session->audio_ms = 0; + session->t_mel_us = 0; + session->t_encode_us = 0; + session->t_decode_us = 0; + transcribe::ScratchReleaseGuard release{ session, true }; + + std::vector logits; + if (const auto st = ops->run(session, pcm_used, static_cast(n_used), logits); st != TRANSCRIBE_OK) { + return st; + } + if (logits.size() != labels.codes.size()) { + transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "transcribe_langid_run: %zu logits for %zu labels", + logits.size(), labels.codes.size()); + return TRANSCRIBE_ERR_BACKEND; + } + for (const float v : logits) { + if (!std::isfinite(v)) { + transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "transcribe_langid_run: non-finite logit"); + return TRANSCRIBE_ERR_BACKEND; + } + } + + const bool restricted = n_allowed != static_cast(labels.codes.size()); + std::vector p_open; + std::vector p_allowed; + masked_softmax(logits, nullptr, p_open); + masked_softmax(logits, mask.data(), p_allowed); + double mass = 0.0; + for (size_t i = 0; i < logits.size(); ++i) { + if (mask[i] != 0) { + mass += p_open[i]; + } + } + + std::vector ranked; + ranked.reserve(static_cast(n_allowed)); + for (size_t i = 0; i < logits.size(); ++i) { + if (mask[i] != 0) { + ranked.push_back({ static_cast(i), p_allowed[i], logits[i] }); + } + } + std::stable_sort(ranked.begin(), ranked.end(), + [](const LangidCandidateEntry & a, const LangidCandidateEntry & b) { return a.p > b.p; }); + if (params->top_k > 0 && static_cast(params->top_k) < ranked.size()) { + ranked.resize(static_cast(params->top_k)); + } + + session->candidates.swap(ranked); + session->n_allowed = n_allowed; + session->allowed_mass = restricted ? static_cast(mass) : 1.0f; + session->audio_ms = n_used / k_samples_per_ms; + return TRANSCRIBE_OK; +} + +extern "C" const char * transcribe_langid_label_code(const struct transcribe_model * model, int32_t i) { + const LangidLabels * labels = langid_labels(model); + if (labels == nullptr || i < 0 || static_cast(i) >= labels->codes.size()) { + return nullptr; + } + return labels->codes[static_cast(i)].c_str(); +} + +extern "C" const char * transcribe_langid_label_name(const struct transcribe_model * model, int32_t i) { + const LangidLabels * labels = langid_labels(model); + if (labels == nullptr || i < 0 || static_cast(i) >= labels->names.size()) { + return nullptr; + } + return labels->names[static_cast(i)].c_str(); +} + +extern "C" void transcribe_langid_set_abort_callback(struct transcribe_langid_session * session, + transcribe_abort_callback cb, + void * user_data) { + if (session != nullptr) { + session->abort_cb = cb; + session->abort_userdata = user_data; + } +} + +extern "C" transcribe_status transcribe_langid_get_result(const struct transcribe_langid_session * session, + struct transcribe_langid_result * out) { + if (session == nullptr || out == nullptr) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + if (const auto st = check_struct_size(out->struct_size, k_min_result_size); st != TRANSCRIBE_OK) { + return st; + } + transcribe_langid_result staged{}; + staged.struct_size = out->struct_size; + staged.n_candidates = static_cast(session->candidates.size()); + staged.n_allowed = session->n_allowed; + staged.allowed_mass = session->allowed_mass; + staged.audio_ms = session->audio_ms; + copy_out_prefix(out, &staged, out->struct_size, sizeof(staged)); + return TRANSCRIBE_OK; +} + +extern "C" transcribe_status transcribe_langid_get_candidate(const struct transcribe_langid_session * session, + int i, + struct transcribe_langid_candidate * out) { + if (out == nullptr) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + if (const auto st = check_struct_size(out->struct_size, k_min_candidate_size); st != TRANSCRIBE_OK) { + return st; + } + transcribe_langid_candidate staged{}; + staged.struct_size = out->struct_size; + if (session != nullptr && i >= 0 && static_cast(i) < session->candidates.size()) { + const LangidCandidateEntry & c = session->candidates[static_cast(i)]; + staged.index = c.index; + staged.code = transcribe_langid_label_code(session->model, c.index); + staged.name = transcribe_langid_label_name(session->model, c.index); + staged.p = c.p; + staged.logit = c.logit; + } + copy_out_prefix(out, &staged, out->struct_size, sizeof(staged)); + return TRANSCRIBE_OK; +} + +extern "C" transcribe_status transcribe_langid_get_timings(const struct transcribe_langid_session * session, + struct transcribe_timings * out) { + if (session == nullptr || out == nullptr) { + return TRANSCRIBE_ERR_INVALID_ARG; + } + if (const auto st = check_struct_size(out->struct_size, k_min_timings_size); st != TRANSCRIBE_OK) { + return st; + } + transcribe::copy_out_timings(session->model->t_load_us, session->t_mel_us, session->t_encode_us, + session->t_decode_us, out); + return TRANSCRIBE_OK; +} + +// C ABI forwarders for the entry points that allocate, compute, transfer +// ownership, or build a lookup key; the ones above are nothrow by +// construction. + +extern "C" transcribe_status transcribe_langid_get_info(const struct transcribe_model * model, + struct transcribe_langid_info * out) { + return api_guard_status("transcribe_langid_get_info", [&] { return langid_get_info_impl(model, out); }); +} + +extern "C" int32_t transcribe_langid_label_index(const struct transcribe_model * model, const char * code_or_alias) { + return transcribe::api_guard_value("transcribe_langid_label_index", int32_t{ -1 }, [&] { + const LangidLabels * labels = langid_labels(model); + return labels != nullptr ? label_index(*labels, code_or_alias) : int32_t{ -1 }; + }); +} + +extern "C" transcribe_status transcribe_langid_session_init(struct transcribe_model * model, + const struct transcribe_langid_session_params * params, + struct transcribe_langid_session ** out) { + return api_guard_status("transcribe_langid_session_init", + [&] { return langid_session_init_impl(model, params, out); }); +} + +extern "C" void transcribe_langid_session_free(struct transcribe_langid_session * session) { + api_guard_void("transcribe_langid_session_free", [&] { delete session; }); +} + +extern "C" transcribe_status transcribe_langid_run(struct transcribe_langid_session * session, + const float * pcm, + int n_samples, + const struct transcribe_langid_params * params) { + return api_guard_status("transcribe_langid_run", [&] { return langid_run_impl(session, pcm, n_samples, params); }); +} diff --git a/src/transcribe-langid.h b/src/transcribe-langid.h new file mode 100644 index 000000000..4df6ad311 --- /dev/null +++ b/src/transcribe-langid.h @@ -0,0 +1,70 @@ +// transcribe-langid.h - internal LANGID role surface: the session base, the +// per-arch ops table, and the shared label table. +// +// Families compute logits over their label table; the role dispatcher +// (transcribe-langid.cpp) crops the input, applies the allowed set, and +// owns the softmax, ranking and top-k. + +#pragma once + +#include "transcribe-session-core.h" +#include "transcribe/langid.h" + +#include +#include +#include +#include +#include + +namespace transcribe { + +struct LangidCandidateEntry { + int32_t index = 0; + float p = 0.0f; + float logit = 0.0f; +}; + +// A model's labels. Built once at load by build_langid_labels, immutable after. +struct LangidLabels { + std::vector codes; + std::vector names; + std::map> index; // every code and alias -> label index +}; + +// Build a label table from parallel code / name lists and "alias=code" +// specs. Returns TRANSCRIBE_ERR_GGUF (and logs `tag`) when the lists are +// empty or differ in length, a code is empty or repeated, or an alias is +// malformed, repeated, collides with a code, or names an unknown code. +transcribe_status build_langid_labels(std::vector codes, + std::vector names, + const std::vector & alias_specs, + const char * tag, + LangidLabels & out); + +// Scored-audio bounds every LANGID model shares. +constexpr int32_t k_langid_min_audio_ms = 500; +constexpr int32_t k_langid_default_audio_ms = 30000; + +struct LangidOps { + const LangidLabels & (*labels)(const transcribe_model * model); + transcribe_langid_session * (*new_session)(); + // Score already validated, already cropped PCM; fill one logit per label. + // Poll session->poll_abort() and return TRANSCRIBE_ERR_ABORTED when it + // fires. Set session->t_mel_us / t_encode_us. + transcribe_status (*run)(transcribe_langid_session * session, + const float * pcm, + int n_samples, + std::vector & logits); +}; + +} // namespace transcribe + +struct transcribe_langid_session : transcribe::SessionCore { + int32_t max_audio_ms = transcribe::k_langid_default_audio_ms; + + // Last successful run's result; empty / zero otherwise. + std::vector candidates; + int32_t n_allowed = 0; + float allowed_mass = 0.0f; + int64_t audio_ms = 0; +}; diff --git a/src/transcribe-mel.cpp b/src/transcribe-mel.cpp index 0893486d6..af4e23cae 100644 --- a/src/transcribe-mel.cpp +++ b/src/transcribe-mel.cpp @@ -140,6 +140,17 @@ void build_hann_window_symmetric_padded(int win_length, int n_fft, bool periodic } } +// Periodic Hamming of length win_length, zero-padded to n_fft on both +// sides: torch.hamming_window(N) with the default periodic=True, +// 0.54 - 0.46*cos(2πk/N). SpeechBrain (ecapa_tdnn). +void build_hamming_window_periodic_padded(int win_length, int n_fft, std::vector & out) { + out.assign(n_fft, 0.0); + const int pad_each = (n_fft - win_length) / 2; + for (int k = 0; k < win_length; ++k) { + out[pad_each + k] = 0.54 - 0.46 * std::cos(2.0 * M_PI * k / static_cast(win_length)); + } +} + // In-place radix-2 Cooley-Tukey FFT, fp64. Operates on n complex // numbers stored as interleaved (re, im, re, im, ...). n must be a // power of 2. For n=512 (Parakeet) this is ~80 lines and runs in @@ -289,6 +300,8 @@ MelFrontend::MelFrontend(const MelConfig & cfg) : cfg_(cfg) { for (int i = 0; i < cfg.win_length && i < static_cast(cfg.window.size()); ++i) { window_[left_pad + i] = static_cast(cfg.window[i]); } + } else if (cfg.window_type == "hamming_periodic") { + build_hamming_window_periodic_padded(cfg.win_length, cfg.n_fft, window_); } else { const bool periodic = (cfg.window_type == "hann_periodic"); build_hann_window_symmetric_padded(cfg.win_length, cfg.n_fft, periodic, window_); @@ -302,6 +315,25 @@ MelFrontend::MelFrontend(const MelConfig & cfg) : cfg_(cfg) { static_cast(cfg.f_max), mel_fb_); } + const int n_rows = n_freq_ > 0 ? static_cast(mel_fb_.size() / static_cast(n_freq_)) : 0; + fb_lo_.assign(n_rows, 0); + fb_hi_.assign(n_rows, 0); + for (int m = 0; m < n_rows; ++m) { + const float * row = mel_fb_.data() + static_cast(m) * n_freq_; + int lo = n_freq_; + int hi = 0; + for (int k = 0; k < n_freq_; ++k) { + if (row[k] != 0.0f) { + lo = std::min(lo, k); + hi = k + 1; + } + } + if (lo < hi) { + fb_lo_[m] = lo - lo % 4; + fb_hi_[m] = hi; + } + } + // Sin/cos LUT for the mixed-radix FFT. Only the non-pow2 path // consumes it; pow2 sizes go through fft_radix2 (Linux) or vDSP // (Apple), both of which carry their own twiddle factors. @@ -480,6 +512,11 @@ transcribe_status MelFrontend::compute(const float * pcm, const bool whisper_mode = (cfg_.normalize == "per_utterance" || cfg_.normalize == "global"); + // SpeechBrain power-to-dB (normalize == "sentence_mean"): + // 10*log10(max(x, amin)) in fp64, one rounding at storage. + const bool db_mode = (cfg_.normalize == "sentence_mean"); + const double db_floor = static_cast(cfg_.log_clamp_min > 0.0f ? cfg_.log_clamp_min : 1.0e-10f); + int stft_threads = n_threads; if (stft_threads <= 0) { stft_threads = default_n_threads(); @@ -519,19 +556,25 @@ transcribe_status MelFrontend::compute(const float * pcm, } for (int m = 0; m < n_mels; ++m) { const float * fb_row = mel_fb_.data() + static_cast(m) * n_freq; + const int k_hi = fb_hi_[static_cast(m)]; double sum = 0.0; - int k = 0; - for (; k < n_freq - 3; k += 4) { + int k = fb_lo_[static_cast(m)]; + for (; k < k_hi && k < n_freq - 3; k += 4) { sum += static_cast(fb_row[k]) * static_cast(power_scratch[k]) + static_cast(fb_row[k + 1]) * static_cast(power_scratch[k + 1]) + static_cast(fb_row[k + 2]) * static_cast(power_scratch[k + 2]) + static_cast(fb_row[k + 3]) * static_cast(power_scratch[k + 3]); } - for (; k < n_freq; ++k) { + for (; k < k_hi; ++k) { sum += static_cast(fb_row[k]) * static_cast(power_scratch[k]); } float result; - if (whisper_mode) { + if (db_mode) { + if (sum < db_floor) { + sum = db_floor; + } + result = static_cast(10.0 * std::log10(sum)); + } else if (whisper_mode) { if (sum < 1.0e-10) { sum = 1.0e-10; } @@ -629,7 +672,15 @@ transcribe_status MelFrontend::compute(const float * pcm, power.data(), n_freq, 0.0f, log_mel.data(), n_frames); { const size_t total = static_cast(n_mels) * static_cast(n_frames); - if (whisper_mode) { + if (db_mode) { + for (size_t i = 0; i < total; ++i) { + double v = static_cast(log_mel[i]); + if (v < db_floor) { + v = db_floor; + } + log_mel[i] = static_cast(10.0 * std::log10(v)); + } + } else if (whisper_mode) { for (size_t i = 0; i < total; ++i) { double v = static_cast(log_mel[i]); if (v < 1.0e-10) { @@ -662,19 +713,25 @@ transcribe_status MelFrontend::compute(const float * pcm, const float * pwr = power.data() + static_cast(t) * n_freq; for (int m = 0; m < n_mels; ++m) { const float * fb_row = mel_fb_.data() + static_cast(m) * n_freq; + const int k_hi = fb_hi_[static_cast(m)]; double sum = 0.0; - int k = 0; - for (; k < n_freq - 3; k += 4) { + int k = fb_lo_[static_cast(m)]; + for (; k < k_hi && k < n_freq - 3; k += 4) { sum += static_cast(fb_row[k]) * static_cast(pwr[k]) + static_cast(fb_row[k + 1]) * static_cast(pwr[k + 1]) + static_cast(fb_row[k + 2]) * static_cast(pwr[k + 2]) + static_cast(fb_row[k + 3]) * static_cast(pwr[k + 3]); } - for (; k < n_freq; ++k) { + for (; k < k_hi; ++k) { sum += static_cast(fb_row[k]) * static_cast(pwr[k]); } float result; - if (whisper_mode) { + if (db_mode) { + if (sum < db_floor) { + sum = db_floor; + } + result = static_cast(10.0 * std::log10(sum)); + } else if (whisper_mode) { if (sum < 1.0e-10) { sum = 1.0e-10; } @@ -698,6 +755,47 @@ transcribe_status MelFrontend::compute(const float * pcm, std::vector().swap(padded_f32); std::vector().swap(window_f32); + // ---- 5s. SpeechBrain sentence-mean normalize ---- + // log_mel already holds 10*log10(max(x, amin)). Top-dB floor over time + // AND frequency (SpeechBrain clamps inside Filterbank), then subtract + // each mel bin's mean over ALL frames (InputNormalization, norm_type= + // "sentence", std_norm=False). No frame is dropped or masked: torch + // keeps the center-pad frames. fp64 accumulators, one rounding at + // storage. + if (db_mode) { + const size_t total = static_cast(n_mels) * static_cast(n_frames); + if (cfg_.top_db > 0.0f) { + double max_all = -std::numeric_limits::infinity(); + for (size_t i = 0; i < total; ++i) { + const double v = static_cast(log_mel[i]); + if (v > max_all) { + max_all = v; + } + } + const double floor_val = max_all - static_cast(cfg_.top_db); + for (size_t i = 0; i < total; ++i) { + if (static_cast(log_mel[i]) < floor_val) { + log_mel[i] = static_cast(floor_val); + } + } + } + for (int m = 0; m < n_mels; ++m) { + float * row = log_mel.data() + static_cast(m) * n_frames; + double sum = 0.0; + for (int t = 0; t < n_frames; ++t) { + sum += static_cast(row[t]); + } + const double mean = sum / static_cast(n_frames); + for (int t = 0; t < n_frames; ++t) { + row[t] = static_cast(static_cast(row[t]) - mean); + } + } + out_mel = std::move(log_mel); + out_n_mels = n_mels; + out_n_frames = n_frames; + return TRANSCRIBE_OK; + } + // ---- 5n. No-op normalize (NeMo "NA"/none) ---- // Emit raw log-mel as-is: streaming Conformer variants (e.g. // nemotron-speech-streaming-en-0.6b) bake feature normalization into diff --git a/src/transcribe-mel.h b/src/transcribe-mel.h index b547435e2..49abaef93 100644 --- a/src/transcribe-mel.h +++ b/src/transcribe-mel.h @@ -59,6 +59,9 @@ struct MelConfig { // "hann_periodic" — torch.hann_window(N, periodic=True): // cos(2*pi*k / N). Used by Whisper (and // Qwen3-ASR's Whisper frontend). + // "hamming_periodic" — torch.hamming_window(N, periodic=True): + // 0.54 - 0.46*cos(2*pi*k / N). Used by + // SpeechBrain (ecapa_tdnn). std::string window_type = "hann_symmetric"; // Normalization mode: @@ -78,6 +81,12 @@ struct MelConfig { // per-utterance maximum, so each frame is // causal/streaming-safe. Still drops the trailing // center-pad STFT frame. + // "sentence_mean" — SpeechBrain Fbank + InputNormalization(norm_type= + // "sentence", std_norm=False): power-to-dB + // 10*log10(max(x, log_clamp_min)), then the + // optional top_db floor, then subtract each mel + // bin's mean over ALL frames (no variance). No frame + // is dropped or masked. Used by ecapa_tdnn. std::string normalize = "per_feature"; // Fixed log-mel maximum for normalize == "global" (Voxtral Realtime @@ -93,8 +102,15 @@ struct MelConfig { // When > 0 AND normalize="none", emit log(max(power, log_clamp_min)) // instead of NeMo's log(power + kLogEps). LASR / MedASR use this with // log_clamp_min = 1e-5; NeMo's frontends leave it at 0.0. + // normalize="sentence_mean" uses it as the dB floor (amin) instead; + // 0.0 there falls back to 1e-10. float log_clamp_min = 0.0f; + // normalize="sentence_mean" only: per-utterance dynamic-range floor + // over time AND frequency, x = max(x, max_all(x) - top_db), applied + // in dB before the mean subtraction. <= 0 disables it. + float top_db = 0.0f; + // Optional checkpoint-provided window [win_length]. When non-empty, // used instead of computing a periodic Hann window. std::vector window; @@ -173,6 +189,11 @@ class MelFrontend { int n_freq_; // n_fft/2 + 1 std::vector window_; // [n_fft], periodic hann zero-padded std::vector mel_fb_; // [num_mels * n_freq] row-major, Slaney + // Per mel row: the scalar matmul runs over [fb_lo_, fb_hi_) only. + // fb_lo_ is rounded down to a multiple of 4 so the 4-wide groups stay + // aligned; the skipped weights are all zero, so the sum is unchanged. + std::vector fb_lo_; + std::vector fb_hi_; // sin/cos LUT for the mixed-radix FFT. Sized to n_fft so that every // recursion-level N (n_fft, n_fft/2, n_fft/4, ..., odd leaf) divides diff --git a/src/transcribe.cpp b/src/transcribe.cpp index 940a9004f..f38a428fb 100644 --- a/src/transcribe.cpp +++ b/src/transcribe.cpp @@ -29,6 +29,7 @@ #include "transcribe-model.h" #include "transcribe-path.h" #include "transcribe/diarize.h" +#include "transcribe/langid.h" #if defined(TRANSCRIBE_GGML_BACKEND_DL) && defined(_WIN32) # ifndef WIN32_LEAN_AND_MEAN @@ -112,6 +113,8 @@ extern "C" const char * transcribe_status_string(int status) { return "output repetition: decode stopped when the output began repeating itself"; case TRANSCRIBE_ERR_UNSUPPORTED_ROLE: return "model does not serve the requested role"; + case TRANSCRIBE_ERR_INPUT_TOO_SHORT: + return "input audio too short"; default: return "unknown status"; } @@ -178,6 +181,16 @@ extern "C" size_t transcribe_abi_struct_size(transcribe_abi_struct which) { return sizeof(struct transcribe_diarize_session_params); case TRANSCRIBE_ABI_DIARIZE_PARAMS: return sizeof(struct transcribe_diarize_params); + case TRANSCRIBE_ABI_LANGID_INFO: + return sizeof(struct transcribe_langid_info); + case TRANSCRIBE_ABI_LANGID_SESSION_PARAMS: + return sizeof(struct transcribe_langid_session_params); + case TRANSCRIBE_ABI_LANGID_PARAMS: + return sizeof(struct transcribe_langid_params); + case TRANSCRIBE_ABI_LANGID_RESULT: + return sizeof(struct transcribe_langid_result); + case TRANSCRIBE_ABI_LANGID_CANDIDATE: + return sizeof(struct transcribe_langid_candidate); } return 0; // unknown id: "cannot verify", never a real size } @@ -222,6 +235,16 @@ extern "C" size_t transcribe_abi_struct_align(transcribe_abi_struct which) { return alignof(struct transcribe_diarize_session_params); case TRANSCRIBE_ABI_DIARIZE_PARAMS: return alignof(struct transcribe_diarize_params); + case TRANSCRIBE_ABI_LANGID_INFO: + return alignof(struct transcribe_langid_info); + case TRANSCRIBE_ABI_LANGID_SESSION_PARAMS: + return alignof(struct transcribe_langid_session_params); + case TRANSCRIBE_ABI_LANGID_PARAMS: + return alignof(struct transcribe_langid_params); + case TRANSCRIBE_ABI_LANGID_RESULT: + return alignof(struct transcribe_langid_result); + case TRANSCRIBE_ABI_LANGID_CANDIDATE: + return alignof(struct transcribe_langid_candidate); } return 0; } diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 87005a152..1946f5f36 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -368,10 +368,10 @@ add_test(NAME transcribe_run_dispatch_unit # ----------------------------------------------------------------------------- # Dispatcher units against fake arches: roles, host thread pool, non-finite -# input, throwing stream hooks, DIARIZE role +# input, throwing stream hooks, DIARIZE and LANGID roles # ----------------------------------------------------------------------------- -foreach(_t role_resolve run_on_threads nonfinite_input stream_hook_throw diarize_dispatch) +foreach(_t role_resolve run_on_threads nonfinite_input stream_hook_throw diarize_dispatch langid_dispatch) add_executable(transcribe_${_t}_unit ${_t}_unit.cpp) target_link_libraries(transcribe_${_t}_unit PRIVATE transcribe) target_include_directories(transcribe_${_t}_unit PRIVATE ${CMAKE_SOURCE_DIR}/src) @@ -539,7 +539,12 @@ set(TRANSCRIBE_FIXTURE_FILES ${CMAKE_CURRENT_SOURCE_DIR}/fixtures/tokenizer_minimal_streaming_buffered.gguf ${CMAKE_CURRENT_SOURCE_DIR}/fixtures/arch_cohere_minimal.gguf ${CMAKE_CURRENT_SOURCE_DIR}/fixtures/arch_granite5_ctc_minimal.gguf - ${CMAKE_CURRENT_SOURCE_DIR}/fixtures/arch_qwen3_asr_minimal.gguf) + ${CMAKE_CURRENT_SOURCE_DIR}/fixtures/arch_qwen3_asr_minimal.gguf + ${CMAKE_CURRENT_SOURCE_DIR}/fixtures/arch_ecapa_tdnn_minimal.gguf + ${CMAKE_CURRENT_SOURCE_DIR}/fixtures/arch_ecapa_tdnn_q8_0.gguf + ${CMAKE_CURRENT_SOURCE_DIR}/fixtures/arch_ecapa_tdnn_q8_0_as_f16.gguf + ${CMAKE_CURRENT_SOURCE_DIR}/fixtures/arch_ecapa_tdnn_bad_hop0.gguf + ${CMAKE_CURRENT_SOURCE_DIR}/fixtures/arch_ecapa_tdnn_bad_win_gt_fft.gguf) find_program(TRANSCRIBE_UV uv) @@ -785,6 +790,20 @@ if(TRANSCRIBE_UV) SKIP_RETURN_CODE 77) endif() +# ----------------------------------------------------------------------------- +# ecapa_tdnn (LANGID role) fixture smoke +# ----------------------------------------------------------------------------- + +if(TRANSCRIBE_UV) + add_executable(transcribe_ecapa_tdnn_smoke ecapa_tdnn_smoke.cpp) + target_link_libraries(transcribe_ecapa_tdnn_smoke PRIVATE transcribe) + target_compile_definitions(transcribe_ecapa_tdnn_smoke PRIVATE + "TRANSCRIBE_TEST_FIXTURES_DIR=\"${CMAKE_CURRENT_SOURCE_DIR}/fixtures\"") + transcribe_apply_warnings(transcribe_ecapa_tdnn_smoke) + add_dependencies(transcribe_ecapa_tdnn_smoke fixtures) + add_test(NAME transcribe_ecapa_tdnn_smoke COMMAND transcribe_ecapa_tdnn_smoke) +endif() + # ----------------------------------------------------------------------------- # Real-model gated test (developer-local, never CI) # ----------------------------------------------------------------------------- @@ -1385,6 +1404,22 @@ if(TRANSCRIBE_BUILD_REAL_MODEL_TESTS) COMMAND transcribe_sortformer_diarize_unit) set_tests_properties(transcribe_sortformer_diarize_unit PROPERTIES SKIP_RETURN_CODE 77) + + # ecapa_tdnn on the LANGID role: label table, aliases, top-1 on the + # committed FLEURS clips, the 30 s crop, and the GPU graph against the + # CPU logits when a GPU backend loads. Gated by + # TRANSCRIBE_ECAPA_TDNN_GGUF (RC 77 skip). + add_executable(transcribe_ecapa_tdnn_real_smoke + ecapa_tdnn_real_smoke.cpp) + target_link_libraries(transcribe_ecapa_tdnn_real_smoke + PRIVATE transcribe transcribe-common-example ggml) + target_compile_definitions(transcribe_ecapa_tdnn_real_smoke PRIVATE + "TRANSCRIBE_TEST_SAMPLES_DIR=\"${CMAKE_SOURCE_DIR}/samples\"") + transcribe_apply_warnings(transcribe_ecapa_tdnn_real_smoke) + add_test(NAME transcribe_ecapa_tdnn_real_smoke + COMMAND transcribe_ecapa_tdnn_real_smoke) + set_tests_properties(transcribe_ecapa_tdnn_real_smoke PROPERTIES + SKIP_RETURN_CODE 77) endif() # ----------------------------------------------------------------------------- @@ -1573,4 +1608,14 @@ if(TRANSCRIBE_BUILD_EXAMPLES) -P ${CMAKE_CURRENT_SOURCE_DIR}/cli_diarize_smoke.cmake) set_tests_properties(transcribe_cli_diarize_smoke PROPERTIES SKIP_REGULAR_EXPRESSION "SKIP:") + + if(TRANSCRIBE_UV) + add_test(NAME transcribe_cli_langid_smoke + COMMAND ${CMAKE_COMMAND} + -DCLI=$ + -DMODEL=${CMAKE_CURRENT_SOURCE_DIR}/fixtures/arch_ecapa_tdnn_minimal.gguf + -DWAV=${CMAKE_SOURCE_DIR}/samples/fleurs-en.wav + -DTEST_DIR=${CMAKE_CURRENT_BINARY_DIR}/cli-langid-smoke + -P ${CMAKE_CURRENT_SOURCE_DIR}/cli_langid_smoke.cmake) + endif() endif() diff --git a/tests/api_smoke.c b/tests/api_smoke.c index dbdb6dca0..0bbaca3e3 100644 --- a/tests/api_smoke.c +++ b/tests/api_smoke.c @@ -78,6 +78,7 @@ static void test_status_string(void) { TRANSCRIBE_ERR_OUTPUT_TRUNCATED, TRANSCRIBE_ERR_OUTPUT_REPETITION, TRANSCRIBE_ERR_UNSUPPORTED_ROLE, + TRANSCRIBE_ERR_INPUT_TOO_SHORT, }; for (size_t i = 0; i < sizeof(all) / sizeof(all[0]); ++i) { const char * s = transcribe_status_string(all[i]); @@ -128,6 +129,8 @@ static void test_abi_metadata(void) { CHECK(transcribe_abi_struct_size(TRANSCRIBE_ABI_SEGMENT) == sizeof(struct transcribe_segment)); CHECK(transcribe_abi_struct_size(TRANSCRIBE_ABI_SPEAKER_SEGMENT) == sizeof(struct transcribe_speaker_segment)); CHECK(transcribe_abi_struct_align(TRANSCRIBE_ABI_SPEAKER_SEGMENT) == _Alignof(struct transcribe_speaker_segment)); + CHECK(transcribe_abi_struct_size(TRANSCRIBE_ABI_LANGID_CANDIDATE) == sizeof(struct transcribe_langid_candidate)); + CHECK(transcribe_abi_struct_align(TRANSCRIBE_ABI_LANGID_CANDIDATE) == _Alignof(struct transcribe_langid_candidate)); CHECK(transcribe_abi_struct_size((transcribe_abi_struct) 9999) == 0); CHECK(transcribe_abi_struct_align((transcribe_abi_struct) 9999) == 0); diff --git a/tests/cli_langid_smoke.cmake b/tests/cli_langid_smoke.cmake new file mode 100644 index 000000000..f1cf1692a --- /dev/null +++ b/tests/cli_langid_smoke.cmake @@ -0,0 +1,49 @@ +# transcribe-cli routes a LANGID model to the language ID path: prints the +# `language:` / `candidate:` lines, writes the candidates with -o, honours +# --allow / --top / --repeat, and refuses --batch and --stream-chunk-ms. +# Structural checks only, on the toy fixture (MODEL). + +file(REMOVE_RECURSE "${TEST_DIR}") +file(MAKE_DIRECTORY "${TEST_DIR}") +set(_output "${TEST_DIR}/candidates.txt") + +function(run_cli out_var rc_var err_var) + execute_process( + COMMAND "${CLI}" -q --backend cpu ${ARGN} + RESULT_VARIABLE rc + OUTPUT_VARIABLE out + ERROR_VARIABLE err) + set(${out_var} "${out}" PARENT_SCOPE) + set(${rc_var} "${rc}" PARENT_SCOPE) + set(${err_var} "${err}" PARENT_SCOPE) +endfunction() + +run_cli(out rc err -m "${MODEL}" -o "${_output}" "${WAV}") +if(NOT rc EQUAL 0 OR NOT out MATCHES "language: (aa|bb|cc|dd|ee) index=[0-4] p=") + message(FATAL_ERROR "default run: rc=${rc}\n${out}\n${err}") +endif() +file(READ "${_output}" written) +if(NOT written MATCHES "^candidate: 1 [a-e][a-e] index=[0-4] p=[0-9.]+ logit=.*candidate: 5 ") + message(FATAL_ERROR "-o wrote unexpected candidates:\n${written}") +endif() + +run_cli(out rc err -m "${MODEL}" --allow bb,dd --repeat 2 "${WAV}") +if(NOT rc EQUAL 0 OR NOT out MATCHES "language: (bb|dd) " OR out MATCHES "candidate: 3 " OR NOT out MATCHES "2 allowed") + message(FATAL_ERROR "--allow bb,dd: rc=${rc}\n${out}\n${err}") +endif() + +run_cli(out rc err -m "${MODEL}" --top 1 "${WAV}") +if(NOT rc EQUAL 0 OR out MATCHES "candidate: 2 ") + message(FATAL_ERROR "--top 1: rc=${rc}\n${out}\n${err}") +endif() + +run_cli(out rc err -m "${MODEL}" --stream-chunk-ms 100 "${WAV}") +if(rc EQUAL 0 OR NOT err MATCHES "no streaming entry point") + message(FATAL_ERROR "--stream-chunk-ms on a langid model: rc=${rc}\n${err}") +endif() + +file(WRITE "${TEST_DIR}/list.txt" "${WAV}\n") +run_cli(out rc err -m "${MODEL}" --batch "${TEST_DIR}/list.txt") +if(rc EQUAL 0 OR NOT err MATCHES "--batch is ASR-only") + message(FATAL_ERROR "--batch on a langid model: rc=${rc}\n${err}") +endif() diff --git a/tests/ecapa_tdnn_real_smoke.cpp b/tests/ecapa_tdnn_real_smoke.cpp new file mode 100644 index 000000000..d931fee12 --- /dev/null +++ b/tests/ecapa_tdnn_real_smoke.cpp @@ -0,0 +1,242 @@ +// ecapa_tdnn_real_smoke.cpp - the real VoxLingua107 ECAPA-TDNN GGUF through +// the public LANGID API: label table, aliases, top-1 on committed FLEURS +// clips, the crop on a long clip, and a clip cut to 800 ms, all on the CPU +// backend; then, when a GPU backend loads, the stock-op GPU graph against the +// CPU graph's logits. Gated by TRANSCRIBE_ECAPA_TDNN_GGUF (RC 77 skip). + +#include "gguf.h" +#include "transcribe.h" +#include "transcribe/langid.h" +#include "wav.h" + +#include +#include +#include +#include +#include +#include + +namespace { + +int g_failures = 0; + +#define CHECK(cond) \ + do { \ + if (!(cond)) { \ + std::fprintf(stderr, "FAIL %s:%d: %s\n", __FILE__, __LINE__, #cond); \ + ++g_failures; \ + } \ + } while (0) + +bool load_sample(const char * name, std::vector & pcm) { + std::string err; + if (!transcribe_cli::load_wav_mono_16k(std::string(TRANSCRIBE_TEST_SAMPLES_DIR) + "/" + name, pcm, err)) { + std::fprintf(stderr, "FAIL: %s: %s\n", name, err.c_str()); + ++g_failures; + return false; + } + return true; +} + +// Top-1 code and p for one clip, or "" when the run fails. +std::string top1(transcribe_langid_session * s, const char * wav, float * p_out, int64_t * audio_ms = nullptr) { + std::vector pcm; + if (!load_sample(wav, pcm)) { + return ""; + } + const transcribe_status st = transcribe_langid_run(s, pcm.data(), static_cast(pcm.size()), nullptr); + if (st != TRANSCRIBE_OK) { + std::fprintf(stderr, "FAIL: %s: run: %s\n", wav, transcribe_status_string(st)); + ++g_failures; + return ""; + } + transcribe_langid_result r; + transcribe_langid_result_init(&r); + transcribe_langid_get_result(s, &r); + if (audio_ms != nullptr) { + *audio_ms = r.audio_ms; + } + transcribe_langid_candidate c; + transcribe_langid_candidate_init(&c); + transcribe_langid_get_candidate(s, 0, &c); + *p_out = c.p; + return c.code != nullptr ? c.code : ""; +} + +// Every label's logit for one clip, by label index; empty when the run fails. +std::vector all_logits(transcribe_langid_session * s, const char * wav) { + std::vector pcm; + if (!load_sample(wav, pcm) || + transcribe_langid_run(s, pcm.data(), static_cast(pcm.size()), nullptr) != TRANSCRIBE_OK) { + return {}; + } + transcribe_langid_result r; + transcribe_langid_result_init(&r); + transcribe_langid_get_result(s, &r); + std::vector out(static_cast(r.n_candidates), 0.0f); + for (int i = 0; i < r.n_candidates; ++i) { + transcribe_langid_candidate c; + transcribe_langid_candidate_init(&c); + transcribe_langid_get_candidate(s, i, &c); + if (c.index >= 0 && c.index < r.n_candidates) { + out[static_cast(c.index)] = c.logit; + } + } + return out; +} + +// general.file_type of the GGUF (0 = all F32, 1 = F16, 7 = Q8_0), or -1. +int gguf_file_type(const char * path) { + gguf_init_params gp{}; + gp.no_alloc = true; + gp.ctx = nullptr; + gguf_context * g = gguf_init_from_file(path, gp); + if (g == nullptr) { + return -1; + } + const int64_t key = gguf_find_key(g, "general.file_type"); + const int ft = + key >= 0 && gguf_get_kv_type(g, key) == GGUF_TYPE_UINT32 ? static_cast(gguf_get_val_u32(g, key)) : -1; + gguf_free(g); + return ft; +} + +// The GPU backends run the stock-op graph, which no other test reaches. +// Their matmuls are not bit-exact F32 (Vulkan on an AMD iGPU lands ~1e-2 off +// the CPU logits, Metal ~7e-3), so this bounds the drift rather than +// demanding parity. F16 / Q8_0 files drift further because the CPU and GPU +// round F16 operands differently (Metal: up to 0.057 on these clips), so they +// get twice that. +void check_gpu_matches_cpu(const char * path, transcribe_model * cpu_model) { + const float bound = gguf_file_type(path) == 0 ? 0.05f : 0.1f; + const transcribe_backend_request gpus[] = { TRANSCRIBE_BACKEND_VULKAN, TRANSCRIBE_BACKEND_METAL }; + transcribe_model * gm = nullptr; + for (const transcribe_backend_request b : gpus) { + transcribe_model_load_params mp; + transcribe_model_load_params_init(&mp); + mp.backend = b; + if (transcribe_model_load_file(path, &mp, &gm) == TRANSCRIBE_OK) { + break; + } + gm = nullptr; + } + if (gm == nullptr) { + std::fprintf(stderr, "ecapa_tdnn_real_smoke: no GPU backend; skipping the GPU graph check.\n"); + return; + } + + transcribe_langid_session * cs = nullptr; + transcribe_langid_session * gs = nullptr; + CHECK(transcribe_langid_session_init(cpu_model, nullptr, &cs) == TRANSCRIBE_OK); + CHECK(transcribe_langid_session_init(gm, nullptr, &gs) == TRANSCRIBE_OK); + for (const char * wav : { "fleurs-en.wav", "fleurs-zh.wav", "ru-long.wav" }) { + const std::vector a = all_logits(cs, wav); + const std::vector b = all_logits(gs, wav); + if (a.empty() || a.size() != b.size()) { + std::fprintf(stderr, "FAIL: %s: GPU run (%zu logits) vs CPU (%zu)\n", wav, b.size(), a.size()); + ++g_failures; + continue; + } + float max_diff = 0.0f; + for (size_t i = 0; i < a.size(); ++i) { + max_diff = std::fmax(max_diff, std::fabs(a[i] - b[i])); + } + std::fprintf(stderr, "ecapa_tdnn_real_smoke: %s: GPU vs CPU max |logit diff| %.4f (bound %.2f)\n", wav, + max_diff, bound); + if (!(max_diff < bound)) { + std::fprintf(stderr, "FAIL: %s: GPU logits differ from CPU by %.4f\n", wav, max_diff); + ++g_failures; + } + } + transcribe_langid_session_free(gs); + transcribe_langid_session_free(cs); + transcribe_model_free(gm); +} + +} // namespace + +int main() { + const char * path = std::getenv("TRANSCRIBE_ECAPA_TDNN_GGUF"); + if (path == nullptr || path[0] == '\0') { + std::fprintf(stderr, "ecapa_tdnn_real_smoke: TRANSCRIBE_ECAPA_TDNN_GGUF not set; skipping.\n"); + return 77; + } + transcribe_log_set(nullptr, nullptr); + + transcribe_model_load_params mp; + transcribe_model_load_params_init(&mp); + mp.backend = TRANSCRIBE_BACKEND_CPU; + transcribe_model * m = nullptr; + const transcribe_status lst = transcribe_model_load_file(path, &mp, &m); + if (lst != TRANSCRIBE_OK) { + std::fprintf(stderr, "FAIL: load %s: %s\n", path, transcribe_status_string(lst)); + return EXIT_FAILURE; + } + + CHECK(transcribe_model_roles(m) == TRANSCRIBE_ROLE_LANGID); + transcribe_langid_info info; + transcribe_langid_info_init(&info); + CHECK(transcribe_langid_get_info(m, &info) == TRANSCRIBE_OK); + CHECK(info.n_labels == 107 && info.sample_rate == 16000 && info.min_audio_ms == 500); + // Modern ISO codes alias the VoxLingua107 legacy labels. + CHECK(transcribe_langid_label_index(m, "he") == transcribe_langid_label_index(m, "iw")); + CHECK(transcribe_langid_label_index(m, "nb") == transcribe_langid_label_index(m, "no")); + CHECK(transcribe_langid_label_index(m, "he") >= 0); + const int en = transcribe_langid_label_index(m, "en"); + CHECK(en >= 0 && std::strcmp(transcribe_langid_label_code(m, en), "en") == 0); + + transcribe_langid_session_params sp; + transcribe_langid_session_params_init(&sp); + transcribe_langid_session * s = nullptr; + CHECK(transcribe_langid_session_init(m, &sp, &s) == TRANSCRIBE_OK); + + const char * cases[][2] = { + { "fleurs-en.wav", "en" }, + { "fleurs-de.wav", "de" }, + { "fleurs-fr.wav", "fr" }, + { "fleurs-es.wav", "es" }, + { "fleurs-ja.wav", "ja" }, + { "fleurs-zh.wav", "zh" }, + { "fleurs-ru.wav", "ru" }, + { "fleurs-id.wav", "id" }, + }; + for (const auto & c : cases) { + float p = 0.0f; + const std::string code = top1(s, c[0], &p); + if (code != c[1] || !(p >= 0.5f)) { + std::fprintf(stderr, "FAIL: %s: top-1 %s (p=%.3f), want %s\n", c[0], code.c_str(), p, c[1]); + ++g_failures; + } + } + + // Long input is scored on its last 30 s; a clip just over the minimum + // still runs. + float p = 0.0f; + int64_t audio_ms = 0; + top1(s, "ru-long.wav", &p, &audio_ms); + CHECK(audio_ms == 30000); + std::vector pcm; + if (load_sample("fleurs-en.wav", pcm)) { + pcm.resize(12800); // 800 ms + CHECK(transcribe_langid_run(s, pcm.data(), static_cast(pcm.size()), nullptr) == TRANSCRIBE_OK); + transcribe_langid_result r; + transcribe_langid_result_init(&r); + transcribe_langid_get_result(s, &r); + transcribe_langid_candidate c; + transcribe_langid_candidate_init(&c); + transcribe_langid_get_candidate(s, 0, &c); + CHECK(r.audio_ms == 800 && std::isfinite(c.p)); + } + + transcribe_langid_session_free(s); + + check_gpu_matches_cpu(path, m); + transcribe_model_free(m); + + if (g_failures != 0) { + std::fprintf(stderr, "%d failure(s)\n", g_failures); + return EXIT_FAILURE; + } + std::printf("ecapa_tdnn_real_smoke: ok\n"); + return EXIT_SUCCESS; +} diff --git a/tests/ecapa_tdnn_smoke.cpp b/tests/ecapa_tdnn_smoke.cpp new file mode 100644 index 000000000..5d149aff0 --- /dev/null +++ b/tests/ecapa_tdnn_smoke.cpp @@ -0,0 +1,241 @@ +// ecapa_tdnn_smoke.cpp - end-to-end smoke of the ecapa_tdnn family (LANGID +// role) against tiny synthetic GGUFs, through the public C API. +// +// arch_ecapa_tdnn_minimal.gguf (tests/fixtures/make_gguf_fixtures.py) is a +// 1/32-width model with the exact metadata and tensor contract of the real +// VoxLingua107 file, so this drives the production loader, front end, ggml +// graph and role dispatcher over the same code path the real checkpoint +// uses. It asserts structure and invariants (statuses, orderings, +// determinism), never specific values, because the weights are random. +// +// Everything runs on the CPU backend: the thread-count comparison below is +// only meaningful there, and it keeps the test hermetic on GPU-equipped CI. + +#include "transcribe.h" +#include "transcribe/langid.h" + +#include +#include +#include +#include +#include +#include +#include + +#ifndef TRANSCRIBE_TEST_FIXTURES_DIR +# error "TRANSCRIBE_TEST_FIXTURES_DIR must be defined by the build" +#endif + +namespace { + +int g_failures = 0; + +#define CHECK(cond) \ + do { \ + if (!(cond)) { \ + std::fprintf(stderr, "FAIL %s:%d: %s\n", __FILE__, __LINE__, #cond); \ + ++g_failures; \ + } \ + } while (0) + +// Deterministic pseudo-noise in [-0.5, 0.5); a hand-rolled LCG so the +// samples are identical on every platform and standard library. +std::vector noise(size_t n, uint32_t seed) { + std::vector out(n); + uint32_t s = seed | 1u; + for (size_t i = 0; i < n; ++i) { + s = s * 1664525u + 1013904223u; + out[i] = static_cast((s >> 8) & 0xFFFFFF) / 16777216.0f - 0.5f; + } + return out; +} + +std::string fixture(const char * name) { + return std::string(TRANSCRIBE_TEST_FIXTURES_DIR) + "/" + name; +} + +transcribe_model * load_cpu(const char * name, transcribe_status * st_out = nullptr) { + transcribe_model_load_params mp; + transcribe_model_load_params_init(&mp); + mp.backend = TRANSCRIBE_BACKEND_CPU; + transcribe_model * m = nullptr; + const transcribe_status st = transcribe_model_load_file(fixture(name).c_str(), &mp, &m); + if (st_out != nullptr) { + *st_out = st; + } + return m; +} + +transcribe_langid_session * open_session(transcribe_model * m, int n_threads, int max_audio_ms = 0) { + transcribe_langid_session_params sp; + transcribe_langid_session_params_init(&sp); + sp.n_threads = n_threads; + sp.max_audio_ms = max_audio_ms; + transcribe_langid_session * s = nullptr; + CHECK(transcribe_langid_session_init(m, &sp, &s) == TRANSCRIBE_OK); + return s; +} + +transcribe_langid_result result_of(const transcribe_langid_session * s) { + transcribe_langid_result r; + transcribe_langid_result_init(&r); + CHECK(transcribe_langid_get_result(s, &r) == TRANSCRIBE_OK); + return r; +} + +std::vector candidates_of(const transcribe_langid_session * s) { + const transcribe_langid_result r = result_of(s); + std::vector out(static_cast(r.n_candidates)); + for (int i = 0; i < r.n_candidates; ++i) { + transcribe_langid_candidate_init(&out[static_cast(i)]); + CHECK(transcribe_langid_get_candidate(s, i, &out[static_cast(i)]) == TRANSCRIBE_OK); + } + return out; +} + +void test_model_surface(transcribe_model * m) { + CHECK(transcribe_model_roles(m) == TRANSCRIBE_ROLE_LANGID); + CHECK(std::strcmp(transcribe_model_arch_string(m), "ecapa_tdnn") == 0); + CHECK(std::strcmp(transcribe_model_variant_string(m), "ecapa-tdnn-toy") == 0); + CHECK(transcribe_model_supports(m, TRANSCRIBE_FEATURE_CANCELLATION)); + + transcribe_langid_info info; + transcribe_langid_info_init(&info); + CHECK(transcribe_langid_get_info(m, &info) == TRANSCRIBE_OK); + CHECK(info.sample_rate == 16000 && info.n_labels == 5 && info.min_audio_ms == 500); + CHECK(std::strcmp(transcribe_langid_label_code(m, 2), "cc") == 0); + CHECK(std::strcmp(transcribe_langid_label_name(m, 2), "Charlie") == 0); + CHECK(transcribe_langid_label_index(m, "xx") == 0); // alias from the GGUF + + // Not an ASR model. + transcribe_capabilities caps; + transcribe_capabilities_init(&caps); + CHECK(transcribe_model_get_capabilities(m, &caps) == TRANSCRIBE_ERR_UNSUPPORTED_ROLE); + transcribe_session * asr = nullptr; + CHECK(transcribe_session_init(m, nullptr, &asr) == TRANSCRIBE_ERR_UNSUPPORTED_ROLE); + CHECK(asr == nullptr); +} + +void test_run(transcribe_model * m) { + transcribe_langid_session * s = open_session(m, 1); + const std::vector pcm = noise(16000, 7); // 1 s + + CHECK(transcribe_langid_run(s, pcm.data(), static_cast(pcm.size()), nullptr) == TRANSCRIBE_OK); + const transcribe_langid_result r = result_of(s); + CHECK(r.n_candidates == 5 && r.n_allowed == 5 && r.allowed_mass == 1.0f && r.audio_ms == 1000); + const auto first = candidates_of(s); + double sum = 0.0; + for (size_t i = 0; i < first.size(); ++i) { + CHECK(std::isfinite(first[i].p) && std::isfinite(first[i].logit)); + CHECK(first[i].code != nullptr && first[i].name != nullptr); + if (i > 0) { + CHECK(first[i - 1].p >= first[i].p); + } + sum += first[i].p; + } + CHECK(std::fabs(sum - 1.0) < 1e-5); + + // Determinism: a second run is bit-identical. + CHECK(transcribe_langid_run(s, pcm.data(), static_cast(pcm.size()), nullptr) == TRANSCRIBE_OK); + const auto again = candidates_of(s); + CHECK(again.size() == first.size()); + for (size_t i = 0; i < first.size() && i < again.size(); ++i) { + CHECK(again[i].index == first[i].index && again[i].logit == first[i].logit); + } + + // Silence is valid input. + const std::vector zeros(16000, 0.0f); + CHECK(transcribe_langid_run(s, zeros.data(), 16000, nullptr) == TRANSCRIBE_OK); + + // Abort before compute. + transcribe_langid_set_abort_callback(s, [](void *) { return true; }, nullptr); + CHECK(transcribe_langid_run(s, pcm.data(), 16000, nullptr) == TRANSCRIBE_ERR_ABORTED); + CHECK(result_of(s).n_candidates == 0); + transcribe_langid_set_abort_callback(s, nullptr, nullptr); + + transcribe_langid_session_free(s); +} + +// The front end and graph give the same logits for any thread count (the +// mel frames and the ggml CPU rows are split without changing any sum). +void test_thread_invariance(transcribe_model * m) { + const std::vector pcm = noise(16000 * 3, 3); + transcribe_langid_session * s1 = open_session(m, 1); + transcribe_langid_session * s4 = open_session(m, 4); + CHECK(transcribe_langid_run(s1, pcm.data(), static_cast(pcm.size()), nullptr) == TRANSCRIBE_OK); + CHECK(transcribe_langid_run(s4, pcm.data(), static_cast(pcm.size()), nullptr) == TRANSCRIBE_OK); + const auto a = candidates_of(s1); + const auto b = candidates_of(s4); + CHECK(a.size() == b.size()); + for (size_t i = 0; i < a.size() && i < b.size(); ++i) { + CHECK(a[i].index == b[i].index && std::fabs(a[i].logit - b[i].logit) < 1e-5f); + } + transcribe_langid_session_free(s1); + transcribe_langid_session_free(s4); +} + +// Q8_0 weights are widened to F16 at load: a Q8_0 model must give exactly +// the logits of the same model stored as F16 holding the dequantized values. +void test_q8_0_widened_to_f16() { + transcribe_model * q8 = load_cpu("arch_ecapa_tdnn_q8_0.gguf"); + transcribe_model * ref = load_cpu("arch_ecapa_tdnn_q8_0_as_f16.gguf"); + CHECK(q8 != nullptr && ref != nullptr); + if (q8 == nullptr || ref == nullptr) { + transcribe_model_free(q8); + transcribe_model_free(ref); + return; + } + const std::vector pcm = noise(16000 * 2, 5); + transcribe_langid_session * sq = open_session(q8, 2); + transcribe_langid_session * sr = open_session(ref, 2); + CHECK(transcribe_langid_run(sq, pcm.data(), static_cast(pcm.size()), nullptr) == TRANSCRIBE_OK); + CHECK(transcribe_langid_run(sr, pcm.data(), static_cast(pcm.size()), nullptr) == TRANSCRIBE_OK); + const auto a = candidates_of(sq); + const auto b = candidates_of(sr); + CHECK(a.size() == 5 && a.size() == b.size()); + for (size_t i = 0; i < a.size() && i < b.size(); ++i) { + CHECK(a[i].index == b[i].index && a[i].logit == b[i].logit); + } + transcribe_langid_session_free(sq); + transcribe_langid_session_free(sr); + transcribe_model_free(q8); + transcribe_model_free(ref); +} + +// A front end the mel code cannot run (zero hop, window wider than n_fft) +// fails the load instead of crashing at run time. +void test_bad_frontend_rejected() { + for (const char * name : { "arch_ecapa_tdnn_bad_hop0.gguf", "arch_ecapa_tdnn_bad_win_gt_fft.gguf" }) { + transcribe_status st = TRANSCRIBE_OK; + transcribe_model * m = load_cpu(name, &st); + CHECK(st == TRANSCRIBE_ERR_GGUF && m == nullptr); + transcribe_model_free(m); + } +} + +} // namespace + +int main() { + transcribe_log_set(nullptr, nullptr); + + transcribe_status st = TRANSCRIBE_OK; + transcribe_model * m = load_cpu("arch_ecapa_tdnn_minimal.gguf", &st); + if (st != TRANSCRIBE_OK || m == nullptr) { + std::fprintf(stderr, "FAIL: load arch_ecapa_tdnn_minimal.gguf: %s\n", transcribe_status_string(st)); + return EXIT_FAILURE; + } + test_model_surface(m); + test_run(m); + test_thread_invariance(m); + transcribe_model_free(m); + + test_q8_0_widened_to_f16(); + test_bad_frontend_rejected(); + + if (g_failures != 0) { + std::fprintf(stderr, "%d failure(s)\n", g_failures); + return EXIT_FAILURE; + } + std::printf("ecapa_tdnn_smoke: ok\n"); + return EXIT_SUCCESS; +} diff --git a/tests/fixtures/make_gguf_fixtures.py b/tests/fixtures/make_gguf_fixtures.py index 2bdd939c8..780f823fe 100644 --- a/tests/fixtures/make_gguf_fixtures.py +++ b/tests/fixtures/make_gguf_fixtures.py @@ -13,7 +13,7 @@ configure time). Tensor data emission uses Python struct only — no numpy dep — because the toy tensors are tiny (~3000 fp32 elements). -Six fixtures are emitted: +Fixtures emitted (see emit_fixtures for the full list): arch_parakeet.gguf -- valid header, KV pairs: general.architecture = "parakeet" @@ -110,6 +110,7 @@ from __future__ import annotations +import math import struct import sys from pathlib import Path @@ -132,11 +133,16 @@ # ggml_type enum values used for tensor data. Pinned here so we are not # at the mercy of upstream renumbering — a mismatch would surface as a # loader test failure (the most useful possible signal). -GGML_TYPE_F32 = 0 +GGML_TYPE_F32 = 0 +GGML_TYPE_F16 = 1 +GGML_TYPE_Q8_0 = 8 -# Bytes per element for each ggml_type we emit. +# (bytes per block, elements per block) for each ggml_type we emit. Q8_0 is +# 32 int8 quants behind one fp16 scale, blocked along ne[0]. GGML_TYPE_SIZE = { - GGML_TYPE_F32: 4, + GGML_TYPE_F32: (4, 1), + GGML_TYPE_F16: (2, 1), + GGML_TYPE_Q8_0: (34, 32), } @@ -262,7 +268,7 @@ def _string_kvs(pairs: list[tuple[str, str]]) -> list[bytes]: # A "Tensor" here is just a (name, ne, dtype, data_bytes) tuple. ne is # fast-to-slow dim order matching ggml_tensor::ne[]. data_bytes is the # raw little-endian bytes of the tensor's elements, length must equal -# product(ne) * GGML_TYPE_SIZE[dtype]. +# product(ne) / block * block_bytes (GGML_TYPE_SIZE[dtype]). # # _build_full_gguf assembles header + KV section + tensor info section # + aligned tensor data blob in one pass. The layout follows @@ -276,10 +282,13 @@ class Tensor: def __init__( self, name: str, ne: list[int], dtype: int, data: bytes ) -> None: + block_bytes, block = GGML_TYPE_SIZE[dtype] + if ne[0] % block != 0: + raise ValueError(f"tensor {name!r}: ne[0]={ne[0]} is not a multiple of {block}") nbytes = 1 for d in ne: nbytes *= d - nbytes *= GGML_TYPE_SIZE[dtype] + nbytes = nbytes // block * block_bytes if len(data) != nbytes: raise ValueError( f"tensor {name!r}: ne={ne} dtype={dtype} expects " @@ -1359,6 +1368,179 @@ def _qwen3_asr_tensors() -> list[Tensor]: ) +# --------------------------------------------------------------------------- +# Toy ecapa_tdnn (LANGID) hparams + tensor catalog +# --------------------------------------------------------------------------- +# +# The same metadata and tensor contract scripts/convert-ecapa_tdnn.py writes +# for speechbrain/lang-id-voxlingua107-ecapa, at 1/32 width: channels +# [32]*4 + [96], the real kernel sizes / dilations / res2net scale, a real +# 60 x 201 front end, five labels aa..ee plus the alias xx=aa. Weights are +# small seeded pseudo-random values so a forward pass stays well inside float +# range; tests assert structure and invariants, never specific values. + +import random as _random + +ECAPA_CHANNELS = [32, 32, 32, 32, 96] +ECAPA_KERNELS = [5, 3, 3, 3, 1] +ECAPA_DILATIONS = [1, 2, 3, 4, 1] +ECAPA_SCALE = 8 +ECAPA_SE = 8 +ECAPA_ATT = 8 +ECAPA_EMB = 16 +ECAPA_HID = 16 +ECAPA_N_MELS = 60 +ECAPA_N_FREQ = 201 +ECAPA_LABEL_CODES = ["aa", "bb", "cc", "dd", "ee"] +ECAPA_LABEL_NAMES = ["Alpha", "Bravo", "Charlie", "Delta", "Echo"] + + +def _ecapa_tdnn_hparams_kv(codes: list[str], names: list[str], aliases: list[str], + hop_length: int = 160, win_length: int = 400) -> list[bytes]: + return [ + _pack_kv_string("stt.frontend.type", "speechbrain_fbank"), + _pack_kv_uint32("stt.frontend.sample_rate", 16000), + _pack_kv_uint32("stt.frontend.n_fft", 400), + _pack_kv_uint32("stt.frontend.hop_length", hop_length), + _pack_kv_uint32("stt.frontend.win_length", win_length), + _pack_kv_uint32("stt.frontend.num_mels", ECAPA_N_MELS), + _pack_kv_string("stt.frontend.window", "hamming_periodic"), + _pack_kv_string("stt.frontend.pad_mode", "constant"), + _pack_kv_float32("stt.frontend.log_clamp_min", 1e-10), + _pack_kv_float32("stt.frontend.top_db", 80.0), + _pack_kv_string("stt.frontend.normalize", "sentence_mean"), + _pack_kv_array_int32("stt.ecapa_tdnn.channels", ECAPA_CHANNELS), + _pack_kv_array_int32("stt.ecapa_tdnn.kernel_sizes", ECAPA_KERNELS), + _pack_kv_array_int32("stt.ecapa_tdnn.dilations", ECAPA_DILATIONS), + _pack_kv_uint32("stt.ecapa_tdnn.res2net_scale", ECAPA_SCALE), + _pack_kv_uint32("stt.ecapa_tdnn.se_channels", ECAPA_SE), + _pack_kv_uint32("stt.ecapa_tdnn.attention_channels", ECAPA_ATT), + _pack_kv_float32("stt.ecapa_tdnn.asp_eps", 1e-12), + _pack_kv_uint32("stt.ecapa_tdnn.embedding_dim", ECAPA_EMB), + _pack_kv_uint32("stt.ecapa_tdnn.classifier_hidden", ECAPA_HID), + _pack_kv_float32("stt.ecapa_tdnn.classifier_leaky_slope", 0.01), + _pack_kv_array_string("stt.langid.labels.codes", codes), + _pack_kv_array_string("stt.langid.labels.names", names), + _pack_kv_array_string("stt.langid.labels.aliases", aliases), + ] + + +def _ecapa_tdnn_tensors(n_labels: int) -> list[Tensor]: + rng = _random.Random(0) + out: list[Tensor] = [] + + def add(name: str, ne: list[int], values: list[float]) -> None: + out.append(Tensor(name, ne, GGML_TYPE_F32, _f32_bytes(values))) + + def weight(name: str, ne: list[int]) -> None: + n = 1 + for d in ne: + n *= d + fan_in = ne[0] * (ne[2] if len(ne) == 3 else 1) + bound = 1.0 / fan_in ** 0.5 + add(name, ne, [rng.uniform(-bound, bound) for _ in range(n)]) + + def vec(name: str, n: int, center: float, spread: float) -> None: + add(name, [n], [center + rng.uniform(-spread, spread) for _ in range(n)]) + + def tdnn(prefix: str, ic: int, oc: int, k: int, conv: bool) -> None: + stem = f"{prefix}.conv" if conv else prefix + weight(f"{stem}.weight", [ic, oc, k] if conv else [ic, oc]) + vec(f"{stem}.bias", oc, 0.0, 0.05) + vec(f"{prefix}.bn.scale", oc, 1.0, 0.1) + vec(f"{prefix}.bn.shift", oc, 0.0, 0.05) + + # Triangular mel-major filterbank, ne = [n_freq, n_mels]. + fb = [0.0] * (ECAPA_N_MELS * ECAPA_N_FREQ) + for m in range(ECAPA_N_MELS): + c = 2 + 3 * m + for k in range(c - 3, c + 4): + if 0 <= k < ECAPA_N_FREQ: + fb[m * ECAPA_N_FREQ + k] = 1.0 - abs(k - c) / 4.0 + add("frontend.mel_filterbank", [ECAPA_N_FREQ, ECAPA_N_MELS], fb) + + c, cm, chunk = ECAPA_CHANNELS[0], ECAPA_CHANNELS[-1], ECAPA_CHANNELS[0] // ECAPA_SCALE + tdnn("blk.0", ECAPA_N_MELS, c, ECAPA_KERNELS[0], conv=True) + for i in (1, 2, 3): + tdnn(f"blk.{i}.tdnn1", c, c, 1, conv=False) + for j in range(ECAPA_SCALE - 1): + tdnn(f"blk.{i}.res2.{j}", chunk, chunk, ECAPA_KERNELS[i], conv=True) + tdnn(f"blk.{i}.tdnn2", c, c, 1, conv=False) + weight(f"blk.{i}.se.c1.weight", [c, ECAPA_SE]) + vec(f"blk.{i}.se.c1.bias", ECAPA_SE, 0.0, 0.05) + weight(f"blk.{i}.se.c2.weight", [ECAPA_SE, c]) + vec(f"blk.{i}.se.c2.bias", c, 0.0, 0.05) + for n in (1, 2, 3): + weight(f"mfa.w{n}.weight", [c, cm]) + vec("mfa.bias", cm, 0.0, 0.05) + vec("mfa.bn.scale", cm, 1.0, 0.1) + vec("mfa.bn.shift", cm, 0.0, 0.05) + for part in ("x", "mean", "std"): + weight(f"asp.tdnn.{part}.weight", [cm, ECAPA_ATT]) + vec("asp.tdnn.bias", ECAPA_ATT, 0.0, 0.05) + vec("asp.tdnn.bn.scale", ECAPA_ATT, 1.0, 0.1) + vec("asp.tdnn.bn.shift", ECAPA_ATT, 0.0, 0.05) + weight("asp.attn.weight", [ECAPA_ATT, cm]) + vec("asp.attn.bias", cm, 0.0, 0.05) + weight("fc.weight", [2 * cm, ECAPA_EMB]) + vec("fc.bias", ECAPA_EMB, 0.0, 0.05) + weight("cls.l1.weight", [ECAPA_EMB, ECAPA_HID]) + vec("cls.l1.bias", ECAPA_HID, 0.0, 0.05) + weight("cls.out.weight", [ECAPA_HID, n_labels]) + vec("cls.out.bias", n_labels, 0.0, 0.05) + return out + + +def _q8_0_blocks(values: list[float]) -> tuple[bytes, list[float]]: + """ggml's quantize_row_q8_0_ref, plus the values its dequantize_row_q8_0 + gives back (fp16 scale times int8, exact in float).""" + data = bytearray() + deq: list[float] = [] + for b in range(0, len(values), 32): + x = values[b:b + 32] + d = max(abs(v) for v in x) / 127.0 + inv = 1.0 / d if d else 0.0 + q = [int(math.copysign(math.floor(abs(v * inv) + 0.5), v)) for v in x] + d16 = struct.pack(" list[Tensor]: + """Every 2-D weight the quantizer would make Q8_0 (ne[0] % 32 == 0), as + Q8_0, or (as_f16) as F16 holding exactly the dequantized Q8_0 values: + what the ecapa_tdnn loader must produce when it widens Q8_0 to F16.""" + out = [] + for t in tensors: + if len(t.ne) != 2 or not t.name.endswith(".weight") or t.ne[0] % 32 != 0 or t.name.startswith("frontend."): + out.append(t) + continue + values = list(struct.unpack(f"<{len(t.data) // 4}f", t.data)) + q8, deq = _q8_0_blocks(values) + if as_f16: + out.append(Tensor(t.name, t.ne, GGML_TYPE_F16, struct.pack(f"<{len(deq)}e", *deq))) + else: + out.append(Tensor(t.name, t.ne, GGML_TYPE_Q8_0, q8)) + return out + + +def _ecapa_tdnn_gguf(codes: list[str], names: list[str], aliases: list[str], q8_0: str = "", + **frontend: int) -> bytes: + tensors = _ecapa_tdnn_tensors(len(codes)) + if q8_0: + tensors = _ecapa_tdnn_q8_0(tensors, as_f16=(q8_0 == "as_f16")) + return _build_full_gguf( + GGUF_MAGIC, + [ + _pack_kv_string("general.architecture", "ecapa_tdnn"), + _pack_kv_string("stt.variant", "ecapa-tdnn-toy"), + *_ecapa_tdnn_hparams_kv(codes, names, aliases, **frontend), + ], + tensors, + ) + + def _write(path: Path, data: bytes) -> None: path.parent.mkdir(parents=True, exist_ok=True) path.write_bytes(data) @@ -1773,6 +1955,24 @@ def emit_fixtures(out_dir: Path) -> None: ) + # ecapa_tdnn (LANGID role): a structurally complete toy model. + _write(out_dir / "arch_ecapa_tdnn_minimal.gguf", + _ecapa_tdnn_gguf(ECAPA_LABEL_CODES, ECAPA_LABEL_NAMES, ["xx=aa"])) + # The toy model with its Q8_0-eligible weights in Q8_0, and the same + # weights as F16 holding the dequantized values. The loader widens Q8_0 + # to F16, so the two must give bit-identical logits on the CPU. + _write(out_dir / "arch_ecapa_tdnn_q8_0.gguf", + _ecapa_tdnn_gguf(ECAPA_LABEL_CODES, ECAPA_LABEL_NAMES, ["xx=aa"], q8_0="q8_0")) + _write(out_dir / "arch_ecapa_tdnn_q8_0_as_f16.gguf", + _ecapa_tdnn_gguf(ECAPA_LABEL_CODES, ECAPA_LABEL_NAMES, ["xx=aa"], q8_0="as_f16")) + # Front ends the loader must reject: a zero hop divides by zero in the + # frame count, and win_length > n_fft overruns the padded window. + _write(out_dir / "arch_ecapa_tdnn_bad_hop0.gguf", + _ecapa_tdnn_gguf(ECAPA_LABEL_CODES, ECAPA_LABEL_NAMES, ["xx=aa"], hop_length=0)) + _write(out_dir / "arch_ecapa_tdnn_bad_win_gt_fft.gguf", + _ecapa_tdnn_gguf(ECAPA_LABEL_CODES, ECAPA_LABEL_NAMES, ["xx=aa"], win_length=512)) + + def main(argv: list[str]) -> int: if len(argv) > 2: print(f"usage: {argv[0]} [output_dir]", file=sys.stderr) diff --git a/tests/golden/ecapa_tdnn/lang-id-voxlingua107-ecapa.manifest.json b/tests/golden/ecapa_tdnn/lang-id-voxlingua107-ecapa.manifest.json new file mode 100644 index 000000000..ffa7762f1 --- /dev/null +++ b/tests/golden/ecapa_tdnn/lang-id-voxlingua107-ecapa.manifest.json @@ -0,0 +1,52 @@ +{ + "schema": "transcribe-golden-manifest-v1", + "family": "ecapa_tdnn", + "variant": "lang-id-voxlingua107-ecapa", + "role": "langid", + "source_model": { + "hf_repo": "speechbrain/lang-id-voxlingua107-ecapa", + "hf_revision": "0253049ae131d6a4be1c4f0d8b0ff483a0f8c8e9" + }, + "reference": { + "kind": "speechbrain", + "source": "https://github.com/speechbrain/speechbrain", + "revision": "v1.1.1", + "entrypoint": "scripts/dump_reference_ecapa_tdnn_speechbrain.py", + "dump_args": [] + }, + "expected_dtype": "float32", + "dtype_source": "manual", + "frontend": { + "sample_rate": 16000, + "n_mels": 60, + "hop_length": 160, + "fft_size": 400, + "win_length": 400, + "window": "hamming_periodic", + "pad_mode": "constant", + "normalization": "sentence_mean", + "top_db": 80 + }, + "tokenizer_summary": { + "type": "other", + "vocab_size": 0, + "special_tokens": {} + }, + "langid": { + "n_labels": 107, + "label_aliases": ["he=iw", "jv=jw", "fil=tl", "nb=no"], + "min_audio_ms": 500, + "default_max_audio_ms": 30000 + }, + "tolerance_file": "tests/tolerances/ecapa_tdnn.json", + "cases": [ + { "audio": "fleurs-en", "language": null, "expected_language": "en" }, + { "audio": "fleurs-de", "language": null, "expected_language": "de" }, + { "audio": "fleurs-fr", "language": null, "expected_language": "fr" }, + { "audio": "fleurs-es", "language": null, "expected_language": "es" }, + { "audio": "fleurs-ja", "language": null, "expected_language": "ja" }, + { "audio": "fleurs-zh", "language": null, "expected_language": "zh" }, + { "audio": "fleurs-ru", "language": null, "expected_language": "ru" }, + { "audio": "fleurs-id", "language": null, "expected_language": "id" } + ] +} diff --git a/tests/langid_dispatch_unit.cpp b/tests/langid_dispatch_unit.cpp new file mode 100644 index 000000000..1ab581df0 --- /dev/null +++ b/tests/langid_dispatch_unit.cpp @@ -0,0 +1,299 @@ +// langid_dispatch_unit.cpp - LANGID role dispatcher (transcribe-langid.cpp) +// against a fake arch: label table validation, role checks, softmax / +// ranking / top-k, the allowed set, crop and minimum length, non-finite +// input and logits, abort. + +#include "transcribe-arch.h" +#include "transcribe-langid.h" +#include "transcribe-model.h" +#include "transcribe.h" +#include "transcribe/langid.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace { + +int g_failures = 0; + +#define CHECK(cond) \ + do { \ + if (!(cond)) { \ + std::fprintf(stderr, "FAIL %s:%d: %s\n", __FILE__, __LINE__, #cond); \ + ++g_failures; \ + } \ + } while (0) + +bool near(double a, double b, double tol = 1e-6) { + return std::fabs(a - b) <= tol; +} + +transcribe::LangidLabels g_labels; // aa..ee, alias xx=aa + +int g_run_calls = 0; +int g_last_n = 0; +const float * g_last_pcm = nullptr; +std::vector g_logits = { 1.0f, 3.0f, 2.0f, 0.0f, -1.0f }; +bool g_check_abort = false; + +const transcribe::LangidLabels & fake_labels(const transcribe_model *) { + return g_labels; +} + +transcribe_langid_session * fake_new_session() { + return new transcribe_langid_session(); +} + +transcribe_status fake_run(transcribe_langid_session * s, const float * pcm, int n, std::vector & logits) { + ++g_run_calls; + g_last_n = n; + g_last_pcm = pcm; + if (g_check_abort && s->poll_abort()) { + return TRANSCRIBE_ERR_ABORTED; + } + logits = g_logits; + return TRANSCRIBE_OK; +} + +const transcribe::LangidOps k_ops = { fake_labels, fake_new_session, fake_run }; + +const transcribe::Arch k_arch = { + /* .name = */ "fake-langid", + /* .load = */ nullptr, + /* .init_context = */ nullptr, + /* .run = */ nullptr, + /* .run_batch = */ nullptr, + /* .stream_validate = */ nullptr, + /* .stream_begin = */ nullptr, + /* .stream_feed = */ nullptr, + /* .stream_finalize = */ nullptr, + /* .stream_reset = */ nullptr, + /* .accepts_ext_kind = */ nullptr, + /* .run_validate = */ nullptr, + /* .diarize = */ nullptr, + /* .langid = */ &k_ops, +}; + +struct Fixture { + transcribe_model model; + transcribe_langid_session * session = nullptr; + std::vector pcm = std::vector(16000, 0.0f); // 1 s + + explicit Fixture(int32_t max_audio_ms = 0) { + g_run_calls = 0; + g_logits = { 1.0f, 3.0f, 2.0f, 0.0f, -1.0f }; + g_check_abort = false; + model.arch = &k_arch; + model.roles = TRANSCRIBE_ROLE_LANGID; + transcribe_langid_session_params sp; + transcribe_langid_session_params_init(&sp); + sp.max_audio_ms = max_audio_ms; + CHECK(transcribe_langid_session_init(&model, &sp, &session) == TRANSCRIBE_OK); + } + + ~Fixture() { transcribe_langid_session_free(session); } + + transcribe_status run(const transcribe_langid_params * p = nullptr) { + return transcribe_langid_run(session, pcm.data(), static_cast(pcm.size()), p); + } + + transcribe_langid_result result() const { + transcribe_langid_result r; + transcribe_langid_result_init(&r); + CHECK(transcribe_langid_get_result(session, &r) == TRANSCRIBE_OK); + return r; + } + + transcribe_langid_candidate candidate(int i) const { + transcribe_langid_candidate c; + transcribe_langid_candidate_init(&c); + CHECK(transcribe_langid_get_candidate(session, i, &c) == TRANSCRIBE_OK); + return c; + } +}; + +// Malformed label metadata fails the load instead of being last-wins. +void test_label_table() { + const std::vector names = { "A", "B" }; + transcribe::LangidLabels out; + CHECK(transcribe::build_langid_labels({ "aa", "bb" }, names, { "xx=aa" }, "test", out) == TRANSCRIBE_OK); + CHECK(out.codes.size() == 2 && out.index.size() == 3 && out.index.at("xx") == 0); + CHECK(transcribe::build_langid_labels({ "aa", "aa" }, names, {}, "test", out) == TRANSCRIBE_ERR_GGUF); + CHECK(transcribe::build_langid_labels({ "aa", "bb" }, names, { "xx=zz" }, "test", out) == TRANSCRIBE_ERR_GGUF); +} + +void test_role_checks_and_labels() { + transcribe_model model; + model.arch = &k_arch; + model.roles = TRANSCRIBE_ROLE_ASR; // no LANGID bit + + transcribe_langid_session * s = reinterpret_cast(0x1); + CHECK(transcribe_langid_session_init(&model, nullptr, &s) == TRANSCRIBE_ERR_UNSUPPORTED_ROLE); + CHECK(s == nullptr); + transcribe_langid_info info; + transcribe_langid_info_init(&info); + CHECK(transcribe_langid_get_info(&model, &info) == TRANSCRIBE_ERR_UNSUPPORTED_ROLE); + + model.roles = TRANSCRIBE_ROLE_LANGID; + CHECK(transcribe_langid_get_info(&model, &info) == TRANSCRIBE_OK); + CHECK(info.sample_rate == 16000 && info.n_labels == 5 && info.min_audio_ms == 500); + CHECK(std::strcmp(transcribe_langid_label_code(&model, 1), "bb") == 0); + CHECK(std::strcmp(transcribe_langid_label_name(&model, 4), "Ee") == 0); + CHECK(transcribe_langid_label_index(&model, "cc") == 2); + CHECK(transcribe_langid_label_index(&model, "xx") == 0); + CHECK(transcribe_langid_label_index(&model, "zz") == -1); + + // A window below the minimum is rejected up front. + transcribe_langid_session_params sp; + transcribe_langid_session_params_init(&sp); + sp.max_audio_ms = 499; + CHECK(transcribe_langid_session_init(&model, &sp, &s) == TRANSCRIBE_ERR_INVALID_ARG); + CHECK(s == nullptr); +} + +void test_ranking() { + Fixture f; + CHECK(f.run() == TRANSCRIBE_OK); + CHECK(g_run_calls == 1 && g_last_n == 16000 && g_last_pcm == f.pcm.data()); + + const transcribe_langid_result r = f.result(); + CHECK(r.n_candidates == 5 && r.n_allowed == 5 && r.allowed_mass == 1.0f && r.audio_ms == 1000); + + // Logits {1, 3, 2, 0, -1} rank bb, cc, aa, dd, ee. + const int want[] = { 1, 2, 0, 3, 4 }; + double sum = 0.0; + const double denom = std::exp(1.0) + std::exp(3.0) + std::exp(2.0) + std::exp(0.0) + std::exp(-1.0); + for (int i = 0; i < 5; ++i) { + const transcribe_langid_candidate c = f.candidate(i); + CHECK(c.index == want[i]); + CHECK(c.code != nullptr && std::strcmp(c.code, g_labels.codes[want[i]].c_str()) == 0); + CHECK(c.logit == g_logits[want[i]]); + CHECK(near(c.p, std::exp(static_cast(g_logits[want[i]])) / denom)); + sum += c.p; + } + CHECK(near(sum, 1.0)); + + // top_k truncates the rows but not n_allowed. + transcribe_langid_params p; + transcribe_langid_params_init(&p); + p.top_k = 2; + CHECK(f.run(&p) == TRANSCRIBE_OK); + CHECK(f.result().n_candidates == 2 && f.result().n_allowed == 5); + CHECK(f.candidate(1).index == 2 && f.candidate(2).code == nullptr); +} + +// The allowed set renormalizes over its members; malformed or unknown lists +// are rejected before the family runs and keep the last result. +void test_allowed_set() { + Fixture f; + transcribe_langid_params p; + transcribe_langid_params_init(&p); + + const char * dd_bb[] = { "dd", "bb" }; + p.allowed = dd_bb; + p.n_allowed = 2; + CHECK(f.run(&p) == TRANSCRIBE_OK); + const transcribe_langid_result r = f.result(); + CHECK(r.n_candidates == 2 && r.n_allowed == 2); + const double e_bb = std::exp(3.0); + const double e_dd = std::exp(0.0); + const double all = std::exp(1.0) + e_bb + std::exp(2.0) + e_dd + std::exp(-1.0); + CHECK(near(r.allowed_mass, (e_bb + e_dd) / all)); + const transcribe_langid_candidate c0 = f.candidate(0); + const transcribe_langid_candidate c1 = f.candidate(1); + CHECK(c0.index == 1 && c1.index == 3); + CHECK(near(c0.p, e_bb / (e_bb + e_dd)) && near(c1.p, e_dd / (e_bb + e_dd))); + + g_run_calls = 0; + p.n_allowed = 0; + CHECK(f.run(&p) == TRANSCRIBE_ERR_INVALID_ARG); + const char * unknown[] = { "aa", "zz" }; + p.allowed = unknown; + p.n_allowed = 2; + CHECK(f.run(&p) == TRANSCRIBE_ERR_UNSUPPORTED_LANGUAGE); + CHECK(g_run_calls == 0); + CHECK(f.result().n_allowed == 2); +} + +// The crop keeps the last max_audio_ms, and the minimum applies to the audio +// actually scored. +void test_crop_and_minimum() { + { + Fixture f(500); + CHECK(f.run() == TRANSCRIBE_OK); // 1 s scored as its last 500 ms + CHECK(g_last_n == 8000 && g_last_pcm == f.pcm.data() + 8000); + CHECK(f.result().audio_ms == 500); + } + { + Fixture f; + CHECK(transcribe_langid_run(f.session, f.pcm.data(), 8000, nullptr) == TRANSCRIBE_OK); // exactly 500 ms + g_run_calls = 0; + CHECK(transcribe_langid_run(f.session, f.pcm.data(), 7999, nullptr) == TRANSCRIBE_ERR_INPUT_TOO_SHORT); + CHECK(g_run_calls == 0); + CHECK(f.result().audio_ms == 500); // previous result kept + } +} + +// Non-finite PCM is rejected up front; non-finite or miscounted logits fail +// the run instead of ranking garbage. +void test_non_finite() { + Fixture f; + CHECK(f.run() == TRANSCRIBE_OK); + g_run_calls = 0; + + std::vector bad(16000, 0.0f); + bad[3] = std::numeric_limits::quiet_NaN(); + CHECK(transcribe_langid_run(f.session, bad.data(), 16000, nullptr) == TRANSCRIBE_ERR_INVALID_ARG); + CHECK(g_run_calls == 0); + CHECK(f.result().n_candidates == 5); + + g_logits = { 1.0f, std::numeric_limits::quiet_NaN(), 0.0f, 0.0f, 0.0f }; + CHECK(f.run() == TRANSCRIBE_ERR_BACKEND); + CHECK(f.result().n_candidates == 0 && f.result().allowed_mass == 0.0f); + g_logits = { 1.0f, 2.0f }; // 2 logits for 5 labels + CHECK(f.run() == TRANSCRIBE_ERR_BACKEND); +} + +// The abort callback reaches the family, and the aborted run clears the result. +void test_abort() { + Fixture f; + CHECK(f.run() == TRANSCRIBE_OK); + g_check_abort = true; + transcribe_langid_set_abort_callback(f.session, [](void *) { return true; }, nullptr); + CHECK(f.run() == TRANSCRIBE_ERR_ABORTED); + CHECK(f.result().n_candidates == 0); + transcribe_langid_set_abort_callback(f.session, nullptr, nullptr); + CHECK(f.run() == TRANSCRIBE_OK); +} + +} // namespace + +int main() { + transcribe_log_set(nullptr, nullptr); + + if (transcribe::build_langid_labels({ "aa", "bb", "cc", "dd", "ee" }, { "Aa", "Bb", "Cc", "Dd", "Ee" }, { "xx=aa" }, + "test", g_labels) != TRANSCRIBE_OK) { + std::fprintf(stderr, "FAIL: fixture label table\n"); + return EXIT_FAILURE; + } + + test_label_table(); + test_role_checks_and_labels(); + test_ranking(); + test_allowed_set(); + test_crop_and_minimum(); + test_non_finite(); + test_abort(); + + if (g_failures != 0) { + std::fprintf(stderr, "%d failure(s)\n", g_failures); + return EXIT_FAILURE; + } + std::printf("langid_dispatch_unit: ok\n"); + return EXIT_SUCCESS; +} diff --git a/tests/mel_unit.cpp b/tests/mel_unit.cpp index 4033c5fa8..218785b7c 100644 --- a/tests/mel_unit.cpp +++ b/tests/mel_unit.cpp @@ -11,17 +11,26 @@ // - wrong mel scale formula (HTK vs Slaney) // - missing Slaney area normalization (would scale by ~50x) // +// Plus the SpeechBrain (ecapa_tdnn) options: the periodic Hamming window +// and the "sentence_mean" normalize (dB, top-dB floor, per-bin mean, no +// frame drop), run end to end on a synthetic one-hot filterbank. +// // All reference values are bit-precise constants captured from // librosa 0.11 and the symmetric-hann formula. Regenerate via the // preflight script if librosa updates and the tolerances drift. #include "transcribe-mel.h" +#include #include #include #include #include +#ifndef M_PI +# define M_PI 3.14159265358979323846 +#endif + namespace { int g_failures = 0; @@ -204,12 +213,112 @@ void test_n_frames_for() { CHECK(mf.n_frames_for(0) == 1); } +// SpeechBrain (ecapa_tdnn) config: periodic Hamming, constant pad, no +// pre-emphasis, 10*log10 with an 80 dB top-dB floor, per-bin mean only. +// The filterbank is a synthetic one-hot bank (mel m reads FFT bin 3*m), +// since the real one is checkpoint-provided. +transcribe::MelConfig speechbrain_config() { + transcribe::MelConfig cfg; + cfg.num_mels = 60; + cfg.n_fft = 400; + cfg.win_length = 400; + cfg.hop_length = 160; + cfg.pre_emphasis = 0.0f; + cfg.pad_mode = "constant"; + cfg.window_type = "hamming_periodic"; + cfg.normalize = "sentence_mean"; + cfg.log_clamp_min = 1e-10f; + cfg.top_db = 80.0f; + cfg.filterbank.assign(static_cast(60) * 201, 0.0f); + for (int m = 0; m < 60; ++m) { + cfg.filterbank[static_cast(m) * 201 + static_cast(m) * 3] = 1.0f; + } + return cfg; +} + +void test_hamming_window() { + transcribe::MelFrontend mf(speechbrain_config()); + const auto & w = mf.window(); + CHECK(w.size() == 400); + if (w.size() != 400) { + return; + } + // Periodic: the peak lands ON sample N/2 (a symmetric Hamming of + // length 400 never reaches 1.0), and w[1] uses 2*pi*n/N, not N-1. + CHECK_NEAR(w[0], 0.08, 1e-15); + CHECK_NEAR(w[100], 0.54, 1e-15); + CHECK_NEAR(w[200], 1.0, 1e-15); + CHECK_NEAR(w[1], 0.0800567490584361, 1e-15); + double sum = 0.0; + for (double v : w) { + sum += v; + } + CHECK_NEAR(sum, 216.0, 1e-12); // 0.54 * N: the cosine term cancels +} + +void test_sentence_mean() { + // 1 kHz sine, 1 s, through the full sentence_mean pipeline. + std::vector pcm(16000); + for (int i = 0; i < 16000; ++i) { + pcm[static_cast(i)] = static_cast(std::sin(2.0 * M_PI * 1000.0 * i / 16000.0)); + } + transcribe::MelFrontend mf(speechbrain_config()); + std::vector out; + int n_mels = 0; + int n_frames = 0; + CHECK(mf.compute(pcm.data(), pcm.size(), out, n_mels, n_frames, 1) == TRANSCRIBE_OK); + // torch center=True framing with no trailing-frame drop. + CHECK(n_mels == 60); + CHECK(n_frames == 101); + if (out.size() != static_cast(60) * 101) { + CHECK(false); + return; + } + // Mean removed per mel bin over all frames ([n_mels, n_frames] rows), + // and nothing else: a variance normalize would also pin each row's + // spread, so require one row with spread well above 1. + double max_spread = 0.0; + for (int m = 0; m < 60; ++m) { + const float * row = out.data() + static_cast(m) * 101; + double sum = 0.0; + float lo = row[0]; + float hi = row[0]; + for (int t = 0; t < 101; ++t) { + sum += static_cast(row[t]); + lo = std::min(lo, row[t]); + hi = std::max(hi, row[t]); + } + CHECK_NEAR(sum / 101.0, 0.0, 1e-4); + max_spread = std::max(max_spread, static_cast(hi - lo)); + // Top-dB floor: before the mean shift every value sat in + // [max_all - 80, max_all], so no row can span more than 80 dB. + CHECK(static_cast(hi - lo) <= 80.0 + 1e-3); + } + CHECK(max_spread > 1.0); + + // With the floor disabled the -100 dB (amin) bins reappear, so some + // row must now span more than 80 dB: proves the floor above did work. + transcribe::MelConfig no_floor_cfg = speechbrain_config(); + no_floor_cfg.top_db = 0.0f; + transcribe::MelFrontend no_floor(no_floor_cfg); + CHECK(no_floor.compute(pcm.data(), pcm.size(), out, n_mels, n_frames, 1) == TRANSCRIBE_OK); + double raw_spread = 0.0; + for (int m = 0; m < 60 && out.size() == static_cast(60) * 101; ++m) { + const float * row = out.data() + static_cast(m) * 101; + const auto mm = std::minmax_element(row, row + 101); + raw_spread = std::max(raw_spread, static_cast(*mm.second - *mm.first)); + } + CHECK(raw_spread > 80.0); +} + } // namespace int main() { test_window(); test_mel_filterbank(); test_n_frames_for(); + test_hamming_window(); + test_sentence_mean(); if (g_failures > 0) { std::fprintf(stderr, "mel_unit: %d failures\n", g_failures); diff --git a/tests/role_resolve_unit.cpp b/tests/role_resolve_unit.cpp index e39917b19..85ef5ca9f 100644 --- a/tests/role_resolve_unit.cpp +++ b/tests/role_resolve_unit.cpp @@ -2,6 +2,7 @@ #include "transcribe-arch.h" #include "transcribe-diarize.h" +#include "transcribe-langid.h" #include "transcribe-model.h" #include "transcribe.h" @@ -59,9 +60,21 @@ transcribe::Arch make_diarize_arch() { const transcribe::Arch k_diarize_arch = make_diarize_arch(); +const transcribe::LangidOps k_langid_ops = {}; + +transcribe::Arch make_langid_arch() { + transcribe::Arch a = {}; + a.name = "fake_langid"; + a.langid = &k_langid_ops; + return a; +} + +const transcribe::Arch k_langid_arch = make_langid_arch(); + void test_resolve_roles() { constexpr uint32_t ASR = TRANSCRIBE_ROLE_ASR; constexpr uint32_t DIA = TRANSCRIBE_ROLE_DIARIZE; + constexpr uint32_t LID = TRANSCRIBE_ROLE_LANGID; struct Case { const transcribe::Arch * arch; @@ -81,6 +94,11 @@ void test_resolve_roles() { { &k_asr_arch, DIA, TRANSCRIBE_ERR_NOT_IMPLEMENTED, 0 }, { &k_asr_arch, ASR | DIA, TRANSCRIBE_ERR_NOT_IMPLEMENTED, 0 }, { &k_diarize_arch, ASR | DIA, TRANSCRIBE_ERR_NOT_IMPLEMENTED, 0 }, + { &k_langid_arch, LID, TRANSCRIBE_OK, LID }, + { &k_langid_arch, 0, TRANSCRIBE_ERR_NOT_IMPLEMENTED, 0 }, + { &k_asr_arch, LID, TRANSCRIBE_ERR_NOT_IMPLEMENTED, 0 }, // no langid ops table + { &k_diarize_arch, DIA | LID, TRANSCRIBE_ERR_NOT_IMPLEMENTED, 0 }, + { &k_langid_arch, ASR | LID, TRANSCRIBE_ERR_NOT_IMPLEMENTED, 0 }, { &k_asr_arch, ASR | (1u << 31), TRANSCRIBE_ERR_NOT_IMPLEMENTED, 0 }, // unknown bit { nullptr, 0, TRANSCRIBE_ERR_NOT_IMPLEMENTED, 0 }, }; diff --git a/tests/tolerances/ecapa_tdnn.json b/tests/tolerances/ecapa_tdnn.json new file mode 100644 index 000000000..20979c704 --- /dev/null +++ b/tests/tolerances/ecapa_tdnn.json @@ -0,0 +1,123 @@ +{ + "_comment": [ + "ECAPA-TDNN / VoxLingua107 per-tensor tolerances for compare_tensors.py.", + "", + "Carried over unchanged from langid.cpp: on the port, every stage tensor the", + "transcribe.cpp ecapa_tdnn family dumps was bit-identical to langid.cpp's", + "(CPU, 1 thread, F32), so the measured drift below still holds.", + "", + "Measured in langid.cpp (handy-computer/langid.cpp @ 96af3a5) against C++ vs", + "reference drift. Every Stage-2 provisional flag has been removed.", + "", + "Regime, identical for every number below:", + " model models/lang-id-voxlingua107-ecapa/lang-id-voxlingua107-ecapa-F32.gguf", + " (ALL_F32; no quantised tensor takes part in these numbers)", + " runtime build/bin/transcribe-cli --backend cpu --threads 1", + " front end the production C++ log-mel front end (src/transcribe-mel.cpp,", + " normalize=sentence_mean; measured on its family-local predecessor,", + " which ran the same fp32 mixed-radix FFT),", + " not a reference feature dump -- fe.mel is measured, not injected", + " host macOS / arm64 (Apple Silicon) and Linux / x86-64 (AVX2+FMA,", + " Ryzen 7 PRO 4750U), ggml CPU backend", + " reference speechbrain 1.1.1 on torch 2.13.0, fp32, CPU, 1 thread,", + " checkpoint speechbrain/lang-id-voxlingua107-ecapa @", + " 0253049ae131d6a4be1c4f0d8b0ff483a0f8c8e9", + " cases all 8 in tests/golden/ecapa_tdnn/lang-id-voxlingua107-ecapa.manifest.json", + " (fleurs-{en,de,fr,es,ja,zh,ru,id}, 7.9 s .. 11.2 s)", + "", + "F16 / Q8_0 / Metal are production variants measured separately and were", + "never used to set these numbers.", + "", + "Recipe, per tensor, over all 8 cases:", + " max_abs = max(1.5 * max_over_cases(observed max_abs), provisional, 1e-6)", + " mean_abs = max(1.5 * max_over_cases(observed mean_abs), provisional, 1e-6)", + "The Stage-2 provisional budgets were max(1e-4*p99_abs, 1e-6) and", + "max(1e-5*rms, 1e-6) from the reference dumps alone.", + "", + "WIDENED ENTRIES (max_abs only, from the x86 host; every other value is", + "the unchanged Stage-2 number):", + " enc.mfa.out observed 2.3806e-04 (fleurs-zh) -> 3.571e-04", + " enc.blk.2.out observed 2.0504e-04 (fleurs-zh) -> 3.076e-04", + "Cause: fp32 FMA reduction order over the 3072- and 1024-wide", + "contractions on x86 AVX2, about 2x the arm64 drift on these tensors.", + "mean_abs on both stays under 7% of budget, and the drift contracts", + "again downstream (enc.emb, cls.logits_raw under 1%).", + "", + "Dominant drift source: fp32 summation ORDER, not algorithm. Two", + "contributors, in order of size:", + " 1. fe.mel -- the C++ mixed-radix 400-point FFT (src/transcribe-mel.cpp)", + " accumulates the STFT differently from torch.stft's FFT, and the", + " difference is then amplified by 10*log10 on near-floor bins. This", + " is the only stage whose drift is not inherited.", + " 2. every matmul -- ggml's blocked/SIMD reduction order over the 1024-", + " and 3072-wide contractions differs from PyTorch's BLAS. The graph", + " itself is exact: an independent fp64 NumPy re-derivation of the", + " forward pass agrees with the ggml graph to 4.5e-07 relative on a", + " synthetic fixture, so no term of the ECAPA algebra is approximated.", + "The residual-add topology means block drift accumulates roughly linearly", + "(blk.0 5.2e-05 -> blk.3 1.5e-04), then the attention softmax and the", + "1/T statistics CONTRACT it again (enc.asp.out 9.5e-06), which is why the", + "downstream vectors sit two orders of magnitude inside their budgets.", + "", + "Tightest entry: enc.blk.0.out, where 1.5*observed (arm64) is 99% of", + "the budget; max_abs on a wide [T, C] activation.", + "", + "cls.logits_raw is the gate." + ], + "fe.mel": { + "max_abs": 0.004079, + "mean_abs": 0.0001782 + }, + "enc.blk.0.out": { + "max_abs": 7.896e-05, + "mean_abs": 1.897e-06 + }, + "enc.blk.1.tdnn1.out": { + "max_abs": 0.0004054, + "mean_abs": 1.008e-05 + }, + "enc.blk.1.res2.out": { + "max_abs": 0.0003865, + "mean_abs": 9.456e-06 + }, + "enc.blk.1.se.out": { + "max_abs": 0.0001077, + "mean_abs": 2.4e-06 + }, + "enc.blk.1.out": { + "max_abs": 0.0001297, + "mean_abs": 2.952e-06 + }, + "enc.blk.2.out": { + "max_abs": 0.0003076, + "mean_abs": 5.964e-06 + }, + "enc.blk.3.out": { + "max_abs": 0.0006391, + "mean_abs": 1.412e-05 + }, + "enc.mfa.out": { + "max_abs": 0.0003571, + "mean_abs": 4.163e-06 + }, + "enc.asp.attn_logits": { + "max_abs": 0.0003274, + "mean_abs": 1.19e-05 + }, + "enc.asp.out": { + "max_abs": 0.0001753, + "mean_abs": 4.176e-06 + }, + "enc.emb": { + "max_abs": 0.006139, + "mean_abs": 0.0002736 + }, + "cls.hidden": { + "max_abs": 0.002295, + "mean_abs": 8.541e-05 + }, + "cls.logits_raw": { + "max_abs": 0.001809, + "mean_abs": 0.0001014 + } +}