23 num_tokens_{slot_mapping.size(0)},
24 num_kv_heads_{k.size(1)},
25 head_size_{k.size(2)},
26 block_size_{k_cache.size(2)},
27 k_src_stride_{k.stride(0)},
28 v_src_stride_{v.stride(0)},
29 k_cache_block_stride_{k_cache.stride(0)},
30 v_cache_block_stride_{v_cache.stride(0)},
31 k_cache_head_stride_{k_cache.stride(1)},
32 v_cache_head_stride_{v_cache.stride(1)},
33 k_cache_slot_stride_{k_cache.stride(2)},
34 v_cache_slot_stride_{v_cache.stride(2)} {
35 assert(k.ndim() == 3 && v.ndim() == 3 &&
36 "`PagedCachingInfinilm` requires `k` and `v` to be 3D");
37 assert(k_cache.ndim() == 4 && v_cache.ndim() == 4 &&
38 "`PagedCachingInfinilm` requires 4D cache tensors");
39 assert(slot_mapping.ndim() == 1 &&
40 "`PagedCachingInfinilm` requires 1D slot mapping");
41 assert((dtype_ == DataType::kFloat16 || dtype_ == DataType::kBFloat16 ||
42 dtype_ == DataType::kFloat32) &&
43 "`PagedCachingInfinilm` supports float16, bfloat16, and float32");
44 assert(v.dtype() == dtype_ && k_cache.dtype() == dtype_ &&
45 v_cache.dtype() == dtype_);
46 assert(slot_mapping.dtype() == DataType::kInt64 &&
47 "`PagedCachingInfinilm` requires int64 slot mapping");
48 assert(k.shape() == v.shape());
49 assert(k_cache.shape() == v_cache.shape());
50 assert(k_cache.size(1) == num_kv_heads_ && k_cache.size(3) == head_size_);
51 assert(k.stride(2) == 1 && v.stride(2) == 1);
52 assert(k_cache.stride(3) == 1 && v_cache.stride(3) == 1);