From b0e97f6b2bf6f89033dae6cfad7f26454434fbe8 Mon Sep 17 00:00:00 2001 From: Aditya Gupta Date: Wed, 2 Sep 2026 14:25:01 -0700 Subject: [PATCH] [JAX SC] Use private _mesh field and public mesh property in SparseCoreEmbed * Replace dataclass `mesh` field with private `_mesh` and `mesh` property. * Update callers to pass `_mesh`. PiperOrigin-RevId: 975328381 --- recml/inference/models/jax/DLRM_DCNv2/dlrm_model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/recml/inference/models/jax/DLRM_DCNv2/dlrm_model.py b/recml/inference/models/jax/DLRM_DCNv2/dlrm_model.py index 23720e5..a38c4c4 100644 --- a/recml/inference/models/jax/DLRM_DCNv2/dlrm_model.py +++ b/recml/inference/models/jax/DLRM_DCNv2/dlrm_model.py @@ -135,7 +135,7 @@ def __call__( sparse_embeddings = embed.SparseCoreEmbed( feature_specs=self.feature_specs, - mesh=self.mesh, + _mesh=self.mesh, sharding_axis=self.sharding_axis, )(embedding_lookups) sparse_embeddings = jax.tree.flatten(sparse_embeddings)