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
2 changes: 2 additions & 0 deletions tests/unit_tests/a2a_overlap/test_fsdp_1f1b_overlap.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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")
Expand Down
29 changes: 28 additions & 1 deletion tests/unit_tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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)
20 changes: 11 additions & 9 deletions tests/unit_tests/data/test_bin_reader.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.

import os
import random
import sys
Expand Down Expand Up @@ -100,9 +102,6 @@ def close(self) -> None:
pass


setattr(boto3, "client", _LocalClient)


##
# Overload ClientError from botocore.exceptions
##
Expand All @@ -114,8 +113,6 @@ class _LocalClientError(Exception):
pass


setattr(exceptions, "ClientError", _LocalClientError)

##
# Mock multistorageclient module
##
Expand All @@ -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")
Expand Down
7 changes: 7 additions & 0 deletions tests/unit_tests/data/test_get_batch.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
1 change: 1 addition & 0 deletions tests/unit_tests/dist_checkpointing/test_fully_parallel.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down
1 change: 1 addition & 0 deletions tests/unit_tests/dist_checkpointing/test_msc.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
5 changes: 5 additions & 0 deletions tests/unit_tests/dist_checkpointing/test_safe_globals.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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)
Expand All @@ -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)
Expand Down
2 changes: 1 addition & 1 deletion tests/unit_tests/distributed/test_param_and_grad_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
25 changes: 16 additions & 9 deletions tests/unit_tests/extension/test_kitchen_sdpa.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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()


# ---------------------------------------------------------------------------
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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
Expand Down Expand Up @@ -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()


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
15 changes: 11 additions & 4 deletions tests/unit_tests/inference/engines/test_dynamic_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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()


Expand Down
Loading
Loading