# 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 copy
import unittest
from dataclasses import dataclass

import torch
import torch.nn as nn

from torchtitan.components.data.types import TokenizedTrainingMicrobatch
from torchtitan.distributed.activation_checkpoint import FullAC, SelectiveAC
from torchtitan.experiments.graph_trainer.llama3 import (
    build_model_config as build_llama3_model_config,
)
from torchtitan.experiments.graph_trainer.tests._trainer_test_utils import (
    build_minimal_trainer,
    single_device_parallelism_context,
)
from torchtitan.experiments.graph_trainer.trainer import GraphTrainer
from torchtitan.trainer import Trainer

DTYPE = torch.bfloat16
NUM_TOKENS = 2 * 2048
MAX_PEAK_MEMORY_RATIO = 1.10
DEBUGMODEL = "debugmodel"


def _set_deterministic() -> None:
    torch.manual_seed(42)
    torch.cuda.manual_seed_all(42)
    torch.use_deterministic_algorithms(True)


def _build_model(model_flavor: str, attn_backend: str = "flex") -> nn.Module:
    model_config = build_llama3_model_config(model_flavor, attn_backend=attn_backend)
    with torch.device("meta"):
        model = model_config.build()
    model.to_empty(device="cuda")
    with torch.no_grad():
        model.init_states(buffer_device=None)
    model.to(dtype=DTYPE)
    model.train()
    return model


@dataclass(frozen=True)
class StepResult:
    loss: torch.Tensor
    grads: list[torch.Tensor]
    reserved_gib: float
    active_gib: float


def _measure_step(
    trainer: Trainer, tokens: torch.Tensor, labels: torch.Tensor
) -> StepResult:
    model = trainer.engine.model_parts[0]
    model.zero_grad(set_to_none=True)
    global_loss_token_counts = torch.tensor(
        labels.numel(), dtype=torch.float, device="cuda"
    )
    # The dataloader always supplies per-document positions, which the trainer
    # requires to build the FlexInnerAttention mask. Reset positions between the
    # packed documents.
    positions = torch.arange(NUM_TOKENS, device="cuda", dtype=torch.int32) % 2048

    torch.cuda.synchronize()
    torch.cuda.reset_peak_memory_stats()
    result = trainer.engine.forward_backward(
        microbatch_groups=[
            [
                TokenizedTrainingMicrobatch(
                    input=tokens,
                    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),
    )
    torch.cuda.synchronize()

    stats = torch.cuda.memory_stats()
    grads = [param.grad.detach().clone() for param in model.parameters()]
    return StepResult(
        loss=result.loss.detach().clone(),
        grads=grads,
        reserved_gib=torch.cuda.max_memory_reserved() / 1e9,
        active_gib=stats["active_bytes.all.peak"] / 1e9,
    )


@unittest.skipUnless(torch.cuda.is_available(), "CUDA required")
class TestGraphSACPeakMemory(unittest.TestCase):
    def setUp(self):
        self.parallelism_context = self.enterContext(
            single_device_parallelism_context()
        )

        _set_deterministic()
        model = _build_model(DEBUGMODEL)
        self.state_dict = {
            key: value.detach().cpu().clone()
            for key, value in model.state_dict().items()
        }
        del model
        torch.cuda.empty_cache()
        self.tokens = torch.randint(0, 2048, (NUM_TOKENS,), device="cuda")
        self.labels = torch.randint(0, 2048, (NUM_TOKENS,), device="cuda")

    def tearDown(self):
        torch.use_deterministic_algorithms(False)

    def test_llama3_debugmodel_peak_memory_matches_eager_selective_ac(self):
        eager_model = _build_model(DEBUGMODEL)
        eager_model.load_state_dict(copy.deepcopy(self.state_dict))
        SelectiveAC.Config().build().apply(eager_model)
        eager_trainer = build_minimal_trainer(
            eager_model,
            build_llama3_model_config(DEBUGMODEL),
            Trainer,
            parallelism_context=self.parallelism_context,
        )

        traced_model = _build_model(DEBUGMODEL)
        traced_model.load_state_dict(copy.deepcopy(self.state_dict))
        traced_trainer = build_minimal_trainer(
            traced_model,
            build_llama3_model_config(DEBUGMODEL),
            GraphTrainer,
            activation_checkpoint_mode="selective",
            parallelism_context=self.parallelism_context,
        )
        # Both policies retain explicitly identified expensive operations while
        # recomputing the surrounding inexpensive work.
        traced_trainer.config.compile.memory_policy = "default"

        # Warm up both paths so allocator and one-time tracing setup do not skew
        # the measured peak memory.
        _measure_step(eager_trainer, self.tokens, self.labels)
        _measure_step(traced_trainer, self.tokens, self.labels)
        torch.cuda.empty_cache()

        eager = _measure_step(eager_trainer, self.tokens, self.labels)
        traced = _measure_step(traced_trainer, self.tokens, self.labels)

        self.assertTrue(
            torch.equal(eager.loss, traced.loss),
            f"loss mismatch: eager={eager.loss.item()} traced={traced.loss.item()}",
        )
        for idx, (eager_grad, traced_grad) in enumerate(
            zip(eager.grads, traced.grads, strict=True)
        ):
            self.assertTrue(
                torch.equal(eager_grad, traced_grad), f"grad[{idx}] mismatch"
            )

        reserved_ratio = traced.reserved_gib / eager.reserved_gib
        active_ratio = traced.active_gib / eager.active_gib
        self.assertLessEqual(
            reserved_ratio,
            MAX_PEAK_MEMORY_RATIO,
            "graph SAC reserved peak memory too high: "
            f"traced={traced.reserved_gib:.3f} GiB, "
            f"eager={eager.reserved_gib:.3f} GiB, "
            f"ratio={reserved_ratio:.3f}",
        )
        self.assertLessEqual(
            active_ratio,
            MAX_PEAK_MEMORY_RATIO,
            "graph SAC active peak memory too high: "
            f"traced={traced.active_gib:.3f} GiB, "
            f"eager={eager.active_gib:.3f} GiB, "
            f"ratio={active_ratio:.3f}",
        )

    def test_llama3_debugmodel_full_ac_numerics_match_eager(self):
        """Graph full recompute produces bitwise-identical loss and grads vs eager."""
        eager_model = _build_model(DEBUGMODEL)
        eager_model.load_state_dict(copy.deepcopy(self.state_dict))
        FullAC.Config().build().apply(eager_model)
        eager_trainer = build_minimal_trainer(
            eager_model,
            build_llama3_model_config(DEBUGMODEL),
            Trainer,
            parallelism_context=self.parallelism_context,
        )

        traced_model = _build_model(DEBUGMODEL)
        traced_model.load_state_dict(copy.deepcopy(self.state_dict))
        traced_trainer = build_minimal_trainer(
            traced_model,
            build_llama3_model_config(DEBUGMODEL),
            GraphTrainer,
            activation_checkpoint_mode="selective",
            parallelism_context=self.parallelism_context,
        )
        traced_trainer.config.compile.memory_policy = "full"

        _measure_step(eager_trainer, self.tokens, self.labels)
        _measure_step(traced_trainer, self.tokens, self.labels)
        torch.cuda.empty_cache()

        eager = _measure_step(eager_trainer, self.tokens, self.labels)
        traced = _measure_step(traced_trainer, self.tokens, self.labels)

        self.assertTrue(
            torch.equal(eager.loss, traced.loss),
            f"loss mismatch: eager={eager.loss.item()} traced={traced.loss.item()}",
        )
        for idx, (eager_grad, traced_grad) in enumerate(
            zip(eager.grads, traced.grads, strict=True)
        ):
            self.assertTrue(
                torch.equal(eager_grad, traced_grad), f"grad[{idx}] mismatch"
            )


if __name__ == "__main__":
    unittest.main()
