# 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.

# Numerical equivalence test for flex attention + context parallelism.
#
# Builds the flex debug model, runs a full-sequence forward as the reference,
# then runs the same weights under CP with the input/positions/BlockMask sharded
# on the sequence axis (as the trainer does). Reconstructs full logits in global
# order (a global-index tensor rides through CP input sharding to undo any
# load-balancer permutation) and compares against the reference. Correct CP attention matches
# to fp32 noise (rel < 1e-4); bf16 mixed precision masks this, so we force fp32.
#
# Run: torchrun --nproc_per_node=2 \
#          -m torchtitan.experiments.transformers_modeling_backend.tests.test_flex_cp_numerical \
#          [--balancer none|headtail|ptrr]

import argparse
import os

import torch
import torch.distributed as dist
from torchtitan_recipes.tests.transformers_modeling_backend import (
    transformers_modeling_backend_debugmodel,
    transformers_modeling_backend_debugmodel_moe,
)

from torchtitan.distributed import context_parallel, ParallelismContext
from torchtitan.distributed.context_parallel import (
    HeadTailCPLoadBalancer,
    PTRRFlexAttentionCPLoadBalancer,
)
from torchtitan.experiments.transformers_modeling_backend import build_model_config
from torchtitan.models.common.attention.cp_attention import (
    KVAllGatherCPFlexInnerAttention,
)
from torchtitan.models.common.decoder_sharding import (
    decoder_input_sharding,
    token_id_placement,
)
from torchtitan.tools import utils


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--hf_model", default="Qwen/Qwen2.5-7B")
    parser.add_argument("--seq_len", type=int, default=256)
    parser.add_argument("--bs", type=int, default=1)
    parser.add_argument(
        "--balancer", default="none", choices=["none", "headtail", "ptrr"]
    )
    parser.add_argument("--moe", action="store_true", help="use the flex MoE model")
    args = parser.parse_args()
    load_balancer_configs = {
        "none": None,
        "headtail": HeadTailCPLoadBalancer.Config(),
        "ptrr": PTRRFlexAttentionCPLoadBalancer.Config(),
    }

    world = int(os.environ["WORLD_SIZE"])
    device = utils.get_local_device()
    torch.cuda.set_device(device)
    dist.init_process_group("nccl")
    rank = dist.get_rank()
    torch.set_default_dtype(torch.float32)

    cp = world  # cp = world_size (dp=tp=pp=1)
    parallelism_context = ParallelismContext(
        dp_replicate=1,
        dp_shard=1,
        cp=cp,
        tp=1,
        pp=1,
        ep=1,
        world_size=world,
        enable_sequence_parallel=False,
    )

    # Build the job config, tweak for a small deterministic run.
    cfg = (
        transformers_modeling_backend_debugmodel_moe(
            seq_len=args.seq_len,
            deterministic=True,
        )
        if args.moe
        else transformers_modeling_backend_debugmodel(
            seq_len=args.seq_len,
            deterministic=True,
        )
    )
    cfg.model = build_model_config(
        "debugmodel_moe" if args.moe else "debugmodel",
        seq_len=args.seq_len,
        hf_model=args.hf_model,
    )
    cfg.training.num_tokens_per_microbatch_per_dp_rank = args.bs * args.seq_len
    # fp32 compute so any CP discrepancy isn't masked by bf16 FSDP mixed precision.
    cfg.training.mixed_precision_param = "float32"
    cfg.parallelism.context_parallel_degree = cp
    cfg.debug.seed = 42

    def build_model(swap_moe=False):
        model_config = cfg.model
        with torch.device(device):
            m = model_config.build()
        m.to(device)
        torch.manual_seed(42)
        m.init_weights(buffer_device=device)
        if swap_moe:
            # The reference is not parallelized, so the titan MoE swap (which
            # runs inside parallelize for the CP model) would not apply. Swap it
            # here too so both sides use the identical MoE implementation; both
            # were seeded identically, so the swap is deterministic.
            from torchtitan.experiments.transformers_modeling_backend.moe_replacement import (
                build_and_swap_native_moe,
            )

            build_and_swap_native_moe(m, parallelism_context)
            m.to(device)  # swap builds Titan MoE experts on CPU; move them back
        return m, model_config

    # --- Reference: full-sequence forward, no parallelism ---
    ref_model, _ = build_model(swap_moe=args.moe)
    ref_model.eval()
    torch.manual_seed(0)
    num_tokens = args.bs * args.seq_len
    input_ids = torch.randint(0, 100, (num_tokens,), device=device)
    positions = torch.arange(args.seq_len, device=device).repeat(args.bs)
    full_mask = ref_model.get_attention_metadata(positions)
    with torch.no_grad():
        ref_logits = ref_model(
            input_ids, positions=positions, attention_metadata=full_mask
        )
    del ref_model
    torch.cuda.empty_cache()

    # --- CP model: same deterministic init (seed 42) built before parallelize,
    # so weights match the reference; parallelize preserves values. ---
    cp_model, model_config = build_model()
    from torchtitan.experiments.transformers_modeling_backend.parallelize import (
        parallelize_hf_transformers,
    )

    parallelize_hf_transformers(
        cp_model,
        parallelism_context=parallelism_context,
        training=cfg.training,
        parallelism=cfg.parallelism,
        ac_config=None,
        dump_folder="/tmp/flex_cp_spike",
    )
    cp_model.eval()

    # Shard input / positions / mask on the sequence axis (trainer's role).
    # A global-index tensor rides along so we can undo any load-balancer
    # permutation when reconstructing full logits for the comparison.
    full_mask_cp = cp_model.get_attention_metadata(positions)
    gidx = torch.arange(num_tokens, device=device)
    input_shardings = {
        **decoder_input_sharding(),
        "global_indices": token_id_placement(),
    }
    batch = {
        "input": input_ids,
        "positions": positions,
        "global_indices": gidx,
        "attention_metadata": full_mask_cp,
    }
    with parallelism_context.activate_spmd():
        load_balancer_config = load_balancer_configs[args.balancer]
        load_balancer = (
            load_balancer_config.build(
                seq_len=context_parallel.get_cp_input_seq_len(
                    batch, input_shardings=input_shardings
                ),
                attention_metadata=batch["attention_metadata"],
            )
            if load_balancer_config is not None
            else None
        )
        permutation = (
            load_balancer.generate_permutation() if load_balancer is not None else None
        )
        batch[
            "attention_metadata"
        ] = KVAllGatherCPFlexInnerAttention.prepare_cp_metadata(
            batch["attention_metadata"],
            permutation=permutation,
        )
        batch = context_parallel.shard_tensors(
            batch,
            input_shardings=input_shardings,
            permutation=permutation,
        )
    loc_input = batch["input"]
    loc_pos = batch["positions"]
    loc_gidx = batch["global_indices"]
    loc_mask = batch["attention_metadata"]
    from torchtitan.distributed.spmd_types import annotate_input_spmd_types

    annotated = annotate_input_spmd_types(
        parallelism_context,
        {"input": loc_input, "positions": loc_pos},
        decoder_input_sharding(),
    )
    loc_input = annotated["input"]
    loc_pos = annotated["positions"]
    _fm = tuple(full_mask_cp.shape) if full_mask_cp is not None else None
    _lm = tuple(loc_mask.shape) if loc_mask is not None else None
    print(
        f"[rank {rank}] balancer={args.balancer} loc_input={tuple(loc_input.shape)} "
        f"full_mask={_fm} loc_mask={_lm}"
    )

    with torch.no_grad(), parallelism_context.activate_spmd():
        loc_logits = cp_model(loc_input, positions=loc_pos, attention_metadata=loc_mask)

    # Reconstruct full logits in global order via all-gather + index scatter.
    gathered_logits = [torch.empty_like(loc_logits) for _ in range(cp)]
    gathered_gidx = [torch.empty_like(loc_gidx) for _ in range(cp)]
    dist.all_gather(gathered_logits, loc_logits.contiguous())
    dist.all_gather(gathered_gidx, loc_gidx.contiguous())
    full = torch.zeros_like(ref_logits)
    for lg, gi in zip(gathered_logits, gathered_gidx):
        full[gi.long(), :] = lg.float()

    max_abs = (full.float() - ref_logits.float()).abs().max().item()
    ref_scale = ref_logits.float().abs().max().item()
    rel = max_abs / max(ref_scale, 1e-9)
    # Dense flex+CP is bit-exact up to fp32 noise. MoE adds grouped_mm f32
    # accumulation-order differences (documented in the add_moe_model skill's
    # scripts/numerical_equivalence.py) plus
    # FSDP expert-shard reduction order, so it uses a looser tolerance; the
    # residual stays flat across CP degrees (a real CP bug would give rel ~O(1)).
    tol = 2e-3 if args.moe else 1e-4
    passed = rel < tol
    if rank == 0:
        verdict = "PASS" if passed else "FAIL"
        print(
            f"\n==== FLEX+CP {verdict} (balancer={args.balancer}): "
            f"max_abs_diff={max_abs:.3e} ref_scale={ref_scale:.3e} rel={rel:.3e} ===="
        )
    dist.destroy_process_group()
    if not passed:
        raise SystemExit(1)


if __name__ == "__main__":
    main()
