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

Method _runner_config

axlearn/common/inference_test.py:229–257  ·  view source on GitHub ↗
(
        self,
        *,
        mesh_shape: tuple[int, int],
        mesh_axis_names: tuple[str, str],
        param_dtype: jnp.dtype,
        inference_dtype: Optional[jnp.dtype],
        data_partition: DataPartitionType,
        ckpt_dir: str,
        use_ema: bool = False,
    )

Source from the content-addressed store, hash-verified

227
228 # pylint: disable-next=no-self-use
229 def _runner_config(
230 self,
231 *,
232 mesh_shape: tuple[int, int],
233 mesh_axis_names: tuple[str, str],
234 param_dtype: jnp.dtype,
235 inference_dtype: Optional[jnp.dtype],
236 data_partition: DataPartitionType,
237 ckpt_dir: str,
238 use_ema: bool = False,
239 ):
240 inference_runner_cfg = InferenceRunner.default_config().set(
241 mesh_shape=mesh_shape,
242 mesh_axis_names=mesh_axis_names,
243 model=DummyModel.default_config().set(dtype=param_dtype),
244 inference_dtype=inference_dtype,
245 input_batch_partition_spec=data_partition,
246 )
247 if use_ema:
248 inference_runner_cfg.init_state_builder = RestoreAndConvertBuilder.default_config().set(
249 name="builder",
250 builder=TensorStoreStateStorageBuilder.default_config().set(
251 dir=ckpt_dir, validation=CheckpointValidationType.CONTAINS_STATE_UP_TO_DTYPE
252 ),
253 converter=EmaParamsConverter.default_config(),
254 )
255 else:
256 inference_runner_cfg.init_state_builder.set(dir=ckpt_dir)
257 return inference_runner_cfg
258
259 # pylint: disable-next=no-self-use
260 def _build_ckpt(

Calls 2

setMethod · 0.45
default_configMethod · 0.45

Tested by

no test coverage detected