parakeet : verify hparams loaded from parakeet model bin file (#3950)

* 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
This commit is contained in:
Bhargav Krish
2026-07-30 06:59:23 +02:00
committed by GitHub
parent a630b35c6f
commit 4523d0ce37
6 changed files with 142 additions and 18 deletions
Binary file not shown.
+10 -3
View File
@@ -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)
+51
View 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},
};
+54 -14
View 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;
+1
View 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")
+26 -1
View 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;
}