diff --git a/export/orbax/export/data_processors/jax_data_processor.py b/export/orbax/export/data_processors/jax_data_processor.py index 358de9386..46a0b08cc 100644 --- a/export/orbax/export/data_processors/jax_data_processor.py +++ b/export/orbax/export/data_processors/jax_data_processor.py @@ -22,6 +22,7 @@ import jaxtyping from orbax.export import constants from orbax.export import obm_configs +from orbax.export import utils from orbax.export.data_processors import data_processor_base from .third_party.neptune.protos import manifest_pb2 @@ -114,7 +115,9 @@ def __init__( input_keys: Set[str] = frozenset(), output_keys: Set[str] = frozenset(), params: Any = None, - options: obm_configs.Jax2ObmOptions | None = None, + options: obm_configs.Jax2ObmOptions = obm_configs.Jax2ObmOptions( + native_serialization_platforms=['cpu'] + ), ): """Initializes the instance. @@ -129,7 +132,16 @@ def __init__( super().__init__(name=name, input_keys=input_keys, output_keys=output_keys) self._processor_callable = processor_callable self._params = params - self._options = obm_configs.Jax2ObmOptions() if options is None else options + platforms = utils.get_lowering_platforms( + options.native_serialization_platforms + ) + if platforms and set(platforms) - {'cpu'}: + raise ValueError( + 'JaxDataProcessor only supports "cpu" for' + ' `native_serialization_platforms`, but got:' + f' {options.native_serialization_platforms}.' + ) + self._options = options self._is_prepared = False def prepare( diff --git a/export/orbax/export/data_processors/jax_data_processor_test.py b/export/orbax/export/data_processors/jax_data_processor_test.py index d5d215b55..b7586d583 100644 --- a/export/orbax/export/data_processors/jax_data_processor_test.py +++ b/export/orbax/export/data_processors/jax_data_processor_test.py @@ -44,6 +44,25 @@ def test_property_access_before_prepare_raises_error(self, property_name): ): _ = getattr(processor, property_name) + @parameterized.named_parameters( + dict(testcase_name='tpu', platforms=['tpu']), + dict(testcase_name='cpu_and_tpu', platforms=['cpu', 'tpu']), + dict(testcase_name='tpu_string', platforms='tpu'), + ) + def test_init_raises_error_with_more_than_cpu_platform(self, platforms): + with self.assertRaisesWithLiteralMatch( + ValueError, + 'JaxDataProcessor only supports "cpu" for' + ' `native_serialization_platforms`, but got:' + f' {platforms}.', + ): + _ = jax_data_processor.JaxDataProcessor( + lambda x: x, + options=obm_configs.Jax2ObmOptions( + native_serialization_platforms=platforms + ), + ) + def test_prepare_fails_with_multiple_calls(self): processor = jax_data_processor.JaxDataProcessor( lambda x: x, name='identity'