mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-26 03:25:43 +08:00
fixed some bugs and rewrote OpenCV resize funcs
This commit is contained in:
parent
b3839ca722
commit
174655006c
@ -1,13 +1,18 @@
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from dinov2 import Dinov2Model, DinoConfig
|
from dinov2 import DinoConfig, Dinov2Model
|
||||||
|
|
||||||
# avoid using torchvision by recreating image processing functions
|
# avoid using torchvision by recreating image processing functions
|
||||||
|
|
||||||
def resize(img: torch.Tensor, size: int) -> torch.Tensor:
|
def resize(img: torch.Tensor, size: int) -> torch.Tensor:
|
||||||
|
|
||||||
_, h, w = img.shape
|
batched = img.ndim == 4
|
||||||
|
|
||||||
|
if not batched:
|
||||||
|
img = img.unsqueeze(0)
|
||||||
|
|
||||||
|
_, _, h, w = img.shape
|
||||||
|
|
||||||
# mantain aspect ratio
|
# mantain aspect ratio
|
||||||
if h < w:
|
if h < w:
|
||||||
@ -17,18 +22,31 @@ def resize(img: torch.Tensor, size: int) -> torch.Tensor:
|
|||||||
new_w = size
|
new_w = size
|
||||||
new_h = int(h * size / w)
|
new_h = int(h * size / w)
|
||||||
|
|
||||||
img = img.unsqueeze(0)
|
|
||||||
img = F.interpolate(img, size = (new_h, new_w), mode = 'bilinear', align_corners = False, antialias = True )
|
img = F.interpolate(img, size = (new_h, new_w), mode = 'bilinear', align_corners = False, antialias = True )
|
||||||
return img.squeeze(0)
|
|
||||||
|
if not batched:
|
||||||
|
img = img.squeeze(0)
|
||||||
|
|
||||||
|
return img
|
||||||
|
|
||||||
|
|
||||||
def center_crop(img: torch.Tensor, size: int) -> torch.Tensor:
|
def center_crop(img: torch.Tensor, size: int) -> torch.Tensor:
|
||||||
|
|
||||||
_, h, w = img.shape
|
batched = img.ndim == 4
|
||||||
|
if not batched:
|
||||||
|
img = img.unsqueeze(0)
|
||||||
|
|
||||||
|
_, _, h, w = img.shape
|
||||||
top = (h - size) // 2
|
top = (h - size) // 2
|
||||||
left = (w - size) // 2
|
left = (w - size) // 2
|
||||||
|
|
||||||
return img[:, top:top + size, left:left + size]
|
|
||||||
|
cropped = img[..., top:top + size, left:left + size]
|
||||||
|
|
||||||
|
if not batched:
|
||||||
|
cropped = cropped.squeeze(0)
|
||||||
|
|
||||||
|
return cropped
|
||||||
|
|
||||||
def normalize(img: torch.Tensor, mean: list, std: list) -> torch.Tensor:
|
def normalize(img: torch.Tensor, mean: list, std: list) -> torch.Tensor:
|
||||||
|
|
||||||
@ -47,8 +65,8 @@ class ImageEncoder(nn.Module):
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
config: DinoConfig,
|
config: DinoConfig,
|
||||||
use_cls_token=True,
|
use_cls_token = True,
|
||||||
image_size=224,
|
image_size = 518,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
@ -77,15 +95,16 @@ class ImageEncoder(nn.Module):
|
|||||||
|
|
||||||
def forward(self, image, value_range=(-1, 1), **kwargs):
|
def forward(self, image, value_range=(-1, 1), **kwargs):
|
||||||
|
|
||||||
|
if image.ndim == 3:
|
||||||
|
image = image.unsqueeze(0)
|
||||||
|
|
||||||
if value_range is not None:
|
if value_range is not None:
|
||||||
low, high = value_range
|
low, high = value_range
|
||||||
image = (image - low) / (high - low)
|
image = (image - low) / (high - low)
|
||||||
|
|
||||||
image = image.to(self.model.device, dtype=self.model.dtype)
|
|
||||||
inputs = self.transform(image)
|
inputs = self.transform(image)
|
||||||
outputs = self.model(inputs)
|
inputs = inputs.to(self.model.device, dtype=self.model.dtype)
|
||||||
|
last_hidden_state = self.model(inputs)
|
||||||
last_hidden_state = outputs.last_hidden_state
|
|
||||||
|
|
||||||
if not self.use_cls_token:
|
if not self.use_cls_token:
|
||||||
last_hidden_state = last_hidden_state[:, 1:, :]
|
last_hidden_state = last_hidden_state[:, 1:, :]
|
||||||
@ -101,16 +120,16 @@ class ImageEncoder(nn.Module):
|
|||||||
batch_size,
|
batch_size,
|
||||||
self.num_patches,
|
self.num_patches,
|
||||||
self.model.config.hidden_size,
|
self.model.config.hidden_size,
|
||||||
device=device,
|
device = device,
|
||||||
dtype=dtype,
|
dtype = dtype,
|
||||||
)
|
)
|
||||||
|
|
||||||
return zero
|
return zero
|
||||||
|
|
||||||
class SingleImageEncoder(nn.Module):
|
class SingleImageEncoder(nn.Module):
|
||||||
def __init__(self):
|
def __init__(self, config):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.main_image_encoder = ImageEncoder()
|
self.main_image_encoder = ImageEncoder(config)
|
||||||
|
|
||||||
def forward(self, image, **kwargs):
|
def forward(self, image, **kwargs):
|
||||||
outputs = {
|
outputs = {
|
||||||
@ -124,22 +143,22 @@ class SingleImageEncoder(nn.Module):
|
|||||||
}
|
}
|
||||||
return outputs
|
return outputs
|
||||||
|
|
||||||
def load_dino2(dino2: Dinov2Model):
|
|
||||||
|
|
||||||
checkpoint = ""
|
|
||||||
dino2.load_state_dict(torch.load(checkpoint))
|
|
||||||
return dino2
|
|
||||||
|
|
||||||
def test_image_encoder():
|
def test_image_encoder():
|
||||||
|
|
||||||
torch.manual_seed(2025)
|
torch.manual_seed(2025)
|
||||||
image_encoder = SingleImageEncoder()
|
config = DinoConfig()
|
||||||
|
image_encoder = SingleImageEncoder(config)
|
||||||
|
|
||||||
image = torch.rand(1, 3, 224, 224)
|
image = torch.rand(3, 224, 224)
|
||||||
|
|
||||||
outputs = image_encoder(image)
|
outputs = image_encoder(image)
|
||||||
|
|
||||||
print(outputs)
|
print(outputs)
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
test_image_encoder()
|
#test_image_encoder()
|
||||||
|
conditioner = SingleImageEncoder(DinoConfig())
|
||||||
|
torch.manual_seed(2025)
|
||||||
|
image = torch.rand(1, 3, 224, 224)
|
||||||
|
outputs = conditioner(image)
|
||||||
|
print(outputs["main"].size())
|
||||||
@ -21,6 +21,8 @@ class DinoConfig():
|
|||||||
qkv_bias: bool = True
|
qkv_bias: bool = True
|
||||||
layerscale_value: float = 1.0
|
layerscale_value: float = 1.0
|
||||||
drop_path_rate: float = 0.0
|
drop_path_rate: float = 0.0
|
||||||
|
device: str = "cuda"
|
||||||
|
dtype = torch.float16
|
||||||
|
|
||||||
class Dinov2Embeddings(nn.Module):
|
class Dinov2Embeddings(nn.Module):
|
||||||
"""
|
"""
|
||||||
@ -372,6 +374,8 @@ class Dinov2Model(nn.Module):
|
|||||||
|
|
||||||
self.embeddings = Dinov2Embeddings(config)
|
self.embeddings = Dinov2Embeddings(config)
|
||||||
self.encoder = Dinov2Encoder(config)
|
self.encoder = Dinov2Encoder(config)
|
||||||
|
self.device = config.device
|
||||||
|
self.dtype = config.dtype
|
||||||
|
|
||||||
self.layernorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
|
self.layernorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
|
||||||
|
|
||||||
@ -391,7 +395,7 @@ class Dinov2Model(nn.Module):
|
|||||||
embedding_output,
|
embedding_output,
|
||||||
head_mask = head_mask,
|
head_mask = head_mask,
|
||||||
)
|
)
|
||||||
sequence_output = encoder_outputs[0]
|
sequence_output = encoder_outputs
|
||||||
sequence_output = self.layernorm(sequence_output)
|
sequence_output = self.layernorm(sequence_output)
|
||||||
|
|
||||||
return sequence_output
|
return sequence_output
|
||||||
@ -31,7 +31,7 @@ class Timesteps(nn.Module):
|
|||||||
|
|
||||||
def forward(self, timesteps: torch.Tensor):
|
def forward(self, timesteps: torch.Tensor):
|
||||||
|
|
||||||
x = timesteps.float().unsqueeze(1) * self.inv_freq.unsqueeze(0)
|
x = timesteps.float().unsqueeze(1) * self.inv_freq.to(timesteps.device).unsqueeze(0)
|
||||||
|
|
||||||
# scale factor
|
# scale factor
|
||||||
if self.scale != 1.0:
|
if self.scale != 1.0:
|
||||||
@ -73,7 +73,7 @@ class TimestepEmbedder(nn.Module):
|
|||||||
cond_embed = self.cond_proj(condition)
|
cond_embed = self.cond_proj(condition)
|
||||||
timestep_embed = timestep_embed + cond_embed
|
timestep_embed = timestep_embed + cond_embed
|
||||||
|
|
||||||
time_conditioned = self.mlp(timestep_embed)
|
time_conditioned = self.mlp(timestep_embed.to(self.mlp[0].weight.device))
|
||||||
|
|
||||||
# for broadcasting with image tokens
|
# for broadcasting with image tokens
|
||||||
return time_conditioned.unsqueeze(1)
|
return time_conditioned.unsqueeze(1)
|
||||||
|
|||||||
@ -5,81 +5,178 @@ from PIL import Image
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
|
||||||
def to_tensor(image_pt):
|
def to_tensor(image_pt):
|
||||||
|
|
||||||
image_pt = image_pt / 255 * 2 - 1
|
image_pt = image_pt / 255 * 2 - 1
|
||||||
if image_pt.dim() == 4:
|
|
||||||
image_pt = image_pt.permute(0, 3, 1, 2)
|
|
||||||
else:
|
|
||||||
image_pt = image_pt.permute(2, 1, 0)
|
|
||||||
|
|
||||||
return image_pt
|
return image_pt
|
||||||
|
|
||||||
def resize_bilinear(img: torch.Tensor, size: int) -> torch.Tensor:
|
def resize_nearest(img: torch.Tensor, size: int) -> torch.Tensor:
|
||||||
# pytorch implementation of cv2.INTER_LINEAR
|
|
||||||
|
|
||||||
batched = (img.ndim == 4)
|
batched = (img.ndim == 4)
|
||||||
if img.ndim == 3:
|
if img.ndim == 3:
|
||||||
img = img.unsqueeze(0)
|
img = img.unsqueeze(0)
|
||||||
|
|
||||||
B, _, H, W = img.shape
|
img = img.permute(0, 3, 1, 2)
|
||||||
H_out = W_out = size
|
|
||||||
|
|
||||||
xs = torch.linspace(0, H_out - 1, H_out, device = img.device)
|
out = F.interpolate(img, size=size, mode='nearest')
|
||||||
ys = torch.linspace(0, W_out - 1, W_out, device = img.device)
|
|
||||||
|
|
||||||
xs = (xs + 0.5) * (H / H_out) - 0.5
|
if not batched:
|
||||||
ys = (ys + 0.5) * (W / W_out) - 0.5
|
out = out.squeeze(0)
|
||||||
|
|
||||||
# normalize
|
return out
|
||||||
xs = 2 * xs / (H - 1) - 1
|
|
||||||
ys = 2 * ys / (W - 1) - 1
|
|
||||||
|
|
||||||
# meshgrid in “ij” order: first rows (xs), then cols (ys)
|
def cubic_kernel(x, a: float = -0.75):
|
||||||
grid_i, grid_j = torch.meshgrid(xs, ys, indexing='ij')
|
absx = x.abs()
|
||||||
|
absx2 = absx ** 2
|
||||||
|
absx3 = absx ** 3
|
||||||
|
|
||||||
# stack into (x,y) where x=columns, y=rows
|
w = (a + 2) * absx3 - (a + 3) * absx2 + 1
|
||||||
grid = torch.stack((grid_j, grid_i), dim=-1)
|
w2 = a * absx3 - 5*a * absx2 + 8*a * absx - 4*a
|
||||||
grid = grid.unsqueeze(0).expand(B, -1, -1, -1)
|
|
||||||
|
|
||||||
out = F.grid_sample(img, grid, mode = 'bilinear',
|
return torch.where(absx <= 1, w, torch.where(absx < 2, w2, torch.zeros_like(x)))
|
||||||
padding_mode = 'zeros', align_corners = True)
|
|
||||||
|
|
||||||
return out if batched else out.squeeze(0)
|
|
||||||
|
|
||||||
def resize_bicubic(img: torch.Tensor, size: int) -> torch.Tensor:
|
def get_indices_weights(in_size, out_size, scale):
|
||||||
# pytorch implementation of INTER_CUBIC
|
# OpenCV-style half-pixel mapping
|
||||||
|
x = torch.arange(out_size, dtype=torch.float32)
|
||||||
|
x = (x + 0.5) / scale - 0.5
|
||||||
|
|
||||||
|
x0 = x.floor().long()
|
||||||
|
dx = x.unsqueeze(1) - (x0.unsqueeze(1) + torch.arange(-1, 3))
|
||||||
|
|
||||||
|
weights = cubic_kernel(dx)
|
||||||
|
weights = weights / weights.sum(dim=1, keepdim=True)
|
||||||
|
|
||||||
|
indices = x0.unsqueeze(1) + torch.arange(-1, 3)
|
||||||
|
indices = indices.clamp(0, in_size - 1)
|
||||||
|
|
||||||
|
return indices, weights
|
||||||
|
|
||||||
|
|
||||||
|
def resize_cubic_1d(x, out_size, dim):
|
||||||
|
b, c, h, w = x.shape
|
||||||
|
in_size = h if dim == 2 else w
|
||||||
|
scale = out_size / in_size
|
||||||
|
|
||||||
|
indices, weights = get_indices_weights(in_size, out_size, scale)
|
||||||
|
|
||||||
|
if dim == 2:
|
||||||
|
x = x.permute(0, 1, 3, 2)
|
||||||
|
x = x.reshape(-1, h)
|
||||||
|
else:
|
||||||
|
x = x.reshape(-1, w)
|
||||||
|
|
||||||
|
gathered = x[:, indices]
|
||||||
|
out = (gathered * weights.unsqueeze(0)).sum(dim=2)
|
||||||
|
|
||||||
|
if dim == 2:
|
||||||
|
out = out.reshape(b, c, w, out_size).permute(0, 1, 3, 2)
|
||||||
|
else:
|
||||||
|
out = out.reshape(b, c, h, out_size)
|
||||||
|
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def resize_cubic(img: torch.Tensor, size: tuple) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
Resize image using OpenCV-equivalent INTER_CUBIC interpolation.
|
||||||
|
Implemented in pure PyTorch
|
||||||
|
"""
|
||||||
|
|
||||||
was_batched = img.ndim == 4
|
|
||||||
if img.ndim == 3:
|
if img.ndim == 3:
|
||||||
img = img.unsqueeze(0)
|
img = img.unsqueeze(0)
|
||||||
|
|
||||||
out = F.interpolate(
|
img = img.permute(0, 3, 1, 2)
|
||||||
img.permute(0, 3, 2, 1),
|
|
||||||
size = (size, size),
|
|
||||||
mode = "bicubic",
|
|
||||||
align_corners = True
|
|
||||||
)
|
|
||||||
|
|
||||||
return out if was_batched else out.squeeze(0)
|
out_h, out_w = size
|
||||||
|
img = resize_cubic_1d(img, out_h, dim=2)
|
||||||
|
img = resize_cubic_1d(img, out_w, dim=3)
|
||||||
|
return img
|
||||||
|
|
||||||
def resize_area(img: torch.Tensor, size: tuple) -> torch.Tensor:
|
def resize_area(img: torch.Tensor, size: tuple) -> torch.Tensor:
|
||||||
# pytorch implementation of INTER_AREA
|
# vectorized implementation for OpenCV's INTER_AREA using pure PyTorch
|
||||||
|
original_shape = img.shape
|
||||||
|
is_hwc = False
|
||||||
|
|
||||||
was_batched = img.ndim == 4
|
|
||||||
if img.ndim == 3:
|
if img.ndim == 3:
|
||||||
img = img.unsqueeze(0)
|
if img.shape[0] <= 4:
|
||||||
|
img = img.unsqueeze(0)
|
||||||
image = F.interpolate(img.permute(0,3,1,2).float(), (size[1], size[0]), mode = "area")
|
else:
|
||||||
|
is_hwc = True
|
||||||
if was_batched:
|
img = img.permute(2, 0, 1).unsqueeze(0)
|
||||||
image = image.permute(0, 2, 3, 1) # return to channel last
|
elif img.ndim == 4:
|
||||||
|
pass
|
||||||
else:
|
else:
|
||||||
image = image.squeeze(0).permute(1, 2, 0)
|
raise ValueError("Expected image with 3 or 4 dims.")
|
||||||
|
|
||||||
|
B, C, H, W = img.shape
|
||||||
|
out_h, out_w = size
|
||||||
|
scale_y = H / out_h
|
||||||
|
scale_x = W / out_w
|
||||||
|
|
||||||
|
device = img.device
|
||||||
|
|
||||||
|
# compute the grid boundries
|
||||||
|
y_start = torch.arange(out_h, device=device).float() * scale_y
|
||||||
|
y_end = y_start + scale_y
|
||||||
|
x_start = torch.arange(out_w, device=device).float() * scale_x
|
||||||
|
x_end = x_start + scale_x
|
||||||
|
|
||||||
|
# for each output pixel, we will compute the range for it
|
||||||
|
y_start_int = torch.floor(y_start).long()
|
||||||
|
y_end_int = torch.ceil(y_end).long()
|
||||||
|
x_start_int = torch.floor(x_start).long()
|
||||||
|
x_end_int = torch.ceil(x_end).long()
|
||||||
|
|
||||||
|
# We will build the weighted sums by iterating over contributing input pixels once
|
||||||
|
output = torch.zeros((B, C, out_h, out_w), dtype=torch.float32, device=device)
|
||||||
|
area = torch.zeros((out_h, out_w), dtype=torch.float32, device=device)
|
||||||
|
|
||||||
|
max_kernel_h = int(torch.max(y_end_int - y_start_int).item())
|
||||||
|
max_kernel_w = int(torch.max(x_end_int - x_start_int).item())
|
||||||
|
|
||||||
|
for dy in range(max_kernel_h):
|
||||||
|
for dx in range(max_kernel_w):
|
||||||
|
# compute the weights for this offset for all output pixels
|
||||||
|
|
||||||
|
y_idx = y_start_int.unsqueeze(1) + dy
|
||||||
|
x_idx = x_start_int.unsqueeze(0) + dx
|
||||||
|
|
||||||
|
# clamp indices to image boundaries
|
||||||
|
y_idx_clamped = torch.clamp(y_idx, 0, H - 1)
|
||||||
|
x_idx_clamped = torch.clamp(x_idx, 0, W - 1)
|
||||||
|
|
||||||
|
# compute weights by broadcasting
|
||||||
|
y_weight = (torch.min(y_end.unsqueeze(1), y_idx_clamped.float() + 1.0) - torch.max(y_start.unsqueeze(1), y_idx_clamped.float())).clamp(min=0)
|
||||||
|
x_weight = (torch.min(x_end.unsqueeze(0), x_idx_clamped.float() + 1.0) - torch.max(x_start.unsqueeze(0), x_idx_clamped.float())).clamp(min=0)
|
||||||
|
|
||||||
|
weight = (y_weight * x_weight)
|
||||||
|
|
||||||
|
y_expand = y_idx_clamped.expand(out_h, out_w)
|
||||||
|
x_expand = x_idx_clamped.expand(out_h, out_w)
|
||||||
|
|
||||||
|
|
||||||
|
pixels = img[:, :, y_expand, x_expand]
|
||||||
|
|
||||||
|
# unsqueeze to broadcast
|
||||||
|
w = weight.unsqueeze(0).unsqueeze(0)
|
||||||
|
|
||||||
|
output += pixels * w
|
||||||
|
area += weight
|
||||||
|
|
||||||
|
# Normalize by area
|
||||||
|
output /= area.unsqueeze(0).unsqueeze(0)
|
||||||
|
|
||||||
|
if is_hwc:
|
||||||
|
return output[0].permute(1, 2, 0)
|
||||||
|
elif img.shape[0] == 1 and original_shape[0] <= 4:
|
||||||
|
return output[0]
|
||||||
|
else:
|
||||||
|
return output
|
||||||
|
|
||||||
return image if was_batched else image.squeeze(0)
|
|
||||||
|
|
||||||
class ImageProcessorV2(nn.Module):
|
class ImageProcessorV2(nn.Module):
|
||||||
def __init__(self, size: int = 512, border_ratio: float = None):
|
def __init__(self, size: int = 512, border_ratio: float = None):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
self.size = size
|
self.size = size
|
||||||
self.border_ratio = border_ratio
|
self.border_ratio = border_ratio
|
||||||
|
|
||||||
@ -98,12 +195,14 @@ class ImageProcessorV2(nn.Module):
|
|||||||
img = torch.from_numpy(img)
|
img = torch.from_numpy(img)
|
||||||
img, mask = self.recenter(img, border_ratio = border_ratio)
|
img, mask = self.recenter(img, border_ratio = border_ratio)
|
||||||
|
|
||||||
img = resize_bicubic(img, size = self.size)
|
img = resize_cubic(img, size = (self.size, self.size))
|
||||||
mask = resize_bilinear(mask.float(), size = self.size)
|
mask = resize_nearest(mask.float(), size = self.size)
|
||||||
mask = mask[..., torch.newaxis]
|
mask = mask[..., torch.newaxis]
|
||||||
|
|
||||||
img = to_tensor(img)
|
img = to_tensor(img)
|
||||||
|
|
||||||
mask = to_tensor(mask)
|
mask = to_tensor(mask)
|
||||||
|
mask = mask.permute(0, 3, 1, 2)
|
||||||
|
|
||||||
return img, mask
|
return img, mask
|
||||||
|
|
||||||
@ -147,7 +246,7 @@ class ImageProcessorV2(nn.Module):
|
|||||||
y2_max = y2_min + w2
|
y2_max = y2_min + w2
|
||||||
|
|
||||||
# note: opencv takes columns first (opposite to pytorch and numpy that take the row first)
|
# note: opencv takes columns first (opposite to pytorch and numpy that take the row first)
|
||||||
result[x2_min:x2_max, y2_min:y2_max] = resize_area(image[x_min:x_max, y_min:y_max], (w2, h2))
|
result[x2_min:x2_max, y2_min:y2_max] = resize_area(image[x_min:x_max, y_min:y_max], (h2, w2))
|
||||||
|
|
||||||
bg = torch.ones((result.shape[0], result.shape[1], 3), dtype = torch.uint8) * 255
|
bg = torch.ones((result.shape[0], result.shape[1], 3), dtype = torch.uint8) * 255
|
||||||
|
|
||||||
@ -175,18 +274,22 @@ class ImageProcessorV2(nn.Module):
|
|||||||
return outputs
|
return outputs
|
||||||
|
|
||||||
def test_image_processor():
|
def test_image_processor():
|
||||||
|
|
||||||
|
"""
|
||||||
|
implementation speed: 0.24465346336364746
|
||||||
|
reference speed: 2.046062469482422
|
||||||
|
atol = 4e-2: True
|
||||||
|
"""
|
||||||
|
|
||||||
|
import time
|
||||||
import matplotlib.pyplot as plt
|
import matplotlib.pyplot as plt
|
||||||
|
|
||||||
image_processor = ImageProcessorV2(size = 224)
|
image_processor = ImageProcessorV2(size = 224)
|
||||||
import time
|
|
||||||
start = time.time()
|
start = time.time()
|
||||||
outputs = image_processor(image = r"C:\Users\yrafa\Work\Hunyuan 3D\cat.jpg")
|
outputs = image_processor(image = r"C:\Users\yrafa\Work\Hunyuan 3D\cat.jpg")
|
||||||
print(time.time() - start)
|
print(time.time() - start)
|
||||||
image = outputs["image"]
|
image = outputs["image"]
|
||||||
print(image.shape)
|
print(image.shape)
|
||||||
plt.imshow(image)
|
plt.imshow(image.squeeze().permute(1, 2, 0).numpy())
|
||||||
plt.axis("off")
|
plt.axis("off")
|
||||||
plt.show()
|
plt.show()
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
test_image_processor()
|
|
||||||
@ -5,6 +5,7 @@ from PIL import Image
|
|||||||
from typing import List, Union
|
from typing import List, Union
|
||||||
from torch.utils._pytree import tree_map
|
from torch.utils._pytree import tree_map
|
||||||
from torch.utils.data._utils.collate import default_collate
|
from torch.utils.data._utils.collate import default_collate
|
||||||
|
from vae import VAE
|
||||||
|
|
||||||
def export_to_trimesh(mesh_output):
|
def export_to_trimesh(mesh_output):
|
||||||
if isinstance(mesh_output, list):
|
if isinstance(mesh_output, list):
|
||||||
@ -23,19 +24,30 @@ def export_to_trimesh(mesh_output):
|
|||||||
return mesh_output
|
return mesh_output
|
||||||
|
|
||||||
class Hunyuan3DDiTFlowMatchingPipeline(nn.Module):
|
class Hunyuan3DDiTFlowMatchingPipeline(nn.Module):
|
||||||
def __init__(self, model, vae, conditioner, image_processor, scheduler):
|
def __init__(self, model, vae, conditioner, image_processor, scheduler, device, dtype):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
self.vae = vae
|
self.vae = vae
|
||||||
self.model = model
|
self.model = model
|
||||||
self.conditioner = conditioner
|
self.conditioner = conditioner
|
||||||
self.image_processor = image_processor
|
self.image_processor = image_processor
|
||||||
self.scheduler = scheduler
|
self.scheduler = scheduler
|
||||||
|
self.device = device
|
||||||
|
self.dtype = dtype
|
||||||
|
|
||||||
def compile(self):
|
def compile(self):
|
||||||
self.vae = torch.compile(self.vae)
|
self.vae = torch.compile(self.vae)
|
||||||
self.model = torch.compile(self.model)
|
self.model = torch.compile(self.model)
|
||||||
self.conditioner = torch.compile(self.conditioner)
|
self.conditioner = torch.compile(self.conditioner)
|
||||||
|
|
||||||
|
def load_ckpt(self, checkpoint_path: str):
|
||||||
|
|
||||||
|
checkpoint = torch.load(checkpoint_path, weights_only = True)
|
||||||
|
self.model.load_state_dict(checkpoint["model"])
|
||||||
|
self.vae.load_state_dict(checkpoint["vae"])
|
||||||
|
self.conditioner.load_state_dict(checkpoint["conditioner"])
|
||||||
|
|
||||||
|
|
||||||
def encode_cond(self, image, additional_cond_inputs, do_classifier_free_guidance):
|
def encode_cond(self, image, additional_cond_inputs, do_classifier_free_guidance):
|
||||||
|
|
||||||
bsz = image.shape[0]
|
bsz = image.shape[0]
|
||||||
@ -51,8 +63,23 @@ class Hunyuan3DDiTFlowMatchingPipeline(nn.Module):
|
|||||||
|
|
||||||
return cond
|
return cond
|
||||||
|
|
||||||
|
def to(self, device=None, dtype=None):
|
||||||
|
if dtype is not None:
|
||||||
|
self.dtype = dtype
|
||||||
|
self.vae.to(dtype=dtype)
|
||||||
|
self.model.to(dtype=dtype)
|
||||||
|
self.conditioner.to(dtype=dtype)
|
||||||
|
if device is not None:
|
||||||
|
self.device = torch.device(device)
|
||||||
|
self.vae.to(device)
|
||||||
|
self.model.to(device)
|
||||||
|
self.conditioner.to(device)
|
||||||
|
|
||||||
def prepare_images(self, images):
|
def prepare_images(self, images):
|
||||||
|
|
||||||
|
if isinstance(images, (str, Image.Image)):
|
||||||
|
return self.image_processor(images)
|
||||||
|
|
||||||
outputs = []
|
outputs = []
|
||||||
for image in images:
|
for image in images:
|
||||||
output = self.image_processor(image)
|
output = self.image_processor(image)
|
||||||
@ -107,7 +134,7 @@ class Hunyuan3DDiTFlowMatchingPipeline(nn.Module):
|
|||||||
self.model.guidance_embed is True
|
self.model.guidance_embed is True
|
||||||
)
|
)
|
||||||
|
|
||||||
cond_inputs = self.prepare_image(image)
|
cond_inputs = self.prepare_images(image)
|
||||||
image = cond_inputs.pop('image')
|
image = cond_inputs.pop('image')
|
||||||
|
|
||||||
cond = self.encode_cond(
|
cond = self.encode_cond(
|
||||||
@ -135,7 +162,7 @@ class Hunyuan3DDiTFlowMatchingPipeline(nn.Module):
|
|||||||
latent_model_input = latents
|
latent_model_input = latents
|
||||||
|
|
||||||
timestep = t.expand(latent_model_input.shape[0]).to(latents.dtype)
|
timestep = t.expand(latent_model_input.shape[0]).to(latents.dtype)
|
||||||
timestep = timestep / self.scheduler.num_train_timesteps
|
timestep = timestep / self.scheduler.num_training_timesteps
|
||||||
noise_pred = self.model(latent_model_input, timestep, cond, guidance=guidance)
|
noise_pred = self.model(latent_model_input, timestep, cond, guidance=guidance)
|
||||||
|
|
||||||
if do_classifier_free_guidance:
|
if do_classifier_free_guidance:
|
||||||
@ -143,7 +170,7 @@ class Hunyuan3DDiTFlowMatchingPipeline(nn.Module):
|
|||||||
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_cond - noise_pred_uncond)
|
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_cond - noise_pred_uncond)
|
||||||
|
|
||||||
# compute the previous noisy sample x_t -> x_t-1
|
# compute the previous noisy sample x_t -> x_t-1
|
||||||
latents = self.scheduler.step(noise_pred, t, latents)
|
latents = self.scheduler.reverse_flow(noise_pred, latents)
|
||||||
|
|
||||||
if callback is not None and i % callback_steps == 0:
|
if callback is not None and i % callback_steps == 0:
|
||||||
step_idx = i // getattr(self.scheduler, "order", 1)
|
step_idx = i // getattr(self.scheduler, "order", 1)
|
||||||
@ -154,3 +181,17 @@ class Hunyuan3DDiTFlowMatchingPipeline(nn.Module):
|
|||||||
|
|
||||||
return export_to_trimesh(mesh)
|
return export_to_trimesh(mesh)
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
from scheduler import EulerScheduler
|
||||||
|
from conditioner import SingleImageEncoder
|
||||||
|
from image_processor import ImageProcessorV2
|
||||||
|
from dinov2 import DinoConfig
|
||||||
|
from hunyuandit import HunYuanDiTPlain
|
||||||
|
|
||||||
|
model = HunYuanDiTPlain(depth = 2)
|
||||||
|
|
||||||
|
pipeline = Hunyuan3DDiTFlowMatchingPipeline(vae = VAE(), scheduler = EulerScheduler(), model = model,
|
||||||
|
conditioner = SingleImageEncoder(DinoConfig()), image_processor = ImageProcessorV2(),
|
||||||
|
device = "cpu", dtype = torch.bfloat16)
|
||||||
|
img = r"C:\Users\yrafa\Work\Hunyuan 3D\cat.jpg"
|
||||||
|
print(pipeline(img))
|
||||||
@ -2,7 +2,7 @@ import torch
|
|||||||
|
|
||||||
class EulerScheduler(torch.nn.Module):
|
class EulerScheduler(torch.nn.Module):
|
||||||
def __init__(self, num_training_timesteps: int = 1_000, shift: float = 1,
|
def __init__(self, num_training_timesteps: int = 1_000, shift: float = 1,
|
||||||
num_inference_timesteps: int = 100, inference: bool = False):
|
num_inference_timesteps: int = 50, inference: bool = True):
|
||||||
super(EulerScheduler, self).__init__()
|
super(EulerScheduler, self).__init__()
|
||||||
|
|
||||||
# compute timestep values so we can index into them later
|
# compute timestep values so we can index into them later
|
||||||
|
|||||||
1010
comfy/ldm/hunyuan3d/model_/vae.py
Normal file
1010
comfy/ldm/hunyuan3d/model_/vae.py
Normal file
File diff suppressed because it is too large
Load Diff
@ -68,7 +68,7 @@ class VAE(nn.Module):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
self.latent_shape = (num_latents, embed_dim)
|
self.latent_shape = (num_latents, embed_dim)
|
||||||
self.scale = scale_factor
|
self.scale_factor = scale_factor
|
||||||
|
|
||||||
self.fourier_embedder = FourierEmbedder(num_freq = num_frequencies, include_pi = include_pi)
|
self.fourier_embedder = FourierEmbedder(num_freq = num_frequencies, include_pi = include_pi)
|
||||||
|
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user