Q-Prefer-D2 / code /Q-Prefer /tests /test_training.py
qgfvadfuvads's picture
Upload Q-Prefer training and inference code
aa7758f verified
Raw
History Blame Contribute Delete
3.71 kB
import json
import math
import tempfile
import unittest
from pathlib import Path
import torch
import torch.nn.functional as F
from qprefer_reward.training.data import (
QPreferPairCollator,
load_manifest,
parse_path_prefix_maps,
resolve_media_path,
)
from qprefer_reward.training.loss import rao_kupper_loss
class TrainingLossTest(unittest.TestCase):
def test_rao_kupper_decisive_tie_and_mask(self):
rewards_a = torch.tensor([[1.0, 0.0, 7.0]])
rewards_b = torch.tensor([[0.0, 0.0, -3.0]])
labels = torch.tensor([[1, 0, 22]])
observed = rao_kupper_loss(rewards_a, rewards_b, labels, k=5.0)
log_k = math.log(5.0)
decisive = -F.logsigmoid(torch.tensor(1.0 - log_k))
tie = (
-F.logsigmoid(torch.tensor(-log_k)) * 2
- math.log(5.0**2 - 1.0)
)
expected = (decisive + tie) / 3
self.assertTrue(torch.allclose(observed, expected))
def test_unit_sample_weights_preserve_original_scale(self):
generator = torch.Generator().manual_seed(7)
rewards_a = torch.randn(4, 3, generator=generator)
rewards_b = torch.randn(4, 3, generator=generator)
labels = torch.tensor([[1, 22, 0], [-1, 22, 1], [0, 22, -1], [1, 22, 1]])
unweighted = rao_kupper_loss(rewards_a, rewards_b, labels)
weighted = rao_kupper_loss(
rewards_a,
rewards_b,
labels,
sample_weight=torch.ones(4),
)
self.assertTrue(torch.allclose(unweighted, weighted))
class TrainingManifestTest(unittest.TestCase):
def test_relative_media_paths_and_summary(self):
with tempfile.TemporaryDirectory() as temporary:
root = Path(temporary)
media = root / "media"
media.mkdir()
for name in ("a.mp4", "b.mp4", "ref.jpg"):
(media / name).touch()
manifest = root / "pairs.json"
manifest.write_text(
json.dumps(
[
{
"pair_id": "pair-1",
"task": "i2v",
"bucket": "near",
"source_version": "unit",
"prompt": "A test prompt",
"path_A": "a.mp4",
"path_B": "b.mp4",
"ref_image": "ref.jpg",
"chosen_label": [1, 22, 0],
}
]
)
)
rows, summary = load_manifest(manifest, media_root=media, check_media=True)
self.assertEqual(rows[0]["path_A"], str((media / "a.mp4").resolve()))
self.assertEqual(summary["rows"], 1)
self.assertEqual(summary["labels"]["MQ"], {"22": 1})
def test_training_prompt_is_the_inference_contract(self):
prompt = QPreferPairCollator.reward_prompt("A red fox runs.")
self.assertIn("Visual Quality: <|VQ_reward|>", prompt)
self.assertIn("Motion Quality: <|MQ_reward|>", prompt)
self.assertIn("Text/Image Alignment: <|TA_reward|>", prompt)
self.assertTrue(prompt.endswith("Textual prompt - A red fox runs.\n"))
def test_absolute_paths_can_be_relocated_without_editing_manifest(self):
mappings = parse_path_prefix_maps(["/old/project=/new/dataset"])
observed = resolve_media_path(
"/old/project/videos/a.mp4",
Path("/tmp/manifest.json"),
None,
mappings,
)
self.assertEqual(observed, "/new/dataset/videos/a.mp4")
if __name__ == "__main__":
unittest.main()