1use crate::cluster::resolve_local_hostname;
23use crate::config::{
24 ClusterConfig, ClusterController, ClusterWorker, DEFAULT_CONTROLLER_PORT, LocalDevices,
25};
26
27#[derive(Debug, Clone, PartialEq, Eq)]
33pub enum GpusSpec {
34 All,
37 List(Vec<u8>),
39}
40
41impl GpusSpec {
42 pub fn parse(raw: &str) -> Result<Self, String> {
45 let trimmed = raw.trim();
46 if trimmed.is_empty() {
47 return Err("--gpus requires a value (e.g. `--gpus 0,1` or `--gpus all`)".to_string());
48 }
49 if trimmed.eq_ignore_ascii_case("all") {
50 return Ok(GpusSpec::All);
51 }
52 let mut out = Vec::new();
53 for part in trimmed.split(',') {
54 let p = part.trim();
55 if p.is_empty() {
56 return Err(format!("--gpus: empty entry in {trimmed:?}"));
57 }
58 let idx: u8 = p
59 .parse()
60 .map_err(|e| format!("--gpus: cannot parse {p:?} as device index: {e}"))?;
61 out.push(idx);
62 }
63 let mut sorted = out.clone();
64 sorted.sort_unstable();
65 for win in sorted.windows(2) {
66 if win[0] == win[1] {
67 return Err(format!(
68 "--gpus: duplicate device index {} in {trimmed:?}",
69 win[0]
70 ));
71 }
72 }
73 Ok(GpusSpec::List(out))
74 }
75
76 pub fn resolve(&self) -> Result<Vec<u8>, String> {
82 match self {
83 GpusSpec::List(v) => Ok(v.clone()),
84 GpusSpec::All => {
85 let devices = local_gpu_count().map_err(|e| format!("--gpus all: {e}"))?;
91 if devices > u8::MAX as usize {
92 return Err(format!(
93 "--gpus all: {devices} GPUs detected, which exceeds \
94 the supported device-index range (0..255). Specify \
95 devices explicitly via --gpus."
96 ));
97 }
98 Ok((0u8..devices as u8).collect())
99 }
100 }
101 }
102}
103
104pub fn local_gpu_count() -> Result<usize, String> {
117 flodl_hw::survey().require_devices().map(|d| d.len())
118}
119
120pub fn synthesize_local_cluster(devices: &[u8]) -> Result<ClusterConfig, String> {
132 if devices.is_empty() {
133 return Err("synthesize_local_cluster: device list is empty".to_string());
134 }
135 let hostname = resolve_local_hostname();
136 let path = std::env::current_dir()
137 .map(|p| p.to_string_lossy().into_owned())
138 .map_err(|e| format!("synthesize_local_cluster: cannot read current_dir: {e}"))?;
139 let port = std::env::var("FLODL_CONTROLLER_PORT")
140 .ok()
141 .and_then(|s| s.parse::<u16>().ok())
142 .unwrap_or(DEFAULT_CONTROLLER_PORT);
143
144 Ok(ClusterConfig {
145 controller: ClusterController {
146 host: "127.0.0.1".to_string(),
147 port,
148 path: path.clone(),
149 docker: None,
150 arch: None,
151 data_path: None,
152 join: None,
153 },
154 workers: vec![ClusterWorker {
155 host: hostname,
156 ranks: (0..devices.len()).collect(),
157 local_devices: LocalDevices::Explicit(devices.to_vec()),
158 nccl_socket_ifname: "lo".to_string(),
159 path,
160 ssh: None,
161 tunnel: false,
162 arch: None,
163 data_path: None,
164 gpu_ram_share: None,
165 docker: None,
166 env: std::collections::BTreeMap::new(),
167 }],
168 env: std::collections::BTreeMap::new(),
169 gpu_ram_share: None,
170 })
171}
172
173pub unsafe fn apply_cuda_visible_devices(devices: &[u8]) {
189 let joined = devices
190 .iter()
191 .map(|d| d.to_string())
192 .collect::<Vec<_>>()
193 .join(",");
194 for key in ["CUDA_VISIBLE_DEVICES", "HIP_VISIBLE_DEVICES"] {
200 if joined.is_empty() {
201 unsafe { std::env::remove_var(key) };
202 } else {
203 unsafe { std::env::set_var(key, &joined) };
204 }
205 }
206}
207
208#[cfg(test)]
209mod tests {
210 use super::*;
211
212 #[test]
213 fn parse_all_case_insensitive() {
214 assert_eq!(GpusSpec::parse("all").unwrap(), GpusSpec::All);
215 assert_eq!(GpusSpec::parse("ALL").unwrap(), GpusSpec::All);
216 assert_eq!(GpusSpec::parse("All").unwrap(), GpusSpec::All);
217 }
218
219 #[test]
220 fn parse_single_index() {
221 assert_eq!(GpusSpec::parse("0").unwrap(), GpusSpec::List(vec![0]));
222 assert_eq!(GpusSpec::parse("3").unwrap(), GpusSpec::List(vec![3]));
223 }
224
225 #[test]
226 fn parse_multiple_indices() {
227 assert_eq!(
228 GpusSpec::parse("0,1,2").unwrap(),
229 GpusSpec::List(vec![0, 1, 2])
230 );
231 assert_eq!(GpusSpec::parse("3,1").unwrap(), GpusSpec::List(vec![3, 1]));
232 }
233
234 #[test]
235 fn parse_tolerates_whitespace() {
236 assert_eq!(
237 GpusSpec::parse(" 0 , 1 ").unwrap(),
238 GpusSpec::List(vec![0, 1])
239 );
240 assert_eq!(GpusSpec::parse(" all ").unwrap(), GpusSpec::All);
241 }
242
243 #[test]
244 fn parse_rejects_empty() {
245 let err = GpusSpec::parse("").unwrap_err();
246 assert!(err.contains("--gpus requires a value"), "got: {err}");
247 let err = GpusSpec::parse(" ").unwrap_err();
248 assert!(err.contains("--gpus requires a value"), "got: {err}");
249 }
250
251 #[test]
252 fn parse_rejects_empty_entry() {
253 let err = GpusSpec::parse("0,,1").unwrap_err();
254 assert!(err.contains("empty entry"), "got: {err}");
255 let err = GpusSpec::parse(",0").unwrap_err();
256 assert!(err.contains("empty entry"), "got: {err}");
257 }
258
259 #[test]
260 fn parse_rejects_non_numeric() {
261 let err = GpusSpec::parse("0,abc").unwrap_err();
262 assert!(err.contains("cannot parse"), "got: {err}");
263 assert!(err.contains("abc"), "got: {err}");
264 }
265
266 #[test]
267 fn parse_rejects_duplicates() {
268 let err = GpusSpec::parse("0,1,0").unwrap_err();
269 assert!(err.contains("duplicate"), "got: {err}");
270 assert!(err.contains("0"), "got: {err}");
271 }
272
273 #[test]
274 fn resolve_list_returns_verbatim() {
275 let r = GpusSpec::List(vec![3, 1]).resolve().unwrap();
276 assert_eq!(r, vec![3, 1]);
277 }
278
279 #[test]
280 fn synthesize_local_cluster_basic_shape() {
281 let c = synthesize_local_cluster(&[0, 1]).unwrap();
283 assert_eq!(c.controller.host, "127.0.0.1");
284 assert_eq!(c.workers.len(), 1);
285 let w = &c.workers[0];
286 assert_eq!(w.ranks, vec![0, 1]);
287 assert_eq!(w.local_devices, LocalDevices::Explicit(vec![0, 1]));
288 assert_eq!(w.nccl_socket_ifname, "lo");
289 assert!(w.arch.is_none());
290 assert!(w.ssh.is_none());
291 assert!(!w.host.trim().is_empty(), "hostname must be non-empty");
292 assert!(!w.path.trim().is_empty(), "path must be non-empty");
293 }
294
295 #[test]
296 fn synthesize_local_cluster_validates() {
297 let c = synthesize_local_cluster(&[0, 1]).unwrap();
300 c.validate()
301 .expect("synthesized cluster must pass validate");
302 }
303
304 #[test]
305 fn synthesize_local_cluster_single_device() {
306 let c = synthesize_local_cluster(&[2]).unwrap();
309 c.validate()
310 .expect("single-device synthesized config validates");
311 assert_eq!(c.workers[0].ranks, vec![0]);
312 assert_eq!(c.workers[0].local_devices, LocalDevices::Explicit(vec![2]));
313 }
314
315 #[test]
316 fn synthesize_local_cluster_rejects_empty() {
317 let err = synthesize_local_cluster(&[]).unwrap_err();
318 assert!(err.contains("empty"), "got: {err}");
319 }
320
321 #[test]
322 fn synthesize_local_cluster_respects_controller_port_env() {
323 unsafe {
329 std::env::set_var("FLODL_CONTROLLER_PORT", "31415");
330 }
331 let c = synthesize_local_cluster(&[0]).unwrap();
332 unsafe {
333 std::env::remove_var("FLODL_CONTROLLER_PORT");
334 }
335 assert_eq!(c.controller.port, 31415);
336 }
337}