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 , live_slice_indices
53+
5354
5455
5556class ElasticSpmdInputDispatcher (BaseInputDispatcher ):
@@ -72,13 +73,20 @@ class Config(BaseInputDispatcher.Config):
7273
7374 @property
7475 def is_in_elastic_mode (self ) -> bool :
76+ print ("In is_in_elastic_mode by lkolluru" )
7577 cfg = self .config
78+ print ("cfg.num_max_slices by lkolluru: " , cfg .num_max_slices )
79+ print ("slice_count by lkolluru: " , slice_count ())
7680 if cfg .num_max_slices is None :
7781 return False
7882 else :
83+ print (
84+ f"Live_slice_count by lkolluru: { live_slice_indices ()} , slices_cnt: { len (live_slice_indices ())} " ,
85+ live_slice_indices (),len (live_slice_indices ())
86+ )
7987 if slice_count () < cfg .num_max_slices :
8088 return True
81- elif slice_count () = = cfg .num_max_slices :
89+ elif slice_count () > = cfg .num_max_slices :
8290 return False
8391 else :
8492 # TODO (jtian22): consider supporting scaling up in the future.
@@ -92,6 +100,7 @@ def __init__(self, cfg: Config, *, parent: Optional[Module]):
92100 cfg : ElasticSpmdInputDispatcher .Config = self .config
93101
94102 mesh = thread_resources .env .physical_mesh
103+ print ("physical mesh by lkolluru: " , mesh )
95104 if mesh .empty :
96105 raise ValueError ("Expected to be initialized within the context of a mesh." )
97106
@@ -117,6 +126,7 @@ def __init__(self, cfg: Config, *, parent: Optional[Module]):
117126 )
118127 if self .is_in_elastic_mode :
119128 num_partitions = num_partitions // slice_count () * cfg .num_max_slices
129+ print (f"num_partitions by lkolluru: { num_partitions } and live_devices: { slice_count ()} " )
120130
121131 if cfg .global_logical_batch_size % num_partitions != 0 :
122132 raise ValueError (
@@ -125,6 +135,7 @@ def __init__(self, cfg: Config, *, parent: Optional[Module]):
125135 )
126136
127137 self ._device_physical_batch_size = cfg .global_logical_batch_size // num_partitions
138+ print ("device_physical_batch_size on init by lkolluru: " , self ._device_physical_batch_size )
128139
129140 # Infer the physical feeds and feed index along dim=0.
130141 _ , _ , pid2fid = get_process_index_and_count_and_mapping (
@@ -156,14 +167,33 @@ def fid2pids(feed_id):
156167
157168 self .feed_count = len (set (pid2fid .values ())) // slice_count () * cfg .num_max_slices
158169 self .feed_index = pid2fid [jax .process_index ()]
170+ print ("feed_count by lkolluru: " , self .feed_count )
171+ print ("global_logical_batch_size by lkolluru: " , cfg .global_logical_batch_size )
172+
173+ # assert cfg.global_logical_batch_size % self.feed_count == 0
174+ if self .feed_count == 0 :
175+ self ._feed_logical_batch_size = cfg .global_logical_batch_size
176+ else :
177+ self ._feed_logical_batch_size = cfg .global_logical_batch_size // self .feed_count
159178
160- assert cfg .global_logical_batch_size % self .feed_count == 0
161- self ._feed_logical_batch_size = cfg .global_logical_batch_size // self .feed_count
179+ adjusted_device_physical_batch_size = math .ceil (
180+ self ._device_physical_batch_size * (cfg .num_max_slices / slice_count ())
181+ )
182+ print (
183+ "adjusted_device_physical_batch_size outside elastic by lkolluru: " ,
184+ adjusted_device_physical_batch_size ,
185+ )
162186
163187 if self .is_in_elastic_mode :
188+ print (" In elastic mode lkolluru" )
189+ # I think this should be len(live_slice_indices())....
164190 adjusted_device_physical_batch_size = math .ceil (
165191 self ._device_physical_batch_size * (cfg .num_max_slices / slice_count ())
166192 )
193+ print (
194+ "adjusted_device_physical_batch_size inside elastic by lkolluru: " ,
195+ adjusted_device_physical_batch_size ,
196+ )
167197 padding_per_device = (
168198 adjusted_device_physical_batch_size - self ._device_physical_batch_size
169199 )
@@ -402,7 +432,14 @@ def _padded_select(path, x, y):
402432
403433def slice_count () -> int :
404434 """Returns the number of slices."""
405- return len (set (d .slice_index for d in jax .devices () if hasattr (d , "slice_index" ))) or 1
435+ # slice_cnt_val=len(set(d.slice_index for d in jax.devices() if hasattr(d, "slice_index"))) or 1
436+ # print("slice_count by lkolluru: ", slice_cnt_val)
437+ # return len(set(d.slice_index for d in jax.devices() if hasattr(d, "slice_index"))) or 1
438+ slice_cnt_val = (
439+ len (set (d .slice_index for d in live_devices () if hasattr (d , "slice_index" ))) or 1
440+ )
441+ print ("slice_count by lkolluru: " , slice_cnt_val )
442+ return len (set (d .slice_index for d in live_devices () if hasattr (d , "slice_index" ))) or 1
406443
407444
408445def process_count_per_slice () -> int :
@@ -411,7 +448,8 @@ def process_count_per_slice() -> int:
411448 len (
412449 set (
413450 d .process_index
414- for d in jax .devices ()
451+ # for d in jax.devices()
452+ for d in live_devices ()
415453 if hasattr (d , "slice_index" ) and d .slice_index == 0
416454 )
417455 )
@@ -502,6 +540,8 @@ def get_process_index_and_count_and_mapping(
502540 # compatible with any mesh with num_devices.
503541 device_map = tensor_sharding .devices_indices_map ((tensor_sharding .num_devices ,) * ndims )
504542
543+ print ("device_map by lkolluru: " , device_map )
544+
505545 # Get the slices for 'dim' for all devices.
506546 global_slice = {k : v [dim ] for k , v in device_map .items ()}
507547
@@ -516,11 +556,15 @@ def get_process_index_and_count_and_mapping(
516556 process_to_slice [d .process_index ].add (key )
517557 all_slices .add (key )
518558
559+ print ("process_to_slice by lkolluru: " , process_to_slice )
560+
519561 # Get the set of slices for the current process which we will use to compute
520562 # the index of the current process.
521563 current_pid = next (iter (tensor_sharding .addressable_devices )).process_index
522564 addressable_slices = frozenset (process_to_slice [current_pid ])
523565
566+ print ("addressable_slices by lkolluru: " , addressable_slices )
567+
524568 # Verify that all processes have the same number of slices.
525569 slices_per_process = len (addressable_slices )
526570 if any (len (x ) != slices_per_process for x in process_to_slice .values ()):
@@ -529,6 +573,7 @@ def get_process_index_and_count_and_mapping(
529573 "different number of slices."
530574 )
531575 unique_processes = list ({frozenset (x ) for x in process_to_slice .values ()})
576+ print ("unique_processes by lkolluru: " , unique_processes )
532577
533578 # After removing duplicate processes all unique slices should
534579 # cover the dimension exactly once. If they don't it means that
@@ -540,6 +585,8 @@ def get_process_index_and_count_and_mapping(
540585 # !!! patch begin
541586 pid2fid = {}
542587 for pid , _ in process_to_slice .items ():
588+ print ("pid by lkolluru: " , pid )
543589 pid2fid [pid ] = unique_processes .index (frozenset (process_to_slice [pid ]))
544590 # !!! patch end
591+ print ("pid2fid by lkolluru: " , pid2fid )
545592 return feed_index , feed_count , pid2fid
0 commit comments