Remove PyTorch dependency for DepthPro (#192)

This commit is contained in:
Filip Strand 2025-05-20 19:32:22 +02:00 committed by GitHub
parent d20f2833da
commit 477af82248
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
5 changed files with 79 additions and 53 deletions

View File

@ -4,11 +4,11 @@ from pathlib import Path
import mlx.core as mx import mlx.core as mx
import numpy as np import numpy as np
import torch
from PIL import Image from PIL import Image
from mflux.models.depth_pro.depth_pro_initializer import DepthProInitializer from mflux.models.depth_pro.depth_pro_initializer import DepthProInitializer
from mflux.models.depth_pro.depth_pro_model import DepthProModel from mflux.models.depth_pro.depth_pro_model import DepthProModel
from mflux.models.depth_pro.depth_pro_util import DepthProUtil
from mflux.post_processing.image_util import ImageUtil from mflux.post_processing.image_util import ImageUtil
@ -29,19 +29,34 @@ class DepthPro:
if not os.path.exists(image_path): if not os.path.exists(image_path):
raise FileNotFoundError(f"Image file not found: {image_path}") raise FileNotFoundError(f"Image file not found: {image_path}")
input_array, height, width = self._pre_process(image_path) input_array, height, width = DepthPro._pre_process(image_path)
depth = self._depth_pro_model(input_array) x0, x1, x2 = DepthPro._create_patches(input_array)
return self._post_process(depth, height=height, width=width) depth = self._depth_pro_model(x0, x1, x2)
return DepthPro._post_process(depth, height=height, width=width)
@staticmethod @staticmethod
def _pre_process(image_path): def _pre_process(image_path: str | Path) -> tuple[mx.array, int, int]:
image = Image.open(image_path).convert("RGB") image = Image.open(image_path).convert("RGB")
input_array = ImageUtil.preprocess_for_depth_pro(image) input_array = ImageUtil.preprocess_for_depth_pro(image)
input_array = DepthPro._resize(input_array) input_array = DepthPro._resize(input_array)
return input_array, image.height, image.width return input_array, image.height, image.width
@staticmethod @staticmethod
def _post_process(depth: mx.array, height: int, width: int): def _create_patches(input_array: mx.array) -> tuple[mx.array, mx.array, mx.array]:
# 1. Create the image pyramid
x0 = input_array
x1 = DepthProUtil.interpolate(x=input_array, scale_factor=0.5)
x2 = DepthProUtil.interpolate(x=input_array, scale_factor=0.25)
# 2: Split to create batched overlapped mini-images at the backbone (BeiT/ViT/Dino) resolution.
x0_patches = DepthProUtil.split(x0, overlap_ratio=0.25)
x1_patches = DepthProUtil.split(x1, overlap_ratio=0.5)
x2_patches = x2
return x0_patches, x1_patches, x2_patches
@staticmethod
def _post_process(depth: mx.array, height: int, width: int) -> DepthResult:
depth_min = mx.min(depth) depth_min = mx.min(depth)
depth_max = mx.max(depth) depth_max = mx.max(depth)
normalized_depth = (depth - depth_min) / (depth_max - depth_min) normalized_depth = (depth - depth_min) / (depth_max - depth_min)
@ -58,13 +73,6 @@ class DepthPro:
@staticmethod @staticmethod
def _resize(x: mx.array) -> mx.array: def _resize(x: mx.array) -> mx.array:
x_np = np.array(x) x = mx.expand_dims(x, 0)
x_torch = torch.from_numpy(x_np) x = DepthProUtil.interpolate(x=x, size=(1536, 1536))
x_torch = x_torch.unsqueeze(0) return x
x_torch = torch.nn.functional.interpolate(
x_torch,
size=(1536, 1536),
mode="bilinear",
align_corners=False,
)
return mx.array(x_torch)

View File

@ -4,7 +4,6 @@ import mlx.core as mx
import mlx.nn as nn import mlx.nn as nn
from mflux.models.depth_pro.conv_utils import ConvUtils from mflux.models.depth_pro.conv_utils import ConvUtils
from mflux.models.depth_pro.depth_pro_util import DepthProUtil
from mflux.models.depth_pro.dino_v2.dino_vision_transformer import DinoVisionTransformer from mflux.models.depth_pro.dino_v2.dino_vision_transformer import DinoVisionTransformer
from mflux.models.depth_pro.upsample_block import UpSampleBlock from mflux.models.depth_pro.upsample_block import UpSampleBlock
@ -22,43 +21,35 @@ class DepthProEncoder(nn.Module):
self.upsample_lowres = nn.ConvTranspose2d(in_channels=1024, out_channels=1024, kernel_size=2, stride=2, padding=0, bias=True) # fmt: off self.upsample_lowres = nn.ConvTranspose2d(in_channels=1024, out_channels=1024, kernel_size=2, stride=2, padding=0, bias=True) # fmt: off
self.fuse_lowres = nn.Conv2d(in_channels=1024 * 2, out_channels=1024, kernel_size=1, stride=1, padding=0, bias=True) # fmt: off self.fuse_lowres = nn.Conv2d(in_channels=1024 * 2, out_channels=1024, kernel_size=1, stride=1, padding=0, bias=True) # fmt: off
def __call__(self, x: mx.array) -> list[mx.array]: def __call__(self, x0: mx.array, x1: mx.array, x2: mx.array) -> list[mx.array]:
# 1. Create the image pyramid # 1: Run the backbone patch encoder model
x0, x1, x2 = DepthProUtil.create_pyramid(x) x_pyramid_patches = mx.concatenate((x0, x1, x2), axis=0)
# 2: Split to create batched overlapped mini-images at the backbone (BeiT/ViT/Dino) resolution.
x0_patches = DepthProUtil.split(x0, overlap_ratio=0.25)
x1_patches = DepthProUtil.split(x1, overlap_ratio=0.5)
x2_patches = x2
# 3: Run the backbone (BeiT) model and get the result of large batch size.
x_pyramid_patches = mx.concatenate((x0_patches, x1_patches, x2_patches), axis=0)
x_pyramid_encodings, backbone_highres_hook0, backbone_highres_hook1 = self.patch_encoder(x_pyramid_patches) x_pyramid_encodings, backbone_highres_hook0, backbone_highres_hook1 = self.patch_encoder(x_pyramid_patches)
x_pyramid_encodings = DepthProEncoder._reshape_feature(x_pyramid_encodings, width=24, height=24) x_pyramid_encodings = DepthProEncoder._reshape_feature(x_pyramid_encodings, width=24, height=24)
# Calculate indices for splitting
x0_encodings = x_pyramid_encodings[: len(x0_patches)]
x1_encodings = x_pyramid_encodings[len(x0_patches) : len(x0_patches) + len(x1_patches)]
x2_encodings = x_pyramid_encodings[len(x0_patches) + len(x1_patches) :]
# 4. Merging
x_latent0_encodings = DepthProEncoder._reshape_feature(backbone_highres_hook0, width=24, height=24) x_latent0_encodings = DepthProEncoder._reshape_feature(backbone_highres_hook0, width=24, height=24)
x_latent1_encodings = DepthProEncoder._reshape_feature(backbone_highres_hook1, width=24, height=24) x_latent1_encodings = DepthProEncoder._reshape_feature(backbone_highres_hook1, width=24, height=24)
# Calculate indices for splitting
x0_encodings = x_pyramid_encodings[: len(x0)]
x1_encodings = x_pyramid_encodings[len(x0) : len(x0) + len(x1)]
x2_encodings = x_pyramid_encodings[len(x0) + len(x1) :]
# 2. Merging
x_latent0_features = DepthProEncoder._merge(x_latent0_encodings[: 1 * 5 * 5], batch_size=1, padding=3) x_latent0_features = DepthProEncoder._merge(x_latent0_encodings[: 1 * 5 * 5], batch_size=1, padding=3)
x_latent1_features = DepthProEncoder._merge(x_latent1_encodings[: 1 * 5 * 5], batch_size=1, padding=3) x_latent1_features = DepthProEncoder._merge(x_latent1_encodings[: 1 * 5 * 5], batch_size=1, padding=3)
x0_features = DepthProEncoder._merge(x0_encodings, batch_size=1, padding=3) x0_features = DepthProEncoder._merge(x0_encodings, batch_size=1, padding=3)
x1_features = DepthProEncoder._merge(x1_encodings, batch_size=1, padding=6) x1_features = DepthProEncoder._merge(x1_encodings, batch_size=1, padding=6)
x2_features = x2_encodings x2_features = x2_encodings
# 5. Upsample feature maps. # 3. Upsample feature maps.
x_latent0_features = self.upsample_latent0(x_latent0_features) x_latent0_features = self.upsample_latent0(x_latent0_features)
x_latent1_features = self.upsample_latent1(x_latent1_features) x_latent1_features = self.upsample_latent1(x_latent1_features)
x0_features = self.upsample0(x0_features) x0_features = self.upsample0(x0_features)
x1_features = self.upsample1(x1_features) x1_features = self.upsample1(x1_features)
x2_features = self.upsample2(x2_features) x2_features = self.upsample2(x2_features)
# 6. Apply the image encoder model. # 4. Apply the image encoder model.
x_global_features, _, _ = self.image_encoder(x2_patches) x_global_features, _, _ = self.image_encoder(x2)
x_global_features = DepthProEncoder._reshape_feature(embeddings=x_global_features, width=24, height=24) x_global_features = DepthProEncoder._reshape_feature(embeddings=x_global_features, width=24, height=24)
x_global_features = ConvUtils.apply_conv(x_global_features, self.upsample_lowres) x_global_features = ConvUtils.apply_conv(x_global_features, self.upsample_lowres)
x_global_features = mx.concatenate((x2_features, x_global_features), axis=1) x_global_features = mx.concatenate((x2_features, x_global_features), axis=1)

View File

@ -13,7 +13,7 @@ class DepthProModel(nn.Module):
self.decoder = MultiresConvDecoder() self.decoder = MultiresConvDecoder()
self.head = FOVHead() self.head = FOVHead()
def __call__(self, x: mx.array) -> tuple[mx.array, mx.array]: def __call__(self, x0: mx.array, x1: mx.array, x2: mx.array) -> tuple[mx.array, mx.array]:
encodings = self.encoder(x) encodings = self.encoder(x0, x1, x2)
features = self.decoder(encodings) features = self.decoder(encodings)
return self.head(features) return self.head(features)

View File

@ -2,22 +2,10 @@ import math
import mlx.core as mx import mlx.core as mx
import numpy as np import numpy as np
import torch from PIL import Image
import torch.nn.functional as F
class DepthProUtil: class DepthProUtil:
@staticmethod
def create_pyramid(x: mx.array) -> tuple[mx.array, mx.array, mx.array]:
x0 = x
x_np = np.array(x)
x_torch = torch.from_numpy(x_np)
x1_torch = F.interpolate(x_torch, size=None, scale_factor=0.5, mode="bilinear", align_corners=False)
x2_torch = F.interpolate(x_torch, size=None, scale_factor=0.25, mode="bilinear", align_corners=False)
x1 = mx.array(x1_torch.numpy())
x2 = mx.array(x2_torch.numpy())
return x0, x1, x2
@staticmethod @staticmethod
def split(x: mx.array, overlap_ratio: float = 0.25) -> mx.array: def split(x: mx.array, overlap_ratio: float = 0.25) -> mx.array:
patch_size = 384 patch_size = 384
@ -37,3 +25,42 @@ class DepthProUtil:
x_patch_list.append(x[..., j0:j1, i0:i1]) x_patch_list.append(x[..., j0:j1, i0:i1])
return mx.concatenate(x_patch_list, axis=0) return mx.concatenate(x_patch_list, axis=0)
@staticmethod
def interpolate(x: mx.array, size=None, scale_factor=None):
x_np = np.array(x)
original_ndim = x_np.ndim
if original_ndim == 3:
C, H_in, W_in = x_np.shape
x_proc = np.expand_dims(x_np, 0)
elif original_ndim == 4:
_, C, H_in, W_in = x_np.shape
x_proc = x_np
else:
raise ValueError(f"Unsupported input shape: {x_np.shape}. Must be 3D (C,H,W) or 4D (B,C,H,W).")
if size is not None:
H_out, W_out = size
elif scale_factor is not None:
H_out, W_out = int(H_in * scale_factor), int(W_in * scale_factor)
else:
return x
B_proc, C_proc, _, _ = x_proc.shape
result_proc = np.zeros((B_proc, C_proc, H_out, W_out), dtype=x_np.dtype)
for b in range(B_proc):
for c_idx in range(C_proc):
channel_img_np = x_proc[b, c_idx]
pil_img = Image.fromarray(channel_img_np)
resized_pil_img = pil_img.resize((W_out, H_out), Image.NEAREST)
result_proc[b, c_idx] = np.array(resized_pil_img)
if original_ndim == 3:
final_result_np = result_proc.squeeze(0)
else:
final_result_np = result_proc
return mx.array(final_result_np)

Binary file not shown.

Before

Width:  |  Height:  |  Size: 48 KiB

After

Width:  |  Height:  |  Size: 46 KiB