#!/usr/bin/env python3
"""Focused algebra tests for hierarchical3_muon_bbp.py."""

from __future__ import annotations

import unittest

import numpy as np
import torch

from experiments.brainstorm.resnetl3.hierarchical3_muon_bbp import (
    Config,
    exact_svd_power,
    link_hessian_blocks,
    make_setup,
    new_model,
    teacher_on,
)
from experiments.muon.lowrank_block_wishart_spectral_experiments import block_wishart


class Hierarchical3MuonBBPTest(unittest.TestCase):
    def setUp(self) -> None:
        self.cfg = Config(steps=1)
        self.setup = make_setup(self.cfg)
        self.model = new_model(self.cfg, self.setup, self.cfg.seed)

    def test_link_hessian_matches_autograd(self) -> None:
        generator = torch.Generator().manual_seed(11)
        z = torch.randn(3, self.cfg.k, generator=generator, dtype=torch.float64)
        target = teacher_on(z, self.setup)
        analytic = link_hessian_blocks(
            self.model, z.numpy(), target.detach().numpy()
        )

        for index in range(len(z)):
            h0 = torch.cat(
                [
                    (z[index] @ self.model.V.T + self.model.b).reshape(-1),
                    (z[index] @ self.model.second_weight().T).reshape(-1),
                ]
            ).detach().requires_grad_(True)

            def loss(h: torch.Tensor) -> torch.Tensor:
                h1 = h[: self.cfg.p]
                hu = h[self.cfg.p :]
                qfeat = h1.square() - 1.0
                t2 = hu + self.model.A @ qfeat + self.model.c
                prediction = (
                    self.model.a0
                    + self.model.alpha @ qfeat
                    + self.model.beta @ (t2.square() - 1.0)
                )
                return 0.5 * (prediction - target[index]).square()

            autodiff = torch.autograd.functional.hessian(loss, h0).detach().numpy()
            np.testing.assert_allclose(analytic[index], autodiff, rtol=2e-11, atol=2e-11)

    def test_sign_svd_has_flat_nonzero_spectrum(self) -> None:
        gradient = torch.tensor(
            [[3.0, -1.0, 0.5], [0.2, 2.0, -0.7]], dtype=torch.float64
        )
        update = exact_svd_power(gradient, power=0.0, eps=1e-8)
        singular = torch.linalg.svdvals(update)
        self.assertTrue(torch.allclose(singular, singular[0].expand_as(singular), atol=2e-6))
        self.assertAlmostEqual(float(update.square().mean().sqrt()), 1.0, places=6)

    def test_orthogonal_hessian_is_principal_block_wishart(self) -> None:
        rng = np.random.default_rng(7)
        n, d, k = 19, 8, 3
        r = self.cfg.p + self.cfg.q
        x = rng.standard_normal((n, d))
        raw = rng.standard_normal((n, r, r))
        blocks = 0.5 * (raw + raw.transpose(0, 2, 1))
        full = block_wishart(x, blocks)
        orthogonal = block_wishart(x[:, k:], blocks)
        selector = np.zeros((d, d - k))
        selector[k:, :] = np.eye(d - k)
        lifted_selector = np.kron(np.eye(r), selector)
        projected = lifted_selector.T @ full @ lifted_selector
        np.testing.assert_allclose(projected, orthogonal, rtol=2e-13, atol=2e-13)


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