2025-08-21 17:00:33 +03:00
#include "llama-kv-cache.h"
2025-06-01 11:39:27 +03:00
#include "llama-impl.h"
2025-06-04 18:58:20 +03:00
#include "llama-io.h"
2025-06-01 11:39:27 +03:00
#include "llama-model.h"
#include "llama-context.h"
#include <algorithm>
#include <cassert>
#include <cmath>
2025-10-28 11:23:54 +01:00
#include <cstring>
2025-06-01 11:39:27 +03:00
#include <limits>
#include <map>
#include <stdexcept>
2026-04-01 16:58:01 +03:00
static bool ggml_is_power_of_2 ( int n ) {
return ( n & ( n - 1 )) == 0 ;
}
// orthonormal Walsh-Hadamard rotation matrix
// note: res^2 == I
static void ggml_gen_hadamard ( ggml_tensor * tensor ) {
assert ( tensor -> type == GGML_TYPE_F32 );
const int n = tensor -> ne [ 0 ];
assert ( ggml_is_power_of_2 ( n ));
assert ( tensor -> ne [ 1 ] == n );
assert ( tensor -> ne [ 2 ] == 1 );
assert ( tensor -> ne [ 3 ] == 1 );
std :: vector < float > data_f32 ;
float * data = ( float * ) tensor -> data ;
if ( tensor -> type != GGML_TYPE_F32 ) {
data_f32 . resize ( n * n );
data = data_f32 . data ();
}
data [ 0 * n + 0 ] = 1.0 / sqrtf ( n );
for ( int s = 1 ; s < n ; s *= 2 ) {
for ( int i = 0 ; i < s ; i ++ ) {
for ( int j = 0 ; j < s ; j ++ ) {
const float val = data [ i * n + j ];
data [( i + s ) * n + ( j )] = val ;
data [( i ) * n + ( j + s )] = val ;
data [( i + s ) * n + ( j + s )] = - val ;
}
}
}
if ( tensor -> type != GGML_TYPE_F32 ) {
ggml_quantize_chunk ( tensor -> type , data , tensor -> data , 0 , 1 , n * n , nullptr );
}
}
2025-06-01 11:39:27 +03:00
//
2025-08-21 17:00:33 +03:00
// llama_kv_cache
2025-06-01 11:39:27 +03:00
//
2025-08-21 17:00:33 +03:00
llama_kv_cache :: llama_kv_cache (
2025-08-24 13:07:07 +03:00
const llama_model & model ,
2026-05-29 10:15:17 +02:00
const llama_hparams & hparams ,
2025-08-24 13:07:07 +03:00
ggml_type type_k ,
ggml_type type_v ,
bool v_trans ,
bool offload ,
bool unified ,
uint32_t kv_size ,
uint32_t n_seq_max ,
uint32_t n_pad ,
uint32_t n_swa ,
2025-09-05 10:39:22 +03:00
llama_swa_type swa_type ,
2026-06-07 20:50:54 +08:00
llama_memory_t mem_other ,
2025-08-24 13:07:07 +03:00
const layer_filter_cb & filter ,
2026-06-07 20:50:54 +08:00
const layer_reuse_cb & reuse ,
const layer_share_cb & share ) :
2026-05-29 10:15:17 +02:00
model ( model ), hparams ( hparams ), v_trans ( v_trans ),
2026-06-07 21:42:54 +03:00
n_seq_max ( n_seq_max ), n_stream ( unified ? 1 : n_seq_max ), n_pad ( n_pad ), n_swa ( n_swa ), swa_type ( swa_type ),
other ( static_cast < llama_kv_cache *> ( mem_other )),
v_cells_impl ( other ? other -> v_cells_impl : std :: make_shared < llama_kv_cells_vec > ()),
v_cells ( * v_cells_impl ) {
2025-06-01 11:39:27 +03:00
2026-06-07 17:33:00 +02:00
// shared cells view the source cache's K/V tensors, so the cell count
// follows the source allocation: a fitted target can be smaller than the
// draft default and oversized views would overflow the source tensors
2026-06-07 21:42:54 +03:00
if ( other ) {
const uint32_t size_other = other -> get_size ();
2026-06-07 17:33:00 +02:00
if ( kv_size != size_other ) {
LLAMA_LOG_WARN ( "%s: kv_size = %u overridden to %u to match the shared source cache \n " , __func__ , kv_size , size_other );
kv_size = size_other ;
}
}
2025-06-01 11:39:27 +03:00
GGML_ASSERT ( kv_size % n_pad == 0 );
2026-06-05 11:09:36 +03:00
const uint32_t n_layer = hparams . n_layer_all ;
2025-06-26 19:34:02 +02:00
2025-10-28 11:23:54 +01:00
// define a comparator for the buft -> ctx map to ensure that the order is well-defined:
struct ggml_backend_buft_comparator {
bool operator ()( const ggml_backend_buffer_type_t & lhs , const ggml_backend_buffer_type_t & rhs ) const {
return strcmp ( ggml_backend_buft_name ( lhs ), ggml_backend_buft_name ( rhs )) < 0 ;
}
};
std :: map < ggml_backend_buffer_type_t , ggml_context_ptr , ggml_backend_buft_comparator > ctx_map ;
2025-06-01 11:39:27 +03:00
// create a context for each buffer type
auto ctx_for_buft = [ & ]( ggml_backend_buffer_type_t buft ) -> ggml_context * {
auto it = ctx_map . find ( buft );
if ( it == ctx_map . end ()) {
ggml_init_params params = {
2026-07-26 19:43:45 +02:00
/*.mem_size =*/ size_t ( 3u * ( 1 + n_stream ) * n_layer * ggml_tensor_overhead ()), //Reserve tensor metadata for up to 3 tensors per layer (K, V, and optional K_idx), plus one view per tensor per stream.
2025-06-01 11:39:27 +03:00
/*.mem_buffer =*/ NULL ,
/*.no_alloc =*/ true ,
};
ggml_context * ctx = ggml_init ( params );
if ( ! ctx ) {
return nullptr ;
}
2025-10-28 11:23:54 +01:00
ctx_map . emplace ( buft , ctx );
2025-06-01 11:39:27 +03:00
return ctx ;
}
2025-10-28 11:23:54 +01:00
return it -> second . get ();
2025-06-01 11:39:27 +03:00
};
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( n_stream == 1 || n_stream == n_seq_max );
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
v_heads . resize ( n_stream );
for ( uint32_t s = 0 ; s < n_stream ; ++ s ) {
v_heads [ s ] = 0 ;
}
v_cells . resize ( n_stream );
for ( uint32_t s = 0 ; s < n_stream ; ++ s ) {
v_cells [ s ]. resize ( kv_size );
}
// by default, all sequence ids are mapped to the 0th stream
seq_to_stream . resize ( LLAMA_MAX_SEQ , 0 );
if ( n_stream > 1 ) {
seq_to_stream . resize ( n_stream , 0 );
for ( uint32_t s = 0 ; s < n_stream ; ++ s ) {
seq_to_stream [ s ] = s ;
}
}
// [TAG_V_CACHE_VARIABLE]
if ( v_trans && hparams . is_n_embd_v_gqa_variable ()) {
LLAMA_LOG_WARN ( "%s: the V embeddings have different sizes across layers and FA is not enabled - padding V cache to %d \n " ,
__func__ , hparams . n_embd_v_gqa_max ());
}
2025-06-01 11:39:27 +03:00
2026-01-25 15:48:56 +02:00
const bool is_mla = hparams . is_mla ();
2026-06-05 11:09:36 +03:00
for ( uint32_t il = 0 ; il < n_layer ; il ++ ) {
2025-08-24 13:07:07 +03:00
if ( ! hparams . has_kv ( il )) {
LLAMA_LOG_DEBUG ( "%s: layer %3d: does not have KV cache \n " , __func__ , il );
continue ;
}
2025-06-01 11:39:27 +03:00
if ( filter && ! filter ( il )) {
2025-08-24 13:07:07 +03:00
LLAMA_LOG_DEBUG ( "%s: layer %3d: filtered \n " , __func__ , il );
2025-06-01 11:39:27 +03:00
continue ;
}
2026-06-07 20:50:54 +08:00
if ( share && other ) {
const int32_t il_share = share ( il );
if ( il_share >= 0 ) {
const auto & layer_share = other -> layers [ other -> map_layer_ids [ il_share ]];
LLAMA_LOG_WARN ( "%s: layer %3d: sharing with layer %d. k = %p, v = %p \n " , __func__ , il , il_share ,
layer_share . k -> data , layer_share . v -> data );
map_layer_ids [ il ] = layers . size ();
layers . push_back ( layer_share );
layers . back (). il = il ;
continue ;
}
}
2026-04-07 20:31:28 +03:00
if ( n_embd_head_k_all == 0 ) {
n_embd_head_k_all = ( int32_t ) hparams . n_embd_head_k ( il );
} else if ( n_embd_head_k_all > 0 && n_embd_head_k_all != ( int32_t ) hparams . n_embd_head_k ( il )) {
n_embd_head_k_all = - 1 ;
}
2026-06-29 16:58:51 +08:00
if ( ! is_mla ) {
if ( n_embd_head_v_all == 0 ) {
n_embd_head_v_all = ( int32_t ) hparams . n_embd_head_v ( il );
} else if ( n_embd_head_v_all > 0 && n_embd_head_v_all != ( int32_t ) hparams . n_embd_head_v ( il )) {
n_embd_head_v_all = - 1 ;
}
2026-04-07 20:31:28 +03:00
}
2025-07-16 16:35:42 +03:00
// [TAG_V_CACHE_VARIABLE]
const uint32_t n_embd_k_gqa = hparams . n_embd_k_gqa ( il );
const uint32_t n_embd_v_gqa = ! v_trans ? hparams . n_embd_v_gqa ( il ) : hparams . n_embd_v_gqa_max ();
2025-06-01 11:39:27 +03:00
const char * dev_name = "CPU" ;
ggml_backend_buffer_type_t buft = ggml_backend_cpu_buffer_type ();
if ( offload ) {
auto * dev = model . dev_layer ( il );
buft = ggml_backend_dev_buffer_type ( dev );
dev_name = ggml_backend_dev_name ( dev );
}
LLAMA_LOG_DEBUG ( "%s: layer %3d: dev = %s \n " , __func__ , il , dev_name );
ggml_context * ctx = ctx_for_buft ( buft );
if ( ! ctx ) {
throw std :: runtime_error ( "failed to create ggml context for kv cache" );
}
2026-01-25 15:48:56 +02:00
const bool has_k = true ;
const bool has_v = ! is_mla ;
2025-06-01 11:39:27 +03:00
2026-01-25 15:48:56 +02:00
ggml_tensor * k = has_k ? ggml_new_tensor_3d ( ctx , type_k , n_embd_k_gqa , kv_size , n_stream ) : nullptr ;
ggml_tensor * v = has_v ? ggml_new_tensor_3d ( ctx , type_v , n_embd_v_gqa , kv_size , n_stream ) : nullptr ;
has_k && ggml_format_name ( k , "cache_k_l%d" , il );
has_v && ggml_format_name ( v , "cache_v_l%d" , il );
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
std :: vector < ggml_tensor *> k_stream ;
std :: vector < ggml_tensor *> v_stream ;
for ( uint32_t s = 0 ; s < n_stream ; ++ s ) {
2026-01-25 15:48:56 +02:00
k_stream . push_back ( has_k ? ggml_view_2d ( ctx , k , n_embd_k_gqa , kv_size , k -> nb [ 1 ], s * k -> nb [ 2 ]) : nullptr );
v_stream . push_back ( has_v ? ggml_view_2d ( ctx , v , n_embd_v_gqa , kv_size , v -> nb [ 1 ], s * v -> nb [ 2 ]) : nullptr );
2025-07-16 16:35:42 +03:00
}
2026-07-26 19:43:45 +02:00
const uint32_t n_embd_k_idx = hparams . n_embd_k_idx ( il );
ggml_tensor * k_idx = n_embd_k_idx > 0
? ggml_new_tensor_3d ( ctx , GGML_TYPE_F32 , n_embd_k_idx , kv_size , n_stream )
: nullptr ;
if ( k_idx ) {
ggml_format_name ( k_idx , "cache_k_idx_l%d" , il );
msa_strict_slots = ( n_stream == n_seq_max );
}
std :: vector < ggml_tensor *> k_idx_stream ;
for ( uint32_t s = 0 ; s < n_stream ; ++ s ) {
k_idx_stream . push_back ( k_idx
? ggml_view_2d ( ctx , k_idx , n_embd_k_idx , kv_size , k_idx -> nb [ 1 ], s * k_idx -> nb [ 2 ])
: nullptr );
}
2025-06-01 11:39:27 +03:00
map_layer_ids [ il ] = layers . size ();
2025-07-16 16:35:42 +03:00
2026-07-26 19:43:45 +02:00
layers . push_back ({ il , k , v , k_idx , k_stream , v_stream , k_idx_stream });
2025-06-01 11:39:27 +03:00
}
2025-08-24 13:07:07 +03:00
if ( reuse ) {
LLAMA_LOG_DEBUG ( "%s: reusing layers: \n " , __func__ );
2025-06-26 19:34:02 +02:00
2026-06-05 11:09:36 +03:00
for ( uint32_t il = 0 ; il < n_layer ; il ++ ) {
2025-08-24 13:07:07 +03:00
const int32_t il_reuse = reuse ( il );
if ( il_reuse < 0 ) {
LLAMA_LOG_DEBUG ( "%s: - layer %3d: no reuse \n " , __func__ , il );
2025-06-26 19:34:02 +02:00
continue ;
}
2025-08-24 13:07:07 +03:00
if ( filter && ! filter ( il )) {
LLAMA_LOG_DEBUG ( "%s: - layer %3d: filtered \n " , __func__ , il );
continue ;
}
2025-06-26 19:34:02 +02:00
GGML_ASSERT ( map_layer_ids . find ( il_reuse ) != map_layer_ids . end ());
2025-08-24 13:07:07 +03:00
2025-06-26 19:34:02 +02:00
map_layer_ids [ il ] = map_layer_ids [ il_reuse ];
2025-08-24 13:07:07 +03:00
LLAMA_LOG_DEBUG ( "%s: - layer %3d: reuse layer %d, is_swa = %d \n " , __func__ , il , il_reuse , hparams . is_swa ( il ));
2025-06-26 19:34:02 +02:00
}
}
2025-06-01 11:39:27 +03:00
// allocate tensors and initialize the buffers to avoid NaNs in the padding
2025-10-28 11:23:54 +01:00
for ( auto & [ buft , ctx ] : ctx_map ) {
2025-12-15 09:24:59 +01:00
ggml_backend_buffer_t buf ;
2026-05-29 10:15:17 +02:00
if ( hparams . no_alloc ) {
2025-12-15 09:24:59 +01:00
buf = ggml_backend_buft_alloc_buffer ( buft , /*size =*/ 0 ); // dummy buffer
for ( ggml_tensor * t = ggml_get_first_tensor ( ctx . get ()); t != nullptr ; t = ggml_get_next_tensor ( ctx . get (), t )) {
t -> buffer = buf ; // set dummy buffer for KV cache so that the backend scheduler won't try to allocate it
}
} else {
buf = ggml_backend_alloc_ctx_tensors_from_buft ( ctx . get (), buft ); // real buffer
}
2025-06-01 11:39:27 +03:00
if ( ! buf ) {
throw std :: runtime_error ( "failed to allocate buffer for kv cache" );
}
LLAMA_LOG_INFO ( "%s: %10s KV buffer size = %8.2f MiB \n " , __func__ , ggml_backend_buffer_name ( buf ), ggml_backend_buffer_get_size ( buf ) / 1024.0 / 1024.0 );
ggml_backend_buffer_clear ( buf , 0 );
2025-10-28 11:23:54 +01:00
ctxs_bufs . emplace_back ( std :: move ( ctx ), buf );
2025-06-01 11:39:27 +03:00
}
{
2026-07-26 19:43:45 +02:00
const size_t memory_size_k = size_k_bytes ();
const size_t memory_size_v = size_v_bytes ();
const size_t memory_size_k_idx = size_k_idx_bytes ();
const size_t memory_size_total = memory_size_k + memory_size_v + memory_size_k_idx ;
2025-06-01 11:39:27 +03:00
2026-07-26 19:43:45 +02:00
constexpr float mib = 1024.0f * 1024.0f ;
const std :: string k_log = format ( ", K (%s): %7.2f MiB" , ggml_type_name ( type_k ), ( float ) memory_size_k / mib );
const std :: string v_log = format ( ", V (%s): %7.2f MiB" , ggml_type_name ( type_v ), ( float ) memory_size_v / mib );
std :: string k_idx_log ;
if ( memory_size_k_idx > 0 ) {
k_idx_log = format ( ", K_idx (%s): %7.2f MiB" , ggml_type_name ( GGML_TYPE_F32 ), ( float ) memory_size_k_idx / mib );
}
LLAMA_LOG_INFO ( "%s: size = %7.2f MiB (%6u cells, %3d layers, %2u/%u seqs)%s%s%s \n " , __func__ ,
( float ) memory_size_total / mib , kv_size , ( int ) layers . size (), n_seq_max , n_stream ,
k_log . c_str (), v_log . c_str (), k_idx_log . c_str ());
2025-06-01 11:39:27 +03:00
}
2025-06-11 12:52:45 +03:00
2026-06-07 20:50:54 +08:00
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
if ( other ) {
n_embd_head_k_all = other -> n_embd_head_k_all ;
n_embd_head_v_all = other -> n_embd_head_v_all ;
attn_rot_k = other -> attn_rot_k ;
attn_rot_v = other -> attn_rot_v ;
} else {
const char * LLAMA_ATTN_ROT_DISABLE = getenv ( "LLAMA_ATTN_ROT_DISABLE" );
const bool attn_rot_disable = LLAMA_ATTN_ROT_DISABLE ? atoi ( LLAMA_ATTN_ROT_DISABLE ) : false ;
if ( attn_rot_disable ) {
LLAMA_LOG_WARN ( "%s: attention rotation force disabled (LLAMA_ATTN_ROT_DISABLE) \n " , __func__ );
}
attn_rot_k =
! attn_rot_disable &&
n_embd_head_k_all > 0 &&
ggml_is_quantized ( type_k ) &&
hparams . n_embd_head_k () % 64 == 0 ;
2026-06-29 16:58:51 +08:00
// always create Hadamard rotation tensors for DeepSeek lightning indexers
2026-07-24 20:55:56 +02:00
if (( model . arch == LLM_ARCH_DEEPSEEK32 || model . arch == LLM_ARCH_DEEPSEEK4 || model . arch == LLM_ARCH_GLM_DSA ) &&
2026-06-29 16:58:51 +08:00
hparams . n_embd_head_k_full == hparams . indexer_head_size ) {
2026-06-07 20:50:54 +08:00
attn_rot_k = true ;
}
attn_rot_v =
! attn_rot_disable &&
n_embd_head_v_all > 0 &&
ggml_is_quantized ( type_v ) &&
hparams . n_embd_head_v () % 64 == 0 ;
2026-04-01 16:58:01 +03:00
}
2026-04-07 20:31:28 +03:00
LLAMA_LOG_INFO ( "%s: attn_rot_k = %d, n_embd_head_k_all = %d \n " , __func__ , attn_rot_k , n_embd_head_k_all );
LLAMA_LOG_INFO ( "%s: attn_rot_v = %d, n_embd_head_k_all = %d \n " , __func__ , attn_rot_v , n_embd_head_v_all );
2026-04-01 16:58:01 +03:00
// pre-compute the haramard matrices and keep them in host memory
// TODO: in the future, we can make copies in the backend buffers to avoid host -> device transfers
if ( attn_rot_k || attn_rot_v ) {
2026-04-07 20:31:28 +03:00
for ( int64_t n = 64 ; n <= std :: max ( n_embd_head_k_all , n_embd_head_v_all ); n *= 2 ) {
2026-04-01 16:58:01 +03:00
attn_rot_hadamard [ n ] = std :: vector < float > ( n * n );
ggml_init_params params = {
/* .mem_size = */ 1 * ggml_tensor_overhead (),
/* .mem_buffer = */ nullptr ,
/* .no_alloc = */ true ,
};
ggml_context_ptr ctx { ggml_init ( params ) };
ggml_tensor * tmp = ggml_new_tensor_2d ( ctx . get (), GGML_TYPE_F32 , n , n );
tmp -> data = attn_rot_hadamard [ n ]. data ();
ggml_gen_hadamard ( tmp );
}
}
2025-06-11 12:52:45 +03:00
const char * LLAMA_KV_CACHE_DEBUG = getenv ( "LLAMA_KV_CACHE_DEBUG" );
debug = LLAMA_KV_CACHE_DEBUG ? atoi ( LLAMA_KV_CACHE_DEBUG ) : 0 ;
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: clear ( bool data ) {
2025-07-16 16:35:42 +03:00
for ( uint32_t s = 0 ; s < n_stream ; ++ s ) {
v_cells [ s ]. reset ();
v_heads [ s ] = 0 ;
}
2025-06-01 11:39:27 +03:00
2025-06-06 14:11:15 +03:00
if ( data ) {
2025-10-28 11:23:54 +01:00
for ( auto & [ _ , buf ] : ctxs_bufs ) {
2025-06-06 14:11:15 +03:00
ggml_backend_buffer_clear ( buf . get (), 0 );
}
2025-06-01 11:39:27 +03:00
}
}
2025-08-21 17:00:33 +03:00
bool llama_kv_cache :: seq_rm ( llama_seq_id seq_id , llama_pos p0 , llama_pos p1 ) {
2026-06-07 20:50:54 +08:00
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
if ( other ) {
return true ;
}
2025-08-11 13:58:24 +03:00
GGML_ASSERT ( seq_id == - 1 || ( seq_id >= 0 && ( size_t ) seq_id < seq_to_stream . size ()));
2025-06-01 11:39:27 +03:00
if ( p0 < 0 ) {
p0 = 0 ;
}
if ( p1 < 0 ) {
p1 = std :: numeric_limits < llama_pos >:: max ();
}
2026-07-26 19:43:45 +02:00
// empty range - nothing to remove
if ( p0 >= p1 ) {
return true ;
}
// MSA anchors block selection to absolute cache slots (slot == position). Tail trim and full removal preserve this invariant, but removing a prefix
// or middle range would free slots while later cells survive, desynchronizing the indexer cache. Reject such removals before modifying the cache.
if ( msa_strict_slots ) {
for ( llama_seq_id sid = 0 ; sid < ( llama_seq_id ) seq_to_stream . size (); ++ sid ) {
if ( seq_id >= 0 && sid != seq_id ) {
continue ;
}
const auto & cells = v_cells [ seq_to_stream [ sid ]];
const llama_pos pmin = cells . seq_pos_min ( sid );
const llama_pos pmax = cells . seq_pos_max ( sid );
if ( pmin < 0 ) {
continue ; // empty sequence
}
const bool overlaps = p0 <= pmax && p1 > pmin ; // the range removes something
const bool leaves_tail = p1 <= pmax ; // cells beyond the range survive
if ( overlaps && leaves_tail ) {
LLAMA_LOG_WARN ( "%s: MSA: partial (non-suffix) removal [%d, %d) for seq %d is not supported "
"(block selection is anchored to cache slots) - rejected \n " , __func__ , p0 , p1 , sid );
return false ;
}
}
}
2025-06-04 09:50:32 +03:00
if ( seq_id >= 0 ) {
2025-08-11 13:58:24 +03:00
auto & cells = v_cells [ seq_to_stream [ seq_id ]];
auto & head = v_heads [ seq_to_stream [ seq_id ]];
uint32_t new_head = cells . size ();
2025-06-04 09:50:32 +03:00
for ( uint32_t i = 0 ; i < cells . size (); ++ i ) {
if ( ! cells . pos_in ( i , p0 , p1 )) {
continue ;
}
if ( cells . seq_has ( i , seq_id ) && cells . seq_rm ( i , seq_id )) {
if ( new_head == cells . size ()) {
new_head = i ;
}
}
2025-06-01 11:39:27 +03:00
}
2025-08-11 13:58:24 +03:00
// If we freed up a slot, set head to it so searching can start there.
if ( new_head != cells . size () && new_head < head ) {
head = new_head ;
}
2025-06-04 09:50:32 +03:00
} else {
// match any sequence
2025-08-11 13:58:24 +03:00
for ( uint32_t s = 0 ; s < n_stream ; ++ s ) {
auto & cells = v_cells [ s ];
auto & head = v_heads [ s ];
uint32_t new_head = cells . size ();
for ( uint32_t i = 0 ; i < cells . size (); ++ i ) {
if ( ! cells . pos_in ( i , p0 , p1 )) {
continue ;
}
cells . rm ( i );
if ( new_head == cells . size ()) {
new_head = i ;
}
2025-06-04 09:50:32 +03:00
}
2025-08-11 13:58:24 +03:00
// If we freed up a slot, set head to it so searching can start there.
if ( new_head != cells . size () && new_head < head ) {
head = new_head ;
2025-06-01 11:39:27 +03:00
}
}
}
return true ;
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: seq_cp ( llama_seq_id seq_id_src , llama_seq_id seq_id_dst , llama_pos p0 , llama_pos p1 ) {
2026-06-07 20:50:54 +08:00
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
if ( other ) {
return ;
}
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( seq_id_src >= 0 && ( size_t ) seq_id_src < seq_to_stream . size ());
GGML_ASSERT ( seq_id_dst >= 0 && ( size_t ) seq_id_dst < seq_to_stream . size ());
const auto s0 = seq_to_stream [ seq_id_src ];
const auto s1 = seq_to_stream [ seq_id_dst ];
if ( s0 == s1 ) {
// since both sequences are in the same stream, no data copy is necessary
// we just have to update the cells meta data
auto & cells = v_cells [ s0 ];
if ( seq_id_src == seq_id_dst ) {
return ;
}
if ( p0 < 0 ) {
p0 = 0 ;
}
if ( p1 < 0 ) {
p1 = std :: numeric_limits < llama_pos >:: max ();
}
for ( uint32_t i = 0 ; i < cells . size (); ++ i ) {
if ( ! cells . pos_in ( i , p0 , p1 )) {
continue ;
}
if ( cells . seq_has ( i , seq_id_src )) {
cells . seq_add ( i , seq_id_dst );
}
}
2025-06-01 11:39:27 +03:00
return ;
}
2025-07-16 16:35:42 +03:00
// cross-stream sequence copies require to copy the actual buffer data
bool is_full = true ;
if ( p0 > 0 && p0 + 1 < ( int ) get_size ()) {
is_full = false ;
2025-06-01 11:39:27 +03:00
}
2025-07-16 16:35:42 +03:00
if ( p1 > 0 && p1 + 1 < ( int ) get_size ()) {
is_full = false ;
2025-06-01 11:39:27 +03:00
}
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( is_full && "seq_cp() is only supported for full KV buffers" );
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
// enqueue the copy operation - the buffer copy will be performed during the next update
sc_info . ssrc . push_back ( s0 );
sc_info . sdst . push_back ( s1 );
v_cells [ s1 ]. reset ();
for ( uint32_t i = 0 ; i < v_cells [ s0 ]. size (); ++ i ) {
if ( v_cells [ s0 ]. seq_has ( i , seq_id_src )) {
llama_pos pos = v_cells [ s0 ]. pos_get ( i );
llama_pos shift = v_cells [ s0 ]. get_shift ( i );
2025-10-29 18:09:18 +01:00
llama_kv_cell_ext ext = v_cells [ s0 ]. ext_get ( i );
2025-07-16 16:35:42 +03:00
if ( shift != 0 ) {
pos -= shift ;
assert ( pos >= 0 );
}
v_cells [ s1 ]. pos_set ( i , pos );
v_cells [ s1 ]. seq_add ( i , seq_id_dst );
if ( shift != 0 ) {
v_cells [ s1 ]. pos_add ( i , shift );
}
2025-10-29 18:09:18 +01:00
v_cells [ s1 ]. ext_set ( i , ext );
2025-06-01 11:39:27 +03:00
}
}
2025-07-16 16:35:42 +03:00
v_heads [ s1 ] = v_heads [ s0 ];
//for (uint32_t s = 0; s < n_stream; ++s) {
// LLAMA_LOG_WARN("%s: seq %d: min = %d, max = %d\n", __func__, s, v_cells[s].seq_pos_min(s), v_cells[s].seq_pos_max(s));
//}
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: seq_keep ( llama_seq_id seq_id ) {
2026-06-07 20:50:54 +08:00
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
if ( other ) {
return ;
}
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( seq_id >= 0 && ( size_t ) seq_id < seq_to_stream . size ());
auto & cells = v_cells [ seq_to_stream [ seq_id ]];
auto & head = v_heads [ seq_to_stream [ seq_id ]];
2025-06-01 11:39:27 +03:00
uint32_t new_head = cells . size ();
for ( uint32_t i = 0 ; i < cells . size (); ++ i ) {
if ( cells . seq_keep ( i , seq_id )) {
if ( new_head == cells . size ()) {
new_head = i ;
}
}
}
// If we freed up a slot, set head to it so searching can start there.
if ( new_head != cells . size () && new_head < head ) {
head = new_head ;
}
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: seq_add ( llama_seq_id seq_id , llama_pos p0 , llama_pos p1 , llama_pos shift ) {
2026-06-07 20:50:54 +08:00
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
if ( other ) {
return ;
}
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( seq_id >= 0 && ( size_t ) seq_id < seq_to_stream . size ());
2025-10-29 18:09:18 +01:00
GGML_ASSERT ( hparams . n_pos_per_embd () == 1 && "seq_add() is only supported for n_pos_per_embd() == 1" );
2025-07-16 16:35:42 +03:00
auto & cells = v_cells [ seq_to_stream [ seq_id ]];
auto & head = v_heads [ seq_to_stream [ seq_id ]];
2025-06-01 11:39:27 +03:00
if ( shift == 0 ) {
return ;
}
uint32_t new_head = cells . size ();
if ( p0 < 0 ) {
p0 = 0 ;
}
if ( p1 < 0 ) {
p1 = std :: numeric_limits < llama_pos >:: max ();
}
// If there is no range then return early to avoid looping over all cells.
if ( p0 == p1 ) {
return ;
}
for ( uint32_t i = 0 ; i < cells . size (); ++ i ) {
if ( ! cells . pos_in ( i , p0 , p1 )) {
continue ;
}
if ( cells . seq_has ( i , seq_id )) {
if ( cells . pos_add ( i , shift )) {
if ( new_head == cells . size ()) {
new_head = i ;
}
}
}
}
// If we freed up a slot, set head to it so searching can start there.
// Otherwise we just start the next search from the beginning.
head = new_head != cells . size () ? new_head : 0 ;
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: seq_div ( llama_seq_id seq_id , llama_pos p0 , llama_pos p1 , int d ) {
2026-06-07 20:50:54 +08:00
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
if ( other ) {
return ;
}
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( seq_id >= 0 && ( size_t ) seq_id < seq_to_stream . size ());
2025-10-29 18:09:18 +01:00
GGML_ASSERT ( hparams . n_pos_per_embd () == 1 && "seq_div() is only supported for n_pos_per_embd() == 1" );
2025-07-16 16:35:42 +03:00
auto & cells = v_cells [ seq_to_stream [ seq_id ]];
2025-06-01 11:39:27 +03:00
if ( d == 1 ) {
return ;
}
if ( p0 < 0 ) {
p0 = 0 ;
}
if ( p1 < 0 ) {
p1 = std :: numeric_limits < llama_pos >:: max ();
}
// If there is no range then return early to avoid looping over the cache.
if ( p0 == p1 ) {
return ;
}
for ( uint32_t i = 0 ; i < cells . size (); ++ i ) {
if ( ! cells . pos_in ( i , p0 , p1 )) {
continue ;
}
if ( cells . seq_has ( i , seq_id )) {
cells . pos_div ( i , d );
}
}
}
2025-08-21 17:00:33 +03:00
llama_pos llama_kv_cache :: seq_pos_min ( llama_seq_id seq_id ) const {
2026-06-07 20:50:54 +08:00
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
if ( other ) {
return other -> seq_pos_min ( seq_id );
}
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( seq_id >= 0 && ( size_t ) seq_id < seq_to_stream . size ());
const auto & cells = v_cells [ seq_to_stream [ seq_id ]];
2025-06-01 11:39:27 +03:00
return cells . seq_pos_min ( seq_id );
}
2025-08-21 17:00:33 +03:00
llama_pos llama_kv_cache :: seq_pos_max ( llama_seq_id seq_id ) const {
2026-06-07 20:50:54 +08:00
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
if ( other ) {
return other -> seq_pos_max ( seq_id );
}
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( seq_id >= 0 && ( size_t ) seq_id < seq_to_stream . size ());
const auto & cells = v_cells [ seq_to_stream [ seq_id ]];
2025-06-01 11:39:27 +03:00
return cells . seq_pos_max ( seq_id );
}
2025-09-24 16:53:48 +02:00
std :: map < ggml_backend_buffer_type_t , size_t > llama_kv_cache :: memory_breakdown () const {
std :: map < ggml_backend_buffer_type_t , size_t > ret ;
2025-12-15 09:24:59 +01:00
for ( const auto & [ ctx , buf ] : ctxs_bufs ) {
ggml_backend_buffer_type_t buft = ggml_backend_buffer_get_type ( buf . get ());
if ( hparams . no_alloc ) {
GGML_ASSERT ( ggml_backend_buffer_get_base ( buf . get ()) == nullptr );
ret [ buft ] += ggml_backend_alloc_ctx_tensors_from_buft_size ( ctx . get (), buft );
} else {
// GGML_ASSERT(ggml_backend_buffer_get_base(buf.get()) != nullptr); // multi_buffer does not have a defined base
ret [ buft ] += ggml_backend_buffer_get_size ( buf . get ());
}
2025-09-24 16:53:48 +02:00
}
2025-12-15 09:24:59 +01:00
2025-09-24 16:53:48 +02:00
return ret ;
}
2025-08-21 17:00:33 +03:00
llama_memory_context_ptr llama_kv_cache :: init_batch (
2025-06-20 10:14:14 +03:00
llama_batch_allocr & balloc ,
2025-06-01 11:39:27 +03:00
uint32_t n_ubatch ,
2025-06-16 14:14:00 +03:00
bool embd_all ) {
GGML_UNUSED ( embd_all );
2025-06-01 11:39:27 +03:00
2025-06-12 10:02:15 +03:00
do {
2025-06-20 10:14:14 +03:00
balloc . split_reset ();
2025-06-01 11:39:27 +03:00
2025-06-12 10:02:15 +03:00
std :: vector < llama_ubatch > ubatches ;
2025-06-20 10:14:14 +03:00
while ( true ) {
2026-07-08 15:55:19 +08:00
auto ubatch = n_stream == 1 ? balloc . split_simple ( n_ubatch ) : balloc . split_equal ( n_ubatch , true , 0 );
2025-06-20 10:14:14 +03:00
if ( ubatch . n_tokens == 0 ) {
break ;
}
ubatches . push_back ( std :: move ( ubatch )); // NOLINT
2025-06-12 10:02:15 +03:00
}
2025-06-01 11:39:27 +03:00
2025-07-04 09:04:59 +03:00
if ( balloc . get_n_used () < balloc . get_n_tokens ()) {
// failed to find a suitable split
break ;
}
2025-07-03 10:53:35 +03:00
auto sinfos = prepare ( ubatches );
if ( sinfos . empty ()) {
2025-06-12 10:02:15 +03:00
break ;
}
2025-06-01 11:39:27 +03:00
2025-08-21 17:00:33 +03:00
return std :: make_unique < llama_kv_cache_context > (
2025-07-03 10:53:35 +03:00
this , std :: move ( sinfos ), std :: move ( ubatches ));
2025-06-12 10:02:15 +03:00
} while ( false );
2025-08-21 17:00:33 +03:00
return std :: make_unique < llama_kv_cache_context > ( LLAMA_MEMORY_STATUS_FAILED_PREPARE );
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
llama_memory_context_ptr llama_kv_cache :: init_full () {
return std :: make_unique < llama_kv_cache_context > ( this );
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
llama_memory_context_ptr llama_kv_cache :: init_update ( llama_context * lctx , bool optimize ) {
2025-08-22 12:22:13 +03:00
GGML_UNUSED ( optimize );
2025-06-04 18:58:20 +03:00
bool do_shift = get_has_shift ();
2025-08-22 12:22:13 +03:00
return std :: make_unique < llama_kv_cache_context > ( this , lctx , do_shift , std :: move ( sc_info ));
2025-06-04 18:58:20 +03:00
}
2025-08-21 17:00:33 +03:00
llama_kv_cache :: slot_info_vec_t llama_kv_cache :: prepare ( const std :: vector < llama_ubatch > & ubatches ) {
llama_kv_cache :: slot_info_vec_t res ;
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
struct state_t {
2025-07-03 10:53:35 +03:00
slot_info sinfo ; // slot info for the ubatch
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
std :: vector < uint32_t > v_heads_old ; // old positions of the heads, before placing the ubatch
2025-08-21 17:00:33 +03:00
std :: vector < llama_kv_cells > v_cells ; // copy of the old cells, before placing the ubatch
2025-06-01 11:39:27 +03:00
};
// remember the old state of the cells so we can restore it in the end
2025-07-16 16:35:42 +03:00
std :: vector < state_t > states ;
2025-06-01 11:39:27 +03:00
bool success = true ;
for ( const auto & ubatch : ubatches ) {
// only find a suitable slot for the ubatch. don't modify the cells yet
2025-08-28 17:09:05 +03:00
const auto sinfo_new = find_slot ( ubatch , false );
2025-07-03 10:53:35 +03:00
if ( sinfo_new . empty ()) {
2025-06-01 11:39:27 +03:00
success = false ;
break ;
}
2026-03-05 08:50:21 +01:00
// remember the position that we found
2025-07-03 10:53:35 +03:00
res . push_back ( sinfo_new );
2025-06-01 11:39:27 +03:00
// store the old state of the cells in the recovery stack
2025-07-16 16:35:42 +03:00
{
state_t state = { sinfo_new , v_heads , {} };
for ( uint32_t s = 0 ; s < sinfo_new . n_stream (); ++ s ) {
auto & cells = v_cells [ sinfo_new . strm [ s ]];
state . v_cells . push_back ( cells . cp ( sinfo_new . idxs [ s ]));
}
states . push_back ( std :: move ( state ));
}
2025-06-01 11:39:27 +03:00
// now emplace the ubatch
2025-07-03 10:53:35 +03:00
apply_ubatch ( sinfo_new , ubatch );
2025-06-01 11:39:27 +03:00
}
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( ! states . empty () || ! success );
2025-06-01 11:39:27 +03:00
// iterate backwards and restore the cells to their original state
for ( auto it = states . rbegin (); it != states . rend (); ++ it ) {
2025-07-16 16:35:42 +03:00
const auto & sinfo = it -> sinfo ;
for ( uint32_t s = 0 ; s < sinfo . n_stream (); ++ s ) {
auto & cells = v_cells [ sinfo . strm [ s ]];
auto & head = v_heads [ sinfo . strm [ s ]];
cells . set ( sinfo . idxs [ s ], it -> v_cells [ s ]);
head = it -> v_heads_old [ s ];
}
2025-06-01 11:39:27 +03:00
}
if ( ! success ) {
return {};
}
return res ;
}
2025-08-22 12:22:13 +03:00
bool llama_kv_cache :: update ( llama_context * lctx , bool do_shift , const stream_copy_info & sc_info ) {
2026-06-07 20:50:54 +08:00
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
if ( other ) {
return true ;
}
2025-06-01 11:39:27 +03:00
bool updated = false ;
2025-06-04 18:58:20 +03:00
auto * sched = lctx -> get_sched ();
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
if ( ! sc_info . empty ()) {
assert ( n_stream > 1 && "stream copy should never happen with a single stream" );
llama_synchronize ( lctx );
const size_t n_copy = sc_info . ssrc . size ();
for ( size_t i = 0 ; i < n_copy ; ++ i ) {
const auto ssrc = sc_info . ssrc [ i ];
const auto sdst = sc_info . sdst [ i ];
assert ( ssrc < n_stream );
assert ( sdst < n_stream );
LLAMA_LOG_DEBUG ( "%s: copying KV buffer: stream %d to stream %d \n " , __func__ , ssrc , sdst );
assert ( ssrc != sdst );
for ( uint32_t il = 0 ; il < layers . size (); ++ il ) {
const auto & layer = layers [ il ];
ggml_backend_tensor_copy ( layer . k_stream [ ssrc ], layer . k_stream [ sdst ]);
2026-01-25 15:48:56 +02:00
if ( layer . v_stream [ ssrc ]) {
ggml_backend_tensor_copy ( layer . v_stream [ ssrc ], layer . v_stream [ sdst ]);
}
2026-07-26 19:43:45 +02:00
if ( layer . k_idx_stream [ ssrc ]) {
GGML_ASSERT ( layer . k_idx_stream [ sdst ]);
ggml_backend_tensor_copy ( layer . k_idx_stream [ ssrc ], layer . k_idx_stream [ sdst ]);
}
2025-07-16 16:35:42 +03:00
}
}
}
2025-06-04 18:58:20 +03:00
if ( do_shift ) {
2025-06-01 11:39:27 +03:00
if ( ! get_can_shift ()) {
GGML_ABORT ( "The current KV cache / model configuration does not support K-shift" );
}
LLAMA_LOG_DEBUG ( "%s: applying K-shift \n " , __func__ );
// apply K-shift if needed
if ( hparams . rope_type != LLAMA_ROPE_TYPE_NONE ) {
ggml_backend_sched_reset ( sched );
2025-07-17 19:08:33 +03:00
auto * res = lctx -> get_gf_res_reserve ();
2025-06-01 11:39:27 +03:00
2025-07-17 19:08:33 +03:00
res -> reset ();
2025-06-01 11:39:27 +03:00
2025-07-17 19:08:33 +03:00
auto * gf = build_graph_shift ( res , lctx );
2025-06-01 11:39:27 +03:00
if ( ! ggml_backend_sched_alloc_graph ( sched , gf )) {
LLAMA_LOG_ERROR ( "%s: failed to allocate compute graph for K-shift \n " , __func__ );
return updated ;
}
res -> set_inputs ( nullptr );
2025-06-04 18:58:20 +03:00
if ( lctx -> graph_compute ( gf , false ) != GGML_STATUS_SUCCESS ) {
2025-06-01 11:39:27 +03:00
LLAMA_LOG_ERROR ( "%s: failed to compute K-shift \n " , __func__ );
return updated ;
}
updated = true ;
}
2025-07-16 16:35:42 +03:00
for ( uint32_t s = 0 ; s < n_stream ; ++ s ) {
auto & cells = v_cells [ s ];
cells . reset_shift ();
}
2025-06-01 11:39:27 +03:00
}
return updated ;
}
2025-08-21 17:00:33 +03:00
llama_kv_cache :: slot_info llama_kv_cache :: find_slot ( const llama_ubatch & ubatch , bool cont ) const {
2025-08-11 11:21:19 +02:00
2025-06-11 12:52:45 +03:00
if ( debug > 0 ) {
2025-08-11 11:21:19 +02:00
for ( uint32_t s = 0 ; s < ubatch . n_seqs_unq ; ++ s ) {
const auto seq_id = ubatch . seq_id_unq [ s ];
const auto stream_id = seq_to_stream [ seq_id ];
const auto & cells = v_cells [ stream_id ];
const uint32_t head_cur = v_heads [ stream_id ];
2025-07-16 16:35:42 +03:00
2025-08-11 11:21:19 +02:00
LLAMA_LOG_DEBUG ( "%s: stream[%d], n = %5d, used = %5d, head = %5d, size = %5d, n_swa = %5d \n " ,
__func__ , stream_id , cells . used_max_p1 (), cells . get_used (), head_cur , get_size (), n_swa );
2025-07-16 16:35:42 +03:00
2025-08-11 11:21:19 +02:00
if (( debug == 2 && n_swa > 0 ) || debug > 2 ) {
std :: string ss ;
for ( uint32_t i = 0 ; i < cells . size (); ++ i ) {
if ( cells . is_empty ( i )) {
ss += '.' ;
2025-06-12 10:02:15 +03:00
} else {
2025-08-11 11:21:19 +02:00
assert ( cells . seq_count ( i ) >= 1 );
if ( cells . seq_count ( i ) == 1 ) {
ss += std :: to_string ( cells . seq_get ( i ));
} else {
ss += 'M' ;
}
}
if ( i % 256 == 255 ) {
ss += " *" ;
ss += '\n' ;
2025-06-12 10:02:15 +03:00
}
2025-06-01 11:39:27 +03:00
}
2025-08-11 11:21:19 +02:00
LLAMA_LOG_DEBUG ( " \n %s \n " , ss . c_str ());
2025-06-11 12:52:45 +03:00
}
2025-08-11 11:21:19 +02:00
if (( debug == 2 && n_swa > 0 ) || debug > 2 ) {
std :: string ss ;
for ( uint32_t i = 0 ; i < cells . size (); ++ i ) {
std :: string cur ;
if ( cells . is_empty ( i )) {
cur = '.' ;
} else {
cur = std :: to_string ( cells . pos_get ( i ));
}
const int n = cur . size ();
for ( int j = 0 ; j < 5 - n ; ++ j ) {
cur += ' ' ;
}
ss += cur ;
if ( i % 256 == 255 ) {
ss += " *" ;
}
if ( i % 64 == 63 ) {
ss += '\n' ;
}
}
LLAMA_LOG_DEBUG ( " \n %s \n " , ss . c_str ());
}
for ( int s = 0 ; s < LLAMA_MAX_SEQ ; ++ s ) {
if ( cells . seq_pos_min ( s ) < 0 ) {
continue ;
}
LLAMA_LOG_DEBUG ( "%s: stream[%d] min[%d] = %5d, max[%d] = %5d \n " , __func__ , stream_id , s , cells . seq_pos_min ( s ), s , cells . seq_pos_max ( s ));
}
2025-06-11 12:52:45 +03:00
}
2025-06-01 11:39:27 +03:00
}
2025-07-16 16:35:42 +03:00
uint32_t n_tokens = ubatch . n_tokens ;
uint32_t n_seqs = 1 ;
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
if ( n_stream > 1 ) {
GGML_ASSERT ( n_tokens % ubatch . n_seqs_unq == 0 );
2025-07-03 10:53:35 +03:00
2025-07-16 16:35:42 +03:00
n_seqs = ubatch . n_seqs_unq ;
n_tokens = n_tokens / n_seqs ;
}
2025-07-03 10:53:35 +03:00
2025-07-16 16:35:42 +03:00
slot_info res = {
/*.s0 =*/ LLAMA_MAX_SEQ ,
/*.s1 =*/ 0 ,
/*.strm =*/ { },
/*.idxs =*/ { },
};
2025-07-03 10:53:35 +03:00
2025-07-16 16:35:42 +03:00
res . resize ( n_seqs );
2025-07-03 10:53:35 +03:00
2025-07-16 16:35:42 +03:00
for ( uint32_t s = 0 ; s < n_seqs ; ++ s ) {
const auto seq_id = ubatch . seq_id_unq [ s ];
if ( n_stream > 1 ) {
GGML_ASSERT ( ubatch . n_seq_id [ s * n_tokens ] == 1 );
GGML_ASSERT ( ubatch . seq_id [ s * n_tokens ][ 0 ] == seq_id );
2025-06-01 11:39:27 +03:00
}
2025-08-27 13:55:12 +03:00
res . s0 = std :: min < uint32_t > ( res . s0 , seq_to_stream [ seq_id ]);
res . s1 = std :: max < uint32_t > ( res . s1 , seq_to_stream [ seq_id ]);
2025-07-03 10:53:35 +03:00
2025-07-16 16:35:42 +03:00
res . strm [ s ] = seq_to_stream [ seq_id ];
res . idxs [ s ]. reserve ( n_tokens );
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
const auto & cells = v_cells [ seq_to_stream [ seq_id ]];
2025-06-01 11:39:27 +03:00
2026-07-26 19:43:45 +02:00
if ( n_tokens > cells . size ()) {
LLAMA_LOG_ERROR ( "%s: n_tokens = %d > size = %u \n " , __func__ , n_tokens , cells . size ());
return { };
}
// MSA block selection assumes slot == logical position (append-only streams).
if ( msa_strict_slots ) {
for ( uint32_t ii = 0 ; ii < n_tokens ; ++ ii ) {
const llama_pos pos = ubatch . pos [ s * n_tokens + ii ];
if ( pos < 0 || ( uint64_t ) pos >= cells . size ()) {
LLAMA_LOG_WARN ( "%s: MSA: position %d is outside the cache range [0, %u) \n " ,
__func__ , pos , cells . size ());
return { };
}
const uint32_t idx = ( uint32_t ) pos ;
if ( ! cells . is_empty ( idx )) {
LLAMA_LOG_WARN ( "%s: MSA: required slot %u is already occupied (stream %u) \n " ,
__func__ , idx , seq_to_stream [ seq_id ]);
return { };
}
// strictly increasing positions, rules out duplicates and, for contiguous requests, is tightened to exact adjacency
if ( ! res . idxs [ s ]. empty () && ( cont ? idx != res . idxs [ s ]. back () + 1
: idx <= res . idxs [ s ]. back ())) {
LLAMA_LOG_WARN ( "%s: MSA: token positions are not %s within the ubatch \n " ,
__func__ , cont ? "contiguous" : "strictly increasing" );
return { };
}
res . idxs [ s ]. push_back ( idx );
}
continue ;
}
2025-07-16 16:35:42 +03:00
uint32_t head_cur = v_heads [ seq_to_stream [ seq_id ]];
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
// if we have enough unused cells before the current head ->
// better to start searching from the beginning of the cache, hoping to fill it
if ( head_cur > cells . get_used () + 2 * n_tokens ) {
head_cur = 0 ;
}
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
uint32_t n_tested = 0 ;
// for continuous slots, we test that all tokens in the ubatch fit, starting from the current head
// for non-continuous slots, we test the tokens one by one
const uint32_t n_test = cont ? n_tokens : 1 ;
while ( true ) {
if ( head_cur + n_test > cells . size ()) {
n_tested += cells . size () - head_cur ;
head_cur = 0 ;
continue ;
}
for ( uint32_t i = 0 ; i < n_test ; i ++ ) {
const auto idx = head_cur ;
head_cur ++ ;
n_tested ++ ;
//const llama_pos pos = ubatch.pos[i];
//const llama_seq_id seq_id = ubatch.seq_id[i][0];
// can we use this cell? either:
// - the cell is empty
// - the cell is occupied only by one sequence:
// - (disabled) mask causally, if the sequence is the same as the one we are inserting
// - mask SWA, using current max pos for that sequence in the cache
// always insert in the cell with minimum pos
bool can_use = cells . is_empty ( idx );
if ( ! can_use && cells . seq_count ( idx ) == 1 ) {
const llama_pos pos_cell = cells . pos_get ( idx );
// (disabled) causal mask
// note: it's better to purge any "future" tokens beforehand
//if (cells.seq_has(idx, seq_id)) {
// can_use = pos_cell >= pos;
//}
if ( ! can_use ) {
const llama_seq_id seq_id_cell = cells . seq_get ( idx );
// SWA mask
2026-01-17 15:42:42 +02:00
if ( llama_hparams :: is_masked_swa ( n_swa , swa_type , pos_cell , cells . seq_pos_max ( seq_id_cell ) + 1 )) {
2025-07-16 16:35:42 +03:00
can_use = true ;
}
}
}
if ( can_use ) {
res . idxs [ s ]. push_back ( idx );
} else {
if ( cont ) {
break ;
2025-06-01 11:39:27 +03:00
}
}
}
2025-07-16 16:35:42 +03:00
if ( res . idxs [ s ]. size () == n_tokens ) {
2025-06-01 11:39:27 +03:00
break ;
}
2025-07-16 16:35:42 +03:00
if ( cont ) {
res . idxs [ s ]. clear ();
}
if ( n_tested >= cells . size ()) {
//LLAMA_LOG_ERROR("%s: failed to find a slot for %d tokens\n", __func__, n_tokens);
return { };
}
2025-06-01 11:39:27 +03:00
}
2025-07-16 16:35:42 +03:00
// we didn't find a suitable slot - return empty result
if ( res . idxs [ s ]. size () < n_tokens ) {
2025-07-03 10:53:35 +03:00
return { };
2025-06-01 11:39:27 +03:00
}
}
2025-07-16 16:35:42 +03:00
assert ( res . s1 >= res . s0 );
2025-07-03 10:53:35 +03:00
return res ;
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: apply_ubatch ( const slot_info & sinfo , const llama_ubatch & ubatch ) {
2026-06-07 20:50:54 +08:00
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
if ( other ) {
return ;
}
2025-06-11 16:48:45 +03:00
// keep track of the max sequence position that we would overwrite with this ubatch
// for non-SWA cache, this would be always empty
2025-06-15 10:08:58 +03:00
llama_seq_id seq_pos_max_rm [ LLAMA_MAX_SEQ ];
2025-07-16 16:35:42 +03:00
for ( uint32_t s = 0 ; s < LLAMA_MAX_SEQ ; ++ s ) {
2025-06-11 16:48:45 +03:00
seq_pos_max_rm [ s ] = - 1 ;
}
2025-07-16 16:35:42 +03:00
assert ( ubatch . n_tokens == sinfo . n_stream () * sinfo . size ());
2025-06-11 16:48:45 +03:00
2025-07-16 16:35:42 +03:00
for ( uint32_t s = 0 ; s < sinfo . n_stream (); ++ s ) {
for ( uint32_t ii = 0 ; ii < sinfo . size (); ++ ii ) {
const uint32_t i = s * sinfo . size () + ii ;
2025-07-03 10:53:35 +03:00
2025-07-16 16:35:42 +03:00
auto & cells = v_cells [ sinfo . strm [ s ]];
2025-07-03 10:53:35 +03:00
2025-07-16 16:35:42 +03:00
const auto idx = sinfo . idxs [ s ][ ii ];
2025-06-11 16:48:45 +03:00
2026-07-26 19:43:45 +02:00
if ( msa_strict_slots && ( llama_pos ) idx != ubatch . pos [ i ]) {
LLAMA_LOG_ERROR ( "%s: MSA slot/position invariant violated: "
"writing pos %d into cell %u (stream %u). The indexer cache "
"would desync and block selection would silently corrupt. "
"This is a bug, please report it with reproduction steps. \n " ,
__func__ , ubatch . pos [ i ], idx , sinfo . strm [ s ]);
GGML_ABORT ( "MSA: slot != pos" );
}
2025-07-16 16:35:42 +03:00
if ( ! cells . is_empty ( idx )) {
assert ( cells . seq_count ( idx ) == 1 );
2025-06-11 16:48:45 +03:00
2025-07-16 16:35:42 +03:00
const llama_seq_id seq_id = cells . seq_get ( idx );
const llama_pos pos = cells . pos_get ( idx );
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
seq_pos_max_rm [ seq_id ] = std :: max ( seq_pos_max_rm [ seq_id ], pos );
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
cells . rm ( idx );
}
cells . pos_set ( idx , ubatch . pos [ i ]);
2025-10-29 18:09:18 +01:00
if ( ubatch . is_pos_2d ()) {
llama_kv_cell_ext ext {
/*.x =*/ ubatch . pos [ i + ubatch . n_tokens * 2 ],
/*.y =*/ ubatch . pos [ i + ubatch . n_tokens ],
};
cells . ext_set ( idx , ext );
}
2025-07-16 16:35:42 +03:00
for ( int32_t s = 0 ; s < ubatch . n_seq_id [ i ]; s ++ ) {
cells . seq_add ( idx , ubatch . seq_id [ i ][ s ]);
}
2025-06-01 11:39:27 +03:00
}
}
2025-06-11 16:48:45 +03:00
// note: we want to preserve the invariant that all positions between [pos_min, pos_max] for each sequence
// will be present in the cache. so we have to purge any position which is less than those we would overwrite
// ref: https://github.com/ggml-org/llama.cpp/pull/13746#issuecomment-2916057092
2025-07-16 16:35:42 +03:00
for ( uint32_t s = 0 ; s < LLAMA_MAX_SEQ ; ++ s ) {
2025-06-11 16:48:45 +03:00
if ( seq_pos_max_rm [ s ] == - 1 ) {
continue ;
}
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( s < seq_to_stream . size ());
auto & cells = v_cells [ seq_to_stream [ s ]];
2025-06-11 16:48:45 +03:00
if ( cells . seq_pos_min ( s ) <= seq_pos_max_rm [ s ]) {
LLAMA_LOG_DEBUG ( "%s: purging positions [%d, %d] of sequence %d from KV cache \n " ,
__func__ , cells . seq_pos_min ( s ), seq_pos_max_rm [ s ], s );
2026-07-26 19:43:45 +02:00
// under MSA strict slots this path should be unreachable, since strict MSA placement never selects occupied cells
GGML_ASSERT ( seq_rm ( s , cells . seq_pos_min ( s ), seq_pos_max_rm [ s ] + 1 ));
2025-06-11 16:48:45 +03:00
}
}
2025-06-20 10:14:14 +03:00
2025-06-01 11:39:27 +03:00
// move the head at the end of the slot
2025-07-16 16:35:42 +03:00
for ( uint32_t s = 0 ; s < sinfo . n_stream (); ++ s ) {
auto & head = v_heads [ sinfo . strm [ s ]];
head = sinfo . idxs [ s ]. back () + 1 ;
}
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
bool llama_kv_cache :: get_can_shift () const {
2026-02-07 04:06:14 +08:00
// Step35 uses per-layer RoPE dims; K-shift assumes a single global n_rot.
if ( model . arch == LLM_ARCH_STEP35 ) {
return false ;
}
2026-02-26 18:08:54 +02:00
if ( hparams . n_pos_per_embd () > 1 ) {
return false ;
}
2026-07-26 19:43:45 +02:00
// shifting would leave k_idx stale
for ( const auto & layer : layers ) {
if ( layer . k_idx ) {
return false ;
}
}
2025-06-01 11:39:27 +03:00
return true ;
}
2025-08-21 17:00:33 +03:00
uint32_t llama_kv_cache :: get_size () const {
2025-07-16 16:35:42 +03:00
const auto & cells = v_cells [ seq_to_stream [ 0 ]];
2025-06-01 11:39:27 +03:00
return cells . size ();
}
2025-08-21 17:00:33 +03:00
uint32_t llama_kv_cache :: get_n_stream () const {
2025-07-16 16:35:42 +03:00
return n_stream ;
}
2025-08-21 17:00:33 +03:00
bool llama_kv_cache :: get_has_shift () const {
2025-07-16 16:35:42 +03:00
bool result = false ;
for ( uint32_t s = 0 ; s < n_stream ; ++ s ) {
result |= v_cells [ s ]. get_has_shift ();
}
return result ;
2025-06-04 18:58:20 +03:00
}
2026-04-01 16:58:01 +03:00
ggml_type llama_kv_cache :: type_k () const {
return layers [ 0 ]. k -> type ;
}
ggml_type llama_kv_cache :: type_v () const {
return layers [ 0 ]. v -> type ;
}
2026-06-29 16:58:51 +08:00
std :: vector < uint32_t > llama_kv_cache :: get_layer_ids () const {
std :: vector < uint32_t > res ;
res . reserve ( layers . size ());
for ( const auto & layer : layers ) {
res . push_back ( layer . il );
}
return res ;
}
ggml_tensor * llama_kv_cache :: get_k_storage ( int32_t il ) const {
const int32_t ikv = map_layer_ids . at ( il );
return layers [ ikv ]. k ;
}
2025-08-27 13:55:12 +03:00
uint32_t llama_kv_cache :: get_n_kv ( const slot_info & sinfo ) const {
2025-07-16 16:35:42 +03:00
uint32_t result = 0 ;
2025-10-28 20:19:44 +02:00
// pad the n_kv value so that the graph remains constant across batches and can be reused
// note: this also helps some backends with performance (f.ex https://github.com/ggml-org/llama.cpp/pull/16812#issuecomment-3455112220)
const uint32_t n_pad_cur = std :: max ( n_pad , 256u );
2025-08-27 13:55:12 +03:00
for ( uint32_t s = 0 ; s < sinfo . n_stream (); ++ s ) {
const auto & cells = v_cells [ sinfo . strm [ s ]];
2025-07-16 16:35:42 +03:00
2025-10-28 20:19:44 +02:00
result = std :: max ( std :: min ( cells . size (), std :: max ( n_pad_cur , GGML_PAD ( cells . used_max_p1 (), n_pad_cur ))), result );
2025-07-16 16:35:42 +03:00
}
return result ;
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
ggml_tensor * llama_kv_cache :: get_k ( ggml_context * ctx , int32_t il , uint32_t n_kv , const slot_info & sinfo ) const {
2025-06-01 11:39:27 +03:00
const int32_t ikv = map_layer_ids . at ( il );
auto * k = layers [ ikv ]. k ;
2025-07-16 16:35:42 +03:00
const uint64_t kv_size = get_size ();
const uint64_t n_embd_k_gqa = k -> ne [ 0 ];
assert ( n_embd_k_gqa == hparams . n_embd_k_gqa ( il ));
const uint32_t ns = sinfo . s1 - sinfo . s0 + 1 ;
return ggml_view_4d ( ctx , k ,
2026-03-09 22:22:39 +01:00
hparams . n_embd_head_k ( il ), hparams . n_head_kv ( il ), n_kv , ns ,
ggml_row_size ( k -> type , hparams . n_embd_head_k ( il )),
2025-07-16 16:35:42 +03:00
ggml_row_size ( k -> type , n_embd_k_gqa ),
ggml_row_size ( k -> type , n_embd_k_gqa * kv_size ),
ggml_row_size ( k -> type , n_embd_k_gqa * kv_size ) * sinfo . s0 );
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
ggml_tensor * llama_kv_cache :: get_v ( ggml_context * ctx , int32_t il , uint32_t n_kv , const slot_info & sinfo ) const {
2025-06-01 11:39:27 +03:00
const int32_t ikv = map_layer_ids . at ( il );
auto * v = layers [ ikv ]. v ;
2025-07-16 16:35:42 +03:00
const uint64_t kv_size = get_size ();
const uint64_t n_embd_v_gqa = v -> ne [ 0 ];
// [TAG_V_CACHE_VARIABLE]
assert ( n_embd_v_gqa >= hparams . n_embd_v_gqa ( il ));
const uint32_t ns = sinfo . s1 - sinfo . s0 + 1 ;
2025-06-01 11:39:27 +03:00
if ( ! v_trans ) {
// note: v->nb[1] <= v->nb[2]
2025-07-16 16:35:42 +03:00
return ggml_view_4d ( ctx , v ,
2026-03-09 22:22:39 +01:00
hparams . n_embd_head_v ( il ), hparams . n_head_kv ( il ), n_kv , ns ,
ggml_row_size ( v -> type , hparams . n_embd_head_v ( il )), // v->nb[1]
2025-08-27 13:55:12 +03:00
ggml_row_size ( v -> type , n_embd_v_gqa ), // v->nb[2]
ggml_row_size ( v -> type , n_embd_v_gqa * kv_size ), // v->nb[3]
2025-07-16 16:35:42 +03:00
ggml_row_size ( v -> type , n_embd_v_gqa * kv_size ) * sinfo . s0 );
2025-06-01 11:39:27 +03:00
}
// note: v->nb[1] > v->nb[2]
2025-07-16 16:35:42 +03:00
return ggml_view_4d ( ctx , v ,
2026-03-09 22:22:39 +01:00
n_kv , hparams . n_head_kv ( il ), hparams . n_embd_head_v ( il ), ns ,
ggml_row_size ( v -> type , kv_size * hparams . n_embd_head_v ( il )), // v->nb[1]
2025-08-27 13:55:12 +03:00
ggml_row_size ( v -> type , kv_size ), // v->nb[2]
ggml_row_size ( v -> type , kv_size * n_embd_v_gqa ), // v->nb[3]
2025-07-16 16:35:42 +03:00
ggml_row_size ( v -> type , kv_size * n_embd_v_gqa ) * sinfo . s0 );
2025-06-01 11:39:27 +03:00
}
2026-07-26 19:43:45 +02:00
ggml_tensor * llama_kv_cache :: get_k_idx ( ggml_context * ctx , int32_t il , uint32_t n_kv , const slot_info & sinfo ) const {
const int32_t ikv = map_layer_ids . at ( il );
auto * k_idx = layers [ ikv ]. k_idx ;
GGML_ASSERT ( k_idx );
const uint64_t kv_size = get_size ();
const int64_t n_idx = k_idx -> ne [ 0 ]; // 128
const uint32_t ns = sinfo . s1 - sinfo . s0 + 1 ;
return ggml_view_4d ( ctx , k_idx ,
n_idx , 1 , n_kv , ns ,
ggml_row_size ( k_idx -> type , n_idx ), // nb1 (single head)
ggml_row_size ( k_idx -> type , n_idx ), // nb2 (per cell)
ggml_row_size ( k_idx -> type , n_idx * kv_size ), // nb3 (per stream)
ggml_row_size ( k_idx -> type , n_idx * kv_size ) * sinfo . s0 );
}
2025-08-21 17:00:33 +03:00
ggml_tensor * llama_kv_cache :: cpy_k ( ggml_context * ctx , ggml_tensor * k_cur , ggml_tensor * k_idxs , int32_t il , const slot_info & sinfo ) const {
2025-08-28 12:27:02 +03:00
GGML_UNUSED ( sinfo );
2025-06-01 11:39:27 +03:00
const int32_t ikv = map_layer_ids . at ( il );
2025-09-08 10:25:33 +03:00
ggml_tensor * k = layers [ ikv ]. k ;
2025-06-01 11:39:27 +03:00
2025-09-08 10:25:33 +03:00
const int64_t n_embd_head = k_cur -> ne [ 0 ];
const int64_t n_head = k_cur -> ne [ 1 ];
const int64_t n_tokens = k_cur -> ne [ 2 ];
2025-06-01 11:39:27 +03:00
2025-09-08 10:25:33 +03:00
const int64_t n_embd_gqa = n_embd_head * n_head ;
2025-07-03 10:53:35 +03:00
2025-09-08 10:25:33 +03:00
// we can merge dims 0 and 1
// TODO: add ggml helper function for this?
GGML_ASSERT ( ggml_row_size ( k_cur -> type , n_embd_head ) == k_cur -> nb [ 1 ]);
k_cur = ggml_view_2d ( ctx , k_cur , n_embd_gqa , n_tokens , k_cur -> nb [ 2 ], 0 );
const int64_t n_stream = k -> ne [ 2 ];
if ( n_stream > 1 ) {
const int64_t kv_size = get_size ();
assert ( n_embd_gqa == k -> ne [ 0 ]);
assert ( kv_size == k -> ne [ 1 ]);
// merge the buffer across all streams because the idxs are global
k = ggml_reshape_2d ( ctx , k , n_embd_gqa , kv_size * n_stream );
2025-07-03 10:53:35 +03:00
}
2025-09-08 10:25:33 +03:00
// store the current K values into the cache
2025-08-28 12:27:02 +03:00
return ggml_set_rows ( ctx , k , k_cur , k_idxs );
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
ggml_tensor * llama_kv_cache :: cpy_v ( ggml_context * ctx , ggml_tensor * v_cur , ggml_tensor * v_idxs , int32_t il , const slot_info & sinfo ) const {
2025-08-28 12:27:02 +03:00
GGML_UNUSED ( sinfo );
2025-06-01 11:39:27 +03:00
const int32_t ikv = map_layer_ids . at ( il );
auto * v = layers [ ikv ]. v ;
2025-09-08 10:25:33 +03:00
const int64_t n_embd_head = v_cur -> ne [ 0 ];
const int64_t n_head = v_cur -> ne [ 1 ];
const int64_t n_tokens = v_cur -> ne [ 2 ];
2025-06-01 11:39:27 +03:00
2025-09-08 10:25:33 +03:00
const int64_t n_embd_gqa = n_embd_head * n_head ;
2025-07-03 10:53:35 +03:00
2025-09-08 10:25:33 +03:00
// we can merge dims 0 and 1
GGML_ASSERT ( ggml_row_size ( v_cur -> type , n_embd_head ) == v_cur -> nb [ 1 ]);
const int64_t n_stream = v -> ne [ 2 ];
// take this branch when FA is enabled (the V cache is not transposed)
2025-06-01 11:39:27 +03:00
if ( ! v_trans ) {
2025-09-08 10:25:33 +03:00
v_cur = ggml_view_2d ( ctx , v_cur , n_embd_gqa , n_tokens , v_cur -> nb [ 2 ], 0 );
if ( n_stream > 1 ) {
const int64_t kv_size = get_size ();
assert ( n_embd_gqa == v -> ne [ 0 ]);
assert ( kv_size == v -> ne [ 1 ]);
// merge the buffer across all streams because the idxs are global
v = ggml_reshape_2d ( ctx , v , n_embd_gqa , kv_size * n_stream );
2025-08-28 12:27:02 +03:00
}
2025-07-03 10:53:35 +03:00
2025-08-28 12:27:02 +03:00
return ggml_set_rows ( ctx , v , v_cur , v_idxs );
2025-06-01 11:39:27 +03:00
}
2025-09-08 10:25:33 +03:00
if ( ggml_row_size ( v_cur -> type , n_embd_gqa ) == v_cur -> nb [ 2 ]) {
// we can merge dims 0, 1 and 2
v_cur = ggml_reshape_2d ( ctx , v_cur , n_embd_gqa , n_tokens );
} else {
// otherwise -> make a copy to get contiguous data
v_cur = ggml_cont_2d ( ctx , v_cur , n_embd_gqa , n_tokens );
2025-08-28 12:27:02 +03:00
}
2025-09-08 10:25:33 +03:00
// [TAG_V_CACHE_VARIABLE]
if ( n_embd_gqa < v -> ne [ 0 ]) {
v_cur = ggml_pad ( ctx , v_cur , v -> ne [ 0 ] - n_embd_gqa , 0 , 0 , 0 );
}
2025-08-28 12:27:02 +03:00
2025-09-08 10:25:33 +03:00
// in this branch the v_idxs are constructed in such a way that each row is a single head element
ggml_tensor * v_view = ggml_reshape_2d ( ctx , v , 1 , ggml_nelements ( v ));
v_cur = ggml_reshape_2d ( ctx , v_cur , 1 , ggml_nelements ( v_cur ));
2025-08-28 12:27:02 +03:00
return ggml_set_rows ( ctx , v_view , v_cur , v_idxs );
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
ggml_tensor * llama_kv_cache :: build_input_k_idxs ( ggml_context * ctx , const llama_ubatch & ubatch ) const {
2025-07-03 10:53:35 +03:00
const uint32_t n_tokens = ubatch . n_tokens ;
ggml_tensor * k_idxs = ggml_new_tensor_1d ( ctx , GGML_TYPE_I64 , n_tokens );
ggml_set_input ( k_idxs );
return k_idxs ;
}
2026-07-26 19:43:45 +02:00
ggml_tensor * llama_kv_cache :: cpy_k_idx ( ggml_context * ctx , ggml_tensor * k_idx_cur , ggml_tensor * k_idxs , int32_t il , const slot_info & sinfo ) const {
GGML_UNUSED ( sinfo );
const int32_t ikv = map_layer_ids . at ( il );
ggml_tensor * k_idx = layers [ ikv ]. k_idx ;
GGML_ASSERT ( k_idx && "cpy_k_idx on a layer with no indexer cache" );
const int64_t n_embd_head = k_idx_cur -> ne [ 0 ]; // 128
const int64_t n_head = k_idx_cur -> ne [ 1 ]; // 1
const int64_t n_tokens = k_idx_cur -> ne [ 2 ];
const int64_t n_embd_gqa = n_embd_head * n_head ; // 128
GGML_ASSERT ( ggml_row_size ( k_idx_cur -> type , n_embd_head ) == k_idx_cur -> nb [ 1 ]);
k_idx_cur = ggml_view_2d ( ctx , k_idx_cur , n_embd_gqa , n_tokens , k_idx_cur -> nb [ 2 ], 0 );
const int64_t n_stream = k_idx -> ne [ 2 ];
if ( n_stream > 1 ) {
const int64_t kv_size = get_size ();
k_idx = ggml_reshape_2d ( ctx , k_idx , n_embd_gqa , kv_size * n_stream );
}
return ggml_set_rows ( ctx , k_idx , k_idx_cur , k_idxs ); // same k_idxs as the K store
}
2025-08-21 17:00:33 +03:00
ggml_tensor * llama_kv_cache :: build_input_v_idxs ( ggml_context * ctx , const llama_ubatch & ubatch ) const {
2025-07-03 10:53:35 +03:00
const uint32_t n_tokens = ubatch . n_tokens ;
2025-07-16 16:35:42 +03:00
ggml_tensor * v_idxs ;
if ( ! v_trans ) {
v_idxs = ggml_new_tensor_1d ( ctx , GGML_TYPE_I64 , n_tokens );
} else {
v_idxs = ggml_new_tensor_1d ( ctx , GGML_TYPE_I64 , n_tokens * hparams . n_embd_v_gqa_max ());
}
2025-07-03 10:53:35 +03:00
ggml_set_input ( v_idxs );
return v_idxs ;
}
2026-04-01 16:58:01 +03:00
ggml_tensor * llama_kv_cache :: build_input_k_rot ( ggml_context * ctx ) const {
ggml_tensor * res = nullptr ;
if ( attn_rot_k ) {
int nrot = 64 ;
// TODO: investigate if using the smallest rotation matrix is beneficial also for K (similar as for V)
// ref: https://github.com/ggml-org/llama.cpp/pull/21038#issuecomment-4141323088
do {
nrot *= 2 ;
2026-04-07 20:31:28 +03:00
} while ( n_embd_head_k_all % nrot == 0 );
2026-04-01 16:58:01 +03:00
nrot /= 2 ;
res = ggml_new_tensor_2d ( ctx , GGML_TYPE_F32 , nrot , nrot );
ggml_set_input ( res );
ggml_set_name ( res , "attn_inp_k_rot" );
}
return res ;
}
ggml_tensor * llama_kv_cache :: build_input_v_rot ( ggml_context * ctx ) const {
ggml_tensor * res = nullptr ;
if ( attn_rot_v ) {
int nrot = 64 ;
// using smaller rotation matrices for V seems beneficial
// ref: https://github.com/ggml-org/llama.cpp/pull/21038#issuecomment-4146397570
//do {
// nrot *= 2;
//} while (hparams.n_embd_head_v() % nrot == 0);
//nrot /= 2;
res = ggml_new_tensor_2d ( ctx , GGML_TYPE_F32 , nrot , nrot );
ggml_set_input ( res );
ggml_set_name ( res , "attn_inp_v_rot" );
}
return res ;
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: set_input_k_idxs ( ggml_tensor * dst , const llama_ubatch * ubatch , const slot_info & sinfo ) const {
2025-07-03 10:53:35 +03:00
const uint32_t n_tokens = ubatch -> n_tokens ;
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( n_tokens == ( int64_t ) sinfo . size () * sinfo . n_stream ());
2025-07-03 10:53:35 +03:00
GGML_ASSERT ( ggml_backend_buffer_is_host ( dst -> buffer ));
int64_t * data = ( int64_t * ) dst -> data ;
2025-07-16 16:35:42 +03:00
for ( uint32_t s = 0 ; s < sinfo . n_stream (); ++ s ) {
const int64_t offs = sinfo . strm [ s ] * get_size ();
for ( uint32_t i = 0 ; i < sinfo . size (); ++ i ) {
data [ s * sinfo . size () + i ] = offs + sinfo . idxs [ s ][ i ];
}
2025-07-03 10:53:35 +03:00
}
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: set_input_v_idxs ( ggml_tensor * dst , const llama_ubatch * ubatch , const slot_info & sinfo ) const {
2025-07-03 10:53:35 +03:00
const uint32_t n_tokens = ubatch -> n_tokens ;
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( n_tokens == ( int64_t ) sinfo . size () * sinfo . n_stream ());
2025-07-03 10:53:35 +03:00
GGML_ASSERT ( ggml_backend_buffer_is_host ( dst -> buffer ));
int64_t * data = ( int64_t * ) dst -> data ;
2025-07-16 16:35:42 +03:00
if ( ! v_trans ) {
for ( uint32_t s = 0 ; s < sinfo . n_stream (); ++ s ) {
const int64_t offs = sinfo . strm [ s ] * get_size ();
for ( uint32_t i = 0 ; i < sinfo . size (); ++ i ) {
data [ s * sinfo . size () + i ] = offs + sinfo . idxs [ s ][ i ];
}
}
} else {
// note: the V cache is transposed when not using flash attention
const int64_t kv_size = get_size ();
const int64_t n_embd_v_gqa = hparams . n_embd_v_gqa_max ();
for ( uint32_t s = 0 ; s < sinfo . n_stream (); ++ s ) {
const int64_t offs = sinfo . strm [ s ] * kv_size * n_embd_v_gqa ;
for ( uint32_t i = 0 ; i < sinfo . size (); ++ i ) {
for ( uint32_t j = 0 ; j < n_embd_v_gqa ; ++ j ) {
data [ s * sinfo . size () * n_embd_v_gqa + i * n_embd_v_gqa + j ] = offs + j * kv_size + sinfo . idxs [ s ][ i ];
}
}
}
}
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: set_input_k_shift ( ggml_tensor * dst ) const {
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( ggml_backend_buffer_is_host ( dst -> buffer ));
int32_t * data = ( int32_t * ) dst -> data ;
for ( uint32_t s = 0 ; s < n_stream ; ++ s ) {
const auto & cells = v_cells [ s ];
for ( uint32_t i = 0 ; i < cells . size (); ++ i ) {
2025-07-17 20:52:33 +03:00
data [ s * cells . size () + i ] = cells . is_empty ( i ) ? 0 : cells . get_shift ( i );
2025-07-16 16:35:42 +03:00
}
2025-07-03 10:53:35 +03:00
}
}
2026-01-17 15:42:42 +02:00
struct args_set_input_kq_mask {
const llama_hparams & hparams ;
const llama_ubatch * ubatch ;
const std :: vector < llama_kv_cells > & v_cells ;
const std :: vector < uint32_t > & seq_to_stream ;
uint32_t n_swa ;
llama_swa_type swa_type ;
int64_t n_kv ;
int64_t n_stream ;
int64_t n_tps ;
};
2026-05-29 15:44:43 +08:00
template < typename T , bool causal , bool swa , bool is_2d , bool alibi >
static void set_input_kq_mask_impl ( const args_set_input_kq_mask & args , T * data ) {
2026-01-17 15:42:42 +02:00
//const auto & hparams = args.hparams;
const auto & ubatch = args . ubatch ;
const auto & v_cells = args . v_cells ;
const auto & seq_to_stream = args . seq_to_stream ;
const uint32_t n_swa = args . n_swa ;
const llama_swa_type swa_type = args . swa_type ;
const int64_t n_kv = args . n_kv ;
const int64_t n_stream = args . n_stream ;
const int64_t n_tps = args . n_tps ;
2026-05-29 15:44:43 +08:00
const T mask_keep = llama_cast < T > ( 0.0f );
const T mask_drop = llama_cast < T > ( - INFINITY );
2026-01-17 15:42:42 +02:00
// the min position in the batch for each sequence
llama_pos seq_pos_min [ LLAMA_MAX_SEQ ];
std :: fill ( seq_pos_min , seq_pos_min + LLAMA_MAX_SEQ , INT32_MAX );
for ( uint32_t i = 0 ; i < ubatch -> n_tokens ; ++ i ) {
const llama_seq_id seq_id = ubatch -> seq_id [ i ][ 0 ];
seq_pos_min [ seq_id ] = std :: min ( seq_pos_min [ seq_id ], ubatch -> pos [ i ]);
}
for ( uint32_t s = 0 ; s < n_stream ; ++ s ) {
2026-03-05 08:50:21 +01:00
// bookkeeping of the KQ mask cells that could change for other tokens of the same sequence
2026-01-17 15:42:42 +02:00
std :: unordered_map < llama_seq_id , uint32_t > seq_srct ;
std :: unordered_map < llama_seq_id , std :: vector < uint32_t >> seq_idxs ;
for ( uint32_t ii = 0 ; ii < n_tps ; ++ ii ) {
const uint32_t i = s * n_tps + ii ;
const llama_seq_id seq_id = ubatch -> seq_id [ i ][ 0 ];
const auto & cells = v_cells . at ( seq_to_stream [ seq_id ]);
llama_pos p0 = - 1 ;
const llama_pos p1 = ubatch -> pos [ i ];
// for M-RoPE
const llama_pos p1_x = is_2d ? ubatch -> pos [ i + ubatch -> n_tokens * 2 ] : 0 ;
const llama_pos p1_y = is_2d ? ubatch -> pos [ i + ubatch -> n_tokens ] : 0 ;
const uint64_t idst = n_kv * i ;
// for tokens of the same sequence, the mask is mostly the same, so we can reuse it
// the only cells that could change are the ones that are with similar positions as the
// ones in the batch (i.e. due to causal masking, SWA, etc.)
// keep track of those cells and shortcut the loop to save time
// note: this optimization is not compatible with Alibi position encoding
// ref: https://github.com/ggml-org/llama.cpp/pull/18842
bool prev = false ;
auto & idxs = seq_idxs [ seq_id ];
if ( ! alibi ) {
if ( seq_srct . find ( seq_id ) != seq_srct . end ()) {
const uint32_t srct = seq_srct [ seq_id ];
const uint64_t idst_prev = n_kv * srct ;
std :: copy ( data + idst_prev , data + idst_prev + n_kv , data + idst );
prev = true ;
} else {
idxs . clear ();
idxs . reserve ( ubatch -> n_tokens + n_swa + 32 );
seq_srct [ seq_id ] = i ;
}
}
for ( uint32_t jj = 0 ; jj < n_kv ; ++ jj ) {
uint32_t j = jj ;
// we have an exiting mask for this sequence -> update just seq_idxs
if ( ! alibi ) {
if ( prev ) {
if ( jj >= idxs . size ()) {
break ;
}
j = idxs [ jj ];
}
}
if ( cells . is_empty ( j )) {
goto skip ;
}
// mask the token if not the same sequence
if ( ! cells . seq_has ( j , seq_id )) {
goto skip ;
}
p0 = cells . pos_get ( j );
if ( ! alibi ) {
if ( ! prev ) {
// record all cells for which: p0 >= seq_pos_min[seq_id] - n_swa - 32
if ( p0 + ( int32_t ) ( n_swa + 32 ) >= seq_pos_min [ seq_id ]) {
idxs . push_back ( j );
}
}
}
if ( causal ) {
// mask future tokens
if ( p0 > p1 ) {
goto skip ;
}
// M-RoPE causal mask
if ( is_2d ) {
if ( p0 == p1 ) {
const auto & p0_ext = cells . ext_get ( j );
if ( p0_ext . is_2d_gt ( p1_x , p1_y )) {
goto skip ;
}
}
}
}
// apply SWA if any
if ( swa ) {
if ( llama_hparams :: is_masked_swa ( n_swa , swa_type , p0 , p1 )) {
goto skip ;
}
}
if ( alibi ) {
2026-05-29 15:44:43 +08:00
data [ idst + j ] = llama_cast < T > ( static_cast < float > ( - std :: abs ( p0 - p1 )));
2026-01-17 15:42:42 +02:00
} else {
2026-05-29 15:44:43 +08:00
data [ idst + j ] = mask_keep ;
2026-01-17 15:42:42 +02:00
}
continue ;
skip :
2026-05-29 15:44:43 +08:00
data [ idst + j ] = mask_drop ;
2026-01-17 15:42:42 +02:00
}
}
}
}
2026-05-29 15:44:43 +08:00
template < typename T , bool causal , bool swa , bool is_2d >
static void set_input_kq_mask_impl ( const args_set_input_kq_mask & args , T * data ) {
2026-01-17 15:42:42 +02:00
const bool alibi = args . hparams . use_alibi ;
if ( alibi ) {
2026-05-29 15:44:43 +08:00
set_input_kq_mask_impl < T , causal , swa , is_2d , true > ( args , data );
2026-01-17 15:42:42 +02:00
} else {
2026-05-29 15:44:43 +08:00
set_input_kq_mask_impl < T , causal , swa , is_2d , false > ( args , data );
2026-01-17 15:42:42 +02:00
}
}
2026-05-29 15:44:43 +08:00
template < typename T , bool causal , bool swa >
static void set_input_kq_mask_impl ( const args_set_input_kq_mask & args , T * data ) {
2026-01-17 15:42:42 +02:00
const bool is_2d = args . ubatch -> is_pos_2d ();
if ( is_2d ) {
2026-05-29 15:44:43 +08:00
set_input_kq_mask_impl < T , causal , swa , true > ( args , data );
2026-01-17 15:42:42 +02:00
} else {
2026-05-29 15:44:43 +08:00
set_input_kq_mask_impl < T , causal , swa , false > ( args , data );
2026-01-17 15:42:42 +02:00
}
}
2026-05-29 15:44:43 +08:00
template < typename T , bool causal >
static void set_input_kq_mask_impl ( const args_set_input_kq_mask & args , T * data ) {
2026-01-17 15:42:42 +02:00
const bool swa = args . swa_type != LLAMA_SWA_TYPE_NONE ;
if ( swa ) {
2026-05-29 15:44:43 +08:00
set_input_kq_mask_impl < T , causal , true > ( args , data );
2026-01-17 15:42:42 +02:00
} else {
2026-05-29 15:44:43 +08:00
set_input_kq_mask_impl < T , causal , false > ( args , data );
}
}
template < typename T >
static void set_input_kq_mask_impl ( const args_set_input_kq_mask & args , T * data , bool causal_attn ) {
if ( causal_attn ) {
set_input_kq_mask_impl < T , true > ( args , data );
} else {
set_input_kq_mask_impl < T , false > ( args , data );
2026-01-17 15:42:42 +02:00
}
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: set_input_kq_mask ( ggml_tensor * dst , const llama_ubatch * ubatch , bool causal_attn ) const {
2025-06-20 10:14:14 +03:00
const uint32_t n_tokens = ubatch -> n_tokens ;
2025-06-01 11:39:27 +03:00
GGML_ASSERT ( ggml_backend_buffer_is_host ( dst -> buffer ));
2025-07-16 16:35:42 +03:00
const int64_t n_kv = dst -> ne [ 0 ];
const int64_t n_stream = dst -> ne [ 3 ]; // num streams in the current ubatch
GGML_ASSERT ( n_tokens % n_stream == 0 );
// n_tps == n_tokens_per_stream
2025-12-10 20:53:16 +02:00
const int64_t n_tps = n_tokens / n_stream ;
2025-06-01 11:39:27 +03:00
2026-01-17 15:42:42 +02:00
//const int64_t t_start = ggml_time_us();
2025-07-17 09:49:15 +03:00
2026-01-17 15:42:42 +02:00
const args_set_input_kq_mask args = {
/*.hparams =*/ hparams ,
/*.ubatch =*/ ubatch ,
/*.v_cells =*/ v_cells ,
/*.seq_to_stream =*/ seq_to_stream ,
/*.n_swa =*/ n_swa ,
/*.swa_type =*/ swa_type ,
/*.n_kv =*/ n_kv ,
/*.n_stream =*/ n_stream ,
/*.n_tps =*/ n_tps ,
};
2025-06-01 11:39:27 +03:00
2026-05-29 15:44:43 +08:00
if ( dst -> type == GGML_TYPE_F16 ) {
set_input_kq_mask_impl < ggml_fp16_t > ( args , ( ggml_fp16_t * ) dst -> data , causal_attn );
2026-01-17 15:42:42 +02:00
} else {
2026-05-29 15:44:43 +08:00
set_input_kq_mask_impl < float > ( args , ( float * ) dst -> data , causal_attn );
2025-06-01 11:39:27 +03:00
}
2026-01-17 15:42:42 +02:00
//const int64_t t_end = ggml_time_us();
//LLAMA_LOG_ERROR("%s: kq mask time: %0.3f ms\n", __func__, (t_end - t_start)/1000.0);
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: set_input_pos_bucket ( ggml_tensor * dst , const llama_ubatch * ubatch ) const {
2025-06-01 11:39:27 +03:00
const int64_t n_tokens = ubatch -> n_tokens ;
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( n_stream == 1 && "TODO: support multiple streams" );
const auto & cells = v_cells [ 0 ];
2025-06-01 11:39:27 +03:00
GGML_ASSERT ( ggml_backend_buffer_is_host ( dst -> buffer ));
2025-07-17 19:08:33 +03:00
GGML_ASSERT ( ! ubatch -> equal_seqs ()); // TODO: use ubatch->n_seqs instead of failing
2025-06-01 11:39:27 +03:00
int32_t * data = ( int32_t * ) dst -> data ;
const int32_t n_kv = dst -> ne [ 0 ];
for ( int h = 0 ; h < 1 ; ++ h ) {
2025-06-20 10:14:14 +03:00
for ( int i = 0 ; i < n_tokens ; ++ i ) {
for ( int j = 0 ; j < n_kv ; ++ j ) {
2025-06-01 11:39:27 +03:00
// the position when the cells is empty is irrelevant - it will be masked out later in the attention
2025-06-20 10:14:14 +03:00
const llama_pos p0 = cells . is_empty ( j ) ? - 1 : cells . pos_get ( j );
2025-06-01 11:39:27 +03:00
2025-06-20 10:14:14 +03:00
data [ h * ( n_kv * n_tokens ) + i * n_kv + j ] = llama_relative_position_bucket ( p0 , ubatch -> pos [ i ], hparams . n_rel_attn_bkts , false );
2025-06-01 11:39:27 +03:00
}
}
}
}
2026-04-01 16:58:01 +03:00
void llama_kv_cache :: set_input_k_rot ( ggml_tensor * dst ) const {
GGML_ASSERT ( ggml_backend_buffer_is_host ( dst -> buffer ));
const auto n_rot = dst -> ne [ 0 ];
GGML_ASSERT ( attn_rot_hadamard . count ( dst -> ne [ 0 ]));
memcpy ( dst -> data , attn_rot_hadamard . at ( n_rot ). data (), ggml_nbytes ( dst ));
}
void llama_kv_cache :: set_input_v_rot ( ggml_tensor * dst ) const {
GGML_ASSERT ( ggml_backend_buffer_is_host ( dst -> buffer ));
const auto n_rot = dst -> ne [ 0 ];
GGML_ASSERT ( attn_rot_hadamard . count ( dst -> ne [ 0 ]));
memcpy ( dst -> data , attn_rot_hadamard . at ( n_rot ). data (), ggml_nbytes ( dst ));
}
2025-08-21 17:00:33 +03:00
size_t llama_kv_cache :: total_size () const {
2025-06-01 11:39:27 +03:00
size_t size = 0 ;
2025-10-28 11:23:54 +01:00
for ( const auto & [ _ , buf ] : ctxs_bufs ) {
2025-06-01 11:39:27 +03:00
size += ggml_backend_buffer_get_size ( buf . get ());
}
return size ;
}
2025-08-21 17:00:33 +03:00
size_t llama_kv_cache :: size_k_bytes () const {
2025-06-01 11:39:27 +03:00
size_t size_k_bytes = 0 ;
for ( const auto & layer : layers ) {
size_k_bytes += ggml_nbytes ( layer . k );
}
return size_k_bytes ;
}
2025-08-21 17:00:33 +03:00
size_t llama_kv_cache :: size_v_bytes () const {
2025-06-01 11:39:27 +03:00
size_t size_v_bytes = 0 ;
for ( const auto & layer : layers ) {
2026-01-25 15:48:56 +02:00
size_v_bytes += layer . v ? ggml_nbytes ( layer . v ) : 0 ;
2025-06-01 11:39:27 +03:00
}
return size_v_bytes ;
}
2026-07-26 19:43:45 +02:00
size_t llama_kv_cache :: size_k_idx_bytes () const {
size_t size_k_idx_bytes = 0 ;
for ( const auto & layer : layers ) {
if ( layer . k_idx ) {
size_k_idx_bytes += ggml_nbytes ( layer . k_idx );
}
}
return size_k_idx_bytes ;
}
2025-08-21 17:00:33 +03:00
ggml_tensor * llama_kv_cache :: build_rope_shift (
2025-06-01 11:39:27 +03:00
const llama_cparams & cparams ,
ggml_context * ctx ,
ggml_tensor * cur ,
ggml_tensor * shift ,
2026-04-01 16:58:01 +03:00
ggml_tensor * rot ,
2025-06-01 11:39:27 +03:00
ggml_tensor * factors ,
float freq_base ,
2026-03-09 22:22:39 +01:00
float freq_scale ,
uint32_t il ) const {
2025-06-01 11:39:27 +03:00
const auto & n_ctx_orig = cparams . n_ctx_orig_yarn ;
2025-12-12 17:12:40 +02:00
const auto & yarn_ext_factor = cparams . yarn_ext_factor ;
const auto & yarn_beta_fast = cparams . yarn_beta_fast ;
const auto & yarn_beta_slow = cparams . yarn_beta_slow ;
2025-12-14 08:34:56 +02:00
const auto & yarn_attn_factor = cparams . yarn_attn_factor ;
2025-06-01 11:39:27 +03:00
2026-03-09 22:22:39 +01:00
const auto & n_rot = hparams . n_rot ( il );
2025-10-30 23:19:14 +08:00
const auto & rope_type = hparams . rope_type == LLAMA_ROPE_TYPE_MROPE || hparams . rope_type == LLAMA_ROPE_TYPE_IMROPE
2025-06-01 11:39:27 +03:00
// @ngxson : this is a workaround
// for M-RoPE, we want to rotate the whole vector when doing KV shift
// a normal RoPE should work, we just need to use the correct ordering
// ref: https://github.com/ggml-org/llama.cpp/pull/13870
? LLAMA_ROPE_TYPE_NEOX
: hparams . rope_type ;
ggml_tensor * tmp ;
if ( ggml_is_quantized ( cur -> type )) {
// dequantize to f32 -> RoPE -> quantize back
tmp = ggml_cast ( ctx , cur , GGML_TYPE_F32 );
2026-04-01 16:58:01 +03:00
// rotate back
2026-07-07 17:46:57 +08:00
tmp = llama_mul_mat_hadamard ( ctx , tmp , rot );
2026-04-01 16:58:01 +03:00
2025-06-01 11:39:27 +03:00
tmp = ggml_rope_ext ( ctx , tmp ,
shift , factors , n_rot , rope_type , n_ctx_orig , freq_base , freq_scale ,
yarn_ext_factor , yarn_attn_factor , yarn_beta_fast , yarn_beta_slow );
2026-04-01 16:58:01 +03:00
// rotate fwd
2026-07-07 17:46:57 +08:00
tmp = llama_mul_mat_hadamard ( ctx , tmp , rot );
2026-04-01 16:58:01 +03:00
2025-06-01 11:39:27 +03:00
tmp = ggml_cpy ( ctx , tmp , cur );
} else {
// we rotate only the first n_rot dimensions
tmp = ggml_rope_ext_inplace ( ctx , cur ,
shift , factors , n_rot , rope_type , n_ctx_orig , freq_base , freq_scale ,
yarn_ext_factor , yarn_attn_factor , yarn_beta_fast , yarn_beta_slow );
}
return tmp ;
}
class llm_graph_input_k_shift : public llm_graph_input_i {
public :
2025-08-21 17:00:33 +03:00
llm_graph_input_k_shift ( const llama_kv_cache * kv_self ) : kv_self ( kv_self ) {}
2025-06-01 11:39:27 +03:00
virtual ~ llm_graph_input_k_shift () = default ;
void set_input ( const llama_ubatch * ubatch ) override ;
2025-07-16 16:35:42 +03:00
ggml_tensor * k_shift ; // I32 [kv_size*n_stream]
2025-06-01 11:39:27 +03:00
2026-04-01 16:58:01 +03:00
// note: assumes k_rot^2 == I
ggml_tensor * k_rot = nullptr ;
2025-08-21 17:00:33 +03:00
const llama_kv_cache * kv_self ;
2025-06-01 11:39:27 +03:00
};
void llm_graph_input_k_shift :: set_input ( const llama_ubatch * ubatch ) {
GGML_UNUSED ( ubatch );
if ( k_shift ) {
kv_self -> set_input_k_shift ( k_shift );
}
2026-04-01 16:58:01 +03:00
if ( k_rot ) {
kv_self -> set_input_k_rot ( k_rot );
}
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
ggml_cgraph * llama_kv_cache :: build_graph_shift ( llm_graph_result * res , llama_context * lctx ) const {
2026-06-07 20:50:54 +08:00
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
GGML_ASSERT ( ! other );
2025-07-17 19:08:33 +03:00
auto * ctx = res -> get_ctx ();
auto * gf = res -> get_gf ();
2025-06-01 11:39:27 +03:00
auto inp = std :: make_unique < llm_graph_input_k_shift > ( this );
2025-07-16 16:35:42 +03:00
inp -> k_shift = ggml_new_tensor_1d ( ctx , GGML_TYPE_I32 , ( int64_t ) get_size () * n_stream );
2025-06-01 11:39:27 +03:00
ggml_set_input ( inp -> k_shift );
2026-04-01 16:58:01 +03:00
inp -> k_rot = build_input_k_rot ( ctx );
2025-07-17 19:08:33 +03:00
const auto & cparams = lctx -> get_cparams ();
2025-06-01 11:39:27 +03:00
for ( const auto & layer : layers ) {
const uint32_t il = layer . il ;
const int64_t n_head_kv = hparams . n_head_kv ( il );
const int64_t n_embd_k_gqa = hparams . n_embd_k_gqa ( il );
2026-03-09 22:22:39 +01:00
const auto n_rot = hparams . n_rot ( il );
const auto n_embd_head_k = hparams . n_embd_head_k ( il );
const auto n_embd_nope = hparams . n_lora_kv > 0 ? n_embd_head_k - n_rot : 0 ;
2025-06-01 11:39:27 +03:00
const float freq_base_l = model . get_rope_freq_base ( cparams , il );
const float freq_scale_l = model . get_rope_freq_scale ( cparams , il );
ggml_tensor * rope_factors = model . get_rope_factors ( cparams , il );
ggml_tensor * k =
ggml_view_3d ( ctx , layer . k ,
2026-01-22 22:09:01 +02:00
n_rot , n_head_kv , get_size () * n_stream ,
2025-06-01 11:39:27 +03:00
ggml_row_size ( layer . k -> type , n_embd_head_k ),
ggml_row_size ( layer . k -> type , n_embd_k_gqa ),
2026-01-22 22:09:01 +02:00
ggml_row_size ( layer . k -> type , n_embd_nope ));
2025-06-01 11:39:27 +03:00
2026-04-01 16:58:01 +03:00
ggml_tensor * cur = build_rope_shift ( cparams , ctx , k , inp -> k_shift , inp -> k_rot , rope_factors , freq_base_l , freq_scale_l , il );
2025-06-01 11:39:27 +03:00
ggml_build_forward_expand ( gf , cur );
}
res -> add_input ( std :: move ( inp ));
2025-07-17 19:08:33 +03:00
return gf ;
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: state_write ( llama_io_write_i & io , llama_seq_id seq_id , llama_state_seq_flags flags ) const {
2026-06-07 20:50:54 +08:00
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
if ( other ) {
return ;
}
2025-08-14 14:59:50 +03:00
GGML_UNUSED ( flags );
2025-07-16 16:35:42 +03:00
io . write ( & n_stream , sizeof ( n_stream ));
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
for ( uint32_t s = 0 ; s < n_stream ; ++ s ) {
cell_ranges_t cr { s , {} };
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
uint32_t cell_count = 0 ;
const auto & cells = v_cells [ s ];
// Count the number of cells with the specified seq_id
// Find all the ranges of cells with this seq id (or all, when -1)
uint32_t cell_range_begin = cells . size ();
for ( uint32_t i = 0 ; i < cells . size (); ++ i ) {
2026-06-02 11:06:29 +03:00
bool add_cell = true ;
add_cell = add_cell && ! cells . is_empty ( i );
add_cell = add_cell && ( seq_id == - 1 || cells . seq_has ( i , seq_id ));
// check the cell is not SWA-masked
if ( add_cell && seq_id != - 1 ) {
const bool is_masked = llama_hparams :: is_masked_swa ( n_swa , swa_type , cells . pos_get ( i ), cells . seq_pos_max ( seq_id ));
add_cell = ! is_masked ;
}
if ( add_cell ) {
2025-07-16 16:35:42 +03:00
++ cell_count ;
if ( cell_range_begin == cells . size ()) {
cell_range_begin = i ;
}
} else {
if ( cell_range_begin != cells . size ()) {
cr . data . emplace_back ( cell_range_begin , i );
cell_range_begin = cells . size ();
}
2025-06-01 11:39:27 +03:00
}
}
2025-07-16 16:35:42 +03:00
if ( cell_range_begin != cells . size ()) {
cr . data . emplace_back ( cell_range_begin , cells . size ());
}
// DEBUG CHECK: Sum of cell counts in ranges should equal the total cell count
uint32_t cell_count_check = 0 ;
for ( const auto & range : cr . data ) {
cell_count_check += range . second - range . first ;
}
GGML_ASSERT ( cell_count == cell_count_check );
io . write ( & cell_count , sizeof ( cell_count ));
// skip empty streams
if ( cell_count == 0 ) {
continue ;
}
state_write_meta ( io , cr , seq_id );
state_write_data ( io , cr );
2025-06-01 11:39:27 +03:00
}
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: state_read ( llama_io_read_i & io , llama_seq_id seq_id , llama_state_seq_flags flags ) {
2026-06-07 20:50:54 +08:00
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
if ( other ) {
return ;
}
2025-08-14 14:59:50 +03:00
GGML_UNUSED ( flags );
2025-07-16 16:35:42 +03:00
GGML_ASSERT ( seq_id == - 1 || ( seq_id >= 0 && ( size_t ) seq_id < seq_to_stream . size ()));
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
uint32_t n_stream_cur ;
2026-05-02 18:03:25 +03:00
io . read ( & n_stream_cur , sizeof ( n_stream_cur ));
2025-07-16 16:35:42 +03:00
if ( n_stream_cur != n_stream ) {
throw std :: runtime_error ( "n_stream mismatch" );
}
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
for ( uint32_t s = 0 ; s < n_stream ; ++ s ) {
uint32_t cell_count ;
2026-05-02 18:03:25 +03:00
io . read ( & cell_count , sizeof ( cell_count ));
2025-07-16 16:35:42 +03:00
if ( cell_count == 0 ) {
continue ;
}
const uint32_t strm = seq_id == - 1 ? s : seq_to_stream [ seq_id ];
2025-12-15 09:28:35 -08:00
slot_info sinfo ;
2025-07-16 16:35:42 +03:00
bool res = true ;
2025-12-15 09:28:35 -08:00
res = res && state_read_meta ( io , strm , cell_count , sinfo , seq_id );
2026-07-24 18:56:42 +02:00
try {
res = res && state_read_data ( io , strm , cell_count , sinfo );
} catch (...) {
res = false ;
}
2025-07-16 16:35:42 +03:00
if ( ! res ) {
if ( seq_id == - 1 ) {
clear ( true );
} else {
seq_rm ( seq_id , - 1 , - 1 );
}
throw std :: runtime_error ( "failed to restore kv cache" );
2025-06-01 11:39:27 +03:00
}
}
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: state_write_meta ( llama_io_write_i & io , const cell_ranges_t & cr , llama_seq_id seq_id ) const {
2025-07-16 16:35:42 +03:00
const auto & cells = v_cells [ cr . strm ];
for ( const auto & range : cr . data ) {
2025-06-01 11:39:27 +03:00
for ( uint32_t i = range . first ; i < range . second ; ++ i ) {
std :: vector < llama_seq_id > seq_ids ;
for ( llama_seq_id cur = 0 ; cur < ( int ) n_seq_max ; ++ cur ) {
if ( cur == seq_id || seq_id == - 1 ) {
if ( cells . seq_has ( i , cur )) {
seq_ids . push_back ( cur );
}
}
}
const llama_pos pos = cells . pos_get ( i );
const uint32_t n_seq_id = seq_ids . size ();
io . write ( & pos , sizeof ( pos ));
io . write ( & n_seq_id , sizeof ( n_seq_id ));
2026-03-06 08:46:51 +02:00
if ( hparams . n_pos_per_embd () > 1 ) {
const llama_kv_cell_ext ext = cells . ext_get ( i );
io . write ( & ext , sizeof ( ext ));
}
2025-10-29 18:09:18 +01:00
2025-06-01 11:39:27 +03:00
for ( const auto & seq_id : seq_ids ) {
io . write ( & seq_id , sizeof ( seq_id ));
}
}
}
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache :: state_write_data ( llama_io_write_i & io , const cell_ranges_t & cr ) const {
2025-07-16 16:35:42 +03:00
const auto & cells = v_cells [ cr . strm ];
2025-06-01 11:39:27 +03:00
const uint32_t v_trans = this -> v_trans ? 1 : 0 ;
const uint32_t n_layer = layers . size ();
io . write ( & v_trans , sizeof ( v_trans ));
io . write ( & n_layer , sizeof ( n_layer ));
// Iterate and write all the keys first, each row is a cell
// Get whole range at a time
for ( const auto & layer : layers ) {
const uint32_t il = layer . il ;
2025-06-19 00:08:14 -05:00
const uint32_t n_embd_k_gqa = hparams . n_embd_k_gqa ( il );
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
auto * k = layer . k_stream [ cr . strm ];
2025-06-01 11:39:27 +03:00
// Write key type
2025-07-16 16:35:42 +03:00
const int32_t k_type_i = ( int32_t ) k -> type ;
2025-06-01 11:39:27 +03:00
io . write ( & k_type_i , sizeof ( k_type_i ));
// Write row size of key
2025-07-16 16:35:42 +03:00
const uint64_t k_size_row = ggml_row_size ( k -> type , n_embd_k_gqa );
2025-06-01 11:39:27 +03:00
io . write ( & k_size_row , sizeof ( k_size_row ));
2026-01-30 10:37:06 +01:00
// Read each range of cells of k_size length and write out
2025-07-16 16:35:42 +03:00
for ( const auto & range : cr . data ) {
2025-06-01 11:39:27 +03:00
const size_t range_size = range . second - range . first ;
const size_t buf_size = range_size * k_size_row ;
2025-07-16 16:35:42 +03:00
io . write_tensor ( k , range . first * k_size_row , buf_size );
2025-06-01 11:39:27 +03:00
}
}
2026-07-26 19:43:45 +02:00
if ( size_k_idx_bytes () > 0 ) {
const uint32_t has_k_idx_u32 = 1 ;
io . write ( & has_k_idx_u32 , sizeof ( has_k_idx_u32 ));
for ( const auto & layer : layers ) {
const uint32_t layer_has_k_idx = layer . k_idx ? 1 : 0 ;
io . write ( & layer_has_k_idx , sizeof ( layer_has_k_idx ));
if ( ! layer_has_k_idx ) {
continue ;
}
GGML_ASSERT ( layer . k_idx_stream [ cr . strm ]);
const int32_t k_idx_type_i = ( int32_t ) layer . k_idx -> type ;
io . write ( & k_idx_type_i , sizeof ( k_idx_type_i ));
const uint64_t k_idx_size_row = ggml_row_size ( layer . k_idx -> type , layer . k_idx -> ne [ 0 ]);
io . write ( & k_idx_size_row , sizeof ( k_idx_size_row ));
for ( const auto & range : cr . data ) {
const size_t range_size = range . second - range . first ;
const size_t buf_size = range_size * k_idx_size_row ;
const size_t offset = range . first * k_idx_size_row ;
io . write_tensor ( layer . k_idx_stream [ cr . strm ], offset , buf_size );
}
}
}
2025-06-01 11:39:27 +03:00
if ( ! v_trans ) {
for ( const auto & layer : layers ) {
const uint32_t il = layer . il ;
2025-06-19 00:08:14 -05:00
const uint32_t n_embd_v_gqa = hparams . n_embd_v_gqa ( il );
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
auto * v = layer . v_stream [ cr . strm ];
2026-01-25 15:48:56 +02:00
if ( ! v ) {
continue ;
}
2025-07-16 16:35:42 +03:00
2025-06-01 11:39:27 +03:00
// Write value type
2025-07-16 16:35:42 +03:00
const int32_t v_type_i = ( int32_t ) v -> type ;
2025-06-01 11:39:27 +03:00
io . write ( & v_type_i , sizeof ( v_type_i ));
// Write row size of value
2025-07-16 16:35:42 +03:00
const uint64_t v_size_row = ggml_row_size ( v -> type , n_embd_v_gqa );
2025-06-01 11:39:27 +03:00
io . write ( & v_size_row , sizeof ( v_size_row ));
2026-01-30 10:37:06 +01:00
// Read each range of cells of v_size length and write out
2025-07-16 16:35:42 +03:00
for ( const auto & range : cr . data ) {
2025-06-01 11:39:27 +03:00
const size_t range_size = range . second - range . first ;
const size_t buf_size = range_size * v_size_row ;
2025-07-16 16:35:42 +03:00
io . write_tensor ( v , range . first * v_size_row , buf_size );
2025-06-01 11:39:27 +03:00
}
}
} else {
// When v is transposed, we also need the element size and get the element ranges from each row
const uint32_t kv_size = cells . size ();
for ( const auto & layer : layers ) {
const uint32_t il = layer . il ;
2025-06-19 00:08:14 -05:00
const uint32_t n_embd_v_gqa = hparams . n_embd_v_gqa ( il );
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
auto * v = layer . v_stream [ cr . strm ];
2026-01-25 15:48:56 +02:00
if ( ! v ) {
continue ;
}
2025-07-16 16:35:42 +03:00
2025-06-01 11:39:27 +03:00
// Write value type
2025-07-16 16:35:42 +03:00
const int32_t v_type_i = ( int32_t ) v -> type ;
2025-06-01 11:39:27 +03:00
io . write ( & v_type_i , sizeof ( v_type_i ));
// Write element size
2025-07-16 16:35:42 +03:00
const uint32_t v_size_el = ggml_type_size ( v -> type );
2025-06-01 11:39:27 +03:00
io . write ( & v_size_el , sizeof ( v_size_el ));
// Write GQA embedding size
io . write ( & n_embd_v_gqa , sizeof ( n_embd_v_gqa ));
// For each row, we get the element values of each cell
for ( uint32_t j = 0 ; j < n_embd_v_gqa ; ++ j ) {
2026-01-30 10:37:06 +01:00
// Read each range of cells of v_size_el length and write out
2025-07-16 16:35:42 +03:00
for ( const auto & range : cr . data ) {
2025-06-01 11:39:27 +03:00
const size_t range_size = range . second - range . first ;
const size_t src_offset = ( range . first + j * kv_size ) * v_size_el ;
const size_t buf_size = range_size * v_size_el ;
2025-07-16 16:35:42 +03:00
io . write_tensor ( v , src_offset , buf_size );
2025-06-01 11:39:27 +03:00
}
}
}
}
}
2025-12-15 09:28:35 -08:00
bool llama_kv_cache :: state_read_meta ( llama_io_read_i & io , uint32_t strm , uint32_t cell_count , slot_info & sinfo , llama_seq_id dest_seq_id ) {
2025-07-16 16:35:42 +03:00
auto & cells = v_cells [ strm ];
auto & head = v_heads [ strm ];
2025-06-01 11:39:27 +03:00
if ( dest_seq_id != - 1 ) {
// single sequence
seq_rm ( dest_seq_id , - 1 , - 1 );
2025-06-20 10:14:14 +03:00
llama_batch_allocr balloc ( hparams . n_pos_per_embd ());
2025-06-01 11:39:27 +03:00
2025-06-20 10:14:14 +03:00
llama_ubatch ubatch = balloc . ubatch_reserve ( cell_count , 1 );
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
ubatch . seq_id_unq [ 0 ] = dest_seq_id ;
2025-06-01 11:39:27 +03:00
for ( uint32_t i = 0 ; i < cell_count ; ++ i ) {
llama_pos pos ;
uint32_t n_seq_id ;
2026-05-02 18:03:25 +03:00
io . read ( & pos , sizeof ( pos ));
io . read ( & n_seq_id , sizeof ( n_seq_id ));
2025-06-01 11:39:27 +03:00
if ( n_seq_id != 1 ) {
LLAMA_LOG_ERROR ( "%s: invalid seq_id-agnostic kv cell \n " , __func__ );
return false ;
}
2026-03-06 08:46:51 +02:00
if ( hparams . n_pos_per_embd () > 1 ) {
llama_kv_cell_ext ext ;
2026-05-02 18:03:25 +03:00
io . read ( & ext , sizeof ( ext ));
2026-03-06 08:46:51 +02:00
ubatch . pos [ i + ubatch . n_tokens ] = ext . y ;
ubatch . pos [ i + ubatch . n_tokens * 2 ] = ext . x ;
}
2025-06-01 11:39:27 +03:00
// read the sequence id, but directly discard it - we will use dest_seq_id instead
{
llama_seq_id seq_id ;
2026-05-02 18:03:25 +03:00
io . read ( & seq_id , sizeof ( seq_id ));
2025-06-01 11:39:27 +03:00
}
2025-06-12 10:02:15 +03:00
ubatch . pos [ i ] = pos ;
ubatch . n_seq_id [ i ] = n_seq_id ;
ubatch . seq_id [ i ] = & dest_seq_id ;
2025-06-01 11:39:27 +03:00
}
2025-12-15 09:28:35 -08:00
sinfo = find_slot ( ubatch , false );
2025-07-03 10:53:35 +03:00
if ( sinfo . empty ()) {
2026-06-02 11:06:29 +03:00
LLAMA_LOG_ERROR ( "%s: failed to find %d available cells in kv cache \n " , __func__ , cell_count );
2025-06-01 11:39:27 +03:00
return false ;
}
2025-10-29 18:09:18 +01:00
// TODO: we cannot yet restore llama_kv_cell_ext as the apply_ubatch() does not support it yet
// see: https://github.com/ggml-org/llama.cpp/pull/16825#issuecomment-3460868350
2025-07-03 10:53:35 +03:00
apply_ubatch ( sinfo , ubatch );
2025-12-15 09:28:35 -08:00
LLAMA_LOG_DEBUG ( "%s: cell_count = %d, dest_seq_id = %d \n " , __func__ , cell_count , dest_seq_id );
2025-06-01 11:39:27 +03:00
2025-12-15 09:28:35 -08:00
// DEBUG CHECK: verify that all cells were allocated and have correct seq_id and pos values
GGML_ASSERT ( sinfo . n_stream () == 1 );
GGML_ASSERT ( sinfo . idxs [ 0 ]. size () == cell_count );
for ( uint32_t i = 0 ; i < cell_count ; ++ i ) {
const uint32_t idx = sinfo . idxs [ 0 ][ i ];
GGML_ASSERT ( cells . pos_get ( idx ) == ubatch . pos [ i ]);
GGML_ASSERT ( cells . seq_has ( idx , dest_seq_id ));
}
2025-06-01 11:39:27 +03:00
} else {
// whole KV cache restore
if ( cell_count > cells . size ()) {
LLAMA_LOG_ERROR ( "%s: not enough cells in kv cache \n " , __func__ );
return false ;
}
2025-06-06 14:11:15 +03:00
clear ( true );
2025-06-01 11:39:27 +03:00
for ( uint32_t i = 0 ; i < cell_count ; ++ i ) {
llama_pos pos ;
uint32_t n_seq_id ;
2026-05-02 18:03:25 +03:00
io . read ( & pos , sizeof ( pos ));
io . read ( & n_seq_id , sizeof ( n_seq_id ));
2025-06-01 11:39:27 +03:00
cells . pos_set ( i , pos );
2026-03-15 07:11:19 +00:00
if ( hparams . n_pos_per_embd () > 1 ) {
llama_kv_cell_ext ext ;
2026-05-02 18:03:25 +03:00
io . read ( & ext , sizeof ( ext ));
2026-03-15 07:11:19 +00:00
cells . ext_set ( i , ext );
}
2025-06-01 11:39:27 +03:00
for ( uint32_t j = 0 ; j < n_seq_id ; ++ j ) {
llama_seq_id seq_id ;
2026-05-02 18:03:25 +03:00
io . read ( & seq_id , sizeof ( seq_id ));
2025-06-01 11:39:27 +03:00
if ( seq_id < 0 || ( uint32_t ) seq_id >= n_seq_max ) {
LLAMA_LOG_ERROR ( "%s: invalid seq_id, %d is out of range [0, %u) \n " , __func__ , seq_id , n_seq_max );
return false ;
}
cells . seq_add ( i , seq_id );
}
}
2025-12-15 09:28:35 -08:00
// Create contiguous slot_info for whole cache restore
sinfo . s0 = strm ;
sinfo . s1 = strm ;
sinfo . resize ( 1 );
sinfo . strm [ 0 ] = strm ;
sinfo . idxs [ 0 ]. resize ( cell_count );
for ( uint32_t i = 0 ; i < cell_count ; ++ i ) {
sinfo . idxs [ 0 ][ i ] = i ;
}
2025-06-01 11:39:27 +03:00
head = 0 ;
}
return true ;
}
2025-12-15 09:28:35 -08:00
bool llama_kv_cache :: state_read_data ( llama_io_read_i & io , uint32_t strm , uint32_t cell_count , const slot_info & sinfo ) {
2025-07-16 16:35:42 +03:00
auto & cells = v_cells [ strm ];
2025-06-01 11:39:27 +03:00
uint32_t v_trans ;
uint32_t n_layer ;
2026-05-02 18:03:25 +03:00
io . read ( & v_trans , sizeof ( v_trans ));
io . read ( & n_layer , sizeof ( n_layer ));
2025-06-01 11:39:27 +03:00
if ( n_layer != layers . size ()) {
LLAMA_LOG_ERROR ( "%s: mismatched layer count (%u instead of %u) \n " , __func__ , n_layer , ( uint32_t ) layers . size ());
return false ;
}
if ( cell_count > cells . size ()) {
LLAMA_LOG_ERROR ( "%s: not enough cells in kv cache to restore state (%u > %u) \n " , __func__ , cell_count , cells . size ());
return false ;
}
if ( this -> v_trans != ( bool ) v_trans ) {
LLAMA_LOG_ERROR ( "%s: incompatible V transposition \n " , __func__ );
return false ;
}
// For each layer, read the keys for each cell, one row is one cell, read as one contiguous block
for ( const auto & layer : layers ) {
const uint32_t il = layer . il ;
2025-06-19 00:08:14 -05:00
const uint32_t n_embd_k_gqa = hparams . n_embd_k_gqa ( il );
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
auto * k = layer . k_stream [ strm ];
2025-06-01 11:39:27 +03:00
// Read type of key
int32_t k_type_i_ref ;
2026-05-02 18:03:25 +03:00
io . read ( & k_type_i_ref , sizeof ( k_type_i_ref ));
2025-07-16 16:35:42 +03:00
const int32_t k_type_i = ( int32_t ) k -> type ;
2025-06-01 11:39:27 +03:00
if ( k_type_i != k_type_i_ref ) {
LLAMA_LOG_ERROR ( "%s: mismatched key type (%d != %d, layer %d) \n " , __func__ , k_type_i , k_type_i_ref , il );
return false ;
}
// Read row size of key
uint64_t k_size_row_ref ;
2026-05-02 18:03:25 +03:00
io . read ( & k_size_row_ref , sizeof ( k_size_row_ref ));
2025-07-16 16:35:42 +03:00
const size_t k_size_row = ggml_row_size ( k -> type , n_embd_k_gqa );
2025-06-01 11:39:27 +03:00
if ( k_size_row != k_size_row_ref ) {
LLAMA_LOG_ERROR ( "%s: mismatched key row size (%zu != %zu, layer %d) \n " , __func__ , k_size_row , ( size_t ) k_size_row_ref , il );
return false ;
}
if ( cell_count ) {
2025-12-15 09:28:35 -08:00
if ( sinfo . is_contiguous ()) {
// Fast path: contiguous cells, single memcpy
2026-05-02 18:03:25 +03:00
io . read_tensor ( k , sinfo . head () * k_size_row , cell_count * k_size_row );
2025-12-15 09:28:35 -08:00
} else {
// Slow path: scatter to non-contiguous positions
for ( uint32_t i = 0 ; i < cell_count ; ++ i ) {
const size_t dst_offset = sinfo . idxs [ 0 ][ i ] * k_size_row ;
2026-05-02 18:03:25 +03:00
io . read_tensor ( k , dst_offset , k_size_row );
2025-12-15 09:28:35 -08:00
}
}
2025-06-01 11:39:27 +03:00
}
}
2026-07-26 19:43:45 +02:00
if ( size_k_idx_bytes () > 0 ) {
uint32_t has_k_idx_u32 = 0 ;
io . read ( & has_k_idx_u32 , sizeof ( has_k_idx_u32 ));
if ( has_k_idx_u32 != 1 ) {
LLAMA_LOG_ERROR ( "%s: missing k_idx data in KV cache state \n " , __func__ );
return false ;
}
for ( const auto & layer : layers ) {
uint32_t layer_has_k_idx = 0 ;
io . read ( & layer_has_k_idx , sizeof ( layer_has_k_idx ));
const uint32_t expected_layer_has_k_idx = layer . k_idx ? 1 : 0 ;
if ( layer_has_k_idx != expected_layer_has_k_idx ) {
LLAMA_LOG_ERROR (
"%s: mismatched k_idx state for layer: got %u, expected %u \n " ,
__func__ , layer_has_k_idx , expected_layer_has_k_idx );
return false ;
}
if ( ! layer_has_k_idx ) {
continue ;
}
GGML_ASSERT ( layer . k_idx_stream [ strm ]);
int32_t k_idx_type_i = - 1 ;
io . read ( & k_idx_type_i , sizeof ( k_idx_type_i ));
if ( k_idx_type_i != ( int32_t ) layer . k_idx -> type ) {
LLAMA_LOG_ERROR (
"%s: mismatched k_idx type: got %d, expected %d \n " ,
__func__ , k_idx_type_i , ( int32_t ) layer . k_idx -> type );
return false ;
}
uint64_t k_idx_size_row = 0 ;
io . read ( & k_idx_size_row , sizeof ( k_idx_size_row ));
const uint64_t expected_k_idx_size_row = ggml_row_size ( layer . k_idx -> type , layer . k_idx -> ne [ 0 ]);
if ( k_idx_size_row != expected_k_idx_size_row ) {
LLAMA_LOG_ERROR (
"%s: mismatched k_idx row size: got %zu, expected %zu \n " ,
__func__ , ( size_t ) k_idx_size_row , ( size_t ) expected_k_idx_size_row );
return false ;
}
if ( cell_count ) {
if ( sinfo . is_contiguous ()) {
io . read_tensor ( layer . k_idx_stream [ strm ], sinfo . head () * k_idx_size_row , cell_count * k_idx_size_row );
} else {
for ( uint32_t i = 0 ; i < cell_count ; ++ i ) {
io . read_tensor ( layer . k_idx_stream [ strm ], sinfo . idxs [ 0 ][ i ] * k_idx_size_row , k_idx_size_row );
}
}
}
}
}
2025-06-01 11:39:27 +03:00
if ( ! this -> v_trans ) {
for ( const auto & layer : layers ) {
const uint32_t il = layer . il ;
2025-06-19 00:08:14 -05:00
const uint32_t n_embd_v_gqa = hparams . n_embd_v_gqa ( il );
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
auto * v = layer . v_stream [ strm ];
2026-01-25 15:48:56 +02:00
if ( ! v ) {
continue ;
}
2025-07-16 16:35:42 +03:00
2025-06-01 11:39:27 +03:00
// Read type of value
int32_t v_type_i_ref ;
2026-05-02 18:03:25 +03:00
io . read ( & v_type_i_ref , sizeof ( v_type_i_ref ));
2025-07-16 16:35:42 +03:00
const int32_t v_type_i = ( int32_t ) v -> type ;
2025-06-01 11:39:27 +03:00
if ( v_type_i != v_type_i_ref ) {
LLAMA_LOG_ERROR ( "%s: mismatched value type (%d != %d, layer %d) \n " , __func__ , v_type_i , v_type_i_ref , il );
return false ;
}
// Read row size of value
uint64_t v_size_row_ref ;
2026-05-02 18:03:25 +03:00
io . read ( & v_size_row_ref , sizeof ( v_size_row_ref ));
2025-07-16 16:35:42 +03:00
const size_t v_size_row = ggml_row_size ( v -> type , n_embd_v_gqa );
2025-06-01 11:39:27 +03:00
if ( v_size_row != v_size_row_ref ) {
LLAMA_LOG_ERROR ( "%s: mismatched value row size (%zu != %zu, layer %d) \n " , __func__ , v_size_row , ( size_t ) v_size_row_ref , il );
return false ;
}
if ( cell_count ) {
2025-12-15 09:28:35 -08:00
if ( sinfo . is_contiguous ()) {
// Fast path: contiguous cells, single memcpy
2026-05-02 18:03:25 +03:00
io . read_tensor ( v , sinfo . head () * v_size_row , cell_count * v_size_row );
2025-12-15 09:28:35 -08:00
} else {
// Slow path: scatter to non-contiguous positions
for ( uint32_t i = 0 ; i < cell_count ; ++ i ) {
const size_t dst_offset = sinfo . idxs [ 0 ][ i ] * v_size_row ;
2026-05-02 18:03:25 +03:00
io . read_tensor ( v , dst_offset , v_size_row );
2025-12-15 09:28:35 -08:00
}
}
2025-06-01 11:39:27 +03:00
}
}
} else {
// For each layer, read the values for each cell (transposed)
for ( const auto & layer : layers ) {
const uint32_t il = layer . il ;
2025-06-19 00:08:14 -05:00
const uint32_t n_embd_v_gqa = hparams . n_embd_v_gqa ( il );
2025-06-01 11:39:27 +03:00
2025-07-16 16:35:42 +03:00
auto * v = layer . v_stream [ strm ];
2026-01-25 15:48:56 +02:00
if ( ! v ) {
continue ;
}
2025-07-16 16:35:42 +03:00
2025-06-01 11:39:27 +03:00
// Read type of value
int32_t v_type_i_ref ;
2026-05-02 18:03:25 +03:00
io . read ( & v_type_i_ref , sizeof ( v_type_i_ref ));
2025-07-16 16:35:42 +03:00
const int32_t v_type_i = ( int32_t ) v -> type ;
2025-06-01 11:39:27 +03:00
if ( v_type_i != v_type_i_ref ) {
LLAMA_LOG_ERROR ( "%s: mismatched value type (%d != %d, layer %d) \n " , __func__ , v_type_i , v_type_i_ref , il );
return false ;
}
// Read element size of value
uint32_t v_size_el_ref ;
2026-05-02 18:03:25 +03:00
io . read ( & v_size_el_ref , sizeof ( v_size_el_ref ));
2025-07-16 16:35:42 +03:00
const size_t v_size_el = ggml_type_size ( v -> type );
2025-06-01 11:39:27 +03:00
if ( v_size_el != v_size_el_ref ) {
LLAMA_LOG_ERROR ( "%s: mismatched value element size (%zu != %zu, layer %d) \n " , __func__ , v_size_el , ( size_t ) v_size_el_ref , il );
return false ;
}
// Read GQA embedding size
uint32_t n_embd_v_gqa_ref ;
2026-05-02 18:03:25 +03:00
io . read ( & n_embd_v_gqa_ref , sizeof ( n_embd_v_gqa_ref ));
2025-06-01 11:39:27 +03:00
if ( n_embd_v_gqa != n_embd_v_gqa_ref ) {
LLAMA_LOG_ERROR ( "%s: mismatched GQA embedding size (%u != %u, layer %d) \n " , __func__ , n_embd_v_gqa , n_embd_v_gqa_ref , il );
return false ;
}
if ( cell_count ) {
2025-12-15 09:28:35 -08:00
if ( sinfo . is_contiguous ()) {
// Fast path: contiguous cells
const uint32_t h = sinfo . head ();
for ( uint32_t j = 0 ; j < n_embd_v_gqa ; ++ j ) {
const size_t dst_offset = ( h + j * cells . size ()) * v_size_el ;
2026-05-02 18:03:25 +03:00
io . read_tensor ( v , dst_offset , cell_count * v_size_el );
2025-12-15 09:28:35 -08:00
}
} else {
// Slow path: scatter to non-contiguous positions
for ( uint32_t j = 0 ; j < n_embd_v_gqa ; ++ j ) {
for ( uint32_t i = 0 ; i < cell_count ; ++ i ) {
const size_t dst_offset = ( sinfo . idxs [ 0 ][ i ] + j * cells . size ()) * v_size_el ;
2026-05-02 18:03:25 +03:00
io . read_tensor ( v , dst_offset , v_size_el );
2025-12-15 09:28:35 -08:00
}
}
2025-06-01 11:39:27 +03:00
}
}
}
}
return true ;
}
//
2025-08-21 17:00:33 +03:00
// llama_kv_cache_context
2025-06-01 11:39:27 +03:00
//
2025-08-21 17:00:33 +03:00
llama_kv_cache_context :: llama_kv_cache_context ( llama_memory_status status ) : status ( status ) {}
2025-06-01 11:39:27 +03:00
2025-08-21 17:00:33 +03:00
llama_kv_cache_context :: llama_kv_cache_context (
llama_kv_cache * kv ) : status ( LLAMA_MEMORY_STATUS_SUCCESS ), kv ( kv ) {
2025-06-04 18:58:20 +03:00
n_kv = kv -> get_size ();
2025-07-03 10:53:35 +03:00
2025-07-16 16:35:42 +03:00
const uint32_t n_stream = kv -> get_n_stream ();
2025-07-03 10:53:35 +03:00
// create a dummy slot info - the actual data is irrelevant. we just need to build the graph
sinfos . resize ( 1 );
2025-07-16 16:35:42 +03:00
sinfos [ 0 ]. s0 = 0 ;
sinfos [ 0 ]. s1 = n_stream - 1 ;
sinfos [ 0 ]. idxs . resize ( n_stream );
for ( uint32_t s = 0 ; s < n_stream ; ++ s ) {
sinfos [ 0 ]. strm . push_back ( s );
sinfos [ 0 ]. idxs [ s ]. resize ( 1 , 0 );
}
2025-06-04 18:58:20 +03:00
}
2025-06-01 11:39:27 +03:00
2025-08-21 17:00:33 +03:00
llama_kv_cache_context :: llama_kv_cache_context (
llama_kv_cache * kv ,
2025-06-04 18:58:20 +03:00
llama_context * lctx ,
bool do_shift ,
2025-08-22 12:22:13 +03:00
stream_copy_info sc_info ) : status ( LLAMA_MEMORY_STATUS_SUCCESS ), kv ( kv ), lctx ( lctx ), do_shift ( do_shift ), sc_info ( std :: move ( sc_info )) {
if ( ! do_shift && this -> sc_info . empty ()) {
2025-06-04 18:58:20 +03:00
status = LLAMA_MEMORY_STATUS_NO_UPDATE ;
2025-06-01 11:39:27 +03:00
}
2025-06-04 18:58:20 +03:00
}
2025-08-21 17:00:33 +03:00
llama_kv_cache_context :: llama_kv_cache_context (
llama_kv_cache * kv ,
llama_kv_cache :: slot_info_vec_t sinfos ,
2025-07-03 10:53:35 +03:00
std :: vector < llama_ubatch > ubatches ) : status ( LLAMA_MEMORY_STATUS_SUCCESS ), kv ( kv ), sinfos ( std :: move ( sinfos )), ubatches ( std :: move ( ubatches )) {
2025-06-04 18:58:20 +03:00
}
2025-06-01 11:39:27 +03:00
2025-08-21 17:00:33 +03:00
llama_kv_cache_context ::~ llama_kv_cache_context () = default ;
2025-06-01 11:39:27 +03:00
2025-08-21 17:00:33 +03:00
bool llama_kv_cache_context :: next () {
2025-06-01 11:39:27 +03:00
assert ( status == LLAMA_MEMORY_STATUS_SUCCESS );
2025-07-03 10:53:35 +03:00
if ( ++ i_cur >= ubatches . size ()) {
2025-06-01 11:39:27 +03:00
return false ;
}
return true ;
}
2025-08-21 17:00:33 +03:00
bool llama_kv_cache_context :: apply () {
2025-06-30 18:03:03 +03:00
assert ( ! llama_memory_status_is_fail ( status ));
2025-06-01 11:39:27 +03:00
2025-06-04 18:58:20 +03:00
// no ubatches -> this is a KV cache update
if ( ubatches . empty ()) {
2025-08-22 12:22:13 +03:00
kv -> update ( lctx , do_shift , sc_info );
2025-06-04 18:58:20 +03:00
return true ;
}
2025-07-03 10:53:35 +03:00
kv -> apply_ubatch ( sinfos [ i_cur ], ubatches [ i_cur ]);
2025-08-27 13:55:12 +03:00
n_kv = kv -> get_n_kv ( sinfos [ i_cur ]);
2025-06-01 11:39:27 +03:00
return true ;
}
2025-08-21 17:00:33 +03:00
llama_memory_status llama_kv_cache_context :: get_status () const {
2025-06-01 11:39:27 +03:00
return status ;
}
2025-08-21 17:00:33 +03:00
const llama_ubatch & llama_kv_cache_context :: get_ubatch () const {
2025-06-01 11:39:27 +03:00
assert ( status == LLAMA_MEMORY_STATUS_SUCCESS );
2025-07-03 10:53:35 +03:00
return ubatches [ i_cur ];
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
uint32_t llama_kv_cache_context :: get_n_kv () const {
2025-06-01 11:39:27 +03:00
return n_kv ;
}
2026-04-01 16:58:01 +03:00
ggml_type llama_kv_cache_context :: type_k () const {
return kv -> type_k ();
}
ggml_type llama_kv_cache_context :: type_v () const {
return kv -> type_v ();
}
2025-08-21 17:00:33 +03:00
ggml_tensor * llama_kv_cache_context :: get_k ( ggml_context * ctx , int32_t il ) const {
2025-07-16 16:35:42 +03:00
return kv -> get_k ( ctx , il , n_kv , sinfos [ i_cur ]);
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
ggml_tensor * llama_kv_cache_context :: get_v ( ggml_context * ctx , int32_t il ) const {
2025-07-16 16:35:42 +03:00
return kv -> get_v ( ctx , il , n_kv , sinfos [ i_cur ]);
2025-06-01 11:39:27 +03:00
}
2026-07-26 19:43:45 +02:00
ggml_tensor * llama_kv_cache_context :: get_k_idx ( ggml_context * ctx , int32_t il ) const {
return kv -> get_k_idx ( ctx , il , n_kv , sinfos [ i_cur ]);
}
2025-08-21 17:00:33 +03:00
ggml_tensor * llama_kv_cache_context :: cpy_k ( ggml_context * ctx , ggml_tensor * k_cur , ggml_tensor * k_idxs , int32_t il ) const {
2025-07-03 10:53:35 +03:00
return kv -> cpy_k ( ctx , k_cur , k_idxs , il , sinfos [ i_cur ]);
2025-06-01 11:39:27 +03:00
}
2025-08-21 17:00:33 +03:00
ggml_tensor * llama_kv_cache_context :: cpy_v ( ggml_context * ctx , ggml_tensor * v_cur , ggml_tensor * v_idxs , int32_t il ) const {
2025-07-03 10:53:35 +03:00
return kv -> cpy_v ( ctx , v_cur , v_idxs , il , sinfos [ i_cur ]);
}
2026-07-26 19:43:45 +02:00
ggml_tensor * llama_kv_cache_context :: cpy_k_idx ( ggml_context * ctx , ggml_tensor * k_idx_cur , ggml_tensor * k_idxs , int32_t il ) const {
return kv -> cpy_k_idx ( ctx , k_idx_cur , k_idxs , il , sinfos [ i_cur ]);
}
2025-08-21 17:00:33 +03:00
ggml_tensor * llama_kv_cache_context :: build_input_k_idxs ( ggml_context * ctx , const llama_ubatch & ubatch ) const {
2025-07-03 10:53:35 +03:00
return kv -> build_input_k_idxs ( ctx , ubatch );
}
2025-08-21 17:00:33 +03:00
ggml_tensor * llama_kv_cache_context :: build_input_v_idxs ( ggml_context * ctx , const llama_ubatch & ubatch ) const {
2025-07-03 10:53:35 +03:00
return kv -> build_input_v_idxs ( ctx , ubatch );
2025-06-01 11:39:27 +03:00
}
2026-04-01 16:58:01 +03:00
ggml_tensor * llama_kv_cache_context :: build_input_k_rot ( ggml_context * ctx ) const {
return kv -> build_input_k_rot ( ctx );
}
ggml_tensor * llama_kv_cache_context :: build_input_v_rot ( ggml_context * ctx ) const {
return kv -> build_input_v_rot ( ctx );
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache_context :: set_input_k_shift ( ggml_tensor * dst ) const {
2025-06-01 11:39:27 +03:00
kv -> set_input_k_shift ( dst );
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache_context :: set_input_k_idxs ( ggml_tensor * dst , const llama_ubatch * ubatch ) const {
2025-07-03 10:53:35 +03:00
kv -> set_input_k_idxs ( dst , ubatch , sinfos [ i_cur ]);
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache_context :: set_input_v_idxs ( ggml_tensor * dst , const llama_ubatch * ubatch ) const {
2025-07-03 10:53:35 +03:00
kv -> set_input_v_idxs ( dst , ubatch , sinfos [ i_cur ]);
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache_context :: set_input_kq_mask ( ggml_tensor * dst , const llama_ubatch * ubatch , bool causal_attn ) const {
2025-06-01 11:39:27 +03:00
kv -> set_input_kq_mask ( dst , ubatch , causal_attn );
}
2025-08-21 17:00:33 +03:00
void llama_kv_cache_context :: set_input_pos_bucket ( ggml_tensor * dst , const llama_ubatch * ubatch ) const {
2025-06-01 11:39:27 +03:00
kv -> set_input_pos_bucket ( dst , ubatch );
}
2026-04-01 16:58:01 +03:00
void llama_kv_cache_context :: set_input_k_rot ( ggml_tensor * dst ) const {
kv -> set_input_k_rot ( dst );
}
void llama_kv_cache_context :: set_input_v_rot ( ggml_tensor * dst ) const {
kv -> set_input_v_rot ( dst );
}