]> git.djapps.eu Git - pkg/ggml/sources/whisper.cpp/commitdiff
parakeet : verify hparams loaded from parakeet model bin file (#3950)
authorBhargav Krish <redacted>
Thu, 30 Jul 2026 04:59:23 +0000 (21:59 -0700)
committerGitHub <redacted>
Thu, 30 Jul 2026 04:59:23 +0000 (06:59 +0200)
* verify hparams loaded from parakeet model bin file

* flexible way to accommodate CI as well security concern & test case addition.

* add bad model for CI tests

* removing whitespaces,couple of nits

models/for-tests-ggml-parakeet-tdt-bad-nfft0.bin [new file with mode: 0644]
models/generate-parakeet-test-model.py
src/parakeet-arch.h
src/parakeet.cpp
tests/CMakeLists.txt
tests/test-parakeet.cpp

diff --git a/models/for-tests-ggml-parakeet-tdt-bad-nfft0.bin b/models/for-tests-ggml-parakeet-tdt-bad-nfft0.bin
new file mode 100644 (file)
index 0000000..bba7d72
Binary files /dev/null and b/models/for-tests-ggml-parakeet-tdt-bad-nfft0.bin differ
index 192a96ce62716af8cf1c9bf509ba8bcf3c006770..8b31a042f5cdefa19f4ef8c2f4d67a17ad38ec45 100755 (executable)
@@ -3,6 +3,7 @@ import struct
 import sys
 import numpy as np
 from pathlib import Path
+import argparse
 
 def write_tensor(fout, name, data):
     n_dims = len(data.shape)
@@ -16,7 +17,7 @@ def write_tensor(fout, name, data):
     fout.write(name_bytes)
     data.tofile(fout)
 
-def generate(output_path):
+def generate(output_path, n_fft_override=None):
     rng = np.random.default_rng(42)
 
     hparams = {
@@ -37,6 +38,9 @@ def generate(output_path):
         'n_max_tokens':           5,
     }
 
+    if n_fft_override is not None:
+        hparams['n_fft'] = n_fft_override
+
     n_vocab    = hparams['n_vocab']
     n_state    = hparams['n_audio_state']
     n_head     = hparams['n_audio_head']
@@ -178,5 +182,8 @@ def generate(output_path):
     print(f"Generated {output_path} ({size / 1024:.1f} KB)")
 
 if __name__ == '__main__':
-    output = sys.argv[1] if len(sys.argv) > 1 else 'models/for-tests-ggml-parakeet-tdt.bin'
-    generate(output)
+    parser = argparse.ArgumentParser()
+    parser.add_argument('output', nargs='?', default='models/for-tests-ggml-parakeet-tdt.bin')
+    parser.add_argument('--n-fft',type=int, default=None)
+    args = parser.parse_args()
+    generate(args.output, args.n_fft)
\ No newline at end of file
index 3407a95c9c7dbf7f522903c48b00f3eef931d3d3..e8c6effe4a35c8a390c774503c6e321247206533 100644 (file)
@@ -65,6 +65,23 @@ enum parakeet_tensor {
     PARAKEET_TENSOR_JOINT_NET_BIAS,
 };
 
+enum parakeet_hparam {
+    PARAKEET_HPARAM_N_VOCAB,
+    PARAKEET_HPARAM_N_AUDIO_CTX,
+    PARAKEET_HPARAM_N_AUDIO_STATE,
+    PARAKEET_HPARAM_N_AUDIO_HEAD,
+    PARAKEET_HPARAM_N_AUDIO_LAYER,
+    PARAKEET_HPARAM_N_MELS,
+    PARAKEET_HPARAM_N_FFT,
+    PARAKEET_HPARAM_SUBSAMPLING_FACTOR,
+    PARAKEET_HPARAM_N_SUBSAMPLING_CHANNELS,
+    PARAKEET_HPARAM_N_CONV_KERNEL,
+    PARAKEET_HPARAM_N_PRED_DIM,
+    PARAKEET_HPARAM_N_PRED_LAYERS,
+    PARAKEET_HPARAM_N_TDT_DURATIONS,
+    PARAKEET_HPARAM_N_MAX_TOKENS,
+};
+
 static const std::map<parakeet_tensor, const char *> PARAKEET_TENSOR_NAMES = {
     // Encoder pre_encode
     {PARAKEET_TENSOR_ENC_PRE_OUT_WEIGHT,          "encoder.pre_encode.out.weight"},
@@ -186,3 +203,37 @@ static const std::map<parakeet_tensor, ggml_op> PARAKEET_TENSOR_INFO = {
     {PARAKEET_TENSOR_JOINT_NET_WEIGHT,            GGML_OP_MUL_MAT},
     {PARAKEET_TENSOR_JOINT_NET_BIAS,              GGML_OP_ADD},
 };
+
+static const std::map<parakeet_hparam, const char *> PARAKEET_HPARAM_NAMES = {
+    {PARAKEET_HPARAM_N_VOCAB,                "n_vocab"},
+    {PARAKEET_HPARAM_N_AUDIO_CTX,            "n_audio_ctx"},
+    {PARAKEET_HPARAM_N_AUDIO_STATE,          "n_audio_state"},
+    {PARAKEET_HPARAM_N_AUDIO_HEAD,           "n_audio_head"},
+    {PARAKEET_HPARAM_N_AUDIO_LAYER,          "n_audio_layer"},
+    {PARAKEET_HPARAM_N_MELS,                 "n_mels"},
+    {PARAKEET_HPARAM_N_FFT,                  "n_fft"},
+    {PARAKEET_HPARAM_SUBSAMPLING_FACTOR,     "subsampling_factor"},
+    {PARAKEET_HPARAM_N_SUBSAMPLING_CHANNELS, "n_subsampling_channels"},
+    {PARAKEET_HPARAM_N_CONV_KERNEL,          "n_conv_kernel"},
+    {PARAKEET_HPARAM_N_PRED_DIM,             "n_pred_dim"},
+    {PARAKEET_HPARAM_N_PRED_LAYERS,          "n_pred_layers"},
+    {PARAKEET_HPARAM_N_TDT_DURATIONS,        "n_tdt_durations"},
+    {PARAKEET_HPARAM_N_MAX_TOKENS,           "n_max_tokens"},
+};
+
+static const std::map<parakeet_hparam, int32_t> PARAKEET_HPARAM_MODEL_VALUES = {
+    {PARAKEET_HPARAM_N_VOCAB,                8192},
+    {PARAKEET_HPARAM_N_AUDIO_CTX,            5000},
+    {PARAKEET_HPARAM_N_AUDIO_STATE,          1024},
+    {PARAKEET_HPARAM_N_AUDIO_HEAD,              8},
+    {PARAKEET_HPARAM_N_AUDIO_LAYER,            24},
+    {PARAKEET_HPARAM_N_MELS,                  128},
+    {PARAKEET_HPARAM_N_FFT,                   512},
+    {PARAKEET_HPARAM_SUBSAMPLING_FACTOR,        8},
+    {PARAKEET_HPARAM_N_SUBSAMPLING_CHANNELS,  256},
+    {PARAKEET_HPARAM_N_CONV_KERNEL,             9},
+    {PARAKEET_HPARAM_N_PRED_DIM,              640},
+    {PARAKEET_HPARAM_N_PRED_LAYERS,             2},
+    {PARAKEET_HPARAM_N_TDT_DURATIONS,           5},
+    {PARAKEET_HPARAM_N_MAX_TOKENS,             10},
+};
index b5da73e985c0b0955e5b7ded4f7bee86f6b06823..178f049b1bff7406dbaf2c5baa53a6224bd160cd 100644 (file)
@@ -685,6 +685,34 @@ static void read_safe(parakeet_model_loader * loader, T & dest) {
     BYTESWAP_VALUE(dest);
 }
 
+
+static bool parakeet_validate_hparams(const std::map<parakeet_hparam, int32_t> & hparam_values) {
+    for (const auto & hparam_expected : PARAKEET_HPARAM_MODEL_VALUES) {
+        const parakeet_hparam hparam = hparam_expected.first;
+        const auto hparam_value = hparam_values.find(hparam);
+        if (hparam_value == hparam_values.end()) {
+            PARAKEET_LOG_ERROR("%s: missing Parakeet metadata: %s\n",
+                    __func__, PARAKEET_HPARAM_NAMES.at(hparam));
+            return false;
+        }
+
+        const int32_t actual = hparam_value->second;
+        const int32_t expected = hparam_expected.second;
+        if(actual <=0 || actual > expected){
+            PARAKEET_LOG_ERROR("%s: invalid Parakeet metadata: %s = %d, expected > 0 and <= %d\n. Unsafe parameter loaded. ",
+                __func__, PARAKEET_HPARAM_NAMES.at(hparam), actual, expected);
+            return false;
+        }
+        if(actual != expected){
+            PARAKEET_LOG_WARN("%s: non-standard Parakeet metadata: %s = %d, expected %d\n. Transcription will be affected. ",
+                __func__, PARAKEET_HPARAM_NAMES.at(hparam), actual, expected);
+        }
+
+    }
+
+    return true;
+}
+
 static bool parakeet_lstm_state_init(
                struct parakeet_state & pstate,
                       ggml_backend_t   backend,
@@ -1003,21 +1031,33 @@ static bool parakeet_model_load(struct parakeet_model_loader * loader, parakeet_
     //load hparams
     parakeet_hparams hparams;
     {
-        read_safe(loader, hparams.n_vocab);
-        read_safe(loader, hparams.n_audio_ctx);
-        read_safe(loader, hparams.n_audio_state);
-        read_safe(loader, hparams.n_audio_head);
-        read_safe(loader, hparams.n_audio_layer);
-        read_safe(loader, hparams.n_mels);
+        std::map<parakeet_hparam, int32_t>hparam_values;
+        auto read_hparam = [&] (parakeet_hparam hparam, int32_t &value){
+            read_safe(loader, value);
+            hparam_values[hparam] = value;
+        };
+        read_hparam(PARAKEET_HPARAM_N_VOCAB, hparams.n_vocab);
+        read_hparam(PARAKEET_HPARAM_N_AUDIO_CTX, hparams.n_audio_ctx);
+        read_hparam(PARAKEET_HPARAM_N_AUDIO_STATE, hparams.n_audio_state);
+        read_hparam(PARAKEET_HPARAM_N_AUDIO_HEAD, hparams.n_audio_head);
+        read_hparam(PARAKEET_HPARAM_N_AUDIO_LAYER, hparams.n_audio_layer);
+        read_hparam(PARAKEET_HPARAM_N_MELS, hparams.n_mels);
+        /*
+        ftype just requires the type check already being done in the loading process.
+        */
         read_safe(loader, hparams.ftype);
-        read_safe(loader, hparams.n_fft);
-        read_safe(loader, hparams.subsampling_factor);
-        read_safe(loader, hparams.n_subsampling_channels);
-        read_safe(loader, hparams.n_conv_kernel);
-        read_safe(loader, hparams.n_pred_dim);
-        read_safe(loader, hparams.n_pred_layers);
-        read_safe(loader, hparams.n_tdt_durations);
-        read_safe(loader, hparams.n_max_tokens);
+        read_hparam(PARAKEET_HPARAM_N_FFT, hparams.n_fft);
+        read_hparam(PARAKEET_HPARAM_SUBSAMPLING_FACTOR, hparams.subsampling_factor);
+        read_hparam(PARAKEET_HPARAM_N_SUBSAMPLING_CHANNELS, hparams.n_subsampling_channels);
+        read_hparam(PARAKEET_HPARAM_N_CONV_KERNEL, hparams.n_conv_kernel);
+        read_hparam(PARAKEET_HPARAM_N_PRED_DIM, hparams.n_pred_dim);
+        read_hparam(PARAKEET_HPARAM_N_PRED_LAYERS, hparams.n_pred_layers);
+        read_hparam(PARAKEET_HPARAM_N_TDT_DURATIONS, hparams.n_tdt_durations);
+        read_hparam(PARAKEET_HPARAM_N_MAX_TOKENS, hparams.n_max_tokens);
+
+        if(!parakeet_validate_hparams(hparam_values)) {
+            return false;
+        }
 
         hparams.arch = PARAKEET_ARCH_TDT;
         wctx.model.hparams = hparams;
index 74a5b142948812f7c71a487adb2efe04e951f0d7..aecc6f3b24aa81b18f6c27b014515e330aee9c79 100644 (file)
@@ -126,6 +126,7 @@ target_include_directories(${PARAKEET_TEST} PRIVATE ../include ../ggml/include .
 target_link_libraries(${PARAKEET_TEST} PRIVATE parakeet common)
 target_compile_definitions(${PARAKEET_TEST} PRIVATE
     PARAKEET_MODEL_PATH="${PROJECT_SOURCE_DIR}/models/for-tests-ggml-parakeet-tdt.bin"
+    PARAKEET_BAD_MODEL_PATH="${PROJECT_SOURCE_DIR}/models/for-tests-ggml-parakeet-tdt-bad-nfft0.bin"
     SAMPLE_PATH="${PROJECT_SOURCE_DIR}/samples/jfk.wav")
 add_test(NAME ${PARAKEET_TEST} COMMAND ${PARAKEET_TEST})
 set_tests_properties(${PARAKEET_TEST} PROPERTIES LABELS "parakeet;gh")
index 83237c600ac6ba9ed9edf55e6789a078f3c77c01..58b64835de29f2f7f6006f4fb25c9e56d9a7748f 100644 (file)
@@ -59,7 +59,19 @@ void segment_callback(parakeet_context * ctx, parakeet_state * state, int n_new,
     printf("\n");
 }
 
-int main() {
+static int test_invalid_model_load(){
+    struct parakeet_context_params ctx_params = parakeet_context_default_params();
+    struct parakeet_context * pctx =
+        parakeet_init_from_file_with_params_no_state(PARAKEET_BAD_MODEL_PATH, ctx_params);
+    if(pctx != nullptr){
+        fprintf(stderr, "Expected invalid Parakeet model to fail loading \n");
+        parakeet_free(pctx);
+        return 1;
+    }
+    return 0;
+}
+
+static int test_valid_model() {
     std::string model_path  = PARAKEET_MODEL_PATH;
     std::string sample_path = SAMPLE_PATH;
 
@@ -97,3 +109,16 @@ int main() {
     printf("\nTest passed: Parakeet model loaded and freed successfully\n");
     return 0;
 }
+
+int main(){
+    if(test_valid_model() != 0){
+        return 1;
+    }
+
+    if(test_invalid_model_load() != 0){
+        return 1;
+    }
+
+    printf("\nTest passed: Parakeet model load tests completed successfully\n");
+    return 0;
+}