File size: 2,273 Bytes
aa7758f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Run a self-contained forward pass with an eight-frame synthetic video."""

from __future__ import annotations

import argparse
import math

import torch

from qprefer_reward import QPreferConfig, QPreferScorer
from qprefer_reward.constants import BASE_MODEL_ID, BASE_MODEL_REVISION


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--adapter", required=True)
    parser.add_argument("--adapter-revision")
    parser.add_argument("--base-model", default=BASE_MODEL_ID)
    parser.add_argument("--base-revision", default=BASE_MODEL_REVISION)
    parser.add_argument("--device", default="cuda")
    parser.add_argument("--expected-vq", type=float)
    parser.add_argument("--expected-ta", type=float)
    parser.add_argument("--tolerance", type=float, default=0.05)
    return parser.parse_args()


def check_close(name: str, observed: float, expected: float | None, tolerance: float) -> None:
    if not math.isfinite(observed):
        raise RuntimeError(f"{name} is not finite: {observed}")
    if expected is not None and abs(observed - expected) > tolerance:
        raise RuntimeError(
            f"{name} mismatch: expected={expected}, observed={observed}, tolerance={tolerance}"
        )


def main() -> None:
    args = parse_args()
    scorer = QPreferScorer(
        QPreferConfig(
            adapter=args.adapter,
            adapter_revision=args.adapter_revision,
            base_model=args.base_model,
            base_revision=args.base_revision,
            device=args.device,
            dtype=torch.bfloat16,
        )
    )
    black_video = torch.zeros(8, 3, 224, 224, dtype=torch.float32)
    scores = scorer.score_batch(
        [black_video],
        ["A static black frame."],
        tensor_value_range="zero_one",
    )
    visual = float(scores.visual_quality[0])
    alignment = float(scores.text_alignment[0])
    check_close("visual_quality", visual, args.expected_vq, args.tolerance)
    check_close("text_alignment", alignment, args.expected_ta, args.tolerance)
    print("Q-Prefer forward smoke test passed")
    print(f"visual_quality={visual:.8f}")
    print(f"text_alignment={alignment:.8f}")


if __name__ == "__main__":
    main()