frox-nano-v2 / src /tests /test_multimodal.py
Hritik045678's picture
Upload folder using huggingface_hub
296a506 verified
Raw
History Blame Contribute Delete
2.67 kB
"""
Frox AI Morph 1.1 — Multimodal Fusion Tests
Run with: pytest tests/test_multimodal.py -v
"""
from __future__ import annotations
import pytest
import torch
from multimodal.fusion.morph_multimodal import MorphMultimodalModel
class TestImageTokenInjection:
def test_single_image_expands_sequence_correctly(self, tiny_config, device):
model = MorphMultimodalModel(tiny_config).to(device).eval()
num_image_tokens = (tiny_config.vision.image_size // tiny_config.vision.patch_size) ** 2
text_len = 6
input_ids = torch.randint(
0, tiny_config.text.vocab_size, (1, text_len), device=device
)
input_ids[0, 3] = model.image_token_id
pixel_values = torch.randn(
1, 3, tiny_config.vision.image_size, tiny_config.vision.image_size, device=device
)
with torch.no_grad():
inputs_embeds, mask = model.prepare_multimodal_inputs(input_ids, pixel_values)
expected_len = (text_len - 1) + num_image_tokens
assert inputs_embeds.shape[1] == expected_len
assert inputs_embeds.shape[2] == tiny_config.text.hidden_size
def test_no_image_falls_back_to_text_only(self, tiny_config, device):
model = MorphMultimodalModel(tiny_config).to(device).eval()
input_ids = torch.randint(0, tiny_config.text.total_vocab_size, (1, 8), device=device)
with torch.no_grad():
embeds, mask = model.prepare_multimodal_inputs(input_ids, pixel_values=None)
assert embeds.shape[1] == 8
assert mask is None
def test_forward_pass_with_image_produces_finite_logits(self, tiny_config, device):
model = MorphMultimodalModel(tiny_config).to(device).eval()
input_ids = torch.randint(0, tiny_config.text.vocab_size, (1, 6), device=device)
input_ids[0, 2] = model.image_token_id
pixel_values = torch.randn(
1, 3, tiny_config.vision.image_size, tiny_config.vision.image_size, device=device
)
with torch.no_grad():
out = model(input_ids=input_ids, pixel_values=pixel_values)
assert not torch.isnan(out.logits).any()
assert not torch.isinf(out.logits).any()
class TestParamCount:
def test_param_count_breakdown_sums_correctly(self, tiny_config, device):
model = MorphMultimodalModel(tiny_config).to(device)
counts = model.param_count()
# lm + vision should roughly account for the total (small rounding tolerance
# from the shared image_sep_embed parameter, which isn't double counted)
assert counts["total_billions"] >= counts["lm_billions"] + counts["vision_billions"] - 1e-6