CoolFace
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 2d agoView on Hugging Face
0likes1.1kdownloads
diffusion-cli.cpp267 linesDownload Raw Back to diffusion
1#include "arg.h"2#include "chat.h"3#include "common.h"4#include "diffusion.h"5#include "llama.h"6#include "log.h"7 8#include <limits.h>9 10#include <clocale>11#include <cstring>12#include <string>13#include <vector>14 15struct callback_data {16    diffusion_params *  diff_params;17    const llama_vocab * vocab;18    int32_t             n_input;19};20 21static bool diffusion_step_callback(int32_t             step,22                                    int32_t             total_steps,23                                    const llama_token * tokens,24                                    int32_t             n_tokens,25                                    void *              user_data) {26    (void) user_data;27 28    callback_data * data = static_cast<callback_data *>(user_data);29 30    auto print_progress_bar = [](int32_t step, int32_t total_steps) {31        int progress_percent = (step * 100) / total_steps;32        int progress_bars    = (step * 50) / total_steps;33        LOG_INF("\rdiffusion step: %d/%d [%s%s] %d%%",34                step,35                total_steps,36                std::string(progress_bars, '=').c_str(),37                std::string(50 - progress_bars, ' ').c_str(),38                progress_percent);39    };40 41    if (data->diff_params->visual_mode) {42        // Visual mode: clear43        LOG_INF("\033[2J\033[H");  // Clear screen and move cursor to top-left44 45        print_progress_bar(step, total_steps);46 47        LOG_INF("\n");48 49        std::string current_text = " ";50 51        for (int32_t i = data->n_input; i < n_tokens; i++) {52            std::string token_str;53            if (tokens[i] != llama_vocab_mask(data->vocab)) {54                char piece[256];55                int  n_chars = llama_token_to_piece(data->vocab, tokens[i], piece, sizeof(piece), 0, false);56                if (n_chars > 0) {57                    piece[n_chars] = '\0';58                    token_str      = piece;59                }60            } else {61                token_str = " ";62            }63 64            current_text += token_str;65        }66 67        LOG_INF("%s\n", current_text.c_str());68    } else {69        print_progress_bar(step, total_steps);70    }71 72    return true;73}74 75static std::string format_input_text(const std::string & prompt, const std::string & system_prompt, bool use_chat_template, llama_model * model) {76    if (!use_chat_template) {77        return prompt;78    }79 80    auto chat_templates = common_chat_templates_init(model, "");81    common_chat_templates_inputs inputs;82    common_chat_msg system_msg;83 84    if (!system_prompt.empty()) {85        system_msg.role = "system";86        system_msg.content = system_prompt;87        inputs.messages.push_back(system_msg);88    }89 90    common_chat_msg user_msg;91    user_msg.role = "user";92    user_msg.content = prompt;93 94    inputs.messages.push_back(user_msg);95    inputs.add_generation_prompt = true;96 97    auto result = common_chat_templates_apply(chat_templates.get(), inputs);98 99    return result.prompt;100}101 102int main(int argc, char ** argv) {103    std::setlocale(LC_NUMERIC, "C");104 105    ggml_time_init();106 107    common_params params;108 109    common_init();110 111    if (!common_params_parse(argc, argv, params, LLAMA_EXAMPLE_DIFFUSION)) {112        return 1;113    }114 115    llama_backend_init();116 117    llama_model_params model_params = llama_model_default_params();118    model_params.n_gpu_layers       = params.n_gpu_layers;119    model_params.devices            = params.devices.data();120    model_params.load_mode          = params.load_mode;121    model_params.check_tensors      = params.check_tensors;122 123    llama_model * model = llama_model_load_from_file(params.model.path.c_str(), model_params);124    if (!model) {125        LOG_ERR("error: failed to load model '%s'\n", params.model.path.c_str());126        return 1;127    }128 129    if (!llama_model_is_diffusion(model)) {130        LOG_ERR("error: unsupported model for diffusion");131        llama_model_free(model);132        return 1;133    }134 135    llama_context_params ctx_params = llama_context_default_params();136    ctx_params.n_ctx                = params.n_ctx;137    ctx_params.n_batch              = params.n_batch;138    ctx_params.n_ubatch             = params.n_ubatch;139    ctx_params.flash_attn_type      = params.flash_attn_type;140    ctx_params.no_perf              = params.no_perf;141    ctx_params.type_k               = params.cache_type_k;142    ctx_params.type_v               = params.cache_type_v;143 144    llama_context * ctx = llama_init_from_model(model, ctx_params);145    if (!ctx) {146        LOG_ERR("error: failed to create context\n");147        llama_model_free(model);148        return 1;149    }150 151    llama_set_n_threads(ctx, params.cpuparams.n_threads, params.cpuparams_batch.n_threads);152 153    const llama_vocab * vocab            = llama_model_get_vocab(model);154 155    std::string         formatted_prompt = format_input_text(params.prompt, params.system_prompt, params.enable_chat_template, model);156 157    std::vector<llama_token> input_tokens = common_tokenize(vocab,158                                                            formatted_prompt,159                                                            /*add special tokens*/ true,160                                                            /*parse special*/ true);161 162    int n_input = input_tokens.size();163 164    if (static_cast<uint32_t>(n_input) >= llama_n_ctx(ctx)) {165        LOG_ERR("error: input too long (%d tokens), max context is %d\n", n_input, llama_n_ctx(ctx));166        llama_free(ctx);167        llama_model_free(model);168        return 1;169    }170 171    llama_token mask_token_id = llama_vocab_mask(vocab);172 173    GGML_ASSERT(mask_token_id != LLAMA_TOKEN_NULL);174 175    bool visual_mode = params.diffusion.visual_mode;176 177    int32_t                  n_generated = 0;178    std::vector<llama_token> output_tokens(params.n_ubatch);179 180    struct diffusion_params diff_params;181 182    char shift_logits_str[8];183    if (llama_model_meta_val_str(model, "diffusion.shift_logits", shift_logits_str, sizeof(shift_logits_str)) >= 0) {184        diff_params.shift_logits = (strcmp(shift_logits_str, "true") == 0);185    } else {186        diff_params.shift_logits = true;187    }188 189    //Use either eps or block length, but not both190    GGML_ASSERT((params.diffusion.eps == 0) ^ (params.diffusion.block_length == 0));191 192    if (params.diffusion.eps) {193        diff_params.schedule = DIFFUSION_TRANSFER_SCHEDULE_TIMESTEP_BASED;194        diff_params.eps      = params.diffusion.eps;195    } else if (params.diffusion.block_length) {196        diff_params.schedule     = DIFFUSION_TRANSFER_SCHEDULE_BLOCK_BASED;197        diff_params.block_length = params.diffusion.block_length;198    }199 200    diff_params.mask_token_id    = mask_token_id;201    diff_params.seed             = params.sampling.seed;202    diff_params.temperature      = params.sampling.temp;203    diff_params.steps            = params.diffusion.steps;204    diff_params.algorithm        = static_cast<diffusion_algorithm>(params.diffusion.algorithm);205    diff_params.max_length       = params.n_ubatch;206    diff_params.top_p            = params.sampling.top_p;207    diff_params.top_k            = params.sampling.top_k;208    diff_params.visual_mode      = params.diffusion.visual_mode;209    diff_params.add_gumbel_noise = params.diffusion.add_gumbel_noise;210 211    diff_params.step_callback           = diffusion_step_callback;212    callback_data cb_data               = { &diff_params, vocab, n_input };213    diff_params.step_callback_user_data = &cb_data;214 215    const char * alg_names[]   = {216        "DIFFUSION_ALGORITHM_ORIGIN",217        "DIFFUSION_ALGORITHM_ENTROPY_BASED",218        "DIFFUSION_ALGORITHM_MARGIN_BASED",219        "DIFFUSION_ALGORITHM_RANDOM",220        "DIFFUSION_ALGORITHM_CONFIDENCE_BASED",221    };222    const char * sched_names[] = {223        "DIFFUSION_TRANSFER_SCHEDULE_TIMESTEP_BASED",224        "DIFFUSION_TRANSFER_SCHEDULE_BLOCK_BASED",225    };226    const char * alg_name =227        (diff_params.algorithm >= 0 && diff_params.algorithm <= 4) ? alg_names[diff_params.algorithm] : "UNKNOWN";228    const char * sched_name =229        (diff_params.schedule >= 0 && diff_params.schedule <= 1) ? sched_names[diff_params.schedule] : "UNKNOWN";230 231    LOG_INF("diffusion_params: - %-25s llama_token      = %d\n", "mask_token_id", mask_token_id);232    LOG_INF("diffusion_params: - %-25s u32              = %d\n", "steps", diff_params.steps);233    LOG_INF("diffusion_params: - %-25s u32              = %d\n", "max_length", diff_params.max_length);234    LOG_INF("diffusion_params: - %-25s enum             = %d (%s)\n", "algorithm", diff_params.algorithm, alg_name);235    LOG_INF("diffusion_params: - %-25s enum             = %d (%s)\n", "schedule", diff_params.schedule, sched_name);236    LOG_INF("diffusion_params: - %-25s f32              = %.3f\n", "temperature", diff_params.temperature);237    if (diff_params.schedule == DIFFUSION_TRANSFER_SCHEDULE_TIMESTEP_BASED) {238        LOG_INF("diffusion_params: - %-25s f32              = %.6f\n", "eps", diff_params.eps);239        LOG_INF("diffusion_params: - %-25s f32              = %.3f\n", "alg_temp", diff_params.alg_temp);240    }241    if (diff_params.schedule == DIFFUSION_TRANSFER_SCHEDULE_BLOCK_BASED) {242        LOG_INF("diffusion_params: - %-25s u32              = %d\n", "block_length", diff_params.block_length);243        LOG_INF("diffusion_params: - %-25s f32              = %.3f\n", "cfg_scale", diff_params.cfg_scale);244    }245 246    diffusion_generate(ctx, input_tokens.data(), output_tokens.data(), n_input, diff_params, n_generated);247 248    if (n_generated > 0) {249        if (visual_mode) {250            //clear screen and move cursor to top-left251            LOG_INF("\033[2J\033[H");252        }253 254        output_tokens.erase(output_tokens.begin(), output_tokens.begin() + n_input);255        std::string output_data = common_detokenize(vocab, output_tokens, false);256        LOG_INF("\n%s\n", output_data.c_str());257    } else {258        LOG_INF("Error: diffusion generation failed\n");259    }260 261    llama_free(ctx);262    llama_model_free(model);263    llama_backend_free();264 265    return 0;266}267