From b93b0e44c1b8058ff25ce3effacb544df5da5cf0 Mon Sep 17 00:00:00 2001 From: Jacky Fang Date: Mon, 3 Aug 2026 10:48:27 +0000 Subject: [PATCH] Fix JAX 0.11.0 sharding compatibility and SFT trainer signature --- src/maxtext/examples/lora_llama3_demo.ipynb | 3 +- .../integration/tunix/tunix_adapter.py | 37 +++++++++++++++++++ .../trainers/post_train/sft/train_sft.py | 7 +++- 3 files changed, 44 insertions(+), 3 deletions(-) diff --git a/src/maxtext/examples/lora_llama3_demo.ipynb b/src/maxtext/examples/lora_llama3_demo.ipynb index d324419093..4bafbde24e 100644 --- a/src/maxtext/examples/lora_llama3_demo.ipynb +++ b/src/maxtext/examples/lora_llama3_demo.ipynb @@ -186,7 +186,7 @@ "metadata": {}, "outputs": [], "source": [ - "MODEL_NAME = \"llama3.1-8b-Instruct\"\n", + "MODEL_NAME = \"llama3.1-8b\"\n", "TOKENIZER_PATH = \"meta-llama/Llama-3.1-8B-Instruct\"\n", "tokenizer = transformers.AutoTokenizer.from_pretrained(TOKENIZER_PATH)\n", "# This is the directory where the fine-tuned model checkpoint will be saved\n", @@ -341,7 +341,6 @@ "outputs": [], "source": [ "run_evaluation = os.environ.get(\"RUN_EVALUATION\", \"false\").lower() == \"true\"\n", - "run_evaluation = True\n", "if run_evaluation:\n", " test_dataset = get_test_dataset(config, tokenizer, DATA_TEMPLATE_PATH)\n", " test_dataset = test_dataset[:NUM_TEST_SAMPLES]\n", diff --git a/src/maxtext/integration/tunix/tunix_adapter.py b/src/maxtext/integration/tunix/tunix_adapter.py index 68da164112..5f5861e365 100644 --- a/src/maxtext/integration/tunix/tunix_adapter.py +++ b/src/maxtext/integration/tunix/tunix_adapter.py @@ -31,6 +31,43 @@ from maxtext.models.models import Transformer +import jax + +# Compatibility shims for JAX 0.11.0+ strict sharding assertions +_orig_wsc = jax.lax.with_sharding_constraint +_orig_top_k = jax.lax.top_k + + +def _compat_wsc(x, shardings): + try: + return _orig_wsc(x, shardings) + except Exception: # pylint: disable=broad-exception-caught + return jax.sharding.reshard(x, shardings) + + +def _compat_top_k(operand, k, axis=-1): + """Compat shim around jax.lax.top_k to reshard sharded reduction operands.""" + try: + return _orig_top_k(operand, k, axis=axis) + except Exception: # pylint: disable=broad-exception-caught + sharding = getattr(operand, "sharding", None) + if sharding is not None and hasattr(sharding, "spec") and hasattr(sharding, "mesh"): # pylint: disable=line-too-long + spec = list(sharding.spec) + idx = axis if axis >= 0 else len(spec) + axis + if 0 <= idx < len(spec): + spec[idx] = None + target_sharding = jax.sharding.NamedSharding(sharding.mesh, jax.sharding.PartitionSpec(*spec)) # pylint: disable=line-too-long + try: + operand = _orig_wsc(operand, target_sharding) + except Exception: # pylint: disable=broad-exception-caught + operand = jax.sharding.reshard(operand, target_sharding) + return _orig_top_k(operand, k, axis=axis) + + +jax.lax.with_sharding_constraint = _compat_wsc +jax.lax.top_k = _compat_top_k + + class TunixMaxTextAdapter(nnx.Module): """Adapter exposing Tunix Trainer call signature over a Transformer model.""" diff --git a/src/maxtext/trainers/post_train/sft/train_sft.py b/src/maxtext/trainers/post_train/sft/train_sft.py index 31709f35b9..461b15b7d8 100644 --- a/src/maxtext/trainers/post_train/sft/train_sft.py +++ b/src/maxtext/trainers/post_train/sft/train_sft.py @@ -107,7 +107,12 @@ def create_train_step_fn(self): nnx.pop(self.model, nnx.Intermediate) graphdef, _, _ = nnx.split(self.model, wrt, ...) - def train_step(model: nnx.Module, optimizer: nnx.Optimizer, inputs: Any): + def train_step( + model: nnx.Module, + optimizer: nnx.Optimizer, + inputs: Any, + grad_accumulator: Any = None, + ): inputs = gen_fn(inputs) # Split model into differentiable params and non-differentiable rest.