import torch from spectral_world_models.models import build_model,StructuredSpectralTransportTransition def test_structured_model_forward(): m=build_model('swm_structured'); x=torch.rand(2,1,32,32); tok=torch.ones(2,5,dtype=torch.long); a=torch.tensor([1,4]); o=m(x,tok,a); assert o['image_pred'].shape==(2,1,32,32) def test_transport_is_orthogonal(): t=StructuredSpectralTransportTransition(24,5,rank=4); q=t.transport_matrix(); eye=torch.eye(24); assert torch.max(torch.abs(q.T@q-eye)).item()<1e-4