import sys
import numpy as np
from pathlib import Path
+import argparse
def write_tensor(fout, name, data):
n_dims = len(data.shape)
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 = {
'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']
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
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"},
{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},
+};
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,
//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;
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")
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;
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;
+}