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

Class MethodRunner

axlearn/common/inference.py:46–137  ·  view source on GitHub ↗

Class for running a pjit-lowered model's method on a batch of data.

Source from the content-addressed store, hash-verified

44
45# pylint: disable-next=too-few-public-methods
46class 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.

Callers 1

create_method_runnerMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected