diff --git a/tests/unit_tests/a2a_overlap/test_fsdp_1f1b_overlap.py b/tests/unit_tests/a2a_overlap/test_fsdp_1f1b_overlap.py index 956ebd2f1f3..6c20a5784b1 100644 --- a/tests/unit_tests/a2a_overlap/test_fsdp_1f1b_overlap.py +++ b/tests/unit_tests/a2a_overlap/test_fsdp_1f1b_overlap.py @@ -8,6 +8,7 @@ from megatron.core.distributed import DistributedDataParallelConfig from megatron.core.distributed.fsdp.mcore_fsdp_adapter import FullyShardedDataParallel from megatron.core.distributed.fsdp.src.megatron_fsdp.fully_shard import fully_shard_optimizer +from megatron.core.pipeline_parallel.fine_grained_activation_offload import PipelineOffloadManager from megatron.core.pipeline_parallel.utils import set_streams from megatron.core.transformer import TransformerLayer from megatron.core.utils import is_te_min_version @@ -49,6 +50,7 @@ def setup_method(self, method): set_streams() def teardown_method(self, method): + PipelineOffloadManager.reset_instance() Utils.destroy_model_parallel() @pytest.mark.skipif(not is_te_min_version("2.3.0"), reason="Requires TE >= 2.3.0") diff --git a/tests/unit_tests/conftest.py b/tests/unit_tests/conftest.py index a856fe98010..5b73e7082ca 100644 --- a/tests/unit_tests/conftest.py +++ b/tests/unit_tests/conftest.py @@ -12,7 +12,12 @@ from megatron.core.utils import is_te_min_version from tests.test_utils.python_scripts.download_unit_tests_dataset import download_and_extract_asset from tests.unit_tests.dist_checkpointing import TempNamedDir -from tests.unit_tests.test_utilities import Utils +from tests.unit_tests.test_utilities import ( + Utils, + reset_transient_process_state, + restore_process_state, + snapshot_process_state, +) def pytest_configure(config): @@ -167,3 +172,25 @@ def reset_env_vars(): # After the test, restore the original environment os.environ.clear() os.environ.update(original_env) + + +@pytest.fixture(scope="module", autouse=True) +def restore_module_process_state(): + snapshot = snapshot_process_state() + yield + restore_process_state(snapshot) + + +@pytest.fixture(scope="class", autouse=True) +def restore_class_process_state(): + snapshot = snapshot_process_state() + yield + restore_process_state(snapshot) + + +@pytest.fixture(autouse=True) +def reset_process_state(): + snapshot = snapshot_process_state() + yield + reset_transient_process_state() + restore_process_state(snapshot) diff --git a/tests/unit_tests/data/test_bin_reader.py b/tests/unit_tests/data/test_bin_reader.py index e479676ac4b..6d370788df5 100644 --- a/tests/unit_tests/data/test_bin_reader.py +++ b/tests/unit_tests/data/test_bin_reader.py @@ -1,3 +1,5 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + import os import random import sys @@ -100,9 +102,6 @@ def close(self) -> None: pass -setattr(boto3, "client", _LocalClient) - - ## # Overload ClientError from botocore.exceptions ## @@ -114,8 +113,6 @@ class _LocalClientError(Exception): pass -setattr(exceptions, "ClientError", _LocalClientError) - ## # Mock multistorageclient module ## @@ -137,14 +134,19 @@ def read(self, path, byte_range): return StorageClient(), path.removeprefix(MSC_PREFIX + "default") -setattr(msc, "open", open) -setattr(msc, "download_file", _msc_download_file) -setattr(msc, "resolve_storage_client", _msc_resolve_storage_client) +@pytest.fixture +def local_object_storage(monkeypatch): + """Point the boto3 and msc clients at the local filesystem for one test.""" + monkeypatch.setattr(boto3, "client", _LocalClient, raising=False) + monkeypatch.setattr(exceptions, "ClientError", _LocalClientError, raising=False) + monkeypatch.setattr(msc, "open", open, raising=False) + monkeypatch.setattr(msc, "download_file", _msc_download_file, raising=False) + monkeypatch.setattr(msc, "resolve_storage_client", _msc_resolve_storage_client, raising=False) @pytest.mark.flaky @pytest.mark.flaky_in_dev -def test_bin_reader(): +def test_bin_reader(local_object_storage): with tempfile.TemporaryDirectory() as temp_dir: # set the default nltk data path os.environ["NLTK_DATA"] = os.path.join(temp_dir, "nltk_data") diff --git a/tests/unit_tests/data/test_get_batch.py b/tests/unit_tests/data/test_get_batch.py index 993959d3cd9..87400777171 100644 --- a/tests/unit_tests/data/test_get_batch.py +++ b/tests/unit_tests/data/test_get_batch.py @@ -27,6 +27,13 @@ from tests.unit_tests.test_utilities import Utils +@pytest.fixture(autouse=True) +def destroy_test_environment(): + yield + destroy_global_vars() + destroy_num_microbatches_calculator() + + def initialize_test_environment( tp_size: int, pp_size: int, diff --git a/tests/unit_tests/dist_checkpointing/test_fully_parallel.py b/tests/unit_tests/dist_checkpointing/test_fully_parallel.py index 2524adaf5a9..11efef2f3e6 100644 --- a/tests/unit_tests/dist_checkpointing/test_fully_parallel.py +++ b/tests/unit_tests/dist_checkpointing/test_fully_parallel.py @@ -424,6 +424,7 @@ def test_only_necessary_exchanges_performed_during_load(self, tmp_path_dist_ckpt Utils.destroy_model_parallel() def test_broadcast_sharded_objects(self, tmp_path_dist_ckpt): + Utils.initialize_distributed() sharded_state_dict = { f'Obj_{i}': ShardedObject(f'Obj_{i}', None, (1,), (0,), replica_id=abs(Utils.rank - i)) diff --git a/tests/unit_tests/dist_checkpointing/test_msc.py b/tests/unit_tests/dist_checkpointing/test_msc.py index c3f3e78133d..c7a4e71633b 100644 --- a/tests/unit_tests/dist_checkpointing/test_msc.py +++ b/tests/unit_tests/dist_checkpointing/test_msc.py @@ -24,6 +24,7 @@ def setup_method(self, method): MultiStorageClientFeature.enable() def teardown_method(self, method): + MultiStorageClientFeature.disable() Utils.destroy_model_parallel() def test_process_save_load(self, tmp_path_dist_ckpt): diff --git a/tests/unit_tests/dist_checkpointing/test_safe_globals.py b/tests/unit_tests/dist_checkpointing/test_safe_globals.py index 324de974a7d..341311aac8c 100755 --- a/tests/unit_tests/dist_checkpointing/test_safe_globals.py +++ b/tests/unit_tests/dist_checkpointing/test_safe_globals.py @@ -11,6 +11,7 @@ from megatron.core.safe_globals import SafeUnpickler from megatron.core.utils import is_torch_min_version +from tests.unit_tests.test_utilities import Utils class UnsafeClass: @@ -23,6 +24,8 @@ def __repr__(self): class TestSafeGlobals: def test_safe_globals(self, tmp_path_dist_ckpt): + if Utils.world_size > 1: + Utils.initialize_distributed() # create dummy checkpoint ckpt_path = tmp_path_dist_ckpt / "test_safe_globals.pt" dummy_obj = Namespace(dummy_value=0) @@ -35,6 +38,8 @@ def test_safe_globals(self, tmp_path_dist_ckpt): @pytest.mark.skipif(not is_torch_min_version("2.6a0"), reason="PyTorch 2.6 is required") def test_unsafe_globals(self, tmp_path_dist_ckpt): + if Utils.world_size > 1: + Utils.initialize_distributed() # create dummy checkpoint ckpt_path = tmp_path_dist_ckpt / "test_safe_globals.pt" dummy_obj = UnsafeClass(123) diff --git a/tests/unit_tests/distributed/test_param_and_grad_buffer.py b/tests/unit_tests/distributed/test_param_and_grad_buffer.py index 8fa01c89ced..1cd7feedf93 100644 --- a/tests/unit_tests/distributed/test_param_and_grad_buffer.py +++ b/tests/unit_tests/distributed/test_param_and_grad_buffer.py @@ -428,7 +428,7 @@ def test_start_param_sync_dp_size_1(): """When dp_size == 1 (e.g., expt_dp_size == 1), start_param_sync should set param_gather_dispatched=True and return immediately without launching any all-gather collective.""" - world_size = torch.distributed.get_world_size() + world_size = Utils.world_size Utils.initialize_model_parallel(tensor_model_parallel_size=world_size) ddp_config = DistributedDataParallelConfig( diff --git a/tests/unit_tests/extension/test_kitchen_sdpa.py b/tests/unit_tests/extension/test_kitchen_sdpa.py index 0660bc226ba..b41c8e3ef48 100644 --- a/tests/unit_tests/extension/test_kitchen_sdpa.py +++ b/tests/unit_tests/extension/test_kitchen_sdpa.py @@ -46,15 +46,22 @@ dot_product_attention = MagicMock() -# Create custom process groups -Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) -model_parallel_cuda_manual_seed(123) - -# Get TP and CP process groups from device mesh -tp_group = parallel_state.get_tensor_model_parallel_group() -cp_group = parallel_state.get_context_parallel_group() - -pg_collection = ProcessGroupCollection(tp=tp_group, cp=cp_group) +pg_collection = None + + +@pytest.fixture(scope="module", autouse=True) +def model_parallel_groups(): + """Create the TP and CP groups for this module's tests, and only when they run.""" + global pg_collection + Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) + model_parallel_cuda_manual_seed(123) + pg_collection = ProcessGroupCollection( + tp=parallel_state.get_tensor_model_parallel_group(), + cp=parallel_state.get_context_parallel_group(), + ) + yield + pg_collection = None + Utils.destroy_model_parallel() def get_attention_implementation( diff --git a/tests/unit_tests/generalized_tensor_parallel/gtp_test_utils.py b/tests/unit_tests/generalized_tensor_parallel/gtp_test_utils.py index dab9fa0790c..d6e37ea8098 100644 --- a/tests/unit_tests/generalized_tensor_parallel/gtp_test_utils.py +++ b/tests/unit_tests/generalized_tensor_parallel/gtp_test_utils.py @@ -9,6 +9,7 @@ from transformer_engine.pytorch.quantization import FP8GlobalStateManager from megatron.core.tensor_parallel.generalized_tensor_parallelism import ( + _GTP_PARAMS, GTPShardedParam, reset_gtp_state, ) @@ -45,6 +46,8 @@ def reset_gtp_globals(): """ yield reset_gtp_state() + # Not part of the production reset: every GTP param ever built stays in it (and allocated). + _GTP_PARAMS.clear() # --------------------------------------------------------------------------- diff --git a/tests/unit_tests/generalized_tensor_parallel/test_gtp_dcp.py b/tests/unit_tests/generalized_tensor_parallel/test_gtp_dcp.py index 79eec626e3d..8971ccc411c 100644 --- a/tests/unit_tests/generalized_tensor_parallel/test_gtp_dcp.py +++ b/tests/unit_tests/generalized_tensor_parallel/test_gtp_dcp.py @@ -66,6 +66,7 @@ from tests.unit_tests.generalized_tensor_parallel.gtp_test_utils import ( # noqa: E402,F401 _requires_mxfp8, _torchrun_dist_init, + reset_gtp_globals, ) diff --git a/tests/unit_tests/inference/contexts/test_dynamic_prefix_caching.py b/tests/unit_tests/inference/contexts/test_dynamic_prefix_caching.py index 5346395de83..77cc684a0a4 100644 --- a/tests/unit_tests/inference/contexts/test_dynamic_prefix_caching.py +++ b/tests/unit_tests/inference/contexts/test_dynamic_prefix_caching.py @@ -39,6 +39,7 @@ from tests.unit_tests.inference.engines.test_dynamic_engine import ( DynamicEngineTestConfig, DynamicInferenceEngineTestBase, + reset_rounder, ) from tests.unit_tests.test_utilities import Utils @@ -56,6 +57,9 @@ def setup_class(cls): def teardown_class(cls): Utils.destroy_model_parallel() + def teardown_method(self, method): + reset_rounder() + @staticmethod def _mamba_config(mamba_chunk_size=128): from megatron.core.inference.config import MambaInferenceStateConfig @@ -3257,9 +3261,7 @@ def test_real_engine_stress_row(self, case): try: self._run_engine_case(case) finally: - DynamicInferenceContext.ROUNDER = 64 - DynamicInferenceContext.TOKEN_ROUNDER = 64 - DynamicInferenceContext.REQUEST_ROUNDER = 64 + reset_rounder() Utils.destroy_model_parallel() diff --git a/tests/unit_tests/inference/engines/request_lifecycle_test_utils.py b/tests/unit_tests/inference/engines/request_lifecycle_test_utils.py index 97a507e98fa..e2e5b859431 100644 --- a/tests/unit_tests/inference/engines/request_lifecycle_test_utils.py +++ b/tests/unit_tests/inference/engines/request_lifecycle_test_utils.py @@ -26,7 +26,7 @@ from tests.unit_tests.inference.engines.test_dynamic_engine import ( DynamicInferenceEngineTestBase as _DynamicInferenceEngineTestBase, ) -from tests.unit_tests.inference.engines.test_dynamic_engine import set_rounder as _set_rounder +from tests.unit_tests.inference.engines.test_dynamic_engine import reset_rounder as _reset_rounder from tests.unit_tests.inference.engines.test_dynamic_engine_async_sched import ( _BASE_PAIR_CONFIG, _assert_request_parity, @@ -456,7 +456,7 @@ def setup_class(cls): @classmethod def teardown_class(cls): delete_cuda_graphs() - _set_rounder(64) + _reset_rounder() Utils.destroy_model_parallel() @staticmethod diff --git a/tests/unit_tests/inference/engines/test_dynamic_engine.py b/tests/unit_tests/inference/engines/test_dynamic_engine.py index 80752f5f8e8..be82cb0413a 100644 --- a/tests/unit_tests/inference/engines/test_dynamic_engine.py +++ b/tests/unit_tests/inference/engines/test_dynamic_engine.py @@ -23,6 +23,7 @@ from megatron.core import parallel_state from megatron.core.activations import squared_relu +from megatron.core.inference.batch_dimensions_utils import TOKEN_ROUNDER from megatron.core.inference.config import ( AsyncScheduleMode, CudaGraphSizingDistribution, @@ -859,6 +860,13 @@ def set_rounder(value): DynamicInferenceContext.REQUEST_ROUNDER = value +def reset_rounder(): + """Restore the production rounders; set_rounder(64) would leave REQUEST_ROUNDER at 64.""" + DynamicInferenceContext.ROUNDER = TOKEN_ROUNDER + DynamicInferenceContext.TOKEN_ROUNDER = TOKEN_ROUNDER + DynamicInferenceContext.REQUEST_ROUNDER = 4 # the default in dynamic_context.py + + def mock_forward(input_ids, position_ids, attention_mask, *args, **kwargs): """Mock forward function to avoid numerics issues with random inputs.""" return torch.randn( @@ -877,8 +885,7 @@ class DynamicEngineTestConfig: random_seed = 123 vocab_size = 100 - set_rounder(4) - num_requests: int = 2 * DynamicInferenceContext.round_up_requests(1, 1) + num_requests: int = 8 min_prompt_length: int = 4 max_prompt_length: int = 16 num_tokens_to_generate: Optional[int] = 4 @@ -2777,7 +2784,7 @@ def setup_class(cls): @classmethod def teardown_class(cls): delete_cuda_graphs() - set_rounder(64) + reset_rounder() Utils.destroy_model_parallel() @pytest.mark.internal @@ -7941,7 +7948,7 @@ def setup_class(cls): @classmethod def teardown_class(cls): delete_cuda_graphs() - set_rounder(64) + reset_rounder() Utils.destroy_model_parallel() @staticmethod diff --git a/tests/unit_tests/inference/engines/test_dynamic_engine_async_sched.py b/tests/unit_tests/inference/engines/test_dynamic_engine_async_sched.py index ef371e00fe6..b2803340878 100644 --- a/tests/unit_tests/inference/engines/test_dynamic_engine_async_sched.py +++ b/tests/unit_tests/inference/engines/test_dynamic_engine_async_sched.py @@ -51,7 +51,7 @@ from tests.unit_tests.inference.engines.test_dynamic_engine import ( DynamicInferenceEngineTestBase as _DynamicInferenceEngineTestBase, ) -from tests.unit_tests.inference.engines.test_dynamic_engine import set_rounder as _set_rounder +from tests.unit_tests.inference.engines.test_dynamic_engine import reset_rounder as _reset_rounder from tests.unit_tests.test_utilities import Utils @@ -1942,7 +1942,7 @@ def setup_class(cls): @classmethod def teardown_class(cls): delete_cuda_graphs() - _set_rounder(64) + _reset_rounder() Utils.destroy_model_parallel() @pytest.mark.parametrize("scenario", _ASYNC_PAIR_SCENARIOS, ids=lambda case: case.name) @@ -2010,7 +2010,7 @@ def test_async_matches_legacy_for_parallel_pair(self, scenario): gc.collect() delete_cuda_graphs() torch.cuda.empty_cache() - _set_rounder(64) + _reset_rounder() Utils.destroy_model_parallel() diff --git a/tests/unit_tests/inference/engines/test_dynamic_engine_parallel_and_chunked_prefill.py b/tests/unit_tests/inference/engines/test_dynamic_engine_parallel_and_chunked_prefill.py index cbbb946afeb..928df1e16df 100644 --- a/tests/unit_tests/inference/engines/test_dynamic_engine_parallel_and_chunked_prefill.py +++ b/tests/unit_tests/inference/engines/test_dynamic_engine_parallel_and_chunked_prefill.py @@ -44,6 +44,7 @@ ) from tests.unit_tests.inference.engines.test_dynamic_engine import ( DynamicInferenceEngineTestBase, + reset_rounder, set_rounder, skip_if_mamba_sequence_packing_not_available, ) @@ -201,8 +202,7 @@ def test_parallel_inference( if tp_size == 1 and pp_size == 1 and ep_size == 1: pytest.skip(reason="Test requires tp_size > 1 or pp_size > 1 or ep_size > 1") - elif not torch.distributed.is_initialized(): - pytest.skip("Distributed not initialized") + Utils.initialize_distributed() world_size = torch.distributed.get_world_size() min_world_size = tp_size * pp_size * ep_size if world_size < min_world_size: @@ -260,8 +260,7 @@ def test_sequence_parallel_fp8_inference(self, materialize_only_last_token_logit @torch.inference_mode() def test_speculative_decoding_pipeline_parallel(self): """Test speculative decoding with pipeline parallelism (pp_size=2).""" - if not torch.distributed.is_initialized(): - pytest.skip("Distributed not initialized") + Utils.initialize_distributed() world_size = torch.distributed.get_world_size() pp_size = 2 if world_size < pp_size: @@ -492,8 +491,7 @@ def test_mtp_kv_cache_decode_refresh_with_forced_acceptance(self): @torch.inference_mode() def test_mtp_kv_cache_chunked_prefill(self, transformer_impl, ep_size, dispatcher, monkeypatch): """Every backend must consume a carried hidden when a prompt spans multiple steps.""" - if not torch.distributed.is_initialized(): - pytest.skip("Distributed not initialized") + Utils.initialize_distributed() if torch.distributed.get_world_size() < ep_size: pytest.skip(f"Test requires at least {ep_size} GPUs") skip_if_mamba_sequence_packing_not_available("hybrid") @@ -1170,8 +1168,7 @@ def step_and_sample(): @torch.inference_mode() def test_mtp_kv_cache_inference_optimized(self, num_cuda_graphs): """MTP draft KV cache on the inference-optimized transformer (TP=1, dense).""" - if not torch.distributed.is_initialized(): - pytest.skip("Distributed not initialized") + Utils.initialize_distributed() skip_if_mamba_sequence_packing_not_available("hybrid") @@ -1207,8 +1204,7 @@ def test_mtp_kv_cache_inference_optimized(self, num_cuda_graphs): @torch.inference_mode() def test_mtp_kv_cache_inference_optimized_sequence_parallel(self): """Same, with TP=2 + SP: the commit pass pads to a TP multiple and scatters.""" - if not torch.distributed.is_initialized(): - pytest.skip("Distributed not initialized") + Utils.initialize_distributed() if torch.distributed.get_world_size() < 2: pytest.skip("Test requires at least 2 GPUs") @@ -1398,7 +1394,7 @@ def setup_class(cls): @classmethod def teardown_class(cls): delete_cuda_graphs() - set_rounder(64) + reset_rounder() Utils.destroy_model_parallel() def _create_model(self, model_provider, num_cuda_graphs, ssm_mixer="mamba"): diff --git a/tests/unit_tests/inference/engines/test_dynamic_engine_request_lifecycle_pause_identity.py b/tests/unit_tests/inference/engines/test_dynamic_engine_request_lifecycle_pause_identity.py index 763504228ca..eb1d09b31d4 100644 --- a/tests/unit_tests/inference/engines/test_dynamic_engine_request_lifecycle_pause_identity.py +++ b/tests/unit_tests/inference/engines/test_dynamic_engine_request_lifecycle_pause_identity.py @@ -25,7 +25,7 @@ _RunResult, _track_checkpoint_calls, ) -from tests.unit_tests.inference.engines.test_dynamic_engine import set_rounder as _set_rounder +from tests.unit_tests.inference.engines.test_dynamic_engine import reset_rounder as _reset_rounder from tests.unit_tests.inference.engines.test_dynamic_engine_async_sched import ( _ASYNC_PARALLEL_SCENARIOS, _instrument_scenario_runtime, @@ -228,5 +228,5 @@ def test_retained_pause_ep2(self): assert treatment.witness["dispatches_total"] > treatment.witness["dispatches_at_pause"] finally: self._cleanup() - _set_rounder(64) + _reset_rounder() Utils.destroy_model_parallel() diff --git a/tests/unit_tests/inference/engines/test_hybrid_prefix_caching_e2e.py b/tests/unit_tests/inference/engines/test_hybrid_prefix_caching_e2e.py index f66eac57e02..9bd43c1a256 100644 --- a/tests/unit_tests/inference/engines/test_hybrid_prefix_caching_e2e.py +++ b/tests/unit_tests/inference/engines/test_hybrid_prefix_caching_e2e.py @@ -40,6 +40,7 @@ import torch from megatron.core import parallel_state +from megatron.core.inference.batch_dimensions_utils import TOKEN_ROUNDER from megatron.core.inference.config import ( AsyncScheduleMode, InferenceConfig, @@ -111,6 +112,12 @@ def set_rounder(value): DynamicInferenceContext.REQUEST_ROUNDER = value +def reset_rounder(): + DynamicInferenceContext.ROUNDER = TOKEN_ROUNDER + DynamicInferenceContext.TOKEN_ROUNDER = TOKEN_ROUNDER + DynamicInferenceContext.REQUEST_ROUNDER = 4 # the default in dynamic_context.py + + @pytest.mark.internal @pytest.mark.skipif(not is_fa_min_version("2.7.3"), reason="need flash attn") class TestMambaPrefixCachingE2E: @@ -144,6 +151,7 @@ def teardown_method(self, method): # Free captured CUDA graphs and their private mempools between tests; # otherwise graph memory accumulates across a --count/whole-file run and OOMs. delete_cuda_graphs() + reset_rounder() @pytest.fixture(params=["mamba", "gdp"], autouse=True) def ssm_mixer(self, request): diff --git a/tests/unit_tests/inference/engines/test_prefix_caching_cuda_graphs.py b/tests/unit_tests/inference/engines/test_prefix_caching_cuda_graphs.py index fe205eb6be4..11fe6f60d0e 100644 --- a/tests/unit_tests/inference/engines/test_prefix_caching_cuda_graphs.py +++ b/tests/unit_tests/inference/engines/test_prefix_caching_cuda_graphs.py @@ -19,6 +19,7 @@ import torch from megatron.core import parallel_state +from megatron.core.inference.batch_dimensions_utils import TOKEN_ROUNDER from megatron.core.inference.config import InferenceConfig, PrefixCachingEvictionPolicy from megatron.core.inference.contexts.dynamic_context import DynamicInferenceContext from megatron.core.inference.engines import DynamicInferenceEngine @@ -57,6 +58,12 @@ def set_rounder(value): DynamicInferenceContext.REQUEST_ROUNDER = value +def reset_rounder(): + DynamicInferenceContext.ROUNDER = TOKEN_ROUNDER + DynamicInferenceContext.TOKEN_ROUNDER = TOKEN_ROUNDER + DynamicInferenceContext.REQUEST_ROUNDER = 4 # the default in dynamic_context.py + + @pytest.mark.internal @pytest.mark.skipif(not is_fa_min_version("2.7.3"), reason="need flash attn") class TestPrefixCachingCudaGraphs: @@ -335,7 +342,7 @@ def setup_class(cls): @classmethod def teardown_class(cls): - set_rounder(64) + reset_rounder() Utils.destroy_model_parallel() def _create_hybrid_model(self, num_cuda_graphs=None, ssm_mixer="mamba"): diff --git a/tests/unit_tests/inference/engines/test_static_engine.py b/tests/unit_tests/inference/engines/test_static_engine.py index a2d19ff2f11..025c197fff7 100644 --- a/tests/unit_tests/inference/engines/test_static_engine.py +++ b/tests/unit_tests/inference/engines/test_static_engine.py @@ -355,8 +355,7 @@ def setup_engine( def test_parallel_inference(self, tp_size, pp_size, ep_size, sequence_parallel): if tp_size == 1 and pp_size == 1 and ep_size == 1: pytest.skip(reason="Test requires tp_size > 1 or pp_size > 1 or ep_size > 1") - elif not torch.distributed.is_initialized(): - pytest.skip("Distributed not initialized") + Utils.initialize_distributed() world_size = torch.distributed.get_world_size() min_world_size = tp_size * pp_size * ep_size if world_size < min_world_size: diff --git a/tests/unit_tests/inference/test_dynamic_sink_attention_e2e.py b/tests/unit_tests/inference/test_dynamic_sink_attention_e2e.py index b75a057e175..613729e12fa 100644 --- a/tests/unit_tests/inference/test_dynamic_sink_attention_e2e.py +++ b/tests/unit_tests/inference/test_dynamic_sink_attention_e2e.py @@ -25,6 +25,7 @@ * The full plumbing through ``Attention.forward()`` → ``flash_decode_and_prefill()`` → sink correction → linear_proj. """ + import pytest import torch @@ -38,7 +39,7 @@ # in a separate edit to ``DynamicEngineTestConfig``. from tests.unit_tests.inference.engines.test_dynamic_engine import ( DynamicInferenceEngineTestBase, - set_rounder, + reset_rounder, ) from tests.unit_tests.test_utilities import Utils @@ -81,7 +82,7 @@ def teardown_class(cls): # Deliberately NOT calling delete_cuda_graphs() — these tests do # not enable CUDA graphs, so there is nothing to clean up, and # avoiding the call sidesteps the known teardown SIGABRT. - set_rounder(64) + reset_rounder() Utils.destroy_model_parallel() @staticmethod diff --git a/tests/unit_tests/inference/test_moe_dispatching_and_routing.py b/tests/unit_tests/inference/test_moe_dispatching_and_routing.py index b5ff4b776ed..8ec408409d0 100644 --- a/tests/unit_tests/inference/test_moe_dispatching_and_routing.py +++ b/tests/unit_tests/inference/test_moe_dispatching_and_routing.py @@ -384,7 +384,11 @@ def setup_class(cls): @classmethod def teardown_class(cls): from megatron.core.inference.symmetric_memory import SymmetricMemoryManager + from megatron.core.transformer.moe.token_dispatcher_inference import ( + NVLSAllGatherVDispatcher, + ) + NVLSAllGatherVDispatcher._delete_buffers() SymmetricMemoryManager.destroy() Utils.destroy_model_parallel() diff --git a/tests/unit_tests/inference/test_wandb_logging.py b/tests/unit_tests/inference/test_wandb_logging.py index 87d4a6dc764..593f73140c5 100644 --- a/tests/unit_tests/inference/test_wandb_logging.py +++ b/tests/unit_tests/inference/test_wandb_logging.py @@ -7,6 +7,7 @@ import pytest import torch +from megatron.core.inference.batch_dimensions_utils import TOKEN_ROUNDER from megatron.core.inference.config import InferenceConfig from megatron.core.inference.contexts.dynamic_context import DynamicInferenceContext from megatron.core.inference.engines import DynamicInferenceEngine @@ -21,12 +22,17 @@ def set_rounder(value): - """Utility function to set the DynamicInferenceContext rounder.""" DynamicInferenceContext.ROUNDER = value # For backwards compatibility DynamicInferenceContext.TOKEN_ROUNDER = value DynamicInferenceContext.REQUEST_ROUNDER = value +def reset_rounder(): + DynamicInferenceContext.ROUNDER = TOKEN_ROUNDER + DynamicInferenceContext.TOKEN_ROUNDER = TOKEN_ROUNDER + DynamicInferenceContext.REQUEST_ROUNDER = 4 # the default in dynamic_context.py + + class TestInferenceWandbLogging: """Test suite for wandb logging in inference.""" @@ -40,7 +46,7 @@ def setup_class(cls): @classmethod def teardown_class(cls): - set_rounder(64) + reset_rounder() Utils.destroy_model_parallel() def _get_dynamic_context( @@ -264,6 +270,7 @@ def test_paused_requests_in_stats(self): # Verify paused request count is included assert 'paused_request_count' in stats assert stats['paused_request_count'] >= 0 + set_rounder(64) # back to the class's pin @pytest.mark.internal def test_metrics_writer_none_handling(self): diff --git a/tests/unit_tests/models/test_gpt_model.py b/tests/unit_tests/models/test_gpt_model.py index 6f22a5f0925..9695ed33e1f 100644 --- a/tests/unit_tests/models/test_gpt_model.py +++ b/tests/unit_tests/models/test_gpt_model.py @@ -461,6 +461,7 @@ def test_gpt_with_te_activation_func(num_experts, gated_linear_unit): class TestGPTModelWithCustomPG: def setup_method(self, method): Utils.destroy_model_parallel() + Utils.initialize_distributed() def teardown_method(self, method): Utils.destroy_model_parallel() diff --git a/tests/unit_tests/pipeline_parallel/test_fine_grained_activation_offloading.py b/tests/unit_tests/pipeline_parallel/test_fine_grained_activation_offloading.py index fe31c7ed52f..df992c9ed31 100644 --- a/tests/unit_tests/pipeline_parallel/test_fine_grained_activation_offloading.py +++ b/tests/unit_tests/pipeline_parallel/test_fine_grained_activation_offloading.py @@ -23,6 +23,7 @@ ) from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.enums import AttnBackend +from megatron.core.transformer.moe.fused_a2a import reset_hybrid_ep_buffer from megatron.core.transformer.transformer_config import MLATransformerConfig, TransformerConfig from megatron.core.utils import is_te_min_version from tests.unit_tests.test_utilities import Utils @@ -46,6 +47,12 @@ def _make_chunk_handler_for_offload_checker(min_offloaded_tensor_size: int = 1): return handler +@pytest.fixture(autouse=True) +def reset_hybrid_ep(): + yield + reset_hybrid_ep_buffer() + + def test_offload_summary_uses_explicit_process_group(monkeypatch): from megatron.core.pipeline_parallel import fine_grained_activation_offload as off_module diff --git a/tests/unit_tests/pipeline_parallel/test_schedules.py b/tests/unit_tests/pipeline_parallel/test_schedules.py index fbfbc323faf..6c12a049a81 100644 --- a/tests/unit_tests/pipeline_parallel/test_schedules.py +++ b/tests/unit_tests/pipeline_parallel/test_schedules.py @@ -27,11 +27,18 @@ convert_schedule_table_to_order, get_overlap_moe_expert_parallel_comm_order, ) +from megatron.core.transformer.experimental_attention_variant.dsa import DSAIndexerLossAutoScaler from tests.unit_tests.test_utilities import Utils rank = Utils.rank +@pytest.fixture(autouse=True) +def reset_dsa_loss_scale(): + yield + DSAIndexerLossAutoScaler.main_loss_backward_scale = None + + def test_reset_activation_offload_uses_language_model_group(mocker): reset = mocker.patch.object(schedule.off_interface, "reset") language_group = object() diff --git a/tests/unit_tests/ssm/ops/mamba2/test_batch_invariant_decode.py b/tests/unit_tests/ssm/ops/mamba2/test_batch_invariant_decode.py index 6167947a5d6..33c0dc3581f 100644 --- a/tests/unit_tests/ssm/ops/mamba2/test_batch_invariant_decode.py +++ b/tests/unit_tests/ssm/ops/mamba2/test_batch_invariant_decode.py @@ -54,6 +54,14 @@ def setUpClass(cls): _pin_mamba_autotuners() + @classmethod + def tearDownClass(cls): + from megatron.core.transformer.custom_layers.batch_invariant_kernels import ( + _unpin_mamba_autotuners, + ) + + _unpin_mamba_autotuners() + def setUp(self): torch.manual_seed(0) # No global flags: batch_invariant_decode_buffered_scan is a pure tensor-ops function diff --git a/tests/unit_tests/ssm/test_gdn_gated_output_norm_fusion.py b/tests/unit_tests/ssm/test_gdn_gated_output_norm_fusion.py index 8311ee4f39e..3dd61ef4fb4 100644 --- a/tests/unit_tests/ssm/test_gdn_gated_output_norm_fusion.py +++ b/tests/unit_tests/ssm/test_gdn_gated_output_norm_fusion.py @@ -250,7 +250,7 @@ def test_post_fusion_distributed_layout( tp, cp, sequence_parallel, batch, packed, pre_fusion, recompute, value_heads ): """Enable the post fusion through CP redistribution and the complete backward.""" - if torch.distributed.get_world_size() < tp * cp: + if Utils.world_size < tp * cp: pytest.skip("This layout requires four distributed GPU ranks") Utils.initialize_model_parallel( tensor_model_parallel_size=tp, pipeline_model_parallel_size=1, context_parallel_size=cp diff --git a/tests/unit_tests/test_hyper_comm_grid.py b/tests/unit_tests/test_hyper_comm_grid.py index 0b15a6e4365..4bc026b4e10 100644 --- a/tests/unit_tests/test_hyper_comm_grid.py +++ b/tests/unit_tests/test_hyper_comm_grid.py @@ -9,6 +9,7 @@ import torch.distributed as dist from megatron.core.hyper_comm_grid import HyperCommGrid +from tests.unit_tests.test_utilities import Utils class TestHyperCommGrid: @@ -522,16 +523,11 @@ class TestHyperCommGridIntegration: @classmethod def setup_class(cls): """Set up distributed environment for the entire test class.""" - if not dist.is_initialized(): - # Initialize PyTorch distributed with NCCL backend - # This assumes proper environment variables are set (RANK, WORLD_SIZE, MASTER_ADDR, MASTER_PORT) - try: - dist.init_process_group(backend="nccl") - cls.distributed_initialized = True - except Exception as e: - pytest.skip(f"Cannot initialize distributed: {e}") - else: + try: + Utils.initialize_distributed() cls.distributed_initialized = True + except Exception as e: + pytest.skip(f"Cannot initialize distributed: {e}") def test_real_distributed_basic_functionality(self): """Test basic HyperCommGrid functionality with real distributed backend.""" diff --git a/tests/unit_tests/test_inference.py b/tests/unit_tests/test_inference.py index 3add3d1bc89..771d6777968 100644 --- a/tests/unit_tests/test_inference.py +++ b/tests/unit_tests/test_inference.py @@ -48,6 +48,7 @@ def mock_forward(*args, **kwargs): controller.inference_wrapped_model.model.forward = mock_forward yield engine_wrapper.static_engine + Utils.destroy_model_parallel() @pytest.fixture(scope="module") diff --git a/tests/unit_tests/test_num_microbatches_calculator.py b/tests/unit_tests/test_num_microbatches_calculator.py index 2faf9a9b9e4..0aa6a8db9b6 100644 --- a/tests/unit_tests/test_num_microbatches_calculator.py +++ b/tests/unit_tests/test_num_microbatches_calculator.py @@ -5,6 +5,12 @@ import megatron.core.num_microbatches_calculator as mb_calculator +@pytest.fixture(autouse=True) +def unset_calculator(): + yield + mb_calculator.unset_num_microbatches_calculator() + + def test_init_num_microbatches_calculator(): mb_calculator._GLOBAL_NUM_MICROBATCHES_CALCULATOR = None mb_calculator.init_num_microbatches_calculator( diff --git a/tests/unit_tests/test_utilities.py b/tests/unit_tests/test_utilities.py index 7a3f516d830..421651d0293 100644 --- a/tests/unit_tests/test_utilities.py +++ b/tests/unit_tests/test_utilities.py @@ -9,12 +9,21 @@ from torch.distributed import rendezvous import megatron.core.parallel_state as ps +from megatron.core.inference import utils as inference_utils +from megatron.core.inference.utils import InferenceMode +from megatron.core.tensor_parallel import random as tp_random +from megatron.core.transformer import cuda_graphs, multi_token_prediction from megatron.training.argument_utils import ( gpt_config_from_args, hybrid_config_from_args, pretrain_cfg_container_from_args, ) +try: + from transformer_engine.pytorch import distributed as te_distributed +except ImportError: + te_distributed = None + _NVTE_ATTN_ENV_VARS = ( 'NVTE_FLASH_ATTN', 'NVTE_FUSED_ATTN', @@ -48,6 +57,105 @@ def clear_nvte_env_vars(): os.environ.pop(name, None) +def reset_cuda_graph_global_state(): + """Reset the process-global CUDA-graph state a passing test can leave behind.""" + record = cuda_graphs._CudagraphGlobalRecord + # TestLLaVACudaGraph is an example of where the pool leaks. + # TestPackedSeqCudagraphs is an example that is affected by a leaked pool. + + # TestMHCWithCudaGraph is an example of where training records leak. + # TestLocalCudagraphPipelineOutput is an example that is affected by leaked training records. + if ( + record.cudagraph_record + or record.cudagraph_inference_record + or record.cudagraph_created + or record._saved_tensors_observer is not None + or cuda_graphs.CudaGraphManager.global_mempool is not None + ): + if torch.cuda.is_available(): + torch.cuda.synchronize() + cuda_graphs.delete_cuda_graphs() + + +def _snapshot_torch_settings(): + return { + "deterministic": ( + torch.are_deterministic_algorithms_enabled(), + torch.is_deterministic_algorithms_warn_only_enabled(), + ), + "fill_uninitialized_memory": torch.utils.deterministic.fill_uninitialized_memory, + "cudnn": ( + torch.backends.cudnn.deterministic, + torch.backends.cudnn.benchmark, + torch.backends.cudnn.allow_tf32, + ), + "matmul": ( + torch.backends.cuda.matmul.allow_tf32, + torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction, + torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction, + ), + } + + +def _restore_torch_settings(settings): + mode, warn_only = settings["deterministic"] + torch.use_deterministic_algorithms(mode, warn_only=warn_only) + torch.utils.deterministic.fill_uninitialized_memory = settings["fill_uninitialized_memory"] + ( + torch.backends.cudnn.deterministic, + torch.backends.cudnn.benchmark, + torch.backends.cudnn.allow_tf32, + ) = settings["cudnn"] + ( + torch.backends.cuda.matmul.allow_tf32, + torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction, + torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction, + ) = settings["matmul"] + + +def snapshot_process_state(): + """Capture process-global settings a test may change; restoring puts back the values.""" + state = { + # Set by every engine, cleared only by suspend(); breaks TestInferenceTopKRouter. + "inference_mode": (InferenceMode._is_active, InferenceMode._use_bounded_mxfp8_rows), + # The first init fixes the tracker type; + # TestPartialCudaGraph forces TE, TestMTPCudaGraphInference the no-op one. + "rng_tracker": ( + tp_random._CUDA_RNG_STATE_TRACKER, + tp_random._CUDA_RNG_STATE_TRACKER_INITIALIZED, + ), + # test_thd_format (deterministic), TestGPTModelBatchInvariant (tf32), + # test_guard_agrees_with_config_resolution (fill_uninitialized_memory). + "torch": _snapshot_torch_settings(), + } + if te_distributed is not None: + # TestParallelAttention fills it with byte tensors; + # test_forward_backward_func_with_full_cuda_graph then expects generators. + state["te_rng_states"] = te_distributed._ALL_ACTIVE_RNG_STATES + return state + + +def restore_process_state(state): + """Return every item captured by `snapshot_process_state` to its captured value.""" + InferenceMode._is_active, InferenceMode._use_bounded_mxfp8_rows = state["inference_mode"] + tp_random._CUDA_RNG_STATE_TRACKER, tp_random._CUDA_RNG_STATE_TRACKER_INITIALIZED = state[ + "rng_tracker" + ] + _restore_torch_settings(state["torch"]) + if "te_rng_states" in state: + te_distributed._ALL_ACTIVE_RNG_STATES = state["te_rng_states"] + + +def reset_transient_process_state(): + """Drop the lazily built caches no later test may inherit.""" + reset_cuda_graph_global_state() + # Sized at the first num_layers seen; TestMTPLossLoggingHelper reads it back. + multi_token_prediction.MTPLossLoggingHelper.tracker.clear() + # Built from the first model (TestMTPCudaGraphExpertParallel resets it by hand). + inference_utils.moe_layer_cache = None + inference_utils._moe_metadata_sync_initialized = False + + def is_nccl_ep_available(): """NCCL EP built into TE, with the ``ep_bootstrap`` signature mcore actually calls. @@ -135,13 +243,15 @@ class Utils: @staticmethod def initialize_distributed(): clear_nvte_env_vars() + if torch.cuda.is_available(): + # Also when another test already created the default group without binding one. + torch.cuda.set_device(Utils.local_rank % torch.cuda.device_count()) if not torch.distributed.is_initialized() and Utils.rank >= 0: print( f'Initializing torch.distributed with rank: {Utils.rank}, ' f'world_size: {Utils.world_size}' ) - torch.cuda.set_device(Utils.local_rank % torch.cuda.device_count()) init_method = 'tcp://' master_ip = os.getenv('MASTER_ADDR', 'localhost') master_port = os.getenv('MASTER_PORT', '29500') diff --git a/tests/unit_tests/test_utils.py b/tests/unit_tests/test_utils.py index 9e4a0c31f3c..24836fcb435 100644 --- a/tests/unit_tests/test_utils.py +++ b/tests/unit_tests/test_utils.py @@ -233,6 +233,7 @@ def _call_nvtx_range(): util.configure_nvtx_profiling(True) _call_nvtx_range() assert execution_tracker['ranges'] + util.configure_nvtx_profiling(False) def test_nvtx_decorator(): diff --git a/tests/unit_tests/tools/checkpoint/test_gpt_hybrid_conversion_parallelism.py b/tests/unit_tests/tools/checkpoint/test_gpt_hybrid_conversion_parallelism.py index 8102b7018ae..6ec28eda3bb 100644 --- a/tests/unit_tests/tools/checkpoint/test_gpt_hybrid_conversion_parallelism.py +++ b/tests/unit_tests/tools/checkpoint/test_gpt_hybrid_conversion_parallelism.py @@ -36,6 +36,8 @@ from gpt_hybrid_conversion import main as conversion_main +from tests.unit_tests.test_utilities import Utils + # These scenarios are SYNTHETIC and single-rank by design: each one writes a # tiny synthetic DCP checkpoint and round-trips it through the converter on @@ -53,7 +55,7 @@ # default PG is already multi-rank. @pytest.fixture(autouse=True) def _skip_when_multi_rank_pg(): - if dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1: + if Utils.world_size > 1: pytest.skip( "Synthetic single-rank tests skipped under a multi-rank default " "process group; multi-rank coverage is in " diff --git a/tests/unit_tests/tools/checkpoint/test_weighted_merge.py b/tests/unit_tests/tools/checkpoint/test_weighted_merge.py index e8fd3c3d460..629c0d3522c 100644 --- a/tests/unit_tests/tools/checkpoint/test_weighted_merge.py +++ b/tests/unit_tests/tools/checkpoint/test_weighted_merge.py @@ -20,6 +20,7 @@ from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_spec from megatron.core.transformer import TransformerConfig from tests.unit_tests.dist_checkpointing import TempNamedDir +from tests.unit_tests.test_utilities import Utils from tools.checkpoint import weighted_merge as weighted_merge_module from tools.checkpoint.weighted_merge import ( WeightedMergeError, @@ -30,12 +31,16 @@ @pytest.fixture def process_group(): - already_initialized = dist.is_available() and dist.is_initialized() + if torch.cuda.is_available(): + Utils.initialize_distributed() + yield + return + # CPU-only run: create a one-rank gloo group and destroy it afterwards. + created = not dist.is_initialized() weighted_merge_module._ensure_process_group() yield - if not already_initialized and dist.is_available() and dist.is_initialized(): - if dist.get_world_size() == 1: - dist.destroy_process_group() + if created: + dist.destroy_process_group() def _configured_world_size(): diff --git a/tests/unit_tests/training/config/test_container_base.py b/tests/unit_tests/training/config/test_container_base.py index b6641665a63..fbe63cb70bf 100644 --- a/tests/unit_tests/training/config/test_container_base.py +++ b/tests/unit_tests/training/config/test_container_base.py @@ -474,11 +474,11 @@ def test_to_yaml_save_to_file(self): assert parsed == config.to_dict() - def test_to_yaml_with_msc_url(self): + def test_to_yaml_with_msc_url(self, monkeypatch): """Test to_yaml with MSC URL.""" config = TestConfigContainer(name="msc_test", value=999) - MultiStorageClientFeature.enable() + monkeypatch.setattr(MultiStorageClientFeature, "_enabled", True) # Verify that the file is created in the temporary directory with tempfile.TemporaryDirectory() as temp_dir: diff --git a/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py b/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py index f7f190fbc62..b7c22b9b38a 100644 --- a/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py +++ b/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py @@ -1969,6 +1969,11 @@ def setup_method(self): yield Utils.destroy_model_parallel() + @pytest.fixture(autouse=True) + def reset_loss_scale(self): + yield + DSAIndexerLossAutoScaler.main_loss_backward_scale = None + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") def test_forward_pass(self): """Test that forward pass preserves output.""" diff --git a/tests/unit_tests/transformer/moe/test_paged_stashing.py b/tests/unit_tests/transformer/moe/test_paged_stashing.py index f8f11743b81..71b06e41249 100644 --- a/tests/unit_tests/transformer/moe/test_paged_stashing.py +++ b/tests/unit_tests/transformer/moe/test_paged_stashing.py @@ -1,5 +1,7 @@ # Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +import gc + import pytest import torch import torch.nn.functional as F @@ -8,14 +10,17 @@ from megatron.core.extensions.transformer_engine import HAVE_TE from megatron.core.fp8_utils import get_fp8_context from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec +from megatron.core.transformer.moe.fused_a2a import is_nccl_ep_bootstrapped, reset_hybrid_ep_buffer from megatron.core.transformer.moe.moe_layer import MoELayer from megatron.core.transformer.moe.moe_utils import get_align_size_for_quantization from megatron.core.transformer.moe.paged_stash import ( + PagedStashManager, _stash_buffer_dtype, check_paged_stash_overflow, paged_stash_init_chunk_handler, paged_stash_reset, ) +from megatron.core.transformer.moe.token_dispatcher import nccl_ep_release_context from megatron.core.transformer.spec_utils import get_submodules from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.utils import is_te_min_version @@ -33,6 +38,13 @@ pytestmark = pytest.mark.launch_on_gb200 +@pytest.fixture(autouse=True) +def reset_paged_stash_manager(): + """The manager's state machine never returns to 'begin' on its own.""" + yield + PagedStashManager.STASH_MGR = None + + def _global_tokens_per_expert_from_local_routing_map(routing_map: torch.Tensor) -> torch.Tensor: """Per-expert token counts from a local routing map, summed across the default process group. @@ -257,6 +269,7 @@ def setup_method(self, method): pass def teardown_method(self, method): + reset_hybrid_ep_buffer() Utils.destroy_model_parallel() @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") @@ -347,6 +360,7 @@ def setup_method(self, method): pass def teardown_method(self, method): + reset_hybrid_ep_buffer() Utils.destroy_model_parallel() @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") @@ -396,6 +410,15 @@ def test_overload_factor_and_over_budget(self): budget = int(num_tokens * capacity_factor) budget += -budget % pad_multiple + # This just captures; needed for setup. + paged_stash_reset(True, config=container.config) + paged_stash_init_chunk_handler(1, 0) + _forward_backward_all_layers(container, hidden_states) + container.zero_grad() + for layer in container.moe_layers: + layer.token_dispatcher.reset_over_budget() + + # This does the stashing needed for the test. paged_stash_reset(True, config=container.config) paged_stash_init_chunk_handler(1, 0) _forward_backward_all_layers(container, hidden_states) @@ -470,10 +493,24 @@ class TestNcclEpPagedStashing: """ def setup_method(self, method): - pass + self.container = None def teardown_method(self, method): - Utils.destroy_model_parallel() + try: + self.container = None + gc.collect() + nccl_ep_release_context() + assert not is_nccl_ep_bootstrapped(), "NCCL EP context survived release" + finally: + Utils.destroy_model_parallel() + + def _run_step(self, hidden_states): + paged_stash_reset(True, config=self.container.config) + paged_stash_init_chunk_handler(1, 0) + out, _, _, _ = _forward_backward_all_layers(self.container, hidden_states) + self.container.zero_grad() + torch.cuda.synchronize() + return out @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") @pytest.mark.internal @@ -498,7 +535,7 @@ def test_over_budget(self, wire_dtype): config.ENABLE_EXPERIMENTAL = True - container = MoEModelTestContainer( + self.container = MoEModelTestContainer( tp_size=1, ep_size=4, pp_size=1, @@ -526,34 +563,38 @@ def test_over_budget(self, wire_dtype): seq_length = 1024 batch_size = 1 - topk = container.config.moe_router_topk - capacity_factor = container.config.moe_expert_rank_capacity_factor + topk = self.container.config.moe_router_topk + capacity_factor = self.container.config.moe_expert_rank_capacity_factor hidden_states = torch.randn( - (seq_length, batch_size, container.config.hidden_size), dtype=torch.bfloat16 + (seq_length, batch_size, self.container.config.hidden_size), dtype=torch.bfloat16 ) num_tokens = seq_length * batch_size * topk - pad_multiple = get_align_size_for_quantization(container.config) + pad_multiple = get_align_size_for_quantization(self.container.config) budget = int(num_tokens * capacity_factor) budget += -budget % pad_multiple - paged_stash_reset(True, config=container.config) + paged_stash_reset(True, config=self.container.config) paged_stash_init_chunk_handler(1, 0) - _forward_backward_all_layers(container, hidden_states) + _forward_backward_all_layers(self.container, hidden_states) # NCCL EP's manager keeps token_probs/token_indices rather than the routing map, and a # rank's received load depends on every rank's routing, so the HybridEP map-derived # cross-check does not transfer. Check the device-side accounting against the budget # instead: required_recv is filled by ep_prepare before any dropping, and over_budget is # the same comparison made on device, so the two must agree with config arithmetic. + budget_report = [ + ( + layer.token_dispatcher._comm_manager._recv_capacity, + bool(layer.token_dispatcher.check_over_budget().item()), + int(layer.token_dispatcher.check_required_capacity().item()), + ) + for layer in self.container.moe_layers + ] any_over_budget = False - for layer_idx, layer in enumerate(container.moe_layers): - comm = layer.token_dispatcher._comm_manager - over_budget = layer.token_dispatcher.check_over_budget().item() - required = layer.token_dispatcher.check_required_capacity().item() - - assert comm._recv_capacity == budget, ( - f"layer {layer_idx}: dispatcher budget ({comm._recv_capacity}) != expected " + for layer_idx, (recv_capacity, over_budget, required) in enumerate(budget_report): + assert recv_capacity == budget, ( + f"layer {layer_idx}: dispatcher budget ({recv_capacity}) != expected " f"({budget}) for capacity factor {capacity_factor}" ) assert required > 0, f"layer {layer_idx}: required capacity was never recorded" @@ -568,19 +609,6 @@ def test_over_budget(self, wire_dtype): "the test is not exercising overflow" ) - # Leave a clean slate. The EP context is process-wide and ep_bootstrap refuses a second - # call, so a later test would otherwise reuse this capacity. Drop the layers before - # finalizing and force a collection: this container has no __del__, so its EpBuffers - # would otherwise be freed at an arbitrary later point -- inside the next test, against - # a context that has since been re-bootstrapped. - import gc - - from megatron.core.transformer.moe.token_dispatcher import nccl_ep_release_context - - del container - gc.collect() - nccl_ep_release_context() - @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") @pytest.mark.internal @pytest.mark.parametrize("zero_copy", [False, True]) @@ -591,11 +619,9 @@ def test_over_budget_recovery(self, zero_copy): if zero_copy and not is_nccl_ep_zero_copy_available(): pytest.skip("NCCL EP zero-copy TE API is not available") - from megatron.core.transformer.moe.token_dispatcher import nccl_ep_release_context - config.ENABLE_EXPERIMENTAL = True - container = MoEModelTestContainer( + self.container = MoEModelTestContainer( tp_size=1, ep_size=4, pp_size=1, @@ -623,39 +649,33 @@ def test_over_budget_recovery(self, zero_copy): seq_length = 1024 batch_size = 1 hidden_states = torch.randn( - (seq_length, batch_size, container.config.hidden_size), dtype=torch.bfloat16 + (seq_length, batch_size, self.container.config.hidden_size), dtype=torch.bfloat16 ) - def run(): - paged_stash_reset(True, config=container.config) - paged_stash_init_chunk_handler(1, 0) - out, _, _, _ = _forward_backward_all_layers(container, hidden_states) - container.zero_grad() - torch.cuda.synchronize() - return out - # 1. Undersized budget: the step drops tokens and reports it. - out_dropped = run() - over_1 = [l.token_dispatcher.check_over_budget().item() for l in container.moe_layers] - req_1 = [l.token_dispatcher.check_required_capacity().item() for l in container.moe_layers] + out_dropped = self._run_step(hidden_states) + over_1 = [l.token_dispatcher.check_over_budget().item() for l in self.container.moe_layers] + req_1 = [ + l.token_dispatcher.check_required_capacity().item() for l in self.container.moe_layers + ] required_t = torch.tensor([max(req_1)], dtype=torch.int64, device="cuda") torch.distributed.all_reduce(required_t, op=torch.distributed.ReduceOp.MAX) required = int(required_t.item()) - budget_before = container.moe_layers[0].token_dispatcher._comm_manager._recv_capacity + budget_before = self.container.moe_layers[0].token_dispatcher._comm_manager._recv_capacity # 2. prepare_for_rerun: clear the capacity factor (-> eager, which has no budget to # exceed), record the peak to grow to, release the EP context, replay dropless. - for layer in container.moe_layers: + for layer in self.container.moe_layers: layer.token_dispatcher.reset_over_budget() layer.token_dispatcher._comm_manager.moe_expert_rank_capacity_factor = None layer.token_dispatcher.grow_ep_recv_capacity(required) layer.token_dispatcher.invalidate_ep_bootstrap() nccl_ep_release_context() - out_replay = run() - eager_2 = [l.token_dispatcher._comm_manager.eager for l in container.moe_layers] - zc_2 = [l.token_dispatcher._comm_manager.zero_copy for l in container.moe_layers] + out_replay = self._run_step(hidden_states) + eager_2 = [l.token_dispatcher._comm_manager.eager for l in self.container.moe_layers] + zc_2 = [l.token_dispatcher._comm_manager.zero_copy for l in self.container.moe_layers] dropped_finite = bool(torch.isfinite(out_dropped).all()) replay_finite = bool(torch.isfinite(out_replay).all()) @@ -665,24 +685,25 @@ def run(): dropped_differs = not torch.allclose(out_dropped, out_replay, rtol=1e-2, atol=0) # 3. Success branch: restore the capacity factor -> static returns at the grown budget. - for layer in container.moe_layers: + for layer in self.container.moe_layers: layer.token_dispatcher.reset_over_budget() layer.token_dispatcher._comm_manager.moe_expert_rank_capacity_factor = ( - container.config.moe_expert_rank_capacity_factor + self.container.config.moe_expert_rank_capacity_factor ) layer.token_dispatcher.invalidate_ep_bootstrap() nccl_ep_release_context() - out_restored = run() - eager_3 = [l.token_dispatcher._comm_manager.eager for l in container.moe_layers] - zc_3 = [l.token_dispatcher._comm_manager.zero_copy for l in container.moe_layers] - caps_3 = [l.token_dispatcher._comm_manager._recv_capacity for l in container.moe_layers] - over_3 = [l.token_dispatcher.check_over_budget().item() for l in container.moe_layers] - nccl_ep_release_context() + out_restored = self._run_step(hidden_states) + eager_3 = [l.token_dispatcher._comm_manager.eager for l in self.container.moe_layers] + zc_3 = [l.token_dispatcher._comm_manager.zero_copy for l in self.container.moe_layers] + caps_3 = [ + l.token_dispatcher._comm_manager._recv_capacity for l in self.container.moe_layers + ] + over_3 = [l.token_dispatcher.check_over_budget().item() for l in self.container.moe_layers] assert required > budget_before, ( f"nothing exceeded budget {budget_before} at capacity factor " - f"{container.config.moe_expert_rank_capacity_factor}; not exercising overflow" + f"{self.container.config.moe_expert_rank_capacity_factor}; not exercising overflow" ) assert all(eager_2), f"replay did not degrade to eager: {eager_2}" assert not any(zc_2), f"replay must drop zero-copy while eager: {zc_2}" @@ -728,7 +749,7 @@ def test_forward_backward_4_layers(self, zero_copy, wire_dtype): config.ENABLE_EXPERIMENTAL = True - container = MoEModelTestContainer( + self.container = MoEModelTestContainer( tp_size=1, ep_size=4, pp_size=1, @@ -757,23 +778,24 @@ def test_forward_backward_4_layers(self, zero_copy, wire_dtype): seq_length = 1024 batch_size = 1 - hidden_size = container.config.hidden_size - hidden_states = torch.randn((seq_length, batch_size, hidden_size), dtype=torch.bfloat16) + hidden_states = torch.randn( + (seq_length, batch_size, self.container.config.hidden_size), dtype=torch.bfloat16 + ) # First iteration: capture schedule, capacity, etc. - paged_stash_reset(True, config=container.config) + paged_stash_reset(True, config=self.container.config) paged_stash_init_chunk_handler(1, 0) output_ref, hidden_states_grad_ref, routing_map_ref, tokens_per_expert_ref = ( - _forward_backward_all_layers(container, hidden_states) + _forward_backward_all_layers(self.container, hidden_states) ) - container.zero_grad() + self.container.zero_grad() # Second iteration: run with paged stash. - paged_stash_reset(True, config=container.config) + paged_stash_reset(True, config=self.container.config) paged_stash_init_chunk_handler(1, 0) output, hidden_states_grad, routing_map, tokens_per_expert = _forward_backward_all_layers( - container, hidden_states + self.container, hidden_states ) overflow = check_paged_stash_overflow() diff --git a/tests/unit_tests/transformer/test_full_cuda_graph.py b/tests/unit_tests/transformer/test_full_cuda_graph.py index 037b9dde287..2e102316d9d 100644 --- a/tests/unit_tests/transformer/test_full_cuda_graph.py +++ b/tests/unit_tests/transformer/test_full_cuda_graph.py @@ -9,12 +9,12 @@ import megatron.core.pipeline_parallel.schedules as schedule from megatron.core import ModelParallelConfig -from megatron.core.full_cuda_graph import FullCudaGraphWrapper, get_shared_capture_stream -from megatron.core.tensor_parallel.random import ( - HAVE_TE, - initialize_rng_tracker, - model_parallel_cuda_manual_seed, +from megatron.core.full_cuda_graph import ( + FullCudaGraphWrapper, + StaticBufferLoader, + get_shared_capture_stream, ) +from megatron.core.tensor_parallel.random import HAVE_TE, model_parallel_cuda_manual_seed from megatron.core.utils import is_te_min_version from megatron.training.models.dist_utils import _ddp_wrap from tests.unit_tests.test_utilities import Utils @@ -22,6 +22,16 @@ rank = Utils.rank +@pytest.fixture(autouse=True) +def reset_full_cuda_graph_state(): + """The wrapper keeps its graph and static buffers on the class.""" + yield + FullCudaGraphWrapper.curr_iteration = {'training': 0, 'validation': 0} + FullCudaGraphWrapper.cuda_graph = {'training': None, 'validation': None} + FullCudaGraphWrapper.result = {'training': None, 'validation': None} + StaticBufferLoader.static_buffers = {'training': [], 'validation': []} + + def test_ddp_grad_accumulators_share_full_cuda_graph_stream(): """Retained DDP AccumulateGrad nodes must use the full-iteration capture stream.""" @@ -97,8 +107,8 @@ def forward(self, inputs): def test_forward_backward_func_with_full_cuda_graph(mocker): from megatron.core.pipeline_parallel import get_forward_backward_func - initialize_rng_tracker(use_te_rng_tracker=True, force_reset=True) Utils.initialize_model_parallel(tensor_model_parallel_size=2, pipeline_model_parallel_size=1) + model_parallel_cuda_manual_seed(123, te_rng_tracker=True, force_reset_rng=True) def forward_step_func(data_iterator, model): import os diff --git a/tests/unit_tests/transformer/test_submodule_callables.py b/tests/unit_tests/transformer/test_submodule_callables.py index 5c4180075b2..029fcb88865 100644 --- a/tests/unit_tests/transformer/test_submodule_callables.py +++ b/tests/unit_tests/transformer/test_submodule_callables.py @@ -7,6 +7,7 @@ from megatron.core.models.gpt.gpt_layer_specs import ( get_gpt_layer_with_transformer_engine_submodules, ) +from megatron.core.transformer.moe.fused_a2a import reset_hybrid_ep_buffer from megatron.core.transformer.transformer_layer import TransformerLayer from megatron.core.utils import is_te_min_version from tests.unit_tests.a2a_overlap.utils import ( @@ -191,7 +192,7 @@ def setup_method(self, method): pass def teardown_method(self, method): - pass + reset_hybrid_ep_buffer() @pytest.mark.skipif(not is_te_min_version("1.9.0.dev0"), reason="Requires TE >= 1.9.0.dev0") @pytest.mark.parametrize("dispatcher_type", get_valid_token_dispatcher_types()) diff --git a/tests/unit_tests/transformer/test_transformer_block_custom_pgs.py b/tests/unit_tests/transformer/test_transformer_block_custom_pgs.py index 27f2bf5a3ff..8632c12edd4 100644 --- a/tests/unit_tests/transformer/test_transformer_block_custom_pgs.py +++ b/tests/unit_tests/transformer/test_transformer_block_custom_pgs.py @@ -215,8 +215,6 @@ def setup_method(self, method): def teardown_method(self, method): Utils.destroy_model_parallel() - torch.backends.cudnn.deterministic = False - torch.backends.cudnn.benchmark = True @pytest.mark.skipif( version.parse(torch.__version__) < version.parse('2.3.0'), diff --git a/tests/unit_tests/utils/test_experimental_log_once.py b/tests/unit_tests/utils/test_experimental_log_once.py index 59d8f3e4880..e6ae0172f9b 100644 --- a/tests/unit_tests/utils/test_experimental_log_once.py +++ b/tests/unit_tests/utils/test_experimental_log_once.py @@ -1,9 +1,12 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + import logging import torch from megatron.core import config from megatron.core import utils as mcore_utils +from tests.unit_tests.test_utilities import Utils # Message emitted by experimental_fn wrapper when EXPERIMENTAL flag is enabled. _LOG_MSG = "ENABLE_EXPERIMENTAL is True, running experimental code." @@ -16,6 +19,8 @@ def _get_test_logger(): def test_experimental_fn_logs_once(caplog): """Ensure the experimental_fn decorator writes the enable message only once.""" + if Utils.world_size > 1: + Utils.initialize_distributed() # Enable experimental features for this test. config.set_experimental_flag(True) @@ -57,6 +62,8 @@ def sample(): # pragma: no cover def test_experimental_cls_logs_once(caplog): """Ensure the experimental_cls decorator writes the enable message only once for classes.""" + if Utils.world_size > 1: + Utils.initialize_distributed() config.set_experimental_flag(True)