Replace cvcuda -> torchvision
This commit is contained in:
+10
-35
@@ -1,7 +1,4 @@
|
||||
import torch
|
||||
import cvcuda
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
class Composite:
|
||||
@@ -35,41 +32,19 @@ class Composite:
|
||||
"mismatch number of backgrounds, foregrounds and foreground masks"
|
||||
)
|
||||
|
||||
if mode == "cpu":
|
||||
inverse_masks = 1.0 - foreground_masks
|
||||
if mode != "cuda" and mode != "cpu":
|
||||
raise Exception("invalid mode")
|
||||
|
||||
fgs = foregrounds * foreground_masks.unsqueeze(-1)
|
||||
bgs = backgrounds * inverse_masks.unsqueeze(-1)
|
||||
|
||||
fgs_np = (fgs * 255.0).clamp(0, 255).to(dtype=torch.uint8).cpu().numpy()
|
||||
bgs_np = (bgs * 255.0).clamp(0, 255).to(dtype=torch.uint8).cpu().numpy()
|
||||
|
||||
results = []
|
||||
for fg, bg in zip(fgs_np, bgs_np):
|
||||
result = cv2.add(fg, bg)
|
||||
results.append(result)
|
||||
|
||||
return (
|
||||
torch.from_numpy(np.stack(results, axis=0).astype(np.float32) / 255.0),
|
||||
)
|
||||
elif mode == "cuda":
|
||||
if mode == "cuda":
|
||||
foregrounds = foregrounds.to("cuda")
|
||||
backgrounds = backgrounds.to("cuda")
|
||||
foreground_masks = foreground_masks.to("cuda")
|
||||
|
||||
fgs = cvcuda.convertto(
|
||||
cvcuda.as_tensor(foregrounds, "NHWC"), np.uint8, scale=255
|
||||
)
|
||||
bgs = cvcuda.convertto(
|
||||
cvcuda.as_tensor(backgrounds, "NHWC"), np.uint8, scale=255
|
||||
)
|
||||
fgmasks = cvcuda.convertto(
|
||||
cvcuda.as_tensor(foreground_masks.unsqueeze(-1), "NHWC"),
|
||||
np.uint8,
|
||||
scale=255,
|
||||
)
|
||||
result = cvcuda.composite(fgs, bgs, fgmasks, 3)
|
||||
inverse_masks = 1.0 - foreground_masks
|
||||
|
||||
return (torch.as_tensor(result.cuda()) / 255.0,)
|
||||
else:
|
||||
raise Exception("invalid mode")
|
||||
fgs = foregrounds * foreground_masks.unsqueeze(-1)
|
||||
bgs = backgrounds * inverse_masks.unsqueeze(-1)
|
||||
|
||||
results = fgs + bgs
|
||||
|
||||
return (results,)
|
||||
|
||||
+13
-26
@@ -1,7 +1,6 @@
|
||||
import torch
|
||||
import cvcuda
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
from torchvision.transforms import v2
|
||||
|
||||
|
||||
class GaussianBlur:
|
||||
@@ -27,29 +26,17 @@ class GaussianBlur:
|
||||
sigma: int = 5,
|
||||
mode: str = "cuda",
|
||||
):
|
||||
if mode == "cpu":
|
||||
images_np = (
|
||||
(images * 255.0).clamp(0, 255).to(dtype=torch.uint8).cpu().numpy()
|
||||
)
|
||||
if mode != "cuda" and mode != "cpu":
|
||||
raise Exception("invalid mode")
|
||||
|
||||
results = []
|
||||
for image in images_np:
|
||||
result = cv2.GaussianBlur(
|
||||
image, (kernel_size, kernel_size), sigma, sigma
|
||||
)
|
||||
results.append(result)
|
||||
|
||||
return (
|
||||
torch.from_numpy(np.stack(results, axis=0).astype(np.float32) / 255.0),
|
||||
)
|
||||
elif mode == "cuda":
|
||||
if mode == "cuda":
|
||||
images = images.to("cuda")
|
||||
|
||||
result = cvcuda.gaussian(
|
||||
cvcuda.as_tensor(images, "NHWC"),
|
||||
(kernel_size, kernel_size),
|
||||
(sigma, sigma),
|
||||
)
|
||||
return (torch.as_tensor(result.cuda()),)
|
||||
else:
|
||||
raise Exception("invalid mode")
|
||||
gaussian_blur = v2.GaussianBlur((kernel_size, kernel_size), (sigma, sigma))
|
||||
# torchvision expects a BCHW tensor
|
||||
# Convert input BHWC -> BCHW
|
||||
blurred = gaussian_blur(images.permute(0, 3, 1, 2))
|
||||
|
||||
# Comfy expects a BHWC tensor
|
||||
# Convert output BCHW -> BHWC
|
||||
return (blurred.permute(0, 2, 3, 1),)
|
||||
|
||||
+1
-3
@@ -4,10 +4,8 @@ description = "ComfyUI nodes for editing background of images/videos with CUDA a
|
||||
version = "1.0.0"
|
||||
license = { file = "LICENSE" }
|
||||
dependencies = [
|
||||
"opencv-python",
|
||||
"numpy",
|
||||
"torch",
|
||||
"https://github.com/CVCUDA/CV-CUDA/releases/download/v0.12.0-beta/cvcuda_cu12-0.12.0b0-cp311-cp311-linux_x86_64.whl",
|
||||
"torchvision"
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
|
||||
+1
-3
@@ -1,4 +1,2 @@
|
||||
opencv-python
|
||||
numpy
|
||||
torch
|
||||
https://github.com/CVCUDA/CV-CUDA/releases/download/v0.12.0-beta/cvcuda_cu12-0.12.0b0-cp311-cp311-linux_x86_64.whl
|
||||
torchvision
|
||||
Reference in New Issue
Block a user