ldov
/

File size: 3,822 Bytes
e46c127
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Small CPU checks for shared-prefix inference; no checkpoint download needed."""
import unittest
from unittest.mock import patch

import numpy as np
import torch
from transformers import (Qwen3_5Config, Qwen3_5ForSequenceClassification,
                          Qwen3_5TextConfig, Qwen3_5VisionConfig)

from modeling_openjev import OpenJevCrossEncoder


class CharacterTokenizer:
    pad_token_id = 0

    def __call__(self, texts, truncation=True, max_length=4096, padding=False, return_tensors=None):
        if isinstance(texts, str):
            return {"input_ids": [ord(c) + 1 for c in texts][:max_length]}
        rows = [[ord(c) + 1 for c in text][:max_length] for text in texts]
        width = max(map(len, rows))
        ids = [row + [0] * (width - len(row)) for row in rows]
        mask = [[1] * len(row) + [0] * (width - len(row)) for row in rows]
        return {"input_ids": torch.tensor(ids), "attention_mask": torch.tensor(mask)}


class SharedPrefixTest(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        torch.manual_seed(7)
        config = Qwen3_5TextConfig(
            vocab_size=256, hidden_size=32, intermediate_size=64,
            num_hidden_layers=4, num_attention_heads=4, num_key_value_heads=2,
            head_dim=8, linear_num_key_heads=2, linear_num_value_heads=4,
            linear_key_head_dim=8, linear_value_head_dim=8,
            layer_types=["linear_attention", "full_attention"] * 2,
            rope_parameters={"rope_type": "default", "rope_theta": 10000,
                             "partial_rotary_factor": 0.25, "mrope_section": [1, 0, 0]},
            pad_token_id=0, num_labels=3, use_cache=False,
        )
        vision = Qwen3_5VisionConfig(depth=1, hidden_size=32, intermediate_size=64,
                                    num_heads=4, out_hidden_size=32)
        model = Qwen3_5ForSequenceClassification(
            Qwen3_5Config(text_config=config, vision_config=vision, num_labels=3,
                           pad_token_id=0)).eval()
        cls.jev = OpenJevCrossEncoder.__new__(OpenJevCrossEncoder)
        cls.jev.model = model
        cls.jev.backbone = model.model
        cls.jev.tok = CharacterTokenizer()
        cls.jev.template = "Premise: {premise}\nHypothesis: {hypothesis}"
        cls.jev.device = "cpu"
        cls.jev.bs = 32
        cls.jev.max_len = 128

    def test_shared_matches_independent_and_isolates_branches(self):
        premise = "A cat sleeps."
        hypotheses = ["A cat rests.", "The dog runs.", "A cat."]
        expected = self.jev.latents([(premise, h) for h in hypotheses])
        actual = self.jev.latents_hypotheses(premise, hypotheses)
        np.testing.assert_allclose(actual, expected, atol=2e-4, rtol=2e-4)
        expected_probs = self.jev.predict([(premise, h) for h in hypotheses])
        actual_probs = self.jev.predict_hypotheses(premise, hypotheses)
        np.testing.assert_allclose(actual_probs, expected_probs, atol=2e-4, rtol=2e-4)

        changed = hypotheses.copy()
        changed[1] = "Birds fly around today."
        other = self.jev.latents_hypotheses(premise, changed)
        np.testing.assert_allclose(other[[0, 2]], actual[[0, 2]], atol=2e-4, rtol=2e-4)

    def test_small_option_sets_use_pairwise_batch(self):
        for hypotheses in (["B"], ["B", "C"]):
            with self.subTest(count=len(hypotheses)):
                with patch.object(self.jev, "_pooled", wraps=self.jev._pooled) as pooled:
                    actual = self.jev.predict_hypotheses("A", hypotheses)
                    self.assertEqual(actual.shape, (len(hypotheses), 3))
                    pooled.assert_called_once()

    def test_empty(self):
        with self.assertRaises(ValueError):
            self.jev.predict_hypotheses("A", [])


if __name__ == "__main__":
    unittest.main()