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
+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;