Q-TensorFormer / tests /test_advanced_features.py
Premchandyadav369
Implement top-tier research features: PID Dual Controller, GQA, Roofline Analyzer, Early-Exit, Meyer-Wallach Entanglement, and Interactive Visual Dashboard
0431133
Raw History Blame Contribute Delete
4.9 kB
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!")