V2.0.0 - Depth blur DA V3 - ok 2

This commit is contained in:
DESKTOP-TVBJISQ\Primere
2026-03-31 16:54:53 +02:00
parent 7e6eac57f9
commit 5f79cbd3ea
2 changed files with 4 additions and 154 deletions
+4 -137
View File
@@ -1,14 +1,11 @@
"""Basic inference nodes for DepthAnythingV3."""
import torch
import torch.nn.functional as F
import comfy.model_management as mm
from comfy.utils import ProgressBar
from comfy_api.latest import io
from .utils import (
IMAGENET_MEAN, IMAGENET_STD, DEFAULT_PATCH_SIZE,
format_camera_params, process_tensor_to_image, process_tensor_to_mask,
resize_to_patch_multiple, logger, check_model_capabilities,
resize_to_patch_multiple, check_model_capabilities,
imagenet_normalize, save_gaussians_to_ply,
)
from .normalization import (
@@ -21,43 +18,18 @@ from .normalization import (
class DepthAnything_V3():
@classmethod
def execute(cls, da3_model, images, normalization_mode="V2-Style", camera_params=None,
resize_method="resize", invert_depth=False, keep_model_size=False):
def execute(cls, da3_model, images, normalization_mode="V2-Style", camera_params=None, resize_method="resize", invert_depth=False, keep_model_size=False):
device = mm.get_torch_device()
# da3_model is now a ModelPatcher — load to GPU via ComfyUI memory management
mm.load_models_gpu([da3_model])
model = da3_model.model
# Get metadata stored by loader
capabilities = da3_model.model_options.get("da3_capabilities", check_model_capabilities(model))
dtype = da3_model.model_options.get("da3_dtype", torch.float16)
if not capabilities["has_sky_segmentation"] and normalization_mode == "V2-Style":
logger.warning(
"WARNING: This model does not support sky segmentation. "
"V2-Style normalization will work but without sky masking. "
"Use Mono/Metric/Nested models for best V2-Style results."
)
B, H, W, C = images.shape
logger.info(f"Input image size: {H}x{W}")
# Convert from ComfyUI format [B, H, W, C] to PyTorch [B, C, H, W]
images_pt = images.permute(0, 3, 1, 2)
# Resize to patch size multiple
images_pt, orig_H, orig_W = resize_to_patch_multiple(images_pt, DEFAULT_PATCH_SIZE, resize_method)
model_H, model_W = images_pt.shape[2], images_pt.shape[3]
logger.info(f"Model input size (after resize): {model_H}x{model_W}")
# Normalize with ImageNet stats (manual, no torchvision dependency)
normalized_images = imagenet_normalize(images_pt)
# Prepare for model: add view dimension [B, N, 3, H, W] where N=1
normalized_images = normalized_images.unsqueeze(1)
# Prepare camera parameters if provided
extrinsics_input = None
intrinsics_input = None
if camera_params is not None:
@@ -67,11 +39,7 @@ class DepthAnything_V3():
if extrinsics_input.shape[0] == 1 and B > 1:
extrinsics_input = extrinsics_input.expand(B, -1, -1, -1)
intrinsics_input = intrinsics_input.expand(B, -1, -1, -1)
logger.info("Using camera-conditioned depth estimation")
else:
logger.warning("Model does not support camera conditioning. Camera params ignored.")
pbar = ProgressBar(B)
depth_out = []
conf_out = []
sky_out = []
@@ -81,22 +49,13 @@ class DepthAnything_V3():
intrinsics_list = []
gaussians_list = []
# Check if model supports 3D Gaussians
infer_gs = capabilities["has_3d_gaussians"]
if infer_gs:
logger.info("Model supports 3D Gaussians - will output raw Gaussians")
for i in range(B):
img = normalized_images[i:i+1].to(device, dtype=dtype)
# Get camera params for this batch item
ext_i = extrinsics_input[i:i+1] if extrinsics_input is not None else None
int_i = intrinsics_input[i:i+1] if intrinsics_input is not None else None
# Run model forward with optional camera conditioning and Gaussians
output = model(img, extrinsics=ext_i, intrinsics=int_i, infer_gs=infer_gs)
# Extract depth
depth = None
if hasattr(output, 'depth'):
depth = output.depth
@@ -106,17 +65,13 @@ class DepthAnything_V3():
if depth is None or not torch.is_tensor(depth):
raise ValueError("Model output does not contain valid depth tensor")
# Extract confidence
conf = None
if hasattr(output, 'depth_conf'):
conf = output.depth_conf
elif isinstance(output, dict) and 'depth_conf' in output:
conf = output['depth_conf']
if conf is None or not torch.is_tensor(conf):
conf = torch.ones_like(depth)
# Extract sky mask
sky = None
if hasattr(output, 'sky'):
sky = output.sky
@@ -126,12 +81,10 @@ class DepthAnything_V3():
if sky is None or not torch.is_tensor(sky):
sky = torch.zeros_like(depth)
else:
# Normalize sky mask to 0-1 range
sky_min, sky_max = sky.min(), sky.max()
if sky_max > sky_min:
sky = (sky - sky_min) / (sky_max - sky_min)
# ===== NORMALIZATION DISPATCH =====
if normalization_mode == "Raw":
depth_processed = apply_raw_normalization(depth, invert_depth)
elif normalization_mode == "V2-Style":
@@ -139,7 +92,6 @@ class DepthAnything_V3():
else: # "Standard"
depth_processed = apply_standard_normalization(depth, invert_depth)
# Normalize confidence
conf_range = conf.max() - conf.min()
if conf_range > 1e-8:
conf = (conf - conf.min()) / conf_range
@@ -150,7 +102,6 @@ class DepthAnything_V3():
conf_out.append(conf.cpu())
sky_out.append(sky.cpu())
# Extract ray maps (if available)
ray = None
if hasattr(output, 'ray'):
ray = output.ray
@@ -167,7 +118,6 @@ class DepthAnything_V3():
ray_origin_out.append(torch.zeros(3, depth.shape[-2], depth.shape[-1]))
ray_dir_out.append(torch.zeros(3, depth.shape[-2], depth.shape[-1]))
# Extract camera parameters (if available)
extr = None
if hasattr(output, 'extrinsics'):
extr = output.extrinsics
@@ -190,7 +140,6 @@ class DepthAnything_V3():
else:
intrinsics_list.append(None)
# Extract 3D Gaussians (only if model supports them and we requested them)
if infer_gs:
gs = None
if hasattr(output, 'gaussians'):
@@ -199,40 +148,11 @@ class DepthAnything_V3():
gs = output['gaussians']
if gs is not None and hasattr(gs, 'means') and torch.is_tensor(gs.means):
# Store raw depth alongside Gaussians for pruning
gaussians_list.append((gs, depth))
pbar.update(1)
# Process outputs based on normalization mode
normalize_depth_output = (normalization_mode != "Raw")
depth_final = process_tensor_to_image(depth_out, orig_H, orig_W, normalize_output=normalize_depth_output, skip_resize=keep_model_size)
depth_final = process_tensor_to_image(depth_out, orig_H, orig_W,
normalize_output=normalize_depth_output,
skip_resize=keep_model_size)
conf_final = process_tensor_to_image(conf_out, orig_H, orig_W,
normalize_output=True,
skip_resize=keep_model_size)
sky_final = process_tensor_to_mask(sky_out, orig_H, orig_W, skip_resize=keep_model_size)
ray_origin_final = cls._process_ray_to_image(ray_origin_out, orig_H, orig_W,
normalize=True, skip_resize=keep_model_size)
ray_dir_final = cls._process_ray_to_image(ray_dir_out, orig_H, orig_W,
normalize=True, skip_resize=keep_model_size)
# Process resized RGB image to match depth output dimensions
rgb_resized = images_pt.permute(0, 2, 3, 1).float().cpu()
if not keep_model_size:
final_H = (orig_H // 2) * 2
final_W = (orig_W // 2) * 2
if rgb_resized.shape[1] != final_H or rgb_resized.shape[2] != final_W:
rgb_resized = F.interpolate(
rgb_resized.permute(0, 3, 1, 2),
size=(final_H, final_W),
mode="bilinear"
).permute(0, 2, 3, 1)
rgb_resized = torch.clamp(rgb_resized, 0, 1)
# Scale intrinsics if we resized back to original dimensions
if not keep_model_size:
final_H = (orig_H // 2) * 2
final_W = (orig_W // 2) * 2
@@ -251,58 +171,10 @@ class DepthAnything_V3():
intr_scaled[1, 2] *= scale_h # cy
intrinsics_list[i] = intr_scaled
# Format camera parameters as strings (for backward compatibility)
extrinsics_str = format_camera_params(extrinsics_list, "extrinsics")
intrinsics_str = format_camera_params(intrinsics_list, "intrinsics")
# Prepare tensor outputs for direct connection to other nodes
if extrinsics_list and extrinsics_list[0] is not None:
extrinsics_tensor = torch.stack([e.squeeze() for e in extrinsics_list if e is not None], dim=0)
else:
extrinsics_tensor = torch.eye(4).unsqueeze(0).expand(len(depth_out), -1, -1)
if intrinsics_list and intrinsics_list[0] is not None:
# Convert 3x3 intrinsics to 4x4 homogeneous (compatible with Sharp)
intr_tensors = []
for i_mat in intrinsics_list:
if i_mat is not None:
k = i_mat.squeeze()
if k.shape == (3, 3):
k4 = torch.eye(4, dtype=k.dtype)
k4[:3, :3] = k
intr_tensors.append(k4)
else:
intr_tensors.append(k)
intrinsics_tensor = torch.stack(intr_tensors, dim=0)
else:
intrinsics_tensor = torch.eye(4).unsqueeze(0).expand(len(depth_out), -1, -1)
# Save Gaussians to PLY file if available (Giant model only)
gaussian_ply_path = ""
if gaussians_list:
import folder_paths
from pathlib import Path
output_dir = Path(folder_paths.get_output_directory())
# Use the first batch item's Gaussians and depth for pruning
gs, raw_depth = gaussians_list[0]
# Raw depth shape: (1, 1, H, W) -> squeeze to (1, H, W) for pruning
depth_for_pruning = raw_depth.squeeze(0) if raw_depth.dim() == 4 else raw_depth
# Get extrinsics for world-to-camera transform (preserves scale/position relationship)
gs_extrinsics = extrinsics_list[0] if extrinsics_list and extrinsics_list[0] is not None else None
filepath = output_dir / "gaussians_worldspace_0000.ply"
gaussian_ply_path = save_gaussians_to_ply(
gs, filepath, depth=depth_for_pruning,
extrinsics=gs_extrinsics,
shift_and_scale=False, save_sh_dc_only=False,
prune_border=True, prune_depth_percent=0.9,
)
return depth_final
# return io.NodeOutput(depth_final, conf_final, rgb_resized, ray_origin_final, ray_dir_final, extrinsics_str, intrinsics_str, sky_final, extrinsics_tensor, intrinsics_tensor, gaussian_ply_path)
@staticmethod
def _process_ray_to_image(ray_list, orig_H, orig_W, normalize=True, skip_resize=False):
"""Convert list of ray tensors to ComfyUI IMAGE format."""
out = torch.cat([r.unsqueeze(0) for r in ray_list], dim=0)
if normalize:
@@ -320,13 +192,8 @@ class DepthAnything_V3():
if not skip_resize:
final_H = (orig_H // 2) * 2
final_W = (orig_W // 2) * 2
if out.shape[1] != final_H or out.shape[2] != final_W:
out = F.interpolate(
out.permute(0, 3, 1, 2),
size=(final_H, final_W),
mode="bilinear"
).permute(0, 2, 3, 1)
out = F.interpolate(out.permute(0, 3, 1, 2), size=(final_H, final_W), mode="bilinear").permute(0, 2, 3, 1)
if normalize:
return torch.clamp(out, 0, 1)
-17
View File
@@ -111,9 +111,6 @@ def _load_depth_model_v3():
return _depth_model_v3
def _load_local_depth_model_v3(model_name):
print('*************************************')
print(model_name)
print('*************************************')
model = load_model.DownloadAndLoadDepthAnythingV3Model.execute(model_name)
return model
@@ -147,16 +144,6 @@ def _predict_depth(arr, imagei, use_v3: bool = False):
depth = nodes_inference.DepthAnything_V3.execute(model, imagei)
h, w, _ = arr.shape
print('=====================')
print(arr.shape)
print('=====================')
# with tempfile.TemporaryDirectory() as tmpdir:
# temp_path = os.path.join(tmpdir, "temp_input.png")
# Image.fromarray((arr * 255.0).astype(np.uint8)).save(temp_path)
# pred = model.inference([temp_path])
# depth = pred.depth[0]
if isinstance(depth, torch.Tensor):
depth = depth.detach().cpu().float()
# DepthAnything_V3 node returns Comfy IMAGE format [B, H, W, C].
@@ -226,10 +213,6 @@ def _erode_focus_mask(depth_blur, erode_px, feather_px):
def _build_protection_mask(raw_depth, focus_depth, protect_sigma, arr):
print('=========== RD =================')
print(raw_depth)
print('============================')
luma = _to_luminance(arr)
sharp = _sharpness_map(luma, radius=1.0)
sharp_mask = np.clip((sharp - PROTECT_SHARPNESS_THR) * PROTECT_SHARPNESS_BIAS, 0.0, 1.0)