1use std::collections::BTreeSet;
30
31#[derive(Debug, Clone, PartialEq)]
33pub enum CpuLayerSpecError {
34 LayerOutOfRange { index: i64, num_layers: usize },
36 CountOutOfRange { count: i64, num_layers: usize },
38 FractionOutOfRange(f64),
40 Unparsable(String),
42}
43
44impl std::fmt::Display for CpuLayerSpecError {
45 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
46 match self {
47 CpuLayerSpecError::LayerOutOfRange { index, num_layers } => {
48 write!(f, "layer {index} is outside a model of {num_layers} layers")
49 }
50 CpuLayerSpecError::CountOutOfRange { count, num_layers } => write!(
51 f,
52 "{count} CPU layers is more than the model's {num_layers}"
53 ),
54 CpuLayerSpecError::FractionOutOfRange(fraction) => {
55 write!(f, "a CPU-layer fraction must be in [0, 1], got {fraction}")
56 }
57 CpuLayerSpecError::Unparsable(text) => write!(
58 f,
59 "could not read {text:?} as a layer list (`3,7,11`), a count (`8`), or a fraction (`0.5`)"
60 ),
61 }
62 }
63}
64
65impl std::error::Error for CpuLayerSpecError {}
66
67pub fn parse_cpu_layers_spec(
79 spec: &str,
80 num_layers: usize,
81) -> Result<BTreeSet<u32>, CpuLayerSpecError> {
82 let spec = spec.trim();
83 if spec.is_empty() {
84 return Ok(BTreeSet::new());
85 }
86 if spec.contains(',') {
87 let mut layers = BTreeSet::new();
88 for part in spec.split(',') {
89 let part = part.trim();
90 let index: i64 = part
91 .parse()
92 .map_err(|_| CpuLayerSpecError::Unparsable(part.to_string()))?;
93 if index < 0 || index >= num_layers as i64 {
94 return Err(CpuLayerSpecError::LayerOutOfRange { index, num_layers });
95 }
96 layers.insert(index as u32);
97 }
98 return Ok(layers);
99 }
100 let count = if spec.contains('.') {
101 let fraction: f64 = spec
102 .parse()
103 .map_err(|_| CpuLayerSpecError::Unparsable(spec.to_string()))?;
104 if !(0.0..=1.0).contains(&fraction) {
105 return Err(CpuLayerSpecError::FractionOutOfRange(fraction));
106 }
107 round_half_even(fraction * num_layers as f64)
108 } else {
109 let count: i64 = spec
110 .parse()
111 .map_err(|_| CpuLayerSpecError::Unparsable(spec.to_string()))?;
112 if count < 0 || count > num_layers as i64 {
113 return Err(CpuLayerSpecError::CountOutOfRange { count, num_layers });
114 }
115 count
116 };
117 Ok(strided_layers(count as usize, num_layers))
118}
119
120pub fn strided_layers(count: usize, num_layers: usize) -> BTreeSet<u32> {
122 if count == 0 {
123 return BTreeSet::new();
124 }
125 (0..count)
126 .map(|i| round_half_even((i * num_layers) as f64 / count as f64) as u32)
127 .collect()
128}
129
130pub fn auto_cpu_layers(
138 num_layers: usize,
139 bank_bytes: u64,
140 pin_budget_bytes: Option<u64>,
141) -> BTreeSet<u32> {
142 let Some(budget) = pin_budget_bytes else {
143 return BTreeSet::new();
144 };
145 if bank_bytes == 0 || bank_bytes <= budget {
146 return BTreeSet::new();
147 }
148 let unpinnable = 1.0 - (budget as f64 / bank_bytes as f64);
149 let n = (unpinnable * num_layers as f64).ceil() as usize;
150 let n = n.min(num_layers);
151 let head = n.div_ceil(2);
152 let mut layers: BTreeSet<u32> = (0..head as u32).collect();
153 layers.extend(((num_layers - (n - head)) as u32)..num_layers as u32);
154 layers
155}
156
157pub fn round_half_even(value: f64) -> i64 {
169 let floor = value.floor();
170 let diff = value - floor;
171 let floor = floor as i64;
172 match diff.partial_cmp(&0.5) {
173 Some(std::cmp::Ordering::Less) => floor,
174 Some(std::cmp::Ordering::Greater) => floor + 1,
175 _ => {
177 if floor % 2 == 0 {
178 floor
179 } else {
180 floor + 1
181 }
182 }
183 }
184}
185
186#[cfg(test)]
187mod tests {
188 use super::*;
189
190 fn set(items: &[u32]) -> BTreeSet<u32> {
191 items.iter().copied().collect()
192 }
193
194 #[test]
195 fn an_explicit_list_names_exactly_those_layers() {
196 assert_eq!(
197 parse_cpu_layers_spec("3,7,11", 40).unwrap(),
198 set(&[3, 7, 11])
199 );
200 assert_eq!(parse_cpu_layers_spec(" 3 , 7 ", 40).unwrap(), set(&[3, 7]));
201 assert_eq!(parse_cpu_layers_spec("", 40).unwrap(), BTreeSet::new());
202 assert_eq!(parse_cpu_layers_spec(" ", 40).unwrap(), BTreeSet::new());
203 }
204
205 #[test]
208 fn a_count_is_spread_evenly_through_the_stack() {
209 assert_eq!(
210 parse_cpu_layers_spec("8", 40).unwrap(),
211 set(&[0, 5, 10, 15, 20, 25, 30, 35])
212 );
213 assert_eq!(parse_cpu_layers_spec("0", 40).unwrap(), BTreeSet::new());
214 assert_eq!(parse_cpu_layers_spec("40", 40).unwrap().len(), 40);
215 }
216
217 #[test]
218 fn a_fraction_is_a_count_of_the_model() {
219 assert_eq!(parse_cpu_layers_spec("0.5", 40).unwrap().len(), 20);
220 assert_eq!(parse_cpu_layers_spec("1.0", 40).unwrap().len(), 40);
221 assert_eq!(parse_cpu_layers_spec("0.0", 40).unwrap(), BTreeSet::new());
222 }
223
224 #[test]
225 fn a_spec_that_names_layers_the_model_lacks_is_refused() {
226 assert!(matches!(
227 parse_cpu_layers_spec("40,1", 40),
228 Err(CpuLayerSpecError::LayerOutOfRange { index: 40, .. })
229 ));
230 assert!(matches!(
231 parse_cpu_layers_spec("-1", 40),
232 Err(CpuLayerSpecError::CountOutOfRange { count: -1, .. })
233 ));
234 assert!(matches!(
235 parse_cpu_layers_spec("99", 40),
236 Err(CpuLayerSpecError::CountOutOfRange { count: 99, .. })
237 ));
238 assert!(matches!(
239 parse_cpu_layers_spec("1.5", 40),
240 Err(CpuLayerSpecError::FractionOutOfRange(_))
241 ));
242 assert!(matches!(
243 parse_cpu_layers_spec("half", 40),
244 Err(CpuLayerSpecError::Unparsable(_))
245 ));
246 }
247
248 #[test]
251 fn a_model_that_fits_the_pin_budget_keeps_every_layer_on_the_gpu() {
252 assert_eq!(auto_cpu_layers(40, 8 << 30, None), BTreeSet::new());
253 assert_eq!(
254 auto_cpu_layers(40, 8 << 30, Some(16 << 30)),
255 BTreeSet::new()
256 );
257 assert_eq!(auto_cpu_layers(40, 0, Some(1 << 30)), BTreeSet::new());
258 }
259
260 #[test]
263 fn an_over_budget_model_gives_up_layers_from_both_ends() {
264 let layers = auto_cpu_layers(40, 16 << 30, Some(8 << 30));
265 assert_eq!(layers.len(), 20);
266 assert!(layers.contains(&0) && layers.contains(&9));
267 assert!(layers.contains(&39) && layers.contains(&30));
268 assert!(
269 !layers.contains(&15) && !layers.contains(&20),
270 "the middle layers keep their GPU residency"
271 );
272 }
273
274 #[test]
275 fn a_model_far_over_budget_moves_every_layer() {
276 let layers = auto_cpu_layers(8, 100 << 30, Some(1 << 30));
277 assert_eq!(layers.len(), 8);
278 }
279
280 #[test]
283 fn an_odd_count_splits_head_heavy_without_overlapping() {
284 let layers = auto_cpu_layers(10, 1024, Some(768));
286 assert_eq!(layers, set(&[0, 1, 9]));
287 }
288
289 #[test]
290 fn halves_round_the_way_python_does() {
291 assert_eq!(round_half_even(0.5), 0);
292 assert_eq!(round_half_even(1.5), 2);
293 assert_eq!(round_half_even(2.5), 2);
294 assert_eq!(round_half_even(2.4), 2);
295 assert_eq!(round_half_even(2.6), 3);
296 }
297}