|
|
|
|
@@ -16,6 +16,13 @@ from toolkit.models.sapiens2 import Sapiens2
|
|
|
|
|
import huggingface_hub
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _fold_frames_to_batch(x: torch.Tensor) -> torch.Tensor:
|
|
|
|
|
"""(B, C, T, H, W) -> (B*T, C, H, W), each sample's frames contiguous -- so the 2D
|
|
|
|
|
feature losses run on EVERY frame of a video instead of only the first."""
|
|
|
|
|
b, c, t, h, w = x.shape
|
|
|
|
|
return x.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class ResBlock(nn.Module):
|
|
|
|
|
def __init__(self, in_channels, out_channels):
|
|
|
|
|
super().__init__()
|
|
|
|
|
@@ -521,18 +528,19 @@ class DiffusionFeatureExtractor4(nn.Module):
|
|
|
|
|
is_video = False
|
|
|
|
|
# stack time for video models on the batch dimension
|
|
|
|
|
if len(noise_pred.shape) == 5:
|
|
|
|
|
# B, C, T, H, W = images.shape
|
|
|
|
|
# only take first time
|
|
|
|
|
noise = noise[:, :, 0, :, :]
|
|
|
|
|
noise_pred = noise_pred[:, :, 0, :, :]
|
|
|
|
|
noisy_latents = noisy_latents[:, :, 0, :, :]
|
|
|
|
|
# (B, C, T, H, W): fold every frame into the batch dim so the loss covers all
|
|
|
|
|
# frames, and repeat the per-sample timestep for each of its frames
|
|
|
|
|
num_frames = noise_pred.shape[2]
|
|
|
|
|
noise = _fold_frames_to_batch(noise)
|
|
|
|
|
noise_pred = _fold_frames_to_batch(noise_pred)
|
|
|
|
|
noisy_latents = _fold_frames_to_batch(noisy_latents)
|
|
|
|
|
timesteps = timesteps.repeat_interleave(num_frames)
|
|
|
|
|
is_video = True
|
|
|
|
|
|
|
|
|
|
if len(tensors.shape) == 5:
|
|
|
|
|
# batch is different
|
|
|
|
|
# (B, T, C, H, W)
|
|
|
|
|
# only take first time
|
|
|
|
|
tensors = tensors[:, 0, :, :, :]
|
|
|
|
|
# batch tensor is frames-first (B, T, C, H, W): fold to (B*T, C, H, W), matching
|
|
|
|
|
# the frame order of the folded predictions above
|
|
|
|
|
tensors = tensors.reshape(-1, *tensors.shape[2:])
|
|
|
|
|
|
|
|
|
|
if model is not None and hasattr(model, 'get_stepped_pred'):
|
|
|
|
|
stepped_latents = model.get_stepped_pred(noise_pred, noise)
|
|
|
|
|
@@ -745,18 +753,19 @@ class DiffusionFeatureExtractor6(nn.Module):
|
|
|
|
|
is_video = False
|
|
|
|
|
# stack time for video models on the batch dimension
|
|
|
|
|
if len(noise_pred.shape) == 5:
|
|
|
|
|
# B, C, T, H, W = images.shape
|
|
|
|
|
# only take first time
|
|
|
|
|
noise = noise[:, :, 0, :, :]
|
|
|
|
|
noise_pred = noise_pred[:, :, 0, :, :]
|
|
|
|
|
noisy_latents = noisy_latents[:, :, 0, :, :]
|
|
|
|
|
# (B, C, T, H, W): fold every frame into the batch dim so the loss covers all
|
|
|
|
|
# frames, and repeat the per-sample timestep for each of its frames
|
|
|
|
|
num_frames = noise_pred.shape[2]
|
|
|
|
|
noise = _fold_frames_to_batch(noise)
|
|
|
|
|
noise_pred = _fold_frames_to_batch(noise_pred)
|
|
|
|
|
noisy_latents = _fold_frames_to_batch(noisy_latents)
|
|
|
|
|
timesteps = timesteps.repeat_interleave(num_frames)
|
|
|
|
|
is_video = True
|
|
|
|
|
|
|
|
|
|
if len(tensors.shape) == 5:
|
|
|
|
|
# batch is different
|
|
|
|
|
# (B, T, C, H, W)
|
|
|
|
|
# only take first time
|
|
|
|
|
tensors = tensors[:, 0, :, :, :]
|
|
|
|
|
# batch tensor is frames-first (B, T, C, H, W): fold to (B*T, C, H, W), matching
|
|
|
|
|
# the frame order of the folded predictions above
|
|
|
|
|
tensors = tensors.reshape(-1, *tensors.shape[2:])
|
|
|
|
|
|
|
|
|
|
with torch.no_grad():
|
|
|
|
|
tv = timesteps.to(noise_pred.device).to(noise_pred.dtype) / 1000.0
|
|
|
|
|
@@ -914,18 +923,19 @@ class DiffusionFeatureExtractor7(nn.Module):
|
|
|
|
|
is_video = False
|
|
|
|
|
# stack time for video models on the batch dimension
|
|
|
|
|
if len(noise_pred.shape) == 5:
|
|
|
|
|
# B, C, T, H, W = images.shape
|
|
|
|
|
# only take first time
|
|
|
|
|
noise = noise[:, :, 0, :, :]
|
|
|
|
|
noise_pred = noise_pred[:, :, 0, :, :]
|
|
|
|
|
noisy_latents = noisy_latents[:, :, 0, :, :]
|
|
|
|
|
# (B, C, T, H, W): fold every frame into the batch dim so the loss covers all
|
|
|
|
|
# frames, and repeat the per-sample timestep for each of its frames
|
|
|
|
|
num_frames = noise_pred.shape[2]
|
|
|
|
|
noise = _fold_frames_to_batch(noise)
|
|
|
|
|
noise_pred = _fold_frames_to_batch(noise_pred)
|
|
|
|
|
noisy_latents = _fold_frames_to_batch(noisy_latents)
|
|
|
|
|
timesteps = timesteps.repeat_interleave(num_frames)
|
|
|
|
|
is_video = True
|
|
|
|
|
|
|
|
|
|
if len(tensors.shape) == 5:
|
|
|
|
|
# batch is different
|
|
|
|
|
# (B, T, C, H, W)
|
|
|
|
|
# only take first time
|
|
|
|
|
tensors = tensors[:, 0, :, :, :]
|
|
|
|
|
# batch tensor is frames-first (B, T, C, H, W): fold to (B*T, C, H, W), matching
|
|
|
|
|
# the frame order of the folded predictions above
|
|
|
|
|
tensors = tensors.reshape(-1, *tensors.shape[2:])
|
|
|
|
|
|
|
|
|
|
with torch.no_grad():
|
|
|
|
|
tv = timesteps.to(noise_pred.device).to(noise_pred.dtype) / 1000.0
|
|
|
|
|
@@ -937,8 +947,11 @@ class DiffusionFeatureExtractor7(nn.Module):
|
|
|
|
|
target_0_1 = (tensors + 1) / 2 # 0 to 1
|
|
|
|
|
|
|
|
|
|
if not self.do_partial_step:
|
|
|
|
|
# step latent
|
|
|
|
|
x0 = noisy_latents - tv * noise_pred
|
|
|
|
|
if getattr(self.sd_ref(), "x0_pred", False):
|
|
|
|
|
x0 = noise_pred
|
|
|
|
|
else:
|
|
|
|
|
# step latent
|
|
|
|
|
x0 = noisy_latents - tv * noise_pred
|
|
|
|
|
stepped_latents = x0
|
|
|
|
|
# min 0.001
|
|
|
|
|
tv = torch.clamp(tv, min=0.001)
|
|
|
|
|
@@ -952,6 +965,9 @@ class DiffusionFeatureExtractor7(nn.Module):
|
|
|
|
|
with torch.no_grad():
|
|
|
|
|
# make a noisy target at next timestep
|
|
|
|
|
target_latents = batch.latents.to(self.sd_ref().vae.device, dtype=self.sd_ref().vae.dtype)
|
|
|
|
|
if target_latents.dim() == 5:
|
|
|
|
|
# fold frames to match the folded noise/predictions
|
|
|
|
|
target_latents = _fold_frames_to_batch(target_latents)
|
|
|
|
|
# add noise
|
|
|
|
|
target_latents = (1.0 - next_step) * target_latents + next_step * noise
|
|
|
|
|
target_n1p1 = self.sd_ref().decode_latents(target_latents)
|
|
|
|
|
@@ -1110,18 +1126,19 @@ class DiffusionFeatureExtractor9(nn.Module):
|
|
|
|
|
is_video = False
|
|
|
|
|
# stack time for video models on the batch dimension
|
|
|
|
|
if len(noise_pred.shape) == 5:
|
|
|
|
|
# B, C, T, H, W = images.shape
|
|
|
|
|
# only take first time
|
|
|
|
|
noise = noise[:, :, 0, :, :]
|
|
|
|
|
noise_pred = noise_pred[:, :, 0, :, :]
|
|
|
|
|
noisy_latents = noisy_latents[:, :, 0, :, :]
|
|
|
|
|
# (B, C, T, H, W): fold every frame into the batch dim so the loss covers all
|
|
|
|
|
# frames, and repeat the per-sample timestep for each of its frames
|
|
|
|
|
num_frames = noise_pred.shape[2]
|
|
|
|
|
noise = _fold_frames_to_batch(noise)
|
|
|
|
|
noise_pred = _fold_frames_to_batch(noise_pred)
|
|
|
|
|
noisy_latents = _fold_frames_to_batch(noisy_latents)
|
|
|
|
|
timesteps = timesteps.repeat_interleave(num_frames)
|
|
|
|
|
is_video = True
|
|
|
|
|
|
|
|
|
|
if len(tensors.shape) == 5:
|
|
|
|
|
# batch is different
|
|
|
|
|
# (B, T, C, H, W)
|
|
|
|
|
# only take first time
|
|
|
|
|
tensors = tensors[:, 0, :, :, :]
|
|
|
|
|
# batch tensor is frames-first (B, T, C, H, W): fold to (B*T, C, H, W), matching
|
|
|
|
|
# the frame order of the folded predictions above
|
|
|
|
|
tensors = tensors.reshape(-1, *tensors.shape[2:])
|
|
|
|
|
|
|
|
|
|
with torch.no_grad():
|
|
|
|
|
tv = timesteps.to(noise_pred.device).to(noise_pred.dtype) / 1000.0
|
|
|
|
|
@@ -1133,8 +1150,11 @@ class DiffusionFeatureExtractor9(nn.Module):
|
|
|
|
|
target_0_1 = (tensors + 1) / 2 # 0 to 1
|
|
|
|
|
|
|
|
|
|
if not self.do_partial_step:
|
|
|
|
|
# step latent
|
|
|
|
|
x0 = noisy_latents - tv * noise_pred
|
|
|
|
|
if getattr(self.sd_ref(), "x0_pred", False):
|
|
|
|
|
x0 = noise_pred
|
|
|
|
|
else:
|
|
|
|
|
# step latent
|
|
|
|
|
x0 = noisy_latents - tv * noise_pred
|
|
|
|
|
stepped_latents = x0
|
|
|
|
|
# min 0.001
|
|
|
|
|
tv = torch.clamp(tv, min=0.001)
|
|
|
|
|
@@ -1144,10 +1164,13 @@ class DiffusionFeatureExtractor9(nn.Module):
|
|
|
|
|
next_step = tv - step
|
|
|
|
|
next_step = torch.clamp(next_step, min=0.0)
|
|
|
|
|
stepped_latents = noisy_latents + (next_step - tv) * noise_pred
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
with torch.no_grad():
|
|
|
|
|
# make a noisy target at next timestep
|
|
|
|
|
target_latents = batch.latents.to(self.sd_ref().vae.device, dtype=self.sd_ref().vae.dtype)
|
|
|
|
|
if target_latents.dim() == 5:
|
|
|
|
|
# fold frames to match the folded noise/predictions
|
|
|
|
|
target_latents = _fold_frames_to_batch(target_latents)
|
|
|
|
|
# add noise
|
|
|
|
|
target_latents = (1.0 - next_step) * target_latents + next_step * noise
|
|
|
|
|
target_n1p1 = self.sd_ref().decode_latents(target_latents)
|
|
|
|
|
@@ -1266,18 +1289,19 @@ class DiffusionFeatureExtractor10(nn.Module):
|
|
|
|
|
is_video = False
|
|
|
|
|
# stack time for video models on the batch dimension
|
|
|
|
|
if len(noise_pred.shape) == 5:
|
|
|
|
|
# B, C, T, H, W = images.shape
|
|
|
|
|
# only take first time
|
|
|
|
|
noise = noise[:, :, 0, :, :]
|
|
|
|
|
noise_pred = noise_pred[:, :, 0, :, :]
|
|
|
|
|
noisy_latents = noisy_latents[:, :, 0, :, :]
|
|
|
|
|
# (B, C, T, H, W): fold every frame into the batch dim so the loss covers all
|
|
|
|
|
# frames, and repeat the per-sample timestep for each of its frames
|
|
|
|
|
num_frames = noise_pred.shape[2]
|
|
|
|
|
noise = _fold_frames_to_batch(noise)
|
|
|
|
|
noise_pred = _fold_frames_to_batch(noise_pred)
|
|
|
|
|
noisy_latents = _fold_frames_to_batch(noisy_latents)
|
|
|
|
|
timesteps = timesteps.repeat_interleave(num_frames)
|
|
|
|
|
is_video = True
|
|
|
|
|
|
|
|
|
|
if len(tensors.shape) == 5:
|
|
|
|
|
# batch is different
|
|
|
|
|
# (B, T, C, H, W)
|
|
|
|
|
# only take first time
|
|
|
|
|
tensors = tensors[:, 0, :, :, :]
|
|
|
|
|
# batch tensor is frames-first (B, T, C, H, W): fold to (B*T, C, H, W), matching
|
|
|
|
|
# the frame order of the folded predictions above
|
|
|
|
|
tensors = tensors.reshape(-1, *tensors.shape[2:])
|
|
|
|
|
|
|
|
|
|
with torch.no_grad():
|
|
|
|
|
tv = timesteps.to(noise_pred.device).to(noise_pred.dtype) / 1000.0
|
|
|
|
|
@@ -1289,8 +1313,11 @@ class DiffusionFeatureExtractor10(nn.Module):
|
|
|
|
|
target_0_1 = (tensors + 1) / 2 # 0 to 1
|
|
|
|
|
|
|
|
|
|
if not self.do_partial_step:
|
|
|
|
|
# step latent
|
|
|
|
|
x0 = noisy_latents - tv * noise_pred
|
|
|
|
|
if getattr(self.sd_ref(), "x0_pred", False):
|
|
|
|
|
x0 = noise_pred
|
|
|
|
|
else:
|
|
|
|
|
# step latent
|
|
|
|
|
x0 = noisy_latents - tv * noise_pred
|
|
|
|
|
stepped_latents = x0
|
|
|
|
|
# min 0.001
|
|
|
|
|
tv = torch.clamp(tv, min=0.001)
|
|
|
|
|
@@ -1304,6 +1331,9 @@ class DiffusionFeatureExtractor10(nn.Module):
|
|
|
|
|
with torch.no_grad():
|
|
|
|
|
# make a noisy target at next timestep
|
|
|
|
|
target_latents = batch.latents.to(self.sd_ref().vae.device, dtype=self.sd_ref().vae.dtype)
|
|
|
|
|
if target_latents.dim() == 5:
|
|
|
|
|
# fold frames to match the folded noise/predictions
|
|
|
|
|
target_latents = _fold_frames_to_batch(target_latents)
|
|
|
|
|
# add noise
|
|
|
|
|
target_latents = (1.0 - next_step) * target_latents + next_step * noise
|
|
|
|
|
target_n1p1 = self.sd_ref().decode_latents(target_latents)
|
|
|
|
|
|