diff --git a/generation/maisi/README.md b/generation/maisi/README.md index e7e4e9e47..f51cdf61f 100644 --- a/generation/maisi/README.md +++ b/generation/maisi/README.md @@ -172,6 +172,16 @@ python -m scripts.inference -c ./configs/config_maisi.json -i ./configs/config_i Please refer to [maisi_inference_tutorial.ipynb](maisi_inference_tutorial.ipynb) for the tutorial for MAISI model inference. + +#### Accelerated Inference with TensorRT: +To run the inference script with TensorRT acceleration, please run: +```bash +export MONAI_DATA_DIRECTORY= +python -m scripts.inference -c ./configs/config_maisi.json -i ./configs/config_infer.json -e ./configs/environment.json -x ./configs/config_trt.json --random-seed 0 +``` +Extra config file, [./configs/config_trt.json](./configs/config_trt.json) is using `trt_compile()` utility from MONAI to convert select modules to TensorRT by overriding their definitions from [./configs/config_infer.json](./configs/config_infer.json). + + #### Quality Check: We have implemented a quality check function for the generated CT images. The main idea behind this function is to ensure that the Hounsfield units (HU) intensity for each organ in the CT images remains within a defined range. For each training image used in the Diffusion network, we computed the median value for a few major organs. Then we summarize the statistics of these median values and save it to [./configs/image_median_statistics.json](./configs/image_median_statistics.json). During inference, for each generated image, we compute the median HU values for the major organs and check whether they fall within the normal range. diff --git a/generation/maisi/configs/config_trt.json b/generation/maisi/configs/config_trt.json index fc52486e3..de6469fe0 100644 --- a/generation/maisi/configs/config_trt.json +++ b/generation/maisi/configs/config_trt.json @@ -18,5 +18,7 @@ "device": "cuda", "controlnet": "$trt_compile(@controlnet_def.to(@device), @trained_controlnet_path, @c_trt_args)", "diffusion_unet": "$trt_compile(@diffusion_unet_def.to(@device), @trained_diffusion_path)", + "autoencoder": "$trt_compile(@autoencoder_def.to(@device), @trained_autoencoder_path, submodule='decoder')", + "mask_generation_autoencoder": "$trt_compile(@mask_generation_autoencoder_def.to(@device), @trained_mask_generation_autoencoder_path, submodule='decoder')", "mask_generation_diffusion": "$trt_compile(@mask_generation_diffusion_def.to(@device), @trained_mask_generation_diffusion_path)" }