CoolFace
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 2d agoView on Hugging Face
0likes1.1kdownloads
diffusion.cpp409 linesDownload Raw Back to diffusion
1#include "diffusion.h"2 3#include "log.h"4 5#include <algorithm>6#include <cstddef>7#include <cmath>8#include <cstring>9#include <random>10#include <utility>11#include <vector>12 13static float calculate_confidence(const llama_token_data_array & cur_p,14                                  diffusion_algorithm            algorithm,15                                  std::mt19937 &                 rng) {16    switch (algorithm) {17        case DIFFUSION_ALGORITHM_CONFIDENCE_BASED:18            return cur_p.data[cur_p.selected].p;  // Selected token probability19 20        case DIFFUSION_ALGORITHM_ENTROPY_BASED:21            {22                float       entropy = 0.0f;23                const float epsilon = 1e-10f;24                for (size_t i = 0; i < cur_p.size; i++) {25                    float prob = cur_p.data[i].p;26                    entropy += prob * logf(prob + epsilon);27                }28                return -entropy;  // Higher entropy = lower confidence29            }30 31        case DIFFUSION_ALGORITHM_MARGIN_BASED:32            return (cur_p.size > 1) ? cur_p.data[0].p - cur_p.data[1].p : cur_p.data[0].p;33 34        case DIFFUSION_ALGORITHM_RANDOM:35            {36                std::uniform_real_distribution<float> uniform(0.0f, 1.0f);37                return uniform(rng);  // Random confidence38            }39 40        case DIFFUSION_ALGORITHM_ORIGIN:41            return cur_p.data[cur_p.selected].p;42 43        default:44            return 0.0f;45    }46}47 48// Unified transfer count calculation function49static int32_t calculate_transfer_count(int32_t                      step,50                                        int32_t                      total_steps,51                                        int32_t                      remaining_masked,52                                        diffusion_transfer_schedule  schedule,53                                        float                        eps,54                                        const std::vector<int32_t> & num_transfer_tokens = {}) {55    switch (schedule) {56        case DIFFUSION_TRANSFER_SCHEDULE_TIMESTEP_BASED:57            {58                float t          = 1.0f - (float) step / total_steps * (1.0f - eps);59                float s          = 1.0f - (float) (step + 1) / total_steps * (1.0f - eps);60                float p_transfer = (step < total_steps - 1) ? (1.0f - s / t) : 1.0f;61                return (int32_t) (remaining_masked * p_transfer);62            }63 64        case DIFFUSION_TRANSFER_SCHEDULE_BLOCK_BASED:65            if (!num_transfer_tokens.empty() && step < (int32_t) num_transfer_tokens.size()) {66                return num_transfer_tokens[step];67            }68            return remaining_masked / (total_steps - step);  // Fallback69 70        default:71            return remaining_masked / (total_steps - step);72    }73}74 75static void add_gumbel_noise(float * logits, int32_t n_vocab, float temperature, std::mt19937 & rng) {76    if (temperature == 0.0f) {77        return;78    }79 80    std::uniform_real_distribution<double> uniform(0.0, 1.0);81    for (int32_t i = 0; i < n_vocab; i++) {82        double noise        = uniform(rng);83        // Prevent log(0)84        noise               = std::max(noise, 1e-20);85        double gumbel_noise = std::pow(-std::log(noise), temperature);86        logits[i]           = std::exp(logits[i]) / gumbel_noise;87    }88}89 90static std::vector<int32_t> get_num_transfer_tokens(int32_t mask_count, int32_t steps) {91    std::vector<int32_t> num_transfer_tokens(steps);92 93    int32_t base      = mask_count / steps;94    int32_t remainder = mask_count % steps;95 96    for (int32_t i = 0; i < steps; i++) {97        num_transfer_tokens[i] = base + (i < remainder ? 1 : 0);98    }99 100    return num_transfer_tokens;101}102 103void diffusion_generate(llama_context *          ctx,104                        const llama_token *      input_tokens,105                        llama_token *            output_tokens,106                        int32_t                  n_input,107                        const diffusion_params & params,108                        int32_t &                n_generated) {109    n_generated = 0;110    if (!ctx || !input_tokens || !output_tokens || n_input <= 0 || params.max_length <= n_input) {111        return;112    }113 114    const llama_model * model = llama_get_model(ctx);115 116    // Initialize with input and pad with mask tokens117    std::copy(input_tokens, input_tokens + n_input, output_tokens);118    std::fill(output_tokens + n_input, output_tokens + params.max_length, params.mask_token_id);119 120    std::mt19937 rng(params.seed);121 122    llama_set_causal_attn(ctx, false);123 124    int32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(model));125 126    std::vector<llama_token_data> candidates(n_vocab);127    std::vector<llama_token_data> conf_candidates;128    conf_candidates.reserve(params.max_length);129    std::vector<int32_t> mask_positions;130    mask_positions.reserve(params.max_length);131 132    // Setup sampler chain133    struct llama_sampler * sampler = llama_sampler_chain_init(llama_sampler_chain_default_params());134    if (params.top_k > 0) {135        llama_sampler_chain_add(sampler, llama_sampler_init_top_k(params.top_k));136    }137    if (params.top_p < 1.0f) {138        llama_sampler_chain_add(sampler, llama_sampler_init_top_p(params.top_p, 1));139    }140    if (params.temperature > 0.0f) {141        llama_sampler_chain_add(sampler, llama_sampler_init_temp(params.temperature));142    }143    llama_sampler_chain_add(sampler, llama_sampler_init_dist(params.seed));144 145    struct llama_sampler * dist_sampler = llama_sampler_init_dist(params.seed);146 147    llama_batch batch = llama_batch_init(params.max_length, 0, 1);148    batch.n_tokens    = params.max_length;149 150    // Pre-allocate buffers for CFG if needed151    int32_t                  logits_size = n_vocab * params.max_length;152    std::vector<float>       cond_logits_buffer;153    std::vector<llama_token> un_x_buffer;154    if (params.cfg_scale > 0.0f) {155        cond_logits_buffer.resize(logits_size);156        un_x_buffer.resize(params.max_length);157    }158 159    // For block-based processing160    std::vector<int32_t> num_transfer_tokens;161    int32_t              num_blocks      = 1;162    int32_t              steps_per_block = params.steps;163 164    if (params.schedule == DIFFUSION_TRANSFER_SCHEDULE_BLOCK_BASED) {165        GGML_ASSERT(params.max_length % params.block_length == 0);166        num_blocks = params.max_length / params.block_length;167        GGML_ASSERT(params.steps % num_blocks == 0);168        steps_per_block = params.steps / num_blocks;169    }170 171    std::vector<float> confidence(params.max_length);172 173    int64_t total_sampling_time = 0;174    int64_t total_time          = 0;175    int64_t time_start          = ggml_time_us();176 177    for (int block_num = 0; block_num < num_blocks; block_num++) {178        int32_t block_start = (params.schedule == DIFFUSION_TRANSFER_SCHEDULE_BLOCK_BASED) ? n_input + block_num * params.block_length : 0;179        int32_t block_end   = (params.schedule == DIFFUSION_TRANSFER_SCHEDULE_BLOCK_BASED) ?180                                  std::min(n_input + (block_num + 1) * params.block_length, params.max_length) :181                                  params.max_length;182 183        // Count masked tokens in current block for block-based processing184        if (params.schedule == DIFFUSION_TRANSFER_SCHEDULE_BLOCK_BASED) {185            int32_t block_mask_count = 0;186            for (int i = block_start; i < block_end; i++) {187                if (output_tokens[i] == params.mask_token_id) {188                    block_mask_count++;189                }190            }191            num_transfer_tokens = get_num_transfer_tokens(block_mask_count, steps_per_block);192        }193 194        for (int32_t step = 0; step < steps_per_block; step++) {195            int32_t global_step = block_num * steps_per_block + step;196 197            if (params.step_callback) {198                if (!params.step_callback(199                        global_step, params.steps, output_tokens, params.max_length, params.step_callback_user_data)) {200                    break;201                }202            }203 204            // Setup batch205            for (int32_t i = 0; i < params.max_length; i++) {206                batch.token[i]     = output_tokens[i];207                batch.pos[i]       = i;208                batch.n_seq_id[i]  = 1;209                batch.seq_id[i][0] = 0;210                batch.logits[i]    = 1;211            }212 213            float * logits = nullptr;214 215            if (params.cfg_scale > 0.0f) {216                int ret = llama_decode(ctx, batch);217                if (ret != 0) {218                    LOG_ERR("Failed to generate conditional");219                    break;220                }221                float * cond_logits_ptr = llama_get_logits(ctx);222                std::memcpy(cond_logits_buffer.data(), cond_logits_ptr, logits_size * sizeof(float));223 224                // Unconditional generation (mask input)225                std::copy(output_tokens, output_tokens + params.max_length, un_x_buffer.begin());226                for (int32_t i = 0; i < n_input; i++) {227                    un_x_buffer[i] = params.mask_token_id;228                }229 230                for (int32_t i = 0; i < params.max_length; i++) {231                    batch.token[i] = un_x_buffer[i];232                }233                ret = llama_decode(ctx, batch);234                if (ret != 0) {235                    LOG_ERR("Failed to generate unconditional");236                    break;237                }238                float * uncond_logits = llama_get_logits(ctx);239 240                // Apply CFG241                for (int32_t i = 0; i < logits_size; i++) {242                    cond_logits_buffer[i] =243                        uncond_logits[i] + (params.cfg_scale + 1.0f) * (cond_logits_buffer[i] - uncond_logits[i]);244                }245                logits = cond_logits_buffer.data();246            } else {247                int ret = llama_decode(ctx, batch);248                if (ret != 0) {249                    LOG_ERR("%s: failed to decode at step %d, ret = %d\n", __func__, global_step, ret);250                    break;251                }252                logits = llama_get_logits(ctx);253            }254 255            if (!logits) {256                LOG_ERR("%s: failed to get logits at step %d\n", __func__, global_step);257                break;258            }259 260            auto get_logits_for_pos = [&](int32_t pos) -> const float * {261                if (params.shift_logits) {262                    return pos == 0 ? logits : logits + (pos - 1) * n_vocab;263                }264                return logits + pos * n_vocab;265            };266 267            int64_t time_start_sampling = ggml_time_us();268 269            mask_positions.clear();270            for (int32_t i = 0; i < params.max_length; i++) {271                if (output_tokens[i] == params.mask_token_id) {272                    // For block-based, only consider current block273                    if (params.schedule != DIFFUSION_TRANSFER_SCHEDULE_BLOCK_BASED || (i >= block_start && i < block_end)) {274                        mask_positions.push_back(i);275                    }276                }277            }278 279            if (mask_positions.empty()) {280                break;281            }282 283            if (params.add_gumbel_noise && params.temperature > 0.0f) {284                add_gumbel_noise(logits, n_vocab, params.temperature, rng);285            }286 287            if (params.algorithm == DIFFUSION_ALGORITHM_ORIGIN) {288                int32_t transfer_count = calculate_transfer_count(289                    step, steps_per_block, mask_positions.size(), params.schedule, params.eps, num_transfer_tokens);290                float p_transfer = (float) transfer_count / mask_positions.size();291 292                for (int32_t pos : mask_positions) {293                    if (std::uniform_real_distribution<float>(0.0f, 1.0f)(rng) < p_transfer) {294                        const float * pos_logits = get_logits_for_pos(pos);295                        for (int32_t token_id = 0; token_id < n_vocab; token_id++) {296                            candidates[token_id].id    = token_id;297                            candidates[token_id].logit = pos_logits[token_id];298                            candidates[token_id].p     = 0.0f;299                        }300 301                        llama_token_data_array cur_p = {302                            candidates.data(),303                            (size_t) n_vocab,304                            -1,305                            false,306                        };307 308                        llama_sampler_apply(sampler, &cur_p);309                        output_tokens[pos] = cur_p.data[cur_p.selected].id;310                    }311                }312            } else {313                std::vector<std::pair<float, int32_t>> confidences;314                std::vector<llama_token>               sampled_tokens(mask_positions.size());315 316                for (size_t i = 0; i < mask_positions.size(); i++) {317                    int32_t       pos        = mask_positions[i];318                    const float * pos_logits = get_logits_for_pos(pos);319 320                    for (int32_t token_id = 0; token_id < n_vocab; token_id++) {321                        candidates[token_id].logit = pos_logits[token_id];322                        candidates[token_id].p     = 0.0f;323                        candidates[token_id].id    = token_id;324                    }325 326                    llama_token_data_array cur_p = {327                        candidates.data(),328                        candidates.size(),329                        -1,330                        false,331                    };332 333                    llama_sampler_apply(sampler, &cur_p);334                    llama_token sampled_token = cur_p.data[cur_p.selected].id;335 336                    float conf = calculate_confidence(cur_p, params.algorithm, rng);337 338                    sampled_tokens[i] = sampled_token;339                    confidences.emplace_back(conf, i);340                }341 342                int32_t transfer_count = calculate_transfer_count(343                    step, steps_per_block, mask_positions.size(), params.schedule, params.eps, num_transfer_tokens);344 345                if (transfer_count > 0) {346                    if (params.alg_temp == 0.0f) {347                        std::partial_sort(confidences.begin(),348                                          confidences.begin() + std::min(transfer_count, (int32_t) confidences.size()),349                                          confidences.end(),350                                          [](const std::pair<float, int32_t> & a, const std::pair<float, int32_t> & b) {351                                              if (a.first != b.first) {352                                                  return a.first > b.first;353                                              }354                                              return a.second < b.second;355                                          });356 357                        for (int32_t i = 0; i < std::min(transfer_count, (int32_t) confidences.size()); i++) {358                            int32_t mask_idx   = confidences[i].second;359                            int32_t pos        = mask_positions[mask_idx];360                            output_tokens[pos] = sampled_tokens[mask_idx];361                        }362                    } else {363                        conf_candidates.clear();364                        for (size_t i = 0; i < confidences.size(); i++) {365                            float conf_logit = confidences[i].first / params.alg_temp;366                            conf_candidates.emplace_back(llama_token_data{ (int32_t) i, conf_logit, 0.0f });367                        }368 369                        llama_token_data_array conf_array = {370                            conf_candidates.data(),371                            conf_candidates.size(),372                            -1,373                            false,374                        };375 376                        for (int32_t i = 0; i < std::min(transfer_count, (int32_t) confidences.size()); i++) {377                            llama_sampler_apply(dist_sampler, &conf_array);378                            int32_t selected_idx = conf_array.selected;379                            int32_t mask_idx     = selected_idx;380                            int32_t pos          = mask_positions[mask_idx];381                            output_tokens[pos]   = sampled_tokens[mask_idx];382 383                            conf_candidates[selected_idx].p = 0.0f;384                            conf_array.selected             = -1;385                        }386                    }387                }388            }389 390            int64_t time_end_sampling = ggml_time_us();391            total_sampling_time += time_end_sampling - time_start_sampling;392        }393    }394 395    int64_t time_end = ggml_time_us();396    total_time += time_end - time_start;397 398    LOG_INF("\ntotal time: %0.2fms, time per step: %0.2fms, sampling time per step: %0.2fms\n",399            total_time / 1000.0,400            total_time / 1000.0 / params.steps,401            total_sampling_time / 1000.0 / params.steps);402 403    llama_batch_free(batch);404    llama_sampler_free(sampler);405    llama_sampler_free(dist_sampler);406 407    n_generated = params.max_length;408}409