19 : k_cache_shape_{k_cache.shape()},
20 k_cache_strides_{k_cache.strides()},
21 v_cache_shape_{v_cache.shape()},
22 v_cache_strides_{v_cache.strides()},
24 k_strides_{k.strides()},
26 v_strides_{v.strides()},
27 past_kv_lengths_shape_{past_kv_lengths.shape()},
28 data_type_{k_cache.dtype()},
29 past_kv_lengths_type_{past_kv_lengths.dtype()},
30 batch_size_{k_cache.size(0)},
31 num_kv_heads_{k_cache.size(1)},
32 max_seq_len_{k_cache.size(2)},
34 hidden_size_{k_cache.size(3)},
35 output_size_{k.numel()},
36 device_index_{k_cache.device().index()} {
37 assert(k_cache.ndim() == 4 && v_cache.ndim() == 4 && k.ndim() == 4 &&
38 v.ndim() == 4 &&
"`KvCachingInfinilm` tensors must be 4D");
39 assert(k_cache_shape_ == v_cache_shape_ &&
40 "`KvCachingInfinilm` cache shapes must match");
41 assert(k_shape_ == v_shape_ &&
42 "`KvCachingInfinilm` source shapes must match");
43 assert(k.size(0) == batch_size_ && k.size(1) == num_kv_heads_ &&
44 k.size(3) == hidden_size_ &&
45 "`KvCachingInfinilm` source shape must match cache "
46 "batch/head/hidden dims");
47 assert(seq_len_ <= max_seq_len_ &&
48 "`KvCachingInfinilm` source sequence length exceeds cache length");
49 assert(k_cache.dtype() == v_cache.dtype() && k_cache.dtype() == k.dtype() &&
50 k_cache.dtype() == v.dtype() &&
51 "`KvCachingInfinilm` K/V tensors must have the same dtype");
53 (data_type_ == DataType::kFloat16 ||
54 data_type_ == DataType::kBFloat16 ||
55 data_type_ == DataType::kFloat32) &&
56 "`KvCachingInfinilm` K/V dtype must be float16, bfloat16, or float32");
57 assert((past_kv_lengths_type_ == DataType::kInt32 ||
58 past_kv_lengths_type_ == DataType::kInt64) &&
59 "`KvCachingInfinilm` past_kv_lengths dtype must be int32 or int64");
60 assert(past_kv_lengths.ndim() == 1 &&
61 past_kv_lengths.size(0) == batch_size_ &&
62 "`KvCachingInfinilm` past_kv_lengths shape must be (batch_size,)");
63 assert(!k_cache.HasBroadcastDim() && !v_cache.HasBroadcastDim() &&
64 "`KvCachingInfinilm` caches must not have broadcasted dimensions");