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.
This commit is contained in:
hugovntr
2024-06-10 12:31:35 +02:00
parent 55344012ea
commit 9b882ca19f
2 changed files with 24 additions and 3 deletions
+2
View File
@@ -4,3 +4,5 @@ Pillow
numpy
scikit-image
opencv-python
cupy
cucim
+22 -3
View File
@@ -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),)