CoolFace
Datasetpublic

MONAI/testing_data

This is testing data for use with MONAI unit tests.

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
1likes8.4kdownloads
create_data.py101 linesDownload Raw Back to root
1import json2import os3import torch4from typing import Any, Dict, Sequence5 6import monai.networks.nets as nets7 8 9def create_model_test_data(10    model_name: str,11    model_params: Dict[str, Any],12    input_shape: Sequence[int],13) -> None:14    """15    Create test data to check model consistency16 17    Args:18        model_class: Name of model to be tested.19        model_params: Dictionary of parameters to construct object.20        input_shape: Tuple of dimensions (B, C, H, W, [D]).21 22    .. code-block:: python23 24        # network params25        unet_params = {26            "dimensions" : 3,27            "in_channels" : 4,28            "out_channels" : 2,29            "channels": (4, 8, 16, 32),30            "strides": (2, 4, 1),31            "kernel_size" : 5,32            "up_kernel_size" : 3,33            "num_res_units": 2,34            "act": "relu",35            "dropout": 0.1,36        }37        # in shape38        input_shape = (1, unet_params["in_channels"], 64, 64, 64)39        # create data40        create_model_test_data("UNet", unet_params, input_shape)41    """42    model_name = model_name.lower()43    base_folder = os.path.dirname(os.path.abspath(__file__))44 45    # get next unused folder46    i=047    while True:48        out_folder = os.path.join(base_folder, f"{model_name}_{i}")49        if not os.path.isdir(out_folder):50            print("\n\nCreating output folder: " + out_folder)51            os.mkdir(out_folder)52            break53        i += 154    out_path_no_ext = os.path.join(out_folder, f"{model_name}_{i}")55 56    # Create model57    model = nets.__dict__[model_name](**model_params)58    model.eval()59 60    # Create input data61    num_elements = int(torch.Tensor(input_shape).prod())62    in_data = torch.arange(num_elements).reshape(input_shape).float()63 64    # Forward pass data65    out_data = model(in_data)66 67    # Save in data, out data and model68    data_path = out_path_no_ext + ".pt"69    to_save = {"in_data": in_data, "out_data": out_data, "model": model.state_dict()}70    print("Writing data output to .pt: " + data_path)71    torch.save(to_save, data_path)72 73    # Save parameters74    json_params = out_path_no_ext + ".json"75    with open(json_params, "w+") as f:76        print("Writing network parameters to .json: " + json_params)77        json.dump(model_params, f)78 79 80 81# default82if __name__ == "__main__":83 84    # network params85    unet_params = {86        "dimensions" : 3,87        "in_channels" : 4,88        "out_channels" : 2,89        "channels": (4, 8, 16, 32),90        "strides": (2, 4, 1),91        "kernel_size" : 5,92        "up_kernel_size" : 3,93        "num_res_units": 2,94        "act": "relu",95        "dropout": 0.1,96    }97    # in shape98    input_shape = (1, unet_params["in_channels"], 64, 64, 64)99    # create data100    create_model_test_data("UNet", unet_params, input_shape)101