File size: 503 Bytes
4fd79a1
 
 
 
 
 
 
1
2
3
4
5
6
7
8
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