Self-Forcing / tests /test_vbench8_protocol.py
Cccccz's picture
Upload tests
ebad435 verified
Raw History Blame Contribute Delete
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()