spark_storage/
group.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2//
3// Group identity for the high-speed-swap layer.
4//
5// A *group* is the unit of NVMe ↔ HBM movement: one (layer, block, kv_head)
6// tuple's K-stripe or V-stripe. Each group is `block_size × head_dim ×
7// elem_bytes` bytes contiguous on disk, sized to round up to the device's
8// optimal I/O block (typically 4 KiB).
9//
10// The bijection (layer, block, kv_head) ⇆ group_id is computed deterministically
11// from the dimensions; we never store the inverse mapping. Group IDs are dense
12// 64-bit integers.
13
14#[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, // KvKind as u8 (no Hash on enum w/o derive on payload)
29}
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    /// `block_size × head_dim × elem_bytes`, rounded up to fs_block_size.
54    pub group_stride: u64,
55    /// Filesystem block size (4 KiB on most NVMe). Group stride is a multiple.
56    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    /// Bytes occupied by one full layer in its file (K + V across all blocks).
81    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    /// File offset for `key` within its layer's file.
86    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    /// Dense `GroupId` for `key`.
96    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    /// Number of bytes a single group occupies on disk (== group_stride).
108    pub fn group_bytes(&self) -> u64 {
109        self.group_stride
110    }
111
112    /// Bytes in one full block: `K` + `V` across all kv-heads, each a
113    /// `group_stride`-pitch group. This is the contiguous unit the
114    /// block-granular `StorageBackend` ops read/write in one operation.
115    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        // BF16: 16 * 128 * 2 = 4096 bytes — already aligned.
128        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        // Hypothetical odd shape: block_size=16, head_dim=96, BF16 → 3 KiB raw.
142        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        // 4 blocks * 2 kv_heads * 4096 bytes * 2 (K+V) = 65536
150        assert_eq!(l.bytes_per_layer(), 65536);
151    }
152}