diff --git a/requirements.txt b/requirements.txt index 62b0e66..f361820 100644 --- a/requirements.txt +++ b/requirements.txt @@ -4,3 +4,5 @@ Pillow numpy scikit-image opencv-python +cupy +cucim diff --git a/targets/colormatch.py b/targets/colormatch.py index 0ddf417..3229d21 100644 --- a/targets/colormatch.py +++ b/targets/colormatch.py @@ -1,6 +1,4 @@ import torch -import numpy as np -from skimage.exposure import match_histograms class Runtime44ColorMatch: @@ -18,7 +16,28 @@ class Runtime44ColorMatch: CATEGORY = "image" def match(self, source: torch.Tensor, target: torch.Tensor): + """ + + Match the color of the `target` image with the `source` image + using the **histogram matching** technique + + Using CuPy (for CUDA) and NumPy (for CPU/Non-CUDA device) + + """ + is_cuda = torch.cuda.is_available() + target = target.numpy() source = source.numpy() - matched = match_histograms(target, source, channel_axis=None) + if is_cuda: + from cucim.skimage.exposure import match_histograms + import cupy + + target = cupy.asarray(target) + source = cupy.asarray(source) + else: + from skimage.exposure import match_histograms + matched = match_histograms(target, source, channel_axis=-1) + matched = cupy.asnumpy(matched) if is_cuda else matched + del source + del target return (torch.from_numpy(matched),)