mirror of
https://github.com/ggml-org/whisper.cpp.git
synced 2026-10-08 07:21:21 +02:00
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:
@@ -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
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user