Download tests/test_v8_generalization.py from kiruluta/Spectral-World-Models-Reproducibility: direct link, hf CLI and curl.
- Browser
- Download file 891 Bytes
-
https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/tests/test_v8_generalization.py
- Command line
-
hf download hf://kiruluta/Spectral-World-Models-Reproducibility/tests/test_v8_generalization.py
-
curl -L -o test_v8_generalization.py https://huggingface.co/kiruluta/Spectral-World-Models-Reproducibility/resolve/main/tests/test_v8_generalization.py
891 Bytes
| import torch | |
| from spectral_world_models.models import build_model | |
| from spectral_world_models.v8_generalization import OOD_CONFIG,_candidate_actions,evaluate_planning,evaluate_counterfactual_cfg | |
| def test_ood_is_shifted_from_training_distribution(): | |
| assert OOD_CONFIG.acceleration > .70 and OOD_CONFIG.max_speed > 3.0 and OOD_CONFIG.occlusion_prob > .12 | |
| def test_candidate_library_has_interventions(): | |
| c=_candidate_actions(8); assert [3]*8 in c and [4]*8 in c and len(c)>=9 | |
| def test_v8_metrics_smoke(): | |
| m=build_model('swm_structured_cf') | |
| p=evaluate_planning(m,torch.device('cpu'),n=2,horizon=2,seed=1) | |
| assert 0 <= p['planning_oracle_agreement'] <= 1 and p['planning_regret'] >= -1e-8 | |
| c=evaluate_counterfactual_cfg(m,torch.device('cpu'),horizon=2,n=2,seed=2) | |
| assert set(c)=={'ood_cf_final_image_mse','ood_cf_effect_vector_cosine','ood_cf_effect_magnitude_ratio'} | |