File size: 4,177 Bytes
ebad435
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
#!/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()