Class for running a pjit-lowered model's method on a batch of data.
| 44 | |
| 45 | # pylint: disable-next=too-few-public-methods |
| 46 | class MethodRunner: |
| 47 | """Class for running a pjit-lowered model's method on a batch of data.""" |
| 48 | |
| 49 | def __init__( |
| 50 | self, |
| 51 | *, |
| 52 | prng_key: Tensor, |
| 53 | mesh: jax.sharding.Mesh, |
| 54 | input_batch_partition_spec: DataPartitionType, |
| 55 | jit_run_on_batch: Callable[ |
| 56 | [Tensor, NestedTensor], |
| 57 | tuple[Tensor, NestedTensor, NestedTensor, NestedTensor], |
| 58 | ], |
| 59 | ): |
| 60 | """Initializes MethodRunner object. |
| 61 | |
| 62 | Args: |
| 63 | prng_key: the random key used for the first run, then updated by each call. |
| 64 | mesh: mesh to be used during method running, same as the one used for pjit. |
| 65 | input_batch_partition_spec: partition spec for input batches. |
| 66 | jit_run_on_batch: callable which takes prng key, input batch and outputs |
| 67 | updated prng key, outputs, summaries, and module outputs. |
| 68 | """ |
| 69 | self._prng_key = prng_key |
| 70 | self._mesh = mesh |
| 71 | self._input_batch_partition_spec = input_batch_partition_spec |
| 72 | self._jit_run_on_batch = jit_run_on_batch |
| 73 | |
| 74 | @dataclass(frozen=True) |
| 75 | class Output: |
| 76 | """Output class of MethodRunner.""" |
| 77 | |
| 78 | # Output batch as a partitioned global array. |
| 79 | output_batch: NestedTensor |
| 80 | # Input batch as a partitioned global array. |
| 81 | input_batch: NestedTensor |
| 82 | # Summaries. |
| 83 | summaries: NestedTensor |
| 84 | # Module outputs. |
| 85 | module_outputs: NestedTensor |
| 86 | |
| 87 | def __call__(self, input_batch: NestedTensor) -> Output: |
| 88 | """Computes outputs and summaries for the given input. |
| 89 | |
| 90 | The convention is for input_batch to be global arrays. |
| 91 | Output batches are global arrays. |
| 92 | This symmetry in global arrays for both input and output batches allows |
| 93 | users to chain multiple `MethodRunner`s together without extra host-device transfer. |
| 94 | If the input_batch is host-local, it will be automatically converted to |
| 95 | global input batch for ease-of-use. |
| 96 | |
| 97 | Args: |
| 98 | input_batch: An input batch of data. By convention, these are global arrays. |
| 99 | Host-local input batches are also accepted and converted to global input batches. |
| 100 | |
| 101 | Returns: |
| 102 | An Output object containing global batch inputs, outputs and summaries. |
| 103 | N.B. the returned input and output batches will have the same partitioning. |
no outgoing calls
no test coverage detected