Download tests/test_v6_structured.py from kiruluta/Spectral-World-Models-Reproducibility: direct link, hf CLI and curl.
- Browser
- Download file 503 Bytes
-
https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/tests/test_v6_structured.py
- Command line
-
hf download hf://kiruluta/Spectral-World-Models-Reproducibility/tests/test_v6_structured.py
-
curl -L -o test_v6_structured.py https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/tests/test_v6_structured.py
503 Bytes
| 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 | |