import sys import os sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))) import torch import pytest from src.resource_allocator import PIDDualSubgradientController from src.hardware_cost_model import HardwareRooflineAnalyzer from src.attention import MultiHeadAttention from src.blocks import HybridBlock from src.quantum_backend import compute_meyer_wallach_entanglement, compute_quantum_expressibility from src.hf_model import QTensorFormerConfig, QTensorFormerForCausalLM def test_pid_dual_subgradient_controller(): controller = PIDDualSubgradientController( target_latency_ms=5.0, target_energy_uj=10.0, kp=0.1, ki=0.01, kd=0.005, ) # Initial multipliers b0 = controller.get_budget() assert b0.lambda_latency == 1.0 # Simulate observed latency higher than target (e.g. 8.0 ms > 5.0 ms) # The controller should penalize latency more heavily (increase lambda_latency) for _ in range(5): b = controller.update(measured_latency_ms=8.0, measured_energy_uj=12.0) assert b.lambda_latency > 1.0, f"Expected lambda_latency to increase, got {b.lambda_latency}" assert b.lambda_energy > 0.2, f"Expected lambda_energy to increase, got {b.lambda_energy}" # Simulate observed latency lower than target (e.g. 3.0 ms < 5.0 ms) for _ in range(15): b = controller.update(measured_latency_ms=3.0, measured_energy_uj=7.0) assert b.lambda_latency < 1.5, f"Expected lambda_latency to relax, got {b.lambda_latency}" print("[SUCCESS] test_pid_dual_subgradient_controller passed") def test_hardware_roofline_analyzer(): analyzer_a100 = HardwareRooflineAnalyzer("gpu_a100") # Low arithmetic intensity (Memory-Bound) res_mem = analyzer_a100.analyze("TTLinear_r2", flops=1000, memory_traffic_bytes=5000) assert res_mem.arithmetic_intensity == 0.2 assert res_mem.regime == "Memory-Bound" # High arithmetic intensity (Compute-Bound) res_comp = analyzer_a100.analyze("DenseLinear", flops=500000, memory_traffic_bytes=2000) assert res_comp.arithmetic_intensity == 250.0 assert res_comp.regime == "Compute-Bound" # Test Apple M2 analyzer_m2 = HardwareRooflineAnalyzer("apple_m2") res_m2 = analyzer_m2.analyze("Attention_SDPA", flops=20000, memory_traffic_bytes=4000) assert res_m2.hardware_name == "Apple M2 Max" print("[SUCCESS] test_hardware_roofline_analyzer passed") def test_grouped_query_attention(): # 8 query heads, 2 KV heads (4:1 GQA) d_model = 64 n_heads = 8 n_kv_heads = 2 gqa = MultiHeadAttention(d_model=d_model, n_heads=n_heads, n_kv_heads=n_kv_heads) x = torch.randn(2, 16, d_model) out, weights = gqa(x) assert out.shape == (2, 16, d_model) assert not torch.isnan(out).any() # Verify MQA (1 KV head) mqa = MultiHeadAttention(d_model=d_model, n_heads=n_heads, n_kv_heads=1) out_mqa, _ = mqa(x) assert out_mqa.shape == (2, 16, d_model) print("[SUCCESS] test_grouped_query_attention passed") def test_early_exit_hybrid_block(): d_model = 64 vocab_size = 500 block = HybridBlock( d_model=d_model, n_heads=4, vocab_size=vocab_size, enable_early_exit=True, early_exit_threshold=0.99, # High threshold to trigger early exit ) x = torch.randn(2, 8, d_model) out, stats = block(x) assert stats.get("early_exit_triggered", False) is True assert "early_exit_logits" in stats assert stats["early_exit_logits"].shape == (2, 8, vocab_size) print("[SUCCESS] test_early_exit_hybrid_block passed") def test_quantum_metrics(): # 1. Product state: |00> -> Meyer-Wallach should be 0.0 psi_product = torch.tensor([1.0, 0.0, 0.0, 0.0]) q_prod = compute_meyer_wallach_entanglement(psi_product) assert q_prod == 0.0, f"Expected 0.0 for product state, got {q_prod}" # 2. Bell state: (|00> + |11>) / sqrt(2) -> Meyer-Wallach should be 1.0 psi_bell = torch.tensor([1.0 / (2**0.5), 0.0, 0.0, 1.0 / (2**0.5)]) q_bell = compute_meyer_wallach_entanglement(psi_bell) assert abs(q_bell - 1.0) < 1e-4, f"Expected 1.0 for Bell state, got {q_bell}" # 3. Expressibility test def circuit_dummy(theta): # 2 qubits ang = theta.squeeze(0) return torch.tensor([torch.cos(ang[0]), torch.sin(ang[0]), torch.cos(ang[1]), torch.sin(ang[1])]) expr = compute_quantum_expressibility(circuit_dummy, n_qubits=2, n_samples=50) assert "expressibility_kl" in expr assert expr["expressibility_kl"] >= 0.0 print("[SUCCESS] test_quantum_metrics passed") if __name__ == "__main__": test_pid_dual_subgradient_controller() test_hardware_roofline_analyzer() test_grouped_query_attention() test_early_exit_hybrid_block() test_quantum_metrics() print("All Advanced Feature Tests Passed Successfully!")