Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 32 additions & 2 deletions include/s2_model.h
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
#include "ggml-backend.h"
#include "ggml-cpu.h"
#include "gguf.h"
#include "s2_sampler.h"
#ifdef GGML_USE_VULKAN
#include "ggml-vulkan.h"
#endif
Expand Down Expand Up @@ -123,11 +124,18 @@ class SlowARModel {
bool step(const std::vector<int32_t> & flat_tokens, int32_t n_threads,
StepResult & result);

bool fast_decode(const std::vector<float> & hidden,
const std::vector<int32_t> & prefix_codes,
bool fast_decode(const std::vector<float> & hidden_in,
const std::vector<int32_t> & prefix_tokens,
int32_t n_threads,
std::vector<float> & logits_out);

bool fast_decode_batch(
const std::vector<float> & hidden_in,
int32_t semantic_code,
int32_t n_threads,
const SamplerParams & sparams,
std::vector<int32_t> & codebooks_out);

const ModelHParams & hparams() const { return hparams_; }

ModelWeights weights_;
Expand All @@ -138,6 +146,7 @@ class SlowARModel {
ggml_backend_t backend_gpu_ = nullptr;
ggml_backend_sched_t sched_ = nullptr;
ggml_backend_sched_t fast_sched_ = nullptr;
ggml_gallocr_t fast_gallocr_ = nullptr;
ggml_context * ctx_kv_ = nullptr;
ggml_backend_buffer_t kv_buf_ = nullptr;
ggml_tensor * memory_k_ = nullptr;
Expand All @@ -152,6 +161,27 @@ class SlowARModel {

std::unordered_set<ggml_tensor *> weight_tensor_set_;

bool fast_decoder_cpu_ = false;

struct FastGraphSlot {
ggml_context * ctx = nullptr;
ggml_cgraph * gf = nullptr;
ggml_tensor * hidden0 = nullptr;
ggml_tensor * prefix_ids = nullptr;
ggml_tensor * positions = nullptr;
ggml_tensor * logits = nullptr;
ggml_tensor * logits_all = nullptr;
ggml_gallocr_t allocr = nullptr;
int32_t n_tokens = 0;
bool valid = false;
};

FastGraphSlot fast_slot_;
std::vector<uint8_t> fast_slot_buf_;
size_t fast_slot_buf_size_ = 0;

std::vector<FastGraphSlot> fast_slots_;

};

}
175 changes: 83 additions & 92 deletions src/s2_generate.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -19,54 +19,50 @@ GenerateResult generate(
out.num_codebooks = model.hparams().num_codebooks;
if (out.num_codebooks <= 0) out.num_codebooks = 1;

const int32_t vocab_size = model.hparams().vocab_size;
const int32_t sem_begin = model.hparams().semantic_begin_id;
const int32_t sem_end = model.hparams().semantic_end_id;
const int32_t vocab_size = model.hparams().vocab_size;
const int32_t sem_begin = model.hparams().semantic_begin_id;
const int32_t sem_end = model.hparams().semantic_end_id;
const int32_t codebook_size = model.hparams().codebook_size;
const int32_t im_end_id = config.im_end_id;
const int32_t num_cb = out.num_codebooks;
const int32_t im_end_id = config.im_end_id;
const int32_t num_cb = out.num_codebooks;

std::vector<float> sem_mask(vocab_size, -std::numeric_limits<float>::infinity());
for (int32_t i = sem_begin; i <= sem_end && i < vocab_size; ++i) {
for (int32_t i = sem_begin; i <= sem_end && i < vocab_size; ++i)
sem_mask[i] = 0.0f;
}
if (im_end_id >= 0 && im_end_id < vocab_size) {
if (im_end_id >= 0 && im_end_id < vocab_size)
sem_mask[im_end_id] = 0.0f;
}

const int32_t rows = prompt.rows;
const int32_t cols = prompt.cols;
std::vector<int32_t> prompt_tm(static_cast<size_t>(rows) * cols);
for (int32_t r = 0; r < rows; ++r) {
for (int32_t c = 0; c < cols; ++c) {
prompt_tm[static_cast<size_t>(c) * rows + r] = prompt.data[static_cast<size_t>(r) * cols + c];
}
}
for (int32_t r = 0; r < rows; ++r)
for (int32_t c = 0; c < cols; ++c)
prompt_tm[static_cast<size_t>(c) * rows + r] =
prompt.data[static_cast<size_t>(r) * cols + c];

StepResult state;
if (params.verbose && log_enabled(LogLevel::Info)) {
if (params.verbose && log_enabled(LogLevel::Info))
std::cout << "[Generate] Prefilling " << prompt.cols << " tokens..." << std::endl;
}

const auto prefill_t0 = std::chrono::steady_clock::now();
if (!model.prefill_fast(prompt_tm, prompt.cols, params.n_threads, state)) {
std::cerr << "[Generate] Prefill failed." << std::endl;
return out;
}
const auto prefill_t1 = std::chrono::steady_clock::now();

SamplerParams sparams;
sparams.temperature = params.temperature;
sparams.top_p = params.top_p;
sparams.top_k = params.top_k;

auto apply_mask_and_sample = [&](const std::vector<float> & logits,
bool block_im_end) -> int32_t {
std::vector<float> biased(vocab_size);
for (int32_t i = 0; i < vocab_size; ++i) {
for (int32_t i = 0; i < vocab_size; ++i)
biased[i] = logits[i] + sem_mask[i];
}
if (block_im_end && im_end_id >= 0 && im_end_id < vocab_size) {
if (block_im_end && im_end_id >= 0 && im_end_id < vocab_size)
biased[im_end_id] = -std::numeric_limits<float>::infinity();
}
SamplerParams sparams;
sparams.temperature = params.temperature;
sparams.top_p = params.top_p;
sparams.top_k = params.top_k;
return sample_token(biased.data(), vocab_size, sparams);
};

Expand All @@ -76,121 +72,118 @@ GenerateResult generate(
out.codes.resize(static_cast<size_t>(num_cb) * params.max_new_tokens, 0);
out.n_frames = 0;

std::vector<float> fast_logits;

SamplerParams sparams;
sparams.temperature = params.temperature;
sparams.top_p = params.top_p;
sparams.top_k = params.top_k;

std::vector<int32_t> ras_window;
int32_t ras_buf[10];
int32_t ras_head = 0;
int32_t ras_count = 0;
const int32_t ras_window_size = 10;
const float ras_high_temp = 1.0f;
const float ras_high_top_p = 0.9f;
const float ras_high_temp = 1.0f;
const float ras_high_top_p = 0.9f;

if (params.verbose && log_enabled(LogLevel::Info)) {
if (params.verbose && log_enabled(LogLevel::Info))
std::cout << "[Generate] Generating (max " << params.max_new_tokens << " tokens)..." << std::endl;
}

int32_t step = 0;
const auto loop_t0 = std::chrono::steady_clock::now();

while (main_token != im_end_id && step < params.max_new_tokens) {

if (!ras_window.empty() &&
std::find(ras_window.begin(), ras_window.end(), main_token) != ras_window.end() &&
main_token >= sem_begin && main_token <= sem_end)
bool ras_hit = false;
if (main_token >= sem_begin && main_token <= sem_end) {
for (int32_t i = 0; i < ras_count; ++i) {
if (ras_buf[i] == main_token) { ras_hit = true; break; }
}
}
if (ras_hit)
{

std::vector<float> biased(vocab_size);
for (int32_t i = 0; i < vocab_size; ++i) {
for (int32_t i = 0; i < vocab_size; ++i)
biased[i] = state.logits[i] + sem_mask[i];
}
if (step < params.min_tokens_before_end && im_end_id >= 0 && im_end_id < vocab_size) {
if (step < params.min_tokens_before_end &&
im_end_id >= 0 && im_end_id < vocab_size)
biased[im_end_id] = -std::numeric_limits<float>::infinity();
}
SamplerParams ras_sparams;
ras_sparams.temperature = ras_high_temp;
ras_sparams.top_p = ras_high_top_p;
ras_sparams.top_k = params.top_k;
main_token = sample_token(biased.data(), vocab_size, ras_sparams);
SamplerParams ras_sp;
ras_sp.temperature = ras_high_temp;
ras_sp.top_p = ras_high_top_p;
ras_sp.top_k = params.top_k;
main_token = sample_token(biased.data(), vocab_size, ras_sp);
}

ras_window.push_back(main_token);
if ((int32_t)ras_window.size() > ras_window_size) {
ras_window.erase(ras_window.begin());
}
ras_buf[ras_head] = main_token;
ras_head = (ras_head + 1) % ras_window_size;
if (ras_count < ras_window_size) ras_count++;

int32_t sem_code = main_token - sem_begin;
if (sem_code < 0) sem_code = 0;
if (sem_code < 0) sem_code = 0;
if (sem_code >= codebook_size) sem_code = codebook_size - 1;

std::vector<int32_t> residual_cbs;
if (!model.fast_decode_batch(state.hidden, sem_code,
params.n_threads, sparams,
residual_cbs))
{
std::cerr << "[Generate] fast_decode_batch failed at step "
<< step << std::endl;
residual_cbs.assign(num_cb - 1, 0);
}

std::vector<int32_t> codebooks_cb;
codebooks_cb.reserve(num_cb);
codebooks_cb.push_back(sem_code);
for (int32_t r : residual_cbs)
codebooks_cb.push_back(r);

for (int32_t cb_idx = 1; cb_idx < num_cb; ++cb_idx) {

if (!model.fast_decode(state.hidden, codebooks_cb, params.n_threads, fast_logits)) {
std::cerr << "[Generate] fast_decode failed at cb " << cb_idx << std::endl;

for (int32_t r = cb_idx; r < num_cb; ++r) {
codebooks_cb.push_back(0);
}
break;
}
int32_t cb_token = sample_token(fast_logits.data(), (int32_t)fast_logits.size(), sparams);
codebooks_cb.push_back(cb_token);
}

for (int32_t cb = 0; cb < num_cb; ++cb) {
out.codes[static_cast<size_t>(cb) * params.max_new_tokens + step] = codebooks_cb[cb];
}
for (int32_t cb = 0; cb < num_cb; ++cb)
out.codes[static_cast<size_t>(cb) * params.max_new_tokens + step] =
codebooks_cb[cb];
out.n_frames++;

if (params.on_frame) {
FrameCallbackData fcd;
fcd.codes = codebooks_cb.data();
fcd.frame_index = step;
fcd.total_frames = out.n_frames;
fcd.codes = codebooks_cb.data();
fcd.frame_index = step;
fcd.total_frames = out.n_frames;
fcd.num_codebooks = num_cb;
if (!params.on_frame(fcd)) {

if (params.verbose && log_enabled(LogLevel::Info)) {
std::cout << "\n[Generate] Aborted by callback at frame " << step << std::endl;
}
if (params.verbose && log_enabled(LogLevel::Info))
std::cout << "\n[Generate] Aborted by callback at frame "
<< step << std::endl;
break;
}
}

std::vector<int32_t> step_input(num_cb + 1);
step_input[0] = main_token;
for (int32_t cb = 0; cb < num_cb; ++cb) {
for (int32_t cb = 0; cb < num_cb; ++cb)
step_input[cb + 1] = codebooks_cb[cb];
}

if (!model.step(step_input, params.n_threads, state)) {
std::cerr << "[Generate] step() failed at step " << step << std::endl;
break;
}

step++;
if (params.verbose && log_enabled(LogLevel::Info) && step % 50 == 0) {
std::cout << "\r[Generate] " << step << " / " << params.max_new_tokens << " tokens..." << std::flush;
}

if (params.verbose && log_enabled(LogLevel::Info) && step % 50 == 0)
std::cout << "\r[Generate] " << step << " / "
<< params.max_new_tokens << " tokens..." << std::flush;

bool block_next_end = (step < params.min_tokens_before_end);
main_token = apply_mask_and_sample(state.logits, block_next_end);
}

if (params.verbose && log_enabled(LogLevel::Info)) {
const auto loop_t1 = std::chrono::steady_clock::now();
const auto loop_t1 = std::chrono::steady_clock::now();
const auto generate_t1 = std::chrono::steady_clock::now();
const double prefill_ms = std::chrono::duration<double, std::milli>(prefill_t1 - prefill_t0).count();
const double loop_ms = std::chrono::duration<double, std::milli>(loop_t1 - loop_t0).count();
const double total_ms = std::chrono::duration<double, std::milli>(generate_t1 - generate_t0).count();
const double prefill_ms = std::chrono::duration<double, std::milli>(
prefill_t1 - prefill_t0).count();
const double loop_ms = std::chrono::duration<double, std::milli>(
loop_t1 - loop_t0).count();
const double total_ms = std::chrono::duration<double, std::milli>(
generate_t1 - generate_t0).count();
const double ms_per_frame = out.n_frames > 0 ? (loop_ms / out.n_frames) : 0.0;
std::cout << std::endl;
std::cout << "[Generate] Done: " << out.n_frames << " frames generated." << std::endl;
std::cout << "[Generate] Done: " << out.n_frames
<< " frames generated." << std::endl;
std::cout << "[Metrics] Generate: prefill=" << prefill_ms
<< " ms, loop=" << loop_ms
<< " ms, total=" << total_ms
Expand All @@ -201,12 +194,10 @@ GenerateResult generate(
const int32_t n_frames = out.n_frames;
if (n_frames < params.max_new_tokens) {
std::vector<int32_t> compacted(static_cast<size_t>(num_cb) * n_frames);
for (int32_t cb = 0; cb < num_cb; ++cb) {
for (int32_t t = 0; t < n_frames; ++t) {
for (int32_t cb = 0; cb < num_cb; ++cb)
for (int32_t t = 0; t < n_frames; ++t)
compacted[static_cast<size_t>(cb) * n_frames + t] =
out.codes[static_cast<size_t>(cb) * params.max_new_tokens + t];
}
}
out.codes = std::move(compacted);
} else {
out.codes.resize(static_cast<size_t>(num_cb) * n_frames);
Expand Down
Loading