Q-TensorFormer / tests /test_nested_tt.py
Premchandyadav369
Transform Q-TensorFormer into an Information-Value Adaptive Resource Allocation Architecture
eaeea8f
Raw History Blame Contribute Delete
1.27 kB
"""
Tests for Nested-Core Tensor-Train Layer.
"""
import sys
import os
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import torch
import pytest
from src.tensor_layers import TTLinear, TTFeedForward
def test_nested_tt_slicing_shapes():
in_dim, out_dim = 64, 128
layer = TTLinear(in_dim, out_dim, max_rank=8)
x = torch.randn(2, 4, in_dim)
# Test all candidate ranks
for r in [1, 2, 4, 8]:
layer.set_rank(r)
assert layer.rank == r
out = layer(x)
assert out.shape == (2, 4, out_dim), f"Expected (2, 4, {out_dim}), got {out.shape}"
assert not torch.isnan(out).any(), f"NaN in output at rank {r}"
traffic = layer.get_memory_traffic()
assert traffic["total_bytes"] > 0
print("✓ test_nested_tt_slicing_shapes passed")
def test_nested_tt_ffn():
ffn = TTFeedForward(hidden_dim=64, ff_multiplier=4, rank=8)
x = torch.randn(2, 64)
for r in [1, 2, 4, 8]:
ffn.set_rank(r)
out = ffn(x)
assert out.shape == (2, 64)
assert ffn.active_params > 0
print("✓ test_nested_tt_ffn passed")
if __name__ == "__main__":
test_nested_tt_slicing_shapes()
test_nested_tt_ffn()
print("All Nested TT tests passed!")