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 index 000000000..bba7d72b2 Binary files /dev/null and b/models/for-tests-ggml-parakeet-tdt-bad-nfft0.bin differ diff --git a/models/generate-parakeet-test-model.py b/models/generate-parakeet-test-model.py index 192a96ce6..8b31a042f 100755 --- a/models/generate-parakeet-test-model.py +++ b/models/generate-parakeet-test-model.py @@ -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 diff --git a/src/parakeet-arch.h b/src/parakeet-arch.h index 3407a95c9..e8c6effe4 100644 --- a/src/parakeet-arch.h +++ b/src/parakeet-arch.h @@ -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_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_INFO = { {PARAKEET_TENSOR_JOINT_NET_WEIGHT, GGML_OP_MUL_MAT}, {PARAKEET_TENSOR_JOINT_NET_BIAS, GGML_OP_ADD}, }; + +static const std::map 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_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}, +}; diff --git a/src/parakeet.cpp b/src/parakeet.cpp index b5da73e98..178f049b1 100644 --- a/src/parakeet.cpp +++ b/src/parakeet.cpp @@ -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 & 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::maphparam_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; diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 74a5b1429..aecc6f3b2 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -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") diff --git a/tests/test-parakeet.cpp b/tests/test-parakeet.cpp index 83237c600..58b64835d 100644 --- a/tests/test-parakeet.cpp +++ b/tests/test-parakeet.cpp @@ -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; +}