Download tests/test_vbench8_protocol.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 4.18 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/tests/test_vbench8_protocol.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/tests/test_vbench8_protocol.py
-
curl -L -o test_vbench8_protocol.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/tests/test_vbench8_protocol.py
4.18 kB
| #!/usr/bin/env python3 | |
| import unittest | |
| from scripts.build_vbench8_extended_mapping import build_mapping | |
| from scripts.summarize_vbench8_generation import ( | |
| denoise_dit_latency_ms, | |
| policy_latency_ms, | |
| speedup_percent, | |
| ) | |
| from scripts.vbench8_protocol import aggregate_selected_score, normalize_scores | |
| class VBench8ProtocolTest(unittest.TestCase): | |
| def test_mapping_counts_and_extended_prompt_alignment(self) -> None: | |
| short = [f"short prompt {index}" for index in range(946)] | |
| extended = [f"extended prompt {index}" for index in range(946)] | |
| info = [{"prompt_en": prompt, "dimension": ["temporal_style"]} for prompt in short] | |
| positions = { | |
| "subject_consistency": range(0, 72), | |
| "overall_consistency": range(72, 165), | |
| "scene": range(165, 251), | |
| } | |
| for suite, indices in positions.items(): | |
| for index in indices: | |
| info[index] = { | |
| "prompt_en": short[index], | |
| "dimension": [suite], | |
| } | |
| if suite == "scene": | |
| info[index]["auxiliary_info"] = { | |
| "scene": {"scene": {"scene": f"scene-{index}"}} | |
| } | |
| mapping = build_mapping( | |
| short_prompts=short, | |
| extended_prompts=extended, | |
| vbench_info=info, | |
| ) | |
| self.assertEqual(len(mapping), 251) | |
| self.assertEqual( | |
| {suite: sum(row["prompt_suite"] == suite for row in mapping) | |
| for suite in positions}, | |
| {"subject_consistency": 72, "overall_consistency": 93, "scene": 86}, | |
| ) | |
| self.assertEqual(mapping[0]["original_prompt"], "short prompt 0") | |
| self.assertEqual(mapping[0]["extended_prompt"], "extended prompt 0") | |
| self.assertIn("auxiliary_info", mapping[-1]) | |
| def test_mapping_rejects_order_mismatch(self) -> None: | |
| short = [f"short prompt {index}" for index in range(946)] | |
| extended = [f"extended prompt {index}" for index in range(946)] | |
| info = [{"prompt_en": prompt, "dimension": ["temporal_style"]} for prompt in short] | |
| info[10]["prompt_en"] = "wrong order" | |
| with self.assertRaises(ValueError): | |
| build_mapping(short_prompts=short, extended_prompts=extended, vbench_info=info) | |
| def test_normalization_motion(self) -> None: | |
| raw = {dimension: 0.0 for dimension in ( | |
| "subject_consistency", "background_consistency", "motion_smoothness", | |
| "dynamic_degree", "aesthetic_quality", "imaging_quality", "scene", | |
| "overall_consistency", | |
| )} | |
| raw["motion_smoothness"] = 0.95 | |
| normalized = normalize_scores(raw) | |
| expected = (0.95 - 0.7060) / (0.9975 - 0.7060) | |
| self.assertAlmostEqual(normalized["motion_smoothness"], expected) | |
| def test_aggregate_all_normalized_one(self) -> None: | |
| raw = { | |
| "subject_consistency": 1.0, | |
| "background_consistency": 1.0, | |
| "motion_smoothness": 0.9975, | |
| "dynamic_degree": 1.0, | |
| "aesthetic_quality": 1.0, | |
| "imaging_quality": 1.0, | |
| "scene": 0.8222, | |
| "overall_consistency": 0.3640, | |
| } | |
| aggregate = aggregate_selected_score(raw) | |
| self.assertAlmostEqual(aggregate["quality_score"], 1.0) | |
| self.assertAlmostEqual(aggregate["semantic_score"], 1.0) | |
| self.assertAlmostEqual(aggregate["selected_vbench_score"], 1.0) | |
| self.assertAlmostEqual(aggregate["selected_vbench_percent"], 100.0) | |
| def test_latency_excludes_context_and_includes_confidence(self) -> None: | |
| generation = { | |
| "full_dit_time_ms": 100.0, | |
| "predictor_time_ms": 20.0, | |
| "confidence_head_time_ms": 3.0, | |
| "context_dit_time_ms": 500.0, | |
| } | |
| self.assertEqual(denoise_dit_latency_ms(generation), 120.0) | |
| self.assertEqual(policy_latency_ms(generation), 123.0) | |
| def test_policy_latency_speedup_uses_ratio_of_means(self) -> None: | |
| self.assertAlmostEqual(speedup_percent([60.0, 100.0], [100.0, 100.0]), 20.0) | |
| if __name__ == "__main__": | |
| unittest.main() | |