DepthPro additional refactor (#193)
This commit is contained in:
parent
477af82248
commit
bc2d73d477
@ -3,7 +3,6 @@ import logging
|
|||||||
import re
|
import re
|
||||||
import subprocess
|
import subprocess
|
||||||
|
|
||||||
|
|
||||||
from mflux.callbacks.callback import BeforeLoopCallback
|
from mflux.callbacks.callback import BeforeLoopCallback
|
||||||
from mflux.error.exceptions import StopImageGenerationException
|
from mflux.error.exceptions import StopImageGenerationException
|
||||||
|
|
||||||
|
|||||||
@ -1,5 +1,4 @@
|
|||||||
from argparse import Namespace
|
from argparse import Namespace
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from mflux import Config, ModelConfig, StopImageGenerationException
|
from mflux import Config, ModelConfig, StopImageGenerationException
|
||||||
|
|||||||
@ -21,7 +21,12 @@ 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, x0: mx.array, x1: mx.array, x2: mx.array) -> list[mx.array]:
|
def __call__(
|
||||||
|
self,
|
||||||
|
x0: mx.array,
|
||||||
|
x1: mx.array,
|
||||||
|
x2: mx.array,
|
||||||
|
) -> tuple[mx.array, mx.array, mx.array, mx.array, mx.array]:
|
||||||
# 1: Run the backbone patch encoder model
|
# 1: Run the backbone patch encoder model
|
||||||
x_pyramid_patches = mx.concatenate((x0, x1, x2), axis=0)
|
x_pyramid_patches = mx.concatenate((x0, x1, x2), 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)
|
||||||
@ -55,13 +60,13 @@ class DepthProEncoder(nn.Module):
|
|||||||
x_global_features = mx.concatenate((x2_features, x_global_features), axis=1)
|
x_global_features = mx.concatenate((x2_features, x_global_features), axis=1)
|
||||||
x_global_features = ConvUtils.apply_conv(x_global_features, self.fuse_lowres)
|
x_global_features = ConvUtils.apply_conv(x_global_features, self.fuse_lowres)
|
||||||
|
|
||||||
return [
|
return (
|
||||||
x_latent0_features,
|
x_latent0_features,
|
||||||
x_latent1_features,
|
x_latent1_features,
|
||||||
x0_features,
|
x0_features,
|
||||||
x1_features,
|
x1_features,
|
||||||
x_global_features,
|
x_global_features,
|
||||||
]
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _reshape_feature(
|
def _reshape_feature(
|
||||||
|
|||||||
@ -14,6 +14,6 @@ class DepthProModel(nn.Module):
|
|||||||
self.head = FOVHead()
|
self.head = FOVHead()
|
||||||
|
|
||||||
def __call__(self, x0: mx.array, x1: mx.array, x2: 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(x0, x1, x2)
|
x0_lat, x1_lat, x0_feat, x1_feat, x_global = self.encoder(x0, x1, x2)
|
||||||
features = self.decoder(encodings)
|
decoded = self.decoder(x0_lat, x1_lat, x0_feat, x1_feat, x_global)
|
||||||
return self.head(features)
|
return self.head(decoded)
|
||||||
|
|||||||
@ -23,16 +23,29 @@ class MultiresConvDecoder(nn.Module):
|
|||||||
FeatureFusionBlock2d(num_features=256, deconv=True),
|
FeatureFusionBlock2d(num_features=256, deconv=True),
|
||||||
]
|
]
|
||||||
|
|
||||||
def __call__(self, encodings: list[mx.array]) -> tuple[mx.array, mx.array]:
|
def __call__(
|
||||||
# Process last layer
|
self,
|
||||||
encodings_last = encodings[4]
|
x0_latent: mx.array,
|
||||||
features = ConvUtils.apply_conv(encodings_last, self.convs[4])
|
x1_latent: mx.array,
|
||||||
|
x0_features: mx.array,
|
||||||
|
x1_features: mx.array,
|
||||||
|
x_global_features: mx.array,
|
||||||
|
) -> mx.array:
|
||||||
|
# Process global features:
|
||||||
|
features = ConvUtils.apply_conv(x_global_features, self.convs[4])
|
||||||
features = self.fusions[4](features)
|
features = self.fusions[4](features)
|
||||||
|
|
||||||
# Process remaining levels with skip connections
|
# Process remaining levels with skip connections:
|
||||||
for i in [3, 2, 1, 0]:
|
x1_skip_features = ConvUtils.apply_conv(x1_features, self.convs[3])
|
||||||
enc = encodings[i]
|
features = self.fusions[3](features, x1_skip_features)
|
||||||
features_i = ConvUtils.apply_conv(enc, self.convs[i])
|
|
||||||
features = self.fusions[i](features, features_i)
|
x0_skip_features = ConvUtils.apply_conv(x0_features, self.convs[2])
|
||||||
|
features = self.fusions[2](features, x0_skip_features)
|
||||||
|
|
||||||
|
x1_skip_latents = ConvUtils.apply_conv(x1_latent, self.convs[1])
|
||||||
|
features = self.fusions[1](features, x1_skip_latents)
|
||||||
|
|
||||||
|
x0_skip_latents = ConvUtils.apply_conv(x0_latent, self.convs[0])
|
||||||
|
features = self.fusions[0](features, x0_skip_latents)
|
||||||
|
|
||||||
return features
|
return features
|
||||||
|
|||||||
@ -11,6 +11,8 @@ from mlx.utils import tree_unflatten
|
|||||||
from mflux.weights.weight_handler import MetaData
|
from mflux.weights.weight_handler import MetaData
|
||||||
from mflux.weights.weight_util import WeightUtil
|
from mflux.weights.weight_util import WeightUtil
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class WeightHandlerDepthPro:
|
class WeightHandlerDepthPro:
|
||||||
def __init__(self, weights: dict, meta_data: MetaData):
|
def __init__(self, weights: dict, meta_data: MetaData):
|
||||||
@ -51,13 +53,13 @@ class WeightHandlerDepthPro:
|
|||||||
|
|
||||||
# 2. Download if model doesn't exist
|
# 2. Download if model doesn't exist
|
||||||
if not model_path.exists():
|
if not model_path.exists():
|
||||||
logging.info("Downloading Depth Pro model from Apple...")
|
logger.info("Downloading Depth Pro model from Apple...")
|
||||||
try:
|
try:
|
||||||
urllib.request.urlretrieve(APPLE_MODEL_URL, model_path)
|
urllib.request.urlretrieve(APPLE_MODEL_URL, model_path)
|
||||||
logging.info(f"Downloaded model to {model_path}")
|
logger.info(f"Downloaded model to {model_path}")
|
||||||
except (urllib.error.URLError, urllib.error.HTTPError) as e:
|
except (urllib.error.URLError, urllib.error.HTTPError) as e:
|
||||||
logging.error(f"Failed to download model: {e}")
|
logger.error(f"Failed to download model: {e}")
|
||||||
logging.info(f"Please manually download from: {APPLE_MODEL_URL}")
|
logger.info(f"Please manually download from: {APPLE_MODEL_URL}")
|
||||||
if not model_path.exists():
|
if not model_path.exists():
|
||||||
raise FileNotFoundError(f"Model file not found at {model_path}")
|
raise FileNotFoundError(f"Model file not found at {model_path}")
|
||||||
|
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user