Skip to content
Merged
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
53 changes: 33 additions & 20 deletions src/arch/canary/model.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -289,7 +289,7 @@ transcribe_status fuse_batch_norm(CanaryModel & m) {
ggml_init_params params = { ctx_size, nullptr, true };
m.bn_fused_ctx = ggml_init(params);
if (m.bn_fused_ctx == nullptr) {
return TRANSCRIBE_ERR_BACKEND;
return TRANSCRIBE_ERR_OOM;
}

for (size_t i = 0; i < n_blocks; ++i) {
Expand All @@ -300,7 +300,7 @@ transcribe_status fuse_batch_norm(CanaryModel & m) {

m.bn_fused_buffer = ggml_backend_alloc_ctx_tensors(m.bn_fused_ctx, m.plan.scheduler_list.back());
if (m.bn_fused_buffer == nullptr) {
return TRANSCRIBE_ERR_BACKEND;
return TRANSCRIBE_ERR_OOM;
}

std::vector<float> bn_w(d), bn_b(d), rm(d), rv(d);
Expand Down Expand Up @@ -498,7 +498,7 @@ transcribe_status load(Loader & loader, const transcribe_model_load_params * par
if (weights_buffer == nullptr) {
gguf_free(gguf_data);
log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "canary: ggml_backend_alloc_ctx_tensors failed");
return TRANSCRIBE_ERR_GGUF;
return TRANSCRIBE_ERR_OOM;
}
m->backend_buffer = weights_buffer;
ggml_backend_buffer_set_usage(weights_buffer, GGML_BACKEND_BUFFER_USAGE_WEIGHTS);
Expand Down Expand Up @@ -779,7 +779,7 @@ transcribe_status run(transcribe_session * session,
static_cast<int>(cm->plan.scheduler_list.size()), 16384, false, true);
if (cc->sched == nullptr) {
log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "canary run: ggml_backend_sched_new failed");
return TRANSCRIBE_ERR_GGUF;
return TRANSCRIBE_ERR_BACKEND;
}
}
ggml_backend_sched_reset(cc->sched);
Expand Down Expand Up @@ -825,7 +825,7 @@ transcribe_status run(transcribe_session * session,
const int64_t t_enc_start = ggml_time_us();
if (const ggml_status gs = ggml_backend_sched_graph_compute(cc->sched, eb.graph); gs != GGML_STATUS_SUCCESS) {
log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "canary run: encoder compute failed (%d)", static_cast<int>(gs));
return TRANSCRIBE_ERR_GGUF;
return TRANSCRIBE_ERR_BACKEND;
}
cc->t_encode_us = ggml_time_us() - t_enc_start;

Expand Down Expand Up @@ -1015,7 +1015,7 @@ transcribe_status run(transcribe_session * session,
if (const ggml_status gs = ggml_backend_sched_graph_compute(cc->sched, cross_db.graph);
gs != GGML_STATUS_SUCCESS) {
log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "canary run: cross_kv compute failed (%d)", static_cast<int>(gs));
return TRANSCRIBE_ERR_GGUF;
return TRANSCRIBE_ERR_BACKEND;
}
cc->kv_cache.cross_populated = true;
}
Expand Down Expand Up @@ -1072,7 +1072,7 @@ transcribe_status run(transcribe_session * session,

if (const ggml_status gs = ggml_backend_sched_graph_compute(cc->sched, db.graph); gs != GGML_STATUS_SUCCESS) {
log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "canary run: decoder prompt compute failed (%d)", static_cast<int>(gs));
return TRANSCRIBE_ERR_GGUF;
return TRANSCRIBE_ERR_BACKEND;
}

// Per-sublayer dumps at layers {0, n_layers/2, n_layers-1}.
Expand Down Expand Up @@ -1264,7 +1264,10 @@ transcribe_status run(transcribe_session * session,

if (const ggml_status gs = ggml_backend_sched_graph_compute(cc->sched, sb.graph);
gs != GGML_STATUS_SUCCESS) {
break;
log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "canary run: decoder step compute failed (%d)",
static_cast<int>(gs));
commit_result();
return TRANSCRIBE_ERR_BACKEND;
}

n_past += 1;
Expand Down Expand Up @@ -1304,19 +1307,26 @@ transcribe_status run(transcribe_session * session,
}

if (!new_compute_ctx(4 * 1024 * 1024)) {
break;
log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "canary run: decoder step ggml_init failed");
commit_result();
return TRANSCRIBE_ERR_OOM;
}

DecoderBuild db_step = build_decoder_graph_kv(cc->compute_ctx, cm->weights, cm->hparams, cc->kv_cache,
/*n_tokens=*/1, n_past, T_enc,
/*skip_log_softmax=*/true, cc->decoder_use_flash);
if (db_step.out == nullptr || db_step.graph == nullptr) {
break;
log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "canary run: decoder step graph build failed");
commit_result();
return TRANSCRIBE_ERR_GGUF;
}

ggml_backend_sched_reset(cc->sched);
if (!ggml_backend_sched_alloc_graph(cc->sched, db_step.graph)) {
break;
transcribe::log_msg(TRANSCRIBE_LOG_LEVEL_ERROR,
"canary run: step graph allocation failed — out of memory.");
commit_result();
return TRANSCRIBE_ERR_OOM;
}

int32_t token_id = next_token;
Expand All @@ -1326,7 +1336,10 @@ transcribe_status run(transcribe_session * session,

if (const ggml_status gs = ggml_backend_sched_graph_compute(cc->sched, db_step.graph);
gs != GGML_STATUS_SUCCESS) {
break;
log_msg(TRANSCRIBE_LOG_LEVEL_ERROR, "canary run: decoder step compute failed (%d)",
static_cast<int>(gs));
commit_result();
return TRANSCRIBE_ERR_BACKEND;
}

n_past += 1;
Expand Down Expand Up @@ -1418,7 +1431,7 @@ transcribe_status encode_one_to_host(CanarySession * cc,
cc->sched = ggml_backend_sched_new(cm->plan.scheduler_list.data(), nullptr,
static_cast<int>(cm->plan.scheduler_list.size()), 16384, false, true);
if (cc->sched == nullptr) {
return TRANSCRIBE_ERR_GGUF;
return TRANSCRIBE_ERR_BACKEND;
}
}
ggml_backend_sched_reset(cc->sched);
Expand Down Expand Up @@ -1456,7 +1469,7 @@ transcribe_status encode_one_to_host(CanarySession * cc,

const int64_t t0 = ggml_time_us();
if (ggml_backend_sched_graph_compute(cc->sched, eb.graph) != GGML_STATUS_SUCCESS) {
return TRANSCRIBE_ERR_GGUF;
return TRANSCRIBE_ERR_BACKEND;
}
enc_us += ggml_time_us() - t0;

Expand Down Expand Up @@ -1701,7 +1714,7 @@ transcribe_status run_batch(transcribe_session * session,
}
ggml_backend_tensor_set(cross.encoder_out_in, packed.data(), 0, packed.size() * sizeof(float));
if (ggml_backend_sched_graph_compute(cc->sched, cross.graph) != GGML_STATUS_SUCCESS) {
return TRANSCRIBE_ERR_GGUF;
return TRANSCRIBE_ERR_BACKEND;
}
cc->kv_cache.cross_populated = true;
}
Expand Down Expand Up @@ -1730,18 +1743,18 @@ transcribe_status run_batch(transcribe_session * session,
}

StepBuildBatched sb{};
auto rebuild = [&](int win, transcribe::EncDecStepIO & io) -> bool {
auto rebuild = [&](int win, transcribe::EncDecStepIO & io) -> transcribe_status {
if (!new_compute_ctx(16 * 1024 * 1024)) {
return false;
return TRANSCRIBE_ERR_OOM;
}
sb = build_step_graph_batched(cc->compute_ctx, cm->weights, hp, cc->kv_cache, win, T_enc_max, n,
cc->decoder_use_flash);
if (sb.graph == nullptr || sb.argmax_out == nullptr) {
return false;
return TRANSCRIBE_ERR_GGUF;
}
ggml_backend_sched_reset(cc->sched);
if (!ggml_backend_sched_alloc_graph(cc->sched, sb.graph)) {
return false;
return TRANSCRIBE_ERR_OOM;
}
ggml_backend_tensor_set(sb.cross_mask_in, cmask.data(), 0, cmask.size() * sizeof(ggml_fp16_t));
io.token_ids = sb.token_ids_in;
Expand All @@ -1750,7 +1763,7 @@ transcribe_status run_batch(transcribe_session * session,
io.self_mask = sb.self_mask_in;
io.argmax = sb.argmax_out;
io.graph = sb.graph;
return true;
return TRANSCRIBE_OK;
};

std::vector<std::vector<int32_t>> generated(n);
Expand Down
Loading
Loading