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])
| 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. |