From 840a5865e88961adfbe9918419b293fb50b2d3a6 Mon Sep 17 00:00:00 2001 From: Oleg Shaldybin Date: Thu, 25 Jun 2026 15:13:51 -0700 Subject: [PATCH] Internal changes (moving code around) PiperOrigin-RevId: 938218966 --- .../orbax/export/data_processors/data_processor_base.py | 1 + export/orbax/export/data_processors/jax_data_processor.py | 8 +++----- .../export/data_processors/jax_data_processor_test.py | 2 +- .../export/data_processors/tf_data_processor_test.py | 4 +++- export/orbax/export/export_testing_utils.py | 1 + export/orbax/export/obm_export_test.py | 2 +- export/orbax/export/utils_test.py | 1 - 7 files changed, 10 insertions(+), 9 deletions(-) diff --git a/export/orbax/export/data_processors/data_processor_base.py b/export/orbax/export/data_processors/data_processor_base.py index fe28e72737..0e09e66f0f 100644 --- a/export/orbax/export/data_processors/data_processor_base.py +++ b/export/orbax/export/data_processors/data_processor_base.py @@ -19,6 +19,7 @@ import abc from collections.abc import Callable, Set from typing import Any +import jaxtyping class DataProcessor(abc.ABC): diff --git a/export/orbax/export/data_processors/jax_data_processor.py b/export/orbax/export/data_processors/jax_data_processor.py index 099e94ec0c..088b61ba99 100644 --- a/export/orbax/export/data_processors/jax_data_processor.py +++ b/export/orbax/export/data_processors/jax_data_processor.py @@ -24,16 +24,14 @@ from orbax.export import obm_configs from orbax.export.data_processors import data_processor_base -from .third_party.neptune.protos import manifest_pb2 - def _jax_spec_from(spec: Any) -> jax.ShapeDtypeStruct: """Converts a ShloTensorSpec to a jax.ShapeDtypeStruct.""" - if isinstance(spec, shlo_function.ShloTensorSpec): - if spec.dtype == shlo_function.ShloDType.bf16: + if isinstance(spec, shlo_type.ShloTensorSpec): + if spec.dtype == shlo_type.ShloDType.bf16: return jax.ShapeDtypeStruct(spec.shape, jax.numpy.bfloat16) return jax.ShapeDtypeStruct( - spec.shape, shlo_function.shlo_dtype_to_np_dtype(spec.dtype) + spec.shape, shlo_type.shlo_dtype_to_np_dtype(spec.dtype) ) if hasattr(spec, 'shape') and hasattr(spec, 'dtype'): return jax.ShapeDtypeStruct( 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 a4eda96f7e..f27056b627 100644 --- a/export/orbax/export/data_processors/jax_data_processor_test.py +++ b/export/orbax/export/data_processors/jax_data_processor_test.py @@ -19,12 +19,12 @@ import jax import jax.numpy as jnp -from orbax.experimental.model.core.python import value from orbax.export import obm_configs from orbax.export.data_processors import jax_data_processor from absl.testing import absltest from .testing.pybase import parameterized +from .third_party.neptune.neptune_model._src.core import value from .third_party.neptune.protos import manifest_pb2 diff --git a/export/orbax/export/data_processors/tf_data_processor_test.py b/export/orbax/export/data_processors/tf_data_processor_test.py index d6a10ed856..d8d42e36be 100644 --- a/export/orbax/export/data_processors/tf_data_processor_test.py +++ b/export/orbax/export/data_processors/tf_data_processor_test.py @@ -13,10 +13,12 @@ # limitations under the License. import os -import orbax.experimental.model.core as obm + from orbax.export.data_processors import tf_data_processor import tensorflow as tf + from absl.testing import absltest +import .third_party.neptune.neptune_model._src.core as obm from .util.task.python import error as google_error diff --git a/export/orbax/export/export_testing_utils.py b/export/orbax/export/export_testing_utils.py index 19ba04832e..1a7379d461 100644 --- a/export/orbax/export/export_testing_utils.py +++ b/export/orbax/export/export_testing_utils.py @@ -27,6 +27,7 @@ from orbax.export import jax_module from orbax.export import obm_configs from orbax.export import serving_config as osc + import tensorflow as tf diff --git a/export/orbax/export/obm_export_test.py b/export/orbax/export/obm_export_test.py index ecb1e3389f..8cca8015f0 100644 --- a/export/orbax/export/obm_export_test.py +++ b/export/orbax/export/obm_export_test.py @@ -24,7 +24,7 @@ from google.protobuf import text_format import jax import jax.numpy as jnp -from jaxtyping import PyTree +import jaxtyping from orbax.export import constants from orbax.export import export_testing_utils from orbax.export import jax_module diff --git a/export/orbax/export/utils_test.py b/export/orbax/export/utils_test.py index 4cc9e77652..e12fd284e4 100644 --- a/export/orbax/export/utils_test.py +++ b/export/orbax/export/utils_test.py @@ -21,7 +21,6 @@ from orbax.export import utils import tensorflow as tf - TensorSpecWithDefault = utils.TensorSpecWithDefault