schp-atr-18 / test_modeling_schp.py
prochet
Add pytest test suite for modeling and image processor
7229448
Raw
History Blame Contribute Delete
1.36 kB
from pathlib import Path
import pytest
import torch
from PIL import Image
from transformers import AutoImageProcessor, AutoModelForSemanticSegmentation
MODEL_PATH = str(Path(__file__).parent)
@pytest.fixture(scope="module")
def model():
return AutoModelForSemanticSegmentation.from_pretrained(
MODEL_PATH, trust_remote_code=True
)
@pytest.fixture(scope="module")
def processor():
return AutoImageProcessor.from_pretrained(MODEL_PATH, trust_remote_code=True)
def test_forward_shape(model):
pixel_values = torch.randn(1, 3, 512, 512)
with torch.no_grad():
outputs = model(pixel_values=pixel_values)
assert outputs.logits.shape == (1, 18, 512, 512)
assert outputs.parsing_logits.shape == (1, 18, 512, 512)
assert outputs.edge_logits.shape == (1, 2, 512, 512)
def test_seg_map(model):
pixel_values = torch.randn(1, 3, 512, 512)
with torch.no_grad():
outputs = model(pixel_values=pixel_values)
seg_map = outputs.logits.argmax(dim=1).squeeze()
assert seg_map.shape == (512, 512)
assert seg_map.min() >= 0
assert seg_map.max() < 18
def test_processor_output_shape(processor):
image = Image.new("RGB", (640, 480))
inputs = processor(images=image, return_tensors="pt")
assert "pixel_values" in inputs
assert inputs["pixel_values"].shape == (1, 3, 512, 512)