Skip to main content

ferrum_cli/
gpu_devices.rs

1use ferrum_types::{Device, FerrumError, Result, RuntimeConfigEntry, RuntimeConfigSource};
2use serde::{Deserialize, Serialize};
3use serde_json::Value;
4use std::collections::{BTreeSet, HashMap};
5
6pub const GPU_DEVICES_RAW_KEY: &str = "FERRUM_GPU_DEVICES_RAW";
7pub const REQUESTED_GPU_DEVICES_KEY: &str = "FERRUM_REQUESTED_GPU_DEVICES";
8pub const SELECTED_GPU_DEVICES_KEY: &str = "FERRUM_SELECTED_GPU_DEVICES";
9pub const SELECTED_DISTRIBUTED_STRATEGY_KEY: &str = "FERRUM_SELECTED_DISTRIBUTED_STRATEGY";
10pub const CUDA_DEVICE_COUNT_KEY: &str = "FERRUM_CUDA_DEVICE_COUNT";
11pub const SELECTED_LAYER_SPLIT_PLAN_KEY: &str = "FERRUM_SELECTED_LAYER_SPLIT_PLAN";
12pub const SELECTED_LAYER_SPLIT_STAGES_KEY: &str = "FERRUM_SELECTED_LAYER_SPLIT_STAGES";
13
14#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
15pub struct CudaLayerSplitStage {
16    pub stage: usize,
17    pub device: usize,
18    pub layer_start: usize,
19    pub layer_end: usize,
20}
21
22#[derive(Debug, Clone, PartialEq, Eq)]
23pub struct GpuDeviceSelection {
24    pub raw_cli_value: String,
25    pub requested_gpu_devices: Vec<usize>,
26    pub selected_gpu_devices: Vec<usize>,
27    pub cuda_device_count: usize,
28    pub selected_distributed_strategy: String,
29    pub selected_layer_split_plan: Option<String>,
30    pub selected_layer_split_stages: Option<Vec<CudaLayerSplitStage>>,
31}
32
33impl GpuDeviceSelection {
34    pub fn primary_device(&self) -> Device {
35        Device::CUDA(self.selected_gpu_devices[0])
36    }
37
38    pub fn requested_csv(&self) -> String {
39        join_gpu_devices(&self.requested_gpu_devices)
40    }
41
42    pub fn selected_csv(&self) -> String {
43        join_gpu_devices(&self.selected_gpu_devices)
44    }
45
46    pub fn apply_model_layer_count(&mut self, num_layers: usize) -> Result<bool> {
47        if self.selected_gpu_devices.len() <= 1 {
48            return Ok(false);
49        }
50        let stages = even_layer_split_stages_for_layers(&self.selected_gpu_devices, num_layers)?;
51        self.selected_layer_split_plan = Some(format_layer_split_plan(&stages));
52        self.selected_layer_split_stages = Some(stages);
53        Ok(true)
54    }
55
56    pub fn runtime_config_entries(&self) -> Vec<RuntimeConfigEntry> {
57        vec![
58            RuntimeConfigEntry::new("FERRUM_BACKEND", "cuda", RuntimeConfigSource::Cli),
59            RuntimeConfigEntry::new(
60                GPU_DEVICES_RAW_KEY,
61                self.raw_cli_value.clone(),
62                RuntimeConfigSource::Cli,
63            ),
64            RuntimeConfigEntry::new(
65                REQUESTED_GPU_DEVICES_KEY,
66                self.requested_csv(),
67                RuntimeConfigSource::Cli,
68            ),
69            RuntimeConfigEntry::new(
70                SELECTED_GPU_DEVICES_KEY,
71                self.selected_csv(),
72                RuntimeConfigSource::Cli,
73            ),
74            RuntimeConfigEntry::new(
75                SELECTED_DISTRIBUTED_STRATEGY_KEY,
76                self.selected_distributed_strategy.clone(),
77                RuntimeConfigSource::Cli,
78            ),
79            RuntimeConfigEntry::new(
80                CUDA_DEVICE_COUNT_KEY,
81                self.cuda_device_count.to_string(),
82                RuntimeConfigSource::Cli,
83            ),
84        ]
85        .into_iter()
86        .chain(self.selected_layer_split_plan.as_ref().map(|plan| {
87            RuntimeConfigEntry::new(
88                SELECTED_LAYER_SPLIT_PLAN_KEY,
89                plan.clone(),
90                RuntimeConfigSource::Cli,
91            )
92        }))
93        .chain(self.selected_layer_split_stages.as_ref().map(|stages| {
94            RuntimeConfigEntry::new(
95                SELECTED_LAYER_SPLIT_STAGES_KEY,
96                serde_json::to_string(stages).expect("serialize layer split stages"),
97                RuntimeConfigSource::Cli,
98            )
99        }))
100        .collect()
101    }
102
103    pub fn insert_backend_options(&self, options: &mut HashMap<String, Value>) {
104        options.insert(
105            "gpu_devices_raw".to_string(),
106            Value::String(self.raw_cli_value.clone()),
107        );
108        options.insert(
109            "requested_gpu_devices".to_string(),
110            serde_json::json!(self.requested_gpu_devices),
111        );
112        options.insert(
113            "selected_gpu_devices".to_string(),
114            serde_json::json!(self.selected_gpu_devices),
115        );
116        options.insert(
117            "selected_distributed_strategy".to_string(),
118            Value::String(self.selected_distributed_strategy.clone()),
119        );
120        options.insert(
121            "cuda_device_count".to_string(),
122            serde_json::json!(self.cuda_device_count),
123        );
124        if let Some(plan) = &self.selected_layer_split_plan {
125            options.insert(
126                "selected_layer_split_plan".to_string(),
127                Value::String(plan.clone()),
128            );
129        }
130        if let Some(stages) = &self.selected_layer_split_stages {
131            options.insert(
132                "selected_layer_split_stages".to_string(),
133                serde_json::to_value(stages).expect("serialize layer split stages"),
134            );
135        }
136    }
137}
138
139pub fn resolve_cuda_gpu_devices(
140    raw: Option<&str>,
141    device: &Device,
142) -> Result<Option<GpuDeviceSelection>> {
143    let Some(raw) = raw else {
144        return Ok(None);
145    };
146    if !matches!(device, Device::CUDA(_)) {
147        return Err(FerrumError::config(format!(
148            "--gpu-devices requires the cuda backend; selected backend device is {device}"
149        )));
150    }
151
152    let requested = parse_gpu_devices(raw)?;
153    let available = available_cuda_device_count()?;
154    validate_gpu_devices_exist(&requested, available)?;
155    let selected_distributed_strategy = if requested.len() > 1 {
156        "layer_split"
157    } else {
158        "single_gpu"
159    };
160    let selected_layer_split_plan =
161        (requested.len() > 1).then(|| even_layer_split_plan_placeholder(&requested));
162
163    Ok(Some(GpuDeviceSelection {
164        raw_cli_value: raw.to_string(),
165        requested_gpu_devices: requested.clone(),
166        selected_gpu_devices: requested,
167        cuda_device_count: available,
168        selected_distributed_strategy: selected_distributed_strategy.to_string(),
169        selected_layer_split_plan,
170        selected_layer_split_stages: None,
171    }))
172}
173
174pub fn parse_gpu_devices(raw: &str) -> Result<Vec<usize>> {
175    let raw = raw.trim();
176    if raw.is_empty() {
177        return Err(FerrumError::config("--gpu-devices cannot be empty"));
178    }
179
180    let mut seen = BTreeSet::new();
181    let mut devices = Vec::new();
182    for (idx, part) in raw.split(',').enumerate() {
183        let value = part.trim();
184        if value.is_empty() {
185            return Err(FerrumError::config(format!(
186                "--gpu-devices has an empty GPU id at position {}",
187                idx + 1
188            )));
189        }
190        if !value.chars().all(|ch| ch.is_ascii_digit()) {
191            return Err(FerrumError::config(format!(
192                "--gpu-devices value {value:?} is invalid; GPU ids must be non-negative integers"
193            )));
194        }
195        let parsed = value.parse::<usize>().map_err(|_| {
196            FerrumError::config(format!(
197                "--gpu-devices value {value:?} is invalid; GPU ids must fit in usize"
198            ))
199        })?;
200        if !seen.insert(parsed) {
201            return Err(FerrumError::config(format!(
202                "--gpu-devices contains duplicate GPU id {parsed}"
203            )));
204        }
205        devices.push(parsed);
206    }
207
208    Ok(devices)
209}
210
211pub fn validate_gpu_devices_exist(requested: &[usize], available_count: usize) -> Result<()> {
212    if available_count == 0 {
213        return Err(FerrumError::device(
214            "--gpu-devices was provided but no CUDA devices are available",
215        ));
216    }
217    for id in requested {
218        if *id >= available_count {
219            return Err(FerrumError::device(format!(
220                "--gpu-devices requested CUDA device {id}, but only {available_count} CUDA device(s) are available"
221            )));
222        }
223    }
224    Ok(())
225}
226
227pub fn join_gpu_devices(devices: &[usize]) -> String {
228    devices
229        .iter()
230        .map(|device| device.to_string())
231        .collect::<Vec<_>>()
232        .join(",")
233}
234
235fn even_layer_split_plan_placeholder(devices: &[usize]) -> String {
236    devices
237        .iter()
238        .enumerate()
239        .map(|(idx, device)| format!("stage{idx}:cuda:{device}:layers=auto"))
240        .collect::<Vec<_>>()
241        .join(";")
242}
243
244pub fn even_layer_split_plan_for_layers(devices: &[usize], num_layers: usize) -> Result<String> {
245    even_layer_split_stages_for_layers(devices, num_layers)
246        .map(|stages| format_layer_split_plan(&stages))
247}
248
249pub fn even_layer_split_stages_for_layers(
250    devices: &[usize],
251    num_layers: usize,
252) -> Result<Vec<CudaLayerSplitStage>> {
253    if devices.is_empty() {
254        return Err(FerrumError::config(
255            "layer split requires at least one CUDA device",
256        ));
257    }
258    if num_layers == 0 {
259        return Err(FerrumError::config(
260            "layer split requires a model with at least one transformer layer",
261        ));
262    }
263    if num_layers < devices.len() {
264        return Err(FerrumError::config(format!(
265            "layer split requires at least as many transformer layers ({num_layers}) as CUDA devices ({})",
266            devices.len()
267        )));
268    }
269
270    let base = num_layers / devices.len();
271    let remainder = num_layers % devices.len();
272    let mut start = 0usize;
273    let mut stages = Vec::with_capacity(devices.len());
274    for (idx, device) in devices.iter().enumerate() {
275        let count = base + usize::from(idx < remainder);
276        let end = start + count - 1;
277        stages.push(CudaLayerSplitStage {
278            stage: idx,
279            device: *device,
280            layer_start: start,
281            layer_end: end,
282        });
283        start = end + 1;
284    }
285    Ok(stages)
286}
287
288pub fn format_layer_split_plan(stages: &[CudaLayerSplitStage]) -> String {
289    stages
290        .iter()
291        .map(|stage| {
292            format!(
293                "stage{}:cuda:{}:layers={}-{}",
294                stage.stage, stage.device, stage.layer_start, stage.layer_end
295            )
296        })
297        .collect::<Vec<_>>()
298        .join(";")
299}
300
301#[cfg(feature = "cuda")]
302fn available_cuda_device_count() -> Result<usize> {
303    ferrum_kernels::cuda_device_count().map_err(FerrumError::device)
304}
305
306#[cfg(not(feature = "cuda"))]
307fn available_cuda_device_count() -> Result<usize> {
308    Err(FerrumError::unsupported(
309        "--gpu-devices requires ferrum built with CUDA support",
310    ))
311}
312
313#[cfg(test)]
314mod tests {
315    use super::*;
316
317    #[test]
318    fn parses_comma_separated_gpu_devices() {
319        assert_eq!(parse_gpu_devices("0,1, 2").unwrap(), vec![0, 1, 2]);
320    }
321
322    #[test]
323    fn rejects_duplicate_gpu_devices() {
324        let err = parse_gpu_devices("0,1,0").unwrap_err().to_string();
325        assert!(err.contains("duplicate GPU id 0"));
326    }
327
328    #[test]
329    fn rejects_negative_gpu_devices() {
330        let err = parse_gpu_devices("0,-1").unwrap_err().to_string();
331        assert!(err.contains("non-negative integers"));
332    }
333
334    #[test]
335    fn rejects_missing_gpu_device_ids() {
336        let err = parse_gpu_devices("0,,1").unwrap_err().to_string();
337        assert!(err.contains("empty GPU id"));
338    }
339
340    #[test]
341    fn validates_requested_gpu_devices_against_available_count() {
342        validate_gpu_devices_exist(&[0, 1], 2).unwrap();
343        let err = validate_gpu_devices_exist(&[2], 2).unwrap_err().to_string();
344        assert!(err.contains("only 2 CUDA device"));
345    }
346
347    #[test]
348    fn layer_split_plan_assigns_contiguous_model_layers() {
349        let plan = even_layer_split_plan_for_layers(&[0, 1], 80).unwrap();
350        assert_eq!(plan, "stage0:cuda:0:layers=0-39;stage1:cuda:1:layers=40-79");
351    }
352
353    #[test]
354    fn layer_split_plan_distributes_remainder_to_earlier_devices() {
355        let plan = even_layer_split_plan_for_layers(&[0, 1, 2], 10).unwrap();
356        assert_eq!(
357            plan,
358            "stage0:cuda:0:layers=0-3;stage1:cuda:1:layers=4-6;stage2:cuda:2:layers=7-9"
359        );
360    }
361
362    #[test]
363    fn layer_split_plan_rejects_more_devices_than_layers() {
364        let err = even_layer_split_plan_for_layers(&[0, 1], 1)
365            .unwrap_err()
366            .to_string();
367        assert!(err.contains("at least as many transformer layers"));
368    }
369
370    #[test]
371    fn runtime_entries_record_raw_requested_selected_and_strategy() {
372        let selection = GpuDeviceSelection {
373            raw_cli_value: "1".to_string(),
374            requested_gpu_devices: vec![1],
375            selected_gpu_devices: vec![1],
376            cuda_device_count: 2,
377            selected_distributed_strategy: "single_gpu".to_string(),
378            selected_layer_split_plan: None,
379            selected_layer_split_stages: None,
380        };
381        let snapshot =
382            ferrum_types::RuntimeConfigSnapshot::from_entries(selection.runtime_config_entries());
383        let entry = |key: &str| {
384            snapshot
385                .entries
386                .iter()
387                .find(|entry| entry.key == key)
388                .unwrap_or_else(|| panic!("missing {key}"))
389        };
390
391        assert_eq!(entry("FERRUM_BACKEND").effective_value, "cuda");
392        assert_eq!(entry(GPU_DEVICES_RAW_KEY).effective_value, "1");
393        assert_eq!(entry(REQUESTED_GPU_DEVICES_KEY).effective_value, "1");
394        assert_eq!(entry(SELECTED_GPU_DEVICES_KEY).effective_value, "1");
395        assert_eq!(
396            entry(SELECTED_DISTRIBUTED_STRATEGY_KEY).effective_value,
397            "single_gpu"
398        );
399        assert_eq!(entry(CUDA_DEVICE_COUNT_KEY).effective_value, "2");
400    }
401
402    #[test]
403    fn runtime_entries_record_layer_split_plan_for_multi_gpu() {
404        let mut selection = GpuDeviceSelection {
405            raw_cli_value: "0,1".to_string(),
406            requested_gpu_devices: vec![0, 1],
407            selected_gpu_devices: vec![0, 1],
408            cuda_device_count: 2,
409            selected_distributed_strategy: "layer_split".to_string(),
410            selected_layer_split_plan: Some(even_layer_split_plan_placeholder(&[0, 1])),
411            selected_layer_split_stages: None,
412        };
413        assert!(selection.apply_model_layer_count(80).unwrap());
414        let snapshot =
415            ferrum_types::RuntimeConfigSnapshot::from_entries(selection.runtime_config_entries());
416        let entry = |key: &str| {
417            snapshot
418                .entries
419                .iter()
420                .find(|entry| entry.key == key)
421                .unwrap_or_else(|| panic!("missing {key}"))
422        };
423
424        assert_eq!(entry(REQUESTED_GPU_DEVICES_KEY).effective_value, "0,1");
425        assert_eq!(entry(SELECTED_GPU_DEVICES_KEY).effective_value, "0,1");
426        assert_eq!(
427            entry(SELECTED_DISTRIBUTED_STRATEGY_KEY).effective_value,
428            "layer_split"
429        );
430        assert_eq!(
431            entry(SELECTED_LAYER_SPLIT_PLAN_KEY).effective_value,
432            "stage0:cuda:0:layers=0-39;stage1:cuda:1:layers=40-79"
433        );
434        let stages: serde_json::Value =
435            serde_json::from_str(&entry(SELECTED_LAYER_SPLIT_STAGES_KEY).effective_value).unwrap();
436        assert_eq!(
437            stages,
438            serde_json::json!([
439                {"stage": 0, "device": 0, "layer_start": 0, "layer_end": 39},
440                {"stage": 1, "device": 1, "layer_start": 40, "layer_end": 79}
441            ])
442        );
443    }
444}