Team Ai
Modelpublic

recursionpharma/OpenPhenom

sourceHugging Faceupdated 7mo agoView on Hugging Face
22likes1.2kdownloads
test_huggingface_mae.py32 linesDownload Raw Back to root
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