recursionpharma/OpenPhenom
221.2k
1import pytest2import torch3 4# huggingface_openphenom_model_dir = "."5huggingface_modelpath = "recursionpharma/OpenPhenom"6 7from .huggingface_mae import MAEModel8 9 10@pytest.fixture11def huggingface_model():12 # This step downloads the model to a local cache, takes a bit to run13 huggingface_model = MAEModel.from_pretrained(huggingface_modelpath)14 huggingface_model.eval()15 return huggingface_model16 17 18@pytest.mark.parametrize("C", [1, 4, 6, 11])19@pytest.mark.parametrize("return_channelwise_embeddings", [True, False])20def test_model_predict(huggingface_model, C, return_channelwise_embeddings):21 example_input_array = torch.randint(22 low=0,23 high=255,24 size=(2, C, 256, 256),25 dtype=torch.uint8,26 device=huggingface_model.device,27 )28 huggingface_model.return_channelwise_embeddings = return_channelwise_embeddings29 embeddings = huggingface_model.predict(example_input_array)30 expected_output_dim = 384 * C if return_channelwise_embeddings else 38431 assert embeddings.shape == (2, expected_output_dim)32 