@@ -475,14 +475,14 @@ def down_nodes_notify_jobs(nodes: List[str], reason: str, resume_data: Optional[
475475
476476
477477
478- def create_placement_request (pg_name : str , region : str , max_distance : Optional [int ], gpu_topology : Optional [str ]):
478+ def create_placement_request (pg_name : str , region : str , max_distance : Optional [int ], accelerator_topology : Optional [str ]):
479479 config = {
480480 "name" : pg_name ,
481481 "region" : region ,
482482 "groupPlacementPolicy" : {
483483 "collocation" : "COLLOCATED" ,
484484 "maxDistance" : max_distance ,
485- "gpuTopology" : gpu_topology
485+ "gpuTopology" : accelerator_topology ,
486486 },
487487 }
488488
@@ -554,14 +554,46 @@ def _allocate_nodes_to_placements(nodes: List[str], excl_job_id:Optional[int], l
554554
555555 return placements
556556
557+ def calculate_hosts_per_topo (accelerator_topology : str , machine_type : NSDict ) -> int :
558+ # Calculate total number of hosts per topology (Assumes format: '1x72')
559+ try :
560+ top_split = [int (x ) for x in accelerator_topology .split ("x" )]
561+ except Exception as e :
562+ log .error (f"Accelerator topology { accelerator_topology } is formatted incorrectly." )
563+ raise e
564+
565+ if len (machine_type .accelerators ) == 0 :
566+ gpus_per_machine = 0
567+ else :
568+ gpus_per_machine = machine_type .accelerators [0 ].count
569+
570+ if len (top_split ) != 2 :
571+ log .error (f"Accelerator topology { accelerator_topology } is formatted incorrectly." )
572+ elif top_split [0 ] <= 0 or top_split [1 ] <= 0 :
573+ log .error (f"Accelerator topology { accelerator_topology } is formatted incorrectly." )
574+ elif gpus_per_machine <= 0 :
575+ log .error (f"The machine type has no accelerators. Cannot use accelerator topology { accelerator_topology } ." )
576+ elif top_split [1 ] % gpus_per_machine :
577+ log .error (f"The GPU count { gpus_per_machine } per node is not a factor of the accelerator topology { accelerator_topology } " )
578+
579+ return (top_split [0 ] * top_split [1 ]) // gpus_per_machine
580+
557581def calculate_chunk_size (nodeset : NSDict , lkp : util .Lookup ) -> int :
558- # Calculates the chunk size based on max distance value received
559- machine_type = lkp .template_info (nodeset .instance_template ).machine_type .family
582+ # Calculates the chunk size based on max distance value received or accelerator topology
583+ # Assuming nodeset is not tpu
584+ machine_type = lkp .template_info (nodeset .instance_template ).machine_type
560585 max_distance = nodeset .placement_max_distance
586+ accelerator_topology = nodeset .accelerator_topology
587+
588+ # Look for accelerator topology first
589+ if accelerator_topology :
590+ hosts_per_topo = calculate_hosts_per_topo (accelerator_topology , machine_type )
591+ return hosts_per_topo
592+
561593 if max_distance == 1 :
562594 return 22
563595 elif max_distance == 2 :
564- if machine_type .startswith ("a3" ):
596+ if machine_type .family . startswith ("a3" ):
565597 return 256
566598 else :
567599 return 150
@@ -574,7 +606,7 @@ def create_nodeset_placements(nodes: List[str], excl_job_id:Optional[int], lkp:
574606 placements = _allocate_nodes_to_placements (nodes , excl_job_id , lkp )
575607 region = lkp .node_region (nodes [0 ])
576608 max_distance = lkp .node_nodeset (nodes [0 ]).get ('placement_max_distance' )
577- gpu_topology = lkp .node_nodeset ( nodes [ 0 ]). get ( 'accelerator_topology' ) if not lkp .node_is_tpu (nodes [0 ]) else None
609+ accelerator_topology = lkp .nodeset_accelerator_topology ( lkp .node_nodeset_name (nodes [0 ]))
578610
579611 if log .isEnabledFor (logging .DEBUG ):
580612 debug_p = {p .placement : to_hostlist (p .nodes ) for p in placements }
@@ -583,7 +615,7 @@ def create_nodeset_placements(nodes: List[str], excl_job_id:Optional[int], lkp:
583615 )
584616
585617 requests = {
586- p .placement : create_placement_request (p .placement , region , max_distance , gpu_topology ) for p in placements if p .placement
618+ p .placement : create_placement_request (p .placement , region , max_distance , accelerator_topology ) for p in placements if p .placement
587619 }
588620 if not requests :
589621 return placements
0 commit comments