@@ -12,7 +12,7 @@ use tako::format_comma_delimited;
1212use tako:: internal:: has_unique_elements;
1313use tako:: resources:: {
1414 ResourceDescriptorItem , ResourceDescriptorKind , ResourceIndex , ResourceLabel ,
15- GPU_RESOURCE_NAME , MEM_RESOURCE_NAME ,
15+ AMD_GPU_RESOURCE_NAME , MEM_RESOURCE_NAME , NVIDIA_GPU_RESOURCE_NAME ,
1616} ;
1717
1818pub fn detect_cpus ( ) -> anyhow:: Result < ResourceDescriptorKind > {
@@ -57,17 +57,17 @@ pub fn detect_additional_resources(items: &mut Vec<ResourceDescriptorItem>) -> a
5757 let has_resource =
5858 |items : & [ ResourceDescriptorItem ] , name : & str | items. iter ( ) . any ( |x| x. name == name) ;
5959
60- if !has_resource ( items, GPU_RESOURCE_NAME ) {
61- if let Some ( gpus ) = detect_gpus_from_env ( ) {
60+ if !has_resource ( items, NVIDIA_GPU_RESOURCE_NAME ) {
61+ if let Some ( detected ) = detect_gpus_from_env ( ) {
6262 items. push ( ResourceDescriptorItem {
63- name : GPU_RESOURCE_NAME . to_string ( ) ,
64- kind : gpus ,
63+ name : detected . resource_name . to_string ( ) ,
64+ kind : detected . resource ,
6565 } ) ;
66- } else if let Ok ( count) = read_linux_gpu_count ( ) {
66+ } else if let Ok ( count) = read_nvidia_linux_gpu_count ( ) {
6767 if count > 0 {
6868 log:: info!( "Detected {} GPUs from procs" , count) ;
6969 items. push ( ResourceDescriptorItem {
70- name : GPU_RESOURCE_NAME . to_string ( ) ,
70+ name : NVIDIA_GPU_RESOURCE_NAME . to_string ( ) ,
7171 kind : ResourceDescriptorKind :: simple_indices ( count as u32 ) ,
7272 } ) ;
7373 }
@@ -86,15 +86,20 @@ pub fn detect_additional_resources(items: &mut Vec<ResourceDescriptorItem>) -> a
8686 Ok ( ( ) )
8787}
8888
89- pub const GPU_ENV_KEYS : & [ & str ; 3 ] = & [
90- "CUDA_VISIBLE_DEVICES" ,
91- "HIP_VISIBLE_DEVICES" ,
92- "ROCR_VISIBLE_DEVICES" ,
89+ pub const GPU_ENV_KEYS : & [ ( & str , & str ) ; 3 ] = & [
90+ ( "CUDA_VISIBLE_DEVICES" , NVIDIA_GPU_RESOURCE_NAME ) ,
91+ ( "HIP_VISIBLE_DEVICES" , AMD_GPU_RESOURCE_NAME ) ,
92+ ( "ROCR_VISIBLE_DEVICES" , AMD_GPU_RESOURCE_NAME ) ,
9393] ;
9494
95+ struct DetectedGpu {
96+ resource_name : & ' static str ,
97+ resource : ResourceDescriptorKind ,
98+ }
99+
95100/// Tries to detect available GPUs from one of the `GPU_ENV_KEYS` environment variables.
96- fn detect_gpus_from_env ( ) -> Option < ResourceDescriptorKind > {
97- GPU_ENV_KEYS . iter ( ) . find_map ( | env_key| {
101+ fn detect_gpus_from_env ( ) -> Option < DetectedGpu > {
102+ for ( env_key, resource_name ) in GPU_ENV_KEYS {
98103 if let Ok ( devices_str) = std:: env:: var ( env_key) {
99104 if let Ok ( devices) = parse_comma_separated_values ( & devices_str) {
100105 log:: info!(
@@ -108,15 +113,18 @@ fn detect_gpus_from_env() -> Option<ResourceDescriptorKind> {
108113
109114 let list =
110115 ResourceDescriptorKind :: list ( devices) . expect ( "List values were not unique" ) ;
111- return Some ( list) ;
116+ return Some ( DetectedGpu {
117+ resource_name,
118+ resource : list,
119+ } ) ;
112120 }
113121 }
114- None
115- } )
122+ }
123+ None
116124}
117125
118126/// Try to find out how many Nvidia GPUs are available on the current node.
119- fn read_linux_gpu_count ( ) -> anyhow:: Result < usize > {
127+ fn read_nvidia_linux_gpu_count ( ) -> anyhow:: Result < usize > {
120128 Ok ( std:: fs:: read_dir ( "/proc/driver/nvidia/gpus" ) ?. count ( ) )
121129}
122130
0 commit comments