mmdiff / models /flux_resizer.py
multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
2267636 verified
Raw
History Blame Contribute Delete
3.52 kB
"""Image resizing to FLUX-compatible resolutions (dimensions divisible by 32)."""
import cv2
import numpy as np
from typing import Tuple
from PIL import Image
class FluxResizer:
"""
Resizer that ensures images are compatible with FLUX requirements.
- FLUX: Dimensions divisible by 32 (due to 2x2 packing on top of 16-stride VAE)
"""
# Predefined optimal resolutions (all divisible by 32)
OPTIMAL_RESOLUTIONS = [
# Square and near-square
(1024, 1024), # 1:1 (64×64, 64×64)
(896, 1152), # ~0.78:1 (56×72)
(1152, 896), # ~1.29:1 (72×56)
(768, 1344), # ~0.57:1 (48×84)
(1344, 768), # ~1.75:1 (84×48)
# Additional common ratios
(832, 1216), # ~0.68:1 (52×76)
(1216, 832), # ~1.46:1 (76×52)
(704, 1408), # 0.5:1 (44×88)
(1408, 704), # 2:1 (88×44)
(960, 1088), # ~0.88:1 (60×68)
(1088, 960), # ~1.13:1 (68×60)
]
def __init__(self):
self.resolution_aspects = [
(h, w, w / h) for h, w in self.OPTIMAL_RESOLUTIONS
]
def select_best_resolution(self, original_h: int, original_w: int) -> Tuple[int, int]:
"""Pick the optimal resolution whose aspect ratio best matches the input."""
original_aspect = original_w / original_h
best_resolution = None
min_aspect_diff = float('inf')
for h, w, aspect in self.resolution_aspects:
aspect_diff = abs(original_aspect - aspect)
if aspect_diff < min_aspect_diff:
min_aspect_diff = aspect_diff
best_resolution = (h, w)
return best_resolution
def resize_image(self, image: np.ndarray) -> Tuple[np.ndarray, Tuple[int, int]]:
"""Resize an (H, W, C) numpy image to the optimal FLUX-compatible resolution."""
original_h, original_w = image.shape[:2]
target_h, target_w = self.select_best_resolution(original_h, original_w)
resized_image = cv2.resize(image, (target_w, target_h), interpolation=cv2.INTER_LINEAR)
return resized_image, (target_h, target_w)
def resize_pil_image(self, image: Image.Image) -> Tuple[Image.Image, Tuple[int, int]]:
"""Resize a PIL image to the optimal FLUX-compatible resolution."""
original_w, original_h = image.size # PIL uses (W, H)
target_h, target_w = self.select_best_resolution(original_h, original_w)
resized_image = image.resize((target_w, target_h), Image.LANCZOS)
return resized_image, (target_h, target_w)
def resize_mask(self, mask: np.ndarray, target_size: Tuple[int, int]) -> np.ndarray:
"""Resize a mask to target_size=(H, W) using nearest-neighbor interpolation."""
target_h, target_w = target_size
if len(mask.shape) == 3 and mask.shape[2] == 1:
mask = mask.squeeze(2)
resized_mask = cv2.resize(mask, (target_w, target_h), interpolation=cv2.INTER_NEAREST)
return resized_mask
def get_compatible_resolutions(self) -> list:
"""Return list of all compatible resolutions."""
return self.OPTIMAL_RESOLUTIONS.copy()
@staticmethod
def verify_compatibility(height: int, width: int) -> bool:
"""True if both dimensions are divisible by 32 (FLUX requirement)."""
return (height % 32 == 0) and (width % 32 == 0)