mirror of
https://github.com/ggml-org/whisper.cpp.git
synced 2026-08-12 14:25:23 +04: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:
Binary file not shown.
@@ -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)
|
||||
@@ -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;
|
||||
|
||||
@@ -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
@@ -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;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user