Add StableAnimator @ 0f3d85ad217c0d3edec89e310bb34c3ecb9eaf9b
https://github.com/Francis-Rings/StableAnimator/commit/0f3d85ad217c0d3edec89e310bb34c3ecb9eaf9b
@@ -0,0 +1,120 @@
|
||||
# Openpose
|
||||
# Original from CMU https://github.com/CMU-Perceptual-Computing-Lab/openpose
|
||||
# 2nd Edited by https://github.com/Hzzone/pytorch-openpose
|
||||
# 3rd Edited by ControlNet
|
||||
# 4th Edited by ControlNet (added face and correct hands)
|
||||
|
||||
import os
|
||||
os.environ["KMP_DUPLICATE_LIB_OK"]="TRUE"
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
from . import util
|
||||
from .wholebody import Wholebody
|
||||
|
||||
def draw_pose(pose, H, W):
|
||||
bodies = pose['bodies']
|
||||
faces = pose['faces']
|
||||
hands = pose['hands']
|
||||
candidate = bodies['candidate']
|
||||
subset = bodies['subset']
|
||||
canvas = np.zeros(shape=(H, W, 3), dtype=np.uint8)
|
||||
|
||||
canvas = util.draw_bodypose(canvas, candidate, subset)
|
||||
|
||||
canvas = util.draw_handpose(canvas, hands)
|
||||
|
||||
if faces is not None:
|
||||
canvas = util.draw_facepose(canvas, faces)
|
||||
|
||||
return canvas
|
||||
|
||||
|
||||
class DWposeDetector:
|
||||
def __init__(self):
|
||||
|
||||
self.pose_estimation = Wholebody()
|
||||
|
||||
def __call__(self, oriImg, remain_face=True):
|
||||
oriImg = oriImg.copy()
|
||||
H, W, C = oriImg.shape
|
||||
with torch.no_grad():
|
||||
candidate, subset = self.pose_estimation(oriImg)
|
||||
nums, keys, locs = candidate.shape
|
||||
candidate[..., 0] /= float(W)
|
||||
candidate[..., 1] /= float(H)
|
||||
body = candidate[:,:18].copy()
|
||||
body = body.reshape(nums*18, locs)
|
||||
score = subset[:,:18]
|
||||
for i in range(len(score)):
|
||||
for j in range(len(score[i])):
|
||||
if score[i][j] > 0.3:
|
||||
score[i][j] = int(18*i+j)
|
||||
else:
|
||||
score[i][j] = -1
|
||||
|
||||
un_visible = subset<0.3
|
||||
candidate[un_visible] = -1
|
||||
|
||||
foot = candidate[:,18:24]
|
||||
|
||||
faces = candidate[:,24:92]
|
||||
|
||||
hands = candidate[:,92:113]
|
||||
hands = np.vstack([hands, candidate[:,113:]])
|
||||
|
||||
bodies = dict(candidate=body, subset=score)
|
||||
if remain_face:
|
||||
pose = dict(bodies=bodies, hands=hands, faces=faces)
|
||||
else:
|
||||
pose = dict(bodies=bodies, hands=hands, faces=None)
|
||||
|
||||
return draw_pose(pose, H, W)
|
||||
|
||||
|
||||
class DWposeDetectorOnlyOnePerson:
|
||||
def __init__(self):
|
||||
|
||||
self.pose_estimation = Wholebody()
|
||||
|
||||
def __call__(self, oriImg, remain_face=True):
|
||||
oriImg = oriImg.copy()
|
||||
H, W, C = oriImg.shape
|
||||
with torch.no_grad():
|
||||
candidate, subset = self.pose_estimation(oriImg)
|
||||
|
||||
if len(subset) > 1:
|
||||
candidate = candidate[0][np.newaxis, :]
|
||||
subset = subset[:1]
|
||||
|
||||
nums, keys, locs = candidate.shape
|
||||
candidate[..., 0] /= float(W)
|
||||
candidate[..., 1] /= float(H)
|
||||
body = candidate[:, :18].copy()
|
||||
body = body.reshape(nums * 18, locs)
|
||||
score = subset[:, :18]
|
||||
|
||||
for i in range(len(score)):
|
||||
for j in range(len(score[i])):
|
||||
if score[i][j] > 0.3:
|
||||
score[i][j] = int(18 * i + j)
|
||||
else:
|
||||
score[i][j] = -1
|
||||
|
||||
# un_visible = subset < 0.3
|
||||
# candidate[un_visible] = -1
|
||||
|
||||
foot = candidate[:, 18:24]
|
||||
|
||||
faces = candidate[:, 24:92]
|
||||
|
||||
hands = candidate[:, 92:113]
|
||||
hands = np.vstack([hands, candidate[:, 113:]])
|
||||
|
||||
bodies = dict(candidate=body, subset=score)
|
||||
if remain_face:
|
||||
pose = dict(bodies=bodies, hands=hands, faces=faces)
|
||||
else:
|
||||
pose = dict(bodies=bodies, hands=hands, faces=None)
|
||||
|
||||
return draw_pose(pose, H, W)
|
||||
@@ -0,0 +1,57 @@
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from .wholebody import Wholebody
|
||||
|
||||
os.environ["KMP_DUPLICATE_LIB_OK"] = "TRUE"
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
class DWposeDetectorAligned:
|
||||
def __init__(self, device='cpu'):
|
||||
self.pose_estimation = Wholebody()
|
||||
|
||||
def release_memory(self):
|
||||
if hasattr(self, 'pose_estimation'):
|
||||
del self.pose_estimation
|
||||
import gc; gc.collect()
|
||||
|
||||
def __call__(self, oriImg):
|
||||
oriImg = oriImg.copy()
|
||||
H, W, C = oriImg.shape
|
||||
with torch.no_grad():
|
||||
candidate, score = self.pose_estimation(oriImg)
|
||||
nums, _, locs = candidate.shape
|
||||
candidate[..., 0] /= float(W)
|
||||
candidate[..., 1] /= float(H)
|
||||
body = candidate[:, :18].copy()
|
||||
body = body.reshape(nums * 18, locs)
|
||||
subset = score[:, :18].copy()
|
||||
for i in range(len(subset)):
|
||||
for j in range(len(subset[i])):
|
||||
if subset[i][j] > 0.3:
|
||||
subset[i][j] = int(18 * i + j)
|
||||
else:
|
||||
subset[i][j] = -1
|
||||
|
||||
# un_visible = subset < 0.3
|
||||
# candidate[un_visible] = -1
|
||||
|
||||
# foot = candidate[:, 18:24]
|
||||
|
||||
faces = candidate[:, 24:92]
|
||||
|
||||
hands = candidate[:, 92:113]
|
||||
hands = np.vstack([hands, candidate[:, 113:]])
|
||||
|
||||
faces_score = score[:, 24:92]
|
||||
hands_score = np.vstack([score[:, 92:113], score[:, 113:]])
|
||||
|
||||
bodies = dict(candidate=body, subset=subset, score=score[:, :18])
|
||||
pose = dict(bodies=bodies, hands=hands, hands_score=hands_score, faces=faces, faces_score=faces_score)
|
||||
|
||||
return pose
|
||||
|
||||
|
||||
dwpose_detector_aligned = DWposeDetectorAligned(device=device)
|
||||
@@ -0,0 +1,125 @@
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
import onnxruntime
|
||||
|
||||
def nms(boxes, scores, nms_thr):
|
||||
"""Single class NMS implemented in Numpy."""
|
||||
x1 = boxes[:, 0]
|
||||
y1 = boxes[:, 1]
|
||||
x2 = boxes[:, 2]
|
||||
y2 = boxes[:, 3]
|
||||
|
||||
areas = (x2 - x1 + 1) * (y2 - y1 + 1)
|
||||
order = scores.argsort()[::-1]
|
||||
|
||||
keep = []
|
||||
while order.size > 0:
|
||||
i = order[0]
|
||||
keep.append(i)
|
||||
xx1 = np.maximum(x1[i], x1[order[1:]])
|
||||
yy1 = np.maximum(y1[i], y1[order[1:]])
|
||||
xx2 = np.minimum(x2[i], x2[order[1:]])
|
||||
yy2 = np.minimum(y2[i], y2[order[1:]])
|
||||
|
||||
w = np.maximum(0.0, xx2 - xx1 + 1)
|
||||
h = np.maximum(0.0, yy2 - yy1 + 1)
|
||||
inter = w * h
|
||||
ovr = inter / (areas[i] + areas[order[1:]] - inter)
|
||||
|
||||
inds = np.where(ovr <= nms_thr)[0]
|
||||
order = order[inds + 1]
|
||||
|
||||
return keep
|
||||
|
||||
def multiclass_nms(boxes, scores, nms_thr, score_thr):
|
||||
"""Multiclass NMS implemented in Numpy. Class-aware version."""
|
||||
final_dets = []
|
||||
num_classes = scores.shape[1]
|
||||
for cls_ind in range(num_classes):
|
||||
cls_scores = scores[:, cls_ind]
|
||||
valid_score_mask = cls_scores > score_thr
|
||||
if valid_score_mask.sum() == 0:
|
||||
continue
|
||||
else:
|
||||
valid_scores = cls_scores[valid_score_mask]
|
||||
valid_boxes = boxes[valid_score_mask]
|
||||
keep = nms(valid_boxes, valid_scores, nms_thr)
|
||||
if len(keep) > 0:
|
||||
cls_inds = np.ones((len(keep), 1)) * cls_ind
|
||||
dets = np.concatenate(
|
||||
[valid_boxes[keep], valid_scores[keep, None], cls_inds], 1
|
||||
)
|
||||
final_dets.append(dets)
|
||||
if len(final_dets) == 0:
|
||||
return None
|
||||
return np.concatenate(final_dets, 0)
|
||||
|
||||
def demo_postprocess(outputs, img_size, p6=False):
|
||||
grids = []
|
||||
expanded_strides = []
|
||||
strides = [8, 16, 32] if not p6 else [8, 16, 32, 64]
|
||||
|
||||
hsizes = [img_size[0] // stride for stride in strides]
|
||||
wsizes = [img_size[1] // stride for stride in strides]
|
||||
|
||||
for hsize, wsize, stride in zip(hsizes, wsizes, strides):
|
||||
xv, yv = np.meshgrid(np.arange(wsize), np.arange(hsize))
|
||||
grid = np.stack((xv, yv), 2).reshape(1, -1, 2)
|
||||
grids.append(grid)
|
||||
shape = grid.shape[:2]
|
||||
expanded_strides.append(np.full((*shape, 1), stride))
|
||||
|
||||
grids = np.concatenate(grids, 1)
|
||||
expanded_strides = np.concatenate(expanded_strides, 1)
|
||||
outputs[..., :2] = (outputs[..., :2] + grids) * expanded_strides
|
||||
outputs[..., 2:4] = np.exp(outputs[..., 2:4]) * expanded_strides
|
||||
|
||||
return outputs
|
||||
|
||||
def preprocess(img, input_size, swap=(2, 0, 1)):
|
||||
if len(img.shape) == 3:
|
||||
padded_img = np.ones((input_size[0], input_size[1], 3), dtype=np.uint8) * 114
|
||||
else:
|
||||
padded_img = np.ones(input_size, dtype=np.uint8) * 114
|
||||
|
||||
r = min(input_size[0] / img.shape[0], input_size[1] / img.shape[1])
|
||||
resized_img = cv2.resize(
|
||||
img,
|
||||
(int(img.shape[1] * r), int(img.shape[0] * r)),
|
||||
interpolation=cv2.INTER_LINEAR,
|
||||
).astype(np.uint8)
|
||||
padded_img[: int(img.shape[0] * r), : int(img.shape[1] * r)] = resized_img
|
||||
|
||||
padded_img = padded_img.transpose(swap)
|
||||
padded_img = np.ascontiguousarray(padded_img, dtype=np.float32)
|
||||
return padded_img, r
|
||||
|
||||
def inference_detector(session, oriImg):
|
||||
input_shape = (640,640)
|
||||
img, ratio = preprocess(oriImg, input_shape)
|
||||
|
||||
ort_inputs = {session.get_inputs()[0].name: img[None, :, :, :]}
|
||||
output = session.run(None, ort_inputs)
|
||||
predictions = demo_postprocess(output[0], input_shape)[0]
|
||||
|
||||
boxes = predictions[:, :4]
|
||||
scores = predictions[:, 4:5] * predictions[:, 5:]
|
||||
|
||||
boxes_xyxy = np.ones_like(boxes)
|
||||
boxes_xyxy[:, 0] = boxes[:, 0] - boxes[:, 2]/2.
|
||||
boxes_xyxy[:, 1] = boxes[:, 1] - boxes[:, 3]/2.
|
||||
boxes_xyxy[:, 2] = boxes[:, 0] + boxes[:, 2]/2.
|
||||
boxes_xyxy[:, 3] = boxes[:, 1] + boxes[:, 3]/2.
|
||||
boxes_xyxy /= ratio
|
||||
dets = multiclass_nms(boxes_xyxy, scores, nms_thr=0.45, score_thr=0.1)
|
||||
if dets is not None:
|
||||
final_boxes, final_scores, final_cls_inds = dets[:, :4], dets[:, 4], dets[:, 5]
|
||||
isscore = final_scores>0.3
|
||||
iscat = final_cls_inds == 0
|
||||
isbbox = [ i and j for (i, j) in zip(isscore, iscat)]
|
||||
final_boxes = final_boxes[isbbox]
|
||||
else:
|
||||
final_boxes = np.array([])
|
||||
|
||||
return final_boxes
|
||||
@@ -0,0 +1,360 @@
|
||||
from typing import List, Tuple
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import onnxruntime as ort
|
||||
|
||||
def preprocess(
|
||||
img: np.ndarray, out_bbox, input_size: Tuple[int, int] = (192, 256)
|
||||
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
|
||||
"""Do preprocessing for RTMPose model inference.
|
||||
|
||||
Args:
|
||||
img (np.ndarray): Input image in shape.
|
||||
input_size (tuple): Input image size in shape (w, h).
|
||||
|
||||
Returns:
|
||||
tuple:
|
||||
- resized_img (np.ndarray): Preprocessed image.
|
||||
- center (np.ndarray): Center of image.
|
||||
- scale (np.ndarray): Scale of image.
|
||||
"""
|
||||
# get shape of image
|
||||
img_shape = img.shape[:2]
|
||||
out_img, out_center, out_scale = [], [], []
|
||||
if len(out_bbox) == 0:
|
||||
out_bbox = [[0, 0, img_shape[1], img_shape[0]]]
|
||||
for i in range(len(out_bbox)):
|
||||
x0 = out_bbox[i][0]
|
||||
y0 = out_bbox[i][1]
|
||||
x1 = out_bbox[i][2]
|
||||
y1 = out_bbox[i][3]
|
||||
bbox = np.array([x0, y0, x1, y1])
|
||||
|
||||
# get center and scale
|
||||
center, scale = bbox_xyxy2cs(bbox, padding=1.25)
|
||||
|
||||
# do affine transformation
|
||||
resized_img, scale = top_down_affine(input_size, scale, center, img)
|
||||
|
||||
# normalize image
|
||||
mean = np.array([123.675, 116.28, 103.53])
|
||||
std = np.array([58.395, 57.12, 57.375])
|
||||
resized_img = (resized_img - mean) / std
|
||||
|
||||
out_img.append(resized_img)
|
||||
out_center.append(center)
|
||||
out_scale.append(scale)
|
||||
|
||||
return out_img, out_center, out_scale
|
||||
|
||||
|
||||
def inference(sess: ort.InferenceSession, img: np.ndarray) -> np.ndarray:
|
||||
"""Inference RTMPose model.
|
||||
|
||||
Args:
|
||||
sess (ort.InferenceSession): ONNXRuntime session.
|
||||
img (np.ndarray): Input image in shape.
|
||||
|
||||
Returns:
|
||||
outputs (np.ndarray): Output of RTMPose model.
|
||||
"""
|
||||
all_out = []
|
||||
# build input
|
||||
for i in range(len(img)):
|
||||
input = [img[i].transpose(2, 0, 1)]
|
||||
|
||||
# build output
|
||||
sess_input = {sess.get_inputs()[0].name: input}
|
||||
sess_output = []
|
||||
for out in sess.get_outputs():
|
||||
sess_output.append(out.name)
|
||||
|
||||
# run model
|
||||
outputs = sess.run(sess_output, sess_input)
|
||||
all_out.append(outputs)
|
||||
|
||||
return all_out
|
||||
|
||||
|
||||
def postprocess(outputs: List[np.ndarray],
|
||||
model_input_size: Tuple[int, int],
|
||||
center: Tuple[int, int],
|
||||
scale: Tuple[int, int],
|
||||
simcc_split_ratio: float = 2.0
|
||||
) -> Tuple[np.ndarray, np.ndarray]:
|
||||
"""Postprocess for RTMPose model output.
|
||||
|
||||
Args:
|
||||
outputs (np.ndarray): Output of RTMPose model.
|
||||
model_input_size (tuple): RTMPose model Input image size.
|
||||
center (tuple): Center of bbox in shape (x, y).
|
||||
scale (tuple): Scale of bbox in shape (w, h).
|
||||
simcc_split_ratio (float): Split ratio of simcc.
|
||||
|
||||
Returns:
|
||||
tuple:
|
||||
- keypoints (np.ndarray): Rescaled keypoints.
|
||||
- scores (np.ndarray): Model predict scores.
|
||||
"""
|
||||
all_key = []
|
||||
all_score = []
|
||||
for i in range(len(outputs)):
|
||||
# use simcc to decode
|
||||
simcc_x, simcc_y = outputs[i]
|
||||
keypoints, scores = decode(simcc_x, simcc_y, simcc_split_ratio)
|
||||
|
||||
# rescale keypoints
|
||||
keypoints = keypoints / model_input_size * scale[i] + center[i] - scale[i] / 2
|
||||
all_key.append(keypoints[0])
|
||||
all_score.append(scores[0])
|
||||
|
||||
return np.array(all_key), np.array(all_score)
|
||||
|
||||
|
||||
def bbox_xyxy2cs(bbox: np.ndarray,
|
||||
padding: float = 1.) -> Tuple[np.ndarray, np.ndarray]:
|
||||
"""Transform the bbox format from (x,y,w,h) into (center, scale)
|
||||
|
||||
Args:
|
||||
bbox (ndarray): Bounding box(es) in shape (4,) or (n, 4), formatted
|
||||
as (left, top, right, bottom)
|
||||
padding (float): BBox padding factor that will be multilied to scale.
|
||||
Default: 1.0
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing center and scale.
|
||||
- np.ndarray[float32]: Center (x, y) of the bbox in shape (2,) or
|
||||
(n, 2)
|
||||
- np.ndarray[float32]: Scale (w, h) of the bbox in shape (2,) or
|
||||
(n, 2)
|
||||
"""
|
||||
# convert single bbox from (4, ) to (1, 4)
|
||||
dim = bbox.ndim
|
||||
if dim == 1:
|
||||
bbox = bbox[None, :]
|
||||
|
||||
# get bbox center and scale
|
||||
x1, y1, x2, y2 = np.hsplit(bbox, [1, 2, 3])
|
||||
center = np.hstack([x1 + x2, y1 + y2]) * 0.5
|
||||
scale = np.hstack([x2 - x1, y2 - y1]) * padding
|
||||
|
||||
if dim == 1:
|
||||
center = center[0]
|
||||
scale = scale[0]
|
||||
|
||||
return center, scale
|
||||
|
||||
|
||||
def _fix_aspect_ratio(bbox_scale: np.ndarray,
|
||||
aspect_ratio: float) -> np.ndarray:
|
||||
"""Extend the scale to match the given aspect ratio.
|
||||
|
||||
Args:
|
||||
scale (np.ndarray): The image scale (w, h) in shape (2, )
|
||||
aspect_ratio (float): The ratio of ``w/h``
|
||||
|
||||
Returns:
|
||||
np.ndarray: The reshaped image scale in (2, )
|
||||
"""
|
||||
w, h = np.hsplit(bbox_scale, [1])
|
||||
bbox_scale = np.where(w > h * aspect_ratio,
|
||||
np.hstack([w, w / aspect_ratio]),
|
||||
np.hstack([h * aspect_ratio, h]))
|
||||
return bbox_scale
|
||||
|
||||
|
||||
def _rotate_point(pt: np.ndarray, angle_rad: float) -> np.ndarray:
|
||||
"""Rotate a point by an angle.
|
||||
|
||||
Args:
|
||||
pt (np.ndarray): 2D point coordinates (x, y) in shape (2, )
|
||||
angle_rad (float): rotation angle in radian
|
||||
|
||||
Returns:
|
||||
np.ndarray: Rotated point in shape (2, )
|
||||
"""
|
||||
sn, cs = np.sin(angle_rad), np.cos(angle_rad)
|
||||
rot_mat = np.array([[cs, -sn], [sn, cs]])
|
||||
return rot_mat @ pt
|
||||
|
||||
|
||||
def _get_3rd_point(a: np.ndarray, b: np.ndarray) -> np.ndarray:
|
||||
"""To calculate the affine matrix, three pairs of points are required. This
|
||||
function is used to get the 3rd point, given 2D points a & b.
|
||||
|
||||
The 3rd point is defined by rotating vector `a - b` by 90 degrees
|
||||
anticlockwise, using b as the rotation center.
|
||||
|
||||
Args:
|
||||
a (np.ndarray): The 1st point (x,y) in shape (2, )
|
||||
b (np.ndarray): The 2nd point (x,y) in shape (2, )
|
||||
|
||||
Returns:
|
||||
np.ndarray: The 3rd point.
|
||||
"""
|
||||
direction = a - b
|
||||
c = b + np.r_[-direction[1], direction[0]]
|
||||
return c
|
||||
|
||||
|
||||
def get_warp_matrix(center: np.ndarray,
|
||||
scale: np.ndarray,
|
||||
rot: float,
|
||||
output_size: Tuple[int, int],
|
||||
shift: Tuple[float, float] = (0., 0.),
|
||||
inv: bool = False) -> np.ndarray:
|
||||
"""Calculate the affine transformation matrix that can warp the bbox area
|
||||
in the input image to the output size.
|
||||
|
||||
Args:
|
||||
center (np.ndarray[2, ]): Center of the bounding box (x, y).
|
||||
scale (np.ndarray[2, ]): Scale of the bounding box
|
||||
wrt [width, height].
|
||||
rot (float): Rotation angle (degree).
|
||||
output_size (np.ndarray[2, ] | list(2,)): Size of the
|
||||
destination heatmaps.
|
||||
shift (0-100%): Shift translation ratio wrt the width/height.
|
||||
Default (0., 0.).
|
||||
inv (bool): Option to inverse the affine transform direction.
|
||||
(inv=False: src->dst or inv=True: dst->src)
|
||||
|
||||
Returns:
|
||||
np.ndarray: A 2x3 transformation matrix
|
||||
"""
|
||||
shift = np.array(shift)
|
||||
src_w = scale[0]
|
||||
dst_w = output_size[0]
|
||||
dst_h = output_size[1]
|
||||
|
||||
# compute transformation matrix
|
||||
rot_rad = np.deg2rad(rot)
|
||||
src_dir = _rotate_point(np.array([0., src_w * -0.5]), rot_rad)
|
||||
dst_dir = np.array([0., dst_w * -0.5])
|
||||
|
||||
# get four corners of the src rectangle in the original image
|
||||
src = np.zeros((3, 2), dtype=np.float32)
|
||||
src[0, :] = center + scale * shift
|
||||
src[1, :] = center + src_dir + scale * shift
|
||||
src[2, :] = _get_3rd_point(src[0, :], src[1, :])
|
||||
|
||||
# get four corners of the dst rectangle in the input image
|
||||
dst = np.zeros((3, 2), dtype=np.float32)
|
||||
dst[0, :] = [dst_w * 0.5, dst_h * 0.5]
|
||||
dst[1, :] = np.array([dst_w * 0.5, dst_h * 0.5]) + dst_dir
|
||||
dst[2, :] = _get_3rd_point(dst[0, :], dst[1, :])
|
||||
|
||||
if inv:
|
||||
warp_mat = cv2.getAffineTransform(np.float32(dst), np.float32(src))
|
||||
else:
|
||||
warp_mat = cv2.getAffineTransform(np.float32(src), np.float32(dst))
|
||||
|
||||
return warp_mat
|
||||
|
||||
|
||||
def top_down_affine(input_size: dict, bbox_scale: dict, bbox_center: dict,
|
||||
img: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
|
||||
"""Get the bbox image as the model input by affine transform.
|
||||
|
||||
Args:
|
||||
input_size (dict): The input size of the model.
|
||||
bbox_scale (dict): The bbox scale of the img.
|
||||
bbox_center (dict): The bbox center of the img.
|
||||
img (np.ndarray): The original image.
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing center and scale.
|
||||
- np.ndarray[float32]: img after affine transform.
|
||||
- np.ndarray[float32]: bbox scale after affine transform.
|
||||
"""
|
||||
w, h = input_size
|
||||
warp_size = (int(w), int(h))
|
||||
|
||||
# reshape bbox to fixed aspect ratio
|
||||
bbox_scale = _fix_aspect_ratio(bbox_scale, aspect_ratio=w / h)
|
||||
|
||||
# get the affine matrix
|
||||
center = bbox_center
|
||||
scale = bbox_scale
|
||||
rot = 0
|
||||
warp_mat = get_warp_matrix(center, scale, rot, output_size=(w, h))
|
||||
|
||||
# do affine transform
|
||||
img = cv2.warpAffine(img, warp_mat, warp_size, flags=cv2.INTER_LINEAR)
|
||||
|
||||
return img, bbox_scale
|
||||
|
||||
|
||||
def get_simcc_maximum(simcc_x: np.ndarray,
|
||||
simcc_y: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
|
||||
"""Get maximum response location and value from simcc representations.
|
||||
|
||||
Note:
|
||||
instance number: N
|
||||
num_keypoints: K
|
||||
heatmap height: H
|
||||
heatmap width: W
|
||||
|
||||
Args:
|
||||
simcc_x (np.ndarray): x-axis SimCC in shape (K, Wx) or (N, K, Wx)
|
||||
simcc_y (np.ndarray): y-axis SimCC in shape (K, Wy) or (N, K, Wy)
|
||||
|
||||
Returns:
|
||||
tuple:
|
||||
- locs (np.ndarray): locations of maximum heatmap responses in shape
|
||||
(K, 2) or (N, K, 2)
|
||||
- vals (np.ndarray): values of maximum heatmap responses in shape
|
||||
(K,) or (N, K)
|
||||
"""
|
||||
N, K, Wx = simcc_x.shape
|
||||
simcc_x = simcc_x.reshape(N * K, -1)
|
||||
simcc_y = simcc_y.reshape(N * K, -1)
|
||||
|
||||
# get maximum value locations
|
||||
x_locs = np.argmax(simcc_x, axis=1)
|
||||
y_locs = np.argmax(simcc_y, axis=1)
|
||||
locs = np.stack((x_locs, y_locs), axis=-1).astype(np.float32)
|
||||
max_val_x = np.amax(simcc_x, axis=1)
|
||||
max_val_y = np.amax(simcc_y, axis=1)
|
||||
|
||||
# get maximum value across x and y axis
|
||||
mask = max_val_x > max_val_y
|
||||
max_val_x[mask] = max_val_y[mask]
|
||||
vals = max_val_x
|
||||
locs[vals <= 0.] = -1
|
||||
|
||||
# reshape
|
||||
locs = locs.reshape(N, K, 2)
|
||||
vals = vals.reshape(N, K)
|
||||
|
||||
return locs, vals
|
||||
|
||||
|
||||
def decode(simcc_x: np.ndarray, simcc_y: np.ndarray,
|
||||
simcc_split_ratio) -> Tuple[np.ndarray, np.ndarray]:
|
||||
"""Modulate simcc distribution with Gaussian.
|
||||
|
||||
Args:
|
||||
simcc_x (np.ndarray[K, Wx]): model predicted simcc in x.
|
||||
simcc_y (np.ndarray[K, Wy]): model predicted simcc in y.
|
||||
simcc_split_ratio (int): The split ratio of simcc.
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing center and scale.
|
||||
- np.ndarray[float32]: keypoints in shape (K, 2) or (n, K, 2)
|
||||
- np.ndarray[float32]: scores in shape (K,) or (n, K)
|
||||
"""
|
||||
keypoints, scores = get_simcc_maximum(simcc_x, simcc_y)
|
||||
keypoints /= simcc_split_ratio
|
||||
|
||||
return keypoints, scores
|
||||
|
||||
|
||||
def inference_pose(session, out_bbox, oriImg):
|
||||
h, w = session.get_inputs()[0].shape[2:]
|
||||
model_input_size = (w, h)
|
||||
resized_img, center, scale = preprocess(oriImg, out_bbox, model_input_size)
|
||||
outputs = inference(session, resized_img)
|
||||
keypoints, scores = postprocess(outputs, model_input_size, center, scale)
|
||||
|
||||
return keypoints, scores
|
||||
@@ -0,0 +1,297 @@
|
||||
import math
|
||||
import numpy as np
|
||||
import matplotlib
|
||||
import cv2
|
||||
|
||||
|
||||
eps = 0.01
|
||||
|
||||
|
||||
def smart_resize(x, s):
|
||||
Ht, Wt = s
|
||||
if x.ndim == 2:
|
||||
Ho, Wo = x.shape
|
||||
Co = 1
|
||||
else:
|
||||
Ho, Wo, Co = x.shape
|
||||
if Co == 3 or Co == 1:
|
||||
k = float(Ht + Wt) / float(Ho + Wo)
|
||||
return cv2.resize(x, (int(Wt), int(Ht)), interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4)
|
||||
else:
|
||||
return np.stack([smart_resize(x[:, :, i], s) for i in range(Co)], axis=2)
|
||||
|
||||
|
||||
def smart_resize_k(x, fx, fy):
|
||||
if x.ndim == 2:
|
||||
Ho, Wo = x.shape
|
||||
Co = 1
|
||||
else:
|
||||
Ho, Wo, Co = x.shape
|
||||
Ht, Wt = Ho * fy, Wo * fx
|
||||
if Co == 3 or Co == 1:
|
||||
k = float(Ht + Wt) / float(Ho + Wo)
|
||||
return cv2.resize(x, (int(Wt), int(Ht)), interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4)
|
||||
else:
|
||||
return np.stack([smart_resize_k(x[:, :, i], fx, fy) for i in range(Co)], axis=2)
|
||||
|
||||
|
||||
def padRightDownCorner(img, stride, padValue):
|
||||
h = img.shape[0]
|
||||
w = img.shape[1]
|
||||
|
||||
pad = 4 * [None]
|
||||
pad[0] = 0 # up
|
||||
pad[1] = 0 # left
|
||||
pad[2] = 0 if (h % stride == 0) else stride - (h % stride) # down
|
||||
pad[3] = 0 if (w % stride == 0) else stride - (w % stride) # right
|
||||
|
||||
img_padded = img
|
||||
pad_up = np.tile(img_padded[0:1, :, :]*0 + padValue, (pad[0], 1, 1))
|
||||
img_padded = np.concatenate((pad_up, img_padded), axis=0)
|
||||
pad_left = np.tile(img_padded[:, 0:1, :]*0 + padValue, (1, pad[1], 1))
|
||||
img_padded = np.concatenate((pad_left, img_padded), axis=1)
|
||||
pad_down = np.tile(img_padded[-2:-1, :, :]*0 + padValue, (pad[2], 1, 1))
|
||||
img_padded = np.concatenate((img_padded, pad_down), axis=0)
|
||||
pad_right = np.tile(img_padded[:, -2:-1, :]*0 + padValue, (1, pad[3], 1))
|
||||
img_padded = np.concatenate((img_padded, pad_right), axis=1)
|
||||
|
||||
return img_padded, pad
|
||||
|
||||
|
||||
def transfer(model, model_weights):
|
||||
transfered_model_weights = {}
|
||||
for weights_name in model.state_dict().keys():
|
||||
transfered_model_weights[weights_name] = model_weights['.'.join(weights_name.split('.')[1:])]
|
||||
return transfered_model_weights
|
||||
|
||||
|
||||
def draw_bodypose(canvas, candidate, subset):
|
||||
H, W, C = canvas.shape
|
||||
candidate = np.array(candidate)
|
||||
subset = np.array(subset)
|
||||
|
||||
stickwidth = 4
|
||||
|
||||
limbSeq = [[2, 3], [2, 6], [3, 4], [4, 5], [6, 7], [7, 8], [2, 9], [9, 10], \
|
||||
[10, 11], [2, 12], [12, 13], [13, 14], [2, 1], [1, 15], [15, 17], \
|
||||
[1, 16], [16, 18], [3, 17], [6, 18]]
|
||||
|
||||
colors = [[255, 0, 0], [255, 85, 0], [255, 170, 0], [255, 255, 0], [170, 255, 0], [85, 255, 0], [0, 255, 0], \
|
||||
[0, 255, 85], [0, 255, 170], [0, 255, 255], [0, 170, 255], [0, 85, 255], [0, 0, 255], [85, 0, 255], \
|
||||
[170, 0, 255], [255, 0, 255], [255, 0, 170], [255, 0, 85]]
|
||||
|
||||
for i in range(17):
|
||||
for n in range(len(subset)):
|
||||
index = subset[n][np.array(limbSeq[i]) - 1]
|
||||
if -1 in index:
|
||||
continue
|
||||
Y = candidate[index.astype(int), 0] * float(W)
|
||||
X = candidate[index.astype(int), 1] * float(H)
|
||||
mX = np.mean(X)
|
||||
mY = np.mean(Y)
|
||||
length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5
|
||||
angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1]))
|
||||
polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stickwidth), int(angle), 0, 360, 1)
|
||||
cv2.fillConvexPoly(canvas, polygon, colors[i])
|
||||
|
||||
canvas = (canvas * 0.6).astype(np.uint8)
|
||||
|
||||
for i in range(18):
|
||||
for n in range(len(subset)):
|
||||
index = int(subset[n][i])
|
||||
if index == -1:
|
||||
continue
|
||||
x, y = candidate[index][0:2]
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
cv2.circle(canvas, (int(x), int(y)), 4, colors[i], thickness=-1)
|
||||
|
||||
return canvas
|
||||
|
||||
|
||||
def draw_handpose(canvas, all_hand_peaks):
|
||||
H, W, C = canvas.shape
|
||||
|
||||
edges = [[0, 1], [1, 2], [2, 3], [3, 4], [0, 5], [5, 6], [6, 7], [7, 8], [0, 9], [9, 10], \
|
||||
[10, 11], [11, 12], [0, 13], [13, 14], [14, 15], [15, 16], [0, 17], [17, 18], [18, 19], [19, 20]]
|
||||
|
||||
for peaks in all_hand_peaks:
|
||||
peaks = np.array(peaks)
|
||||
|
||||
for ie, e in enumerate(edges):
|
||||
x1, y1 = peaks[e[0]]
|
||||
x2, y2 = peaks[e[1]]
|
||||
x1 = int(x1 * W)
|
||||
y1 = int(y1 * H)
|
||||
x2 = int(x2 * W)
|
||||
y2 = int(y2 * H)
|
||||
if x1 > eps and y1 > eps and x2 > eps and y2 > eps:
|
||||
cv2.line(canvas, (x1, y1), (x2, y2), matplotlib.colors.hsv_to_rgb([ie / float(len(edges)), 1.0, 1.0]) * 255, thickness=2)
|
||||
|
||||
for i, keyponit in enumerate(peaks):
|
||||
x, y = keyponit
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
if x > eps and y > eps:
|
||||
cv2.circle(canvas, (x, y), 4, (0, 0, 255), thickness=-1)
|
||||
return canvas
|
||||
|
||||
|
||||
def draw_facepose(canvas, all_lmks):
|
||||
H, W, C = canvas.shape
|
||||
for lmks in all_lmks:
|
||||
lmks = np.array(lmks)
|
||||
for lmk in lmks:
|
||||
x, y = lmk
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
if x > eps and y > eps:
|
||||
cv2.circle(canvas, (x, y), 3, (255, 255, 255), thickness=-1)
|
||||
return canvas
|
||||
|
||||
|
||||
# detect hand according to body pose keypoints
|
||||
# please refer to https://github.com/CMU-Perceptual-Computing-Lab/openpose/blob/master/src/openpose/hand/handDetector.cpp
|
||||
def handDetect(candidate, subset, oriImg):
|
||||
# right hand: wrist 4, elbow 3, shoulder 2
|
||||
# left hand: wrist 7, elbow 6, shoulder 5
|
||||
ratioWristElbow = 0.33
|
||||
detect_result = []
|
||||
image_height, image_width = oriImg.shape[0:2]
|
||||
for person in subset.astype(int):
|
||||
# if any of three not detected
|
||||
has_left = np.sum(person[[5, 6, 7]] == -1) == 0
|
||||
has_right = np.sum(person[[2, 3, 4]] == -1) == 0
|
||||
if not (has_left or has_right):
|
||||
continue
|
||||
hands = []
|
||||
#left hand
|
||||
if has_left:
|
||||
left_shoulder_index, left_elbow_index, left_wrist_index = person[[5, 6, 7]]
|
||||
x1, y1 = candidate[left_shoulder_index][:2]
|
||||
x2, y2 = candidate[left_elbow_index][:2]
|
||||
x3, y3 = candidate[left_wrist_index][:2]
|
||||
hands.append([x1, y1, x2, y2, x3, y3, True])
|
||||
# right hand
|
||||
if has_right:
|
||||
right_shoulder_index, right_elbow_index, right_wrist_index = person[[2, 3, 4]]
|
||||
x1, y1 = candidate[right_shoulder_index][:2]
|
||||
x2, y2 = candidate[right_elbow_index][:2]
|
||||
x3, y3 = candidate[right_wrist_index][:2]
|
||||
hands.append([x1, y1, x2, y2, x3, y3, False])
|
||||
|
||||
for x1, y1, x2, y2, x3, y3, is_left in hands:
|
||||
# pos_hand = pos_wrist + ratio * (pos_wrist - pos_elbox) = (1 + ratio) * pos_wrist - ratio * pos_elbox
|
||||
# handRectangle.x = posePtr[wrist*3] + ratioWristElbow * (posePtr[wrist*3] - posePtr[elbow*3]);
|
||||
# handRectangle.y = posePtr[wrist*3+1] + ratioWristElbow * (posePtr[wrist*3+1] - posePtr[elbow*3+1]);
|
||||
# const auto distanceWristElbow = getDistance(poseKeypoints, person, wrist, elbow);
|
||||
# const auto distanceElbowShoulder = getDistance(poseKeypoints, person, elbow, shoulder);
|
||||
# handRectangle.width = 1.5f * fastMax(distanceWristElbow, 0.9f * distanceElbowShoulder);
|
||||
x = x3 + ratioWristElbow * (x3 - x2)
|
||||
y = y3 + ratioWristElbow * (y3 - y2)
|
||||
distanceWristElbow = math.sqrt((x3 - x2) ** 2 + (y3 - y2) ** 2)
|
||||
distanceElbowShoulder = math.sqrt((x2 - x1) ** 2 + (y2 - y1) ** 2)
|
||||
width = 1.5 * max(distanceWristElbow, 0.9 * distanceElbowShoulder)
|
||||
# x-y refers to the center --> offset to topLeft point
|
||||
# handRectangle.x -= handRectangle.width / 2.f;
|
||||
# handRectangle.y -= handRectangle.height / 2.f;
|
||||
x -= width / 2
|
||||
y -= width / 2 # width = height
|
||||
# overflow the image
|
||||
if x < 0: x = 0
|
||||
if y < 0: y = 0
|
||||
width1 = width
|
||||
width2 = width
|
||||
if x + width > image_width: width1 = image_width - x
|
||||
if y + width > image_height: width2 = image_height - y
|
||||
width = min(width1, width2)
|
||||
# the max hand box value is 20 pixels
|
||||
if width >= 20:
|
||||
detect_result.append([int(x), int(y), int(width), is_left])
|
||||
|
||||
'''
|
||||
return value: [[x, y, w, True if left hand else False]].
|
||||
width=height since the network require squared input.
|
||||
x, y is the coordinate of top left
|
||||
'''
|
||||
return detect_result
|
||||
|
||||
|
||||
# Written by Lvmin
|
||||
def faceDetect(candidate, subset, oriImg):
|
||||
# left right eye ear 14 15 16 17
|
||||
detect_result = []
|
||||
image_height, image_width = oriImg.shape[0:2]
|
||||
for person in subset.astype(int):
|
||||
has_head = person[0] > -1
|
||||
if not has_head:
|
||||
continue
|
||||
|
||||
has_left_eye = person[14] > -1
|
||||
has_right_eye = person[15] > -1
|
||||
has_left_ear = person[16] > -1
|
||||
has_right_ear = person[17] > -1
|
||||
|
||||
if not (has_left_eye or has_right_eye or has_left_ear or has_right_ear):
|
||||
continue
|
||||
|
||||
head, left_eye, right_eye, left_ear, right_ear = person[[0, 14, 15, 16, 17]]
|
||||
|
||||
width = 0.0
|
||||
x0, y0 = candidate[head][:2]
|
||||
|
||||
if has_left_eye:
|
||||
x1, y1 = candidate[left_eye][:2]
|
||||
d = max(abs(x0 - x1), abs(y0 - y1))
|
||||
width = max(width, d * 3.0)
|
||||
|
||||
if has_right_eye:
|
||||
x1, y1 = candidate[right_eye][:2]
|
||||
d = max(abs(x0 - x1), abs(y0 - y1))
|
||||
width = max(width, d * 3.0)
|
||||
|
||||
if has_left_ear:
|
||||
x1, y1 = candidate[left_ear][:2]
|
||||
d = max(abs(x0 - x1), abs(y0 - y1))
|
||||
width = max(width, d * 1.5)
|
||||
|
||||
if has_right_ear:
|
||||
x1, y1 = candidate[right_ear][:2]
|
||||
d = max(abs(x0 - x1), abs(y0 - y1))
|
||||
width = max(width, d * 1.5)
|
||||
|
||||
x, y = x0, y0
|
||||
|
||||
x -= width
|
||||
y -= width
|
||||
|
||||
if x < 0:
|
||||
x = 0
|
||||
|
||||
if y < 0:
|
||||
y = 0
|
||||
|
||||
width1 = width * 2
|
||||
width2 = width * 2
|
||||
|
||||
if x + width > image_width:
|
||||
width1 = image_width - x
|
||||
|
||||
if y + width > image_height:
|
||||
width2 = image_height - y
|
||||
|
||||
width = min(width1, width2)
|
||||
|
||||
if width >= 20:
|
||||
detect_result.append([int(x), int(y), int(width)])
|
||||
|
||||
return detect_result
|
||||
|
||||
|
||||
# get max index of 2d array
|
||||
def npmax(array):
|
||||
arrayindex = array.argmax(1)
|
||||
arrayvalue = array.max(1)
|
||||
i = arrayvalue.argmax()
|
||||
j = arrayindex[i]
|
||||
return i, j
|
||||
@@ -0,0 +1,50 @@
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
import onnxruntime as ort
|
||||
from .onnxdet import inference_detector
|
||||
from .onnxpose import inference_pose
|
||||
|
||||
class Wholebody:
|
||||
def __init__(self):
|
||||
device = 'cuda:0'
|
||||
# device = 'cuda'
|
||||
providers = ['CPUExecutionProvider'
|
||||
] if device == 'cpu' else ['CUDAExecutionProvider']
|
||||
onnx_det = 'checkpoints/DWPose/yolox_l.onnx'
|
||||
onnx_pose = 'checkpoints/DWPose/dw-ll_ucoco_384.onnx'
|
||||
|
||||
self.session_det = ort.InferenceSession(path_or_bytes=onnx_det, providers=providers)
|
||||
self.session_pose = ort.InferenceSession(path_or_bytes=onnx_pose, providers=providers)
|
||||
|
||||
def __call__(self, oriImg):
|
||||
det_result = inference_detector(self.session_det, oriImg)
|
||||
keypoints, scores = inference_pose(self.session_pose, det_result, oriImg)
|
||||
|
||||
keypoints_info = np.concatenate(
|
||||
(keypoints, scores[..., None]), axis=-1)
|
||||
# compute neck joint
|
||||
neck = np.mean(keypoints_info[:, [5, 6]], axis=1)
|
||||
# neck score when visualizing pred
|
||||
neck[:, 2:4] = np.logical_and(
|
||||
keypoints_info[:, 5, 2:4] > 0.3,
|
||||
keypoints_info[:, 6, 2:4] > 0.3).astype(int)
|
||||
new_keypoints_info = np.insert(
|
||||
keypoints_info, 17, neck, axis=1)
|
||||
mmpose_idx = [
|
||||
17, 6, 8, 10, 7, 9, 12, 14, 16, 13, 15, 2, 1, 4, 3
|
||||
]
|
||||
openpose_idx = [
|
||||
1, 2, 3, 4, 6, 7, 8, 9, 10, 12, 13, 14, 15, 16, 17
|
||||
]
|
||||
new_keypoints_info[:, openpose_idx] = \
|
||||
new_keypoints_info[:, mmpose_idx]
|
||||
keypoints_info = new_keypoints_info
|
||||
|
||||
keypoints, scores = keypoints_info[
|
||||
..., :2], keypoints_info[..., 2]
|
||||
|
||||
return keypoints, scores
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,205 @@
|
||||
import math
|
||||
import matplotlib
|
||||
import cv2
|
||||
import os
|
||||
import numpy as np
|
||||
from dwpose_utils.dwpose_detector import dwpose_detector_aligned
|
||||
import argparse
|
||||
|
||||
eps = 0.01
|
||||
|
||||
def alpha_blend_color(color, alpha):
|
||||
"""blend color according to point conf
|
||||
"""
|
||||
return [int(c * alpha) for c in color]
|
||||
|
||||
def draw_bodypose(canvas, candidate, subset, score):
|
||||
H, W, C = canvas.shape
|
||||
candidate = np.array(candidate)
|
||||
subset = np.array(subset)
|
||||
|
||||
stickwidth = 4
|
||||
|
||||
limbSeq = [[2, 3], [2, 6], [3, 4], [4, 5], [6, 7], [7, 8], [2, 9], [9, 10], \
|
||||
[10, 11], [2, 12], [12, 13], [13, 14], [2, 1], [1, 15], [15, 17], \
|
||||
[1, 16], [16, 18], [3, 17], [6, 18]]
|
||||
|
||||
colors = [[255, 0, 0], [255, 85, 0], [255, 170, 0], [255, 255, 0], [170, 255, 0], [85, 255, 0], [0, 255, 0], \
|
||||
[0, 255, 85], [0, 255, 170], [0, 255, 255], [0, 170, 255], [0, 85, 255], [0, 0, 255], [85, 0, 255], \
|
||||
[170, 0, 255], [255, 0, 255], [255, 0, 170], [255, 0, 85]]
|
||||
|
||||
for i in range(17):
|
||||
for n in range(len(subset)):
|
||||
index = subset[n][np.array(limbSeq[i]) - 1]
|
||||
conf = score[n][np.array(limbSeq[i]) - 1]
|
||||
if conf[0] < 0.3 or conf[1] < 0.3:
|
||||
continue
|
||||
Y = candidate[index.astype(int), 0] * float(W)
|
||||
X = candidate[index.astype(int), 1] * float(H)
|
||||
mX = np.mean(X)
|
||||
mY = np.mean(Y)
|
||||
length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5
|
||||
angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1]))
|
||||
polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stickwidth), int(angle), 0, 360, 1)
|
||||
cv2.fillConvexPoly(canvas, polygon, alpha_blend_color(colors[i], conf[0] * conf[1]))
|
||||
|
||||
canvas = (canvas * 0.6).astype(np.uint8)
|
||||
|
||||
for i in range(18):
|
||||
for n in range(len(subset)):
|
||||
index = int(subset[n][i])
|
||||
if index == -1:
|
||||
continue
|
||||
x, y = candidate[index][0:2]
|
||||
conf = score[n][i]
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
cv2.circle(canvas, (int(x), int(y)), 4, alpha_blend_color(colors[i], conf), thickness=-1)
|
||||
|
||||
return canvas
|
||||
|
||||
def draw_handpose(canvas, all_hand_peaks, all_hand_scores):
|
||||
H, W, C = canvas.shape
|
||||
|
||||
edges = [[0, 1], [1, 2], [2, 3], [3, 4], [0, 5], [5, 6], [6, 7], [7, 8], [0, 9], [9, 10], \
|
||||
[10, 11], [11, 12], [0, 13], [13, 14], [14, 15], [15, 16], [0, 17], [17, 18], [18, 19], [19, 20]]
|
||||
|
||||
for peaks, scores in zip(all_hand_peaks, all_hand_scores):
|
||||
|
||||
for ie, e in enumerate(edges):
|
||||
x1, y1 = peaks[e[0]]
|
||||
x2, y2 = peaks[e[1]]
|
||||
x1 = int(x1 * W)
|
||||
y1 = int(y1 * H)
|
||||
x2 = int(x2 * W)
|
||||
y2 = int(y2 * H)
|
||||
score = int(scores[e[0]] * scores[e[1]] * 255)
|
||||
if x1 > eps and y1 > eps and x2 > eps and y2 > eps:
|
||||
cv2.line(canvas, (x1, y1), (x2, y2),
|
||||
matplotlib.colors.hsv_to_rgb([ie / float(len(edges)), 1.0, 1.0]) * score, thickness=2)
|
||||
|
||||
for i, keyponit in enumerate(peaks):
|
||||
x, y = keyponit
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
score = int(scores[i] * 255)
|
||||
if x > eps and y > eps:
|
||||
cv2.circle(canvas, (x, y), 4, (0, 0, score), thickness=-1)
|
||||
return canvas
|
||||
|
||||
def draw_facepose(canvas, all_lmks, all_scores):
|
||||
H, W, C = canvas.shape
|
||||
for lmks, scores in zip(all_lmks, all_scores):
|
||||
for lmk, score in zip(lmks, scores):
|
||||
x, y = lmk
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
conf = int(score * 255)
|
||||
if x > eps and y > eps:
|
||||
cv2.circle(canvas, (x, y), 3, (conf, conf, conf), thickness=-1)
|
||||
return canvas
|
||||
|
||||
def draw_pose(pose, H, W, ref_w=2160):
|
||||
"""vis dwpose outputs
|
||||
|
||||
Args:
|
||||
pose (List): DWposeDetector outputs in dwpose_detector.py
|
||||
H (int): height
|
||||
W (int): width
|
||||
ref_w (int, optional) Defaults to 2160.
|
||||
|
||||
Returns:
|
||||
np.ndarray: image pixel value in RGB mode
|
||||
"""
|
||||
bodies = pose['bodies']
|
||||
faces = pose['faces']
|
||||
hands = pose['hands']
|
||||
candidate = bodies['candidate']
|
||||
subset = bodies['subset']
|
||||
|
||||
sz = min(H, W)
|
||||
sr = (ref_w / sz) if sz != ref_w else 1
|
||||
|
||||
########################################## create zero canvas ##################################################
|
||||
canvas = np.zeros(shape=(int(H*sr), int(W*sr), 3), dtype=np.uint8)
|
||||
|
||||
########################################### draw body pose #####################################################
|
||||
canvas = draw_bodypose(canvas, candidate, subset, score=bodies['score'])
|
||||
|
||||
########################################### draw hand pose #####################################################
|
||||
canvas = draw_handpose(canvas, hands, pose['hands_score'])
|
||||
|
||||
########################################### draw face pose #####################################################
|
||||
canvas = draw_facepose(canvas, faces, pose['faces_score'])
|
||||
|
||||
return cv2.cvtColor(cv2.resize(canvas, (W, H)), cv2.COLOR_BGR2RGB).transpose(2, 0, 1)
|
||||
|
||||
def get_video_pose(video_path, ref_image_path, poses_folder_path=None):
|
||||
|
||||
ref_image = cv2.imread(ref_image_path)
|
||||
ref_image = cv2.cvtColor(ref_image, cv2.COLOR_BGR2RGB)
|
||||
height, width, _ = ref_image.shape
|
||||
ref_pose = dwpose_detector_aligned(ref_image)
|
||||
ref_keypoint_id = [0, 1, 2, 5, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17]
|
||||
ref_keypoint_id = [i for i in ref_keypoint_id \
|
||||
if len(ref_pose['bodies']['subset']) > 0 and ref_pose['bodies']['subset'][0][i] >= .0]
|
||||
ref_body = ref_pose['bodies']['candidate'][ref_keypoint_id]
|
||||
|
||||
os.makedirs(poses_folder_path, exist_ok=True)
|
||||
detected_poses = []
|
||||
files = os.listdir(video_path)
|
||||
png_files = [f for f in files if f.endswith('.png')]
|
||||
png_files.sort(key=lambda x: int(x.split('_')[1].split('.')[0]))
|
||||
for sub_name in png_files:
|
||||
sub_driven_image_path = os.path.join(video_path, sub_name)
|
||||
driven_image = cv2.imread(sub_driven_image_path)
|
||||
driven_image = cv2.cvtColor(driven_image, cv2.COLOR_BGR2RGB)
|
||||
driven_pose = dwpose_detector_aligned(driven_image)
|
||||
detected_poses.append(driven_pose)
|
||||
|
||||
detected_bodies = np.stack(
|
||||
[p['bodies']['candidate'] for p in detected_poses if p['bodies']['candidate'].shape[0] == 18])[:,
|
||||
ref_keypoint_id]
|
||||
ay, by = np.polyfit(detected_bodies[:, :, 1].flatten(), np.tile(ref_body[:, 1], len(detected_bodies)), 1)
|
||||
fh = height
|
||||
fw = width
|
||||
ax = ay / (fh / fw / height * width)
|
||||
bx = np.mean(np.tile(ref_body[:, 0], len(detected_bodies)) - detected_bodies[:, :, 0].flatten() * ax)
|
||||
a = np.array([ax, ay])
|
||||
b = np.array([bx, by])
|
||||
output_pose = []
|
||||
# pose rescale
|
||||
for detected_pose in detected_poses:
|
||||
detected_pose['bodies']['candidate'] = detected_pose['bodies']['candidate'] * a + b
|
||||
detected_pose['faces'] = detected_pose['faces'] * a + b
|
||||
detected_pose['hands'] = detected_pose['hands'] * a + b
|
||||
im = draw_pose(detected_pose, height, width)
|
||||
output_pose.append(np.array(im))
|
||||
return np.stack(output_pose)
|
||||
|
||||
|
||||
def get_image_pose(ref_image_path):
|
||||
ref_image = cv2.imread(ref_image_path)
|
||||
ref_image = cv2.cvtColor(ref_image, cv2.COLOR_BGR2RGB)
|
||||
height, width, _ = ref_image.shape
|
||||
ref_pose = dwpose_detector_aligned(ref_image)
|
||||
pose_img = draw_pose(ref_pose, height, width)
|
||||
return np.array(pose_img)
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser(description="Skeleton extraction from images.")
|
||||
parser.add_argument('--target_image_folder_path', type=str, required=True, help='Path to the folder containing target images.')
|
||||
parser.add_argument('--ref_image_path', type=str, required=True, help='Path to the reference image.')
|
||||
parser.add_argument('--poses_folder_path', type=str, required=True, help='Path to save the extracted poses.')
|
||||
args = parser.parse_args()
|
||||
|
||||
video_path = args.target_image_folder_path
|
||||
ref_image_path = args.ref_image_path
|
||||
poses_folder_path = args.poses_folder_path
|
||||
detected_maps = get_video_pose(video_path, ref_image_path, poses_folder_path=poses_folder_path)
|
||||
for i in range(detected_maps.shape[0]):
|
||||
pose_image = np.transpose(detected_maps[i], (1, 2, 0))
|
||||
# pose_image = detected_maps[i]
|
||||
pose_save_path = os.path.join(poses_folder_path, f"frame_{i}.png")
|
||||
cv2.imwrite(pose_save_path, pose_image)
|
||||
print(f"save the pose image in {pose_save_path}")
|
||||
@@ -0,0 +1,167 @@
|
||||
import math
|
||||
import argparse
|
||||
import cv2
|
||||
import os
|
||||
from tqdm import tqdm
|
||||
import decord
|
||||
import numpy as np
|
||||
import matplotlib
|
||||
|
||||
from dwpose_utils.dwpose_detector import dwpose_detector_aligned
|
||||
|
||||
eps = 0.01
|
||||
|
||||
def alpha_blend_color(color, alpha):
|
||||
"""blend color according to point conf
|
||||
"""
|
||||
return [int(c * alpha) for c in color]
|
||||
|
||||
def draw_bodypose_aligned(canvas, candidate, subset, score):
|
||||
H, W, C = canvas.shape
|
||||
candidate = np.array(candidate)
|
||||
subset = np.array(subset)
|
||||
stickwidth = 4
|
||||
limbSeq = [[2, 3], [2, 6], [3, 4], [4, 5], [6, 7], [7, 8], [2, 9], [9, 10], \
|
||||
[10, 11], [2, 12], [12, 13], [13, 14], [2, 1], [1, 15], [15, 17], \
|
||||
[1, 16], [16, 18], [3, 17], [6, 18]]
|
||||
colors = [[255, 0, 0], [255, 85, 0], [255, 170, 0], [255, 255, 0], [170, 255, 0], [85, 255, 0], [0, 255, 0], \
|
||||
[0, 255, 85], [0, 255, 170], [0, 255, 255], [0, 170, 255], [0, 85, 255], [0, 0, 255], [85, 0, 255], \
|
||||
[170, 0, 255], [255, 0, 255], [255, 0, 170], [255, 0, 85]]
|
||||
for i in range(17):
|
||||
for n in range(len(subset)):
|
||||
index = subset[n][np.array(limbSeq[i]) - 1]
|
||||
conf = score[n][np.array(limbSeq[i]) - 1]
|
||||
if conf[0] < 0.3 or conf[1] < 0.3:
|
||||
continue
|
||||
Y = candidate[index.astype(int), 0] * float(W)
|
||||
X = candidate[index.astype(int), 1] * float(H)
|
||||
mX = np.mean(X)
|
||||
mY = np.mean(Y)
|
||||
length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5
|
||||
angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1]))
|
||||
polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stickwidth), int(angle), 0, 360, 1)
|
||||
cv2.fillConvexPoly(canvas, polygon, alpha_blend_color(colors[i], conf[0] * conf[1]))
|
||||
|
||||
canvas = (canvas * 0.6).astype(np.uint8)
|
||||
|
||||
for i in range(18):
|
||||
for n in range(len(subset)):
|
||||
index = int(subset[n][i])
|
||||
if index == -1:
|
||||
continue
|
||||
x, y = candidate[index][0:2]
|
||||
conf = score[n][i]
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
cv2.circle(canvas, (int(x), int(y)), 4, alpha_blend_color(colors[i], conf), thickness=-1)
|
||||
|
||||
return canvas
|
||||
|
||||
def draw_handpose_aligned(canvas, all_hand_peaks, all_hand_scores):
|
||||
H, W, C = canvas.shape
|
||||
|
||||
edges = [[0, 1], [1, 2], [2, 3], [3, 4], [0, 5], [5, 6], [6, 7], [7, 8], [0, 9], [9, 10], \
|
||||
[10, 11], [11, 12], [0, 13], [13, 14], [14, 15], [15, 16], [0, 17], [17, 18], [18, 19], [19, 20]]
|
||||
|
||||
for peaks, scores in zip(all_hand_peaks, all_hand_scores):
|
||||
|
||||
for ie, e in enumerate(edges):
|
||||
x1, y1 = peaks[e[0]]
|
||||
x2, y2 = peaks[e[1]]
|
||||
x1 = int(x1 * W)
|
||||
y1 = int(y1 * H)
|
||||
x2 = int(x2 * W)
|
||||
y2 = int(y2 * H)
|
||||
score = int(scores[e[0]] * scores[e[1]] * 255)
|
||||
if x1 > eps and y1 > eps and x2 > eps and y2 > eps:
|
||||
cv2.line(canvas, (x1, y1), (x2, y2),
|
||||
matplotlib.colors.hsv_to_rgb([ie / float(len(edges)), 1.0, 1.0]) * score, thickness=2)
|
||||
|
||||
for i, keyponit in enumerate(peaks):
|
||||
x, y = keyponit
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
score = int(scores[i] * 255)
|
||||
if x > eps and y > eps:
|
||||
cv2.circle(canvas, (x, y), 4, (0, 0, score), thickness=-1)
|
||||
return canvas
|
||||
|
||||
def draw_facepose_aligned(canvas, all_lmks, all_scores):
|
||||
H, W, C = canvas.shape
|
||||
for lmks, scores in zip(all_lmks, all_scores):
|
||||
for lmk, score in zip(lmks, scores):
|
||||
x, y = lmk
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
conf = int(score * 255)
|
||||
if x > eps and y > eps:
|
||||
cv2.circle(canvas, (x, y), 3, (conf, conf, conf), thickness=-1)
|
||||
return canvas
|
||||
|
||||
def draw_pose_aligned(pose, H, W, ref_w=2160):
|
||||
bodies = pose['bodies']
|
||||
faces = pose['faces']
|
||||
hands = pose['hands']
|
||||
candidate = bodies['candidate']
|
||||
subset = bodies['subset']
|
||||
sz = min(H, W)
|
||||
sr = (ref_w / sz) if sz != ref_w else 1
|
||||
canvas = np.zeros(shape=(int(H*sr), int(W*sr), 3), dtype=np.uint8)
|
||||
canvas = draw_bodypose_aligned(canvas, candidate, subset, score=bodies['score'])
|
||||
canvas = draw_handpose_aligned(canvas, hands, pose['hands_score'])
|
||||
canvas = draw_facepose_aligned(canvas, faces, pose['faces_score'])
|
||||
|
||||
return cv2.cvtColor(cv2.resize(canvas, (W, H)), cv2.COLOR_BGR2RGB).transpose(2, 0, 1)
|
||||
|
||||
|
||||
def get_image_pose(ref_image_path):
|
||||
ref_image = cv2.imread(ref_image_path)
|
||||
ref_image = cv2.cvtColor(ref_image, cv2.COLOR_BGR2RGB)
|
||||
height, width, _ = ref_image.shape
|
||||
ref_pose = dwpose_detector_aligned(ref_image)
|
||||
pose_img = draw_pose_aligned(ref_pose, height, width)
|
||||
return np.array(pose_img)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser("Training Skeleton Poses Extraction", add_help=True)
|
||||
parser.add_argument("--start", type=int, help="Specify the value of start")
|
||||
parser.add_argument("--end", type=int, help="Specify the value of end")
|
||||
parser.add_argument("--name", type=str, help="Specify the name of dataset")
|
||||
parser.add_argument("--root_path", type=str, help="Specify the root path of dataset")
|
||||
args = parser.parse_args()
|
||||
|
||||
start = args.start
|
||||
end = args.end
|
||||
dataset_name = args.name
|
||||
|
||||
image_root = os.path.join(args.root_path, dataset_name)
|
||||
for idx in range(start, end+1):
|
||||
subfolder = str(idx).zfill(5)
|
||||
subfolder_path = os.path.join(image_root, subfolder)
|
||||
images_subfolder_path = os.path.join(subfolder_path, "images")
|
||||
print(f"images subfolder path: {images_subfolder_path}")
|
||||
|
||||
pose_subfolder_path = os.path.join(subfolder_path, "poses")
|
||||
if not os.path.exists(pose_subfolder_path):
|
||||
os.makedirs(pose_subfolder_path)
|
||||
print(f"Folder created: {pose_subfolder_path}")
|
||||
else:
|
||||
print(f"Folder already exists: {pose_subfolder_path}")
|
||||
for root, dirs, files in os.walk(images_subfolder_path):
|
||||
for file in files:
|
||||
if file.endswith('.png'):
|
||||
file_path = os.path.join(root, file)
|
||||
print(file_path)
|
||||
file_name = os.path.splitext(file)[0]
|
||||
image_name = file_name + '.png'
|
||||
image_legal_path = os.path.join(images_subfolder_path, image_name)
|
||||
if os.path.exists(os.path.join(pose_subfolder_path, file_name + '.png')):
|
||||
existed_path = os.path.join(pose_subfolder_path, file_name + '.png')
|
||||
print(f"{existed_path} already exists!")
|
||||
continue
|
||||
detected_map = get_image_pose(image_legal_path)
|
||||
detected_map = np.transpose(detected_map, (1, 2, 0))
|
||||
pose_save_path = os.path.join(pose_subfolder_path, file_name + '.png')
|
||||
cv2.imwrite(pose_save_path, detected_map)
|
||||
print(f"Finish Pose Extraction: {pose_save_path}")
|
||||
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) Shuyuan Tu.
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE
|
||||
@@ -0,0 +1,384 @@
|
||||
# StableAnimator
|
||||
|
||||
<a href='https://francis-rings.github.io/StableAnimator'><img src='https://img.shields.io/badge/Project-Page-Green'></a> <a href='https://arxiv.org/abs/2411.17697'><img src='https://img.shields.io/badge/Paper-Arxiv-red'></a> <a href='https://huggingface.co/FrancisRing/StableAnimator/tree/main'><img src='https://img.shields.io/badge/HuggingFace-Model-orange'></a> <a href='https://www.youtube.com/watch?v=7fwFyFDzQgg'><img src='https://img.shields.io/badge/YouTube-Watch-red?style=flat-square&logo=youtube'></a> <a href='https://www.bilibili.com/video/BV1X5zyYUEuD'><img src='https://img.shields.io/badge/Bilibili-Watch-blue?style=flat-square&logo=bilibili'></a>
|
||||
|
||||
StableAnimator: High-Quality Identity-Preserving Human Image Animation
|
||||
<br/>
|
||||
*Shuyuan Tu<sup>1</sup>, Zhen Xing<sup>1</sup>, Xintong Han<sup>3</sup>, Zhi-Qi Cheng<sup>4</sup>, Qi Dai<sup>2</sup>, Chong Luo<sup>2</sup>, Zuxuan Wu<sup>1</sup>*
|
||||
<br/>
|
||||
[<sup>1</sup>Fudan University; <sup>2</sup>Microsoft Research Asia; <sup>3</sup>Huya Inc; <sup>4</sup>Carnegie Mellon University]
|
||||
|
||||
<p align="center">
|
||||
<img src="assets/figures/case-47.gif" width="256" />
|
||||
<img src="assets/figures/case-61.gif" width="256" />
|
||||
<img src="assets/figures/case-45.gif" width="256" />
|
||||
<img src="assets/figures/case-46.gif" width="256" />
|
||||
<img src="assets/figures/case-5.gif" width="256" />
|
||||
<img src="assets/figures/case-17.gif" width="256" />
|
||||
<br/>
|
||||
<span>Pose-driven Human image animations generated by StableAnimator, showing its power to synthesize <b>high-fidelity</b> and <b>ID-preserving videos</b>. All animations are <b>directly synthesized by StableAnimator without the use of any face-related post-processing tools</b>, such as the face-swapping tool FaceFusion or face restoration models like GFP-GAN and CodeFormer.</span>
|
||||
</p>
|
||||
|
||||
<p align="center">
|
||||
<img src="assets/figures/case-35.gif" width="384" />
|
||||
<img src="assets/figures/case-42.gif" width="384" />
|
||||
<img src="assets/figures/case-18.gif" width="384" />
|
||||
<img src="assets/figures/case-24.gif" width="384" />
|
||||
<br/>
|
||||
<span>Comparison results between StableAnimator and state-of-the-art (SOTA) human image animation models highlight the superior performance of StableAnimator in delivering <b>high-fidelity, identity-preserving human image animation</b>.</span>
|
||||
</p>
|
||||
|
||||
|
||||
## Overview
|
||||
|
||||
<p align="center">
|
||||
<img src="assets/figures/framework.jpg" alt="model architecture" width="1280"/>
|
||||
</br>
|
||||
<i>The overview of the framework of StableAnimator.</i>
|
||||
</p>
|
||||
|
||||
Current diffusion models for human image animation struggle to ensure identity (ID) consistency. This paper presents StableAnimator, <b>the first end-to-end ID-preserving video diffusion framework, which synthesizes high-quality videos without any post-processing, conditioned on a reference image and a sequence of poses.</b> Building upon a video diffusion model, StableAnimator contains carefully designed modules for both training and inference striving for identity consistency. In particular, StableAnimator begins by computing image and face embeddings with off-the-shelf extractors, respectively and face embeddings are further refined by interacting with image embeddings using a global content-aware Face Encoder. Then, StableAnimator introduces a novel distribution-aware ID Adapter that prevents interference caused by temporal layers while preserving ID via alignment. During inference, we propose a novel Hamilton-Jacobi-Bellman (HJB) equation-based optimization to further enhance the face quality. We demonstrate that solving the HJB equation can be integrated into the diffusion denoising process, and the resulting solution constrains the denoising path and thus benefits ID preservation. Experiments on multiple benchmarks show the effectiveness of StableAnimator both qualitatively and quantitatively.
|
||||
|
||||
## News
|
||||
* `[2024-12-13]`:🔥 The training code and training tutorial are released! You can train/finetune your own StableAnimator on your own collected datasets! Other codes will be released very soon. Stay tuned!
|
||||
* `[2024-12-10]`:🔥 The gradio interface is released! Many thanks to [@gluttony-10](https://space.bilibili.com/893892) for his contribution! Other codes will be released very soon. Stay tuned!
|
||||
* `[2024-12-6]`:🔥 All data preprocessing codes (human skeleton extraction and human face mask extraction) are released! The training code and detailed training tutorial will be released before 2024.12.13. Stay tuned!
|
||||
* `[2024-12-4]`:🔥 We are thrilled to release an interesting dance demo (🔥🔥APT Dance🔥🔥)! The generated video can be seen on [YouTube](https://www.youtube.com/watch?v=KNPoAsWr_sk) and [Bilibili](https://www.bilibili.com/video/BV1KczXYhER7).
|
||||
* `[2024-11-28]`:🔥 The data pre-processing codes (human skeleton extraction) are available! Other codes will be released very soon. Stay tuned!
|
||||
* `[2024-11-26]`:🔥 The project page, code, technical report and [a basic model checkpoint](https://huggingface.co/FrancisRing/StableAnimator/tree/main) are released. Further training codes, data pre-processing codes, the evaluation dataset and StableAnimator-pro will be released very soon. Stay tuned!
|
||||
|
||||
## To-Do List
|
||||
- [x] StableAnimator-basic
|
||||
- [x] Inference Code
|
||||
- [x] Evaluation Samples
|
||||
- [x] Data Pre-Processing Code (Skeleton Extraction)
|
||||
- [x] Data Pre-Processing Code (Human Face Mask Extraction)
|
||||
- [x] Training Code
|
||||
- [ ] Evaluation Dataset
|
||||
- [ ] StableAnimator-pro
|
||||
- [ ] Inference Code with HJB-based Face Optimization
|
||||
|
||||
## Quickstart
|
||||
|
||||
For the basic version of the model checkpoint, it supports generating videos at a 576x1024 or 512x512 resolution. If you encounter insufficient memory issues, you can appropriately reduce the number of animated frames.
|
||||
|
||||
### Environment setup
|
||||
|
||||
```
|
||||
pip install torch==2.5.1 torchvision==0.20.1 torchaudio==2.5.1 --index-url https://download.pytorch.org/whl/cu124
|
||||
pip install torch==2.5.1+cu124 xformers --index-url https://download.pytorch.org/whl/cu124
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
### Download weights
|
||||
If you encounter connection issues with Hugging Face, you can utilize the mirror endpoint by setting the environment variable: `export HF_ENDPOINT=https://hf-mirror.com`.
|
||||
Please download weights manually as follows:
|
||||
```
|
||||
cd StableAnimator
|
||||
git lfs install
|
||||
git clone https://huggingface.co/FrancisRing/StableAnimator checkpoints
|
||||
```
|
||||
All the weights should be organized in models as follows
|
||||
The overall file structure of this project should be organized as follows:
|
||||
```
|
||||
StableAnimator/
|
||||
├── DWPose
|
||||
├── animation
|
||||
├── checkpoints
|
||||
│ ├── DWPose
|
||||
│ │ ├── dw-ll_ucoco_384.onnx
|
||||
│ │ └── yolox_l.onnx
|
||||
│ ├── Animation
|
||||
│ │ ├── pose_net.pth
|
||||
│ │ ├── face_encoder.pth
|
||||
│ │ └── unet.pth
|
||||
│ ├── SVD
|
||||
│ │ ├── feature_extractor
|
||||
│ │ ├── image_encoder
|
||||
│ │ ├── scheduler
|
||||
│ │ ├── unet
|
||||
│ │ ├── vae
|
||||
│ │ ├── model_index.json
|
||||
│ │ ├── svd_xt.safetensors
|
||||
│ │ └── svd_xt_image_decoder.safetensors
|
||||
│ └── inference.zip
|
||||
├── models
|
||||
│ │ └── antelopev2
|
||||
│ │ ├── 1k3d68.onnx
|
||||
│ │ ├── 2d106det.onnx
|
||||
│ │ ├── genderage.onnx
|
||||
│ │ ├── glintr100.onnx
|
||||
│ │ └── scrfd_10g_bnkps.onnx
|
||||
├── app.py
|
||||
├── command_basic_infer.sh
|
||||
├── inference_basic.py
|
||||
├── requirement.txt
|
||||
```
|
||||
<b>Notably, there is a bug in the automatic download process of Antelopev2, with the error details described as follows:</b>
|
||||
```
|
||||
Traceback (most recent call last):
|
||||
File "/home/StableAnimator/inference_normal.py", line 243, in <module>
|
||||
face_model = FaceModel()
|
||||
File "/home/StableAnimator/animation/modules/face_model.py", line 11, in __init__
|
||||
self.app = FaceAnalysis(
|
||||
File "/opt/conda/lib/python3.10/site-packages/insightface/app/face_analysis.py", line 43, in __init__
|
||||
assert 'detection' in self.models
|
||||
AssertionError
|
||||
```
|
||||
This issue is related to the incorrect path of Antelopev2, which is automatically downloaded into the `models/antelopev2/antelopev2` directory. The correct path of Antelopev2 should be `models/antelopev2`. You can run the following commands to tackle this issue:
|
||||
```
|
||||
cd StableAnimator
|
||||
mv ./models/antelopev2/antelopev2 ./models/tmp
|
||||
rm -rf ./models/antelopev2
|
||||
mv ./models/tmp ./models/antelopev2
|
||||
```
|
||||
|
||||
### Evaluation Samples
|
||||
The evaluation samples presented in the paper can be downloaded from [OneDrive](https://1drv.ms/f/c/becb962aad1a1f95/EubdzCAI7BFLhJff2LrHkt8BC9mOiwJ5V67t-ypxRnCK4Q?e=ElEmcn) or `inference.zip` in checkpoints. Please download evaluation samples manually as follows:
|
||||
```
|
||||
cd StableAnimator
|
||||
mkdir inference
|
||||
```
|
||||
All the evaluation samples should be organized as follows:
|
||||
```
|
||||
inference/
|
||||
├── case-1
|
||||
│ ├── poses
|
||||
│ ├── faces
|
||||
│ └── reference.png
|
||||
├── case-2
|
||||
│ ├── poses
|
||||
│ ├── faces
|
||||
│ └── reference.png
|
||||
├── case-3
|
||||
│ ├── poses
|
||||
│ ├── faces
|
||||
│ └── reference.png
|
||||
```
|
||||
|
||||
### Human Skeleton Extraction
|
||||
We leverage the pre-trained DWPose to extract the human skeletons. In the initialization of DWPose, the pretrained weights should be configured in `/DWPose/dwpose_utils/wholebody.py`:
|
||||
```
|
||||
onnx_det = 'path/checkpoints/DWPose/yolox_l.onnx'
|
||||
onnx_pose = 'path/checkpoints/DWPose/dw-ll_ucoco_384.onnx'
|
||||
```
|
||||
Given the target image folder containing multiple .png files, you can use the following command to obtain the corresponding human skeleton images:
|
||||
```
|
||||
python DWPose/skeleton_extraction.py --target_image_folder_path="path/test/target_images" --ref_image_path="path/test/reference.png" --poses_folder_path="path/test/poses"
|
||||
```
|
||||
It is worth noting that the .png files in the target image folder are named in the format `frame_i.png`, such as `frame_0.png`, `frame_1.png`, and so on.
|
||||
`--ref_image_path` refers to the path of the given reference image. The obtained human skeleton images are saved in `path/test/poses`. It is particularly significant that the target skeleton images should be aligned with the reference image regarding the body shape.
|
||||
|
||||
If you only have the target MP4 file (target.mp4), we recommend you to use `ffmpeg` to convert the MP4 file to multiple frames (.png files) without any quality loss.
|
||||
```
|
||||
ffmpeg -i target.mp4 -q:v 1 -start_number 0 path/test/target_images/frame_%d.png
|
||||
```
|
||||
The obtained frames are saved in `path/test/target_images`.
|
||||
|
||||
### Human Face Mask Extraction
|
||||
Given the path to an image folder containing multiple RGB `.png` files, you can run the following command to extract the corresponding human face masks:
|
||||
```
|
||||
python face_mask_extraction.py --image_folder="path/StableAnimator/inference/your_case/target_images"
|
||||
```
|
||||
`path/StableAnimator/inference/your_case/target_images` contains multiple `.png` files. The obtained masks are saved in `path/StableAnimator/inference/your_case/faces`.
|
||||
|
||||
### Model inference
|
||||
A sample configuration for testing is provided as `command_basic_infer.sh`. You can also easily modify the various configurations according to your needs.
|
||||
|
||||
```
|
||||
bash command_basic_infer.sh
|
||||
```
|
||||
StableAnimator supports human image animation at two different resolution settings: 512x512 and 576x1024. You can modify "--width" and "--height" in `command_basic_infer.sh` to set the resolution of the animation. "--output_dir" in `command_basic_infer.sh` refers to the saved path of the generated animation. "--validation_control_folder" and "--validation_image" in `command_basic_infer.sh` refer to the paths of the given pose sequence and the reference image, respectively.
|
||||
"--pretrained_model_name_or_path" in `command_basic_infer.sh` is the path of pretrained SVD. "posenet_model_name_or_path", "face_encoder_model_name_or_path", and "unet_model_name_or_path" in `command_basic_infer.sh` refer to paths of pretrained StableAnimator weights.
|
||||
If you have enough GPU resources, you can increase the value (4=>8=>16) of "--decode_chunk_size" in `command_basic_infer.sh` to promote the temporal smoothness of the animation.
|
||||
|
||||
Tips: if your GPU memory is limited, you can reduce the number of animated frames. This command will generate two files: <b>animated_images</b> and <b>animated_images.gif</b>.
|
||||
If you want to obtain the high quality MP4 file, we recommend you to leverage ffmpeg on the <b>animated_images</b> as follows:
|
||||
```
|
||||
cd animated_images
|
||||
ffmpeg -framerate 20 -i frame_%d.png -c:v libx264 -crf 10 -pix_fmt yuv420p /path/animation.mp4
|
||||
```
|
||||
"-framerate" refers to the fps setting. "-crf" indicates the quality of the generated MP4 file, with smaller values corresponding to higher quality.
|
||||
Additionally, you can also run the following command to launch a Gradio interface:
|
||||
```
|
||||
python app.py
|
||||
```
|
||||
|
||||
### Model Training
|
||||
<b>🔥It’s worth noting that if you’re looking to train a conditioned Stable Video Diffusion (SVD) model, this training tutorial will also be helpful.🔥</b>
|
||||
For the training dataset, it has to be organized as follows:
|
||||
|
||||
```
|
||||
animation_data/
|
||||
├── rec
|
||||
│ │ ├──00001
|
||||
│ │ │ ├──images
|
||||
│ │ │ │ ├──frame_0.png
|
||||
│ │ │ │ ├──frame_1.png
|
||||
│ │ │ │ ├──frame_2.png
|
||||
│ │ │ │ └──...
|
||||
│ │ │ ├──faces
|
||||
│ │ │ │ ├──frame_0.png
|
||||
│ │ │ │ ├──frame_1.png
|
||||
│ │ │ │ ├──frame_2.png
|
||||
│ │ │ │ └──...
|
||||
│ │ │ └──poses
|
||||
│ │ │ │ ├──frame_0.png
|
||||
│ │ │ │ ├──frame_1.png
|
||||
│ │ │ │ ├──frame_2.png
|
||||
│ │ │ │ └──...
|
||||
│ │ ├──00002
|
||||
│ │ │ ├──images
|
||||
│ │ │ ├──faces
|
||||
│ │ │ └──poses
|
||||
│ │ ├──00003
|
||||
│ │ │ ├──images
|
||||
│ │ │ ├──faces
|
||||
│ │ │ └──poses
|
||||
│ │ └──...
|
||||
├── vec
|
||||
│ │ ├──00001
|
||||
│ │ │ ├──images
|
||||
│ │ │ ├──faces
|
||||
│ │ │ └──poses
|
||||
│ │ ├──00002
|
||||
│ │ │ ├──images
|
||||
│ │ │ ├──faces
|
||||
│ │ │ └──poses
|
||||
│ │ ├──00003
|
||||
│ │ │ ├──images
|
||||
│ │ │ ├──faces
|
||||
│ │ │ └──poses
|
||||
│ │ └──...
|
||||
├── video_rec_path.txt
|
||||
└── video_vec_path.txt
|
||||
```
|
||||
StableAnimator is trained on mixed-resolution videos, with 512x512 videos stored in `animation_data/rec` and 576x1024 videos stored in `animation_data/vec`. Each folder in `animation_data/rec` or `animation_data/vec` contains three subfolders which contains multiple `.png` image files.
|
||||
All `.png` image files are named in the format `frame_i.png`, such as `frame_0.png`, `frame_1.png`, and so on.
|
||||
`00001`, `00002`, `00003` indicate individual video information.
|
||||
In terms of three subfolders, `images`, `faces`, and `poses` store RGB frames, corresponding human face masks, and corresponding human skeleton poses, respectively.
|
||||
`video_rec_path.txt` and `video_vec_path.txt` record folder paths of `animation_data/rec` and `animation_data/vec`, respectively.
|
||||
For example, the content of `video_rec_path.txt` is shown as follows:
|
||||
```
|
||||
path/StableAnimator/animation_data/rec/00001
|
||||
path/StableAnimator/animation_data/rec/00002
|
||||
path/StableAnimator/animation_data/rec/00003
|
||||
path/StableAnimator/animation_data/rec/00004
|
||||
path/StableAnimator/animation_data/rec/00005
|
||||
path/StableAnimator/animation_data/rec/00006
|
||||
...
|
||||
```
|
||||
If you only have raw videos, you can leverage `ffmpeg` to extract frames from raw videos and store them in the subfolder `images`.
|
||||
```
|
||||
ffmpeg -i raw_video_1.mp4 -q:v 1 -start_number 0 path/StableAnimator/animation_data/rec/00001/images/frame_%d.png
|
||||
```
|
||||
The obtained frames are saved in `path/StableAnimator/animation_data/rec/00001/images`.
|
||||
|
||||
For extracting the human skeleton poses, you can run the following command:
|
||||
```
|
||||
python DWPose/training_skeleton_extraction.py --root_path="path/StableAnimator/animation_data" --name="rec" --start=1 --end=500
|
||||
```
|
||||
`--root_path` and `--name` refer to the root path of training datasets and the name of the dataset.
|
||||
`--start` and `--end` specify the starting and ending indices of the selected training dataset. For example, `--name="rec" --start=1 --end=500` indicates that the skeleton extraction will start at `path/StableAnimator/animation_data/rec/00001` and end at `path/StableAnimator/animation_data/rec/00500`.
|
||||
|
||||
For extraction details of corresponding face masks, please refer to the Human Face Mask Extraction section.
|
||||
When your dataset is organized exactly as outlined above, you can easily train your StableAnimator by running the following command:
|
||||
```
|
||||
bash command_train.sh
|
||||
```
|
||||
For the parameter details of `command_train.sh`, `CUDA_VISIBLE_DEVICES` refers to gpu devices. In my setting, I use 4 NVIDIA A100 80G to train StableAnimator (`CUDA_VISIBLE_DEVICES=3,2,1,0`).
|
||||
`--pretrained_model_name_or_path` and `--output_dir` refer to the pretrained SVD path and the checkpoint saved path of the trained StableAnimator.
|
||||
`--data_root_path`, `--rec_data_path`, and `--vec_data_path` are the root path of datasets, the path of `video_rec_path.txt`, and the path of `video_vec_path.txt`, respectively.
|
||||
`validation_image_folder`, `validation_control_folder`, and `validation_image` are paths of validation ground truths, validation driven skeleton poses, and the validation reference image.
|
||||
`--sample_n_frames` is the number of frames that StableAnimator processes in a single batch.
|
||||
`--num_train_epochs` is the training epoch number. It is worth noting that the default number of training epochs is set to infinite. You can manually terminate the training process once you observe that your StableAnimator has reached its peak performance.
|
||||
The overall file structure of StableAnimator at training is shown as follows:
|
||||
```
|
||||
StableAnimator/
|
||||
├── DWPose
|
||||
├── animation
|
||||
├── animation_data
|
||||
│ ├── rec
|
||||
│ ├── vec
|
||||
│ ├── video_rec_path.txt
|
||||
│ └── video_vec_path.txt
|
||||
├── validation
|
||||
│ ├── ground_truth
|
||||
│ │ ├── frame_0.png
|
||||
│ │ ├── frame_1.png
|
||||
│ │ ├── frame_2.png
|
||||
│ │ └── ...
|
||||
│ ├── poses
|
||||
│ │ ├── frame_0.png
|
||||
│ │ ├── frame_1.png
|
||||
│ │ ├── frame_2.png
|
||||
│ │ └── ...
|
||||
│ └── reference.png
|
||||
├── checkpoints
|
||||
│ ├── DWPose
|
||||
│ │ ├── dw-ll_ucoco_384.onnx
|
||||
│ │ └── yolox_l.onnx
|
||||
│ ├── Animation
|
||||
│ │ ├── pose_net.pth
|
||||
│ │ ├── face_encoder.pth
|
||||
│ │ └── unet.pth
|
||||
│ ├── SVD
|
||||
│ │ ├── feature_extractor
|
||||
│ │ ├── image_encoder
|
||||
│ │ ├── scheduler
|
||||
│ │ ├── unet
|
||||
│ │ ├── vae
|
||||
│ │ ├── model_index.json
|
||||
│ │ ├── svd_xt.safetensors
|
||||
│ │ └── svd_xt_image_decoder.safetensors
|
||||
│ └── inference.zip
|
||||
├── models
|
||||
│ │ └── antelopev2
|
||||
│ │ ├── 1k3d68.onnx
|
||||
│ │ ├── 2d106det.onnx
|
||||
│ │ ├── genderage.onnx
|
||||
│ │ ├── glintr100.onnx
|
||||
│ │ └── scrfd_10g_bnkps.onnx
|
||||
├── app.py
|
||||
├── command_basic_infer.sh
|
||||
├── inference_basic.py
|
||||
├── train.py
|
||||
├── command_train.sh
|
||||
└── requirement.txt
|
||||
```
|
||||
<b>It is worth noting that training StableAnimator requires approximately 70GB of VRAM due to the mixed-resolution (512x512 and 576x1024) training pipeline.
|
||||
However, if you train StableAnimator exclusively on 512x512 videos, the VRAM requirement is reduced to approximately 40GB.</b>
|
||||
Additionally, The backgrounds of the selected training videos should remain static, as this helps the diffusion model calculate accurate reconstruction loss.
|
||||
|
||||
If you want to train StableAnimator on a single resolution, you can use the following command:
|
||||
```
|
||||
bash command_train_single.sh
|
||||
```
|
||||
You can customize the resolution by modifying `--dataset_width` and `--dataset_height`, both of which default to 512.
|
||||
|
||||
Regarding finetuning StableAnimator, you can run the following command:
|
||||
```
|
||||
bash command_finetune.sh
|
||||
```
|
||||
`posenet_model_finetune_path`, `face_encoder_finetune_path`, and `unet_model_finetune_path` in `command_finetune.sh` refer to paths of pretrained StableAnimator weights.
|
||||
|
||||
### VRAM requirement and Runtime
|
||||
|
||||
For the 15s demo video (512x512, fps=30), the 16-frame basic model requires 8GB VRAM and finishes in 5 minutes on a 4090 GPU.
|
||||
|
||||
The minimum VRAM requirement for the 16-frame U-Net of the pro model is 10GB (576x1024, fps=30); however, the VAE decoder demands 16GB. You have the option to run the VAE decoder on CPU.
|
||||
|
||||
## Contact
|
||||
If you have any suggestions or find our work helpful, feel free to contact me
|
||||
|
||||
Email: francisshuyuan@gmail.com
|
||||
|
||||
If you find our work useful, <b>please consider giving a star to this github repository and citing it</b>:
|
||||
```bib
|
||||
@article{tu2024stableanimator,
|
||||
title={StableAnimator: High-Quality Identity-Preserving Human Image Animation},
|
||||
author={Shuyuan Tu and Zhen Xing and Xintong Han and Zhi-Qi Cheng and Qi Dai and Chong Luo and Zuxuan Wu},
|
||||
journal={arXiv preprint arXiv:2411.17697},
|
||||
year={2024}
|
||||
}
|
||||
```
|
||||
@@ -0,0 +1,226 @@
|
||||
import os, io, csv, math, random
|
||||
from importlib.metadata import files
|
||||
import os.path as osp
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
from torch.utils.data.dataset import Dataset
|
||||
from einops import rearrange
|
||||
import cv2
|
||||
import warnings
|
||||
|
||||
|
||||
class LargeScaleAnimationVideos(Dataset):
|
||||
def __init__(self, root_path, txt_path, width, height, n_sample_frames, sample_frame_rate, sample_margin=30,
|
||||
app=None, handler_ante=None, face_helper=None):
|
||||
self.root_path = root_path
|
||||
self.txt_path = txt_path
|
||||
self.width = width
|
||||
self.height = height
|
||||
self.n_sample_frames = n_sample_frames
|
||||
self.sample_frame_rate = sample_frame_rate
|
||||
self.sample_margin = sample_margin
|
||||
|
||||
self.video_files = self._read_txt_file_images()
|
||||
|
||||
self.app = app
|
||||
self.handler_ante = handler_ante
|
||||
self.face_helper = face_helper
|
||||
|
||||
def _read_txt_file_images(self):
|
||||
with open(self.txt_path, 'r') as file:
|
||||
lines = file.readlines()
|
||||
video_files = []
|
||||
for line in lines:
|
||||
video_file = line.strip()
|
||||
video_files.append(video_file)
|
||||
return video_files
|
||||
|
||||
def __len__(self):
|
||||
return len(self.video_files)
|
||||
|
||||
def frame_count(self, frames_path):
|
||||
files = os.listdir(frames_path)
|
||||
png_files = [file for file in files if file.endswith('.png') or file.endswith('.jpg')]
|
||||
png_files_count = len(png_files)
|
||||
return png_files_count
|
||||
|
||||
def find_frames_list(self, frames_path):
|
||||
files = os.listdir(frames_path)
|
||||
image_files = [file for file in files if file.endswith('.png') or file.endswith('.jpg')]
|
||||
if image_files[0].startswith('frame_'):
|
||||
image_files.sort(key=lambda x: int(x.split('_')[1].split('.')[0]))
|
||||
else:
|
||||
image_files.sort(key=lambda x: int(x.split('.')[0]))
|
||||
return image_files
|
||||
|
||||
def get_face_masks(self, pil_img):
|
||||
rgb_image = np.array(pil_img)
|
||||
bgr_image = cv2.cvtColor(rgb_image, cv2.COLOR_RGB2BGR)
|
||||
image_info = self.app.get(bgr_image)
|
||||
mask = np.zeros((self.height, self.width), dtype=np.uint8)
|
||||
|
||||
if len(image_info) > 0:
|
||||
for info in image_info:
|
||||
x_1 = info['bbox'][0]
|
||||
y_1 = info['bbox'][1]
|
||||
x_2 = info['bbox'][2]
|
||||
y_2 = info['bbox'][3]
|
||||
cv2.rectangle(mask, (int(x_1), int(y_1)), (int(x_2), int(y_2)), (255), thickness=cv2.FILLED)
|
||||
mask = mask.astype(np.float64) / 255.0
|
||||
else:
|
||||
self.face_helper.clean_all()
|
||||
with torch.no_grad():
|
||||
bboxes = self.face_helper.face_det.detect_faces(bgr_image, 0.97)
|
||||
if len(bboxes) > 0:
|
||||
for bbox in bboxes:
|
||||
cv2.rectangle(mask, (int(bbox[0]), int(bbox[1])), (int(bbox[2]), int(bbox[3])), (255),
|
||||
thickness=cv2.FILLED)
|
||||
mask = mask.astype(np.float64) / 255.0
|
||||
else:
|
||||
mask = np.ones((self.height, self.width), dtype=np.uint8)
|
||||
return mask
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
warnings.filterwarnings('ignore', category=DeprecationWarning)
|
||||
warnings.filterwarnings('ignore', category=FutureWarning)
|
||||
|
||||
frames_path = osp.join(self.video_files[idx], "images")
|
||||
poses_path = osp.join(self.video_files[idx], "poses")
|
||||
face_masks_path = osp.join(self.video_files[idx], "faces")
|
||||
video_length = self.frame_count(frames_path)
|
||||
frames_list = self.find_frames_list(frames_path)
|
||||
|
||||
clip_length = min(video_length, (self.n_sample_frames - 1) * self.sample_frame_rate + 1)
|
||||
|
||||
# print("-------------------------------")
|
||||
# print(clip_length)
|
||||
# print(video_length)
|
||||
# print(type(random))
|
||||
# print("-------------------------------")
|
||||
|
||||
start_idx = random.randint(0, video_length - clip_length)
|
||||
batch_index = np.linspace(
|
||||
start_idx, start_idx + clip_length - 1, self.n_sample_frames, dtype=int
|
||||
).tolist()
|
||||
all_indices = list(range(0, video_length))
|
||||
available_indices = [i for i in all_indices if i not in batch_index]
|
||||
reference_frame_idx = None
|
||||
if available_indices:
|
||||
reference_frame_idx = random.choice(available_indices)
|
||||
else:
|
||||
print("There is no available frame")
|
||||
extreme_sample_frame_rate = 2
|
||||
extreme_clip_length = min(video_length, (self.n_sample_frames - 1) * extreme_sample_frame_rate + 1)
|
||||
extreme_start_idx = random.randint(0, video_length - extreme_clip_length)
|
||||
extreme_batch_index = np.linspace(
|
||||
extreme_start_idx, extreme_start_idx + extreme_clip_length - 1, self.n_sample_frames, dtype=int
|
||||
).tolist()
|
||||
extreme_available_indices = [i for i in all_indices if i not in extreme_batch_index]
|
||||
if extreme_available_indices:
|
||||
reference_frame_idx = random.choice(extreme_available_indices)
|
||||
else:
|
||||
print("There is no available frame in the extreme circumstance")
|
||||
print(frames_path)
|
||||
print(1 / 0)
|
||||
|
||||
pose_pil_image_list = []
|
||||
tgt_pil_image_list = []
|
||||
tgt_face_masks_list = []
|
||||
|
||||
reference_frame_path = osp.join(frames_path, frames_list[reference_frame_idx])
|
||||
reference_pil_image = Image.open(reference_frame_path).convert('RGB')
|
||||
reference_pil_image = reference_pil_image.resize((self.width, self.height))
|
||||
reference_pil_image = torch.from_numpy(np.array(reference_pil_image)).float()
|
||||
reference_pil_image = reference_pil_image / 127.5 - 1
|
||||
|
||||
self.face_helper.clean_all()
|
||||
reference_frame_face = cv2.imread(reference_frame_path)
|
||||
reference_frame_face = cv2.resize(reference_frame_face, (self.width, self.height))
|
||||
# reference_frame_face_bgr = cv2.cvtColor(reference_frame_face, cv2.COLOR_RGB2BGR)
|
||||
reference_frame_face_info = self.app.get(reference_frame_face)
|
||||
if len(reference_frame_face_info) > 0:
|
||||
reference_frame_face_info = sorted(reference_frame_face_info, key=lambda x: (x['bbox'][2] - x['bbox'][0]) * (x['bbox'][3] - x['bbox'][1]))[-1]
|
||||
reference_frame_id_ante_embedding = reference_frame_face_info['embedding']
|
||||
else:
|
||||
reference_frame_id_ante_embedding = None
|
||||
|
||||
if reference_frame_id_ante_embedding is None:
|
||||
self.face_helper.read_image(reference_frame_face)
|
||||
self.face_helper.get_face_landmarks_5(only_center_face=True)
|
||||
self.face_helper.align_warp_face()
|
||||
|
||||
if len(self.face_helper.cropped_faces) == 0:
|
||||
reference_frame_id_ante_embedding = np.zeros((512,))
|
||||
else:
|
||||
reference_frame_align_face = self.face_helper.cropped_faces[0]
|
||||
print('fail to detect face using insightface, extract embedding on align face')
|
||||
reference_frame_id_ante_embedding = self.handler_ante.get_feat(reference_frame_align_face)
|
||||
|
||||
for index in batch_index:
|
||||
tgt_img_path = osp.join(frames_path, frames_list[index])
|
||||
pose_name = os.path.splitext(os.path.basename(tgt_img_path))[0]
|
||||
pose_name = pose_name + '.png'
|
||||
face_name = pose_name
|
||||
pose_path = osp.join(poses_path, pose_name)
|
||||
face_mask_path = osp.join(face_masks_path, face_name)
|
||||
|
||||
try:
|
||||
tgt_img_pil = Image.open(tgt_img_path).convert('RGB')
|
||||
except Exception as e:
|
||||
print(f"Fail loading the image: {tgt_img_path}")
|
||||
|
||||
# tgt_face_mask = self.get_face_masks(tgt_img_pil)
|
||||
|
||||
try:
|
||||
tgt_face_mask = Image.open(face_mask_path)
|
||||
tgt_face_mask = tgt_face_mask.resize((self.width, self.height))
|
||||
tgt_face_mask = torch.from_numpy(np.array(tgt_face_mask)).float()
|
||||
tgt_face_mask = tgt_face_mask / 255
|
||||
except Exception as e:
|
||||
print(f"Fail loading the face masks: {face_mask_path}")
|
||||
tgt_face_mask = torch.ones(self.height, self.width, 1)
|
||||
tgt_face_masks_list.append(tgt_face_mask)
|
||||
|
||||
tgt_img_pil = tgt_img_pil.resize((self.width, self.height))
|
||||
tgt_img_tensor = torch.from_numpy(np.array(tgt_img_pil)).float()
|
||||
tgt_img_normalized = tgt_img_tensor / 127.5 - 1
|
||||
tgt_pil_image_list.append(tgt_img_normalized)
|
||||
|
||||
try:
|
||||
pose = Image.open(pose_path).convert('RGB')
|
||||
pose = pose.resize((self.width, self.height))
|
||||
pose = torch.from_numpy(np.array(pose)).float()
|
||||
pose = pose / 127.5 - 1
|
||||
except Exception as e:
|
||||
print(f"Fail loading the poses: {pose_path}")
|
||||
pose = torch.zeros_like(reference_pil_image)
|
||||
pose_pil_image_list.append(pose)
|
||||
|
||||
pose_pil_image_list = torch.stack(pose_pil_image_list, dim=0)
|
||||
tgt_pil_image_list = torch.stack(tgt_pil_image_list, dim=0)
|
||||
tgt_pil_image_list = rearrange(tgt_pil_image_list, "f h w c -> f c h w")
|
||||
reference_pil_image = rearrange(reference_pil_image, "h w c -> c h w")
|
||||
pose_pil_image_list = rearrange(pose_pil_image_list, "f h w c -> f c h w")
|
||||
|
||||
tgt_face_masks_list = torch.stack(tgt_face_masks_list, dim=0)
|
||||
tgt_face_masks_list = torch.unsqueeze(tgt_face_masks_list, dim=-1)
|
||||
tgt_face_masks_list = rearrange(tgt_face_masks_list, "f h w c -> f c h w")
|
||||
|
||||
sample = dict(
|
||||
pixel_values=tgt_pil_image_list,
|
||||
reference_image=reference_pil_image,
|
||||
pose_pixels=pose_pil_image_list,
|
||||
faceid_embeds=reference_frame_id_ante_embedding,
|
||||
tgt_face_masks=tgt_face_masks_list,
|
||||
)
|
||||
|
||||
return sample
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,378 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import torch
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.models.attention import BasicTransformerBlock, TemporalBasicTransformerBlock
|
||||
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.models.resnet import AlphaBlender
|
||||
from diffusers.utils import BaseOutput
|
||||
from torch import nn
|
||||
|
||||
|
||||
@dataclass
|
||||
class TransformerTemporalModelOutput(BaseOutput):
|
||||
"""
|
||||
The output of [`TransformerTemporalModel`].
|
||||
|
||||
Args:
|
||||
sample (`torch.FloatTensor` of shape `(batch_size x num_frames, num_channels, height, width)`):
|
||||
The hidden states output conditioned on `encoder_hidden_states` input.
|
||||
"""
|
||||
|
||||
sample: torch.FloatTensor
|
||||
|
||||
|
||||
class TransformerTemporalModel(ModelMixin, ConfigMixin):
|
||||
"""
|
||||
A Transformer model for video-like data.
|
||||
|
||||
Parameters:
|
||||
num_attention_heads (`int`, *optional*, defaults to 16): The number of heads to use for multi-head attention.
|
||||
attention_head_dim (`int`, *optional*, defaults to 88): The number of channels in each head.
|
||||
in_channels (`int`, *optional*):
|
||||
The number of channels in the input and output (specify if the input is **continuous**).
|
||||
num_layers (`int`, *optional*, defaults to 1): The number of layers of Transformer blocks to use.
|
||||
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
|
||||
cross_attention_dim (`int`, *optional*): The number of `encoder_hidden_states` dimensions to use.
|
||||
attention_bias (`bool`, *optional*):
|
||||
Configure if the `TransformerBlock` attention should contain a bias parameter.
|
||||
sample_size (`int`, *optional*): The width of the latent images (specify if the input is **discrete**).
|
||||
This is fixed during training since it is used to learn a number of position embeddings.
|
||||
activation_fn (`str`, *optional*, defaults to `"geglu"`):
|
||||
Activation function to use in feed-forward. See `diffusers.models.activations.get_activation` for supported
|
||||
activation functions.
|
||||
norm_elementwise_affine (`bool`, *optional*):
|
||||
Configure if the `TransformerBlock` should use learnable elementwise affine parameters for normalization.
|
||||
double_self_attention (`bool`, *optional*):
|
||||
Configure if each `TransformerBlock` should contain two self-attention layers.
|
||||
positional_embeddings: (`str`, *optional*):
|
||||
The type of positional embeddings to apply to the sequence input before passing use.
|
||||
num_positional_embeddings: (`int`, *optional*):
|
||||
The maximum length of the sequence over which to apply positional embeddings.
|
||||
"""
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int = 16,
|
||||
attention_head_dim: int = 88,
|
||||
in_channels: Optional[int] = None,
|
||||
out_channels: Optional[int] = None,
|
||||
num_layers: int = 1,
|
||||
dropout: float = 0.0,
|
||||
norm_num_groups: int = 32,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
attention_bias: bool = False,
|
||||
sample_size: Optional[int] = None,
|
||||
activation_fn: str = "geglu",
|
||||
norm_elementwise_affine: bool = True,
|
||||
double_self_attention: bool = True,
|
||||
positional_embeddings: Optional[str] = None,
|
||||
num_positional_embeddings: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.attention_head_dim = attention_head_dim
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = torch.nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True)
|
||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||
|
||||
# 3. Define transformers blocks
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
BasicTransformerBlock(
|
||||
inner_dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
dropout=dropout,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
activation_fn=activation_fn,
|
||||
attention_bias=attention_bias,
|
||||
double_self_attention=double_self_attention,
|
||||
norm_elementwise_affine=norm_elementwise_affine,
|
||||
positional_embeddings=positional_embeddings,
|
||||
num_positional_embeddings=num_positional_embeddings,
|
||||
)
|
||||
for d in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
self.proj_out = nn.Linear(inner_dim, in_channels)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.FloatTensor,
|
||||
encoder_hidden_states: Optional[torch.LongTensor] = None,
|
||||
timestep: Optional[torch.LongTensor] = None,
|
||||
class_labels: torch.LongTensor = None,
|
||||
num_frames: int = 1,
|
||||
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
return_dict: bool = True,
|
||||
) -> TransformerTemporalModelOutput:
|
||||
"""
|
||||
The [`TransformerTemporal`] forward method.
|
||||
|
||||
Args:
|
||||
hidden_states (`torch.LongTensor` of shape `(batch size, num latent pixels)` if discrete,
|
||||
`torch.FloatTensor` of shape `(batch size, channel, height, width)`if continuous): Input hidden_states.
|
||||
encoder_hidden_states ( `torch.LongTensor` of shape `(batch size, encoder_hidden_states dim)`, *optional*):
|
||||
Conditional embeddings for cross attention layer. If not given, cross-attention defaults to
|
||||
self-attention.
|
||||
timestep ( `torch.LongTensor`, *optional*):
|
||||
Used to indicate denoising step. Optional timestep to be applied as an embedding in `AdaLayerNorm`.
|
||||
class_labels ( `torch.LongTensor` of shape `(batch size, num classes)`, *optional*):
|
||||
Used to indicate class labels conditioning. Optional class labels to be applied as an embedding in
|
||||
`AdaLayerZeroNorm`.
|
||||
num_frames (`int`, *optional*, defaults to 1):
|
||||
The number of frames to be processed per batch. This is used to reshape the hidden states.
|
||||
cross_attention_kwargs (`dict`, *optional*):
|
||||
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||
`self.processor` in [diffusers.models.attention_processor](
|
||||
https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~models.unets.unet_2d_condition.UNet2DConditionOutput`] instead of a plain
|
||||
tuple.
|
||||
|
||||
Returns:
|
||||
[`~models.transformer_temporal.TransformerTemporalModelOutput`] or `tuple`:
|
||||
If `return_dict` is True, an [`~models.transformer_temporal.TransformerTemporalModelOutput`] is
|
||||
returned, otherwise a `tuple` where the first element is the sample tensor.
|
||||
"""
|
||||
# 1. Input
|
||||
batch_frames, channel, height, width = hidden_states.shape
|
||||
batch_size = batch_frames // num_frames
|
||||
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = hidden_states[None, :].reshape(batch_size, num_frames, channel, height, width)
|
||||
hidden_states = hidden_states.permute(0, 2, 1, 3, 4)
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
hidden_states = hidden_states.permute(0, 3, 4, 2, 1).reshape(batch_size * height * width, num_frames, channel)
|
||||
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
|
||||
# 2. Blocks
|
||||
for block in self.transformer_blocks:
|
||||
hidden_states = block(
|
||||
hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
cross_attention_kwargs=cross_attention_kwargs,
|
||||
class_labels=class_labels,
|
||||
)
|
||||
|
||||
# 3. Output
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
hidden_states = (
|
||||
hidden_states[None, None, :]
|
||||
.reshape(batch_size, height, width, num_frames, channel)
|
||||
.permute(0, 3, 4, 1, 2)
|
||||
.contiguous()
|
||||
)
|
||||
hidden_states = hidden_states.reshape(batch_frames, channel, height, width)
|
||||
|
||||
output = hidden_states + residual
|
||||
|
||||
if not return_dict:
|
||||
return (output,)
|
||||
|
||||
return TransformerTemporalModelOutput(sample=output)
|
||||
|
||||
|
||||
class TransformerSpatioTemporalModel(nn.Module):
|
||||
"""
|
||||
A Transformer model for video-like data.
|
||||
|
||||
Parameters:
|
||||
num_attention_heads (`int`, *optional*, defaults to 16): The number of heads to use for multi-head attention.
|
||||
attention_head_dim (`int`, *optional*, defaults to 88): The number of channels in each head.
|
||||
in_channels (`int`, *optional*):
|
||||
The number of channels in the input and output (specify if the input is **continuous**).
|
||||
out_channels (`int`, *optional*):
|
||||
The number of channels in the output (specify if the input is **continuous**).
|
||||
num_layers (`int`, *optional*, defaults to 1): The number of layers of Transformer blocks to use.
|
||||
cross_attention_dim (`int`, *optional*): The number of `encoder_hidden_states` dimensions to use.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int = 16,
|
||||
attention_head_dim: int = 88,
|
||||
in_channels: int = 320,
|
||||
out_channels: Optional[int] = None,
|
||||
num_layers: int = 1,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.attention_head_dim = attention_head_dim
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
self.inner_dim = inner_dim
|
||||
|
||||
# 2. Define input layers
|
||||
self.in_channels = in_channels
|
||||
self.norm = torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6)
|
||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||
|
||||
# 3. Define transformers blocks
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
BasicTransformerBlock(
|
||||
inner_dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
)
|
||||
for d in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
time_mix_inner_dim = inner_dim
|
||||
self.temporal_transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
TemporalBasicTransformerBlock(
|
||||
inner_dim,
|
||||
time_mix_inner_dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
)
|
||||
for _ in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
time_embed_dim = in_channels * 4
|
||||
self.time_pos_embed = TimestepEmbedding(in_channels, time_embed_dim, out_dim=in_channels)
|
||||
self.time_proj = Timesteps(in_channels, True, 0)
|
||||
self.time_mixer = AlphaBlender(alpha=0.5, merge_strategy="learned_with_images")
|
||||
|
||||
# 4. Define output layers
|
||||
self.out_channels = in_channels if out_channels is None else out_channels
|
||||
# TODO: should use out_channels for continuous projections
|
||||
self.proj_out = nn.Linear(inner_dim, in_channels)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
image_only_indicator: Optional[torch.Tensor] = None,
|
||||
return_dict: bool = True,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
hidden_states (`torch.FloatTensor` of shape `(batch size, channel, height, width)`):
|
||||
Input hidden_states.
|
||||
num_frames (`int`):
|
||||
The number of frames to be processed per batch. This is used to reshape the hidden states.
|
||||
encoder_hidden_states ( `torch.LongTensor` of shape `(batch size, encoder_hidden_states dim)`, *optional*):
|
||||
Conditional embeddings for cross attention layer. If not given, cross-attention defaults to
|
||||
self-attention.
|
||||
image_only_indicator (`torch.LongTensor` of shape `(batch size, num_frames)`, *optional*):
|
||||
A tensor indicating whether the input contains only images. 1 indicates that the input contains only
|
||||
images, 0 indicates that the input contains video frames.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~models.transformer_temporal.TransformerTemporalModelOutput`]
|
||||
instead of a plain tuple.
|
||||
|
||||
Returns:
|
||||
[`~models.transformer_temporal.TransformerTemporalModelOutput`] or `tuple`:
|
||||
If `return_dict` is True, an [`~models.transformer_temporal.TransformerTemporalModelOutput`] is
|
||||
returned, otherwise a `tuple` where the first element is the sample tensor.
|
||||
"""
|
||||
# 1. Input
|
||||
batch_frames, _, height, width = hidden_states.shape
|
||||
num_frames = image_only_indicator.shape[-1]
|
||||
batch_size = batch_frames // num_frames
|
||||
|
||||
time_context = encoder_hidden_states
|
||||
time_context_first_timestep = time_context[None, :].reshape(
|
||||
batch_size, num_frames, -1, time_context.shape[-1]
|
||||
)[:, 0]
|
||||
time_context = time_context_first_timestep[None, :].broadcast_to(
|
||||
height * width, batch_size, 1, time_context.shape[-1]
|
||||
)
|
||||
time_context = time_context.reshape(height * width * batch_size, 1, time_context.shape[-1])
|
||||
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
inner_dim = hidden_states.shape[1]
|
||||
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch_frames, height * width, inner_dim)
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(self.proj_in, hidden_states)
|
||||
|
||||
num_frames_emb = torch.arange(num_frames, device=hidden_states.device)
|
||||
num_frames_emb = num_frames_emb.repeat(batch_size, 1)
|
||||
num_frames_emb = num_frames_emb.reshape(-1)
|
||||
t_emb = self.time_proj(num_frames_emb)
|
||||
|
||||
# `Timesteps` does not contain any weights and will always return f32 tensors
|
||||
# but time_embedding might actually be running in fp16. so we need to cast here.
|
||||
# there might be better ways to encapsulate this.
|
||||
t_emb = t_emb.to(dtype=hidden_states.dtype)
|
||||
|
||||
emb = self.time_pos_embed(t_emb)
|
||||
emb = emb[:, None, :]
|
||||
|
||||
# 2. Blocks
|
||||
for block, temporal_block in zip(self.transformer_blocks, self.temporal_transformer_blocks):
|
||||
if self.gradient_checkpointing:
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(
|
||||
block,
|
||||
hidden_states,
|
||||
None,
|
||||
encoder_hidden_states,
|
||||
None,
|
||||
use_reentrant=False,
|
||||
)
|
||||
else:
|
||||
hidden_states = block(
|
||||
hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
)
|
||||
|
||||
hidden_states_mix = hidden_states
|
||||
hidden_states_mix = hidden_states_mix + emb
|
||||
|
||||
if self.gradient_checkpointing:
|
||||
hidden_states_mix = torch.utils.checkpoint.checkpoint(
|
||||
temporal_block,
|
||||
hidden_states_mix,
|
||||
num_frames,
|
||||
time_context,
|
||||
)
|
||||
hidden_states = self.time_mixer(
|
||||
x_spatial=hidden_states,
|
||||
x_temporal=hidden_states_mix,
|
||||
image_only_indicator=image_only_indicator,
|
||||
)
|
||||
else:
|
||||
hidden_states_mix = temporal_block(
|
||||
hidden_states_mix,
|
||||
num_frames=num_frames,
|
||||
encoder_hidden_states=time_context,
|
||||
)
|
||||
hidden_states = self.time_mixer(
|
||||
x_spatial=hidden_states,
|
||||
x_temporal=hidden_states_mix,
|
||||
image_only_indicator=image_only_indicator,
|
||||
)
|
||||
|
||||
# 3. Output
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(self.proj_out, hidden_states)
|
||||
hidden_states = hidden_states.reshape(batch_frames, height, width, inner_dim).permute(0, 3, 1, 2).contiguous()
|
||||
|
||||
output = hidden_states + residual
|
||||
|
||||
if not return_dict:
|
||||
return (output,)
|
||||
|
||||
return TransformerTemporalModelOutput(sample=output)
|
||||
@@ -0,0 +1,278 @@
|
||||
from time import process_time_ns
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from diffusers.models.lora import LoRALinearLayer
|
||||
from diffusers.utils.import_utils import is_xformers_available
|
||||
if is_xformers_available():
|
||||
import xformers
|
||||
else:
|
||||
print(1/0)
|
||||
|
||||
class AnimationAttnProcessor(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size=None,
|
||||
cross_attention_dim=None,
|
||||
rank=4,
|
||||
network_alpha=None,
|
||||
lora_scale=1.0,):
|
||||
super().__init__()
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
# self.rank = rank
|
||||
# self.lora_scale = lora_scale
|
||||
#
|
||||
# self.to_q_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha)
|
||||
# self.to_k_lora = LoRALinearLayer(cross_attention_dim or hidden_size, hidden_size, rank, network_alpha)
|
||||
# self.to_v_lora = LoRALinearLayer(cross_attention_dim or hidden_size, hidden_size, rank, network_alpha)
|
||||
# self.to_out_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha)
|
||||
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
):
|
||||
|
||||
# hidden_states = hidden_states.to(dtype=torch.float16)
|
||||
|
||||
residual = hidden_states
|
||||
|
||||
# print("-----------------------------")
|
||||
# print("This is AnimationAttnProcessor")
|
||||
# print(hidden_states.dtype)
|
||||
# print("-----------------------------")
|
||||
|
||||
if attn.spatial_norm is not None:
|
||||
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
|
||||
if attention_mask is not None:
|
||||
_, query_tokens, _ = hidden_states.shape
|
||||
attention_mask = attention_mask.expand(-1, query_tokens, -1)
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
# query = attn.to_q(hidden_states) + self.lora_scale * self.to_q_lora(hidden_states)
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
elif attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
# key = attn.to_k(encoder_hidden_states) + self.lora_scale * self.to_k_lora(encoder_hidden_states)
|
||||
# value = attn.to_v(encoder_hidden_states) + self.lora_scale * self.to_v_lora(encoder_hidden_states)
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
query = attn.head_to_batch_dim(query).contiguous()
|
||||
key = attn.head_to_batch_dim(key).contiguous()
|
||||
value = attn.head_to_batch_dim(value).contiguous()
|
||||
|
||||
if is_xformers_available():
|
||||
### xformers
|
||||
hidden_states = xformers.ops.memory_efficient_attention(query, key, value, attn_bias=attention_mask)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
else:
|
||||
attention_probs = attn.get_attention_scores(query, key, attention_mask)
|
||||
hidden_states = torch.bmm(attention_probs, value)
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# linear proj
|
||||
# hidden_states = attn.to_out[0](hidden_states) + self.lora_scale * self.to_out_lora(hidden_states)
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class AnimationIDAttnProcessor(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
cross_attention_dim=None,
|
||||
rank=4,
|
||||
network_alpha=None,
|
||||
lora_scale=1.0,
|
||||
scale=1.0,
|
||||
num_tokens=4):
|
||||
super().__init__()
|
||||
|
||||
# self.to_q_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha)
|
||||
# self.to_k_lora = LoRALinearLayer(cross_attention_dim or hidden_size, hidden_size, rank, network_alpha)
|
||||
# self.to_v_lora = LoRALinearLayer(cross_attention_dim or hidden_size, hidden_size, rank, network_alpha)
|
||||
# self.to_out_lora = LoRALinearLayer(hidden_size, hidden_size, rank, network_alpha)
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.scale = scale
|
||||
|
||||
self.id_to_k = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False)
|
||||
self.id_to_v = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False)
|
||||
|
||||
self.lora_scale = lora_scale
|
||||
self.num_tokens = num_tokens
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
scale=1.0,
|
||||
):
|
||||
|
||||
# hidden_states = hidden_states.to(encoder_hidden_states.dtype)
|
||||
|
||||
residual = hidden_states
|
||||
|
||||
# print("-----------------------------")
|
||||
# print("This is AnimationIDAttnProcessor")
|
||||
# print(hidden_states.dtype)
|
||||
# print("-----------------------------")
|
||||
|
||||
if attn.spatial_norm is not None:
|
||||
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
# query = attn.to_q(hidden_states) + self.lora_scale * self.to_q_lora(hidden_states)
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
|
||||
# print(attn.heads) # 5
|
||||
# print(batch_size) # 21
|
||||
# print(encoder_hidden_states.size()) # [21, 5, 1024]
|
||||
# print(self.num_tokens) # 4
|
||||
|
||||
encoder_hidden_states = encoder_hidden_states.to(hidden_states.dtype)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
else:
|
||||
end_pos = encoder_hidden_states.shape[1] - self.num_tokens
|
||||
encoder_hidden_states, ip_hidden_states = (
|
||||
encoder_hidden_states[:, :end_pos, :],
|
||||
encoder_hidden_states[:, end_pos:, :],
|
||||
)
|
||||
if attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
# key = attn.to_k(encoder_hidden_states) + self.lora_scale * self.to_k_lora(encoder_hidden_states)
|
||||
# value = attn.to_v(encoder_hidden_states) + self.lora_scale * self.to_v_lora(encoder_hidden_states)
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
# print(query.size()) # [21, 4096, 320]
|
||||
# print(key.size()) # [21, 1, 320]
|
||||
# print(value.size()) # [21, 1, 320]
|
||||
|
||||
inner_dim = key.shape[-1]
|
||||
head_dim = inner_dim // attn.heads
|
||||
|
||||
# query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
# key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
# value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
query = attn.head_to_batch_dim(query).contiguous()
|
||||
key = attn.head_to_batch_dim(key).contiguous()
|
||||
value = attn.head_to_batch_dim(value).contiguous()
|
||||
|
||||
key = key.to(query.dtype)
|
||||
value = value.to(query.dtype)
|
||||
|
||||
|
||||
if is_xformers_available():
|
||||
### xformers
|
||||
hidden_states = xformers.ops.memory_efficient_attention(query, key, value, attn_bias=attention_mask)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
else:
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# print("==========================This is AnimationIDAttnProcessor==========================")
|
||||
# print(hidden_states.size()) # [21, 4096, 320]
|
||||
|
||||
ip_key = self.id_to_k(ip_hidden_states)
|
||||
ip_value = self.id_to_v(ip_hidden_states)
|
||||
|
||||
# ip_key = ip_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
# ip_value = ip_value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
|
||||
|
||||
ip_key = attn.head_to_batch_dim(ip_key).contiguous()
|
||||
ip_value = attn.head_to_batch_dim(ip_value).contiguous()
|
||||
ip_key = ip_key.to(query.dtype)
|
||||
ip_value = ip_value.to(query.dtype)
|
||||
|
||||
if is_xformers_available():
|
||||
### xformers
|
||||
ip_hidden_states = xformers.ops.memory_efficient_attention(query, ip_key, ip_value, attn_bias=None)
|
||||
ip_hidden_states = ip_hidden_states.to(query.dtype)
|
||||
else:
|
||||
ip_hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False)
|
||||
ip_hidden_states = ip_hidden_states.to(query.dtype)
|
||||
|
||||
# print(ip_hidden_states.size()) # [105, 4096, 64]
|
||||
# ip_hidden_states = ip_hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||
ip_hidden_states = attn.batch_to_head_dim(ip_hidden_states)
|
||||
# print(ip_hidden_states.size()) # [21, 4096, 320]
|
||||
hidden_states = hidden_states + self.scale * ip_hidden_states
|
||||
# linear proj
|
||||
# hidden_states = attn.to_out[0](hidden_states) + self.lora_scale * self.to_out_lora(hidden_states)
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
@@ -0,0 +1,152 @@
|
||||
from time import process_time_ns
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from diffusers.models.lora import LoRALinearLayer
|
||||
from diffusers.utils.import_utils import is_xformers_available
|
||||
|
||||
if is_xformers_available():
|
||||
import xformers
|
||||
else:
|
||||
print(1 / 0)
|
||||
|
||||
class AnimationIDAttnNormalizedProcessor(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
cross_attention_dim=None,
|
||||
rank=4,
|
||||
network_alpha=None,
|
||||
lora_scale=1.0,
|
||||
scale=1.0,
|
||||
num_tokens=4):
|
||||
super().__init__()
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.scale = scale
|
||||
|
||||
self.id_to_k = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False)
|
||||
self.id_to_v = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False)
|
||||
|
||||
self.lora_scale = lora_scale
|
||||
self.num_tokens = num_tokens
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
scale=1.0,
|
||||
):
|
||||
|
||||
# hidden_states = hidden_states.to(encoder_hidden_states.dtype)
|
||||
|
||||
residual = hidden_states
|
||||
|
||||
if attn.spatial_norm is not None:
|
||||
hidden_states = attn.spatial_norm(hidden_states, temb)
|
||||
|
||||
input_ndim = hidden_states.ndim
|
||||
|
||||
if input_ndim == 4:
|
||||
batch_size, channel, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
|
||||
|
||||
batch_size, sequence_length, _ = (
|
||||
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
|
||||
)
|
||||
|
||||
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
|
||||
|
||||
if attn.group_norm is not None:
|
||||
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
|
||||
# print(attn.heads) # 5
|
||||
# print(batch_size) # 21
|
||||
# print(encoder_hidden_states.size()) # [21, 5, 1024]
|
||||
# print(self.num_tokens) # 4
|
||||
|
||||
encoder_hidden_states = encoder_hidden_states.to(hidden_states.dtype)
|
||||
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
else:
|
||||
end_pos = encoder_hidden_states.shape[1] - self.num_tokens
|
||||
encoder_hidden_states, ip_hidden_states = (
|
||||
encoder_hidden_states[:, :end_pos, :],
|
||||
encoder_hidden_states[:, end_pos:, :],
|
||||
)
|
||||
if attn.norm_cross:
|
||||
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
|
||||
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
# print(query.size()) # [21, 4096, 320]
|
||||
# print(key.size()) # [21, 1, 320]
|
||||
# print(value.size()) # [21, 1, 320]
|
||||
|
||||
inner_dim = key.shape[-1]
|
||||
head_dim = inner_dim // attn.heads
|
||||
|
||||
query = attn.head_to_batch_dim(query).contiguous()
|
||||
key = attn.head_to_batch_dim(key).contiguous()
|
||||
value = attn.head_to_batch_dim(value).contiguous()
|
||||
|
||||
key = key.to(query.dtype)
|
||||
value = value.to(query.dtype)
|
||||
|
||||
if is_xformers_available():
|
||||
### xformers
|
||||
hidden_states = xformers.ops.memory_efficient_attention(query, key, value, attn_bias=attention_mask)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
else:
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask=attention_mask, dropout_p=0.0,
|
||||
is_causal=False)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
hidden_states = attn.batch_to_head_dim(hidden_states)
|
||||
|
||||
# print("==========================This is AnimationIDAttnProcessor==========================")
|
||||
# print(hidden_states.size()) # [21, 4096, 320]
|
||||
|
||||
ip_key = self.id_to_k(ip_hidden_states)
|
||||
ip_value = self.id_to_v(ip_hidden_states)
|
||||
|
||||
ip_key = attn.head_to_batch_dim(ip_key).contiguous()
|
||||
ip_value = attn.head_to_batch_dim(ip_value).contiguous()
|
||||
ip_key = ip_key.to(query.dtype)
|
||||
ip_value = ip_value.to(query.dtype)
|
||||
|
||||
if is_xformers_available():
|
||||
### xformers
|
||||
ip_hidden_states = xformers.ops.memory_efficient_attention(query, ip_key, ip_value, attn_bias=None)
|
||||
ip_hidden_states = ip_hidden_states.to(query.dtype)
|
||||
else:
|
||||
ip_hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask=None, dropout_p=0.0,
|
||||
is_causal=False)
|
||||
ip_hidden_states = ip_hidden_states.to(query.dtype)
|
||||
|
||||
ip_hidden_states = attn.batch_to_head_dim(ip_hidden_states)
|
||||
mean_latents, std_latents = torch.mean(hidden_states, dim=(1, 2), keepdim=True), torch.std(hidden_states, dim=(1, 2), keepdim=True)
|
||||
mean_ip, std_ip = torch.mean(ip_hidden_states, dim=(1, 2), keepdim=True), torch.std(ip_hidden_states, dim=(1, 2), keepdim=True)
|
||||
ip_hidden_states = (ip_hidden_states - mean_ip) * (std_latents / (std_ip + 1e-5)) + mean_latents
|
||||
hidden_states = hidden_states + self.scale * ip_hidden_states
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if input_ndim == 4:
|
||||
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
|
||||
|
||||
if attn.residual_connection:
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states = hidden_states / attn.rescale_output_factor
|
||||
|
||||
return hidden_states
|
||||
@@ -0,0 +1,26 @@
|
||||
from torch import nn
|
||||
from facexlib.parsing import init_parsing_model
|
||||
from facexlib.utils.face_restoration_helper import FaceRestoreHelper
|
||||
import insightface
|
||||
from insightface.app import FaceAnalysis
|
||||
|
||||
|
||||
class FaceModel(nn.Module):
|
||||
def __init__(self):
|
||||
super(FaceModel, self).__init__()
|
||||
self.app = FaceAnalysis(
|
||||
name='antelopev2', root='.', providers=['CUDAExecutionProvider', 'CPUExecutionProvider', ]
|
||||
)
|
||||
self.app.prepare(ctx_id=-1, det_size=(640, 640))
|
||||
self.handler_ante = insightface.model_zoo.get_model('models/antelopev2/glintr100.onnx')
|
||||
self.handler_ante.prepare(ctx_id=-1)
|
||||
|
||||
self.face_helper = FaceRestoreHelper(
|
||||
upscale_factor=1,
|
||||
face_size=512,
|
||||
crop_ratio=(1, 1),
|
||||
det_model='retinaface_resnet50',
|
||||
save_ext='png',
|
||||
device="cpu",
|
||||
)
|
||||
self.face_helper.face_parse = init_parsing_model(model_name='bisenet', device="cpu")
|
||||
@@ -0,0 +1,141 @@
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
|
||||
def reshape_tensor(x, heads):
|
||||
bs, length, width = x.shape
|
||||
x = x.view(bs, length, heads, -1)
|
||||
x = x.transpose(1, 2)
|
||||
x = x.reshape(bs, heads, length, -1)
|
||||
return x
|
||||
|
||||
class PerceiverAttention(nn.Module):
|
||||
def __init__(self, *, dim, dim_head=64, heads=8):
|
||||
super().__init__()
|
||||
self.scale = dim_head**-0.5
|
||||
self.dim_head = dim_head
|
||||
self.heads = heads
|
||||
inner_dim = dim_head * heads
|
||||
|
||||
self.norm1 = nn.LayerNorm(dim)
|
||||
self.norm2 = nn.LayerNorm(dim)
|
||||
|
||||
self.to_q = nn.Linear(dim, inner_dim, bias=False)
|
||||
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
|
||||
self.to_out = nn.Linear(inner_dim, dim, bias=False)
|
||||
|
||||
def forward(self, x, latents):
|
||||
"""
|
||||
Args:
|
||||
x (torch.Tensor): image features
|
||||
shape (b, n1, D)
|
||||
latent (torch.Tensor): latent features
|
||||
shape (b, n2, D)
|
||||
"""
|
||||
|
||||
x = self.norm1(x)
|
||||
latents = self.norm2(latents)
|
||||
|
||||
b, l, _ = latents.shape
|
||||
|
||||
q = self.to_q(latents)
|
||||
kv_input = torch.cat((x, latents), dim=-2)
|
||||
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
|
||||
|
||||
q = reshape_tensor(q, self.heads)
|
||||
k = reshape_tensor(k, self.heads)
|
||||
v = reshape_tensor(v, self.heads)
|
||||
|
||||
# attention
|
||||
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
|
||||
weight = (q * scale) @ (k * scale).transpose(-2, -1)
|
||||
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
|
||||
out = weight @ v
|
||||
|
||||
out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
def FeedForward(dim, mult=4):
|
||||
inner_dim = int(dim * mult)
|
||||
return nn.Sequential(
|
||||
nn.LayerNorm(dim),
|
||||
nn.Linear(dim, inner_dim, bias=False),
|
||||
nn.GELU(),
|
||||
nn.Linear(inner_dim, dim, bias=False),
|
||||
)
|
||||
|
||||
class FacePerceiver(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim=768,
|
||||
depth=4,
|
||||
dim_head=64,
|
||||
heads=16,
|
||||
embedding_dim=1280,
|
||||
output_dim=768,
|
||||
ff_mult=4,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.proj_in = torch.nn.Linear(embedding_dim, dim)
|
||||
self.proj_out = torch.nn.Linear(dim, output_dim)
|
||||
self.norm_out = torch.nn.LayerNorm(output_dim)
|
||||
self.layers = torch.nn.ModuleList([])
|
||||
for _ in range(depth):
|
||||
self.layers.append(
|
||||
torch.nn.ModuleList(
|
||||
[
|
||||
PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads),
|
||||
FeedForward(dim=dim, mult=ff_mult),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
nn.init.constant_(self.proj_out.weight, 0)
|
||||
if self.proj_out.bias is not None:
|
||||
nn.init.constant_(self.proj_out.bias, 0)
|
||||
|
||||
def forward(self, latents, x):
|
||||
x = self.proj_in(x)
|
||||
for attn, ff in self.layers:
|
||||
latents = attn(x, latents) + latents
|
||||
latents = ff(latents) + latents
|
||||
latents = self.proj_out(latents)
|
||||
return self.norm_out(latents)
|
||||
|
||||
|
||||
class FusionFaceId(ModelMixin):
|
||||
def __init__(self, cross_attention_dim=768, id_embeddings_dim=512, clip_embeddings_dim=1024, num_tokens=4):
|
||||
super().__init__()
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.num_tokens = num_tokens
|
||||
|
||||
self.proj = torch.nn.Sequential(
|
||||
torch.nn.Linear(id_embeddings_dim, id_embeddings_dim*2),
|
||||
torch.nn.GELU(),
|
||||
torch.nn.Linear(id_embeddings_dim*2, cross_attention_dim*num_tokens),
|
||||
)
|
||||
|
||||
self.norm = torch.nn.LayerNorm(cross_attention_dim)
|
||||
|
||||
self.fusion_model = FacePerceiver(
|
||||
dim=cross_attention_dim,
|
||||
depth=4,
|
||||
dim_head=64,
|
||||
heads=cross_attention_dim // 64,
|
||||
embedding_dim=clip_embeddings_dim,
|
||||
output_dim=cross_attention_dim,
|
||||
ff_mult=4,
|
||||
)
|
||||
|
||||
|
||||
def forward(self, id_embeds, clip_embeds, shortcut=False, scale=1.0):
|
||||
x = self.proj(id_embeds)
|
||||
x = x.reshape(-1, self.num_tokens, self.cross_attention_dim)
|
||||
x = self.norm(x)
|
||||
out = self.fusion_model(x, clip_embeds)
|
||||
if shortcut:
|
||||
out = x + scale * out
|
||||
return out
|
||||
@@ -0,0 +1,78 @@
|
||||
from pathlib import Path
|
||||
|
||||
import einops
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.init as init
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
|
||||
class PoseNet(ModelMixin):
|
||||
def __init__(self, noise_latent_channels=320):
|
||||
super().__init__()
|
||||
# multiple convolution layers
|
||||
self.conv_layers = nn.Sequential(
|
||||
nn.Conv2d(in_channels=3, out_channels=3, kernel_size=3, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(in_channels=3, out_channels=16, kernel_size=4, stride=2, padding=1),
|
||||
nn.SiLU(),
|
||||
|
||||
nn.Conv2d(in_channels=16, out_channels=16, kernel_size=3, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(in_channels=16, out_channels=32, kernel_size=4, stride=2, padding=1),
|
||||
nn.SiLU(),
|
||||
|
||||
nn.Conv2d(in_channels=32, out_channels=32, kernel_size=3, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(in_channels=32, out_channels=64, kernel_size=4, stride=2, padding=1),
|
||||
nn.SiLU(),
|
||||
|
||||
nn.Conv2d(in_channels=64, out_channels=64, kernel_size=3, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(in_channels=64, out_channels=128, kernel_size=3, stride=1, padding=1),
|
||||
nn.SiLU()
|
||||
)
|
||||
|
||||
# Final projection layer
|
||||
self.final_proj = nn.Conv2d(in_channels=128, out_channels=noise_latent_channels, kernel_size=1)
|
||||
|
||||
# Initialize layers
|
||||
self._initialize_weights()
|
||||
|
||||
self.scale = nn.Parameter(torch.ones(1) * 2)
|
||||
|
||||
def _initialize_weights(self):
|
||||
"""Initialize weights with He. initialization and zero out the biases
|
||||
"""
|
||||
for m in self.conv_layers:
|
||||
if isinstance(m, nn.Conv2d):
|
||||
n = m.kernel_size[0] * m.kernel_size[1] * m.in_channels
|
||||
init.normal_(m.weight, mean=0.0, std=np.sqrt(2. / n))
|
||||
if m.bias is not None:
|
||||
init.zeros_(m.bias)
|
||||
init.zeros_(self.final_proj.weight)
|
||||
if self.final_proj.bias is not None:
|
||||
init.zeros_(self.final_proj.bias)
|
||||
|
||||
def forward(self, x):
|
||||
if x.ndim == 5:
|
||||
x = einops.rearrange(x, "b f c h w -> (b f) c h w")
|
||||
x = self.conv_layers(x)
|
||||
x = self.final_proj(x)
|
||||
|
||||
return x * self.scale
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, pretrained_model_path):
|
||||
"""load pretrained pose-net weights
|
||||
"""
|
||||
if not Path(pretrained_model_path).exists():
|
||||
print(f"There is no model file in {pretrained_model_path}")
|
||||
print(f"loaded PoseNet's pretrained weights from {pretrained_model_path}.")
|
||||
|
||||
state_dict = torch.load(pretrained_model_path, map_location="cpu")
|
||||
model = PoseNet(noise_latent_channels=320)
|
||||
|
||||
model.load_state_dict(state_dict, strict=True)
|
||||
|
||||
return model
|
||||
@@ -0,0 +1,390 @@
|
||||
from typing import Dict, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.utils import is_torch_version
|
||||
from diffusers.utils.accelerate_utils import apply_forward_hook
|
||||
from diffusers.models.attention_processor import CROSS_ATTENTION_PROCESSORS, AttentionProcessor, AttnProcessor
|
||||
from diffusers.models.modeling_outputs import AutoencoderKLOutput
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.models.unets.unet_3d_blocks import MidBlockTemporalDecoder, UpBlockTemporalDecoder
|
||||
from diffusers.models.autoencoders.vae import DecoderOutput, DiagonalGaussianDistribution, Encoder
|
||||
|
||||
|
||||
|
||||
class TemporalDecoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 4,
|
||||
out_channels: int = 3,
|
||||
block_out_channels: Tuple[int] = (128, 256, 512, 512),
|
||||
layers_per_block: int = 2,
|
||||
):
|
||||
super().__init__()
|
||||
self.layers_per_block = layers_per_block
|
||||
|
||||
self.conv_in = nn.Conv2d(in_channels, block_out_channels[-1], kernel_size=3, stride=1, padding=1)
|
||||
self.mid_block = MidBlockTemporalDecoder(
|
||||
num_layers=self.layers_per_block,
|
||||
in_channels=block_out_channels[-1],
|
||||
out_channels=block_out_channels[-1],
|
||||
attention_head_dim=block_out_channels[-1],
|
||||
)
|
||||
|
||||
# up
|
||||
self.up_blocks = nn.ModuleList([])
|
||||
reversed_block_out_channels = list(reversed(block_out_channels))
|
||||
output_channel = reversed_block_out_channels[0]
|
||||
for i in range(len(block_out_channels)):
|
||||
prev_output_channel = output_channel
|
||||
output_channel = reversed_block_out_channels[i]
|
||||
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
up_block = UpBlockTemporalDecoder(
|
||||
num_layers=self.layers_per_block + 1,
|
||||
in_channels=prev_output_channel,
|
||||
out_channels=output_channel,
|
||||
add_upsample=not is_final_block,
|
||||
)
|
||||
self.up_blocks.append(up_block)
|
||||
prev_output_channel = output_channel
|
||||
|
||||
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=32, eps=1e-6)
|
||||
|
||||
self.conv_act = nn.SiLU()
|
||||
self.conv_out = torch.nn.Conv2d(
|
||||
in_channels=block_out_channels[0],
|
||||
out_channels=out_channels,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
)
|
||||
|
||||
conv_out_kernel_size = (3, 1, 1)
|
||||
padding = [int(k // 2) for k in conv_out_kernel_size]
|
||||
self.time_conv_out = torch.nn.Conv3d(
|
||||
in_channels=out_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=conv_out_kernel_size,
|
||||
padding=padding,
|
||||
)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
sample: torch.Tensor,
|
||||
image_only_indicator: torch.Tensor,
|
||||
num_frames: int = 1,
|
||||
) -> torch.Tensor:
|
||||
r"""The forward method of the `Decoder` class."""
|
||||
|
||||
sample = self.conv_in(sample)
|
||||
|
||||
upscale_dtype = next(iter(self.up_blocks.parameters())).dtype
|
||||
# if self.training and self.gradient_checkpointing:
|
||||
if self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
return module(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
if is_torch_version(">=", "1.11.0"):
|
||||
# middle
|
||||
sample = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(self.mid_block),
|
||||
sample,
|
||||
image_only_indicator,
|
||||
use_reentrant=False,
|
||||
)
|
||||
sample = sample.to(upscale_dtype)
|
||||
|
||||
# up
|
||||
for up_block in self.up_blocks:
|
||||
sample = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(up_block),
|
||||
sample,
|
||||
image_only_indicator,
|
||||
use_reentrant=False,
|
||||
)
|
||||
else:
|
||||
# middle
|
||||
sample = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(self.mid_block),
|
||||
sample,
|
||||
image_only_indicator,
|
||||
)
|
||||
sample = sample.to(upscale_dtype)
|
||||
|
||||
# up
|
||||
for up_block in self.up_blocks:
|
||||
sample = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(up_block),
|
||||
sample,
|
||||
image_only_indicator,
|
||||
)
|
||||
else:
|
||||
# middle
|
||||
sample = self.mid_block(sample, image_only_indicator=image_only_indicator)
|
||||
sample = sample.to(upscale_dtype)
|
||||
|
||||
# up
|
||||
for up_block in self.up_blocks:
|
||||
sample = up_block(sample, image_only_indicator=image_only_indicator)
|
||||
|
||||
# post-process
|
||||
sample = self.conv_norm_out(sample)
|
||||
sample = self.conv_act(sample)
|
||||
sample = self.conv_out(sample)
|
||||
|
||||
batch_frames, channels, height, width = sample.shape
|
||||
batch_size = batch_frames // num_frames
|
||||
sample = sample[None, :].reshape(batch_size, num_frames, channels, height, width).permute(0, 2, 1, 3, 4)
|
||||
sample = self.time_conv_out(sample)
|
||||
|
||||
sample = sample.permute(0, 2, 1, 3, 4).reshape(batch_frames, channels, height, width)
|
||||
|
||||
return sample
|
||||
|
||||
|
||||
class RefinedAutoencoderKLTemporalDecoder(ModelMixin, ConfigMixin):
|
||||
r"""
|
||||
A VAE model with KL loss for encoding images into latents and decoding latent representations into images.
|
||||
|
||||
This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented
|
||||
for all models (such as downloading or saving).
|
||||
|
||||
Parameters:
|
||||
in_channels (int, *optional*, defaults to 3): Number of channels in the input image.
|
||||
out_channels (int, *optional*, defaults to 3): Number of channels in the output.
|
||||
down_block_types (`Tuple[str]`, *optional*, defaults to `("DownEncoderBlock2D",)`):
|
||||
Tuple of downsample block types.
|
||||
block_out_channels (`Tuple[int]`, *optional*, defaults to `(64,)`):
|
||||
Tuple of block output channels.
|
||||
layers_per_block: (`int`, *optional*, defaults to 1): Number of layers per block.
|
||||
latent_channels (`int`, *optional*, defaults to 4): Number of channels in the latent space.
|
||||
sample_size (`int`, *optional*, defaults to `32`): Sample input size.
|
||||
scaling_factor (`float`, *optional*, defaults to 0.18215):
|
||||
The component-wise standard deviation of the trained latent space computed using the first batch of the
|
||||
training set. This is used to scale the latent space to have unit variance when training the diffusion
|
||||
model. The latents are scaled with the formula `z = z * scaling_factor` before being passed to the
|
||||
diffusion model. When decoding, the latents are scaled back to the original scale with the formula: `z = 1
|
||||
/ scaling_factor * z`. For more details, refer to sections 4.3.2 and D.1 of the [High-Resolution Image
|
||||
Synthesis with Latent Diffusion Models](https://arxiv.org/abs/2112.10752) paper.
|
||||
force_upcast (`bool`, *optional*, default to `True`):
|
||||
If enabled it will force the VAE to run in float32 for high image resolution pipelines, such as SD-XL. VAE
|
||||
can be fine-tuned / trained to a lower range without loosing too much precision in which case
|
||||
`force_upcast` can be set to `False` - see: https://huggingface.co/madebyollin/sdxl-vae-fp16-fix
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 3,
|
||||
out_channels: int = 3,
|
||||
down_block_types: Tuple[str] = ("DownEncoderBlock2D",),
|
||||
block_out_channels: Tuple[int] = (64,),
|
||||
layers_per_block: int = 1,
|
||||
latent_channels: int = 4,
|
||||
sample_size: int = 32,
|
||||
scaling_factor: float = 0.18215,
|
||||
force_upcast: float = True,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# pass init params to Encoder
|
||||
self.encoder = Encoder(
|
||||
in_channels=in_channels,
|
||||
out_channels=latent_channels,
|
||||
down_block_types=down_block_types,
|
||||
block_out_channels=block_out_channels,
|
||||
layers_per_block=layers_per_block,
|
||||
double_z=True,
|
||||
)
|
||||
|
||||
# pass init params to Decoder
|
||||
self.decoder = TemporalDecoder(
|
||||
in_channels=latent_channels,
|
||||
out_channels=out_channels,
|
||||
block_out_channels=block_out_channels,
|
||||
layers_per_block=layers_per_block,
|
||||
)
|
||||
|
||||
self.quant_conv = nn.Conv2d(2 * latent_channels, 2 * latent_channels, 1)
|
||||
|
||||
sample_size = (
|
||||
self.config.sample_size[0]
|
||||
if isinstance(self.config.sample_size, (list, tuple))
|
||||
else self.config.sample_size
|
||||
)
|
||||
self.tile_latent_min_size = int(sample_size / (2 ** (len(self.config.block_out_channels) - 1)))
|
||||
self.tile_overlap_factor = 0.25
|
||||
|
||||
def _set_gradient_checkpointing(self, module, value=False):
|
||||
if isinstance(module, (Encoder, TemporalDecoder)):
|
||||
module.gradient_checkpointing = value
|
||||
|
||||
@property
|
||||
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.attn_processors
|
||||
def attn_processors(self) -> Dict[str, AttentionProcessor]:
|
||||
r"""
|
||||
Returns:
|
||||
`dict` of attention processors: A dictionary containing all attention processors used in the model with
|
||||
indexed by its weight name.
|
||||
"""
|
||||
# set recursively
|
||||
processors = {}
|
||||
|
||||
def fn_recursive_add_processors(name: str, module: torch.nn.Module, processors: Dict[str, AttentionProcessor]):
|
||||
if hasattr(module, "get_processor"):
|
||||
processors[f"{name}.processor"] = module.get_processor()
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
|
||||
|
||||
return processors
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_add_processors(name, module, processors)
|
||||
|
||||
return processors
|
||||
|
||||
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attn_processor
|
||||
def set_attn_processor(self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]]):
|
||||
r"""
|
||||
Sets the attention processor to use to compute attention.
|
||||
|
||||
Parameters:
|
||||
processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
|
||||
The instantiated processor class or a dictionary of processor classes that will be set as the processor
|
||||
for **all** `Attention` layers.
|
||||
|
||||
If `processor` is a dict, the key needs to define the path to the corresponding cross attention
|
||||
processor. This is strongly recommended when setting trainable attention processors.
|
||||
|
||||
"""
|
||||
count = len(self.attn_processors.keys())
|
||||
|
||||
if isinstance(processor, dict) and len(processor) != count:
|
||||
raise ValueError(
|
||||
f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
|
||||
f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
|
||||
)
|
||||
|
||||
def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
|
||||
if hasattr(module, "set_processor"):
|
||||
if not isinstance(processor, dict):
|
||||
module.set_processor(processor)
|
||||
else:
|
||||
module.set_processor(processor.pop(f"{name}.processor"))
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_attn_processor(name, module, processor)
|
||||
|
||||
def set_default_attn_processor(self):
|
||||
"""
|
||||
Disables custom attention processors and sets the default attention implementation.
|
||||
"""
|
||||
if all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()):
|
||||
processor = AttnProcessor()
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}"
|
||||
)
|
||||
|
||||
self.set_attn_processor(processor)
|
||||
|
||||
@apply_forward_hook
|
||||
def encode(
|
||||
self, x: torch.Tensor, return_dict: bool = True
|
||||
) -> Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]:
|
||||
"""
|
||||
Encode a batch of images into latents.
|
||||
|
||||
Args:
|
||||
x (`torch.Tensor`): Input batch of images.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether to return a [`~models.autoencoders.autoencoder_kl.AutoencoderKLOutput`] instead of a plain
|
||||
tuple.
|
||||
|
||||
Returns:
|
||||
The latent representations of the encoded images. If `return_dict` is True, a
|
||||
[`~models.autoencoders.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is
|
||||
returned.
|
||||
"""
|
||||
h = self.encoder(x)
|
||||
moments = self.quant_conv(h)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
|
||||
if not return_dict:
|
||||
return (posterior,)
|
||||
|
||||
return AutoencoderKLOutput(latent_dist=posterior)
|
||||
|
||||
@apply_forward_hook
|
||||
def decode(
|
||||
self,
|
||||
z: torch.Tensor,
|
||||
num_frames: int,
|
||||
return_dict: bool = True,
|
||||
) -> Union[DecoderOutput, torch.Tensor]:
|
||||
"""
|
||||
Decode a batch of images.
|
||||
|
||||
Args:
|
||||
z (`torch.Tensor`): Input batch of latent vectors.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
|
||||
|
||||
Returns:
|
||||
[`~models.vae.DecoderOutput`] or `tuple`:
|
||||
If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
|
||||
returned.
|
||||
|
||||
"""
|
||||
batch_size = z.shape[0] // num_frames
|
||||
image_only_indicator = torch.zeros(batch_size, num_frames, dtype=z.dtype, device=z.device)
|
||||
decoded = self.decoder(z, num_frames=num_frames, image_only_indicator=image_only_indicator)
|
||||
|
||||
if not return_dict:
|
||||
return (decoded,)
|
||||
|
||||
return DecoderOutput(sample=decoded)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
sample: torch.Tensor,
|
||||
sample_posterior: bool = False,
|
||||
return_dict: bool = True,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
num_frames: int = 1,
|
||||
) -> Union[DecoderOutput, torch.Tensor]:
|
||||
r"""
|
||||
Args:
|
||||
sample (`torch.Tensor`): Input sample.
|
||||
sample_posterior (`bool`, *optional*, defaults to `False`):
|
||||
Whether to sample from the posterior.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`DecoderOutput`] instead of a plain tuple.
|
||||
"""
|
||||
x = sample
|
||||
posterior = self.encode(x).latent_dist
|
||||
if sample_posterior:
|
||||
z = posterior.sample(generator=generator)
|
||||
else:
|
||||
z = posterior.mode()
|
||||
|
||||
dec = self.decode(z, num_frames=num_frames).sample
|
||||
|
||||
if not return_dict:
|
||||
return (dec,)
|
||||
|
||||
return DecoderOutput(sample=dec)
|
||||
@@ -0,0 +1,390 @@
|
||||
# Copyright 2024 The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.utils import BaseOutput
|
||||
from diffusers.models.attention import BasicTransformerBlock, TemporalBasicTransformerBlock
|
||||
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.models.resnet import AlphaBlender
|
||||
|
||||
|
||||
@dataclass
|
||||
class TransformerTemporalModelOutput(BaseOutput):
|
||||
"""
|
||||
The output of [`TransformerTemporalModel`].
|
||||
|
||||
Args:
|
||||
sample (`torch.Tensor` of shape `(batch_size x num_frames, num_channels, height, width)`):
|
||||
The hidden states output conditioned on `encoder_hidden_states` input.
|
||||
"""
|
||||
|
||||
sample: torch.Tensor
|
||||
|
||||
|
||||
class TransformerTemporalModel(ModelMixin, ConfigMixin):
|
||||
"""
|
||||
A Transformer model for video-like data.
|
||||
|
||||
Parameters:
|
||||
num_attention_heads (`int`, *optional*, defaults to 16): The number of heads to use for multi-head attention.
|
||||
attention_head_dim (`int`, *optional*, defaults to 88): The number of channels in each head.
|
||||
in_channels (`int`, *optional*):
|
||||
The number of channels in the input and output (specify if the input is **continuous**).
|
||||
num_layers (`int`, *optional*, defaults to 1): The number of layers of Transformer blocks to use.
|
||||
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
|
||||
cross_attention_dim (`int`, *optional*): The number of `encoder_hidden_states` dimensions to use.
|
||||
attention_bias (`bool`, *optional*):
|
||||
Configure if the `TransformerBlock` attention should contain a bias parameter.
|
||||
sample_size (`int`, *optional*): The width of the latent images (specify if the input is **discrete**).
|
||||
This is fixed during training since it is used to learn a number of position embeddings.
|
||||
activation_fn (`str`, *optional*, defaults to `"geglu"`):
|
||||
Activation function to use in feed-forward. See `diffusers.models.activations.get_activation` for supported
|
||||
activation functions.
|
||||
norm_elementwise_affine (`bool`, *optional*):
|
||||
Configure if the `TransformerBlock` should use learnable elementwise affine parameters for normalization.
|
||||
double_self_attention (`bool`, *optional*):
|
||||
Configure if each `TransformerBlock` should contain two self-attention layers.
|
||||
positional_embeddings: (`str`, *optional*):
|
||||
The type of positional embeddings to apply to the sequence input before passing use.
|
||||
num_positional_embeddings: (`int`, *optional*):
|
||||
The maximum length of the sequence over which to apply positional embeddings.
|
||||
"""
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int = 16,
|
||||
attention_head_dim: int = 88,
|
||||
in_channels: Optional[int] = None,
|
||||
out_channels: Optional[int] = None,
|
||||
num_layers: int = 1,
|
||||
dropout: float = 0.0,
|
||||
norm_num_groups: int = 32,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
attention_bias: bool = False,
|
||||
sample_size: Optional[int] = None,
|
||||
activation_fn: str = "geglu",
|
||||
norm_elementwise_affine: bool = True,
|
||||
double_self_attention: bool = True,
|
||||
positional_embeddings: Optional[str] = None,
|
||||
num_positional_embeddings: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.attention_head_dim = attention_head_dim
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
self.in_channels = in_channels
|
||||
|
||||
self.norm = torch.nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True)
|
||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||
|
||||
# 3. Define transformers blocks
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
BasicTransformerBlock(
|
||||
inner_dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
dropout=dropout,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
activation_fn=activation_fn,
|
||||
attention_bias=attention_bias,
|
||||
double_self_attention=double_self_attention,
|
||||
norm_elementwise_affine=norm_elementwise_affine,
|
||||
positional_embeddings=positional_embeddings,
|
||||
num_positional_embeddings=num_positional_embeddings,
|
||||
)
|
||||
for d in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
self.proj_out = nn.Linear(inner_dim, in_channels)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Optional[torch.LongTensor] = None,
|
||||
timestep: Optional[torch.LongTensor] = None,
|
||||
class_labels: torch.LongTensor = None,
|
||||
num_frames: int = 1,
|
||||
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
return_dict: bool = True,
|
||||
) -> TransformerTemporalModelOutput:
|
||||
"""
|
||||
The [`TransformerTemporal`] forward method.
|
||||
|
||||
Args:
|
||||
hidden_states (`torch.LongTensor` of shape `(batch size, num latent pixels)` if discrete, `torch.Tensor` of shape `(batch size, channel, height, width)` if continuous):
|
||||
Input hidden_states.
|
||||
encoder_hidden_states ( `torch.LongTensor` of shape `(batch size, encoder_hidden_states dim)`, *optional*):
|
||||
Conditional embeddings for cross attention layer. If not given, cross-attention defaults to
|
||||
self-attention.
|
||||
timestep ( `torch.LongTensor`, *optional*):
|
||||
Used to indicate denoising step. Optional timestep to be applied as an embedding in `AdaLayerNorm`.
|
||||
class_labels ( `torch.LongTensor` of shape `(batch size, num classes)`, *optional*):
|
||||
Used to indicate class labels conditioning. Optional class labels to be applied as an embedding in
|
||||
`AdaLayerZeroNorm`.
|
||||
num_frames (`int`, *optional*, defaults to 1):
|
||||
The number of frames to be processed per batch. This is used to reshape the hidden states.
|
||||
cross_attention_kwargs (`dict`, *optional*):
|
||||
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||
`self.processor` in
|
||||
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~models.transformers.transformer_temporal.TransformerTemporalModelOutput`]
|
||||
instead of a plain tuple.
|
||||
|
||||
Returns:
|
||||
[`~models.transformers.transformer_temporal.TransformerTemporalModelOutput`] or `tuple`:
|
||||
If `return_dict` is True, an
|
||||
[`~models.transformers.transformer_temporal.TransformerTemporalModelOutput`] is returned, otherwise a
|
||||
`tuple` where the first element is the sample tensor.
|
||||
"""
|
||||
# 1. Input
|
||||
batch_frames, channel, height, width = hidden_states.shape
|
||||
batch_size = batch_frames // num_frames
|
||||
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = hidden_states[None, :].reshape(batch_size, num_frames, channel, height, width)
|
||||
hidden_states = hidden_states.permute(0, 2, 1, 3, 4)
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
hidden_states = hidden_states.permute(0, 3, 4, 2, 1).reshape(batch_size * height * width, num_frames, channel)
|
||||
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
|
||||
# 2. Blocks
|
||||
for block in self.transformer_blocks:
|
||||
hidden_states = block(
|
||||
hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
cross_attention_kwargs=cross_attention_kwargs,
|
||||
class_labels=class_labels,
|
||||
)
|
||||
|
||||
# 3. Output
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
hidden_states = (
|
||||
hidden_states[None, None, :]
|
||||
.reshape(batch_size, height, width, num_frames, channel)
|
||||
.permute(0, 3, 4, 1, 2)
|
||||
.contiguous()
|
||||
)
|
||||
hidden_states = hidden_states.reshape(batch_frames, channel, height, width)
|
||||
|
||||
output = hidden_states + residual
|
||||
|
||||
if not return_dict:
|
||||
return (output,)
|
||||
|
||||
return TransformerTemporalModelOutput(sample=output)
|
||||
|
||||
|
||||
class TransformerSpatioTemporalModel(nn.Module):
|
||||
"""
|
||||
A Transformer model for video-like data.
|
||||
|
||||
Parameters:
|
||||
num_attention_heads (`int`, *optional*, defaults to 16): The number of heads to use for multi-head attention.
|
||||
attention_head_dim (`int`, *optional*, defaults to 88): The number of channels in each head.
|
||||
in_channels (`int`, *optional*):
|
||||
The number of channels in the input and output (specify if the input is **continuous**).
|
||||
out_channels (`int`, *optional*):
|
||||
The number of channels in the output (specify if the input is **continuous**).
|
||||
num_layers (`int`, *optional*, defaults to 1): The number of layers of Transformer blocks to use.
|
||||
cross_attention_dim (`int`, *optional*): The number of `encoder_hidden_states` dimensions to use.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int = 16,
|
||||
attention_head_dim: int = 88,
|
||||
in_channels: int = 320,
|
||||
out_channels: Optional[int] = None,
|
||||
num_layers: int = 1,
|
||||
cross_attention_dim: Optional[int] = None,
|
||||
num_tokens=4,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.attention_head_dim = attention_head_dim
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
self.inner_dim = inner_dim
|
||||
|
||||
# 2. Define input layers
|
||||
self.in_channels = in_channels
|
||||
self.norm = torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6)
|
||||
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||
|
||||
# 3. Define transformers blocks
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
BasicTransformerBlock(
|
||||
inner_dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
)
|
||||
for d in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
time_mix_inner_dim = inner_dim
|
||||
self.temporal_transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
TemporalBasicTransformerBlock(
|
||||
inner_dim,
|
||||
time_mix_inner_dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
cross_attention_dim=cross_attention_dim,
|
||||
)
|
||||
for _ in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
time_embed_dim = in_channels * 4
|
||||
self.time_pos_embed = TimestepEmbedding(in_channels, time_embed_dim, out_dim=in_channels)
|
||||
self.time_proj = Timesteps(in_channels, True, 0)
|
||||
self.time_mixer = AlphaBlender(alpha=0.5, merge_strategy="learned_with_images")
|
||||
|
||||
# 4. Define output layers
|
||||
self.out_channels = in_channels if out_channels is None else out_channels
|
||||
# TODO: should use out_channels for continuous projections
|
||||
self.proj_out = nn.Linear(inner_dim, in_channels)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
self.num_tokens = num_tokens
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
image_only_indicator: Optional[torch.Tensor] = None,
|
||||
return_dict: bool = True,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
hidden_states (`torch.Tensor` of shape `(batch size, channel, height, width)`):
|
||||
Input hidden_states.
|
||||
num_frames (`int`):
|
||||
The number of frames to be processed per batch. This is used to reshape the hidden states.
|
||||
encoder_hidden_states ( `torch.LongTensor` of shape `(batch size, encoder_hidden_states dim)`, *optional*):
|
||||
Conditional embeddings for cross attention layer. If not given, cross-attention defaults to
|
||||
self-attention.
|
||||
image_only_indicator (`torch.LongTensor` of shape `(batch size, num_frames)`, *optional*):
|
||||
A tensor indicating whether the input contains only images. 1 indicates that the input contains only
|
||||
images, 0 indicates that the input contains video frames.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~models.transformers.transformer_temporal.TransformerTemporalModelOutput`]
|
||||
instead of a plain tuple.
|
||||
|
||||
Returns:
|
||||
[`~models.transformers.transformer_temporal.TransformerTemporalModelOutput`] or `tuple`:
|
||||
If `return_dict` is True, an
|
||||
[`~models.transformers.transformer_temporal.TransformerTemporalModelOutput`] is returned, otherwise a
|
||||
`tuple` where the first element is the sample tensor.
|
||||
"""
|
||||
# 1. Input
|
||||
batch_frames, _, height, width = hidden_states.shape
|
||||
num_frames = image_only_indicator.shape[-1]
|
||||
batch_size = batch_frames // num_frames
|
||||
|
||||
# time_context = encoder_hidden_states
|
||||
|
||||
end_pos = encoder_hidden_states.shape[1] - self.num_tokens
|
||||
time_context = encoder_hidden_states[:, :end_pos, :]
|
||||
|
||||
time_context_first_timestep = time_context[None, :].reshape(
|
||||
batch_size, num_frames, -1, time_context.shape[-1]
|
||||
)[:, 0]
|
||||
time_context = time_context_first_timestep[:, None].broadcast_to(
|
||||
batch_size, height * width, time_context.shape[-2], time_context.shape[-1]
|
||||
)
|
||||
time_context = time_context.reshape(batch_size * height * width, -1, time_context.shape[-1])
|
||||
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = self.norm(hidden_states)
|
||||
inner_dim = hidden_states.shape[1]
|
||||
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch_frames, height * width, inner_dim)
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
|
||||
num_frames_emb = torch.arange(num_frames, device=hidden_states.device)
|
||||
num_frames_emb = num_frames_emb.repeat(batch_size, 1)
|
||||
num_frames_emb = num_frames_emb.reshape(-1)
|
||||
t_emb = self.time_proj(num_frames_emb)
|
||||
|
||||
# `Timesteps` does not contain any weights and will always return f32 tensors
|
||||
# but time_embedding might actually be running in fp16. so we need to cast here.
|
||||
# there might be better ways to encapsulate this.
|
||||
t_emb = t_emb.to(dtype=hidden_states.dtype)
|
||||
|
||||
emb = self.time_pos_embed(t_emb)
|
||||
emb = emb[:, None, :]
|
||||
|
||||
# 2. Blocks
|
||||
for block, temporal_block in zip(self.transformer_blocks, self.temporal_transformer_blocks):
|
||||
# print(hidden_states.dtype) # torch.float16
|
||||
# print(encoder_hidden_states.dtype) # torch.float32
|
||||
if self.training and self.gradient_checkpointing:
|
||||
hidden_states = torch.utils.checkpoint.checkpoint(
|
||||
block,
|
||||
hidden_states,
|
||||
None,
|
||||
encoder_hidden_states,
|
||||
None,
|
||||
use_reentrant=False,
|
||||
)
|
||||
else:
|
||||
hidden_states = block(
|
||||
hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
)
|
||||
|
||||
hidden_states_mix = hidden_states
|
||||
hidden_states_mix = hidden_states_mix + emb
|
||||
|
||||
hidden_states_mix = temporal_block(
|
||||
hidden_states_mix,
|
||||
num_frames=num_frames,
|
||||
encoder_hidden_states=time_context,
|
||||
)
|
||||
hidden_states = self.time_mixer(
|
||||
x_spatial=hidden_states,
|
||||
x_temporal=hidden_states_mix,
|
||||
image_only_indicator=image_only_indicator,
|
||||
)
|
||||
|
||||
# 3. Output
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
hidden_states = hidden_states.reshape(batch_frames, height, width, inner_dim).permute(0, 3, 1, 2).contiguous()
|
||||
|
||||
output = hidden_states + residual
|
||||
|
||||
if not return_dict:
|
||||
return (output,)
|
||||
|
||||
return TransformerTemporalModelOutput(sample=output)
|
||||
@@ -0,0 +1,509 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import UNet2DConditionLoadersMixin
|
||||
from diffusers.models.attention_processor import CROSS_ATTENTION_PROCESSORS, AttentionProcessor, AttnProcessor
|
||||
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
|
||||
from animation.modules.unet_3d_blocks import get_down_block, UNetMidBlockSpatioTemporal, get_up_block
|
||||
# from diffusers.models.unets.unet_3d_blocks import get_down_block, get_up_block, UNetMidBlockSpatioTemporal
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
@dataclass
|
||||
class UNetSpatioTemporalConditionOutput(BaseOutput):
|
||||
"""
|
||||
The output of [`UNetSpatioTemporalConditionModel`].
|
||||
|
||||
Args:
|
||||
sample (`torch.FloatTensor` of shape `(batch_size, num_frames, num_channels, height, width)`):
|
||||
The hidden states output conditioned on `encoder_hidden_states` input. Output of last layer of model.
|
||||
"""
|
||||
|
||||
sample: torch.FloatTensor = None
|
||||
|
||||
|
||||
class UNetSpatioTemporalConditionModel(ModelMixin, ConfigMixin, UNet2DConditionLoadersMixin):
|
||||
r"""
|
||||
A conditional Spatio-Temporal UNet model that takes a noisy video frames, conditional state,
|
||||
and a timestep and returns a sample shaped output.
|
||||
|
||||
This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented
|
||||
for all models (such as downloading or saving).
|
||||
|
||||
Parameters:
|
||||
sample_size (`int` or `Tuple[int, int]`, *optional*, defaults to `None`):
|
||||
Height and width of input/output sample.
|
||||
in_channels (`int`, *optional*, defaults to 8): Number of channels in the input sample.
|
||||
out_channels (`int`, *optional*, defaults to 4): Number of channels in the output.
|
||||
down_block_types (`Tuple[str]`, *optional*, defaults to `("CrossAttnDownBlockSpatioTemporal",
|
||||
"CrossAttnDownBlockSpatioTemporal", "CrossAttnDownBlockSpatioTemporal", "DownBlockSpatioTemporal")`):
|
||||
The tuple of downsample blocks to use.
|
||||
up_block_types (`Tuple[str]`, *optional*, defaults to `("UpBlockSpatioTemporal",
|
||||
"CrossAttnUpBlockSpatioTemporal", "CrossAttnUpBlockSpatioTemporal", "CrossAttnUpBlockSpatioTemporal")`):
|
||||
The tuple of upsample blocks to use.
|
||||
block_out_channels (`Tuple[int]`, *optional*, defaults to `(320, 640, 1280, 1280)`):
|
||||
The tuple of output channels for each block.
|
||||
addition_time_embed_dim: (`int`, defaults to 256):
|
||||
Dimension to to encode the additional time ids.
|
||||
projection_class_embeddings_input_dim (`int`, defaults to 768):
|
||||
The dimension of the projection of encoded `added_time_ids`.
|
||||
layers_per_block (`int`, *optional*, defaults to 2): The number of layers per block.
|
||||
cross_attention_dim (`int` or `Tuple[int]`, *optional*, defaults to 1280):
|
||||
The dimension of the cross attention features.
|
||||
transformer_layers_per_block (`int`, `Tuple[int]`, or `Tuple[Tuple]` , *optional*, defaults to 1):
|
||||
The number of transformer blocks of type [`~models.attention.BasicTransformerBlock`]. Only relevant for
|
||||
[`~models.unet_3d_blocks.CrossAttnDownBlockSpatioTemporal`],
|
||||
[`~models.unet_3d_blocks.CrossAttnUpBlockSpatioTemporal`],
|
||||
[`~models.unet_3d_blocks.UNetMidBlockSpatioTemporal`].
|
||||
num_attention_heads (`int`, `Tuple[int]`, defaults to `(5, 10, 10, 20)`):
|
||||
The number of attention heads.
|
||||
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
sample_size: Optional[int] = None,
|
||||
in_channels: int = 8,
|
||||
out_channels: int = 4,
|
||||
down_block_types: Tuple[str] = (
|
||||
"CrossAttnDownBlockSpatioTemporal",
|
||||
"CrossAttnDownBlockSpatioTemporal",
|
||||
"CrossAttnDownBlockSpatioTemporal",
|
||||
"DownBlockSpatioTemporal",
|
||||
),
|
||||
up_block_types: Tuple[str] = (
|
||||
"UpBlockSpatioTemporal",
|
||||
"CrossAttnUpBlockSpatioTemporal",
|
||||
"CrossAttnUpBlockSpatioTemporal",
|
||||
"CrossAttnUpBlockSpatioTemporal",
|
||||
),
|
||||
block_out_channels: Tuple[int] = (320, 640, 1280, 1280),
|
||||
addition_time_embed_dim: int = 256,
|
||||
projection_class_embeddings_input_dim: int = 768,
|
||||
layers_per_block: Union[int, Tuple[int]] = 2,
|
||||
cross_attention_dim: Union[int, Tuple[int]] = 1024,
|
||||
transformer_layers_per_block: Union[int, Tuple[int], Tuple[Tuple]] = 1,
|
||||
num_attention_heads: Union[int, Tuple[int]] = (5, 10, 10, 20),
|
||||
num_frames: int = 25,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.sample_size = sample_size
|
||||
|
||||
# Check inputs
|
||||
if len(down_block_types) != len(up_block_types):
|
||||
raise ValueError(
|
||||
f"Must provide the same number of `down_block_types` as `up_block_types`. " \
|
||||
f"`down_block_types`: {down_block_types}. `up_block_types`: {up_block_types}."
|
||||
)
|
||||
|
||||
if len(block_out_channels) != len(down_block_types):
|
||||
raise ValueError(
|
||||
f"Must provide the same number of `block_out_channels` as `down_block_types`. " \
|
||||
f"`block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}."
|
||||
)
|
||||
|
||||
if not isinstance(num_attention_heads, int) and len(num_attention_heads) != len(down_block_types):
|
||||
raise ValueError(
|
||||
f"Must provide the same number of `num_attention_heads` as `down_block_types`. " \
|
||||
f"`num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}."
|
||||
)
|
||||
|
||||
if isinstance(cross_attention_dim, list) and len(cross_attention_dim) != len(down_block_types):
|
||||
raise ValueError(
|
||||
f"Must provide the same number of `cross_attention_dim` as `down_block_types`. " \
|
||||
f"`cross_attention_dim`: {cross_attention_dim}. `down_block_types`: {down_block_types}."
|
||||
)
|
||||
|
||||
if not isinstance(layers_per_block, int) and len(layers_per_block) != len(down_block_types):
|
||||
raise ValueError(
|
||||
f"Must provide the same number of `layers_per_block` as `down_block_types`. " \
|
||||
f"`layers_per_block`: {layers_per_block}. `down_block_types`: {down_block_types}."
|
||||
)
|
||||
|
||||
# input
|
||||
self.conv_in = nn.Conv2d(
|
||||
in_channels,
|
||||
block_out_channels[0],
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
)
|
||||
|
||||
# time
|
||||
time_embed_dim = block_out_channels[0] * 4
|
||||
|
||||
self.time_proj = Timesteps(block_out_channels[0], True, downscale_freq_shift=0)
|
||||
timestep_input_dim = block_out_channels[0]
|
||||
|
||||
self.time_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim)
|
||||
|
||||
self.add_time_proj = Timesteps(addition_time_embed_dim, True, downscale_freq_shift=0)
|
||||
self.add_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim)
|
||||
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
self.up_blocks = nn.ModuleList([])
|
||||
|
||||
if isinstance(num_attention_heads, int):
|
||||
num_attention_heads = (num_attention_heads,) * len(down_block_types)
|
||||
|
||||
if isinstance(cross_attention_dim, int):
|
||||
cross_attention_dim = (cross_attention_dim,) * len(down_block_types)
|
||||
|
||||
if isinstance(layers_per_block, int):
|
||||
layers_per_block = [layers_per_block] * len(down_block_types)
|
||||
|
||||
if isinstance(transformer_layers_per_block, int):
|
||||
transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types)
|
||||
|
||||
blocks_time_embed_dim = time_embed_dim
|
||||
|
||||
# down
|
||||
output_channel = block_out_channels[0]
|
||||
for i, down_block_type in enumerate(down_block_types):
|
||||
input_channel = output_channel
|
||||
output_channel = block_out_channels[i]
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
|
||||
down_block = get_down_block(
|
||||
down_block_type,
|
||||
num_layers=layers_per_block[i],
|
||||
transformer_layers_per_block=transformer_layers_per_block[i],
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
temb_channels=blocks_time_embed_dim,
|
||||
add_downsample=not is_final_block,
|
||||
resnet_eps=1e-5,
|
||||
cross_attention_dim=cross_attention_dim[i],
|
||||
num_attention_heads=num_attention_heads[i],
|
||||
resnet_act_fn="silu",
|
||||
)
|
||||
self.down_blocks.append(down_block)
|
||||
|
||||
# mid
|
||||
self.mid_block = UNetMidBlockSpatioTemporal(
|
||||
block_out_channels[-1],
|
||||
temb_channels=blocks_time_embed_dim,
|
||||
transformer_layers_per_block=transformer_layers_per_block[-1],
|
||||
cross_attention_dim=cross_attention_dim[-1],
|
||||
num_attention_heads=num_attention_heads[-1],
|
||||
)
|
||||
|
||||
# count how many layers upsample the images
|
||||
self.num_upsamplers = 0
|
||||
|
||||
# up
|
||||
reversed_block_out_channels = list(reversed(block_out_channels))
|
||||
reversed_num_attention_heads = list(reversed(num_attention_heads))
|
||||
reversed_layers_per_block = list(reversed(layers_per_block))
|
||||
reversed_cross_attention_dim = list(reversed(cross_attention_dim))
|
||||
reversed_transformer_layers_per_block = list(reversed(transformer_layers_per_block))
|
||||
|
||||
output_channel = reversed_block_out_channels[0]
|
||||
for i, up_block_type in enumerate(up_block_types):
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
|
||||
prev_output_channel = output_channel
|
||||
output_channel = reversed_block_out_channels[i]
|
||||
input_channel = reversed_block_out_channels[min(i + 1, len(block_out_channels) - 1)]
|
||||
|
||||
# add upsample block for all BUT final layer
|
||||
if not is_final_block:
|
||||
add_upsample = True
|
||||
self.num_upsamplers += 1
|
||||
else:
|
||||
add_upsample = False
|
||||
|
||||
up_block = get_up_block(
|
||||
up_block_type,
|
||||
num_layers=reversed_layers_per_block[i] + 1,
|
||||
transformer_layers_per_block=reversed_transformer_layers_per_block[i],
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
prev_output_channel=prev_output_channel,
|
||||
temb_channels=blocks_time_embed_dim,
|
||||
add_upsample=add_upsample,
|
||||
resnet_eps=1e-5,
|
||||
resolution_idx=i,
|
||||
cross_attention_dim=reversed_cross_attention_dim[i],
|
||||
num_attention_heads=reversed_num_attention_heads[i],
|
||||
resnet_act_fn="silu",
|
||||
)
|
||||
self.up_blocks.append(up_block)
|
||||
prev_output_channel = output_channel
|
||||
|
||||
# out
|
||||
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=32, eps=1e-5)
|
||||
self.conv_act = nn.SiLU()
|
||||
|
||||
self.conv_out = nn.Conv2d(
|
||||
block_out_channels[0],
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
padding=1,
|
||||
)
|
||||
|
||||
@property
|
||||
def attn_processors(self) -> Dict[str, AttentionProcessor]:
|
||||
r"""
|
||||
Returns:
|
||||
`dict` of attention processors: A dictionary containing all attention processors used in the model with
|
||||
indexed by its weight name.
|
||||
"""
|
||||
# set recursively
|
||||
processors = {}
|
||||
|
||||
def fn_recursive_add_processors(
|
||||
name: str,
|
||||
module: torch.nn.Module,
|
||||
processors: Dict[str, AttentionProcessor],
|
||||
):
|
||||
if hasattr(module, "get_processor"):
|
||||
processors[f"{name}.processor"] = module.get_processor(return_deprecated_lora=True)
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
|
||||
|
||||
return processors
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_add_processors(name, module, processors)
|
||||
|
||||
return processors
|
||||
|
||||
def set_attn_processor(self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]]):
|
||||
r"""
|
||||
Sets the attention processor to use to compute attention.
|
||||
|
||||
Parameters:
|
||||
processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
|
||||
The instantiated processor class or a dictionary of processor classes that will be set as the processor
|
||||
for **all** `Attention` layers.
|
||||
|
||||
If `processor` is a dict, the key needs to define the path to the corresponding cross attention
|
||||
processor. This is strongly recommended when setting trainable attention processors.
|
||||
|
||||
"""
|
||||
count = len(self.attn_processors.keys())
|
||||
|
||||
if isinstance(processor, dict) and len(processor) != count:
|
||||
raise ValueError(
|
||||
f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
|
||||
f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
|
||||
)
|
||||
|
||||
def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
|
||||
if hasattr(module, "set_processor"):
|
||||
if not isinstance(processor, dict):
|
||||
module.set_processor(processor)
|
||||
else:
|
||||
module.set_processor(processor.pop(f"{name}.processor"))
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_attn_processor(name, module, processor)
|
||||
|
||||
def set_default_attn_processor(self):
|
||||
"""
|
||||
Disables custom attention processors and sets the default attention implementation.
|
||||
"""
|
||||
if all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()):
|
||||
processor = AttnProcessor()
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Cannot call `set_default_attn_processor` " \
|
||||
f"when attention processors are of type {next(iter(self.attn_processors.values()))}"
|
||||
)
|
||||
|
||||
self.set_attn_processor(processor)
|
||||
|
||||
def _set_gradient_checkpointing(self, module, value=False):
|
||||
if hasattr(module, "gradient_checkpointing"):
|
||||
module.gradient_checkpointing = value
|
||||
|
||||
# Copied from diffusers.models.unets.unet_3d_condition.UNet3DConditionModel.enable_forward_chunking
|
||||
def enable_forward_chunking(self, chunk_size: Optional[int] = None, dim: int = 0) -> None:
|
||||
"""
|
||||
Sets the attention processor to use [feed forward
|
||||
chunking](https://huggingface.co/blog/reformer#2-chunked-feed-forward-layers).
|
||||
|
||||
Parameters:
|
||||
chunk_size (`int`, *optional*):
|
||||
The chunk size of the feed-forward layers. If not specified, will run feed-forward layer individually
|
||||
over each tensor of dim=`dim`.
|
||||
dim (`int`, *optional*, defaults to `0`):
|
||||
The dimension over which the feed-forward computation should be chunked. Choose between dim=0 (batch)
|
||||
or dim=1 (sequence length).
|
||||
"""
|
||||
if dim not in [0, 1]:
|
||||
raise ValueError(f"Make sure to set `dim` to either 0 or 1, not {dim}")
|
||||
|
||||
# By default chunk size is 1
|
||||
chunk_size = chunk_size or 1
|
||||
|
||||
def fn_recursive_feed_forward(module: torch.nn.Module, chunk_size: int, dim: int):
|
||||
if hasattr(module, "set_chunk_feed_forward"):
|
||||
module.set_chunk_feed_forward(chunk_size=chunk_size, dim=dim)
|
||||
|
||||
for child in module.children():
|
||||
fn_recursive_feed_forward(child, chunk_size, dim)
|
||||
|
||||
for module in self.children():
|
||||
fn_recursive_feed_forward(module, chunk_size, dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
sample: torch.FloatTensor,
|
||||
timestep: Union[torch.Tensor, float, int],
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
added_time_ids: torch.Tensor,
|
||||
pose_latents: torch.Tensor = None,
|
||||
image_only_indicator: bool = False,
|
||||
return_dict: bool = True,
|
||||
) -> Union[UNetSpatioTemporalConditionOutput, Tuple]:
|
||||
r"""
|
||||
The [`UNetSpatioTemporalConditionModel`] forward method.
|
||||
|
||||
Args:
|
||||
sample (`torch.FloatTensor`):
|
||||
The noisy input tensor with the following shape `(batch, num_frames, channel, height, width)`.
|
||||
timestep (`torch.FloatTensor` or `float` or `int`): The number of timesteps to denoise an input.
|
||||
encoder_hidden_states (`torch.FloatTensor`):
|
||||
The encoder hidden states with shape `(batch, sequence_length, cross_attention_dim)`.
|
||||
added_time_ids: (`torch.FloatTensor`):
|
||||
The additional time ids with shape `(batch, num_additional_ids)`. These are encoded with sinusoidal
|
||||
embeddings and added to the time embeddings.
|
||||
pose_latents: (`torch.FloatTensor`):
|
||||
The additional latents for pose sequences.
|
||||
image_only_indicator (`bool`, *optional*, defaults to `False`):
|
||||
Whether or not training with all images.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~models.unet_slatio_temporal.UNetSpatioTemporalConditionOutput`]
|
||||
instead of a plain tuple.
|
||||
Returns:
|
||||
[`~models.unet_slatio_temporal.UNetSpatioTemporalConditionOutput`] or `tuple`:
|
||||
If `return_dict` is True,
|
||||
an [`~models.unet_slatio_temporal.UNetSpatioTemporalConditionOutput`] is returned,
|
||||
otherwise a `tuple` is returned where the first element is the sample tensor.
|
||||
"""
|
||||
# 1. time
|
||||
timesteps = timestep
|
||||
if not torch.is_tensor(timesteps):
|
||||
# TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can
|
||||
# This would be a good case for the `match` statement (Python 3.10+)
|
||||
is_mps = sample.device.type == "mps"
|
||||
if isinstance(timestep, float):
|
||||
dtype = torch.float32 if is_mps else torch.float64
|
||||
else:
|
||||
dtype = torch.int32 if is_mps else torch.int64
|
||||
timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device)
|
||||
elif len(timesteps.shape) == 0:
|
||||
timesteps = timesteps[None].to(sample.device)
|
||||
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
batch_size, num_frames = sample.shape[:2]
|
||||
timesteps = timesteps.expand(batch_size)
|
||||
|
||||
t_emb = self.time_proj(timesteps)
|
||||
|
||||
# `Timesteps` does not contain any weights and will always return f32 tensors
|
||||
# but time_embedding might actually be running in fp16. so we need to cast here.
|
||||
# there might be better ways to encapsulate this.
|
||||
t_emb = t_emb.to(dtype=sample.dtype)
|
||||
|
||||
emb = self.time_embedding(t_emb)
|
||||
|
||||
time_embeds = self.add_time_proj(added_time_ids.flatten())
|
||||
time_embeds = time_embeds.reshape((batch_size, -1))
|
||||
time_embeds = time_embeds.to(emb.dtype)
|
||||
aug_emb = self.add_embedding(time_embeds)
|
||||
emb = emb + aug_emb
|
||||
|
||||
# Flatten the batch and frames dimensions
|
||||
# sample: [batch, frames, channels, height, width] -> [batch * frames, channels, height, width]
|
||||
sample = sample.flatten(0, 1)
|
||||
# Repeat the embeddings num_video_frames times
|
||||
# emb: [batch, channels] -> [batch * frames, channels]
|
||||
emb = emb.repeat_interleave(num_frames, dim=0)
|
||||
# encoder_hidden_states: [batch, 1, channels] -> [batch * frames, 1, channels]
|
||||
encoder_hidden_states = encoder_hidden_states.repeat_interleave(num_frames, dim=0)
|
||||
|
||||
# 2. pre-process
|
||||
sample = self.conv_in(sample)
|
||||
if pose_latents is not None:
|
||||
sample = sample + pose_latents
|
||||
|
||||
image_only_indicator = torch.ones(batch_size, num_frames, dtype=sample.dtype, device=sample.device) \
|
||||
if image_only_indicator else torch.zeros(batch_size, num_frames, dtype=sample.dtype, device=sample.device)
|
||||
|
||||
down_block_res_samples = (sample,)
|
||||
for downsample_block in self.down_blocks:
|
||||
if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention:
|
||||
sample, res_samples = downsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
image_only_indicator=image_only_indicator,
|
||||
)
|
||||
else:
|
||||
sample, res_samples = downsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
image_only_indicator=image_only_indicator,
|
||||
)
|
||||
|
||||
down_block_res_samples += res_samples
|
||||
|
||||
# 4. mid
|
||||
sample = self.mid_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
image_only_indicator=image_only_indicator,
|
||||
)
|
||||
|
||||
# 5. up
|
||||
for i, upsample_block in enumerate(self.up_blocks):
|
||||
res_samples = down_block_res_samples[-len(upsample_block.resnets):]
|
||||
down_block_res_samples = down_block_res_samples[: -len(upsample_block.resnets)]
|
||||
|
||||
if hasattr(upsample_block, "has_cross_attention") and upsample_block.has_cross_attention:
|
||||
sample = upsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
res_hidden_states_tuple=res_samples,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
image_only_indicator=image_only_indicator,
|
||||
)
|
||||
else:
|
||||
sample = upsample_block(
|
||||
hidden_states=sample,
|
||||
temb=emb,
|
||||
res_hidden_states_tuple=res_samples,
|
||||
image_only_indicator=image_only_indicator,
|
||||
)
|
||||
|
||||
# 6. post-process
|
||||
sample = self.conv_norm_out(sample)
|
||||
sample = self.conv_act(sample)
|
||||
sample = self.conv_out(sample)
|
||||
|
||||
# 7. Reshape back to original shape
|
||||
sample = sample.reshape(batch_size, num_frames, *sample.shape[1:])
|
||||
|
||||
if not return_dict:
|
||||
return (sample,)
|
||||
|
||||
return UNetSpatioTemporalConditionOutput(sample=sample)
|
||||
@@ -0,0 +1,694 @@
|
||||
import inspect
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Dict, List, Optional, Union
|
||||
|
||||
import PIL.Image
|
||||
import einops
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers.image_processor import VaeImageProcessor, PipelineImageInput
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion import retrieve_timesteps
|
||||
from diffusers.pipelines.stable_video_diffusion.pipeline_stable_video_diffusion \
|
||||
import _resize_with_antialiasing, _append_dims
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
from diffusers.utils.torch_utils import is_compiled_module, randn_tensor
|
||||
|
||||
from animation.modules.attention_processor import AnimationAttnProcessor, AnimationIDAttnProcessor
|
||||
from einops import rearrange
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
def _append_dims(x, target_dims):
|
||||
"""Appends dimensions to the end of a tensor until it has target_dims dimensions."""
|
||||
dims_to_append = target_dims - x.ndim
|
||||
if dims_to_append < 0:
|
||||
raise ValueError(f"input has {x.ndim} dims but target_dims is {target_dims}, which is less")
|
||||
return x[(...,) + (None,) * dims_to_append]
|
||||
|
||||
|
||||
def tensor2vid(video: torch.Tensor, processor: "VaeImageProcessor", output_type: str = "np"):
|
||||
batch_size, channels, num_frames, height, width = video.shape
|
||||
outputs = []
|
||||
for batch_idx in range(batch_size):
|
||||
batch_vid = video[batch_idx].permute(1, 0, 2, 3)
|
||||
batch_output = processor.postprocess(batch_vid, output_type)
|
||||
outputs.append(batch_output)
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
@dataclass
|
||||
class InferenceAnimationPipelineOutput(BaseOutput):
|
||||
r"""
|
||||
Output class for mimicmotion pipeline.
|
||||
|
||||
Args:
|
||||
frames (`[List[List[PIL.Image.Image]]`, `np.ndarray`, `torch.Tensor`]):
|
||||
List of denoised PIL images of length `batch_size` or numpy array or torch tensor of shape `(batch_size,
|
||||
num_frames, height, width, num_channels)`.
|
||||
"""
|
||||
|
||||
frames: Union[List[List[PIL.Image.Image]], np.ndarray, torch.Tensor]
|
||||
|
||||
|
||||
class InferenceAnimationPipeline(DiffusionPipeline):
|
||||
r"""
|
||||
Pipeline to generate video from an input image using Stable Video Diffusion.
|
||||
|
||||
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods
|
||||
implemented for all pipelines (downloading, saving, running on a particular device, etc.).
|
||||
|
||||
Args:
|
||||
vae ([`AutoencoderKLTemporalDecoder`]):
|
||||
Variational Auto-Encoder (VAE) model to encode and decode images to and from latent representations.
|
||||
image_encoder ([`~transformers.CLIPVisionModelWithProjection`]):
|
||||
Frozen CLIP image-encoder ([laion/CLIP-ViT-H-14-laion2B-s32B-b79K]
|
||||
(https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K)).
|
||||
unet ([`UNetSpatioTemporalConditionModel`]):
|
||||
A `UNetSpatioTemporalConditionModel` to denoise the encoded image latents.
|
||||
scheduler ([`EulerDiscreteScheduler`]):
|
||||
A scheduler to be used in combination with `unet` to denoise the encoded image latents.
|
||||
feature_extractor ([`~transformers.CLIPImageProcessor`]):
|
||||
A `CLIPImageProcessor` to extract features from generated images.
|
||||
pose_net ([`PoseNet`]):
|
||||
A `` to inject pose signals into unet.
|
||||
"""
|
||||
|
||||
model_cpu_offload_seq = "image_encoder->unet->vae"
|
||||
_callback_tensor_inputs = ["latents"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vae,
|
||||
image_encoder,
|
||||
unet,
|
||||
scheduler,
|
||||
feature_extractor,
|
||||
pose_net,
|
||||
face_encoder,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
image_encoder=image_encoder,
|
||||
unet=unet,
|
||||
scheduler=scheduler,
|
||||
feature_extractor=feature_extractor,
|
||||
pose_net=pose_net,
|
||||
face_encoder=face_encoder,
|
||||
)
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
|
||||
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
|
||||
|
||||
self.num_tokens = 4
|
||||
|
||||
# self.app = FaceAnalysis(name="buffalo_l", providers=['CUDAExecutionProvider', 'CPUExecutionProvider'])
|
||||
# self.app.prepare(ctx_id=0, det_size=(640, 640))
|
||||
# self.lora_rank = 128
|
||||
# self.set_ip_adapter()
|
||||
|
||||
def get_prepare_faceid(self, face_image):
|
||||
faceid_image = np.array(face_image)
|
||||
faces = self.app.get(faceid_image)
|
||||
if faces == []:
|
||||
faceid_embeds = torch.zeros_like(torch.empty((1, 512)))
|
||||
else:
|
||||
faceid_embeds = torch.from_numpy(faces[0].normed_embedding).unsqueeze(0)
|
||||
return faceid_embeds
|
||||
|
||||
def set_ip_adapter(self):
|
||||
unet = self.unet
|
||||
attn_procs = {}
|
||||
for name in unet.attn_processors.keys():
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") else unet.config.cross_attention_dim
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = unet.config.block_out_channels[-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(unet.config.block_out_channels))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = unet.config.block_out_channels[block_id]
|
||||
if cross_attention_dim is None:
|
||||
attn_procs[name] = AnimationAttnProcessor(
|
||||
hidden_size=hidden_size, cross_attention_dim=cross_attention_dim, rank=self.lora_rank,
|
||||
).to(self.device, dtype=self.torch_dtype)
|
||||
else:
|
||||
attn_procs[name] = AnimationIDAttnProcessor(
|
||||
hidden_size=hidden_size, cross_attention_dim=cross_attention_dim, scale=1.0, rank=self.lora_rank,
|
||||
num_tokens=self.num_tokens,
|
||||
).to(self.device, dtype=self.torch_dtype)
|
||||
|
||||
unet.set_attn_processor(attn_procs)
|
||||
|
||||
def _encode_image(
|
||||
self,
|
||||
image: PipelineImageInput,
|
||||
device: Union[str, torch.device],
|
||||
num_videos_per_prompt: int,
|
||||
do_classifier_free_guidance: bool):
|
||||
dtype = next(self.image_encoder.parameters()).dtype
|
||||
|
||||
if not isinstance(image, torch.Tensor):
|
||||
image = self.image_processor.pil_to_numpy(image)
|
||||
image = self.image_processor.numpy_to_pt(image)
|
||||
|
||||
# We normalize the image before resizing to match with the original implementation.
|
||||
# Then we unnormalize it after resizing.
|
||||
image = image * 2.0 - 1.0
|
||||
image = _resize_with_antialiasing(image, (224, 224))
|
||||
image = (image + 1.0) / 2.0
|
||||
|
||||
# Normalize the image with for CLIP input
|
||||
image = self.feature_extractor(
|
||||
images=image,
|
||||
do_normalize=True,
|
||||
do_center_crop=False,
|
||||
do_resize=False,
|
||||
do_rescale=False,
|
||||
return_tensors="pt",
|
||||
).pixel_values
|
||||
|
||||
image = image.to(device=device, dtype=dtype)
|
||||
image_embeddings = self.image_encoder(image).image_embeds
|
||||
image_embeddings = image_embeddings.unsqueeze(1)
|
||||
|
||||
# duplicate image embeddings for each generation per prompt, using mps friendly method
|
||||
bs_embed, seq_len, _ = image_embeddings.shape
|
||||
image_embeddings = image_embeddings.repeat(1, num_videos_per_prompt, 1)
|
||||
image_embeddings = image_embeddings.view(bs_embed * num_videos_per_prompt, seq_len, -1)
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
negative_image_embeddings = torch.zeros_like(image_embeddings)
|
||||
|
||||
# For classifier free guidance, we need to do two forward passes.
|
||||
# Here we concatenate the unconditional and text embeddings into a single batch
|
||||
# to avoid doing two forward passes
|
||||
image_embeddings = torch.cat([negative_image_embeddings, image_embeddings])
|
||||
|
||||
return image_embeddings
|
||||
|
||||
def _encode_vae_image(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
device: Union[str, torch.device],
|
||||
num_videos_per_prompt: int,
|
||||
do_classifier_free_guidance: bool,
|
||||
):
|
||||
image = image.to(device=device, dtype=self.vae.dtype)
|
||||
image_latents = self.vae.encode(image).latent_dist.mode()
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
negative_image_latents = torch.zeros_like(image_latents)
|
||||
|
||||
# For classifier free guidance, we need to do two forward passes.
|
||||
# Here we concatenate the unconditional and text embeddings into a single batch
|
||||
# to avoid doing two forward passes
|
||||
image_latents = torch.cat([negative_image_latents, image_latents])
|
||||
|
||||
# duplicate image_latents for each generation per prompt, using mps friendly method
|
||||
image_latents = image_latents.repeat(num_videos_per_prompt, 1, 1, 1)
|
||||
|
||||
return image_latents
|
||||
|
||||
def _get_add_time_ids(
|
||||
self,
|
||||
fps: int,
|
||||
motion_bucket_id: int,
|
||||
noise_aug_strength: float,
|
||||
dtype: torch.dtype,
|
||||
batch_size: int,
|
||||
num_videos_per_prompt: int,
|
||||
do_classifier_free_guidance: bool,
|
||||
):
|
||||
add_time_ids = [fps, motion_bucket_id, noise_aug_strength]
|
||||
|
||||
passed_add_embed_dim = self.unet.config.addition_time_embed_dim * len(add_time_ids)
|
||||
expected_add_embed_dim = self.unet.add_embedding.linear_1.in_features
|
||||
|
||||
if expected_add_embed_dim != passed_add_embed_dim:
|
||||
raise ValueError(
|
||||
f"Model expects an added time embedding vector of length {expected_add_embed_dim}, " \
|
||||
f"but a vector of {passed_add_embed_dim} was created. The model has an incorrect config. " \
|
||||
f"Please check `unet.config.time_embedding_type` and `text_encoder_2.config.projection_dim`."
|
||||
)
|
||||
|
||||
add_time_ids = torch.tensor([add_time_ids], dtype=dtype)
|
||||
add_time_ids = add_time_ids.repeat(batch_size * num_videos_per_prompt, 1)
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
add_time_ids = torch.cat([add_time_ids, add_time_ids])
|
||||
|
||||
return add_time_ids
|
||||
|
||||
def decode_latents(
|
||||
self,
|
||||
latents: torch.Tensor,
|
||||
num_frames: int,
|
||||
decode_chunk_size: int = 8):
|
||||
# [batch, frames, channels, height, width] -> [batch*frames, channels, height, width]
|
||||
latents = latents.flatten(0, 1)
|
||||
|
||||
latents = 1 / self.vae.config.scaling_factor * latents
|
||||
|
||||
forward_vae_fn = self.vae._orig_mod.forward if is_compiled_module(self.vae) else self.vae.forward
|
||||
accepts_num_frames = "num_frames" in set(inspect.signature(forward_vae_fn).parameters.keys())
|
||||
|
||||
# decode decode_chunk_size frames at a time to avoid OOM
|
||||
frames = []
|
||||
for i in range(0, latents.shape[0], decode_chunk_size):
|
||||
num_frames_in = latents[i: i + decode_chunk_size].shape[0]
|
||||
decode_kwargs = {}
|
||||
if accepts_num_frames:
|
||||
# we only pass num_frames_in if it's expected
|
||||
decode_kwargs["num_frames"] = num_frames_in
|
||||
|
||||
frame = self.vae.decode(latents[i: i + decode_chunk_size], **decode_kwargs).sample
|
||||
frames.append(frame.cpu())
|
||||
frames = torch.cat(frames, dim=0)
|
||||
|
||||
# [batch*frames, channels, height, width] -> [batch, channels, frames, height, width]
|
||||
frames = frames.reshape(-1, num_frames, *frames.shape[1:]).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloat16
|
||||
frames = frames.float()
|
||||
return frames
|
||||
|
||||
def check_inputs(self, image, height, width):
|
||||
if (
|
||||
not isinstance(image, torch.Tensor)
|
||||
and not isinstance(image, PIL.Image.Image)
|
||||
and not isinstance(image, list)
|
||||
):
|
||||
raise ValueError(
|
||||
"`image` has to be of type `torch.FloatTensor` or `PIL.Image.Image` or `List[PIL.Image.Image]` but is"
|
||||
f" {type(image)}"
|
||||
)
|
||||
|
||||
if height % 8 != 0 or width % 8 != 0:
|
||||
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
batch_size: int,
|
||||
num_frames: int,
|
||||
num_channels_latents: int,
|
||||
height: int,
|
||||
width: int,
|
||||
dtype: torch.dtype,
|
||||
device: Union[str, torch.device],
|
||||
generator: torch.Generator,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
):
|
||||
shape = (
|
||||
batch_size,
|
||||
num_frames,
|
||||
num_channels_latents // 2,
|
||||
height // self.vae_scale_factor,
|
||||
width // self.vae_scale_factor,
|
||||
)
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
|
||||
if latents is None:
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
else:
|
||||
latents = latents.to(device)
|
||||
|
||||
# scale the initial noise by the standard deviation required by the scheduler
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
return latents
|
||||
|
||||
@property
|
||||
def guidance_scale(self):
|
||||
return self._guidance_scale
|
||||
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
@property
|
||||
def do_classifier_free_guidance(self):
|
||||
if isinstance(self.guidance_scale, (int, float)):
|
||||
return self.guidance_scale > 1
|
||||
return self.guidance_scale.max() > 1
|
||||
|
||||
@property
|
||||
def num_timesteps(self):
|
||||
return self._num_timesteps
|
||||
|
||||
def prepare_extra_step_kwargs(self, generator, eta):
|
||||
# prepare extra kwargs for the scheduler step, since not all schedulers have the same signature
|
||||
# eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers.
|
||||
# eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502
|
||||
# and should be between [0, 1]
|
||||
|
||||
accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||
extra_step_kwargs = {}
|
||||
if accepts_eta:
|
||||
extra_step_kwargs["eta"] = eta
|
||||
|
||||
# check if the scheduler accepts generator
|
||||
accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||
if accepts_generator:
|
||||
extra_step_kwargs["generator"] = generator
|
||||
return extra_step_kwargs
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
image: Union[PIL.Image.Image, List[PIL.Image.Image], torch.FloatTensor],
|
||||
image_pose: Union[torch.FloatTensor],
|
||||
height: int = 576,
|
||||
width: int = 1024,
|
||||
num_frames: Optional[int] = None,
|
||||
tile_size: Optional[int] = 16,
|
||||
tile_overlap: Optional[int] = 4,
|
||||
num_inference_steps: int = 25,
|
||||
min_guidance_scale: float = 1.0,
|
||||
max_guidance_scale: float = 3.0,
|
||||
fps: int = 7,
|
||||
motion_bucket_id: int = 127,
|
||||
noise_aug_strength: float = 0.02,
|
||||
image_only_indicator: bool = False,
|
||||
decode_chunk_size: Optional[int] = None,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
validation_image_id_ante_embedding=None,
|
||||
output_type: Optional[str] = "pil",
|
||||
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
return_dict: bool = True,
|
||||
):
|
||||
r"""
|
||||
The call function to the pipeline for generation.
|
||||
|
||||
Args:
|
||||
image (`PIL.Image.Image` or `List[PIL.Image.Image]` or `torch.FloatTensor`):
|
||||
Image or images to guide image generation. If you provide a tensor, it needs to be compatible with
|
||||
[`CLIPImageProcessor`](https://huggingface.co/lambdalabs/sd-image-variations-diffusers/blob/main/
|
||||
feature_extractor/preprocessor_config.json).
|
||||
height (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`):
|
||||
The height in pixels of the generated image.
|
||||
width (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`):
|
||||
The width in pixels of the generated image.
|
||||
num_frames (`int`, *optional*):
|
||||
The number of video frames to generate. Defaults to 14 for `stable-video-diffusion-img2vid`
|
||||
and to 25 for `stable-video-diffusion-img2vid-xt`
|
||||
num_inference_steps (`int`, *optional*, defaults to 25):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference. This parameter is modulated by `strength`.
|
||||
min_guidance_scale (`float`, *optional*, defaults to 1.0):
|
||||
The minimum guidance scale. Used for the classifier free guidance with first frame.
|
||||
max_guidance_scale (`float`, *optional*, defaults to 3.0):
|
||||
The maximum guidance scale. Used for the classifier free guidance with last frame.
|
||||
fps (`int`, *optional*, defaults to 7):
|
||||
Frames per second.The rate at which the generated images shall be exported to a video after generation.
|
||||
Note that Stable Diffusion Video's UNet was micro-conditioned on fps-1 during training.
|
||||
motion_bucket_id (`int`, *optional*, defaults to 127):
|
||||
The motion bucket ID. Used as conditioning for the generation.
|
||||
The higher the number the more motion will be in the video.
|
||||
noise_aug_strength (`float`, *optional*, defaults to 0.02):
|
||||
The amount of noise added to the init image,
|
||||
the higher it is the less the video will look like the init image. Increase it for more motion.
|
||||
image_only_indicator (`bool`, *optional*, defaults to False):
|
||||
Whether to treat the inputs as batch of images instead of videos.
|
||||
decode_chunk_size (`int`, *optional*):
|
||||
The number of frames to decode at a time.The higher the chunk size, the higher the temporal consistency
|
||||
between frames, but also the higher the memory consumption.
|
||||
By default, the decoder will decode all frames at once for maximal quality.
|
||||
Reduce `decode_chunk_size` to reduce memory usage.
|
||||
num_videos_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
|
||||
generation deterministic.
|
||||
latents (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor is generated by sampling using the supplied random `generator`.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
|
||||
callback_on_step_end (`Callable`, *optional*):
|
||||
A function that calls at the end of each denoising steps during the inference. The function is called
|
||||
with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
|
||||
callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
|
||||
`callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeline class.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether to return a [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] instead of a
|
||||
plain tuple.
|
||||
device:
|
||||
On which device the pipeline runs on.
|
||||
|
||||
Returns:
|
||||
[`~pipelines.stable_diffusion.StableVideoDiffusionPipelineOutput`] or `tuple`:
|
||||
If `return_dict` is `True`,
|
||||
[`~pipelines.stable_diffusion.StableVideoDiffusionPipelineOutput`] is returned,
|
||||
otherwise a `tuple` is returned where the first element is a list of list with the generated frames.
|
||||
|
||||
Examples:
|
||||
|
||||
```py
|
||||
from diffusers import StableVideoDiffusionPipeline
|
||||
from diffusers.utils import load_image, export_to_video
|
||||
|
||||
pipe = StableVideoDiffusionPipeline.from_pretrained(
|
||||
"stabilityai/stable-video-diffusion-img2vid-xt", torch_dtype=torch.float16, variant="fp16")
|
||||
pipe.to("cuda")
|
||||
|
||||
image = load_image(
|
||||
"https://lh3.googleusercontent.com/y-iFOHfLTwkuQSUegpwDdgKmOjRSTvPxat63dQLB25xkTs4lhIbRUFeNBWZzYf370g=s1200")
|
||||
image = image.resize((1024, 576))
|
||||
|
||||
frames = pipe(image, num_frames=25, decode_chunk_size=8).frames[0]
|
||||
export_to_video(frames, "generated.mp4", fps=7)
|
||||
```
|
||||
"""
|
||||
# 0. Default height and width to unet
|
||||
height = height or self.unet.config.sample_size * self.vae_scale_factor
|
||||
width = width or self.unet.config.sample_size * self.vae_scale_factor
|
||||
|
||||
num_frames = num_frames if num_frames is not None else self.unet.config.num_frames
|
||||
decode_chunk_size = decode_chunk_size if decode_chunk_size is not None else num_frames
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(image, height, width)
|
||||
|
||||
# 2. Define call parameters
|
||||
if isinstance(image, PIL.Image.Image):
|
||||
batch_size = 1
|
||||
elif isinstance(image, list):
|
||||
batch_size = len(image)
|
||||
else:
|
||||
batch_size = image.shape[0]
|
||||
device = self._execution_device
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
do_classifier_free_guidance = max_guidance_scale >= 1.0
|
||||
self._guidance_scale = max_guidance_scale
|
||||
|
||||
# 3. Encode input image
|
||||
image_embeddings = self._encode_image(image, device, num_videos_per_prompt, do_classifier_free_guidance)
|
||||
# self.image_encoder.cpu()
|
||||
|
||||
# NOTE: Stable Diffusion Video was conditioned on fps - 1, which
|
||||
# is why it is reduced here.
|
||||
fps = fps - 1
|
||||
|
||||
# 4. Encode input image using VAE
|
||||
|
||||
# print(image_embeddings.size()) # [2, 1, 1024]
|
||||
validation_image_id_ante_embedding = torch.from_numpy(validation_image_id_ante_embedding).unsqueeze(0)
|
||||
validation_image_id_ante_embedding = validation_image_id_ante_embedding.to(device=device, dtype=image_embeddings.dtype)
|
||||
|
||||
faceid_latents = self.face_encoder(validation_image_id_ante_embedding, image_embeddings[1:])
|
||||
# print(faceid_latents.size()) # [1, 4, 1024]
|
||||
uncond_image_embeddings = image_embeddings[:1]
|
||||
uncond_faceid_latents = torch.zeros_like(faceid_latents)
|
||||
uncond_image_embeddings = torch.cat([uncond_image_embeddings, uncond_faceid_latents], dim=1)
|
||||
cond_image_embeddings = image_embeddings[1:]
|
||||
cond_image_embeddings = torch.cat([cond_image_embeddings, faceid_latents], dim=1)
|
||||
image_embeddings = torch.cat([uncond_image_embeddings, cond_image_embeddings])
|
||||
|
||||
image = self.image_processor.preprocess(image, height=height, width=width).to(device)
|
||||
noise = randn_tensor(image.shape, generator=generator, device=device, dtype=image.dtype)
|
||||
image = image + noise_aug_strength * noise
|
||||
|
||||
needs_upcasting = (self.vae.dtype == torch.float16 or self.vae.dtype == torch.bfloat16) and self.vae.config.force_upcast
|
||||
if needs_upcasting:
|
||||
self_vae_dtype = self.vae.dtype
|
||||
self.vae.to(dtype=torch.float32)
|
||||
|
||||
image_latents = self._encode_vae_image(
|
||||
image,
|
||||
device=device,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
do_classifier_free_guidance=do_classifier_free_guidance,
|
||||
)
|
||||
image_latents = image_latents.to(image_embeddings.dtype)
|
||||
|
||||
if needs_upcasting:
|
||||
self.vae.to(dtype=self_vae_dtype)
|
||||
# self.vae.cpu()
|
||||
|
||||
# Repeat the image latents for each frame so we can concatenate them with the noise
|
||||
# image_latents [batch, channels, height, width] ->[batch, num_frames, channels, height, width]
|
||||
image_latents = image_latents.unsqueeze(1).repeat(1, num_frames, 1, 1, 1)
|
||||
|
||||
# 5. Get Added Time IDs
|
||||
added_time_ids = self._get_add_time_ids(
|
||||
fps,
|
||||
motion_bucket_id,
|
||||
noise_aug_strength,
|
||||
image_embeddings.dtype,
|
||||
batch_size,
|
||||
num_videos_per_prompt,
|
||||
self.do_classifier_free_guidance,
|
||||
)
|
||||
added_time_ids = added_time_ids.to(device)
|
||||
|
||||
# 4. Prepare timesteps
|
||||
timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, None)
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = self.unet.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
tile_size,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
image_embeddings.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
latents = latents.repeat(1, num_frames // tile_size + 1, 1, 1, 1)[:, :num_frames]
|
||||
|
||||
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
|
||||
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, 0.0)
|
||||
|
||||
# 7. Prepare guidance scale
|
||||
guidance_scale = torch.linspace(min_guidance_scale, max_guidance_scale, num_frames).unsqueeze(0)
|
||||
guidance_scale = guidance_scale.to(device, latents.dtype)
|
||||
guidance_scale = guidance_scale.repeat(batch_size * num_videos_per_prompt, 1)
|
||||
guidance_scale = _append_dims(guidance_scale, latents.ndim)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
|
||||
# 8. Denoising loop
|
||||
self._num_timesteps = len(timesteps)
|
||||
indices = [[0, *range(i + 1, min(i + tile_size, num_frames))] for i in
|
||||
range(0, num_frames - tile_size + 1, tile_size - tile_overlap)]
|
||||
if indices[-1][-1] < num_frames - 1:
|
||||
indices.append([0, *range(num_frames - tile_size + 1, num_frames)])
|
||||
|
||||
pose_pil_image_list = []
|
||||
for pose in image_pose:
|
||||
pose = torch.from_numpy(np.array(pose)).float()
|
||||
pose = pose / 127.5 - 1
|
||||
pose_pil_image_list.append(pose)
|
||||
pose_pil_image_list = torch.stack(pose_pil_image_list, dim=0)
|
||||
pose_pil_image_list = rearrange(pose_pil_image_list, "f h w c -> f c h w")
|
||||
|
||||
|
||||
# print(indices) # [[0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15]]
|
||||
# print(pose_pil_image_list.size()) # [16, 3, 512, 512]
|
||||
|
||||
self.pose_net.to(device)
|
||||
self.unet.to(device)
|
||||
|
||||
with torch.cuda.device(device):
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
with self.progress_bar(total=len(timesteps) * len(indices)) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
|
||||
# Concatenate image_latents over channels dimension
|
||||
latent_model_input = torch.cat([latent_model_input, image_latents], dim=2)
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred = torch.zeros_like(image_latents)
|
||||
noise_pred_cnt = image_latents.new_zeros((num_frames,))
|
||||
weight = (torch.arange(tile_size, device=device) + 0.5) * 2. / tile_size
|
||||
weight = torch.minimum(weight, 2 - weight)
|
||||
for idx in indices:
|
||||
# classification-free inference
|
||||
pose_latents = self.pose_net(pose_pil_image_list[idx].to(device=device, dtype=latent_model_input.dtype))
|
||||
_noise_pred = self.unet(
|
||||
latent_model_input[:1, idx],
|
||||
t,
|
||||
encoder_hidden_states=image_embeddings[:1],
|
||||
added_time_ids=added_time_ids[:1],
|
||||
pose_latents=None,
|
||||
image_only_indicator=image_only_indicator,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred[:1, idx] += _noise_pred * weight[:, None, None, None]
|
||||
|
||||
# normal inference
|
||||
_noise_pred = self.unet(
|
||||
latent_model_input[1:, idx],
|
||||
t,
|
||||
encoder_hidden_states=image_embeddings[1:],
|
||||
added_time_ids=added_time_ids[1:],
|
||||
pose_latents=pose_latents,
|
||||
image_only_indicator=image_only_indicator,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred[1:, idx] += _noise_pred * weight[:, None, None, None]
|
||||
|
||||
noise_pred_cnt[idx] += weight
|
||||
progress_bar.update()
|
||||
noise_pred.div_(noise_pred_cnt[:, None, None, None])
|
||||
|
||||
# perform guidance
|
||||
if self.do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_cond - noise_pred_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
for k in callback_on_step_end_tensor_inputs:
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
|
||||
# self.pose_net.cpu()
|
||||
# self.unet.cpu()
|
||||
# self.face_encoder.cpu()
|
||||
|
||||
if not output_type == "latent":
|
||||
self.vae.decoder.to(device)
|
||||
frames = self.decode_latents(latents, num_frames, decode_chunk_size)
|
||||
# print(frames.size()) # [1, 3, 16, 512, 512]
|
||||
# print(latents.size()) # [1, 16, 4, 64, 64]
|
||||
frames = tensor2vid(frames, self.image_processor, output_type=output_type)
|
||||
# print(frames[0].size()) # [16, 3, 512, 512]
|
||||
else:
|
||||
frames = latents
|
||||
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return frames
|
||||
|
||||
return InferenceAnimationPipelineOutput(frames=frames)
|
||||
@@ -0,0 +1,703 @@
|
||||
import inspect
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Dict, List, Optional, Union
|
||||
|
||||
import PIL.Image
|
||||
import einops
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers.image_processor import VaeImageProcessor, PipelineImageInput
|
||||
from diffusers.models import AutoencoderKLTemporalDecoder, UNetSpatioTemporalConditionModel
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion import retrieve_timesteps
|
||||
from diffusers.pipelines.stable_video_diffusion.pipeline_stable_video_diffusion \
|
||||
import _resize_with_antialiasing, _append_dims
|
||||
from diffusers.schedulers import EulerDiscreteScheduler
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
from diffusers.utils.torch_utils import is_compiled_module, randn_tensor
|
||||
from animation.modules.attention_processor import AnimationAttnProcessor, AnimationIDAttnProcessor
|
||||
from animation.modules.id_encoder import FusionFaceId
|
||||
from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection
|
||||
|
||||
from insightface.app import FaceAnalysis
|
||||
|
||||
from ..modules.pose_net import PoseNet
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
def _append_dims(x, target_dims):
|
||||
"""Appends dimensions to the end of a tensor until it has target_dims dimensions."""
|
||||
dims_to_append = target_dims - x.ndim
|
||||
if dims_to_append < 0:
|
||||
raise ValueError(f"input has {x.ndim} dims but target_dims is {target_dims}, which is less")
|
||||
return x[(...,) + (None,) * dims_to_append]
|
||||
|
||||
|
||||
# Copied from diffusers.pipelines.animatediff.pipeline_animatediff.tensor2vid
|
||||
def tensor2vid(video: torch.Tensor, processor: "VaeImageProcessor", output_type: str = "np"):
|
||||
batch_size, channels, num_frames, height, width = video.shape
|
||||
outputs = []
|
||||
for batch_idx in range(batch_size):
|
||||
batch_vid = video[batch_idx].permute(1, 0, 2, 3)
|
||||
batch_output = processor.postprocess(batch_vid, output_type)
|
||||
|
||||
outputs.append(batch_output)
|
||||
|
||||
if output_type == "np":
|
||||
outputs = np.stack(outputs)
|
||||
|
||||
elif output_type == "pt":
|
||||
outputs = torch.stack(outputs)
|
||||
|
||||
elif not output_type == "pil":
|
||||
raise ValueError(f"{output_type} does not exist. Please choose one of ['np', 'pt', 'pil]")
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
@dataclass
|
||||
class AnimationPipelineOutput(BaseOutput):
|
||||
r"""
|
||||
Output class for mimicmotion pipeline.
|
||||
|
||||
Args:
|
||||
frames (`[List[List[PIL.Image.Image]]`, `np.ndarray`, `torch.Tensor`]):
|
||||
List of denoised PIL images of length `batch_size` or numpy array or torch tensor of shape `(batch_size,
|
||||
num_frames, height, width, num_channels)`.
|
||||
"""
|
||||
|
||||
frames: Union[List[List[PIL.Image.Image]], np.ndarray, torch.Tensor]
|
||||
|
||||
|
||||
class AnimationPipeline(DiffusionPipeline):
|
||||
r"""
|
||||
Pipeline to generate video from an input image using Stable Video Diffusion.
|
||||
|
||||
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods
|
||||
implemented for all pipelines (downloading, saving, running on a particular device, etc.).
|
||||
|
||||
Args:
|
||||
vae ([`AutoencoderKLTemporalDecoder`]):
|
||||
Variational Auto-Encoder (VAE) model to encode and decode images to and from latent representations.
|
||||
image_encoder ([`~transformers.CLIPVisionModelWithProjection`]):
|
||||
Frozen CLIP image-encoder ([laion/CLIP-ViT-H-14-laion2B-s32B-b79K]
|
||||
(https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K)).
|
||||
unet ([`UNetSpatioTemporalConditionModel`]):
|
||||
A `UNetSpatioTemporalConditionModel` to denoise the encoded image latents.
|
||||
scheduler ([`EulerDiscreteScheduler`]):
|
||||
A scheduler to be used in combination with `unet` to denoise the encoded image latents.
|
||||
feature_extractor ([`~transformers.CLIPImageProcessor`]):
|
||||
A `CLIPImageProcessor` to extract features from generated images.
|
||||
pose_net ([`PoseNet`]):
|
||||
A `` to inject pose signals into unet.
|
||||
"""
|
||||
|
||||
model_cpu_offload_seq = "image_encoder->unet->vae"
|
||||
_callback_tensor_inputs = ["latents"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vae,
|
||||
image_encoder,
|
||||
unet,
|
||||
scheduler: EulerDiscreteScheduler,
|
||||
feature_extractor: CLIPImageProcessor,
|
||||
pose_net: PoseNet,
|
||||
face_encoder: FusionFaceId,
|
||||
torch_dtype,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
image_encoder=image_encoder,
|
||||
unet=unet,
|
||||
scheduler=scheduler,
|
||||
feature_extractor=feature_extractor,
|
||||
pose_net=pose_net,
|
||||
face_encoder=face_encoder,
|
||||
)
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
|
||||
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
|
||||
|
||||
print(1/0)
|
||||
|
||||
self.app = FaceAnalysis(name="buffalo_l", providers=['CUDAExecutionProvider', 'CPUExecutionProvider'])
|
||||
self.app.prepare(ctx_id=0, det_size=(640, 640))
|
||||
|
||||
self.lora_rank = 128
|
||||
self.torch_dtype = torch_dtype
|
||||
self.num_tokens = 4
|
||||
self.set_ip_adapter()
|
||||
|
||||
|
||||
def get_prepare_faceid(self, face_image):
|
||||
faceid_image = np.array(face_image)
|
||||
faces = self.app.get(faceid_image)
|
||||
if faces == []:
|
||||
faceid_embeds = torch.zeros_like(torch.empty((1, 512)))
|
||||
else:
|
||||
faceid_embeds = torch.from_numpy(faces[0].normed_embedding).unsqueeze(0)
|
||||
return faceid_embeds
|
||||
|
||||
def set_ip_adapter(self):
|
||||
unet = self.unet
|
||||
attn_procs = {}
|
||||
for name in unet.attn_processors.keys():
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") else unet.config.cross_attention_dim
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = unet.config.block_out_channels[-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(unet.config.block_out_channels))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = unet.config.block_out_channels[block_id]
|
||||
if cross_attention_dim is None:
|
||||
attn_procs[name] = AnimationAttnProcessor(
|
||||
hidden_size=hidden_size, cross_attention_dim=cross_attention_dim, rank=self.lora_rank,
|
||||
).to(self.device, dtype=self.torch_dtype)
|
||||
else:
|
||||
attn_procs[name] = AnimationIDAttnProcessor(
|
||||
hidden_size=hidden_size, cross_attention_dim=cross_attention_dim, scale=1.0, rank=self.lora_rank, num_tokens=self.num_tokens,
|
||||
).to(self.device, dtype=self.torch_dtype)
|
||||
|
||||
unet.set_attn_processor(attn_procs)
|
||||
|
||||
def _encode_image(
|
||||
self,
|
||||
image: PipelineImageInput,
|
||||
device: Union[str, torch.device],
|
||||
num_videos_per_prompt: int,
|
||||
do_classifier_free_guidance: bool):
|
||||
dtype = next(self.image_encoder.parameters()).dtype
|
||||
|
||||
if not isinstance(image, torch.Tensor):
|
||||
image = self.image_processor.pil_to_numpy(image)
|
||||
image = self.image_processor.numpy_to_pt(image)
|
||||
|
||||
# We normalize the image before resizing to match with the original implementation.
|
||||
# Then we unnormalize it after resizing.
|
||||
image = image * 2.0 - 1.0
|
||||
image = _resize_with_antialiasing(image, (224, 224))
|
||||
image = (image + 1.0) / 2.0
|
||||
|
||||
# Normalize the image with for CLIP input
|
||||
image = self.feature_extractor(
|
||||
images=image,
|
||||
do_normalize=True,
|
||||
do_center_crop=False,
|
||||
do_resize=False,
|
||||
do_rescale=False,
|
||||
return_tensors="pt",
|
||||
).pixel_values
|
||||
|
||||
image = image.to(device=device, dtype=dtype)
|
||||
image_embeddings = self.image_encoder(image).image_embeds
|
||||
image_embeddings = image_embeddings.unsqueeze(1)
|
||||
|
||||
# duplicate image embeddings for each generation per prompt, using mps friendly method
|
||||
bs_embed, seq_len, _ = image_embeddings.shape
|
||||
image_embeddings = image_embeddings.repeat(1, num_videos_per_prompt, 1)
|
||||
image_embeddings = image_embeddings.view(bs_embed * num_videos_per_prompt, seq_len, -1)
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
negative_image_embeddings = torch.zeros_like(image_embeddings)
|
||||
|
||||
# For classifier free guidance, we need to do two forward passes.
|
||||
# Here we concatenate the unconditional and text embeddings into a single batch
|
||||
# to avoid doing two forward passes
|
||||
image_embeddings = torch.cat([negative_image_embeddings, image_embeddings])
|
||||
|
||||
return image_embeddings
|
||||
|
||||
def _encode_vae_image(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
device: Union[str, torch.device],
|
||||
num_videos_per_prompt: int,
|
||||
do_classifier_free_guidance: bool,
|
||||
):
|
||||
image = image.to(device=device, dtype=self.vae.dtype)
|
||||
image_latents = self.vae.encode(image).latent_dist.mode()
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
negative_image_latents = torch.zeros_like(image_latents)
|
||||
|
||||
# For classifier free guidance, we need to do two forward passes.
|
||||
# Here we concatenate the unconditional and text embeddings into a single batch
|
||||
# to avoid doing two forward passes
|
||||
image_latents = torch.cat([negative_image_latents, image_latents])
|
||||
|
||||
# duplicate image_latents for each generation per prompt, using mps friendly method
|
||||
image_latents = image_latents.repeat(num_videos_per_prompt, 1, 1, 1)
|
||||
|
||||
return image_latents
|
||||
|
||||
def _get_add_time_ids(
|
||||
self,
|
||||
fps: int,
|
||||
motion_bucket_id: int,
|
||||
noise_aug_strength: float,
|
||||
dtype: torch.dtype,
|
||||
batch_size: int,
|
||||
num_videos_per_prompt: int,
|
||||
do_classifier_free_guidance: bool,
|
||||
):
|
||||
add_time_ids = [fps, motion_bucket_id, noise_aug_strength]
|
||||
|
||||
passed_add_embed_dim = self.unet.config.addition_time_embed_dim * len(add_time_ids)
|
||||
expected_add_embed_dim = self.unet.add_embedding.linear_1.in_features
|
||||
|
||||
if expected_add_embed_dim != passed_add_embed_dim:
|
||||
raise ValueError(
|
||||
f"Model expects an added time embedding vector of length {expected_add_embed_dim}, " \
|
||||
f"but a vector of {passed_add_embed_dim} was created. The model has an incorrect config. " \
|
||||
f"Please check `unet.config.time_embedding_type` and `text_encoder_2.config.projection_dim`."
|
||||
)
|
||||
|
||||
add_time_ids = torch.tensor([add_time_ids], dtype=dtype)
|
||||
add_time_ids = add_time_ids.repeat(batch_size * num_videos_per_prompt, 1)
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
add_time_ids = torch.cat([add_time_ids, add_time_ids])
|
||||
|
||||
return add_time_ids
|
||||
|
||||
def decode_latents(
|
||||
self,
|
||||
latents: torch.Tensor,
|
||||
num_frames: int,
|
||||
decode_chunk_size: int = 8):
|
||||
# [batch, frames, channels, height, width] -> [batch*frames, channels, height, width]
|
||||
latents = latents.flatten(0, 1)
|
||||
|
||||
latents = 1 / self.vae.config.scaling_factor * latents
|
||||
|
||||
forward_vae_fn = self.vae._orig_mod.forward if is_compiled_module(self.vae) else self.vae.forward
|
||||
accepts_num_frames = "num_frames" in set(inspect.signature(forward_vae_fn).parameters.keys())
|
||||
|
||||
# decode decode_chunk_size frames at a time to avoid OOM
|
||||
frames = []
|
||||
for i in range(0, latents.shape[0], decode_chunk_size):
|
||||
num_frames_in = latents[i: i + decode_chunk_size].shape[0]
|
||||
decode_kwargs = {}
|
||||
if accepts_num_frames:
|
||||
# we only pass num_frames_in if it's expected
|
||||
decode_kwargs["num_frames"] = num_frames_in
|
||||
|
||||
frame = self.vae.decode(latents[i: i + decode_chunk_size], **decode_kwargs).sample
|
||||
frames.append(frame.cpu())
|
||||
frames = torch.cat(frames, dim=0)
|
||||
|
||||
# [batch*frames, channels, height, width] -> [batch, channels, frames, height, width]
|
||||
frames = frames.reshape(-1, num_frames, *frames.shape[1:]).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloat16
|
||||
frames = frames.float()
|
||||
return frames
|
||||
|
||||
def check_inputs(self, image, height, width):
|
||||
if (
|
||||
not isinstance(image, torch.Tensor)
|
||||
and not isinstance(image, PIL.Image.Image)
|
||||
and not isinstance(image, list)
|
||||
):
|
||||
raise ValueError(
|
||||
"`image` has to be of type `torch.FloatTensor` or `PIL.Image.Image` or `List[PIL.Image.Image]` but is"
|
||||
f" {type(image)}"
|
||||
)
|
||||
|
||||
if height % 8 != 0 or width % 8 != 0:
|
||||
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
batch_size: int,
|
||||
num_frames: int,
|
||||
num_channels_latents: int,
|
||||
height: int,
|
||||
width: int,
|
||||
dtype: torch.dtype,
|
||||
device: Union[str, torch.device],
|
||||
generator: torch.Generator,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
):
|
||||
shape = (
|
||||
batch_size,
|
||||
num_frames,
|
||||
num_channels_latents // 2,
|
||||
height // self.vae_scale_factor,
|
||||
width // self.vae_scale_factor,
|
||||
)
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
|
||||
if latents is None:
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
else:
|
||||
latents = latents.to(device)
|
||||
|
||||
# scale the initial noise by the standard deviation required by the scheduler
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
return latents
|
||||
|
||||
@property
|
||||
def guidance_scale(self):
|
||||
return self._guidance_scale
|
||||
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
@property
|
||||
def do_classifier_free_guidance(self):
|
||||
if isinstance(self.guidance_scale, (int, float)):
|
||||
return self.guidance_scale > 1
|
||||
return self.guidance_scale.max() > 1
|
||||
|
||||
@property
|
||||
def num_timesteps(self):
|
||||
return self._num_timesteps
|
||||
|
||||
def prepare_extra_step_kwargs(self, generator, eta):
|
||||
# prepare extra kwargs for the scheduler step, since not all schedulers have the same signature
|
||||
# eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers.
|
||||
# eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502
|
||||
# and should be between [0, 1]
|
||||
|
||||
accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||
extra_step_kwargs = {}
|
||||
if accepts_eta:
|
||||
extra_step_kwargs["eta"] = eta
|
||||
|
||||
# check if the scheduler accepts generator
|
||||
accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||
if accepts_generator:
|
||||
extra_step_kwargs["generator"] = generator
|
||||
return extra_step_kwargs
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
image: Union[PIL.Image.Image, List[PIL.Image.Image], torch.FloatTensor],
|
||||
image_pose: Union[torch.FloatTensor],
|
||||
height: int = 576,
|
||||
width: int = 1024,
|
||||
num_frames: Optional[int] = None,
|
||||
tile_size: Optional[int] = 16,
|
||||
tile_overlap: Optional[int] = 4,
|
||||
num_inference_steps: int = 25,
|
||||
min_guidance_scale: float = 1.0,
|
||||
max_guidance_scale: float = 3.0,
|
||||
fps: int = 7,
|
||||
motion_bucket_id: int = 127,
|
||||
noise_aug_strength: float = 0.02,
|
||||
image_only_indicator: bool = False,
|
||||
decode_chunk_size: Optional[int] = None,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
return_dict: bool = True,
|
||||
):
|
||||
r"""
|
||||
The call function to the pipeline for generation.
|
||||
|
||||
Args:
|
||||
image (`PIL.Image.Image` or `List[PIL.Image.Image]` or `torch.FloatTensor`):
|
||||
Image or images to guide image generation. If you provide a tensor, it needs to be compatible with
|
||||
[`CLIPImageProcessor`](https://huggingface.co/lambdalabs/sd-image-variations-diffusers/blob/main/
|
||||
feature_extractor/preprocessor_config.json).
|
||||
height (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`):
|
||||
The height in pixels of the generated image.
|
||||
width (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`):
|
||||
The width in pixels of the generated image.
|
||||
num_frames (`int`, *optional*):
|
||||
The number of video frames to generate. Defaults to 14 for `stable-video-diffusion-img2vid`
|
||||
and to 25 for `stable-video-diffusion-img2vid-xt`
|
||||
num_inference_steps (`int`, *optional*, defaults to 25):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference. This parameter is modulated by `strength`.
|
||||
min_guidance_scale (`float`, *optional*, defaults to 1.0):
|
||||
The minimum guidance scale. Used for the classifier free guidance with first frame.
|
||||
max_guidance_scale (`float`, *optional*, defaults to 3.0):
|
||||
The maximum guidance scale. Used for the classifier free guidance with last frame.
|
||||
fps (`int`, *optional*, defaults to 7):
|
||||
Frames per second.The rate at which the generated images shall be exported to a video after generation.
|
||||
Note that Stable Diffusion Video's UNet was micro-conditioned on fps-1 during training.
|
||||
motion_bucket_id (`int`, *optional*, defaults to 127):
|
||||
The motion bucket ID. Used as conditioning for the generation.
|
||||
The higher the number the more motion will be in the video.
|
||||
noise_aug_strength (`float`, *optional*, defaults to 0.02):
|
||||
The amount of noise added to the init image,
|
||||
the higher it is the less the video will look like the init image. Increase it for more motion.
|
||||
image_only_indicator (`bool`, *optional*, defaults to False):
|
||||
Whether to treat the inputs as batch of images instead of videos.
|
||||
decode_chunk_size (`int`, *optional*):
|
||||
The number of frames to decode at a time.The higher the chunk size, the higher the temporal consistency
|
||||
between frames, but also the higher the memory consumption.
|
||||
By default, the decoder will decode all frames at once for maximal quality.
|
||||
Reduce `decode_chunk_size` to reduce memory usage.
|
||||
num_videos_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
|
||||
generation deterministic.
|
||||
latents (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor is generated by sampling using the supplied random `generator`.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
|
||||
callback_on_step_end (`Callable`, *optional*):
|
||||
A function that calls at the end of each denoising steps during the inference. The function is called
|
||||
with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
|
||||
callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
|
||||
`callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeline class.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether to return a [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] instead of a
|
||||
plain tuple.
|
||||
device:
|
||||
On which device the pipeline runs on.
|
||||
|
||||
Returns:
|
||||
[`~pipelines.stable_diffusion.StableVideoDiffusionPipelineOutput`] or `tuple`:
|
||||
If `return_dict` is `True`,
|
||||
[`~pipelines.stable_diffusion.StableVideoDiffusionPipelineOutput`] is returned,
|
||||
otherwise a `tuple` is returned where the first element is a list of list with the generated frames.
|
||||
|
||||
Examples:
|
||||
|
||||
```py
|
||||
from diffusers import StableVideoDiffusionPipeline
|
||||
from diffusers.utils import load_image, export_to_video
|
||||
|
||||
pipe = StableVideoDiffusionPipeline.from_pretrained(
|
||||
"stabilityai/stable-video-diffusion-img2vid-xt", torch_dtype=torch.float16, variant="fp16")
|
||||
pipe.to("cuda")
|
||||
|
||||
image = load_image(
|
||||
"https://lh3.googleusercontent.com/y-iFOHfLTwkuQSUegpwDdgKmOjRSTvPxat63dQLB25xkTs4lhIbRUFeNBWZzYf370g=s1200")
|
||||
image = image.resize((1024, 576))
|
||||
|
||||
frames = pipe(image, num_frames=25, decode_chunk_size=8).frames[0]
|
||||
export_to_video(frames, "generated.mp4", fps=7)
|
||||
```
|
||||
"""
|
||||
# 0. Default height and width to unet
|
||||
height = height or self.unet.config.sample_size * self.vae_scale_factor
|
||||
width = width or self.unet.config.sample_size * self.vae_scale_factor
|
||||
|
||||
num_frames = num_frames if num_frames is not None else self.unet.config.num_frames
|
||||
decode_chunk_size = decode_chunk_size if decode_chunk_size is not None else num_frames
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(image, height, width)
|
||||
|
||||
# 2. Define call parameters
|
||||
if isinstance(image, PIL.Image.Image):
|
||||
batch_size = 1
|
||||
elif isinstance(image, list):
|
||||
batch_size = len(image)
|
||||
else:
|
||||
batch_size = image.shape[0]
|
||||
device = self._execution_device
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
do_classifier_free_guidance = max_guidance_scale >= 1.0
|
||||
self._guidance_scale = max_guidance_scale
|
||||
|
||||
# 3. Encode input image
|
||||
image_embeddings = self._encode_image(image, device, num_videos_per_prompt, do_classifier_free_guidance)
|
||||
self.image_encoder.cpu()
|
||||
|
||||
# NOTE: Stable Diffusion Video was conditioned on fps - 1, which
|
||||
# is why it is reduced here.
|
||||
fps = fps - 1
|
||||
|
||||
# 4. Encode input image using VAE
|
||||
|
||||
faceid_embeds = self.get_prepare_faceid(face_image=image).to(device)
|
||||
faceid_latents = self.face_encoder(faceid_embeds, image_embeddings[1:])
|
||||
print("This is in the animation pipeline")
|
||||
print(image_embeddings.size())
|
||||
print(image_embeddings[1:].size())
|
||||
print(faceid_latents.size())
|
||||
print(1/0)
|
||||
uncond_image_embeddings = image_embeddings[:1]
|
||||
uncond_faceid_latents = torch.zeros_like(faceid_latents)
|
||||
uncond_image_embeddings = torch.cat([uncond_image_embeddings, uncond_faceid_latents])
|
||||
cond_image_embeddings = image_embeddings[1:]
|
||||
cond_image_embeddings = torch.cat([cond_image_embeddings, faceid_latents])
|
||||
image_embeddings = torch.cat([uncond_image_embeddings, cond_image_embeddings])
|
||||
|
||||
image = self.image_processor.preprocess(image, height=height, width=width).to(device)
|
||||
noise = randn_tensor(image.shape, generator=generator, device=device, dtype=image.dtype)
|
||||
image = image + noise_aug_strength * noise
|
||||
|
||||
needs_upcasting = (self.vae.dtype == torch.float16 or self.vae.dtype == torch.bfloat16) and self.vae.config.force_upcast
|
||||
if needs_upcasting:
|
||||
self_vae_dtype = self.vae.dtype
|
||||
self.vae.to(dtype=torch.float32)
|
||||
|
||||
image_latents = self._encode_vae_image(
|
||||
image,
|
||||
device=device,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
do_classifier_free_guidance=do_classifier_free_guidance,
|
||||
)
|
||||
image_latents = image_latents.to(image_embeddings.dtype)
|
||||
|
||||
if needs_upcasting:
|
||||
self.vae.to(dtype=self_vae_dtype)
|
||||
self.vae.cpu()
|
||||
|
||||
# Repeat the image latents for each frame so we can concatenate them with the noise
|
||||
# image_latents [batch, channels, height, width] ->[batch, num_frames, channels, height, width]
|
||||
image_latents = image_latents.unsqueeze(1).repeat(1, num_frames, 1, 1, 1)
|
||||
|
||||
# 5. Get Added Time IDs
|
||||
added_time_ids = self._get_add_time_ids(
|
||||
fps,
|
||||
motion_bucket_id,
|
||||
noise_aug_strength,
|
||||
image_embeddings.dtype,
|
||||
batch_size,
|
||||
num_videos_per_prompt,
|
||||
self.do_classifier_free_guidance,
|
||||
)
|
||||
added_time_ids = added_time_ids.to(device)
|
||||
|
||||
# 4. Prepare timesteps
|
||||
timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, None)
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = self.unet.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
tile_size,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
image_embeddings.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
latents = latents.repeat(1, num_frames // tile_size + 1, 1, 1, 1)[:, :num_frames]
|
||||
|
||||
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
|
||||
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, 0.0)
|
||||
|
||||
# 7. Prepare guidance scale
|
||||
guidance_scale = torch.linspace(min_guidance_scale, max_guidance_scale, num_frames).unsqueeze(0)
|
||||
guidance_scale = guidance_scale.to(device, latents.dtype)
|
||||
guidance_scale = guidance_scale.repeat(batch_size * num_videos_per_prompt, 1)
|
||||
guidance_scale = _append_dims(guidance_scale, latents.ndim)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
|
||||
# 8. Denoising loop
|
||||
self._num_timesteps = len(timesteps)
|
||||
indices = [[0, *range(i + 1, min(i + tile_size, num_frames))] for i in
|
||||
range(0, num_frames - tile_size + 1, tile_size - tile_overlap)]
|
||||
if indices[-1][-1] < num_frames - 1:
|
||||
indices.append([0, *range(num_frames - tile_size + 1, num_frames)])
|
||||
|
||||
# self.pose_net.to(device)
|
||||
# self.unet.to(device)
|
||||
# self.face_encoder.to(device)
|
||||
|
||||
with torch.cuda.device(device):
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
with self.progress_bar(total=len(timesteps) * len(indices)) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
|
||||
# Concatenate image_latents over channels dimension
|
||||
latent_model_input = torch.cat([latent_model_input, image_latents], dim=2)
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred = torch.zeros_like(image_latents)
|
||||
noise_pred_cnt = image_latents.new_zeros((num_frames,))
|
||||
weight = (torch.arange(tile_size, device=device) + 0.5) * 2. / tile_size
|
||||
weight = torch.minimum(weight, 2 - weight)
|
||||
for idx in indices:
|
||||
|
||||
# classification-free inference
|
||||
pose_latents = self.pose_net(image_pose[idx].to(device))
|
||||
_noise_pred = self.unet(
|
||||
latent_model_input[:1, idx],
|
||||
t,
|
||||
encoder_hidden_states=image_embeddings[:1],
|
||||
added_time_ids=added_time_ids[:1],
|
||||
pose_latents=None,
|
||||
image_only_indicator=image_only_indicator,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred[:1, idx] += _noise_pred * weight[:, None, None, None]
|
||||
|
||||
# normal inference
|
||||
_noise_pred = self.unet(
|
||||
latent_model_input[1:, idx],
|
||||
t,
|
||||
encoder_hidden_states=image_embeddings[1:],
|
||||
added_time_ids=added_time_ids[1:],
|
||||
pose_latents=pose_latents,
|
||||
image_only_indicator=image_only_indicator,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred[1:, idx] += _noise_pred * weight[:, None, None, None]
|
||||
|
||||
noise_pred_cnt[idx] += weight
|
||||
progress_bar.update()
|
||||
noise_pred.div_(noise_pred_cnt[:, None, None, None])
|
||||
|
||||
# perform guidance
|
||||
if self.do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_cond - noise_pred_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
for k in callback_on_step_end_tensor_inputs:
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
|
||||
self.pose_net.cpu()
|
||||
self.unet.cpu()
|
||||
self.face_encoder.cpu()
|
||||
|
||||
if not output_type == "latent":
|
||||
self.vae.decoder.to(device)
|
||||
frames = self.decode_latents(latents, num_frames, decode_chunk_size)
|
||||
frames = tensor2vid(frames, self.image_processor, output_type=output_type)
|
||||
else:
|
||||
frames = latents
|
||||
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return frames
|
||||
|
||||
return AnimationPipelineOutput(frames=frames)
|
||||
@@ -0,0 +1,721 @@
|
||||
import inspect
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Dict, List, Optional, Union
|
||||
|
||||
import PIL.Image
|
||||
import einops
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers.image_processor import VaeImageProcessor, PipelineImageInput
|
||||
from diffusers.models import AutoencoderKLTemporalDecoder, UNetSpatioTemporalConditionModel
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion import retrieve_timesteps
|
||||
from diffusers.pipelines.stable_video_diffusion.pipeline_stable_video_diffusion \
|
||||
import _resize_with_antialiasing, _append_dims
|
||||
from diffusers.schedulers import EulerDiscreteScheduler
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
from diffusers.utils.torch_utils import is_compiled_module, randn_tensor
|
||||
|
||||
from animation.modules.attention_processor import AnimationAttnProcessor, AnimationIDAttnProcessor
|
||||
from animation.modules.id_encoder import FusionFaceId
|
||||
from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection
|
||||
from einops import rearrange
|
||||
from insightface.app import FaceAnalysis
|
||||
|
||||
from ..modules.pose_net import PoseNet
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
def _append_dims(x, target_dims):
|
||||
"""Appends dimensions to the end of a tensor until it has target_dims dimensions."""
|
||||
dims_to_append = target_dims - x.ndim
|
||||
if dims_to_append < 0:
|
||||
raise ValueError(f"input has {x.ndim} dims but target_dims is {target_dims}, which is less")
|
||||
return x[(...,) + (None,) * dims_to_append]
|
||||
|
||||
|
||||
# Copied from diffusers.pipelines.animatediff.pipeline_animatediff.tensor2vid
|
||||
# def tensor2vid(video: torch.Tensor, processor: "VaeImageProcessor", output_type: str = "np"):
|
||||
# batch_size, channels, num_frames, height, width = video.shape
|
||||
# outputs = []
|
||||
# for batch_idx in range(batch_size):
|
||||
# batch_vid = video[batch_idx].permute(1, 0, 2, 3)
|
||||
# batch_output = processor.postprocess(batch_vid, output_type)
|
||||
#
|
||||
# outputs.append(batch_output)
|
||||
#
|
||||
# if output_type == "np":
|
||||
# outputs = np.stack(outputs)
|
||||
#
|
||||
# elif output_type == "pt":
|
||||
# outputs = torch.stack(outputs)
|
||||
#
|
||||
# elif not output_type == "pil":
|
||||
# raise ValueError(f"{output_type} does not exist. Please choose one of ['np', 'pt', 'pil]")
|
||||
#
|
||||
# return outputs
|
||||
|
||||
def tensor2vid(video: torch.Tensor, processor: "VaeImageProcessor", output_type: str = "np"):
|
||||
batch_size, channels, num_frames, height, width = video.shape
|
||||
outputs = []
|
||||
for batch_idx in range(batch_size):
|
||||
batch_vid = video[batch_idx].permute(1, 0, 2, 3)
|
||||
batch_output = processor.postprocess(batch_vid, output_type)
|
||||
outputs.append(batch_output)
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
@dataclass
|
||||
class ValidationAnimationPipelineOutput(BaseOutput):
|
||||
r"""
|
||||
Output class for mimicmotion pipeline.
|
||||
|
||||
Args:
|
||||
frames (`[List[List[PIL.Image.Image]]`, `np.ndarray`, `torch.Tensor`]):
|
||||
List of denoised PIL images of length `batch_size` or numpy array or torch tensor of shape `(batch_size,
|
||||
num_frames, height, width, num_channels)`.
|
||||
"""
|
||||
|
||||
frames: Union[List[List[PIL.Image.Image]], np.ndarray, torch.Tensor]
|
||||
|
||||
|
||||
class ValidationAnimationPipeline(DiffusionPipeline):
|
||||
r"""
|
||||
Pipeline to generate video from an input image using Stable Video Diffusion.
|
||||
|
||||
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods
|
||||
implemented for all pipelines (downloading, saving, running on a particular device, etc.).
|
||||
|
||||
Args:
|
||||
vae ([`AutoencoderKLTemporalDecoder`]):
|
||||
Variational Auto-Encoder (VAE) model to encode and decode images to and from latent representations.
|
||||
image_encoder ([`~transformers.CLIPVisionModelWithProjection`]):
|
||||
Frozen CLIP image-encoder ([laion/CLIP-ViT-H-14-laion2B-s32B-b79K]
|
||||
(https://huggingface.co/laion/CLIP-ViT-H-14-laion2B-s32B-b79K)).
|
||||
unet ([`UNetSpatioTemporalConditionModel`]):
|
||||
A `UNetSpatioTemporalConditionModel` to denoise the encoded image latents.
|
||||
scheduler ([`EulerDiscreteScheduler`]):
|
||||
A scheduler to be used in combination with `unet` to denoise the encoded image latents.
|
||||
feature_extractor ([`~transformers.CLIPImageProcessor`]):
|
||||
A `CLIPImageProcessor` to extract features from generated images.
|
||||
pose_net ([`PoseNet`]):
|
||||
A `` to inject pose signals into unet.
|
||||
"""
|
||||
|
||||
model_cpu_offload_seq = "image_encoder->unet->vae"
|
||||
_callback_tensor_inputs = ["latents"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vae,
|
||||
image_encoder,
|
||||
unet,
|
||||
scheduler,
|
||||
feature_extractor,
|
||||
pose_net,
|
||||
face_encoder,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
image_encoder=image_encoder,
|
||||
unet=unet,
|
||||
scheduler=scheduler,
|
||||
feature_extractor=feature_extractor,
|
||||
pose_net=pose_net,
|
||||
face_encoder=face_encoder,
|
||||
)
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
|
||||
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
|
||||
|
||||
self.num_tokens = 4
|
||||
|
||||
# self.app = FaceAnalysis(name="buffalo_l", providers=['CUDAExecutionProvider', 'CPUExecutionProvider'])
|
||||
# self.app.prepare(ctx_id=0, det_size=(640, 640))
|
||||
# self.lora_rank = 128
|
||||
# self.set_ip_adapter()
|
||||
|
||||
def get_prepare_faceid(self, face_image):
|
||||
faceid_image = np.array(face_image)
|
||||
faces = self.app.get(faceid_image)
|
||||
if faces == []:
|
||||
faceid_embeds = torch.zeros_like(torch.empty((1, 512)))
|
||||
else:
|
||||
faceid_embeds = torch.from_numpy(faces[0].normed_embedding).unsqueeze(0)
|
||||
return faceid_embeds
|
||||
|
||||
def set_ip_adapter(self):
|
||||
unet = self.unet
|
||||
attn_procs = {}
|
||||
for name in unet.attn_processors.keys():
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") else unet.config.cross_attention_dim
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = unet.config.block_out_channels[-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(unet.config.block_out_channels))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = unet.config.block_out_channels[block_id]
|
||||
if cross_attention_dim is None:
|
||||
attn_procs[name] = AnimationAttnProcessor(
|
||||
hidden_size=hidden_size, cross_attention_dim=cross_attention_dim, rank=self.lora_rank,
|
||||
).to(self.device, dtype=self.torch_dtype)
|
||||
else:
|
||||
attn_procs[name] = AnimationIDAttnProcessor(
|
||||
hidden_size=hidden_size, cross_attention_dim=cross_attention_dim, scale=1.0, rank=self.lora_rank,
|
||||
num_tokens=self.num_tokens,
|
||||
).to(self.device, dtype=self.torch_dtype)
|
||||
|
||||
unet.set_attn_processor(attn_procs)
|
||||
|
||||
def _encode_image(
|
||||
self,
|
||||
image: PipelineImageInput,
|
||||
device: Union[str, torch.device],
|
||||
num_videos_per_prompt: int,
|
||||
do_classifier_free_guidance: bool):
|
||||
dtype = next(self.image_encoder.parameters()).dtype
|
||||
|
||||
if not isinstance(image, torch.Tensor):
|
||||
image = self.image_processor.pil_to_numpy(image)
|
||||
image = self.image_processor.numpy_to_pt(image)
|
||||
|
||||
# We normalize the image before resizing to match with the original implementation.
|
||||
# Then we unnormalize it after resizing.
|
||||
image = image * 2.0 - 1.0
|
||||
image = _resize_with_antialiasing(image, (224, 224))
|
||||
image = (image + 1.0) / 2.0
|
||||
|
||||
# Normalize the image with for CLIP input
|
||||
image = self.feature_extractor(
|
||||
images=image,
|
||||
do_normalize=True,
|
||||
do_center_crop=False,
|
||||
do_resize=False,
|
||||
do_rescale=False,
|
||||
return_tensors="pt",
|
||||
).pixel_values
|
||||
|
||||
image = image.to(device=device, dtype=dtype)
|
||||
image_embeddings = self.image_encoder(image).image_embeds
|
||||
image_embeddings = image_embeddings.unsqueeze(1)
|
||||
|
||||
# duplicate image embeddings for each generation per prompt, using mps friendly method
|
||||
bs_embed, seq_len, _ = image_embeddings.shape
|
||||
image_embeddings = image_embeddings.repeat(1, num_videos_per_prompt, 1)
|
||||
image_embeddings = image_embeddings.view(bs_embed * num_videos_per_prompt, seq_len, -1)
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
negative_image_embeddings = torch.zeros_like(image_embeddings)
|
||||
|
||||
# For classifier free guidance, we need to do two forward passes.
|
||||
# Here we concatenate the unconditional and text embeddings into a single batch
|
||||
# to avoid doing two forward passes
|
||||
image_embeddings = torch.cat([negative_image_embeddings, image_embeddings])
|
||||
|
||||
return image_embeddings
|
||||
|
||||
def _encode_vae_image(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
device: Union[str, torch.device],
|
||||
num_videos_per_prompt: int,
|
||||
do_classifier_free_guidance: bool,
|
||||
):
|
||||
image = image.to(device=device, dtype=self.vae.dtype)
|
||||
image_latents = self.vae.encode(image).latent_dist.mode()
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
negative_image_latents = torch.zeros_like(image_latents)
|
||||
|
||||
# For classifier free guidance, we need to do two forward passes.
|
||||
# Here we concatenate the unconditional and text embeddings into a single batch
|
||||
# to avoid doing two forward passes
|
||||
image_latents = torch.cat([negative_image_latents, image_latents])
|
||||
|
||||
# duplicate image_latents for each generation per prompt, using mps friendly method
|
||||
image_latents = image_latents.repeat(num_videos_per_prompt, 1, 1, 1)
|
||||
|
||||
return image_latents
|
||||
|
||||
def _get_add_time_ids(
|
||||
self,
|
||||
fps: int,
|
||||
motion_bucket_id: int,
|
||||
noise_aug_strength: float,
|
||||
dtype: torch.dtype,
|
||||
batch_size: int,
|
||||
num_videos_per_prompt: int,
|
||||
do_classifier_free_guidance: bool,
|
||||
):
|
||||
add_time_ids = [fps, motion_bucket_id, noise_aug_strength]
|
||||
|
||||
passed_add_embed_dim = self.unet.config.addition_time_embed_dim * len(add_time_ids)
|
||||
expected_add_embed_dim = self.unet.add_embedding.linear_1.in_features
|
||||
|
||||
if expected_add_embed_dim != passed_add_embed_dim:
|
||||
raise ValueError(
|
||||
f"Model expects an added time embedding vector of length {expected_add_embed_dim}, " \
|
||||
f"but a vector of {passed_add_embed_dim} was created. The model has an incorrect config. " \
|
||||
f"Please check `unet.config.time_embedding_type` and `text_encoder_2.config.projection_dim`."
|
||||
)
|
||||
|
||||
add_time_ids = torch.tensor([add_time_ids], dtype=dtype)
|
||||
add_time_ids = add_time_ids.repeat(batch_size * num_videos_per_prompt, 1)
|
||||
|
||||
if do_classifier_free_guidance:
|
||||
add_time_ids = torch.cat([add_time_ids, add_time_ids])
|
||||
|
||||
return add_time_ids
|
||||
|
||||
def decode_latents(
|
||||
self,
|
||||
latents: torch.Tensor,
|
||||
num_frames: int,
|
||||
decode_chunk_size: int = 8):
|
||||
# [batch, frames, channels, height, width] -> [batch*frames, channels, height, width]
|
||||
latents = latents.flatten(0, 1)
|
||||
|
||||
latents = 1 / self.vae.config.scaling_factor * latents
|
||||
|
||||
forward_vae_fn = self.vae._orig_mod.forward if is_compiled_module(self.vae) else self.vae.forward
|
||||
accepts_num_frames = "num_frames" in set(inspect.signature(forward_vae_fn).parameters.keys())
|
||||
|
||||
# decode decode_chunk_size frames at a time to avoid OOM
|
||||
frames = []
|
||||
for i in range(0, latents.shape[0], decode_chunk_size):
|
||||
num_frames_in = latents[i: i + decode_chunk_size].shape[0]
|
||||
decode_kwargs = {}
|
||||
if accepts_num_frames:
|
||||
# we only pass num_frames_in if it's expected
|
||||
decode_kwargs["num_frames"] = num_frames_in
|
||||
|
||||
frame = self.vae.decode(latents[i: i + decode_chunk_size], **decode_kwargs).sample
|
||||
frames.append(frame.cpu())
|
||||
frames = torch.cat(frames, dim=0)
|
||||
|
||||
# [batch*frames, channels, height, width] -> [batch, channels, frames, height, width]
|
||||
frames = frames.reshape(-1, num_frames, *frames.shape[1:]).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloat16
|
||||
frames = frames.float()
|
||||
return frames
|
||||
|
||||
def check_inputs(self, image, height, width):
|
||||
if (
|
||||
not isinstance(image, torch.Tensor)
|
||||
and not isinstance(image, PIL.Image.Image)
|
||||
and not isinstance(image, list)
|
||||
):
|
||||
raise ValueError(
|
||||
"`image` has to be of type `torch.FloatTensor` or `PIL.Image.Image` or `List[PIL.Image.Image]` but is"
|
||||
f" {type(image)}"
|
||||
)
|
||||
|
||||
if height % 8 != 0 or width % 8 != 0:
|
||||
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
batch_size: int,
|
||||
num_frames: int,
|
||||
num_channels_latents: int,
|
||||
height: int,
|
||||
width: int,
|
||||
dtype: torch.dtype,
|
||||
device: Union[str, torch.device],
|
||||
generator: torch.Generator,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
):
|
||||
shape = (
|
||||
batch_size,
|
||||
num_frames,
|
||||
num_channels_latents // 2,
|
||||
height // self.vae_scale_factor,
|
||||
width // self.vae_scale_factor,
|
||||
)
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
|
||||
if latents is None:
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
else:
|
||||
latents = latents.to(device)
|
||||
|
||||
# scale the initial noise by the standard deviation required by the scheduler
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
return latents
|
||||
|
||||
@property
|
||||
def guidance_scale(self):
|
||||
return self._guidance_scale
|
||||
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
@property
|
||||
def do_classifier_free_guidance(self):
|
||||
if isinstance(self.guidance_scale, (int, float)):
|
||||
return self.guidance_scale > 1
|
||||
return self.guidance_scale.max() > 1
|
||||
|
||||
@property
|
||||
def num_timesteps(self):
|
||||
return self._num_timesteps
|
||||
|
||||
def prepare_extra_step_kwargs(self, generator, eta):
|
||||
# prepare extra kwargs for the scheduler step, since not all schedulers have the same signature
|
||||
# eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers.
|
||||
# eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502
|
||||
# and should be between [0, 1]
|
||||
|
||||
accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||
extra_step_kwargs = {}
|
||||
if accepts_eta:
|
||||
extra_step_kwargs["eta"] = eta
|
||||
|
||||
# check if the scheduler accepts generator
|
||||
accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||
if accepts_generator:
|
||||
extra_step_kwargs["generator"] = generator
|
||||
return extra_step_kwargs
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
image: Union[PIL.Image.Image, List[PIL.Image.Image], torch.FloatTensor],
|
||||
image_pose: Union[torch.FloatTensor],
|
||||
height: int = 576,
|
||||
width: int = 1024,
|
||||
num_frames: Optional[int] = None,
|
||||
tile_size: Optional[int] = 16,
|
||||
tile_overlap: Optional[int] = 4,
|
||||
num_inference_steps: int = 25,
|
||||
min_guidance_scale: float = 1.0,
|
||||
max_guidance_scale: float = 3.0,
|
||||
fps: int = 7,
|
||||
motion_bucket_id: int = 127,
|
||||
noise_aug_strength: float = 0.02,
|
||||
image_only_indicator: bool = False,
|
||||
decode_chunk_size: Optional[int] = None,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.FloatTensor] = None,
|
||||
validation_image_id_ante_embedding=None,
|
||||
output_type: Optional[str] = "pil",
|
||||
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
return_dict: bool = True,
|
||||
):
|
||||
r"""
|
||||
The call function to the pipeline for generation.
|
||||
|
||||
Args:
|
||||
image (`PIL.Image.Image` or `List[PIL.Image.Image]` or `torch.FloatTensor`):
|
||||
Image or images to guide image generation. If you provide a tensor, it needs to be compatible with
|
||||
[`CLIPImageProcessor`](https://huggingface.co/lambdalabs/sd-image-variations-diffusers/blob/main/
|
||||
feature_extractor/preprocessor_config.json).
|
||||
height (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`):
|
||||
The height in pixels of the generated image.
|
||||
width (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`):
|
||||
The width in pixels of the generated image.
|
||||
num_frames (`int`, *optional*):
|
||||
The number of video frames to generate. Defaults to 14 for `stable-video-diffusion-img2vid`
|
||||
and to 25 for `stable-video-diffusion-img2vid-xt`
|
||||
num_inference_steps (`int`, *optional*, defaults to 25):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference. This parameter is modulated by `strength`.
|
||||
min_guidance_scale (`float`, *optional*, defaults to 1.0):
|
||||
The minimum guidance scale. Used for the classifier free guidance with first frame.
|
||||
max_guidance_scale (`float`, *optional*, defaults to 3.0):
|
||||
The maximum guidance scale. Used for the classifier free guidance with last frame.
|
||||
fps (`int`, *optional*, defaults to 7):
|
||||
Frames per second.The rate at which the generated images shall be exported to a video after generation.
|
||||
Note that Stable Diffusion Video's UNet was micro-conditioned on fps-1 during training.
|
||||
motion_bucket_id (`int`, *optional*, defaults to 127):
|
||||
The motion bucket ID. Used as conditioning for the generation.
|
||||
The higher the number the more motion will be in the video.
|
||||
noise_aug_strength (`float`, *optional*, defaults to 0.02):
|
||||
The amount of noise added to the init image,
|
||||
the higher it is the less the video will look like the init image. Increase it for more motion.
|
||||
image_only_indicator (`bool`, *optional*, defaults to False):
|
||||
Whether to treat the inputs as batch of images instead of videos.
|
||||
decode_chunk_size (`int`, *optional*):
|
||||
The number of frames to decode at a time.The higher the chunk size, the higher the temporal consistency
|
||||
between frames, but also the higher the memory consumption.
|
||||
By default, the decoder will decode all frames at once for maximal quality.
|
||||
Reduce `decode_chunk_size` to reduce memory usage.
|
||||
num_videos_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
|
||||
generation deterministic.
|
||||
latents (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor is generated by sampling using the supplied random `generator`.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
|
||||
callback_on_step_end (`Callable`, *optional*):
|
||||
A function that calls at the end of each denoising steps during the inference. The function is called
|
||||
with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
|
||||
callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
|
||||
`callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeline class.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether to return a [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] instead of a
|
||||
plain tuple.
|
||||
device:
|
||||
On which device the pipeline runs on.
|
||||
|
||||
Returns:
|
||||
[`~pipelines.stable_diffusion.StableVideoDiffusionPipelineOutput`] or `tuple`:
|
||||
If `return_dict` is `True`,
|
||||
[`~pipelines.stable_diffusion.StableVideoDiffusionPipelineOutput`] is returned,
|
||||
otherwise a `tuple` is returned where the first element is a list of list with the generated frames.
|
||||
|
||||
Examples:
|
||||
|
||||
```py
|
||||
from diffusers import StableVideoDiffusionPipeline
|
||||
from diffusers.utils import load_image, export_to_video
|
||||
|
||||
pipe = StableVideoDiffusionPipeline.from_pretrained(
|
||||
"stabilityai/stable-video-diffusion-img2vid-xt", torch_dtype=torch.float16, variant="fp16")
|
||||
pipe.to("cuda")
|
||||
|
||||
image = load_image(
|
||||
"https://lh3.googleusercontent.com/y-iFOHfLTwkuQSUegpwDdgKmOjRSTvPxat63dQLB25xkTs4lhIbRUFeNBWZzYf370g=s1200")
|
||||
image = image.resize((1024, 576))
|
||||
|
||||
frames = pipe(image, num_frames=25, decode_chunk_size=8).frames[0]
|
||||
export_to_video(frames, "generated.mp4", fps=7)
|
||||
```
|
||||
"""
|
||||
# 0. Default height and width to unet
|
||||
height = height or self.unet.config.sample_size * self.vae_scale_factor
|
||||
width = width or self.unet.config.sample_size * self.vae_scale_factor
|
||||
|
||||
num_frames = num_frames if num_frames is not None else self.unet.config.num_frames
|
||||
decode_chunk_size = decode_chunk_size if decode_chunk_size is not None else num_frames
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(image, height, width)
|
||||
|
||||
# 2. Define call parameters
|
||||
if isinstance(image, PIL.Image.Image):
|
||||
batch_size = 1
|
||||
elif isinstance(image, list):
|
||||
batch_size = len(image)
|
||||
else:
|
||||
batch_size = image.shape[0]
|
||||
device = self._execution_device
|
||||
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||
# corresponds to doing no classifier free guidance.
|
||||
do_classifier_free_guidance = max_guidance_scale >= 1.0
|
||||
self._guidance_scale = max_guidance_scale
|
||||
|
||||
# 3. Encode input image
|
||||
image_embeddings = self._encode_image(image, device, num_videos_per_prompt, do_classifier_free_guidance)
|
||||
# self.image_encoder.cpu()
|
||||
|
||||
# NOTE: Stable Diffusion Video was conditioned on fps - 1, which
|
||||
# is why it is reduced here.
|
||||
fps = fps - 1
|
||||
|
||||
# 4. Encode input image using VAE
|
||||
|
||||
# print(image_embeddings.size()) # [2, 1, 1024]
|
||||
validation_image_id_ante_embedding = torch.from_numpy(validation_image_id_ante_embedding).unsqueeze(0)
|
||||
# print(device)
|
||||
validation_image_id_ante_embedding = validation_image_id_ante_embedding.to(device)
|
||||
faceid_latents = self.face_encoder(validation_image_id_ante_embedding, image_embeddings[1:])
|
||||
# print(faceid_latents.size()) # [1, 4, 1024]
|
||||
uncond_image_embeddings = image_embeddings[:1]
|
||||
uncond_faceid_latents = torch.zeros_like(faceid_latents)
|
||||
uncond_image_embeddings = torch.cat([uncond_image_embeddings, uncond_faceid_latents], dim=1)
|
||||
cond_image_embeddings = image_embeddings[1:]
|
||||
cond_image_embeddings = torch.cat([cond_image_embeddings, faceid_latents], dim=1)
|
||||
image_embeddings = torch.cat([uncond_image_embeddings, cond_image_embeddings])
|
||||
|
||||
image = self.image_processor.preprocess(image, height=height, width=width).to(device)
|
||||
noise = randn_tensor(image.shape, generator=generator, device=device, dtype=image.dtype)
|
||||
image = image + noise_aug_strength * noise
|
||||
|
||||
needs_upcasting = (self.vae.dtype == torch.float16 or self.vae.dtype == torch.bfloat16) and self.vae.config.force_upcast
|
||||
if needs_upcasting:
|
||||
self_vae_dtype = self.vae.dtype
|
||||
self.vae.to(dtype=torch.float32)
|
||||
|
||||
image_latents = self._encode_vae_image(
|
||||
image,
|
||||
device=device,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
do_classifier_free_guidance=do_classifier_free_guidance,
|
||||
)
|
||||
image_latents = image_latents.to(image_embeddings.dtype)
|
||||
|
||||
if needs_upcasting:
|
||||
self.vae.to(dtype=self_vae_dtype)
|
||||
# self.vae.cpu()
|
||||
|
||||
# Repeat the image latents for each frame so we can concatenate them with the noise
|
||||
# image_latents [batch, channels, height, width] ->[batch, num_frames, channels, height, width]
|
||||
image_latents = image_latents.unsqueeze(1).repeat(1, num_frames, 1, 1, 1)
|
||||
|
||||
# 5. Get Added Time IDs
|
||||
added_time_ids = self._get_add_time_ids(
|
||||
fps,
|
||||
motion_bucket_id,
|
||||
noise_aug_strength,
|
||||
image_embeddings.dtype,
|
||||
batch_size,
|
||||
num_videos_per_prompt,
|
||||
self.do_classifier_free_guidance,
|
||||
)
|
||||
added_time_ids = added_time_ids.to(device)
|
||||
|
||||
# 4. Prepare timesteps
|
||||
timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, None)
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = self.unet.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
tile_size,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
image_embeddings.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
latents = latents.repeat(1, num_frames // tile_size + 1, 1, 1, 1)[:, :num_frames]
|
||||
|
||||
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
|
||||
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, 0.0)
|
||||
|
||||
# 7. Prepare guidance scale
|
||||
guidance_scale = torch.linspace(min_guidance_scale, max_guidance_scale, num_frames).unsqueeze(0)
|
||||
guidance_scale = guidance_scale.to(device, latents.dtype)
|
||||
guidance_scale = guidance_scale.repeat(batch_size * num_videos_per_prompt, 1)
|
||||
guidance_scale = _append_dims(guidance_scale, latents.ndim)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
|
||||
# 8. Denoising loop
|
||||
self._num_timesteps = len(timesteps)
|
||||
indices = [[0, *range(i + 1, min(i + tile_size, num_frames))] for i in
|
||||
range(0, num_frames - tile_size + 1, tile_size - tile_overlap)]
|
||||
if indices[-1][-1] < num_frames - 1:
|
||||
indices.append([0, *range(num_frames - tile_size + 1, num_frames)])
|
||||
|
||||
pose_pil_image_list = []
|
||||
for pose in image_pose:
|
||||
pose = torch.from_numpy(np.array(pose)).float()
|
||||
pose = pose / 127.5 - 1
|
||||
pose_pil_image_list.append(pose)
|
||||
pose_pil_image_list = torch.stack(pose_pil_image_list, dim=0)
|
||||
pose_pil_image_list = rearrange(pose_pil_image_list, "f h w c -> f c h w")
|
||||
|
||||
# print(indices) # [[0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15]]
|
||||
# print(pose_pil_image_list.size()) # [16, 3, 512, 512]
|
||||
|
||||
self.pose_net.to(device)
|
||||
self.unet.to(device)
|
||||
|
||||
with torch.cuda.device(device):
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
with self.progress_bar(total=len(timesteps) * len(indices)) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
|
||||
# Concatenate image_latents over channels dimension
|
||||
latent_model_input = torch.cat([latent_model_input, image_latents], dim=2)
|
||||
|
||||
# predict the noise residual
|
||||
noise_pred = torch.zeros_like(image_latents)
|
||||
noise_pred_cnt = image_latents.new_zeros((num_frames,))
|
||||
weight = (torch.arange(tile_size, device=device) + 0.5) * 2. / tile_size
|
||||
weight = torch.minimum(weight, 2 - weight)
|
||||
for idx in indices:
|
||||
# classification-free inference
|
||||
pose_latents = self.pose_net(pose_pil_image_list[idx].to(device))
|
||||
_noise_pred = self.unet(
|
||||
latent_model_input[:1, idx],
|
||||
t,
|
||||
encoder_hidden_states=image_embeddings[:1],
|
||||
added_time_ids=added_time_ids[:1],
|
||||
pose_latents=None,
|
||||
image_only_indicator=image_only_indicator,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred[:1, idx] += _noise_pred * weight[:, None, None, None]
|
||||
|
||||
# normal inference
|
||||
_noise_pred = self.unet(
|
||||
latent_model_input[1:, idx],
|
||||
t,
|
||||
encoder_hidden_states=image_embeddings[1:],
|
||||
added_time_ids=added_time_ids[1:],
|
||||
pose_latents=pose_latents,
|
||||
image_only_indicator=image_only_indicator,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred[1:, idx] += _noise_pred * weight[:, None, None, None]
|
||||
|
||||
noise_pred_cnt[idx] += weight
|
||||
progress_bar.update()
|
||||
noise_pred.div_(noise_pred_cnt[:, None, None, None])
|
||||
|
||||
# perform guidance
|
||||
if self.do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_cond - noise_pred_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
for k in callback_on_step_end_tensor_inputs:
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
|
||||
# self.pose_net.cpu()
|
||||
# self.unet.cpu()
|
||||
# self.face_encoder.cpu()
|
||||
|
||||
if not output_type == "latent":
|
||||
self.vae.decoder.to(device)
|
||||
frames = self.decode_latents(latents, num_frames, decode_chunk_size)
|
||||
# print(frames.size()) # [1, 3, 16, 512, 512]
|
||||
# print(latents.size()) # [1, 16, 4, 64, 64]
|
||||
frames = tensor2vid(frames, self.image_processor, output_type=output_type)
|
||||
# print(frames[0].size()) # [16, 3, 512, 512]
|
||||
else:
|
||||
frames = latents
|
||||
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return frames
|
||||
|
||||
return ValidationAnimationPipelineOutput(frames=frames)
|
||||
@@ -0,0 +1,9 @@
|
||||
import diffusers.models.activations
|
||||
|
||||
|
||||
def patch_geglu_inplace():
|
||||
"""Patch GEGLU with inplace multiplication to save GPU memory."""
|
||||
def forward(self, hidden_states):
|
||||
hidden_states, gate = self.proj(hidden_states).chunk(2, dim=-1)
|
||||
return hidden_states.mul_(self.gelu(gate))
|
||||
diffusers.models.activations.GEGLU.forward = forward
|
||||
@@ -0,0 +1,53 @@
|
||||
import logging
|
||||
|
||||
import torch
|
||||
import torch.utils.checkpoint
|
||||
from diffusers.models import AutoencoderKLTemporalDecoder
|
||||
from diffusers.schedulers import EulerDiscreteScheduler
|
||||
from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection
|
||||
|
||||
from ..modules.unet import UNetSpatioTemporalConditionModel
|
||||
from ..modules.pose_net import PoseNet
|
||||
from ..pipelines.pipeline_animation import MimicMotionPipeline
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
class MimicMotionModel(torch.nn.Module):
|
||||
def __init__(self, base_model_path):
|
||||
"""construnct base model components and load pretrained svd model except pose-net
|
||||
Args:
|
||||
base_model_path (str): pretrained svd model path
|
||||
"""
|
||||
super().__init__()
|
||||
self.unet = UNetSpatioTemporalConditionModel.from_config(
|
||||
UNetSpatioTemporalConditionModel.load_config(base_model_path, subfolder="unet"))
|
||||
self.vae = AutoencoderKLTemporalDecoder.from_pretrained(
|
||||
base_model_path, subfolder="vae", torch_dtype=torch.float16, variant="fp16")
|
||||
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(
|
||||
base_model_path, subfolder="image_encoder", torch_dtype=torch.float16, variant="fp16")
|
||||
self.noise_scheduler = EulerDiscreteScheduler.from_pretrained(
|
||||
base_model_path, subfolder="scheduler")
|
||||
self.feature_extractor = CLIPImageProcessor.from_pretrained(
|
||||
base_model_path, subfolder="feature_extractor")
|
||||
# pose_net
|
||||
self.pose_net = PoseNet(noise_latent_channels=self.unet.config.block_out_channels[0])
|
||||
|
||||
def create_pipeline(infer_config, device):
|
||||
"""create mimicmotion pipeline and load pretrained weight
|
||||
|
||||
Args:
|
||||
infer_config (str):
|
||||
device (str or torch.device): "cpu" or "cuda:{device_id}"
|
||||
"""
|
||||
mimicmotion_models = MimicMotionModel(infer_config.base_model_path)
|
||||
mimicmotion_models.load_state_dict(torch.load(infer_config.ckpt_path, map_location="cpu"), strict=False)
|
||||
pipeline = MimicMotionPipeline(
|
||||
vae=mimicmotion_models.vae,
|
||||
image_encoder=mimicmotion_models.image_encoder,
|
||||
unet=mimicmotion_models.unet,
|
||||
scheduler=mimicmotion_models.noise_scheduler,
|
||||
feature_extractor=mimicmotion_models.feature_extractor,
|
||||
pose_net=mimicmotion_models.pose_net
|
||||
)
|
||||
return pipeline
|
||||
|
||||
@@ -0,0 +1,107 @@
|
||||
import logging
|
||||
from pathlib import Path
|
||||
import torch
|
||||
from requests.packages import target
|
||||
from torchvision.io import write_video
|
||||
from diffusers.utils.torch_utils import is_compiled_module
|
||||
import inspect
|
||||
import numpy as np
|
||||
from torchvision.transforms import Compose, ToTensor, Normalize
|
||||
import accelerate
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
def decode_latents(
|
||||
vae,
|
||||
latents,
|
||||
num_frames,
|
||||
decode_chunk_size=8):
|
||||
# [batch, frames, channels, height, width] -> [batch*frames, channels, height, width]
|
||||
latents = latents.flatten(0, 1)
|
||||
|
||||
latents = 1 / vae.config.scaling_factor * latents
|
||||
|
||||
forward_vae_fn = vae._orig_mod.forward if is_compiled_module(vae) else vae.forward
|
||||
accepts_num_frames = "num_frames" in set(inspect.signature(forward_vae_fn).parameters.keys())
|
||||
|
||||
# decode decode_chunk_size frames at a time to avoid OOM
|
||||
frames = []
|
||||
for i in range(0, latents.shape[0], decode_chunk_size):
|
||||
num_frames_in = latents[i: i + decode_chunk_size].shape[0]
|
||||
decode_kwargs = {}
|
||||
if accepts_num_frames:
|
||||
# we only pass num_frames_in if it's expected
|
||||
decode_kwargs["num_frames"] = num_frames_in
|
||||
|
||||
frame = vae.decode(latents[i: i + decode_chunk_size], **decode_kwargs).sample
|
||||
frames.append(frame.cpu())
|
||||
frames = torch.cat(frames, dim=0)
|
||||
|
||||
# [batch*frames, channels, height, width] -> [batch, channels, frames, height, width]
|
||||
frames = frames.reshape(-1, num_frames, *frames.shape[1:]).permute(0, 2, 1, 3, 4)
|
||||
|
||||
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloat16
|
||||
frames = frames.float()
|
||||
return frames
|
||||
|
||||
def tensor2vid(video, processor, output_type="np"):
|
||||
batch_size, channels, num_frames, height, width = video.shape
|
||||
outputs = []
|
||||
for batch_idx in range(batch_size):
|
||||
batch_vid = video[batch_idx].permute(1, 0, 2, 3)
|
||||
batch_output = processor.postprocess(batch_vid, output_type)
|
||||
|
||||
outputs.append(batch_output)
|
||||
|
||||
if output_type == "np":
|
||||
outputs = np.stack(outputs)
|
||||
|
||||
elif output_type == "pt":
|
||||
outputs = torch.stack(outputs)
|
||||
|
||||
elif not output_type == "pil":
|
||||
# raise ValueError(f"{output_type} does not exist. Please choose one of ['np', 'pt', 'pil]")
|
||||
return outputs
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
def get_aligned_face(face_loss_helper, frames, device):
|
||||
print("This is the process of getting aligned faces")
|
||||
print("Please check the number of detected faces")
|
||||
print(frames.shape)
|
||||
print(1/0)
|
||||
face_loss_helper.clean_all()
|
||||
faces = mtcnn.align(frames)
|
||||
transfroms = Compose(
|
||||
[ToTensor(), Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5])])
|
||||
return transfroms(faces).to(device)
|
||||
|
||||
|
||||
def save_to_mp4(frames, save_path, fps=7):
|
||||
frames = frames.permute((0, 2, 3, 1)) # (f, c, h, w) to (f, h, w, c)
|
||||
Path(save_path).parent.mkdir(parents=True, exist_ok=True)
|
||||
write_video(save_path, frames, fps=fps)
|
||||
|
||||
def faceid_loss_compute(vae, latents, target_images, num_frames, decode_chunk_size, image_processor, device, face_loss_model=None, face_loss_helper=None):
|
||||
|
||||
print("--------------------------")
|
||||
while True:
|
||||
x = 1 + 1
|
||||
frames = decode_latents(vae, latents, num_frames, decode_chunk_size)
|
||||
frames = tensor2vid(frames, image_processor, output_type="np")
|
||||
print("This is faceid loss computation")
|
||||
print(frames.shape)
|
||||
print(type(target_images))
|
||||
print(target_images.size())
|
||||
print(1/0)
|
||||
|
||||
pred_faces = get_aligned_face(face_loss_helper, frames, device)
|
||||
target_faces = get_aligned_face(face_loss_helper, target_images, device)
|
||||
|
||||
pred_embed = face_loss_model(pred_faces)[0]
|
||||
target_embed = face_loss_model(target_faces)[0]
|
||||
|
||||
face_loss = pred_embed.dot(target_embed).item()
|
||||
face_loss = face_loss.mean()
|
||||
return face_loss
|
||||
@@ -0,0 +1,349 @@
|
||||
import os
|
||||
import cv2
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from diffusers.models.attention_processor import XFormersAttnProcessor
|
||||
from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection
|
||||
import torch
|
||||
from diffusers import AutoencoderKLTemporalDecoder, EulerDiscreteScheduler
|
||||
|
||||
from animation.modules.attention_processor import AnimationAttnProcessor
|
||||
from animation.modules.attention_processor_normalized import AnimationIDAttnNormalizedProcessor
|
||||
from animation.modules.face_model import FaceModel
|
||||
from animation.modules.id_encoder import FusionFaceId
|
||||
from animation.modules.pose_net import PoseNet
|
||||
from animation.modules.unet import UNetSpatioTemporalConditionModel
|
||||
from animation.pipelines.inference_pipeline_animation import InferenceAnimationPipeline
|
||||
import random
|
||||
|
||||
import gradio as gr
|
||||
import gc
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
pretrained_model_name_or_path = "checkpoints/stable-video-diffusion-img2vid-xt"
|
||||
revision = None
|
||||
posenet_model_name_or_path = "checkpoints/Animation/pose_net.pth"
|
||||
face_encoder_model_name_or_path = "checkpoints/Animation/face_encoder.pth"
|
||||
unet_model_name_or_path = "checkpoints/Animation/unet.pth"
|
||||
|
||||
|
||||
def load_images_from_folder(folder, width, height):
|
||||
images = []
|
||||
files = os.listdir(folder)
|
||||
png_files = [f for f in files if f.endswith('.png')]
|
||||
png_files.sort(key=lambda x: int(x.split('_')[1].split('.')[0]))
|
||||
for filename in png_files:
|
||||
img = Image.open(os.path.join(folder, filename)).convert('RGB')
|
||||
img = img.resize((width, height))
|
||||
images.append(img)
|
||||
|
||||
return images
|
||||
|
||||
|
||||
def save_frames_as_png(frames, output_path):
|
||||
pil_frames = [Image.fromarray(frame) if isinstance(frame, np.ndarray) else frame for frame in frames]
|
||||
num_frames = len(pil_frames)
|
||||
for i in range(num_frames):
|
||||
pil_frame = pil_frames[i]
|
||||
save_path = os.path.join(output_path, f'frame_{i}.png')
|
||||
pil_frame.save(save_path)
|
||||
|
||||
|
||||
def save_frames_as_mp4(frames, output_mp4_path, fps):
|
||||
print("Starting saving the frames as mp4")
|
||||
height, width, _ = frames[0].shape
|
||||
fourcc = cv2.VideoWriter_fourcc(*'mp4v') # 'H264' for better quality
|
||||
out = cv2.VideoWriter(output_mp4_path, fourcc, fps, (width, height))
|
||||
for frame in frames:
|
||||
frame_bgr = cv2.cvtColor(frame, cv2.COLOR_RGB2BGR)
|
||||
out.write(frame_bgr)
|
||||
out.release()
|
||||
|
||||
|
||||
def export_to_gif(frames, output_gif_path, fps):
|
||||
"""
|
||||
Export a list of frames to a GIF.
|
||||
|
||||
Args:
|
||||
- frames (list): List of frames (as numpy arrays or PIL Image objects).
|
||||
- output_gif_path (str): Path to save the output GIF.
|
||||
- duration_ms (int): Duration of each frame in milliseconds.
|
||||
|
||||
"""
|
||||
# Convert numpy arrays to PIL Images if needed
|
||||
pil_frames = [Image.fromarray(frame) if isinstance(
|
||||
frame, np.ndarray) else frame for frame in frames]
|
||||
|
||||
pil_frames[0].save(output_gif_path.replace('.mp4', '.gif'),
|
||||
format='GIF',
|
||||
append_images=pil_frames[1:],
|
||||
save_all=True,
|
||||
duration=125,
|
||||
loop=0)
|
||||
|
||||
|
||||
def generate(
|
||||
image_input: str,
|
||||
pose_input: str,
|
||||
width: int,
|
||||
height: int,
|
||||
guidance_scale: float,
|
||||
num_inference_steps: int,
|
||||
fps: int,
|
||||
frames_overlap: int,
|
||||
tile_size: int,
|
||||
noise_aug_strength: float,
|
||||
decode_chunk_size: int,
|
||||
seed: int,
|
||||
):
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
output_dir = Path("outputs")
|
||||
output_dir = os.path.join(output_dir, timestamp)
|
||||
if seed == -1:
|
||||
seed = random.randint(1, 2**20 - 1)
|
||||
generator = torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
pipeline = InferenceAnimationPipeline(
|
||||
vae=vae,
|
||||
image_encoder=image_encoder,
|
||||
unet=unet,
|
||||
scheduler=noise_scheduler,
|
||||
feature_extractor=feature_extractor,
|
||||
pose_net=pose_net,
|
||||
face_encoder=face_encoder,
|
||||
).to(device=device, dtype=dtype)
|
||||
|
||||
validation_image_path = image_input
|
||||
validation_image = Image.open(image_input).convert('RGB')
|
||||
validation_control_images = load_images_from_folder(pose_input, width=width, height=height)
|
||||
|
||||
num_frames = len(validation_control_images)
|
||||
face_model.face_helper.clean_all()
|
||||
validation_face = cv2.imread(validation_image_path)
|
||||
validation_image_bgr = cv2.cvtColor(validation_face, cv2.COLOR_RGB2BGR)
|
||||
validation_image_face_info = face_model.app.get(validation_image_bgr)
|
||||
if len(validation_image_face_info) > 0:
|
||||
validation_image_face_info = sorted(validation_image_face_info, key=lambda x: (x['bbox'][2] - x['bbox'][0]) * (x['bbox'][3] - x['bbox'][1]))[-1]
|
||||
validation_image_id_ante_embedding = validation_image_face_info['embedding']
|
||||
else:
|
||||
validation_image_id_ante_embedding = None
|
||||
|
||||
if validation_image_id_ante_embedding is None:
|
||||
face_model.face_helper.read_image(validation_image_bgr)
|
||||
face_model.face_helper.get_face_landmarks_5(only_center_face=True)
|
||||
face_model.face_helper.align_warp_face()
|
||||
|
||||
if len(face_model.face_helper.cropped_faces) == 0:
|
||||
validation_image_id_ante_embedding = np.zeros((512,))
|
||||
else:
|
||||
validation_image_align_face = face_model.face_helper.cropped_faces[0]
|
||||
print('fail to detect face using insightface, extract embedding on align face')
|
||||
validation_image_id_ante_embedding = face_model.handler_ante.get_feat(validation_image_align_face)
|
||||
|
||||
# generator = torch.Generator(device=accelerator.device).manual_seed(23123134)
|
||||
|
||||
decode_chunk_size = decode_chunk_size
|
||||
video_frames = pipeline(
|
||||
image=validation_image,
|
||||
image_pose=validation_control_images,
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
tile_size=tile_size,
|
||||
tile_overlap=frames_overlap,
|
||||
decode_chunk_size=decode_chunk_size,
|
||||
motion_bucket_id=127.,
|
||||
fps=7,
|
||||
min_guidance_scale=guidance_scale,
|
||||
max_guidance_scale=guidance_scale,
|
||||
noise_aug_strength=noise_aug_strength,
|
||||
num_inference_steps=num_inference_steps,
|
||||
generator=generator,
|
||||
output_type="pil",
|
||||
validation_image_id_ante_embedding=validation_image_id_ante_embedding,
|
||||
).frames[0]
|
||||
|
||||
out_file = os.path.join(
|
||||
output_dir,
|
||||
f"animation_video.mp4",
|
||||
)
|
||||
for i in range(num_frames):
|
||||
img = video_frames[i]
|
||||
video_frames[i] = np.array(img)
|
||||
|
||||
png_out_file = os.path.join(output_dir, "animated_images")
|
||||
os.makedirs(png_out_file, exist_ok=True)
|
||||
|
||||
save_frames_as_mp4(video_frames, out_file, fps)
|
||||
export_to_gif(video_frames, out_file, fps)
|
||||
save_frames_as_png(video_frames, png_out_file)
|
||||
|
||||
seed_update = gr.update(visible=True, value=seed)
|
||||
|
||||
return out_file, seed_update
|
||||
|
||||
|
||||
with gr.Blocks(theme=gr.themes.Soft()) as demo:
|
||||
gr.Markdown("""
|
||||
<div>
|
||||
<h2 style="font-size: 30px;text-align: center;">StableAnimator</h2>
|
||||
</div>
|
||||
<div style="text-align: center;">
|
||||
<a href="https://github.com/Francis-Rings/StableAnimator">🌐 Github</a> |
|
||||
<a href="https://arxiv.org/abs/2411.17697">📜 arXiv </a>
|
||||
</div>
|
||||
<div style="text-align: center; font-weight: bold; color: red;">
|
||||
⚠️ This demo is for academic research and experiential use only.
|
||||
</div>
|
||||
""")
|
||||
with gr.Row():
|
||||
with gr.Column():
|
||||
with gr.Group():
|
||||
image_input = gr.Image(label="Reference Image", type="filepath")
|
||||
pose_input = gr.Textbox(label="Driven Poses", placeholder="Please enter your driven pose directory here.")
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
width = gr.Number(label="Width (supports only 512×512 and 576×1024)", value=512)
|
||||
height = gr.Number(label="Height (supports only 512×512 and 576×1024)", value=512)
|
||||
with gr.Row():
|
||||
guidance_scale = gr.Number(label="Guidance scale (recommended 3.0)", value=3.0, step=0.1, precision=1)
|
||||
num_inference_steps = gr.Number(label="Inference steps (recommended 25)", value=20)
|
||||
with gr.Row():
|
||||
fps = gr.Number(label="FPS", value=8)
|
||||
frames_overlap = gr.Number(label="Overlap Frames (recommended 4)", value=4)
|
||||
with gr.Row():
|
||||
tile_size = gr.Number(label="Tile Size (recommended 16)", value=16)
|
||||
noise_aug_strength = gr.Number(label="Noise Augmentation Strength (recommended 0.02)", value=0.02, step=0.01, precision=2)
|
||||
with gr.Row():
|
||||
decode_chunk_size = gr.Number(label="Decode Chunk Size (recommended 4 or 16)", value=4)
|
||||
seed = gr.Number(label="Random Seed (Enter a positive number, -1 for random)", value=-1)
|
||||
generate_button = gr.Button("🎬 Generate The Video")
|
||||
with gr.Column():
|
||||
video_output = gr.Video(label="Generate The Video")
|
||||
with gr.Row():
|
||||
seed_text = gr.Number(label="Video Generation Seed", visible=False, interactive=False)
|
||||
gr.Examples([
|
||||
["inference/case-1/reference.png","inference/case-1/poses",512,512],
|
||||
["inference/case-2/reference.png","inference/case-2/poses",512,512],
|
||||
["inference/case-3/reference.png","inference/case-3/poses",512,512],
|
||||
["inference/case-4/reference.png","inference/case-4/poses",512,512],
|
||||
["inference/case-5/reference.png","inference/case-5/poses",576,1024],
|
||||
], inputs=[image_input, pose_input, width, height])
|
||||
|
||||
|
||||
generate_button.click(
|
||||
generate,
|
||||
inputs=[image_input, pose_input, width, height, guidance_scale, num_inference_steps, fps, frames_overlap, tile_size, noise_aug_strength, decode_chunk_size, seed],
|
||||
outputs=[video_output, seed_text],
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
feature_extractor = CLIPImageProcessor.from_pretrained(pretrained_model_name_or_path, subfolder="feature_extractor", revision=revision)
|
||||
noise_scheduler = EulerDiscreteScheduler.from_pretrained(pretrained_model_name_or_path, subfolder="scheduler")
|
||||
image_encoder = CLIPVisionModelWithProjection.from_pretrained(pretrained_model_name_or_path, subfolder="image_encoder", revision=revision)
|
||||
vae = AutoencoderKLTemporalDecoder.from_pretrained(pretrained_model_name_or_path, subfolder="vae", revision=revision)
|
||||
unet = UNetSpatioTemporalConditionModel.from_pretrained(
|
||||
pretrained_model_name_or_path,
|
||||
subfolder="unet",
|
||||
low_cpu_mem_usage=True,
|
||||
)
|
||||
pose_net = PoseNet(noise_latent_channels=unet.config.block_out_channels[0])
|
||||
face_encoder = FusionFaceId(
|
||||
cross_attention_dim=1024,
|
||||
id_embeddings_dim=512,
|
||||
# clip_embeddings_dim=image_encoder.config.hidden_size,
|
||||
clip_embeddings_dim=1024,
|
||||
num_tokens=4, )
|
||||
face_model = FaceModel()
|
||||
|
||||
lora_rank = 128
|
||||
attn_procs = {}
|
||||
unet_svd = unet.state_dict()
|
||||
|
||||
for name in unet.attn_processors.keys():
|
||||
if "transformer_blocks" in name and "temporal_transformer_blocks" not in name:
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") else unet.config.cross_attention_dim
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = unet.config.block_out_channels[-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(unet.config.block_out_channels))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = unet.config.block_out_channels[block_id]
|
||||
if cross_attention_dim is None:
|
||||
# print(f"This is AnimationAttnProcessor: {name}")
|
||||
attn_procs[name] = AnimationAttnProcessor(hidden_size=hidden_size, cross_attention_dim=cross_attention_dim, rank=lora_rank)
|
||||
else:
|
||||
# print(f"This is AnimationIDAttnProcessor: {name}")
|
||||
layer_name = name.split(".processor")[0]
|
||||
weights = {
|
||||
"to_k_ip.weight": unet_svd[layer_name + ".to_k.weight"],
|
||||
"to_v_ip.weight": unet_svd[layer_name + ".to_v.weight"],
|
||||
}
|
||||
attn_procs[name] = AnimationIDAttnNormalizedProcessor(hidden_size=hidden_size, cross_attention_dim=cross_attention_dim, rank=lora_rank)
|
||||
attn_procs[name].load_state_dict(weights, strict=False)
|
||||
elif "temporal_transformer_blocks" in name:
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") else unet.config.cross_attention_dim
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = unet.config.block_out_channels[-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(unet.config.block_out_channels))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = unet.config.block_out_channels[block_id]
|
||||
if cross_attention_dim is None:
|
||||
attn_procs[name] = XFormersAttnProcessor()
|
||||
else:
|
||||
attn_procs[name] = XFormersAttnProcessor()
|
||||
unet.set_attn_processor(attn_procs)
|
||||
|
||||
# resume the previous checkpoint
|
||||
if posenet_model_name_or_path is not None and face_encoder_model_name_or_path is not None and unet_model_name_or_path is not None:
|
||||
print("Loading existing posenet weights, face_encoder weights and unet weights.")
|
||||
if posenet_model_name_or_path.endswith(".pth"):
|
||||
pose_net_state_dict = torch.load(posenet_model_name_or_path, map_location="cpu")
|
||||
pose_net.load_state_dict(pose_net_state_dict, strict=True)
|
||||
else:
|
||||
print("posenet weights loading fail")
|
||||
print(1/0)
|
||||
if face_encoder_model_name_or_path.endswith(".pth"):
|
||||
face_encoder_state_dict = torch.load(face_encoder_model_name_or_path, map_location="cpu")
|
||||
face_encoder.load_state_dict(face_encoder_state_dict, strict=True)
|
||||
else:
|
||||
print("face_encoder weights loading fail")
|
||||
print(1/0)
|
||||
if unet_model_name_or_path.endswith(".pth"):
|
||||
unet_state_dict = torch.load(unet_model_name_or_path, map_location="cpu")
|
||||
unet.load_state_dict(unet_state_dict, strict=True)
|
||||
else:
|
||||
print("unet weights loading fail")
|
||||
print(1/0)
|
||||
|
||||
vae.requires_grad_(False)
|
||||
image_encoder.requires_grad_(False)
|
||||
unet.requires_grad_(False)
|
||||
pose_net.requires_grad_(False)
|
||||
face_encoder.requires_grad_(False)
|
||||
|
||||
total_vram_in_gb = torch.cuda.get_device_properties(0).total_memory / 1073741824
|
||||
print(f'\033[32mCUDA version:{torch.version.cuda}\033[0m')
|
||||
print(f'\033[32mPytorch version:{torch.__version__}\033[0m')
|
||||
print(f'\033[32mGPU Type:{torch.cuda.get_device_name()}\033[0m')
|
||||
print(f'\033[32mGPU Memory:{total_vram_in_gb:.2f}GB\033[0m')
|
||||
if torch.cuda.get_device_capability()[0] >= 8:
|
||||
print(f'\033[32mSupports BF16, use BF16\033[0m')
|
||||
dtype = torch.bfloat16
|
||||
else:
|
||||
print(f'\033[32mBF16 is not supported, use FP16. The 5B model is not recommended\033[0m')
|
||||
dtype = torch.float16
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
demo.queue()
|
||||
demo.launch(inbrowser=True)
|
||||
|
After Width: | Height: | Size: 1.7 MiB |
|
After Width: | Height: | Size: 21 MiB |
|
After Width: | Height: | Size: 23 MiB |
|
After Width: | Height: | Size: 54 MiB |
|
After Width: | Height: | Size: 76 MiB |
|
After Width: | Height: | Size: 27 MiB |
|
After Width: | Height: | Size: 35 MiB |
|
After Width: | Height: | Size: 32 MiB |
|
After Width: | Height: | Size: 1.4 MiB |
|
After Width: | Height: | Size: 14 MiB |
|
After Width: | Height: | Size: 3.0 MiB |
@@ -0,0 +1,18 @@
|
||||
CUDA_VISIBLE_DEVICES=0 python inference_basic.py \
|
||||
--pretrained_model_name_or_path="path/checkpoints/SVD/stable-video-diffusion-img2vid-xt" \
|
||||
--output_dir="path/basic_infer" \
|
||||
--validation_control_folder="path/inference/case-1/poses" \
|
||||
--validation_image="path/inference/case-1/reference.png" \
|
||||
--width=576 \
|
||||
--height=1024 \
|
||||
--guidance_scale=3.0 \
|
||||
--num_inference_steps=25 \
|
||||
--posenet_model_name_or_path="path/checkpoints/Animation/pose_net.pth" \
|
||||
--face_encoder_model_name_or_path="path/checkpoints/Animation/face_encoder.pth" \
|
||||
--unet_model_name_or_path="path/checkpoints/Animation/unet.pth" \
|
||||
--tile_size=16 \
|
||||
--overlap=4 \
|
||||
--noise_aug_strength=0.02 \
|
||||
--frames_overlap=4 \
|
||||
--decode_chunk_size=4 \
|
||||
--gradient_checkpointing
|
||||
@@ -0,0 +1,26 @@
|
||||
CUDA_VISIBLE_DEVICES=3,2,1,0 accelerate launch train.py \
|
||||
--pretrained_model_name_or_path="path/checkpoints/SVD/stable-video-diffusion-img2vid-xt" \
|
||||
--finetune_mode=True \
|
||||
--posenet_model_finetune_path="path/checkpoints/Animation/pose_net.pth" \
|
||||
--face_encoder_finetune_path="path/checkpoints/Animation/face_encoder.pth" \
|
||||
--unet_model_finetune_path="path/checkpoints/Animation/unet.pth" \
|
||||
--output_dir="path/checkpoints/Animation" \
|
||||
--data_root_path="path/animation_data" \
|
||||
--rec_data_path="path/animation_data/video_rec_path.txt" \
|
||||
--vec_data_path="path/animation_data/video_vec_path.txt" \
|
||||
--validation_image_folder="path/validation/ground_truth" \
|
||||
--validation_control_folder="path/validation/poses" \
|
||||
--validation_image="path/validation/reference.png" \
|
||||
--num_workers=8 \
|
||||
--lr_warmup_steps=500 \
|
||||
--sample_n_frames=16 \
|
||||
--learning_rate=1e-5 \
|
||||
--per_gpu_batch_size=1 \
|
||||
--num_train_epochs=6000 \
|
||||
--mixed_precision="fp16" \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--checkpointing_steps=2000 \
|
||||
--validation_steps=500 \
|
||||
--gradient_checkpointing \
|
||||
--checkpoints_total_limit=5000 \
|
||||
--resume_from_checkpoint="latest"
|
||||
@@ -0,0 +1,22 @@
|
||||
CUDA_VISIBLE_DEVICES=3,2,1,0 accelerate launch train.py \
|
||||
--pretrained_model_name_or_path="path/checkpoints/SVD/stable-video-diffusion-img2vid-xt" \
|
||||
--output_dir="path/checkpoints/Animation" \
|
||||
--data_root_path="path/animation_data" \
|
||||
--rec_data_path="path/animation_data/video_rec_path.txt" \
|
||||
--vec_data_path="path/animation_data/video_vec_path.txt" \
|
||||
--validation_image_folder="path/validation/ground_truth" \
|
||||
--validation_control_folder="path/validation/poses" \
|
||||
--validation_image="path/validation/reference.png" \
|
||||
--num_workers=8 \
|
||||
--lr_warmup_steps=500 \
|
||||
--sample_n_frames=16 \
|
||||
--learning_rate=1e-5 \
|
||||
--per_gpu_batch_size=1 \
|
||||
--num_train_epochs=6000 \
|
||||
--mixed_precision="fp16" \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--checkpointing_steps=2000 \
|
||||
--validation_steps=500 \
|
||||
--gradient_checkpointing \
|
||||
--checkpoints_total_limit=5000 \
|
||||
--resume_from_checkpoint="latest"
|
||||
@@ -0,0 +1,23 @@
|
||||
CUDA_VISIBLE_DEVICES=3,2,1,0 accelerate launch train_single.py \
|
||||
--pretrained_model_name_or_path="path/checkpoints/SVD/stable-video-diffusion-img2vid-xt" \
|
||||
--output_dir="path/checkpoints/Animation" \
|
||||
--data_root_path="path/animation_data" \
|
||||
--data_path="path/animation_data/video_path.txt" \
|
||||
--dataset_width=512 \
|
||||
--dataset_height=512 \
|
||||
--validation_image_folder="path/validation/ground_truth" \
|
||||
--validation_control_folder="path/validation/poses" \
|
||||
--validation_image="path/validation/reference.png" \
|
||||
--num_workers=8 \
|
||||
--lr_warmup_steps=500 \
|
||||
--sample_n_frames=16 \
|
||||
--learning_rate=1e-5 \
|
||||
--per_gpu_batch_size=1 \
|
||||
--num_train_epochs=6000 \
|
||||
--mixed_precision="fp16" \
|
||||
--gradient_accumulation_steps=1 \
|
||||
--checkpointing_steps=2000 \
|
||||
--validation_steps=500 \
|
||||
--gradient_checkpointing \
|
||||
--checkpoints_total_limit=5000 \
|
||||
--resume_from_checkpoint="latest"
|
||||
@@ -0,0 +1,85 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
from facexlib.parsing import init_parsing_model
|
||||
from facexlib.utils.face_restoration_helper import FaceRestoreHelper
|
||||
from insightface.app import FaceAnalysis
|
||||
import cv2
|
||||
import argparse
|
||||
import os
|
||||
|
||||
def get_face_masks(image_path, save_path, app, face_helper, height=904, width=512):
|
||||
image_1 = cv2.imread(image_path)
|
||||
height, width = image_1.shape[:2]
|
||||
image_bgr_1 = cv2.cvtColor(image_1, cv2.COLOR_RGB2BGR)
|
||||
image_info_1 = app.get(image_bgr_1)
|
||||
|
||||
mask_1 = np.zeros((height, width), dtype=np.uint8)
|
||||
if len(image_info_1) > 0:
|
||||
print("This is FaceAnalysis")
|
||||
for info in image_info_1:
|
||||
x_1 = info['bbox'][0]
|
||||
y_1 = info['bbox'][1]
|
||||
x_2 = info['bbox'][2]
|
||||
y_2 = info['bbox'][3]
|
||||
cv2.rectangle(mask_1, (int(x_1), int(y_1)), (int(x_2), int(y_2)), (255), thickness=cv2.FILLED)
|
||||
cv2.imwrite(save_path, mask_1)
|
||||
else:
|
||||
face_helper.clean_all()
|
||||
with torch.no_grad():
|
||||
bboxes = face_helper.face_det.detect_faces(image_bgr_1, 0.97)
|
||||
if len(bboxes) > 0:
|
||||
print("This is FaceRestoreHelper")
|
||||
for bbox in bboxes:
|
||||
cv2.rectangle(mask_1, (int(bbox[0]), int(bbox[1])), (int(bbox[2]), int(bbox[3])), (255), thickness=cv2.FILLED)
|
||||
cv2.imwrite(save_path, mask_1)
|
||||
else:
|
||||
print("This is no detected face")
|
||||
mask_1[:] = 255
|
||||
cv2.imwrite(save_path, mask_1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
parser = argparse.ArgumentParser("Human Face Mask Extraction", add_help=True)
|
||||
parser.add_argument("--image_folder", type=str, help="Specify a path of a image folder")
|
||||
args = parser.parse_args()
|
||||
|
||||
image_folder = args.image_folder
|
||||
|
||||
app = FaceAnalysis(
|
||||
name='antelopev2', root='.', providers=['CUDAExecutionProvider', 'CPUExecutionProvider']
|
||||
)
|
||||
app.prepare(ctx_id=0, det_size=(640, 640))
|
||||
face_helper = FaceRestoreHelper(
|
||||
upscale_factor=1,
|
||||
face_size=512,
|
||||
crop_ratio=(1, 1),
|
||||
det_model='retinaface_resnet50',
|
||||
save_ext='png',
|
||||
device="cuda",
|
||||
)
|
||||
face_helper.face_parse = init_parsing_model(model_name='bisenet', device="cuda")
|
||||
|
||||
print(f"images subfolder path: {image_folder}")
|
||||
face_subfolder_path = os.path.join(os.path.dirname(image_folder), "faces")
|
||||
if not os.path.exists(face_subfolder_path):
|
||||
os.makedirs(face_subfolder_path)
|
||||
print(f"Folder created: {face_subfolder_path}")
|
||||
else:
|
||||
print(f"Folder already exists: {face_subfolder_path}")
|
||||
for root, dirs, files in os.walk(image_folder):
|
||||
for file in files:
|
||||
if file.endswith('.png'):
|
||||
file_path = os.path.join(root, file)
|
||||
print(file_path)
|
||||
file_name = os.path.splitext(file)[0]
|
||||
image_name = file_name + '.png'
|
||||
image_legal_path = os.path.join(image_folder, image_name)
|
||||
if os.path.exists(os.path.join(face_subfolder_path, file_name + '.png')):
|
||||
existed_path = os.path.join(face_subfolder_path, file_name + '.png')
|
||||
print(f"{existed_path} already exists!")
|
||||
continue
|
||||
|
||||
face_save_path = os.path.join(face_subfolder_path, file_name + '.png')
|
||||
get_face_masks(image_path=image_legal_path, save_path=face_save_path, app=app, face_helper=face_helper)
|
||||
print(f"Finish face Extraction: {face_save_path}")
|
||||
@@ -0,0 +1,400 @@
|
||||
import argparse
|
||||
import os
|
||||
import cv2
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from diffusers.models.attention_processor import XFormersAttnProcessor
|
||||
from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection
|
||||
import torch
|
||||
from diffusers import AutoencoderKLTemporalDecoder, EulerDiscreteScheduler
|
||||
|
||||
from animation.modules.attention_processor import AnimationAttnProcessor
|
||||
from animation.modules.attention_processor_normalized import AnimationIDAttnNormalizedProcessor
|
||||
from animation.modules.face_model import FaceModel
|
||||
from animation.modules.id_encoder import FusionFaceId
|
||||
from animation.modules.pose_net import PoseNet
|
||||
from animation.modules.unet import UNetSpatioTemporalConditionModel
|
||||
from animation.pipelines.inference_pipeline_animation import InferenceAnimationPipeline
|
||||
import random
|
||||
|
||||
def seed_everything(seed):
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
np.random.seed(seed % (2**32))
|
||||
random.seed(seed)
|
||||
|
||||
|
||||
def load_images_from_folder(folder, width, height):
|
||||
images = []
|
||||
files = os.listdir(folder)
|
||||
png_files = [f for f in files if f.endswith('.png')]
|
||||
png_files.sort(key=lambda x: int(x.split('_')[1].split('.')[0]))
|
||||
for filename in png_files:
|
||||
img = Image.open(os.path.join(folder, filename)).convert('RGB')
|
||||
img = img.resize((width, height))
|
||||
images.append(img)
|
||||
|
||||
return images
|
||||
|
||||
def save_frames_as_png(frames, output_path):
|
||||
pil_frames = [Image.fromarray(frame) if isinstance(frame, np.ndarray) else frame for frame in frames]
|
||||
num_frames = len(pil_frames)
|
||||
for i in range(num_frames):
|
||||
pil_frame = pil_frames[i]
|
||||
save_path = os.path.join(output_path, f'frame_{i}.png')
|
||||
pil_frame.save(save_path)
|
||||
|
||||
def save_frames_as_mp4(frames, output_mp4_path, fps):
|
||||
print("Starting saving the frames as mp4")
|
||||
height, width, _ = frames[0].shape
|
||||
fourcc = cv2.VideoWriter_fourcc(*'mp4v') # 'H264' for better quality
|
||||
out = cv2.VideoWriter(output_mp4_path, fourcc, fps, (width, height))
|
||||
for frame in frames:
|
||||
frame_bgr = frame if frame.shape[2] == 3 else cv2.cvtColor(frame, cv2.COLOR_RGB2BGR)
|
||||
out.write(frame_bgr)
|
||||
out.release()
|
||||
|
||||
|
||||
def export_to_gif(frames, output_gif_path, fps):
|
||||
"""
|
||||
Export a list of frames to a GIF.
|
||||
|
||||
Args:
|
||||
- frames (list): List of frames (as numpy arrays or PIL Image objects).
|
||||
- output_gif_path (str): Path to save the output GIF.
|
||||
- duration_ms (int): Duration of each frame in milliseconds.
|
||||
|
||||
"""
|
||||
# Convert numpy arrays to PIL Images if needed
|
||||
pil_frames = [Image.fromarray(frame) if isinstance(
|
||||
frame, np.ndarray) else frame for frame in frames]
|
||||
|
||||
pil_frames[0].save(output_gif_path.replace('.mp4', '.gif'),
|
||||
format='GIF',
|
||||
append_images=pil_frames[1:],
|
||||
save_all=True,
|
||||
duration=125,
|
||||
loop=0)
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Script to train Stable Diffusion XL for InstructPix2Pix."
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--pretrained_model_name_or_path",
|
||||
type=str,
|
||||
default=None,
|
||||
required=True
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--validation_image",
|
||||
type=str,
|
||||
default=None,
|
||||
help=(
|
||||
"A set of paths to the controlnext conditioning image be evaluated every `--validation_steps`"
|
||||
" and logged to `--report_to`. Provide either a matching number of `--validation_prompt`s, a"
|
||||
" a single `--validation_prompt` to be used with all `--validation_image`s, or a single"
|
||||
" `--validation_image` that will be used with all `--validation_prompt`s."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--validation_control_folder",
|
||||
type=str,
|
||||
default=None,
|
||||
help=(
|
||||
"the validation control image"
|
||||
),
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
required=True
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--height",
|
||||
type=int,
|
||||
default=768,
|
||||
required=False
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--width",
|
||||
type=int,
|
||||
default=512,
|
||||
required=False
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--guidance_scale",
|
||||
type=float,
|
||||
default=2.0,
|
||||
required=False
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--num_inference_steps",
|
||||
type=int,
|
||||
default=25,
|
||||
required=False
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--posenet_model_name_or_path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path to pretrained posenet model",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--face_encoder_model_name_or_path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path to pretrained face encoder model",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--unet_model_name_or_path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path to pretrained unet model",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--tile_size",
|
||||
type=int,
|
||||
default=16,
|
||||
required=False
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--overlap",
|
||||
type=int,
|
||||
default=4,
|
||||
required=False
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--noise_aug_strength",
|
||||
type=float,
|
||||
default=0.0, # or set to 0.02
|
||||
required=False
|
||||
)
|
||||
parser.add_argument(
|
||||
"--frames_overlap",
|
||||
type=int,
|
||||
default=4,
|
||||
required=False
|
||||
)
|
||||
parser.add_argument(
|
||||
"--gradient_checkpointing",
|
||||
action="store_true",
|
||||
help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--revision",
|
||||
type=str,
|
||||
default=None,
|
||||
required=False,
|
||||
help="Revision of pretrained model identifier from huggingface.co/models.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--decode_chunk_size",
|
||||
type=int,
|
||||
default=None,
|
||||
required=False
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = parse_args()
|
||||
|
||||
# torch.set_default_dtype(torch.float16)
|
||||
seed = 23123134
|
||||
# seed = 42
|
||||
# seed = 123
|
||||
seed_everything(seed)
|
||||
generator = torch.Generator(device='cuda').manual_seed(seed)
|
||||
|
||||
feature_extractor = CLIPImageProcessor.from_pretrained(args.pretrained_model_name_or_path, subfolder="feature_extractor", revision=args.revision)
|
||||
noise_scheduler = EulerDiscreteScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler")
|
||||
image_encoder = CLIPVisionModelWithProjection.from_pretrained(
|
||||
args.pretrained_model_name_or_path, subfolder="image_encoder", revision=args.revision
|
||||
)
|
||||
vae = AutoencoderKLTemporalDecoder.from_pretrained(
|
||||
args.pretrained_model_name_or_path, subfolder="vae", revision=args.revision)
|
||||
unet = UNetSpatioTemporalConditionModel.from_pretrained(
|
||||
args.pretrained_model_name_or_path,
|
||||
subfolder="unet",
|
||||
low_cpu_mem_usage=True,
|
||||
)
|
||||
pose_net = PoseNet(noise_latent_channels=unet.config.block_out_channels[0])
|
||||
face_encoder = FusionFaceId(
|
||||
cross_attention_dim=1024,
|
||||
id_embeddings_dim=512,
|
||||
# clip_embeddings_dim=image_encoder.config.hidden_size,
|
||||
clip_embeddings_dim=1024,
|
||||
num_tokens=4, )
|
||||
face_model = FaceModel()
|
||||
|
||||
lora_rank = 128
|
||||
attn_procs = {}
|
||||
unet_svd = unet.state_dict()
|
||||
|
||||
for name in unet.attn_processors.keys():
|
||||
if "transformer_blocks" in name and "temporal_transformer_blocks" not in name:
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") else unet.config.cross_attention_dim
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = unet.config.block_out_channels[-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(unet.config.block_out_channels))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = unet.config.block_out_channels[block_id]
|
||||
if cross_attention_dim is None:
|
||||
# print(f"This is AnimationAttnProcessor: {name}")
|
||||
attn_procs[name] = AnimationAttnProcessor(hidden_size=hidden_size, cross_attention_dim=cross_attention_dim, rank=lora_rank)
|
||||
else:
|
||||
# print(f"This is AnimationIDAttnProcessor: {name}")
|
||||
layer_name = name.split(".processor")[0]
|
||||
weights = {
|
||||
"to_k_ip.weight": unet_svd[layer_name + ".to_k.weight"],
|
||||
"to_v_ip.weight": unet_svd[layer_name + ".to_v.weight"],
|
||||
}
|
||||
attn_procs[name] = AnimationIDAttnNormalizedProcessor(hidden_size=hidden_size, cross_attention_dim=cross_attention_dim, rank=lora_rank)
|
||||
attn_procs[name].load_state_dict(weights, strict=False)
|
||||
elif "temporal_transformer_blocks" in name:
|
||||
cross_attention_dim = None if name.endswith("attn1.processor") else unet.config.cross_attention_dim
|
||||
if name.startswith("mid_block"):
|
||||
hidden_size = unet.config.block_out_channels[-1]
|
||||
elif name.startswith("up_blocks"):
|
||||
block_id = int(name[len("up_blocks.")])
|
||||
hidden_size = list(reversed(unet.config.block_out_channels))[block_id]
|
||||
elif name.startswith("down_blocks"):
|
||||
block_id = int(name[len("down_blocks.")])
|
||||
hidden_size = unet.config.block_out_channels[block_id]
|
||||
if cross_attention_dim is None:
|
||||
attn_procs[name] = XFormersAttnProcessor()
|
||||
else:
|
||||
attn_procs[name] = XFormersAttnProcessor()
|
||||
unet.set_attn_processor(attn_procs)
|
||||
|
||||
# resume the previous checkpoint
|
||||
if args.posenet_model_name_or_path is not None and args.face_encoder_model_name_or_path is not None and args.unet_model_name_or_path is not None:
|
||||
print("Loading existing posenet weights, face_encoder weights and unet weights.")
|
||||
if args.posenet_model_name_or_path.endswith(".pth"):
|
||||
pose_net_state_dict = torch.load(args.posenet_model_name_or_path, map_location="cpu")
|
||||
pose_net.load_state_dict(pose_net_state_dict, strict=True)
|
||||
else:
|
||||
print("posenet weights loading fail")
|
||||
print(1/0)
|
||||
if args.face_encoder_model_name_or_path.endswith(".pth"):
|
||||
face_encoder_state_dict = torch.load(args.face_encoder_model_name_or_path, map_location="cpu")
|
||||
face_encoder.load_state_dict(face_encoder_state_dict, strict=True)
|
||||
else:
|
||||
print("face_encoder weights loading fail")
|
||||
print(1/0)
|
||||
if args.unet_model_name_or_path.endswith(".pth"):
|
||||
unet_state_dict = torch.load(args.unet_model_name_or_path, map_location="cpu")
|
||||
unet.load_state_dict(unet_state_dict, strict=True)
|
||||
else:
|
||||
print("unet weights loading fail")
|
||||
print(1/0)
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
vae.requires_grad_(False)
|
||||
image_encoder.requires_grad_(False)
|
||||
unet.requires_grad_(False)
|
||||
pose_net.requires_grad_(False)
|
||||
face_encoder.requires_grad_(False)
|
||||
|
||||
if args.gradient_checkpointing:
|
||||
unet.enable_gradient_checkpointing()
|
||||
|
||||
weight_dtype = torch.float16
|
||||
# weight_dtype = torch.float32
|
||||
# weight_dtype = torch.bfloat16
|
||||
|
||||
pipeline = InferenceAnimationPipeline(
|
||||
vae=vae,
|
||||
image_encoder=image_encoder,
|
||||
unet=unet,
|
||||
scheduler=noise_scheduler,
|
||||
feature_extractor=feature_extractor,
|
||||
pose_net=pose_net,
|
||||
face_encoder=face_encoder,
|
||||
).to(device='cuda', dtype=weight_dtype)
|
||||
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
validation_image_path = args.validation_image
|
||||
validation_image = Image.open(args.validation_image).convert('RGB')
|
||||
validation_control_images = load_images_from_folder(args.validation_control_folder, width=args.width, height=args.height)
|
||||
|
||||
num_frames = len(validation_control_images)
|
||||
face_model.face_helper.clean_all()
|
||||
validation_face = cv2.imread(validation_image_path)
|
||||
validation_image_bgr = cv2.cvtColor(validation_face, cv2.COLOR_RGB2BGR)
|
||||
validation_image_face_info = face_model.app.get(validation_image_bgr)
|
||||
if len(validation_image_face_info) > 0:
|
||||
validation_image_face_info = sorted(validation_image_face_info, key=lambda x: (x['bbox'][2] - x['bbox'][0]) * (x['bbox'][3] - x['bbox'][1]))[-1]
|
||||
validation_image_id_ante_embedding = validation_image_face_info['embedding']
|
||||
else:
|
||||
validation_image_id_ante_embedding = None
|
||||
|
||||
if validation_image_id_ante_embedding is None:
|
||||
face_model.face_helper.read_image(validation_image_bgr)
|
||||
face_model.face_helper.get_face_landmarks_5(only_center_face=True)
|
||||
face_model.face_helper.align_warp_face()
|
||||
|
||||
if len(face_model.face_helper.cropped_faces) == 0:
|
||||
validation_image_id_ante_embedding = np.zeros((512,))
|
||||
else:
|
||||
validation_image_align_face = face_model.face_helper.cropped_faces[0]
|
||||
print('fail to detect face using insightface, extract embedding on align face')
|
||||
validation_image_id_ante_embedding = face_model.handler_ante.get_feat(validation_image_align_face)
|
||||
|
||||
# generator = torch.Generator(device=accelerator.device).manual_seed(23123134)
|
||||
|
||||
decode_chunk_size = args.decode_chunk_size
|
||||
video_frames = pipeline(
|
||||
image=validation_image,
|
||||
image_pose=validation_control_images,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=num_frames,
|
||||
tile_size=args.tile_size,
|
||||
tile_overlap=args.frames_overlap,
|
||||
decode_chunk_size=decode_chunk_size,
|
||||
motion_bucket_id=127.,
|
||||
fps=7,
|
||||
min_guidance_scale=args.guidance_scale,
|
||||
max_guidance_scale=args.guidance_scale,
|
||||
noise_aug_strength=args.noise_aug_strength,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
generator=generator,
|
||||
output_type="pil",
|
||||
validation_image_id_ante_embedding=validation_image_id_ante_embedding,
|
||||
).frames[0]
|
||||
|
||||
out_file = os.path.join(
|
||||
args.output_dir,
|
||||
f"animation_video.mp4",
|
||||
)
|
||||
for i in range(num_frames):
|
||||
img = video_frames[i]
|
||||
video_frames[i] = np.array(img)
|
||||
|
||||
png_out_file = os.path.join(args.output_dir, "animated_images")
|
||||
os.makedirs(png_out_file, exist_ok=True)
|
||||
export_to_gif(video_frames, out_file, 8)
|
||||
save_frames_as_png(video_frames, png_out_file)
|
||||
|
||||
|
||||
# bash command_basic_infer.sh
|
||||
@@ -0,0 +1,20 @@
|
||||
diffusers
|
||||
transformers==4.35.2
|
||||
accelerate==0.25.0
|
||||
timm==0.4.12
|
||||
decord
|
||||
einops
|
||||
scipy
|
||||
pandas
|
||||
coloredlogs
|
||||
flatbuffers
|
||||
numpy==1.26.4
|
||||
packaging
|
||||
protobuf
|
||||
sympy
|
||||
imageio-ffmpeg
|
||||
insightface
|
||||
facexlib
|
||||
opencv-python-headless
|
||||
gradio
|
||||
onnxruntime-gpu
|
||||