4949from axlearn .common .config import REQUIRED , Required , config_class , maybe_set_config
5050from axlearn .common .input_dispatch import BaseInputDispatcher , _validate_logical_feed_shapes
5151from axlearn .common .module import Module
52- from axlearn .common .utils import Nested , Tensor
52+ from axlearn .common .utils import Nested , Tensor , live_devices
5353
5454
5555class ElasticSpmdInputDispatcher (BaseInputDispatcher ):
@@ -72,7 +72,10 @@ class Config(BaseInputDispatcher.Config):
7272
7373 @property
7474 def is_in_elastic_mode (self ) -> bool :
75+ print ("In is_in_elastic_mode by lkolluru" )
7576 cfg = self .config
77+ print ("cfg.num_max_slices by lkolluru: " , cfg .num_max_slices )
78+ print ("slice_count by lkolluru: " , slice_count ())
7679 if cfg .num_max_slices is None :
7780 return False
7881 else :
@@ -165,11 +168,23 @@ def fid2pids(feed_id):
165168 else :
166169 self ._feed_logical_batch_size = cfg .global_logical_batch_size // self .feed_count
167170
171+ adjusted_device_physical_batch_size = math .ceil (
172+ self ._device_physical_batch_size * (cfg .num_max_slices / slice_count ())
173+ )
174+ print (
175+ "adjusted_device_physical_batch_size outside elastic by lkolluru: " ,
176+ adjusted_device_physical_batch_size ,
177+ )
178+
168179 if self .is_in_elastic_mode :
169180 print (" In elastic mode lkolluru" )
170181 adjusted_device_physical_batch_size = math .ceil (
171182 self ._device_physical_batch_size * (cfg .num_max_slices / slice_count ())
172183 )
184+ print (
185+ "adjusted_device_physical_batch_size inside elastic by lkolluru: " ,
186+ adjusted_device_physical_batch_size ,
187+ )
173188 padding_per_device = (
174189 adjusted_device_physical_batch_size - self ._device_physical_batch_size
175190 )
@@ -408,9 +423,14 @@ def _padded_select(path, x, y):
408423
409424def slice_count () -> int :
410425 """Returns the number of slices."""
411- slice_cnt_val = len (set (d .slice_index for d in jax .devices () if hasattr (d , "slice_index" ))) or 1
426+ # slice_cnt_val=len(set(d.slice_index for d in jax.devices() if hasattr(d, "slice_index"))) or 1
427+ # print("slice_count by lkolluru: ", slice_cnt_val)
428+ # return len(set(d.slice_index for d in jax.devices() if hasattr(d, "slice_index"))) or 1
429+ slice_cnt_val = (
430+ len (set (d .slice_index for d in live_devices () if hasattr (d , "slice_index" ))) or 1
431+ )
412432 print ("slice_count by lkolluru: " , slice_cnt_val )
413- return len (set (d .slice_index for d in jax . devices () if hasattr (d , "slice_index" ))) or 1
433+ return len (set (d .slice_index for d in live_devices () if hasattr (d , "slice_index" ))) or 1
414434
415435
416436def process_count_per_slice () -> int :
@@ -419,7 +439,8 @@ def process_count_per_slice() -> int:
419439 len (
420440 set (
421441 d .process_index
422- for d in jax .devices ()
442+ # for d in jax.devices()
443+ for d in live_devices ()
423444 if hasattr (d , "slice_index" ) and d .slice_index == 0
424445 )
425446 )
0 commit comments