Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 1 addition & 2 deletions src/maxtext/examples/lora_llama3_demo.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -186,7 +186,7 @@
"metadata": {},
"outputs": [],
"source": [
"MODEL_NAME = \"llama3.1-8b-Instruct\"\n",
"MODEL_NAME = \"llama3.1-8b\"\n",

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Why do we need to use non-instruct version?

@RexBearIU RexBearIU Aug 4, 2026 •

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We fixed this in PR #4417. We accidentally removed Instruct while testing the new dependencies.

"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",
Expand Down Expand Up @@ -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",
Expand Down
37 changes: 37 additions & 0 deletions src/maxtext/integration/tunix/tunix_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down
7 changes: 6 additions & 1 deletion src/maxtext/trainers/post_train/sft/train_sft.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
Loading