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:
@@ -4,3 +4,5 @@ Pillow
|
||||
numpy
|
||||
scikit-image
|
||||
opencv-python
|
||||
cupy
|
||||
cucim
|
||||
|
||||
+22
-3
@@ -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),)
|
||||
|
||||
Reference in New Issue
Block a user