From 9b882ca19f110688ac0e4bbb8d13ce6893c656b7 Mon Sep 17 00:00:00 2001 From: hugovntr <26794193+hugovntr@users.noreply.github.com> Date: Mon, 10 Jun 2024 12:31:35 +0200 Subject: [PATCH] feat(runtime/colormatch): add CUDA support for histogram matching - Added cupy and cucim to requirements.txt for CUDA operations. - Updated Runtime44ColorMatch class to handle CUDA devices. - Included conditional import for match_histograms from cucim. - Converted tensors to cupy arrays when CUDA is available. - Removed unnecessary imports and optimized memory usage. --- requirements.txt | 2 ++ targets/colormatch.py | 25 ++++++++++++++++++++++--- 2 files changed, 24 insertions(+), 3 deletions(-) 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),)