1#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]
15pub struct GroupId(pub u64);
16
17#[derive(Clone, Copy, Debug, PartialEq, Eq)]
18pub enum KvKind {
19 K = 0,
20 V = 1,
21}
22
23#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
24pub struct GroupKey {
25 pub layer: u32,
26 pub block: u32,
27 pub kv_head: u16,
28 pub kv_kind: u8, }
30
31impl GroupKey {
32 pub fn new(layer: u32, block: u32, kv_head: u16, kv_kind: KvKind) -> Self {
33 Self {
34 layer,
35 block,
36 kv_head,
37 kv_kind: kv_kind as u8,
38 }
39 }
40 pub fn kind(self) -> KvKind {
41 match self.kv_kind {
42 0 => KvKind::K,
43 _ => KvKind::V,
44 }
45 }
46}
47
48#[derive(Clone, Copy, Debug)]
49pub struct GroupLayout {
50 pub num_layers: u32,
51 pub num_blocks: u32,
52 pub num_kv_heads: u16,
53 pub group_stride: u64,
55 pub fs_block_size: u64,
57}
58
59impl GroupLayout {
60 pub fn new(
61 num_layers: u32,
62 num_blocks: u32,
63 num_kv_heads: u16,
64 block_size: u32,
65 head_dim: u32,
66 elem_bytes: u32,
67 fs_block_size: u64,
68 ) -> Self {
69 let raw = (block_size as u64) * (head_dim as u64) * (elem_bytes as u64);
70 let group_stride = raw.div_ceil(fs_block_size) * fs_block_size;
71 Self {
72 num_layers,
73 num_blocks,
74 num_kv_heads,
75 group_stride,
76 fs_block_size,
77 }
78 }
79
80 pub fn bytes_per_layer(&self) -> u64 {
82 2 * (self.num_blocks as u64) * (self.num_kv_heads as u64) * self.group_stride
83 }
84
85 pub fn file_offset(&self, key: GroupKey) -> u64 {
87 debug_assert!(key.block < self.num_blocks);
88 debug_assert!(key.kv_head < self.num_kv_heads);
89 let kv_stride = (self.num_kv_heads as u64) * self.group_stride;
90 (key.block as u64) * (2 * kv_stride)
91 + (key.kv_kind as u64) * kv_stride
92 + (key.kv_head as u64) * self.group_stride
93 }
94
95 pub fn group_id(&self, key: GroupKey) -> GroupId {
97 let per_layer = 2 * (self.num_blocks as u64) * (self.num_kv_heads as u64);
98 let per_block = 2 * (self.num_kv_heads as u64);
99 GroupId(
100 (key.layer as u64) * per_layer
101 + (key.block as u64) * per_block
102 + (key.kv_kind as u64) * (self.num_kv_heads as u64)
103 + (key.kv_head as u64),
104 )
105 }
106
107 pub fn group_bytes(&self) -> u64 {
109 self.group_stride
110 }
111
112 pub fn block_bytes(&self) -> u64 {
116 2 * (self.num_kv_heads as u64) * self.group_stride
117 }
118}
119
120#[cfg(test)]
121mod tests {
122 use super::*;
123
124 #[test]
125 fn offset_and_id_are_consistent() {
126 let l = GroupLayout::new(80, 4096, 8, 16, 128, 2, 4096);
127 assert_eq!(l.group_stride, 4096);
129 let k0 = GroupKey::new(0, 0, 0, KvKind::K);
130 assert_eq!(l.file_offset(k0), 0);
131 let k1 = GroupKey::new(0, 0, 0, KvKind::V);
132 assert_eq!(l.file_offset(k1), (8u64) * 4096);
133 let k2 = GroupKey::new(0, 1, 0, KvKind::K);
134 assert_eq!(l.file_offset(k2), 2 * 8 * 4096);
135 let k3 = GroupKey::new(0, 0, 7, KvKind::V);
136 assert_eq!(l.file_offset(k3), 8 * 4096 + 7 * 4096);
137 }
138
139 #[test]
140 fn rounds_up_to_fs_block() {
141 let l = GroupLayout::new(1, 1, 1, 16, 96, 2, 4096);
143 assert_eq!(l.group_stride, 4096);
144 }
145
146 #[test]
147 fn bytes_per_layer_correct() {
148 let l = GroupLayout::new(1, 4, 2, 16, 128, 2, 4096);
149 assert_eq!(l.bytes_per_layer(), 65536);
151 }
152}