#include "tts_transformer.h"
#include "transformer/transformer_state_internal.h"
#include "transformer/transformer_internal.h"
#include "transformer/transformer_sampling.h"

#include <algorithm>
#include <cmath>
#include <cstdio>
#include <vector>

namespace qwen3_tts {

bool TTSTransformer::generate(const int32_t * text_tokens, int32_t n_tokens,
                              const float * speaker_embd, int32_t max_len,
                              std::vector<int32_t> & output,
                              int32_t language_id,
                              float repetition_penalty,
                              float temperature,
                              int32_t top_k,
                              float top_p,
                              int64_t seed,
                              const int32_t * instruct_tokens,
                              int32_t n_instruct_tokens,
                              const int32_t * reference_tokens,
                              int32_t n_reference_tokens,
                              const int32_t * reference_codes,
                              int32_t n_reference_frames,
                              int32_t n_reference_codebooks,
                              const tts_code_frame_callback_t * frame_callback) {
#ifdef QWEN3_TTS_TIMING
    using clk = std::chrono::high_resolution_clock;
    tts_timing timing = {};
    auto t_gen_start = clk::now();
    auto t0 = t_gen_start, t1 = t_gen_start;
    impl_->timing = &timing;
#endif

    if (!impl_->model.ctx) {
        error_msg_ = "Model not loaded";
        return false;
    }
    if (!text_tokens) {
        error_msg_ = "text_tokens is null";
        return false;
    }
    if (n_tokens < 4) {
        error_msg_ = "Need at least 4 text tokens for generation";
        return false;
    }
    if (max_len <= 0) {
        output.clear();
        return true;
    }

    const auto & cfg = impl_->model.config;
    const auto & trace_cfg = transformer_internal::get_debug_trace_config();
    if (trace_cfg.enabled) {
        transformer_internal::debug_trace_write_text_line(trace_cfg, "hidden_size=" + std::to_string(cfg.hidden_size));
        transformer_internal::debug_trace_write_text_line(trace_cfg, "codec_vocab_size=" + std::to_string(cfg.codec_vocab_size));
        transformer_internal::debug_trace_write_text_line(trace_cfg, "code_pred_vocab_size=" + std::to_string(cfg.code_pred_vocab_size));
        transformer_internal::debug_trace_write_text_line(trace_cfg, "n_codebooks=" + std::to_string(cfg.n_codebooks));
        transformer_internal::debug_trace_write_text_line(trace_cfg, "n_tokens=" + std::to_string(n_tokens));
        transformer_internal::debug_trace_write_text_line(trace_cfg, "max_len=" + std::to_string(max_len));
    }

    std::vector<float> prefill_embd;
    std::vector<float> trailing_text_hidden;
    std::vector<float> tts_pad_embed;

#ifdef QWEN3_TTS_TIMING
    t0 = clk::now();
#endif
    if (!transformer_internal::ops::build_prefill_graph(*this, text_tokens, n_tokens, speaker_embd, language_id,
                             prefill_embd, trailing_text_hidden, tts_pad_embed,
                             instruct_tokens, n_instruct_tokens,
                             reference_tokens, n_reference_tokens,
                             reference_codes, n_reference_frames,
                             n_reference_codebooks)) {
        return false;
    }
#ifdef QWEN3_TTS_TIMING
    t1 = clk::now();
    timing.t_prefill_build_ms = std::chrono::duration<double, std::milli>(t1 - t0).count();
#endif

    const int32_t prefill_len = (int32_t) (prefill_embd.size() / cfg.hidden_size);
    const int32_t trailing_len = (int32_t) (trailing_text_hidden.size() / cfg.hidden_size);
#ifdef QWEN3_TTS_TIMING
    timing.n_prefill_tokens = prefill_len;
    timing.n_trailing_tokens = trailing_len;
#endif

    if (trace_cfg.enabled) {
        transformer_internal::debug_trace_write_text_line(trace_cfg, "prefill_len=" + std::to_string(prefill_len));
        transformer_internal::debug_trace_write_text_line(trace_cfg, "trailing_len=" + std::to_string(trailing_len));
        transformer_internal::debug_trace_write_bin(trace_cfg, "input_text_tokens.i32.bin", text_tokens,
                                                    (size_t) n_tokens, "i32", {(int64_t) n_tokens});
        if (!prefill_embd.empty()) {
            transformer_internal::debug_trace_write_bin(trace_cfg, "prefill_embd.f32.bin", prefill_embd.data(),
                                                        prefill_embd.size(), "f32",
                                                        {(int64_t) prefill_len, (int64_t) cfg.hidden_size});
        }
        if (speaker_embd) {
            transformer_internal::debug_trace_write_bin(trace_cfg, "speaker_embd.f32.bin", speaker_embd,
                                                        (size_t) cfg.hidden_size, "f32",
                                                        {(int64_t) cfg.hidden_size});
        }
    }

    const int32_t required_ctx = prefill_len + max_len + 8;
    if (impl_->state.cache.n_ctx < required_ctx || impl_->state.cache.n_ctx > std::max<int32_t>(required_ctx * 2, 512)) {
        if (!init_kv_cache(required_ctx)) {
            return false;
        }
    }
    clear_kv_cache();

    if (impl_->state.code_pred_cache.n_ctx < 16) {
        if (!init_code_pred_kv_cache(16)) {
            return false;
        }
    }
    transformer_internal::ops::maybe_reserve_scheduler_graphs(*this, prefill_len, required_ctx);

    std::vector<float> hidden_out;
    std::vector<float> logits;

#ifdef QWEN3_TTS_TIMING
    t0 = clk::now();
#endif
    if (!forward_prefill(prefill_embd.data(), prefill_len, 0, hidden_out, &logits)) {
        return false;
    }
#ifdef QWEN3_TTS_TIMING
    t1 = clk::now();
    timing.t_prefill_forward_ms = std::chrono::duration<double, std::milli>(t1 - t0).count();
#endif

    output.clear();
    output.reserve(max_len * cfg.n_codebooks);

    int32_t n_past = prefill_len;
    std::vector<int32_t> frame_codes(cfg.n_codebooks);
    std::vector<int32_t> generated_cb0_tokens;
    transformer_sampling_state sampling{resolve_sampling_seed(seed), 0};
    const int32_t suppress_start = std::min(cfg.code_pred_vocab_size, cfg.codec_vocab_size);

    for (int frame = 0; frame < max_len; ++frame) {
        const bool trace_frame = transformer_internal::debug_trace_should_dump_frame(trace_cfg, frame);
        if (trace_frame) {
            char raw_name[128];
            snprintf(raw_name, sizeof(raw_name), "frame%03d_cb0_logits_raw.f32.bin", frame);
            transformer_internal::debug_trace_write_bin(trace_cfg, raw_name, logits.data(),
                                                        (size_t) cfg.codec_vocab_size, "f32",
                                                        {(int64_t) cfg.codec_vocab_size});
        }

        for (int32_t i = suppress_start; i < cfg.codec_vocab_size; ++i) {
            if (i != cfg.codec_eos_id) {
                logits[i] = -INFINITY;
            }
        }

        transformer_apply_repetition_penalty(logits.data(), cfg.codec_vocab_size,
                                             generated_cb0_tokens.data(),
                                             (int32_t) generated_cb0_tokens.size(),
                                             repetition_penalty);

        if (trace_frame) {
            char post_rules_name[128];
            snprintf(post_rules_name, sizeof(post_rules_name), "frame%03d_cb0_logits_post_rules.f32.bin", frame);
            transformer_internal::debug_trace_write_bin(trace_cfg, post_rules_name, logits.data(),
                                                        (size_t) cfg.codec_vocab_size, "f32",
                                                        {(int64_t) cfg.codec_vocab_size});
        }

        int32_t next_token;
        if (temperature <= 0.0f) {
            next_token = transformer_argmax(logits.data(), cfg.codec_vocab_size);
        } else {
            next_token = transformer_sample_top_k_p(logits.data(), cfg.codec_vocab_size,
                                                    temperature, top_k, top_p,
                                                    1.0f, nullptr, 0, sampling);
        }

        if (next_token == cfg.codec_eos_id) {
            if (trace_frame) {
                int32_t eos_token = next_token;
                char eos_name[128];
                snprintf(eos_name, sizeof(eos_name), "frame%03d_cb0_token.i32.bin", frame);
                transformer_internal::debug_trace_write_bin(trace_cfg, eos_name, &eos_token, 1, "i32", {1});
            }
            break;
        }

        const bool is_thinking = (next_token >= cfg.codec_think_id && next_token <= cfg.codec_think_eos_id);
        if (is_thinking) {
            fprintf(stderr, "  [frame %d] Filtering thinking token: %d\n", frame, next_token);
        }

        frame_codes[0] = next_token;
        generated_cb0_tokens.push_back(next_token);
        if (trace_frame) {
            char token_name[128];
            snprintf(token_name, sizeof(token_name), "frame%03d_cb0_token.i32.bin", frame);
            transformer_internal::debug_trace_write_bin(trace_cfg, token_name, &frame_codes[0], 1, "i32", {1});

            if (!last_hidden_.empty()) {
                char hidden_name[128];
                snprintf(hidden_name, sizeof(hidden_name), "frame%03d_talker_hidden.f32.bin", frame);
                transformer_internal::debug_trace_write_bin(trace_cfg, hidden_name, last_hidden_.data(),
                                                            last_hidden_.size(), "f32",
                                                            {(int64_t) last_hidden_.size()});
            }
        }

#ifdef QWEN3_TTS_TIMING
        t0 = clk::now();
#endif
        std::vector<int32_t> codes_1_15;
        const bool need_host_hidden_for_predictor =
            impl_->use_coreml_code_predictor || trace_cfg.enabled;
        const float * predictor_hidden =
            (need_host_hidden_for_predictor && !last_hidden_.empty()) ? last_hidden_.data() : nullptr;
        if (!predict_codes_autoregressive(predictor_hidden, frame_codes[0], codes_1_15,
                                          temperature, top_k, top_p,
                                          sampling.seed, &sampling.subseq, frame)) {
            return false;
        }
#ifdef QWEN3_TTS_TIMING
        t1 = clk::now();
        timing.t_code_pred_ms += std::chrono::duration<double, std::milli>(t1 - t0).count();
#endif

        for (int cb = 1; cb < cfg.n_codebooks; ++cb) {
            frame_codes[cb] = codes_1_15[cb - 1];
        }
        if (trace_frame) {
            char frame_codes_name[128];
            snprintf(frame_codes_name, sizeof(frame_codes_name),
                     "frame%03d_codec_tokens_cb0_15.i32.bin", frame);
            transformer_internal::debug_trace_write_bin(trace_cfg, frame_codes_name,
                                                        frame_codes.data(), frame_codes.size(), "i32",
                                                        {(int64_t) frame_codes.size()});
        }

        if (!is_thinking) {
            for (int cb = 0; cb < cfg.n_codebooks; ++cb) {
                output.push_back(frame_codes[cb]);
            }
            if (frame_callback && !(*frame_callback)(frame_codes.data(), cfg.n_codebooks,
                                                     (int32_t) generated_cb0_tokens.size() - 1)) {
                error_msg_ = "Generation aborted by frame callback";
                return false;
            }
        }

#ifdef QWEN3_TTS_TIMING
        timing.n_frames = frame + 1;
#endif

        if (frame + 1 >= max_len) {
            break;
        }

        const float * trailing_row = (frame < trailing_len)
            ? trailing_text_hidden.data() + (size_t) frame * cfg.hidden_size
            : tts_pad_embed.data();

#ifdef QWEN3_TTS_TIMING
        t0 = clk::now();
#endif
        const bool read_hidden_after_step =
            impl_->use_coreml_code_predictor || trace_cfg.enabled;
        if (!forward_step_internal(nullptr, frame_codes.data(), trailing_row,
                                   n_past, logits, nullptr, read_hidden_after_step)) {
            return false;
        }
#ifdef QWEN3_TTS_TIMING
        t1 = clk::now();
        timing.t_talker_forward_ms += std::chrono::duration<double, std::milli>(t1 - t0).count();
#endif

        n_past++;
    }

#ifdef QWEN3_TTS_TIMING
    timing.t_generate_total_ms = std::chrono::duration<double, std::milli>(clk::now() - t_gen_start).count();
    impl_->timing = nullptr;
    const auto & t = timing;
    int nf = t.n_frames;
    fprintf(stderr, "\n=== Detailed Generation Timing (%d frames) ===\n", nf);
    fprintf(stderr, "  Prefill rows:       %8d\n", t.n_prefill_tokens);
    fprintf(stderr, "  Trailing rows:      %8d\n", t.n_trailing_tokens);
    fprintf(stderr, "\n  Prefill:\n");
    fprintf(stderr, "    Build prompt:     %8.1f ms\n", t.t_prefill_build_ms);
    fprintf(stderr, "      Special text:   %8.1f ms\n", t.t_prefill_special_text_proj_ms);
    fprintf(stderr, "      Instruct text:  %8.1f ms\n", t.t_prefill_instruct_text_proj_ms);
    fprintf(stderr, "      Role text:      %8.1f ms\n", t.t_prefill_role_text_proj_ms);
    fprintf(stderr, "      Body text:      %8.1f ms\n", t.t_prefill_body_text_proj_ms);
    fprintf(stderr, "      Codec lookups:  %8.1f ms\n", t.t_prefill_codec_lookup_ms);
    fprintf(stderr, "      Ref code embed: %8.1f ms\n", t.t_prefill_ref_code_embed_ms);
    fprintf(stderr, "      Compose CPU:    %8.1f ms\n", t.t_prefill_compose_ms);
    fprintf(stderr, "    Forward total:    %8.1f ms\n", t.t_prefill_forward_ms);
    fprintf(stderr, "      Graph build:    %8.1f ms\n", t.t_prefill_graph_build_ms);
    fprintf(stderr, "      Graph alloc:    %8.1f ms\n", t.t_prefill_graph_alloc_ms);
    fprintf(stderr, "      Compute:        %8.1f ms\n", t.t_prefill_compute_ms);
    fprintf(stderr, "      Data I/O:       %8.1f ms\n", t.t_prefill_data_ms);
    fprintf(stderr, "\n  Talker forward_step (total / per-frame):\n");
    fprintf(stderr, "    Total:            %8.1f ms   (%.1f ms/frame)\n", t.t_talker_forward_ms, nf > 0 ? t.t_talker_forward_ms / nf : 0.0);
    fprintf(stderr, "      Graph build:    %8.1f ms   (%.1f ms/frame)\n", t.t_talker_graph_build_ms, nf > 0 ? t.t_talker_graph_build_ms / nf : 0.0);
    fprintf(stderr, "      Graph alloc:    %8.1f ms   (%.1f ms/frame)\n", t.t_talker_graph_alloc_ms, nf > 0 ? t.t_talker_graph_alloc_ms / nf : 0.0);
    fprintf(stderr, "      Compute:        %8.1f ms   (%.1f ms/frame)\n", t.t_talker_compute_ms, nf > 0 ? t.t_talker_compute_ms / nf : 0.0);
    fprintf(stderr, "      Data I/O:       %8.1f ms   (%.1f ms/frame)\n", t.t_talker_data_ms, nf > 0 ? t.t_talker_data_ms / nf : 0.0);
    fprintf(stderr, "        Input upload: %8.1f ms   (%.1f ms/frame)\n", t.t_talker_input_upload_ms, nf > 0 ? t.t_talker_input_upload_ms / nf : 0.0);
    fprintf(stderr, "        Hidden read:  %8.1f ms   (%.1f ms/frame)\n", t.t_talker_hidden_read_ms, nf > 0 ? t.t_talker_hidden_read_ms / nf : 0.0);
    fprintf(stderr, "        Logits read:  %8.1f ms   (%.1f ms/frame)\n", t.t_talker_logits_read_ms, nf > 0 ? t.t_talker_logits_read_ms / nf : 0.0);
    fprintf(stderr, "        Sched reset:  %8.1f ms   (%.1f ms/frame)\n", t.t_talker_sched_reset_ms, nf > 0 ? t.t_talker_sched_reset_ms / nf : 0.0);
    fprintf(stderr, "\n  Code predictor (total / per-frame):\n");
    fprintf(stderr, "    Backend:          %s\n", impl_->use_coreml_code_predictor ? "CoreML (CPU+NE)" : "GGML");
    if (impl_->use_coreml_code_predictor && !impl_->coreml_code_predictor_path.empty()) {
        fprintf(stderr, "    CoreML model:     %s\n", impl_->coreml_code_predictor_path.c_str());
    }
    fprintf(stderr, "    Total:            %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_ms, nf > 0 ? t.t_code_pred_ms / nf : 0.0);
    fprintf(stderr, "      Init/KV/embed:  %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_init_ms, nf > 0 ? t.t_code_pred_init_ms / nf : 0.0);
    fprintf(stderr, "      Prefill (2tok): %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_prefill_ms, nf > 0 ? t.t_code_pred_prefill_ms / nf : 0.0);
    fprintf(stderr, "      Steps (14):     %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_steps_ms, nf > 0 ? t.t_code_pred_steps_ms / nf : 0.0);
    fprintf(stderr, "      Graph build:    %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_graph_build_ms, nf > 0 ? t.t_code_pred_graph_build_ms / nf : 0.0);
    fprintf(stderr, "        Prefill:      %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_prefill_graph_build_ms, nf > 0 ? t.t_code_pred_prefill_graph_build_ms / nf : 0.0);
    fprintf(stderr, "        Steps:        %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_steps_graph_build_ms, nf > 0 ? t.t_code_pred_steps_graph_build_ms / nf : 0.0);
    fprintf(stderr, "      Graph alloc:    %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_graph_alloc_ms, nf > 0 ? t.t_code_pred_graph_alloc_ms / nf : 0.0);
    fprintf(stderr, "        Prefill:      %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_prefill_graph_alloc_ms, nf > 0 ? t.t_code_pred_prefill_graph_alloc_ms / nf : 0.0);
    fprintf(stderr, "        Steps:        %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_steps_graph_alloc_ms, nf > 0 ? t.t_code_pred_steps_graph_alloc_ms / nf : 0.0);
    fprintf(stderr, "      Compute:        %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_compute_ms, nf > 0 ? t.t_code_pred_compute_ms / nf : 0.0);
    fprintf(stderr, "        Prefill:      %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_prefill_compute_ms, nf > 0 ? t.t_code_pred_prefill_compute_ms / nf : 0.0);
    fprintf(stderr, "        Steps:        %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_steps_compute_ms, nf > 0 ? t.t_code_pred_steps_compute_ms / nf : 0.0);
    fprintf(stderr, "      Data I/O:       %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_data_ms, nf > 0 ? t.t_code_pred_data_ms / nf : 0.0);
    fprintf(stderr, "        Prefill:      %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_prefill_data_ms, nf > 0 ? t.t_code_pred_prefill_data_ms / nf : 0.0);
    fprintf(stderr, "        Steps:        %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_steps_data_ms, nf > 0 ? t.t_code_pred_steps_data_ms / nf : 0.0);
    fprintf(stderr, "        Input upload: %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_input_upload_ms, nf > 0 ? t.t_code_pred_input_upload_ms / nf : 0.0);
    fprintf(stderr, "        Logits read:  %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_logits_read_ms, nf > 0 ? t.t_code_pred_logits_read_ms / nf : 0.0);
    fprintf(stderr, "        Sched reset:  %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_sched_reset_ms, nf > 0 ? t.t_code_pred_sched_reset_ms / nf : 0.0);
    fprintf(stderr, "      Sampling:       %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_sampling_ms, nf > 0 ? t.t_code_pred_sampling_ms / nf : 0.0);
    fprintf(stderr, "      CoreML total:   %8.1f ms   (%.1f ms/frame)\n", t.t_code_pred_coreml_ms, nf > 0 ? t.t_code_pred_coreml_ms / nf : 0.0);
    fprintf(stderr, "\n  Embed lookups:      %8.1f ms   (%.1f ms/frame)\n", t.t_embed_lookup_ms, nf > 0 ? t.t_embed_lookup_ms / nf : 0.0);
    double accounted = t.t_prefill_build_ms + t.t_prefill_forward_ms + t.t_talker_forward_ms + t.t_code_pred_ms + t.t_embed_lookup_ms;
    fprintf(stderr, "  Other/overhead:     %8.1f ms\n", t.t_generate_total_ms - accounted);
    fprintf(stderr, "  ─────────────────────────────────────────\n");
    fprintf(stderr, "  Total generate:     %8.1f ms\n", t.t_generate_total_ms);
    if (nf > 0) {
        fprintf(stderr, "  Throughput:         %8.1f ms/frame (%.1f frames/s)\n",
                t.t_generate_total_ms / nf, 1000.0 * nf / t.t_generate_total_ms);
    }
#endif

    return true;
}

} // namespace qwen3_tts
