ldov
/

openjevv / test_shared_prefix.py
ldov's picture AlexWortega's picture
Duplicate from AlexWortega/openjev
e46c127
Raw History Blame Contribute Delete
3.82 kB
"""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()