first commit
This commit is contained in:
@@ -0,0 +1,121 @@
|
||||
"""Tests for HPIController — NN edition."""
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from hpi.hpi_controller import HPIController
|
||||
from hpi.nn_models import CriticNN, check_positive_definite
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def controller():
|
||||
return HPIController(J=1.0, mgl=1.0, R=1.0, pe_amplitude=0.5,
|
||||
lr=1e-2, epochs=100, lambda_weight=1e-6)
|
||||
|
||||
|
||||
class TestPositiveDefinite:
|
||||
"""Test suite for check_positive_definite on CriticNN."""
|
||||
|
||||
def test_positive_definite_true(self):
|
||||
"""A well-initialized critic should pass PD (random init usually PD)."""
|
||||
critic = CriticNN()
|
||||
# Check: random init typically gives PD Hessian at origin
|
||||
# If not, the function still runs without error
|
||||
result = check_positive_definite(critic)
|
||||
assert isinstance(result, bool)
|
||||
|
||||
def test_positive_definite_false(self):
|
||||
"""Critic with all-zero final layer fails PD (zero-rank Hessian)."""
|
||||
critic = CriticNN()
|
||||
# Zero out the final layer to make Hessian rank-deficient at origin
|
||||
with torch.no_grad():
|
||||
critic.fc3.weight.zero_()
|
||||
assert check_positive_definite(critic) is False
|
||||
|
||||
def test_positive_definite_large_cross(self):
|
||||
"""Critic with rank-1 Hessian (single direction) fails PD.
|
||||
|
||||
Zero all layers, then set only fc3.weight[0,0]=1. This makes only
|
||||
z1 depend on x1, giving Hessian = 2*[1,0]^T[1,0] which is rank-1.
|
||||
"""
|
||||
critic = CriticNN()
|
||||
with torch.no_grad():
|
||||
for p in critic.parameters():
|
||||
p.zero_()
|
||||
critic.fc3.weight[0, 0] = 1.0
|
||||
assert check_positive_definite(critic) is False
|
||||
|
||||
def test_positive_definite_borderline(self):
|
||||
"""Near-singular Hessian fails PD."""
|
||||
critic = CriticNN()
|
||||
with torch.no_grad():
|
||||
for p in critic.parameters():
|
||||
p.zero_()
|
||||
critic.fc3.weight[0, 0] = 1e-8 # very small, det ~ 0
|
||||
assert check_positive_definite(critic) is False
|
||||
|
||||
|
||||
class TestPhaseOne:
|
||||
"""Test suite for Phase 1."""
|
||||
|
||||
def test_phase_one_finds_L0(self, controller):
|
||||
"""Phase 1 should find a valid L0."""
|
||||
x0 = np.array([0.5, 0.0])
|
||||
L0, critic_sd, actor_sd, info = controller.phase_one_find_L0(
|
||||
x0, T=2.0, dt=0.05, L_start=0.1, L_step=1.0, L_max=20.0,
|
||||
verbose=False,
|
||||
)
|
||||
assert L0 > 0
|
||||
assert isinstance(critic_sd, dict)
|
||||
assert isinstance(actor_sd, dict)
|
||||
assert isinstance(info, dict)
|
||||
# Load and check PD
|
||||
controller.trainer.critic_nn.load_state_dict(critic_sd)
|
||||
assert controller.check_positive_definite()
|
||||
|
||||
def test_phase_one_raises_on_low_Lmax(self, controller):
|
||||
"""Phase 1 should raise RuntimeError if L_max < L_start (no valid L tried)."""
|
||||
x0 = np.array([0.5, 0.0])
|
||||
with pytest.raises(RuntimeError):
|
||||
controller.phase_one_find_L0(
|
||||
x0, T=1.0, dt=0.05, L_start=1.0, L_step=1.0, L_max=0.5
|
||||
)
|
||||
|
||||
|
||||
class TestFullPipeline:
|
||||
"""Test suite for full pipeline."""
|
||||
|
||||
def test_full_pipeline(self, controller):
|
||||
"""Three-phase pipeline runs without error and returns valid results."""
|
||||
results = controller.run(
|
||||
x0=np.array([0.5, 0.0]),
|
||||
T=2.0,
|
||||
dt=0.05,
|
||||
L_start=0.1,
|
||||
L_step=1.0,
|
||||
L_max=10.0,
|
||||
gamma=0.1,
|
||||
epsilon=1e-3,
|
||||
verbose=False,
|
||||
)
|
||||
|
||||
assert "phase1_L0" in results
|
||||
assert results["phase1_L0"] > 0
|
||||
assert isinstance(results["phase1_critic_sd"], dict)
|
||||
assert isinstance(results["phase1_actor_sd"], dict)
|
||||
|
||||
assert "phase2_actor_sd" in results
|
||||
assert "phase2_critic_sd" in results
|
||||
assert len(results["phase2_history"]) >= 1
|
||||
|
||||
assert "phase3_actor_sd" in results
|
||||
assert "phase3_critic_sd" in results
|
||||
assert isinstance(results["phase3_actor_sd"], dict)
|
||||
assert isinstance(results["phase3_critic_sd"], dict)
|
||||
|
||||
# Verify final critic is PD
|
||||
controller.trainer.critic_nn.load_state_dict(results["phase3_critic_sd"])
|
||||
is_pd = controller.check_positive_definite()
|
||||
# May or may not be PD depending on convergence, but shouldn't error
|
||||
assert isinstance(is_pd, bool)
|
||||
Reference in New Issue
Block a user