MONAI/testing_data
This is testing data for use with MONAI unit tests.
18.4k
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 