(
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,
)
| 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( |
no test coverage detected