Felipe97/llama-cpp-compiled
01.1k
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 