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
12 changes: 12 additions & 0 deletions common/arg.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2457,6 +2457,18 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
params.kv_mean_center_path = value;
}
).set_env("LLAMA_ARG_KV_MEAN_CENTER"));
add_opt(common_arg(
{"--kv-vram-cells"}, "N",
"tiered KV cache: keep the first N cells of each layer's K/V in VRAM and the rest in pinned system RAM\n"
"mapped into the same device range (CUDA VMM), so the context can exceed VRAM. Positions past N are\n"
"read over PCIe once a sequence is that deep; output is identical to an all-VRAM cache. 0 = off (default)",
[](common_params & params, int value) {
if (value < 0) {
throw std::invalid_argument("invalid value");
}
params.n_kv_vram_cells = value;
}
).set_env("LLAMA_ARG_KV_VRAM_CELLS"));
add_opt(common_arg(
{"--hellaswag"},
"compute HellaSwag score over random tasks from datafile supplied with -f",
Expand Down
1 change: 1 addition & 0 deletions common/common.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1764,6 +1764,7 @@ struct llama_context_params common_context_params_to_llama(const common_params &
// note: params (and therefore params.kv_mean_center_path) is kept alive by the caller for
// at least as long as it takes to call llama_init_from_model() with the returned cparams
cparams.path_kv_mean_center = params.kv_mean_center_path.empty() ? nullptr : params.kv_mean_center_path.c_str();
cparams.n_kv_vram_cells = (uint32_t) std::max(0, params.n_kv_vram_cells);

return cparams;
}
Expand Down
3 changes: 3 additions & 0 deletions common/common.h
Original file line number Diff line number Diff line change
Expand Up @@ -588,6 +588,9 @@ struct common_params {
// only takes effect when cache_type_k == GGML_TYPE_Q4_0; see docs/kv-mean-center.md
std::string kv_mean_center_path = "";

// tiered KV cache: cells past this many live in pinned host memory (0 = all in device memory)
int32_t n_kv_vram_cells = 0;

common_conversation_mode conversation_mode = COMMON_CONVERSATION_MODE_AUTO;

// multimodal models (see tools/mtmd)
Expand Down
4 changes: 4 additions & 0 deletions ggml/include/ggml-cuda.h
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,10 @@ GGML_BACKEND_API bool ggml_backend_is_cuda(ggml_backend_t backend);
// device buffer
GGML_BACKEND_API ggml_backend_buffer_type_t ggml_backend_cuda_buffer_type(int device);

// device buffer whose buffers keep vram_frac of each of n_parts equal parts in VRAM and the remainder in
// pinned host memory mapped into the same device address range (CUDA VMM). nullptr without VMM.
GGML_BACKEND_API ggml_backend_buffer_type_t ggml_backend_cuda_tier_buffer_type(int device, double vram_frac, int n_parts, const char * tag);

// conduct allreduce operation between devices
GGML_BACKEND_API bool ggml_backend_cuda_allreduce_tensor(ggml_backend_t * backends, struct ggml_tensor ** tensors, size_t n_backends);

Expand Down
30 changes: 30 additions & 0 deletions ggml/src/ggml-cuda/fattn.cu
Original file line number Diff line number Diff line change
Expand Up @@ -595,8 +595,31 @@ size_t ggml_cuda_flash_attn_ext_get_alloc_size(int device, const ggml_tensor * d
return f16_extra.end - (uintptr_t) dst->data;
}

// ggml-cuda.cu: host-tail staging of tiered buffers (ggml_backend_cuda_tier_buffer_type)
void * ggml_cuda_tier_stage(const void * ptr, size_t nbytes, cudaStream_t stream);

void ggml_cuda_flash_attn_ext(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
ggml_cuda_set_device(ctx.device);

// K/V in a tiered buffer whose host tail this op reaches: copy the used host rows into the VRAM staging
// buffer with the copy engine and point K/V at the all-VRAM alias of the same range for this op. Prefill
// kernels read each K/V row once per query tile, which for host rows would be PCIe traffic every time;
// decode reads each row once, but DMA moves it faster than SMs reading host memory from inside the
// kernel (RTX 4070, 180k context, 86k cells in host memory: +28% decode). Nothing is copied for ops
// that stay below the tier line. The staged bytes are the same bytes, so results are unchanged.
ggml_tensor * K = dst->src[1];
ggml_tensor * V = dst->src[2];
void * K_data = K ? K->data : nullptr;
void * V_data = V ? V->data : nullptr;
if (K && V && V != K) {
if (void * a = ggml_cuda_tier_stage(K->data, ggml_nbytes(K), ctx.stream())) {
K->data = a;
}
if (void * a = ggml_cuda_tier_stage(V->data, ggml_nbytes(V), ctx.stream())) {
V->data = a;
}
}

switch (ggml_cuda_get_best_fattn_kernel(ggml_cuda_get_device(), dst)) {
case BEST_FATTN_KERNEL_NONE:
GGML_ABORT("fatal error");
Expand All @@ -610,6 +633,13 @@ void ggml_cuda_flash_attn_ext(ggml_backend_cuda_context & ctx, ggml_tensor * dst
ggml_cuda_flash_attn_ext_mma_f16(ctx, dst);
break;
}

if (K) {
K->data = K_data;
}
if (V) {
V->data = V_data;
}
}

bool ggml_cuda_flash_attn_ext_supported(int device, const ggml_tensor * dst) {
Expand Down
Loading
Loading