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}