From dbdc6a24811889814c0956a991565ebd1235efce Mon Sep 17 00:00:00 2001 From: 0xShug0 <231717474+0xShug0@users.noreply.github.com> Date: Thu, 2 Jul 2026 22:20:10 -0400 Subject: [PATCH 1/5] Add Kokoro preview integration --- CMakeLists.txt | 8 + README.md | 3 + include/engine/models/kokoro_tts/assets.h | 283 +++ include/engine/models/kokoro_tts/decoder.h | 42 + include/engine/models/kokoro_tts/frontend.h | 50 + include/engine/models/kokoro_tts/g2p_en.h | 159 ++ include/engine/models/kokoro_tts/loader.h | 33 + include/engine/models/kokoro_tts/plbert.h | 35 + include/engine/models/kokoro_tts/predictor.h | 57 + include/engine/models/kokoro_tts/session.h | 73 + src/framework/runtime/registry.cpp | 4 +- src/models/kokoro_tts/assets.cpp | 835 +++++++ src/models/kokoro_tts/decoder.cpp | 1352 ++++++++++++ src/models/kokoro_tts/frontend.cpp | 285 +++ src/models/kokoro_tts/g2p_en.cpp | 1866 ++++++++++++++++ src/models/kokoro_tts/loader.cpp | 121 ++ src/models/kokoro_tts/plbert.cpp | 442 ++++ src/models/kokoro_tts/predictor.cpp | 2054 ++++++++++++++++++ src/models/kokoro_tts/session.cpp | 481 ++++ tools/model_manager.py | 17 +- 20 files changed, 8194 insertions(+), 6 deletions(-) create mode 100644 include/engine/models/kokoro_tts/assets.h create mode 100644 include/engine/models/kokoro_tts/decoder.h create mode 100644 include/engine/models/kokoro_tts/frontend.h create mode 100644 include/engine/models/kokoro_tts/g2p_en.h create mode 100644 include/engine/models/kokoro_tts/loader.h create mode 100644 include/engine/models/kokoro_tts/plbert.h create mode 100644 include/engine/models/kokoro_tts/predictor.h create mode 100644 include/engine/models/kokoro_tts/session.h create mode 100644 src/models/kokoro_tts/assets.cpp create mode 100644 src/models/kokoro_tts/decoder.cpp create mode 100644 src/models/kokoro_tts/frontend.cpp create mode 100644 src/models/kokoro_tts/g2p_en.cpp create mode 100644 src/models/kokoro_tts/loader.cpp create mode 100644 src/models/kokoro_tts/plbert.cpp create mode 100644 src/models/kokoro_tts/predictor.cpp create mode 100644 src/models/kokoro_tts/session.cpp diff --git a/CMakeLists.txt b/CMakeLists.txt index 9d35a92d5..6a2736964 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -195,6 +195,14 @@ add_library(engine_runtime STATIC src/models/omnivoice/postprocess.cpp src/models/omnivoice/session.cpp src/models/omnivoice/loader.cpp + src/models/kokoro_tts/assets.cpp + src/models/kokoro_tts/loader.cpp + src/models/kokoro_tts/frontend.cpp + src/models/kokoro_tts/g2p_en.cpp + src/models/kokoro_tts/predictor.cpp + src/models/kokoro_tts/decoder.cpp + src/models/kokoro_tts/plbert.cpp + src/models/kokoro_tts/session.cpp src/models/pocket_tts/assets.cpp src/models/pocket_tts/acoustic_model.cpp src/models/pocket_tts/audio_decoder.cpp diff --git a/README.md b/README.md index f3d5491b4..9576ae1c5 100644 --- a/README.md +++ b/README.md @@ -1,5 +1,8 @@ # audio.cpp +> [!IMPORTANT] +> **Kokoro 82M preview:** the framework path is wired and runs in real time through the normal CLI, but it is not faster than the Python implementation yet. Treat this branch as a correctness and integration preview while optimization continues. + `audio.cpp` is a high-performance C++ audio inference framework built on top of `ggml`, designed to make modern local audio models practical, portable, and fast. Tired of juggling a dozen Conda environments, hundreds of Python packages, and dependency conflicts just to try a few audio models? audio.cpp gives those paths a shared native runtime instead. diff --git a/include/engine/models/kokoro_tts/assets.h b/include/engine/models/kokoro_tts/assets.h new file mode 100644 index 000000000..480515320 --- /dev/null +++ b/include/engine/models/kokoro_tts/assets.h @@ -0,0 +1,283 @@ +#pragma once + +#include "engine/framework/assets/tensor_source.h" +#include "engine/framework/core/backend_weight_store.h" +#include "engine/framework/core/module.h" +#include "engine/framework/io/json.h" + +#include + +#include +#include +#include +#include +#include +#include +#include + +namespace kokoro_ggml { +struct KokoroWeights { + struct LinearWeights { + engine::core::TensorValue weight; + std::optional bias; + int64_t out_features = 0; + int64_t in_features = 0; + bool use_bias = false; + }; + + struct HostAffineWeights { + std::vector weight; + std::vector bias; + int64_t out_features = 0; + int64_t in_features = 0; + bool use_bias = false; + }; + + struct EmbeddingWeights { + engine::core::TensorValue weight; + int64_t num_embeddings = 0; + int64_t embedding_dim = 0; + }; + + struct LayerNormWeights { + engine::core::TensorValue weight; + engine::core::TensorValue bias; + int64_t channels = 0; + float eps = 1.0e-5f; + }; + + struct WeightNormConv1dWeights { + engine::core::TensorValue weight; + std::optional bias; + int64_t out_channels = 0; + int64_t in_channels = 0; + int64_t kernel = 0; + int64_t stride = 1; + int64_t padding = 0; + int64_t dilation = 1; + int64_t groups = 1; + bool use_bias = false; + }; + + struct Conv1dWeights { + engine::core::TensorValue weight; + std::optional bias; + int64_t out_channels = 0; + int64_t in_channels = 0; + int64_t kernel = 0; + int64_t stride = 1; + int64_t padding = 0; + int64_t dilation = 1; + int64_t groups = 1; + bool use_bias = false; + }; + + struct WeightNormConvTranspose1dWeights { + engine::core::TensorValue weight; + engine::core::TensorValue dense_weight; + std::optional bias; + std::shared_ptr phase_shuffle_conv; + int64_t in_channels = 0; + int64_t out_channels = 0; + int64_t kernel = 0; + int64_t stride = 1; + int64_t padding = 0; + int64_t output_padding = 0; + int64_t groups = 1; + bool use_bias = false; + }; + + struct LstmWeights { + engine::core::TensorValue weight_ih_l0; + engine::core::TensorValue weight_hh_l0; + engine::core::TensorValue bias_ih_l0; + engine::core::TensorValue bias_hh_l0; + engine::core::TensorValue combined_bias_l0; + engine::core::TensorValue weight_ih_l0_reverse; + engine::core::TensorValue weight_hh_l0_reverse; + engine::core::TensorValue bias_ih_l0_reverse; + engine::core::TensorValue bias_hh_l0_reverse; + engine::core::TensorValue combined_bias_l0_reverse; + int64_t input_size = 0; + int64_t hidden_size = 0; + }; + + struct AdaLayerNormWeights { + LinearWeights fc; + int64_t channels = 0; + float eps = 1.0e-5f; + }; + + struct AdaIn1dWeights { + LinearWeights fc; + int64_t channels = 0; + float eps = 1.0e-5f; + }; + + struct AlbertEmbeddingsWeights { + EmbeddingWeights word_embeddings; + EmbeddingWeights position_embeddings; + EmbeddingWeights token_type_embeddings; + LayerNormWeights layer_norm; + }; + + struct AlbertAttentionWeights { + LinearWeights query; + LinearWeights key; + LinearWeights value; + LinearWeights dense; + LayerNormWeights layer_norm; + }; + + struct AlbertLayerWeights { + AlbertAttentionWeights attention; + LinearWeights ffn; + LinearWeights ffn_output; + LayerNormWeights full_layer_layer_norm; + }; + + struct AlbertWeights { + AlbertEmbeddingsWeights embeddings; + LinearWeights embedding_hidden_mapping_in; + AlbertLayerWeights shared_layer; + LinearWeights pooler; + int64_t hidden_size = 768; + int64_t embedding_size = 128; + int64_t intermediate_size = 2048; + int64_t max_position_embeddings = 512; + int64_t num_hidden_layers = 12; + int64_t num_attention_heads = 12; + float layer_norm_eps = 1.0e-12f; + }; + + struct TextEncoderBlockWeights { + WeightNormConv1dWeights conv; + LayerNormWeights layer_norm; + }; + + struct TextEncoderWeights { + EmbeddingWeights embedding; + std::vector cnn; + LstmWeights lstm; + }; + + struct DurationEncoderWeights { + std::vector lstms; + std::vector ada_layer_norms; + }; + + struct AdainResBlock1dWeights { + WeightNormConv1dWeights conv1; + WeightNormConv1dWeights conv2; + WeightNormConv1dWeights conv1x1; + WeightNormConvTranspose1dWeights pool; + AdaIn1dWeights norm1; + AdaIn1dWeights norm2; + bool learned_sc = false; + bool use_pool = false; + bool upsample = false; + }; + + struct GeneratorResBlockWeights { + std::vector convs1; + std::vector convs2; + std::vector adain1; + std::vector adain2; + std::vector alpha1; + std::vector alpha2; + }; + + struct GeneratorWeights { + std::vector ups; + std::vector noise_convs; + std::vector noise_res; + std::vector resblocks; + WeightNormConv1dWeights conv_post; + HostAffineWeights source_linear; + int64_t harmonic_num = 8; + int64_t sampling_rate = 24000; + float sine_amp = 0.1f; + float noise_std = 0.003f; + float voiced_threshold = 10.0f; + int64_t gen_istft_n_fft = 20; + int64_t gen_istft_hop_size = 5; + }; + + struct ProsodyPredictorWeights { + DurationEncoderWeights duration_encoder; + LstmWeights lstm; + LinearWeights duration_proj; + LstmWeights shared; + std::vector f0_blocks; + std::vector n_blocks; + Conv1dWeights f0_proj; + Conv1dWeights n_proj; + }; + + struct DecoderWeights { + AdainResBlock1dWeights encode; + std::vector decode; + WeightNormConv1dWeights f0_conv; + WeightNormConv1dWeights n_conv; + WeightNormConv1dWeights asr_res; + GeneratorWeights generator; + }; + + AlbertWeights bert; + LinearWeights bert_encoder; + ProsodyPredictorWeights predictor; + TextEncoderWeights text_encoder; + DecoderWeights decoder; + + int64_t n_token = 178; + int64_t hidden_dim = 512; + int64_t style_dim = 128; + int64_t n_layer = 3; + int64_t max_dur = 50; + int64_t n_mels = 80; + int64_t context_length = 512; + int64_t max_output_tokens = 512 * 50; + float dropout = 0.2f; + float lrelu_slope = 0.1f; + float post_lrelu_slope = 0.01f; + + std::shared_ptr store; +}; +namespace g2p_en { +class EnglishG2P; +} +} + +namespace engine::models::kokoro_tts { + +struct KokoroVoicePack { + std::string id; + std::string language_code; + int64_t rows = 0; + int64_t cols = 0; + std::vector values; +}; + +struct KokoroAssets { + std::filesystem::path model_root; + engine::io::json::Value config; + std::shared_ptr model_weights; + int64_t context_length = 512; + std::unordered_map vocab; + std::unordered_map voices; + std::filesystem::path english_lexicon_dir; + std::shared_ptr english_g2p_us; + std::shared_ptr english_g2p_gb; +}; + +std::shared_ptr load_kokoro_assets(const std::filesystem::path & model_root); + +std::shared_ptr load_kokoro_backend_weights( + const KokoroAssets & assets, + ggml_backend_t backend, + engine::core::BackendType backend_type, + engine::assets::TensorStorageType matmul_storage_type, + engine::assets::TensorStorageType conv_storage_type, + size_t weight_context_bytes); + +} // namespace engine::models::kokoro_tts diff --git a/include/engine/models/kokoro_tts/decoder.h b/include/engine/models/kokoro_tts/decoder.h new file mode 100644 index 000000000..20f942cb3 --- /dev/null +++ b/include/engine/models/kokoro_tts/decoder.h @@ -0,0 +1,42 @@ +#pragma once + +#include "engine/models/kokoro_tts/predictor.h" + +#include +#include + +typedef struct ggml_backend * ggml_backend_t; + +namespace kokoro_ggml { + +struct KokoroDecoderCapacityContract { + int64_t decoder_frames = 0; + int64_t conditioning_frames = 0; +}; + +class KokoroDecoderRuntime { +public: + KokoroDecoderRuntime( + std::shared_ptr weights, + ggml_backend_t backend, + int n_threads, + bool use_device_backend, + uint64_t rng_seed, + KokoroDecoderCapacityContract contract); + ~KokoroDecoderRuntime(); + + KokoroDecoderRuntime(const KokoroDecoderRuntime &) = delete; + KokoroDecoderRuntime & operator=(const KokoroDecoderRuntime &) = delete; + + void prepare(KokoroDecoderCapacityContract contract); + + std::vector decode( + const PredictorOutputs & predictor, + const std::vector & ref_s); + +private: + struct Impl; + std::unique_ptr impl_; +}; + +} // namespace kokoro_ggml diff --git a/include/engine/models/kokoro_tts/frontend.h b/include/engine/models/kokoro_tts/frontend.h new file mode 100644 index 000000000..d092f207f --- /dev/null +++ b/include/engine/models/kokoro_tts/frontend.h @@ -0,0 +1,50 @@ +#pragma once + +#include "engine/framework/runtime/session.h" +#include "engine/models/kokoro_tts/assets.h" + +#include +#include +#include +#include + +namespace engine::models::kokoro_tts { + +struct KokoroSynthesisInput { + std::string voice_id; + std::string language_code; + std::string phonemes; + std::vector input_ids; + std::vector style; + float speaking_rate = 1.0f; +}; + +struct KokoroFrontendSessionState { + std::string voice_id; + std::string language_code; + const KokoroVoicePack * voice_pack = nullptr; + float speaking_rate = 1.0f; +}; + +KokoroFrontendSessionState resolve_kokoro_frontend_session_state( + const std::optional & text, + const std::optional & voice, + const KokoroAssets & assets); + +void validate_kokoro_frontend_session_state( + const runtime::Transcript & text, + const std::optional & voice, + const KokoroFrontendSessionState & state, + const KokoroAssets & assets); + +KokoroSynthesisInput build_kokoro_synthesis_input( + const runtime::Transcript & text, + const KokoroFrontendSessionState & state, + const KokoroAssets & assets); + +int64_t estimate_kokoro_request_tokens( + const runtime::SessionPreparationRequest & request, + const KokoroFrontendSessionState & state, + const KokoroAssets & assets); + +} // namespace engine::models::kokoro_tts diff --git a/include/engine/models/kokoro_tts/g2p_en.h b/include/engine/models/kokoro_tts/g2p_en.h new file mode 100644 index 000000000..c24777759 --- /dev/null +++ b/include/engine/models/kokoro_tts/g2p_en.h @@ -0,0 +1,159 @@ +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace kokoro_ggml::g2p_en { + +struct TokenContext { + std::optional future_vowel; + bool future_to = false; +}; + +struct LexiconResult { + std::string phonemes; + int rating = 0; +}; + +struct LexiconEntry { + bool has_default = false; + std::string default_phonemes; + std::unordered_map> by_tag; +}; + +struct TokenMeta { + bool is_head = true; + bool prespace = false; + std::optional stress; + std::optional alias; + std::optional currency; + std::string num_flags; + std::optional rating; +}; + +struct MToken { + std::string text; + std::string tag; + std::string whitespace; + std::optional phonemes; + TokenMeta meta; +}; + +using FeatureValue = std::variant; + +struct PreprocessResult { + std::string text; + std::unordered_map features; +}; + +class Lexicon { +public: + explicit Lexicon(bool british = false); + Lexicon(std::filesystem::path gold_tsv, std::filesystem::path silver_tsv, bool british = false); + + std::optional lookup( + const std::string & word, + const std::optional & tag, + const std::optional & stress, + const TokenContext & ctx) const; + + std::optional operator()(const MToken & token, const TokenContext & ctx) const; + static std::string apply_stress(const std::string & phonemes, const std::optional & stress); + static bool is_number_token(const std::string & word, bool is_head); + +private: + std::optional get_special_case( + const std::string & word, + const std::optional & tag, + const std::optional & stress, + const TokenContext & ctx) const; + + std::optional get_nnp(const std::string & word) const; + std::optional lookup_raw( + const std::string & word, + const std::optional & tag, + const std::optional & stress) const; + + std::optional stem_s( + const std::string & word, + const std::optional & tag, + const std::optional & stress, + const TokenContext & ctx) const; + std::optional stem_ed( + const std::string & word, + const std::optional & tag, + const std::optional & stress, + const TokenContext & ctx) const; + std::optional stem_ing( + const std::string & word, + const std::optional & tag, + const std::optional & stress, + const TokenContext & ctx) const; + + std::optional get_word( + const std::string & word, + const std::optional & tag, + const std::optional & stress, + const TokenContext & ctx) const; + + std::optional get_number( + const std::string & word, + const std::optional & currency, + bool is_head) const; + + bool is_known(const std::string & word) const; + + static std::optional parent_tag(const std::optional & tag); + static bool is_currency_shape(const std::string & word); + + std::optional s_suffix(const std::string & stem) const; + std::optional ed_suffix(const std::string & stem) const; + std::optional ing_suffix(const std::string & stem) const; + + void load_tsv( + const std::filesystem::path & path, + std::unordered_map & target); + + bool british_ = false; + std::unordered_map golds_; + std::unordered_map silvers_; +}; + +class EnglishG2P { +public: + explicit EnglishG2P(bool british = false); + EnglishG2P(std::filesystem::path lexicon_dir, bool british = false); + + static PreprocessResult preprocess(const std::string & text); + std::vector tokenize( + const std::string & text, + const std::unordered_map & features) const; + + std::pair> operator()(const std::string & text, bool enable_preprocess = true) const; + +private: + static std::vector fold_left(const std::vector & tokens); + static std::vector>> retokenize(const std::vector & tokens); + static void resolve_tokens(std::vector & tokens); + static TokenContext token_context(const TokenContext & ctx, const std::optional & phonemes, const MToken & token); + static std::string tokens_to_phonemes(const std::vector & tokens); + static MToken merge_token_pair( + const MToken & left, + const MToken & right, + const std::optional & unk); + static MToken merge_tokens( + const std::vector & tokens, + size_t begin, + size_t end, + const std::optional & unk); + + Lexicon lexicon_; +}; + +} // namespace kokoro_ggml::g2p_en diff --git a/include/engine/models/kokoro_tts/loader.h b/include/engine/models/kokoro_tts/loader.h new file mode 100644 index 000000000..a4d39328f --- /dev/null +++ b/include/engine/models/kokoro_tts/loader.h @@ -0,0 +1,33 @@ +#pragma once + +#include "engine/framework/runtime/model.h" +#include "engine/models/kokoro_tts/assets.h" + +#include +#include + +namespace engine::models::kokoro_tts { + +class KokoroTTSLoadedModel final : public runtime::ILoadedVoiceModel { +public: + KokoroTTSLoadedModel( + runtime::ModelMetadata metadata, + runtime::CapabilitySet capabilities, + std::shared_ptr assets); + + const runtime::ModelMetadata & metadata() const noexcept override; + const runtime::CapabilitySet & capabilities() const noexcept override; + std::unique_ptr create_task_session( + const runtime::TaskSpec & task, + const runtime::SessionOptions & options) const override; + +private: + runtime::ModelMetadata metadata_; + runtime::CapabilitySet capabilities_; + std::shared_ptr assets_; +}; + +std::unique_ptr load_kokoro_tts_model(const std::filesystem::path & model_path); +std::shared_ptr make_kokoro_tts_loader(); + +} // namespace engine::models::kokoro_tts diff --git a/include/engine/models/kokoro_tts/plbert.h b/include/engine/models/kokoro_tts/plbert.h new file mode 100644 index 000000000..30d24f111 --- /dev/null +++ b/include/engine/models/kokoro_tts/plbert.h @@ -0,0 +1,35 @@ +#pragma once + +#include +#include +#include + +typedef struct ggml_backend * ggml_backend_t; + +namespace kokoro_ggml { + +struct KokoroWeights; + +int64_t kokoro_plbert_output_dim(std::shared_ptr weights, bool project_hidden); + +class KokoroPlbertRuntime { +public: + KokoroPlbertRuntime( + std::shared_ptr weights, + ggml_backend_t backend, + int n_threads, + bool use_device_backend, + int64_t fixed_token_capacity = 0); + ~KokoroPlbertRuntime(); + + KokoroPlbertRuntime(const KokoroPlbertRuntime &) = delete; + KokoroPlbertRuntime & operator=(const KokoroPlbertRuntime &) = delete; + + std::vector encode(const std::vector & input_ids, bool project_hidden = true); + +private: + struct Impl; + std::unique_ptr impl_; +}; + +} // namespace kokoro_ggml diff --git a/include/engine/models/kokoro_tts/predictor.h b/include/engine/models/kokoro_tts/predictor.h new file mode 100644 index 000000000..fb6ac7fdc --- /dev/null +++ b/include/engine/models/kokoro_tts/predictor.h @@ -0,0 +1,57 @@ +#pragma once + +#include +#include +#include +#include + +struct ggml_tensor; +typedef struct ggml_backend * ggml_backend_t; + +namespace kokoro_ggml { + +struct KokoroWeights; + +struct KokoroPredictorGraphConfig { + size_t duration_graph_bytes = 384ull * 1024ull * 1024ull; + size_t text_graph_bytes = 256ull * 1024ull * 1024ull; + size_t tail_graph_bytes = 640ull * 1024ull * 1024ull; + int graph_node_capacity = 131072; +}; + +struct PredictorOutputs { + std::vector durations; + std::vector f0_curve; + std::vector decoder_x; + int64_t decoder_x_rows = 0; + int64_t decoder_x_cols = 0; + const ggml_tensor * decoder_x_tensor = nullptr; + bool decoder_x_on_backend = false; +}; + +class KokoroPredictorRuntime { +public: + KokoroPredictorRuntime( + std::shared_ptr weights, + ggml_backend_t backend, + int n_threads, + bool use_device_backend, + int64_t plbert_fixed_token_capacity = 0, + int64_t pre_tail_token_capacity = 0, + KokoroPredictorGraphConfig graph_config = {}); + ~KokoroPredictorRuntime(); + + KokoroPredictorRuntime(const KokoroPredictorRuntime &) = delete; + KokoroPredictorRuntime & operator=(const KokoroPredictorRuntime &) = delete; + + PredictorOutputs predict( + const std::vector & input_ids, + const std::vector & ref_s, + float speed); + +private: + struct Impl; + std::unique_ptr impl_; +}; + +} // namespace kokoro_ggml diff --git a/include/engine/models/kokoro_tts/session.h b/include/engine/models/kokoro_tts/session.h new file mode 100644 index 000000000..1ae330f4a --- /dev/null +++ b/include/engine/models/kokoro_tts/session.h @@ -0,0 +1,73 @@ +#pragma once + +#include "engine/framework/runtime/session_base.h" +#include "engine/models/kokoro_tts/assets.h" + +#include +#include + +namespace kokoro_ggml { +class KokoroDecoderRuntime; +} + +namespace engine::models::kokoro_tts { + +struct KokoroSynthesisInput; +struct KokoroFrontendSessionState; + +class KokoroTTSSession final + : public runtime::RuntimeSessionBase + , public runtime::IOfflineVoiceTaskSession { +public: + KokoroTTSSession( + runtime::TaskSpec task, + runtime::SessionOptions options, + std::shared_ptr assets); + ~KokoroTTSSession() override; + + std::string family() const override; + runtime::VoiceTaskKind task_kind() const override; + runtime::RunMode run_mode() const override; + void prepare(const runtime::SessionPreparationRequest & request) override; + runtime::TaskResult run(const runtime::TaskRequest & request) override; + +private: + struct DecoderCapacityContract { + int64_t decoder_frame_capacity = 0; + int64_t conditioning_sample_capacity = 0; + int64_t conditioning_frame_capacity = 0; + }; + + struct PreparedRuntime; + + runtime::MappedGraphCapacityAdapter make_graph_capacity_adapter(); + int64_t base_graph_capacity_tokens() const; + std::vector prepared_graph_capacities() const; + DecoderCapacityContract make_decoder_capacity_contract(int64_t decoder_frame_capacity) const; + void prepare_graph_capacity(int64_t capacity); + void prepare_decoder_graph_capacity(int64_t capacity); + + runtime::TaskSpec task_; + std::shared_ptr assets_; + std::shared_ptr weights_; + runtime::GraphCapacityController graph_capacity_controller_; + int64_t fixed_token_capacity_ = 0; + int64_t pre_tail_token_capacity_ = 0; + uint64_t rng_seed_ = 0; + engine::assets::TensorStorageType matmul_weight_storage_type_ = engine::assets::TensorStorageType::Native; + engine::assets::TensorStorageType conv_weight_storage_type_ = engine::assets::TensorStorageType::Native; + size_t weight_context_bytes_ = 512ull * 1024ull * 1024ull; + size_t predictor_duration_graph_bytes_ = 384ull * 1024ull * 1024ull; + size_t predictor_text_graph_bytes_ = 256ull * 1024ull * 1024ull; + size_t predictor_tail_graph_bytes_ = 640ull * 1024ull * 1024ull; + std::string cached_request_key_; + std::unique_ptr cached_input_; + std::unique_ptr frontend_session_state_; + int64_t prepared_decoder_capacity_ = 0; + std::unique_ptr prepared_decoder_; + DecoderCapacityContract prepared_decoder_context_ = {}; + int64_t prepared_session_capacity_ = 0; + std::unique_ptr prepared_session_; +}; + +} // namespace engine::models::kokoro_tts diff --git a/src/framework/runtime/registry.cpp b/src/framework/runtime/registry.cpp index c2b1414aa..a6425e81c 100644 --- a/src/framework/runtime/registry.cpp +++ b/src/framework/runtime/registry.cpp @@ -3,9 +3,9 @@ #include "engine/framework/debug/trace.h" #include "engine/framework/io/config.h" #include "engine/framework/io/filesystem.h" +#include "engine/models/kokoro_tts/loader.h" // Development registry entries from Share/AudioCPP that are not present in this release tree yet: // #include "engine/models/higgs_tts/loader.h" -// #include "engine/models/kokoro_tts/loader.h" // #include "engine/models/moss_tts/loader.h" // #include "engine/models/parakeet_tdt/loader.h" #include "engine/models/ace_step/loader.h" @@ -205,8 +205,8 @@ ModelRegistry make_registry_from_config( ModelRegistry make_default_registry(const std::optional & config_path) { const std::vector> available_loaders = { + engine::models::kokoro_tts::make_kokoro_tts_loader(), // Development registry entries from Share/AudioCPP that are not present in this release tree yet: - // engine::models::kokoro_tts::make_kokoro_tts_loader(), // engine::models::moss_tts::make_moss_tts_loader(), // engine::models::higgs_tts::make_higgs_tts_loader(), // engine::models::parakeet_tdt::make_parakeet_tdt_loader(), diff --git a/src/models/kokoro_tts/assets.cpp b/src/models/kokoro_tts/assets.cpp new file mode 100644 index 000000000..bbbd54090 --- /dev/null +++ b/src/models/kokoro_tts/assets.cpp @@ -0,0 +1,835 @@ +#include "engine/framework/assets/resource_bundle.h" +#include "engine/framework/core/backend_weight_store.h" +#include "engine/framework/io/binary.h" +#include "engine/framework/io/json.h" +#include "engine/models/kokoro_tts/assets.h" + +#include "engine/models/kokoro_tts/g2p_en.h" + +#include +#include +#include +#include +#include +#include + +namespace kokoro_ggml { + +namespace { + +struct KokoroConfigMetadata { + int64_t n_token = 178; + int64_t hidden_dim = 512; + int64_t style_dim = 128; + int64_t n_layer = 3; + int64_t max_dur = 50; + int64_t n_mels = 80; + int64_t text_encoder_kernel_size = 5; + int64_t plbert_hidden_size = 768; + int64_t plbert_num_attention_heads = 12; + int64_t plbert_intermediate_size = 2048; + int64_t plbert_max_position_embeddings = 512; + int64_t plbert_num_hidden_layers = 12; + int64_t gen_istft_n_fft = 20; + int64_t gen_istft_hop_size = 5; + std::vector upsample_rates = {10, 6}; + std::vector upsample_kernel_sizes = {20, 12}; + std::vector resblock_kernel_sizes = {3, 7, 11}; + std::vector> resblock_dilation_sizes = {{1, 3, 5}, {1, 3, 5}, {1, 3, 5}}; + int64_t upsample_initial_channel = 512; +}; + +} // namespace + +std::vector apply_weight_norm( + const std::vector & g, + const std::vector & v, + int64_t leading, + int64_t inner_size) { + std::vector weight(v.size(), 0.0f); + for (int64_t i = 0; i < leading; ++i) { + double norm = 0.0; + for (int64_t j = 0; j < inner_size; ++j) { + const float value = v[static_cast(i * inner_size + j)]; + norm += static_cast(value) * static_cast(value); + } + const float scale = g[static_cast(i)] / std::sqrt(static_cast(norm) + 1.0e-12f); + for (int64_t j = 0; j < inner_size; ++j) { + weight[static_cast(i * inner_size + j)] = v[static_cast(i * inner_size + j)] * scale; + } + } + return weight; +} + +std::vector combine_lstm_biases(const std::vector & lhs, const std::vector & rhs) { + if (lhs.size() != rhs.size()) { + throw std::runtime_error("Kokoro LSTM bias size mismatch"); + } + std::vector out(lhs.size(), 0.0f); + for (size_t i = 0; i < lhs.size(); ++i) { + out[i] = lhs[i] + rhs[i]; + } + return out; +} + +engine::assets::TensorStorageType vector_storage_type(engine::assets::TensorStorageType storage_type) { + if (storage_type == engine::assets::TensorStorageType::Q4_0 || + storage_type == engine::assets::TensorStorageType::Q4_1 || + storage_type == engine::assets::TensorStorageType::Q5_0 || + storage_type == engine::assets::TensorStorageType::Q5_1 || + storage_type == engine::assets::TensorStorageType::Q8_0) { + return engine::assets::TensorStorageType::F32; + } + return storage_type; +} + +KokoroWeights::LinearWeights load_linear( + engine::core::BackendWeightStore & store, + const engine::assets::TensorSource & source, + const std::string & prefix, + int64_t out_features, + int64_t in_features, + bool use_bias, + engine::assets::TensorStorageType storage_type) { + KokoroWeights::LinearWeights linear; + linear.out_features = out_features; + linear.in_features = in_features; + linear.use_bias = use_bias; + linear.weight = store.load_tensor(source, prefix + ".weight", storage_type, {out_features, in_features}); + if (use_bias) { + linear.bias = store.load_tensor(source, prefix + ".bias", vector_storage_type(storage_type), {out_features}); + } + return linear; +} + +KokoroWeights::HostAffineWeights load_host_affine( + const engine::assets::TensorSource & source, + const std::string & prefix, + int64_t out_features, + int64_t in_features, + bool use_bias) { + KokoroWeights::HostAffineWeights affine; + affine.out_features = out_features; + affine.in_features = in_features; + affine.use_bias = use_bias; + affine.weight = source.require_f32(prefix + ".weight", {out_features, in_features}); + if (use_bias) { + affine.bias = source.require_f32(prefix + ".bias", {out_features}); + } + return affine; +} + +KokoroWeights::EmbeddingWeights load_embedding( + engine::core::BackendWeightStore & store, + const engine::assets::TensorSource & source, + const std::string & name, + int64_t num_embeddings, + int64_t embedding_dim, + engine::assets::TensorStorageType storage_type) { + KokoroWeights::EmbeddingWeights embedding; + embedding.num_embeddings = num_embeddings; + embedding.embedding_dim = embedding_dim; + embedding.weight = store.load_tensor(source, name, storage_type, {num_embeddings, embedding_dim}); + return embedding; +} + +KokoroWeights::LayerNormWeights load_layer_norm( + engine::core::BackendWeightStore & store, + const engine::assets::TensorSource & source, + const std::string & prefix, + int64_t channels, + float eps) { + KokoroWeights::LayerNormWeights norm; + norm.channels = channels; + norm.eps = eps; + norm.weight = store.load_tensor(source, prefix + ".weight", engine::assets::TensorStorageType::Native, {channels}); + norm.bias = store.load_tensor(source, prefix + ".bias", engine::assets::TensorStorageType::Native, {channels}); + return norm; +} + +KokoroWeights::WeightNormConv1dWeights load_weight_norm_conv1d( + engine::core::BackendWeightStore & store, + const engine::assets::TensorSource & source, + const std::string & prefix, + int64_t out_channels, + int64_t in_channels, + int64_t kernel, + int64_t stride, + int64_t padding, + int64_t dilation, + int64_t groups, + bool use_bias, + engine::assets::TensorStorageType storage_type) { + const int64_t grouped_in_channels = in_channels / groups; + auto g = source.require_f32( prefix + ".weight_g", {out_channels, 1, 1}); + auto v = source.require_f32( prefix + ".weight_v", {out_channels, grouped_in_channels, kernel}); + KokoroWeights::WeightNormConv1dWeights conv; + conv.out_channels = out_channels; + conv.in_channels = in_channels; + conv.kernel = kernel; + conv.stride = stride; + conv.padding = padding; + conv.dilation = dilation; + conv.groups = groups; + conv.use_bias = use_bias; + conv.weight = store.make_from_f32( + engine::core::TensorShape::from_dims({out_channels, grouped_in_channels, kernel}), + storage_type, + apply_weight_norm(g, v, out_channels, grouped_in_channels * kernel)); + if (use_bias) { + conv.bias = store.load_tensor(source, prefix + ".bias", vector_storage_type(storage_type), {out_channels}); + } + return conv; +} + +KokoroWeights::Conv1dWeights load_conv1d( + engine::core::BackendWeightStore & store, + const engine::assets::TensorSource & source, + const std::string & prefix, + int64_t out_channels, + int64_t in_channels, + int64_t kernel, + int64_t stride, + int64_t padding, + int64_t dilation, + int64_t groups, + bool use_bias, + engine::assets::TensorStorageType storage_type) { + KokoroWeights::Conv1dWeights conv; + conv.out_channels = out_channels; + conv.in_channels = in_channels; + conv.kernel = kernel; + conv.stride = stride; + conv.padding = padding; + conv.dilation = dilation; + conv.groups = groups; + conv.use_bias = use_bias; + conv.weight = store.load_tensor(source, prefix + ".weight", storage_type, {out_channels, in_channels / groups, kernel}); + if (use_bias) { + conv.bias = store.load_tensor(source, prefix + ".bias", vector_storage_type(storage_type), {out_channels}); + } + return conv; +} + +KokoroWeights::WeightNormConvTranspose1dWeights load_weight_norm_conv_transpose1d( + engine::core::BackendWeightStore & store, + const engine::assets::TensorSource & source, + const std::string & prefix, + int64_t in_channels, + int64_t out_channels, + int64_t kernel, + int64_t stride, + int64_t padding, + int64_t output_padding, + int64_t groups, + bool use_bias, + engine::assets::TensorStorageType storage_type) { + const int64_t grouped_out_channels = out_channels / groups; + auto g = source.require_f32( prefix + ".weight_g", {in_channels, 1, 1}); + auto v = source.require_f32( prefix + ".weight_v", {in_channels, grouped_out_channels, kernel}); + KokoroWeights::WeightNormConvTranspose1dWeights conv; + conv.in_channels = in_channels; + conv.out_channels = out_channels; + conv.kernel = kernel; + conv.stride = stride; + conv.padding = padding; + conv.output_padding = output_padding; + conv.groups = groups; + conv.use_bias = use_bias; + auto normalized = apply_weight_norm(g, v, in_channels, grouped_out_channels * kernel); + conv.weight = store.make_from_f32( + engine::core::TensorShape::from_dims({in_channels, grouped_out_channels, kernel}), + storage_type, + normalized); + std::vector bias_values; + if (use_bias) { + bias_values = source.require_f32(prefix + ".bias", {groups == in_channels ? in_channels : out_channels}); + conv.bias = store.load_tensor( + source, + prefix + ".bias", + vector_storage_type(storage_type), + {groups == in_channels ? in_channels : out_channels}); + } + if (groups == in_channels && in_channels == out_channels && grouped_out_channels == 1) { + std::vector dense(static_cast(in_channels * out_channels * kernel), 0.0f); + for (int64_t channel = 0; channel < in_channels; ++channel) { + const float * src = normalized.data() + static_cast(channel * kernel); + float * dst = dense.data() + static_cast((channel * out_channels + channel) * kernel); + std::memcpy(dst, src, static_cast(kernel) * sizeof(float)); + } + conv.dense_weight = store.make_from_f32( + engine::core::TensorShape::from_dims({in_channels, out_channels, kernel}), + storage_type, + std::move(dense)); + } + if (groups == 1 && output_padding == 0 && kernel == stride * 2 && padding > 0 && padding < stride) { + auto phase_conv = std::make_shared(); + phase_conv->in_channels = in_channels; + phase_conv->out_channels = out_channels * stride; + phase_conv->kernel = 3; + phase_conv->stride = 1; + phase_conv->padding = 1; + phase_conv->dilation = 1; + phase_conv->groups = 1; + phase_conv->use_bias = use_bias; + std::vector phase_weight( + static_cast(phase_conv->out_channels * phase_conv->in_channels * phase_conv->kernel), + 0.0f); + std::vector phase_bias; + if (use_bias) { + phase_bias.assign(static_cast(phase_conv->out_channels), 0.0f); + } + auto source_weight = [&](int64_t in_channel, int64_t out_channel, int64_t kernel_index) -> float { + const size_t offset = static_cast((in_channel * out_channels + out_channel) * kernel + kernel_index); + return normalized[offset]; + }; + auto target_weight = [&](int64_t phase_channel, int64_t in_channel, int64_t kernel_index) -> float & { + const size_t offset = static_cast( + (phase_channel * phase_conv->in_channels + in_channel) * phase_conv->kernel + kernel_index); + return phase_weight[offset]; + }; + const int64_t split_phase = stride - padding; + for (int64_t out_channel = 0; out_channel < out_channels; ++out_channel) { + for (int64_t phase = 0; phase < stride; ++phase) { + const int64_t phase_channel = out_channel * stride + phase; + if (use_bias) { + phase_bias[static_cast(phase_channel)] = bias_values[static_cast(out_channel)]; + } + for (int64_t in_channel = 0; in_channel < in_channels; ++in_channel) { + const int64_t center_kernel = padding + phase; + target_weight(phase_channel, in_channel, 1) = + source_weight(in_channel, out_channel, center_kernel); + if (phase < split_phase) { + const int64_t previous_kernel = stride + padding + phase; + target_weight(phase_channel, in_channel, 0) = + source_weight(in_channel, out_channel, previous_kernel); + } else { + const int64_t next_kernel = phase + padding - stride; + target_weight(phase_channel, in_channel, 2) = + source_weight(in_channel, out_channel, next_kernel); + } + } + } + } + phase_conv->weight = store.make_from_f32( + engine::core::TensorShape::from_dims({phase_conv->out_channels, phase_conv->in_channels, phase_conv->kernel}), + storage_type, + std::move(phase_weight)); + if (use_bias) { + phase_conv->bias = store.make_from_f32( + engine::core::TensorShape::from_dims({phase_conv->out_channels}), + vector_storage_type(storage_type), + std::move(phase_bias)); + } + conv.phase_shuffle_conv = std::move(phase_conv); + } + return conv; +} + +KokoroWeights::LstmWeights load_lstm( + engine::core::BackendWeightStore & store, + const engine::assets::TensorSource & source, + const std::string & prefix, + int64_t input_size, + int64_t hidden_size, + engine::assets::TensorStorageType storage_type) { + KokoroWeights::LstmWeights lstm; + lstm.input_size = input_size; + lstm.hidden_size = hidden_size; + lstm.weight_ih_l0 = store.load_tensor(source, prefix + ".weight_ih_l0", storage_type, {4 * hidden_size, input_size}); + lstm.weight_hh_l0 = store.load_tensor(source, prefix + ".weight_hh_l0", storage_type, {4 * hidden_size, hidden_size}); + lstm.bias_ih_l0 = store.load_tensor(source, prefix + ".bias_ih_l0", vector_storage_type(storage_type), {4 * hidden_size}); + lstm.bias_hh_l0 = store.load_tensor(source, prefix + ".bias_hh_l0", vector_storage_type(storage_type), {4 * hidden_size}); + auto bias_ih = source.require_f32(prefix + ".bias_ih_l0", {4 * hidden_size}); + auto bias_hh = source.require_f32(prefix + ".bias_hh_l0", {4 * hidden_size}); + lstm.combined_bias_l0 = store.make_from_f32( + engine::core::TensorShape::from_dims({4 * hidden_size}), + vector_storage_type(storage_type), + combine_lstm_biases(bias_ih, bias_hh)); + lstm.weight_ih_l0_reverse = + store.load_tensor(source, prefix + ".weight_ih_l0_reverse", storage_type, {4 * hidden_size, input_size}); + lstm.weight_hh_l0_reverse = + store.load_tensor(source, prefix + ".weight_hh_l0_reverse", storage_type, {4 * hidden_size, hidden_size}); + lstm.bias_ih_l0_reverse = + store.load_tensor(source, prefix + ".bias_ih_l0_reverse", vector_storage_type(storage_type), {4 * hidden_size}); + lstm.bias_hh_l0_reverse = + store.load_tensor(source, prefix + ".bias_hh_l0_reverse", vector_storage_type(storage_type), {4 * hidden_size}); + auto bias_ih_reverse = source.require_f32(prefix + ".bias_ih_l0_reverse", {4 * hidden_size}); + auto bias_hh_reverse = source.require_f32(prefix + ".bias_hh_l0_reverse", {4 * hidden_size}); + lstm.combined_bias_l0_reverse = store.make_from_f32( + engine::core::TensorShape::from_dims({4 * hidden_size}), + vector_storage_type(storage_type), + combine_lstm_biases(bias_ih_reverse, bias_hh_reverse)); + return lstm; +} + +namespace { + +KokoroWeights::LayerNormWeights load_beta_gamma_layer_norm( + engine::core::BackendWeightStore & store, + const engine::assets::TensorSource & source, + const std::string & prefix, + int64_t channels, + float eps) { + KokoroWeights::LayerNormWeights norm; + norm.channels = channels; + norm.eps = eps; + norm.weight = store.load_tensor(source, prefix + ".gamma", engine::assets::TensorStorageType::Native, {channels}); + norm.bias = store.load_tensor(source, prefix + ".beta", engine::assets::TensorStorageType::Native, {channels}); + return norm; +} + +KokoroWeights::AdaLayerNormWeights load_ada_layer_norm( + engine::core::BackendWeightStore & store, + const engine::assets::TensorSource & source, + const std::string & prefix, + int64_t channels, + engine::assets::TensorStorageType storage_type) { + KokoroWeights::AdaLayerNormWeights norm; + norm.channels = channels; + norm.fc = load_linear(store, source, prefix + ".fc", channels * 2, 128, true, storage_type); + return norm; +} + +KokoroWeights::AdaIn1dWeights load_adain1d( + engine::core::BackendWeightStore & store, + const engine::assets::TensorSource & source, + const std::string & prefix, + int64_t channels, + engine::assets::TensorStorageType storage_type) { + KokoroWeights::AdaIn1dWeights norm; + norm.channels = channels; + norm.fc = load_linear(store, source, prefix + ".fc", channels * 2, 128, true, storage_type); + return norm; +} + +KokoroWeights::AlbertLayerWeights load_albert_layer( + engine::core::BackendWeightStore & store, + const engine::assets::TensorSource & source, + const std::string & prefix, + int64_t hidden_size, + int64_t intermediate_size, + engine::assets::TensorStorageType storage_type) { + KokoroWeights::AlbertLayerWeights layer; + layer.attention.query = load_linear(store, source, prefix + ".attention.query", hidden_size, hidden_size, true, storage_type); + layer.attention.key = load_linear(store, source, prefix + ".attention.key", hidden_size, hidden_size, true, storage_type); + layer.attention.value = load_linear(store, source, prefix + ".attention.value", hidden_size, hidden_size, true, storage_type); + layer.attention.dense = load_linear(store, source, prefix + ".attention.dense", hidden_size, hidden_size, true, storage_type); + layer.attention.layer_norm = load_layer_norm(store, source, prefix + ".attention.LayerNorm", hidden_size, 1.0e-12f); + layer.ffn = load_linear(store, source, prefix + ".ffn", intermediate_size, hidden_size, true, storage_type); + layer.ffn_output = load_linear(store, source, prefix + ".ffn_output", hidden_size, intermediate_size, true, storage_type); + layer.full_layer_layer_norm = load_layer_norm(store, source, prefix + ".full_layer_layer_norm", hidden_size, 1.0e-12f); + return layer; +} + +KokoroWeights::TextEncoderBlockWeights load_text_encoder_block( + engine::core::BackendWeightStore & store, + const engine::assets::TensorSource & source, + int64_t block_index, + engine::assets::TensorStorageType conv_storage_type) { + const std::string prefix = "text_encoder.cnn." + std::to_string(block_index); + KokoroWeights::TextEncoderBlockWeights block; + block.conv = load_weight_norm_conv1d(store, source, prefix + ".0", 512, 512, 5, 1, 2, 1, 1, true, conv_storage_type); + block.layer_norm = load_beta_gamma_layer_norm(store, source, prefix + ".1", 512, 1.0e-5f); + return block; +} + +KokoroWeights::AdainResBlock1dWeights load_adain_resblock( + engine::core::BackendWeightStore & store, + const engine::assets::TensorSource & source, + const std::string & prefix, + int64_t dim_in, + int64_t dim_out, + bool upsample, + engine::assets::TensorStorageType matmul_storage_type, + engine::assets::TensorStorageType conv_storage_type) { + KokoroWeights::AdainResBlock1dWeights block; + block.learned_sc = dim_in != dim_out; + block.use_pool = upsample; + block.upsample = upsample; + block.conv1 = load_weight_norm_conv1d(store, source, prefix + ".conv1", dim_out, dim_in, 3, 1, 1, 1, 1, true, conv_storage_type); + block.conv2 = load_weight_norm_conv1d(store, source, prefix + ".conv2", dim_out, dim_out, 3, 1, 1, 1, 1, true, conv_storage_type); + block.norm1 = load_adain1d(store, source, prefix + ".norm1", dim_in, matmul_storage_type); + block.norm2 = load_adain1d(store, source, prefix + ".norm2", dim_out, matmul_storage_type); + if (block.learned_sc) { + block.conv1x1 = + load_weight_norm_conv1d(store, source, prefix + ".conv1x1", dim_out, dim_in, 1, 1, 0, 1, 1, false, conv_storage_type); + } + if (block.use_pool) { + block.pool = + load_weight_norm_conv_transpose1d(store, source, prefix + ".pool", dim_in, dim_in, 3, 2, 1, 1, dim_in, true, conv_storage_type); + } + return block; +} + +KokoroWeights::GeneratorResBlockWeights load_generator_resblock( + engine::core::BackendWeightStore & store, + const engine::assets::TensorSource & source, + const std::string & prefix, + int64_t channels, + int64_t kernel, + const std::vector & dilations, + engine::assets::TensorStorageType matmul_storage_type, + engine::assets::TensorStorageType conv_storage_type) { + KokoroWeights::GeneratorResBlockWeights block; + for (size_t i = 0; i < dilations.size(); ++i) { + block.convs1.push_back(load_weight_norm_conv1d( + store, + source, + prefix + ".convs1." + std::to_string(i), + channels, + channels, + kernel, + 1, + static_cast((kernel * dilations[i] - dilations[i]) / 2), + dilations[i], + 1, + true, + conv_storage_type)); + block.convs2.push_back(load_weight_norm_conv1d( + store, + source, + prefix + ".convs2." + std::to_string(i), + channels, + channels, + kernel, + 1, + (kernel - 1) / 2, + 1, + 1, + true, + conv_storage_type)); + block.adain1.push_back(load_adain1d(store, source, prefix + ".adain1." + std::to_string(i), channels, matmul_storage_type)); + block.adain2.push_back(load_adain1d(store, source, prefix + ".adain2." + std::to_string(i), channels, matmul_storage_type)); + block.alpha1.push_back( + store.load_tensor(source, prefix + ".alpha1." + std::to_string(i), engine::assets::TensorStorageType::Native, {1, channels, 1})); + block.alpha2.push_back( + store.load_tensor(source, prefix + ".alpha2." + std::to_string(i), engine::assets::TensorStorageType::Native, {1, channels, 1})); + } + return block; +} + +} // namespace + +KokoroConfigMetadata parse_kokoro_config_metadata(const engine::io::json::Value & root) { + if (!root.is_object()) { + throw std::runtime_error("kokoro config root must be an object"); + } + KokoroConfigMetadata config; + const auto * n_token = root.find("n_token"); + const auto * hidden_dim = root.find("hidden_dim"); + const auto * style_dim = root.find("style_dim"); + const auto * n_layer = root.find("n_layer"); + const auto * max_dur = root.find("max_dur"); + const auto * n_mels = root.find("n_mels"); + const auto * text_encoder_kernel_size = root.find("text_encoder_kernel_size"); + if (n_token) config.n_token = n_token->as_i64(); + if (hidden_dim) config.hidden_dim = hidden_dim->as_i64(); + if (style_dim) config.style_dim = style_dim->as_i64(); + if (n_layer) config.n_layer = n_layer->as_i64(); + if (max_dur) config.max_dur = max_dur->as_i64(); + if (n_mels) config.n_mels = n_mels->as_i64(); + if (text_encoder_kernel_size) config.text_encoder_kernel_size = text_encoder_kernel_size->as_i64(); + + if (const auto * plbert = root.find("plbert")) { + const auto * hidden_size = plbert->find("hidden_size"); + const auto * num_attention_heads = plbert->find("num_attention_heads"); + const auto * intermediate_size = plbert->find("intermediate_size"); + const auto * max_position_embeddings = plbert->find("max_position_embeddings"); + const auto * num_hidden_layers = plbert->find("num_hidden_layers"); + if (hidden_size) config.plbert_hidden_size = hidden_size->as_i64(); + if (num_attention_heads) config.plbert_num_attention_heads = num_attention_heads->as_i64(); + if (intermediate_size) config.plbert_intermediate_size = intermediate_size->as_i64(); + if (max_position_embeddings) config.plbert_max_position_embeddings = max_position_embeddings->as_i64(); + if (num_hidden_layers) config.plbert_num_hidden_layers = num_hidden_layers->as_i64(); + } + + if (const auto * istftnet = root.find("istftnet")) { + const auto * gen_istft_n_fft = istftnet->find("gen_istft_n_fft"); + const auto * gen_istft_hop_size = istftnet->find("gen_istft_hop_size"); + const auto * upsample_rates = istftnet->find("upsample_rates"); + const auto * upsample_kernel_sizes = istftnet->find("upsample_kernel_sizes"); + const auto * resblock_kernel_sizes = istftnet->find("resblock_kernel_sizes"); + const auto * upsample_initial_channel = istftnet->find("upsample_initial_channel"); + if (gen_istft_n_fft) config.gen_istft_n_fft = gen_istft_n_fft->as_i64(); + if (gen_istft_hop_size) config.gen_istft_hop_size = gen_istft_hop_size->as_i64(); + if (upsample_rates) { + config.upsample_rates.clear(); + for (const auto & value : upsample_rates->as_array()) { + config.upsample_rates.push_back(value.as_i64()); + } + } + if (upsample_kernel_sizes) { + config.upsample_kernel_sizes.clear(); + for (const auto & value : upsample_kernel_sizes->as_array()) { + config.upsample_kernel_sizes.push_back(value.as_i64()); + } + } + if (resblock_kernel_sizes) { + config.resblock_kernel_sizes.clear(); + for (const auto & value : resblock_kernel_sizes->as_array()) { + config.resblock_kernel_sizes.push_back(value.as_i64()); + } + } + if (upsample_initial_channel) config.upsample_initial_channel = upsample_initial_channel->as_i64(); + } + + return config; +} + +std::shared_ptr load_kokoro_weights( + const engine::io::json::Value & config_root, + const engine::assets::TensorSource & source, + ggml_backend_t backend, + engine::core::BackendType backend_type, + engine::assets::TensorStorageType matmul_storage_type, + engine::assets::TensorStorageType conv_storage_type, + size_t weight_context_bytes) { + const auto config = parse_kokoro_config_metadata(config_root); + auto weights = std::make_shared(); + weights->store = std::make_shared( + backend, + backend_type, + "kokoro.weights", + weight_context_bytes); + weights->n_token = config.n_token; + weights->hidden_dim = config.hidden_dim; + weights->style_dim = config.style_dim; + weights->n_layer = config.n_layer; + weights->max_dur = config.max_dur; + weights->n_mels = config.n_mels; + weights->context_length = config.plbert_max_position_embeddings; + weights->max_output_tokens = weights->context_length * weights->max_dur; + weights->bert.hidden_size = config.plbert_hidden_size; + weights->bert.embedding_size = 128; + weights->bert.intermediate_size = config.plbert_intermediate_size; + weights->bert.max_position_embeddings = config.plbert_max_position_embeddings; + weights->bert.num_hidden_layers = config.plbert_num_hidden_layers; + weights->bert.num_attention_heads = config.plbert_num_attention_heads; + weights->decoder.generator.gen_istft_n_fft = config.gen_istft_n_fft; + weights->decoder.generator.gen_istft_hop_size = config.gen_istft_hop_size; + + auto & store = *weights->store; + + weights->bert.embeddings.word_embeddings = load_embedding(store, source, "bert.embeddings.word_embeddings.weight", weights->n_token, 128, matmul_storage_type); + weights->bert.embeddings.position_embeddings = load_embedding(store, source, "bert.embeddings.position_embeddings.weight", weights->context_length, 128, matmul_storage_type); + weights->bert.embeddings.token_type_embeddings = load_embedding(store, source, "bert.embeddings.token_type_embeddings.weight", 2, 128, matmul_storage_type); + weights->bert.embeddings.layer_norm = load_layer_norm(store, source, "bert.embeddings.LayerNorm", 128, 1.0e-12f); + weights->bert.embedding_hidden_mapping_in = load_linear(store, source, "bert.encoder.embedding_hidden_mapping_in", weights->bert.hidden_size, 128, true, matmul_storage_type); + weights->bert.shared_layer = load_albert_layer( + store, + source, + "bert.encoder.albert_layer_groups.0.albert_layers.0", + weights->bert.hidden_size, + weights->bert.intermediate_size, + matmul_storage_type); + weights->bert.pooler = load_linear(store, source, "bert.pooler", weights->bert.hidden_size, weights->bert.hidden_size, true, matmul_storage_type); + + weights->bert_encoder = load_linear(store, source, "bert_encoder", weights->hidden_dim, weights->bert.hidden_size, true, matmul_storage_type); + + for (int64_t i = 0; i < weights->n_layer; ++i) { + weights->text_encoder.cnn.push_back(load_text_encoder_block(store, source, i, conv_storage_type)); + } + weights->text_encoder.embedding = load_embedding(store, source, "text_encoder.embedding.weight", weights->n_token, weights->hidden_dim, matmul_storage_type); + weights->text_encoder.lstm = load_lstm(store, source, "text_encoder.lstm", weights->hidden_dim, weights->hidden_dim / 2, matmul_storage_type); + + const std::string predictor_duration_encoder_prefix = "predictor.text_encoder"; + for (int64_t i = 0; i < weights->n_layer; ++i) { + weights->predictor.duration_encoder.lstms.push_back(load_lstm( + store, + source, + predictor_duration_encoder_prefix + ".lstms." + std::to_string(i * 2), + weights->hidden_dim + weights->style_dim, + weights->hidden_dim / 2, + matmul_storage_type)); + weights->predictor.duration_encoder.ada_layer_norms.push_back(load_ada_layer_norm(store, source, predictor_duration_encoder_prefix + ".lstms." + std::to_string(i * 2 + 1), weights->hidden_dim, matmul_storage_type)); + } + weights->predictor.lstm = + load_lstm(store, source, "predictor.lstm", weights->hidden_dim + weights->style_dim, weights->hidden_dim / 2, matmul_storage_type); + weights->predictor.duration_proj = load_linear(store, source, "predictor.duration_proj.linear_layer", weights->max_dur, weights->hidden_dim, true, matmul_storage_type); + weights->predictor.shared = + load_lstm(store, source, "predictor.shared", weights->hidden_dim + weights->style_dim, weights->hidden_dim / 2, matmul_storage_type); + weights->predictor.f0_blocks.push_back( + load_adain_resblock(store, source, "predictor.F0.0", weights->hidden_dim, weights->hidden_dim, false, matmul_storage_type, conv_storage_type)); + weights->predictor.f0_blocks.push_back(load_adain_resblock( + store, source, "predictor.F0.1", weights->hidden_dim, weights->hidden_dim / 2, true, matmul_storage_type, conv_storage_type)); + weights->predictor.f0_blocks.push_back(load_adain_resblock( + store, source, "predictor.F0.2", weights->hidden_dim / 2, weights->hidden_dim / 2, false, matmul_storage_type, conv_storage_type)); + weights->predictor.n_blocks.push_back( + load_adain_resblock(store, source, "predictor.N.0", weights->hidden_dim, weights->hidden_dim, false, matmul_storage_type, conv_storage_type)); + weights->predictor.n_blocks.push_back(load_adain_resblock( + store, source, "predictor.N.1", weights->hidden_dim, weights->hidden_dim / 2, true, matmul_storage_type, conv_storage_type)); + weights->predictor.n_blocks.push_back(load_adain_resblock( + store, source, "predictor.N.2", weights->hidden_dim / 2, weights->hidden_dim / 2, false, matmul_storage_type, conv_storage_type)); + weights->predictor.f0_proj = load_conv1d(store, source, "predictor.F0_proj", 1, weights->hidden_dim / 2, 1, 1, 0, 1, 1, true, conv_storage_type); + weights->predictor.n_proj = load_conv1d(store, source, "predictor.N_proj", 1, weights->hidden_dim / 2, 1, 1, 0, 1, 1, true, conv_storage_type); + + weights->decoder.encode = load_adain_resblock(store, source, "decoder.encode", 514, 1024, false, matmul_storage_type, conv_storage_type); + weights->decoder.decode.push_back( + load_adain_resblock(store, source, "decoder.decode.0", 1090, 1024, false, matmul_storage_type, conv_storage_type)); + weights->decoder.decode.push_back( + load_adain_resblock(store, source, "decoder.decode.1", 1090, 1024, false, matmul_storage_type, conv_storage_type)); + weights->decoder.decode.push_back( + load_adain_resblock(store, source, "decoder.decode.2", 1090, 1024, false, matmul_storage_type, conv_storage_type)); + weights->decoder.decode.push_back( + load_adain_resblock(store, source, "decoder.decode.3", 1090, 512, true, matmul_storage_type, conv_storage_type)); + weights->decoder.f0_conv = load_weight_norm_conv1d(store, source, "decoder.F0_conv", 1, 1, 3, 2, 1, 1, 1, true, conv_storage_type); + weights->decoder.n_conv = load_weight_norm_conv1d(store, source, "decoder.N_conv", 1, 1, 3, 2, 1, 1, 1, true, conv_storage_type); + weights->decoder.asr_res = load_weight_norm_conv1d(store, source, "decoder.asr_res.0", 64, 512, 1, 1, 0, 1, 1, true, conv_storage_type); + + weights->decoder.generator.ups.push_back( + load_weight_norm_conv_transpose1d(store, source, "decoder.generator.ups.0", 512, 256, 20, 10, 5, 0, 1, true, conv_storage_type)); + weights->decoder.generator.ups.push_back( + load_weight_norm_conv_transpose1d(store, source, "decoder.generator.ups.1", 256, 128, 12, 6, 3, 0, 1, true, conv_storage_type)); + weights->decoder.generator.noise_convs.push_back(load_conv1d(store, source, "decoder.generator.noise_convs.0", 256, 22, 12, 6, 3, 1, 1, true, conv_storage_type)); + weights->decoder.generator.noise_convs.push_back(load_conv1d(store, source, "decoder.generator.noise_convs.1", 128, 22, 1, 1, 0, 1, 1, true, conv_storage_type)); + weights->decoder.generator.noise_res.push_back( + load_generator_resblock(store, source, "decoder.generator.noise_res.0", 256, 7, {1, 3, 5}, matmul_storage_type, conv_storage_type)); + weights->decoder.generator.noise_res.push_back( + load_generator_resblock(store, source, "decoder.generator.noise_res.1", 128, 11, {1, 3, 5}, matmul_storage_type, conv_storage_type)); + weights->decoder.generator.resblocks.push_back( + load_generator_resblock(store, source, "decoder.generator.resblocks.0", 256, 3, {1, 3, 5}, matmul_storage_type, conv_storage_type)); + weights->decoder.generator.resblocks.push_back( + load_generator_resblock(store, source, "decoder.generator.resblocks.1", 256, 7, {1, 3, 5}, matmul_storage_type, conv_storage_type)); + weights->decoder.generator.resblocks.push_back( + load_generator_resblock(store, source, "decoder.generator.resblocks.2", 256, 11, {1, 3, 5}, matmul_storage_type, conv_storage_type)); + weights->decoder.generator.resblocks.push_back( + load_generator_resblock(store, source, "decoder.generator.resblocks.3", 128, 3, {1, 3, 5}, matmul_storage_type, conv_storage_type)); + weights->decoder.generator.resblocks.push_back( + load_generator_resblock(store, source, "decoder.generator.resblocks.4", 128, 7, {1, 3, 5}, matmul_storage_type, conv_storage_type)); + weights->decoder.generator.resblocks.push_back( + load_generator_resblock(store, source, "decoder.generator.resblocks.5", 128, 11, {1, 3, 5}, matmul_storage_type, conv_storage_type)); + weights->decoder.generator.conv_post = + load_weight_norm_conv1d(store, source, "decoder.generator.conv_post", 22, 128, 7, 1, 3, 1, 1, true, conv_storage_type); + weights->decoder.generator.source_linear = load_host_affine(source, "decoder.generator.m_source.l_linear", 1, 9, true); + + weights->store->upload(); + + return weights; +} + +} // namespace kokoro_ggml + +namespace engine::models::kokoro_tts { + +namespace { + +struct KokoroAssetResources { + assets::ResourceBundle bundle; + io::json::Value config; + io::json::Value voices; + std::shared_ptr weights; +}; + +KokoroAssetResources load_asset_resources(const std::filesystem::path & model_root) { + KokoroAssetResources resources; + resources.bundle = assets::ResourceBundle(std::filesystem::weakly_canonical(model_root)); + resources.bundle.add_model_files({ + {"config", "config.json"}, + {"voices", "voices.json"}, + {"weights", "kokoro-v1_0.safetensors"}, + }); + resources.config = resources.bundle.parse_json("config"); + if (!resources.config.is_object()) { + throw std::runtime_error("Kokoro config root must be an object"); + } + resources.voices = resources.bundle.parse_json("voices"); + if (!resources.voices.is_object()) { + throw std::runtime_error("Kokoro voices.json root must be an object"); + } + resources.weights = resources.bundle.open_tensor_source("weights"); + return resources; +} + +std::string language_code_from_voice_id(const std::string & voice_id) { + if (voice_id.size() < 2 || (voice_id[1] != 'f' && voice_id[1] != 'm')) { + throw std::runtime_error("invalid Kokoro voice id: " + voice_id); + } + return std::string(1, voice_id[0]); +} + +std::vector read_f32_file_exact(const std::filesystem::path & path, size_t expected_values) { + const auto blob = engine::io::read_binary_blob(path); + const size_t expected_bytes = expected_values * sizeof(float); + if (blob.size() != expected_bytes) { + throw std::runtime_error( + "unexpected Kokoro voice pack size for " + path.string() + + ": expected " + std::to_string(expected_bytes) + + " bytes, got " + std::to_string(blob.size())); + } + std::vector values(expected_values, 0.0f); + std::memcpy(values.data(), blob.data(), expected_bytes); + return values; +} + +} // namespace + +std::shared_ptr load_kokoro_assets(const std::filesystem::path & model_root) { + auto resources = load_asset_resources(model_root); + const auto & root = resources.bundle.model_root(); + + auto assets = std::make_shared(); + assets->model_root = root; + assets->config = std::move(resources.config); + assets->model_weights = std::move(resources.weights); + assets->context_length = kokoro_ggml::parse_kokoro_config_metadata(assets->config).plbert_max_position_embeddings; + assets->english_lexicon_dir = root / "misaki_en"; + assets->english_g2p_us = std::make_shared(assets->english_lexicon_dir, false); + assets->english_g2p_gb = std::make_shared(assets->english_lexicon_dir, true); + + const auto * vocab = assets->config.find("vocab"); + if (vocab == nullptr || !vocab->is_object()) { + throw std::runtime_error("Kokoro config.json is missing vocab object"); + } + for (const auto & [symbol, value] : vocab->as_object()) { + assets->vocab[symbol] = static_cast(value.as_i64()); + } + + for (const auto & [voice_id, value] : resources.voices.as_object()) { + if (!value.is_object()) { + throw std::runtime_error("Kokoro voice entry must be an object: " + voice_id); + } + KokoroVoicePack pack; + pack.id = voice_id; + pack.language_code = language_code_from_voice_id(voice_id); + const auto & object = value.as_object(); + const auto rows_it = object.find("rows"); + const auto cols_it = object.find("cols"); + const auto path_it = object.find("path"); + if (rows_it == object.end() || cols_it == object.end() || path_it == object.end()) { + throw std::runtime_error("Kokoro voice entry is missing rows/cols/path: " + voice_id); + } + pack.rows = rows_it->second.as_i64(); + pack.cols = cols_it->second.as_i64(); + pack.values = read_f32_file_exact( + root / "voices" / path_it->second.as_string(), + static_cast(pack.rows * pack.cols)); + assets->voices.emplace(voice_id, std::move(pack)); + } + + return assets; +} + +std::shared_ptr load_kokoro_backend_weights( + const KokoroAssets & assets, + ggml_backend_t backend, + engine::core::BackendType backend_type, + engine::assets::TensorStorageType matmul_storage_type, + engine::assets::TensorStorageType conv_storage_type, + size_t weight_context_bytes) { + if (!assets.model_weights) { + throw std::runtime_error("Kokoro assets are missing model weights"); + } + auto weights = kokoro_ggml::load_kokoro_weights( + assets.config, + *assets.model_weights, + backend, + backend_type, + matmul_storage_type, + conv_storage_type, + weight_context_bytes); + assets.model_weights->release_storage(); + return weights; +} + +} // namespace engine::models::kokoro_tts diff --git a/src/models/kokoro_tts/decoder.cpp b/src/models/kokoro_tts/decoder.cpp new file mode 100644 index 000000000..9beacd56f --- /dev/null +++ b/src/models/kokoro_tts/decoder.cpp @@ -0,0 +1,1352 @@ +#include "engine/models/kokoro_tts/decoder.h" + +#include "engine/models/kokoro_tts/assets.h" + +#include "engine/framework/audio/dsp.h" +#include "engine/framework/core/backend.h" +#include "engine/framework/debug/profiler.h" +#include "engine/framework/debug/trace.h" +#include "engine/framework/modules/activation_modules.h" +#include "engine/framework/modules/conv_modules.h" +#include "engine/framework/modules/linear_module.h" +#include "engine/framework/modules/norm_modules.h" +#include "engine/framework/modules/structural_modules.h" + +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace kokoro_ggml { + +namespace { + +namespace audio = engine::audio; +namespace core = engine::core; +namespace modules = engine::modules; + +constexpr float kPi = 3.14159265358979323846f; + +using engine::debug::measure_ms; + +void set_graph_output(ggml_tensor * tensor) { + ggml_set_output(tensor); + for (ggml_tensor * backing = tensor->view_src; backing != nullptr; backing = backing->view_src) { + ggml_set_output(backing); + } +} + +size_t tensor_offset_2d(int64_t row, int64_t col, int64_t cols) { + return static_cast(row * cols + col); +} + +std::vector make_zero_tensor_2d(int64_t rows, int64_t cols) { + if (rows < 0 || cols < 0) { + throw std::runtime_error("kokoro decoder tensor dimensions must be non-negative"); + } + return std::vector(static_cast(rows * cols), 0.0f); +} + +float & tensor_at(std::vector & values, int64_t row, int64_t col, int64_t cols) { + return values[tensor_offset_2d(row, col, cols)]; +} + +struct TimeMaskInputs { + ggml_tensor * keep = nullptr; + ggml_tensor * norm = nullptr; + int64_t frame_capacity = 0; +}; + +const TimeMaskInputs & add_time_mask_inputs( + ggml_context * ctx, + std::vector & masks, + int64_t frame_capacity) { + if (frame_capacity <= 0) { + throw std::runtime_error("kokoro decoder mask capacity must be positive"); + } + TimeMaskInputs mask = {}; + mask.keep = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, frame_capacity, 1, 1); + mask.norm = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, frame_capacity, 1, 1); + mask.frame_capacity = frame_capacity; + ggml_set_input(mask.keep); + ggml_set_input(mask.norm); + masks.push_back(mask); + return masks.back(); +} + +ggml_tensor * repeat_mask_like(ggml_context * ctx, ggml_tensor * mask, ggml_tensor * like) { + return ggml_repeat(ctx, mask, like); +} + +core::TensorValue broadcast_channel_3d( + core::ModuleBuildContext & ctx, + const core::TensorValue & value, + ggml_tensor * like, + int64_t channels) { + core::validate_shape(value, core::TensorShape::from_dims({channels}), "kokoro decoder adain channel value"); + ggml_tensor * reshaped = ggml_reshape_3d(ctx.ggml, value.tensor, 1, channels, 1); + return core::wrap_tensor( + ggml_repeat(ctx.ggml, reshaped, like), + core::TensorShape::from_dims({like->ne[2], channels, like->ne[0]}), + GGML_TYPE_F32); +} + +ggml_tensor * build_masked_adain_bct( + core::ModuleBuildContext & build_ctx, + ggml_tensor * x, + const core::TensorValue & gamma, + const core::TensorValue & beta, + int64_t channels, + float eps, + std::vector & masks) { + const TimeMaskInputs & mask = add_time_mask_inputs(build_ctx.ggml, masks, x->ne[0]); + ggml_tensor * keep = repeat_mask_like(build_ctx.ggml, mask.keep, x); + ggml_tensor * norm = repeat_mask_like(build_ctx.ggml, mask.norm, x); + ggml_tensor * masked = ggml_mul(build_ctx.ggml, x, norm); + ggml_tensor * mean = ggml_mean(build_ctx.ggml, masked); + ggml_tensor * centered = ggml_sub(build_ctx.ggml, x, ggml_repeat(build_ctx.ggml, mean, x)); + ggml_tensor * centered_for_variance = ggml_mul(build_ctx.ggml, centered, norm); + ggml_tensor * squared = ggml_mul(build_ctx.ggml, centered, centered_for_variance); + ggml_tensor * variance = ggml_mean(build_ctx.ggml, squared); + ggml_tensor * stddev = ggml_sqrt(build_ctx.ggml, ggml_scale_bias(build_ctx.ggml, variance, 1.0f, eps)); + ggml_tensor * normalized = ggml_div(build_ctx.ggml, centered, ggml_repeat(build_ctx.ggml, stddev, x)); + normalized = ggml_mul(build_ctx.ggml, normalized, keep); + const auto gamma_rep = broadcast_channel_3d(build_ctx, gamma, x, channels); + const auto beta_rep = broadcast_channel_3d(build_ctx, beta, x, channels); + ggml_tensor * out = ggml_add( + build_ctx.ggml, + ggml_mul(build_ctx.ggml, normalized, gamma_rep.tensor), + beta_rep.tensor); + return ggml_mul(build_ctx.ggml, out, keep); +} + +ggml_tensor * build_adain_bct( + core::ModuleBuildContext & build_ctx, + ggml_tensor * x, + const core::TensorValue & gamma, + const core::TensorValue & beta, + int64_t channels, + float eps) { + ggml_tensor * x_for_stats = ggml_cont(build_ctx.ggml, x); + ggml_tensor * mean = ggml_mean(build_ctx.ggml, x_for_stats); + ggml_tensor * centered = ggml_sub(build_ctx.ggml, x, ggml_repeat(build_ctx.ggml, mean, x)); + ggml_tensor * squared = ggml_mul(build_ctx.ggml, centered, centered); + ggml_tensor * variance = ggml_mean(build_ctx.ggml, ggml_cont(build_ctx.ggml, squared)); + ggml_tensor * stddev = ggml_sqrt(build_ctx.ggml, ggml_scale_bias(build_ctx.ggml, variance, 1.0f, eps)); + ggml_tensor * normalized = ggml_div(build_ctx.ggml, centered, ggml_repeat(build_ctx.ggml, stddev, x)); + const auto gamma_rep = broadcast_channel_3d(build_ctx, gamma, x, channels); + const auto beta_rep = broadcast_channel_3d(build_ctx, beta, x, channels); + return ggml_add( + build_ctx.ggml, + ggml_mul(build_ctx.ggml, normalized, gamma_rep.tensor), + beta_rep.tensor); +} + +void upload_time_masks( + const std::vector & masks, + int64_t valid_base_frames, + int64_t base_frame_capacity) { + if (valid_base_frames <= 0 || valid_base_frames > base_frame_capacity) { + throw std::runtime_error("kokoro decoder valid frame count exceeds prepared capacity"); + } + for (const TimeMaskInputs & mask : masks) { + if (mask.frame_capacity <= 0 || mask.frame_capacity < base_frame_capacity) { + std::ostringstream message; + message << "kokoro decoder mask capacity is not aligned with graph capacity" + << " mask_frames=" << mask.frame_capacity + << " valid_base_frames=" << valid_base_frames + << " base_frame_capacity=" << base_frame_capacity; + throw std::runtime_error(message.str()); + } + const int64_t scale = mask.frame_capacity / base_frame_capacity; + const int64_t offset = mask.frame_capacity % base_frame_capacity; + if (scale <= 0 || offset > 1) { + std::ostringstream message; + message << "kokoro decoder mask capacity is not represented by the generator length schedule" + << " mask_frames=" << mask.frame_capacity + << " base_frame_capacity=" << base_frame_capacity + << " scale=" << scale + << " offset=" << offset; + throw std::runtime_error(message.str()); + } + const int64_t valid_frames = valid_base_frames * scale + offset; + if (valid_frames <= 0 || valid_frames > mask.frame_capacity) { + throw std::runtime_error("kokoro decoder mask valid frame count is invalid"); + } + std::vector keep(static_cast(mask.frame_capacity), 0.0f); + std::vector norm(static_cast(mask.frame_capacity), 0.0f); + std::fill(keep.begin(), keep.begin() + valid_frames, 1.0f); + const float norm_value = static_cast(mask.frame_capacity) / static_cast(valid_frames); + std::fill(norm.begin(), norm.begin() + valid_frames, norm_value); + ggml_backend_tensor_set(mask.keep, keep.data(), 0, ggml_nbytes(mask.keep)); + if (mask.norm != nullptr) { + ggml_backend_tensor_set(mask.norm, norm.data(), 0, ggml_nbytes(mask.norm)); + } + } +} + +ggml_tensor * reflect_pad_left_1_bct_decoder(ggml_context * ctx, ggml_tensor * x) { + core::ModuleBuildContext build_ctx = {}; + build_ctx.ggml = ctx; + const auto input = core::wrap_tensor( + x, + core::TensorShape::from_dims({x->ne[2], x->ne[1], x->ne[0]}), + GGML_TYPE_F32); + const auto output = modules::ReflectPad1dModule({1, 0}).build(build_ctx, input); + return output.tensor; +} + +ggml_tensor * build_snake1d_bct_decoder( + ggml_context * ctx, + ggml_tensor * x, + const core::TensorValue & alpha) { + core::ModuleBuildContext build_ctx = {}; + build_ctx.ggml = ctx; + const int64_t batch = x->ne[2]; + const int64_t channels = x->ne[1]; + const int64_t frames = x->ne[0]; + const auto input = core::wrap_tensor(x, core::TensorShape::from_dims({batch, channels, frames}), GGML_TYPE_F32); + modules::Snake1dWeights weights = {}; + weights.alpha = core::wrap_tensor( + ggml_reshape_1d(ctx, alpha.tensor, channels), + core::TensorShape::from_dims({channels}), + alpha.type); + const auto output = modules::Snake1dModule({channels}).build(build_ctx, input, weights); + return output.tensor; +} + +ggml_tensor * build_adaptive_instance_norm_bct_decoder( + ggml_context * ctx, + ggml_tensor * x, + const KokoroWeights::AdaIn1dWeights & weights, + const core::TensorValue & style, + std::vector & masks, + bool use_time_masks) { + const int64_t batch = x->ne[2]; + const int64_t channels = x->ne[1]; + const int64_t frames = x->ne[0]; + core::ModuleBuildContext build_ctx = {}; + build_ctx.ggml = ctx; + const auto affine = modules::LinearModule({ + weights.fc.in_features, + weights.fc.out_features, + weights.fc.use_bias}).build( + build_ctx, + style, + { + weights.fc.weight, + weights.fc.bias, + }); + ggml_tensor * affine_contiguous = ggml_cont(ctx, affine.tensor); + const auto scale_delta = core::wrap_tensor( + ggml_view_2d(ctx, affine_contiguous, channels, 1, affine_contiguous->nb[1], 0), + core::TensorShape::from_dims({1, channels}), + GGML_TYPE_F32); + const auto shift = core::wrap_tensor( + ggml_view_2d( + ctx, + affine_contiguous, + channels, + 1, + affine_contiguous->nb[1], + static_cast(channels) * affine_contiguous->nb[0]), + core::TensorShape::from_dims({1, channels}), + GGML_TYPE_F32); + const auto gamma_2d = core::wrap_tensor( + ggml_scale_bias(ctx, scale_delta.tensor, 1.0f, 1.0f), + scale_delta.shape, + GGML_TYPE_F32); + const auto gamma = core::wrap_tensor( + ggml_reshape_1d(ctx, gamma_2d.tensor, channels), + core::TensorShape::from_dims({channels}), + GGML_TYPE_F32); + const auto beta = core::wrap_tensor( + ggml_reshape_1d(ctx, shift.tensor, channels), + core::TensorShape::from_dims({channels}), + GGML_TYPE_F32); + (void)batch; + (void)frames; + if (use_time_masks) { + return build_masked_adain_bct(build_ctx, x, gamma, beta, channels, weights.eps, masks); + } + return build_adain_bct(build_ctx, x, gamma, beta, channels, weights.eps); +} + +template +modules::Conv1dWeights make_conv1d_weights(core::ModuleBuildContext & ctx, const ConvWeightsT & conv) { + (void)ctx; + modules::Conv1dWeights weights = {}; + weights.weight = conv.weight; + if (conv.use_bias) { + weights.bias = conv.bias; + } + return weights; +} + +modules::ConvTranspose1dWeights make_conv_transpose1d_weights( + core::ModuleBuildContext & ctx, + const KokoroWeights::WeightNormConvTranspose1dWeights & conv) { + (void)ctx; + modules::ConvTranspose1dWeights weights = {}; + weights.weight = conv.groups == 1 ? conv.weight : conv.dense_weight; + if (conv.use_bias) { + weights.bias = conv.bias; + } + return weights; +} + +template +ggml_tensor * build_decoder_conv1d_bct( + ggml_context * ctx, + ggml_tensor * input, + const ConvWeightsT & conv, + bool allow_pointwise_fastpath) { + if (conv.groups != 1) { + throw std::runtime_error("kokoro decoder conv1d requires groups == 1"); + } + if (allow_pointwise_fastpath && + conv.kernel == 1 && + conv.stride == 1 && + conv.padding == 0 && + conv.dilation == 1) { + ggml_tensor * x = ggml_cont(ctx, input); + ggml_tensor * x_2d = ggml_reshape_2d(ctx, x, input->ne[0], input->ne[1]); + ggml_tensor * x_t = ggml_cont(ctx, ggml_transpose(ctx, x_2d)); + ggml_tensor * w = ggml_reshape_2d(ctx, conv.weight.tensor, conv.in_channels, conv.out_channels); + ggml_tensor * y_t = ggml_mul_mat(ctx, w, x_t); + ggml_tensor * y_2d = ggml_cont(ctx, ggml_transpose(ctx, y_t)); + if (conv.use_bias) { + ggml_tensor * b = ggml_reshape_2d(ctx, conv.bias->tensor, 1, conv.out_channels); + y_2d = ggml_add(ctx, y_2d, b); + } + return ggml_reshape_3d(ctx, y_2d, y_2d->ne[0], y_2d->ne[1], 1); + } + core::ModuleBuildContext build_ctx = {}; + build_ctx.ggml = ctx; + const auto input_bct = core::wrap_tensor( + ggml_cont(ctx, input), + core::TensorShape::from_dims({input->ne[2], input->ne[1], input->ne[0]}), + GGML_TYPE_F32); + return modules::Conv1dModule({ + conv.in_channels, + conv.out_channels, + conv.kernel, + static_cast(conv.stride), + static_cast(conv.padding), + static_cast(conv.dilation), + conv.use_bias}).build(build_ctx, input_bct, make_conv1d_weights(build_ctx, conv)).tensor; +} + +ggml_tensor * build_conv_transpose1d_bct_decoder( + ggml_context * ctx, + ggml_tensor * input, + const KokoroWeights::WeightNormConvTranspose1dWeights & conv) { + if (conv.groups != 1) { + throw std::runtime_error("kokoro decoder conv_transpose1d requires groups == 1"); + } + core::ModuleBuildContext build_ctx = {}; + build_ctx.ggml = ctx; + const auto input_bct = core::wrap_tensor( + ggml_cont(ctx, input), + core::TensorShape::from_dims({input->ne[2], input->ne[1], input->ne[0]}), + GGML_TYPE_F32); + auto output_bct = modules::ConvTranspose1dModule({ + conv.in_channels, + conv.out_channels, + conv.kernel, + static_cast(conv.stride), + 0, + 1, + conv.use_bias}).build(build_ctx, input_bct, make_conv_transpose1d_weights(build_ctx, conv)); + const int64_t cropped_len = + (input->ne[0] - 1) * conv.stride - 2 * conv.padding + conv.kernel + conv.output_padding; + const auto cropped = core::wrap_tensor( + ggml_cont( + ctx, + ggml_view_3d( + ctx, + output_bct.tensor, + cropped_len, + conv.out_channels, + 1, + output_bct.tensor->nb[1], + output_bct.tensor->nb[2], + static_cast(conv.padding) * sizeof(float))), + core::TensorShape::from_dims({1, conv.out_channels, cropped_len}), + GGML_TYPE_F32); + return cropped.tensor; +} + +bool can_use_phase_shuffle_conv_transpose1d(const KokoroWeights::WeightNormConvTranspose1dWeights & conv) { + return conv.groups == 1 && + conv.output_padding == 0 && + conv.kernel == conv.stride * 2 && + conv.padding > 0 && + conv.padding < conv.stride; +} + +ggml_tensor * build_phase_shuffle_conv_transpose1d_bct_decoder( + ggml_context * ctx, + ggml_tensor * input, + const KokoroWeights::WeightNormConvTranspose1dWeights & conv) { + if (!can_use_phase_shuffle_conv_transpose1d(conv)) { + throw std::runtime_error("kokoro decoder phase-shuffle conv transpose shape is unsupported"); + } + if (!conv.phase_shuffle_conv) { + throw std::runtime_error("kokoro decoder phase-shuffle conv was not prepared"); + } + + ggml_tensor * phases = build_decoder_conv1d_bct(ctx, input, *conv.phase_shuffle_conv, false); + phases = ggml_reshape_4d(ctx, phases, phases->ne[0], conv.stride, conv.out_channels, 1); + ggml_tensor * interleaved = ggml_cont(ctx, ggml_permute(ctx, phases, 1, 0, 2, 3)); + return ggml_reshape_3d(ctx, interleaved, input->ne[0] * conv.stride, conv.out_channels, 1); +} + +struct DeterministicRng { + uint64_t state = 0; + bool has_spare = false; + float spare = 0.0f; + + explicit DeterministicRng(uint64_t seed) : state(seed) {} + + uint32_t next_u32() { + state = state * 6364136223846793005ULL + 1ULL; + return static_cast(state >> 32); + } + + float uniform01() { + return (static_cast(next_u32()) + 0.5f) / 4294967296.0f; + } + + float normal() { + if (has_spare) { + has_spare = false; + return spare; + } + const float u1 = std::max(uniform01(), 1.0e-12f); + const float u2 = uniform01(); + const float radius = std::sqrt(-2.0f * std::log(u1)); + const float theta = 2.0f * kPi * u2; + spare = radius * std::sin(theta); + has_spare = true; + return radius * std::cos(theta); + } +}; + +struct HarmonicConditioning { + std::vector features; + int64_t feature_rows = 0; + int64_t feature_cols = 0; + int64_t valid_feature_cols = 0; +}; + +class SourceSignalGraphRuntime { +public: + SourceSignalGraphRuntime( + const KokoroWeights::GeneratorWeights & generator, + ggml_backend_t backend, + int n_threads) + : generator_(&generator), + backend_(backend), + n_threads_(std::max(1, n_threads)) {} + + void run( + const std::vector & phase, + const std::vector & noise, + const std::vector & voiced, + int64_t sample_count, + int64_t harmonic_dims, + std::vector & output) { + if (sample_count <= 0 || harmonic_dims <= 0) { + throw std::runtime_error("Kokoro source graph dimensions must be positive"); + } + if (static_cast(phase.size()) != sample_count * harmonic_dims || + static_cast(noise.size()) != sample_count * harmonic_dims || + static_cast(voiced.size()) != sample_count) { + throw std::runtime_error("Kokoro source graph input shape mismatch"); + } + prepare(sample_count, harmonic_dims); + session_->run(phase, noise, voiced, output); + } + +private: + struct Session { + const KokoroWeights::GeneratorWeights * generator = nullptr; + ggml_backend_t backend = nullptr; + int n_threads = 1; + int64_t sample_count = 0; + int64_t harmonic_dims = 0; + ggml_context * ctx = nullptr; + ggml_tensor * phase_in = nullptr; + ggml_tensor * noise_in = nullptr; + ggml_tensor * voiced_in = nullptr; + ggml_tensor * output = nullptr; + ggml_cgraph * graph = nullptr; + ggml_gallocr_t gallocr = nullptr; + + Session( + const KokoroWeights::GeneratorWeights & generator_in, + ggml_backend_t backend_in, + int n_threads_in, + int64_t sample_count_in, + int64_t harmonic_dims_in) + : generator(&generator_in), + backend(backend_in), + n_threads(std::max(1, n_threads_in)), + sample_count(sample_count_in), + harmonic_dims(harmonic_dims_in) { + if (generator->source_linear.weight.empty() || + static_cast(generator->source_linear.weight.size()) != harmonic_dims || + (generator->source_linear.use_bias && generator->source_linear.bias.empty())) { + throw std::runtime_error("kokoro decoder source affine weights are missing"); + } + ggml_init_params params{ + /*.mem_size =*/ 64ull * 1024ull * 1024ull, + /*.mem_buffer =*/ nullptr, + /*.no_alloc =*/ true, + }; + ctx = ggml_init(params); + if (!ctx) { + throw std::runtime_error("failed to initialize ggml context for Kokoro source graph"); + } + try { + phase_in = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, sample_count, harmonic_dims); + noise_in = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, sample_count, harmonic_dims); + voiced_in = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, sample_count); + ggml_set_input(phase_in); + ggml_set_input(noise_in); + ggml_set_input(voiced_in); + + ggml_tensor * unvoiced = ggml_scale_bias(ctx, voiced_in, -1.0f, 1.0f); + ggml_tensor * noise_amp = ggml_add( + ctx, + ggml_scale(ctx, voiced_in, generator->noise_std), + ggml_scale(ctx, unvoiced, generator->sine_amp / 3.0f)); + const float source_bias = + generator->source_linear.use_bias ? generator->source_linear.bias[0] : 0.0f; + ggml_tensor * merged = ggml_scale_bias( + ctx, + ggml_view_1d(ctx, phase_in, sample_count, 0), + 0.0f, + source_bias); + for (int64_t h = 0; h < harmonic_dims; ++h) { + const size_t row_offset = static_cast(h * sample_count) * sizeof(float); + ggml_tensor * phase_h = ggml_view_1d(ctx, phase_in, sample_count, row_offset); + ggml_tensor * noise_h = ggml_view_1d(ctx, noise_in, sample_count, row_offset); + ggml_tensor * base_sine = ggml_scale(ctx, ggml_sin(ctx, phase_h), generator->sine_amp); + ggml_tensor * voiced_sine = ggml_mul(ctx, base_sine, voiced_in); + ggml_tensor * noisy_sine = ggml_add(ctx, voiced_sine, ggml_mul(ctx, noise_amp, noise_h)); + merged = ggml_add( + ctx, + merged, + ggml_scale(ctx, noisy_sine, generator->source_linear.weight[static_cast(h)])); + } + output = ggml_tanh(ctx, merged); + output = ggml_cont(ctx, output); + set_graph_output(output); + + graph = ggml_new_graph_custom(ctx, 4096, false); + ggml_build_forward_expand(graph, output); + core::set_backend_threads(backend, n_threads); + gallocr = ggml_gallocr_new(ggml_backend_get_default_buffer_type(backend)); + if (gallocr == nullptr || !ggml_gallocr_alloc_graph(gallocr, graph)) { + throw std::runtime_error("failed to allocate Kokoro source graph tensors"); + } + } catch (...) { + if (gallocr) { + ggml_gallocr_free(gallocr); + } + if (ctx) { + ggml_free(ctx); + } + ctx = nullptr; + throw; + } + } + + ~Session() { + if (gallocr) { + ggml_gallocr_free(gallocr); + } + if (ctx) { + ggml_free(ctx); + } + } + + void run( + const std::vector & phase, + const std::vector & noise, + const std::vector & voiced, + std::vector & out) { + const double upload_ms = measure_ms([&]() { + ggml_backend_tensor_set(phase_in, phase.data(), 0, ggml_nbytes(phase_in)); + ggml_backend_tensor_set(noise_in, noise.data(), 0, ggml_nbytes(noise_in)); + ggml_backend_tensor_set(voiced_in, voiced.data(), 0, ggml_nbytes(voiced_in)); + }); + engine::debug::timing_log_scalar("kokoro.decoder.conditioning_source.graph.upload_ms", upload_ms); + core::set_backend_threads(backend, n_threads); + ggml_status status = GGML_STATUS_SUCCESS; + const double compute_ms = measure_ms([&]() { + status = engine::core::compute_backend_graph(backend, graph); + }); + engine::debug::timing_log_scalar("kokoro.decoder.conditioning_source.graph.compute_ms", compute_ms); + if (status != GGML_STATUS_SUCCESS) { + throw std::runtime_error(std::string("kokoro source graph compute failed: ") + ggml_status_to_string(status)); + } + out.resize(static_cast(sample_count)); + const double read_ms = measure_ms([&]() { + ggml_backend_tensor_get(output, out.data(), 0, ggml_nbytes(output)); + }); + engine::debug::timing_log_scalar("kokoro.decoder.conditioning_source.graph.read_ms", read_ms); + } + }; + + void prepare(int64_t sample_count, int64_t harmonic_dims) { + if (session_ && + session_->sample_count == sample_count && + session_->harmonic_dims == harmonic_dims) { + return; + } + session_.reset(); + session_ = std::make_unique( + *generator_, + backend_, + n_threads_, + sample_count, + harmonic_dims); + } + + const KokoroWeights::GeneratorWeights * generator_ = nullptr; + ggml_backend_t backend_ = nullptr; + int n_threads_ = 1; + std::unique_ptr session_; +}; + +class ConditioningGraphRuntime { +public: + ConditioningGraphRuntime( + const KokoroWeights::GeneratorWeights & generator, + ggml_backend_t backend, + int64_t decoder_frame_capacity, + int64_t conditioning_frame_capacity, + int n_threads, + bool use_device_backend) + : generator_(&generator), + n_threads_(std::max(1, n_threads)), + session_(std::make_unique( + generator, + backend, + decoder_frame_capacity, + conditioning_frame_capacity, + n_threads_, + use_device_backend)) {} + + const HarmonicConditioning & run(const std::vector & f0_curve, DeterministicRng & rng) { + return session_->run(f0_curve, rng); + } + +private: + struct Session { + const KokoroWeights::GeneratorWeights * generator = nullptr; + ggml_backend_t backend = nullptr; + int64_t decoder_frame_capacity = 0; + int64_t conditioning_frame_capacity = 0; + int64_t upsample_scale = 300; + int64_t conditioning_sample_capacity = 0; + int64_t harmonic_dims = 0; + int n_threads = 1; + bool use_device_backend = false; + audio::STFTConfig stft_config = {}; + std::vector source_signal; + std::vector valid_phase_coarse; + std::vector valid_phase_upsampled; + std::vector voiced_mask; + std::vector harmonic_noise; + std::vector valid_source_signal; + HarmonicConditioning conditioning; + std::unique_ptr source_graph; + + Session( + const KokoroWeights::GeneratorWeights & generator_in, + ggml_backend_t backend_in, + int64_t decoder_frame_capacity_in, + int64_t conditioning_frame_capacity_in, + int n_threads_in, + bool use_device_backend_in) + : generator(&generator_in), + backend(backend_in), + decoder_frame_capacity(decoder_frame_capacity_in), + conditioning_frame_capacity(conditioning_frame_capacity_in), + conditioning_sample_capacity(decoder_frame_capacity_in * upsample_scale), + harmonic_dims(generator_in.harmonic_num + 1), + n_threads(n_threads_in), + use_device_backend(use_device_backend_in), + stft_config{ + generator_in.gen_istft_n_fft, + generator_in.gen_istft_hop_size, + generator_in.gen_istft_n_fft, + true, + audio::STFTPadMode::Reflect, + audio::STFTFamily::Kokoro, + }, + source_signal(static_cast(conditioning_sample_capacity), 0.0f) { + conditioning.feature_rows = (stft_config.n_fft / 2 + 1) * 2; + conditioning.feature_cols = conditioning_frame_capacity; + conditioning.valid_feature_cols = 0; + conditioning.features = make_zero_tensor_2d(conditioning.feature_rows, conditioning.feature_cols); + if (use_device_backend) { + source_graph = std::make_unique(generator_in, backend, n_threads); + } + } + + const HarmonicConditioning & run(const std::vector & f0_curve, DeterministicRng & rng) { + const int64_t valid_decoder_frames = static_cast(f0_curve.size()); + if (valid_decoder_frames <= 0 || valid_decoder_frames > decoder_frame_capacity) { + throw std::runtime_error("Kokoro conditioning graph input frame count exceeds prepared capacity"); + } + const double source_ms = measure_ms([&]() { + synthesize_source_signal(f0_curve, valid_decoder_frames, rng); + }); + const double stft_ms = measure_ms([&]() { + compute_conditioning_features(valid_decoder_frames); + }); + engine::debug::timing_log_scalar("kokoro.decoder.conditioning_source_ms", source_ms); + engine::debug::timing_log_scalar("kokoro.decoder.conditioning_stft_ms", stft_ms); + return conditioning; + } + + void synthesize_source_signal( + const std::vector & f0_curve, + int64_t valid_decoder_frames, + DeterministicRng & rng) { + const int64_t valid_sample_count = valid_decoder_frames * upsample_scale; + for (int64_t h = 1; h < harmonic_dims; ++h) { + (void) rng.uniform01(); + } + valid_phase_coarse.resize(static_cast(harmonic_dims * valid_decoder_frames)); +#ifdef _OPENMP + #pragma omp parallel for if(harmonic_dims >= 4) +#endif + for (int64_t h = 0; h < harmonic_dims; ++h) { + const float harmonic = static_cast(h + 1); + float accumulated_phase = 0.0f; + for (int64_t t = 0; t < valid_decoder_frames; ++t) { + accumulated_phase += + std::fmod( + (f0_curve[static_cast(t)] * harmonic) / + static_cast(generator->sampling_rate), + 1.0f); + const float phase = accumulated_phase * 2.0f * kPi; + tensor_at(valid_phase_coarse, h, t, valid_decoder_frames) = + phase * static_cast(upsample_scale); + } + } + harmonic_noise.resize(static_cast(harmonic_dims * valid_sample_count)); + for (float & value : harmonic_noise) { + value = rng.normal(); + } + if (generator->source_linear.weight.empty() || + (generator->source_linear.use_bias && generator->source_linear.bias.empty())) { + throw std::runtime_error("kokoro decoder source affine weights are missing"); + } + valid_phase_upsampled.resize(static_cast(harmonic_dims * valid_sample_count)); + voiced_mask.resize(static_cast(valid_sample_count)); +#ifdef _OPENMP + #pragma omp parallel for if(valid_sample_count >= 4096) +#endif + for (int64_t t = 0; t < valid_sample_count; ++t) { + const int64_t frame = std::min(t / upsample_scale, valid_decoder_frames - 1); + const float uv_t = + f0_curve[static_cast(frame)] > generator->voiced_threshold ? 1.0f : 0.0f; + voiced_mask[static_cast(t)] = uv_t; + const float src = (static_cast(t) + 0.5f) / static_cast(upsample_scale) - 0.5f; + int64_t left = 0; + int64_t right = 0; + float frac = 0.0f; + if (src >= static_cast(valid_decoder_frames - 1)) { + left = valid_decoder_frames - 1; + right = valid_decoder_frames - 1; + } else if (src > 0.0f) { + left = static_cast(std::floor(src)); + right = left + 1; + frac = src - static_cast(left); + } + for (int64_t h = 0; h < harmonic_dims; ++h) { + tensor_at(valid_phase_upsampled, h, t, valid_sample_count) = + tensor_at(valid_phase_coarse, h, left, valid_decoder_frames) * (1.0f - frac) + + tensor_at(valid_phase_coarse, h, right, valid_decoder_frames) * frac; + } + } + if (source_graph) { + source_graph->run( + valid_phase_upsampled, + harmonic_noise, + voiced_mask, + valid_sample_count, + harmonic_dims, + source_signal); + return; + } + const float source_bias = + generator->source_linear.use_bias ? generator->source_linear.bias[0] : 0.0f; +#ifdef _OPENMP + #pragma omp parallel for if(valid_sample_count >= 4096) +#endif + for (int64_t t = 0; t < valid_sample_count; ++t) { + float merged = source_bias; + const float uv_t = voiced_mask[static_cast(t)]; + const float noise_amp = uv_t * generator->noise_std + (1.0f - uv_t) * generator->sine_amp / 3.0f; + for (int64_t h = 0; h < harmonic_dims; ++h) { + const float base_sine = + std::sin(tensor_at(valid_phase_upsampled, h, t, valid_sample_count)) * generator->sine_amp; + const float noisy_sine = + base_sine * uv_t + + noise_amp * tensor_at(harmonic_noise, h, t, valid_sample_count); + merged += generator->source_linear.weight[static_cast(h)] * noisy_sine; + } + source_signal[static_cast(t)] = std::tanh(merged); + } + } + + void compute_conditioning_features(int64_t valid_decoder_frames) { + const int64_t valid_sample_count = valid_decoder_frames * upsample_scale; + const auto & window = audio::get_cached_stft_window(stft_config); + valid_source_signal.assign(source_signal.begin(), source_signal.begin() + valid_sample_count); + const audio::AudioTensor complex = audio::STFT().compute_complex( + valid_source_signal, + window, + 1, + valid_sample_count, + stft_config, + static_cast(n_threads)); + const int64_t bins = complex.shape[1]; + const int64_t frames = complex.shape[2]; + if (conditioning.feature_rows != bins * 2 || frames > conditioning.feature_cols) { + throw std::runtime_error("kokoro conditioning STFT output exceeds prepared capacity"); + } + conditioning.valid_feature_cols = frames; + std::fill(conditioning.features.begin(), conditioning.features.end(), 0.0f); +#ifdef _OPENMP + #pragma omp parallel for collapse(2) if(bins * frames >= 4096) +#endif + for (int64_t bin = 0; bin < bins; ++bin) { + for (int64_t frame = 0; frame < frames; ++frame) { + const size_t base = static_cast(((bin * frames) + frame) * 2); + const float re = complex.values[base]; + const float im = complex.values[base + 1]; + tensor_at(conditioning.features, bin, frame, conditioning.feature_cols) = std::sqrt(re * re + im * im); + tensor_at(conditioning.features, bin + bins, frame, conditioning.feature_cols) = std::atan2(im, re); + } + } + } + }; + + const KokoroWeights::GeneratorWeights * generator_ = nullptr; + int n_threads_ = 1; + std::unique_ptr session_; +}; + +struct GeneratorGraphStage { + const KokoroWeights::Conv1dWeights * noise_conv = nullptr; + const KokoroWeights::GeneratorResBlockWeights * noise_res = nullptr; + const KokoroWeights::WeightNormConvTranspose1dWeights * up = nullptr; + std::array resblocks{}; + bool reflect_pad = false; +}; + +ggml_tensor * build_generator_resblock( + ggml_context * ctx, + ggml_tensor * x, + const KokoroWeights::GeneratorResBlockWeights & block, + const core::TensorValue & style, + bool allow_cpu_pointwise_fastpath, + std::vector & time_masks, + bool use_time_masks) { + ggml_tensor * current = x; + for (size_t i = 0; i < block.convs1.size(); ++i) { + ggml_tensor * xt = + build_adaptive_instance_norm_bct_decoder(ctx, current, block.adain1[i], style, time_masks, use_time_masks); + xt = build_snake1d_bct_decoder(ctx, xt, block.alpha1[i]); + xt = build_decoder_conv1d_bct(ctx, xt, block.convs1[i], allow_cpu_pointwise_fastpath); + xt = build_adaptive_instance_norm_bct_decoder(ctx, xt, block.adain2[i], style, time_masks, use_time_masks); + xt = build_snake1d_bct_decoder(ctx, xt, block.alpha2[i]); + xt = build_decoder_conv1d_bct(ctx, xt, block.convs2[i], allow_cpu_pointwise_fastpath); + current = ggml_add(ctx, xt, current); + } + return current; +} + +struct GeneratorGraphSession { + const KokoroWeights::GeneratorWeights * weights = nullptr; + int64_t decoder_frame_capacity = 0; + int64_t conditioning_frame_capacity = 0; + int64_t style_dim = 0; + int n_threads = 1; + bool use_device_backend = false; + ggml_context * ctx = nullptr; + ggml_tensor * decoder_x_in = nullptr; + ggml_tensor * conditioning_in = nullptr; + ggml_tensor * style_in = nullptr; + ggml_tensor * output = nullptr; + std::vector time_masks; + ggml_cgraph * graph = nullptr; + ggml_backend_t backend = nullptr; + ggml_gallocr_t gallocr = nullptr; + + GeneratorGraphSession( + const KokoroWeights::GeneratorWeights & weights_in, + ggml_backend_t backend_in, + const std::vector & stages, + int64_t decoder_frame_capacity_in, + int64_t conditioning_frame_capacity_in, + int64_t style_dim_in, + int n_threads_in, + bool use_device_backend_in) + : weights(&weights_in), + decoder_frame_capacity(decoder_frame_capacity_in), + conditioning_frame_capacity(conditioning_frame_capacity_in), + style_dim(style_dim_in), + n_threads(n_threads_in), + use_device_backend(use_device_backend_in), + backend(backend_in) { + ggml_init_params params{ + /*.mem_size =*/ 256ull * 1024ull * 1024ull, + /*.mem_buffer =*/ nullptr, + /*.no_alloc =*/ true, + }; + ctx = ggml_init(params); + if (!ctx) { + throw std::runtime_error("failed to initialize ggml context for Kokoro generator graph"); + } + + try { + const bool allow_cpu_pointwise_fastpath = !use_device_backend; + + decoder_x_in = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, decoder_frame_capacity, 512, 1); + conditioning_in = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, conditioning_frame_capacity, 22, 1); + style_in = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, style_dim, 1); + ggml_set_input(decoder_x_in); + ggml_set_input(conditioning_in); + ggml_set_input(style_in); + const auto style = core::wrap_tensor( + style_in, + core::TensorShape::from_dims({1, style_dim}), + GGML_TYPE_F32); + + ggml_tensor * current = decoder_x_in; + for (size_t i = 0; i < stages.size(); ++i) { + const GeneratorGraphStage & stage = stages[i]; + ggml_tensor * source = build_decoder_conv1d_bct(ctx, conditioning_in, *stage.noise_conv, allow_cpu_pointwise_fastpath); + source = build_generator_resblock( + ctx, + source, + *stage.noise_res, + style, + allow_cpu_pointwise_fastpath, + time_masks, + false); + + ggml_tensor * x = ggml_leaky_relu(ctx, current, 0.1f, false); + x = use_device_backend + ? build_phase_shuffle_conv_transpose1d_bct_decoder(ctx, x, *stage.up) + : build_conv_transpose1d_bct_decoder(ctx, x, *stage.up); + if (stage.reflect_pad) { + x = reflect_pad_left_1_bct_decoder(ctx, x); + } + x = ggml_add(ctx, x, source); + + ggml_tensor * stage_sum = nullptr; + for (size_t j = 0; j < stage.resblocks.size(); ++j) { + ggml_tensor * block = build_generator_resblock( + ctx, + x, + *stage.resblocks[j], + style, + allow_cpu_pointwise_fastpath, + time_masks, + false); + stage_sum = stage_sum == nullptr ? block : ggml_add(ctx, stage_sum, block); + } + current = ggml_scale(ctx, stage_sum, 1.0f / 3.0f); + } + + current = ggml_leaky_relu(ctx, current, 0.01f, false); + output = build_decoder_conv1d_bct(ctx, current, weights->conv_post, allow_cpu_pointwise_fastpath); + output = ggml_cont(ctx, output); + set_graph_output(output); + + graph = ggml_new_graph_custom(ctx, 65536, false); + ggml_build_forward_expand(graph, output); + + core::set_backend_threads(backend, n_threads); + const double alloc_ms = measure_ms([&]() { + gallocr = ggml_gallocr_new(ggml_backend_get_default_buffer_type(backend)); + if (gallocr != nullptr) { + if (!ggml_gallocr_alloc_graph(gallocr, graph)) { + ggml_gallocr_free(gallocr); + gallocr = nullptr; + } + } + }); + if (gallocr == nullptr) { + throw std::runtime_error( + std::string("failed to allocate Kokoro generator graph ") + + (use_device_backend ? "device" : "host") + + " tensors"); + } + const double materialize_ms = measure_ms([&]() { + std::vector decoder(static_cast(512 * decoder_frame_capacity), 0.0f); + std::vector conditioning(static_cast(22 * conditioning_frame_capacity), 0.0f); + std::vector style(static_cast(style_dim), 0.0f); + ggml_backend_tensor_set(decoder_x_in, decoder.data(), 0, ggml_nbytes(decoder_x_in)); + ggml_backend_tensor_set(conditioning_in, conditioning.data(), 0, ggml_nbytes(conditioning_in)); + ggml_backend_tensor_set(style_in, style.data(), 0, ggml_nbytes(style_in)); + upload_time_masks(time_masks, decoder_frame_capacity, decoder_frame_capacity); + }); + engine::debug::timing_log_scalar("kokoro.graph.build.decoder_generator_alloc_ms", alloc_ms); + engine::debug::timing_log_scalar("kokoro.graph.build.decoder_generator_materialize_ms", materialize_ms); + } catch (...) { + if (gallocr) { + ggml_gallocr_free(gallocr); + } + if (ctx) { + ggml_free(ctx); + } + ctx = nullptr; + throw; + } + } + + ~GeneratorGraphSession() { + if (gallocr) { + ggml_gallocr_free(gallocr); + } + if (ctx) { + ggml_free(ctx); + } + } + + std::vector run( + const std::vector * decoder_x, + int64_t decoder_x_rows, + int64_t decoder_x_cols, + const ggml_tensor * decoder_x_tensor, + const HarmonicConditioning & conditioning, + const std::vector & style) { + if (decoder_x_tensor == nullptr) { + if (decoder_x == nullptr || decoder_x_rows != 512 || decoder_x_cols != decoder_frame_capacity) { + throw std::runtime_error("Kokoro generator decoder input shape does not match exact graph capacity"); + } + } else if (decoder_x_tensor->ne[1] != 512 || + decoder_x_tensor->ne[0] != decoder_frame_capacity || + decoder_x_cols != decoder_frame_capacity) { + throw std::runtime_error("Kokoro generator decoder backend tensor shape does not match exact graph capacity"); + } + + if (conditioning.feature_rows != 22 || + conditioning.feature_cols != conditioning_frame_capacity || + conditioning.valid_feature_cols != conditioning_frame_capacity) { + throw std::runtime_error("Kokoro generator conditioning shape does not match exact graph capacity"); + } + + double decoder_upload_ms = 0.0; + if (decoder_x_tensor != nullptr && + decoder_x_tensor->ne[0] == decoder_frame_capacity && + decoder_x_cols == decoder_frame_capacity) { + decoder_upload_ms = measure_ms([&]() { + ggml_backend_tensor_copy(decoder_x_tensor, decoder_x_in); + }); + } else { + const int64_t valid_decoder_frames = decoder_x_cols; + std::vector padded_decoder(static_cast(512 * decoder_frame_capacity), 0.0f); + if (decoder_x_tensor != nullptr) { + const int64_t source_frames = decoder_x_tensor->ne[0]; + std::vector source_decoder(static_cast(512 * source_frames), 0.0f); + ggml_backend_tensor_get( + decoder_x_tensor, + source_decoder.data(), + 0, + static_cast(512 * source_frames) * sizeof(float)); + for (int64_t channel = 0; channel < 512; ++channel) { + std::memcpy( + padded_decoder.data() + static_cast(channel * decoder_frame_capacity), + source_decoder.data() + static_cast(channel * source_frames), + static_cast(valid_decoder_frames) * sizeof(float)); + } + } else { + for (int64_t channel = 0; channel < 512; ++channel) { + std::memcpy( + padded_decoder.data() + static_cast(channel * decoder_frame_capacity), + decoder_x->data() + static_cast(channel * valid_decoder_frames), + static_cast(valid_decoder_frames) * sizeof(float)); + } + } + decoder_upload_ms = measure_ms([&]() { + ggml_backend_tensor_set(decoder_x_in, padded_decoder.data(), 0, ggml_nbytes(decoder_x_in)); + }); + } + engine::debug::timing_log_scalar("kokoro.decoder.generator.decoder_upload_ms", decoder_upload_ms); + const double conditioning_upload_ms = measure_ms([&]() { + ggml_backend_tensor_set(conditioning_in, conditioning.features.data(), 0, ggml_nbytes(conditioning_in)); + }); + engine::debug::timing_log_scalar("kokoro.decoder.generator.conditioning_upload_ms", conditioning_upload_ms); + const double style_upload_ms = measure_ms([&]() { + ggml_backend_tensor_set(style_in, style.data(), 0, ggml_nbytes(style_in)); + upload_time_masks(time_masks, decoder_x_cols, decoder_frame_capacity); + }); + engine::debug::timing_log_scalar("kokoro.decoder.generator.style_upload_ms", style_upload_ms); + core::set_backend_threads(backend, n_threads); + ggml_status status = GGML_STATUS_SUCCESS; + const double compute_ms = measure_ms([&]() { + status = engine::core::compute_backend_graph(backend, graph); + }); + engine::debug::timing_log_scalar("kokoro.decoder.generator.graph.compute_ms", compute_ms); + if (status != GGML_STATUS_SUCCESS) { + throw std::runtime_error(std::string("kokoro generator graph compute failed: ") + ggml_status_to_string(status)); + } + + std::vector out(static_cast(ggml_nelements(output)), 0.0f); + const double output_read_ms = measure_ms([&]() { + ggml_backend_tensor_get(output, out.data(), 0, ggml_nbytes(output)); + }); + engine::debug::timing_log_scalar("kokoro.decoder.generator.output_read_ms", output_read_ms); + return out; + } +}; + +class GeneratorGraphRuntime { +public: + GeneratorGraphRuntime( + const KokoroWeights::GeneratorWeights & generator, + int64_t style_dim, + ggml_backend_t backend, + int n_threads, + bool use_device_backend) + : generator_(&generator), + style_dim_(style_dim), + backend_(backend), + n_threads_(std::max(1, n_threads)), + use_device_backend_(use_device_backend) { + if (generator.noise_convs.size() != generator.ups.size() || + generator.noise_res.size() != generator.ups.size() || + generator.resblocks.size() != generator.ups.size() * 3) { + throw std::runtime_error("Kokoro generator weight layout is inconsistent"); + } + stages_.reserve(generator.ups.size()); + for (size_t i = 0; i < generator.ups.size(); ++i) { + GeneratorGraphStage stage; + stage.noise_conv = &generator.noise_convs[i]; + stage.noise_res = &generator.noise_res[i]; + stage.up = &generator.ups[i]; + stage.resblocks = { + &generator.resblocks[i * 3 + 0], + &generator.resblocks[i * 3 + 1], + &generator.resblocks[i * 3 + 2], + }; + stage.reflect_pad = (i + 1 == generator.ups.size()); + stages_.push_back(stage); + } + } + + void prepare(int64_t decoder_frame_capacity, int64_t conditioning_frame_capacity) { + if (decoder_frame_capacity <= 0 || conditioning_frame_capacity <= 0) { + throw std::runtime_error("Kokoro generator graph capacity must be positive"); + } + if (decoder_frame_capacity_ == decoder_frame_capacity && + conditioning_frame_capacity_ == conditioning_frame_capacity && + session_ != nullptr) { + return; + } + session_.reset(); + decoder_frame_capacity_ = decoder_frame_capacity; + conditioning_frame_capacity_ = conditioning_frame_capacity; + session_ = std::make_unique( + *generator_, + backend_, + stages_, + decoder_frame_capacity_, + conditioning_frame_capacity_, + style_dim_, + n_threads_, + use_device_backend_); + } + + std::vector run( + const std::vector & decoder_x, + int64_t decoder_x_rows, + int64_t decoder_x_cols, + const ggml_tensor * decoder_x_tensor, + const HarmonicConditioning & conditioning, + const std::vector & style) { + return session_->run( + decoder_x_tensor ? nullptr : &decoder_x, + decoder_x_rows, + decoder_x_cols, + decoder_x_tensor, + conditioning, + style); + } + +private: + const KokoroWeights::GeneratorWeights * generator_ = nullptr; + int64_t style_dim_ = 0; + ggml_backend_t backend_ = nullptr; + int64_t decoder_frame_capacity_ = 0; + int64_t conditioning_frame_capacity_ = 0; + int n_threads_ = 1; + bool use_device_backend_ = false; + std::vector stages_; + std::unique_ptr session_; +}; + +} // namespace + +struct KokoroDecoderRuntime::Impl { + std::shared_ptr weights; + int n_threads = 1; + bool use_device_backend = false; + uint64_t rng_seed = 0; + ggml_backend_t backend = nullptr; + audio::STFTConfig inverse_stft_config; + std::unique_ptr conditioning_graph; + GeneratorGraphRuntime generator_graph; + + Impl( + std::shared_ptr weights_in, + ggml_backend_t backend_in, + int n_threads_in, + bool use_device_backend_in, + uint64_t rng_seed_in, + KokoroDecoderCapacityContract contract_in) + : weights(std::move(weights_in)), + n_threads(std::max(1, n_threads_in)), + use_device_backend(use_device_backend_in), + rng_seed(rng_seed_in), + backend(backend_in), + inverse_stft_config({ + weights->decoder.generator.gen_istft_n_fft, + weights->decoder.generator.gen_istft_hop_size, + weights->decoder.generator.gen_istft_n_fft, + true, + audio::STFTPadMode::Reflect, + audio::STFTFamily::Kokoro, + }), + generator_graph(weights->decoder.generator, weights->style_dim, backend, n_threads, use_device_backend) { + prepare(contract_in); + } + + void prepare(KokoroDecoderCapacityContract contract) { + conditioning_graph.reset(); + conditioning_graph = std::make_unique( + weights->decoder.generator, + backend, + contract.decoder_frames, + contract.conditioning_frames, + n_threads, + use_device_backend); + generator_graph.prepare( + contract.decoder_frames, + contract.conditioning_frames); + } +}; + +KokoroDecoderRuntime::KokoroDecoderRuntime( + std::shared_ptr weights, + ggml_backend_t backend, + int n_threads, + bool use_device_backend, + uint64_t rng_seed, + KokoroDecoderCapacityContract contract) + : impl_(std::make_unique( + std::move(weights), + backend, + n_threads, + use_device_backend, + rng_seed, + contract)) {} + +KokoroDecoderRuntime::~KokoroDecoderRuntime() = default; + +void KokoroDecoderRuntime::prepare(KokoroDecoderCapacityContract contract) { + impl_->prepare(contract); +} + +std::vector KokoroDecoderRuntime::decode( + const PredictorOutputs & predictor, + const std::vector & ref_s) { + if (static_cast(ref_s.size()) != 256) { + throw std::runtime_error("Kokoro decoder requires ref_s with 256 elements"); + } + + const std::vector style_decoder(ref_s.begin(), ref_s.begin() + 128); + DeterministicRng rng(impl_->rng_seed); + double conditioning_ms = 0.0; + double generator_ms = 0.0; + double istft_ms = 0.0; + const HarmonicConditioning * conditioning = nullptr; + conditioning_ms = measure_ms([&]() { + conditioning = &impl_->conditioning_graph->run(predictor.f0_curve, rng); + }); + + std::vector generator_output; + generator_ms = measure_ms([&]() { + generator_output = impl_->generator_graph.run( + predictor.decoder_x, + predictor.decoder_x_rows, + predictor.decoder_x_cols, + predictor.decoder_x_tensor, + *conditioning, + style_decoder); + }); + const int64_t bins = impl_->inverse_stft_config.n_fft / 2 + 1; + const int64_t generated_frames = static_cast(generator_output.size()) / (bins * 2); + const int64_t frames = conditioning->valid_feature_cols; + if (frames <= 0 || frames > generated_frames) { + throw std::runtime_error("Kokoro decoder valid spectrogram frame count exceeds generated output"); + } + std::vector waveform; + istft_ms = measure_ms([&]() { + std::vector complex_spec(static_cast(bins * frames * 2), 0.0f); + for (int64_t f = 0; f < bins; ++f) { + for (int64_t c = 0; c < frames; ++c) { + const float mag = std::exp(generator_output[static_cast(f * generated_frames + c)]); + const float phase = std::sin(generator_output[static_cast((f + bins) * generated_frames + c)]); + const size_t base = static_cast((f * frames + c) * 2); + complex_spec[base] = mag * std::cos(phase); + complex_spec[base + 1] = mag * std::sin(phase); + } + } + const auto & window = audio::get_cached_stft_window(impl_->inverse_stft_config); + const int64_t samples = impl_->inverse_stft_config.hop_length * (frames - 1); + waveform = audio::ISTFT().compute( + complex_spec, + window, + 1, + bins, + frames, + samples, + impl_->inverse_stft_config).values; + }); + const int64_t valid_waveform_samples = + impl_->inverse_stft_config.hop_length * std::max(conditioning->valid_feature_cols - 1, 0); + if (valid_waveform_samples < 0 || valid_waveform_samples > static_cast(waveform.size())) { + throw std::runtime_error("Kokoro decoder valid waveform length exceeds generated output"); + } + waveform.resize(static_cast(valid_waveform_samples)); + engine::debug::timing_log_scalar("kokoro.decoder.conditioning_ms", conditioning_ms); + engine::debug::timing_log_scalar("kokoro.decoder.generator_ms", generator_ms); + engine::debug::timing_log_scalar("kokoro.decoder.istft_ms", istft_ms); + + return waveform; +} + +} // namespace kokoro_ggml diff --git a/src/models/kokoro_tts/frontend.cpp b/src/models/kokoro_tts/frontend.cpp new file mode 100644 index 000000000..db2846848 --- /dev/null +++ b/src/models/kokoro_tts/frontend.cpp @@ -0,0 +1,285 @@ +#include "engine/models/kokoro_tts/frontend.h" + +#include "engine/models/kokoro_tts/g2p_en.h" + +#include +#include +#include +#include + +namespace engine::models::kokoro_tts { + +namespace { + +std::string lower_ascii(std::string value) { + std::transform(value.begin(), value.end(), value.begin(), [](unsigned char ch) { + return static_cast(std::tolower(ch)); + }); + return value; +} + +std::string trim_ascii(std::string value) { + const auto is_space = [](unsigned char ch) { return std::isspace(ch) != 0; }; + auto begin = value.begin(); + while (begin != value.end() && is_space(static_cast(*begin))) { + ++begin; + } + auto end = value.end(); + while (end != begin && is_space(static_cast(*(end - 1)))) { + --end; + } + return std::string(begin, end); +} + +std::string default_voice_id() { + return "af_heart"; +} + +std::string resolve_language_code_alias(const std::string & value) { + const std::string normalized = lower_ascii(trim_ascii(value)); + if (normalized.empty()) { + return {}; + } + if (normalized == "a" || normalized == "en" || normalized == "en-us" || normalized == "us" || normalized == "american" || normalized == "american english") { + return "a"; + } + if (normalized == "b" || normalized == "en-gb" || normalized == "gb" || normalized == "uk" || normalized == "british" || normalized == "british english") { + return "b"; + } + throw std::runtime_error("unsupported Kokoro language: " + value + " (English only: en-us/a or en-gb/b)"); +} + +std::string voice_language_code(const std::string & voice_id) { + if (voice_id.size() < 2 || (voice_id[1] != 'f' && voice_id[1] != 'm')) { + throw std::runtime_error("invalid Kokoro voice id: " + voice_id); + } + return std::string(1, voice_id[0]); +} + +std::string resolve_voice_id( + const std::optional & voice, + const KokoroAssets & assets) { + std::string voice_id = default_voice_id(); + if (voice.has_value() && voice->speaker.has_value() && voice->speaker->cached_voice_id.has_value()) { + voice_id = *voice->speaker->cached_voice_id; + } + if (assets.voices.find(voice_id) == assets.voices.end()) { + throw std::runtime_error("unknown Kokoro voice id: " + voice_id); + } + const std::string language_code = voice_language_code(voice_id); + if (language_code != "a" && language_code != "b") { + throw std::runtime_error( + "Kokoro currently supports only English voices; got voice id " + voice_id + + " with lang_code=" + language_code); + } + return voice_id; +} + +std::string resolve_language_code( + const runtime::Transcript & text, + const std::optional & voice, + const std::string & voice_id) { + std::string language_code; + if (voice.has_value() && voice->style.has_value() && voice->style->language.has_value()) { + language_code = resolve_language_code_alias(*voice->style->language); + } else if (!text.language.empty()) { + language_code = resolve_language_code_alias(text.language); + } else { + language_code = voice_language_code(voice_id); + } + const std::string expected = voice_language_code(voice_id); + if (language_code != expected) { + throw std::runtime_error( + "Kokoro voice/language mismatch: voice " + voice_id + + " requires lang_code=" + expected + + " but request resolved to " + language_code); + } + return language_code; +} + +std::string phonemize_text( + const runtime::Transcript & text, + const std::string & language_code, + const KokoroAssets & assets) { + if (text.text.empty()) { + throw std::runtime_error("Kokoro TTS requires non-empty text"); + } + if (language_code == "a" || language_code == "b") { + const auto & g2p = language_code == "b" ? assets.english_g2p_gb : assets.english_g2p_us; + if (!g2p) { + throw std::runtime_error("Kokoro English G2P assets were not prepared"); + } + return (*g2p)(text.text).first; + } + throw std::runtime_error( + "unsupported Kokoro language code: " + language_code + + " (English only: a/en-us or b/en-gb)"); +} + +struct EncodedInputIds { + std::vector ids; + size_t phoneme_count = 0; +}; + +EncodedInputIds encode_input_ids_and_count( + const std::string & phonemes, + const KokoroAssets & assets) { + EncodedInputIds encoded; + encoded.ids.reserve(phonemes.size() + 2); + encoded.ids.push_back(0); + for (size_t i = 0; i < phonemes.size();) { + const unsigned char lead = static_cast(phonemes[i]); + size_t width = 0; + if ((lead & 0x80u) == 0) { + width = 1; + } else if ((lead & 0xE0u) == 0xC0u) { + width = 2; + } else if ((lead & 0xF0u) == 0xE0u) { + width = 3; + } else if ((lead & 0xF8u) == 0xF0u) { + width = 4; + } else { + throw std::runtime_error("invalid UTF-8 lead byte in Kokoro phoneme string"); + } + if (i + width > phonemes.size()) { + throw std::runtime_error("truncated UTF-8 codepoint in Kokoro phoneme string"); + } + for (size_t j = 1; j < width; ++j) { + const unsigned char byte = static_cast(phonemes[i + j]); + if ((byte & 0xC0u) != 0x80u) { + throw std::runtime_error("invalid UTF-8 continuation byte in Kokoro phoneme string"); + } + } + const auto it = assets.vocab.find(phonemes.substr(i, width)); + if (it == assets.vocab.end()) { + throw std::runtime_error("Kokoro vocab is missing phoneme symbol: " + phonemes.substr(i, width)); + } + encoded.ids.push_back(it->second); + ++encoded.phoneme_count; + i += width; + } + encoded.ids.push_back(0); + return encoded; +} + +std::vector style_for_phoneme_count( + const KokoroVoicePack & pack, + size_t phoneme_count) { + if (pack.cols != 256) { + throw std::runtime_error("Kokoro voice pack must have 256 columns: " + pack.id); + } + if (phoneme_count == 0) { + throw std::runtime_error("Kokoro phoneme string must not be empty"); + } + if (phoneme_count > static_cast(pack.rows)) { + throw std::runtime_error( + "Kokoro phoneme count exceeds voice style rows: " + std::to_string(phoneme_count)); + } + const size_t row = phoneme_count - 1; + const size_t offset = row * static_cast(pack.cols); + std::vector style(static_cast(pack.cols)); + std::memcpy( + style.data(), + pack.values.data() + offset, + static_cast(pack.cols) * sizeof(float)); + return style; +} + +float resolve_speaking_rate(const std::optional & voice) { + if (!voice.has_value() || !voice->style.has_value() || !voice->style->speaking_rate.has_value()) { + return 1.0f; + } + const float rate = *voice->style->speaking_rate; + if (!(rate > 0.0f)) { + throw std::runtime_error("Kokoro speaking_rate must be positive"); + } + return rate; +} + +} // namespace + +KokoroFrontendSessionState resolve_kokoro_frontend_session_state( + const std::optional & text, + const std::optional & voice, + const KokoroAssets & assets) { + runtime::Transcript transcript; + if (text.has_value()) { + transcript = *text; + } + KokoroFrontendSessionState state; + state.voice_id = resolve_voice_id(voice, assets); + state.language_code = resolve_language_code(transcript, voice, state.voice_id); + const auto voice_it = assets.voices.find(state.voice_id); + if (voice_it == assets.voices.end()) { + throw std::runtime_error("unknown Kokoro voice id: " + state.voice_id); + } + state.voice_pack = &voice_it->second; + state.speaking_rate = resolve_speaking_rate(voice); + return state; +} + +void validate_kokoro_frontend_session_state( + const runtime::Transcript & text, + const std::optional & voice, + const KokoroFrontendSessionState & state, + const KokoroAssets & assets) { + const std::string resolved_voice_id = resolve_voice_id(voice, assets); + if (resolved_voice_id != state.voice_id) { + throw std::runtime_error( + "Kokoro session voice_id changed after launch: " + + state.voice_id + " -> " + resolved_voice_id); + } + const std::string resolved_language_code = resolve_language_code(text, voice, state.voice_id); + if (resolved_language_code != state.language_code) { + throw std::runtime_error( + "Kokoro session language_code changed after launch: " + + state.language_code + " -> " + resolved_language_code); + } + const auto voice_it = assets.voices.find(state.voice_id); + if (voice_it == assets.voices.end() || &voice_it->second != state.voice_pack) { + throw std::runtime_error("Kokoro session voice pack changed after launch"); + } + const float resolved_speaking_rate = resolve_speaking_rate(voice); + if (resolved_speaking_rate != state.speaking_rate) { + throw std::runtime_error("Kokoro session speaking_rate changed after launch"); + } +} + +KokoroSynthesisInput build_kokoro_synthesis_input( + const runtime::Transcript & text, + const KokoroFrontendSessionState & state, + const KokoroAssets & assets) { + if (state.voice_pack == nullptr) { + throw std::runtime_error("Kokoro frontend session voice pack was not prepared"); + } + const std::string phonemes = phonemize_text(text, state.language_code, assets); + const EncodedInputIds encoded = encode_input_ids_and_count(phonemes, assets); + if (encoded.phoneme_count > 510) { + throw std::runtime_error( + "Kokoro phoneme string exceeds 510 symbols; segmenting is not implemented in the framework path yet"); + } + KokoroSynthesisInput input; + input.voice_id = state.voice_id; + input.language_code = state.language_code; + input.phonemes = phonemes; + input.input_ids = encoded.ids; + input.style = style_for_phoneme_count(*state.voice_pack, encoded.phoneme_count); + input.speaking_rate = state.speaking_rate; + if (static_cast(input.input_ids.size()) > assets.context_length) { + throw std::runtime_error("Kokoro tokenized input exceeds model context length"); + } + return input; +} + +int64_t estimate_kokoro_request_tokens( + const runtime::SessionPreparationRequest & request, + const KokoroFrontendSessionState & state, + const KokoroAssets & assets) { + if (!request.text.has_value()) { + return 0; + } + const auto input = build_kokoro_synthesis_input(*request.text, state, assets); + return static_cast(input.input_ids.size()); +} + +} // namespace engine::models::kokoro_tts diff --git a/src/models/kokoro_tts/g2p_en.cpp b/src/models/kokoro_tts/g2p_en.cpp new file mode 100644 index 000000000..2ba430bb6 --- /dev/null +++ b/src/models/kokoro_tts/g2p_en.cpp @@ -0,0 +1,1866 @@ +#include "engine/models/kokoro_tts/g2p_en.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace kokoro_ggml::g2p_en { +namespace { + +constexpr char kPrimaryStress[] = "ˈ"; +constexpr char kSecondaryStress[] = "ˌ"; + +const std::string kDiphthongs = "AIOQWYʤʧ"; +const std::string kVowels = "AIOQWYaiuæɑɒɔəɛɜɪʊʌᵻ"; +const std::string kConsonants = "bdfhjklmnpstvwzðŋɡɹɾʃʒʤʧθ"; +const std::string kUsTaus = "AIOWYiuæɑəɛɪɹʊʌ"; + +const std::unordered_set kPunctTags = {".", ",", "-LRB-", "-RRB-", "``", "\"\"", "''", ":", "$", "#", "NFP"}; +const std::unordered_map kPunctTagPhonemes = { + {"-LRB-", "("}, {"-RRB-", ")"}, {"``", "\xE2\x80\x9C"}, {"\"\"", "\xE2\x80\x9D"}, {"''", "\xE2\x80\x9D"}, +}; +const std::unordered_map kAddSymbols = {{'.', "dot"}, {'/', "slash"}}; +const std::unordered_map kSymbols = {{'%', "percent"}, {'&', "and"}, {'+', "plus"}, {'@', "at"}}; +const std::unordered_map> kCurrencies = { + {"$", {"dollar", "cent"}}, + {"£", {"pound", "pence"}}, + {"€", {"euro", "cent"}}, +}; +const std::unordered_set kOrdinals = {"st", "nd", "rd", "th"}; +const std::unordered_set kSubtokenJunk = {'\'', ',', '-', '.', '_', '/', static_cast(0xe2)}; +const std::string kPuncts = ";:,.!?—…\"“”"; +const std::string kNonQuotePuncts = ";:,.!?—…"; + +const std::unordered_set kDeterminers = {"a", "an", "the", "this", "that", "these", "those"}; +const std::unordered_set kPronouns = {"i", "you", "he", "she", "it", "we", "they", "me", "him", "her", "us", "them"}; +const std::unordered_set kPrepositions = { + "to", "in", "on", "at", "by", "for", "from", "with", "without", "within", "into", "onto", "over", + "under", "before", "after", "through", "across", "between", "among", "versus", "vs", "of" +}; +const std::unordered_set kConjunctions = {"and", "or", "but", "nor", "so", "yet"}; +const std::unordered_set kAdverbs = {"not", "never", "always", "often", "rarely", "here", "there", "why", "how", "when"}; + +struct NumberWords { + std::string cardinal; + std::string ordinal; +}; + +template +bool contains_key(const Container & container, const Key & key) { + return container.find(key) != container.end(); +} + +bool contains_codepoint(const std::string & text, const std::string & symbol) { + return text.find(symbol) != std::string::npos; +} + +bool starts_with(const std::string & text, const std::string & prefix) { + return text.size() >= prefix.size() && text.compare(0, prefix.size(), prefix) == 0; +} + +bool ends_with(const std::string & text, const std::string & suffix) { + return text.size() >= suffix.size() && text.compare(text.size() - suffix.size(), suffix.size(), suffix) == 0; +} + +bool ends_with_any(const std::string & text, const std::initializer_list & suffixes) { + for (const char * suffix : suffixes) { + if (ends_with(text, suffix)) { + return true; + } + } + return false; +} + +bool is_ascii_alpha(char ch) { + return std::isalpha(static_cast(ch)) != 0; +} + +bool is_all_ascii_alpha(const std::string & text) { + return !text.empty() && std::all_of(text.begin(), text.end(), is_ascii_alpha); +} + +bool is_all_upper_ascii(const std::string & text) { + return !text.empty() && std::all_of(text.begin(), text.end(), [](char ch) { + return !std::isalpha(static_cast(ch)) || std::isupper(static_cast(ch)) != 0; + }); +} + +bool is_capitalized_ascii(const std::string & text) { + return !text.empty() && + std::isupper(static_cast(text.front())) != 0 && + std::all_of(text.begin() + 1, text.end(), [](char ch) { + return !std::isalpha(static_cast(ch)) || std::islower(static_cast(ch)) != 0; + }); +} + +std::string lower_ascii(std::string text) { + std::transform(text.begin(), text.end(), text.begin(), [](unsigned char ch) { + return static_cast(std::tolower(ch)); + }); + return text; +} + +std::string upper_ascii(std::string text) { + std::transform(text.begin(), text.end(), text.begin(), [](unsigned char ch) { + return static_cast(std::toupper(ch)); + }); + return text; +} + +std::string capitalize_ascii(std::string text) { + if (!text.empty()) { + text[0] = static_cast(std::toupper(static_cast(text[0]))); + for (size_t i = 1; i < text.size(); ++i) { + text[i] = static_cast(std::tolower(static_cast(text[i]))); + } + } + return text; +} + +std::vector split_tab(const std::string & line) { + std::vector fields; + size_t begin = 0; + while (true) { + const size_t tab = line.find('\t', begin); + if (tab == std::string::npos) { + fields.push_back(line.substr(begin)); + return fields; + } + fields.push_back(line.substr(begin, tab - begin)); + begin = tab + 1; + } +} + +std::optional normalize_stress(const std::optional & stress, const std::string & phonemes) { + if (!stress.has_value()) { + return stress; + } + if (*stress == 0.0f && contains_codepoint(phonemes, kPrimaryStress)) { + return -1.0f; + } + return stress; +} + +int stress_weight(const std::string & phonemes) { + int sum = 0; + for (char ch : phonemes) { + sum += kDiphthongs.find(ch) != std::string::npos ? 2 : 1; + } + return sum; +} + +bool is_digit_text(const std::string & text) { + return !text.empty() && std::all_of(text.begin(), text.end(), [](char ch) { + return std::isdigit(static_cast(ch)) != 0; + }); +} + +bool is_number_token_impl(const std::string & word, bool is_head) { + if (std::none_of(word.begin(), word.end(), [](char ch) { return std::isdigit(static_cast(ch)) != 0; })) { + return false; + } + std::string trimmed = word; + for (const char * suffix : {"ing", "'d", "ed", "'s", "st", "nd", "rd", "th", "s"}) { + if (ends_with(trimmed, suffix)) { + trimmed = trimmed.substr(0, trimmed.size() - std::char_traits::length(suffix)); + break; + } + } + for (size_t i = 0; i < trimmed.size(); ++i) { + const char ch = trimmed[i]; + if (std::isdigit(static_cast(ch)) != 0 || ch == ',' || ch == '.') { + continue; + } + if (is_head && i == 0 && ch == '-') { + continue; + } + return false; + } + return !trimmed.empty(); +} + +bool contains_whitespace(const std::string & text) { + return std::any_of(text.begin(), text.end(), [](unsigned char ch) { + return std::isspace(ch) != 0; + }); +} + +int english_g2p_char_class(char ch) { + if (std::isalpha(static_cast(ch)) != 0) { + return 0; + } + if (is_digit_text(std::string(1, ch))) { + return 1; + } + return 2; +} + +size_t count_words(const std::string & text) { + size_t count = 0; + bool in_word = false; + for (unsigned char ch : text) { + if (std::isspace(ch) != 0) { + in_word = false; + continue; + } + if (!in_word) { + ++count; + in_word = true; + } + } + return count; +} + +bool all_subtoken_junk(const std::string & text) { + for (char ch : text) { + if (ch == '\'' || ch == ',' || ch == '-' || ch == '.' || ch == '_' || ch == '/') { + continue; + } + return false; + } + return !text.empty(); +} + +size_t utf8_codepoint_width(const std::string & text, size_t offset) { + const unsigned char lead = static_cast(text[offset]); + size_t width = 0; + if ((lead & 0x80u) == 0) { + width = 1; + } else if ((lead & 0xE0u) == 0xC0u) { + width = 2; + } else if ((lead & 0xF0u) == 0xE0u) { + width = 3; + } else if ((lead & 0xF8u) == 0xF0u) { + width = 4; + } else { + throw std::runtime_error("invalid UTF-8 lead byte in English G2P phoneme string"); + } + if (offset + width > text.size()) { + throw std::runtime_error("truncated UTF-8 codepoint in English G2P phoneme string"); + } + for (size_t j = 1; j < width; ++j) { + const unsigned char byte = static_cast(text[offset + j]); + if ((byte & 0xC0u) != 0x80u) { + throw std::runtime_error("invalid UTF-8 continuation byte in English G2P phoneme string"); + } + } + return width; +} + +bool utf8_inventory_contains( + const std::string & inventory, + const char * symbol, + size_t symbol_width) { + for (size_t i = 0; i < inventory.size();) { + const size_t width = utf8_codepoint_width(inventory, i); + if (width == symbol_width && std::memcmp(inventory.data() + i, symbol, symbol_width) == 0) { + return true; + } + i += width; + } + return false; +} + +std::string join_phoneme_pieces(const std::vector & pieces) { + std::string out; + for (size_t i = 0; i < pieces.size(); ++i) { + if (i > 0) { + out.push_back(' '); + } + out += pieces[i]; + } + return out; +} + +std::optional compound_fallback_impl( + const Lexicon & lexicon, + const std::string & word, + std::unordered_map> & memo) { + if (const auto found = memo.find(word); found != memo.end()) { + return found->second; + } + if (word.size() < 4 || !is_all_ascii_alpha(word)) { + memo[word] = std::nullopt; + return std::nullopt; + } + + const TokenContext ctx{}; + if (auto direct = lexicon.lookup(word, std::nullopt, std::nullopt, ctx)) { + memo[word] = direct; + return direct; + } + + for (size_t split = 2; split + 2 <= word.size(); ++split) { + const std::string left = word.substr(0, split); + const std::string right = word.substr(split); + auto left_result = lexicon.lookup(left, std::nullopt, std::nullopt, ctx); + if (!left_result.has_value()) { + continue; + } + auto right_result = compound_fallback_impl(lexicon, right, memo); + if (!right_result.has_value()) { + right_result = lexicon.lookup(right, std::nullopt, std::nullopt, ctx); + } + if (!right_result.has_value()) { + continue; + } + const int rating = std::min(left_result->rating, right_result->rating); + memo[word] = LexiconResult{left_result->phonemes + " " + right_result->phonemes, rating}; + return memo[word]; + } + + memo[word] = std::nullopt; + return std::nullopt; +} + +std::optional compound_fallback( + const Lexicon & lexicon, + const std::string & text, + const std::optional & stress) { + const std::string word = lower_ascii(text); + std::unordered_map> memo; + auto result = compound_fallback_impl(lexicon, word, memo); + if (!result.has_value()) { + return std::nullopt; + } + result->phonemes = Lexicon::apply_stress(result->phonemes, stress); + return result; +} + +std::optional spell_out_fallback( + const Lexicon & lexicon, + const std::string & text, + const std::optional & stress) { + std::vector pieces; + for (char ch : text) { + if (std::isspace(static_cast(ch)) != 0) { + continue; + } + MToken token; + token.text = std::string(1, is_ascii_alpha(ch) ? static_cast(std::toupper(static_cast(ch))) : ch); + token.tag = std::isdigit(static_cast(ch)) != 0 ? "CD" : "NNP"; + token.meta.is_head = true; + if (auto result = lexicon(token, TokenContext{})) { + pieces.push_back(result->phonemes); + } + } + if (pieces.empty()) { + return std::nullopt; + } + return LexiconResult{Lexicon::apply_stress(join_phoneme_pieces(pieces), stress), 1}; +} + +LexiconResult resolve_with_safe_fallback( + const Lexicon & lexicon, + const MToken & token) { + if (auto compound = compound_fallback(lexicon, token.text, token.meta.stress)) { + return *compound; + } + if (auto spelled = spell_out_fallback(lexicon, token.text, token.meta.stress)) { + return *spelled; + } + return LexiconResult{Lexicon::apply_stress("ə", token.meta.stress), 1}; +} + +std::string trim_copy(const std::string & text) { + size_t begin = 0; + while (begin < text.size() && std::isspace(static_cast(text[begin])) != 0) { + ++begin; + } + size_t end = text.size(); + while (end > begin && std::isspace(static_cast(text[end - 1])) != 0) { + --end; + } + return text.substr(begin, end - begin); +} + +void replace_all(std::string & text, const std::string & from, const std::string & to) { + size_t pos = 0; + while ((pos = text.find(from, pos)) != std::string::npos) { + text.replace(pos, from.size(), to); + pos += to.size(); + } +} + +std::vector split_words(const std::string & text) { + std::vector words; + std::string current; + for (char ch : text) { + if (std::isspace(static_cast(ch)) != 0) { + if (!current.empty()) { + words.push_back(current); + current.clear(); + } + continue; + } + current.push_back(ch); + } + if (!current.empty()) { + words.push_back(current); + } + return words; +} + +std::vector split_non_alpha(const std::string & text) { + std::vector out; + std::string current; + for (char ch : text) { + if (std::isalpha(static_cast(ch)) != 0) { + current.push_back(ch); + continue; + } + if (!current.empty()) { + out.push_back(current); + current.clear(); + } + } + if (!current.empty()) { + out.push_back(current); + } + return out; +} + +std::filesystem::path resolve_default_lexicon_dir() { + const std::array candidates = { + std::filesystem::path("models/kokoro-82m-v1_0-ggml/misaki_en"), + std::filesystem::path("models/misaki_en"), + std::filesystem::current_path() / "models/kokoro-82m-v1_0-ggml/misaki_en", + std::filesystem::current_path() / "models/misaki_en", + }; + for (const auto & candidate : candidates) { + if (std::filesystem::exists(candidate / "us_gold.tsv") || std::filesystem::exists(candidate / "gb_gold.tsv")) { + return candidate; + } + } + throw std::runtime_error("failed to locate Kokoro misaki_en lexicon assets"); +} + +std::vector subtokenize(const std::string & word) { + std::vector tokens; + size_t i = 0; + while (i < word.size()) { + const unsigned char ch = static_cast(word[i]); + if (word[i] == '\'') { + size_t j = i + 1; + while (j < word.size() && word[j] == '\'') { + ++j; + } + tokens.push_back(word.substr(i, j - i)); + i = j; + continue; + } + if (std::isdigit(ch) != 0 || ((word[i] == '-' || word[i] == '+') && i + 1 < word.size() && std::isdigit(static_cast(word[i + 1])) != 0)) { + size_t j = i + 1; + while (j < word.size()) { + const unsigned char cj = static_cast(word[j]); + if (std::isdigit(cj) != 0) { + ++j; + continue; + } + if ((word[j] == ',' || word[j] == '.') && + j + 1 < word.size() && + std::isdigit(static_cast(word[j + 1])) != 0) { + ++j; + continue; + } + break; + } + tokens.push_back(word.substr(i, j - i)); + i = j; + continue; + } + if (std::isalpha(ch) != 0) { + size_t j = i + 1; + while (j < word.size()) { + const unsigned char pj = static_cast(word[j]); + if (std::isalpha(pj) != 0 || word[j] == '\'') { + if (j + 1 < word.size() && std::islower(pj) != 0 && std::isupper(static_cast(word[j + 1])) != 0) { + ++j; + break; + } + ++j; + continue; + } + break; + } + tokens.push_back(word.substr(i, j - i)); + i = j; + continue; + } + if (word[i] == '-' || word[i] == '_') { + size_t j = i + 1; + while (j < word.size() && (word[j] == '-' || word[j] == '_')) { + ++j; + } + tokens.push_back(word.substr(i, j - i)); + i = j; + continue; + } + tokens.push_back(word.substr(i, 1)); + ++i; + } + return tokens; +} + +std::string guess_tag(const std::string & token) { + if (token.empty()) { + return "NN"; + } + if (token == "$" || token == "£" || token == "€") { + return "$"; + } + if (token == "(") return "-LRB-"; + if (token == ")") return "-RRB-"; + if (token == "\"" || token == "“") return "``"; + if (token == "”") return "''"; + if (token == ":" || token == ";" || token == "-" || token == "–" || token == "—") return ":"; + if (token == "," || token == "." || token == "!" || token == "?" || token == "…") return std::string(1, token[0]); + + const std::string lowered = lower_ascii(token); + if (contains_key(kDeterminers, lowered)) return "DT"; + if (contains_key(kPronouns, lowered)) return "PRP"; + if (contains_key(kPrepositions, lowered)) return lowered == "to" ? "TO" : "IN"; + if (contains_key(kConjunctions, lowered)) return "CC"; + if (contains_key(kAdverbs, lowered) || ends_with(lowered, "ly")) return "RB"; + if (is_number_token_impl(token, true)) return "CD"; + if (is_all_upper_ascii(token) || is_capitalized_ascii(token)) return "NNP"; + return "NN"; +} + +std::string join_words(const std::vector & words, const std::string & delimiter) { + std::ostringstream out; + for (size_t i = 0; i < words.size(); ++i) { + if (i > 0) { + out << delimiter; + } + out << words[i]; + } + return out.str(); +} + +std::string two_digit_cardinal(int value) { + static const std::array below_20 = { + "zero", "one", "two", "three", "four", "five", "six", "seven", "eight", "nine", + "ten", "eleven", "twelve", "thirteen", "fourteen", "fifteen", "sixteen", "seventeen", "eighteen", "nineteen" + }; + static const std::array tens = { + "", "", "twenty", "thirty", "forty", "fifty", "sixty", "seventy", "eighty", "ninety" + }; + if (value < 20) { + return below_20[static_cast(value)]; + } + const int ten = value / 10; + const int one = value % 10; + if (one == 0) { + return tens[static_cast(ten)]; + } + return std::string(tens[static_cast(ten)]) + "-" + below_20[static_cast(one)]; +} + +std::string integer_to_cardinal(int64_t value) { + if (value == 0) { + return "zero"; + } + if (value < 0) { + return "minus " + integer_to_cardinal(-value); + } + static const std::array scales = { + "", "thousand", "million", "billion", "trillion", "quadrillion", "quintillion" + }; + std::vector chunks; + size_t scale = 0; + while (value > 0) { + const int part = static_cast(value % 1000); + if (part != 0) { + std::string chunk; + const int hundreds = part / 100; + const int rest = part % 100; + if (hundreds > 0) { + chunk += two_digit_cardinal(hundreds) + " hundred"; + if (rest != 0) { + chunk += " "; + } + } + if (rest != 0) { + chunk += two_digit_cardinal(rest); + } + if (scale > 0) { + chunk += " "; + chunk += scales[scale]; + } + chunks.push_back(chunk); + } + value /= 1000; + ++scale; + } + std::reverse(chunks.begin(), chunks.end()); + return join_words(chunks, " "); +} + +std::string integer_to_ordinal(int64_t value) { + static const std::unordered_map replacements = { + {"one", "first"}, + {"two", "second"}, + {"three", "third"}, + {"five", "fifth"}, + {"eight", "eighth"}, + {"nine", "ninth"}, + {"twelve", "twelfth"}, + {"twenty", "twentieth"}, + {"thirty", "thirtieth"}, + {"forty", "fortieth"}, + {"fifty", "fiftieth"}, + {"sixty", "sixtieth"}, + {"seventy", "seventieth"}, + {"eighty", "eightieth"}, + {"ninety", "ninetieth"}, + {"hundred", "hundredth"}, + {"thousand", "thousandth"}, + {"million", "millionth"}, + {"billion", "billionth"}, + {"trillion", "trillionth"}, + }; + std::string cardinal = integer_to_cardinal(value); + size_t pos = cardinal.find_last_of(" -"); + const size_t begin = pos == std::string::npos ? 0 : pos + 1; + std::string tail = cardinal.substr(begin); + const auto it = replacements.find(tail); + if (it != replacements.end()) { + return cardinal.substr(0, begin) + it->second; + } + if (ends_with(tail, "y")) { + return cardinal.substr(0, cardinal.size() - 1) + "ieth"; + } + return cardinal + "th"; +} + +std::string integer_to_year_words(int64_t value) { + if (value < 1000 || value > 9999) { + return integer_to_cardinal(value); + } + if (value >= 2000 && value <= 2009) { + return "two thousand " + integer_to_cardinal(value % 1000); + } + const int first = static_cast(value / 100); + const int second = static_cast(value % 100); + if (second == 0) { + return integer_to_cardinal(first) + " hundred"; + } + if (second < 10) { + return integer_to_cardinal(first) + " oh " + integer_to_cardinal(second); + } + return integer_to_cardinal(first) + " " + integer_to_cardinal(second); +} + +std::vector spell_digit_string(const std::string & digits) { + static const std::array digit_words = { + "zero", "one", "two", "three", "four", "five", "six", "seven", "eight", "nine" + }; + std::vector out; + for (char ch : digits) { + if (std::isdigit(static_cast(ch)) != 0) { + out.emplace_back(digit_words[static_cast(ch - '0')]); + } + } + return out; +} + +} // namespace + +Lexicon::Lexicon(bool british) + : Lexicon( + resolve_default_lexicon_dir() / (british ? "gb_gold.tsv" : "us_gold.tsv"), + resolve_default_lexicon_dir() / (british ? "gb_silver.tsv" : "us_silver.tsv"), + british) {} + +Lexicon::Lexicon(std::filesystem::path gold_tsv, std::filesystem::path silver_tsv, bool british) + : british_(british) { + load_tsv(gold_tsv, golds_); + load_tsv(silver_tsv, silvers_); +} + +void Lexicon::load_tsv( + const std::filesystem::path & path, + std::unordered_map & target) { + std::ifstream input(path); + if (!input) { + throw std::runtime_error("failed to open misaki lexicon TSV: " + path.string()); + } + std::string line; + while (std::getline(input, line)) { + if (line.empty()) { + continue; + } + const auto fields = split_tab(line); + if (fields.size() != 3) { + throw std::runtime_error("invalid misaki lexicon TSV line: " + line); + } + auto & entry = target[fields[0]]; + if (fields[1] == "*") { + entry.has_default = true; + entry.default_phonemes = fields[2]; + } else { + entry.by_tag[fields[1]] = fields[2] == "~" ? std::optional() : std::optional(fields[2]); + } + if (fields[0].size() >= 2) { + if (fields[0] == lower_ascii(fields[0])) { + const auto capitalized = capitalize_ascii(fields[0]); + if (capitalized != fields[0] && target.find(capitalized) == target.end()) { + target[capitalized] = entry; + } + } else if (fields[0] == capitalize_ascii(fields[0])) { + const auto lowered = lower_ascii(fields[0]); + if (target.find(lowered) == target.end()) { + target[lowered] = entry; + } + } + } + } +} + +std::optional Lexicon::parent_tag(const std::optional & tag) { + if (!tag.has_value()) { + return std::nullopt; + } + if (tag->rfind("VB", 0) == 0) return std::string("VERB"); + if (tag->rfind("NN", 0) == 0) return std::string("NOUN"); + if (tag->rfind("ADV", 0) == 0 || tag->rfind("RB", 0) == 0) return std::string("ADV"); + if (tag->rfind("ADJ", 0) == 0 || tag->rfind("JJ", 0) == 0) return std::string("ADJ"); + return tag; +} + +std::string Lexicon::apply_stress(const std::string & phonemes, const std::optional & stress) { + if (phonemes.empty() || !stress.has_value()) { + return phonemes; + } + std::string out = phonemes; + if (*stress < -1.0f) { + while (true) { + size_t pos = out.find(kPrimaryStress); + if (pos == std::string::npos) break; + out.erase(pos, std::char_traits::length(kPrimaryStress)); + } + while (true) { + size_t pos = out.find(kSecondaryStress); + if (pos == std::string::npos) break; + out.erase(pos, std::char_traits::length(kSecondaryStress)); + } + return out; + } + if (*stress == -1.0f || ((*stress == 0.0f || *stress == -0.5f) && contains_codepoint(out, kPrimaryStress))) { + replace_all(out, kSecondaryStress, ""); + replace_all(out, kPrimaryStress, kSecondaryStress); + return out; + } + if ((*stress == 0.0f || *stress == 0.5f || *stress == 1.0f) && + !contains_codepoint(out, kPrimaryStress) && !contains_codepoint(out, kSecondaryStress)) { + for (char ch : out) { + if (kVowels.find(ch) != std::string::npos) { + return std::string(kSecondaryStress) + out; + } + } + return out; + } + if (*stress >= 1.0f && !contains_codepoint(out, kPrimaryStress) && contains_codepoint(out, kSecondaryStress)) { + replace_all(out, kSecondaryStress, kPrimaryStress); + return out; + } + if (*stress > 1.0f && !contains_codepoint(out, kPrimaryStress) && !contains_codepoint(out, kSecondaryStress)) { + for (char ch : out) { + if (kVowels.find(ch) != std::string::npos) { + return std::string(kPrimaryStress) + out; + } + } + } + return out; +} + +std::optional Lexicon::get_nnp(const std::string & word) const { + std::string phonemes; + for (char ch : word) { + if (!is_ascii_alpha(ch)) { + continue; + } + const std::string key(1, static_cast(std::toupper(static_cast(ch)))); + const auto it = golds_.find(key); + if (it == golds_.end() || !it->second.has_default) { + return std::nullopt; + } + phonemes += it->second.default_phonemes; + } + if (phonemes.empty()) { + return std::nullopt; + } + phonemes = apply_stress(phonemes, 0.0f); + const size_t split = phonemes.rfind(kSecondaryStress); + if (split != std::string::npos) { + phonemes.replace(split, std::char_traits::length(kSecondaryStress), kPrimaryStress); + } + return LexiconResult{phonemes, 3}; +} + +bool Lexicon::is_known(const std::string & word) const { + if (contains_key(golds_, word) || contains_key(silvers_, word) || (word.size() == 1 && contains_key(kSymbols, word[0]))) { + return true; + } + if (!is_all_ascii_alpha(word)) { + return false; + } + if (word.size() == 1) { + return true; + } + if (is_all_upper_ascii(word) && contains_key(golds_, lower_ascii(word))) { + return true; + } + return word.size() > 1 && std::all_of(word.begin() + 1, word.end(), [](char ch) { + return std::isupper(static_cast(ch)) != 0; + }); +} + +std::optional Lexicon::lookup_raw( + const std::string & word, + const std::optional & tag, + const std::optional & stress) const { + std::optional selected; + int rating = 4; + auto it = golds_.find(word); + if (it != golds_.end()) { + const auto & entry = it->second; + if (!entry.by_tag.empty()) { + std::optional query = tag; + if (!query.has_value()) { + query = "DEFAULT"; + } else if (!contains_key(entry.by_tag, *query)) { + query = parent_tag(query); + } + auto jt = entry.by_tag.find(query.value_or("DEFAULT")); + if (jt == entry.by_tag.end()) { + jt = entry.by_tag.find("DEFAULT"); + } + if (jt != entry.by_tag.end()) { + if (!jt->second.has_value()) { + return std::nullopt; + } + selected = *jt->second; + } + } else if (entry.has_default) { + selected = entry.default_phonemes; + } + } + if (!selected.has_value()) { + auto st = silvers_.find(word); + if (st != silvers_.end() && st->second.has_default) { + selected = st->second.default_phonemes; + rating = 3; + } + } + if (!selected.has_value() && is_all_upper_ascii(word)) { + return get_nnp(word); + } + if (!selected.has_value()) { + return std::nullopt; + } + return LexiconResult{apply_stress(*selected, normalize_stress(stress, *selected)), rating}; +} + +std::optional Lexicon::s_suffix(const std::string & stem) const { + if (stem.empty()) return std::nullopt; + if (ends_with_any(stem, {"p", "t", "k", "f", "θ"})) return stem + "s"; + if (ends_with_any(stem, {"s", "z", "ʃ", "ʒ", "ʧ", "ʤ"})) { + return stem + (british_ ? "ɪz" : "ᵻz"); + } + return stem + "z"; +} + +std::optional Lexicon::ed_suffix(const std::string & stem) const { + if (stem.empty()) return std::nullopt; + if (ends_with_any(stem, {"p", "k", "f", "θ", "ʃ", "s", "ʧ"})) return stem + "t"; + if (ends_with(stem, "d")) return stem + (british_ ? "ɪd" : "ᵻd"); + if (!ends_with(stem, "t")) return stem + "d"; + if (british_ || stem.size() < 2) return stem + "ɪd"; + if (kUsTaus.find(stem[stem.size() - 2]) != std::string::npos) return stem.substr(0, stem.size() - 1) + "ɾᵻd"; + return stem + "ᵻd"; +} + +std::optional Lexicon::ing_suffix(const std::string & stem) const { + if (stem.empty()) return std::nullopt; + if (british_) { + if (ends_with_any(stem, {"ə", "ː"})) return std::nullopt; + } else if (stem.size() > 1 && stem.back() == 't' && kUsTaus.find(stem[stem.size() - 2]) != std::string::npos) { + return stem.substr(0, stem.size() - 1) + "ɾɪŋ"; + } + return stem + "ɪŋ"; +} + +std::optional Lexicon::stem_s( + const std::string & word, + const std::optional & tag, + const std::optional & stress, + const TokenContext & ctx) const { + if (word.size() < 3 || word.back() != 's') return std::nullopt; + std::string stem; + if (!(word.size() >= 2 && word.substr(word.size() - 2) == "ss") && is_known(word.substr(0, word.size() - 1))) { + stem = word.substr(0, word.size() - 1); + } else if (((word.size() >= 2 && word.substr(word.size() - 2) == "'s") || + (word.size() > 4 && word.size() >= 2 && word.substr(word.size() - 2) == "es" && + word.substr(word.size() - 3) != "ies")) && + is_known(word.substr(0, word.size() - 2))) { + stem = word.substr(0, word.size() - 2); + } else if (word.size() > 4 && word.substr(word.size() - 3) == "ies" && is_known(word.substr(0, word.size() - 3) + "y")) { + stem = word.substr(0, word.size() - 3) + "y"; + } else { + return std::nullopt; + } + auto base = lookup(stem, tag, stress, ctx); + if (!base.has_value()) return std::nullopt; + auto phonemes = s_suffix(base->phonemes); + if (!phonemes.has_value()) return std::nullopt; + return LexiconResult{*phonemes, base->rating}; +} + +std::optional Lexicon::stem_ed( + const std::string & word, + const std::optional & tag, + const std::optional & stress, + const TokenContext & ctx) const { + if (word.size() < 4 || word.back() != 'd') return std::nullopt; + std::string stem; + if (!(word.size() >= 2 && word.substr(word.size() - 2) == "dd") && is_known(word.substr(0, word.size() - 1))) { + stem = word.substr(0, word.size() - 1); + } else if (word.size() > 4 && word.substr(word.size() - 2) == "ed" && + word.substr(word.size() - 3) != "eed" && is_known(word.substr(0, word.size() - 2))) { + stem = word.substr(0, word.size() - 2); + } else { + return std::nullopt; + } + auto base = lookup(stem, tag, stress, ctx); + if (!base.has_value()) return std::nullopt; + auto phonemes = ed_suffix(base->phonemes); + if (!phonemes.has_value()) return std::nullopt; + return LexiconResult{*phonemes, base->rating}; +} + +std::optional Lexicon::stem_ing( + const std::string & word, + const std::optional & tag, + const std::optional & stress, + const TokenContext & ctx) const { + if (word.size() < 5 || word.substr(word.size() - 3) != "ing") return std::nullopt; + std::string stem; + if (word.size() > 5 && is_known(word.substr(0, word.size() - 3))) { + stem = word.substr(0, word.size() - 3); + } else if (is_known(word.substr(0, word.size() - 3) + "e")) { + stem = word.substr(0, word.size() - 3) + "e"; + } else if (word.size() > 5) { + const std::string doubled = word.substr(word.size() - 6); + if ((doubled.size() >= 6 && + ((doubled[0] == doubled[1] && std::string("bcdgklmnprstvxz").find(doubled[0]) != std::string::npos) || + doubled.rfind("cking") != std::string::npos)) && + is_known(word.substr(0, word.size() - 4))) { + stem = word.substr(0, word.size() - 4); + } else { + return std::nullopt; + } + } else { + return std::nullopt; + } + auto base = lookup(stem, tag, stress.has_value() ? stress : std::optional(0.5f), ctx); + if (!base.has_value()) return std::nullopt; + auto phonemes = ing_suffix(base->phonemes); + if (!phonemes.has_value()) return std::nullopt; + return LexiconResult{*phonemes, base->rating}; +} + +std::optional Lexicon::get_special_case( + const std::string & word, + const std::optional & tag, + const std::optional & stress, + const TokenContext & ctx) const { + (void) stress; + if (tag.has_value() && *tag == "ADD" && word.size() == 1 && contains_key(kAddSymbols, word[0])) { + return lookup(kAddSymbols.at(word[0]), std::nullopt, -0.5f, ctx); + } + if (word.size() == 1 && contains_key(kSymbols, word[0])) { + return lookup(kSymbols.at(word[0]), std::nullopt, std::nullopt, ctx); + } + if ((word == "a" || word == "A")) { + if (tag.has_value() && *tag == "DT") return LexiconResult{"ɐ", 4}; + return LexiconResult{"ˈA", 4}; + } + if (word == "am" || word == "Am" || word == "AM") { + if (tag.has_value() && tag->rfind("NN", 0) == 0) { + return get_nnp(word); + } + if (!ctx.future_vowel.has_value() || word != "am" || (stress.has_value() && *stress > 0.0f)) { + return lookup_raw("am", std::nullopt, stress); + } + return LexiconResult{"ɐm", 4}; + } + if (word == "an" || word == "An" || word == "AN") { + if (word == "AN" && tag.has_value() && tag->rfind("NN", 0) == 0) { + return get_nnp(word); + } + return LexiconResult{"ɐn", 4}; + } + if (word == "I" && tag.has_value() && *tag == "PRP") { + return LexiconResult{std::string(kSecondaryStress) + "I", 4}; + } + if ((word == "by" || word == "By" || word == "BY") && parent_tag(tag) == std::optional("ADV")) { + return LexiconResult{"bˈI", 4}; + } + if ((word == "to" || word == "To") || (word == "TO" && tag.has_value() && (*tag == "TO" || *tag == "IN"))) { + if (!ctx.future_vowel.has_value()) return lookup_raw("to", std::nullopt, stress); + return LexiconResult{*ctx.future_vowel ? "tʊ" : "tə", 4}; + } + if ((word == "in" || word == "In") || (word == "IN" && (!tag.has_value() || *tag != "NNP"))) { + return LexiconResult{(ctx.future_vowel.has_value() && tag == std::optional("IN")) ? "ɪn" : std::string(kPrimaryStress) + "ɪn", 4}; + } + if ((word == "the" || word == "The") || (word == "THE" && tag.has_value() && *tag == "DT")) { + return LexiconResult{ctx.future_vowel == true ? "ði" : "ðə", 4}; + } + if (tag == std::optional("IN")) { + const std::string lowered = lower_ascii(word); + if (lowered == "vs" || lowered == "vs.") { + return lookup("versus", std::nullopt, std::nullopt, ctx); + } + } + if (word == "used" || word == "Used" || word == "USED") { + if (tag == std::optional("VBD") && ctx.future_to) { + return lookup_raw("used", std::optional("VBD"), stress); + } + return lookup_raw("used", std::nullopt, stress); + } + if (word.find('.') != std::string::npos) { + std::string stripped = word; + stripped.erase(std::remove(stripped.begin(), stripped.end(), '.'), stripped.end()); + if (!stripped.empty() && is_all_ascii_alpha(stripped)) { + const auto parts = split_non_alpha(word); + size_t longest = 0; + for (const auto & part : parts) longest = std::max(longest, part.size()); + if (longest < 3) { + return get_nnp(word); + } + } + } + return std::nullopt; +} + +std::optional Lexicon::lookup( + const std::string & word, + const std::optional & tag, + const std::optional & stress, + const TokenContext & ctx) const { + if (const auto special = get_special_case(word, tag, stress, ctx)) { + return special; + } + std::string lookup_word = word; + const std::string lowered = lower_ascii(word); + if (word.size() > 1 && + std::all_of(word.begin(), word.end(), [](unsigned char ch) { return std::isalpha(ch) != 0 || ch == '\''; }) && + word != lowered && + (!tag.has_value() || *tag != "NNP" || word.size() > 7) && + !contains_key(golds_, word) && !contains_key(silvers_, word) && + (is_all_upper_ascii(word) || is_capitalized_ascii(word)) && + (contains_key(golds_, lowered) || contains_key(silvers_, lowered))) { + lookup_word = lowered; + } + if (is_known(lookup_word)) { + if (const auto direct = lookup_raw(lookup_word, tag, stress)) { + return direct; + } + } + if (lookup_word.size() >= 2 && lookup_word.substr(lookup_word.size() - 2) == "s'") { + if (const auto direct = lookup_raw(lookup_word.substr(0, lookup_word.size() - 2) + "'s", tag, stress)) { + return direct; + } + } + if (!lookup_word.empty() && lookup_word.back() == '\'') { + if (const auto direct = lookup_raw(lookup_word.substr(0, lookup_word.size() - 1), tag, stress)) { + return direct; + } + } + if (const auto stem = stem_s(lookup_word, tag, stress, ctx)) return stem; + if (const auto stem = stem_ed(lookup_word, tag, stress, ctx)) return stem; + if (const auto stem = stem_ing(lookup_word, tag, stress, ctx)) return stem; + return std::nullopt; +} + +std::optional Lexicon::get_word( + const std::string & word, + const std::optional & tag, + const std::optional & stress, + const TokenContext & ctx) const { + std::string maybe_lower = word; + const std::string lowered = lower_ascii(word); + if (word.size() > 1 && word.find_first_not_of("ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz'") == std::string::npos && + word != lowered && + (!tag.has_value() || *tag != "NNP" || word.size() > 7) && + !contains_key(golds_, word) && !contains_key(silvers_, word) && + (is_all_upper_ascii(word) || is_capitalized_ascii(word)) && + (contains_key(golds_, lowered) || contains_key(silvers_, lowered) || + stem_s(lowered, tag, stress, ctx).has_value() || + stem_ed(lowered, tag, stress, ctx).has_value() || + stem_ing(lowered, tag, stress, ctx).has_value())) { + maybe_lower = lowered; + } + if (is_known(maybe_lower)) { + return lookup(maybe_lower, tag, stress, ctx); + } + if (ends_with(maybe_lower, "s'") && is_known(maybe_lower.substr(0, maybe_lower.size() - 2) + "'s")) { + return lookup(maybe_lower.substr(0, maybe_lower.size() - 2) + "'s", tag, stress, ctx); + } + if (ends_with(maybe_lower, "'") && is_known(maybe_lower.substr(0, maybe_lower.size() - 1))) { + return lookup(maybe_lower.substr(0, maybe_lower.size() - 1), tag, stress, ctx); + } + if (const auto stem = stem_s(maybe_lower, tag, stress, ctx)) return stem; + if (const auto stem = stem_ed(maybe_lower, tag, stress, ctx)) return stem; + if (const auto stem = stem_ing(maybe_lower, tag, stress.has_value() ? stress : std::optional(0.5f), ctx)) return stem; + return std::nullopt; +} + +bool Lexicon::is_currency_shape(const std::string & word) { + const size_t dot_count = static_cast(std::count(word.begin(), word.end(), '.')); + if (dot_count == 0) { + return true; + } + if (dot_count > 1) { + return false; + } + const size_t dot = word.find('.'); + const std::string cents = word.substr(dot + 1); + if (cents.size() < 3) { + return true; + } + return std::all_of(cents.begin(), cents.end(), [](char ch) { return ch == '0'; }); +} + +bool Lexicon::is_number_token(const std::string & word, bool is_head) { + return is_number_token_impl(word, is_head); +} + +std::optional Lexicon::get_number( + const std::string & word, + const std::optional & currency, + bool is_head) const { + std::string suffix; + std::string raw = word; + size_t suffix_begin = raw.size(); + while (suffix_begin > 0 && std::isalpha(static_cast(raw[suffix_begin - 1])) != 0) { + --suffix_begin; + } + suffix = raw.substr(suffix_begin); + raw = raw.substr(0, suffix_begin); + + std::vector> pieces; + auto append_lookup = [&](const std::string & token, std::optional stress = std::nullopt) { + auto found = lookup(token, std::nullopt, stress, TokenContext{}); + if (!found.has_value()) { + throw std::runtime_error("english g2p number expansion missing lexicon entry for: " + token); + } + pieces.emplace_back(found->phonemes, found->rating); + }; + auto append_phrase = [&](const std::string & phrase) { + for (const std::string & chunk : split_words(phrase)) { + const auto parts = split_non_alpha(chunk); + if (parts.empty()) { + continue; + } + for (const std::string & part : parts) { + if (part.empty()) continue; + append_lookup(part, part == "point" ? std::optional(-2.0f) : std::nullopt); + } + } + }; + + if (!raw.empty() && raw.front() == '-') { + append_lookup("minus"); + raw.erase(raw.begin()); + } + + const std::string no_commas = [&]() { + std::string s = raw; + s.erase(std::remove(s.begin(), s.end(), ','), s.end()); + return s; + }(); + + if (currency.has_value() && contains_key(kCurrencies, *currency) && is_currency_shape(no_commas)) { + const auto [major_name, minor_name] = kCurrencies.at(*currency); + const size_t dot = no_commas.find('.'); + const int64_t major = dot == std::string::npos || no_commas.substr(0, dot).empty() ? 0 : std::stoll(no_commas.substr(0, dot)); + std::string minor_text = dot == std::string::npos ? "" : no_commas.substr(dot + 1); + if (minor_text.size() > 2) { + minor_text = minor_text.substr(0, 2); + } + while (!minor_text.empty() && minor_text.size() < 2) { + minor_text.push_back('0'); + } + const int64_t minor = minor_text.empty() ? 0 : std::stoll(minor_text); + if (!(major == 0 && minor > 0)) { + append_phrase(integer_to_cardinal(major)); + append_phrase(major == 1 ? major_name : major_name + "s"); + } + if (minor > 0) { + if (major > 0) { + append_lookup("and"); + } + append_phrase(integer_to_cardinal(minor)); + append_phrase((minor == 1 || minor_name == "pence") ? minor_name : minor_name + "s"); + } + } else if (is_digit_text(no_commas) && contains_key(kOrdinals, lower_ascii(suffix))) { + append_phrase(integer_to_ordinal(std::stoll(no_commas))); + } else if (currency.has_value() == false && no_commas.find('.') == std::string::npos && no_commas.size() == 4 && is_digit_text(no_commas)) { + append_phrase(integer_to_year_words(std::stoll(no_commas))); + } else if (!is_head && no_commas.find('.') == std::string::npos) { + if ((no_commas.size() > 1 && no_commas.front() == '0') || no_commas.size() > 3) { + for (const std::string & digit_word : spell_digit_string(no_commas)) { + append_lookup(digit_word); + } + } else if (no_commas.size() == 3 && no_commas.substr(1) != "00") { + append_phrase(integer_to_cardinal(no_commas[0] - '0')); + if (no_commas[1] == '0') { + append_lookup("O", -2.0f); + append_phrase(integer_to_cardinal(no_commas[2] - '0')); + } else { + append_phrase(integer_to_cardinal(std::stoi(no_commas.substr(1)))); + } + } else { + append_phrase(integer_to_cardinal(std::stoll(no_commas))); + } + } else if (std::count(no_commas.begin(), no_commas.end(), '.') > 1 || !is_head) { + std::stringstream ss(no_commas); + std::string segment; + while (std::getline(ss, segment, '.')) { + if (segment.empty()) continue; + if ((segment.size() > 1 && segment.front() == '0') || (segment.size() != 2 && std::any_of(segment.begin() + 1, segment.end(), [](char ch) { return ch != '0'; }))) { + for (const std::string & digit_word : spell_digit_string(segment)) { + append_lookup(digit_word); + } + } else { + append_phrase(integer_to_cardinal(std::stoll(segment))); + } + } + } else if (no_commas.find('.') != std::string::npos) { + const size_t dot = no_commas.find('.'); + const std::string left = no_commas.substr(0, dot); + const std::string right = no_commas.substr(dot + 1); + if (left.empty()) { + append_lookup("point"); + } else { + append_phrase(integer_to_cardinal(std::stoll(left))); + append_lookup("point"); + } + for (const std::string & digit_word : spell_digit_string(right)) { + append_lookup(digit_word); + } + } else { + append_phrase(integer_to_cardinal(std::stoll(no_commas))); + } + + if (pieces.empty()) { + return std::nullopt; + } + + std::string phonemes; + int rating = pieces.front().second; + for (size_t i = 0; i < pieces.size(); ++i) { + if (i > 0) phonemes.push_back(' '); + phonemes += pieces[i].first; + rating = std::min(rating, pieces[i].second); + } + + const std::string lowered_suffix = lower_ascii(suffix); + if (lowered_suffix == "s" || lowered_suffix == "'s") { + auto updated = s_suffix(phonemes); + if (updated.has_value()) phonemes = *updated; + } else if (lowered_suffix == "ed" || lowered_suffix == "'d") { + auto updated = ed_suffix(phonemes); + if (updated.has_value()) phonemes = *updated; + } else if (lowered_suffix == "ing") { + auto updated = ing_suffix(phonemes); + if (updated.has_value()) phonemes = *updated; + } + return LexiconResult{phonemes, rating}; +} + +std::optional Lexicon::operator()(const MToken & token, const TokenContext & ctx) const { + std::string word = token.meta.alias.value_or(token.text); + std::replace(word.begin(), word.end(), static_cast(0x91), '\''); + std::replace(word.begin(), word.end(), static_cast(0x92), '\''); + + std::optional stress; + if (word != lower_ascii(word)) { + stress = word == upper_ascii(word) ? 2.0f : 0.5f; + } + if (const auto result = get_word(word, token.tag, stress, ctx)) { + std::string phonemes = result->phonemes; + if (token.meta.currency.has_value()) { + const auto it = kCurrencies.find(*token.meta.currency); + if (it != kCurrencies.end()) { + auto unit = get_word(it->second.first + "s", std::nullopt, std::nullopt, TokenContext{}); + if (unit.has_value()) { + phonemes += " " + unit->phonemes; + } + } + } + return LexiconResult{apply_stress(phonemes, token.meta.stress), result->rating}; + } + if (is_number_token(word, token.meta.is_head)) { + if (const auto result = get_number(word, token.meta.currency, token.meta.is_head)) { + return LexiconResult{apply_stress(result->phonemes, token.meta.stress), result->rating}; + } + } + if (!std::all_of(word.begin(), word.end(), [](unsigned char ch) { return std::isalpha(ch) != 0; })) { + return std::nullopt; + } + return std::nullopt; +} + +EnglishG2P::EnglishG2P(bool british) + : lexicon_(british) {} + +EnglishG2P::EnglishG2P(std::filesystem::path lexicon_dir, bool british) + : lexicon_( + lexicon_dir / (british ? "gb_gold.tsv" : "us_gold.tsv"), + lexicon_dir / (british ? "gb_silver.tsv" : "us_silver.tsv"), + british) {} + +PreprocessResult EnglishG2P::preprocess(const std::string & text) { + PreprocessResult result; + result.text = trim_copy(text); + std::string rebuilt; + size_t token_index = 0; + for (size_t i = 0; i < result.text.size();) { + if (result.text[i] == '[') { + const size_t close_text = result.text.find(']', i + 1); + const size_t open_feat = close_text == std::string::npos ? std::string::npos : result.text.find('(', close_text + 1); + const size_t close_feat = open_feat == std::string::npos ? std::string::npos : result.text.find(')', open_feat + 1); + if (close_text != std::string::npos && open_feat == close_text + 1 && close_feat != std::string::npos) { + const std::string visible = result.text.substr(i + 1, close_text - i - 1); + const std::string feature = result.text.substr(open_feat + 1, close_feat - open_feat - 1); + rebuilt += visible; + const size_t emitted_count = count_words(visible); + FeatureValue value; + if (!feature.empty()) { + int64_t as_int = 0; + const auto int_result = std::from_chars(feature.data(), feature.data() + feature.size(), as_int); + if (int_result.ec == std::errc() && int_result.ptr == feature.data() + feature.size()) { + value = as_int; + } else if (feature == "0.5" || feature == "+0.5") { + value = 0.5f; + } else if (feature == "-0.5") { + value = -0.5f; + } else if (feature.size() > 1 && feature.front() == '/' && feature.back() == '/') { + value = feature.substr(0, feature.size() - 1); + } else if (feature.size() > 1 && feature.front() == '#' && feature.back() == '#') { + value = feature.substr(0, feature.size() - 1); + } else { + value = std::string(); + } + if (emitted_count != 0) { + result.features[token_index] = value; + } + } + token_index += emitted_count; + i = close_feat + 1; + continue; + } + } + rebuilt.push_back(result.text[i]); + ++i; + } + result.text = rebuilt; + return result; +} + +std::vector EnglishG2P::tokenize( + const std::string & text, + const std::unordered_map & features) const { + std::vector tokens; + size_t i = 0; + while (i < text.size()) { + while (i < text.size() && std::isspace(static_cast(text[i])) != 0) { + ++i; + } + if (i >= text.size()) { + break; + } + size_t begin = i; + if (std::isalpha(static_cast(text[i])) != 0 || std::isdigit(static_cast(text[i])) != 0) { + while (i < text.size() && (std::isalnum(static_cast(text[i])) != 0 || text[i] == '\'' || text[i] == '.' || text[i] == ',' || text[i] == '-' || text[i] == '/' || text[i] == '_' || static_cast(text[i]) >= 0x80)) { + ++i; + } + } else { + ++i; + } + size_t end = i; + size_t ws_begin = i; + while (i < text.size() && std::isspace(static_cast(text[i])) != 0) { + ++i; + } + MToken token; + token.text = text.substr(begin, end - begin); + token.tag = guess_tag(token.text); + token.whitespace = text.substr(ws_begin, i - ws_begin); + token.meta.num_flags = ""; + tokens.push_back(std::move(token)); + } + for (const auto & [index, value] : features) { + if (index >= tokens.size()) { + continue; + } + auto & token = tokens[index]; + if (const auto * as_int = std::get_if(&value)) { + token.meta.stress = static_cast(*as_int); + } else if (const auto * as_float = std::get_if(&value)) { + token.meta.stress = *as_float; + } else if (const auto * as_string = std::get_if(&value)) { + if (starts_with(*as_string, "/")) { + token.phonemes = as_string->substr(1); + token.meta.rating = 5; + } else if (starts_with(*as_string, "#")) { + token.meta.num_flags = as_string->substr(1); + } + } + } + for (size_t idx = 0; idx < tokens.size(); ++idx) { + auto & token = tokens[idx]; + const std::string lowered = lower_ascii(token.text); + const std::string next = idx + 1 < tokens.size() ? lower_ascii(tokens[idx + 1].text) : std::string(); + if (lowered == "used" && next == "to") { + token.tag = "VBD"; + } else if (lowered == "by" && next != "the" && next != "way") { + token.tag = "RB"; + } else if (lowered == "read" && idx > 0 && lower_ascii(tokens[idx - 1].text) == "to") { + token.tag = "VB"; + } + } + return tokens; +} + +MToken EnglishG2P::merge_token_pair( + const MToken & left, + const MToken & right, + const std::optional & unk) { + MToken merged; + merged.text.reserve(left.text.size() + left.whitespace.size() + right.text.size()); + merged.text += left.text; + merged.text += left.whitespace; + merged.text += right.text; + merged.tag = left.tag; + merged.whitespace = right.whitespace; + merged.meta.is_head = left.meta.is_head; + merged.meta.prespace = left.meta.prespace; + merged.meta.num_flags = left.meta.num_flags; + if (unk.has_value()) { + std::string phonemes; + phonemes += left.phonemes.value_or(*unk); + if (right.meta.prespace && !phonemes.empty() && + !std::isspace(static_cast(phonemes.back())) && + right.phonemes.has_value()) { + phonemes.push_back(' '); + } + phonemes += right.phonemes.value_or(*unk); + merged.phonemes = std::move(phonemes); + } + if (left.meta.rating.has_value() && right.meta.rating.has_value()) { + merged.meta.rating = std::min(*left.meta.rating, *right.meta.rating); + } else if (left.meta.rating.has_value()) { + merged.meta.rating = left.meta.rating; + } else { + merged.meta.rating = right.meta.rating; + } + if (left.meta.stress.has_value() && right.meta.stress.has_value()) { + if (*left.meta.stress == *right.meta.stress) { + merged.meta.stress = left.meta.stress; + } + } else if (left.meta.stress.has_value()) { + merged.meta.stress = left.meta.stress; + } else { + merged.meta.stress = right.meta.stress; + } + merged.meta.currency = right.meta.currency.has_value() ? right.meta.currency : left.meta.currency; + return merged; +} + +MToken EnglishG2P::merge_tokens( + const std::vector & tokens, + size_t begin, + size_t end, + const std::optional & unk) { + MToken merged; + merged.text.reserve((end - begin) * 4); + for (size_t i = begin; i < end; ++i) { + merged.text += tokens[i].text; + if (i + 1 != end) { + merged.text += tokens[i].whitespace; + } + } + merged.tag = tokens[begin].tag; + merged.whitespace = tokens[end - 1].whitespace; + merged.meta.is_head = tokens[begin].meta.is_head; + merged.meta.prespace = tokens[begin].meta.prespace; + merged.meta.num_flags = tokens[begin].meta.num_flags; + std::optional rating; + std::optional stress; + std::optional currency; + if (unk.has_value()) { + std::string phonemes; + for (size_t i = begin; i < end; ++i) { + const auto & token = tokens[i]; + if (token.meta.prespace && !phonemes.empty() && !std::isspace(static_cast(phonemes.back())) && token.phonemes.has_value()) { + phonemes.push_back(' '); + } + phonemes += token.phonemes.value_or(*unk); + } + merged.phonemes = phonemes; + } + for (size_t i = begin; i < end; ++i) { + const auto & token = tokens[i]; + if (!rating.has_value()) { + rating = token.meta.rating; + } else if (token.meta.rating.has_value()) { + rating = std::min(*rating, *token.meta.rating); + } + if (token.meta.stress.has_value()) { + if (stress.has_value() && *stress != *token.meta.stress) { + stress = std::nullopt; + } else if (!stress.has_value()) { + stress = token.meta.stress; + } + } + if (token.meta.currency.has_value()) { + currency = token.meta.currency; + } + } + merged.meta.rating = rating; + merged.meta.stress = stress; + merged.meta.currency = currency; + return merged; +} + +std::vector EnglishG2P::fold_left(const std::vector & tokens) { + std::vector result; + result.reserve(tokens.size()); + for (const auto & token : tokens) { + if (!result.empty() && !token.meta.is_head) { + result.back() = merge_token_pair(result.back(), token, std::string()); + } else { + result.push_back(token); + } + } + return result; +} + +std::vector>> EnglishG2P::retokenize(const std::vector & tokens) { + std::vector>> words; + std::optional currency; + for (size_t i = 0; i < tokens.size(); ++i) { + const auto & token = tokens[i]; + std::vector subtokens; + if (!token.meta.alias.has_value() && !token.phonemes.has_value()) { + const auto parts = subtokenize(token.text); + subtokens.reserve(parts.size()); + for (const auto & part : parts) { + MToken sub = token; + sub.text = part; + sub.tag = part == token.text ? token.tag : guess_tag(part); + sub.whitespace.clear(); + sub.meta.is_head = true; + sub.meta.prespace = false; + subtokens.push_back(std::move(sub)); + } + } else { + subtokens.push_back(token); + } + if (!subtokens.empty()) { + subtokens.back().whitespace = token.whitespace; + } + for (size_t j = 0; j < subtokens.size(); ++j) { + auto & sub = subtokens[j]; + if (sub.meta.alias.has_value() || sub.phonemes.has_value()) { + } else if (sub.tag == "$" && contains_key(kCurrencies, sub.text)) { + currency = sub.text; + sub.phonemes = ""; + sub.meta.rating = 4; + } else if (sub.tag == ":" && (sub.text == "-" || sub.text == "–")) { + sub.phonemes = "—"; + sub.meta.rating = 3; + } else if (contains_key(kPunctTags, sub.tag) && + !std::all_of(sub.text.begin(), sub.text.end(), [](unsigned char ch) { return std::isalpha(ch) != 0; })) { + const auto it = kPunctTagPhonemes.find(sub.tag); + if (it != kPunctTagPhonemes.end()) { + sub.phonemes = it->second; + } else { + std::string kept; + for (char ch : sub.text) { + if (kPuncts.find(ch) != std::string::npos) { + kept.push_back(ch); + } + } + sub.phonemes = kept; + } + sub.meta.rating = 4; + } else if (currency.has_value()) { + if (sub.tag != "CD") { + currency.reset(); + } else if (j + 1 == subtokens.size() && (i + 1 == tokens.size() || tokens[i + 1].tag != "CD")) { + sub.meta.currency = currency; + } + } else if (0 < j && j + 1 < subtokens.size() && sub.text == "2" && + !subtokens[j - 1].text.empty() && !subtokens[j + 1].text.empty() && + std::isalpha(static_cast(subtokens[j - 1].text.back())) != 0 && + std::isalpha(static_cast(subtokens[j + 1].text.front())) != 0) { + sub.meta.alias = "to"; + } + + if (sub.meta.alias.has_value() || sub.phonemes.has_value()) { + words.push_back(sub); + } else if (!words.empty() && std::holds_alternative>(words.back()) && + std::get>(words.back()).back().whitespace.empty()) { + sub.meta.is_head = false; + std::get>(words.back()).push_back(sub); + } else { + if (sub.whitespace.empty()) { + words.push_back(std::vector{sub}); + } else { + words.push_back(sub); + } + } + } + } + for (auto & word : words) { + if (std::holds_alternative>(word)) { + auto & group = std::get>(word); + if (group.size() == 1) { + word = group.front(); + } + } + } + return words; +} + +void EnglishG2P::resolve_tokens(std::vector & tokens) { + std::string text; + for (size_t i = 0; i < tokens.size(); ++i) { + text += tokens[i].text; + if (i + 1 != tokens.size()) { + text += tokens[i].whitespace; + } + } + bool saw_alpha = false; + bool saw_digit = false; + bool saw_other = false; + for (char ch : text) { + if (contains_key(kSubtokenJunk, ch)) { + continue; + } + const int klass = english_g2p_char_class(ch); + saw_alpha = saw_alpha || klass == 0; + saw_digit = saw_digit || klass == 1; + saw_other = saw_other || klass == 2; + } + const bool prespace = + contains_whitespace(text) || + text.find('/') != std::string::npos || + static_cast(saw_alpha) + static_cast(saw_digit) + static_cast(saw_other) > 1; + for (size_t i = 0; i < tokens.size(); ++i) { + auto & token = tokens[i]; + if (!token.phonemes.has_value()) { + if (i + 1 == tokens.size() && token.text.size() == 1 && kNonQuotePuncts.find(token.text[0]) != std::string::npos) { + token.phonemes = token.text; + token.meta.rating = 3; + } else if (all_subtoken_junk(token.text)) { + token.phonemes = ""; + token.meta.rating = 3; + } + } else if (i > 0) { + token.meta.prespace = prespace; + } + } + if (prespace) { + return; + } + struct WeightedIndex { + bool primary = false; + int weight = 0; + size_t index = 0; + }; + std::vector indices; + for (size_t i = 0; i < tokens.size(); ++i) { + if (tokens[i].phonemes.has_value() && !tokens[i].phonemes->empty()) { + indices.push_back({tokens[i].phonemes->find(kPrimaryStress) != std::string::npos, stress_weight(*tokens[i].phonemes), i}); + } + } + int primary_count = 0; + for (const auto & entry : indices) primary_count += entry.primary ? 1 : 0; + if (indices.size() == 2 && tokens[indices[0].index].text.size() == 1) { + tokens[indices[1].index].phonemes = Lexicon::apply_stress(*tokens[indices[1].index].phonemes, -0.5f); + return; + } + if (indices.size() < 2 || primary_count <= static_cast((indices.size() + 1) / 2)) { + return; + } + std::sort(indices.begin(), indices.end(), [](const WeightedIndex & lhs, const WeightedIndex & rhs) { + if (lhs.primary != rhs.primary) return lhs.primary < rhs.primary; + if (lhs.weight != rhs.weight) return lhs.weight < rhs.weight; + return lhs.index < rhs.index; + }); + indices.resize(indices.size() / 2); + for (const auto & entry : indices) { + tokens[entry.index].phonemes = Lexicon::apply_stress(*tokens[entry.index].phonemes, -0.5f); + } +} + +TokenContext EnglishG2P::token_context(const TokenContext & ctx, const std::optional & phonemes, const MToken & token) { + std::optional vowel = ctx.future_vowel; + if (phonemes.has_value()) { + for (size_t i = 0; i < phonemes->size();) { + const size_t width = utf8_codepoint_width(*phonemes, i); + const char * symbol = phonemes->data() + i; + if (utf8_inventory_contains(kVowels, symbol, width)) { + vowel = true; + break; + } + if (utf8_inventory_contains(kConsonants, symbol, width) || + (width == 1 && kNonQuotePuncts.find(*symbol) != std::string::npos)) { + vowel = false; + break; + } + i += width; + } + } + const bool future_to = token.text == "to" || token.text == "To" || (token.text == "TO" && (token.tag == "TO" || token.tag == "IN")); + return TokenContext{vowel, future_to}; +} + +std::string EnglishG2P::tokens_to_phonemes(const std::vector & tokens) { + std::string result; + size_t reserve = 0; + for (const auto & token : tokens) { + if (token.phonemes.has_value()) { + reserve += token.phonemes->size(); + if (!token.whitespace.empty()) { + ++reserve; + } + } + } + result.reserve(reserve); + bool pending_space = false; + for (const auto & token : tokens) { + if (!token.phonemes.has_value()) { + continue; + } + if (pending_space && !result.empty()) { + result.push_back(' '); + } + result += *token.phonemes; + pending_space = !token.whitespace.empty(); + } + return result; +} + +std::pair> EnglishG2P::operator()(const std::string & text, bool enable_preprocess) const { + const PreprocessResult pre = enable_preprocess ? preprocess(text) : PreprocessResult{text, {}}; + auto tokens = tokenize(pre.text, pre.features); + tokens = fold_left(tokens); + auto words = retokenize(tokens); + TokenContext ctx; + std::vector resolved; + resolved.reserve(words.size()); + for (auto it = words.rbegin(); it != words.rend(); ++it) { + if (std::holds_alternative(*it)) { + MToken token = std::move(std::get(*it)); + if (!token.phonemes.has_value()) { + auto result = lexicon_(token, ctx); + if (!result.has_value()) { + result = resolve_with_safe_fallback(lexicon_, token); + } + token.phonemes = result->phonemes; + token.meta.rating = result->rating; + } + ctx = token_context(ctx, token.phonemes, token); + resolved.push_back(std::move(token)); + continue; + } + + auto group = std::move(std::get>(*it)); + bool should_fallback = false; + size_t left = 0; + size_t right = group.size(); + while (left < right) { + bool blocked = false; + for (size_t i = left; i < right; ++i) { + if (group[i].meta.alias.has_value() || group[i].phonemes.has_value()) { + blocked = true; + break; + } + } + if (!blocked) { + MToken merged = merge_tokens(group, left, right, std::nullopt); + if (auto result = lexicon_(merged, ctx)) { + group[left].phonemes = result->phonemes; + group[left].meta.rating = result->rating; + for (size_t i = left + 1; i < right; ++i) { + group[i].phonemes = ""; + group[i].meta.rating = result->rating; + } + ctx = token_context(ctx, group[left].phonemes, merged); + right = left; + left = 0; + continue; + } + } + if (left + 1 < right) { + ++left; + continue; + } + --right; + if (!group[right].phonemes.has_value()) { + if (all_subtoken_junk(group[right].text)) { + group[right].phonemes = ""; + group[right].meta.rating = 3; + } else { + should_fallback = true; + break; + } + } + left = 0; + } + if (should_fallback) { + MToken merged = merge_tokens(group, 0, group.size(), std::nullopt); + const LexiconResult fallback_result = resolve_with_safe_fallback(lexicon_, merged); + group[0].phonemes = fallback_result.phonemes; + group[0].meta.rating = fallback_result.rating; + for (size_t i = 1; i < group.size(); ++i) { + group[i].phonemes = ""; + group[i].meta.rating = fallback_result.rating; + } + } else { + resolve_tokens(group); + } + for (auto rit = group.rbegin(); rit != group.rend(); ++rit) { + ctx = token_context(ctx, rit->phonemes, *rit); + } + std::reverse(group.begin(), group.end()); + for (auto & token : group) { + resolved.push_back(std::move(token)); + } + } + std::reverse(resolved.begin(), resolved.end()); + for (auto & token : resolved) { + if (token.phonemes.has_value()) { + replace_all(*token.phonemes, "ɾ", "T"); + replace_all(*token.phonemes, "ʔ", "t"); + } + } + return {tokens_to_phonemes(resolved), resolved}; +} + +} // namespace kokoro_ggml::g2p_en diff --git a/src/models/kokoro_tts/loader.cpp b/src/models/kokoro_tts/loader.cpp new file mode 100644 index 000000000..2f1184529 --- /dev/null +++ b/src/models/kokoro_tts/loader.cpp @@ -0,0 +1,121 @@ +#include "engine/models/kokoro_tts/loader.h" +#include "engine/models/kokoro_tts/session.h" + +#include "engine/framework/io/filesystem.h" + +#include + +namespace engine::models::kokoro_tts { + +namespace { + +std::filesystem::path resolve_model_root(const std::filesystem::path & model_path) { + if (engine::io::is_existing_directory(model_path)) { + return std::filesystem::weakly_canonical(model_path); + } + throw std::runtime_error("Kokoro TTS expects a model directory: " + model_path.string()); +} + +std::vector discover_config_assets(const runtime::ModelLoadRequest & request) { + return runtime::discover_named_assets(resolve_model_root(request.model_path), {"config.json", "voices.json", "vocab.tsv"}); +} + +std::vector discover_weight_assets(const runtime::ModelLoadRequest & request) { + return runtime::discover_named_assets(resolve_model_root(request.model_path), {"kokoro-v1_0.safetensors"}); +} + +class KokoroTTSLoader final : public runtime::IVoiceModelLoader { +public: + std::string family() const override { + return "kokoro_tts"; + } + + bool can_load(const runtime::ModelLoadRequest & request) const override { + try { + const auto root = resolve_model_root(request.model_path); + return engine::io::is_existing_file(root / "config.json") + && engine::io::is_existing_file(root / "voices.json") + && engine::io::is_existing_file(root / "kokoro-v1_0.safetensors") + && (!request.family_hint.has_value() || *request.family_hint == family()); + } catch (...) { + return false; + } + } + + runtime::ModelInspection inspect(const runtime::ModelLoadRequest & request) const override { + const auto root = resolve_model_root(request.model_path); + runtime::ModelInspection inspection; + inspection.model_root = root; + inspection.metadata.family = family(); + inspection.metadata.variant = root.filename().string(); + inspection.metadata.description = "Kokoro TTS loaded from local extracted assets."; + inspection.metadata.config_candidates = {"config.json", "voices.json", "vocab.tsv"}; + inspection.metadata.weight_candidates = {"kokoro-v1_0.safetensors"}; + inspection.capabilities.supported_tasks = { + {runtime::VoiceTaskKind::Tts, {runtime::RunMode::Offline}}, + }; + inspection.capabilities.languages = {"a", "b"}; + inspection.capabilities.supports_style_condition = true; + inspection.discovered_configs = discover_config_assets(request); + inspection.discovered_weights = discover_weight_assets(request); + return inspection; + } + + std::unique_ptr load(const runtime::ModelLoadRequest & request) const override { + return load_kokoro_tts_model(resolve_model_root(request.model_path)); + } +}; + +} // namespace + +KokoroTTSLoadedModel::KokoroTTSLoadedModel( + runtime::ModelMetadata metadata, + runtime::CapabilitySet capabilities, + std::shared_ptr assets) + : metadata_(std::move(metadata)), + capabilities_(std::move(capabilities)), + assets_(std::move(assets)) {} + +const runtime::ModelMetadata & KokoroTTSLoadedModel::metadata() const noexcept { + return metadata_; +} + +const runtime::CapabilitySet & KokoroTTSLoadedModel::capabilities() const noexcept { + return capabilities_; +} + +std::unique_ptr KokoroTTSLoadedModel::create_task_session( + const runtime::TaskSpec & task, + const runtime::SessionOptions & options) const { + if (task.task != runtime::VoiceTaskKind::Tts) { + throw std::runtime_error("Kokoro TTS only supports VoiceTaskKind::Tts"); + } + return std::make_unique(task, options, assets_); +} + +std::unique_ptr load_kokoro_tts_model(const std::filesystem::path & model_path) { + const auto root = resolve_model_root(model_path); + auto assets = load_kokoro_assets(root); + + runtime::ModelMetadata metadata; + metadata.family = "kokoro_tts"; + metadata.variant = root.filename().string(); + metadata.description = "Kokoro TTS loaded from local extracted assets."; + metadata.config_candidates = {"config.json", "voices.json", "vocab.tsv"}; + metadata.weight_candidates = {"kokoro-v1_0.safetensors"}; + + runtime::CapabilitySet capabilities; + capabilities.supported_tasks = { + {runtime::VoiceTaskKind::Tts, {runtime::RunMode::Offline}}, + }; + capabilities.languages = {"a", "b"}; + capabilities.supports_style_condition = true; + + return std::make_unique(std::move(metadata), std::move(capabilities), std::move(assets)); +} + +std::shared_ptr make_kokoro_tts_loader() { + return std::make_shared(); +} + +} // namespace engine::models::kokoro_tts diff --git a/src/models/kokoro_tts/plbert.cpp b/src/models/kokoro_tts/plbert.cpp new file mode 100644 index 000000000..932bbc44d --- /dev/null +++ b/src/models/kokoro_tts/plbert.cpp @@ -0,0 +1,442 @@ +#include "engine/models/kokoro_tts/plbert.h" + +#include "engine/models/kokoro_tts/assets.h" + +#include "engine/framework/core/backend.h" +#include "engine/framework/debug/profiler.h" +#include "engine/framework/debug/trace.h" + +#include +#include + +#include +#include +#include +#include +#include +#include +#include + +namespace kokoro_ggml { + +constexpr size_t kPlbertCtxBytes = 128ull * 1024ull * 1024ull; + +namespace core = engine::core; + +namespace { + +using engine::debug::measure_ms; + +ggml_tensor * add_bias_3d( + ggml_context * ctx, + ggml_tensor * x, + const core::TensorValue & bias, + int64_t channels) { + ggml_tensor * bias_tensor = ggml_reshape_3d(ctx, bias.tensor, channels, 1, 1); + return ggml_add(ctx, x, bias_tensor); +} + +ggml_tensor * linear_3d( + ggml_context * ctx, + ggml_tensor * x, + const KokoroWeights::LinearWeights & linear) { + ggml_tensor * y = ggml_mul_mat(ctx, linear.weight.tensor, x); + if (linear.use_bias) { + y = add_bias_3d(ctx, y, *linear.bias, linear.out_features); + } + return y; +} + +ggml_tensor * layer_norm_3d( + ggml_context * ctx, + ggml_tensor * x, + const KokoroWeights::LayerNormWeights & norm) { + ggml_tensor * y = ggml_norm(ctx, x, norm.eps); + ggml_tensor * gamma = ggml_reshape_3d(ctx, norm.weight.tensor, norm.channels, 1, 1); + ggml_tensor * beta = ggml_reshape_3d(ctx, norm.bias.tensor, norm.channels, 1, 1); + y = ggml_mul(ctx, y, gamma); + y = ggml_add(ctx, y, beta); + return y; +} + +ggml_tensor * gelu_new_3d(ggml_context * ctx, ggml_tensor * x) { + constexpr float kCubeCoeff = 0.044715f; + constexpr float kScale = 0.7978845608028654f; // sqrt(2 / pi) + ggml_tensor * x_sq = ggml_mul(ctx, x, x); + ggml_tensor * x_cube = ggml_mul(ctx, x_sq, x); + ggml_tensor * inner = ggml_add(ctx, x, ggml_scale(ctx, x_cube, kCubeCoeff)); + ggml_tensor * tanh_term = ggml_tanh(ctx, ggml_scale(ctx, inner, kScale)); + ggml_tensor * shifted = ggml_scale_bias(ctx, tanh_term, 1.0f, 1.0f); + ggml_tensor * scaled = ggml_scale(ctx, ggml_mul(ctx, x, shifted), 0.5f); + return scaled; +} + +ggml_tensor * embedding_lookup( + ggml_context * ctx, + const KokoroWeights::EmbeddingWeights & embedding, + ggml_tensor * ids) { + return ggml_get_rows(ctx, embedding.weight.tensor, ids); +} + +core::TensorValue permute_tensor( + ggml_context * ctx, + const core::TensorValue & input, + const std::array & axes) { + core::TensorShape output_shape = {}; + output_shape.rank = input.shape.rank; + std::array ggml_axes = {0, 1, 2, 3}; + for (size_t out_axis = 0; out_axis < input.shape.rank; ++out_axis) { + const int in_axis = axes[out_axis]; + if (in_axis < 0 || in_axis >= static_cast(input.shape.rank)) { + throw std::runtime_error("Kokoro PL-BERT permute axis out of range"); + } + output_shape.dims[out_axis] = input.shape.dims[in_axis]; + const int out_ggml_axis = static_cast(input.shape.rank) - 1 - static_cast(out_axis); + ggml_axes[out_ggml_axis] = core::logical_axis_to_ggml_axis(input.shape.rank, in_axis); + } + return core::wrap_tensor( + ggml_permute(ctx, input.tensor, ggml_axes[0], ggml_axes[1], ggml_axes[2], ggml_axes[3]), + output_shape, + input.type); +} + +core::TensorValue ensure_contiguous(ggml_context * ctx, const core::TensorValue & input) { + return core::has_backend_addressable_layout(input.tensor) + ? input + : core::wrap_tensor(ggml_cont(ctx, input.tensor), input.shape, input.type); +} + +core::TensorValue heads_from_linear( + ggml_context * ctx, + ggml_tensor * value, + int64_t batch, + int64_t token_count, + int64_t num_heads, + int64_t head_dim) { + auto logical = core::wrap_tensor( + ggml_reshape_4d(ctx, value, head_dim, num_heads, token_count, batch), + core::TensorShape::from_dims({batch, token_count, num_heads, head_dim}), + GGML_TYPE_F32); + auto heads = permute_tensor(ctx, logical, {0, 2, 1, 3}); + return ensure_contiguous(ctx, heads); +} + +std::vector build_attention_mask(const std::vector & token_validity, int64_t token_count, int64_t batch) { + constexpr float kMaskedAttentionBias = -65504.0f; + std::vector mask(static_cast(token_count * token_count * batch), ggml_fp32_to_fp16(0.0f)); + for (int64_t b = 0; b < batch; ++b) { + for (int64_t q = 0; q < token_count; ++q) { + for (int64_t k = 0; k < token_count; ++k) { + const bool keep = + token_validity[static_cast(b * token_count + q)] != 0 && + token_validity[static_cast(b * token_count + k)] != 0; + const float value = keep ? 0.0f : kMaskedAttentionBias; + const size_t offset = static_cast(b * token_count * token_count + q * token_count + k); + mask[offset] = ggml_fp32_to_fp16(value); + } + } + } + return mask; +} + +ggml_tensor * albert_attention( + ggml_context * ctx, + ggml_tensor * x, + ggml_tensor * mask, + const KokoroWeights::AlbertWeights & bert, + const KokoroWeights::AlbertAttentionWeights & attention) { + const int64_t token_count = x->ne[1]; + const int64_t batch = x->ne[2]; + const int64_t head_dim = bert.hidden_size / bert.num_attention_heads; + + ggml_tensor * q = linear_3d(ctx, x, attention.query); + ggml_tensor * k = linear_3d(ctx, x, attention.key); + ggml_tensor * v = linear_3d(ctx, x, attention.value); + + auto q_heads = heads_from_linear(ctx, q, batch, token_count, bert.num_attention_heads, head_dim); + auto k_heads = heads_from_linear(ctx, k, batch, token_count, bert.num_attention_heads, head_dim); + auto v_heads = heads_from_linear(ctx, v, batch, token_count, bert.num_attention_heads, head_dim); + + auto scores = core::wrap_tensor( + ggml_mul_mat(ctx, k_heads.tensor, q_heads.tensor), + core::TensorShape::from_dims({batch, bert.num_attention_heads, token_count, token_count}), + GGML_TYPE_F32); + scores = core::wrap_tensor( + ggml_scale(ctx, scores.tensor, 1.0f / std::sqrt(static_cast(head_dim))), + scores.shape, + GGML_TYPE_F32); + if (mask != nullptr) { + scores = core::wrap_tensor( + ggml_add(ctx, scores.tensor, ggml_cast(ctx, mask, GGML_TYPE_F32)), + scores.shape, + GGML_TYPE_F32); + } + scores = ensure_contiguous(ctx, scores); + auto probs = core::wrap_tensor(ggml_soft_max(ctx, scores.tensor), scores.shape, GGML_TYPE_F32); + + const auto v_transposed = ensure_contiguous(ctx, permute_tensor(ctx, v_heads, {0, 1, 3, 2})); + auto context = core::wrap_tensor( + ggml_mul_mat(ctx, v_transposed.tensor, probs.tensor), + core::TensorShape::from_dims({batch, bert.num_attention_heads, token_count, head_dim}), + GGML_TYPE_F32); + context = permute_tensor(ctx, context, {0, 2, 1, 3}); + context = ensure_contiguous(ctx, context); + auto attn_input = core::wrap_tensor( + ggml_reshape_3d(ctx, context.tensor, bert.hidden_size, token_count, batch), + core::TensorShape::from_dims({batch, token_count, bert.hidden_size}), + GGML_TYPE_F32); + + ggml_tensor * attn = linear_3d(ctx, attn_input.tensor, attention.dense); + attn = ggml_add(ctx, attn, x); + return layer_norm_3d(ctx, attn, attention.layer_norm); +} + +ggml_tensor * albert_layer( + ggml_context * ctx, + ggml_tensor * x, + ggml_tensor * mask, + const KokoroWeights::AlbertWeights & bert, + const KokoroWeights::AlbertLayerWeights & layer) { + ggml_tensor * y = albert_attention(ctx, x, mask, bert, layer.attention); + ggml_tensor * ffn = linear_3d(ctx, y, layer.ffn); + ffn = gelu_new_3d(ctx, ffn); + ffn = linear_3d(ctx, ffn, layer.ffn_output); + ffn = ggml_add(ctx, ffn, y); + return layer_norm_3d(ctx, ffn, layer.full_layer_layer_norm); +} + +ggml_tensor * plbert_last_hidden_state( + ggml_context * ctx, + ggml_tensor * input_ids, + ggml_tensor * attn_mask, + ggml_tensor * position_ids, + ggml_tensor * token_type_ids, + const KokoroWeights & weights) { + ggml_tensor * word = embedding_lookup(ctx, weights.bert.embeddings.word_embeddings, input_ids); + ggml_tensor * pos = embedding_lookup(ctx, weights.bert.embeddings.position_embeddings, position_ids); + ggml_tensor * tok = embedding_lookup(ctx, weights.bert.embeddings.token_type_embeddings, token_type_ids); + ggml_tensor * x = ggml_add(ctx, ggml_add(ctx, word, pos), tok); + x = layer_norm_3d(ctx, x, weights.bert.embeddings.layer_norm); + x = linear_3d(ctx, x, weights.bert.embedding_hidden_mapping_in); + + for (int64_t i = 0; i < weights.bert.num_hidden_layers; ++i) { + x = albert_layer(ctx, x, attn_mask, weights.bert, weights.bert.shared_layer); + } + return x; +} + +} // namespace + +int64_t kokoro_plbert_output_dim(std::shared_ptr weights, bool project_hidden) { + if (!weights) { + throw std::runtime_error("Kokoro weights are null"); + } + return project_hidden ? weights->hidden_dim : weights->bert.hidden_size; +} + +struct PlbertSession { + std::shared_ptr weights; + ggml_backend_t backend = nullptr; + int64_t token_count = 0; + bool project_hidden = true; + int n_threads = 1; + bool use_device_backend = false; + ggml_context * ctx = nullptr; + ggml_tensor * ids = nullptr; + ggml_tensor * attn_mask = nullptr; + ggml_tensor * position_ids = nullptr; + ggml_tensor * token_type_ids = nullptr; + ggml_tensor * output = nullptr; + ggml_cgraph * graph = nullptr; + ggml_backend_buffer_t buffer = nullptr; + + PlbertSession( + std::shared_ptr weights_in, + ggml_backend_t backend_in, + int64_t token_count_in, + bool project_hidden_in, + int n_threads_in, + bool use_device_backend_in) + : weights(std::move(weights_in)), + backend(backend_in), + token_count(token_count_in), + project_hidden(project_hidden_in), + n_threads(n_threads_in), + use_device_backend(use_device_backend_in) { + ggml_init_params params{ + /*.mem_size =*/ kPlbertCtxBytes, + /*.mem_buffer =*/ nullptr, + /*.no_alloc =*/ true, + }; + ctx = ggml_init(params); + if (!ctx) { + throw std::runtime_error("failed to initialize ggml context for kokoro_plbert_encode"); + } + try { + ids = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, token_count, 1); + attn_mask = ggml_new_tensor_4d(ctx, GGML_TYPE_F16, token_count, token_count, 1, 1); + position_ids = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, token_count, 1); + token_type_ids = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, token_count, 1); + ggml_set_input(ids); + ggml_set_input(attn_mask); + ggml_set_input(position_ids); + ggml_set_input(token_type_ids); + + ggml_tensor * hidden = plbert_last_hidden_state(ctx, ids, attn_mask, position_ids, token_type_ids, *weights); + output = project_hidden ? linear_3d(ctx, hidden, weights->bert_encoder) : hidden; + graph = ggml_new_graph_custom(ctx, 4096, false); + ggml_build_forward_expand(graph, output); + + core::set_backend_threads(backend, n_threads); + const double alloc_ms = measure_ms([&]() { + buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + }); + if (buffer == nullptr) { + throw std::runtime_error("failed to allocate Kokoro PL-BERT tensors"); + } + const double fixed_input_upload_ms = measure_ms([&]() { + std::vector positions(static_cast(token_count), 0); + std::vector token_types(static_cast(token_count), 0); + for (int64_t t = 0; t < token_count; ++t) { + positions[static_cast(t)] = static_cast(t); + } + ggml_backend_tensor_set(position_ids, positions.data(), 0, ggml_nbytes(position_ids)); + ggml_backend_tensor_set(token_type_ids, token_types.data(), 0, ggml_nbytes(token_type_ids)); + }); + engine::debug::timing_log_scalar("kokoro.graph.build.plbert_alloc_ms", alloc_ms); + engine::debug::timing_log_scalar("kokoro.graph.build.plbert_fixed_input_upload_ms", fixed_input_upload_ms); + } catch (...) { + if (buffer) { + ggml_backend_buffer_free(buffer); + buffer = nullptr; + } + ggml_free(ctx); + ctx = nullptr; + throw; + } + } + + ~PlbertSession() { + if (buffer) { + ggml_backend_buffer_free(buffer); + } + if (ctx) { + ggml_free(ctx); + } + } + + bool matches( + const std::shared_ptr & weights_in, + int64_t token_count_in, + bool project_hidden_in, + int n_threads_in, + bool use_device_backend_in) const { + return weights.get() == weights_in.get() && + token_count == token_count_in && + project_hidden == project_hidden_in && + n_threads == n_threads_in && + use_device_backend == use_device_backend_in; + } + + std::vector run(const std::vector & input_ids) { + if (input_ids.empty() || static_cast(input_ids.size()) > token_count) { + throw std::runtime_error("Kokoro PL-BERT input length exceeds prepared capacity"); + } + std::vector padded_ids(static_cast(token_count), 0); + std::memcpy( + padded_ids.data(), + input_ids.data(), + static_cast(input_ids.size()) * sizeof(int32_t)); + std::vector valid_mask(static_cast(token_count), 0); + std::fill_n(valid_mask.begin(), input_ids.size(), 1); + const std::vector attn_mask_host = build_attention_mask(valid_mask, token_count, 1); + + ggml_backend_tensor_set(ids, padded_ids.data(), 0, ggml_nbytes(ids)); + ggml_backend_tensor_set(attn_mask, attn_mask_host.data(), 0, ggml_nbytes(attn_mask)); + core::set_backend_threads(backend, n_threads); + const ggml_status status = engine::core::compute_backend_graph(backend, graph); + if (status != GGML_STATUS_SUCCESS) { + throw std::runtime_error(std::string("kokoro_plbert_encode graph compute failed: ") + ggml_status_to_string(status)); + } + std::vector full(static_cast(ggml_nelements(output))); + ggml_backend_tensor_get(output, full.data(), 0, ggml_nbytes(output)); + const int64_t hidden_dim = output->ne[0]; + std::vector result(static_cast(hidden_dim * static_cast(input_ids.size()))); + for (int64_t row = 0; row < static_cast(input_ids.size()); ++row) { + std::memcpy( + result.data() + static_cast(row * hidden_dim), + full.data() + static_cast(row * hidden_dim), + static_cast(hidden_dim) * sizeof(float)); + } + return result; + } +}; + +struct KokoroPlbertRuntime::Impl { + std::shared_ptr weights; + ggml_backend_t backend = nullptr; + int n_threads = 1; + bool use_device_backend = false; + int64_t fixed_token_capacity = 0; + std::unique_ptr session; + + Impl( + std::shared_ptr weights_in, + ggml_backend_t backend_in, + int n_threads_in, + bool use_device_backend_in, + int64_t fixed_token_capacity_in) + : weights(std::move(weights_in)), + backend(backend_in), + n_threads(std::max(1, n_threads_in)), + use_device_backend(use_device_backend_in), + fixed_token_capacity(fixed_token_capacity_in) {} + + PlbertSession & session_for(int64_t token_count, bool project_hidden) { + const int64_t effective_token_count = fixed_token_capacity > 0 ? fixed_token_capacity : token_count; + if (token_count > effective_token_count) { + throw std::runtime_error("Kokoro PL-BERT request length exceeds fixed token capacity"); + } + if (session && session->matches(weights, effective_token_count, project_hidden, n_threads, use_device_backend)) { + return *session; + } + session.reset(); + const double build_ms = measure_ms([&]() { + session = std::make_unique( + weights, + backend, + effective_token_count, + project_hidden, + n_threads, + use_device_backend); + }); + engine::debug::timing_log_scalar("kokoro.graph.build.plbert_ms", build_ms); + return *session; + } +}; + +KokoroPlbertRuntime::KokoroPlbertRuntime( + std::shared_ptr weights, + ggml_backend_t backend, + int n_threads, + bool use_device_backend, + int64_t fixed_token_capacity) + : impl_(std::make_unique(std::move(weights), backend, n_threads, use_device_backend, fixed_token_capacity)) { + if (!impl_->weights) { + throw std::runtime_error("Kokoro weights are null"); + } +} + +KokoroPlbertRuntime::~KokoroPlbertRuntime() = default; + +std::vector KokoroPlbertRuntime::encode(const std::vector & input_ids, bool project_hidden) { + if (input_ids.empty()) { + throw std::runtime_error("kokoro_plbert_encode requires non-empty input_ids"); + } + if (static_cast(input_ids.size()) > impl_->weights->context_length) { + throw std::runtime_error("kokoro_plbert_encode input length exceeds PL-BERT context length"); + } + return impl_->session_for(static_cast(input_ids.size()), project_hidden).run(input_ids); +} + +} // namespace kokoro_ggml diff --git a/src/models/kokoro_tts/predictor.cpp b/src/models/kokoro_tts/predictor.cpp new file mode 100644 index 000000000..0eadc2c12 --- /dev/null +++ b/src/models/kokoro_tts/predictor.cpp @@ -0,0 +1,2054 @@ +#include "engine/models/kokoro_tts/predictor.h" + +#include "engine/models/kokoro_tts/plbert.h" +#include "engine/models/kokoro_tts/assets.h" + +#include "engine/framework/core/backend.h" +#include "engine/framework/debug/profiler.h" +#include "engine/framework/debug/trace.h" +#include "engine/framework/modules/activation_modules.h" +#include "engine/framework/modules/conditioning_modules.h" +#include "engine/framework/modules/conv_modules.h" +#include "engine/framework/modules/linear_module.h" +#include "engine/framework/modules/lookup_modules.h" +#include "engine/framework/modules/norm_modules.h" +#include "engine/framework/modules/recurrent_modules.h" +#include "engine/framework/modules/structural_modules.h" + +#include +#include + +#include +#include +#include +#include +#include +#include +#include + +namespace kokoro_ggml { + +namespace core = engine::core; +namespace modules = engine::modules; + +namespace { + +using engine::debug::measure_ms; + +int64_t checked_capacity(int64_t value, int64_t max_value) { + if (value <= 0 || max_value <= 0) { + throw std::runtime_error("kokoro capacity must be positive"); + } + if (value > max_value) { + throw std::runtime_error("kokoro requested capacity exceeds maximum"); + } + return value; +} + +core::TensorValue view_2d_last_dim_slice( + core::ModuleBuildContext & ctx, + const core::TensorValue & value, + int64_t start, + int64_t width) { + if (value.shape.rank != 2 || value.type != GGML_TYPE_F32) { + throw std::runtime_error("kokoro predictor slice expects an F32 rank-2 tensor"); + } + if (start < 0 || width <= 0 || start + width > value.shape.dims[1]) { + throw std::runtime_error("kokoro predictor slice is out of bounds"); + } + ggml_tensor * contiguous = ggml_cont(ctx.ggml, value.tensor); + return core::wrap_tensor( + ggml_view_2d( + ctx.ggml, + contiguous, + width, + value.shape.dims[0], + contiguous->nb[1], + static_cast(start) * contiguous->nb[0]), + core::TensorShape::from_dims({value.shape.dims[0], width}), + GGML_TYPE_F32); +} + +struct TimeMaskInputs { + ggml_tensor * keep = nullptr; + ggml_tensor * norm = nullptr; + int64_t frame_capacity = 0; +}; + +const TimeMaskInputs & add_time_mask_inputs( + ggml_context * ctx, + std::vector & masks, + int64_t frame_capacity) { + if (frame_capacity <= 0) { + throw std::runtime_error("kokoro predictor mask capacity must be positive"); + } + TimeMaskInputs mask = {}; + mask.keep = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, frame_capacity, 1); + mask.norm = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, frame_capacity, 1); + mask.frame_capacity = frame_capacity; + ggml_set_input(mask.keep); + ggml_set_input(mask.norm); + masks.push_back(mask); + return masks.back(); +} + +const TimeMaskInputs & add_keep_mask_input( + ggml_context * ctx, + std::vector & masks, + int64_t frame_capacity) { + if (frame_capacity <= 0) { + throw std::runtime_error("kokoro predictor keep-mask capacity must be positive"); + } + TimeMaskInputs mask = {}; + mask.keep = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, frame_capacity, 1); + mask.frame_capacity = frame_capacity; + ggml_set_input(mask.keep); + masks.push_back(mask); + return masks.back(); +} + +ggml_tensor * repeat_mask_like(ggml_context * ctx, ggml_tensor * mask, ggml_tensor * like) { + return ggml_repeat(ctx, mask, like); +} + +core::TensorValue broadcast_channel_2d( + core::ModuleBuildContext & ctx, + const core::TensorValue & value, + ggml_tensor * like, + int64_t channels) { + core::validate_shape(value, core::TensorShape::from_dims({channels}), "kokoro adain channel value"); + ggml_tensor * reshaped = ggml_reshape_2d(ctx.ggml, value.tensor, 1, channels); + return core::wrap_tensor( + ggml_repeat(ctx.ggml, reshaped, like), + core::TensorShape::from_dims({channels, like->ne[0]}), + GGML_TYPE_F32); +} + +ggml_tensor * build_masked_adain_ct( + core::ModuleBuildContext & build_ctx, + ggml_tensor * x, + const core::TensorValue & gamma, + const core::TensorValue & beta, + int64_t channels, + float eps, + std::vector & masks) { + const TimeMaskInputs & mask = add_time_mask_inputs(build_ctx.ggml, masks, x->ne[0]); + ggml_tensor * keep = repeat_mask_like(build_ctx.ggml, mask.keep, x); + ggml_tensor * norm = repeat_mask_like(build_ctx.ggml, mask.norm, x); + ggml_tensor * masked = ggml_mul(build_ctx.ggml, x, norm); + ggml_tensor * mean = ggml_mean(build_ctx.ggml, masked); + ggml_tensor * centered = ggml_sub(build_ctx.ggml, x, ggml_repeat(build_ctx.ggml, mean, x)); + ggml_tensor * centered_for_variance = ggml_mul(build_ctx.ggml, centered, norm); + ggml_tensor * squared = ggml_mul(build_ctx.ggml, centered, centered_for_variance); + ggml_tensor * variance = ggml_mean(build_ctx.ggml, squared); + ggml_tensor * stddev = ggml_sqrt(build_ctx.ggml, ggml_scale_bias(build_ctx.ggml, variance, 1.0f, eps)); + ggml_tensor * normalized = ggml_div(build_ctx.ggml, centered, ggml_repeat(build_ctx.ggml, stddev, x)); + normalized = ggml_mul(build_ctx.ggml, normalized, keep); + const auto gamma_rep = broadcast_channel_2d(build_ctx, gamma, x, channels); + const auto beta_rep = broadcast_channel_2d(build_ctx, beta, x, channels); + ggml_tensor * out = ggml_add( + build_ctx.ggml, + ggml_mul(build_ctx.ggml, normalized, gamma_rep.tensor), + beta_rep.tensor); + return ggml_mul(build_ctx.ggml, out, keep); +} + +ggml_tensor * build_adain_ct( + core::ModuleBuildContext & build_ctx, + ggml_tensor * x, + const core::TensorValue & gamma, + const core::TensorValue & beta, + int64_t channels, + float eps) { + ggml_tensor * x_for_stats = ggml_cont(build_ctx.ggml, x); + ggml_tensor * mean = ggml_mean(build_ctx.ggml, x_for_stats); + ggml_tensor * centered = ggml_sub(build_ctx.ggml, x, ggml_repeat(build_ctx.ggml, mean, x)); + ggml_tensor * squared = ggml_mul(build_ctx.ggml, centered, centered); + ggml_tensor * variance = ggml_mean(build_ctx.ggml, ggml_cont(build_ctx.ggml, squared)); + ggml_tensor * stddev = ggml_sqrt(build_ctx.ggml, ggml_scale_bias(build_ctx.ggml, variance, 1.0f, eps)); + ggml_tensor * normalized = ggml_div(build_ctx.ggml, centered, ggml_repeat(build_ctx.ggml, stddev, x)); + const auto gamma_rep = broadcast_channel_2d(build_ctx, gamma, x, channels); + const auto beta_rep = broadcast_channel_2d(build_ctx, beta, x, channels); + return ggml_add( + build_ctx.ggml, + ggml_mul(build_ctx.ggml, normalized, gamma_rep.tensor), + beta_rep.tensor); +} + +void upload_time_masks( + const std::vector & masks, + int64_t valid_base_frames, + int64_t base_frame_capacity) { + if (valid_base_frames <= 0 || valid_base_frames > base_frame_capacity) { + throw std::runtime_error("kokoro predictor valid frame count exceeds prepared capacity"); + } + for (const TimeMaskInputs & mask : masks) { + if (mask.frame_capacity <= 0 || + (valid_base_frames * mask.frame_capacity) % base_frame_capacity != 0) { + throw std::runtime_error("kokoro predictor mask capacity is not aligned with graph capacity"); + } + const int64_t valid_frames = (valid_base_frames * mask.frame_capacity) / base_frame_capacity; + if (valid_frames <= 0 || valid_frames > mask.frame_capacity) { + throw std::runtime_error("kokoro predictor mask valid frame count is invalid"); + } + std::vector keep(static_cast(mask.frame_capacity), 0.0f); + std::vector norm(static_cast(mask.frame_capacity), 0.0f); + std::fill(keep.begin(), keep.begin() + valid_frames, 1.0f); + const float norm_value = static_cast(mask.frame_capacity) / static_cast(valid_frames); + std::fill(norm.begin(), norm.begin() + valid_frames, norm_value); + ggml_backend_tensor_set(mask.keep, keep.data(), 0, ggml_nbytes(mask.keep)); + if (mask.norm != nullptr) { + ggml_backend_tensor_set(mask.norm, norm.data(), 0, ggml_nbytes(mask.norm)); + } + } +} + +std::vector expand_tc_by_durations( + const std::vector & values, + int64_t rows, + int64_t cols, + const std::vector & durations, + int64_t total_rows) { + if (rows != static_cast(durations.size())) { + throw std::runtime_error("kokoro predictor tc expansion shape mismatch"); + } + std::vector out(static_cast(total_rows * cols), 0.0f); + int64_t dst_row = 0; + for (int64_t src_row = 0; src_row < rows; ++src_row) { + const int32_t repeat = durations[static_cast(src_row)]; + const float * src = values.data() + static_cast(src_row * cols); + for (int32_t i = 0; i < repeat; ++i) { + std::memcpy( + out.data() + static_cast(dst_row * cols), + src, + static_cast(cols) * sizeof(float)); + ++dst_row; + } + } + if (dst_row != total_rows) { + throw std::runtime_error("kokoro predictor tc expansion produced unexpected frame count"); + } + return out; +} + +std::vector expand_ct_by_durations( + const std::vector & values, + int64_t rows, + int64_t cols, + const std::vector & durations, + int64_t total_cols) { + if (cols != static_cast(durations.size())) { + throw std::runtime_error("kokoro predictor ct expansion shape mismatch"); + } + std::vector out(static_cast(rows * total_cols), 0.0f); + for (int64_t row = 0; row < rows; ++row) { + const float * src_row = values.data() + static_cast(row * cols); + float * dst_row = out.data() + static_cast(row * total_cols); + int64_t dst_col = 0; + for (int64_t src_col = 0; src_col < cols; ++src_col) { + const float value = src_row[src_col]; + const int32_t repeat = durations[static_cast(src_col)]; + for (int32_t i = 0; i < repeat; ++i) { + dst_row[dst_col++] = value; + } + } + if (dst_col != total_cols) { + throw std::runtime_error("kokoro predictor ct expansion produced unexpected frame count"); + } + } + return out; +} + +core::TensorValue make_zero_lstm_state( + core::ModuleBuildContext & ctx, + int64_t hidden_size, + std::vector & zero_state_inputs) { + auto tensor = core::make_tensor(ctx, GGML_TYPE_F32, core::TensorShape::from_dims({1, hidden_size})); + ggml_set_input(tensor.tensor); + zero_state_inputs.push_back(tensor.tensor); + return tensor; +} + +void upload_zero_state_inputs(const std::vector & zero_state_inputs) { + for (ggml_tensor * tensor : zero_state_inputs) { + std::vector zeros(static_cast(ggml_nelements(tensor)), 0.0f); + ggml_backend_tensor_set(tensor, zeros.data(), 0, ggml_nbytes(tensor)); + } +} + +modules::LSTMSequenceOutputs build_lstm_sequence_combined_input_bias( + core::ModuleBuildContext & ctx, + const core::TensorValue & input, + const KokoroWeights::LstmWeights & lstm, + bool reverse, + const core::TensorValue & initial_hidden, + const core::TensorValue & initial_cell, + const core::TensorValue * valid_mask, + bool preserve_inactive_state) { + core::validate_shape(input, core::TensorShape::from_dims({input.shape.dims[0], lstm.input_size}), "input"); + if (valid_mask != nullptr) { + core::validate_shape(*valid_mask, core::TensorShape::from_dims({input.shape.dims[0], 1}), "valid_mask"); + } + const int64_t frames = input.shape.dims[0]; + const int64_t gates = 4 * lstm.hidden_size; + const auto & weight_ih = reverse ? lstm.weight_ih_l0_reverse : lstm.weight_ih_l0; + const auto & weight_hh = reverse ? lstm.weight_hh_l0_reverse : lstm.weight_hh_l0; + const auto & combined_bias = reverse ? lstm.combined_bias_l0_reverse : lstm.combined_bias_l0; + const auto projected_inputs = modules::LinearModule({lstm.input_size, gates, true}).build( + ctx, + input, + { + weight_ih, + combined_bias, + }); + + std::vector steps(static_cast(frames)); + auto hidden = initial_hidden; + auto cell = initial_cell; + const modules::ConcatModule concat_rows({0}); + + for (int64_t step = 0; step < frames; ++step) { + const int64_t t = reverse ? (frames - 1 - step) : step; + const auto previous_hidden = hidden; + const auto previous_cell = cell; + const auto projected_x_t = modules::SliceModule({0, t, 1}).build(ctx, projected_inputs); + const auto projected_hidden = modules::LinearModule({lstm.hidden_size, gates, false}).build( + ctx, + hidden, + { + weight_hh, + core::TensorValue{}, + }); + const auto gate_values = core::wrap_tensor( + ggml_add(ctx.ggml, projected_x_t.tensor, projected_hidden.tensor), + projected_x_t.shape, + GGML_TYPE_F32); + + const auto input_gate = modules::SigmoidModule().build(ctx, modules::SliceModule({1, 0, lstm.hidden_size}).build(ctx, gate_values)); + const auto forget_gate = modules::SigmoidModule().build(ctx, modules::SliceModule({1, lstm.hidden_size, lstm.hidden_size}).build(ctx, gate_values)); + const auto candidate = modules::TanhModule().build(ctx, modules::SliceModule({1, 2 * lstm.hidden_size, lstm.hidden_size}).build(ctx, gate_values)); + const auto output_gate = modules::SigmoidModule().build(ctx, modules::SliceModule({1, 3 * lstm.hidden_size, lstm.hidden_size}).build(ctx, gate_values)); + const auto kept_cell = core::wrap_tensor(ggml_mul(ctx.ggml, forget_gate.tensor, cell.tensor), cell.shape, GGML_TYPE_F32); + const auto written_cell = core::wrap_tensor(ggml_mul(ctx.ggml, input_gate.tensor, candidate.tensor), cell.shape, GGML_TYPE_F32); + cell = core::wrap_tensor(ggml_add(ctx.ggml, kept_cell.tensor, written_cell.tensor), cell.shape, GGML_TYPE_F32); + const auto activated_cell = modules::TanhModule().build(ctx, cell); + hidden = core::wrap_tensor(ggml_mul(ctx.ggml, output_gate.tensor, activated_cell.tensor), hidden.shape, GGML_TYPE_F32); + if (valid_mask != nullptr) { + const auto active = modules::SliceModule({0, t, 1}).build(ctx, *valid_mask); + const auto active_hidden = core::wrap_tensor( + ggml_repeat(ctx.ggml, active.tensor, hidden.tensor), + hidden.shape, + GGML_TYPE_F32); + if (preserve_inactive_state) { + const auto inactive_hidden = core::wrap_tensor( + ggml_scale_bias(ctx.ggml, active_hidden.tensor, -1.0f, 1.0f), + hidden.shape, + GGML_TYPE_F32); + const auto kept_hidden = core::wrap_tensor( + ggml_mul(ctx.ggml, previous_hidden.tensor, inactive_hidden.tensor), + hidden.shape, + GGML_TYPE_F32); + const auto written_hidden = core::wrap_tensor( + ggml_mul(ctx.ggml, hidden.tensor, active_hidden.tensor), + hidden.shape, + GGML_TYPE_F32); + hidden = core::wrap_tensor(ggml_add(ctx.ggml, written_hidden.tensor, kept_hidden.tensor), hidden.shape, GGML_TYPE_F32); + + const auto active_cell = core::wrap_tensor( + ggml_repeat(ctx.ggml, active.tensor, cell.tensor), + cell.shape, + GGML_TYPE_F32); + const auto inactive_cell = core::wrap_tensor( + ggml_scale_bias(ctx.ggml, active_cell.tensor, -1.0f, 1.0f), + cell.shape, + GGML_TYPE_F32); + const auto kept_cell = core::wrap_tensor( + ggml_mul(ctx.ggml, previous_cell.tensor, inactive_cell.tensor), + cell.shape, + GGML_TYPE_F32); + const auto written_cell = core::wrap_tensor( + ggml_mul(ctx.ggml, cell.tensor, active_cell.tensor), + cell.shape, + GGML_TYPE_F32); + cell = core::wrap_tensor(ggml_add(ctx.ggml, written_cell.tensor, kept_cell.tensor), cell.shape, GGML_TYPE_F32); + } else { + hidden = core::wrap_tensor( + ggml_mul(ctx.ggml, hidden.tensor, active_hidden.tensor), + hidden.shape, + GGML_TYPE_F32); + cell = core::wrap_tensor( + ggml_mul(ctx.ggml, cell.tensor, ggml_repeat(ctx.ggml, active.tensor, cell.tensor)), + cell.shape, + GGML_TYPE_F32); + } + } + steps[static_cast(t)] = hidden; + } + + auto sequence = steps[0]; + for (int64_t t = 1; t < frames; ++t) { + sequence = concat_rows.build(ctx, sequence, steps[static_cast(t)]); + } + return modules::LSTMSequenceOutputs{sequence, hidden, cell}; +} + +modules::LSTMSequenceOutputs build_lstm_sequence_combined_input_bias( + core::ModuleBuildContext & ctx, + const core::TensorValue & input, + const KokoroWeights::LstmWeights & lstm, + bool reverse, + const core::TensorValue * valid_mask, + std::vector & zero_state_inputs) { + return build_lstm_sequence_combined_input_bias( + ctx, + input, + lstm, + reverse, + make_zero_lstm_state(ctx, lstm.hidden_size, zero_state_inputs), + make_zero_lstm_state(ctx, lstm.hidden_size, zero_state_inputs), + valid_mask, + false); +} + +modules::BidirectionalLSTMOutputs build_bidirectional_lstm_combined_input_bias( + core::ModuleBuildContext & ctx, + const core::TensorValue & input, + const KokoroWeights::LstmWeights & lstm, + const core::TensorValue * valid_mask, + std::vector & zero_state_inputs) { + const auto forward = build_lstm_sequence_combined_input_bias(ctx, input, lstm, false, valid_mask, zero_state_inputs); + const auto reverse = build_lstm_sequence_combined_input_bias(ctx, input, lstm, true, valid_mask, zero_state_inputs); + const auto sequence = modules::ConcatModule({1}).build(ctx, forward.sequence, reverse.sequence); + return {sequence, forward.hidden, forward.cell, reverse.hidden, reverse.cell}; +} + +core::TensorValue apply_time_mask( + core::ModuleBuildContext & ctx, + const core::TensorValue & value, + const core::TensorValue & mask) { + return core::wrap_tensor( + ggml_mul(ctx.ggml, value.tensor, ggml_repeat(ctx.ggml, mask.tensor, value.tensor)), + value.shape, + GGML_TYPE_F32); +} + +template +modules::Conv1dWeights make_conv1d_weights(core::ModuleBuildContext & ctx, const ConvWeightsT & conv) { + modules::Conv1dWeights weights = {}; + (void) ctx; + weights.weight = conv.weight; + if (conv.use_bias) { + weights.bias = conv.bias; + } + return weights; +} + +modules::ConvTranspose1dWeights make_conv_transpose1d_weights( + core::ModuleBuildContext & ctx, + const KokoroWeights::WeightNormConvTranspose1dWeights & conv) { + modules::ConvTranspose1dWeights weights = {}; + (void) ctx; + weights.weight = conv.groups == 1 ? conv.weight : conv.dense_weight; + if (conv.use_bias) { + weights.bias = conv.bias; + } + return weights; +} + +template +ggml_tensor * build_standard_conv1d_ct_predictor(ggml_context * ctx, ggml_tensor * input, const ConvWeightsT & conv) { + if (conv.groups != 1) { + throw std::runtime_error("kokoro predictor conv1d requires groups == 1"); + } + core::ModuleBuildContext build_ctx = {}; + build_ctx.ggml = ctx; + const auto input_ct = core::wrap_tensor( + ggml_cont(ctx, input), + core::TensorShape::from_dims({input->ne[1], input->ne[0]}), + GGML_TYPE_F32); + const auto input_bct = modules::ReshapeModule({ + core::TensorShape::from_dims({1, conv.in_channels, input->ne[0]})}) + .build(build_ctx, input_ct); + const auto output_bct = modules::Conv1dModule({ + conv.in_channels, + conv.out_channels, + conv.kernel, + static_cast(conv.stride), + static_cast(conv.padding), + static_cast(conv.dilation), + conv.use_bias}).build(build_ctx, input_bct, make_conv1d_weights(build_ctx, conv)); + return modules::ReshapeModule({ + core::TensorShape::from_dims({conv.out_channels, output_bct.shape.dims[2]})}) + .build(build_ctx, output_bct) + .tensor; +} + +ggml_tensor * build_grouped_conv_transpose1d_ct_predictor( + ggml_context * ctx, + ggml_tensor * input, + const KokoroWeights::WeightNormConvTranspose1dWeights & conv) { + const int64_t grouped_in = conv.in_channels / conv.groups; + const int64_t grouped_out = conv.out_channels / conv.groups; + if (conv.groups == conv.in_channels && conv.in_channels == conv.out_channels && grouped_in == 1 && grouped_out == 1) { + core::ModuleBuildContext build_ctx = {}; + build_ctx.ggml = ctx; + const auto input_ct = core::wrap_tensor( + ggml_cont(ctx, input), + core::TensorShape::from_dims({input->ne[1], input->ne[0]}), + GGML_TYPE_F32); + const auto input_bct = modules::ReshapeModule({ + core::TensorShape::from_dims({1, conv.in_channels, input->ne[0]})}) + .build(build_ctx, input_ct); + modules::ConvTranspose1dWeights weights = {}; + weights.weight = conv.dense_weight; + if (conv.use_bias) { + weights.bias = conv.bias; + } + auto output_bct = modules::ConvTranspose1dModule({ + conv.in_channels, + conv.out_channels, + conv.kernel, + static_cast(conv.stride), + 0, + 1, + conv.use_bias}).build(build_ctx, input_bct, weights); + const int64_t cropped_len = + (input->ne[0] - 1) * conv.stride - 2 * conv.padding + conv.kernel + conv.output_padding; + const auto cropped = core::wrap_tensor( + ggml_cont( + ctx, + ggml_view_3d( + ctx, + output_bct.tensor, + cropped_len, + conv.out_channels, + 1, + output_bct.tensor->nb[1], + output_bct.tensor->nb[2], + static_cast(conv.padding) * sizeof(float))), + core::TensorShape::from_dims({1, conv.out_channels, cropped_len}), + GGML_TYPE_F32); + return modules::ReshapeModule({ + core::TensorShape::from_dims({conv.out_channels, cropped.shape.dims[2]})}) + .build(build_ctx, cropped) + .tensor; + } + throw std::runtime_error("kokoro predictor grouped conv_transpose1d supports only depthwise layout"); +} + +ggml_tensor * build_conv_transpose1d_ct_predictor( + ggml_context * ctx, + ggml_tensor * input, + const KokoroWeights::WeightNormConvTranspose1dWeights & conv) { + if (conv.groups > 1) { + return build_grouped_conv_transpose1d_ct_predictor(ctx, input, conv); + } + core::ModuleBuildContext build_ctx = {}; + build_ctx.ggml = ctx; + const auto input_ct = core::wrap_tensor( + ggml_cont(ctx, input), + core::TensorShape::from_dims({input->ne[1], input->ne[0]}), + GGML_TYPE_F32); + const auto input_bct = modules::ReshapeModule({ + core::TensorShape::from_dims({1, conv.in_channels, input->ne[0]})}) + .build(build_ctx, input_ct); + auto output_bct = modules::ConvTranspose1dModule({ + conv.in_channels, + conv.out_channels, + conv.kernel, + static_cast(conv.stride), + 0, + 1, + conv.use_bias}).build(build_ctx, input_bct, make_conv_transpose1d_weights(build_ctx, conv)); + const int64_t cropped_len = + (input->ne[0] - 1) * conv.stride - 2 * conv.padding + conv.kernel + conv.output_padding; + const auto cropped = core::wrap_tensor( + ggml_cont( + ctx, + ggml_view_3d( + ctx, + output_bct.tensor, + cropped_len, + conv.out_channels, + 1, + output_bct.tensor->nb[1], + output_bct.tensor->nb[2], + static_cast(conv.padding) * sizeof(float))), + core::TensorShape::from_dims({1, conv.out_channels, cropped_len}), + GGML_TYPE_F32); + return modules::ReshapeModule({ + core::TensorShape::from_dims({conv.out_channels, cropped.shape.dims[2]})}) + .build(build_ctx, cropped) + .tensor; +} + +ggml_tensor * build_adaptive_instance_norm_ct_predictor( + ggml_context * ctx, + ggml_tensor * x, + const KokoroWeights::AdaIn1dWeights & weights, + const core::TensorValue & style, + std::vector & masks, + bool use_time_masks) { + const int64_t channels = x->ne[1]; + core::ModuleBuildContext build_ctx = {}; + build_ctx.ggml = ctx; + const auto affine = modules::LinearModule({ + weights.fc.in_features, + weights.fc.out_features, + weights.fc.use_bias}).build( + build_ctx, + style, + { + weights.fc.weight, + weights.fc.bias, + }); + const auto scale_delta = view_2d_last_dim_slice(build_ctx, affine, 0, channels); + const auto shift = view_2d_last_dim_slice(build_ctx, affine, channels, channels); + const auto gamma_2d = core::wrap_tensor( + ggml_scale_bias(ctx, scale_delta.tensor, 1.0f, 1.0f), + scale_delta.shape, + GGML_TYPE_F32); + const auto gamma = core::wrap_tensor( + ggml_reshape_1d(ctx, gamma_2d.tensor, channels), + core::TensorShape::from_dims({channels}), + GGML_TYPE_F32); + const auto beta = core::wrap_tensor( + ggml_reshape_1d(ctx, shift.tensor, channels), + core::TensorShape::from_dims({channels}), + GGML_TYPE_F32); + if (use_time_masks) { + return build_masked_adain_ct(build_ctx, x, gamma, beta, channels, weights.eps, masks); + } + return build_adain_ct(build_ctx, x, gamma, beta, channels, weights.eps); +} + +ggml_tensor * build_adain_resblock_ct_predictor( + ggml_context * ctx, + ggml_tensor * x, + const KokoroWeights::AdainResBlock1dWeights & block, + const core::TensorValue & style_predictor, + const core::TensorValue & style_decoder, + bool use_decoder_style, + std::vector & masks, + bool use_time_masks) { + const auto & style = use_decoder_style ? style_decoder : style_predictor; + ggml_tensor * shortcut = x; + if (block.upsample) { + core::ModuleBuildContext build_ctx = {}; + build_ctx.ggml = ctx; + const auto input = core::wrap_tensor( + x, + core::TensorShape::from_dims({x->ne[1], x->ne[0]}), + GGML_TYPE_F32); + const auto upsampled = modules::Interpolate1dModule({x->ne[0] * 2, modules::Interpolate1dMode::Nearest}).build(build_ctx, input); + shortcut = upsampled.tensor; + } + if (block.learned_sc) { + shortcut = build_standard_conv1d_ct_predictor(ctx, shortcut, block.conv1x1); + } + + ggml_tensor * residual = build_adaptive_instance_norm_ct_predictor(ctx, x, block.norm1, style, masks, use_time_masks); + residual = ggml_leaky_relu(ctx, residual, 0.2f, false); + if (block.use_pool) { + residual = build_conv_transpose1d_ct_predictor(ctx, residual, block.pool); + } + residual = build_standard_conv1d_ct_predictor(ctx, residual, block.conv1); + residual = build_adaptive_instance_norm_ct_predictor( + ctx, + residual, + block.norm2, + style, + masks, + use_time_masks); + residual = ggml_leaky_relu(ctx, residual, 0.2f, false); + residual = build_standard_conv1d_ct_predictor(ctx, residual, block.conv2); + ggml_tensor * out = ggml_scale(ctx, ggml_add(ctx, residual, shortcut), 0.7071067811865475f); + if (!use_time_masks) { + return out; + } + const TimeMaskInputs & mask = add_keep_mask_input(ctx, masks, out->ne[0]); + return ggml_mul(ctx, out, repeat_mask_like(ctx, mask.keep, out)); +} + +class PlbertGraphRuntime { +public: + PlbertGraphRuntime( + std::shared_ptr weights, + ggml_backend_t backend, + int n_threads, + bool use_device_backend, + int64_t fixed_token_capacity) + : runtime_(std::move(weights), backend, n_threads, use_device_backend, fixed_token_capacity) {} + + std::vector run(const std::vector & input_ids) { + return runtime_.encode(input_ids, true); + } + +private: + KokoroPlbertRuntime runtime_; +}; + +struct PredictorPreTailOutputs { + std::vector durations; + std::vector expanded_encoder_tc; + int64_t expanded_encoder_rows = 0; + int64_t expanded_encoder_cols = 0; + std::vector asr_ct; + int64_t asr_rows = 0; + int64_t asr_cols = 0; +}; + +class PredictorPreTailGraphRuntime { +public: + PredictorPreTailGraphRuntime( + const KokoroWeights & weights, + ggml_backend_t backend, + int n_threads, + bool use_device_backend, + KokoroPredictorGraphConfig graph_config) + : weights_(&weights), + backend_(backend), + duration_backend_(backend), + text_backend_(backend), + n_threads_(std::max(1, n_threads)), + graph_config_(graph_config) { + (void)use_device_backend; + } + + ~PredictorPreTailGraphRuntime() = default; + + PredictorPreTailOutputs run( + const std::vector & plbert_hidden_tc, + const std::vector & input_ids, + const std::vector & style_predictor, + float speed, + int64_t token_capacity) { + const int64_t valid_token_count = static_cast(input_ids.size()); + if (valid_token_count <= 0 || valid_token_count > token_capacity) { + throw std::runtime_error("kokoro predictor pre-tail input length exceeds prepared capacity"); + } + auto & prepared = session_for_capacity(token_capacity, token_capacity != valid_token_count); + const int64_t graph_token_count = prepared.token_count; + const int64_t predictor_feature_cols = weights_->predictor.lstm.input_size; + const int64_t text_feature_rows = weights_->text_encoder.lstm.hidden_size * 2; + + std::vector padded_hidden; + const std::vector * hidden_for_graph = &plbert_hidden_tc; + if (valid_token_count != graph_token_count) { + padded_hidden.assign(static_cast(graph_token_count * weights_->hidden_dim), 0.0f); + std::memcpy( + padded_hidden.data(), + plbert_hidden_tc.data(), + static_cast(valid_token_count * weights_->hidden_dim) * sizeof(float)); + hidden_for_graph = &padded_hidden; + } + std::vector padded_ids; + const std::vector * ids_for_graph = &input_ids; + if (valid_token_count != graph_token_count) { + padded_ids.assign(static_cast(graph_token_count), 0); + std::memcpy( + padded_ids.data(), + input_ids.data(), + static_cast(valid_token_count) * sizeof(int32_t)); + ids_for_graph = &padded_ids; + } + + PredictorPreTailOutputs out = {}; + std::vector duration_features_tc; + auto & duration = duration_session(prepared); + const double duration_ms = measure_ms([&]() { + duration.run(*hidden_for_graph, valid_token_count, style_predictor, speed, out.durations, duration_features_tc); + }); + engine::debug::timing_log_scalar("kokoro.predictor.duration_compute_ms", duration_ms); + auto & text = text_session(prepared); + std::vector text_features_ct; + const double text_ms = measure_ms([&]() { + text_features_ct = text.run(*ids_for_graph, valid_token_count); + }); + engine::debug::timing_log_scalar("kokoro.predictor.text_compute_ms", text_ms); + if (valid_token_count != graph_token_count) { + out.durations.resize(static_cast(valid_token_count)); + std::vector valid_duration_features(static_cast(valid_token_count * predictor_feature_cols), 0.0f); + for (int64_t row = 0; row < valid_token_count; ++row) { + std::memcpy( + valid_duration_features.data() + static_cast(row * predictor_feature_cols), + duration_features_tc.data() + static_cast(row * predictor_feature_cols), + static_cast(predictor_feature_cols) * sizeof(float)); + } + duration_features_tc = std::move(valid_duration_features); + std::vector valid_text_features(static_cast(text_feature_rows * valid_token_count), 0.0f); + for (int64_t row = 0; row < text_feature_rows; ++row) { + std::memcpy( + valid_text_features.data() + static_cast(row * valid_token_count), + text_features_ct.data() + static_cast(row * graph_token_count), + static_cast(valid_token_count) * sizeof(float)); + } + text_features_ct = std::move(valid_text_features); + } + + int64_t total_frames = 0; + for (int32_t duration : out.durations) { + total_frames += duration; + } + const int64_t max_output_frames = valid_token_count * weights_->max_dur; + if (total_frames > max_output_frames) { + throw std::runtime_error("kokoro predictor produced durations beyond prepared capacity"); + } + + const double expand_ms = measure_ms([&]() { + out.expanded_encoder_tc = expand_tc_by_durations( + duration_features_tc, + valid_token_count, + weights_->predictor.lstm.input_size, + out.durations, + total_frames); + out.asr_ct = expand_ct_by_durations( + text_features_ct, + text_feature_rows, + valid_token_count, + out.durations, + total_frames); + }); + engine::debug::timing_log_scalar("kokoro.predictor.expand_ms", expand_ms); + out.expanded_encoder_rows = total_frames; + out.expanded_encoder_cols = predictor_feature_cols; + out.asr_rows = text_feature_rows; + out.asr_cols = total_frames; + return out; + } + + int64_t token_count() const { + return session_.token_count; + } + + void prepare_capacity(int64_t token_count) { + auto & prepared = session_for_capacity(token_count, true); + (void) duration_session(prepared); + (void) text_session(prepared); + } + +private: + struct DurationSession { + const KokoroWeights * weights = nullptr; + int64_t token_count = 0; + bool use_valid_mask = false; + int n_threads = 1; + KokoroPredictorGraphConfig graph_config = {}; + ggml_backend_t backend = nullptr; + ggml_context * ctx = nullptr; + ggml_cgraph * graph = nullptr; + ggml_backend_buffer_t buffer = nullptr; + ggml_backend_graph_plan_t plan = nullptr; + core::TensorValue plbert_hidden_tc = {}; + core::TensorValue valid_mask_tc = {}; + core::TensorValue style_predictor = {}; + core::TensorValue speed = {}; + core::TensorValue durations_clamped = {}; + core::TensorValue duration_features_tc = {}; + + DurationSession( + const KokoroWeights & weights_in, + ggml_backend_t backend_in, + int64_t token_count_in, + bool use_valid_mask_in, + int n_threads_in, + KokoroPredictorGraphConfig graph_config_in) + : weights(&weights_in), + token_count(token_count_in), + use_valid_mask(use_valid_mask_in), + n_threads(n_threads_in), + graph_config(graph_config_in), + backend(backend_in) { + ggml_init_params params = {}; + params.mem_size = graph_config.duration_graph_bytes; + params.mem_buffer = nullptr; + params.no_alloc = true; + ctx = ggml_init(params); + if (!ctx) { + throw std::runtime_error("failed to initialize ggml context for Kokoro predictor pre-tail"); + } + + try { + std::vector zero_state_inputs; + + core::ModuleBuildContext build_ctx = {}; + build_ctx.ggml = ctx; + build_ctx.module_instance_name = "kokoro_predictor_pretail"; + + plbert_hidden_tc = core::make_tensor( + build_ctx, + GGML_TYPE_F32, + core::TensorShape::from_dims({token_count, weights->hidden_dim})); + if (use_valid_mask) { + valid_mask_tc = core::make_tensor( + build_ctx, + GGML_TYPE_F32, + core::TensorShape::from_dims({token_count, 1})); + ggml_set_input(valid_mask_tc.tensor); + } + style_predictor = core::make_tensor( + build_ctx, + GGML_TYPE_F32, + core::TensorShape::from_dims({1, weights->style_dim})); + speed = core::make_tensor(build_ctx, GGML_TYPE_F32, core::TensorShape::from_dims({1})); + ggml_set_input(plbert_hidden_tc.tensor); + ggml_set_input(style_predictor.tensor); + ggml_set_input(speed.tensor); + + const auto concat_with_style = [&](const core::TensorValue & x_tc) { + const auto repeated_style = modules::RepeatModule({ + core::TensorShape::from_dims({x_tc.shape.dims[0], style_predictor.shape.dims[1]})}) + .build(build_ctx, style_predictor); + return modules::ConcatModule({1}).build(build_ctx, x_tc, repeated_style); + }; + + const auto apply_adaptive_layer_norm = [&](const core::TensorValue & x_tc, const KokoroWeights::AdaLayerNormWeights & ada_weights) { + const int64_t hidden = ada_weights.fc.out_features / 2; + const int64_t cond = ada_weights.fc.in_features; + const auto normed = modules::LayerNormModule({hidden, ada_weights.eps, false, false}).build(build_ctx, x_tc, {}); + const auto affine = modules::LinearModule({cond, hidden * 2, ada_weights.fc.use_bias}).build( + build_ctx, + style_predictor, + {ada_weights.fc.weight, ada_weights.fc.bias}); + const auto scale_delta = view_2d_last_dim_slice(build_ctx, affine, 0, hidden); + const auto shift_vec = view_2d_last_dim_slice(build_ctx, affine, hidden, hidden); + const auto scale_vec = + core::wrap_tensor(ggml_scale_bias(ctx, scale_delta.tensor, 1.0f, 1.0f), scale_delta.shape, GGML_TYPE_F32); + const auto scale = modules::RepeatModule({core::TensorShape::from_dims({x_tc.shape.dims[0], hidden})}) + .build(build_ctx, scale_vec); + const auto shift = modules::RepeatModule({core::TensorShape::from_dims({x_tc.shape.dims[0], hidden})}) + .build(build_ctx, shift_vec); + const auto scaled = core::wrap_tensor(ggml_mul(ctx, normed.tensor, scale.tensor), normed.shape, GGML_TYPE_F32); + return core::wrap_tensor(ggml_add(ctx, scaled.tensor, shift.tensor), normed.shape, GGML_TYPE_F32); + }; + + const auto build_duration_bidir_lstm = [&](const core::TensorValue & input, + const KokoroWeights::LstmWeights & lstm) { + return build_bidirectional_lstm_combined_input_bias( + build_ctx, + input, + lstm, + use_valid_mask ? &valid_mask_tc : nullptr, + zero_state_inputs); + }; + + duration_features_tc = concat_with_style(plbert_hidden_tc); + for (size_t i = 0; i < weights->predictor.duration_encoder.lstms.size(); ++i) { + const auto & lstm = weights->predictor.duration_encoder.lstms[i]; + const auto lstm_outs = build_duration_bidir_lstm(duration_features_tc, lstm); + duration_features_tc = concat_with_style( + apply_adaptive_layer_norm(lstm_outs.sequence, weights->predictor.duration_encoder.ada_layer_norms[i])); + } + + const auto predictor_hidden_tc = build_duration_bidir_lstm(duration_features_tc, weights->predictor.lstm) + .sequence; + + auto logits = modules::LinearModule({ + weights->predictor.duration_proj.in_features, + weights->predictor.duration_proj.out_features, + weights->predictor.duration_proj.use_bias}).build( + build_ctx, + predictor_hidden_tc, + { + weights->predictor.duration_proj.weight, + weights->predictor.duration_proj.use_bias + ? *weights->predictor.duration_proj.bias + : core::TensorValue{}, + }); + + auto sig = modules::SigmoidModule().build(build_ctx, logits); + auto flat = core::reshape_tensor(build_ctx, sig, core::TensorShape::from_dims({token_count, weights->max_dur})); + auto summed = core::wrap_tensor(ggml_sum_rows(ctx, flat.tensor), core::TensorShape::from_dims({token_count, 1}), GGML_TYPE_F32); + auto scaled = core::wrap_tensor(ggml_div(ctx, summed.tensor, speed.tensor), summed.shape, GGML_TYPE_F32); + auto rounded = core::wrap_tensor(ggml_round(ctx, scaled.tensor), scaled.shape, GGML_TYPE_F32); + durations_clamped = core::wrap_tensor( + ggml_clamp(ctx, rounded.tensor, 1.0f, static_cast(weights->max_dur)), + rounded.shape, + GGML_TYPE_F32); + + graph = ggml_new_graph_custom(ctx, graph_config.graph_node_capacity, false); + ggml_build_forward_expand(graph, durations_clamped.tensor); + ggml_build_forward_expand(graph, duration_features_tc.tensor); + + core::set_backend_threads(backend, n_threads); + const double alloc_ms = measure_ms([&]() { + buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + }); + if (!buffer) { + throw std::runtime_error("failed to allocate Kokoro predictor pre-tail tensors"); + } + if (engine::core::uses_host_graph_plan(backend)) { + plan = engine::core::create_backend_graph_plan_if_host(backend, graph); + if (!plan) { + throw std::runtime_error("failed to create Kokoro predictor pre-tail plan"); + } + } + const double zero_state_upload_ms = measure_ms([&]() { upload_zero_state_inputs(zero_state_inputs); }); + const double input_init_ms = measure_ms([&]() { + std::vector hidden(static_cast(token_count * weights->hidden_dim), 0.0f); + std::vector style(static_cast(weights->style_dim), 0.0f); + float speed_value = 1.0f; + core::write_tensor_f32(plbert_hidden_tc, hidden); + if (use_valid_mask) { + std::vector valid_mask(static_cast(token_count), 1.0f); + core::write_tensor_f32(valid_mask_tc, valid_mask); + } + core::write_tensor_f32(style_predictor, style); + core::write_tensor_f32(speed, &speed_value, 1); + }); + engine::debug::timing_log_scalar("kokoro.graph.build.predictor_duration_alloc_ms", alloc_ms); + engine::debug::timing_log_scalar( + "kokoro.graph.build.predictor_duration_zero_state_upload_ms", + zero_state_upload_ms); + engine::debug::timing_log_scalar("kokoro.graph.build.predictor_duration_input_init_ms", input_init_ms); + } catch (...) { + if (buffer) { + ggml_backend_buffer_free(buffer); + } + if (ctx) { + ggml_free(ctx); + } + ctx = nullptr; + throw; + } + } + + ~DurationSession() { + if (plan) { + engine::core::free_backend_graph_plan(backend, plan); + } + if (buffer) { + ggml_backend_buffer_free(buffer); + } + if (ctx) { + ggml_free(ctx); + } + } + + void run( + const std::vector & plbert_hidden_tc_host, + int64_t valid_token_count, + const std::vector & style_predictor_host, + float speed_host, + std::vector & durations_out, + std::vector & duration_features_out) { + if (valid_token_count <= 0 || valid_token_count > token_count) { + throw std::runtime_error("kokoro predictor duration valid token count exceeds prepared capacity"); + } + const double input_upload_ms = measure_ms([&]() { + core::write_tensor_f32(plbert_hidden_tc, plbert_hidden_tc_host); + if (use_valid_mask) { + std::vector valid_mask(static_cast(token_count), 0.0f); + std::fill(valid_mask.begin(), valid_mask.begin() + valid_token_count, 1.0f); + core::write_tensor_f32(valid_mask_tc, valid_mask); + } + core::write_tensor_f32(style_predictor, style_predictor_host); + core::write_tensor_f32(speed, &speed_host, 1); + }); + core::set_backend_threads(backend, n_threads); + ggml_status status = GGML_STATUS_SUCCESS; + const double graph_compute_ms = measure_ms([&]() { + status = core::compute_backend_graph(backend, graph, plan); + }); + if (status != GGML_STATUS_SUCCESS) { + throw std::runtime_error(std::string("kokoro predictor pre-tail compute failed: ") + ggml_status_to_string(status)); + } + + std::vector durations_f; + const double duration_read_ms = measure_ms([&]() { + durations_f = core::read_tensor_f32(durations_clamped.tensor); + }); + durations_out.resize(static_cast(token_count), 1); + for (int64_t i = 0; i < token_count; ++i) { + const int32_t duration = std::max(1, static_cast(std::lround(durations_f[static_cast(i)]))); + durations_out[static_cast(i)] = duration; + } + const double feature_read_ms = measure_ms([&]() { + duration_features_out = core::read_tensor_f32(duration_features_tc.tensor); + }); + engine::debug::timing_log_scalar("kokoro.predictor_duration.input_upload_ms", input_upload_ms); + engine::debug::timing_log_scalar("kokoro.predictor_duration.graph.compute_ms", graph_compute_ms); + engine::debug::timing_log_scalar("kokoro.predictor_duration.duration_read_ms", duration_read_ms); + engine::debug::timing_log_scalar("kokoro.predictor_duration.feature_read_ms", feature_read_ms); + } + }; + + struct TextSession { + const KokoroWeights * weights = nullptr; + int64_t token_count = 0; + bool use_valid_mask = false; + int n_threads = 1; + KokoroPredictorGraphConfig graph_config = {}; + ggml_backend_t backend = nullptr; + ggml_context * ctx = nullptr; + ggml_cgraph * graph = nullptr; + ggml_backend_buffer_t buffer = nullptr; + ggml_backend_graph_plan_t plan = nullptr; + core::TensorValue input_ids = {}; + core::TensorValue valid_mask_tc = {}; + core::TensorValue valid_mask_bct = {}; + core::TensorValue text_features_ct = {}; + + TextSession( + const KokoroWeights & weights_in, + ggml_backend_t backend_in, + int64_t token_count_in, + bool use_valid_mask_in, + int n_threads_in, + KokoroPredictorGraphConfig graph_config_in) + : weights(&weights_in), + token_count(token_count_in), + use_valid_mask(use_valid_mask_in), + n_threads(n_threads_in), + graph_config(graph_config_in), + backend(backend_in) { + ggml_init_params params = {}; + params.mem_size = graph_config.text_graph_bytes; + params.mem_buffer = nullptr; + params.no_alloc = true; + ctx = ggml_init(params); + if (!ctx) { + throw std::runtime_error("failed to initialize ggml context for Kokoro predictor text path"); + } + + try { + std::vector zero_state_inputs; + + core::ModuleBuildContext build_ctx = {}; + build_ctx.ggml = ctx; + build_ctx.module_instance_name = "kokoro_predictor_text"; + + input_ids = core::make_tensor( + build_ctx, + GGML_TYPE_I32, + core::TensorShape::from_dims({token_count})); + ggml_set_input(input_ids.tensor); + if (use_valid_mask) { + valid_mask_tc = core::make_tensor( + build_ctx, + GGML_TYPE_F32, + core::TensorShape::from_dims({token_count, 1})); + valid_mask_bct = core::make_tensor( + build_ctx, + GGML_TYPE_F32, + core::TensorShape::from_dims({1, 1, token_count})); + ggml_set_input(valid_mask_tc.tensor); + ggml_set_input(valid_mask_bct.tensor); + } + + auto text_x_tc = modules::EmbeddingModule({ + weights->text_encoder.embedding.num_embeddings, + weights->text_encoder.embedding.embedding_dim}).build(build_ctx, input_ids, weights->text_encoder.embedding.weight); + if (use_valid_mask) { + text_x_tc = apply_time_mask(build_ctx, text_x_tc, valid_mask_tc); + } + auto text_x_ct = modules::TransposeModule({{1, 0, 2, 3}, 2}).build(build_ctx, text_x_tc); + text_x_ct = core::wrap_tensor(ggml_cont(ctx, text_x_ct.tensor), text_x_ct.shape, GGML_TYPE_F32); + auto text_x_bct = modules::ReshapeModule({ + core::TensorShape::from_dims({1, text_x_ct.shape.dims[0], text_x_ct.shape.dims[1]})}) + .build(build_ctx, text_x_ct); + for (size_t i = 0; i < weights->text_encoder.cnn.size(); ++i) { + const auto & block = weights->text_encoder.cnn[i]; + text_x_bct = modules::Conv1dModule({ + block.conv.in_channels, + block.conv.out_channels, + block.conv.kernel, + static_cast(block.conv.stride), + static_cast(block.conv.padding), + static_cast(block.conv.dilation), + block.conv.use_bias}).build(build_ctx, text_x_bct, make_conv1d_weights(build_ctx, block.conv)); + auto norm_in = modules::TransposeModule({{0, 2, 1, 3}, 3}).build(build_ctx, text_x_bct); + auto normed = modules::LayerNormModule({ + block.layer_norm.channels, + block.layer_norm.eps, + true, + true}).build( + build_ctx, + norm_in, + { + block.layer_norm.weight, + block.layer_norm.bias, + }); + text_x_bct = modules::TransposeModule({{0, 2, 1, 3}, 3}).build(build_ctx, normed); + text_x_bct = core::wrap_tensor(ggml_cont(ctx, text_x_bct.tensor), text_x_bct.shape, GGML_TYPE_F32); + text_x_bct = core::wrap_tensor(ggml_leaky_relu(ctx, text_x_bct.tensor, 0.2f, false), text_x_bct.shape, GGML_TYPE_F32); + if (use_valid_mask) { + text_x_bct = apply_time_mask(build_ctx, text_x_bct, valid_mask_bct); + } + } + auto text_x_btc = modules::TransposeModule({{0, 2, 1, 3}, 3}).build(build_ctx, text_x_bct); + text_x_btc = core::wrap_tensor(ggml_cont(ctx, text_x_btc.tensor), text_x_btc.shape, GGML_TYPE_F32); + auto text_x_tc_for_lstm = modules::ReshapeModule({ + core::TensorShape::from_dims({token_count, text_x_btc.shape.dims[2]})}) + .build(build_ctx, text_x_btc); + const auto text_lstm_sequence = build_bidirectional_lstm_combined_input_bias( + build_ctx, + text_x_tc_for_lstm, + weights->text_encoder.lstm, + use_valid_mask ? &valid_mask_tc : nullptr, + zero_state_inputs).sequence; + text_features_ct = modules::TransposeModule({{1, 0, 2, 3}, 2}).build(build_ctx, text_lstm_sequence); + + graph = ggml_new_graph_custom(ctx, graph_config.graph_node_capacity, false); + ggml_build_forward_expand(graph, text_features_ct.tensor); + + core::set_backend_threads(backend, n_threads); + const double alloc_ms = measure_ms([&]() { + buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + }); + if (!buffer) { + throw std::runtime_error("failed to allocate Kokoro predictor text tensors"); + } + if (engine::core::uses_host_graph_plan(backend)) { + plan = engine::core::create_backend_graph_plan_if_host(backend, graph); + if (!plan) { + throw std::runtime_error("failed to create Kokoro predictor text plan"); + } + } + const double zero_state_upload_ms = measure_ms([&]() { upload_zero_state_inputs(zero_state_inputs); }); + const double input_init_ms = measure_ms([&]() { + std::vector ids(static_cast(token_count), 0); + core::write_tensor_i32(input_ids, ids); + if (use_valid_mask) { + std::vector valid_mask(static_cast(token_count), 1.0f); + core::write_tensor_f32(valid_mask_tc, valid_mask); + core::write_tensor_f32(valid_mask_bct, valid_mask); + } + }); + engine::debug::timing_log_scalar("kokoro.graph.build.predictor_text_alloc_ms", alloc_ms); + engine::debug::timing_log_scalar( + "kokoro.graph.build.predictor_text_zero_state_upload_ms", + zero_state_upload_ms); + engine::debug::timing_log_scalar("kokoro.graph.build.predictor_text_input_init_ms", input_init_ms); + } catch (...) { + if (buffer) { + ggml_backend_buffer_free(buffer); + } + if (ctx) { + ggml_free(ctx); + } + ctx = nullptr; + throw; + } + } + + ~TextSession() { + if (plan) { + engine::core::free_backend_graph_plan(backend, plan); + } + if (buffer) { + ggml_backend_buffer_free(buffer); + } + if (ctx) { + ggml_free(ctx); + } + } + + std::vector run(const std::vector & input_ids_host, int64_t valid_token_count) { + if (valid_token_count <= 0 || valid_token_count > token_count) { + throw std::runtime_error("kokoro predictor text valid token count exceeds prepared capacity"); + } + const double input_upload_ms = measure_ms([&]() { + core::write_tensor_i32(input_ids, input_ids_host); + if (use_valid_mask) { + std::vector valid_mask(static_cast(token_count), 0.0f); + std::fill(valid_mask.begin(), valid_mask.begin() + valid_token_count, 1.0f); + core::write_tensor_f32(valid_mask_tc, valid_mask); + core::write_tensor_f32(valid_mask_bct, valid_mask); + } + }); + core::set_backend_threads(backend, n_threads); + ggml_status status = GGML_STATUS_SUCCESS; + const double graph_compute_ms = measure_ms([&]() { + status = core::compute_backend_graph(backend, graph, plan); + }); + if (status != GGML_STATUS_SUCCESS) { + throw std::runtime_error(std::string("kokoro predictor text compute failed: ") + ggml_status_to_string(status)); + } + std::vector text_features; + const double output_read_ms = measure_ms([&]() { + text_features = core::read_tensor_f32(text_features_ct.tensor); + }); + engine::debug::timing_log_scalar("kokoro.predictor_text.input_upload_ms", input_upload_ms); + engine::debug::timing_log_scalar("kokoro.predictor_text.graph.compute_ms", graph_compute_ms); + engine::debug::timing_log_scalar("kokoro.predictor_text.output_read_ms", output_read_ms); + return text_features; + } + }; + + const KokoroWeights * weights_ = nullptr; + ggml_backend_t backend_ = nullptr; + ggml_backend_t duration_backend_ = nullptr; + ggml_backend_t text_backend_ = nullptr; + int n_threads_ = 1; + KokoroPredictorGraphConfig graph_config_ = {}; + + struct PreparedPreTailSession { + int64_t token_count = 0; + bool use_valid_mask = false; + std::unique_ptr duration; + std::unique_ptr text; + }; + + PreparedPreTailSession session_; + + PreparedPreTailSession & session_for_capacity(int64_t token_count, bool use_valid_mask) { + if (token_count <= 0) { + throw std::runtime_error("kokoro predictor pre-tail capacity must be positive"); + } + if (session_.token_count == token_count && session_.use_valid_mask == use_valid_mask) { + return session_; + } + PreparedPreTailSession session; + session.token_count = token_count; + session.use_valid_mask = use_valid_mask; + session_ = std::move(session); + return session_; + } + + DurationSession & duration_session(PreparedPreTailSession & prepared) { + if (!prepared.duration) { + const double build_ms = measure_ms([&]() { + prepared.duration = std::make_unique( + *weights_, + duration_backend_, + prepared.token_count, + prepared.use_valid_mask, + n_threads_, + graph_config_); + }); + engine::debug::timing_log_scalar("kokoro.graph.build.predictor_duration_ms", build_ms); + } + return *prepared.duration; + } + + TextSession & text_session(PreparedPreTailSession & prepared) { + if (!prepared.text) { + const double build_ms = measure_ms([&]() { + prepared.text = std::make_unique( + *weights_, + text_backend_, + prepared.token_count, + prepared.use_valid_mask, + n_threads_, + graph_config_); + }); + engine::debug::timing_log_scalar("kokoro.graph.build.predictor_text_ms", build_ms); + } + return *prepared.text; + } +}; + +class TailSharedLstmBlockRuntime { +public: + TailSharedLstmBlockRuntime( + const KokoroWeights & weights, + ggml_backend_t backend, + int n_threads, + KokoroPredictorGraphConfig graph_config) + : weights_(&weights), + backend_(backend), + n_threads_(std::max(1, n_threads)), + graph_config_(graph_config) {} + + std::vector run( + const std::vector & encoder_tc, + int64_t frames, + int64_t cols) { + if (frames <= 0) { + throw std::runtime_error("kokoro tail shared LSTM requires positive frame count"); + } + if (cols != weights_->predictor.shared.input_size) { + throw std::runtime_error("kokoro tail shared LSTM input width changed"); + } + const int64_t hidden_size = weights_->predictor.shared.hidden_size; + std::vector forward(static_cast(frames * hidden_size), 0.0f); + std::vector reverse(static_cast(frames * hidden_size), 0.0f); + run_direction(false, encoder_tc, frames, cols, forward); + run_direction(true, encoder_tc, frames, cols, reverse); + + std::vector shared_ct(static_cast(2 * hidden_size * frames), 0.0f); + for (int64_t t = 0; t < frames; ++t) { + for (int64_t h = 0; h < hidden_size; ++h) { + shared_ct[static_cast(h * frames + t)] = + forward[static_cast(t * hidden_size + h)]; + shared_ct[static_cast((hidden_size + h) * frames + t)] = + reverse[static_cast(t * hidden_size + h)]; + } + } + return shared_ct; + } + +private: + static constexpr int64_t kBlockFrames = 64; + + struct DirectionResult { + std::vector sequence; + std::vector hidden; + std::vector cell; + }; + + struct DirectionSession { + const KokoroWeights * weights = nullptr; + ggml_backend_t backend = nullptr; + int n_threads = 1; + bool reverse = false; + KokoroPredictorGraphConfig graph_config = {}; + ggml_context * ctx = nullptr; + ggml_backend_buffer_t buffer = nullptr; + ggml_backend_graph_plan_t plan = nullptr; + ggml_cgraph * graph = nullptr; + core::TensorValue input = {}; + core::TensorValue valid_mask = {}; + core::TensorValue hidden_in = {}; + core::TensorValue cell_in = {}; + core::TensorValue sequence_out = {}; + core::TensorValue hidden_out = {}; + core::TensorValue cell_out = {}; + + DirectionSession( + const KokoroWeights & weights_in, + ggml_backend_t backend_in, + int n_threads_in, + bool reverse_in, + KokoroPredictorGraphConfig graph_config_in) + : weights(&weights_in), + backend(backend_in), + n_threads(n_threads_in), + reverse(reverse_in), + graph_config(graph_config_in) { + ggml_init_params params{ + /*.mem_size =*/ graph_config.tail_graph_bytes, + /*.mem_buffer =*/ nullptr, + /*.no_alloc =*/ true, + }; + ctx = ggml_init(params); + if (!ctx) { + throw std::runtime_error("failed to initialize Kokoro tail shared LSTM block context"); + } + try { + core::ModuleBuildContext build_ctx = {}; + build_ctx.ggml = ctx; + input = core::make_tensor( + build_ctx, + GGML_TYPE_F32, + core::TensorShape::from_dims({kBlockFrames, weights->predictor.shared.input_size})); + valid_mask = core::make_tensor( + build_ctx, + GGML_TYPE_F32, + core::TensorShape::from_dims({kBlockFrames, 1})); + hidden_in = core::make_tensor( + build_ctx, + GGML_TYPE_F32, + core::TensorShape::from_dims({1, weights->predictor.shared.hidden_size})); + cell_in = core::make_tensor( + build_ctx, + GGML_TYPE_F32, + core::TensorShape::from_dims({1, weights->predictor.shared.hidden_size})); + ggml_set_input(input.tensor); + ggml_set_input(valid_mask.tensor); + ggml_set_input(hidden_in.tensor); + ggml_set_input(cell_in.tensor); + + const auto outputs = build_lstm_sequence_combined_input_bias( + build_ctx, + input, + weights->predictor.shared, + reverse, + hidden_in, + cell_in, + &valid_mask, + true); + sequence_out = outputs.sequence; + hidden_out = outputs.hidden; + cell_out = outputs.cell; + graph = ggml_new_graph_custom(ctx, graph_config.graph_node_capacity, false); + ggml_build_forward_expand(graph, sequence_out.tensor); + ggml_build_forward_expand(graph, hidden_out.tensor); + ggml_build_forward_expand(graph, cell_out.tensor); + + core::set_backend_threads(backend, n_threads); + buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + if (!buffer) { + throw std::runtime_error("failed to allocate Kokoro tail shared LSTM block tensors"); + } + if (engine::core::uses_host_graph_plan(backend)) { + plan = engine::core::create_backend_graph_plan_if_host(backend, graph); + if (!plan) { + throw std::runtime_error("failed to create Kokoro tail shared LSTM block plan"); + } + } + std::vector zeros_input(static_cast(kBlockFrames * weights->predictor.shared.input_size), 0.0f); + std::vector zeros_mask(static_cast(kBlockFrames), 0.0f); + std::vector zeros_state(static_cast(weights->predictor.shared.hidden_size), 0.0f); + core::write_tensor_f32(input, zeros_input); + core::write_tensor_f32(valid_mask, zeros_mask); + core::write_tensor_f32(hidden_in, zeros_state); + core::write_tensor_f32(cell_in, zeros_state); + } catch (...) { + if (plan) { + engine::core::free_backend_graph_plan(backend, plan); + plan = nullptr; + } + if (buffer) { + ggml_backend_buffer_free(buffer); + buffer = nullptr; + } + if (ctx) { + ggml_free(ctx); + ctx = nullptr; + } + throw; + } + } + + ~DirectionSession() { + if (plan) { + engine::core::free_backend_graph_plan(backend, plan); + } + if (buffer) { + ggml_backend_buffer_free(buffer); + } + if (ctx) { + ggml_free(ctx); + } + } + + DirectionResult run( + const std::vector & block, + int64_t valid_frames, + const std::vector & hidden, + const std::vector & cell) { + if (valid_frames <= 0 || valid_frames > kBlockFrames) { + throw std::runtime_error("kokoro tail shared LSTM block valid frame count is out of range"); + } + std::vector mask(static_cast(kBlockFrames), 0.0f); + std::fill(mask.begin(), mask.begin() + valid_frames, 1.0f); + core::write_tensor_f32(input, block); + core::write_tensor_f32(valid_mask, mask); + core::write_tensor_f32(hidden_in, hidden); + core::write_tensor_f32(cell_in, cell); + + core::set_backend_threads(backend, n_threads); + ggml_status status = GGML_STATUS_SUCCESS; + status = core::compute_backend_graph(backend, graph, plan); + if (status != GGML_STATUS_SUCCESS) { + throw std::runtime_error(std::string("kokoro tail shared LSTM block compute failed: ") + ggml_status_to_string(status)); + } + return { + core::read_tensor_f32(sequence_out.tensor), + core::read_tensor_f32(hidden_out.tensor), + core::read_tensor_f32(cell_out.tensor), + }; + } + }; + + DirectionSession & direction_session(bool reverse) { + auto & session = reverse ? reverse_session_ : forward_session_; + if (!session) { + const double build_ms = measure_ms([&]() { + session = std::make_unique( + *weights_, + backend_, + n_threads_, + reverse, + graph_config_); + }); + engine::debug::timing_log_scalar( + reverse ? "kokoro.graph.build.predictor_tail_shared_lstm_reverse_ms" + : "kokoro.graph.build.predictor_tail_shared_lstm_forward_ms", + build_ms); + } + return *session; + } + + void run_direction( + bool reverse, + const std::vector & encoder_tc, + int64_t frames, + int64_t cols, + std::vector & output_tc) { + DirectionSession & session = direction_session(reverse); + const int64_t hidden_size = weights_->predictor.shared.hidden_size; + std::vector hidden(static_cast(hidden_size), 0.0f); + std::vector cell(static_cast(hidden_size), 0.0f); + const int64_t chunks = (frames + kBlockFrames - 1) / kBlockFrames; + + for (int64_t chunk_index = 0; chunk_index < chunks; ++chunk_index) { + const int64_t logical_chunk = reverse ? (chunks - 1 - chunk_index) : chunk_index; + const int64_t start = logical_chunk * kBlockFrames; + const int64_t valid = std::min(kBlockFrames, frames - start); + std::vector block(static_cast(kBlockFrames * cols), 0.0f); + for (int64_t t = 0; t < valid; ++t) { + std::memcpy( + block.data() + static_cast(t * cols), + encoder_tc.data() + static_cast((start + t) * cols), + static_cast(cols) * sizeof(float)); + } + DirectionResult result = session.run(block, valid, hidden, cell); + hidden = std::move(result.hidden); + cell = std::move(result.cell); + for (int64_t t = 0; t < valid; ++t) { + std::memcpy( + output_tc.data() + static_cast((start + t) * hidden_size), + result.sequence.data() + static_cast(t * hidden_size), + static_cast(hidden_size) * sizeof(float)); + } + } + } + + const KokoroWeights * weights_ = nullptr; + ggml_backend_t backend_ = nullptr; + int n_threads_ = 1; + KokoroPredictorGraphConfig graph_config_ = {}; + std::unique_ptr forward_session_; + std::unique_ptr reverse_session_; +}; + +class PredictorTailGraphRuntime { +public: + PredictorTailGraphRuntime( + const KokoroWeights & weights, + ggml_backend_t backend, + int n_threads, + bool use_device_backend, + KokoroPredictorGraphConfig graph_config) + : weights_(&weights), + backend_(backend), + n_threads_(std::max(1, n_threads)), + use_device_backend_(use_device_backend), + graph_config_(graph_config), + shared_lstm_(weights, backend, n_threads, graph_config) {} + + PredictorOutputs run( + const PredictorPreTailOutputs & pre_tail, + const std::vector & style_predictor, + const std::vector & style_decoder) { + std::vector shared_ct; + const double shared_lstm_ms = measure_ms([&]() { + shared_ct = shared_lstm_.run( + pre_tail.expanded_encoder_tc, + pre_tail.expanded_encoder_rows, + pre_tail.expanded_encoder_cols); + }); + engine::debug::timing_log_scalar("kokoro.predictor_tail.shared_lstm_block_ms", shared_lstm_ms); + return session_for(pre_tail.expanded_encoder_rows).run( + shared_ct, + pre_tail.expanded_encoder_rows, + pre_tail.asr_ct, + pre_tail.asr_rows, + pre_tail.asr_cols, + pre_tail.durations, + style_predictor, + style_decoder); + } + +private: + struct Session { + const KokoroWeights * weights = nullptr; + int64_t frames = 0; + int n_threads = 1; + bool use_device_backend = false; + KokoroPredictorGraphConfig graph_config = {}; + ggml_context * ctx = nullptr; + ggml_tensor * shared_in = nullptr; + ggml_tensor * asr_in = nullptr; + ggml_tensor * style_predictor_in = nullptr; + ggml_tensor * style_decoder_in = nullptr; + ggml_tensor * f0_out = nullptr; + ggml_tensor * decoder_x_out = nullptr; + std::vector time_masks; + ggml_cgraph * graph = nullptr; + ggml_backend_t backend = nullptr; + ggml_backend_buffer_t buffer = nullptr; + ggml_backend_graph_plan_t plan = nullptr; + + Session( + const KokoroWeights & weights_in, + ggml_backend_t backend_in, + int64_t frames_in, + int n_threads_in, + bool use_device_backend_in, + KokoroPredictorGraphConfig graph_config_in) + : weights(&weights_in), + frames(frames_in), + n_threads(n_threads_in), + use_device_backend(use_device_backend_in), + graph_config(graph_config_in), + backend(backend_in) { + ggml_init_params params{ + /*.mem_size =*/ graph_config.tail_graph_bytes, + /*.mem_buffer =*/ nullptr, + /*.no_alloc =*/ true, + }; + ctx = ggml_init(params); + if (!ctx) { + throw std::runtime_error("failed to initialize ggml context for Kokoro predictor tail"); + } + + try { + double define_prosody_ms = 0.0; + double define_decoder_ms = 0.0; + double define_expand_ms = 0.0; + const double define_ms = measure_ms([&]() { + shared_in = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, frames, weights->predictor.shared.hidden_size * 2); + asr_in = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, frames, 512); + style_predictor_in = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, weights->style_dim, 1); + style_decoder_in = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, weights->style_dim, 1); + ggml_set_input(shared_in); + ggml_set_input(asr_in); + ggml_set_input(style_predictor_in); + ggml_set_input(style_decoder_in); + core::ModuleBuildContext build_ctx = {}; + build_ctx.ggml = ctx; + const auto style_predictor = core::wrap_tensor( + style_predictor_in, + core::TensorShape::from_dims({1, weights->style_dim}), + GGML_TYPE_F32); + const auto style_decoder = core::wrap_tensor( + style_decoder_in, + core::TensorShape::from_dims({1, weights->style_dim}), + GGML_TYPE_F32); + + ggml_tensor * shared_ct = shared_in; + + ggml_tensor * f0 = shared_ct; + ggml_tensor * n = shared_ct; + ggml_tensor * f0_down = nullptr; + ggml_tensor * n_down = nullptr; + define_prosody_ms = measure_ms([&]() { + for (size_t i = 0; i < weights->predictor.f0_blocks.size(); ++i) { + f0 = build_adain_resblock_ct_predictor( + ctx, + f0, + weights->predictor.f0_blocks[i], + style_predictor, + style_decoder, + false, + time_masks, + false); + } + for (size_t i = 0; i < weights->predictor.n_blocks.size(); ++i) { + n = build_adain_resblock_ct_predictor( + ctx, + n, + weights->predictor.n_blocks[i], + style_predictor, + style_decoder, + false, + time_masks, + false); + } + + f0_out = build_standard_conv1d_ct_predictor(ctx, f0, weights->predictor.f0_proj); + ggml_tensor * n_out = build_standard_conv1d_ct_predictor(ctx, n, weights->predictor.n_proj); + f0_down = build_standard_conv1d_ct_predictor(ctx, f0_out, weights->decoder.f0_conv); + n_down = build_standard_conv1d_ct_predictor(ctx, n_out, weights->decoder.n_conv); + }); + const auto asr_in_tv = core::wrap_tensor(asr_in, core::TensorShape::from_dims({asr_in->ne[1], asr_in->ne[0]}), GGML_TYPE_F32); + const auto f0_down_tv = core::wrap_tensor(f0_down, core::TensorShape::from_dims({f0_down->ne[1], f0_down->ne[0]}), GGML_TYPE_F32); + const auto n_down_tv = core::wrap_tensor(n_down, core::TensorShape::from_dims({n_down->ne[1], n_down->ne[0]}), GGML_TYPE_F32); + define_decoder_ms = measure_ms([&]() { + auto asr_f0 = modules::ConcatModule({0}).build(build_ctx, asr_in_tv, f0_down_tv); + auto pre_encode = modules::ConcatModule({0}).build(build_ctx, asr_f0, n_down_tv); + ggml_tensor * x = pre_encode.tensor; + x = build_adain_resblock_ct_predictor( + ctx, + x, + weights->decoder.encode, + style_predictor, + style_decoder, + true, + time_masks, + false); + ggml_tensor * asr_res = build_standard_conv1d_ct_predictor(ctx, asr_in, weights->decoder.asr_res); + const auto asr_res_tv = + core::wrap_tensor(asr_res, core::TensorShape::from_dims({asr_res->ne[1], asr_res->ne[0]}), GGML_TYPE_F32); + auto f0_n = modules::ConcatModule({0}).build(build_ctx, f0_down_tv, n_down_tv); + auto cond = modules::ConcatModule({0}).build(build_ctx, asr_res_tv, f0_n); + bool use_residual_conditioning = true; + for (size_t i = 0; i < weights->decoder.decode.size(); ++i) { + if (use_residual_conditioning) { + const auto x_tv = core::wrap_tensor(x, core::TensorShape::from_dims({x->ne[1], x->ne[0]}), GGML_TYPE_F32); + auto decode_in = modules::ConcatModule({0}).build(build_ctx, x_tv, cond); + x = decode_in.tensor; + } + x = build_adain_resblock_ct_predictor( + ctx, + x, + weights->decoder.decode[i], + style_predictor, + style_decoder, + true, + time_masks, + false); + if (weights->decoder.decode[i].upsample) { + use_residual_conditioning = false; + } + } + decoder_x_out = x; + }); + define_expand_ms = measure_ms([&]() { + graph = ggml_new_graph_custom(ctx, graph_config.graph_node_capacity, false); + ggml_build_forward_expand(graph, f0_out); + ggml_build_forward_expand(graph, decoder_x_out); + }); + }); + + core::set_backend_threads(backend, n_threads); + const double alloc_ms = measure_ms([&]() { + buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + }); + if (!buffer) { + throw std::runtime_error("failed to allocate Kokoro predictor tail tensors"); + } + if (engine::core::uses_host_graph_plan(backend)) { + plan = engine::core::create_backend_graph_plan_if_host(backend, graph); + if (!plan) { + throw std::runtime_error("failed to create Kokoro predictor tail plan"); + } + } + const double materialize_ms = measure_ms([&]() { + std::vector shared( + static_cast(weights->predictor.shared.hidden_size * 2 * frames), + 0.0f); + std::vector asr(static_cast(512 * frames), 0.0f); + std::vector style_predictor(static_cast(weights->style_dim), 0.0f); + std::vector style_decoder(static_cast(weights->style_dim), 0.0f); + ggml_backend_tensor_set(shared_in, shared.data(), 0, ggml_nbytes(shared_in)); + ggml_backend_tensor_set(asr_in, asr.data(), 0, ggml_nbytes(asr_in)); + ggml_backend_tensor_set(style_predictor_in, style_predictor.data(), 0, ggml_nbytes(style_predictor_in)); + ggml_backend_tensor_set(style_decoder_in, style_decoder.data(), 0, ggml_nbytes(style_decoder_in)); + upload_time_masks(time_masks, frames, frames); + }); + engine::debug::timing_log_scalar("kokoro.graph.build.predictor_tail_define_ms", define_ms); + engine::debug::timing_log_scalar("kokoro.graph.build.predictor_tail_define_prosody_ms", define_prosody_ms); + engine::debug::timing_log_scalar("kokoro.graph.build.predictor_tail_define_decoder_ms", define_decoder_ms); + engine::debug::timing_log_scalar("kokoro.graph.build.predictor_tail_define_expand_ms", define_expand_ms); + engine::debug::timing_log_scalar("kokoro.graph.build.predictor_tail_alloc_ms", alloc_ms); + engine::debug::timing_log_scalar("kokoro.graph.build.predictor_tail_materialize_ms", materialize_ms); + } catch (...) { + if (buffer) { + ggml_backend_buffer_free(buffer); + } + if (ctx) { + ggml_free(ctx); + } + ctx = nullptr; + throw; + } + } + + ~Session() { + if (plan) { + engine::core::free_backend_graph_plan(backend, plan); + } + if (buffer) { + ggml_backend_buffer_free(buffer); + } + if (ctx) { + ggml_free(ctx); + } + } + + PredictorOutputs run( + const std::vector & shared_ct, + int64_t expanded_encoder_rows, + const std::vector & asr_ct, + int64_t asr_rows, + int64_t asr_cols, + const std::vector & durations, + const std::vector & style_predictor, + const std::vector & style_decoder) { + if (expanded_encoder_rows <= 0 || expanded_encoder_rows > frames || asr_cols != expanded_encoder_rows) { + throw std::runtime_error("kokoro predictor tail input frame count exceeds prepared capacity"); + } + if (static_cast(shared_ct.size()) != weights->predictor.shared.hidden_size * 2 * expanded_encoder_rows || + asr_rows != 512) { + throw std::runtime_error("kokoro predictor tail input shape changed"); + } + const double upload_ms = measure_ms([&]() { + std::vector padded_shared(static_cast(weights->predictor.shared.hidden_size * 2 * frames), 0.0f); + const int64_t shared_rows = weights->predictor.shared.hidden_size * 2; + for (int64_t row = 0; row < shared_rows; ++row) { + std::memcpy( + padded_shared.data() + static_cast(row * frames), + shared_ct.data() + static_cast(row * expanded_encoder_rows), + static_cast(expanded_encoder_rows) * sizeof(float)); + } + std::vector padded_asr(static_cast(512 * frames), 0.0f); + for (int64_t row = 0; row < 512; ++row) { + std::memcpy( + padded_asr.data() + static_cast(row * frames), + asr_ct.data() + static_cast(row * expanded_encoder_rows), + static_cast(expanded_encoder_rows) * sizeof(float)); + } + ggml_backend_tensor_set(shared_in, padded_shared.data(), 0, ggml_nbytes(shared_in)); + ggml_backend_tensor_set(asr_in, padded_asr.data(), 0, ggml_nbytes(asr_in)); + ggml_backend_tensor_set(style_predictor_in, style_predictor.data(), 0, ggml_nbytes(style_predictor_in)); + ggml_backend_tensor_set(style_decoder_in, style_decoder.data(), 0, ggml_nbytes(style_decoder_in)); + upload_time_masks(time_masks, expanded_encoder_rows, frames); + }); + engine::debug::timing_log_scalar("kokoro.predictor_tail.input_upload_ms", upload_ms); + core::set_backend_threads(backend, n_threads); + ggml_status status = GGML_STATUS_SUCCESS; + const double compute_ms = measure_ms([&]() { + status = core::compute_backend_graph(backend, graph, plan); + }); + engine::debug::timing_log_scalar("kokoro.predictor_tail.compute_ms", compute_ms); + if (status != GGML_STATUS_SUCCESS) { + throw std::runtime_error(std::string("kokoro predictor tail compute failed: ") + ggml_status_to_string(status)); + } + std::vector f0_curve; + const int64_t decoder_x_capacity_cols = decoder_x_out->ne[0]; + if (decoder_x_capacity_cols <= 0 || decoder_x_capacity_cols % frames != 0) { + throw std::runtime_error("kokoro predictor tail decoder output capacity is inconsistent"); + } + const int64_t decoder_frame_scale = decoder_x_capacity_cols / frames; + const int64_t decoder_x_cols = expanded_encoder_rows * decoder_frame_scale; + if (decoder_x_cols <= 0 || decoder_x_cols > decoder_x_capacity_cols) { + throw std::runtime_error("kokoro predictor tail decoder output length is inconsistent"); + } + const double output_read_ms = measure_ms([&]() { + f0_curve = core::read_tensor_f32(f0_out); + }); + if (static_cast(f0_curve.size()) < decoder_x_cols) { + throw std::runtime_error("kokoro predictor tail f0 output is shorter than decoder output"); + } + if (static_cast(f0_curve.size()) > decoder_x_cols) { + f0_curve.resize(static_cast(decoder_x_cols)); + } + engine::debug::timing_log_scalar("kokoro.predictor_tail.output_read_ms", output_read_ms); + PredictorOutputs result; + result.durations = durations; + result.f0_curve = std::move(f0_curve); + result.decoder_x = std::vector{}; + result.decoder_x_rows = decoder_x_out->ne[1]; + result.decoder_x_cols = decoder_x_cols; + result.decoder_x_tensor = decoder_x_out; + result.decoder_x_on_backend = true; + return result; + } + }; + + Session & session_for(int64_t frames) { + if (session_ && + session_->weights == weights_ && + session_->frames == frames && + session_->n_threads == n_threads_ && + session_->use_device_backend == use_device_backend_) { + return *session_; + } + session_.reset(); + const double build_ms = measure_ms([&]() { + session_ = std::make_unique( + *weights_, + backend_, + frames, + n_threads_, + use_device_backend_, + graph_config_); + }); + engine::debug::timing_log_scalar("kokoro.graph.build.predictor_tail_ms", build_ms); + return *session_; + } + + const KokoroWeights * weights_ = nullptr; + ggml_backend_t backend_ = nullptr; + int n_threads_ = 1; + bool use_device_backend_ = false; + KokoroPredictorGraphConfig graph_config_ = {}; + TailSharedLstmBlockRuntime shared_lstm_; + std::unique_ptr session_; +}; + +} // namespace + +struct KokoroPredictorRuntime::Impl { + std::shared_ptr weights; + int n_threads = 1; + bool use_device_backend = false; + ggml_backend_t backend = nullptr; + int64_t pre_tail_token_capacity = 0; + KokoroPredictorGraphConfig graph_config = {}; + PlbertGraphRuntime plbert; + PredictorPreTailGraphRuntime pre_tail_graph; + PredictorTailGraphRuntime tail_graph; + + Impl( + std::shared_ptr weights_in, + ggml_backend_t backend_in, + int n_threads_in, + bool use_device_backend_in, + int64_t plbert_fixed_token_capacity, + int64_t pre_tail_token_capacity_in, + KokoroPredictorGraphConfig graph_config_in) + : weights(std::move(weights_in)), + n_threads(std::max(1, n_threads_in)), + use_device_backend(use_device_backend_in), + backend(backend_in), + pre_tail_token_capacity(pre_tail_token_capacity_in), + graph_config(graph_config_in), + plbert(weights, backend, n_threads, use_device_backend, plbert_fixed_token_capacity), + pre_tail_graph(*weights, backend, n_threads, use_device_backend, graph_config), + tail_graph(*weights, backend, n_threads, use_device_backend, graph_config) { + if (pre_tail_token_capacity > 0) { + pre_tail_graph.prepare_capacity(pre_tail_token_capacity); + } + } +}; + +KokoroPredictorRuntime::KokoroPredictorRuntime( + std::shared_ptr weights, + ggml_backend_t backend, + int n_threads, + bool use_device_backend, + int64_t plbert_fixed_token_capacity, + int64_t pre_tail_token_capacity, + KokoroPredictorGraphConfig graph_config) + : impl_(std::make_unique( + std::move(weights), + backend, + n_threads, + use_device_backend, + plbert_fixed_token_capacity, + pre_tail_token_capacity, + graph_config)) {} + +KokoroPredictorRuntime::~KokoroPredictorRuntime() = default; + +PredictorOutputs KokoroPredictorRuntime::predict( + const std::vector & input_ids, + const std::vector & ref_s, + float speed) { + if (input_ids.empty()) { + throw std::runtime_error("kokoro_predict requires non-empty input_ids"); + } + if (static_cast(ref_s.size()) != 256) { + throw std::runtime_error("kokoro_predict requires ref_s with 256 elements"); + } + const int64_t token_count = static_cast(input_ids.size()); + const std::vector style_decoder(ref_s.begin(), ref_s.begin() + 128); + const std::vector style_predictor(ref_s.begin() + 128, ref_s.end()); + std::vector plbert_hidden_tc; + const double plbert_ms = measure_ms([&]() { + plbert_hidden_tc = impl_->plbert.run(input_ids); + }); + engine::debug::timing_log_scalar("kokoro.predictor.plbert_ms", plbert_ms); + + const int64_t pre_tail_capacity = + checked_capacity(std::max(token_count, impl_->pre_tail_token_capacity), impl_->weights->context_length); + PredictorPreTailOutputs pre_tail; + const double pretail_ms = measure_ms([&]() { + pre_tail = impl_->pre_tail_graph.run(plbert_hidden_tc, input_ids, style_predictor, speed, pre_tail_capacity); + }); + engine::debug::timing_log_scalar("kokoro.predictor.pretail_ms", pretail_ms); + PredictorOutputs output; + const double tail_ms = measure_ms([&]() { + output = impl_->tail_graph.run(pre_tail, style_predictor, style_decoder); + }); + engine::debug::timing_log_scalar("kokoro.predictor.tail_ms", tail_ms); + return output; +} + +} // namespace kokoro_ggml diff --git a/src/models/kokoro_tts/session.cpp b/src/models/kokoro_tts/session.cpp new file mode 100644 index 000000000..98279685e --- /dev/null +++ b/src/models/kokoro_tts/session.cpp @@ -0,0 +1,481 @@ +#include "engine/models/kokoro_tts/session.h" + +#include "engine/framework/debug/profiler.h" +#include "engine/framework/text/chunking.h" +#include "engine/framework/debug/trace.h" +#include "engine/framework/runtime/options.h" + +#include "engine/models/kokoro_tts/decoder.h" +#include "engine/models/kokoro_tts/frontend.h" +#include "engine/models/kokoro_tts/predictor.h" + +#include +#include +#include +#include +#include + +namespace engine::models::kokoro_tts { + +namespace { +using engine::debug::measure_ms; +constexpr int64_t kDefaultTextChunkSize = 240; + +int64_t parse_positive_i64_option( + const runtime::SessionOptions & options, + std::initializer_list keys, + int64_t fallback) { + for (const char * key : keys) { + const auto it = options.options.find(key); + if (it == options.options.end() || it->second.empty()) { + continue; + } + const int64_t value = std::stoll(it->second); + if (value <= 0) { + throw std::runtime_error(std::string(key) + " must be positive"); + } + return value; + } + return fallback; +} + +uint64_t parse_u64_option( + const runtime::SessionOptions & options, + std::initializer_list keys, + uint64_t fallback) { + for (const char * key : keys) { + const auto it = options.options.find(key); + if (it == options.options.end() || it->second.empty()) { + continue; + } + return static_cast(std::stoull(it->second)); + } + return fallback; +} + +size_t parse_size_mb_option( + const runtime::SessionOptions & options, + std::initializer_list keys, + size_t fallback) { + for (const char * key : keys) { + const auto it = options.options.find(key); + if (it == options.options.end() || it->second.empty()) { + continue; + } + const auto mb = std::stoull(it->second); + if (mb == 0 || mb > std::numeric_limits::max() / (1024ull * 1024ull)) { + throw std::runtime_error(std::string(key) + " is out of range"); + } + return static_cast(mb) * 1024ull * 1024ull; + } + return fallback; +} + +engine::assets::TensorStorageType parse_storage_option( + const runtime::SessionOptions & options, + std::initializer_list keys, + engine::assets::TensorStorageType fallback) { + for (const char * key : keys) { + const auto it = options.options.find(key); + if (it == options.options.end() || it->second.empty()) { + continue; + } + return engine::assets::parse_tensor_storage_type(it->second); + } + return fallback; +} + +void validate_matmul_storage(engine::assets::TensorStorageType storage_type, const char * option_name) { + if (storage_type == engine::assets::TensorStorageType::Native || + storage_type == engine::assets::TensorStorageType::F32 || + storage_type == engine::assets::TensorStorageType::F16 || + storage_type == engine::assets::TensorStorageType::BF16 || + storage_type == engine::assets::TensorStorageType::Q8_0) { + return; + } + throw std::runtime_error(std::string(option_name) + " supports only native, f32, f16, bf16, and q8_0"); +} + +void validate_conv_storage(engine::assets::TensorStorageType storage_type, const char * option_name) { + if (storage_type == engine::assets::TensorStorageType::Native || + storage_type == engine::assets::TensorStorageType::F32 || + storage_type == engine::assets::TensorStorageType::F16) { + return; + } + throw std::runtime_error(std::string(option_name) + " supports only native, f32, and f16"); +} + +std::string request_cache_key(const runtime::Transcript & text) { + std::ostringstream out; + out << text.text.size() << ":" << text.text; + return out.str(); +} + +} // namespace + +struct KokoroTTSSession::PreparedRuntime { + std::unique_ptr predictor; +}; + +KokoroTTSSession::KokoroTTSSession( + runtime::TaskSpec task, + runtime::SessionOptions options, + std::shared_ptr assets) + : RuntimeSessionBase(options), + task_(std::move(task)), + assets_(std::move(assets)) { + if (!assets_ || !assets_->model_weights) { + throw std::runtime_error("Kokoro TTS session requires loaded assets"); + } + matmul_weight_storage_type_ = parse_storage_option( + RuntimeSessionBase::options(), + {"kokoro_tts.weight_type", "kokoro.weight_type"}, + matmul_weight_storage_type_); + conv_weight_storage_type_ = parse_storage_option( + RuntimeSessionBase::options(), + {"kokoro_tts.conv_weight_type", "kokoro.conv_weight_type"}, + conv_weight_storage_type_); + matmul_weight_storage_type_ = parse_storage_option( + RuntimeSessionBase::options(), + {"kokoro_tts.matmul_weight_type", "kokoro.matmul_weight_type"}, + matmul_weight_storage_type_); + validate_matmul_storage(matmul_weight_storage_type_, "kokoro_tts.weight_type"); + validate_conv_storage(conv_weight_storage_type_, "kokoro_tts.conv_weight_type"); + weight_context_bytes_ = parse_size_mb_option( + RuntimeSessionBase::options(), + {"kokoro_tts.weight_context_mb", "kokoro_weight_context_mb"}, + weight_context_bytes_); + predictor_duration_graph_bytes_ = parse_size_mb_option( + RuntimeSessionBase::options(), + {"kokoro_tts.predictor_duration_graph_mb", "kokoro_predictor_duration_graph_mb"}, + predictor_duration_graph_bytes_); + predictor_text_graph_bytes_ = parse_size_mb_option( + RuntimeSessionBase::options(), + {"kokoro_tts.predictor_text_graph_mb", "kokoro_predictor_text_graph_mb"}, + predictor_text_graph_bytes_); + predictor_tail_graph_bytes_ = parse_size_mb_option( + RuntimeSessionBase::options(), + {"kokoro_tts.predictor_tail_graph_mb", "kokoro_predictor_tail_graph_mb"}, + predictor_tail_graph_bytes_); + weights_ = load_kokoro_backend_weights( + *assets_, + execution_context().backend(), + execution_context().backend_type(), + matmul_weight_storage_type_, + conv_weight_storage_type_, + weight_context_bytes_); + const auto graph_capacity_mode = runtime::resolve_graph_capacity_mode( + RuntimeSessionBase::options(), + runtime::GraphCapacityMode::Fixed, + {"offline_graph_capacity_mode", "graph_capacity_mode"}); + if (graph_capacity_mode == runtime::GraphCapacityMode::Unsupported) { + throw std::runtime_error("Kokoro TTS graph_capacity_mode=unsupported is not implemented"); + } + graph_capacity_controller_ = runtime::GraphCapacityController(graph_capacity_mode); + fixed_token_capacity_ = parse_positive_i64_option( + RuntimeSessionBase::options(), + {"max_input_tokens", "offline_max_input_tokens", "kokoro_max_input_tokens"}, + std::min(512, weights_->context_length)); + pre_tail_token_capacity_ = parse_positive_i64_option( + RuntimeSessionBase::options(), + {"kokoro_pretail_tokens", "pre_tail_tokens"}, + 0); + rng_seed_ = parse_u64_option( + RuntimeSessionBase::options(), + {"kokoro_rng_seed", "rng_seed"}, + runtime::random_u64_seed()); + if (fixed_token_capacity_ > weights_->context_length) { + throw std::runtime_error("Kokoro fixed token capacity exceeds model context length"); + } + if (pre_tail_token_capacity_ > weights_->context_length) { + throw std::runtime_error("Kokoro pre-tail token capacity exceeds model context length"); + } +} + +KokoroTTSSession::~KokoroTTSSession() = default; + +std::string KokoroTTSSession::family() const { + return "kokoro_tts"; +} + +runtime::VoiceTaskKind KokoroTTSSession::task_kind() const { + return task_.task; +} + +runtime::RunMode KokoroTTSSession::run_mode() const { + return task_.mode; +} + +int64_t KokoroTTSSession::base_graph_capacity_tokens() const { + return fixed_token_capacity_; +} + +runtime::MappedGraphCapacityAdapter KokoroTTSSession::make_graph_capacity_adapter() { + return runtime::MappedGraphCapacityAdapter( + base_graph_capacity_tokens(), + base_graph_capacity_tokens(), + [this](int64_t request_size) { + if (request_size <= 0) { + throw std::runtime_error("Kokoro graph capacity request size must be positive"); + } + if (request_size > weights_->context_length) { + throw std::runtime_error("Kokoro request exceeds model context length"); + } + return request_size; + }, + [this]() { return prepared_graph_capacities(); }, + [this](int64_t capacity) { prepare_graph_capacity(capacity); }); +} + +std::vector KokoroTTSSession::prepared_graph_capacities() const { + std::vector capacities; + if (prepared_session_ && prepared_session_->predictor && prepared_session_capacity_ > 0) { + capacities.push_back(prepared_session_capacity_); + } + return capacities; +} + +KokoroTTSSession::DecoderCapacityContract KokoroTTSSession::make_decoder_capacity_contract(int64_t decoder_frame_capacity) const { + if (decoder_frame_capacity <= 0) { + throw std::runtime_error("Kokoro decoder frame capacity must be positive"); + } + DecoderCapacityContract contract = {}; + contract.decoder_frame_capacity = decoder_frame_capacity; + contract.conditioning_sample_capacity = contract.decoder_frame_capacity * 300; + const int64_t pad = weights_->decoder.generator.gen_istft_n_fft / 2; + contract.conditioning_frame_capacity = + 1 + (contract.conditioning_sample_capacity + 2 * pad - weights_->decoder.generator.gen_istft_n_fft) / + weights_->decoder.generator.gen_istft_hop_size; + if (contract.conditioning_sample_capacity <= 0 || contract.conditioning_frame_capacity <= 0) { + throw std::runtime_error("Kokoro decoder capacity contract overflowed"); + } + return contract; +} + +void KokoroTTSSession::prepare_graph_capacity(int64_t capacity) { + if (capacity <= 0) { + throw std::runtime_error("Kokoro graph capacity must be positive"); + } + if (capacity > weights_->context_length) { + throw std::runtime_error("Kokoro graph capacity exceeds model context length"); + } + if (prepared_session_ && prepared_session_capacity_ >= capacity) { + return; + } + const int threads = std::max(1, execution_context().config().threads); + const bool use_device_backend = !execution_context().uses_host_graph_plan(); + ggml_backend_t backend = execution_context().backend(); + const int64_t plbert_fixed_token_capacity = 0; + const int64_t predictor_pre_tail_capacity = pre_tail_token_capacity_; + prepared_session_.reset(); + prepared_session_capacity_ = 0; + auto prepared = std::make_unique(); + double build_ms = 0.0; + kokoro_ggml::KokoroPredictorGraphConfig predictor_graph_config; + predictor_graph_config.duration_graph_bytes = predictor_duration_graph_bytes_; + predictor_graph_config.text_graph_bytes = predictor_text_graph_bytes_; + predictor_graph_config.tail_graph_bytes = predictor_tail_graph_bytes_; + build_ms = measure_ms([&]() { + prepared->predictor = std::make_unique( + weights_, + backend, + threads, + use_device_backend, + plbert_fixed_token_capacity, + predictor_pre_tail_capacity, + predictor_graph_config); + }); + prepared_session_ = std::move(prepared); + prepared_session_capacity_ = capacity; + engine::debug::timing_log_scalar("kokoro.prepare.predictor.graph.build_ms", build_ms); +} + +void KokoroTTSSession::prepare_decoder_graph_capacity(int64_t capacity) { + if (capacity <= 0) { + throw std::runtime_error("Kokoro decoder graph capacity must be positive"); + } + if (prepared_decoder_ && prepared_decoder_capacity_ == capacity) { + return; + } + const int threads = std::max(1, execution_context().config().threads); + const bool use_device_backend = !execution_context().uses_host_graph_plan(); + ggml_backend_t backend = execution_context().backend(); + const DecoderCapacityContract contract = make_decoder_capacity_contract(capacity); + double build_ms = 0.0; + kokoro_ggml::KokoroDecoderCapacityContract decoder_contract; + decoder_contract.decoder_frames = contract.decoder_frame_capacity; + decoder_contract.conditioning_frames = contract.conditioning_frame_capacity; + if (prepared_decoder_) { + build_ms = measure_ms([&]() { + prepared_decoder_->prepare(decoder_contract); + }); + } else { + build_ms = measure_ms([&]() { + prepared_decoder_ = std::make_unique( + weights_, + backend, + threads, + use_device_backend, + rng_seed_, + decoder_contract); + }); + } + prepared_decoder_capacity_ = capacity; + prepared_decoder_context_ = contract; + engine::debug::timing_log_scalar("kokoro.prepare.decoder_runtime_build_ms", build_ms); +} + +void KokoroTTSSession::prepare(const runtime::SessionPreparationRequest & request) { + if (task_.task != runtime::VoiceTaskKind::Tts) { + throw std::runtime_error("Kokoro TTS session only supports VoiceTaskKind::Tts"); + } + if (task_.mode != runtime::RunMode::Offline) { + throw std::runtime_error("Kokoro TTS session only supports offline mode"); + } + if (const auto seed = runtime::parse_u64_option(request.options, {"seed", "kokoro_rng_seed", "rng_seed"})) { + if (rng_seed_ != *seed) { + rng_seed_ = *seed; + prepared_decoder_.reset(); + prepared_decoder_capacity_ = 0; + prepared_decoder_context_ = {}; + } + } + if (!frontend_session_state_) { + frontend_session_state_ = + std::make_unique( + resolve_kokoro_frontend_session_state(request.text, request.voice, *assets_)); + } else if (request.text.has_value() || request.voice.has_value()) { + runtime::Transcript transcript; + if (request.text.has_value()) { + transcript = *request.text; + } + validate_kokoro_frontend_session_state( + transcript, + request.voice, + *frontend_session_state_, + *assets_); + } + auto adapter = make_graph_capacity_adapter(); + int64_t request_size = 0; + if (request.text.has_value()) { + const int64_t text_chunk_size = + engine::text::parse_text_chunk_size_override(request.options).value_or(kDefaultTextChunkSize); + const auto text_chunks = engine::text::split_text_chunks(request.text->text, text_chunk_size); + for (const auto & chunk : text_chunks) { + runtime::SessionPreparationRequest chunk_request = request; + chunk_request.text = runtime::Transcript{chunk, request.text->language}; + request_size = std::max( + request_size, + estimate_kokoro_request_tokens(chunk_request, *frontend_session_state_, *assets_)); + } + } + graph_capacity_controller_.ensure_prepared(adapter, request_size); + mark_prepared(); +} + +runtime::TaskResult KokoroTTSSession::run(const runtime::TaskRequest & request) { + require_prepared("Kokoro TTS run()"); + if (task_.task != runtime::VoiceTaskKind::Tts) { + throw std::runtime_error("Kokoro TTS session only supports VoiceTaskKind::Tts"); + } + if (task_.mode != runtime::RunMode::Offline) { + throw std::runtime_error("Kokoro TTS session only supports offline mode"); + } + if (!request.text_input.has_value()) { + throw std::runtime_error("Kokoro TTS run requires text_input"); + } + + const int64_t text_chunk_size = + engine::text::parse_text_chunk_size_override(request.options).value_or(kDefaultTextChunkSize); + const auto chunk_requests = runtime::chunk_text_request(request, text_chunk_size); + engine::debug::trace_log_scalar("kokoro.text_chunk_size", text_chunk_size); + engine::debug::trace_log_scalar("kokoro.text_chunk_count", static_cast(chunk_requests.size())); + double frontend_ms = 0.0; + double inference_ms = 0.0; + double predictor_ms = 0.0; + double decoder_ms = 0.0; + runtime::AudioBuffer merged_audio; + for (const auto & chunk_request : chunk_requests) { + if (!frontend_session_state_) { + frontend_session_state_ = + std::make_unique( + resolve_kokoro_frontend_session_state(chunk_request.text_input, chunk_request.voice, *assets_)); + } + validate_kokoro_frontend_session_state( + *chunk_request.text_input, + chunk_request.voice, + *frontend_session_state_, + *assets_); + const std::string cache_key = request_cache_key(*chunk_request.text_input); + KokoroSynthesisInput input; + frontend_ms += measure_ms([&]() { + if (!cache_key.empty() && cached_input_ && cache_key == cached_request_key_) { + input = *cached_input_; + return; + } + input = build_kokoro_synthesis_input(*chunk_request.text_input, *frontend_session_state_, *assets_); + if (!cache_key.empty()) { + cached_request_key_ = cache_key; + cached_input_ = std::make_unique(input); + } + }); + + const auto inference_started = std::chrono::steady_clock::now(); + auto adapter = make_graph_capacity_adapter(); + const int64_t request_size = static_cast(input.input_ids.size()); + graph_capacity_controller_.ensure_prepared(adapter, request_size); + const int64_t selected_capacity = graph_capacity_controller_.select_capacity_for_run(adapter, request_size); + if (!prepared_session_ || + prepared_session_capacity_ != selected_capacity || + !prepared_session_->predictor) { + throw std::runtime_error("Kokoro selected graph capacity was not prepared"); + } + + kokoro_ggml::PredictorOutputs predictor; + predictor_ms += measure_ms([&]() { + predictor = prepared_session_->predictor->predict( + input.input_ids, + input.style, + input.speaking_rate); + }); + const int64_t decoder_request_size = predictor.decoder_x_cols; + if (predictor.decoder_x_on_backend && predictor.decoder_x_tensor == nullptr) { + throw std::runtime_error("Kokoro predictor reported backend decoder features without a tensor"); + } + if (decoder_request_size <= 0) { + throw std::runtime_error("Kokoro predictor produced invalid decoder request size"); + } + if (static_cast(predictor.f0_curve.size()) != decoder_request_size) { + throw std::runtime_error("Kokoro predictor decoder and f0 frame counts diverged"); + } + if (!prepared_decoder_ || prepared_decoder_capacity_ != decoder_request_size) { + prepare_decoder_graph_capacity(decoder_request_size); + } + if (!prepared_decoder_ || + prepared_decoder_context_.decoder_frame_capacity <= 0 || + decoder_request_size != prepared_decoder_context_.decoder_frame_capacity) { + throw std::runtime_error("Kokoro decoder runtime was not prepared"); + } + std::vector audio; + decoder_ms += measure_ms([&]() { + audio = prepared_decoder_->decode( + predictor, + input.style); + }); + const auto inference_ended = std::chrono::steady_clock::now(); + inference_ms += std::chrono::duration(inference_ended - inference_started).count(); + runtime::append_audio_buffer(merged_audio, runtime::AudioBuffer{24000, 1, std::move(audio)}); + } + + runtime::TaskResult result; + result.audio_output = std::move(merged_audio); + const double wall_ms = frontend_ms + inference_ms; + engine::debug::timing_log_scalar("kokoro.frontend_ms", frontend_ms); + engine::debug::timing_log_scalar("kokoro.inference_ms", inference_ms); + engine::debug::timing_log_scalar("kokoro.predictor_ms", predictor_ms); + engine::debug::timing_log_scalar("kokoro.decoder_ms", decoder_ms); + engine::debug::timing_log_scalar("session.wall_ms", wall_ms); + return result; +} + +} // namespace engine::models::kokoro_tts diff --git a/tools/model_manager.py b/tools/model_manager.py index b692eb77f..49aa94e29 100644 --- a/tools/model_manager.py +++ b/tools/model_manager.py @@ -182,10 +182,19 @@ def package_usage_examples(package: ModelPackage) -> list[str]: ), ModelPackage( id="kokoro_82m_bf16", - display_name="Kokoro 82M bf16", - target_directory="Kokoro-82M-bf16", - source=SnapshotSource(repo_id="mlx-community/Kokoro-82M-bf16"), - required_files=("config.json", "kokoro-v1_0.safetensors", "voices/af_heart.safetensors"), + display_name="Kokoro 82M GGML", + target_directory="kokoro-82m-v1_0-ggml", + source=SnapshotSource(repo_id="mlx-community/kokoro_mlx"), + required_files=( + "config.json", + "kokoro-v1_0.safetensors", + "voices.json", + "misaki_en/gb_gold.tsv", + "misaki_en/gb_silver.tsv", + "misaki_en/us_gold.tsv", + "misaki_en/us_silver.tsv", + "voices/af_heart.f32", + ), ), ModelPackage( id="moss_tts", From 4ec1b613913c452b6313cf13182f13d448e62cb9 Mon Sep 17 00:00:00 2001 From: mirek190 Date: Wed, 9 Sep 2026 18:48:15 +0100 Subject: [PATCH 2/5] perf(kokoro): accelerate CPU inference --- CMakeLists.txt | 8 ++ src/models/kokoro_tts/cpu_kernels.h | 103 +++++++++++++++++++ src/models/kokoro_tts/decoder.cpp | 66 ++++++++---- src/models/kokoro_tts/predictor.cpp | 24 +++-- tests/kokoro_tts/kokoro_cpu_kernel_test.cpp | 105 ++++++++++++++++++++ tests/kokoro_tts/kokoro_tts_warm_bench.cpp | 8 ++ 6 files changed, 289 insertions(+), 25 deletions(-) create mode 100644 src/models/kokoro_tts/cpu_kernels.h create mode 100644 tests/kokoro_tts/kokoro_cpu_kernel_test.cpp diff --git a/CMakeLists.txt b/CMakeLists.txt index 6a2736964..dfc6afe0b 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -508,6 +508,7 @@ if (ENGINE_BUILD_WARMBENCH) endfunction() add_engine_warmbench(chatterbox_warm_bench tests/chatterbox/chatterbox_warm_bench.cpp) + add_engine_warmbench(kokoro_tts_warm_bench tests/kokoro_tts/kokoro_tts_warm_bench.cpp) add_engine_warmbench(citrinet_asr_warm_bench tests/citrinet_asr/citrinet_asr_warm_bench.cpp) add_engine_warmbench(marblenet_vad_warm_bench tests/marblenet_vad/marblenet_vad_warm_bench.cpp) add_engine_warmbench(miocodec_warm_bench tests/miocodec/miocodec_warm_bench.cpp) @@ -554,6 +555,13 @@ if (ENGINE_BUILD_TESTS) COMMAND audio_dsp_test ) + add_engine_unittest(kokoro_cpu_kernel_test tests/kokoro_tts/kokoro_cpu_kernel_test.cpp) + + add_test( + NAME kokoro_cpu_kernel_test + COMMAND kokoro_cpu_kernel_test + ) + add_engine_unittest(rnnoise_utility_test tests/unittests/test_rnnoise_utility.cpp) target_compile_definitions(rnnoise_utility_test PRIVATE ENGINE_REPO_ROOT="${CMAKE_CURRENT_SOURCE_DIR}" diff --git a/src/models/kokoro_tts/cpu_kernels.h b/src/models/kokoro_tts/cpu_kernels.h new file mode 100644 index 000000000..8e0012143 --- /dev/null +++ b/src/models/kokoro_tts/cpu_kernels.h @@ -0,0 +1,103 @@ +#pragma once + +#include +#include + +#include + +namespace kokoro_ggml::cpu_detail { + +// Match ggml's F32 mean (sequential double accumulation, then float division) +// while evaluating independent channels in parallel and reusing the output as +// scratch for centered values. Keep every F32 rounding step explicit. +inline void kokoro_adain_cpu( + ggml_tensor * dst, + const ggml_tensor * x, + const ggml_tensor * gamma, + const ggml_tensor * beta, + int ith, + int nth, + void * userdata) { + const float eps = *static_cast(userdata); + const int64_t frames = x->ne[0]; + const int64_t rows = x->ne[1] * x->ne[2]; + for (int64_t row = rows * ith / nth; row < rows * (ith + 1) / nth; ++row) { + const auto * src = static_cast(x->data) + row * frames; + auto * out = static_cast(dst->data) + row * frames; + double sum = 0.0; + for (int64_t t = 0; t < frames; ++t) sum += static_cast(src[t]); + const float mean = static_cast(sum) / static_cast(frames); + double squares = 0.0; + for (int64_t t = 0; t < frames; ++t) { + const float centered = src[t] - mean; + out[t] = centered; + const float square = centered * centered; + squares += static_cast(square); + } + const float variance = static_cast(squares) / static_cast(frames); + const float stddev = std::sqrt(variance + eps); + const float scale = static_cast(gamma->data)[row % x->ne[1]]; + const float shift = static_cast(beta->data)[row % x->ne[1]]; + for (int64_t t = 0; t < frames; ++t) { + const float normalized = out[t] / stddev; + const float scaled = normalized * scale; + out[t] = scaled + shift; + } + } +} + +inline void kokoro_snake_cpu( + ggml_tensor * dst, + const ggml_tensor * x, + const ggml_tensor * alpha, + int ith, + int nth, + void *) { + const int64_t rows = x->ne[1] * x->ne[2]; + for (int64_t row = rows * ith / nth; row < rows * (ith + 1) / nth; ++row) { + const float a = static_cast(alpha->data)[row % x->ne[1]]; + const auto * src = reinterpret_cast(static_cast(x->data) + + (row % x->ne[1]) * x->nb[1] + (row / x->ne[1]) * x->nb[2]); + auto * out = static_cast(dst->data) + row * x->ne[0]; + for (int64_t t = 0; t < x->ne[0]; ++t) { + const float ax = src[t] * a; + const float s = std::sin(ax); + const float square = s * s; + const float fraction = square / a; + out[t] = src[t] + fraction; + } + } +} + +// Partition complete output rows, avoiding inter-thread cache-line sharing. +// This only copies/pads F32 samples; convolution reduction order is unchanged. +template +void kokoro_im2col_rows(ggml_tensor * dst, int ith, int nth, void * userdata) { + const auto & conv = *static_cast(userdata); + const ggml_tensor * input = dst->src[0]; + const int64_t begin = dst->ne[1] * ith / nth; + const int64_t end = dst->ne[1] * (ith + 1) / nth; + auto * out = static_cast(dst->data); + for (int64_t frame = begin; frame < end; ++frame) { + const int64_t base = frame * conv.stride - conv.padding; + const bool interior = base >= 0 && base + (conv.kernel - 1) * conv.dilation < input->ne[0]; + for (int64_t channel = 0; channel < conv.in_channels; ++channel) { + const auto * src = reinterpret_cast( + static_cast(input->data) + channel * input->nb[1]); + float * row = out + frame * dst->ne[0] + channel * conv.kernel; + if (conv.dilation == 1 && interior) { + std::memcpy(row, src + base, conv.kernel * sizeof(float)); + } else if (interior) { + for (int64_t k = 0; k < conv.kernel; ++k) { + row[k] = src[base + k * conv.dilation]; + } + } else { + for (int64_t k = 0; k < conv.kernel; ++k) { + const int64_t index = base + k * conv.dilation; + row[k] = index >= 0 && index < input->ne[0] ? src[index] : 0.0f; + } + } + } + } +} +} // namespace kokoro_ggml::cpu_detail diff --git a/src/models/kokoro_tts/decoder.cpp b/src/models/kokoro_tts/decoder.cpp index 9beacd56f..7d07323ac 100644 --- a/src/models/kokoro_tts/decoder.cpp +++ b/src/models/kokoro_tts/decoder.cpp @@ -1,4 +1,5 @@ #include "engine/models/kokoro_tts/decoder.h" +#include "cpu_kernels.h" #include "engine/models/kokoro_tts/assets.h" @@ -205,10 +206,16 @@ ggml_tensor * reflect_pad_left_1_bct_decoder(ggml_context * ctx, ggml_tensor * x return output.tensor; } + ggml_tensor * build_snake1d_bct_decoder( ggml_context * ctx, ggml_tensor * x, - const core::TensorValue & alpha) { + const core::TensorValue & alpha, + bool use_cpu_fastpath) { + if (use_cpu_fastpath && x->type == GGML_TYPE_F32 && + alpha.type == GGML_TYPE_F32 && ggml_is_contiguous(x) && ggml_is_contiguous(alpha.tensor)) { + return ggml_map_custom2(ctx, x, alpha.tensor, cpu_detail::kokoro_snake_cpu, GGML_N_TASKS_MAX, nullptr); + } core::ModuleBuildContext build_ctx = {}; build_ctx.ggml = ctx; const int64_t batch = x->ne[2]; @@ -230,7 +237,8 @@ ggml_tensor * build_adaptive_instance_norm_bct_decoder( const KokoroWeights::AdaIn1dWeights & weights, const core::TensorValue & style, std::vector & masks, - bool use_time_masks) { + bool use_time_masks, + bool use_cpu_fastpath) { const int64_t batch = x->ne[2]; const int64_t channels = x->ne[1]; const int64_t frames = x->ne[0]; @@ -275,6 +283,10 @@ ggml_tensor * build_adaptive_instance_norm_bct_decoder( GGML_TYPE_F32); (void)batch; (void)frames; + if (use_cpu_fastpath && !use_time_masks && x->type == GGML_TYPE_F32 && ggml_is_contiguous(x)) { + return ggml_map_custom3(ctx, x, gamma.tensor, beta.tensor, + cpu_detail::kokoro_adain_cpu, GGML_N_TASKS_MAX, const_cast(&weights.eps)); + } if (use_time_masks) { return build_masked_adain_bct(build_ctx, x, gamma, beta, channels, weights.eps, masks); } @@ -304,21 +316,22 @@ modules::ConvTranspose1dWeights make_conv_transpose1d_weights( return weights; } + template ggml_tensor * build_decoder_conv1d_bct( ggml_context * ctx, ggml_tensor * input, const ConvWeightsT & conv, - bool allow_pointwise_fastpath) { + bool allow_cpu_fastpath) { if (conv.groups != 1) { throw std::runtime_error("kokoro decoder conv1d requires groups == 1"); } - if (allow_pointwise_fastpath && + if (allow_cpu_fastpath && conv.kernel == 1 && conv.stride == 1 && conv.padding == 0 && conv.dilation == 1) { - ggml_tensor * x = ggml_cont(ctx, input); + ggml_tensor * x = ggml_is_contiguous(input) ? input : ggml_cont(ctx, input); ggml_tensor * x_2d = ggml_reshape_2d(ctx, x, input->ne[0], input->ne[1]); ggml_tensor * x_t = ggml_cont(ctx, ggml_transpose(ctx, x_2d)); ggml_tensor * w = ggml_reshape_2d(ctx, conv.weight.tensor, conv.in_channels, conv.out_channels); @@ -330,6 +343,25 @@ ggml_tensor * build_decoder_conv1d_bct( } return ggml_reshape_3d(ctx, y_2d, y_2d->ne[0], y_2d->ne[1], 1); } + // Keep the existing device/non-F32 paths. Weight objects outlive this graph. + if (allow_cpu_fastpath && input->ne[2] == 1 && + input->type == GGML_TYPE_F32 && conv.weight.tensor->type == GGML_TYPE_F32) { + const int64_t frames = (input->ne[0] + 2 * conv.padding - + conv.dilation * (conv.kernel - 1) - 1) / conv.stride + 1; + // im2col understands channel strides; only the time axis must be packed. + ggml_tensor * args[] = {input->nb[0] == sizeof(float) ? input : ggml_cont(ctx, input)}; + ggml_tensor * columns = ggml_custom_4d(ctx, GGML_TYPE_F32, + conv.in_channels * conv.kernel, frames, 1, 1, args, 1, + cpu_detail::kokoro_im2col_rows, GGML_N_TASKS_MAX, const_cast(&conv)); + ggml_tensor * weights = ggml_reshape_2d(ctx, conv.weight.tensor, + conv.in_channels * conv.kernel, conv.out_channels); + ggml_tensor * y = ggml_reshape_3d(ctx, ggml_mul_mat(ctx, columns, weights), + frames, conv.out_channels, 1); + if (conv.use_bias) { + y = ggml_add(ctx, y, ggml_reshape_3d(ctx, conv.bias->tensor, 1, conv.out_channels, 1)); + } + return y; + } core::ModuleBuildContext build_ctx = {}; build_ctx.ggml = ctx; const auto input_bct = core::wrap_tensor( @@ -868,18 +900,18 @@ ggml_tensor * build_generator_resblock( ggml_tensor * x, const KokoroWeights::GeneratorResBlockWeights & block, const core::TensorValue & style, - bool allow_cpu_pointwise_fastpath, + bool allow_cpu_fastpath, std::vector & time_masks, bool use_time_masks) { ggml_tensor * current = x; for (size_t i = 0; i < block.convs1.size(); ++i) { ggml_tensor * xt = - build_adaptive_instance_norm_bct_decoder(ctx, current, block.adain1[i], style, time_masks, use_time_masks); - xt = build_snake1d_bct_decoder(ctx, xt, block.alpha1[i]); - xt = build_decoder_conv1d_bct(ctx, xt, block.convs1[i], allow_cpu_pointwise_fastpath); - xt = build_adaptive_instance_norm_bct_decoder(ctx, xt, block.adain2[i], style, time_masks, use_time_masks); - xt = build_snake1d_bct_decoder(ctx, xt, block.alpha2[i]); - xt = build_decoder_conv1d_bct(ctx, xt, block.convs2[i], allow_cpu_pointwise_fastpath); + build_adaptive_instance_norm_bct_decoder(ctx, current, block.adain1[i], style, time_masks, use_time_masks, allow_cpu_fastpath); + xt = build_snake1d_bct_decoder(ctx, xt, block.alpha1[i], allow_cpu_fastpath); + xt = build_decoder_conv1d_bct(ctx, xt, block.convs1[i], allow_cpu_fastpath); + xt = build_adaptive_instance_norm_bct_decoder(ctx, xt, block.adain2[i], style, time_masks, use_time_masks, allow_cpu_fastpath); + xt = build_snake1d_bct_decoder(ctx, xt, block.alpha2[i], allow_cpu_fastpath); + xt = build_decoder_conv1d_bct(ctx, xt, block.convs2[i], allow_cpu_fastpath); current = ggml_add(ctx, xt, current); } return current; @@ -929,7 +961,7 @@ struct GeneratorGraphSession { } try { - const bool allow_cpu_pointwise_fastpath = !use_device_backend; + const bool allow_cpu_fastpath = !use_device_backend; decoder_x_in = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, decoder_frame_capacity, 512, 1); conditioning_in = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, conditioning_frame_capacity, 22, 1); @@ -945,13 +977,13 @@ struct GeneratorGraphSession { ggml_tensor * current = decoder_x_in; for (size_t i = 0; i < stages.size(); ++i) { const GeneratorGraphStage & stage = stages[i]; - ggml_tensor * source = build_decoder_conv1d_bct(ctx, conditioning_in, *stage.noise_conv, allow_cpu_pointwise_fastpath); + ggml_tensor * source = build_decoder_conv1d_bct(ctx, conditioning_in, *stage.noise_conv, allow_cpu_fastpath); source = build_generator_resblock( ctx, source, *stage.noise_res, style, - allow_cpu_pointwise_fastpath, + allow_cpu_fastpath, time_masks, false); @@ -971,7 +1003,7 @@ struct GeneratorGraphSession { x, *stage.resblocks[j], style, - allow_cpu_pointwise_fastpath, + allow_cpu_fastpath, time_masks, false); stage_sum = stage_sum == nullptr ? block : ggml_add(ctx, stage_sum, block); @@ -980,7 +1012,7 @@ struct GeneratorGraphSession { } current = ggml_leaky_relu(ctx, current, 0.01f, false); - output = build_decoder_conv1d_bct(ctx, current, weights->conv_post, allow_cpu_pointwise_fastpath); + output = build_decoder_conv1d_bct(ctx, current, weights->conv_post, allow_cpu_fastpath); output = ggml_cont(ctx, output); set_graph_output(output); diff --git a/src/models/kokoro_tts/predictor.cpp b/src/models/kokoro_tts/predictor.cpp index 0eadc2c12..86624a4c0 100644 --- a/src/models/kokoro_tts/predictor.cpp +++ b/src/models/kokoro_tts/predictor.cpp @@ -1,4 +1,5 @@ #include "engine/models/kokoro_tts/predictor.h" +#include "cpu_kernels.h" #include "engine/models/kokoro_tts/plbert.h" #include "engine/models/kokoro_tts/assets.h" @@ -588,7 +589,8 @@ ggml_tensor * build_adaptive_instance_norm_ct_predictor( const KokoroWeights::AdaIn1dWeights & weights, const core::TensorValue & style, std::vector & masks, - bool use_time_masks) { + bool use_time_masks, + bool use_cpu_fastpath) { const int64_t channels = x->ne[1]; core::ModuleBuildContext build_ctx = {}; build_ctx.ggml = ctx; @@ -616,6 +618,10 @@ ggml_tensor * build_adaptive_instance_norm_ct_predictor( ggml_reshape_1d(ctx, shift.tensor, channels), core::TensorShape::from_dims({channels}), GGML_TYPE_F32); + if (use_cpu_fastpath && !use_time_masks && x->type == GGML_TYPE_F32 && ggml_is_contiguous(x)) { + return ggml_map_custom3(ctx, x, gamma.tensor, beta.tensor, + cpu_detail::kokoro_adain_cpu, GGML_N_TASKS_MAX, const_cast(&weights.eps)); + } if (use_time_masks) { return build_masked_adain_ct(build_ctx, x, gamma, beta, channels, weights.eps, masks); } @@ -630,7 +636,8 @@ ggml_tensor * build_adain_resblock_ct_predictor( const core::TensorValue & style_decoder, bool use_decoder_style, std::vector & masks, - bool use_time_masks) { + bool use_time_masks, + bool use_cpu_fastpath) { const auto & style = use_decoder_style ? style_decoder : style_predictor; ggml_tensor * shortcut = x; if (block.upsample) { @@ -647,7 +654,7 @@ ggml_tensor * build_adain_resblock_ct_predictor( shortcut = build_standard_conv1d_ct_predictor(ctx, shortcut, block.conv1x1); } - ggml_tensor * residual = build_adaptive_instance_norm_ct_predictor(ctx, x, block.norm1, style, masks, use_time_masks); + ggml_tensor * residual = build_adaptive_instance_norm_ct_predictor(ctx, x, block.norm1, style, masks, use_time_masks, use_cpu_fastpath); residual = ggml_leaky_relu(ctx, residual, 0.2f, false); if (block.use_pool) { residual = build_conv_transpose1d_ct_predictor(ctx, residual, block.pool); @@ -659,7 +666,8 @@ ggml_tensor * build_adain_resblock_ct_predictor( block.norm2, style, masks, - use_time_masks); + use_time_masks, + use_cpu_fastpath); residual = ggml_leaky_relu(ctx, residual, 0.2f, false); residual = build_standard_conv1d_ct_predictor(ctx, residual, block.conv2); ggml_tensor * out = ggml_scale(ctx, ggml_add(ctx, residual, shortcut), 0.7071067811865475f); @@ -1726,7 +1734,7 @@ class PredictorTailGraphRuntime { style_decoder, false, time_masks, - false); + false, !use_device_backend); } for (size_t i = 0; i < weights->predictor.n_blocks.size(); ++i) { n = build_adain_resblock_ct_predictor( @@ -1737,7 +1745,7 @@ class PredictorTailGraphRuntime { style_decoder, false, time_masks, - false); + false, !use_device_backend); } f0_out = build_standard_conv1d_ct_predictor(ctx, f0, weights->predictor.f0_proj); @@ -1760,7 +1768,7 @@ class PredictorTailGraphRuntime { style_decoder, true, time_masks, - false); + false, !use_device_backend); ggml_tensor * asr_res = build_standard_conv1d_ct_predictor(ctx, asr_in, weights->decoder.asr_res); const auto asr_res_tv = core::wrap_tensor(asr_res, core::TensorShape::from_dims({asr_res->ne[1], asr_res->ne[0]}), GGML_TYPE_F32); @@ -1781,7 +1789,7 @@ class PredictorTailGraphRuntime { style_decoder, true, time_masks, - false); + false, !use_device_backend); if (weights->decoder.decode[i].upsample) { use_residual_conditioning = false; } diff --git a/tests/kokoro_tts/kokoro_cpu_kernel_test.cpp b/tests/kokoro_tts/kokoro_cpu_kernel_test.cpp new file mode 100644 index 000000000..0cc6a603b --- /dev/null +++ b/tests/kokoro_tts/kokoro_cpu_kernel_test.cpp @@ -0,0 +1,105 @@ +#include "../../src/models/kokoro_tts/cpu_kernels.h" + +#include + +#include +#include +#include + +struct Conv { + int64_t in_channels; + int64_t kernel; + int64_t stride; + int64_t padding; + int64_t dilation; +}; + +void compare(ggml_context * ctx, ggml_tensor * expected, ggml_tensor * actual, int threads) { + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, expected); + ggml_build_forward_expand(graph, actual); + for (int repeat = 0; repeat < 2; ++repeat) { + if (ggml_graph_compute_with_ctx(ctx, graph, threads) != GGML_STATUS_SUCCESS || + ggml_nbytes(expected) != ggml_nbytes(actual) || + std::memcmp(expected->data, actual->data, ggml_nbytes(expected)) != 0) { + throw std::runtime_error("CPU kernel is not bit-exact with ggml reference"); + } + } +} + +int main() { + size_t cases = 0; + for (int threads : {1, 8}) { + for (int width : {1, 7, 33, 129}) for (int channels : {1, 3, 16}) + for (int kernel : {1, 3, 7}) for (int stride : {1, 2}) + for (int dilation : {1, 3}) for (int padding : {0, 8}) for (int gap : {0, 5}) { + const int numerator = width + 2 * padding - dilation * (kernel - 1) - 1; + if (numerator < 0) { + continue; + } + ggml_context * ctx = ggml_init({16 * 1024 * 1024, nullptr, false}); + Conv conv{channels, kernel, stride, padding, dilation}; + ggml_tensor * storage = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, width + gap, channels, 1); + for (int64_t i = 0; i < ggml_nelements(storage); ++i) + static_cast(storage->data)[i] = static_cast((i * 17) % 101 - 50) / 13.0f; + ggml_tensor * input = ggml_view_3d(ctx, storage, width, channels, 1, + storage->nb[1], storage->nb[2], 0); + ggml_tensor * weight = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, kernel, channels, 2); + ggml_tensor * expected = ggml_im2col(ctx, weight, input, stride, 0, padding, 0, + dilation, 0, false, GGML_TYPE_F32); + ggml_tensor * args[] = {input}; + ggml_tensor * actual = ggml_custom_4d(ctx, GGML_TYPE_F32, kernel * channels, + numerator / stride + 1, 1, 1, args, 1, + kokoro_ggml::cpu_detail::kokoro_im2col_rows, GGML_N_TASKS_MAX, &conv); + compare(ctx, expected, actual, threads); + ggml_free(ctx); + ++cases; + } + for (int width : {1, 31, 127}) for (int channels : {1, 3, 16}) { + ggml_context * ctx = ggml_init({16 * 1024 * 1024, nullptr, false}); + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, width, channels, 2); + ggml_tensor * alpha = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 1, channels, 1); + for (int64_t i = 0; i < ggml_nelements(input); ++i) + static_cast(input->data)[i] = static_cast((i * 17) % 101 - 50) / 13.0f; + for (int i = 0; i < channels; ++i) { + static_cast(alpha->data)[i] = 0.1f + i * 0.37f; + } + ggml_tensor * sine = ggml_sin(ctx, ggml_mul(ctx, input, alpha)); + ggml_tensor * expected = ggml_add(ctx, input, ggml_div(ctx, ggml_mul(ctx, sine, sine), alpha)); + ggml_tensor * actual = ggml_map_custom2(ctx, input, alpha, + kokoro_ggml::cpu_detail::kokoro_snake_cpu, GGML_N_TASKS_MAX, nullptr); + compare(ctx, expected, actual, threads); + ggml_free(ctx); + ++cases; + } + } + for (int threads : {1, 8}) for (int width : {1, 31, 127, 4097}) + for (int channels : {1, 3, 16}) for (int mode : {0, 1, 2}) { + ggml_context * ctx = ggml_init({16 * 1024 * 1024, nullptr, false}); + ggml_tensor * input = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, width, channels, 2); + ggml_tensor * gamma = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 1, channels, 1); + ggml_tensor * beta = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, 1, channels, 1); + for (int64_t i = 0; i < ggml_nelements(input); ++i) { + static_cast(input->data)[i] = mode == 0 ? 1.0f : + static_cast((i * 17) % 101 - 50) * (mode == 1 ? 0.17f : 1.e-8f); + } + for (int i = 0; i < channels; ++i) { + static_cast(gamma->data)[i] = 0.1f + i * 0.37f; + static_cast(beta->data)[i] = -0.2f + i * 0.03f; + } + float eps = 1.e-5f; + ggml_tensor * mean = ggml_mean(ctx, input); + ggml_tensor * centered = ggml_sub(ctx, input, ggml_repeat(ctx, mean, input)); + ggml_tensor * variance = ggml_mean(ctx, ggml_mul(ctx, centered, centered)); + ggml_tensor * stddev = ggml_sqrt(ctx, ggml_scale_bias(ctx, variance, 1.0f, eps)); + ggml_tensor * normalized = ggml_div(ctx, centered, ggml_repeat(ctx, stddev, input)); + ggml_tensor * expected = ggml_add(ctx, ggml_mul(ctx, normalized, + ggml_repeat(ctx, gamma, input)), ggml_repeat(ctx, beta, input)); + ggml_tensor * actual = ggml_map_custom3(ctx, input, gamma, beta, + kokoro_ggml::cpu_detail::kokoro_adain_cpu, GGML_N_TASKS_MAX, &eps); + compare(ctx, expected, actual, threads); + ggml_free(ctx); + ++cases; + } + std::cout << "PASS: " << cases << " bit-exact kernel cases, each executed twice (1/8 threads).\n"; +} diff --git a/tests/kokoro_tts/kokoro_tts_warm_bench.cpp b/tests/kokoro_tts/kokoro_tts_warm_bench.cpp index 3d7a72e4b..ddd018a2d 100644 --- a/tests/kokoro_tts/kokoro_tts_warm_bench.cpp +++ b/tests/kokoro_tts/kokoro_tts_warm_bench.cpp @@ -226,10 +226,13 @@ std::vector read_timing_lines(const std::filesystem::path & path) { } void clear_file(const std::filesystem::path & path) { + engine::debug::reset_logging(); std::ofstream output(path, std::ios::trunc); if (!output.is_open()) { throw std::runtime_error("failed to clear timing log: " + path.string()); } + output.close(); + engine::debug::configure_logging({true, path.string()}); } void write_sectioned_timing_log( @@ -424,7 +427,11 @@ int main(int argc, char ** argv) { for (size_t request_index = 0; request_index < requests.size(); ++request_index) { for (int i = 0; i < iterations; ++i) { clear_file(timing_path); + const auto request_start = std::chrono::steady_clock::now(); last_results[request_index] = session->run(requests[request_index]); + const double request_ms = std::chrono::duration( + std::chrono::steady_clock::now() - request_start).count(); + engine::debug::timing_log_scalar("benchmark.request_wall_ms", request_ms); const auto metrics = parse_timing_file(timing_path); auto lines = read_timing_lines(timing_path); log_sections.push_back({ @@ -437,6 +444,7 @@ int main(int argc, char ** argv) { } } + engine::debug::reset_logging(); write_sectioned_timing_log(timing_path, log_sections); for (size_t request_index = 0; request_index < requests.size(); ++request_index) { From 5631034e4bd99c7c10febe35e6c63ce245a93937 Mon Sep 17 00:00:00 2001 From: mirek190 Date: Thu, 10 Sep 2026 10:55:29 +0100 Subject: [PATCH 3/5] feat(kokoro): add standalone multilingual GGUF packages --- CMakeLists.txt | 8 +- include/engine/models/kokoro_tts/assets.h | 2 + .../models/kokoro_tts/g2p_multilingual.h | 15 + include/engine/models/kokoro_tts/package.h | 17 + src/models/kokoro_tts/assets.cpp | 11 +- src/models/kokoro_tts/frontend.cpp | 20 +- src/models/kokoro_tts/g2p_multilingual.cpp | 401 ++++++++++++++++++ src/models/kokoro_tts/loader.cpp | 11 +- src/models/kokoro_tts/package.cpp | 178 ++++++++ tests/kokoro_tts/MULTILINGUAL_GGUF.md | 84 ++++ tests/kokoro_tts/compare_gguf_quality.py | 148 +++++++ tests/kokoro_tts/compare_multilingual_g2p.py | 34 ++ tests/kokoro_tts/kokoro_g2p_probe.cpp | 17 + tests/kokoro_tts/multilingual_cases.json | 14 + .../validate_multilingual_packages.py | 55 +++ tools/prepare_kokoro_gguf.py | 141 ++++++ 16 files changed, 1140 insertions(+), 16 deletions(-) create mode 100644 include/engine/models/kokoro_tts/g2p_multilingual.h create mode 100644 include/engine/models/kokoro_tts/package.h create mode 100644 src/models/kokoro_tts/g2p_multilingual.cpp create mode 100644 src/models/kokoro_tts/package.cpp create mode 100644 tests/kokoro_tts/MULTILINGUAL_GGUF.md create mode 100644 tests/kokoro_tts/compare_gguf_quality.py create mode 100644 tests/kokoro_tts/compare_multilingual_g2p.py create mode 100644 tests/kokoro_tts/kokoro_g2p_probe.cpp create mode 100644 tests/kokoro_tts/multilingual_cases.json create mode 100644 tests/kokoro_tts/validate_multilingual_packages.py create mode 100644 tools/prepare_kokoro_gguf.py diff --git a/CMakeLists.txt b/CMakeLists.txt index dfc6afe0b..e6955d923 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -196,6 +196,8 @@ add_library(engine_runtime STATIC src/models/omnivoice/session.cpp src/models/omnivoice/loader.cpp src/models/kokoro_tts/assets.cpp + src/models/kokoro_tts/package.cpp + src/models/kokoro_tts/g2p_multilingual.cpp src/models/kokoro_tts/loader.cpp src/models/kokoro_tts/frontend.cpp src/models/kokoro_tts/g2p_en.cpp @@ -402,7 +404,10 @@ target_include_directories(engine_runtime PRIVATE ) target_link_libraries(engine_runtime PUBLIC ggml) -target_link_libraries(engine_runtime PRIVATE sentencepiece cjson_vendor yaml_vendor) +target_link_libraries(engine_runtime PRIVATE sentencepiece cjson_vendor yaml_vendor ${CMAKE_DL_LIBS}) +if (MSVC) + set_source_files_properties(src/models/kokoro_tts/g2p_multilingual.cpp PROPERTIES COMPILE_OPTIONS /utf-8) +endif() if (ENGINE_ENABLE_OPENMP) target_link_libraries(engine_runtime PRIVATE OpenMP::OpenMP_CXX) endif() @@ -556,6 +561,7 @@ if (ENGINE_BUILD_TESTS) ) add_engine_unittest(kokoro_cpu_kernel_test tests/kokoro_tts/kokoro_cpu_kernel_test.cpp) + add_engine_unittest(kokoro_g2p_probe tests/kokoro_tts/kokoro_g2p_probe.cpp) add_test( NAME kokoro_cpu_kernel_test diff --git a/include/engine/models/kokoro_tts/assets.h b/include/engine/models/kokoro_tts/assets.h index 480515320..8251820cd 100644 --- a/include/engine/models/kokoro_tts/assets.h +++ b/include/engine/models/kokoro_tts/assets.h @@ -259,6 +259,8 @@ struct KokoroVoicePack { }; struct KokoroAssets { + std::shared_ptr package; + std::shared_ptr multilingual_g2p; std::filesystem::path model_root; engine::io::json::Value config; std::shared_ptr model_weights; diff --git a/include/engine/models/kokoro_tts/g2p_multilingual.h b/include/engine/models/kokoro_tts/g2p_multilingual.h new file mode 100644 index 000000000..e8da12969 --- /dev/null +++ b/include/engine/models/kokoro_tts/g2p_multilingual.h @@ -0,0 +1,15 @@ +#pragma once +#include +#include +#include +namespace engine::models::kokoro_tts { +class MultilingualG2P { +public: + explicit MultilingualG2P(const std::filesystem::path & root); + ~MultilingualG2P(); + std::string phonemize(const std::string & text, const std::string & language) const; +private: + struct Impl; + std::unique_ptr impl_; +}; +} diff --git a/include/engine/models/kokoro_tts/package.h b/include/engine/models/kokoro_tts/package.h new file mode 100644 index 000000000..22c6d4a0d --- /dev/null +++ b/include/engine/models/kokoro_tts/package.h @@ -0,0 +1,17 @@ +#pragma once + +#include "engine/framework/assets/tensor_source.h" +#include +#include + +namespace engine::models::kokoro_tts { +// Owns materialized resources for the lifetime of a standalone model. +struct KokoroPackage { + std::filesystem::path root; + std::shared_ptr weights; + bool temporary = false; + ~KokoroPackage(); +}; +std::shared_ptr open_kokoro_package(const std::filesystem::path & path); +bool is_kokoro_gguf(const std::filesystem::path & path) noexcept; +} diff --git a/src/models/kokoro_tts/assets.cpp b/src/models/kokoro_tts/assets.cpp index bbbd54090..19fa28282 100644 --- a/src/models/kokoro_tts/assets.cpp +++ b/src/models/kokoro_tts/assets.cpp @@ -3,6 +3,8 @@ #include "engine/framework/io/binary.h" #include "engine/framework/io/json.h" #include "engine/models/kokoro_tts/assets.h" +#include "engine/models/kokoro_tts/package.h" +#include "engine/models/kokoro_tts/g2p_multilingual.h" #include "engine/models/kokoro_tts/g2p_en.h" @@ -715,6 +717,7 @@ namespace engine::models::kokoro_tts { namespace { struct KokoroAssetResources { + std::shared_ptr package; assets::ResourceBundle bundle; io::json::Value config; io::json::Value voices; @@ -723,11 +726,11 @@ struct KokoroAssetResources { KokoroAssetResources load_asset_resources(const std::filesystem::path & model_root) { KokoroAssetResources resources; - resources.bundle = assets::ResourceBundle(std::filesystem::weakly_canonical(model_root)); + resources.package = open_kokoro_package(model_root); + resources.bundle = assets::ResourceBundle(resources.package->root); resources.bundle.add_model_files({ {"config", "config.json"}, {"voices", "voices.json"}, - {"weights", "kokoro-v1_0.safetensors"}, }); resources.config = resources.bundle.parse_json("config"); if (!resources.config.is_object()) { @@ -737,7 +740,7 @@ KokoroAssetResources load_asset_resources(const std::filesystem::path & model_ro if (!resources.voices.is_object()) { throw std::runtime_error("Kokoro voices.json root must be an object"); } - resources.weights = resources.bundle.open_tensor_source("weights"); + resources.weights = resources.package->weights; return resources; } @@ -769,6 +772,8 @@ std::shared_ptr load_kokoro_assets(const std::filesystem::pa const auto & root = resources.bundle.model_root(); auto assets = std::make_shared(); + assets->package = resources.package; + assets->multilingual_g2p = std::make_shared(root); assets->model_root = root; assets->config = std::move(resources.config); assets->model_weights = std::move(resources.weights); diff --git a/src/models/kokoro_tts/frontend.cpp b/src/models/kokoro_tts/frontend.cpp index db2846848..b59fb398c 100644 --- a/src/models/kokoro_tts/frontend.cpp +++ b/src/models/kokoro_tts/frontend.cpp @@ -1,6 +1,7 @@ #include "engine/models/kokoro_tts/frontend.h" #include "engine/models/kokoro_tts/g2p_en.h" +#include "engine/models/kokoro_tts/g2p_multilingual.h" #include #include @@ -46,7 +47,13 @@ std::string resolve_language_code_alias(const std::string & value) { if (normalized == "b" || normalized == "en-gb" || normalized == "gb" || normalized == "uk" || normalized == "british" || normalized == "british english") { return "b"; } - throw std::runtime_error("unsupported Kokoro language: " + value + " (English only: en-us/a or en-gb/b)"); + for (const auto & pair : std::vector>{ + {"e", "e"}, {"es", "e"}, {"es-es", "e"}, {"f", "f"}, {"fr", "f"}, {"fr-fr", "f"}, + {"h", "h"}, {"hi", "h"}, {"hi-in", "h"}, {"i", "i"}, {"it", "i"}, {"it-it", "i"}, + {"j", "j"}, {"ja", "j"}, {"ja-jp", "j"}, {"p", "p"}, {"pt", "p"}, {"pt-br", "p"}, + {"z", "z"}, {"zh", "z"}, {"zh-cn", "z"}, {"cmn", "z"}}) + if (normalized == pair.first) return pair.second; + throw std::runtime_error("unsupported Kokoro language: " + value); } std::string voice_language_code(const std::string & voice_id) { @@ -67,11 +74,7 @@ std::string resolve_voice_id( throw std::runtime_error("unknown Kokoro voice id: " + voice_id); } const std::string language_code = voice_language_code(voice_id); - if (language_code != "a" && language_code != "b") { - throw std::runtime_error( - "Kokoro currently supports only English voices; got voice id " + voice_id + - " with lang_code=" + language_code); - } + (void) resolve_language_code_alias(language_code); return voice_id; } @@ -111,9 +114,8 @@ std::string phonemize_text( } return (*g2p)(text.text).first; } - throw std::runtime_error( - "unsupported Kokoro language code: " + language_code + - " (English only: a/en-us or b/en-gb)"); + if (!assets.multilingual_g2p) throw std::runtime_error("Kokoro multilingual resources were not prepared"); + return assets.multilingual_g2p->phonemize(text.text, language_code); } struct EncodedInputIds { diff --git a/src/models/kokoro_tts/g2p_multilingual.cpp b/src/models/kokoro_tts/g2p_multilingual.cpp new file mode 100644 index 000000000..d58ae1176 --- /dev/null +++ b/src/models/kokoro_tts/g2p_multilingual.cpp @@ -0,0 +1,401 @@ +#include "engine/models/kokoro_tts/g2p_multilingual.h" +#include "engine/framework/io/json.h" +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#ifdef _WIN32 +#define NOMINMAX +#include +#else +#include +#endif + +namespace engine::models::kokoro_tts { +namespace { +using engine::io::json::Value; +using U = std::u32string; +U decode(const std::string & s) { return std::wstring_convert, char32_t>{}.from_bytes(s); } +std::string encode(const U & s) { return std::wstring_convert, char32_t>{}.to_bytes(s); } +void replace(std::string & s, const std::string & from, const std::string & to) { + size_t p = 0; + while ((p = s.find(from, p)) != std::string::npos) { s.replace(p, from.size(), to); p += to.size(); } +} +std::string spaces(const std::string & s) { + auto out = std::regex_replace(s, std::regex("[ \\t\\r\\n]+"), " "); + auto start = out.find_first_not_of(' '); + return start == std::string::npos ? "" : out.substr(start, out.find_last_not_of(' ') - start + 1); +} +bool han(char32_t c) { return c >= 0x4e00 && c <= 0x9fff; } + +class Library { +#ifdef _WIN32 + HMODULE handle_ = nullptr; +#else + void * handle_ = nullptr; +#endif +public: + Library(const char * env, const char * windows_name, const char * unix_name) { + const char * override_path = std::getenv(env); +#ifdef _WIN32 + std::filesystem::path path; + if (override_path && *override_path) path = std::filesystem::u8path(override_path); + else { + std::wstring exe(32768, L'\0'); + const auto n = GetModuleFileNameW(nullptr, exe.data(), static_cast(exe.size())); + if (!n || n >= exe.size()) throw std::runtime_error("Cannot locate the audio.cpp executable"); + exe.resize(n); + path = std::filesystem::path(exe).parent_path() / windows_name; + } + handle_ = LoadLibraryExW(path.c_str(), nullptr, LOAD_LIBRARY_SEARCH_DLL_LOAD_DIR | LOAD_LIBRARY_SEARCH_DEFAULT_DIRS); +#else + (void) windows_name; + handle_ = dlopen(override_path && *override_path ? override_path : unix_name, RTLD_NOW | RTLD_LOCAL); +#endif + if (!handle_) throw std::runtime_error(std::string("Missing Kokoro pronunciation library; set ") + env); + } + ~Library() { +#ifdef _WIN32 + if (handle_) FreeLibrary(handle_); +#else + if (handle_) dlclose(handle_); +#endif + } + template F symbol(const char * name) { +#ifdef _WIN32 + auto p = GetProcAddress(handle_, name); +#else + auto p = dlsym(handle_, name); +#endif + if (!p) throw std::runtime_error(std::string("Missing pronunciation API: ") + name); + return reinterpret_cast(p); + } +}; + +std::string espeak_text(const std::string & text, const std::string & language, const std::filesystem::path & root) { + // eSpeak's selected voice and dictionary state are process-global. + static std::mutex mutex; + const std::lock_guard lock(mutex); + static Library lib("AUDIOCPP_ESPEAK_LIBRARY", "espeak-ng.dll", "libespeak-ng.so.1"); + auto initialize = lib.symbol("espeak_Initialize"); + auto terminate = lib.symbol("espeak_Terminate"); + auto voice = lib.symbol("espeak_SetVoiceByName"); + auto phonemes = lib.symbol("espeak_TextToPhonemes"); + if (initialize(2, 0, root.u8string().c_str(), 0) < 0) throw std::runtime_error("Cannot initialize Kokoro eSpeak data"); + struct End { int (*fn)(); ~End() { fn(); } } end{terminate}; + if (voice((language == "fr-fr" ? "fr" : language).c_str()) != 0) + throw std::runtime_error("Missing eSpeak language: " + language); + // Preserve punctuation ourselves: TextToPhonemes consumes clause punctuation. + const U punctuation = U";:,.!?¡¿—…\"«»“”()"; + std::string out, chunk; + auto flush = [&] { + if (chunk.empty()) return; + const void * cursor = chunk.c_str(); + while (cursor) { + const void * previous = cursor; + const auto * ps = phonemes(&cursor, 1, 2 | (1 << 7) | ('^' << 8)); + if (ps) { + if (!out.empty() && out.back() != ' ' && out != u8"¿" && out != u8"¡") out += ' '; + out += ps; + } + if (cursor == previous) throw std::runtime_error("eSpeak made no progress"); + } + chunk.clear(); + }; + for (char32_t c : decode(text)) { + if (punctuation.find(c) != U::npos) { + const bool spaced = !chunk.empty() && chunk.back() == ' '; + flush(); if (spaced && !out.empty() && out.back() != ' ') out += ' '; + out += encode(U(1, c)); + } + else chunk += encode(U(1, c)); + } + flush(); + out = std::regex_replace(out, std::regex("\\([a-z-]+\\)"), ""); + for (const auto & pair : std::vector>{ + {u8"a^ɪ", "I"}, {u8"a^ʊ", "W"}, {"d^z", u8"ʣ"}, {u8"d^ʒ", u8"ʤ"}, + {u8"e^ɪ", "A"}, {u8"o^ʊ", "O"}, {u8"ə^ʊ", "Q"}, {"s^s", "S"}, + {"t^s", u8"ʦ"}, {u8"t^ʃ", u8"ʧ"}, {u8"ɔ^ɪ", "Y"}}) replace(out, pair.first, pair.second); + replace(out, "^", ""); replace(out, "-", ""); + return spaces(out); +} + +std::vector split(const std::string & s, char delim) { + std::vector out; + std::istringstream in(s); std::string item; + while (std::getline(in, item, delim)) out.push_back(item); + return out; +} +std::string cjk_number(uint64_t n, bool ja) { + const U digits = ja ? U"零一二三四五六七八九" : U"零一二三四五六七八九"; + if (n < 10) return encode(U(1, digits[n])); + const std::vector> units = { + {100000000, ja ? U"億" : U"亿"}, {10000, U"万"}, {1000, U"千"}, {100, U"百"}, {10, U"十"}}; + for (const auto & [unit, name] : units) if (n >= unit) { + auto head = n / unit; auto rest = n % unit; + auto out = ((head == 1 && (ja ? unit < 10000 : unit == 10)) ? "" : cjk_number(head, ja)) + encode(name); + if (rest) { + if (!ja && unit >= 100 && rest < unit / 10) out += u8"零"; + out += cjk_number(rest, ja); + } + return out; + } + return ""; +} +std::string normalize_numbers(const std::string & text, bool ja) { + std::string out; + const auto s = decode(text); + for (size_t i = 0; i < s.size();) { + char32_t c = s[i]; + if (c >= 0xff01 && c <= 0xff5e) c -= 0xfee0; + if (c < U'0' || c > U'9') { out += encode(U(1, c)); ++i; continue; } + std::string number; + while (i < s.size()) { + c = s[i]; if (c >= 0xff10 && c <= 0xff19) c -= 0xfee0; + if (c < U'0' || c > U'9') break; + number += static_cast(c); ++i; + } + if (number.size() > 12) throw std::runtime_error("Kokoro CJK number exceeds 12 digits; spell it out"); + out += cjk_number(std::stoull(number), ja); + } + return out; +} +} + +struct MultilingualG2P::Impl { + std::filesystem::path root; + mutable std::once_flag ja_once, zh_once; + mutable std::unordered_map kana; + mutable std::set ja_words; + mutable Value zh; + explicit Impl(std::filesystem::path path) : root(std::move(path)) {} + void load_ja() const { + std::call_once(ja_once, [&] { + auto value = engine::io::json::parse_file(root / "g2p/ja.json"); + for (const auto & [k, v] : value.require("kana").as_object()) kana.emplace(decode(k), decode(v.as_string())); + for (const auto & v : value.require("words").as_array()) ja_words.insert(v.as_string()); + }); + } + std::string japanese(const std::string & text) const { + load_ja(); + static Library lib("AUDIOCPP_MECAB_LIBRARY", "libmecab.dll", "libmecab.so.2"); + auto create = lib.symbol("mecab_new"); + auto destroy = lib.symbol("mecab_destroy"); + // Prefix of MeCab's stable C node ABI; later cost fields are not accessed. + struct Node { + Node * prev; Node * next; Node * enext; Node * bnext; + void * rpath; void * lpath; + const char * surface; const char * feature; + unsigned int id; + unsigned short length, rlength, rcAttr, lcAttr, posid; + unsigned char char_type, stat; + }; + auto parse = lib.symbol("mecab_sparse_tonode"); + auto error = lib.symbol("mecab_strerror"); + std::vector args = {"mecab", "-r", (root / "unidic/dicrc").u8string(), + "-d", (root / "unidic").u8string()}; + std::vector argv; for (auto & arg : args) argv.push_back(arg.data()); + void * tagger = create(static_cast(argv.size()), argv.data()); + if (!tagger) throw std::runtime_error(std::string("Cannot load Japanese dictionary: ") + error(nullptr)); + struct End { void * p; void (*fn)(void *); ~End() { fn(p); } } end{tagger, destroy}; + const auto normalized = normalize_numbers(text, true); + const Node * result = parse(tagger, normalized.c_str()); + if (!result) throw std::runtime_error("Japanese tokenization failed"); + struct Word { std::string surface; U reading; int type; }; + std::vector words; + for (auto * node = result; node; node = node->next) { + if (node->stat >= 2) continue; + std::string surface(node->surface, node->length); + auto fields = split(node->feature, ','); + auto reading_field = [&](size_t index) { + return index < fields.size() ? fields[index] : std::string{}; + }; + auto pron = reading_field(9); if (pron.empty()) pron = reading_field(17); + if (pron.empty()) pron = surface; + if (pron == "*") continue; // UniDic's no-pronunciation marker, as in Misaki. + U reading = decode(pron); + for (auto & c : reading) if (c >= 0x30a1 && c <= 0x30f6) c -= 0x60; + int type = node->char_type; + words.push_back({surface, std::move(reading), type == 7 || node->stat == 0 ? 6 : type}); + } + // Preserve Misaki's lexical grouping before inserting spaces. + for (size_t i = 0; i < words.size(); ++i) { + std::string combined; size_t last = i; + for (size_t j = i; j < words.size() && words[j].type == words[i].type; ++j) { + combined += words[j].surface; + if (ja_words.count(combined)) last = j; + } + while (last > i) { + words[i].surface += words[i + 1].surface; + words[i].reading += words[i + 1].reading; + words.erase(words.begin() + i + 1); --last; + } + } + auto mapping = [&](const U & key) { auto it = kana.find(key); return it == kana.end() ? U{} : it->second; }; + U output; + for (const auto & w : words) { + U phonemes; + auto surface = decode(w.surface); + bool ascii = std::all_of(surface.begin(), surface.end(), [](char32_t c) { return c < 128; }); + if (ascii) phonemes = surface; + else if (w.type != 6 && w.type != 3) continue; + else for (size_t i = 0; i < w.reading.size(); ++i) { + auto c = w.reading[i]; U key(1, c); + auto prev = i ? w.reading.substr(i - 1, 2) : U{}; + auto next = w.reading.substr(i, 2); + if (prev.size() == 2 && kana.count(prev)) { phonemes += mapping(prev); continue; } + if (next.size() == 2 && kana.count(next)) continue; + if (c == U'ー') { phonemes += U"ː"; continue; } + if (c == U'っ') { phonemes += U"ʔ"; continue; } + if (c == U'ん') { + auto after = i + 1 < w.reading.size() ? mapping(w.reading.substr(i + 1, 1)) : U{}; + phonemes += after.empty() ? U"ɴ" : U(U"mpb").find(after[0]) != U::npos ? U"m" : + U(U"kɡ").find(after[0]) != U::npos ? U"ŋ" : U(U"ɲʨʥ").find(after[0]) != U::npos ? U"ɲ" : + U(U"ntdɾz").find(after[0]) != U::npos ? U"n" : U"ɴ"; + continue; + } + if (U(U"ゃゅょぁぃぅぇぉ").find(c) != U::npos) continue; + auto ps = mapping(key); + if (ps.empty() && c != U' ' && c != 0x3099 && c != 0x309a) + throw std::runtime_error("Japanese pronunciation missing for: " + encode(key)); + phonemes += ps; + } + if (phonemes.empty()) continue; + if (phonemes.size() == 1 && U(U"]).,?!:").find(phonemes[0]) != U::npos) + while (!output.empty() && output.back() == U' ') output.pop_back(); + output += phonemes; + if (!(phonemes.size() == 1 && U(U"([“").find(phonemes[0]) != U::npos)) output += U' '; + } + auto out = spaces(encode(output)); + replace(out, "(", u8"«"); replace(out, ")", u8"»"); + return out; + } + + std::string chinese(const std::string & text) const { + std::call_once(zh_once, [&] { zh = engine::io::json::parse_file(root / "g2p/zh.json"); }); + const auto & frequencies = zh.require("frequency").as_object(); + const auto & characters = zh.require("chars").as_object(); + const auto & phrases = zh.require("phrases").as_object(); + const auto & ipa = zh.require("ipa").as_object(); + const double log_total = std::log(static_cast(zh.require("total").as_i64())); + U text32 = decode(normalize_numbers(text, false)); + std::string out; + for (size_t start = 0; start < text32.size();) { + if (!han(text32[start])) { + const U punct = U"、,。.!:;?《》「」()"; + const U mapped = U",,..!:;?“”“”()"; + auto p = punct.find(text32[start]); + out += encode(U(1, p == U::npos ? text32[start] : mapped[p])); + if (p != U::npos || U(U",.!?:;").find(text32[start]) != U::npos) out += ' '; + ++start; continue; + } + size_t end = start; while (end < text32.size() && han(text32[end])) ++end; + U span = text32.substr(start, end - start); + // Jieba's maximum-probability dictionary segmentation. + std::vector score(span.size() + 1, 0); + std::vector next(span.size()); + for (size_t i = span.size(); i-- > 0;) { + score[i] = -1e300; next[i] = i + 1; + for (size_t j = i + 1; j <= span.size() && j - i <= 32; ++j) { + auto found = frequencies.find(encode(span.substr(i, j - i))); + if (j != i + 1 && found == frequencies.end()) continue; + double frequency = found == frequencies.end() ? 1 : static_cast(found->second.as_i64()); + double candidate = std::log(frequency) - log_total + score[j]; + if (candidate >= score[i]) { score[i] = candidate; next[i] = j; } + } + } + // Run the same four-state Jieba HMM on consecutive singleton DAG words. + std::vector words; + auto flush_singletons = [&](const U & buffer) { + if (buffer.empty()) return; + if (buffer.size() == 1 || frequencies.count(encode(buffer))) { + for (char32_t c : buffer) words.push_back(U(1, c)); + return; + } + const std::string states = "BMES"; + const std::array previous = {"ES", "MB", "BM", "SE"}; + std::vector> probs(buffer.size()); + std::vector> back(buffer.size()); + auto probability = [](const Value & table, const std::string & key) { + const auto * v = table.find(key); return v ? v->as_number() : -3.14e100; + }; + const auto & emission = *zh.find("emission"); + const auto & transition = *zh.find("transition"); + for (int y = 0; y < 4; ++y) + probs[0][y] = probability(*zh.find("start"), states.substr(y, 1)) + + probability(*emission.find(states.substr(y, 1)), encode(buffer.substr(0, 1))); + for (size_t t = 1; t < buffer.size(); ++t) for (int y = 0; y < 4; ++y) { + const auto state = states.substr(y, 1); + const auto emit = probability(*emission.find(state), encode(buffer.substr(t, 1))); + double best = -1e300; char best_char = 0; int best_index = 0; + for (char prev : previous[y]) { + int k = static_cast(states.find(prev)); + double value = probs[t - 1][k] + probability(*transition.find(std::string(1, prev)), state) + emit; + if (value > best || (value == best && prev > best_char)) { + best = value; best_char = prev; best_index = k; + } + } + probs[t][y] = best; back[t][y] = best_index; + } + int state = probs.back()[2] > probs.back()[3] ? 2 : 3; + std::vector path(buffer.size()); + for (size_t t = buffer.size(); t-- > 0;) { path[t] = state; if (t) state = back[t][state]; } + size_t begin = 0, consumed = 0; + for (size_t t = 0; t < path.size(); ++t) { + if (states[path[t]] == 'B') begin = t; + else if (states[path[t]] == 'E') { words.push_back(buffer.substr(begin, t - begin + 1)); consumed = t + 1; } + else if (states[path[t]] == 'S') { words.push_back(buffer.substr(t, 1)); consumed = t + 1; } + } + if (consumed < buffer.size()) words.push_back(buffer.substr(consumed)); + }; + U buffer; + for (size_t i = 0; i < span.size();) { + U word = span.substr(i, next[i] - i); + if (word.size() == 1) buffer += word; + else { flush_singletons(buffer); buffer.clear(); words.push_back(word); } + i = next[i]; + } + flush_singletons(buffer); + for (size_t i = 0; i < words.size(); ++i) { + const U & word = words[i]; + const auto phrase = phrases.find(encode(word)); + for (size_t k = 0; k < word.size(); ++k) { + std::string syllable; + if (phrase != phrases.end()) syllable = phrase->second.as_array().at(k).as_string(); + else { + auto ch = characters.find(encode(U(1, word[k]))); + if (ch == characters.end()) throw std::runtime_error("Missing Chinese pronunciation: " + encode(U(1, word[k]))); + syllable = ch->second.as_string(); + } + auto found = ipa.find(syllable); + if (found == ipa.end()) throw std::runtime_error("Missing Chinese pinyin mapping: " + syllable); + out += found->second.as_string(); + } + if (i + 1 < words.size()) out += ' '; + } + start = end; + } + return spaces(out); + } +}; +MultilingualG2P::MultilingualG2P(const std::filesystem::path & root) : impl_(std::make_unique(root)) {} +MultilingualG2P::~MultilingualG2P() = default; +std::string MultilingualG2P::phonemize(const std::string & text, const std::string & language) const { + if (language == "j") return impl_->japanese(text); + if (language == "z") return impl_->chinese(text); + static const std::map langs = {{"e", "es"}, {"f", "fr-fr"}, + {"h", "hi"}, {"i", "it"}, {"p", "pt-br"}}; + auto it = langs.find(language); + if (it == langs.end()) throw std::runtime_error("Unsupported Kokoro language: " + language); + return espeak_text(text, it->second, impl_->root); +} +} diff --git a/src/models/kokoro_tts/loader.cpp b/src/models/kokoro_tts/loader.cpp index 2f1184529..e00387e76 100644 --- a/src/models/kokoro_tts/loader.cpp +++ b/src/models/kokoro_tts/loader.cpp @@ -1,4 +1,5 @@ #include "engine/models/kokoro_tts/loader.h" +#include "engine/models/kokoro_tts/package.h" #include "engine/models/kokoro_tts/session.h" #include "engine/framework/io/filesystem.h" @@ -13,7 +14,9 @@ std::filesystem::path resolve_model_root(const std::filesystem::path & model_pat if (engine::io::is_existing_directory(model_path)) { return std::filesystem::weakly_canonical(model_path); } - throw std::runtime_error("Kokoro TTS expects a model directory: " + model_path.string()); + if (engine::io::is_existing_file(model_path) && model_path.extension() == ".gguf") + return std::filesystem::weakly_canonical(model_path); + throw std::runtime_error("Kokoro TTS expects a model directory or GGUF: " + model_path.string()); } std::vector discover_config_assets(const runtime::ModelLoadRequest & request) { @@ -33,6 +36,8 @@ class KokoroTTSLoader final : public runtime::IVoiceModelLoader { bool can_load(const runtime::ModelLoadRequest & request) const override { try { const auto root = resolve_model_root(request.model_path); + if (root.extension() == ".gguf" && engine::io::is_existing_file(root)) + return (!request.family_hint.has_value() || *request.family_hint == family()) && is_kokoro_gguf(root); return engine::io::is_existing_file(root / "config.json") && engine::io::is_existing_file(root / "voices.json") && engine::io::is_existing_file(root / "kokoro-v1_0.safetensors") @@ -54,7 +59,7 @@ class KokoroTTSLoader final : public runtime::IVoiceModelLoader { inspection.capabilities.supported_tasks = { {runtime::VoiceTaskKind::Tts, {runtime::RunMode::Offline}}, }; - inspection.capabilities.languages = {"a", "b"}; + inspection.capabilities.languages = {"a", "b", "e", "f", "h", "i", "j", "p", "z"}; inspection.capabilities.supports_style_condition = true; inspection.discovered_configs = discover_config_assets(request); inspection.discovered_weights = discover_weight_assets(request); @@ -108,7 +113,7 @@ std::unique_ptr load_kokoro_tts_model(const std::filesyste capabilities.supported_tasks = { {runtime::VoiceTaskKind::Tts, {runtime::RunMode::Offline}}, }; - capabilities.languages = {"a", "b"}; + capabilities.languages = {"a", "b", "e", "f", "h", "i", "j", "p", "z"}; capabilities.supports_style_condition = true; return std::make_unique(std::move(metadata), std::move(capabilities), std::move(assets)); diff --git a/src/models/kokoro_tts/package.cpp b/src/models/kokoro_tts/package.cpp new file mode 100644 index 000000000..4c61c3fea --- /dev/null +++ b/src/models/kokoro_tts/package.cpp @@ -0,0 +1,178 @@ +#include "engine/models/kokoro_tts/package.h" +#include "engine/framework/io/binary.h" +#include +#include +#include +#include +#include +#include +#include +#include + +namespace engine::models::kokoro_tts { +namespace { +using namespace engine::assets; +class GgufSource final : public TensorSource { + std::filesystem::path path_; + std::unique_ptr gguf_{nullptr, gguf_free}; + std::unique_ptr tensors_{nullptr, ggml_free}; + std::map names_; +public: + explicit GgufSource(const std::filesystem::path & path) : path_(path) { + ggml_context * tensors = nullptr; + gguf_.reset(gguf_init_from_file(path.string().c_str(), {true, &tensors})); + tensors_.reset(tensors); + if (!gguf_ || !tensors_) throw std::runtime_error("Invalid Kokoro GGUF: " + path.string()); + const auto architecture = gguf_find_key(gguf_.get(), "general.architecture"); + if (architecture < 0 || gguf_get_kv_type(gguf_.get(), architecture) != GGUF_TYPE_STRING || + std::string(gguf_get_val_str(gguf_.get(), architecture)) != "kokoro_tts") + throw std::runtime_error("Expected a kokoro_tts GGUF architecture"); + const auto names = gguf_find_key(gguf_.get(), "kokoro.tensor_names"); + if (names < 0 || gguf_get_kv_type(gguf_.get(), names) != GGUF_TYPE_ARRAY || + gguf_get_arr_type(gguf_.get(), names) != GGUF_TYPE_STRING || + gguf_get_arr_n(gguf_.get(), names) != static_cast(gguf_get_n_tensors(gguf_.get()))) + throw std::runtime_error("Missing Kokoro tensor name mapping"); + for (int64_t i = 0; i < gguf_get_n_tensors(gguf_.get()); ++i) + if (!names_.emplace(gguf_get_arr_str(gguf_.get(), names, i), gguf_get_tensor_name(gguf_.get(), i)).second) + throw std::runtime_error("Duplicate Kokoro tensor mapping"); + for (auto * t = ggml_get_first_tensor(tensors_.get()); t; t = ggml_get_next_tensor(tensors_.get(), t)) { + if (t->type != GGML_TYPE_F32 && t->type != GGML_TYPE_F16 && + t->type != GGML_TYPE_BF16 && t->type != GGML_TYPE_Q8_0) + throw std::runtime_error("Unsupported Kokoro GGUF tensor type: " + std::string(t->name)); + } + } + const gguf_context * metadata() const { return gguf_.get(); } + const std::filesystem::path & source_path() const noexcept override { return path_; } + bool has_tensor(std::string_view name) const noexcept override { + return names_.find(std::string(name)) != names_.end(); + } + TensorMetadata require_metadata(std::string_view name) const override { + auto found = names_.find(std::string(name)); + auto * t = found == names_.end() ? nullptr : ggml_get_tensor(tensors_.get(), found->second.c_str()); + if (!t) throw std::runtime_error("Missing Kokoro tensor: " + std::string(name)); + TensorMetadata m{std::string(name), ggml_type_name(t->type), {}}; + // GGUF dimensions use the reverse of the source checkpoint order. + for (int i = ggml_n_dims(t) - 1; i >= 0; --i) m.shape.push_back(t->ne[i]); + // GGUF elides trailing singleton dimensions; retain the source shape metadata. + auto key = "kokoro.tensor_shape." + std::string(name); + auto id = gguf_find_key(gguf_.get(), key.c_str()); + if (id >= 0) { + if (gguf_get_kv_type(gguf_.get(), id) != GGUF_TYPE_ARRAY || + gguf_get_arr_type(gguf_.get(), id) != GGUF_TYPE_INT64) + throw std::runtime_error("Invalid Kokoro tensor shape metadata"); + const auto n = gguf_get_arr_n(gguf_.get(), id); + const auto * dims = static_cast(gguf_get_arr_data(gguf_.get(), id)); + if (!n || n > 4) throw std::runtime_error("Invalid Kokoro tensor rank"); + m.shape.assign(dims, dims + n); + for (size_t i = 0; i < n; ++i) + if (dims[i] <= 0 || dims[i] != t->ne[n - i - 1]) + throw std::runtime_error("Inconsistent Kokoro tensor shape"); + } + return m; + } + std::vector tensors() const override { + std::vector out; + for (const auto & item : names_) out.push_back(require_metadata(item.first)); + return out; + } + RawTensorData require_tensor_data(std::string_view name) const override { + RawTensorData out{require_metadata(name), {}}; + const auto id = gguf_find_tensor(gguf_.get(), names_.at(std::string(name)).c_str()); + out.bytes.resize(gguf_get_tensor_size(gguf_.get(), id)); + std::ifstream in(path_, std::ios::binary); + in.seekg(gguf_get_data_offset(gguf_.get()) + gguf_get_tensor_offset(gguf_.get(), id)); + if (!in.read(reinterpret_cast(out.bytes.data()), out.bytes.size())) + throw std::runtime_error("Truncated Kokoro GGUF tensor: " + std::string(name)); + return out; + } + std::vector require_f32(std::string_view name, + const std::optional> & expected = std::nullopt) const override { + auto raw = require_tensor_data(name); + if (expected && *expected != raw.metadata.shape) throw std::runtime_error("Kokoro tensor shape mismatch"); + core::TensorShape shape; + shape.rank = static_cast(raw.metadata.shape.size()); + for (size_t i = 0; i < raw.metadata.shape.size(); ++i) shape.dims[i] = raw.metadata.shape[i]; + return tensor_data_to_f32(name, {shape, + ggml_type_for_tensor_storage(tensor_storage_type_for_dtype(raw.metadata.dtype)), std::move(raw.bytes)}); + } + std::optional> optional_f32(std::string_view name, + const std::optional> & expected = std::nullopt) const override { + if (!has_tensor(name)) return std::nullopt; + return require_f32(name, expected); + } + int64_t require_i64_scalar(std::string_view) const override { + throw std::runtime_error("Kokoro GGUF does not contain integer scalars"); + } +}; + +void extract_resources(const gguf_context * g, const std::filesystem::path & root) { + auto names = gguf_find_key(g, "audiocpp.embedded_files.names"); + auto offsets = gguf_find_key(g, "audiocpp.embedded_files.offsets"); + auto data = gguf_find_key(g, "audiocpp.embedded_files.data"); + if (names < 0 || offsets < 0 || data < 0 || + gguf_get_kv_type(g, names) != GGUF_TYPE_ARRAY || gguf_get_arr_type(g, names) != GGUF_TYPE_STRING || + gguf_get_kv_type(g, offsets) != GGUF_TYPE_ARRAY || gguf_get_arr_type(g, offsets) != GGUF_TYPE_UINT64 || + gguf_get_kv_type(g, data) != GGUF_TYPE_ARRAY || gguf_get_arr_type(g, data) != GGUF_TYPE_UINT8 || + gguf_get_arr_n(g, offsets) != gguf_get_arr_n(g, names) + 1) + throw std::runtime_error("Kokoro GGUF is missing valid embedded resources"); + const auto * begin = static_cast(gguf_get_arr_data(g, data)); + const auto * off = static_cast(gguf_get_arr_data(g, offsets)); + const auto size = gguf_get_arr_n(g, data); + std::set seen; + for (size_t i = 0; i < gguf_get_arr_n(g, names); ++i) { + std::string name = gguf_get_arr_str(g, names, i); + std::replace(name.begin(), name.end(), '\\', '/'); + const auto relative = std::filesystem::u8path(name).lexically_normal(); + if (name.empty() || name.find(':') != std::string::npos || relative.has_root_path() || + relative.empty() || *relative.begin() == ".." || !seen.insert(relative).second || + off[i] > off[i + 1] || off[i + 1] > size) + throw std::runtime_error("Unsafe Kokoro embedded resource"); + auto path = root / relative; + std::filesystem::create_directories(path.parent_path()); + std::ofstream out(path, std::ios::binary); + if (!out.write(begin + off[i], off[i + 1] - off[i])) + throw std::runtime_error("Cannot extract Kokoro resource: " + name); + } +} +} + +bool is_kokoro_gguf(const std::filesystem::path & path) noexcept { + try { + gguf_init_params params{true, nullptr}; + std::unique_ptr metadata( + gguf_init_from_file(path.string().c_str(), params), gguf_free); + if (!metadata) return false; + const auto architecture = gguf_find_key(metadata.get(), "general.architecture"); + return architecture >= 0 && gguf_get_kv_type(metadata.get(), architecture) == GGUF_TYPE_STRING && + std::string(gguf_get_val_str(metadata.get(), architecture)) == "kokoro_tts"; + } catch (...) { + return false; + } +} + +KokoroPackage::~KokoroPackage() { + if (temporary) { + std::error_code ec; + std::filesystem::remove_all(root, ec); + } +} +std::shared_ptr open_kokoro_package(const std::filesystem::path & path) { + auto package = std::make_shared(); + if (std::filesystem::is_directory(path)) { + package->root = std::filesystem::canonical(path); + package->weights = assets::open_tensor_source(package->root / "kokoro-v1_0.safetensors"); + return package; + } + auto source = std::make_shared(std::filesystem::canonical(path)); + std::random_device rng; + for (int tries = 0; tries < 20; ++tries) { + package->root = std::filesystem::temp_directory_path() / + ("audiocpp-kokoro-" + std::to_string(rng()) + "-" + std::to_string(rng())); + if (std::filesystem::create_directory(package->root)) { package->temporary = true; break; } + } + if (!package->temporary) throw std::runtime_error("Cannot create Kokoro resource directory"); + extract_resources(source->metadata(), package->root); + package->weights = source; + return package; +} +} diff --git a/tests/kokoro_tts/MULTILINGUAL_GGUF.md b/tests/kokoro_tts/MULTILINGUAL_GGUF.md new file mode 100644 index 000000000..8fbaeefb4 --- /dev/null +++ b/tests/kokoro_tts/MULTILINGUAL_GGUF.md @@ -0,0 +1,84 @@ +# Kokoro multilingual GGUF preparation and validation + +The Kokoro loader accepts the original extracted model directory or a standalone +GGUF produced by `tools/prepare_kokoro_gguf.py`. Specify `--family kokoro_tts` for +GGUF loading. This is a Kokoro-specific container; arbitrary third-party Kokoro +GGUF layouts are not supported. + +Each package includes all 54 source voices, configuration, English pronunciation +resources, eSpeak data, the full UniDic dictionary, and Chinese pronunciation and +segmentation dictionaries. Executable libraries are not embedded in the model. + +| Language | CLI language | Example voice | +| --- | --- | --- | +| American English | en-us | af_heart | +| British English | en-gb | bf_emma | +| Spanish | es | ef_dora | +| French | fr-fr | ff_siwis | +| Hindi | hi | hf_alpha | +| Italian | it | if_sara | +| Japanese | ja | jf_alpha | +| Brazilian Portuguese | pt-br | pf_dora | +| Mandarin Chinese | zh | zf_xiaobei | + +## Runtime dependencies + +English retains the existing native frontend. Spanish, French, Hindi, Italian, +and Portuguese use the native eSpeak library. Japanese uses native MeCab with +embedded UniDic data. Chinese uses native dictionary/DAG/HMM processing. +Python is used only for conversion and upstream comparison, never inference. + +On Windows, place `espeak-ng.dll` and `libmecab.dll` beside the executable, or set +`AUDIOCPP_ESPEAK_LIBRARY` / `AUDIOCPP_MECAB_LIBRARY` to absolute library paths. +On other platforms these variables can select installed native libraries. +Embedded data is extracted to a temporary directory for the loaded model's +lifetime and removed when its assets are released. + +## Conversion + +Use a preparation environment containing numpy, safetensors, gguf (tested with +0.19), misaki[ja,zh] (tested with 0.9.4), espeakng-loader, unidic, and Jieba. +Run `python -m unidic download` first. Start from the original extracted Kokoro +directory containing weights, config, all voices, and English resources. + +```powershell +python tools/prepare_kokoro_gguf.py --source ../models_v3_test/kokoro-82m-v1_0-ggml --resources ../models_v3_test/Kokoro-multilingual-resources --output-dir ../models_v3_test/Kokoro-GGUF +``` + +Outputs are `kokoro-v1.0-q8_0.gguf` and `kokoro-v1.0-bf16.gguf`. Q8 quantizes +eligible matrices; unsupported weight layouts remain BF16, while sensitive +small tensors remain F32 in both packages. Q8 does not mean every tensor is Q8. +The full embedded dictionaries dominate file size: approximately 943 MB for Q8 +and 965 MB for BF16. Language resources are identical in both. + +## Synthesis + +```powershell +.\build\windows-cpu-release\bin\audiocpp_cli.exe --task tts --family kokoro_tts --model ..\models_v3_test\Kokoro-GGUF\kokoro-v1.0-q8_0.gguf --backend cpu --threads 8 --language en-us --voice-id af_heart --text "Hello, this is a native Kokoro TTS test." --out kokoro-q8.wav +``` + +For non-Latin text on Windows, use a UTF-8 text file with +`--batch-text-file input.txt --batch-merge-audio concat` instead of `--text`. + +CPU thread count should be tuned to the machine. On a Ryzen 9 7950X3D, +`--threads 16` reduced Q8 request times by 18–23% versus eight threads in two +opposite-order runs of the seven-request English benchmark. All corresponding +WAV hashes were identical. This is a runtime configuration improvement; it +does not change the model file or establish an optimal setting for other CPUs. + +## Validation + +`compare_multilingual_g2p.py` compares `kokoro_g2p_probe` against installed Misaki. +`validate_multilingual_packages.py` synthesizes each of the nine language variants +with both precisions and saves WAV files, command logs, and `validation.json`. + +Initial Windows CPU results: 12/12 pronunciation cases matched upstream exactly; +18/18 synthesis cases produced non-silent 24 kHz audio. These are smoke tests, +not perceptual-equivalence measurements or GPU coverage. Reported process times +include model loading and dictionary extraction, not just generation. + +The new Japanese/Chinese frontend is not a claim of complete upstream text +normalization parity: unusual numbers, mixed scripts, and Unicode normalization +edge cases need broader coverage. Review all bundled resource and library +licenses before redistributing a package; model data does not replace the +separate native-library redistribution requirements. diff --git a/tests/kokoro_tts/compare_gguf_quality.py b/tests/kokoro_tts/compare_gguf_quality.py new file mode 100644 index 000000000..9d065241e --- /dev/null +++ b/tests/kokoro_tts/compare_gguf_quality.py @@ -0,0 +1,148 @@ +"""Compare Kokoro GGUF output with matching extracted-F32 reference WAVs.""" +import argparse +import json +import math +import wave +from pathlib import Path + +import numpy as np +from pesq import pesq +from pystoi import stoi +from scipy.signal import resample, resample_poly, stft +from scipy.spatial.distance import cdist + + +def load(path: Path): + with wave.open(str(path), "rb") as wav: + if wav.getnchannels() != 1 or wav.getsampwidth() != 2: + raise ValueError(f"expected mono PCM16: {path}") + rate = wav.getframerate() + audio = np.frombuffer(wav.readframes(wav.getnframes()), dtype=" left: + filters[band, left:center] = np.arange(center - left) / (center - left) + if right > center: + filters[band, center:right] = np.arange(right - center, 0, -1) / (right - center) + return np.log(np.maximum(filters @ power, 1e-10)).T + + +def dtw_log_mel_cosine(reference, estimate, rate): + ref = log_mel(reference, rate) + est = log_mel(estimate, rate) + # Per-frame cosine distance emphasizes spectral shape over loudness. + ref -= ref.mean(axis=1, keepdims=True) + est -= est.mean(axis=1, keepdims=True) + costs = cdist(ref, est, metric="cosine").astype(np.float32) + rows, cols = costs.shape + accumulated = np.empty_like(costs) + steps = np.empty((rows, cols), dtype=np.uint8) + accumulated[0, 0] = costs[0, 0] + accumulated[1:, 0] = np.cumsum(costs[1:, 0]) + accumulated[0, 0] + accumulated[0, 1:] = np.cumsum(costs[0, 1:]) + accumulated[0, 0] + steps[1:, 0] = 1 + steps[0, 1:] = 2 + for row in range(1, rows): + previous = accumulated[row - 1] + for col in range(1, cols): + options = (previous[col - 1], previous[col], accumulated[row, col - 1]) + step = int(np.argmin(options)) + accumulated[row, col] = costs[row, col] + options[step] + steps[row, col] = step + row, col = rows - 1, cols - 1 + path_cost = 0.0 + length = 0 + while row or col: + path_cost += float(costs[row, col]) + length += 1 + step = steps[row, col] + if step == 0: + row -= 1; col -= 1 + elif step == 1: + row -= 1 + else: + col -= 1 + path_cost += float(costs[0, 0]) + return 1.0 - path_cost / (length + 1) + + +def compare(reference_path, candidate_path): + rate, reference = load(reference_path) + candidate_rate, candidate = load(candidate_path) + if candidate_rate != rate: + raise ValueError("sample-rate mismatch") + original_candidate_size = len(candidate) + # Duration prediction changes by a few frames under quantization. Normalize the + # global duration before signal metrics so they measure acoustic similarity. + candidate = resample(candidate, len(reference)) + reference_16k = resample_poly(reference, 2, 3) + candidate_16k = resample_poly(candidate, 2, 3) + eps = 1e-20 + return { + "reference_seconds": len(reference) / rate, + "candidate_seconds": original_candidate_size / rate, + "duration_delta_percent": 100 * (original_candidate_size - len(reference)) / len(reference), + "rms_delta_db": 20 * math.log10(max(np.sqrt(np.mean(candidate * candidate)), eps) / + max(np.sqrt(np.mean(reference * reference)), eps)), + "correlation": float(np.corrcoef(reference, candidate)[0, 1]), + "si_sdr_db": float(si_sdr(reference, candidate)), + "log_spectral_distance_db": log_spectral_distance(reference, candidate, rate), + "dtw_log_mel_cosine": dtw_log_mel_cosine(reference, candidate, rate), + "stoi": float(stoi(reference_16k, candidate_16k, 16000, extended=False)), + "pesq_wb": float(pesq(16000, reference_16k, candidate_16k, "wb")), + } + + +parser = argparse.ArgumentParser() +parser.add_argument("--root", type=Path, required=True) +args = parser.parse_args() +records = [] +for precision in ("bf16-gguf", "q8-gguf"): + for index in range(7): + record = {"precision": precision, "request": index} + record.update(compare(args.root / "f32-directory" / f"request_{index}.wav", + args.root / precision / f"request_{index}.wav")) + records.append(record) + +(args.root / "quality.json").write_text(json.dumps(records, indent=2), encoding="utf-8") +for precision in ("bf16-gguf", "q8-gguf"): + subset = [r for r in records if r["precision"] == precision] + print(precision) + for key in ("duration_delta_percent", "rms_delta_db", "correlation", "si_sdr_db", + "log_spectral_distance_db", "dtw_log_mel_cosine", "stoi", "pesq_wb"): + values = np.array([r[key] for r in subset]) + print(f" {key}: mean={values.mean():.6f}, min={values.min():.6f}, max={values.max():.6f}") diff --git a/tests/kokoro_tts/compare_multilingual_g2p.py b/tests/kokoro_tts/compare_multilingual_g2p.py new file mode 100644 index 000000000..629004d1f --- /dev/null +++ b/tests/kokoro_tts/compare_multilingual_g2p.py @@ -0,0 +1,34 @@ +"""Compare native pronunciation with the upstream Misaki frontends.""" +import argparse +import json +import os +from pathlib import Path +import subprocess +import espeakng_loader +from phonemizer.backend.espeak.wrapper import EspeakWrapper +EspeakWrapper.set_library(espeakng_loader.get_library_path()) +EspeakWrapper.set_data_path(espeakng_loader.get_data_path()) +from misaki import espeak, ja, zh + +p = argparse.ArgumentParser() +p.add_argument('--probe', type=Path, required=True) +p.add_argument('--resources', type=Path, required=True) +p.add_argument('--output', type=Path, required=True) +a = p.parse_args() +cases_path = Path(__file__).with_name('multilingual_cases.json') +cases = json.loads(cases_path.read_text(encoding='utf-8')) +frontends = {k: espeak.EspeakG2P(language=v) for k,v in + {'e':'es', 'f':'fr-fr', 'h':'hi', 'i':'it', 'p':'pt-br'}.items()} +frontends['j'], frontends['z'] = ja.JAG2P(), zh.ZHG2P() +native = subprocess.run([str(a.probe.resolve()), str(a.resources.resolve()), str(cases_path.resolve())], + capture_output=True, encoding='utf-8', check=True) +lines = native.stdout.splitlines() +if len(lines) != len(cases): raise RuntimeError(native.stdout + native.stderr) +results = [] +for case, actual in zip(cases, lines): + expected = frontends[case['language']](case['text'])[0] + results.append({**case, 'reference':expected, 'native':actual, 'match':expected == actual}) +a.output.parent.mkdir(parents=True, exist_ok=True) +a.output.write_text(json.dumps(results, ensure_ascii=False, indent=2), encoding='utf-8') +for i, r in enumerate(results): print(i, r['language'], r['match'], flush=True) +print('Report:', a.output) diff --git a/tests/kokoro_tts/kokoro_g2p_probe.cpp b/tests/kokoro_tts/kokoro_g2p_probe.cpp new file mode 100644 index 000000000..d7f5279c1 --- /dev/null +++ b/tests/kokoro_tts/kokoro_g2p_probe.cpp @@ -0,0 +1,17 @@ +#include "engine/models/kokoro_tts/g2p_multilingual.h" +#include "engine/framework/io/json.h" +#include +#include +int main(int argc, char ** argv) { + try { + if (argc != 3) throw std::runtime_error("usage: kokoro_g2p_probe "); + engine::models::kokoro_tts::MultilingualG2P g2p(argv[1]); + auto cases = engine::io::json::parse_file(argv[2]); + for (const auto & item : cases.as_array()) { + try { + std::cout << g2p.phonemize(item.find("text")->as_string(), item.find("language")->as_string()) << '\n'; + } catch (const std::exception & e) { std::cout << "ERROR: " << e.what() << '\n'; } + } + return 0; + } catch (const std::exception & e) { std::cerr << e.what() << '\n'; return 1; } +} diff --git a/tests/kokoro_tts/multilingual_cases.json b/tests/kokoro_tts/multilingual_cases.json new file mode 100644 index 000000000..061f3fff3 --- /dev/null +++ b/tests/kokoro_tts/multilingual_cases.json @@ -0,0 +1,14 @@ +[ + {"language":"e", "voice":"ef_dora", "text":"Hola, esta es una prueba de síntesis de voz. Buenos días a todos."}, + {"language":"f", "voice":"ff_siwis", "text":"Bonjour, ceci est un test de synthèse vocale. Comment allez-vous aujourd'hui ?"}, + {"language":"h", "voice":"hf_alpha", "text":"नमस्ते, यह आवाज़ का एक परीक्षण है। आपका दिन शुभ हो।"}, + {"language":"i", "voice":"if_sara", "text":"Ciao, questa è una prova di sintesi vocale. Buongiorno a tutti."}, + {"language":"p", "voice":"pf_dora", "text":"Olá, este é um teste de síntese de voz. Bom dia a todos."}, + {"language":"j", "voice":"jf_alpha", "text":"こんにちは。これは音声合成のテストです。今日はいい天気ですね。"}, + {"language":"z", "voice":"zf_xiaobei", "text":"你好,这是一个语音合成测试。今天天气很好,祝你有愉快的一天。"}, + {"language":"j", "voice":"jf_alpha", "text":"東京から大阪まで電車で行きます。明日は学校で日本語を勉強します。"}, + {"language":"z", "voice":"zf_xiaobei", "text":"银行今天开门了吗?重庆是一个美丽的城市。我们正在学习中文。"}, + {"language":"e", "voice":"ef_dora", "text":"¿Cómo estás? Tengo 25 años y me gusta la música."}, + {"language":"j", "voice":"jf_alpha", "text":"私は25歳です。学校には100人の学生がいます。"}, + {"language":"z", "voice":"zf_xiaobei", "text":"我今年25岁,买了12个苹果。"} +] diff --git a/tests/kokoro_tts/validate_multilingual_packages.py b/tests/kokoro_tts/validate_multilingual_packages.py new file mode 100644 index 000000000..40b13b034 --- /dev/null +++ b/tests/kokoro_tts/validate_multilingual_packages.py @@ -0,0 +1,55 @@ +"""Synthesize one UTF-8 sample for every Kokoro language from both standalone GGUFs.""" +import argparse +import json +from pathlib import Path +import subprocess +import time +import wave +import numpy as np + +p = argparse.ArgumentParser() +p.add_argument('--cli', type=Path, required=True) +p.add_argument('--models', type=Path, required=True) +p.add_argument('--output', type=Path, required=True) +p.add_argument('--types', nargs='+', default=['q8_0', 'bf16']) +a = p.parse_args() +cases = json.loads(Path(__file__).with_name('multilingual_cases.json').read_text(encoding='utf-8'))[:7] +cases = [ + {'language':'en-us','voice':'af_heart','text':'Hello, this is a native Kokoro TTS test. Good morning to everyone.'}, + {'language':'en-gb','voice':'bf_emma','text':'Hello, this is a native Kokoro TTS test. Good morning to everyone.'} +] + cases +codes = {'e':'es', 'f':'fr-fr', 'h':'hi', 'i':'it', 'p':'pt-br', 'j':'ja', 'z':'zh'} +a.output.mkdir(parents=True, exist_ok=True) +report = [] +for precision in a.types: + model = a.models.resolve() / ('kokoro-v1.0-' + precision + '.gguf') + for case in cases: + lang = codes.get(case['language'], case['language']) + dest = a.output.resolve() / (precision + '-' + lang) + dest.mkdir(exist_ok=True) + textfile = dest / 'text.txt' + textfile.write_text(case['text'] + '\n', encoding='utf-8') + wav = a.output.resolve() / (precision + '-' + lang + '.wav') + command = [str(a.cli.resolve()), '--task', 'tts', '--family', 'kokoro_tts', + '--model', str(model), '--backend', 'cpu', '--threads', '8', '--seed', '1234', + '--voice-id', case['voice'], '--language', lang, '--batch-text-file', str(textfile), + '--batch-merge-audio', 'concat', '--out', str(wav)] + start = time.perf_counter() + run = subprocess.run(command, capture_output=True, encoding='utf-8', errors='replace', timeout=180) + elapsed = time.perf_counter() - start + (dest / 'console.log').write_text(run.stdout + run.stderr, encoding='utf-8') + result = {'precision':precision, **case, 'exit_code':run.returncode, 'process_seconds':elapsed} + if run.returncode == 0: + with wave.open(str(wav), 'rb') as f: + assert f.getsampwidth() == 2 + pcm = np.frombuffer(f.readframes(f.getnframes()), dtype=' 1 and result['rms'] > 0.002 + else: + result['error'] = (run.stdout + run.stderr)[-2000:] + report.append(result) + (a.output / 'validation.json').write_text(json.dumps(report, ensure_ascii=False, indent=2), encoding='utf-8') + print(precision, lang, 'PASS' if result.get('valid_audio') else result.get('error', 'INVALID AUDIO'), flush=True) +if any(not r.get('valid_audio') for r in report): raise SystemExit(1) diff --git a/tools/prepare_kokoro_gguf.py b/tools/prepare_kokoro_gguf.py new file mode 100644 index 000000000..cb341fd2d --- /dev/null +++ b/tools/prepare_kokoro_gguf.py @@ -0,0 +1,141 @@ +"""Create standalone Kokoro GGUFs, preserving every voice and pronunciation resource. + +Conversion dependencies: numpy, safetensors, gguf, misaki[ja,zh], espeakng-loader, +unidic (run `python -m unidic download` once). Python is not used for inference. +""" +import argparse +import importlib.metadata +import json +import struct +import shutil +from pathlib import Path +import numpy as np +import gguf +from safetensors import safe_open + +class ResourceWriter(gguf.GGUFWriter): + def _pack_val(self, val, vtype, add_vtype, sub_type=None): + # Avoid a Python function call per byte for large embedded dictionaries. + if vtype == gguf.GGUFValueType.ARRAY and isinstance(val, bytes): + return (struct.pack(' 0}, 'total': jieba.dt.total, + 'start': start_P, 'transition': trans_P, 'emission': emit_P}) + # Keep redistribution notices with the resources. + notices = root / 'licenses' + notices.mkdir(exist_ok=True) + shutil.copytree(Path(unidic.DICDIR) / 'licenses', notices / 'unidic-dictionary', dirs_exist_ok=True) + shutil.copy2(Path(unidic.DICDIR) / 'README', notices / 'unidic-dictionary-README') + for dist_name in ['misaki', 'unidic', 'espeakng-loader', 'jieba', 'pypinyin', 'mecab-python3']: + dist = importlib.metadata.distribution(dist_name) + for file in dist.files or []: + if any(x in file.name.lower() for x in ['license', 'copying', 'notice']): + src = Path(dist.locate_file(file)) + if src.is_file(): shutil.copy2(src, notices / (dist_name + '-' + file.name)) + dump(root / 'g2p' / 'versions.json', {n: importlib.metadata.version(n) + for n in ['misaki', 'unidic', 'espeakng-loader', 'jieba', 'pypinyin']}) + + +def convert(root, output, precision, overwrite=False): + if output.exists() and not overwrite: raise FileExistsError(output) + output.parent.mkdir(parents=True, exist_ok=True) + temporary = output.with_suffix('.gguf.partial') + writer = ResourceWriter(str(temporary), 'kokoro_tts') + writer.add_name('Kokoro v1.0 multilingual ' + precision) + writer.add_array('kokoro.languages', ['en-us', 'en-gb', 'es', 'fr-fr', 'hi', 'it', 'ja', 'pt-br', 'zh']) + writer.add_array('kokoro.voices', sorted(json.loads((root / 'voices.json').read_text()))) + names, offsets, data = [], [0], bytearray() + for file in sorted(root.rglob('*')): + if not file.is_file() or file.suffix in ['.safetensors', '.gguf']: continue + names.append(file.relative_to(root).as_posix()) + data.extend(file.read_bytes()) + offsets.append(len(data)) + writer.add_array('audiocpp.embedded_files.names', names) + writer.add_key_value('audiocpp.embedded_files.offsets', offsets, gguf.GGUFValueType.ARRAY, + sub_type=gguf.GGUFValueType.UINT64) + writer.add_key_value('audiocpp.embedded_files.data', bytes(data), gguf.GGUFValueType.ARRAY, + sub_type=gguf.GGUFValueType.UINT8) + counts = {} + with safe_open(root / 'kokoro-v1_0.safetensors', framework='numpy') as model: + writer.add_array('kokoro.tensor_names', list(model.keys())) + for index, name in enumerate(model.keys()): + array = model.get_tensor(name).astype(np.float32) + writer.add_key_value('kokoro.tensor_shape.' + name, list(array.shape), + gguf.GGUFValueType.ARRAY, sub_type=gguf.GGUFValueType.INT64) + # Retain scalar/vector/norm and Snake parameters in F32. Weight-norm + # convolution tensors use BF16 on disk and are reconstructed in F32. + kind = gguf.GGMLQuantizationType.F32 + if array.ndim >= 2 and all(d > 1 for d in array.shape): + kind = gguf.GGMLQuantizationType.BF16 + if precision == 'q8_0' and array.ndim == 2 and array.shape[-1] % 32 == 0: + kind = gguf.GGMLQuantizationType.Q8_0 + encoded = gguf.quants.quantize(array, kind) + writer.add_tensor('kokoro.' + str(index), encoded, raw_dtype=kind) + counts[kind.name] = counts.get(kind.name, 0) + 1 + output.parent.mkdir(parents=True, exist_ok=True) + writer.write_header_to_file() + writer.write_kv_data_to_file() + writer.write_tensors_to_file() + writer.close() + temporary.replace(output) + print(json.dumps({'path': str(output), 'bytes': output.stat().st_size, + 'tensors': counts, 'resources': len(names), 'resource_bytes': len(data)}), flush=True) + + +if __name__ == '__main__': + p = argparse.ArgumentParser(description=__doc__) + p.add_argument('--source', type=Path, required=True) + p.add_argument('--resources', type=Path, required=True) + p.add_argument('--output-dir', type=Path, required=True) + p.add_argument('--skip-prepare', action='store_true') + p.add_argument('--overwrite', action='store_true') + p.add_argument('--type', choices=['q8_0', 'bf16', 'both'], default='both') + args = p.parse_args() + if not args.skip_prepare: prepare(args.source, args.resources) + for precision in (['q8_0', 'bf16'] if args.type == 'both' else [args.type]): + convert(args.resources, args.output_dir / ('kokoro-v1.0-' + precision + '.gguf'), precision, args.overwrite) From f3dae551cc0a7c365602891e1b6aa5014212ae72 Mon Sep 17 00:00:00 2001 From: mirek190 Date: Fri, 11 Sep 2026 14:42:54 +0100 Subject: [PATCH 4/5] fix(kokoro): add model spec for registry sync --- model_specs/kokoro_tts.json | 87 +++++++++++++++++++++++++++++++++++++ 1 file changed, 87 insertions(+) create mode 100644 model_specs/kokoro_tts.json diff --git a/model_specs/kokoro_tts.json b/model_specs/kokoro_tts.json new file mode 100644 index 000000000..305f4885c --- /dev/null +++ b/model_specs/kokoro_tts.json @@ -0,0 +1,87 @@ +{ + "schema_version": 1, + "family": "kokoro_tts", + "display_name": "Kokoro 82M", + "description": "Kokoro 82M multilingual text-to-speech with 54 preset voices, optimized native CPU inference, and shared eSpeak-ng phonemization.", + "category": "tts", + "status": "community", + "tasks": ["tts"], + "modes": ["offline"], + "languages": ["en-us", "en-gb", "es", "fr-fr", "hi", "it", "ja", "pt-br", "zh"], + "runtime": { + "tags": ["gguf"] + }, + "capabilities": { + "tts": ["built_in_voices", "long_form"] + }, + "options": { + "request": [ + { + "name": "seed", + "type": "int", + "description": "Decoder noise seed; omitted requests choose a random seed.", + "required": false, + "min": 0 + } + ], + "session": [], + "load": [ + { + "name": "matmul_weight_type", + "type": "enum", + "preset": "weight_type_full", + "required": false, + "default": "native", + "description": "Storage type for matrix multiplication weights." + }, + { + "name": "conv_weight_type", + "type": "enum", + "preset": "weight_type_conv", + "required": false, + "default": "f32", + "description": "Storage type for convolution weights." + } + ] + }, + "packages": [], + "dependencies": [], + "ui": { + "tags": ["TTS", "Preset voices"], + "docs": ["tests/kokoro_tts/MULTILINGUAL_GGUF.md"] + }, + "sources": [ + { + "format": "gguf", + "roots": { + "model": ".", + "weights": "$gguf" + }, + "files": { + "config": "model:config.json", + "voices": "model:voices.json", + "vocabulary": "model:vocab.tsv" + }, + "tensors": { + "weights": { + "source": "weights:", + "prefix": "kokoro" + } + } + }, + { + "format": "safetensors", + "roots": { + "model": "." + }, + "files": { + "config": "model:config.json", + "voices": "model:voices.json", + "vocabulary": "model:vocab.tsv" + }, + "tensors": { + "weights": "model:kokoro-v1_0.safetensors" + } + } + ] +} From 4cab2ebb9d35e33783ead67380a3add9052f8328 Mon Sep 17 00:00:00 2001 From: mirek190 Date: Fri, 11 Sep 2026 14:55:04 +0100 Subject: [PATCH 5/5] test(kokoro): allow cross-platform floating-point rounding --- tests/kokoro_tts/kokoro_cpu_kernel_test.cpp | 37 ++++++++++++++++----- 1 file changed, 29 insertions(+), 8 deletions(-) diff --git a/tests/kokoro_tts/kokoro_cpu_kernel_test.cpp b/tests/kokoro_tts/kokoro_cpu_kernel_test.cpp index 0cc6a603b..908bfa731 100644 --- a/tests/kokoro_tts/kokoro_cpu_kernel_test.cpp +++ b/tests/kokoro_tts/kokoro_cpu_kernel_test.cpp @@ -3,8 +3,10 @@ #include #include +#include #include #include +#include struct Conv { int64_t in_channels; @@ -14,15 +16,32 @@ struct Conv { int64_t dilation; }; -void compare(ggml_context * ctx, ggml_tensor * expected, ggml_tensor * actual, int threads) { +void compare(ggml_context * ctx, ggml_tensor * expected, ggml_tensor * actual, int threads, + const char * kernel, bool bit_exact) { ggml_cgraph * graph = ggml_new_graph(ctx); ggml_build_forward_expand(graph, expected); ggml_build_forward_expand(graph, actual); for (int repeat = 0; repeat < 2; ++repeat) { if (ggml_graph_compute_with_ctx(ctx, graph, threads) != GGML_STATUS_SUCCESS || - ggml_nbytes(expected) != ggml_nbytes(actual) || - std::memcmp(expected->data, actual->data, ggml_nbytes(expected)) != 0) { - throw std::runtime_error("CPU kernel is not bit-exact with ggml reference"); + ggml_nbytes(expected) != ggml_nbytes(actual)) { + throw std::runtime_error(std::string(kernel) + " CPU kernel execution failed"); + } + if (bit_exact) { + if (std::memcmp(expected->data, actual->data, ggml_nbytes(expected)) != 0) { + throw std::runtime_error(std::string(kernel) + " CPU kernel is not bit-exact with ggml reference"); + } + continue; + } + const auto * want = static_cast(expected->data); + const auto * got = static_cast(actual->data); + for (int64_t i = 0; i < ggml_nelements(expected); ++i) { + const float abs_error = std::abs(want[i] - got[i]); + const float tolerance = 2.e-5f + 2.e-5f * std::abs(want[i]); + if (!std::isfinite(want[i]) || !std::isfinite(got[i]) || abs_error > tolerance) { + throw std::runtime_error(std::string(kernel) + " CPU kernel differs from ggml reference at element " + + std::to_string(i) + ": expected " + std::to_string(want[i]) + + ", got " + std::to_string(got[i])); + } } } } @@ -51,7 +70,7 @@ int main() { ggml_tensor * actual = ggml_custom_4d(ctx, GGML_TYPE_F32, kernel * channels, numerator / stride + 1, 1, 1, args, 1, kokoro_ggml::cpu_detail::kokoro_im2col_rows, GGML_N_TASKS_MAX, &conv); - compare(ctx, expected, actual, threads); + compare(ctx, expected, actual, threads, "im2col", true); ggml_free(ctx); ++cases; } @@ -68,7 +87,7 @@ int main() { ggml_tensor * expected = ggml_add(ctx, input, ggml_div(ctx, ggml_mul(ctx, sine, sine), alpha)); ggml_tensor * actual = ggml_map_custom2(ctx, input, alpha, kokoro_ggml::cpu_detail::kokoro_snake_cpu, GGML_N_TASKS_MAX, nullptr); - compare(ctx, expected, actual, threads); + compare(ctx, expected, actual, threads, "snake", false); ggml_free(ctx); ++cases; } @@ -97,9 +116,11 @@ int main() { ggml_repeat(ctx, gamma, input)), ggml_repeat(ctx, beta, input)); ggml_tensor * actual = ggml_map_custom3(ctx, input, gamma, beta, kokoro_ggml::cpu_detail::kokoro_adain_cpu, GGML_N_TASKS_MAX, &eps); - compare(ctx, expected, actual, threads); + compare(ctx, expected, actual, threads, "AdaIN", false); ggml_free(ctx); ++cases; } - std::cout << "PASS: " << cases << " bit-exact kernel cases, each executed twice (1/8 threads).\n"; + std::cout << "PASS: " << cases + << " kernel parity cases, each executed twice (1/8 threads); im2col is bit-exact and " + "floating-point arithmetic uses cross-platform tolerance.\n"; }