MCPcopy Create free account
hub / github.com/apple/axlearn / dispatch_global_batch

Method dispatch_global_batch

axlearn/common/input_base.py:227–276  ·  view source on GitHub ↗

Converts a global physical batch to a global logical batch. The leaves of the output logical batch are partitioned across `batch_axis_names` along the 0th (batch) dimension. This should be invoked from within `pjit` so that the sharding constraints can be applied. I

(self, global_physical_batch: Nested[Tensor])

Source from the content-addressed store, hash-verified

225 yield input_batch
226
227 def dispatch_global_batch(self, global_physical_batch: Nested[Tensor]) -> Nested[Tensor]:
228 """Converts a global physical batch to a global logical batch.
229
230 The leaves of the output logical batch are partitioned across `batch_axis_names` along the
231 0th (batch) dimension. This should be invoked from within `pjit` so that the sharding
232 constraints can be applied.
233
234 If `cfg.input_partitioner` is not None, it will be applied to each logical batch after
235 constraining `batch_axis_names`.
236 """
237
238 def constrain_batch_axis(path: str, value: Tensor):
239 mesh = thread_resources.env.physical_mesh
240 batch_partitions = math.prod(
241 mesh.shape[axis] for axis in jax.tree.leaves(self._partition_spec[0])
242 )
243 # Warn if an invalid constraint is applied, since by default this can silently be
244 # ignored, potentially leading to unexpected OOMs.
245 if value.shape[0] % batch_partitions != 0:
246 logging.warning(
247 "Attempting to constrain path=%s (with batch dim %d) over %d partitions (%s).",
248 path,
249 value.shape[0],
250 batch_partitions,
251 self._partition_spec,
252 )
253 return maybe_shard(value, self._partition_spec)
254
255 if "input_dispatcher" in self.children:
256 global_logical_batch = self.input_dispatcher.physical_to_logical_batch(
257 jax.tree.map(
258 constrain_batch_axis,
259 tree_paths(global_physical_batch),
260 global_physical_batch,
261 )
262 )
263 else:
264 global_logical_batch = dispatch_input_batch(
265 global_physical_batch, batch_axis_names=self._partition_spec[0]
266 )
267
268 global_logical_batch = jax.tree.map(
269 constrain_batch_axis, tree_paths(global_logical_batch), global_logical_batch
270 )
271
272 # Further constrain based on user-configured partitioning rules.
273 if self._input_partitioner is not None:
274 global_logical_batch = self._input_partitioner(global_logical_batch)
275
276 return global_logical_batch
277
278 def element_spec(self) -> Nested[jax.ShapeDtypeStruct]:
279 """Returns the per-feed logical batch spec.

Callers 4

test_dispatch_tpuMethod · 0.80
fnFunction · 0.80
_train_stepMethod · 0.80

Calls 4

tree_pathsFunction · 0.90
dispatch_input_batchFunction · 0.90
mapMethod · 0.80

Tested by 2

test_dispatch_tpuMethod · 0.64
fnFunction · 0.64