#include "llama.h"
#include "ggml.h"
#include <cstdio>
#include <cstdlib>
#include <cstring>
#include <string>
#include <vector>
#include <fstream>
struct dump_state {
int target_pos; int current_pos; std::string out_dir;
bool active; };
static dump_state g_dump;
static bool eval_callback(struct ggml_tensor * t, bool ask, void * user_data) {
if (ask) return true;
if (!g_dump.active) return true;
const char * name = ggml_get_name(t);
if (!name) return true;
bool is_l_out = (strncmp(name, "l_out", 5) == 0);
bool is_attn_out = (strncmp(name, "attn_out", 8) == 0);
bool is_kqv_out = (strncmp(name, "kqv_out", 7) == 0);
bool is_kqv = (!is_kqv_out && strncmp(name, "kqv", 3) == 0 && (name[3] == '-' || name[3] == '\0'));
bool is_qcur_pos = (strncmp(name, "Qcur_pos", 8) == 0);
bool is_kcur_pos = (strncmp(name, "Kcur_pos", 8) == 0);
bool is_attn_norm = (strncmp(name, "attn_norm", 9) == 0
&& (name[9] == '-' || name[9] == '\0'));
bool is_qcur_normed = (strncmp(name, "Qcur_normed", 11) == 0);
bool is_kcur_normed = (strncmp(name, "Kcur_normed", 11) == 0);
bool is_vcur_normed = (strncmp(name, "Vcur_normed", 11) == 0);
bool is_qcur = (!is_qcur_pos && !is_qcur_normed
&& strncmp(name, "Qcur", 4) == 0 && (name[4] == '-' || name[4] == '\0'));
bool is_kcur = (!is_kcur_pos && !is_kcur_normed
&& strncmp(name, "Kcur", 4) == 0 && (name[4] == '-' || name[4] == '\0'));
bool is_vcur = (!is_vcur_normed
&& strncmp(name, "Vcur", 4) == 0 && (name[4] == '-' || name[4] == '\0'));
bool is_cache_k = (strncmp(name, "cache_k_l", 9) == 0);
bool is_cache_v = (strncmp(name, "cache_v_l", 9) == 0);
bool is_ffn_norm_1 = (strncmp(name, "ffn_norm_1-", 11) == 0);
bool is_ffn_norm_2 = (strncmp(name, "ffn_norm_2-", 11) == 0);
bool is_ffn_mlp = (strncmp(name, "ffn_mlp-", 8) == 0);
bool is_ffn_moe_logits = (strncmp(name, "ffn_moe_logits-", 15) == 0);
bool is_ffn_moe_combined = (strncmp(name, "ffn_moe_combined-", 17) == 0);
bool is_ffn_moe = (!is_ffn_moe_logits && !is_ffn_moe_combined
&& strncmp(name, "ffn_moe-", 8) == 0);
bool is_ffn_norm = (!is_ffn_norm_1 && !is_ffn_norm_2
&& strncmp(name, "ffn_norm-", 9) == 0);
bool is_ffn_out = (strncmp(name, "ffn_out-", 8) == 0);
bool is_ffn_post_norm = (strncmp(name, "ffn_post_norm-", 14) == 0);
bool is_out_scaled = (strncmp(name, "out_scaled-", 11) == 0);
if (!is_l_out && !is_attn_out && !is_kqv_out && !is_kqv
&& !is_qcur_pos && !is_kcur_pos
&& !is_attn_norm && !is_qcur_normed && !is_kcur_normed && !is_vcur_normed
&& !is_qcur && !is_kcur && !is_vcur
&& !is_ffn_norm_1 && !is_ffn_norm_2 && !is_ffn_mlp
&& !is_ffn_moe_logits && !is_ffn_moe_combined && !is_ffn_moe
&& !is_ffn_norm && !is_ffn_out && !is_ffn_post_norm && !is_out_scaled
&& !is_cache_k && !is_cache_v) return true;
int layer = -1;
if (is_cache_k || is_cache_v) {
const char * l_marker = strstr(name, "_l");
if (l_marker) layer = atoi(l_marker + 2);
} else {
const char * dash = strrchr(name, '-');
if (dash) layer = atoi(dash + 1);
}
const char * prefix = is_l_out ? "l_out"
: is_attn_out ? "attn_out"
: is_kqv_out ? "kqv_out"
: is_kqv ? "kqv"
: is_qcur_pos ? "q_normed"
: is_kcur_pos ? "k_normed"
: is_attn_norm ? "attn_norm_out"
: is_qcur_normed ? "qcur_normed"
: is_kcur_normed ? "kcur_normed"
: is_vcur_normed ? "vcur_normed"
: is_qcur ? "qcur"
: is_kcur ? "kcur"
: is_vcur ? "vcur"
: is_ffn_norm_1 ? "ffn_norm_1"
: is_ffn_norm_2 ? "ffn_norm_2"
: is_ffn_mlp ? "ffn_mlp"
: is_ffn_moe_logits ? "ffn_moe_logits"
: is_ffn_moe_combined ? "ffn_moe_combined"
: is_ffn_moe ? "ffn_moe"
: is_ffn_norm ? "ffn_norm"
: is_ffn_out ? "ffn_out"
: is_ffn_post_norm ? "ffn_post_norm"
: is_out_scaled ? "out_scaled"
: is_cache_k ? "cache_k"
: is_cache_v ? "cache_v"
: "unknown";
static const bool dump_all_cache = (getenv("HF2Q_DUMP_ALL_CACHE") != nullptr);
if ((is_cache_k || is_cache_v) && !dump_all_cache && layer != 24) return true;
int64_t n_elements = ggml_nelements(t);
size_t elem_bytes = ggml_type_size(t->type);
size_t n_bytes = n_elements * elem_bytes;
std::vector<char> raw_data(n_bytes);
ggml_backend_tensor_get(t, raw_data.data(), 0, n_bytes);
std::vector<float> data(n_elements);
if (t->type == GGML_TYPE_F32) {
memcpy(data.data(), raw_data.data(), n_bytes);
} else if (t->type == GGML_TYPE_F16) {
const uint16_t * src = (const uint16_t *)raw_data.data();
for (int64_t i = 0; i < n_elements; i++) {
uint16_t h = src[i];
uint32_t sign = (h & 0x8000) << 16;
uint32_t exp = (h & 0x7C00) >> 10;
uint32_t frac = (h & 0x03FF);
uint32_t f32_bits;
if (exp == 0) {
if (frac == 0) f32_bits = sign;
else {
while (!(frac & 0x0400)) { frac <<= 1; exp--; }
exp++;
frac &= 0x03FF;
f32_bits = sign | ((exp + 112) << 23) | (frac << 13);
}
} else if (exp == 0x1F) {
f32_bits = sign | 0x7F800000 | (frac << 13);
} else {
f32_bits = sign | ((exp + 112) << 23) | (frac << 13);
}
memcpy(&data[i], &f32_bits, 4);
}
} else {
fprintf(stderr, "[DUMP] skipping %s: unsupported dtype %d\n", name, (int)t->type);
return true;
}
char path[512];
snprintf(path, sizeof(path), "%s/llama_%s_layer%02d_pos%d.bin",
g_dump.out_dir.c_str(), prefix, layer, g_dump.current_pos);
FILE * f = fopen(path, "wb");
if (f) {
fwrite(data.data(), sizeof(float), n_elements, f);
fclose(f);
fprintf(stderr, "[DUMP] %s: %lld f32 (src %s) shape=[%lld,%lld,%lld,%lld] nb=[%zu,%zu,%zu,%zu] -> %s\n",
name, (long long)n_elements, ggml_type_name(t->type),
(long long)t->ne[0], (long long)t->ne[1], (long long)t->ne[2], (long long)t->ne[3],
(size_t)t->nb[0], (size_t)t->nb[1], (size_t)t->nb[2], (size_t)t->nb[3],
path);
}
return true;
}
int main(int argc, char ** argv) {
if (argc < 5) {
fprintf(stderr, "Usage: %s <gguf> <prompt_file> <target_decode_token> <output_dir>\n", argv[0]);
return 1;
}
const char * model_path = argv[1];
const char * prompt_file = argv[2];
int target_decode_token = atoi(argv[3]);
const char * out_dir = argv[4];
std::ifstream pf(prompt_file);
std::string prompt((std::istreambuf_iterator<char>(pf)),
std::istreambuf_iterator<char>());
fprintf(stderr, "Prompt: %zu bytes\n", prompt.size());
llama_backend_init();
auto mparams = llama_model_default_params();
mparams.n_gpu_layers = 999;
auto * model = llama_model_load_from_file(model_path, mparams);
if (!model) {
fprintf(stderr, "Failed to load model\n");
return 1;
}
auto cparams = llama_context_default_params();
cparams.n_ctx = 2048;
cparams.n_batch = 512;
auto * ctx = llama_init_from_model(model, cparams);
if (!ctx) {
fprintf(stderr, "Failed to create context\n");
return 1;
}
g_dump.out_dir = out_dir;
g_dump.active = false;
g_dump.target_pos = -1;
llama_free(ctx);
cparams.cb_eval = eval_callback;
cparams.cb_eval_user_data = nullptr;
ctx = llama_init_from_model(model, cparams);
if (!ctx) {
fprintf(stderr, "Failed to create context with callback\n");
return 1;
}
const auto * vocab = llama_model_get_vocab(model);
int n_prompt_max = prompt.size() + 256;
std::vector<llama_token> tokens(n_prompt_max);
int n_tokens = llama_tokenize(vocab, prompt.c_str(), prompt.size(),
tokens.data(), n_prompt_max, true, true);
if (n_tokens < 0) {
fprintf(stderr, "Tokenization failed\n");
return 1;
}
tokens.resize(n_tokens);
fprintf(stderr, "Tokens: %d\n", n_tokens);
const char * prefill_pos_env = getenv("HF2Q_PREFILL_DUMP_POS");
int prefill_dump_pos = prefill_pos_env ? atoi(prefill_pos_env) : -1;
const bool per_token_prefill = (getenv("HF2Q_PER_TOKEN_PREFILL") != nullptr);
if (prefill_dump_pos >= 0 || per_token_prefill) {
fprintf(stderr, "Prefilling %d tokens one-by-one%s...\n",
n_tokens,
prefill_dump_pos >= 0 ? " (with dump)" : "");
for (int i = 0; i < n_tokens; i++) {
if (i == prefill_dump_pos) {
g_dump.active = true;
g_dump.current_pos = i;
fprintf(stderr, "Activating dump at prefill pos %d\n", i);
} else {
g_dump.active = false;
}
llama_token t = tokens[i];
llama_batch b = llama_batch_get_one(&t, 1);
if (llama_decode(ctx, b) != 0) {
fprintf(stderr, "Prefill decode failed at pos %d\n", i);
return 1;
}
}
g_dump.active = false;
if (prefill_dump_pos >= 0 && !per_token_prefill) {
fprintf(stderr, "Prefill dump complete.\n");
llama_free(ctx);
llama_model_free(model);
llama_backend_free();
return 0;
}
} else {
const bool batched_dump = (getenv("HF2Q_BATCHED_DUMP_POS") != nullptr);
if (batched_dump) {
int pos_flag = atoi(getenv("HF2Q_BATCHED_DUMP_POS"));
fprintf(stderr, "Batched prefill with dump (target pos %d).\n", pos_flag);
g_dump.active = true;
g_dump.current_pos = pos_flag;
} else {
fprintf(stderr, "Prefilling %d tokens in one batch...\n", n_tokens);
}
llama_batch batch = llama_batch_get_one(tokens.data(), n_tokens);
if (llama_decode(ctx, batch) != 0) {
fprintf(stderr, "Prefill failed\n");
return 1;
}
g_dump.active = false;
if (batched_dump) {
fprintf(stderr, "Batched dump complete.\n");
llama_free(ctx);
llama_model_free(model);
llama_backend_free();
return 0;
}
}
const char * predict_env = getenv("HF2Q_PREDICT");
int n_predict = predict_env ? atoi(predict_env) : (target_decode_token + 5);
const char * output_file = getenv("HF2Q_OUTPUT_FILE");
FILE * output_fp = output_file ? fopen(output_file, "w") : nullptr;
int n_decoded = 0;
llama_token prev_token = -1;
for (int i = 0; i < n_predict; i++) {
auto * logits = llama_get_logits_ith(ctx, -1);
llama_token best = 0;
float best_logit = logits[0];
int n_vocab = llama_vocab_n_tokens(vocab);
for (int v = 1; v < n_vocab; v++) {
if (logits[v] > best_logit) {
best_logit = logits[v];
best = v;
}
}
if (llama_vocab_is_eog(vocab, best)) {
fprintf(stderr, "EOS reached at step %d\n", i);
break;
}
if (output_fp) {
char piece[128];
int len = llama_token_to_piece(vocab, best, piece, sizeof(piece), 0, false);
if (len > 0) {
fwrite(piece, 1, len, output_fp);
fflush(output_fp);
}
}
int seq_pos = n_tokens + i; if (i == target_decode_token - 1) {
g_dump.active = true;
g_dump.current_pos = n_tokens + i;
fprintf(stderr, "Activating dump at decode step %d (seq_pos=%d)\n", i+1, n_tokens + i + 1);
} else {
g_dump.active = false;
}
llama_batch next_batch = llama_batch_get_one(&best, 1);
if (llama_decode(ctx, next_batch) != 0) {
fprintf(stderr, "Decode failed at step %d\n", i);
return 1;
}
if (i == target_decode_token) {
auto * target_logits = llama_get_logits_ith(ctx, -1);
char lpath[512];
snprintf(lpath, sizeof(lpath), "%s/llama_logits_pos%d.bin", out_dir, n_tokens + i);
FILE * f = fopen(lpath, "wb");
if (f) {
fwrite(target_logits, sizeof(float), n_vocab, f);
fclose(f);
fprintf(stderr, "[DUMP] logits (%d f32) -> %s\n", n_vocab, lpath);
}
std::vector<std::pair<int, float>> indexed;
for (int v = 0; v < n_vocab; v++) {
indexed.push_back({v, target_logits[v]});
}
std::sort(indexed.begin(), indexed.end(),
[](auto & a, auto & b) { return a.second > b.second; });
fprintf(stderr, "Top-10 logits at seq_pos=%d:\n", n_tokens + i);
for (int k = 0; k < 10; k++) {
fprintf(stderr, " tok=%6d logit=%.6f\n", indexed[k].first, indexed[k].second);
}
g_dump.active = false;
}
prev_token = best;
n_decoded++;
}
fprintf(stderr, "Decoded %d tokens\n", n_decoded);
if (output_fp) fclose(output_fp);
llama_free(ctx);
llama_model_free(model);
llama_backend_free();
return 0;
}