diff --git a/src/dependencies/extra_deps/post_train_github_deps.txt b/src/dependencies/extra_deps/post_train_github_deps.txt index 437d934bbe..b67a0d791c 100644 --- a/src/dependencies/extra_deps/post_train_github_deps.txt +++ b/src/dependencies/extra_deps/post_train_github_deps.txt @@ -1,3 +1,3 @@ -google-tunix @ https://github.com/google/tunix/archive/c4ec573d29e4c3a3955b348256d464b119c8a6d1.zip +google-tunix @ https://github.com/google/tunix/archive/1b0e3c5e89058d4dddf0ec68ae8be06c127f68ac.zip tpu-inference @ https://github.com/vllm-project/tpu-inference/archive/b67ae5f8f234fd559cf4b376840e7bb9ae6d3275.zip vllm @ git+https://github.com/vllm-project/vllm@d626108b1841888ec90aced33367149a6bbc7e4b diff --git a/src/maxtext/trainers/post_train/distillation/scripts/run_distill_xpk.sh b/src/maxtext/trainers/post_train/distillation/scripts/run_distill_xpk.sh index dbf9a67da0..8a6678b34f 100644 --- a/src/maxtext/trainers/post_train/distillation/scripts/run_distill_xpk.sh +++ b/src/maxtext/trainers/post_train/distillation/scripts/run_distill_xpk.sh @@ -164,7 +164,7 @@ require_env() { : "${DISTILL_LAYER_INDICES:=[0,1,2,3,4,5,6,7]}" # Image pinning (used by prep_image). -: "${TUNIX_SOURCE:=git+https://github.com/google/tunix@348959d18a4a09c75e58a7d49aec9d8b0eb4a8b6}" +: "${TUNIX_SOURCE:=git+https://github.com/google/tunix@1b0e3c5e89058d4dddf0ec68ae8be06c127f68ac}" : "${JAX_PIN:=0.10.0}" : "${JAXLIB_PIN:=0.10.0}" : "${LIBTPU_PIN:=0.0.39}" diff --git a/src/maxtext/training_engine/maxtext_engine.py b/src/maxtext/training_engine/maxtext_engine.py index e397b3f4a5..356b07eeb1 100644 --- a/src/maxtext/training_engine/maxtext_engine.py +++ b/src/maxtext/training_engine/maxtext_engine.py @@ -28,9 +28,8 @@ from absl import logging from flax import nnx +from flax import struct from flax.linen import partitioning as nn_partitioning -from flax.traverse_util import flatten_dict -from flax.traverse_util import unflatten_dict import jax import jax.numpy as jnp from jax.typing import ArrayLike # pylint: disable=g-importing-member @@ -158,7 +157,7 @@ def _batch_signature(dynamic_batch: Any, static_batch: dict[str, Any]) -> Any: ) -@dataclasses.dataclass(kw_only=True) +@struct.dataclass(frozen=True, kw_only=True) class RouterReplayTrainerPayload(abstract_engine.TrainerPayload): """A TrainerPayload extension carrying forced router-replay expert decisions. @@ -352,7 +351,6 @@ def __init__( self._compile_requested = False self._compiled_signature: Any = None self._signature_compare_warned: bool = False - self._raiden_syncs: Any = None if not training_config.model_name: raise ValueError("training_config.model_name must be specified") model_or_model_mesh_pair = model_creation_utils.from_pretrained( @@ -400,7 +398,7 @@ def __init__( ) self._metrics_recorder = metrics_module.MetricsRecorder() self._throttler = inflight_throttler.InflightThrottler(config=self._config) - self._raiden_syncs: Any = None + self._raiden_sync: Any = None @property def model(self) -> Any: @@ -1407,38 +1405,6 @@ def _get_trainable_params_state(self) -> Any: return nnx.state(model, nnx.Param) return self.model - def _split_into_chunks(self, nested_state: Any, num_chunks: int) -> list[Any]: - """Splits a nested param dict into `num_chunks` nested dicts of near-equal leaf count. - - Rebinding raiden's native WeightSynchronizer with a full new array list - only releases its hold on the PREVIOUS bind's buffers atomically with - acquiring the new ones (BindWeights: "releases the holds on the - previously bound buffers and acquires holds on the new ones") -- so the - complete new state must already be host-staged before the old one can be - dropped, and every rebind after the first needs ~2x one copy's worth of - host memory, not ~1x. Splitting the state across `num_chunks` independent - RaidenSynchronizer instances -- each bound, D2H'd, and released one at a - time -- bounds that overlap to ~(num_chunks+1)/num_chunks of one copy - instead of ~2x. Confirmed against a live OOM at num_chunks=1: main hit - the 420G container limit on a Qwen3-30B-A3B (~245GB bf16) trainer's - SECOND weight-sync cycle (the first has no stale buffer to overlap with). - The orchestrator and rollout's RaidenSamplerAdapter already support a - source contributing multiple WorkUnitMetadata entries (pooled by exact - variable name in manifest preflight, not by unit count), so no changes - are needed outside this trainer-side split. - """ - if hasattr(nested_state, "to_pure_dict"): - pure_state = nested_state.to_pure_dict() - elif hasattr(nested_state, "to_dict"): - pure_state = nested_state.to_dict() - else: - pure_state = nested_state - flat = flatten_dict(pure_state) - chunk_flats = [{} for _ in range(num_chunks)] - for i, key in enumerate(flat): - chunk_flats[i % num_chunks][key] = flat[key] - return [unflatten_dict(cf) for cf in chunk_flats] - def prepare_weight_sync( self, staging_transport: str = "raiden", @@ -1455,11 +1421,18 @@ def prepare_weight_sync( """ if staging_transport == "raiden": try: - # pylint: disable=g-import-not-at-top,import-outside-toplevel - from tunix.experimental.weight_sync import raiden_synchronizer - except ImportError: - logging.warning("tunix.experimental.weight_sync.raiden_synchronizer not found; returning empty metadata.") - return [] + from tunix.experimental.weight_sync import raiden_synchronizer # pylint: disable=g-import-not-at-top,import-outside-toplevel + except ImportError as exc: + # Fatal, not a warning: Raiden staging was explicitly requested and cannot be + # provided. Returning empty metadata instead defers the failure to the caller -- + # `WeightSyncCoordinator` eventually raises "metadata collection returned an empty + # side", which reports a count from another process and never mentions the missing + # module, leaving the real cause in this worker's log on another host. + raise RuntimeError( + "staging_transport='raiden' requires tunix.experimental.weight_sync." + "raiden_synchronizer, which the installed tunix does not provide. Install a" + " tunix build that ships it, or select a different staging_transport." + ) from exc # 1. Drain all in-flight TPU computations to ensure weights are fully updated self._throttler.wait_for_all() @@ -1490,80 +1463,70 @@ def prepare_weight_sync( scan_axis=self._config.param_scan_axis, ) - # 3. Bind parameters to the Raiden transport, one chunk at a time (see - # _split_into_chunks) -- construct the per-chunk synchronizers once, - # matching the persistent-instance-per-cycle pattern the rebind - # optimization (fewer stale holds) depends on. - num_chunks = max(1, int(os.environ.get("RAIDEN_WEIGHT_SYNC_CHUNKS", "1"))) - if self._raiden_syncs is None: - # Under Pathways (JAX_PLATFORMS=proxy + JAX_BACKEND_TARGET set, same - # detection tunix's K8sJaxContext.initialize() uses), trainer params - # are proxy-backed and Raiden can't bind them in place -- host_stage - # pulls them to client host memory first. Direct-TPU trainers skip - # that extra copy since their params already live on TPU. - is_pathways = bool("proxy" in os.environ.get("JAX_PLATFORMS", "") and os.environ.get("JAX_BACKEND_TARGET")) - # worker_index must be unique per chunk (it seeds WorkUnitId's - # job_replica_id) -- otherwise every chunk's work unit collides under - # the same id in the handler's registry and only one survives - # registration. - self._raiden_syncs = [ - raiden_synchronizer.RaidenSynchronizer( - job_name="trainer", - worker_index=jax.process_index() if num_chunks == 1 else (jax.process_index() * num_chunks + i + 1), - auto_h2d=False, - host_stage=is_pathways, - parallelism=4, - ) - for i in range(num_chunks) - ] - - chunks = self._split_into_chunks(params_state, num_chunks) if num_chunks > 1 else [params_state] - del params_state + # 3. Bind parameters to the Raiden transport. Construct the synchronizer + # once, matching the persistent-instance-per-cycle pattern the rebind + # optimization depends on. + # + # Under Pathways (JAX_PLATFORMS=proxy + JAX_BACKEND_TARGET set, same + # detection tunix's K8sJaxContext.initialize() uses), trainer params + # are proxy-backed. Raiden must use FFI (weight_synchronizer_ffi) to bind + # directly to device arrays on Pathways TPU workers without host CPU staging, + # avoiding client host OOM and multi-minute proxy transfer timeouts. + is_pathways = bool("proxy" in os.environ.get("JAX_PLATFORMS", "") and os.environ.get("JAX_BACKEND_TARGET")) + if is_pathways and getattr(raiden_synchronizer, "_raiden_ffi", None) is None: + raise RuntimeError( + "Under Pathways (JAX_PLATFORMS=proxy), Raiden weight synchronization " + "requires weight_synchronizer_ffi (from tpu_raiden_jax) to avoid client host OOM " + "and proxy staging timeouts. However, _raiden_ffi is not available in " + "tunix.experimental.weight_sync.raiden_synchronizer. Please ensure a " + "compatible tpu_raiden_jax wheel with FFI support is installed." + ) - verify_weights = os.environ.get("VERIFY_WEIGHTS", "").lower() == "true" - all_metadata = [] - total_variables = 0 - for chunk_idx, (sync, chunk_state) in enumerate(zip(self._raiden_syncs, chunks)): - sync.bind(chunk_state) + if self._raiden_sync is None: + self._raiden_sync = raiden_synchronizer.RaidenSynchronizer( + job_name="trainer", + worker_index=jax.process_index(), + auto_h2d=False, + parallelism=4, + ) - # 4. Initiate Device-to-Host transfer to stage this chunk for network - # transfer before moving on to the next chunk. - if sync.active: - sync.d2h() + self._raiden_sync.bind(params_state) + del params_state - if verify_weights: - logging.info("Source weights checksums (chunk %d): %s", chunk_idx, sync.checksums()) + # 4. Initiate Device-to-Host transfer to stage weights for network transfer. + if is_pathways or self._raiden_sync.active: + self._raiden_sync.d2h() - metadata = sync.work_unit_metadata() - total_variables += len(metadata.variables) - all_metadata.append(metadata) + verify_weights = os.environ.get("VERIFY_WEIGHTS", "").lower() == "true" + if verify_weights: + logging.info("Source weights checksums: %s", self._raiden_sync.checksums()) + metadata = self._raiden_sync.work_unit_metadata() logging.info( - "Trainer prepared weight sync for step %d: registered %d variables across %d chunk(s) on mesh %s", + "Trainer prepared weight sync for step %d: registered %d variables on mesh %s", self.train_step, - total_variables, - num_chunks, - all_metadata[0].mesh_axes if all_metadata else None, + len(metadata.variables), + metadata.mesh_axes, ) - return all_metadata + return [metadata] - return [] + # Unknown transport: raise rather than return empty metadata. A typo would otherwise + # surface only as the coordinator's "empty side" error, with nothing logged anywhere + # naming the transport that was actually asked for. + raise ValueError(f"unknown staging_transport {staging_transport!r}; expected 'raiden'.") def release_weight_sync(self, **kwargs: Any) -> Any: """Releases staged weight buffers after transfer completion.""" - if self._raiden_syncs: - for sync in self._raiden_syncs: - logging.vlog(1, "Trainer Raiden metrics: %s", sync.metrics()) - sync.release_host_arrays() + if self._raiden_sync: + logging.vlog(1, "Trainer Raiden metrics: %s", self._raiden_sync.metrics()) return True def close(self) -> None: """Closes the trainer, writes buffered metrics and final checkpoint.""" - if self._raiden_syncs: - for sync in self._raiden_syncs: - if hasattr(sync, "close"): - sync.close() - self._raiden_syncs = None + if self._raiden_sync: + if hasattr(self._raiden_sync, "close"): + self._raiden_sync.close() + self._raiden_sync = None self.save_checkpoint(metadata=None, force=True) self._checkpoint_manager.close() diff --git a/tests/end_to_end/tpu/compare_training_engine.py b/tests/end_to_end/tpu/compare_training_engine.py index ffc0f436c7..753b12eb1c 100644 --- a/tests/end_to_end/tpu/compare_training_engine.py +++ b/tests/end_to_end/tpu/compare_training_engine.py @@ -34,6 +34,7 @@ from typing import Any from flax import nnx +from flax import struct import jax import jax.numpy as jnp from maxtext.common import common_types @@ -120,7 +121,7 @@ def __call__( return self.proj(x) -@dataclasses.dataclass(kw_only=True) +@struct.dataclass(frozen=True, kw_only=True) class DummyPayload(abstract_engine.TrainerPayload): """Dummy payload for training engine parity comparisons.""" diff --git a/tests/post_training/unit/maxtext_engine_e2e_test.py b/tests/post_training/unit/maxtext_engine_e2e_test.py index ca04ed3896..1f2bd0e432 100644 --- a/tests/post_training/unit/maxtext_engine_e2e_test.py +++ b/tests/post_training/unit/maxtext_engine_e2e_test.py @@ -21,6 +21,7 @@ from absl.testing import absltest from flax import nnx +from flax import struct import jax import jax.numpy as jnp from maxtext.configs import pyconfig @@ -42,7 +43,7 @@ def __init__(self): self.weights = nnx.Param(jnp.array([1.0, 2.0])) -@dataclasses.dataclass(kw_only=True) +@struct.dataclass(frozen=True, kw_only=True) class DummyPayload(abstract_engine.TrainerPayload): """Dummy payload for testing.""" @@ -121,6 +122,8 @@ def setup_config(self, enable_checkpointing: bool = False, **kwargs): "tensorboard_dir": self.create_tempdir().full_path, "skip_jax_distributed_system": True, "enable_checkpointing": enable_checkpointing, + # Disable scan_layers to prevent prepare_weight_sync from trying to unscan layers on DummyNNXModel + "scan_layers": False, } if enable_checkpointing: overrides.update( diff --git a/tests/post_training/unit/maxtext_engine_test.py b/tests/post_training/unit/maxtext_engine_test.py index c3821b9b8c..5ee877bbb2 100644 --- a/tests/post_training/unit/maxtext_engine_test.py +++ b/tests/post_training/unit/maxtext_engine_test.py @@ -22,6 +22,7 @@ from absl.testing import absltest from flax import nnx +from flax import struct import jax import jax.numpy as jnp from maxtext.configs import pyconfig @@ -61,7 +62,7 @@ def __init__(self): self.calls = nnx.BatchStat(jnp.array(0.0)) -@dataclasses.dataclass(kw_only=True) +@struct.dataclass(frozen=True, kw_only=True) class DummyPayload(abstract_engine.TrainerPayload): token_ids: Any = dataclasses.field(default_factory=lambda: jnp.ones((2, 2))) token_mask: Any = dataclasses.field(default_factory=lambda: jnp.ones((2, 2)))