Skip to content
Closed
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
2 changes: 2 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -153,3 +153,5 @@ a.out.*

AGENTS.local.md
.pi/SYSTEM.md

/scratch/
1 change: 1 addition & 0 deletions examples/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ llama_add_compile_flags()
if (EMSCRIPTEN)
else()
add_subdirectory(batched)
add_subdirectory(bonsai-logits)
add_subdirectory(debug)
add_subdirectory(embedding)
add_subdirectory(eval-callback)
Expand Down
4 changes: 4 additions & 0 deletions examples/bonsai-logits/CMakeLists.txt
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
set(TARGET llama-bonsai-logits)
add_executable(${TARGET} bonsai-logits.cpp)
target_link_libraries(${TARGET} PRIVATE llama ${CMAKE_THREAD_LIBS_INIT})
target_compile_features(${TARGET} PRIVATE cxx_std_17)
95 changes: 95 additions & 0 deletions examples/bonsai-logits/bonsai-logits.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,95 @@
#include "llama.h"
#include "ggml-backend.h"
#include <algorithm>
#include <climits>
#include <chrono>
#include <cmath>
#include <cstdint>
#include <cstdio>
#include <exception>
#include <cstdlib>
#include <fstream>
#include <iterator>
#include <string>
#include <vector>

static void log_message(ggml_log_level level, const char * text, void *) {
if (level <= GGML_LOG_LEVEL_WARN) std::fputs(text, stderr);
}

int main(int argc, char ** argv) {
if (argc != 8 && argc != 9) {
std::fprintf(stderr, "usage: llama-bonsai-logits model backend-dir corpus output prompt-tokens decode-tokens interval [batch]\n");
return 2;
}
int prompt_count, decode_count, interval, batch;
try {
prompt_count = std::stoi(argv[5]);
decode_count = std::stoi(argv[6]);
interval = std::stoi(argv[7]);
batch = argc == 9 ? std::stoi(argv[8]) : 512;
} catch (const std::exception &) {
std::fprintf(stderr, "token counts, interval, and batch must be integers\n");
return 2;
}
if (batch < 1 || batch > 4096) return 2;
if (prompt_count < 1 || decode_count < 1 || interval < 1 || decode_count % interval ||
prompt_count > INT_MAX - decode_count) return 2;
std::ifstream corpus_file(argv[3], std::ios::binary);
if (!corpus_file) return 3;
const std::string corpus((std::istreambuf_iterator<char>(corpus_file)), std::istreambuf_iterator<char>());
if (corpus.size() > INT_MAX) return 3;
llama_log_set(log_message, nullptr);
ggml_backend_load_all_from_path(argv[2]);
llama_backend_init();
auto mp = llama_model_default_params();
// BONSAI_LOGITS_NGL selects the offload depth; zero also withholds every device so the
// scheduler cannot route prompt matrix products to an accelerator.
const char * ngl_env = std::getenv("BONSAI_LOGITS_NGL");
mp.n_gpu_layers = ngl_env ? std::stoi(ngl_env) : 99;
static ggml_backend_dev_t no_devices[] = { nullptr };
if (mp.n_gpu_layers == 0) { mp.devices = no_devices; }
auto * model = llama_model_load_from_file(argv[1], mp);
if (!model) return 4;
auto cp = llama_context_default_params();
cp.n_ctx = prompt_count + decode_count; cp.n_batch = batch; cp.n_ubatch = batch;
cp.n_threads = 8; cp.n_threads_batch = 8;
cp.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_ENABLED;
auto * ctx = llama_init_from_model(model, cp);
if (!ctx) return 5;
const auto * vocab = llama_model_get_vocab(model);
const int needed = -llama_tokenize(vocab, corpus.data(), int(corpus.size()), nullptr, 0, true, true);
if (needed < prompt_count + decode_count) return 6;
std::vector<llama_token> tokens(needed);
if (llama_tokenize(vocab, corpus.data(), int(corpus.size()), tokens.data(), needed, true, true) != needed) return 6;
auto * out = std::fopen(argv[4], "wb");
if (!out) return 7;
const int32_t nv = llama_vocab_n_tokens(vocab), steps = 1 + decode_count / interval;
// The output contains two int32 dimensions followed by full FP32 vocabulary rows.
if (std::fwrite(&nv, sizeof(nv), 1, out) != 1 || std::fwrite(&steps, sizeof(steps), 1, out) != 1) return 9;
auto dump = [&]() {
const float * logits = llama_get_logits(ctx);
for (int i = 0; i < nv; ++i) if (!std::isfinite(logits[i])) return false;
return std::fwrite(logits, sizeof(float), nv, out) == size_t(nv);
};
std::fprintf(stderr, "prefill batch=%d\n", batch);
const auto begin = std::chrono::steady_clock::now();
for (int offset = 0; offset < prompt_count; offset += batch) {
const int count = std::min(batch, prompt_count - offset);
if (llama_decode(ctx, llama_batch_get_one(tokens.data() + offset, count))) return 8;
}
if (!dump()) return 9;
for (int i = 0; i < decode_count; ++i) {
if (llama_decode(ctx, llama_batch_get_one(tokens.data() + prompt_count + i, 1))) return 8;
if ((i + 1) % interval == 0) {
if (!dump()) return 9;
std::fprintf(stderr, "checked position=%d elapsed=%.3f seconds\n", prompt_count + i + 1,
std::chrono::duration<double>(std::chrono::steady_clock::now() - begin).count());
std::fflush(stderr);
}
}
if (std::fclose(out)) return 10;
std::fprintf(stderr, "complete prompt=%d decode=%d positions=%d vocabulary=%d\n", prompt_count, decode_count, steps, nv);
llama_free(ctx); llama_model_free(model); llama_backend_free();
return 0;
}
8 changes: 8 additions & 0 deletions ggml/src/ggml-sycl/fattn.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -247,6 +247,14 @@ static best_fattn_kernel ggml_sycl_get_best_fattn_kernel(const int device, const
const bool can_use_vector_kernel = Q->ne[0] <= 512 && Q->ne[0] % 64 == 0 && K->ne[1] % FATTN_KQ_STRIDE == 0
&& !has_bf16;

// Serve single-query decode with the vector kernel on request. The grouped-query rule below assumes a
// matrix-engine kernel exists for those shapes, and this backend has none.
static const int decode_vec = ggml_sycl_get_env("GGML_SYCL_FA_DECODE_VEC", 0);
if (decode_vec && Q->ne[1] == 1 && can_use_vector_kernel &&
!ggml_is_quantized(K->type) && !ggml_is_quantized(V->type)) {
return BEST_FATTN_KERNEL_VEC;
}

// Fused-XMX path: oneDNN Graph SDPA (flash attention). Strictly
// additive -- taken only when statically supported, otherwise falls through to VEC/TILE below.
if (ggml_sycl_flash_attn_ext_onednn_supported(dst)) {
Expand Down
31 changes: 31 additions & 0 deletions ggml/src/ggml-sycl/getrows.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -214,6 +214,37 @@ static void get_rows_sycl_float(ggml_backend_sycl_context & ctx, const ggml_tens

GGML_TENSOR_BINARY_OP_LOCALS

static const int single_row_copy = ggml_sycl_get_env("GGML_SYCL_SINGLE_ROW_COPY", 0);
if constexpr (std::is_same_v<src0_t, float> && std::is_same_v<dst_t, float>) {
if (single_row_copy && ne01 == 1 && ne02 == 1 && ne03 == 1 &&
ggml_nelements(src1) == 1 && ggml_is_contiguous(src0) && ggml_is_contiguous(dst)) {
// The only valid row index is zero, so src1 is never read.
if (dst_dd == src0_dd) {
return;
}
if (single_row_copy == 2) {
const int64_t vectors = (ne00 + 3) / 4;
const sycl::range<1> local(SYCL_GET_ROWS_BLOCK_SIZE);
const sycl::range<1> global(((vectors + local[0] - 1) / local[0]) * local[0]);
stream->parallel_for(sycl::nd_range<1>(global, local), [=](sycl::nd_item<1> item) {
const int64_t i = item.get_global_id(0);
if (4 * i + 3 < ne00) {
sycl::vec<src0_t, 4> values;
values.load(i, src0_dd);
values.store(i, dst_dd);
} else {
for (int64_t j = 4 * i; j < ne00; ++j) {
dst_dd[j] = src0_dd[j];
}
}
});
} else {
stream->memcpy(dst_dd, src0_dd, ne00 * sizeof(src0_t));
}
return;
}
}

const sycl::range<3> block_dims(1, 1, SYCL_GET_ROWS_BLOCK_SIZE);
const int block_num_x = (ne00 + SYCL_GET_ROWS_BLOCK_SIZE - 1) / SYCL_GET_ROWS_BLOCK_SIZE;
const sycl::range<3> block_nums(ne11 * ne12, ne10, block_num_x);
Expand Down
7 changes: 6 additions & 1 deletion ggml/src/ggml-sycl/ggml-sycl.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2682,8 +2682,13 @@ inline void ggml_sycl_op_mul_mat_sycl(
#ifdef GGML_SYCL_F16
bool use_fp16 = true; // TODO(Yu) SYCL capability check
#else
bool use_fp16 = false;
bool use_fp16 = src0->type == GGML_TYPE_PQ2_0 || src0->type == GGML_TYPE_PTQ1_0;
#endif
const bool bonsai = src0->type == GGML_TYPE_PQ2_0 || src0->type == GGML_TYPE_PTQ1_0;
static const int bonsai_fp16 = ggml_sycl_get_env("GGML_SYCL_BONSAI_F16", 1);
if (bonsai && !bonsai_fp16) {
use_fp16 = false;
}

#if GGML_SYCL_DNNL && defined(GGML_SYCL_HAS_BF16)
// Fast path for bf16 src0
Expand Down
Loading