Team Ai
Datasetpublic

echodict/llama.cpp

version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes479downloads
test-llama-archs.cpp664 linesDownload Raw Back to tests
1#include "common.h"2#include "log.h"3#include "ggml-backend.h"4#include "ggml.h"5#include "gguf.h"6#include "ggml-cpp.h"7#include "llama.h"8#include "llama-cpp.h"9 10// TODO: replace with #include "llama-ext.h" in the future11#include "../src/llama-arch.h"12#include "../src/llama-model-saver.h"13 14#include <cinttypes>15#include <cstdio>16#include <cstring>17#include <cstdint>18#include <random>19#include <stdexcept>20#include <string>21#include <utility>22#include <vector>23 24// normalized mean squared error = mse(a, b) / mse(a, 0)25static double nmse(const std::vector<float> & a, const std::vector<float> & b) {26    GGML_ASSERT(a.size() == b.size());27    double mse_a_b = 0.0;28    double mse_a_0 = 0.0;29 30    for (size_t i = 0; i < a.size(); i++) {31        float a_i = a[i];32        float b_i = b[i];33 34        mse_a_b += (a_i - b_i) * (a_i - b_i);35        mse_a_0 += a_i * a_i;36    }37 38    return mse_a_b / mse_a_0;39}40 41static void set_tensor_data(struct ggml_tensor * tensor, void * userdata) {42    std::hash<std::string> hasher;43    std::mt19937 gen(hasher(tensor->name) + *(const size_t *) userdata);44    std::normal_distribution<float> dis(0.0f, 1.0e-2f);45 46    const int64_t ne = ggml_nelements(tensor);47    if (tensor->type == GGML_TYPE_F32) {48        std::vector<float> tmp(ne);49        for (int64_t i = 0; i < ne; i++) {50            tmp[i] = dis(gen);51        }52        ggml_backend_tensor_set(tensor, tmp.data(), 0, ggml_nbytes(tensor));53    } else if (tensor->type == GGML_TYPE_F16) {54        std::vector<ggml_fp16_t> tmp(ne);55        for (int64_t i = 0; i < ne; i++) {56            tmp[i] = ggml_fp32_to_fp16(dis(gen));57        }58        ggml_backend_tensor_set(tensor, tmp.data(), 0, ggml_nbytes(tensor));59    } else {60        GGML_ABORT("fatal error");61    }62}63 64static void usage(char ** argv) {65    printf("Usage: %s [-a/--arch arch] [-s/--seed seed] [-v/--verbose]\n", argv[0]);66}67 68static std::vector<llama_token> get_tokens(const uint32_t n_tokens, const uint32_t n_vocab, const size_t seed){69    std::mt19937 gen(seed);70    std::uniform_int_distribution<> dis(0, n_vocab - 1);71    std::vector<llama_token> ret;72    ret.reserve(n_tokens);73    for (uint32_t i = 0; i < n_tokens; i++) {74        ret.push_back(dis(gen));75    }76    return ret;77}78 79static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) {80    gguf_context_ptr ret(gguf_init_empty());81    llama_model_saver ms(arch, ret.get());82    const uint32_t n_ctx = 128;83 84    uint32_t n_vocab = 128;85    uint32_t n_embd  = 256;86    uint32_t n_head  = 2;87    uint32_t n_ff    = 384;88    uint32_t n_layer = 2;89    if (arch == LLM_ARCH_LLAMA4) {90        n_layer = 4; // hparams.n_no_rope_layer_step is hard-coded to 491    } else if (arch == LLM_ARCH_GEMMA4) {92        n_embd = 128;93        n_head = 2;94        n_ff   = 192;95        n_layer = 5; // need at least 5 for swa_pattern (every 5th is full_attention)96    } else if (arch == LLM_ARCH_GEMMA3N) {97        n_embd = 64;98        n_head = 1;99        n_ff   = 96;100        n_layer = 22; // hparams.n_layer_kv_from_start = 20 is hardcoded101    } else if (arch == LLM_ARCH_DEEPSEEK2102            || arch == LLM_ARCH_GLM_DSA103            || arch == LLM_ARCH_KIMI_LINEAR104            || arch == LLM_ARCH_MISTRAL4) {105        n_embd = 128;106        n_head = 1;107        n_ff   = 192;108    } else if (arch == LLM_ARCH_NEMOTRON_H || arch == LLM_ARCH_NEMOTRON_H_MOE) {109        n_layer = 3;110    } else if (arch == LLM_ARCH_CHAMELEON) {111        n_vocab = 10240;112    }113 114    const uint32_t n_embd_head = n_embd / n_head;115 116    ms.add_kv(LLM_KV_GENERAL_ARCHITECTURE,      llm_arch_name(arch));117    ms.add_kv(LLM_KV_VOCAB_SIZE,                n_vocab);118    ms.add_kv(LLM_KV_CONTEXT_LENGTH,            n_ctx);119    ms.add_kv(LLM_KV_EMBEDDING_LENGTH,          n_embd);120    ms.add_kv(LLM_KV_FEATURES_LENGTH,           n_embd);121    ms.add_kv(LLM_KV_BLOCK_COUNT,               n_layer);122    ms.add_kv(LLM_KV_LEADING_DENSE_BLOCK_COUNT, uint32_t(1));123 124    if (arch == LLM_ARCH_NEMOTRON_H || arch == LLM_ARCH_NEMOTRON_H_MOE) {125        std::vector<uint32_t> n_ff_per_layer;126        n_ff_per_layer.reserve(n_layer);127        for (uint32_t il = 0; il < n_layer; il++) {128            n_ff_per_layer.push_back(il <= 1 ? 0 : n_ff);129        }130        ms.add_kv(LLM_KV_FEED_FORWARD_LENGTH, n_ff_per_layer);131    } else {132        ms.add_kv(LLM_KV_FEED_FORWARD_LENGTH, n_ff);133    }134 135    ms.add_kv(LLM_KV_USE_PARALLEL_RESIDUAL,   false);136    ms.add_kv(LLM_KV_LOGIT_SCALE,             1.0f);137    ms.add_kv(LLM_KV_TIME_MIX_EXTRA_DIM,      uint32_t(64));138    ms.add_kv(LLM_KV_TIME_DECAY_EXTRA_DIM,    uint32_t(128));139    ms.add_kv(LLM_KV_FULL_ATTENTION_INTERVAL, uint32_t(2));140 141    if (arch == LLM_ARCH_PLAMO2 || arch == LLM_ARCH_JAMBA || arch == LLM_ARCH_NEMOTRON_H || arch == LLM_ARCH_NEMOTRON_H_MOE ||142            arch == LLM_ARCH_GRANITE_HYBRID || arch == LLM_ARCH_LFM2 || arch == LLM_ARCH_LFM2MOE || arch == LLM_ARCH_KIMI_LINEAR) {143        GGML_ASSERT(n_layer >= 2);144        std::vector<uint32_t> n_head_per_layer;145        n_head_per_layer.reserve(n_layer);146        for (uint32_t il = 0; il < n_layer; il++) {147            n_head_per_layer.push_back(il == 1 ? 0 : n_head);148        }149        ms.add_kv(LLM_KV_ATTENTION_HEAD_COUNT, n_head_per_layer);150        ms.add_kv(LLM_KV_ATTENTION_HEAD_COUNT_KV, n_head_per_layer);151    } else {152        ms.add_kv(LLM_KV_ATTENTION_HEAD_COUNT, n_head);153        ms.add_kv(LLM_KV_ATTENTION_HEAD_COUNT_KV, n_head);154    }155 156    ms.add_kv(LLM_KV_ATTENTION_MAX_ALIBI_BIAS, 8.0f);157    if (arch == LLM_ARCH_DEEPSEEK2158            || arch == LLM_ARCH_GLM_DSA159            || arch == LLM_ARCH_KIMI_LINEAR160            || arch == LLM_ARCH_MISTRAL4) {161        ms.add_kv(LLM_KV_ATTENTION_KEY_LENGTH,       uint32_t(576));162        ms.add_kv(LLM_KV_ATTENTION_VALUE_LENGTH,     uint32_t(512));163        ms.add_kv(LLM_KV_ROPE_DIMENSION_COUNT,       uint32_t(64));164        ms.add_kv(LLM_KV_ATTENTION_KEY_LENGTH_MLA,   uint32_t(192));165        ms.add_kv(LLM_KV_ATTENTION_VALUE_LENGTH_MLA, uint32_t(128));166    }167    ms.add_kv(LLM_KV_ATTENTION_CLAMP_KQV,              1.0f);168    ms.add_kv(LLM_KV_ATTENTION_LAYERNORM_EPS,          1e-5f);169    ms.add_kv(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS,      1e-5f);170    ms.add_kv(LLM_KV_ATTENTION_GROUPNORM_EPS,          1e-5f);171    ms.add_kv(LLM_KV_ATTENTION_GROUPNORM_GROUPS,       uint32_t(8));172    ms.add_kv(LLM_KV_ATTENTION_Q_LORA_RANK,            uint32_t(512));173    ms.add_kv(LLM_KV_ATTENTION_KV_LORA_RANK,           uint32_t(512));174    ms.add_kv(LLM_KV_ATTENTION_RELATIVE_BUCKETS_COUNT, uint32_t(8));175    ms.add_kv(LLM_KV_ATTENTION_SLIDING_WINDOW,         n_ctx/8);176 177    if (arch == LLM_ARCH_GEMMA4) {178        ms.add_kv(LLM_KV_EMBEDDING_LENGTH_PER_LAYER,      n_embd/2);179        ms.add_kv(LLM_KV_ATTENTION_SHARED_KV_LAYERS,      uint32_t(0));180        ms.add_kv(LLM_KV_ATTENTION_KEY_LENGTH_SWA,        n_embd_head);181        ms.add_kv(LLM_KV_ATTENTION_VALUE_LENGTH_SWA,      n_embd_head);182        ms.add_kv(LLM_KV_ROPE_FREQ_BASE_SWA,              10000.0f);183        // SWA pattern: every 5th layer is full attention (matches E2B layer_types)184        ms.add_kv(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, uint32_t(5));185    } else if (arch == LLM_ARCH_MIMO2 || arch == LLM_ARCH_STEP35) {186        std::vector<uint32_t> pattern;187        pattern.reserve(n_layer);188        for (uint32_t il = 0; il < n_layer; il++) {189            pattern.push_back(il % 2);190        }191        ms.add_kv(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, pattern);192    } else {193        ms.add_kv(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, uint32_t(2));194    }195 196    ms.add_kv(LLM_KV_ATTENTION_INDEXER_HEAD_COUNT, uint32_t(1));197    ms.add_kv(LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, uint32_t(64));198    ms.add_kv(LLM_KV_ATTENTION_INDEXER_TOP_K,      uint32_t(8));199    ms.add_kv(LLM_KV_ROPE_DIMENSION_SECTIONS, std::vector<uint32_t>({n_embd_head/4, n_embd_head/4, n_embd_head/4, n_embd_head/4}));200    ms.add_kv(LLM_KV_TOKENIZER_MODEL,         "no_vocab");201    // ms.add_kv(LLM_KV_DENSE_2_FEAT_OUT,     n_embd);202    // ms.add_kv(LLM_KV_DENSE_3_FEAT_IN,      n_embd);203 204    if (moe) {205        ms.add_kv(LLM_KV_EXPERT_FEED_FORWARD_LENGTH, n_ff);206        ms.add_kv(LLM_KV_INTERLEAVE_MOE_LAYER_STEP,  uint32_t(2));207        ms.add_kv(LLM_KV_EXPERT_COUNT,               uint32_t(2));208        ms.add_kv(LLM_KV_EXPERT_USED_COUNT,          uint32_t(1));209        ms.add_kv(LLM_KV_EXPERT_SHARED_COUNT,        uint32_t(1));210        ms.add_kv(LLM_KV_EXPERT_GATING_FUNC,         uint32_t(2)); // sigmoid211        ms.add_kv(LLM_KV_EXPERT_GROUP_SCALE,         1.0f);212        ms.add_kv(LLM_KV_EXPERTS_PER_GROUP,          uint32_t(1));213    }214 215    ms.add_kv(LLM_KV_POSNET_EMBEDDING_LENGTH,   n_embd);216    ms.add_kv(LLM_KV_POSNET_BLOCK_COUNT,        n_layer);217    ms.add_kv(LLM_KV_CONVNEXT_EMBEDDING_LENGTH, n_embd);218    ms.add_kv(LLM_KV_CONVNEXT_BLOCK_COUNT,      n_layer);219    ms.add_kv(LLM_KV_XIELU_ALPHA_N,             1.0f);220    ms.add_kv(LLM_KV_XIELU_ALPHA_P,             1.0f);221    ms.add_kv(LLM_KV_XIELU_BETA,                1.0f);222    ms.add_kv(LLM_KV_XIELU_EPS,                 1.0e-7f);223    ms.add_kv(LLM_KV_SSM_INNER_SIZE,            arch == LLM_ARCH_QWEN3NEXT || arch == LLM_ARCH_QWEN35 || arch == LLM_ARCH_QWEN35MOE ? 256 : 2*n_embd);224    ms.add_kv(LLM_KV_SSM_CONV_KERNEL,           uint32_t(4));225    ms.add_kv(LLM_KV_SSM_STATE_SIZE,            uint32_t(128));226    ms.add_kv(LLM_KV_SSM_TIME_STEP_RANK,        n_head);227    ms.add_kv(LLM_KV_SSM_GROUP_COUNT,           arch == LLM_ARCH_PLAMO2 ? 0 : uint32_t(2));228    ms.add_kv(LLM_KV_KDA_HEAD_DIM,              uint32_t(128));229    ms.add_kv(LLM_KV_WKV_HEAD_SIZE,             n_embd/n_head);230    ms.add_kv(LLM_KV_SHORTCONV_L_CACHE,         uint32_t(3));231 232    for (uint32_t il = 0; il < n_layer; il++) {233        ggml_tensor t;234        memset(&t, 0, sizeof(ggml_tensor));235        t.type = GGML_TYPE_F16;236        ggml_format_name(&t, "conv%" PRIu32 "d.weight", il);237        gguf_add_tensor(ms.gguf_ctx, &t);238        ggml_format_name(&t, "posnet.%" PRIu32 ".conv1.weight", il);239        gguf_add_tensor(ms.gguf_ctx, &t);240        ggml_format_name(&t, "posnet.%" PRIu32 ".conv2.weight", il);241        gguf_add_tensor(ms.gguf_ctx, &t);242        ggml_format_name(&t, "convnext.%" PRIu32 ".dw.weight", il);243        gguf_add_tensor(ms.gguf_ctx, &t);244    }245    return ret;246}247 248static bool silent_model_load_progress(float /*progress*/, void * /*user_data*/) {249    return true;250}251 252static std::pair<llama_model_ptr, llama_context_ptr> get_model_and_ctx(253        struct gguf_context * gguf_ctx, FILE * file, const size_t seed, const std::vector<ggml_backend_dev_t> & devs,254        const llama_split_mode split_mode = LLAMA_SPLIT_MODE_LAYER, bool encode = false) {255    GGML_ASSERT((gguf_ctx == nullptr) != (file == nullptr));256    llama_model_params model_params = llama_model_default_params();257    model_params.progress_callback = silent_model_load_progress;258    std::vector<ggml_backend_dev_t> devs_copy = devs;259    devs_copy.push_back(nullptr);260    model_params.devices = devs_copy.data();261    model_params.split_mode = split_mode;262 263    llama_context_params ctx_params = llama_context_default_params();264    ctx_params.n_ctx = 0;265    ctx_params.n_threads = 4;266    ctx_params.n_threads_batch = 4;267    if (!encode) {268        ctx_params.n_ubatch = 64;269    }270 271    size_t tmp = seed;272    llama_model_ptr model(gguf_ctx != nullptr ?273        llama_model_init_from_user(gguf_ctx, set_tensor_data, &tmp, model_params) :274        llama_model_load_from_file_ptr(file, model_params));275    if (!model) {276        throw std::runtime_error("failed to create llama model");277    }278    llama_context_ptr lctx(llama_init_from_model(model.get(), ctx_params));279    if (!lctx) {280        throw std::runtime_error("failed to create llama context");281    }282    return std::make_pair(std::move(model), std::move(lctx));283}284 285static std::vector<float> get_logits(286        llama_model * model, llama_context * lctx, const std::vector<llama_token> & tokens, bool encode = false) {287    const uint32_t n_vocab  = llama_vocab_n_tokens(llama_model_get_vocab(model));288    const uint32_t n_ctx    = llama_n_ctx(lctx);289    const uint32_t n_tokens = tokens.size();290    llama_batch batch = llama_batch_init(n_ctx, 0, 1);291    GGML_ASSERT(n_tokens <= n_ctx);292    for (uint32_t pos = 0; pos < n_tokens; pos++) {293        common_batch_add(batch, tokens[pos], pos, {0}, true);294    }295    batch.n_tokens = n_tokens;296    if (encode) {297        if (llama_encode(lctx, batch)) {298            llama_batch_free(batch);299            throw std::runtime_error("failed to encode batch");300        }301    }302    if (llama_decode(lctx, batch)) {303        llama_batch_free(batch);304        throw std::runtime_error("failed to decode batch");305    }306 307    std::vector<float> ret;308    ret.reserve(n_tokens*n_vocab);309    for (uint32_t i = 0; i < n_tokens; i++) {310        const float * logits_ith = llama_get_logits_ith(lctx, i);311        for (uint32_t j = 0; j < n_vocab; j++) {312            ret.push_back(logits_ith[j]);313        }314    }315    llama_batch_free(batch);316    return ret;317}318 319static bool moe_mandatory(const llm_arch arch) {320    switch (arch) {321        case LLM_ARCH_LLAMA4:322        case LLM_ARCH_GROK:323        case LLM_ARCH_QWEN2MOE:324        case LLM_ARCH_QWEN3MOE:325        case LLM_ARCH_QWEN3NEXT:326        case LLM_ARCH_QWEN3VLMOE:327        case LLM_ARCH_QWEN35MOE:328        case LLM_ARCH_PHIMOE:329        case LLM_ARCH_DBRX:330        case LLM_ARCH_OLMOE:331        case LLM_ARCH_ARCTIC:332        case LLM_ARCH_DEEPSEEK:333        case LLM_ARCH_DEEPSEEK2:334        case LLM_ARCH_GLM4_MOE:335        case LLM_ARCH_GLM_DSA:336        case LLM_ARCH_EXAONE_MOE:337        case LLM_ARCH_BAILINGMOE:338        case LLM_ARCH_BAILINGMOE2:339        case LLM_ARCH_DOTS1:340        case LLM_ARCH_AFMOE:341        case LLM_ARCH_ERNIE4_5:342        case LLM_ARCH_ERNIE4_5_MOE:343        case LLM_ARCH_HUNYUAN_MOE:344        case LLM_ARCH_OPENAI_MOE:345        case LLM_ARCH_LFM2MOE:346        case LLM_ARCH_SMALLTHINKER:347        case LLM_ARCH_LLADA_MOE:348        case LLM_ARCH_GROVEMOE:349        case LLM_ARCH_MINIMAX_M2:350        case LLM_ARCH_RND1:351        case LLM_ARCH_PADDLEOCR:352        case LLM_ARCH_MIMO2:353        case LLM_ARCH_KIMI_LINEAR:354        case LLM_ARCH_STEP35:355        case LLM_ARCH_MISTRAL4:356            return true;357        default:358            return false;359    }360}361 362static bool moe_implemented(const llm_arch arch) {363    if (moe_mandatory(arch)) {364        return true;365    }366    switch (arch) {367        case LLM_ARCH_LLAMA:368        case LLM_ARCH_REFACT:369        case LLM_ARCH_MINICPM:370        case LLM_ARCH_GRANITE:371        case LLM_ARCH_GRANITE_MOE:372        case LLM_ARCH_MISTRAL3:373        case LLM_ARCH_LLAMA_EMBED:374            return true;375        default:376            return false;377    }378}379 380static bool arch_supported(const llm_arch arch) {381    if (arch == LLM_ARCH_CLIP || arch == LLM_ARCH_GPTJ || arch == LLM_ARCH_UNKNOWN) {382        return false; // These models don't have usable implementations.383    }384    if (arch == LLM_ARCH_CHAMELEON) {385        return false; // Only half-implemented and to be removed in the future.386    }387    if (arch == LLM_ARCH_WAVTOKENIZER_DEC) {388        return false; // FIXME CUDA backend crashes.389    }390    if (arch == LLM_ARCH_GEMMA4) {391        return false; // FIXME @ngxson392    }393    if (arch == LLM_ARCH_LLAMA_EMBED || arch == LLM_ARCH_GEMMA_EMBEDDING || arch == LLM_ARCH_T5ENCODER) {394        return false; // FIXME Embedding (?) models produce inconsistent results.395    }396    if (arch == LLM_ARCH_RWKV6 || arch == LLM_ARCH_RWKV6QWEN2 || arch == LLM_ARCH_RWKV7 || arch == LLM_ARCH_ARWKV7) {397        return false; // FIXME RWKV models hang indefinitely.398    }399    if (arch == LLM_ARCH_BERT || arch == LLM_ARCH_MODERN_BERT || arch == LLM_ARCH_NOMIC_BERT || arch == LLM_ARCH_NOMIC_BERT_MOE ||400            arch == LLM_ARCH_NEO_BERT || arch == LLM_ARCH_JINA_BERT_V2 || arch == LLM_ARCH_JINA_BERT_V3 || arch == LLM_ARCH_EUROBERT) {401        return false; // TODO vocab402    }403    if (arch == LLM_ARCH_PLM) {404        return false; // TODO tensor shapes405    }406    if (arch == LLM_ARCH_DEEPSEEK2OCR) {407        return false;408    }409 410    // FIXME some models are segfaulting with WebGPU:411#ifdef GGML_USE_WEBGPU412    if (arch == LLM_ARCH_QWEN3NEXT || arch == LLM_ARCH_QWEN35 || arch == LLM_ARCH_QWEN35MOE || arch == LLM_ARCH_KIMI_LINEAR) {413        return false;414    }415#endif // GGML_USE_WEBGPU416 417    return true;418}419 420static int save_models(const llm_arch target_arch, const size_t seed, const ggml_log_level log_level, const std::string & dir) {421    struct user_data_t {422        struct {423            ggml_log_callback callback;424            void * user_data;425        } original_logger;426        ggml_log_level min_level; // prints below this log level go to debug log427    };428    user_data_t ud;429    llama_log_get(&ud.original_logger.callback, &ud.original_logger.user_data);430    ud.min_level = log_level;431 432    llama_log_set([](ggml_log_level level, const char * text, void * user_data) {433        const user_data_t * ud = (const user_data_t *) user_data;434        const ggml_log_level level_eff = level >= ud->min_level ? level : GGML_LOG_LEVEL_DEBUG;435        ud->original_logger.callback(level_eff, text, ud->original_logger.user_data);436    }, &ud);437 438    for (const llm_arch & arch : llm_arch_all()) {439        if (arch == LLM_ARCH_UNKNOWN) {440            continue;441        }442        if (target_arch != LLM_ARCH_UNKNOWN && arch != target_arch) {443            continue;444        }445        if (arch == LLM_ARCH_GEMMA4) {446            continue; // FIXME: ISWA KV cache initialization needs more fixture params447        }448        for (bool moe : {false, true}) {449            if (moe && !moe_implemented(arch)) {450                continue;451            }452            if (!moe && moe_mandatory(arch)) {453                continue;454            }455            if (!llama_model_saver_supports_arch(arch)) {456                LOG_INF("%s: %s model (%s) is unsupported, skipping\n", __func__, llm_arch_name(arch), moe ? "MoE" : "dense");457                continue;458            }459            gguf_context_ptr gguf_ctx = get_gguf_ctx(arch, moe);460            auto model_and_ctx = get_model_and_ctx(gguf_ctx.get(), nullptr, seed, {});461            const std::string path = dir + "/" + llm_arch_name(arch) + (moe ? "-moe.gguf" : "-dense.gguf");462            LOG_INF("%s: Saving %s model (%s) to %s...\n", __func__, llm_arch_name(arch), moe ? "MoE" : "dense", path.c_str());463            llama_model_save_to_file(model_and_ctx.first.get(), path.c_str());464        }465    }466    llama_log_set(ud.original_logger.callback, ud.original_logger.user_data);467    return 0;468}469 470static int test_backends(const llm_arch target_arch, const size_t seed, const ggml_log_level log_level) {471    struct user_data_t {472        struct {473            ggml_log_callback callback;474            void * user_data;475        } original_logger;476        ggml_log_level min_level; // prints below this log level go to debug log477    };478    user_data_t ud;479    llama_log_get(&ud.original_logger.callback, &ud.original_logger.user_data);480    ud.min_level = log_level;481 482    llama_log_set([](ggml_log_level level, const char * text, void * user_data) {483        const user_data_t * ud = (const user_data_t *) user_data;484        const ggml_log_level level_eff = level >= ud->min_level ? level : GGML_LOG_LEVEL_DEBUG;485        ud->original_logger.callback(level_eff, text, ud->original_logger.user_data);486    }, &ud);487 488    const std::vector<llama_token> tokens = get_tokens(128, 128, seed);489 490    struct device_config {491        std::vector<ggml_backend_dev_t> devs;492        std::string                     label;493        llama_split_mode                split_mode;494 495        device_config(std::vector<ggml_backend_dev_t> devs, std::string name, llama_split_mode split_mode)496            : devs(std::move(devs)), label(std::move(name)), split_mode(split_mode) {}497    };498 499    std::vector<device_config> dev_configs;500    {501        std::vector<ggml_backend_dev_t> devices_meta;502        {503            const size_t device_count = ggml_backend_dev_count();504            for (size_t i = 0; i < device_count; i++) {505                ggml_backend_dev_t dev = ggml_backend_dev_get(i);506                dev_configs.emplace_back(std::vector<ggml_backend_dev_t>{dev}, ggml_backend_dev_description(dev), LLAMA_SPLIT_MODE_LAYER);507 508                // cpu-based devices cannot be used in tensor split mode509                if (ggml_backend_dev_buffer_type(dev) != ggml_backend_cpu_buffer_type()) {510                    devices_meta.push_back(dev);511                }512            }513        }514 515        dev_configs.emplace_back(devices_meta, "Meta", LLAMA_SPLIT_MODE_TENSOR);516    }517 518    bool all_ok = true;519    common_log_flush(common_log_main());520    printf("|%16s|%30s|%6s|%15s|%9s|\n", "Model arch.", "Device", "Config", "NMSE vs. CPU", "Roundtrip");521    printf("|----------------|------------------------------|------|---------------|---------|\n");522    for (const llm_arch & arch : llm_arch_all()) {523        if (arch == LLM_ARCH_UNKNOWN) {524            continue;525        }526        if (target_arch != LLM_ARCH_UNKNOWN && arch != target_arch) {527            continue;528        }529        if (arch == LLM_ARCH_GEMMA4) {530            continue; // FIXME: ISWA KV cache initialization needs more fixture params531        }532 533        const bool encode = arch == LLM_ARCH_T5 || arch == LLM_ARCH_DREAM || arch == LLM_ARCH_LLADA || arch == LLM_ARCH_LLADA_MOE || arch == LLM_ARCH_RND1;534        for (bool moe : {false, true}) {535            if (moe && !moe_implemented(arch)) {536                continue;537            }538            if (!moe && moe_mandatory(arch)) {539                continue;540            }541            const std::string config_name = moe ? "MoE" : "Dense";542            gguf_context_ptr gguf_ctx = get_gguf_ctx(arch, moe);543            std::pair<llama_model_ptr, llama_context_ptr> model_and_ctx_cpu;544            std::vector<float> logits_cpu;545            for (device_config & dc : dev_configs) {546                std::pair<llama_model_ptr, llama_context_ptr> model_and_ctx_dev;547                std::vector<float> logits_dev;548                std::string status_nmse      = "\033[1;33mSKIP\033[0m";549                std::string status_roundtrip = "\033[1;33mSKIP\033[0m";550                char nmse_str[12] = {0};551                bool skip = !arch_supported(arch) || (dc.split_mode == LLAMA_SPLIT_MODE_TENSOR && dc.devs.empty());552#if defined(GGML_USE_WEBGPU)553                skip = true; // FIXME554#endif // GGML_USE_WEBGPU555                if (!skip) {556                    if (logits_cpu.empty()) {557                        model_and_ctx_cpu = get_model_and_ctx(gguf_ctx.get(), nullptr, seed, {}, LLAMA_SPLIT_MODE_LAYER, encode);558                        logits_cpu = get_logits(model_and_ctx_cpu.first.get(), model_and_ctx_cpu.second.get(), tokens, encode);559                    }560                    if (dc.split_mode != LLAMA_SPLIT_MODE_TENSOR || llm_arch_supports_sm_tensor(arch)) {561                        model_and_ctx_dev = get_model_and_ctx(gguf_ctx.get(), nullptr, seed, dc.devs, dc.split_mode, encode);562                        logits_dev = get_logits(model_and_ctx_dev.first.get(), model_and_ctx_dev.second.get(), tokens, encode);563                        const double nmse_val = nmse(logits_cpu, logits_dev);564                        snprintf(nmse_str, sizeof(nmse_str), "(%.2e)", nmse_val);565                        status_nmse = "\033[1;32mOK\033[0m";566                        if (nmse_val > 1e-4) {567                            all_ok = false;568                            status_nmse = "\033[1;31mFAIL\033[0m";569                        }570                    }571 572                    FILE * file = tmpfile(); // Can be null on Windows without administrator privileges.573                    // FIXME: when adding a tensor to a gguf_context a copy is made, this changes the pointer which the meta backend574                    //     in turn uses to map the tensors to their simple equivalents - this is fundamentally incompatible575                    if (file != nullptr && llama_model_saver_supports_arch(arch) && dc.split_mode != LLAMA_SPLIT_MODE_TENSOR) {576                        GGML_ASSERT(model_and_ctx_dev.first && model_and_ctx_dev.second);577                        llama_model_saver ms = llama_model_saver(model_and_ctx_dev.first.get());578                        ms.add_kv_from_model();579                        ms.add_tensors_from_model();580                        ms.save(file);581                        rewind(file);582 583                        auto model_and_ctx_roundtrip = get_model_and_ctx(nullptr, file, seed, dc.devs, dc.split_mode, encode);584                        const std::vector<float> logits_roundtrip = get_logits(585                            model_and_ctx_roundtrip.first.get(), model_and_ctx_roundtrip.second.get(), tokens, encode);586                        status_roundtrip = "\033[1;32mOK\033[0m";587                        GGML_ASSERT(logits_roundtrip.size() == logits_dev.size());588                        for (size_t i = 0; i < logits_roundtrip.size(); i++) {589                            if (logits_roundtrip[i] != logits_dev[i]) {590                                all_ok = false;591                                status_roundtrip = "\033[1;31mFAIL\033[0m";592                                break;593                            }594                        }595                    }596                }597 598                printf("|%16s|%30s|%6s|%15s %10s|%20s|\n", llm_arch_name(arch), dc.label.c_str(),599                    config_name.c_str(), status_nmse.c_str(), nmse_str, status_roundtrip.c_str());600            }601        }602    }603    llama_log_set(ud.original_logger.callback, ud.original_logger.user_data);604    return all_ok ? 0 : 1;605}606 607int main(int argc, char ** argv) {608    // FIXME these tests are disabled in the CI for macOS-latest-cmake-arm64 because they are segfaulting609    common_init();610    std::random_device rd;611 612    llm_arch arch = LLM_ARCH_UNKNOWN;613    size_t seed = rd();614    ggml_log_level log_level = GGML_LOG_LEVEL_ERROR;615    std::string out;616 617    for (int i = 1; i < argc; i++) {618        if (strcmp(argv[i], "-a") == 0 || strcmp(argv[i], "--arch") == 0) {619            if (i + 1 < argc) {620                const std::string arch_name = argv[++i];621                arch = llm_arch_from_string(arch_name);622                if (arch == LLM_ARCH_UNKNOWN) {623                    LOG_ERR("%s: unkown LLM architecture: %s\n", __func__, arch_name.c_str());624                    return 1;625                }626            } else {627                usage(argv);628                return 1;629            }630        }631        if (strcmp(argv[i], "-s") == 0 || strcmp(argv[i], "--seed") == 0) {632            if (i + 1 < argc) {633                seed = std::stoull(argv[++i]);634            } else {635                usage(argv);636                return 1;637            }638        }639        if (strcmp(argv[i], "-v") == 0 || strcmp(argv[i], "--verbose") == 0) {640            log_level = GGML_LOG_LEVEL_INFO;641            continue;642        }643        if (strcmp(argv[i], "-o") == 0 || strcmp(argv[i], "--out") == 0) {644            if (i + 1 < argc) {645                out = argv[++i];646            } else {647                usage(argv);648                return 1;649            }650        }651    }652    printf("%s: using seed %zu\n", __func__, seed);653 654    try {655        if (!out.empty()) {656            return save_models(arch, seed, log_level, out);657        }658        return test_backends(arch, seed, log_level);659    } catch (const std::exception & err) {660        fprintf(stderr, "encountered runtime error: %s\n", err.what());661        return -1;662    }663}664