# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

import inspect
import operator
import sys
from copy import deepcopy
from types import SimpleNamespace
from unittest.mock import patch

import torch
from torch._decomp import get_decompositions
from torch._inductor.fx_passes.bucketing import (
    is_all_gather_into_tensor as is_all_gather,
)
from torch.cuda._graph_annotations import _is_tools_id_unavailable
from torch.fx.experimental.proxy_tensor import make_fx
from torch.fx.experimental.symbolic_shapes import ShapeEnv
from torch.fx.passes.fake_tensor_prop import FakeTensorProp
from torch.fx.traceback import preserve_node_meta
from torch.testing._internal.common_fsdp import FSDPTest
from torch.testing._internal.common_utils import TestCase
from torch.utils.checkpoint import checkpoint, CheckpointPolicy

from torchtitan.components.data.types import TokenizedTrainingMicrobatch
from torchtitan.distributed import ParallelismContext
from torchtitan.experiments.graph_trainer.common_utils import (
    _EP_TOKEN_COUNT_EXCHANGE,
    _EP_TOKEN_COUNT_SYNC,
    _EP_TOKEN_EXCHANGE,
    _EP_TOKEN_EXCHANGE_WAIT,
    _MODULE_FQN,
    annotate_module_fqns,
    annotate_moe_ep_regions,
    get_default_transformer_block_buckets,
)
from torchtitan.experiments.graph_trainer.configs import (
    EpOverlapConfig,
    GraphTrainerCompileConfig,
)
from torchtitan.experiments.graph_trainer.cuda_graph import (
    insert_kernel_annotations_pass,
    is_cuda_graph_fully_compatible,
    is_cuda_graph_node_compatible,
)
from torchtitan.experiments.graph_trainer.decompositions import (
    apply_decompositions_pass,
)
from torchtitan.experiments.graph_trainer.ep_eager_chunk import (
    maybe_apply_ep_overlap_eager_chunking,
    populate_eager_chunk_metadata_pass,
)
from torchtitan.experiments.graph_trainer.ep_overlap_pass import (
    _apply_schedule,
    _schedule_ep_overlap_regions,
    _ScheduledRegion,
)
from torchtitan.experiments.graph_trainer.ep_pass_utils import (
    _chunk_owner,
    ChunkBody,
    ChunkedRegion,
    ChunkOwner,
)
from torchtitan.experiments.graph_trainer.ep_process_group_pass import (
    isolate_ep_process_group_pass,
)
from torchtitan.experiments.graph_trainer.fsdp_passes import (
    _FSDP_BUCKET_META,
    _validate_transformer_block_bucket_counts,
    deduplicate_fsdp_unshard_chains_pass,
    get_transformer_block_bucket_counts,
    get_transformer_block_layer_ids,
    reassign_collective_pgs_pass,
    schedule_fsdp_comms_to_dense_regions_pass,
)
from torchtitan.experiments.graph_trainer.make_fx_tracer import (
    GraphStateSpec,
    minimal_fx_tracer,
)
from torchtitan.experiments.graph_trainer.memory_policy import (
    _backward_side_nodes,
    _default_memory_policy_pass,
    _make_default_memory_policy,
    _make_full_memory_policy,
    tag_min_cut_saved_values,
    tag_sac_policy,
    tag_with_memory_policy_pass,
    validate_memory_policy_config,
)
from torchtitan.experiments.graph_trainer.passes import (
    compile_time_passes,
    selective_activation_remat_pass,
)
from torchtitan.experiments.graph_trainer.remove_noop_passes import (
    canonicalize_graph_pass,
    eliminate_dead_code_pass,
    normalize_view_ops_as_reshape,
    remove_b2b_transpose_pass,
    remove_detach_pass,
    remove_identity_slice_pass,
    remove_identity_view_pass,
)
from torchtitan.experiments.graph_trainer.simple_fsdp import (
    data_parallel,
    FSDP_PARAM_FQNS_META,
)
from torchtitan.experiments.graph_trainer.subgraph_regions import (
    apply_subgraph_region_annotations_pass,
    SUBGRAPH_REGION,
    SUBGRAPH_REGION_ROLE,
)
from torchtitan.experiments.graph_trainer.tests.test_cpu_offload import (  # noqa: F401
    TestCpuOffloadPass,
)
from torchtitan.experiments.graph_trainer.tests.test_custom_codegen import (  # noqa: F401
    TestCustomCodegenPass,
)
from torchtitan.experiments.graph_trainer.tests.test_performance_passes import (  # noqa: F401
    TestAnnotateRMSNormForRegionalInductorPass,
)
from torchtitan.models.common.linear import Linear
from torchtitan.protocols.module import Module, ModuleList


@torch.library.custom_op(
    "torchtitan_graph_trainer_test::ordered_identity", mutates_args=()
)
def _ordered_identity(x: torch.Tensor) -> torch.Tensor:
    return x.clone()


@_ordered_identity.register_fake
def _ordered_identity_fake(x: torch.Tensor) -> torch.Tensor:
    return torch.empty_like(x)


_ordered_identity.register_effect(torch.library.EffectType.ORDERED)


class TestDefaultTransformerBlockBuckets(TestCase):
    def test_compile_time_passes_enable_chunked_loss_bucket_only_when_needed(self):
        from torchtitan.components.loss import ChunkedLossWrapper, CrossEntropyLoss
        from torchtitan.experiments.graph_trainer.configs import (
            GraphTrainerCompileConfig,
        )
        from torchtitan.experiments.graph_trainer.passes import compile_time_passes

        def make_config(loss):
            return SimpleNamespace(
                compile=GraphTrainerCompileConfig(inductor_compilation="full"),
                loss=loss,
                model=SimpleNamespace(layers=[0, 1]),
                parallelism=SimpleNamespace(),
            )

        traced_result = SimpleNamespace(
            state_fqns=[],
            graph_state=GraphStateSpec(),
        )
        with patch(
            "torchtitan.experiments.graph_trainer.common_utils."
            "get_default_transformer_block_buckets",
            return_value=[],
        ) as mock_bucket_plan:
            compile_time_passes(traced_result, make_config(CrossEntropyLoss.Config()))
            compile_time_passes(traced_result, make_config(ChunkedLossWrapper.Config()))

        self.assertEqual(
            [
                call.kwargs["chunked_loss_enabled"]
                for call in mock_bucket_plan.call_args_list
            ],
            [False, True],
        )


class TestFSDPUnshardDedupPass(TestCase):
    def _duplicate_unshard_graph(self) -> torch.fx.GraphModule:
        graph = torch.fx.Graph()
        sharded_param = graph.placeholder("sharded_param")
        x = graph.placeholder("x")

        all_gather_0 = graph.call_function(
            torch.ops._c10d_functional.all_gather_into_tensor.default,
            args=(sharded_param, 1, "fsdp_pg"),
        )
        wait_0 = graph.call_function(
            torch.ops._c10d_functional.wait_tensor.default,
            args=(all_gather_0,),
        )
        unsharded_0 = graph.call_function(
            torch.ops.aten.view.default,
            args=(wait_0, [4]),
        )

        all_gather_1 = graph.call_function(
            torch.ops._c10d_functional.all_gather_into_tensor.default,
            args=(sharded_param, 1, "fsdp_pg"),
        )
        wait_1 = graph.call_function(
            torch.ops._c10d_functional.wait_tensor.default,
            args=(all_gather_1,),
        )
        unsharded_1 = graph.call_function(
            torch.ops.aten.view.default,
            args=(wait_1, [4]),
        )
        fsdp_meta = {FSDP_PARAM_FQNS_META: ("linear.weight",)}
        for node in (
            all_gather_0,
            wait_0,
            unsharded_0,
            all_gather_1,
            wait_1,
            unsharded_1,
        ):
            node.meta["custom"] = fsdp_meta

        params = graph.call_function(
            torch.ops.aten.add.Tensor,
            args=(unsharded_0, unsharded_1),
        )
        out = graph.call_function(torch.ops.aten.add.Tensor, args=(params, x))
        graph.output(out)

        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        gm.graph.lint()
        gm.recompile()
        return gm

    def test_duplicate_fsdp_unshard_chains_are_canonicalized(self) -> None:
        gm = self._duplicate_unshard_graph()

        self.assertEqual(
            sum(1 for node in gm.graph.nodes if is_all_gather(node)),
            2,
        )

        deduplicate_fsdp_unshard_chains_pass(gm)

        self.assertEqual(
            sum(1 for node in gm.graph.nodes if is_all_gather(node)),
            1,
        )
        gm.graph.lint()

    def test_duplicate_fsdp_unshard_prefers_live_output(self) -> None:
        gm = self._duplicate_unshard_graph()
        unsharded_0, unsharded_1 = gm.graph.find_nodes(
            op="call_function",
            target=torch.ops.aten.view.default,
        )
        params = next(iter(unsharded_0.users))
        params.replace_input_with(unsharded_0, unsharded_1)
        gm.graph.eliminate_dead_code()

        deduplicate_fsdp_unshard_chains_pass(gm)

        self.assertEqual(
            sum(1 for node in gm.graph.nodes if is_all_gather(node)),
            1,
        )
        self.assertEqual(params.args, (unsharded_1, unsharded_1))
        gm.graph.lint()

    def test_unshard_shared_by_chunks_has_no_chunk_owner(self) -> None:
        gm = self._duplicate_unshard_graph()
        fsdp_nodes = [
            node
            for node in gm.graph.nodes
            if node.meta.get("custom", {}).get(FSDP_PARAM_FQNS_META)
        ]
        self.assertEqual(len(fsdp_nodes), 6)
        for chunk_id, nodes in enumerate((fsdp_nodes[:3], fsdp_nodes[3:])):
            for node in nodes:
                custom = dict(node.meta["custom"])
                custom.update(
                    {
                        "chunk_id": chunk_id,
                        "chunked_region_fqn": "layers.0.moe",
                        "chunked_region_role": "body",
                    }
                )
                node.meta["custom"] = custom

        deduplicate_fsdp_unshard_chains_pass(gm)

        all_gathers = [node for node in gm.graph.nodes if is_all_gather(node)]
        self.assertEqual(len(all_gathers), 1)
        all_gather = all_gathers[0]
        self.assertIsNone(_chunk_owner(all_gather))

        self.assertEqual(len(all_gather.users), 1)
        wait = next(iter(all_gather.users))
        self.assertIsNone(_chunk_owner(wait))
        self.assertEqual(
            {_chunk_owner(node) for node in wait.users},
            {
                ChunkOwner("layers.0.moe", False, 0),
                ChunkOwner("layers.0.moe", False, 1),
            },
        )


class ToyModel(Module):
    """A small toy model with multiple linear layers and activation
    checkpointing so that the backward graph recomputes the forward
    all-gathers."""

    def __init__(self, dim=16, n_layers=3):
        super().__init__()

        def _make_linear():
            cfg = Linear.Config(in_features=dim, out_features=dim, bias=True)
            return cfg.build()

        self.layers = ModuleList([_make_linear() for _ in range(n_layers)])

    def forward(self, x):
        for layer in self.layers:
            x = checkpoint(
                lambda m, inp: torch.relu(m(inp)),
                layer,
                x,
                use_reentrant=False,
            )
        return x


class TestReassignCollectivePgsPass(FSDPTest):
    """Integration tests: toy model + simple_fsdp + minimal_fx_tracer + reassign_collective_pgs_pass."""

    def _setup(self):
        """Set up ParallelismContext and device mesh for FSDP."""
        self.parallelism_context = ParallelismContext(
            dp_shard=-1,
            dp_replicate=1,
            cp=1,
            tp=1,
            pp=1,
            ep=1,
            world_size=self.world_size,
            enable_sequence_parallel=False,
        )

    def _make_fsdp_model(self, dim=16, n_layers=3):
        """Create a toy model and apply simple_fsdp data_parallel."""
        model = ToyModel(dim, n_layers).cuda()
        from torchtitan.experiments.graph_trainer.common_utils import (
            get_simple_fsdp_mesh,
        )

        fsdp_mesh = get_simple_fsdp_mesh(self.parallelism_context)
        model = data_parallel(model, device_mesh=fsdp_mesh, mode="fully_shard")
        return model

    def _get_fsdp_pg_name(self):
        """Get the FSDP process group name from the mesh."""
        from torchtitan.experiments.graph_trainer.common_utils import (
            get_simple_fsdp_mesh,
        )

        fsdp_mesh = get_simple_fsdp_mesh(self.parallelism_context)
        return fsdp_mesh.get_group().group_name

    def _trace_joint_graph(self, model, inputs):
        """Trace the joint fwd+bwd graph with minimal_fx_tracer."""

        def fwd_bwd_step(x):
            loss = model(x).sum()
            params = [p for p in model.parameters() if p.requires_grad]
            grads = torch.autograd.grad(loss, params)
            return [loss, *grads]

        traced = minimal_fx_tracer(fwd_bwd_step, module=model)(inputs)
        return traced.gm, traced.example_inputs

    def _count_ag_nodes_with_pg(self, gm, pg_name):
        """Count all-gather nodes in the graph that use the given PG name."""
        count = 0
        for node in gm.graph.nodes:
            if is_all_gather(node) and node.args[2] == pg_name:
                count += 1
        return count

    def _count_all_ag_nodes(self, gm):
        """Count all all-gather nodes in the graph regardless of PG."""
        count = 0
        for node in gm.graph.nodes:
            if is_all_gather(node):
                count += 1
        return count

    def _count_rs_nodes_with_pg(self, gm, pg_name):
        return sum(
            1
            for node in gm.graph.nodes
            if node.op == "call_function"
            and node.target is torch.ops._c10d_functional.reduce_scatter_tensor.default
            and node.args[3] == pg_name
        )

    def _count_ep_a2a_nodes_with_pg(self, gm, pg_name):
        return sum(
            1
            for node in gm.graph.nodes
            if node.op == "call_function"
            and "all_to_all_single" in str(node.target)
            and node.args[3] == pg_name
        )

    def test_overlap_rewrites_ag_nodes(self):
        """Apply reassign_collective_pgs_pass on the traced joint graph and verify
        that FSDP AG nodes are rewritten to the auto-created extra PG."""
        from torchtitan.experiments.graph_trainer.fsdp_passes import (
            _EXTRA_FSDP_PG_REGISTRY,
        )

        self._setup()
        model = self._make_fsdp_model()
        inputs = torch.randn(4, 16).cuda()
        fsdp_pg_name = self._get_fsdp_pg_name()

        gm, example_inputs = self._trace_joint_graph(model, inputs)

        # Before: all AG nodes should use the FSDP PG
        ag_before = self._count_ag_nodes_with_pg(gm, fsdp_pg_name)
        self.assertGreater(ag_before, 0, "Expected AG nodes with FSDP PG name")

        _EXTRA_FSDP_PG_REGISTRY.pop(fsdp_pg_name, None)
        reassign_collective_pgs_pass(gm, example_inputs)

        extra_pg_name = _EXTRA_FSDP_PG_REGISTRY[fsdp_pg_name]
        ag_with_old = self._count_ag_nodes_with_pg(gm, fsdp_pg_name)
        ag_with_new = self._count_ag_nodes_with_pg(gm, extra_pg_name)

        self.assertEqual(ag_with_old, 0, "No AG nodes should still use the old PG")
        self.assertEqual(
            ag_with_new,
            ag_before,
            "All AG nodes should now use the extra PG",
        )

    def test_overlap_preserves_total_ag_count(self):
        """The pass should not add or remove AG nodes, only rewrite PG names."""
        self._setup()
        model = self._make_fsdp_model()
        inputs = torch.randn(4, 16).cuda()

        gm, example_inputs = self._trace_joint_graph(model, inputs)

        total_before = self._count_all_ag_nodes(gm)
        reassign_collective_pgs_pass(gm, example_inputs)
        total_after = self._count_all_ag_nodes(gm)

        self.assertEqual(total_before, total_after)

    def test_overlap_rewrites_multiple_pgs(self):
        """When the graph has AG nodes from multiple FSDP PGs (e.g. FSDP +
        expert-FSDP), each source PG should be mapped to its own extra PG."""
        import torch.distributed as dist

        from torchtitan.experiments.graph_trainer.fsdp_passes import (
            _EXTRA_FSDP_PG_REGISTRY,
        )

        self._setup()
        model = self._make_fsdp_model()
        inputs = torch.randn(4, 16).cuda()
        fsdp_pg_name = self._get_fsdp_pg_name()

        gm, example_inputs = self._trace_joint_graph(model, inputs)

        # Create a second PG to simulate expert-FSDP.
        second_pg = dist.new_group(
            ranks=list(range(self.world_size)),
            use_local_synchronization=True,
        )
        second_pg_name = second_pg.group_name

        # Rewrite half the AG nodes to use the second PG.
        ag_nodes = [n for n in gm.graph.nodes if is_all_gather(n)]
        self.assertGreater(len(ag_nodes), 1)
        half = len(ag_nodes) // 2
        for node in ag_nodes[:half]:
            node.args = (node.args[0], node.args[1], second_pg_name)

        ag_pg1_before = self._count_ag_nodes_with_pg(gm, fsdp_pg_name)
        ag_pg2_before = self._count_ag_nodes_with_pg(gm, second_pg_name)
        self.assertGreater(ag_pg1_before, 0)
        self.assertGreater(ag_pg2_before, 0)

        _EXTRA_FSDP_PG_REGISTRY.pop(fsdp_pg_name, None)
        _EXTRA_FSDP_PG_REGISTRY.pop(second_pg_name, None)
        reassign_collective_pgs_pass(gm, example_inputs)

        # Both source PGs should have their own extra PG.
        self.assertIn(fsdp_pg_name, _EXTRA_FSDP_PG_REGISTRY)
        self.assertIn(second_pg_name, _EXTRA_FSDP_PG_REGISTRY)
        extra_pg1 = _EXTRA_FSDP_PG_REGISTRY[fsdp_pg_name]
        extra_pg2 = _EXTRA_FSDP_PG_REGISTRY[second_pg_name]
        self.assertNotEqual(
            extra_pg1, extra_pg2, "Each source PG must map to a distinct extra PG"
        )

        # No AG nodes should still use original PGs.
        self.assertEqual(self._count_ag_nodes_with_pg(gm, fsdp_pg_name), 0)
        self.assertEqual(self._count_ag_nodes_with_pg(gm, second_pg_name), 0)

        # All AG nodes should use their respective extra PGs.
        self.assertEqual(self._count_ag_nodes_with_pg(gm, extra_pg1), ag_pg1_before)
        self.assertEqual(self._count_ag_nodes_with_pg(gm, extra_pg2), ag_pg2_before)

    def test_overlap_rewrites_ep_a2a_on_fsdp_pg_to_separate_pg(self):
        from torchtitan.experiments.graph_trainer.ep_process_group_pass import (
            _EXTRA_EP_PG_REGISTRY,
        )
        from torchtitan.experiments.graph_trainer.fsdp_passes import (
            _EXTRA_FSDP_PG_REGISTRY,
        )

        self._setup()
        fsdp_pg_name = self._get_fsdp_pg_name()
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional
        ag = graph.call_function(
            c10d.all_gather_into_tensor.default, args=(x, 1, fsdp_pg_name)
        )
        wait = graph.call_function(c10d.wait_tensor.default, args=(ag,))
        rs = graph.call_function(
            c10d.reduce_scatter_tensor.default, args=(x, "sum", 1, fsdp_pg_name)
        )
        rs_wait = graph.call_function(c10d.wait_tensor.default, args=(rs,))
        a2a = graph.call_function(
            c10d.all_to_all_single.default, args=(x, [], [], fsdp_pg_name)
        )
        a2a.meta["custom"] = {
            _MODULE_FQN: "layers.0.moe",
            "EP": "dispatch",
            _EP_TOKEN_EXCHANGE: "dispatch",
        }
        graph.output((wait, rs_wait, a2a))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        _EXTRA_FSDP_PG_REGISTRY.pop(fsdp_pg_name, None)
        _EXTRA_EP_PG_REGISTRY.pop(fsdp_pg_name, None)
        reassign_collective_pgs_pass(gm, ())
        isolate_ep_process_group_pass(gm, ())

        fsdp_extra_pg = _EXTRA_FSDP_PG_REGISTRY[fsdp_pg_name]
        ep_extra_pg = _EXTRA_EP_PG_REGISTRY[fsdp_pg_name]
        self.assertNotEqual(fsdp_extra_pg, ep_extra_pg)
        self.assertEqual(self._count_ag_nodes_with_pg(gm, fsdp_extra_pg), 1)
        self.assertEqual(self._count_ep_a2a_nodes_with_pg(gm, ep_extra_pg), 1)


class TestFsdpDenseSchedulerPass(TestCase):
    """Pure FX tests for FSDP dense scheduling; no process group required."""

    def _tag_fsdp_schedule_node(self, node, fqn, *, backward=False):
        node.meta["custom"] = {_MODULE_FQN: fqn}
        if backward:
            node.meta["autograd_backward"] = True
        return node

    def _tag_fsdp_chain(self, node, param_fqns, direction):
        custom = dict(node.meta.get("custom", {}))
        custom[FSDP_PARAM_FQNS_META] = tuple(param_fqns)
        node.meta["custom"] = custom
        if direction == "bwd":
            node.meta["autograd_backward"] = True
        return node

    def _tag_fsdp_bucket(self, node, plan_fqns, direction):
        self._tag_fsdp_chain(node, plan_fqns, direction)
        node.meta[_FSDP_BUCKET_META] = {
            "plan_fqns": tuple(plan_fqns),
            "direction": direction,
        }
        return node

    def _node_order(self, gm):
        return {node: i for i, node in enumerate(gm.graph.nodes)}

    def _add_bucketed_ag_with_late_input(
        self,
        graph,
        source,
        *,
        plan_fqn,
        direction,
        group_name="pg",
    ):
        prep = graph.call_function(torch.ops.aten.detach.default, args=(source,))
        bucket = graph.call_function(
            torch.ops.bucketing._pre_bucket_all_gather.default,
            args=([prep], 2, torch.float32, [0], 0),
        )
        shard = graph.call_function(torch.ops.aten.slice.Tensor, args=(bucket, 0, 0, 1))
        coll = graph.call_function(
            torch.ops._c10d_functional.all_gather_into_tensor_out.default,
            args=(shard, 2, group_name),
            kwargs={"out": bucket},
        )
        wait = graph.call_function(
            torch.ops._c10d_functional.wait_tensor.default, args=(coll,)
        )
        self._tag_fsdp_chain(prep, [plan_fqn], direction)
        for node in (bucket, shard, coll, wait):
            self._tag_fsdp_bucket(node, [plan_fqn], direction)
        return prep, coll, wait

    def test_transformer_block_bucket_counts_follow_bucket_plan(self):
        counts = get_transformer_block_bucket_counts(
            [
                "tok_embeddings",
                "layers.0",
                [
                    "layers.1.attention_norm",
                    "layers.1.attention",
                    "layers.1.ffn_norm",
                    "layers.1.moe.router",
                    "layers.1.moe.shared_experts",
                ],
                "layers.1.moe.routed_experts",
                ["norm", "lm_head"],
            ],
            n_layers=2,
        )

        self.assertEqual(counts, {0: 1, 1: 2})

    def test_transformer_block_layer_ids_follow_traced_state(self):
        layer_ids = get_transformer_block_layer_ids(
            [
                "tok_embeddings.weight",
                "layers.1.attention.wq.weight",
                "layers.1.ffn_norm.weight",
                "layers.4.moe.router.weight",
                "norm.weight",
            ],
            n_layers=6,
        )

        self.assertEqual(layer_ids, frozenset({1, 4}))

    def test_transformer_block_layer_ids_reject_out_of_range_layer(self):
        with self.assertRaisesRegex(ValueError, r"references layers \[6\]"):
            get_transformer_block_layer_ids(
                ["layers.1.attention.weight", "layers.6.moe.router.weight"],
                n_layers=6,
            )

    def test_fsdp_dense_scheduler_accepts_empty_local_stage(self):
        _validate_transformer_block_bucket_counts(
            {},
            n_layers=2,
            expected_bucket_counts={0: 1, 1: 1},
            local_layer_ids=frozenset(),
            require_backward_all_gathers=True,
        )

    def test_fsdp_dense_scheduler_rejects_collective_outside_local_stage(self):
        graph = torch.fx.Graph()
        node = graph.placeholder("x")
        comms = {
            1: {
                "fwd_ag": [(node, node)],
                "bwd_ag": [(node, node)],
                "bwd_rs": [(node, node)],
            }
        }

        with self.assertRaisesRegex(ValueError, r"outside local layers \[1\]"):
            _validate_transformer_block_bucket_counts(
                comms,
                n_layers=2,
                expected_bucket_counts={0: 1, 1: 1},
                local_layer_ids=frozenset(),
                require_backward_all_gathers=True,
            )

    def test_fsdp_dense_scheduler_prefetches_full_ac_ag_input_chains(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")

        fwd_dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.0.attention",
        )
        fwd_buckets = [
            self._add_bucketed_ag_with_late_input(
                graph,
                x,
                plan_fqn=plan_fqn,
                direction="fwd",
            )
            for plan_fqn in (
                "layers.1.attention",
                "layers.1.moe.routed_experts.w13",
            )
        ]
        fwd_dense1 = self._tag_fsdp_schedule_node(
            graph.call_function(
                torch.ops.aten.add.Tensor,
                args=(fwd_buckets[0][2], fwd_buckets[1][2]),
            ),
            "layers.1.attention",
        )
        bwd_dense1 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(fwd_dense1,)),
            "layers.1.attention",
            backward=True,
        )
        bwd_dense1.name = f"{bwd_dense1.name}_recomputed"
        bwd_buckets = [
            self._add_bucketed_ag_with_late_input(
                graph,
                x,
                plan_fqn=plan_fqn,
                direction="bwd",
            )
            for plan_fqn in (
                "layers.0.attention",
                "layers.0.moe.routed_experts.w13",
            )
        ]
        bwd_dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(
                torch.ops.aten.add.Tensor,
                args=(bwd_buckets[0][2], bwd_buckets[1][2]),
            ),
            "layers.0.attention",
            backward=True,
        )
        bwd_dense0.name = f"{bwd_dense0.name}_recomputed"
        graph.output((fwd_dense0, fwd_dense1, bwd_dense1, bwd_dense0))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm,
            moe_layer_ids=frozenset({0, 1}),
            n_layers=2,
            strict=True,
        )

        gm.graph.lint()
        order = self._node_order(gm)
        for prep, coll, wait in fwd_buckets:
            self.assertLess(order[prep], order[coll])
            self.assertLess(order[coll], order[fwd_dense0])
            self.assertLess(order[fwd_dense0], order[wait])
        self.assertLess(order[fwd_buckets[0][1]], order[fwd_buckets[1][1]])
        for prep, coll, wait in bwd_buckets:
            self.assertLess(order[prep], order[coll])
            self.assertLess(order[coll], order[bwd_dense1])
            self.assertLess(order[bwd_dense1], order[wait])
        self.assertLess(order[bwd_buckets[0][1]], order[bwd_buckets[1][1]])

    def test_fsdp_dense_scheduler_prefetches_padded_ag_input_chain(self):
        graph = torch.fx.Graph()
        parameter = graph.placeholder("parameter")
        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(parameter,)),
            "layers.0.attention",
        )
        padded = graph.call_function(
            torch.ops.aten.constant_pad_nd.default,
            args=(parameter, [0, 1], 0.0),
        )
        bucket = graph.call_function(
            torch.ops.bucketing._pre_bucket_all_gather.default,
            args=([padded], 2, torch.float32, [0], 0),
        )
        shard = graph.call_function(torch.ops.aten.slice.Tensor, args=(bucket, 0, 0, 1))
        coll = self._tag_fsdp_bucket(
            graph.call_function(
                torch.ops._c10d_functional.all_gather_into_tensor_out.default,
                args=(shard, 2, "pg"),
                kwargs={"out": bucket},
            ),
            ["layers.1.attention"],
            "fwd",
        )
        wait = graph.call_function(
            torch.ops._c10d_functional.wait_tensor.default, args=(coll,)
        )
        for node in (padded, bucket, shard, wait):
            self._tag_fsdp_chain(node, ["layers.1.attention"], "fwd")
        dense1 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(wait,)),
            "layers.1.attention",
        )
        graph.output((dense0, dense1))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm,
            moe_layer_ids=frozenset(),
            n_layers=2,
            strict=True,
        )

        gm.graph.lint()
        order = self._node_order(gm)
        self.assertLess(order[padded], order[bucket])
        self.assertLess(order[bucket], order[shard])
        self.assertLess(order[shard], order[coll])
        self.assertLess(order[coll], order[dense0])
        self.assertLess(order[dense0], order[wait])

    def test_fsdp_dense_scheduler_prefetches_shared_ag_input_once(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")

        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.0.attention",
        )
        shared_prep = graph.call_function(torch.ops.aten.detach.default, args=(x,))
        self._tag_fsdp_chain(
            shared_prep,
            [
                "layers.1.attention",
                "layers.1.moe.routed_experts.w13",
            ],
            "fwd",
        )
        buckets = []
        for plan_fqn in (
            "layers.1.attention",
            "layers.1.moe.routed_experts.w13",
        ):
            bucket = graph.call_function(
                torch.ops.bucketing._pre_bucket_all_gather.default,
                args=([shared_prep], 2, torch.float32, [0], 0),
            )
            shard = graph.call_function(
                torch.ops.aten.slice.Tensor, args=(bucket, 0, 0, 1)
            )
            coll = graph.call_function(
                torch.ops._c10d_functional.all_gather_into_tensor_out.default,
                args=(shard, 2, "pg"),
                kwargs={"out": bucket},
            )
            wait = graph.call_function(
                torch.ops._c10d_functional.wait_tensor.default, args=(coll,)
            )
            for node in (bucket, shard, coll, wait):
                self._tag_fsdp_bucket(node, [plan_fqn], "fwd")
            buckets.append((coll, wait))
        dense1 = self._tag_fsdp_schedule_node(
            graph.call_function(
                torch.ops.aten.add.Tensor,
                args=(buckets[0][1], buckets[1][1]),
            ),
            "layers.1.attention",
        )
        graph.output((dense0, dense1))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm,
            moe_layer_ids=frozenset({1}),
            n_layers=2,
            strict=True,
        )

        gm.graph.lint()
        order = self._node_order(gm)
        for coll, wait in buckets:
            self.assertLess(order[shared_prep], order[coll])
            self.assertLess(order[coll], order[dense0])
            self.assertLess(order[dense0], order[wait])
        self.assertLess(order[buckets[0][0]], order[buckets[1][0]])

    def test_fsdp_dense_scheduler_prefetches_ag_with_group_placeholder(self):
        graph = torch.fx.Graph()
        parameter = graph.placeholder("parameter")
        group_name = graph.placeholder("group_name")
        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(parameter,)),
            "layers.0.attention",
        )
        prep, coll, wait = self._add_bucketed_ag_with_late_input(
            graph,
            parameter,
            plan_fqn="layers.1.attention",
            direction="fwd",
            group_name=group_name,
        )
        dense1 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(wait,)),
            "layers.1.attention",
        )
        graph.output((dense0, dense1))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm,
            moe_layer_ids=frozenset(),
            n_layers=2,
            strict=True,
        )

        gm.graph.lint()
        order = self._node_order(gm)
        self.assertLess(order[group_name], order[prep])
        self.assertLess(order[prep], order[coll])
        self.assertLess(order[coll], order[dense0])
        self.assertLess(order[dense0], order[wait])

    def test_fsdp_dense_scheduler_rejects_late_getattr_without_reordering(self):
        root = torch.nn.Module()
        root.register_buffer("parameter", torch.ones(1))
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.0.attention",
        )
        parameter = graph.get_attr("parameter")
        prep, coll, wait = self._add_bucketed_ag_with_late_input(
            graph,
            parameter,
            plan_fqn="layers.1.attention",
            direction="fwd",
        )
        dense1 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(wait,)),
            "layers.1.attention",
        )
        graph.output((dense0, dense1))
        gm = torch.fx.GraphModule(root, graph)
        original_nodes = list(gm.graph.nodes)

        with self.assertRaisesRegex(ValueError, "launch inputs are after dense region"):
            schedule_fsdp_comms_to_dense_regions_pass(
                gm,
                moe_layer_ids=frozenset(),
                n_layers=2,
                strict=True,
            )

        self.assertEqual(list(gm.graph.nodes), original_nodes)
        order = self._node_order(gm)
        self.assertLess(order[dense0], order[parameter])
        self.assertLess(order[parameter], order[prep])
        self.assertLess(order[prep], order[coll])
        self.assertLess(order[coll], order[wait])

    def test_fsdp_dense_scheduler_does_not_move_independent_rs_producer(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        dense2 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.2.attention",
            backward=True,
        )
        dense1_early = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense2,)),
            "layers.1.attention",
            backward=True,
        )
        grad = graph.call_function(torch.ops.aten.neg.default, args=(x,))
        grad.meta["autograd_backward"] = True
        dense1_late = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense1_early,)),
            "layers.1.attention",
            backward=True,
        )
        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense1_late,)),
            "layers.0.attention",
            backward=True,
        )
        rs2 = self._tag_fsdp_bucket(
            graph.call_function(
                c10d.reduce_scatter_tensor.default,
                args=(grad, "sum", 1, "pg"),
            ),
            ["layers.2"],
            "bwd",
        )
        rs2_wait = graph.call_function(c10d.wait_tensor.default, args=(rs2,))
        graph.output((dense0, rs2_wait))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm,
            moe_layer_ids=frozenset(),
            n_layers=3,
            strict=True,
        )

        gm.graph.lint()
        order = self._node_order(gm)
        self.assertLess(order[dense1_early], order[grad])
        self.assertLess(order[grad], order[rs2])
        self.assertLess(order[rs2], order[dense1_late])

    def test_fsdp_dense_scheduler_accepts_expected_bucket_counts(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        fwd_ag0 = graph.call_function(
            c10d.all_gather_into_tensor.default, args=(x, 1, "pg")
        )
        fwd_ag0_wait = graph.call_function(c10d.wait_tensor.default, args=(fwd_ag0,))
        fwd0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(fwd_ag0_wait,)),
            "layers.0.attention",
        )
        bwd_ag0 = graph.call_function(
            c10d.all_gather_into_tensor.default, args=(x, 1, "pg")
        )
        bwd_ag0_wait = graph.call_function(c10d.wait_tensor.default, args=(bwd_ag0,))
        bwd0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(bwd_ag0_wait,)),
            "layers.0.attention",
            backward=True,
        )
        rs0 = self._tag_fsdp_schedule_node(
            graph.call_function(
                c10d.reduce_scatter_tensor.default, args=(bwd0, "sum", 0, "pg")
            ),
            "layers.0",
            backward=True,
        )
        rs0_wait = graph.call_function(c10d.wait_tensor.default, args=(rs0,))
        graph.output((fwd0, rs0_wait))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        optional_bwd_ag_gm = deepcopy(gm)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm,
            moe_layer_ids=frozenset(),
            n_layers=1,
            transformer_bucket_counts_by_layer={0: 1},
            strict=True,
        )

        gm.graph.lint()
        schedule_fsdp_comms_to_dense_regions_pass(
            optional_bwd_ag_gm,
            moe_layer_ids=frozenset(),
            n_layers=1,
            transformer_bucket_counts_by_layer={0: 1},
            require_backward_all_gathers=False,
            strict=True,
        )
        optional_bwd_ag_gm.graph.lint()

    def test_fsdp_dense_scheduler_accepts_stage_local_no_reshard_counts(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        fwd_ag = self._tag_fsdp_bucket(
            graph.call_function(c10d.all_gather_into_tensor.default, args=(x, 1, "pg")),
            ["layers.1"],
            "fwd",
        )
        fwd_wait = graph.call_function(c10d.wait_tensor.default, args=(fwd_ag,))
        fwd = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(fwd_wait,)),
            "layers.1.attention",
        )
        bwd = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(fwd,)),
            "layers.1.attention",
            backward=True,
        )
        rs = self._tag_fsdp_bucket(
            graph.call_function(
                c10d.reduce_scatter_tensor.default, args=(bwd, "sum", 0, "pg")
            ),
            ["layers.1"],
            "bwd",
        )
        rs_wait = graph.call_function(c10d.wait_tensor.default, args=(rs,))
        graph.output(rs_wait)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm,
            moe_layer_ids=frozenset(),
            n_layers=3,
            transformer_bucket_counts_by_layer={0: 1, 1: 1, 2: 1},
            local_layer_ids=frozenset({1}),
            require_backward_all_gathers=False,
            strict=True,
        )

        gm.graph.lint()

    def test_fsdp_dense_scheduler_validates_transformer_bucket_counts(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        ag0 = graph.call_function(
            c10d.all_gather_into_tensor.default, args=(x, 1, "pg")
        )
        ag0_wait = graph.call_function(c10d.wait_tensor.default, args=(ag0,))
        fwd0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(ag0_wait,)),
            "layers.0.attention",
        )
        bwd0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(fwd0,)),
            "layers.0.attention",
            backward=True,
        )
        graph.output(bwd0)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        with self.assertRaisesRegex(ValueError, "layer 0"):
            schedule_fsdp_comms_to_dense_regions_pass(
                gm,
                moe_layer_ids=frozenset(),
                n_layers=1,
                transformer_bucket_counts_by_layer={0: 1},
            )

    def test_fsdp_dense_scheduler_ignores_non_transformer_buckets(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        ag = graph.call_function(c10d.all_gather_into_tensor.default, args=(x, 1, "pg"))
        ag_wait = self._tag_fsdp_schedule_node(
            graph.call_function(c10d.wait_tensor.default, args=(ag,)),
            "norm",
        )
        rs = self._tag_fsdp_schedule_node(
            graph.call_function(
                c10d.reduce_scatter_tensor.default, args=(ag_wait, "sum", 0, "pg")
            ),
            "lm_head",
            backward=True,
        )
        rs_wait = graph.call_function(c10d.wait_tensor.default, args=(rs,))
        graph.output(rs_wait)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        with self.assertRaisesRegex(ValueError, "layer 0"):
            schedule_fsdp_comms_to_dense_regions_pass(
                gm,
                moe_layer_ids=frozenset(),
                n_layers=1,
                transformer_bucket_counts_by_layer={0: 1},
            )

    def test_fsdp_dense_scheduler_skips_transformer_edge_buckets(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.0.attention",
        )
        fwd_ag0 = graph.call_function(
            c10d.all_gather_into_tensor.default, args=(x, 1, "pg")
        )
        fwd_ag0_wait = graph.call_function(c10d.wait_tensor.default, args=(fwd_ag0,))
        fwd_ag0_use = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(fwd_ag0_wait,)),
            "layers.0.attention",
        )
        bwd_dense2 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(fwd_ag0_use,)),
            "layers.2.attention",
            backward=True,
        )
        bwd_ag2 = graph.call_function(
            c10d.all_gather_into_tensor.default, args=(x, 1, "pg")
        )
        bwd_ag2_wait = graph.call_function(c10d.wait_tensor.default, args=(bwd_ag2,))
        bwd_ag2_use = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(bwd_ag2_wait,)),
            "layers.2.attention",
            backward=True,
        )
        rs0 = self._tag_fsdp_schedule_node(
            graph.call_function(
                c10d.reduce_scatter_tensor.default, args=(bwd_ag2_use, "sum", 0, "pg")
            ),
            "layers.0",
            backward=True,
        )
        rs0_wait = graph.call_function(c10d.wait_tensor.default, args=(rs0,))
        graph.output((dense0, bwd_dense2, rs0_wait))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm,
            moe_layer_ids=frozenset(),
            n_layers=3,
            strict=True,
        )

        order = self._node_order(gm)
        self.assertLess(order[dense0], order[fwd_ag0])
        self.assertLess(order[bwd_dense2], order[bwd_ag2])
        self.assertLess(order[bwd_ag2_use], order[rs0])

    def test_fsdp_dense_scheduler_places_top_level_edge_buckets(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.0.attention",
        )
        top_ag = self._tag_fsdp_bucket(
            graph.call_function(c10d.all_gather_into_tensor.default, args=(x, 1, "pg")),
            ["norm", "lm_head"],
            "fwd",
        )
        top_ag_wait = graph.call_function(c10d.wait_tensor.default, args=(top_ag,))
        top_fwd = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(top_ag_wait,)),
            "norm",
        )
        top_bwd = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(top_fwd,)),
            "lm_head",
            backward=True,
        )
        bwd_ag0 = self._tag_fsdp_bucket(
            graph.call_function(c10d.all_gather_into_tensor.default, args=(x, 1, "pg")),
            ["layers.0.attention"],
            "bwd",
        )
        bwd_ag0_wait = graph.call_function(c10d.wait_tensor.default, args=(bwd_ag0,))
        bwd0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(bwd_ag0_wait,)),
            "layers.0.attention",
            backward=True,
        )
        top_rs = self._tag_fsdp_bucket(
            graph.call_function(
                c10d.reduce_scatter_tensor.default, args=(top_bwd, "sum", 0, "pg")
            ),
            ["norm", "lm_head"],
            "bwd",
        )
        top_rs_wait = graph.call_function(c10d.wait_tensor.default, args=(top_rs,))
        graph.output((dense0, bwd0, top_rs_wait))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm,
            moe_layer_ids=frozenset(),
            n_layers=1,
            strict=True,
        )

        order = self._node_order(gm)
        self.assertLess(order[top_ag], order[dense0])
        self.assertLess(order[dense0], order[top_ag_wait])
        self.assertLess(order[bwd_ag0], order[top_bwd])
        self.assertLess(order[top_bwd], order[bwd_ag0_wait])
        self.assertLess(order[top_bwd], order[top_rs])
        self.assertLess(order[top_rs], order[bwd0])
        self.assertGreater(order[top_rs_wait], order[bwd0])

    def test_fsdp_dense_scheduler_places_backward_ag_before_rs(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        ag0 = self._tag_fsdp_schedule_node(
            graph.call_function(c10d.all_gather_into_tensor.default, args=(x, 1, "pg")),
            "",
        )
        dense2 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.2.attention",
            backward=True,
        )
        dense1 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense2,)),
            "layers.1.attention",
            backward=True,
        )
        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense1,)),
            "layers.0.attention",
            backward=True,
        )
        ag0_wait = graph.call_function(c10d.wait_tensor.default, args=(ag0,))
        ag0_wait_user = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(ag0_wait,)),
            "layers.0",
            backward=True,
        )
        rs2 = self._tag_fsdp_schedule_node(
            graph.call_function(
                c10d.reduce_scatter_tensor.default, args=(dense2, "sum", 0, "pg")
            ),
            "layers.2",
            backward=True,
        )
        rs2_wait = graph.call_function(c10d.wait_tensor.default, args=(rs2,))
        graph.output((dense0, ag0_wait_user, rs2_wait))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm,
            moe_layer_ids=frozenset(),
            n_layers=3,
            strict=True,
        )

        order = self._node_order(gm)
        self.assertLess(order[dense2], order[ag0])
        self.assertLess(order[ag0], order[rs2])
        self.assertLess(order[rs2], order[dense1])
        self.assertLess(order[dense1], order[dense0])
        self.assertGreater(order[ag0_wait], order[dense0])
        self.assertGreater(order[rs2_wait], order[dense0])

    def test_fsdp_dense_scheduler_keeps_forward_ag_with_backward_descendants(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.0.attention",
        )
        ag1 = self._tag_fsdp_schedule_node(
            graph.call_function(c10d.all_gather_into_tensor.default, args=(x, 1, "pg")),
            "",
        )
        ag1_wait = graph.call_function(c10d.wait_tensor.default, args=(ag1,))
        fwd_use = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(ag1_wait,)),
            "layers.1.attention",
        )
        bwd2 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(fwd_use,)),
            "layers.2.attention",
            backward=True,
        )
        bwd1 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(bwd2,)),
            "layers.1.attention",
            backward=True,
        )
        graph.output((dense0, fwd_use, bwd1))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm,
            moe_layer_ids=frozenset(),
            n_layers=3,
            strict=True,
        )

        order = self._node_order(gm)
        self.assertLess(order[ag1], order[dense0])
        self.assertLess(order[dense0], order[ag1_wait])

    def test_fsdp_dense_scheduler_uses_compute_not_fsdp_unpack_as_anchor(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        ag0 = graph.call_function(
            c10d.all_gather_into_tensor.default, args=(x, 1, "pg")
        )
        wait0 = self._tag_fsdp_schedule_node(
            graph.call_function(c10d.wait_tensor.default, args=(ag0,)),
            "layers.0.attention_norm",
        )
        unpack0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.view.default, args=(wait0, [1])),
            "layers.0.attention_norm",
        )
        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(unpack0,)),
            "layers.0.attention",
        )
        ag1 = graph.call_function(
            c10d.all_gather_into_tensor.default, args=(x, 1, "pg")
        )
        ag1_wait = graph.call_function(c10d.wait_tensor.default, args=(ag1,))
        dense1 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(ag1_wait,)),
            "layers.1.attention",
        )
        graph.output((dense0, dense1))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm,
            moe_layer_ids=frozenset(),
            n_layers=2,
            strict=True,
        )

        order = self._node_order(gm)
        self.assertLess(order[wait0], order[ag1])
        self.assertLess(order[unpack0], order[ag1])
        self.assertLess(order[ag1], order[dense0])
        self.assertLess(order[dense0], order[ag1_wait])

    def test_fsdp_dense_scheduler_moves_padded_all_gather_chain(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")

        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.0.attention",
        )
        padded = graph.call_function(
            torch.ops.aten.constant_pad_nd.default, args=(x, [0, 1], 0.0)
        )
        bucket = graph.call_function(
            torch.ops.bucketing._pre_bucket_all_gather.default,
            args=([padded], 2, torch.float32, [0], 0),
        )
        shard = graph.call_function(torch.ops.aten.slice.Tensor, args=(bucket, 0, 0, 1))
        ag1 = self._tag_fsdp_bucket(
            graph.call_function(
                torch.ops._c10d_functional.all_gather_into_tensor_out.default,
                args=(shard, 2, "pg"),
                kwargs={"out": bucket},
            ),
            ["layers.1"],
            "fwd",
        )
        ag1_wait = graph.call_function(
            torch.ops._c10d_functional.wait_tensor.default, args=(ag1,)
        )
        for node in (padded, bucket, shard, ag1_wait):
            self._tag_fsdp_chain(node, ["layers.1"], "fwd")
        dense1 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(ag1_wait,)),
            "layers.1.attention",
        )
        graph.output((dense0, dense1))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm, moe_layer_ids=frozenset(), n_layers=2, strict=True
        )

        order = self._node_order(gm)
        self.assertLess(order[padded], order[bucket])
        self.assertLess(order[bucket], order[shard])
        self.assertLess(order[shard], order[ag1])
        self.assertLess(order[ag1], order[dense0])
        self.assertLess(order[dense0], order[ag1_wait])

    def test_fsdp_dense_scheduler_treats_recomputed_ag_use_as_backward(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        ag0 = self._tag_fsdp_schedule_node(
            graph.call_function(c10d.all_gather_into_tensor.default, args=(x, 1, "pg")),
            "",
        )
        dense2 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.2.attention",
            backward=True,
        )
        dense1 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense2,)),
            "layers.1.attention",
            backward=True,
        )
        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense1,)),
            "layers.0.attention",
            backward=True,
        )
        ag0_wait = graph.call_function(c10d.wait_tensor.default, args=(ag0,))
        view = graph.call_function(torch.ops.aten.view.default, args=(ag0_wait, [1]))
        recomputed = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(view,)),
            "layers.0.ffn_norm",
        )
        recomputed.name = "relu_recomputed"
        graph.output((dense0, recomputed))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm, moe_layer_ids=frozenset(), n_layers=3, strict=True
        )

        order = self._node_order(gm)
        self.assertLess(order[dense2], order[ag0])
        self.assertLess(order[ag0], order[dense1])
        self.assertLess(order[ag0_wait], order[recomputed])

    def test_fsdp_dense_scheduler_keeps_rs_after_gradient_producer(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        dense2 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.2.attention",
            backward=True,
        )
        dense1a = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense2,)),
            "layers.1.attention",
            backward=True,
        )
        grad_cast = graph.call_function(
            torch.ops.aten._to_copy.default,
            args=(dense1a,),
            kwargs={"dtype": torch.float32},
        )
        dense1b = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(grad_cast,)),
            "layers.1.attention",
            backward=True,
        )
        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense1b,)),
            "layers.0.attention",
            backward=True,
        )
        rs2 = self._tag_fsdp_schedule_node(
            graph.call_function(
                c10d.reduce_scatter_tensor.default, args=(grad_cast, "sum", 0, "pg")
            ),
            "layers.2",
            backward=True,
        )
        rs2_wait = graph.call_function(c10d.wait_tensor.default, args=(rs2,))
        graph.output((dense0, rs2_wait))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm, moe_layer_ids=frozenset(), n_layers=3, strict=True
        )

        order = self._node_order(gm)
        self.assertLess(order[dense1a], order[grad_cast])
        self.assertLess(order[grad_cast], order[rs2])
        self.assertLess(order[rs2], order[dense1b])

    def test_fsdp_dense_scheduler_sinks_output_only_rs_wait(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        dense2 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.2.attention",
            backward=True,
        )
        rs2 = self._tag_fsdp_schedule_node(
            graph.call_function(
                c10d.reduce_scatter_tensor.default, args=(dense2, "sum", 0, "pg")
            ),
            "layers.2",
            backward=True,
        )
        rs2_wait = graph.call_function(c10d.wait_tensor.default, args=(rs2,))
        dense1 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense2,)),
            "layers.1.attention",
            backward=True,
        )
        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense1,)),
            "layers.0.attention",
            backward=True,
        )
        graph.output((dense0, rs2_wait))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm, moe_layer_ids=frozenset(), n_layers=3, strict=True
        )

        order = self._node_order(gm)
        self.assertLess(order[rs2], order[dense1])
        self.assertGreater(order[rs2_wait], order[dense0])

    def test_fsdp_dense_scheduler_sinks_unscheduled_rs_wait(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        rs0 = self._tag_fsdp_bucket(
            graph.call_function(
                c10d.reduce_scatter_tensor.default, args=(x, "sum", 0, "pg")
            ),
            ["layers.0"],
            "bwd",
        )
        rs0_wait = graph.call_function(c10d.wait_tensor.default, args=(rs0,))
        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.0.attention",
            backward=True,
        )
        graph.output((dense0, rs0_wait))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm, moe_layer_ids=frozenset(), n_layers=1, strict=True
        )

        order = self._node_order(gm)
        self.assertLess(order[rs0], order[dense0])
        self.assertGreater(order[rs0_wait], order[dense0])

    def test_fsdp_dense_scheduler_sinks_rs_wait_output_unpack(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        dense2 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.2.attention",
            backward=True,
        )
        rs2 = self._tag_fsdp_schedule_node(
            graph.call_function(
                c10d.reduce_scatter_tensor.default, args=(dense2, "sum", 0, "pg")
            ),
            "layers.2",
            backward=True,
        )
        rs2_wait = graph.call_function(c10d.wait_tensor.default, args=(rs2,))
        alias = graph.call_function(torch.ops.aten.alias.default, args=(rs2_wait,))
        split = graph.call_function(
            torch.ops.aten.split_with_sizes.default, args=(alias, [1])
        )
        shard = graph.call_function(operator.getitem, args=(split, 0))
        dense1 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense2,)),
            "layers.1.attention",
            backward=True,
        )
        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense1,)),
            "layers.0.attention",
            backward=True,
        )
        graph.output((dense0, shard))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm, moe_layer_ids=frozenset(), n_layers=3, strict=True
        )

        order = self._node_order(gm)
        self.assertLess(order[rs2], order[dense1])
        self.assertGreater(order[rs2_wait], order[dense0])
        self.assertGreater(order[alias], order[dense0])
        self.assertGreater(order[split], order[dense0])
        self.assertGreater(order[shard], order[dense0])

    def test_fsdp_dense_scheduler_sinks_rs_wait_grad_accum_chain(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        dense2 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.2.attention",
            backward=True,
        )
        rs_a = self._tag_fsdp_bucket(
            graph.call_function(
                c10d.reduce_scatter_tensor.default, args=(dense2, "sum", 0, "pg")
            ),
            ["loss"],
            "bwd",
        )
        wait_a = graph.call_function(c10d.wait_tensor.default, args=(rs_a,))
        detached = graph.call_function(torch.ops.aten.detach.default, args=(wait_a,))
        rs_b = self._tag_fsdp_bucket(
            graph.call_function(
                c10d.reduce_scatter_tensor.default, args=(dense2, "sum", 0, "pg")
            ),
            ["loss"],
            "bwd",
        )
        wait_b = graph.call_function(c10d.wait_tensor.default, args=(rs_b,))
        accum = graph.call_function(torch.ops.aten.add_.Tensor, args=(detached, wait_b))
        grad_out = graph.call_function(torch.ops.aten.detach.default, args=(accum,))
        dense1 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense2,)),
            "layers.1.attention",
            backward=True,
        )
        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense1,)),
            "layers.0.attention",
            backward=True,
        )
        graph.output((dense0, grad_out))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm, moe_layer_ids=frozenset(), n_layers=3, strict=True
        )

        order = self._node_order(gm)
        self.assertLess(order[dense0], order[wait_a])
        self.assertLess(order[dense0], order[wait_b])
        self.assertLess(order[wait_a], order[detached])
        self.assertLess(order[detached], order[accum])
        self.assertLess(order[wait_b], order[accum])
        self.assertLess(order[accum], order[grad_out])

    def test_fsdp_dense_scheduler_places_rs_after_moe_backward(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        dense2 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.2.attention",
            backward=True,
        )
        layer1_boundary = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(dense2,)),
            "layers.1",
            backward=True,
        )
        ffn_norm = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(layer1_boundary,)),
            "layers.1.ffn_norm",
            backward=True,
        )
        moe_dispatch = self._tag_fsdp_schedule_node(
            graph.call_function(
                c10d.all_to_all_single.default,
                args=(ffn_norm, [], [], "ep_pg"),
            ),
            "layers.1.moe.routed_experts",
            backward=True,
        )
        moe_dispatch_wait = self._tag_fsdp_schedule_node(
            graph.call_function(c10d.wait_tensor.default, args=(moe_dispatch,)),
            "layers.1.moe.routed_experts",
            backward=True,
        )
        dense1_attention = self._tag_fsdp_schedule_node(
            graph.call_function(
                torch.ops.aten.relu.default,
                args=(moe_dispatch_wait,),
            ),
            "layers.1.attention",
            backward=True,
        )
        rs2 = self._tag_fsdp_schedule_node(
            graph.call_function(
                c10d.reduce_scatter_tensor.default,
                args=(dense2, "sum", 0, "pg"),
            ),
            "layers.2",
            backward=True,
        )
        rs2_wait = graph.call_function(c10d.wait_tensor.default, args=(rs2,))
        graph.output((dense1_attention, rs2_wait))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm, moe_layer_ids=frozenset({1}), n_layers=3, strict=True
        )

        order = self._node_order(gm)
        self.assertLess(order[layer1_boundary], order[moe_dispatch])
        self.assertLess(order[moe_dispatch], order[moe_dispatch_wait])
        self.assertLess(order[moe_dispatch_wait], order[rs2])
        self.assertLess(order[rs2], order[dense1_attention])

    def test_fsdp_dense_scheduler_excludes_moe_nodes_from_dense_region(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        moe_node = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(x,)),
            "layers.0.moe.router",
        )
        dense0 = self._tag_fsdp_schedule_node(
            graph.call_function(torch.ops.aten.relu.default, args=(moe_node,)),
            "layers.0.attention",
        )
        ag1 = self._tag_fsdp_schedule_node(
            graph.call_function(c10d.all_gather_into_tensor.default, args=(x, 1, "pg")),
            "layers.1",
        )
        ag1_wait = graph.call_function(c10d.wait_tensor.default, args=(ag1,))
        graph.output((dense0, ag1_wait))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        schedule_fsdp_comms_to_dense_regions_pass(
            gm,
            moe_layer_ids=frozenset({0}),
            n_layers=2,
            strict=True,
        )

        order = self._node_order(gm)
        self.assertLess(order[moe_node], order[ag1])
        self.assertLess(order[ag1], order[dense0])


class TestOverlapPgIsolationPass(FSDPTest):
    def _setup(self):
        self.parallelism_context = ParallelismContext(
            dp_shard=-1,
            dp_replicate=1,
            cp=1,
            tp=1,
            pp=1,
            ep=1,
            world_size=self.world_size,
            enable_sequence_parallel=False,
        )

    def _get_fsdp_pg_name(self):
        from torchtitan.experiments.graph_trainer.common_utils import (
            get_simple_fsdp_mesh,
        )

        fsdp_mesh = get_simple_fsdp_mesh(self.parallelism_context)
        return fsdp_mesh.get_group().group_name

    def _count_all_ag_nodes(self, gm):
        return sum(1 for node in gm.graph.nodes if is_all_gather(node))

    def _count_ep_a2a_nodes_with_pg(self, gm, pg_name):
        return sum(
            1
            for node in gm.graph.nodes
            if node.op == "call_function"
            and "all_to_all_single" in str(node.target)
            and node.args[3] == pg_name
        )

    def test_overlap_preserves_distinct_ep_pg_with_same_fsdp_ranks(self):
        import torch.distributed as dist

        from torchtitan.experiments.graph_trainer.ep_process_group_pass import (
            _EXTRA_EP_PG_REGISTRY,
        )
        from torchtitan.experiments.graph_trainer.fsdp_passes import (
            _EXTRA_FSDP_PG_REGISTRY,
        )

        self._setup()
        fsdp_pg_name = self._get_fsdp_pg_name()
        ep_pg = dist.new_group(
            ranks=list(range(self.world_size)),
            use_local_synchronization=True,
        )
        ep_pg_name = ep_pg.group_name

        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional
        ag = graph.call_function(
            c10d.all_gather_into_tensor.default, args=(x, 1, fsdp_pg_name)
        )
        wait = graph.call_function(c10d.wait_tensor.default, args=(ag,))
        rs = graph.call_function(
            c10d.reduce_scatter_tensor.default, args=(x, "sum", 1, fsdp_pg_name)
        )
        rs_wait = graph.call_function(c10d.wait_tensor.default, args=(rs,))
        a2a = graph.call_function(
            c10d.all_to_all_single.default, args=(x, [], [], ep_pg_name)
        )
        a2a.meta["custom"] = {
            _MODULE_FQN: "layers.0.moe",
            "EP": "combine",
            _EP_TOKEN_EXCHANGE: "combine",
        }
        graph.output((wait, rs_wait, a2a))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        _EXTRA_FSDP_PG_REGISTRY.pop(fsdp_pg_name, None)
        _EXTRA_EP_PG_REGISTRY.pop(ep_pg_name, None)
        reassign_collective_pgs_pass(gm, ())
        isolate_ep_process_group_pass(gm, ())

        self.assertNotIn(ep_pg_name, _EXTRA_EP_PG_REGISTRY)
        self.assertEqual(self._count_ep_a2a_nodes_with_pg(gm, ep_pg_name), 1)

    def test_ep_pg_pass_rewrites_all_ep_a2a_on_tp_pg_to_separate_pg(self):
        from torchtitan.experiments.graph_trainer.ep_process_group_pass import (
            _EXTRA_EP_PG_REGISTRY,
        )

        self._setup()
        tp_pg_name = self._get_fsdp_pg_name()
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional
        ag = graph.call_function(
            c10d.all_gather_into_tensor.default, args=(x, 1, tp_pg_name)
        )
        a2a = graph.call_function(
            c10d.all_to_all_single.default, args=(x, [], [], tp_pg_name)
        )
        a2a.meta["custom"] = {
            _MODULE_FQN: "layers.0.moe",
            "EP": "dispatch",
        }
        generic_ep_a2a = graph.call_function(
            c10d.all_to_all_single.default, args=(x, [], [], tp_pg_name)
        )
        generic_ep_a2a.meta["custom"] = {
            _MODULE_FQN: "layers.0.moe",
            "EP": "combine",
        }
        graph.output((ag, a2a, generic_ep_a2a))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        _EXTRA_EP_PG_REGISTRY.pop(tp_pg_name, None)
        isolate_ep_process_group_pass(gm, ())

        ep_extra_pg = _EXTRA_EP_PG_REGISTRY[tp_pg_name]
        self.assertEqual(self._count_ep_a2a_nodes_with_pg(gm, ep_extra_pg), 2)

    def test_overlap_is_noop_when_no_fsdp_ag(self):
        """If the graph has no FSDP all-gathers, the pass is a no-op."""
        self._setup()
        # Plain (non-FSDP) module: a graph without FSDP all-gathers.
        gm = torch.fx.symbolic_trace(torch.nn.Linear(4, 4))
        ag_before = self._count_all_ag_nodes(gm)
        reassign_collective_pgs_pass(gm, ())
        ag_after = self._count_all_ag_nodes(gm)
        self.assertEqual(ag_before, 0)
        self.assertEqual(ag_after, 0)


class TestApplySACPass(TestCase):
    """Unit tests for the tag_sac_policy joint graph pass."""

    def _build_gm(self, op_targets):
        """Build a GraphModule with a chain of call_function nodes.

        Each op in op_targets becomes a call_function node. The graph
        structure is: placeholder(x), placeholder(y) -> op1 -> op2 -> ... -> output.
        """
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        y = graph.placeholder("y")
        last = x
        for i, target in enumerate(op_targets):
            if target is operator.getitem:
                last = graph.call_function(target, args=(last, 0))
            else:
                last = graph.call_function(target, args=(last, y))
                # If the next op is getitem, wrap in a tuple so getitem has
                # a proper tuple/list input.
                if i + 1 < len(op_targets) and op_targets[i + 1] is operator.getitem:

                    def _make_tuple(x):
                        return (x, x)

                    last = graph.call_function(_make_tuple, args=(last,))
        graph.output(last)
        return torch.fx.GraphModule(torch.nn.Module(), graph)

    def _get_call_function_nodes(self, gm):
        """Return all call_function nodes from the graph."""
        return [n for n in gm.graph.nodes if n.op == "call_function"]

    def test_none_policy_disables_activation_rematerialization(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        fwd = graph.call_function(torch.ops.aten.add.Tensor, args=(x, x))
        bwd = graph.call_function(torch.ops.aten.mul.Tensor, args=(fwd, 2))
        bwd.meta["autograd_backward"] = True
        graph.output(bwd)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        config = SimpleNamespace(
            compile=GraphTrainerCompileConfig(memory_policy="none")
        )
        tag_with_memory_policy_pass(gm, config=config)
        selective_activation_remat_pass(gm)

        self.assertEqual(fwd.meta["recompute"], CheckpointPolicy.MUST_SAVE)
        self.assertFalse(
            any(
                node.name.endswith("_recomputed")
                for node in gm.graph.nodes
                if node.op == "call_function"
            )
        )

    def test_non_save_ops_marked_recompute(self):
        """Ops not in the save list should be marked PREFER_RECOMPUTE."""
        gm = self._build_gm(
            [
                torch.ops.aten.add.Tensor,
                torch.ops.aten.relu.default,
            ]
        )
        tag_sac_policy(gm)
        for node in self._get_call_function_nodes(gm):
            self.assertEqual(node.meta["recompute"], CheckpointPolicy.PREFER_RECOMPUTE)

    def test_save_ops_marked_must_save(self):
        """Non-mm ops in the save list should be marked MUST_SAVE."""
        custom_save = {torch.ops.aten.add.Tensor}
        gm = self._build_gm([torch.ops.aten.add.Tensor])
        tag_sac_policy(gm, policy_fn=_make_default_memory_policy(custom_save))
        nodes = self._get_call_function_nodes(gm)
        self.assertEqual(len(nodes), 1)
        self.assertEqual(nodes[0].meta["recompute"], CheckpointPolicy.MUST_SAVE)

    def test_sym_size_ops_always_saved(self):
        """Sym-int nodes are forced MUST_SAVE regardless of policy: recomputing a
        shape read would pin the parent tensor alive just to reread its size."""
        # sym_size produces a SymInt only for symbolic dims, so tag the node's
        # meta["val"] with a real SymInt — that is what is_sym_node keys off.
        shape_env = ShapeEnv()
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        relu = graph.call_function(torch.ops.aten.relu.default, args=(x,))
        sym = graph.call_function(torch.ops.aten.sym_size.int, args=(relu, 0))
        sym.meta["val"] = shape_env.create_unbacked_symint()
        graph.output(relu)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        # Even a recompute-everything policy must save sym_size.
        tag_sac_policy(gm, policy_fn=_make_full_memory_policy())

        tags = {
            n.target: n.meta["recompute"]
            for n in gm.graph.nodes
            if n.op == "call_function"
        }
        self.assertEqual(tags[torch.ops.aten.sym_size.int], CheckpointPolicy.MUST_SAVE)
        self.assertEqual(
            tags[torch.ops.aten.relu.default], CheckpointPolicy.MUST_RECOMPUTE
        )

    def test_effectful_ops_are_saved_and_not_rematerialized(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        ordered_identity = (
            torch.ops.torchtitan_graph_trainer_test.ordered_identity.default
        )
        effect = graph.call_function(ordered_identity, args=(x,))
        backward = graph.call_function(torch.ops.aten.mul.Tensor, args=(effect, 2))
        backward.meta["autograd_backward"] = True
        graph.output(backward)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        tag_sac_policy(gm, policy_fn=_make_full_memory_policy())
        self.assertEqual(effect.meta["recompute"], CheckpointPolicy.MUST_SAVE)

        selective_activation_remat_pass(gm)
        effect_nodes = [
            node
            for node in gm.graph.nodes
            if node.op == "call_function" and node.target is ordered_identity
        ]
        self.assertEqual(effect_nodes, [effect])

    def test_getitem_propagates_parent_tags(self):
        """operator.getitem nodes should inherit the parent's recompute tag."""
        gm = self._build_gm(
            [
                torch.ops.aten.add.Tensor,
                operator.getitem,
                torch.ops.aten.relu.default,
            ]
        )
        nodes = self._get_call_function_nodes(gm)
        # nodes: [add, make_tuple, getitem, relu]
        # make_tuple is the tuple-returning parent of getitem
        self.assertEqual(nodes[0].target, torch.ops.aten.add.Tensor)
        self.assertEqual(nodes[2].target, operator.getitem)

        tag_sac_policy(gm)

        tuple_node = nodes[1]
        getitem_node = nodes[2]
        self.assertEqual(getitem_node.meta["recompute"], tuple_node.meta["recompute"])

    def test_wait_tensor_propagates_parent_tags(self):
        """wait_tensor nodes should inherit the parent's recompute tag."""
        custom_save = {torch.ops._c10d_functional.reduce_scatter_tensor.default}
        gm = self._build_gm(
            [
                torch.ops._c10d_functional.reduce_scatter_tensor.default,
                torch.ops._c10d_functional.wait_tensor.default,
            ]
        )
        nodes = self._get_call_function_nodes(gm)
        nodes[0].meta["custom"] = {_MODULE_FQN: "layers.3.attention"}

        tag_sac_policy(gm, policy_fn=_make_default_memory_policy(custom_save))

        rs_node = nodes[0]
        wait_node = nodes[1]
        self.assertEqual(rs_node.meta["recompute"], CheckpointPolicy.MUST_SAVE)
        self.assertEqual(wait_node.meta["recompute"], CheckpointPolicy.MUST_SAVE)

    def test_default_policy_saves_fsdp_unshard_when_not_resharding(self):
        """Saves the helper-selected FSDP unshard output only when needed."""
        cases = (
            ("never", 1, CheckpointPolicy.MUST_SAVE),
            # Under PP, the default FSDP policy keeps params unsharded across
            # forward/backward, so SAC must save the same unshard boundary.
            ("default", 2, CheckpointPolicy.MUST_SAVE),
            ("always", 1, CheckpointPolicy.PREFER_RECOMPUTE),
        )

        for reshard_after_forward, pp_degree, expected_wait_policy in cases:
            with self.subTest(
                reshard_after_forward=reshard_after_forward,
                pp_degree=pp_degree,
            ):
                (
                    gm,
                    all_gather,
                    wait,
                    view,
                ) = self._fsdp_unshard_test_graph()
                config = SimpleNamespace(
                    parallelism=SimpleNamespace(
                        fsdp_reshard_after_forward=reshard_after_forward,
                        pipeline_parallel_degree=pp_degree,
                    )
                )

                _default_memory_policy_pass(gm, config=config)

                self.assertEqual(
                    all_gather.meta["recompute"],
                    CheckpointPolicy.PREFER_RECOMPUTE,
                )
                self.assertEqual(wait.meta["recompute"], expected_wait_policy)
                self.assertEqual(
                    view.meta["recompute"],
                    CheckpointPolicy.PREFER_RECOMPUTE,
                )

    def _fsdp_unshard_test_graph(self):
        graph = torch.fx.Graph()
        param = graph.placeholder("param")
        x = graph.placeholder("x")
        all_gather = graph.call_function(
            torch.ops._c10d_functional.all_gather_into_tensor.default,
            args=(param, 1, "0"),
        )
        wait = graph.call_function(
            torch.ops._c10d_functional.wait_tensor.default,
            args=(all_gather,),
        )
        view = graph.call_function(torch.ops.aten.view.default, args=(wait, [4]))
        fsdp_meta = {FSDP_PARAM_FQNS_META: ("linear.weight",)}
        for node in (all_gather, wait, view):
            node.meta["custom"] = fsdp_meta
        out = graph.call_function(torch.ops.aten.add.Tensor, args=(view, x))
        graph.output(out)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        return gm, all_gather, wait, view

    def test_boundary_nodes_forced_to_must_save(self):
        """Nodes at AC region boundaries should be forced to MUST_SAVE."""
        gm = self._build_gm(
            [
                torch.ops.aten.add.Tensor,
                torch.ops.aten.relu.default,
            ]
        )
        nodes = self._get_call_function_nodes(gm)
        nodes[0].meta["custom"] = {_MODULE_FQN: "layers.0.feed_forward"}
        nodes[1].meta["custom"] = {_MODULE_FQN: "layers.1.attention"}

        tag_sac_policy(gm)

        # add is at the boundary (layer 0 -> layer 1), forced to MUST_SAVE
        self.assertEqual(nodes[0].meta["recompute"], CheckpointPolicy.MUST_SAVE)
        self.assertEqual(nodes[1].meta["recompute"], CheckpointPolicy.PREFER_RECOMPUTE)

    def test_custom_op_list_to_save(self):
        """A custom op_list_to_save should override the defaults."""
        custom_save = {torch.ops.aten.relu.default}
        gm = self._build_gm(
            [
                torch.ops.aten.add.Tensor,
                torch.ops.aten.relu.default,
            ]
        )
        tag_sac_policy(gm, policy_fn=_make_default_memory_policy(custom_save))
        policies = {
            n.target: n.meta["recompute"] for n in self._get_call_function_nodes(gm)
        }
        self.assertEqual(
            policies[torch.ops.aten.add.Tensor], CheckpointPolicy.PREFER_RECOMPUTE
        )
        self.assertEqual(
            policies[torch.ops.aten.relu.default], CheckpointPolicy.MUST_SAVE
        )

    def test_mixed_mm_and_save_ops(self):
        """Graph with both mm and other save ops are annotated correctly."""
        custom_save = {torch.ops.aten.mm.default, torch.ops.aten.max.default}
        gm = self._build_gm(
            [
                torch.ops.aten.mm.default,  # in save list -> MUST_SAVE
                torch.ops.aten.max.default,  # in save list -> MUST_SAVE
                torch.ops.aten.mm.default,  # in save list -> MUST_SAVE
                torch.ops.aten.add.Tensor,  # not in save list -> PREFER_RECOMPUTE
                torch.ops.aten.mm.default,  # in save list -> MUST_SAVE
            ]
        )
        tag_sac_policy(gm, policy_fn=_make_default_memory_policy(custom_save))
        nodes = self._get_call_function_nodes(gm)
        expected = [
            (torch.ops.aten.mm.default, CheckpointPolicy.MUST_SAVE),
            (torch.ops.aten.max.default, CheckpointPolicy.MUST_SAVE),
            (torch.ops.aten.mm.default, CheckpointPolicy.MUST_SAVE),
            (torch.ops.aten.add.Tensor, CheckpointPolicy.PREFER_RECOMPUTE),
            (torch.ops.aten.mm.default, CheckpointPolicy.MUST_SAVE),
        ]
        self.assertEqual(len(nodes), len(expected))
        for node, (target, policy) in zip(nodes, expected):
            self.assertEqual(node.target, target)
            self.assertEqual(node.meta["recompute"], policy, f"node {node.name}")

    def test_remat_uses_autograd_backward_without_phase_annotation(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        fwd = graph.call_function(torch.ops.aten.add.Tensor, args=(x, x))
        bwd = graph.call_function(torch.ops.aten.mul.Tensor, args=(fwd, 2))
        graph.output(bwd)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        fwd.meta["recompute"] = CheckpointPolicy.PREFER_RECOMPUTE
        bwd.meta["autograd_backward"] = True

        gm = selective_activation_remat_pass(gm)

        recomputed_nodes = [
            node for node in gm.graph.nodes if node.name == "add_tensor_recomputed"
        ]
        self.assertEqual(len(recomputed_nodes), 1)
        self.assertTrue(recomputed_nodes[0].meta["autograd_backward"])

    def test_remat_dup_gets_independent_custom_meta(self):
        # fx.Graph.node_copy shallow-copies node.meta, so without intervention a
        # recompute dup shares the SAME nested meta["custom"] dict as its forward
        # original -- annotating one would silently mutate the other. The pass must
        # give the dup its own copy (preserving the values).
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        fwd = graph.call_function(torch.ops.aten.add.Tensor, args=(x, x))
        # A forward consumer keeps the original alive (remat erases originals whose
        # consumers are all backward) so we can compare it against the dup.
        fwd_use = graph.call_function(torch.ops.aten.relu.default, args=(fwd,))
        bwd = graph.call_function(torch.ops.aten.mul.Tensor, args=(fwd, 2))
        graph.output((fwd_use, bwd))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        fwd.meta["recompute"] = CheckpointPolicy.PREFER_RECOMPUTE
        fwd.meta["custom"] = {_MODULE_FQN: "layers.0.attention_norm"}
        bwd.meta["autograd_backward"] = True

        gm = selective_activation_remat_pass(gm)

        fwd_node = next(n for n in gm.graph.nodes if n.name == "add_tensor")
        dup = next(n for n in gm.graph.nodes if n.name == "add_tensor_recomputed")
        # Independent dict object, same values preserved.
        self.assertIsNot(dup.meta["custom"], fwd_node.meta["custom"])
        self.assertEqual(dup.meta["custom"][_MODULE_FQN], "layers.0.attention_norm")
        # Mutating the dup's annotation must not leak into the original.
        dup.meta["custom"]["cuda_graph_partition"] = "cuda_graph_9"
        self.assertNotIn("cuda_graph_partition", fwd_node.meta["custom"])


class TestFullMemoryPolicy(TestCase):
    """Unit tests for the full recompute memory policy."""

    def _build_gm(self, op_targets, layer_fqns=None):
        """Build a GraphModule with call_function nodes and optional layer FQNs."""
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        y = graph.placeholder("y")
        last = x
        nodes = []
        for target in op_targets:
            last = graph.call_function(target, args=(last, y))
            nodes.append(last)
        graph.output(last)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        if layer_fqns:
            for node, fqn in zip(nodes, layer_fqns):
                if fqn is not None:
                    node.meta["custom"] = {_MODULE_FQN: fqn}
        return gm

    def _get_call_function_nodes(self, gm):
        return [n for n in gm.graph.nodes if n.op == "call_function"]

    def test_all_ops_marked_recompute(self):
        """All ops should be marked MUST_RECOMPUTE with full policy."""
        gm = self._build_gm(
            [
                torch.ops.aten.mm.default,
                torch.ops.aten.add.Tensor,
                torch.ops.aten.relu.default,
            ]
        )
        tag_sac_policy(gm, policy_fn=_make_full_memory_policy())
        for node in self._get_call_function_nodes(gm):
            self.assertEqual(
                node.meta["recompute"],
                CheckpointPolicy.MUST_RECOMPUTE,
                f"node {node.name} should be MUST_RECOMPUTE",
            )

    def test_save_ops_also_recomputed(self):
        """Compute-intensive ops (linear, max) are recomputed under full."""
        gm = self._build_gm(
            [
                torch.ops.aten.linear.default,
                torch.ops.aten.max.default,
            ]
        )
        tag_sac_policy(gm, policy_fn=_make_full_memory_policy())
        for node in self._get_call_function_nodes(gm):
            self.assertEqual(
                node.meta["recompute"],
                CheckpointPolicy.MUST_RECOMPUTE,
            )

    def test_rng_ops_saved(self):
        """RNG ops (nondeterministic_seeded) are saved: the remat pass cannot
        replay their random state, unlike eager AC's preserve_rng_state."""
        policy_fn = _make_full_memory_policy()
        gm = self._build_gm([torch.ops.aten.native_dropout.default])
        node = self._get_call_function_nodes(gm)[0]
        self.assertEqual(policy_fn(node), CheckpointPolicy.MUST_SAVE)

    def test_higher_order_ops_recomputed(self):
        """Higher-order ops (flex_attention) are recomputed, not saved: the
        remat pass duplicates the HOP together with its subgraph get_attrs."""
        policy_fn = _make_full_memory_policy()
        gm = self._build_gm([torch.ops.higher_order.flex_attention])
        node = self._get_call_function_nodes(gm)[0]
        self.assertEqual(policy_fn(node), CheckpointPolicy.MUST_RECOMPUTE)

    def test_selected_module_ops_saved(self):
        policy_fn = _make_full_memory_policy(
            "layers.*.moe.router.gate :: aten.mm.default | "
            "layers.*.attention.inner_attention :: higher_order.flex_attention"
        )
        cases = (
            (
                torch.ops.aten.mm.default,
                "layers.3.moe.router.gate",
                CheckpointPolicy.MUST_SAVE,
            ),
            (
                torch.ops.aten.mm.default,
                "layers.3.attention.wkv_a",
                CheckpointPolicy.MUST_RECOMPUTE,
            ),
            (
                torch.ops.aten._to_copy.default,
                "layers.3.moe.router.gate",
                CheckpointPolicy.MUST_RECOMPUTE,
            ),
            (
                torch.ops.higher_order.flex_attention,
                "layers.3.attention.inner_attention",
                CheckpointPolicy.MUST_SAVE,
            ),
        )

        for target, fqn, expected in cases:
            with self.subTest(target=target, fqn=fqn):
                gm = self._build_gm([target], layer_fqns=[fqn])
                node = self._get_call_function_nodes(gm)[0]
                self.assertEqual(policy_fn(node), expected)

    def test_selected_module_op_can_target_one_layer(self):
        policy_fn = _make_full_memory_policy(
            "layers.7.attention.wo :: torch.ops.aten.mm.default"
        )
        for layer_id, expected in (
            (7, CheckpointPolicy.MUST_SAVE),
            (8, CheckpointPolicy.MUST_RECOMPUTE),
        ):
            gm = self._build_gm(
                [torch.ops.aten.mm.default],
                layer_fqns=[f"layers.{layer_id}.attention.wo"],
            )
            node = self._get_call_function_nodes(gm)[0]
            self.assertEqual(policy_fn(node), expected)

    def test_full_policy_uses_configured_save_ops(self):
        gm = self._build_gm(
            [torch.ops.aten.mm.default],
            layer_fqns=["layers.3.moe.router.gate"],
        )
        config = SimpleNamespace(
            compile=GraphTrainerCompileConfig(
                memory_policy="full",
                full_recompute_save_ops=("layers.*.moe.router.gate :: aten.mm.default"),
            )
        )

        tag_with_memory_policy_pass(gm, config=config)

        node = self._get_call_function_nodes(gm)[0]
        self.assertEqual(node.meta["recompute"], CheckpointPolicy.MUST_SAVE)

    def test_save_ops_rejected_for_other_memory_policies(self):
        compile_config = GraphTrainerCompileConfig(
            memory_policy="default",
            full_recompute_save_ops="layers.*.moe.router.gate::aten.mm.default",
        )

        with self.assertRaisesRegex(
            ValueError, r"requires compile\.memory_policy='full'"
        ):
            validate_memory_policy_config(compile_config)

    def test_invalid_save_op_selectors_rejected(self):
        invalid_values = (
            "layers.*.moe.router.gate",
            ":: aten.mm.default",
            "layers.*.moe.router.gate ::",
            "layers.*.moe.router.gate :: aten.not_an_op.default",
            "layers.*.moe.router.gate :: aten.mm",
        )
        for value in invalid_values:
            with self.subTest(value=value), self.assertRaises(ValueError):
                _make_full_memory_policy(value)

    def test_layer_boundary_forced_to_must_save(self):
        """Nodes at layer boundaries should still be forced to MUST_SAVE."""
        gm = self._build_gm(
            [
                torch.ops.aten.add.Tensor,
                torch.ops.aten.relu.default,
            ],
            layer_fqns=[
                "layers.0.feed_forward",
                "layers.1.attention",
            ],
        )
        tag_sac_policy(gm, policy_fn=_make_full_memory_policy())
        nodes = self._get_call_function_nodes(gm)
        # add crosses from layer 0 to layer 1 — forced to MUST_SAVE
        self.assertEqual(nodes[0].meta["recompute"], CheckpointPolicy.MUST_SAVE)
        # relu has no higher-layer consumer — stays MUST_RECOMPUTE
        self.assertEqual(nodes[1].meta["recompute"], CheckpointPolicy.MUST_RECOMPUTE)

    def test_same_layer_nodes_all_recomputed(self):
        """Within a single layer, all ops should be recomputed."""
        gm = self._build_gm(
            [
                torch.ops.aten.mm.default,
                torch.ops.aten.linear.default,
                torch.ops.aten.add.Tensor,
                torch.ops.aten.relu.default,
            ],
            layer_fqns=[
                "layers.0.attention",
                "layers.0.attention",
                "layers.0.feed_forward",
                "layers.0.feed_forward",
            ],
        )
        tag_sac_policy(gm, policy_fn=_make_full_memory_policy())
        for node in self._get_call_function_nodes(gm):
            self.assertEqual(
                node.meta["recompute"],
                CheckpointPolicy.MUST_RECOMPUTE,
                f"node {node.name} in single layer should be MUST_RECOMPUTE",
            )


class TestMinCutMemoryPolicy(TestCase):
    @staticmethod
    def _config():
        return SimpleNamespace(
            compile=GraphTrainerCompileConfig(memory_policy="min_cut")
        )

    @staticmethod
    def _fake_prop(gm, *inputs):
        with torch._subclasses.FakeTensorMode() as fake_mode:
            fake_inputs = [
                torch.empty(shape, device="cuda", dtype=dtype)
                for shape, dtype in inputs
            ]
            FakeTensorProp(gm, mode=fake_mode).propagate_dont_convert_inputs(
                *fake_inputs
            )

    @staticmethod
    def _recomputed_nodes(gm):
        return [node for node in gm.graph.nodes if node.name.endswith("_recomputed")]

    @staticmethod
    def _log_softmax_decomposition_table():
        return get_decompositions([torch.ops.aten._log_softmax.default])

    def test_view_cut_saves_its_base(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        weight = graph.placeholder("weight")
        grad = graph.placeholder("grad")
        mm = graph.call_function(torch.ops.aten.mm.default, args=(x, weight))
        view = graph.call_function(torch.ops.aten.view.default, args=(mm, [4, 4]))
        bwd = graph.call_function(torch.ops.aten.mm.default, args=(view, grad))
        bwd.meta["autograd_backward"] = True
        graph.output(bwd)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        tag_min_cut_saved_values(gm, _backward_side_nodes(gm), {view})

        self.assertEqual(mm.meta["recompute"], CheckpointPolicy.MUST_SAVE)
        self.assertEqual(view.meta["recompute"], CheckpointPolicy.MUST_RECOMPUTE)

    def test_applies_to_whole_graph(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        a = graph.call_function(torch.ops.aten.sin.default, args=(x,))
        b = graph.call_function(torch.ops.aten.cos.default, args=(a,))
        loss = graph.call_function(torch.ops.aten.sum.default, args=(b,))
        bwd = graph.call_function(torch.ops.aten.neg.default, args=(b,))
        bwd.meta["autograd_backward"] = True
        graph.output((loss, bwd))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        self._fake_prop(gm, ((64, 64), torch.float32))

        tag_with_memory_policy_pass(gm, config=self._config())
        self.assertTrue(
            any(
                node.meta.get("recompute") == CheckpointPolicy.MUST_RECOMPUTE
                for node in gm.graph.nodes
            )
        )
        self.assertEqual(len(self._recomputed_nodes(gm)), 0)
        selective_activation_remat_pass(gm)

        self.assertGreaterEqual(len(self._recomputed_nodes(gm)), 1)

    def test_decomposition_is_a_standalone_pass_before_min_cut(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        grad = graph.placeholder("grad")
        log_probs = graph.call_function(
            torch.ops.aten._log_softmax.default, args=(x, -1, False)
        )
        loss = graph.call_function(torch.ops.aten.sum.default, args=(log_probs,))
        bwd = graph.call_function(
            torch.ops.aten._log_softmax_backward_data.default,
            args=(grad, log_probs, -1, torch.float32),
        )
        bwd.meta["autograd_backward"] = True
        graph.output((loss, bwd))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        self._fake_prop(gm, ((128, 1024), torch.float32), ((128, 1024), torch.float32))

        apply_decompositions_pass(
            gm,
            decomposition_table=self._log_softmax_decomposition_table(),
        )
        self.assertFalse(
            any(
                node.target == torch.ops.aten._log_softmax.default
                for node in gm.graph.nodes
            )
        )
        self.assertEqual(len(self._recomputed_nodes(gm)), 0)

        tag_with_memory_policy_pass(gm, config=self._config())
        bwd = next(
            node
            for node in gm.graph.nodes
            if node.target == torch.ops.aten._log_softmax_backward_data.default
        )
        self.assertIsInstance(bwd.args[1], torch.fx.Node)
        self.assertEqual(bwd.args[1].meta["recompute"], CheckpointPolicy.MUST_SAVE)
        selective_activation_remat_pass(gm)

        bwd = next(
            node
            for node in gm.graph.nodes
            if node.target == torch.ops.aten._log_softmax_backward_data.default
        )
        self.assertIsInstance(bwd.args[1], torch.fx.Node)
        self.assertFalse(bwd.args[1].name.endswith("_recomputed"))

    def test_min_cut_policy_respects_existing_checkpoint_policy(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        saved = graph.call_function(torch.ops.aten.sin.default, args=(x,))
        recompute = graph.call_function(torch.ops.aten.cos.default, args=(saved,))
        loss = graph.call_function(torch.ops.aten.sum.default, args=(recompute,))
        bwd = graph.call_function(torch.ops.aten.neg.default, args=(recompute,))
        bwd.meta["autograd_backward"] = True
        graph.output((loss, bwd))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        self._fake_prop(gm, ((64, 64), torch.float32))
        saved.meta["recompute"] = CheckpointPolicy.MUST_SAVE
        recompute.meta["recompute"] = CheckpointPolicy.PREFER_RECOMPUTE

        tag_with_memory_policy_pass(gm, config=self._config())

        self.assertEqual(saved.meta["recompute"], CheckpointPolicy.MUST_SAVE)
        self.assertIn(
            recompute.meta["recompute"],
            (CheckpointPolicy.PREFER_RECOMPUTE, CheckpointPolicy.MUST_RECOMPUTE),
        )
        selective_activation_remat_pass(gm)
        recomputed_targets = {node.target for node in self._recomputed_nodes(gm)}
        self.assertIn(torch.ops.aten.cos.default, recomputed_targets)
        self.assertNotIn(torch.ops.aten.sin.default, recomputed_targets)

    def test_explicit_subgraph_decomposition_and_min_cut_policy(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        grad = graph.placeholder("grad")
        log_probs = graph.call_function(
            torch.ops.aten._log_softmax.default, args=(x, -1, False)
        )
        loss = graph.call_function(torch.ops.aten.sum.default, args=(log_probs,))
        bwd = graph.call_function(
            torch.ops.aten._log_softmax_backward_data.default,
            args=(grad, log_probs, -1, torch.float32),
        )
        graph.output((loss, bwd))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        self._fake_prop(gm, ((128, 1024), torch.float32), ((128, 1024), torch.float32))

        for node in (log_probs, loss, bwd):
            node.meta.setdefault("custom", {})
            node.meta["custom"][SUBGRAPH_REGION] = "region"
            node.meta["custom"][SUBGRAPH_REGION_ROLE] = "fw_bw_grad_accum"
        bwd.meta["autograd_backward"] = True

        apply_subgraph_region_annotations_pass(gm)
        apply_decompositions_pass(
            gm,
            decomposition_table=self._log_softmax_decomposition_table(),
            recurse=True,
            apply_to_root=False,
        )
        submods = [
            module
            for module in gm.modules()
            if isinstance(module, torch.fx.GraphModule) and module is not gm
        ]
        self.assertEqual(len(submods), 1)
        submod = submods[0]
        tag_with_memory_policy_pass(submod, config=self._config())
        self.assertFalse(
            any(
                node.target == torch.ops.aten._log_softmax.default
                for node in submod.graph.nodes
            )
        )
        bwd = next(
            node
            for node in submod.graph.nodes
            if node.target == torch.ops.aten._log_softmax_backward_data.default
        )
        self.assertIsInstance(bwd.args[1], torch.fx.Node)
        self.assertEqual(bwd.args[1].meta["recompute"], CheckpointPolicy.MUST_SAVE)

        selective_activation_remat_pass(submod)
        bwd = next(
            node
            for node in submod.graph.nodes
            if node.target == torch.ops.aten._log_softmax_backward_data.default
        )
        self.assertIsInstance(bwd.args[1], torch.fx.Node)
        self.assertFalse(bwd.args[1].name.endswith("_recomputed"))


class TestBucketingPrefetchOrder(FSDPTest):
    """Guard that SAC + bucketing produces correct all_gather prefetch order.

    Uses the real Llama3 debug model with FSDP via the GraphTrainer path.
    Verifies that bucketed all_gather starts follow forward layer order
    (0, 1, 2, ...) and not reverse order (which was a prior bug).
    """

    BATCH_SIZE = 4
    SEQ_LEN = 128

    @staticmethod
    def _get_bucketed_ag_layer_order(gm):
        """Extract layer IDs from bucketed all_gather_into_tensor_out nodes.

        For each bucketed all_gather, searches its transitive users for
        a node with module_fqn under ``layers.<N>`` and records N.
        Returns deduplicated layer IDs in graph order.
        """
        layer_ids = []
        for node in gm.graph.nodes:
            if node.op != "call_function":
                continue
            if "all_gather_into_tensor_out" not in str(node.target):
                continue
            # BFS through users to find a node with layers.N FQN
            visited = set()
            queue = list(node.users)
            found_lid = None
            while queue and found_lid is None:
                u = queue.pop(0)
                if u in visited:
                    continue
                visited.add(u)
                fqn = u.meta.get("custom", {}).get(_MODULE_FQN, "")
                parts = fqn.split(".")
                if parts[0] == "layers" and len(parts) >= 2:
                    try:
                        found_lid = int(parts[1])
                    except ValueError:
                        pass
                else:
                    queue.extend(u.users)
            if found_lid is not None and (not layer_ids or layer_ids[-1] != found_lid):
                layer_ids.append(found_lid)
        return layer_ids

    def _run_and_get_layer_ids(self, fsdp_reshard_after_forward: str):
        """Run a single forward+backward step and return bucketed AG layer ids."""
        from torchtitan.components.tokenizer import HuggingFaceTokenizer
        from torchtitan.experiments.graph_trainer.common_utils import (
            annotate_graph_trainer_model,
        )
        from torchtitan.experiments.graph_trainer.llama3 import (
            build_model_config as build_llama3_model_config,
        )
        from torchtitan.experiments.graph_trainer.simple_fsdp import (
            data_parallel,
            MixedPrecisionPolicy,
        )
        from torchtitan.experiments.graph_trainer.tests._trainer_test_utils import (
            build_minimal_trainer,
        )
        from torchtitan.experiments.graph_trainer.trainer import GraphTrainer

        parallelism_context = ParallelismContext(
            dp_shard=-1,
            dp_replicate=1,
            cp=1,
            tp=1,
            pp=1,
            ep=1,
            world_size=self.world_size,
            enable_sequence_parallel=False,
        )

        model_config = build_llama3_model_config("debugmodel")
        vocab_size = model_config.vocab_size

        with torch.device("meta"):
            model = model_config.build()

        annotate_graph_trainer_model(model)
        from torchtitan.experiments.graph_trainer.common_utils import (
            get_simple_fsdp_mesh,
        )

        fsdp_mesh = get_simple_fsdp_mesh(parallelism_context)
        mp_policy = MixedPrecisionPolicy(
            param_dtype=torch.bfloat16,
            reduce_dtype=torch.float32,
        )
        model = data_parallel(
            model, device_mesh=fsdp_mesh, mode="fully_shard", mp_policy=mp_policy
        )
        model.to_empty(device="cuda")
        with torch.no_grad():
            model.init_states(buffer_device=None)
        model.train()

        # Use GraphTrainer's full path: trace + construct_default_graph_passes
        trainer = build_minimal_trainer(
            model,
            model_config,
            GraphTrainer,
            tokenizer=HuggingFaceTokenizer(tokenizer_path="./tests/assets/tokenizer"),
            fsdp_reshard_after_forward=fsdp_reshard_after_forward,
            parallelism_context=parallelism_context,
        )

        num_tokens = self.BATCH_SIZE * self.SEQ_LEN
        inputs = torch.randint(0, vocab_size, (num_tokens,), device="cuda")
        labels = torch.randint(0, vocab_size, (num_tokens,), device="cuda")
        # The dataloader supplies per-document positions, which the trainer
        # requires to build the block-causal FlexInnerAttention mask.
        positions = torch.arange(self.SEQ_LEN, device="cuda", dtype=torch.int32).repeat(
            self.BATCH_SIZE
        )
        global_loss_token_counts = torch.tensor(
            num_tokens, dtype=torch.float, device="cuda"
        )

        # One accumulation step traces the model and applies all graph passes.
        trainer.engine.forward_backward(
            microbatch_groups=[
                [
                    TokenizedTrainingMicrobatch(
                        input=inputs,
                        positions=positions,
                        labels=labels,
                        padding_mask=torch.zeros_like(labels, dtype=torch.bool),
                        loss_token_counts=torch.tensor(labels.numel()),
                        routing_token_counts=torch.tensor([labels.numel()]),
                    )
                ]
            ],
            global_loss_token_counts=global_loss_token_counts,
            global_routing_token_counts=global_loss_token_counts.unsqueeze(0),
        )

        layer_ids = self._get_bucketed_ag_layer_order(trainer.engine._traced_step.gm)
        self.assertGreater(len(layer_ids), 0, "No layer all_gather nodes found")
        return layer_ids

    def test_forward_allgather_prefetch_follows_layer_order(self):
        """Without reshard-after-forward, all all_gathers are in forward and
        must appear in non-decreasing layer order 0 → N."""
        layer_ids = self._run_and_get_layer_ids(fsdp_reshard_after_forward="never")

        for i in range(1, len(layer_ids)):
            self.assertGreaterEqual(
                layer_ids[i],
                layer_ids[i - 1],
                f"Forward all_gather prefetch order violated: "
                f"layer {layer_ids[i]} before layer {layer_ids[i - 1]} "
                f"(full order: {layer_ids})",
            )

    def test_allgather_prefetch_with_reshard_after_forward(self):
        """With reshard-after-forward, backward also issues all_gathers.
        The graph-order sequence must be forward (0 → N) then backward (N → 0)."""
        layer_ids = self._run_and_get_layer_ids(fsdp_reshard_after_forward="always")

        # Split forward (ascending) and backward (descending) at the peak.
        peak = max(range(len(layer_ids)), key=lambda i: layer_ids[i])
        forward_ids = layer_ids[: peak + 1]
        backward_ids = layer_ids[peak:]

        # Backward all_gathers should exist when reshard-after-forward is on.
        self.assertGreater(
            len(backward_ids),
            1,
            f"Expected backward all_gathers with reshard-after-forward, "
            f"got order: {layer_ids}",
        )

        for i in range(1, len(forward_ids)):
            self.assertGreaterEqual(
                forward_ids[i],
                forward_ids[i - 1],
                f"Forward all_gather prefetch order violated: "
                f"layer {forward_ids[i]} before layer {forward_ids[i - 1]} "
                f"(full order: {layer_ids})",
            )

        for i in range(1, len(backward_ids)):
            self.assertLessEqual(
                backward_ids[i],
                backward_ids[i - 1],
                f"Backward all_gather prefetch order violated: "
                f"layer {backward_ids[i]} after layer {backward_ids[i - 1]} "
                f"(full order: {layer_ids})",
            )

    def test_drops_assert_async_and_dead_chain(self):
        # _assert_async is side-effectful, so plain DCE keeps it (and its whole
        # le/all condition chain). The pass erases the assert, then DCE reaps the
        # now-orphaned chain; unrelated live nodes are untouched.
        aten = torch.ops.aten
        g = torch.fx.Graph()
        x = g.placeholder("x")
        le = g.call_function(aten.le.Scalar, (x, 5))
        reduced = g.call_function(aten.all.default, (le,))
        g.call_function(aten._assert_async.msg, (reduced, "cond"))  # side-effect
        out = g.call_function(aten.relu.default, (x,))
        g.output(out)
        gm = torch.fx.GraphModule(torch.nn.Module(), g)

        eliminate_dead_code_pass(gm)
        targets = [n.target for n in gm.graph.nodes if n.op == "call_function"]
        self.assertNotIn(aten._assert_async.msg, targets)
        self.assertNotIn(aten.le.Scalar, targets)
        self.assertNotIn(aten.all.default, targets)
        self.assertIn(aten.relu.default, targets)


class TestRemoveDetachPass(TestCase):
    """Unit tests for the remove_detach_pass graph pass."""

    def _build_detach_gm(self, op_targets):
        """Build a GraphModule with a chain of call_function nodes.

        Each op in op_targets becomes a call_function node chained sequentially:
        placeholder(x) -> op1(x) -> op2(...) -> ... -> output.
        """
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        last = x
        for target in op_targets:
            last = graph.call_function(target, args=(last,))
        graph.output(last)
        return torch.fx.GraphModule(torch.nn.Module(), graph)

    def _count_detach_nodes(self, gm):
        """Count aten.detach.default call_function nodes."""
        return sum(
            1
            for n in gm.graph.nodes
            if n.op == "call_function" and n.target is torch.ops.aten.detach.default
        )

    def _count_call_function_nodes(self, gm):
        """Count all call_function nodes."""
        return sum(1 for n in gm.graph.nodes if n.op == "call_function")

    def test_detach_nodes_removed(self):
        """Detach nodes are removed from a simple graph containing them."""
        gm = self._build_detach_gm(
            [
                torch.ops.aten.relu.default,
                torch.ops.aten.detach.default,
                torch.ops.aten.neg.default,
            ]
        )
        self.assertEqual(self._count_detach_nodes(gm), 1)

        result = remove_detach_pass(gm)

        self.assertEqual(self._count_detach_nodes(result), 0)
        # relu and neg should remain
        self.assertEqual(self._count_call_function_nodes(result), 2)

    def test_graph_without_detach_unchanged(self):
        """Graphs without detach nodes are returned unchanged."""
        gm = self._build_detach_gm(
            [
                torch.ops.aten.relu.default,
                torch.ops.aten.neg.default,
            ]
        )
        num_nodes_before = len(list(gm.graph.nodes))

        result = remove_detach_pass(gm)

        self.assertIs(result, gm)
        self.assertEqual(len(list(result.graph.nodes)), num_nodes_before)

    def test_numerics_preserved(self):
        """Forward outputs are preserved after removing detach nodes."""
        gm = self._build_detach_gm(
            [
                torch.ops.aten.relu.default,
                torch.ops.aten.detach.default,
                torch.ops.aten.neg.default,
            ]
        )
        x = torch.randn(4, 4)
        expected = torch.neg(torch.detach_copy(torch.relu(x)))

        remove_detach_pass(gm)
        actual = gm(x)

        self.assertEqual(actual, expected)

    def test_detach_with_multiple_users(self):
        """Detach node with multiple users: all uses are replaced."""
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        detach = graph.call_function(torch.ops.aten.detach.default, args=(x,))
        relu = graph.call_function(torch.ops.aten.relu.default, args=(detach,))
        neg = graph.call_function(torch.ops.aten.neg.default, args=(detach,))
        graph.output((relu, neg))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        self.assertEqual(self._count_detach_nodes(gm), 1)

        remove_detach_pass(gm)

        self.assertEqual(self._count_detach_nodes(gm), 0)

        # Both relu and neg should now consume the placeholder directly
        for node in gm.graph.nodes:
            if node.op == "call_function" and node.target in (
                torch.ops.aten.relu.default,
                torch.ops.aten.neg.default,
            ):
                self.assertEqual(node.args[0].op, "placeholder")

        # Verify numerics
        x = torch.randn(4, 4)
        relu_out, neg_out = gm(x)
        self.assertEqual(relu_out, torch.relu(x))
        self.assertEqual(neg_out, torch.neg(x))

    def test_nested_detach_chain(self):
        """Nested detach chain (detach -> detach -> detach) is fully removed."""
        gm = self._build_detach_gm(
            [
                torch.ops.aten.relu.default,
                torch.ops.aten.detach.default,
                torch.ops.aten.detach.default,
                torch.ops.aten.detach.default,
                torch.ops.aten.neg.default,
            ]
        )
        self.assertEqual(self._count_detach_nodes(gm), 3)

        remove_detach_pass(gm)

        self.assertEqual(self._count_detach_nodes(gm), 0)
        self.assertEqual(self._count_call_function_nodes(gm), 2)

        # Verify numerics
        x = torch.randn(4, 4)
        expected = torch.neg(torch.relu(x))
        self.assertEqual(gm(x), expected)


class TestChunkPasses(TestCase):
    def _mark_chunk_body(
        self,
        node,
        *,
        fqn: str = "layers.0.moe",
        chunk_id: int,
        backward: bool = False,
        ep: str | None = None,
        token_exchange: bool = False,
        producer: str | None = None,
    ):
        custom = dict(node.meta.get("custom", {}))
        custom[_MODULE_FQN] = fqn
        if ep is not None:
            custom["EP"] = ep
            if token_exchange:
                custom[_EP_TOKEN_EXCHANGE] = ep
        node.meta["custom"] = custom
        node.meta["chunk_id"] = chunk_id
        node.meta["chunked_region_fqn"] = fqn
        node.meta["chunked_region_role"] = "body"
        if producer is not None:
            node.meta["chunked_region_producer"] = producer
        if backward:
            node.meta["autograd_backward"] = True
        return node

    def _build_ep_overlap_schedule_gm(self, *, backward: bool = False):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional
        outputs = []
        for chunk_id in (1, 0) if backward else (0, 1):
            pre = graph.call_function(torch.ops.aten.relu.default, args=(x,))
            first_launch = graph.call_function(
                c10d.all_to_all_single.default,
                args=(pre, [], [], "ep"),
            )
            first_wait = graph.call_function(
                c10d.wait_tensor.default, args=(first_launch,)
            )
            compute = graph.call_function(
                torch.ops.aten.neg.default, args=(first_wait,)
            )
            second_launch = graph.call_function(
                c10d.all_to_all_single.default,
                args=(compute, [], [], "ep"),
            )
            second_wait = graph.call_function(
                c10d.wait_tensor.default, args=(second_launch,)
            )
            tail = graph.call_function(torch.ops.aten.neg.default, args=(second_wait,))
            outputs.append(tail)

            first_ep = "combine" if backward else "dispatch"
            second_ep = "dispatch" if backward else "combine"
            self._mark_chunk_body(
                pre, chunk_id=chunk_id, backward=backward, ep=first_ep
            )
            self._mark_chunk_body(
                first_launch,
                chunk_id=chunk_id,
                backward=backward,
                ep=first_ep,
                token_exchange=True,
            )
            self._mark_chunk_body(
                first_wait, chunk_id=chunk_id, backward=backward, ep=first_ep
            )
            self._mark_chunk_body(compute, chunk_id=chunk_id, backward=backward)
            self._mark_chunk_body(
                second_launch,
                chunk_id=chunk_id,
                backward=backward,
                ep=second_ep,
                token_exchange=True,
            )
            self._mark_chunk_body(
                second_wait, chunk_id=chunk_id, backward=backward, ep=second_ep
            )
            self._mark_chunk_body(tail, chunk_id=chunk_id, backward=backward)

        graph.output(tuple(outputs))
        return torch.fx.GraphModule(torch.nn.Module(), graph)

    def _schedule_ep_overlap_and_order(
        self,
        gm,
        *,
        module_pattern: str = "layers.*.moe",
        pair_first_token_exchange: bool = True,
    ):
        _schedule_ep_overlap_regions(
            gm,
            module_pattern=module_pattern,
            require_all_to_all=True,
            pair_first_token_exchange=pair_first_token_exchange,
        )
        return {node: idx for idx, node in enumerate(gm.graph.nodes)}

    def _assert_nodes_in_order(self, order, nodes):
        for before, after in zip(nodes, nodes[1:]):
            self.assertLess(order[before], order[after])

    def _build_ep_sync_copy_schedule_gm(
        self,
        *,
        fqn: str = "layers.0.moe",
        copies_per_chunk: int = 2,
        cpu_destination: bool = True,
    ):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional
        refs = {}
        for chunk_id in (0, 1):
            router = graph.call_function(torch.ops.aten.relu.default, args=(x,))
            count = graph.call_function(
                c10d.all_to_all_single.default, args=(router, [], [], "ep")
            )
            body_nodes = [router, count]
            copies = []
            consumers = []
            consumer_value = None
            for copy_idx in range(copies_per_chunk):
                producer = count
                if copy_idx:
                    producer = graph.call_function(
                        torch.ops.aten.neg.default, args=(count,)
                    )
                    body_nodes.append(producer)

                copy_kwargs = {"non_blocking": False}
                if cpu_destination:
                    copy_kwargs["device"] = torch.device("cpu")
                copy = graph.call_function(
                    torch.ops.aten._to_copy.default,
                    args=(producer,),
                    kwargs=copy_kwargs,
                )
                if cpu_destination:
                    copy.meta["val"] = torch.empty(2, device="cpu")
                consumer = graph.call_function(
                    torch.ops.aten._local_scalar_dense.default, args=(copy,)
                )
                copies.append(copy)
                consumers.append(consumer)
                body_nodes.extend((copy, consumer))
                consumer_value = (
                    consumer
                    if consumer_value is None
                    else graph.call_function(
                        torch.ops.aten.add.Tensor, args=(consumer_value, consumer)
                    )
                )
                if consumer_value is not consumer:
                    body_nodes.append(consumer_value)

            if consumer_value is None:
                consumer_value = count
            dispatch = graph.call_function(
                c10d.all_to_all_single.default, args=(consumer_value, [], [], "ep")
            )
            wait = graph.call_function(c10d.wait_tensor.default, args=(dispatch,))
            body_nodes.extend((dispatch, wait))
            for node in body_nodes:
                self._mark_chunk_body(
                    node,
                    fqn=fqn,
                    chunk_id=chunk_id,
                    ep="dispatch",
                    token_exchange=node is dispatch,
                )
            for copy in copies:
                copy.meta["custom"][_EP_TOKEN_COUNT_SYNC] = "dispatch"
            refs[chunk_id] = {
                "copies": tuple(copies),
                "consumers": tuple(consumers),
                "dispatch": dispatch,
                "wait": wait,
            }

        graph.output((refs[0]["wait"], refs[1]["wait"]))
        return torch.fx.GraphModule(torch.nn.Module(), graph), refs

    def _compile_config_for_ep_overlap_test(self):
        from types import SimpleNamespace

        traced_result = SimpleNamespace(
            num_static_inputs=2,
            state_fqns=[],
            graph_state=GraphStateSpec(),
        )
        config = SimpleNamespace(
            model=SimpleNamespace(layers=[object()]),
            parallelism=SimpleNamespace(
                expert_parallel_degree=1,
                fsdp_reshard_after_forward="default",
                pipeline_parallel_degree=1,
            ),
            compile=GraphTrainerCompileConfig(
                ep_overlap=EpOverlapConfig(
                    enabled=True,
                    chunk_dim="batch",
                    module_fqn="layers.*",
                ),
                cpu_offload_prefetch_n_layers=1,
                cpu_offload_defer_n_layers=1,
                cpu_offload_budget_gb=1.0,
                memory_policy="default",
                inductor_compilation="full",
                numerics_changing_optim=False,
                enable_fsdp_ag_rs_overlap=False,
                enable_fsdp_dense_region_overlap=False,
                precompile_artifact_dir="",
            ),
        )
        return traced_result, config

    def _compile_pass_names(self, traced_result, config):
        def pass_name(pass_fn):
            return (
                pass_fn.func.__name__ if hasattr(pass_fn, "func") else pass_fn.__name__
            )

        return [
            pass_name(pass_fn)
            for pass_fn in compile_time_passes(
                traced_result, config, use_cuda_graph=False
            )
        ]

    def test_ep_overlap_pass_pipeline_order(self):
        traced_result, config = self._compile_config_for_ep_overlap_test()
        names = self._compile_pass_names(traced_result, config)
        dead_code_indices = [
            i for i, name in enumerate(names) if name == "eliminate_dead_code_pass"
        ]
        populate_idx = names.index("populate_eager_chunk_metadata_pass")
        post_chunk_dce = min(i for i in dead_code_indices if i > populate_idx)

        self.assertLess(
            names.index("canonicalize_graph_pass"),
            names.index("deduplicate_fsdp_unshard_chains_pass"),
        )
        self.assertLess(
            names.index("deduplicate_fsdp_unshard_chains_pass"),
            names.index("tag_with_memory_policy_pass"),
        )
        self.assertLess(
            names.index("apply_cpu_offload_pass"),
            names.index("selective_activation_remat_pass"),
        )
        self.assertLess(names.index("selective_activation_remat_pass"), populate_idx)
        self.assertLess(populate_idx, names.index("isolate_ep_process_group_pass"))
        self.assertLess(
            names.index("isolate_ep_process_group_pass"),
            post_chunk_dce,
        )
        self.assertLess(
            post_chunk_dce,
            names.index("joint_transformer_block_bucketing_reordering_pass"),
        )
        self.assertLess(
            names.index("joint_transformer_block_bucketing_reordering_pass"),
            names.index("ep_overlap_schedule_pass"),
        )
        self.assertLess(
            names.index("ep_overlap_schedule_pass"),
            names.index("full_inductor_compilation_pass"),
        )

    def test_fsdp_dense_region_scheduler_pass_gating(self):
        def transformer_batch_default(config):
            pass

        def transformer_batch_explicit(config):
            config.compile.enable_fsdp_dense_region_overlap = True

        def moe_ep_default(config):
            config.compile.ep_overlap.module_fqn = "layers.*.moe"

        def moe_ep_explicit(config):
            config.compile.ep_overlap.module_fqn = "layers.*.moe"
            config.compile.enable_fsdp_dense_region_overlap = True

        def fsdp_dense_without_ep(config):
            config.compile.ep_overlap.enabled = False
            config.compile.enable_fsdp_dense_region_overlap = True

        cases = (
            ("transformer_batch_default", transformer_batch_default, True, False, None),
            (
                "transformer_batch_explicit",
                transformer_batch_explicit,
                True,
                False,
                "ignored when compile.ep_overlap.enabled is set",
            ),
            ("moe_ep_default", moe_ep_default, True, False, None),
            (
                "moe_ep_explicit",
                moe_ep_explicit,
                True,
                False,
                "ignored when compile.ep_overlap.enabled is set",
            ),
            ("fsdp_dense_without_ep", fsdp_dense_without_ep, False, True, None),
        )
        for (
            name,
            configure,
            expects_ep_schedule,
            expects_fsdp_schedule,
            warning,
        ) in cases:
            with self.subTest(name=name):
                traced_result, config = self._compile_config_for_ep_overlap_test()
                configure(config)
                if warning is None:
                    names = self._compile_pass_names(traced_result, config)
                else:
                    with self.assertWarnsRegex(UserWarning, warning):
                        names = self._compile_pass_names(traced_result, config)

                self.assertEqual(
                    "ep_overlap_schedule_pass" in names,
                    expects_ep_schedule,
                )
                self.assertEqual(
                    "schedule_fsdp_comms_to_dense_regions_pass" in names,
                    expects_fsdp_schedule,
                )

    def test_moe_edp_shard_bucket_plan_splits_expert_buckets(self):
        buckets = get_default_transformer_block_buckets(
            3,
            moe_layer_ids=frozenset({1}),
            split_moe_expert_buckets=True,
        )

        self.assertIn(
            [
                "layers.1.attention_norm",
                "layers.1.attention",
                "layers.1.ffn_norm",
                "layers.1.moe.router",
                "layers.1.moe.shared_experts",
            ],
            buckets,
        )
        self.assertIn("layers.1.moe.routed_experts", buckets)
        self.assertNotIn("layers.1", buckets)

    def test_moe_ep_annotations_cover_all_to_all_dispatcher(self):
        from torchtitan.models.common.token_dispatcher import AllToAllTokenDispatcher

        annotate_moe_ep_regions()

        expected_annotations = [
            (AllToAllTokenDispatcher.dispatch, {"EP": "dispatch"}),
            (
                AllToAllTokenDispatcher._token_count_exchange,
                {_EP_TOKEN_COUNT_EXCHANGE: "dispatch"},
            ),
            (
                AllToAllTokenDispatcher._sync_token_count_exchange,
                {_EP_TOKEN_COUNT_SYNC: "dispatch"},
            ),
            (
                AllToAllTokenDispatcher._dispatch_token_exchange,
                {_EP_TOKEN_EXCHANGE: "dispatch"},
            ),
            (
                AllToAllTokenDispatcher._combine_token_exchange,
                {_EP_TOKEN_EXCHANGE: "combine"},
            ),
            (AllToAllTokenDispatcher.combine, {"EP": "combine"}),
        ]
        for method, annotation in expected_annotations:
            self.assertEqual(
                inspect.getclosurevars(method).nonlocals["annotation_dict"],
                annotation,
            )

    def test_ep_overlap_reorders_forward_and_backward_token_exchange_blocks(self):
        for backward, first_chunk in ((False, 0), (True, 1)):
            with self.subTest(backward=backward):
                gm = self._build_ep_overlap_schedule_gm(backward=backward)
                order = self._schedule_ep_overlap_and_order(gm)
                nodes = list(gm.graph.nodes)

                def named(chunk_id, target, ep=None):
                    matches = [
                        node
                        for node in nodes
                        if node.meta.get("chunk_id") == chunk_id
                        and node.target == target
                        and (
                            ep is None
                            or node.meta.get("custom", {}).get(_EP_TOKEN_EXCHANGE) == ep
                            or node.meta.get("custom", {}).get("EP") == ep
                        )
                    ]
                    self.assertEqual(len(matches), 1)
                    return matches[0]

                second_chunk = 1 - first_chunk
                c10d = torch.ops._c10d_functional
                first_dispatch_launch = named(
                    first_chunk,
                    c10d.all_to_all_single.default,
                    ep="dispatch",
                )
                second_dispatch_launch = named(
                    second_chunk,
                    c10d.all_to_all_single.default,
                    ep="dispatch",
                )
                first_dispatch_wait = next(iter(first_dispatch_launch.users))
                first_combine_launch = named(
                    first_chunk,
                    c10d.all_to_all_single.default,
                    ep="combine",
                )
                second_combine_launch = named(
                    second_chunk,
                    c10d.all_to_all_single.default,
                    ep="combine",
                )

                self.assertLess(
                    order[first_dispatch_launch], order[second_dispatch_launch]
                )
                self.assertLess(
                    order[second_dispatch_launch], order[first_dispatch_wait]
                )
                self.assertLess(
                    order[first_combine_launch], order[second_combine_launch]
                )
                if not backward:
                    self.assertLess(
                        order[first_combine_launch],
                        order[next(iter(second_dispatch_launch.users))],
                    )
                else:
                    self.assertLess(
                        order[second_dispatch_launch],
                        order[next(iter(first_dispatch_launch.users))],
                    )

    def test_ep_overlap_schedules_moe_shaped_forward_and_backward_order(self):
        c10d = torch.ops._c10d_functional

        def build_forward_gm():
            graph = torch.fx.Graph()
            x = graph.placeholder("x")
            refs = {}
            for chunk_id in (0, 1):
                router = graph.call_function(torch.ops.aten.relu.default, args=(x,))
                count = graph.call_function(
                    c10d.all_to_all_single.default, args=(router, [], [], "ep")
                )
                sync = graph.call_function(
                    torch.ops.aten._local_scalar_dense.default, args=(count,)
                )
                dispatch = graph.call_function(
                    c10d.all_to_all_single.default, args=(sync, [], [], "ep")
                )
                shared = graph.call_function(torch.ops.aten.neg.default, args=(x,))
                dispatch_wait = graph.call_function(
                    c10d.wait_tensor.default, args=(dispatch,)
                )
                grouped_mm = graph.call_function(
                    torch.ops.aten.relu.default, args=(dispatch_wait,)
                )
                combine = graph.call_function(
                    c10d.all_to_all_single.default, args=(grouped_mm, [], [], "ep")
                )
                combine_wait = graph.call_function(
                    c10d.wait_tensor.default, args=(combine,)
                )
                refs[chunk_id] = {
                    "router": router,
                    "count": count,
                    "sync": sync,
                    "dispatch": dispatch,
                    "shared": shared,
                    "dispatch_wait": dispatch_wait,
                    "grouped_mm": grouped_mm,
                    "combine": combine,
                    "combine_wait": combine_wait,
                }
                for node in (router, count, sync, dispatch, dispatch_wait):
                    self._mark_chunk_body(
                        node,
                        chunk_id=chunk_id,
                        ep="dispatch",
                        token_exchange=node is dispatch,
                    )
                for node in (shared, grouped_mm):
                    self._mark_chunk_body(node, chunk_id=chunk_id)
                for node in (combine, combine_wait):
                    self._mark_chunk_body(
                        node,
                        chunk_id=chunk_id,
                        ep="combine",
                        token_exchange=node is combine,
                    )
            graph.output((refs[0]["combine_wait"], refs[1]["combine_wait"]))
            return torch.fx.GraphModule(torch.nn.Module(), graph), refs

        def build_backward_gm():
            graph = torch.fx.Graph()
            x = graph.placeholder("x")
            refs = {}
            for chunk_id in (0, 1):
                combine = graph.call_function(
                    c10d.all_to_all_single.default, args=(x, [], [], "ep")
                )
                remat_grouped_mm = graph.call_function(
                    torch.ops.aten.neg.default, args=(x,)
                )
                combine_wait = graph.call_function(
                    c10d.wait_tensor.default, args=(combine,)
                )
                input_grad = graph.call_function(
                    torch.ops.aten.relu.default, args=(combine_wait,)
                )
                dispatch = graph.call_function(
                    c10d.all_to_all_single.default, args=(input_grad, [], [], "ep")
                )
                wgrad = graph.call_function(
                    torch.ops.aten.neg.default, args=(dispatch,)
                )
                dispatch_wait = graph.call_function(
                    c10d.wait_tensor.default, args=(dispatch,)
                )
                refs[chunk_id] = {
                    "combine": combine,
                    "remat_grouped_mm": remat_grouped_mm,
                    "combine_wait": combine_wait,
                    "input_grad": input_grad,
                    "dispatch": dispatch,
                    "wgrad": wgrad,
                    "dispatch_wait": dispatch_wait,
                }
                for node in (combine, combine_wait):
                    self._mark_chunk_body(
                        node,
                        chunk_id=chunk_id,
                        backward=True,
                        ep="combine",
                        token_exchange=node is combine,
                    )
                for node in (remat_grouped_mm, input_grad, wgrad):
                    self._mark_chunk_body(node, chunk_id=chunk_id, backward=True)
                for node in (dispatch, dispatch_wait):
                    self._mark_chunk_body(
                        node,
                        chunk_id=chunk_id,
                        backward=True,
                        ep="dispatch",
                        token_exchange=node is dispatch,
                    )
            graph.output((refs[1]["dispatch_wait"], refs[0]["dispatch_wait"]))
            return torch.fx.GraphModule(torch.nn.Module(), graph), refs

        fwd_gm, fwd = build_forward_gm()
        fwd_order = self._schedule_ep_overlap_and_order(fwd_gm)
        self._assert_nodes_in_order(
            fwd_order,
            [
                fwd[0]["router"],
                fwd[0]["count"],
                fwd[0]["sync"],
                fwd[1]["router"],
                fwd[1]["count"],
                fwd[1]["sync"],
                fwd[0]["dispatch"],
                fwd[1]["dispatch"],
                fwd[0]["shared"],
                fwd[1]["shared"],
                fwd[0]["dispatch_wait"],
                fwd[0]["grouped_mm"],
                fwd[0]["combine"],
                fwd[1]["dispatch_wait"],
                fwd[1]["grouped_mm"],
                fwd[1]["combine"],
                fwd[0]["combine_wait"],
                fwd[1]["combine_wait"],
            ],
        )

        bwd_gm, bwd = build_backward_gm()
        bwd_order = self._schedule_ep_overlap_and_order(bwd_gm)
        self._assert_nodes_in_order(
            bwd_order,
            [
                bwd[1]["combine"],
                bwd[0]["combine"],
                bwd[1]["remat_grouped_mm"],
                bwd[0]["remat_grouped_mm"],
                bwd[1]["combine_wait"],
                bwd[1]["input_grad"],
                bwd[1]["dispatch"],
                bwd[0]["combine_wait"],
                bwd[0]["input_grad"],
                bwd[0]["dispatch"],
                bwd[1]["wgrad"],
                bwd[0]["wgrad"],
                bwd[1]["dispatch_wait"],
                bwd[0]["dispatch_wait"],
            ],
        )

    def test_ep_overlap_keeps_transformer_batch_first_marker_wait_gated(self):
        c10d = torch.ops._c10d_functional

        def build_forward_gm():
            graph = torch.fx.Graph()
            x = graph.placeholder("x")
            refs = {}
            for chunk_id in (0, 1):
                dense_prefix = graph.call_function(
                    torch.ops.aten.relu.default, args=(x,)
                )
                router = graph.call_function(
                    torch.ops.aten.neg.default, args=(dense_prefix,)
                )
                count = graph.call_function(
                    c10d.all_to_all_single.default, args=(router, [], [], "ep")
                )
                sync = graph.call_function(
                    torch.ops.aten._local_scalar_dense.default, args=(count,)
                )
                dispatch = graph.call_function(
                    c10d.all_to_all_single.default, args=(sync, [], [], "ep")
                )
                dispatch_wait = graph.call_function(
                    c10d.wait_tensor.default, args=(dispatch,)
                )
                grouped_mm = graph.call_function(
                    torch.ops.aten.relu.default, args=(dispatch_wait,)
                )
                combine = graph.call_function(
                    c10d.all_to_all_single.default, args=(grouped_mm, [], [], "ep")
                )
                combine_wait = graph.call_function(
                    c10d.wait_tensor.default, args=(combine,)
                )
                refs[chunk_id] = {
                    "dense_prefix": dense_prefix,
                    "router": router,
                    "count": count,
                    "sync": sync,
                    "dispatch": dispatch,
                    "dispatch_wait": dispatch_wait,
                    "grouped_mm": grouped_mm,
                    "combine": combine,
                    "combine_wait": combine_wait,
                }
                for node in (dense_prefix, router, count, sync, dispatch):
                    self._mark_chunk_body(
                        node,
                        fqn="layers.0",
                        chunk_id=chunk_id,
                        ep="dispatch",
                        token_exchange=node is dispatch,
                    )
                for node in (dispatch_wait, grouped_mm, combine, combine_wait):
                    self._mark_chunk_body(
                        node,
                        fqn="layers.0",
                        chunk_id=chunk_id,
                        ep="combine",
                        token_exchange=node is combine,
                    )
            graph.output((refs[0]["combine_wait"], refs[1]["combine_wait"]))
            return torch.fx.GraphModule(torch.nn.Module(), graph), refs

        def build_backward_gm():
            graph = torch.fx.Graph()
            x = graph.placeholder("x")
            refs = {}
            for chunk_id in (0, 1):
                remat_prefix = graph.call_function(
                    torch.ops.aten.relu.default, args=(x,)
                )
                combine = graph.call_function(
                    c10d.all_to_all_single.default, args=(remat_prefix, [], [], "ep")
                )
                combine_wait = graph.call_function(
                    c10d.wait_tensor.default, args=(combine,)
                )
                input_grad = graph.call_function(
                    torch.ops.aten.relu.default, args=(combine_wait,)
                )
                dispatch = graph.call_function(
                    c10d.all_to_all_single.default, args=(input_grad, [], [], "ep")
                )
                dispatch_wait = graph.call_function(
                    c10d.wait_tensor.default, args=(dispatch,)
                )
                refs[chunk_id] = {
                    "remat_prefix": remat_prefix,
                    "combine": combine,
                    "combine_wait": combine_wait,
                    "input_grad": input_grad,
                    "dispatch": dispatch,
                    "dispatch_wait": dispatch_wait,
                }
                for node in (remat_prefix, combine, combine_wait):
                    self._mark_chunk_body(
                        node,
                        fqn="layers.0",
                        chunk_id=chunk_id,
                        backward=True,
                        ep="combine",
                        token_exchange=node is combine,
                    )
                for node in (input_grad, dispatch, dispatch_wait):
                    self._mark_chunk_body(
                        node,
                        fqn="layers.0",
                        chunk_id=chunk_id,
                        backward=True,
                        ep="dispatch",
                        token_exchange=node is dispatch,
                    )
            graph.output((refs[1]["dispatch_wait"], refs[0]["dispatch_wait"]))
            return torch.fx.GraphModule(torch.nn.Module(), graph), refs

        fwd_gm, fwd = build_forward_gm()
        fwd_order = self._schedule_ep_overlap_and_order(
            fwd_gm,
            module_pattern="layers.*",
            pair_first_token_exchange=False,
        )
        self._assert_nodes_in_order(
            fwd_order,
            [
                fwd[0]["dense_prefix"],
                fwd[0]["dispatch"],
                fwd[1]["dense_prefix"],
                fwd[1]["dispatch"],
                fwd[0]["dispatch_wait"],
                fwd[0]["combine"],
                fwd[1]["dispatch_wait"],
                fwd[1]["combine"],
            ],
        )

        bwd_gm, bwd = build_backward_gm()
        bwd_order = self._schedule_ep_overlap_and_order(
            bwd_gm,
            module_pattern="layers.*",
            pair_first_token_exchange=False,
        )
        self._assert_nodes_in_order(
            bwd_order,
            [
                bwd[1]["remat_prefix"],
                bwd[1]["combine"],
                bwd[0]["remat_prefix"],
                bwd[0]["combine"],
                bwd[1]["combine_wait"],
                bwd[1]["dispatch"],
                bwd[0]["combine_wait"],
                bwd[0]["dispatch"],
                bwd[1]["dispatch_wait"],
                bwd[0]["dispatch_wait"],
            ],
        )

    def test_ep_overlap_pairs_token_count_sync_cpu_copies_before_consumers(self):
        gm, refs = self._build_ep_sync_copy_schedule_gm()

        order = self._schedule_ep_overlap_and_order(gm)

        copies = refs[0]["copies"] + refs[1]["copies"]
        consumers = refs[0]["consumers"] + refs[1]["consumers"]
        self.assertEqual(
            [copy.kwargs["non_blocking"] for copy in copies],
            [True, True, True, False],
        )
        self.assertLess(
            max(order[copy] for copy in copies),
            min(order[consumer] for consumer in consumers),
        )
        self.assertLess(
            max(order[consumer] for consumer in consumers),
            min(order[refs[chunk_id]["dispatch"]] for chunk_id in (0, 1)),
        )

    def test_ep_overlap_hoists_token_count_sync_cpu_copies_per_chunk(self):
        gm, refs = self._build_ep_sync_copy_schedule_gm(fqn="layers.0")

        order = self._schedule_ep_overlap_and_order(
            gm,
            module_pattern="layers.*",
            pair_first_token_exchange=False,
        )

        for chunk_id in (0, 1):
            self.assertEqual(
                [copy.kwargs["non_blocking"] for copy in refs[chunk_id]["copies"]],
                [True, False],
            )
            self.assertLess(
                max(order[copy] for copy in refs[chunk_id]["copies"]),
                min(order[consumer] for consumer in refs[chunk_id]["consumers"]),
            )
            self.assertLess(
                max(order[consumer] for consumer in refs[chunk_id]["consumers"]),
                order[refs[chunk_id]["dispatch"]],
            )
        self.assertLess(order[refs[0]["dispatch"]], order[refs[1]["copies"][0]])

    def test_ep_overlap_rejects_malformed_token_count_sync_copy_count(self):
        gm, _refs = self._build_ep_sync_copy_schedule_gm(copies_per_chunk=1)

        with self.assertRaisesRegex(ValueError, "exactly two token-count sync"):
            _schedule_ep_overlap_regions(
                gm,
                module_pattern="layers.*.moe",
                require_all_to_all=True,
                pair_first_token_exchange=True,
            )

    def test_ep_overlap_rejects_ambiguous_token_count_sync_copy_destination(self):
        gm, _refs = self._build_ep_sync_copy_schedule_gm(cpu_destination=False)

        with self.assertRaisesRegex(ValueError, "CPU _to_copy destinations"):
            _schedule_ep_overlap_regions(
                gm,
                module_pattern="layers.*.moe",
                require_all_to_all=True,
                pair_first_token_exchange=True,
            )

    def test_ep_overlap_rejects_unannotated_all_to_all_markers(self):
        gm = self._build_ep_overlap_schedule_gm()
        c10d = torch.ops._c10d_functional
        launches = [
            node
            for node in gm.graph.nodes
            if node.op == "call_function"
            and node.target == c10d.all_to_all_single.default
        ]
        for node in launches:
            custom = dict(node.meta.get("custom", {}))
            custom.pop(_EP_TOKEN_EXCHANGE, None)
            node.meta["custom"] = custom

        with self.assertRaises(ValueError, msg="did not find any chunked EP"):
            _schedule_ep_overlap_regions(
                gm,
                module_pattern="layers.*.moe",
                require_all_to_all=True,
                pair_first_token_exchange=True,
            )

    def test_ep_overlap_rejects_token_exchange_metadata_on_compute(self):
        gm = self._build_ep_overlap_schedule_gm()
        compute = next(
            node
            for node in gm.graph.nodes
            if node.target == torch.ops.aten.neg.default
            and node.meta.get("chunked_region_role") == "body"
        )
        compute.meta.setdefault("custom", {})[_EP_TOKEN_EXCHANGE] = "dispatch"

        with self.assertRaisesRegex(ValueError, "non-marker node"):
            _schedule_ep_overlap_regions(
                gm,
                module_pattern="layers.*.moe",
                require_all_to_all=True,
                pair_first_token_exchange=True,
            )

    def test_ep_overlap_rejects_mismatched_token_exchange_labels(self):
        gm = self._build_ep_overlap_schedule_gm()
        launch = next(
            node
            for node in gm.graph.nodes
            if node.op == "call_function"
            and node.target == torch.ops._c10d_functional.all_to_all_single.default
            and node.meta.get("chunk_id") == 0
            and node.meta.get("custom", {}).get(_EP_TOKEN_EXCHANGE) == "dispatch"
        )
        launch.meta["custom"][_EP_TOKEN_EXCHANGE] = "combine"

        with self.assertRaisesRegex(ValueError, "matching token-exchange labels"):
            _schedule_ep_overlap_regions(
                gm,
                module_pattern="layers.*.moe",
                require_all_to_all=True,
                pair_first_token_exchange=True,
            )

    def test_ep_overlap_rejects_wait_before_token_exchange_launch(self):
        gm = self._build_ep_overlap_schedule_gm()
        launch = next(
            node
            for node in gm.graph.nodes
            if node.op == "call_function"
            and node.target == torch.ops._c10d_functional.all_to_all_single.default
            and node.meta.get("custom", {}).get(_EP_TOKEN_EXCHANGE) == "dispatch"
        )
        wait = next(
            user
            for user in launch.users
            if user.target == torch.ops._c10d_functional.wait_tensor.default
        )
        launch.prepend(wait)

        with self.assertRaisesRegex(ValueError, "wait to appear after its launch"):
            _schedule_ep_overlap_regions(
                gm,
                module_pattern="layers.*.moe",
                require_all_to_all=True,
                pair_first_token_exchange=True,
            )

    def test_ep_overlap_uses_future_closure_prefix_before_wait(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional
        refs = {}

        # Build the graph in the opposite order from the desired backward
        # schedule. This catches accidental sorting by original graph order when
        # a ready filler phase intentionally contains chunk1 work before chunk0.
        for chunk_id in (0, 1):
            combine_pre = graph.call_function(torch.ops.aten.relu.default, args=(x,))
            combine = graph.call_function(
                c10d.all_to_all_single.default,
                args=(combine_pre, [], [], "ep"),
            )
            dispatch_prefix = graph.call_function(torch.ops.aten.neg.default, args=(x,))
            combine_wait = graph.call_function(
                c10d.wait_tensor.default, args=(combine,)
            )
            dispatch_pre = graph.call_function(
                torch.ops.aten.add.Tensor, args=(dispatch_prefix, combine_wait)
            )
            dispatch = graph.call_function(
                c10d.all_to_all_single.default,
                args=(dispatch_pre, [], [], "ep"),
            )
            dispatch_wait = graph.call_function(
                c10d.wait_tensor.default, args=(dispatch,)
            )
            tail = graph.call_function(
                torch.ops.aten.relu.default, args=(dispatch_wait,)
            )
            refs[chunk_id] = {
                "combine": combine,
                "dispatch_prefix": dispatch_prefix,
                "combine_wait": combine_wait,
                "dispatch_pre": dispatch_pre,
                "dispatch": dispatch,
                "dispatch_wait": dispatch_wait,
                "tail": tail,
            }

            for node in (combine_pre, combine, combine_wait):
                node.meta["custom"] = {_MODULE_FQN: "layers.0.moe", "EP": "combine"}
            combine.meta["custom"][_EP_TOKEN_EXCHANGE] = "combine"
            dispatch_prefix.meta["custom"] = {_MODULE_FQN: "layers.0.moe"}
            dispatch_pre.meta["custom"] = {_MODULE_FQN: "layers.0.moe"}
            for node in (dispatch, dispatch_wait):
                node.meta["custom"] = {_MODULE_FQN: "layers.0.moe", "EP": "dispatch"}
            dispatch.meta["custom"][_EP_TOKEN_EXCHANGE] = "dispatch"
            tail.meta["custom"] = {_MODULE_FQN: "layers.0.moe"}
            for node in (
                combine_pre,
                combine,
                dispatch_prefix,
                combine_wait,
                dispatch_pre,
                dispatch,
                dispatch_wait,
                tail,
            ):
                node.meta["chunk_id"] = chunk_id
                node.meta["chunked_region_fqn"] = "layers.0.moe"
                node.meta["chunked_region_role"] = "body"
                node.meta["autograd_backward"] = True

        graph.output((refs[1]["tail"], refs[0]["tail"]))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        order = self._schedule_ep_overlap_and_order(gm)

        self.assertEqual(
            refs[1]["combine_wait"].meta["custom"][_EP_TOKEN_EXCHANGE_WAIT],
            "combine",
        )
        self.assertLess(order[refs[1]["combine"]], order[refs[0]["combine"]])
        self.assertLess(order[refs[0]["combine"]], order[refs[1]["dispatch_prefix"]])
        self.assertLess(
            order[refs[1]["dispatch_prefix"]],
            order[refs[0]["dispatch_prefix"]],
        )
        self.assertLess(
            order[refs[0]["dispatch_prefix"]], order[refs[1]["combine_wait"]]
        )
        self.assertLess(order[refs[1]["combine_wait"]], order[refs[1]["dispatch"]])
        self.assertLess(order[refs[1]["dispatch"]], order[refs[0]["combine_wait"]])
        self.assertLess(order[refs[0]["combine_wait"]], order[refs[0]["dispatch"]])
        self.assertLess(order[refs[0]["dispatch"]], order[refs[1]["dispatch_wait"]])
        self.assertLess(order[refs[1]["dispatch_wait"]], order[refs[1]["tail"]])
        self.assertLess(order[refs[1]["tail"]], order[refs[0]["dispatch_wait"]])
        self.assertLess(order[refs[0]["dispatch_wait"]], order[refs[0]["tail"]])

    def test_ep_overlap_does_not_hoist_non_token_collectives_as_filler(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional
        refs = {}

        for chunk_id in (0, 1):
            dispatch_pre = graph.call_function(torch.ops.aten.relu.default, args=(x,))
            dispatch = graph.call_function(
                c10d.all_to_all_single.default,
                args=(dispatch_pre, [], [], "ep"),
            )
            dispatch_wait = graph.call_function(
                c10d.wait_tensor.default, args=(dispatch,)
            )
            all_reduce = graph.call_function(
                c10d.all_reduce.default, args=(x, "sum", "dp")
            )
            combine_pre = graph.call_function(
                torch.ops.aten.add.Tensor, args=(dispatch_wait, all_reduce)
            )
            combine = graph.call_function(
                c10d.all_to_all_single.default,
                args=(combine_pre, [], [], "ep"),
            )
            combine_wait = graph.call_function(
                c10d.wait_tensor.default, args=(combine,)
            )
            refs[chunk_id] = {
                "dispatch": dispatch,
                "dispatch_wait": dispatch_wait,
                "all_reduce": all_reduce,
                "combine": combine,
                "combine_wait": combine_wait,
            }

            for node in (dispatch_pre, dispatch, dispatch_wait):
                node.meta["custom"] = {_MODULE_FQN: "layers.0.moe", "EP": "dispatch"}
            dispatch.meta["custom"][_EP_TOKEN_EXCHANGE] = "dispatch"
            all_reduce.meta["custom"] = {_MODULE_FQN: "layers.0.moe"}
            combine_pre.meta["custom"] = {_MODULE_FQN: "layers.0.moe"}
            for node in (combine, combine_wait):
                node.meta["custom"] = {_MODULE_FQN: "layers.0.moe", "EP": "combine"}
            combine.meta["custom"][_EP_TOKEN_EXCHANGE] = "combine"

            for node in (
                dispatch_pre,
                dispatch,
                dispatch_wait,
                all_reduce,
                combine_pre,
                combine,
                combine_wait,
            ):
                node.meta["chunk_id"] = chunk_id
                node.meta["chunked_region_fqn"] = "layers.0.moe"
                node.meta["chunked_region_role"] = "body"

        graph.output((refs[0]["combine_wait"], refs[1]["combine_wait"]))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        order = self._schedule_ep_overlap_and_order(gm)

        self.assertLess(order[refs[0]["dispatch"]], order[refs[1]["dispatch"]])
        self.assertLess(order[refs[1]["dispatch"]], order[refs[0]["dispatch_wait"]])
        for chunk_id in (0, 1):
            self.assertLess(
                order[refs[chunk_id]["dispatch_wait"]],
                order[refs[chunk_id]["all_reduce"]],
            )
            self.assertLess(
                order[refs[chunk_id]["all_reduce"]],
                order[refs[chunk_id]["combine"]],
            )

    def test_ep_overlap_rejects_mismatched_all_to_all_count(self):
        gm = self._build_ep_overlap_schedule_gm(backward=True)
        c10d = torch.ops._c10d_functional
        for node in gm.graph.nodes:
            if (
                node.op == "call_function"
                and node.target == c10d.all_to_all_single.default
                and node.meta.get("chunk_id") == 0
            ):
                node.target = c10d.all_reduce.default
                node.args = (node.args[0], "sum", "ep")
                node.meta["custom"].pop(_EP_TOKEN_EXCHANGE, None)
                break
        with self.assertRaisesRegex(ValueError, "matching EP all-to-all counts"):
            _schedule_ep_overlap_regions(
                gm,
                module_pattern="layers.*.moe",
                require_all_to_all=True,
                pair_first_token_exchange=True,
            )

    def test_ep_overlap_schedules_arbitrary_matching_token_exchange_sequence(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional
        outputs = []
        launches = {}
        exchange_kinds = ("combine", "dispatch", "combine")
        for chunk_id in (1, 0):
            value = x
            body_nodes = []
            for idx, kind in enumerate(exchange_kinds):
                pre = graph.call_function(torch.ops.aten.relu.default, args=(value,))
                launch = graph.call_function(
                    c10d.all_to_all_single.default,
                    args=(pre, [], [], "ep"),
                )
                wait = graph.call_function(c10d.wait_tensor.default, args=(launch,))
                launches[(chunk_id, idx)] = launch
                for node in (pre, launch, wait):
                    node.meta["custom"] = {_MODULE_FQN: "layers.0.moe", "EP": kind}
                launch.meta["custom"][_EP_TOKEN_EXCHANGE] = kind
                value = graph.call_function(torch.ops.aten.neg.default, args=(wait,))
                value.meta["custom"] = {_MODULE_FQN: "layers.0.moe"}
                body_nodes.extend((pre, launch, wait, value))
            outputs.append(value)
            for node in body_nodes:
                node.meta["chunk_id"] = chunk_id
                node.meta["chunked_region_fqn"] = "layers.0.moe"
                node.meta["chunked_region_role"] = "body"
                node.meta["autograd_backward"] = True

        graph.output(tuple(outputs))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        order = self._schedule_ep_overlap_and_order(gm)
        for idx in range(len(exchange_kinds)):
            self.assertLess(order[launches[(1, idx)]], order[launches[(0, idx)]])
            if idx + 1 < len(exchange_kinds):
                self.assertLess(
                    order[launches[(0, idx)]], order[launches[(1, idx + 1)]]
                )

    def test_ep_overlap_reorders_non_body_setup_dependencies(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional
        outputs = []
        setup_nodes = {}
        first_waits = {}
        for chunk_id in (0, 1):
            setup = graph.call_function(torch.ops.aten.clone.default, args=(x,))
            pre = graph.call_function(torch.ops.aten.relu.default, args=(setup,))
            dispatch = graph.call_function(
                c10d.all_to_all_single.default,
                args=(pre, [], [], "ep"),
            )
            dispatch_wait = graph.call_function(
                c10d.wait_tensor.default, args=(dispatch,)
            )
            compute = graph.call_function(
                torch.ops.aten.neg.default, args=(dispatch_wait,)
            )
            combine = graph.call_function(
                c10d.all_to_all_single.default,
                args=(compute, [], [], "ep"),
            )
            combine_wait = graph.call_function(
                c10d.wait_tensor.default, args=(combine,)
            )
            outputs.append(combine_wait)
            setup_nodes[chunk_id] = setup
            first_waits[chunk_id] = dispatch_wait

            setup.meta["custom"] = {_MODULE_FQN: "layers.0.moe"}
            for node in (pre, dispatch, dispatch_wait):
                node.meta["custom"] = {_MODULE_FQN: "layers.0.moe", "EP": "dispatch"}
            dispatch.meta["custom"][_EP_TOKEN_EXCHANGE] = "dispatch"
            compute.meta["custom"] = {_MODULE_FQN: "layers.0.moe"}
            for node in (combine, combine_wait):
                node.meta["custom"] = {_MODULE_FQN: "layers.0.moe", "EP": "combine"}
            combine.meta["custom"][_EP_TOKEN_EXCHANGE] = "combine"
            for node in (pre, dispatch, dispatch_wait, compute, combine, combine_wait):
                node.meta["chunk_id"] = chunk_id
                node.meta["chunked_region_fqn"] = "layers.0.moe"
                node.meta["chunked_region_role"] = "body"

        graph.output(tuple(outputs))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        order = self._schedule_ep_overlap_and_order(gm)

        self.assertLess(order[setup_nodes[1]], order[first_waits[0]])
        self.assertLess(order[first_waits[0]], order[first_waits[1]])

    def test_ep_overlap_apply_schedule_interleaves_cross_region_boundaries(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        region0_start = graph.call_function(torch.ops.aten.relu.default, args=(x,))
        region1_start = graph.call_function(
            torch.ops.aten.neg.default, args=(region0_start,)
        )
        boundary = graph.call_function(
            torch.ops.aten.clone.default, args=(region1_start,)
        )
        region0_tail = graph.call_function(
            torch.ops.aten.relu.default, args=(boundary,)
        )
        region1_tail = graph.call_function(
            torch.ops.aten.add.Tensor, args=(region1_start, region0_tail)
        )
        graph.output(region1_tail)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        def scheduled_region(
            root_fqn: str,
            phases: tuple[tuple[torch.fx.Node, ...], ...],
        ) -> _ScheduledRegion:
            bodies = {
                idx: ChunkBody(
                    owner=ChunkOwner(root_fqn, False, idx),
                    nodes=(),
                    node_set=frozenset(),
                    live_ins=frozenset(),
                    producer="graph",
                )
                for idx in (0, 1)
            }
            return _ScheduledRegion(
                region=ChunkedRegion(root_fqn, False, bodies),
                phases=phases,
            )

        _apply_schedule(
            gm,
            [
                scheduled_region("layers.0", ((region0_start,), (region0_tail,))),
                scheduled_region("layers.1", ((region1_start,), (region1_tail,))),
            ],
        )

        order = {node: idx for idx, node in enumerate(gm.graph.nodes)}
        self.assertLess(order[region0_start], order[region0_tail])
        self.assertLess(order[region1_start], order[region1_tail])
        self.assertLess(order[region1_start], order[boundary])
        self.assertLess(order[boundary], order[region0_tail])
        self.assertLess(order[region0_tail], order[region1_tail])

    def test_ep_overlap_moves_owned_remat_with_chunk_body(self):
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        c10d = torch.ops._c10d_functional

        c0_pre = graph.call_function(torch.ops.aten.relu.default, args=(x,))
        c0_dispatch = graph.call_function(
            c10d.all_to_all_single.default, args=(c0_pre, [], [], "ep")
        )
        c0_dispatch_wait = graph.call_function(
            c10d.wait_tensor.default, args=(c0_dispatch,)
        )
        c0_compute = graph.call_function(
            torch.ops.aten.neg.default, args=(c0_dispatch_wait,)
        )
        c0_combine = graph.call_function(
            c10d.all_to_all_single.default, args=(c0_compute, [], [], "ep")
        )
        c0_combine_wait = graph.call_function(
            c10d.wait_tensor.default, args=(c0_combine,)
        )

        c1_setup = graph.call_function(torch.ops.aten.clone.default, args=(x,))
        c1_setup.meta["autograd_backward"] = True
        c1_setup.meta["recompute"] = CheckpointPolicy.PREFER_RECOMPUTE
        c1_pre = graph.call_function(torch.ops.aten.relu.default, args=(c1_setup,))
        c1_dispatch = graph.call_function(
            c10d.all_to_all_single.default, args=(c1_pre, [], [], "ep")
        )
        c1_dispatch_wait = graph.call_function(
            c10d.wait_tensor.default, args=(c1_dispatch,)
        )
        c1_compute = graph.call_function(
            torch.ops.aten.neg.default, args=(c1_dispatch_wait,)
        )
        c1_combine = graph.call_function(
            c10d.all_to_all_single.default, args=(c1_compute, [], [], "ep")
        )
        c1_combine_wait = graph.call_function(
            c10d.wait_tensor.default, args=(c1_combine,)
        )
        graph.output((c0_combine_wait, c1_combine_wait))

        for chunk_id, nodes in {
            0: (
                c0_pre,
                c0_dispatch,
                c0_dispatch_wait,
                c0_compute,
                c0_combine,
                c0_combine_wait,
            ),
            1: (
                c1_setup,
                c1_pre,
                c1_dispatch,
                c1_dispatch_wait,
                c1_compute,
                c1_combine,
                c1_combine_wait,
            ),
        }.items():
            for node in nodes:
                node.meta["chunk_id"] = chunk_id
                node.meta["chunked_region_fqn"] = "layers.0.moe"
                node.meta["chunked_region_role"] = "body"
                node.meta["autograd_backward"] = True
            ep_nodes = nodes[1:] if chunk_id == 1 else nodes
            for node in ep_nodes[:3]:
                node.meta["custom"] = {_MODULE_FQN: "layers.0.moe", "EP": "combine"}
            ep_nodes[1].meta["custom"][_EP_TOKEN_EXCHANGE] = "combine"
            ep_nodes[3].meta["custom"] = {_MODULE_FQN: "layers.0.moe"}
            for node in ep_nodes[4:]:
                node.meta["custom"] = {_MODULE_FQN: "layers.0.moe", "EP": "dispatch"}
            ep_nodes[4].meta["custom"][_EP_TOKEN_EXCHANGE] = "dispatch"

        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        order = self._schedule_ep_overlap_and_order(gm)
        self.assertLess(order[c1_setup], order[c1_pre])
        self.assertLess(order[c1_dispatch], order[c0_dispatch_wait])

    def test_eager_chunking_traces_overlap_metadata(self):
        class Block(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.linear = torch.nn.Linear(3, 3)

            def forward(self, x):
                return torch.relu(self.linear(x))

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.layers = torch.nn.ModuleList([Block()])

            def forward(self, x):
                return self.layers[0](x)

        model = Model()
        annotate_module_fqns(model)
        maybe_apply_ep_overlap_eager_chunking(
            model,
            GraphTrainerCompileConfig(
                ep_overlap=EpOverlapConfig(
                    enabled=True,
                    module_fqn="layers.*",
                ),
            ),
        )

        def step(inputs):
            y = model(inputs)
            loss = y.sum()
            params = [p for p in model.parameters() if p.requires_grad]
            return [loss] + list(torch.autograd.grad(loss, params))

        traced = minimal_fx_tracer(step, module=model)(torch.randn(4, 3))
        gm = populate_eager_chunk_metadata_pass(traced.gm)

        body_chunks = {
            node.meta.get("chunk_id")
            for node in gm.graph.nodes
            if node.meta.get("chunked_region_role") == "body"
        }
        backward_body_chunks = {
            node.meta.get("chunk_id")
            for node in gm.graph.nodes
            if node.meta.get("chunked_region_role") == "body"
            and node.meta.get("autograd_backward", False)
        }
        boundary_roles = {
            node.meta.get("chunked_region_role")
            for node in gm.graph.nodes
            if node.meta.get("chunked_region_fqn") == "layers.0"
        }
        self.assertEqual(body_chunks, {0, 1})
        self.assertEqual(backward_body_chunks, {0, 1})
        self.assertIn("split_boundary", boundary_roles)
        self.assertIn("materialization", boundary_roles)


class TestRemoveIdentityViewPass(TestCase):
    """Unit tests for the remove_identity_view_pass graph pass."""

    _VIEW_TARGETS = [
        torch.ops.aten.view.default,
        torch.ops.aten.reshape.default,
        torch.ops.aten._unsafe_view.default,
    ]

    def _build_view_gm(self, op_targets, *, shapes=None):
        """Build a GraphModule with a chain of call_function nodes.

        Each op in ``op_targets`` becomes a call_function node chained
        sequentially: placeholder(x) -> op1(x, shape) -> op2(..., shape) -> output.

        If ``shapes`` is provided it must have the same length as
        ``op_targets`` and supplies the shape argument for each view-like
        node.  Non-view nodes ignore the corresponding entry.
        """
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        last = x
        for i, target in enumerate(op_targets):
            if target in (
                torch.ops.aten.view.default,
                torch.ops.aten.reshape.default,
                torch.ops.aten._unsafe_view.default,
            ):
                shape = shapes[i] if shapes else [4, 4]
                last = graph.call_function(target, args=(last, shape))
            else:
                last = graph.call_function(target, args=(last,))
        graph.output(last)
        return torch.fx.GraphModule(torch.nn.Module(), graph)

    def _attach_fake_meta(self, gm, input_shape):
        """Attach fake tensor metadata to all nodes based on op semantics."""
        fake_mode = torch._subclasses.FakeTensorMode()
        with fake_mode:
            fake_input = torch.randn(input_shape)
        for node in gm.graph.nodes:
            if node.op == "placeholder":
                node.meta["val"] = fake_input
            elif node.op == "call_function":
                if node.target in (
                    torch.ops.aten.view.default,
                    torch.ops.aten.reshape.default,
                    torch.ops.aten._unsafe_view.default,
                ):
                    target_shape = node.args[1]
                    with fake_mode:
                        node.meta["val"] = torch.randn(target_shape)
                else:
                    # For unary ops like relu/neg, output shape == input shape.
                    node.meta["val"] = node.args[0].meta.get("val")

    def _count_view_nodes(self, gm):
        """Count view/reshape/_unsafe_view call_function nodes."""
        targets = {
            torch.ops.aten.view.default,
            torch.ops.aten.reshape.default,
            torch.ops.aten._unsafe_view.default,
        }
        return sum(
            1 for n in gm.graph.nodes if n.op == "call_function" and n.target in targets
        )

    def _count_call_function_nodes(self, gm):
        """Count all call_function nodes."""
        return sum(1 for n in gm.graph.nodes if n.op == "call_function")

    def test_identity_view_removed(self):
        """Identity view (same shape in and out) is removed for each op type."""
        for target in self._VIEW_TARGETS:
            with self.subTest(target=target):
                gm = self._build_view_gm(
                    [torch.ops.aten.relu.default, target, torch.ops.aten.neg.default],
                    shapes=[None, [4, 4], None],
                )
                self._attach_fake_meta(gm, (4, 4))
                self.assertEqual(self._count_view_nodes(gm), 1)

                result = remove_identity_view_pass(gm)

                self.assertEqual(self._count_view_nodes(result), 0)
                self.assertEqual(self._count_call_function_nodes(result), 2)

    def test_non_identity_view_preserved(self):
        """Non-identity view (shape changes) is kept."""
        gm = self._build_view_gm(
            [torch.ops.aten.view.default],
            shapes=[[2, 8]],
        )
        self._attach_fake_meta(gm, (4, 4))
        self.assertEqual(self._count_view_nodes(gm), 1)

        remove_identity_view_pass(gm)

        self.assertEqual(self._count_view_nodes(gm), 1)

    def test_view_without_metadata_skipped(self):
        """Nodes without tensor metadata are skipped safely."""
        gm = self._build_view_gm(
            [torch.ops.aten.view.default],
            shapes=[[4, 4]],
        )
        # Do NOT attach fake meta — nodes have no "val" in meta.

        # Should not raise and should not modify the graph.
        remove_identity_view_pass(gm)

        self.assertEqual(self._count_view_nodes(gm), 1)

    def test_numerics_preserved(self):
        """Forward outputs are preserved after removing identity views."""
        gm = self._build_view_gm(
            [
                torch.ops.aten.relu.default,
                torch.ops.aten.view.default,
                torch.ops.aten.neg.default,
            ],
            shapes=[None, [4, 4], None],
        )
        self._attach_fake_meta(gm, (4, 4))

        x = torch.randn(4, 4)
        expected = torch.neg(torch.relu(x).view(4, 4))

        remove_identity_view_pass(gm)
        actual = gm(x)

        self.assertEqual(actual, expected)

    def test_view_with_multiple_users(self):
        """Identity view with multiple users: all uses are replaced."""
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        view = graph.call_function(torch.ops.aten.view.default, args=(x, [4, 4]))
        relu = graph.call_function(torch.ops.aten.relu.default, args=(view,))
        neg = graph.call_function(torch.ops.aten.neg.default, args=(view,))
        graph.output((relu, neg))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        # Attach metadata
        fake_mode = torch._subclasses.FakeTensorMode()
        with fake_mode:
            fake_input = torch.randn(4, 4)
        for node in gm.graph.nodes:
            if node.op == "placeholder":
                node.meta["val"] = fake_input
            elif node.target is torch.ops.aten.view.default:
                node.meta["val"] = fake_input  # same shape
            elif node.op == "call_function":
                node.meta["val"] = fake_input

        self.assertEqual(self._count_view_nodes(gm), 1)

        remove_identity_view_pass(gm)

        self.assertEqual(self._count_view_nodes(gm), 0)

        # Both relu and neg should now consume the placeholder directly
        for node in gm.graph.nodes:
            if node.op == "call_function" and node.target in (
                torch.ops.aten.relu.default,
                torch.ops.aten.neg.default,
            ):
                self.assertEqual(node.args[0].op, "placeholder")

        # Verify numerics
        x = torch.randn(4, 4)
        relu_out, neg_out = gm(x)
        self.assertEqual(relu_out, torch.relu(x))
        self.assertEqual(neg_out, torch.neg(x))

    def test_chain_of_identity_views(self):
        """Chain of identity views (view -> view -> view) is fully removed."""
        gm = self._build_view_gm(
            [
                torch.ops.aten.relu.default,
                torch.ops.aten.view.default,
                torch.ops.aten.reshape.default,
                torch.ops.aten._unsafe_view.default,
                torch.ops.aten.neg.default,
            ],
            shapes=[None, [4, 4], [4, 4], [4, 4], None],
        )
        self._attach_fake_meta(gm, (4, 4))
        self.assertEqual(self._count_view_nodes(gm), 3)

        remove_identity_view_pass(gm)

        self.assertEqual(self._count_view_nodes(gm), 0)
        self.assertEqual(self._count_call_function_nodes(gm), 2)

        # Verify numerics
        x = torch.randn(4, 4)
        expected = torch.neg(torch.relu(x))
        self.assertEqual(gm(x), expected)

    def test_graph_without_views_unchanged(self):
        """Graphs without view nodes are returned unchanged."""
        gm = self._build_view_gm(
            [torch.ops.aten.relu.default, torch.ops.aten.neg.default],
            shapes=[None, None],
        )
        self._attach_fake_meta(gm, (4, 4))
        num_nodes_before = len(list(gm.graph.nodes))

        result = remove_identity_view_pass(gm)

        self.assertIs(result, gm)
        self.assertEqual(len(list(result.graph.nodes)), num_nodes_before)


class TestRemoveB2BTransposePass(TestCase):
    """Unit tests for the remove_b2b_transpose_pass graph pass."""

    def _count_t_nodes(self, gm):
        """Count aten.t.default call_function nodes."""
        return sum(
            1
            for n in gm.graph.nodes
            if n.op == "call_function" and n.target is torch.ops.aten.t.default
        )

    def test_b2b_transpose_pair_removed(self):
        """``t(t(x))`` collapses: both transpose nodes are removed and the
        consumer reads the original tensor directly."""
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        t1 = graph.call_function(torch.ops.aten.t.default, args=(x,))
        t2 = graph.call_function(torch.ops.aten.t.default, args=(t1,))
        relu = graph.call_function(torch.ops.aten.relu.default, args=(t2,))
        graph.output(relu)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        self.assertEqual(self._count_t_nodes(gm), 2)

        remove_b2b_transpose_pass(gm)

        self.assertEqual(self._count_t_nodes(gm), 0)
        # relu now consumes the placeholder directly.
        relu_node = next(
            n for n in gm.graph.nodes if n.target is torch.ops.aten.relu.default
        )
        self.assertEqual(relu_node.args[0].op, "placeholder")

        # Numerics preserved: relu(t(t(x))) == relu(x).
        x = torch.randn(3, 4)
        self.assertEqual(gm(x), torch.relu(x))

    def test_single_transpose_preserved(self):
        """A lone transpose is not a back-to-back pair and must be kept."""
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        t = graph.call_function(torch.ops.aten.t.default, args=(x,))
        graph.output(t)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        remove_b2b_transpose_pass(gm)

        self.assertEqual(self._count_t_nodes(gm), 1)
        x = torch.randn(3, 4)
        self.assertEqual(gm(x), x.t())

    def test_inner_transpose_with_other_user_kept(self):
        """When the inner transpose feeds another consumer, only the outer
        transpose is removed; the inner one stays for its other user."""
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        t1 = graph.call_function(torch.ops.aten.t.default, args=(x,))
        t2 = graph.call_function(torch.ops.aten.t.default, args=(t1,))
        # t1 also feeds a relu, so it cannot be erased.
        relu = graph.call_function(torch.ops.aten.relu.default, args=(t1,))
        graph.output((t2, relu))
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        self.assertEqual(self._count_t_nodes(gm), 2)

        remove_b2b_transpose_pass(gm)

        # Outer transpose removed, inner one kept (still used by relu).
        self.assertEqual(self._count_t_nodes(gm), 1)

        x = torch.randn(3, 4)
        out_t2, out_relu = gm(x)
        self.assertEqual(out_t2, x)  # t(t(x)) == x
        self.assertEqual(out_relu, torch.relu(x.t()))

    def test_chain_of_transposes(self):
        """An odd-length chain ``t(t(t(x)))`` collapses to a single ``t(x)``."""
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        t1 = graph.call_function(torch.ops.aten.t.default, args=(x,))
        t2 = graph.call_function(torch.ops.aten.t.default, args=(t1,))
        t3 = graph.call_function(torch.ops.aten.t.default, args=(t2,))
        graph.output(t3)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        self.assertEqual(self._count_t_nodes(gm), 3)

        remove_b2b_transpose_pass(gm)

        self.assertEqual(self._count_t_nodes(gm), 1)
        x = torch.randn(3, 4)
        self.assertEqual(gm(x), x.t())

    def test_graph_without_transpose_unchanged(self):
        """Graphs without transpose nodes are returned unchanged."""
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        relu = graph.call_function(torch.ops.aten.relu.default, args=(x,))
        neg = graph.call_function(torch.ops.aten.neg.default, args=(relu,))
        graph.output(neg)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        num_nodes_before = len(list(gm.graph.nodes))

        result = remove_b2b_transpose_pass(gm)

        self.assertIs(result, gm)
        self.assertEqual(len(list(result.graph.nodes)), num_nodes_before)


class TestRemoveIdentitySlicePass(TestCase):
    """Unit tests for the remove_identity_slice_pass graph pass."""

    def _build_slice_gm(self, input_shape, dim, start, end, step=1):
        """Build a GraphModule with a single aten.slice.Tensor node.

        Creates: placeholder(x) -> slice(x, dim, start, end, step) -> output.
        The placeholder is annotated with fake tensor metadata of the given shape.
        """
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        sliced = graph.call_function(
            torch.ops.aten.slice.Tensor, args=(x, dim, start, end, step)
        )
        graph.output(sliced)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        # Annotate placeholder with fake tensor metadata
        from torch._subclasses.fake_tensor import FakeTensorMode

        with FakeTensorMode() as fake_mode:
            fake_val = fake_mode.from_tensor(torch.empty(*input_shape))
        for node in gm.graph.nodes:
            if node.op == "placeholder":
                node.meta["val"] = fake_val
        return gm

    def _count_slice_nodes(self, gm):
        """Count aten.slice.Tensor nodes in the graph."""
        return sum(
            1
            for n in gm.graph.nodes
            if n.op == "call_function" and n.target is torch.ops.aten.slice.Tensor
        )

    def test_full_dim_slice_is_removed(self):
        """A slice selecting the full dimension (start=0, end>=dim_size, step=1)
        should be removed."""
        gm = self._build_slice_gm(input_shape=(8, 16), dim=0, start=0, end=8, step=1)
        self.assertEqual(self._count_slice_nodes(gm), 1)

        remove_identity_slice_pass(gm)
        self.assertEqual(self._count_slice_nodes(gm), 0)

    def test_full_dim_slice_large_end_is_removed(self):
        """A slice with end > dim_size should also be removed (identity)."""
        import sys

        gm = self._build_slice_gm(
            input_shape=(8, 16), dim=0, start=0, end=sys.maxsize, step=1
        )
        self.assertEqual(self._count_slice_nodes(gm), 1)

        remove_identity_slice_pass(gm)
        self.assertEqual(self._count_slice_nodes(gm), 0)

    def test_partial_slice_start_preserved(self):
        """A slice with start > 0 is not an identity and should be preserved."""
        gm = self._build_slice_gm(input_shape=(8, 16), dim=0, start=2, end=8, step=1)
        remove_identity_slice_pass(gm)
        self.assertEqual(self._count_slice_nodes(gm), 1)

    def test_partial_slice_end_preserved(self):
        """A slice with end < dim_size is not an identity and should be preserved."""
        gm = self._build_slice_gm(input_shape=(8, 16), dim=0, start=0, end=4, step=1)
        remove_identity_slice_pass(gm)
        self.assertEqual(self._count_slice_nodes(gm), 1)

    def test_partial_slice_step_preserved(self):
        """A slice with step > 1 is not an identity and should be preserved."""
        gm = self._build_slice_gm(input_shape=(8, 16), dim=0, start=0, end=8, step=2)
        remove_identity_slice_pass(gm)
        self.assertEqual(self._count_slice_nodes(gm), 1)

    def test_no_metadata_skipped(self):
        """Slice nodes without fake tensor metadata should be skipped safely."""
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        sliced = graph.call_function(
            torch.ops.aten.slice.Tensor, args=(x, 0, 0, 100, 1)
        )
        graph.output(sliced)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        # No metadata set -- pass should not crash
        remove_identity_slice_pass(gm)
        self.assertEqual(self._count_slice_nodes(gm), 1)

    def test_multi_dim_slice(self):
        """Identity slice on a non-zero dimension should be removed."""
        gm = self._build_slice_gm(
            input_shape=(8, 16, 32), dim=2, start=0, end=32, step=1
        )
        remove_identity_slice_pass(gm)
        self.assertEqual(self._count_slice_nodes(gm), 0)

    def test_numerics_preserved(self):
        """The pass should not change the numerical output of the graph."""
        gm = self._build_slice_gm(input_shape=(4, 8), dim=0, start=0, end=4, step=1)

        # Run before the pass
        x = torch.randn(4, 8)
        out_before = gm(x)

        remove_identity_slice_pass(gm)

        out_after = gm(x)
        self.assertTrue(torch.equal(out_before, out_after))

    def test_chained_identity_slices(self):
        """Multiple chained identity slices should all be removed."""
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        s1 = graph.call_function(torch.ops.aten.slice.Tensor, args=(x, 0, 0, 8, 1))
        s2 = graph.call_function(torch.ops.aten.slice.Tensor, args=(s1, 1, 0, 16, 1))
        graph.output(s2)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        from torch._subclasses.fake_tensor import FakeTensorMode

        with FakeTensorMode() as fake_mode:
            fake_val = fake_mode.from_tensor(torch.empty(8, 16))
        for node in gm.graph.nodes:
            if node.op == "placeholder":
                node.meta["val"] = fake_val

        # Also annotate s1 with metadata so s2 can check its input's shape
        for node in gm.graph.nodes:
            if (
                node.op == "call_function"
                and node.target is torch.ops.aten.slice.Tensor
            ):
                node.meta["val"] = fake_val
                break  # Only need the first slice node (s1)

        self.assertEqual(self._count_slice_nodes(gm), 2)
        remove_identity_slice_pass(gm)
        self.assertEqual(self._count_slice_nodes(gm), 0)

    def _build_dynamic_slice_gm(
        self,
        *,
        dynamic_arg: str,
    ) -> torch.fx.GraphModule:
        """Build a slice graph where one of start/end/step is a Node."""
        from torch._subclasses.fake_tensor import FakeTensorMode

        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        # sym_size is just a convenient way to produce a Node — its concrete
        # value is irrelevant, what matters is that the arg is an FX Node.
        sym_node = graph.call_function(torch.ops.aten.sym_size.int, args=(x, 0))
        start = sym_node if dynamic_arg == "start" else 0
        end = sym_node if dynamic_arg == "end" else sys.maxsize
        step = sym_node if dynamic_arg == "step" else 1
        sliced = graph.call_function(
            torch.ops.aten.slice.Tensor, args=(x, 0, start, end, step)
        )
        graph.output(sliced)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        with FakeTensorMode() as fake_mode:
            fake_val = fake_mode.from_tensor(torch.empty(8, 16))
        for node in gm.graph.nodes:
            if node.op == "placeholder":
                node.meta["val"] = fake_val
        return gm

    def test_dynamic_end_skipped(self):
        """Slices whose ``end`` is an FX Node (dynamic shape at runtime) must
        be left alone — we can't prove identity at pass time."""
        gm = self._build_dynamic_slice_gm(dynamic_arg="end")
        # Must not raise and must not remove the slice (can't prove identity).
        remove_identity_slice_pass(gm)
        self.assertEqual(self._count_slice_nodes(gm), 1)

    def test_dynamic_start_skipped(self):
        """Slices whose ``start`` is an FX Node must be left alone."""
        gm = self._build_dynamic_slice_gm(dynamic_arg="start")
        remove_identity_slice_pass(gm)
        self.assertEqual(self._count_slice_nodes(gm), 1)

    def test_dynamic_step_skipped(self):
        """Slices whose ``step`` is an FX Node must be left alone."""
        gm = self._build_dynamic_slice_gm(dynamic_arg="step")
        remove_identity_slice_pass(gm)
        self.assertEqual(self._count_slice_nodes(gm), 1)


class TestAnnotateModuleFqns(TestCase):
    """Unit tests for annotate_module_fqns and insert_kernel_annotations_pass."""

    def _trace_and_get_fqns(self, model, *args):
        """Trace fwd+bwd via minimal_fx_tracer and return module_fqn annotations."""

        def fwd_step(*inputs):
            pred = model(inputs[0])
            loss = pred.sum()
            params = [p for p in model.parameters() if p.requires_grad]
            grads = torch.autograd.grad(loss, params)
            return [loss] + list(grads)

        traced = minimal_fx_tracer(fwd_step, module=model)(*args)
        fqns = set()
        for node in traced.gm.graph.nodes:
            fqn = (node.meta.get("custom") or {}).get(_MODULE_FQN)
            if fqn:
                fqns.add(fqn)
        return fqns

    def test_annotate_transformer_like_model(self):
        """Module FQNs survive minimal_fx_tracer for a transformer-like model
        with distinct submodule classes (norm, attention, ffn)."""

        class Norm(torch.nn.Module):
            def __init__(self, dim):
                super().__init__()
                self.norm = torch.nn.LayerNorm(dim)

            def forward(self, x):
                return self.norm(x)

        class Attention(torch.nn.Module):
            def __init__(self, dim):
                super().__init__()
                self.wq = torch.nn.Linear(dim, dim, bias=False)
                self.wo = torch.nn.Linear(dim, dim, bias=False)

            def forward(self, x):
                return self.wo(torch.relu(self.wq(x)))

        class FFN(torch.nn.Module):
            def __init__(self, dim):
                super().__init__()
                self.w1 = torch.nn.Linear(dim, dim * 2, bias=False)
                self.w2 = torch.nn.Linear(dim * 2, dim, bias=False)

            def forward(self, x):
                return self.w2(torch.relu(self.w1(x)))

        class TransformerBlock(torch.nn.Module):
            def __init__(self, dim):
                super().__init__()
                self.attention_norm = Norm(dim)
                self.attention = Attention(dim)
                self.ffn_norm = Norm(dim)
                self.feed_forward = FFN(dim)

            def forward(self, x):
                h = x + self.attention(self.attention_norm(x))
                return h + self.feed_forward(self.ffn_norm(h))

        class Model(torch.nn.Module):
            def __init__(self, dim):
                super().__init__()
                self.layer = TransformerBlock(dim)

            def forward(self, x):
                return self.layer(x)

        dim = 16
        model = Model(dim)
        annotate_module_fqns(model)
        fqns = self._trace_and_get_fqns(model, torch.randn(4, dim))

        # Verify key module paths are present.  The Norm wrapper has no
        # ops of its own, so its inner LayerNorm gets the deepest path.
        self.assertIn("layer.attention_norm.norm", fqns)
        self.assertIn("layer.attention", fqns)
        self.assertIn("layer.attention.wq", fqns)
        self.assertIn("layer.attention.wo", fqns)
        self.assertIn("layer.ffn_norm.norm", fqns)
        self.assertIn("layer.feed_forward", fqns)
        self.assertIn("layer.feed_forward.w1", fqns)
        self.assertIn("layer.feed_forward.w2", fqns)

    def test_same_class_instances_get_distinct_fqns(self):
        """Two parameterless instances of the same class get distinct fqns.

        Uses minimal_fx_tracer directly (not trace_train_step) because
        parameterless models cannot produce gradients via autograd.grad.
        """

        class Block(torch.nn.Module):
            def forward(self, x):
                return x + 1

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.a = Block()
                self.b = Block()

            def forward(self, x):
                return self.b(self.a(x))

        model = Model()
        annotate_module_fqns(model)

        def fwd_only(state, x):
            return model(x)

        traced = minimal_fx_tracer(fwd_only)({}, torch.randn(4))
        fqns = set()
        for node in traced.gm.graph.nodes:
            fqn = (node.meta.get("custom") or {}).get(_MODULE_FQN)
            if fqn:
                fqns.add(fqn)

        self.assertIn("a", fqns)
        self.assertIn("b", fqns)

    def test_same_class_parameterless_works_with_make_fx(self):
        """Same-class parameterless instances get distinct fqns with plain make_fx."""

        class Block(torch.nn.Module):
            def forward(self, x):
                return x + 1

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.a = Block()
                self.b = Block()

            def forward(self, x):
                return self.b(self.a(x))

        model = Model()
        annotate_module_fqns(model)

        with preserve_node_meta():
            gm = make_fx(model)(torch.randn(4))

        fqns = set()
        for node in gm.graph.nodes:
            fqn = (node.meta.get("custom") or {}).get(_MODULE_FQN)
            if fqn:
                fqns.add(fqn)

        self.assertIn("a", fqns)
        self.assertIn("b", fqns)

    def test_same_class_instances_with_params_get_distinct_fqns(self):
        """Two instances of the same class with parameters get distinct fqns."""

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.a = torch.nn.Linear(4, 4)
                self.b = torch.nn.Linear(4, 4)

            def forward(self, x):
                return self.b(self.a(x))

        model = Model()
        annotate_module_fqns(model)
        fqns = self._trace_and_get_fqns(model, torch.randn(2, 4))

        self.assertIn("a", fqns)
        self.assertIn("b", fqns)

    def _build_annotated_gm(self):
        """Build a GraphModule with module_fqn annotations on its nodes."""
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        n1 = graph.call_function(torch.relu, (x,))
        n1.meta["custom"] = {_MODULE_FQN: "attn"}
        n2 = graph.call_function(torch.sigmoid, (n1,))
        n2.meta["custom"] = {_MODULE_FQN: "attn"}
        n3 = graph.call_function(torch.tanh, (n2,))
        n3.meta["custom"] = {_MODULE_FQN: "ffn"}
        graph.output(n3)
        return torch.fx.GraphModule(torch.nn.Module(), graph)

    def test_insert_kernel_annotations_pass_inserts_calls(self):
        """When tools ID is available, the pass inserts enter/exit calls."""
        if _is_tools_id_unavailable():
            self.skipTest("cudaGraphNodeGetToolsId not available")

        gm = self._build_annotated_gm()
        num_before = sum(1 for n in gm.graph.nodes if n.op == "call_function")

        insert_kernel_annotations_pass(gm)

        num_after = sum(1 for n in gm.graph.nodes if n.op == "call_function")
        # 2 scopes (attn, ffn) = 2 enters + 2 exits = 4 new nodes
        self.assertEqual(num_after - num_before, 4)

    def test_insert_kernel_annotations_pass_noop_when_unavailable(self):
        """When tools ID is unavailable, the pass leaves the graph unchanged."""
        gm = self._build_annotated_gm()
        num_before = len(list(gm.graph.nodes))

        with patch(
            "torch.cuda._graph_annotations._is_tools_id_unavailable",
            return_value=True,
        ):
            insert_kernel_annotations_pass(gm)

        num_after = len(list(gm.graph.nodes))
        self.assertEqual(num_before, num_after)

    def test_insert_kernel_annotations_pass_noop_without_metadata(self):
        """The pass should not insert anything when no custom metadata exists."""
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        n1 = graph.call_function(torch.relu, (x,))
        graph.output(n1)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        num_before = len(list(gm.graph.nodes))
        insert_kernel_annotations_pass(gm)
        num_after = len(list(gm.graph.nodes))

        self.assertEqual(num_before, num_after)


class TestNormalizeViewOpsAsReshape(TestCase):
    def test_replaces_view_and_unsafe_view(self):
        aten = torch.ops.aten
        g = torch.fx.Graph()
        x = g.placeholder("x")
        g.output(
            g.call_function(
                aten._unsafe_view.default,
                args=(
                    g.call_function(aten.view.default, args=(x, [4, 4])),
                    [2, 8],
                ),
            )
        )
        gm = torch.fx.GraphModule(torch.nn.Module(), g)
        normalize_view_ops_as_reshape(gm)
        for n in gm.graph.nodes:
            self.assertNotIn(n.target, {aten.view.default, aten._unsafe_view.default})


class TestCanonicalizeGraphPass(TestCase):
    """Unit tests for the combined canonicalize_graph_pass entry."""

    def test_runs_all_subpasses(self):
        """A single call drops detach + back-to-back transpose nodes and
        normalizes the surviving view op to reshape."""
        aten = torch.ops.aten
        graph = torch.fx.Graph()
        x = graph.placeholder("x")
        d = graph.call_function(aten.detach.default, args=(x,))
        t1 = graph.call_function(aten.t.default, args=(d,))
        t2 = graph.call_function(aten.t.default, args=(t1,))
        # Non-identity view (shape changes), left without fake meta so the
        # identity-view removal skips it and only normalization applies.
        v = graph.call_function(aten.view.default, args=(t2, [2, 8]))
        graph.output(v)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        canonicalize_graph_pass(gm)

        targets = [n.target for n in gm.graph.nodes if n.op == "call_function"]
        self.assertNotIn(aten.detach.default, targets)
        self.assertNotIn(aten.t.default, targets)
        self.assertNotIn(aten.view.default, targets)
        self.assertIn(aten.reshape.default, targets)

        # Numerics preserved: reshape(t(t(detach(x))), [2, 8]) == x.reshape(2, 8).
        x = torch.randn(4, 4)
        self.assertEqual(gm(x), x.reshape(2, 8))


class TestStandaloneInductorCompilationPass(TestCase):
    def test_prunes_unused_placeholders_and_restores_outer_signature(self):
        from torchtitan.experiments.graph_trainer.inductor_passes import (
            standalone_inductor_compilation_pass,
        )

        graph = torch.fx.Graph()
        x_node = graph.placeholder("x")
        y_node = graph.placeholder("y")
        unused_node = graph.placeholder("unused")
        x_node.meta["marker"] = "x"
        unused_node.meta["marker"] = "unused"
        result = graph.call_function(torch.ops.aten.add.Tensor, args=(x_node, y_node))
        graph.output(result)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)

        x = torch.ones(4)
        y = torch.full((4,), 2.0)
        unused = torch.full((4,), 3.0)
        seen = {}

        def fake_standalone_compile(compile_gm, compile_inputs, **kwargs):
            import torch._inductor.config as ic

            seen["placeholder_targets"] = [
                node.target for node in compile_gm.graph.find_nodes(op="placeholder")
            ]
            seen["placeholder_meta"] = [
                node.meta for node in compile_gm.graph.find_nodes(op="placeholder")
            ]
            seen["compile_inputs"] = compile_inputs
            seen["kwargs"] = kwargs
            seen["reorder_for_peak_memory"] = ic.reorder_for_peak_memory

            def compiled_fn(a, b):
                return a - b

            seen["compiled_fn"] = compiled_fn
            return compiled_fn

        with patch(
            "torch._inductor.standalone_compile", side_effect=fake_standalone_compile
        ):
            wrapped = standalone_inductor_compilation_pass(
                gm,
                (x, y, unused),
                inductor_configs={"reorder_for_peak_memory": False},
            )

        self.assertEqual(seen["placeholder_targets"], ["x", "y"])
        self.assertEqual(seen["placeholder_meta"][0]["marker"], "x")
        self.assertEqual(len(seen["compile_inputs"]), 2)
        self.assertIs(seen["compile_inputs"][0], x)
        self.assertIs(seen["compile_inputs"][1], y)
        self.assertEqual(
            seen["kwargs"],
            {
                "dynamic_shapes": "from_tracing_context",
                "aot": True,
                "donate_graph_module": False,
            },
        )
        self.assertFalse(seen["reorder_for_peak_memory"])

        wrapped_placeholders = [
            node.target for node in wrapped.graph.find_nodes(op="placeholder")
        ]
        self.assertEqual(wrapped_placeholders, ["x", "y", "unused"])
        wrapped_call = next(
            node for node in wrapped.graph.nodes if node.op == "call_function"
        )
        self.assertIs(wrapped_call.target.__wrapped__, seen["compiled_fn"])

        torch.testing.assert_close(wrapped(x, y, unused), x - y)


class TestAsyncTensorParallelPass(FSDPTest):
    """Verify async_tensor_parallel_pass produces fused ops."""

    @property
    def world_size(self):
        return 2

    def test_ag_mm_becomes_fused_op(self):
        from torch.distributed._symmetric_memory import _test_mode

        from torchtitan.experiments.graph_trainer.passes import (
            async_tensor_parallel_pass,
        )

        pg = torch.distributed.distributed_c10d._get_default_group().group_name
        aten, c10d = torch.ops.aten, torch.ops._c10d_functional

        # shard[2048,4096] -> all_gather -> wait -> mm(w[4096,1024])
        g = torch.fx.Graph()
        s, w = g.placeholder("shard"), g.placeholder("weight")
        ag = g.call_function(c10d.all_gather_into_tensor.default, args=(s, 2, pg))
        wait = g.call_function(c10d.wait_tensor.default, args=(ag,))
        g.output(g.call_function(aten.mm.default, args=(wait, w)))

        # Shapes: shard, weight, ag, wait, mm
        shapes = [(2048, 4096), (4096, 1024), (4096, 4096), (4096, 4096), (4096, 1024)]
        with torch._subclasses.FakeTensorMode():
            for node, shape in zip(g.nodes, shapes):
                node.meta["val"] = torch.randn(shape)

        gm = torch.fx.GraphModule(torch.nn.Module(), g)
        with _test_mode({pg}):
            async_tensor_parallel_pass(gm, ())

        fused = torch.ops.symm_mem.fused_all_gather_matmul.default
        self.assertTrue(any(n.target == fused for n in gm.graph.nodes))

    def test_mm_rs_becomes_fused_op(self):
        from torch.distributed._symmetric_memory import _test_mode

        from torchtitan.experiments.graph_trainer.passes import (
            async_tensor_parallel_pass,
        )

        pg = torch.distributed.distributed_c10d._get_default_group().group_name
        aten, c10d = torch.ops.aten, torch.ops._c10d_functional

        # mm(input[4096,4096], w[4096,1024]) -> reduce_scatter -> wait
        g = torch.fx.Graph()
        x, w = g.placeholder("x"), g.placeholder("w")
        mm = g.call_function(aten.mm.default, args=(x, w))
        rs = g.call_function(
            c10d.reduce_scatter_tensor.default,
            args=(mm, "sum", 2, pg),
        )
        g.output(g.call_function(c10d.wait_tensor.default, args=(rs,)))

        # Shapes: x, w, mm, rs, wait
        shapes = [
            (4096, 4096),
            (4096, 1024),
            (4096, 1024),
            (2048, 1024),
            (2048, 1024),
        ]
        with torch._subclasses.FakeTensorMode():
            for node, shape in zip(g.nodes, shapes):
                node.meta["val"] = torch.randn(shape)

        gm = torch.fx.GraphModule(torch.nn.Module(), g)
        with _test_mode({pg}):
            async_tensor_parallel_pass(gm, ())

        fused = torch.ops.symm_mem.fused_matmul_reduce_scatter.default
        self.assertTrue(any(n.target == fused for n in gm.graph.nodes))


class TestSelectiveActivationRematPass(TestCase):
    """Unit tests for ``selective_activation_remat_pass``."""

    def test_topological_insertion_order(self):
        """
        When multiple independent ``must_recompute`` deps share a downstream
        consumer, duplicates must be inserted in graph (topological) order so
        each dup's args reference upstream dups rather than the originals.
        Without that ordering (e.g. naive DFS or unordered set iteration), a
        downstream dup created before its upstream dup would fall back to the
        original ``must_recompute`` node, defeating recompute.

            a = clone(inp1)        # must_recompute
            b = clone(inp2)        # must_recompute
            d = clone(inp3)        # must_recompute
            c = a + b              # must_recompute
            e = c + d              # must_recompute
            bwd = e + e            # autograd_backward
        """
        from torchtitan.experiments.graph_trainer.selective_activation_remat import (
            selective_activation_remat_pass,
        )

        graph = torch.fx.Graph()
        inp1 = graph.placeholder("inp1")
        inp2 = graph.placeholder("inp2")
        inp3 = graph.placeholder("inp3")
        a = graph.call_function(torch.ops.aten.clone.default, args=(inp1,))
        b = graph.call_function(torch.ops.aten.clone.default, args=(inp2,))
        d = graph.call_function(torch.ops.aten.clone.default, args=(inp3,))
        c = graph.call_function(torch.ops.aten.add.Tensor, args=(a, b))
        e = graph.call_function(torch.ops.aten.add.Tensor, args=(c, d))
        bwd = graph.call_function(torch.ops.aten.add.Tensor, args=(e, e))
        graph.output(bwd)
        for n in (a, b, c, d, e):
            n.meta["recompute"] = CheckpointPolicy.MUST_RECOMPUTE
        bwd.meta["autograd_backward"] = True

        original_names_in_order = [n.name for n in (a, b, d, c, e)]
        e_name = e.name

        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        result = selective_activation_remat_pass(gm)

        nodes = list(result.graph.nodes)
        dups = [n for n in nodes if n.name.endswith("_recomputed")]
        # All 5 must_recompute nodes are transitive deps of bwd.
        self.assertEqual(len(dups), 5)

        # Dup graph order matches the forward order of the originals
        # (a, b, d, c, e).
        self.assertEqual(
            [n.name for n in dups],
            [name + "_recomputed" for name in original_names_in_order],
        )

        # The backward node's must_recompute input was redirected to the dup
        # of e; the original e (now dead) was erased. Use the Python ``bwd``
        # reference rather than searching by ``autograd_backward`` because
        # dups also carry that flag.
        for inp in bwd.all_input_nodes:
            self.assertEqual(inp.name, e_name + "_recomputed")
        self.assertNotIn(e_name, [n.name for n in nodes])

    def test_loss_region_excluded_from_remat(self):
        # Two disjoint backward regions, each consuming a must_recompute forward
        # node: one is a chunked-loss region (module_fqn 'loss') and one is a
        # model-layer AC region (module_fqn 'layers.0.*'). Only the model-layer
        # region drives remat; the loss region is left untouched (no error).
        graph = torch.fx.Graph()
        inp1 = graph.placeholder("inp1")
        inp2 = graph.placeholder("inp2")
        a = graph.call_function(torch.ops.aten.clone.default, args=(inp1,))
        b = graph.call_function(torch.ops.aten.clone.default, args=(inp2,))
        bwd_loss = graph.call_function(torch.ops.aten.add.Tensor, args=(a, a))
        sep = graph.call_function(torch.ops.aten.neg.default, args=(inp1,))
        bwd_layer = graph.call_function(torch.ops.aten.mul.Tensor, args=(b, b))
        graph.output((bwd_loss, sep, bwd_layer))
        for node in (a, b):
            node.meta["recompute"] = CheckpointPolicy.MUST_RECOMPUTE
        bwd_loss.meta["autograd_backward"] = True
        bwd_loss.meta["custom"] = {"module_fqn": "loss"}
        bwd_layer.meta["autograd_backward"] = True
        bwd_layer.meta["custom"] = {"module_fqn": "layers.0.attention"}
        a_name, b_name = a.name, b.name

        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        result = selective_activation_remat_pass(gm)

        nodes = list(result.graph.nodes)
        node_names = [n.name for n in nodes]
        dups = {n.name for n in nodes if n.name.endswith("_recomputed")}
        # Only the model-layer region's input is rematerialized.
        self.assertEqual(dups, {b_name + "_recomputed"})
        for inp in bwd_layer.all_input_nodes:
            self.assertEqual(inp.name, b_name + "_recomputed")
        self.assertNotIn(b_name, node_names)
        # The loss region is untouched: its recompute input stays and is still
        # read directly (not remat'd, not erased).
        self.assertIn(a_name, node_names)
        for inp in bwd_loss.all_input_nodes:
            self.assertEqual(inp.name, a_name)

    def test_multiple_model_layer_regions_recompute_errors(self):
        # Two disjoint *model-layer* backward regions that both need remat is
        # still unsupported and must error (remat handles a single such region).
        graph = torch.fx.Graph()
        inp1 = graph.placeholder("inp1")
        inp2 = graph.placeholder("inp2")
        a = graph.call_function(torch.ops.aten.clone.default, args=(inp1,))
        b = graph.call_function(torch.ops.aten.clone.default, args=(inp2,))
        bwd1 = graph.call_function(torch.ops.aten.add.Tensor, args=(a, a))
        sep = graph.call_function(torch.ops.aten.neg.default, args=(inp1,))
        bwd2 = graph.call_function(torch.ops.aten.mul.Tensor, args=(b, b))
        graph.output((bwd1, sep, bwd2))
        for node in (a, b):
            node.meta["recompute"] = CheckpointPolicy.MUST_RECOMPUTE
        bwd1.meta["autograd_backward"] = True
        bwd1.meta["custom"] = {"module_fqn": "layers.0.attention"}
        bwd2.meta["autograd_backward"] = True
        bwd2.meta["custom"] = {"module_fqn": "layers.1.attention"}

        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        with self.assertRaisesRegex(RuntimeError, "disjoint backward regions"):
            selective_activation_remat_pass(gm)

    def test_forward_consumer_keeps_original(self):
        """When a must_recompute node has both forward and backward
        consumers, the original stays (forward needs it) and a dup is
        inserted for the backward consumer. The original is not erased.

            a = clone(inp)              # must_recompute, used by both fwd + bwd
            fwd_use = a + a             # forward consumer
            bwd = a * a                 # autograd_backward consumer
        """
        from torchtitan.experiments.graph_trainer.selective_activation_remat import (
            selective_activation_remat_pass,
        )

        graph = torch.fx.Graph()
        inp = graph.placeholder("inp")
        a = graph.call_function(torch.ops.aten.clone.default, args=(inp,))
        fwd_use = graph.call_function(torch.ops.aten.add.Tensor, args=(a, a))
        bwd = graph.call_function(torch.ops.aten.mul.Tensor, args=(a, a))
        graph.output((fwd_use, bwd))
        a.meta["recompute"] = CheckpointPolicy.MUST_RECOMPUTE
        bwd.meta["autograd_backward"] = True

        a_name = a.name

        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        result = selective_activation_remat_pass(gm)

        names = [n.name for n in result.graph.nodes]
        # Original kept (forward consumer still needs it) and dup inserted.
        self.assertIn(a_name, names)
        self.assertIn(a_name + "_recomputed", names)

        # bwd's args go to the dup; fwd_use still points to the original.
        bwd_node = next(
            n for n in result.graph.nodes if n.target is torch.ops.aten.mul.Tensor
        )
        for inp_node in bwd_node.all_input_nodes:
            self.assertEqual(inp_node.name, a_name + "_recomputed")
        fwd_use_node = next(
            n for n in result.graph.nodes if n.target is torch.ops.aten.add.Tensor
        )
        for inp_node in fwd_use_node.all_input_nodes:
            self.assertEqual(inp_node.name, a_name)

    def test_subgraph_get_attr_duplicated_for_recompute(self):
        """A recomputed node with a subgraph (GraphModule) get_attr input gets
        a PRIVATE copy of that get_attr, pointing at the same submodule.

        This mirrors how ``flex_attention`` references its score_mod / mask_mod
        subgraphs via get_attr nodes. Without private copies, the later
        regional_inductor pass — which partitions each flex node together with
        its subgraph get_attrs — would place a shared get_attr in only one
        region, leaving the other flex with the subgraph passed as a raw
        GraphModule arg (which fails to compile).

            sub = get_attr("subgraph")     # GraphModule attribute
            a   = some_op(inp, sub)        # must_recompute (HOP-like)
            bwd = a + a                    # autograd_backward consumer

        Plain-tensor get_attrs are left shared (they are never region-annotated),
        so this only fires for GraphModule-valued attributes.
        """
        from torchtitan.experiments.graph_trainer.selective_activation_remat import (
            selective_activation_remat_pass,
        )

        # A trivial GraphModule used as the subgraph attribute, plus a plain
        # tensor constant attribute that must stay shared (not duplicated).
        sub_graph = torch.fx.Graph()
        sub_graph.output(sub_graph.placeholder("x"))
        root = torch.nn.Module()
        root.subgraph = torch.fx.GraphModule(torch.nn.Module(), sub_graph)
        root.const = torch.nn.Buffer(torch.zeros(1))

        graph = torch.fx.Graph()
        inp = graph.placeholder("inp")
        sub = graph.get_attr("subgraph")
        const = graph.get_attr("const")
        a = graph.call_function(torch.ops.aten.add.Tensor, args=(inp, sub))
        a.kwargs = {"const": const}
        bwd = graph.call_function(torch.ops.aten.add.Tensor, args=(a, a))
        graph.output(bwd)
        a.meta["recompute"] = CheckpointPolicy.MUST_RECOMPUTE
        bwd.meta["autograd_backward"] = True

        gm = torch.fx.GraphModule(root, graph)
        result = selective_activation_remat_pass(gm)

        dups = [n for n in result.graph.nodes if n.name.endswith("_recomputed")]
        self.assertEqual(len(dups), 1)
        dup = dups[0]

        dup_subgraph_attrs = [
            n
            for n in dup.all_input_nodes
            if n.op == "get_attr" and n.target == "subgraph"
        ]
        self.assertEqual(len(dup_subgraph_attrs), 1)
        # Private copy: a distinct node from the original, same submodule target.
        self.assertIsNot(dup_subgraph_attrs[0], sub)

        # The plain-tensor const get_attr is shared, not duplicated.
        const_attrs = [
            n for n in result.graph.nodes if n.op == "get_attr" and n.target == "const"
        ]
        self.assertEqual(len(const_attrs), 1)

    def test_offload_reload_chain_hoisted(self):
        """Mirrors the graph the CPU-offload pass produces: a forward
        offload chain (``ao.offload`` -> ``ao.wait_tensor``) and a backward
        reload chain (``ao.reload`` -> ``ao.wait_tensor``). When a
        recomputed node references the offloaded forward node F, the dup
        must read from the backward wait_tensor on GPU, not from F's
        freed-GPU storage. The remat pass discovers the offload chain
        through graph structure and hoists the backward reload chain in
        front of the dup's target.

            # Forward (autograd_backward=False)
            F           = clone(inp1)
            offload_op  = ao.offload(F)
            fwd_wait    = ao.wait_tensor(offload_op, F)
            N           = add(F, inp2)             # must_recompute

            # Backward (autograd_backward=True), placed after bwd_use so
            # the hoist actually has work to do:
            bwd_use     = mul(N, N)
            reload_op   = ao.reload(fwd_wait, "cuda")
            bwd_wait    = ao.wait_tensor(reload_op)
            bwd_other   = mul(bwd_wait, bwd_wait)
        """
        # Importing this module registers the ao::offload / ao::reload /
        # ao::wait_tensor ops with torch.ops.
        import torch._functorch._activation_offloading.offload_ops  # noqa: F401

        from torchtitan.experiments.graph_trainer.selective_activation_remat import (
            selective_activation_remat_pass,
        )

        graph = torch.fx.Graph()
        inp1 = graph.placeholder("inp1")
        inp2 = graph.placeholder("inp2")
        f = graph.call_function(torch.ops.aten.clone.default, args=(inp1,))
        offload_op = graph.call_function(torch.ops.ao.offload.default, args=(f,))
        fwd_wait = graph.call_function(
            torch.ops.ao.wait_tensor.default, args=(offload_op, f)
        )
        n = graph.call_function(torch.ops.aten.add.Tensor, args=(f, inp2))
        bwd_use = graph.call_function(torch.ops.aten.mul.Tensor, args=(n, n))
        reload_op = graph.call_function(
            torch.ops.ao.reload.default, args=(fwd_wait, "cuda")
        )
        bwd_wait = graph.call_function(
            torch.ops.ao.wait_tensor.default, args=(reload_op,)
        )
        bwd_other = graph.call_function(
            torch.ops.aten.mul.Tensor, args=(bwd_wait, bwd_wait)
        )
        graph.output((bwd_use, bwd_other))

        n.meta["recompute"] = CheckpointPolicy.MUST_RECOMPUTE
        bwd_use.meta["autograd_backward"] = True
        reload_op.meta["autograd_backward"] = True
        bwd_wait.meta["autograd_backward"] = True
        bwd_other.meta["autograd_backward"] = True

        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        result = selective_activation_remat_pass(gm)

        nodes = list(result.graph.nodes)

        # Backward reload chain has been moved in front of the dup's target
        # (bwd_use) in topological order (reload_op before bwd_wait).
        reload_idx = nodes.index(reload_op)
        wait_idx = nodes.index(bwd_wait)
        bwd_use_idx = nodes.index(bwd_use)
        self.assertLess(reload_idx, wait_idx)
        self.assertLess(wait_idx, bwd_use_idx)

        # The forward offload chain stayed in forward (no hoist needed).
        offload_idx = nodes.index(offload_op)
        fwd_wait_idx = nodes.index(fwd_wait)
        self.assertLess(offload_idx, fwd_wait_idx)
        # Forward chain is also before the (hoisted) backward chain.
        self.assertLess(fwd_wait_idx, reload_idx)

        # The dup of N references bwd_wait (via the offload chain
        # redirect), not the original offloaded forward node F.
        dup = next(d for d in nodes if d.name.endswith("_recomputed"))
        self.assertIn(bwd_wait, dup.all_input_nodes)
        self.assertNotIn(f, dup.all_input_nodes)
        # The dup itself is positioned after the hoisted chain and before
        # its target.
        dup_idx = nodes.index(dup)
        self.assertLess(wait_idx, dup_idx)
        self.assertLess(dup_idx, bwd_use_idx)

        # bwd_use's args were redirected to the dup.
        for inp in bwd_use.all_input_nodes:
            self.assertIs(inp, dup)

        # bwd_other still consumes the (now hoisted) bwd_wait.
        for inp in bwd_other.all_input_nodes:
            self.assertIs(inp, bwd_wait)

    def test_offload_reload_chain_already_in_front_not_hoisted(self):
        """The CPU offload pass deliberately places ``ao.reload`` well before
        its ``ao.wait_tensor`` (via ``prefetch_reloads``) so the async H2D
        overlaps with backward compute. If the reload chain is already in
        front of the dup that needs it, ``ensure_offload_chain_before`` must
        leave it alone — re-hoisting collapses that prefetch gap and
        serializes the H2D against compute.

            # Forward (autograd_backward=False):
            F           = clone(inp1)
            offload_op  = ao.offload(F)
            fwd_wait    = ao.wait_tensor(offload_op, F)
            N           = add(F, inp2)              # must_recompute

            # Backward (autograd_backward=True), reload chain placed
            # EARLY — before the dup's target — exactly as
            # ``prefetch_reloads`` would arrange it:
            early_bwd   = mul(inp1, inp1)
            reload_op   = ao.reload(fwd_wait, "cuda")
            bwd_wait    = ao.wait_tensor(reload_op)
            middle_bwd  = mul(bwd_wait, bwd_wait)   # uses reload chain too
            bwd_use     = mul(N, N)                 # consumes N (dup target)
        """
        import torch._functorch._activation_offloading.offload_ops  # noqa: F401

        from torchtitan.experiments.graph_trainer.selective_activation_remat import (
            selective_activation_remat_pass,
        )

        graph = torch.fx.Graph()
        inp1 = graph.placeholder("inp1")
        inp2 = graph.placeholder("inp2")
        f = graph.call_function(torch.ops.aten.clone.default, args=(inp1,))
        offload_op = graph.call_function(torch.ops.ao.offload.default, args=(f,))
        fwd_wait = graph.call_function(
            torch.ops.ao.wait_tensor.default, args=(offload_op, f)
        )
        n = graph.call_function(torch.ops.aten.add.Tensor, args=(f, inp2))
        early_bwd = graph.call_function(torch.ops.aten.mul.Tensor, args=(inp1, inp1))
        reload_op = graph.call_function(
            torch.ops.ao.reload.default, args=(fwd_wait, "cuda")
        )
        bwd_wait = graph.call_function(
            torch.ops.ao.wait_tensor.default, args=(reload_op,)
        )
        middle_bwd = graph.call_function(
            torch.ops.aten.mul.Tensor, args=(bwd_wait, bwd_wait)
        )
        bwd_use = graph.call_function(torch.ops.aten.mul.Tensor, args=(n, n))
        graph.output((middle_bwd, bwd_use))

        n.meta["recompute"] = CheckpointPolicy.MUST_RECOMPUTE
        early_bwd.meta["autograd_backward"] = True
        reload_op.meta["autograd_backward"] = True
        bwd_wait.meta["autograd_backward"] = True
        middle_bwd.meta["autograd_backward"] = True
        bwd_use.meta["autograd_backward"] = True

        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        result = selective_activation_remat_pass(gm)

        nodes = list(result.graph.nodes)
        early_idx = nodes.index(early_bwd)
        reload_idx = nodes.index(reload_op)
        wait_idx = nodes.index(bwd_wait)
        middle_idx = nodes.index(middle_bwd)
        bwd_use_idx = nodes.index(bwd_use)

        # The reload chain stayed at its original position (between early_bwd
        # and middle_bwd), preserving the prefetch gap. If the pass had
        # collapsed it next to bwd_use, reload_op/bwd_wait would land after
        # middle_bwd — which would also be a topology violation since
        # middle_bwd consumes bwd_wait.
        self.assertLess(early_idx, reload_idx)
        self.assertLess(reload_idx, wait_idx)
        self.assertLess(wait_idx, middle_idx)
        self.assertLess(middle_idx, bwd_use_idx)

        # The dup of N references bwd_wait (at its original position) and
        # is itself inserted right before bwd_use.
        dup = next(d for d in nodes if d.name.endswith("_recomputed"))
        self.assertIn(bwd_wait, dup.all_input_nodes)
        dup_idx = nodes.index(dup)
        self.assertLess(wait_idx, dup_idx)
        self.assertLess(dup_idx, bwd_use_idx)

        # middle_bwd still consumes bwd_wait at its original location.
        for inp in middle_bwd.all_input_nodes:
            self.assertIs(inp, bwd_wait)


class TestEliminateDeadCodePass(TestCase):
    """Unit tests for eliminate_dead_code_pass."""

    def test_removes_dead_pure_node_keeps_live(self):
        g = torch.fx.Graph()
        x = g.placeholder("x")
        live = g.call_function(torch.ops.aten.relu.default, (x,))
        g.call_function(torch.ops.aten.add.Tensor, (x, x))  # dead: no users
        g.output(live)
        gm = torch.fx.GraphModule(torch.nn.Module(), g)

        eliminate_dead_code_pass(gm)
        targets = [n.target for n in gm.graph.nodes if n.op == "call_function"]
        self.assertIn(torch.ops.aten.relu.default, targets)
        self.assertNotIn(torch.ops.aten.add.Tensor, targets)

    def test_removes_unreachable_bad_node(self):
        fake_mode = torch._subclasses.FakeTensorMode(allow_non_fake_inputs=True)
        with fake_mode:
            a_meta = torch.empty(256, 8192, device="cuda")
            b_meta = torch.empty(16384, 512, device="cuda")
            out_meta = torch.empty(1, device="cuda")

        graph = torch.fx.Graph()
        out = graph.placeholder("out")
        a = graph.placeholder("a")
        b = graph.placeholder("b")
        graph.call_function(torch.ops.aten.mm.default, args=(a, b))
        graph.output(out)
        gm = torch.fx.GraphModule(torch.nn.Module(), graph)
        out.meta["val"] = out_meta
        a.meta["val"] = a_meta
        b.meta["val"] = b_meta

        eliminate_dead_code_pass(gm)
        self.assertNotIn(torch.ops.aten.mm.default, {n.target for n in gm.graph.nodes})

    def test_keeps_impure_node_with_no_users(self):
        # copy_ mutates its first arg (impure); DCE must keep it even though its
        # own result is unused.
        g = torch.fx.Graph()
        x = g.placeholder("x")
        y = g.placeholder("y")
        g.call_function(torch.ops.aten.copy_.default, (x, y))  # impure, unused result
        out = g.call_function(torch.ops.aten.relu.default, (x,))
        g.output(out)
        gm = torch.fx.GraphModule(torch.nn.Module(), g)

        eliminate_dead_code_pass(gm)
        targets = [n.target for n in gm.graph.nodes if n.op == "call_function"]
        self.assertIn(torch.ops.aten.copy_.default, targets)


class TestIsFullCudaGraphCompatible(TestCase):
    """Pure-CPU tests for the per-node CUDA-graph-safety predicate and the
    whole-graph gate built on it."""

    def test_clean_graph_is_cuda_graph_fully_compatible(self):
        g = torch.fx.Graph()
        x = g.placeholder("x")
        relu = g.call_function(torch.ops.aten.relu.default, (x,))
        g.output(relu)
        gm = torch.fx.GraphModule(torch.nn.Module(), g)
        self.assertTrue(is_cuda_graph_node_compatible(relu))
        self.assertTrue(is_cuda_graph_fully_compatible(gm))

    def test_local_scalar_dense_is_unsafe(self):
        # _local_scalar_dense (.item()/.tolist()) extracts a host scalar a CUDA
        # graph replay can't reproduce -> unsafe, so the graph is not one piece.
        g = torch.fx.Graph()
        x = g.placeholder("x")
        s = g.call_function(torch.ops.aten._local_scalar_dense.default, (x,))
        g.output(s)
        gm = torch.fx.GraphModule(torch.nn.Module(), g)
        self.assertFalse(is_cuda_graph_node_compatible(s))
        self.assertFalse(is_cuda_graph_fully_compatible(gm))


class TestEagerChunking(TestCase):
    def _config(
        self,
        *,
        chunk_dim: str = "batch",
        module_fqn: str = "layers.*",
    ) -> GraphTrainerCompileConfig:
        return GraphTrainerCompileConfig(
            ep_overlap=EpOverlapConfig(
                enabled=True,
                chunk_dim=chunk_dim,
                module_fqn=module_fqn,
            ),
        )

    def test_eager_chunking_is_idempotent(self):
        class Block(torch.nn.Module):
            def forward(self, x):
                return x.sin()

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.layers = torch.nn.ModuleList([Block()])

            def forward(self, x):
                return self.layers[0](x)

        model = Model()
        config = self._config()
        maybe_apply_ep_overlap_eager_chunking(model, config)
        wrapped_forward = model.layers[0].forward
        maybe_apply_ep_overlap_eager_chunking(model, config)

        self.assertIs(model.layers[0].forward, wrapped_forward)

    def test_eager_chunking_compile_disabled_is_noop(self):
        class Block(torch.nn.Module):
            def forward(self, x):
                return x.sin()

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.layers = torch.nn.ModuleList([Block()])

            def forward(self, x):
                return self.layers[0](x)

        model = Model()
        forward = model.layers[0].forward
        config = self._config()
        config.ep_overlap.enabled = False

        maybe_apply_ep_overlap_eager_chunking(model, config)

        self.assertIs(model.layers[0].forward.__func__, forward.__func__)

    def test_eager_chunking_rejects_unsupported_output_type(self):
        class Block(torch.nn.Module):
            def forward(self, x):
                return {"x": x}

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.layers = torch.nn.ModuleList([Block()])

            def forward(self, x):
                return self.layers[0](x)

        model = Model()
        maybe_apply_ep_overlap_eager_chunking(model, self._config())

        with self.assertRaisesRegex(TypeError, "layers.0.*dict"):
            model(torch.randn(4, 3))

    def test_transformer_batch_chunking_splits_token_metadata_by_batch(self):
        seen_positions = []
        seen_padding_masks = []

        class Block(torch.nn.Module):
            def forward(
                self,
                x,
                attention_metadata=None,
                positions=None,
                *,
                padding_mask=None,
            ):
                seen_positions.append(positions)
                seen_padding_masks.append(padding_mask)
                return x

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.layers = torch.nn.ModuleList([Block()])

            def forward(self, x, positions, padding_mask):
                return self.layers[0](x, None, positions, padding_mask=padding_mask)

        model = Model()
        maybe_apply_ep_overlap_eager_chunking(model, self._config())
        x = torch.randn(4, 4, 2)
        positions = torch.arange(16).view(4, 4)
        padding_mask = torch.tensor(
            [
                [False, False, False, True],
                [False, False, True, True],
                [False, False, False, False],
                [False, True, True, True],
            ]
        )

        self.assertEqual(model(x, positions, padding_mask), x)
        self.assertEqual([tuple(pos.shape) for pos in seen_positions], [(2, 4), (2, 4)])
        self.assertEqual(seen_positions[0], positions[:2])
        self.assertEqual(seen_positions[1], positions[2:])
        self.assertEqual(seen_padding_masks[0], padding_mask[:2])
        self.assertEqual(seen_padding_masks[1], padding_mask[2:])

    def test_transformer_batch_chunking_rejects_same_extent_tensor_mask(self):
        class Block(torch.nn.Module):
            def forward(self, x, attention_metadata):
                return x + attention_metadata

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.layers = torch.nn.ModuleList([Block()])

            def forward(self, x, attention_metadata):
                return self.layers[0](x, attention_metadata)

        model = Model()
        maybe_apply_ep_overlap_eager_chunking(model, self._config())

        with self.assertRaisesRegex(
            ValueError,
            "attention_metadata must be None, BlockMask.*upstream .*TransformerBlock",
        ):
            model(torch.randn(4, 3), torch.randn(4, 3))

    def test_moe_chunking_splits_activation_only(self):
        seen_shapes = []

        class Moe(torch.nn.Module):
            def forward(self, x):
                seen_shapes.append(tuple(x.shape))
                return x

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.layers = torch.nn.ModuleList([torch.nn.Module()])
                self.layers[0].moe = Moe()

            def forward(self, x):
                return self.layers[0].moe(x)

        model = Model()
        maybe_apply_ep_overlap_eager_chunking(
            model,
            self._config(chunk_dim="seq", module_fqn="layers.*.moe"),
        )
        x = torch.randn(8, 3)

        self.assertEqual(model(x), x)
        self.assertEqual(seen_shapes, [(4, 3), (4, 3)])

    def test_moe_chunking_splits_padding_mask_with_activation(self):
        seen_inputs = []

        class Moe(torch.nn.Module):
            def forward(self, x, *, padding_mask_T=None):
                seen_inputs.append((x, padding_mask_T))
                return x.masked_fill(padding_mask_T.unsqueeze(-1), 0)

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.layers = torch.nn.ModuleList([torch.nn.Module()])
                self.layers[0].moe = Moe()

            def forward(self, x, padding_mask_T):
                return self.layers[0].moe(x, padding_mask_T=padding_mask_T)

        model = Model()
        maybe_apply_ep_overlap_eager_chunking(
            model,
            self._config(chunk_dim="seq", module_fqn="layers.*.moe"),
        )
        x = torch.randn(8, 3)
        padding_mask_T = torch.tensor(
            [False, True, False, True, True, False, True, False]
        )

        self.assertEqual(
            model(x, padding_mask_T),
            x.masked_fill(padding_mask_T.unsqueeze(-1), 0),
        )
        self.assertEqual(
            [tuple(chunk_x.shape) for chunk_x, _ in seen_inputs],
            [(4, 3), (4, 3)],
        )
        self.assertEqual(seen_inputs[0][1], padding_mask_T[:4])
        self.assertEqual(seen_inputs[1][1], padding_mask_T[4:])

    def test_moe_chunking_shares_aux_loss_denominator(self):
        seen_denominators = []

        class Moe(torch.nn.Module):
            def forward(self, x, *, aux_loss_denominator):
                seen_denominators.append(aux_loss_denominator)
                return x

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.layers = torch.nn.ModuleList([torch.nn.Module()])
                self.layers[0].moe = Moe()

            def forward(self, x, aux_loss_denominator):
                return self.layers[0].moe(x, aux_loss_denominator=aux_loss_denominator)

        model = Model()
        maybe_apply_ep_overlap_eager_chunking(
            model,
            self._config(chunk_dim="seq", module_fqn="layers.*.moe"),
        )
        x = torch.randn(8, 3)
        aux_loss_denominator = torch.tensor(8)

        self.assertEqual(model(x, aux_loss_denominator), x)
        self.assertEqual(seen_denominators, [aux_loss_denominator] * 2)

    def test_moe_chunking_rejects_extra_tensor_input(self):
        class Moe(torch.nn.Module):
            def forward(self, x, aux):
                return x + aux

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.layers = torch.nn.ModuleList([torch.nn.Module()])
                self.layers[0].moe = Moe()

            def forward(self, x, aux):
                return self.layers[0].moe(x, aux)

        model = Model()
        maybe_apply_ep_overlap_eager_chunking(
            model,
            self._config(module_fqn="layers.*.moe"),
        )

        with self.assertRaisesRegex(
            ValueError,
            "expected exactly one positional activation tensor.*upstream MoE.forward",
        ):
            model(torch.randn(4, 3), torch.randn(4, 3))

    def test_eager_chunking_traces_overlap_metadata(self):
        class Block(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.linear = torch.nn.Linear(3, 3)

            def forward(self, x):
                return torch.relu(self.linear(x))

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.layers = torch.nn.ModuleList([Block()])

            def forward(self, x):
                return self.layers[0](x)

        model = Model()
        annotate_module_fqns(model)
        maybe_apply_ep_overlap_eager_chunking(model, self._config())

        traced = minimal_fx_tracer(lambda inputs: model(inputs), module=model)(
            torch.randn(4, 3)
        )
        gm = populate_eager_chunk_metadata_pass(traced.gm)

        body_nodes = [
            node
            for node in gm.graph.nodes
            if node.meta.get("chunked_region_role") == "body"
        ]
        roles = {
            node.meta.get("chunked_region_role")
            for node in gm.graph.nodes
            if node.meta.get("chunked_region_fqn") == "layers.0"
        }
        self.assertEqual({node.meta.get("chunk_id") for node in body_nodes}, {0, 1})
        self.assertEqual(
            {node.meta.get("chunked_region_fqn") for node in body_nodes},
            {"layers.0"},
        )
        self.assertIn("split_boundary", roles)
        self.assertIn("materialization", roles)

    def test_eager_chunking_splits_block_mask_batch_metadata(self):
        from torch.nn.attention.flex_attention import create_block_mask

        seen_masks = []

        def mask_mod(b, h, q_idx, kv_idx):
            return (b == 2) & (q_idx >= kv_idx)

        class Block(torch.nn.Module):
            def forward(self, x, attention_metadata, positions):
                seen_masks.append(attention_metadata)
                return x

        class Model(torch.nn.Module):
            def __init__(self):
                super().__init__()
                self.layers = torch.nn.ModuleList([Block()])

            def forward(self, x, attention_metadata, positions):
                return self.layers[0](x, attention_metadata, positions)

        model = Model()
        maybe_apply_ep_overlap_eager_chunking(model, self._config())
        block_mask = create_block_mask(
            mask_mod,
            B=4,
            H=None,
            Q_LEN=128,
            KV_LEN=128,
            device="cpu",
        )
        x = torch.randn(4, 128, 8)
        positions = torch.arange(128).repeat(4, 1)

        self.assertEqual(model(x, block_mask, positions).shape, x.shape)
        self.assertEqual(len(seen_masks), 2)
        self.assertEqual([mask.kv_num_blocks.size(0) for mask in seen_masks], [2, 2])

        b = torch.tensor(0)
        h = torch.tensor(0)
        q_idx = torch.tensor(1)
        kv_idx = torch.tensor(0)
        self.assertFalse(seen_masks[0].mask_mod(b, h, q_idx, kv_idx).item())
        self.assertTrue(seen_masks[1].mask_mod(b, h, q_idx, kv_idx).item())


if __name__ == "__main__":
    from torch.testing._internal.common_utils import run_tests

    run_tests()
