# 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 unittest
from unittest.mock import MagicMock

import torch
from torch.optim import Adam

from torchtitan.components.optim import OptimizersContainer
from torchtitan.components.optim.lr_scheduler import LRSchedulersContainer
from torchtitan_recipes.tests.models.llama3 import llama3_debugmodel


class TestLRScheduler(unittest.TestCase):
    def test_optimizer_package_import_path(self):
        from torchtitan.components.optim import (
            LRSchedulersContainer as PackageLRSchedulersContainer,
        )

        self.assertIs(PackageLRSchedulersContainer, LRSchedulersContainer)

    def setUp(self):
        # Create a simple model with parameters
        self.model = torch.nn.Linear(10, 10)
        # Create an optimizer
        self.optimizer = Adam(self.model.parameters(), lr=0.1)

        # We don't actually call `optimizer.step()` which will cause a warning
        # from PyTorch. Avoid the warnings that may confuse people.
        self.optimizer._opt_called = True

        # Create an optimizer container
        self.optimizer_container = MagicMock(spec=OptimizersContainer)
        self.optimizer_container.__iter__.return_value = iter([self.optimizer])
        self.optimizer_container.__len__.return_value = 1

    def create_trainer_config(
        self,
        training_steps=10,
        warmup_steps=None,
        decay_ratio=None,
        decay_type=None,
        min_lr_factor=None,
    ):
        config = llama3_debugmodel()
        config.training.steps = training_steps
        if warmup_steps is not None:
            config.optim.lr_scheduler.warmup_steps = warmup_steps
        if decay_ratio is not None:
            config.optim.lr_scheduler.decay_ratio = decay_ratio
        if decay_type is not None:
            config.optim.lr_scheduler.decay_type = decay_type
        if min_lr_factor is not None:
            config.optim.lr_scheduler.min_lr_factor = min_lr_factor
        return config

    def test_linear_warmup_decay(self):
        """Test the linear warmup followed by linear decay schedule."""
        # Create a job config with 10 steps, 2 warmup steps, and linear decay
        config = self.create_trainer_config(
            training_steps=10,
            warmup_steps=2,
            decay_ratio=None,  # Use default decay: start decay immediately
            decay_type=None,
            min_lr_factor=None,
        )

        # Build the lr scheduler
        lr_scheduler = config.optim.lr_scheduler.build(
            optimizers=self.optimizer_container,
            training_steps=config.training.steps,
        )

        # Expected adjustment factors for each step
        expected_factors = [
            0.5,  # Step 0: 50% of max LR (warmup)
            1.0,  # Step 1: 100% of max LR (warmup complete)
            1.0,  # Step 2: We maunally added step of stable phase, to prevent LR from dropping to 0 at last step
            7.0 / 8.0,  # Step 3: 7/8 of max LR
            6.0 / 8.0,  # Step 4: 3/4 of max LR
            5.0 / 8.0,  # Step 5: 5/8 of max LR
            4.0 / 8.0,  # Step 6: 1/2 of max LR
            3.0 / 8.0,  # Step 7: 3/8 of max LR
            2.0 / 8.0,  # Step 8: 1/4 of max LR
            1.0 / 8.0,  # Step 9: 1/8 of max LR
        ]

        # Check the learning rate at each step
        for i, factor in enumerate(expected_factors):
            # The LambdaLR multiplies the base lr by the factor
            expected_lr = 0.1 * factor
            self.assertAlmostEqual(
                self.optimizer.param_groups[0]["lr"],
                expected_lr,
                places=6,
                msg=f"Step {i}: Expected LR {expected_lr}, got {self.optimizer.param_groups[0]['lr']}",
            )
            lr_scheduler.step()

    def test_warmup_stable_decay(self):
        """Test warmup followed by stable phase and then decay."""
        # Create a job config with 10 steps, 2 warmup steps, 3 stable steps, and 5 decay steps
        config = self.create_trainer_config(
            training_steps=10,
            warmup_steps=2,
            decay_ratio=0.5,  # 50% of steps for decay
            decay_type="linear",
            min_lr_factor=0.0,
        )

        # Build the lr scheduler
        lr_scheduler = config.optim.lr_scheduler.build(
            optimizers=self.optimizer_container,
            training_steps=config.training.steps,
        )

        # Expected adjustment factors for each step
        expected_factors = [
            0.5,  # Step 0: 50% of max LR (warmup)
            1.0,  # Step 1: 100% of max LR (warmup complete)
            1.0,  # Step 2: Stable phase
            1.0,  # Step 3: Stable phase
            1.0,  # Step 4: Stable phase
            1.0,  # Step 5: We maunally added step of stable phase, to prevent LR from dropping to 0 at last step
            0.8,  # Step 6: Linear decay starts (80% of max LR)
            0.6,  # Step 7: 60% of max LR
            0.4,  # Step 8: 40% of max LR
            0.2,  # Step 9: 20% of max LR
        ]

        # Check the learning rate at each step
        for i, factor in enumerate(expected_factors):
            expected_lr = 0.1 * factor
            self.assertAlmostEqual(
                self.optimizer.param_groups[0]["lr"],
                expected_lr,
                places=6,
                msg=f"Step {i}: Expected LR {expected_lr}, got {self.optimizer.param_groups[0]['lr']}",
            )
            lr_scheduler.step()

    def test_min_lr(self):
        """Test that the learning rate doesn't go below the minimum."""
        # Create a job config with a minimum learning rate
        config = self.create_trainer_config(
            training_steps=10,
            warmup_steps=2,
            decay_ratio=None,
            decay_type="linear",
            min_lr_factor=0.2,  # 20% of base LR as minimum
        )

        # Build the lr scheduler
        lr_scheduler = config.optim.lr_scheduler.build(
            optimizers=self.optimizer_container,
            training_steps=config.training.steps,
        )

        # Step through all steps
        for _ in range(10):
            lr_scheduler.step()

        # After all steps, LR should be at minimum (0.1 * 0.2 = 0.02)
        self.assertAlmostEqual(self.optimizer.param_groups[0]["lr"], 0.02, places=6)

    def test_steps_beyond_total_steps_hold_final_lr(self):
        """Training past lr_scheduler.total_steps keeps the final LR instead of
        decaying below min_lr_factor (or rising again for cosine), and leaves the
        steps inside the schedule unchanged."""

        def run(decay_type, decay_ratio, training_steps, total_steps):
            optimizer = Adam(torch.nn.Linear(10, 10).parameters(), lr=0.1)
            optimizer._opt_called = True
            container = MagicMock(spec=OptimizersContainer)
            container.__iter__.return_value = iter([optimizer])
            container.__len__.return_value = 1
            config = self.create_trainer_config(
                training_steps=training_steps,
                warmup_steps=2,
                decay_ratio=decay_ratio,
                decay_type=decay_type,
                min_lr_factor=0.1,
            )
            config.optim.lr_scheduler.total_steps = total_steps
            lr_scheduler = config.optim.lr_scheduler.build(
                optimizers=container, training_steps=config.training.steps
            )
            lrs = []
            for _ in range(training_steps):
                lr_scheduler.step()
                lrs.append(optimizer.param_groups[0]["lr"])
            return lrs

        for decay_type, decay_ratio, final_lr in (
            ("linear", None, 0.01),
            ("sqrt", None, 0.01),
            ("cosine", None, 0.01),
            # No decay phase: stays at the base LR instead of failing an assert.
            ("linear", 0.0, 0.1),
        ):
            with self.subTest(decay_type=decay_type, decay_ratio=decay_ratio):
                lrs = run(decay_type, decay_ratio, training_steps=10, total_steps=6)
                within = run(
                    decay_type, decay_ratio, training_steps=6, total_steps=None
                )
                for lr, expected in zip(lrs[:6], within):
                    self.assertAlmostEqual(lr, expected, places=6)
                for lr in lrs[5:]:
                    self.assertAlmostEqual(lr, final_lr, places=6)

    def test_warmup_exceeds_training(self):
        """Test when warmup steps exceed training steps."""
        # Create a job config where warmup steps > training steps
        config = self.create_trainer_config(
            training_steps=5,
            warmup_steps=10,  # More than training steps
            decay_ratio=None,
            decay_type="linear",
            min_lr_factor=0.0,
        )

        # Build the lr scheduler - should adjust warmup steps
        lr_scheduler = config.optim.lr_scheduler.build(
            optimizers=self.optimizer_container,
            training_steps=config.training.steps,
        )

        # Expected adjustment factors for each step
        expected_factors = [
            0.2,  # Step 0: 50% of max LR (warmup)
            0.4,  # Step 1: 100% of max LR (warmup complete)
            0.6,  # Step 2: Stable phase
            0.8,  # Step 3: Stable phase
            1.0,  # Step 4: Stable phase
        ]

        # Check the learning rate at each step
        for i, factor in enumerate(expected_factors):
            expected_lr = 0.1 * factor
            self.assertAlmostEqual(
                self.optimizer.param_groups[0]["lr"],
                expected_lr,
                places=6,
                msg=f"Step {i}: Expected LR {expected_lr}, got {self.optimizer.param_groups[0]['lr']}",
            )
            lr_scheduler.step()

    def test_warmup_stable_only(self):
        """Test warmup followed by stable phase only, with no decay phase."""
        # Create a job config with 10 steps, 2 warmup steps, and no decay phase
        config = self.create_trainer_config(
            training_steps=10,
            warmup_steps=2,
            decay_ratio=0.0,  # 0% of steps for decay (no decay)
            decay_type="linear",
            min_lr_factor=0.0,
        )

        # Build the lr scheduler
        lr_scheduler = config.optim.lr_scheduler.build(
            optimizers=self.optimizer_container,
            training_steps=config.training.steps,
        )

        # Expected adjustment factors for each step
        expected_factors = [
            0.5,  # Step 0: 50% of max LR (warmup)
            1.0,  # Step 1: 100% of max LR (warmup complete)
            1.0,  # Step 2: We maunally added step of stable phase, to prevent LR from dropping to 0 at last step
            1.0,  # Step 3: Stable phase
            1.0,  # Step 4: Stable phase
            1.0,  # Step 5: Stable phase
            1.0,  # Step 6: Stable phase
            1.0,  # Step 7: Stable phase
            1.0,  # Step 8: Stable phase
            1.0,  # Step 9: Stable phase
        ]

        # Check the learning rate at each step
        for i, factor in enumerate(expected_factors):
            expected_lr = 0.1 * factor
            self.assertAlmostEqual(
                self.optimizer.param_groups[0]["lr"],
                expected_lr,
                places=6,
                msg=f"Step {i}: Expected LR {expected_lr}, got {self.optimizer.param_groups[0]['lr']}",
            )
            lr_scheduler.step()

    def test_warmup_plus_decay_exceeds_training(self):
        """Test when warmup + decay steps exceed training steps."""
        # Create a job config where warmup + decay steps > training steps
        # Expected behavior: warmup steps = 5, decay steps = 5
        config = self.create_trainer_config(
            training_steps=10,
            warmup_steps=5,
            decay_ratio=0.8,  # 80% of steps for decay (8 steps)
            decay_type="linear",
            min_lr_factor=0.0,
        )

        # Build the lr scheduler - should adjust warmup steps
        lr_scheduler = config.optim.lr_scheduler.build(
            optimizers=self.optimizer_container,
            training_steps=config.training.steps,
        )

        # Expected adjustment factors for each step
        expected_factors = [
            0.2,  # Step 0: 50% of max LR (warmup)
            0.4,  # Step 1: 100% of max LR (warmup complete)
            0.6,  # Step 2: Stable phase
            0.8,  # Step 3: Stable phase
            1.0,  # Step 4: Stable phase
            1.0,  # Step 5: We maunally added step of stable phase, to prevent LR from dropping to 0 at last step
            0.8,  # Step 6: Linear decay starts (80% of max LR)
            0.6,  # Step 7: 60% of max LR
            0.4,  # Step 8: 40% of max LR
            0.2,  # Step 9: 20% of max LR
        ]

        # Check the learning rate at each step
        for i, factor in enumerate(expected_factors):
            expected_lr = 0.1 * factor
            self.assertAlmostEqual(
                self.optimizer.param_groups[0]["lr"],
                expected_lr,
                places=6,
                msg=f"Step {i}: Expected LR {expected_lr}, got {self.optimizer.param_groups[0]['lr']}",
            )
            lr_scheduler.step()

    def test_config_rejects_out_of_range_values(self):
        # Each of these used to build silently: a negative min_lr_factor drives
        # the LR below zero, and a negative decay_ratio or warmup_steps distorts
        # the schedule without any warning.
        invalid = [
            {"warmup_steps": -1},
            {"total_steps": 0},
            {"decay_ratio": -0.5},
            {"decay_ratio": 1.5},
            {"min_lr_factor": -0.5},
            {"min_lr_factor": 1.5},
        ]
        for kwargs in invalid:
            with self.subTest(**kwargs):
                with self.assertRaises(ValueError):
                    LRSchedulersContainer.Config(**kwargs)


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