diff --git a/animation/StableAnimator/DWPose/__init__.py b/animation/StableAnimator/DWPose/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/animation/StableAnimator/DWPose/dwpose_utils/__init__.py b/animation/StableAnimator/DWPose/dwpose_utils/__init__.py new file mode 100644 index 0000000..8172784 --- /dev/null +++ b/animation/StableAnimator/DWPose/dwpose_utils/__init__.py @@ -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) \ No newline at end of file diff --git a/animation/StableAnimator/DWPose/dwpose_utils/dwpose_detector.py b/animation/StableAnimator/DWPose/dwpose_utils/dwpose_detector.py new file mode 100644 index 0000000..e663f55 --- /dev/null +++ b/animation/StableAnimator/DWPose/dwpose_utils/dwpose_detector.py @@ -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) \ No newline at end of file diff --git a/animation/StableAnimator/DWPose/dwpose_utils/onnxdet.py b/animation/StableAnimator/DWPose/dwpose_utils/onnxdet.py new file mode 100644 index 0000000..e0411c9 --- /dev/null +++ b/animation/StableAnimator/DWPose/dwpose_utils/onnxdet.py @@ -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 diff --git a/animation/StableAnimator/DWPose/dwpose_utils/onnxpose.py b/animation/StableAnimator/DWPose/dwpose_utils/onnxpose.py new file mode 100644 index 0000000..79cd4a0 --- /dev/null +++ b/animation/StableAnimator/DWPose/dwpose_utils/onnxpose.py @@ -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 \ No newline at end of file diff --git a/animation/StableAnimator/DWPose/dwpose_utils/util.py b/animation/StableAnimator/DWPose/dwpose_utils/util.py new file mode 100644 index 0000000..73d7d01 --- /dev/null +++ b/animation/StableAnimator/DWPose/dwpose_utils/util.py @@ -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 diff --git a/animation/StableAnimator/DWPose/dwpose_utils/wholebody.py b/animation/StableAnimator/DWPose/dwpose_utils/wholebody.py new file mode 100644 index 0000000..51dd282 --- /dev/null +++ b/animation/StableAnimator/DWPose/dwpose_utils/wholebody.py @@ -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 + + + diff --git a/animation/StableAnimator/DWPose/skeleton_extraction.py b/animation/StableAnimator/DWPose/skeleton_extraction.py new file mode 100644 index 0000000..2833d15 --- /dev/null +++ b/animation/StableAnimator/DWPose/skeleton_extraction.py @@ -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}") diff --git a/animation/StableAnimator/DWPose/training_skeleton_extraction.py b/animation/StableAnimator/DWPose/training_skeleton_extraction.py new file mode 100644 index 0000000..a0f7e1a --- /dev/null +++ b/animation/StableAnimator/DWPose/training_skeleton_extraction.py @@ -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}") diff --git a/animation/StableAnimator/LICENSE b/animation/StableAnimator/LICENSE new file mode 100644 index 0000000..89282bc --- /dev/null +++ b/animation/StableAnimator/LICENSE @@ -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 \ No newline at end of file diff --git a/animation/StableAnimator/README.md b/animation/StableAnimator/README.md new file mode 100644 index 0000000..8a3676d --- /dev/null +++ b/animation/StableAnimator/README.md @@ -0,0 +1,384 @@ +# StableAnimator + + + +StableAnimator: High-Quality Identity-Preserving Human Image Animation +
+*Shuyuan Tu1, Zhen Xing1, Xintong Han3, Zhi-Qi Cheng4, Qi Dai2, Chong Luo2, Zuxuan Wu1* +
+[1Fudan University; 2Microsoft Research Asia; 3Huya Inc; 4Carnegie Mellon University] + +

+ + + + + + +
+ Pose-driven Human image animations generated by StableAnimator, showing its power to synthesize high-fidelity and ID-preserving videos. All animations are directly synthesized by StableAnimator without the use of any face-related post-processing tools, such as the face-swapping tool FaceFusion or face restoration models like GFP-GAN and CodeFormer. +

+ +

+ + + + +
+ Comparison results between StableAnimator and state-of-the-art (SOTA) human image animation models highlight the superior performance of StableAnimator in delivering high-fidelity, identity-preserving human image animation. +

+ + +## Overview + +

+ model architecture +
+ The overview of the framework of StableAnimator. +

+ +Current diffusion models for human image animation struggle to ensure identity (ID) consistency. This paper presents StableAnimator, 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. 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 +``` +Notably, there is a bug in the automatic download process of Antelopev2, with the error details described as follows: +``` +Traceback (most recent call last): + File "/home/StableAnimator/inference_normal.py", line 243, in + 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: animated_images and animated_images.gif. +If you want to obtain the high quality MP4 file, we recommend you to leverage ffmpeg on the animated_images 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 +πŸ”₯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.πŸ”₯ +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 +``` +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. +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, please consider giving a star to this github repository and citing it: +```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} +} +``` diff --git a/animation/StableAnimator/animation/__init__.py b/animation/StableAnimator/animation/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/animation/StableAnimator/animation/dataset/animation_dataset.py b/animation/StableAnimator/animation/dataset/animation_dataset.py new file mode 100644 index 0000000..49c8e00 --- /dev/null +++ b/animation/StableAnimator/animation/dataset/animation_dataset.py @@ -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 + + + + + + diff --git a/animation/StableAnimator/animation/modules/__init__.py b/animation/StableAnimator/animation/modules/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/animation/StableAnimator/animation/modules/attention.py b/animation/StableAnimator/animation/modules/attention.py new file mode 100644 index 0000000..851177c --- /dev/null +++ b/animation/StableAnimator/animation/modules/attention.py @@ -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) diff --git a/animation/StableAnimator/animation/modules/attention_processor.py b/animation/StableAnimator/animation/modules/attention_processor.py new file mode 100644 index 0000000..333a9a7 --- /dev/null +++ b/animation/StableAnimator/animation/modules/attention_processor.py @@ -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 diff --git a/animation/StableAnimator/animation/modules/attention_processor_normalized.py b/animation/StableAnimator/animation/modules/attention_processor_normalized.py new file mode 100644 index 0000000..399b5a8 --- /dev/null +++ b/animation/StableAnimator/animation/modules/attention_processor_normalized.py @@ -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 diff --git a/animation/StableAnimator/animation/modules/face_model.py b/animation/StableAnimator/animation/modules/face_model.py new file mode 100644 index 0000000..ef92d96 --- /dev/null +++ b/animation/StableAnimator/animation/modules/face_model.py @@ -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") diff --git a/animation/StableAnimator/animation/modules/id_encoder.py b/animation/StableAnimator/animation/modules/id_encoder.py new file mode 100644 index 0000000..c78a9c9 --- /dev/null +++ b/animation/StableAnimator/animation/modules/id_encoder.py @@ -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 diff --git a/animation/StableAnimator/animation/modules/pose_net.py b/animation/StableAnimator/animation/modules/pose_net.py new file mode 100644 index 0000000..daecab4 --- /dev/null +++ b/animation/StableAnimator/animation/modules/pose_net.py @@ -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 diff --git a/animation/StableAnimator/animation/modules/refined_vae.py b/animation/StableAnimator/animation/modules/refined_vae.py new file mode 100644 index 0000000..3522c71 --- /dev/null +++ b/animation/StableAnimator/animation/modules/refined_vae.py @@ -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) diff --git a/animation/StableAnimator/animation/modules/transformer_temporal.py b/animation/StableAnimator/animation/modules/transformer_temporal.py new file mode 100644 index 0000000..33de50c --- /dev/null +++ b/animation/StableAnimator/animation/modules/transformer_temporal.py @@ -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) diff --git a/animation/StableAnimator/animation/modules/unet.py b/animation/StableAnimator/animation/modules/unet.py new file mode 100644 index 0000000..52a8725 --- /dev/null +++ b/animation/StableAnimator/animation/modules/unet.py @@ -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) diff --git a/animation/StableAnimator/animation/modules/unet_3d_blocks.py b/animation/StableAnimator/animation/modules/unet_3d_blocks.py new file mode 100644 index 0000000..7ab636b --- /dev/null +++ b/animation/StableAnimator/animation/modules/unet_3d_blocks.py @@ -0,0 +1,1551 @@ +# 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 typing import Any, Dict, Optional, Tuple, Union + +import torch +from torch import nn + +from diffusers.utils import deprecate, is_torch_version, logging +from diffusers.utils.torch_utils import apply_freeu +from diffusers.models.attention import Attention +from diffusers.models.resnet import ( + Downsample2D, + ResnetBlock2D, + SpatioTemporalResBlock, + TemporalConvLayer, + Upsample2D, +) +from diffusers.models.transformers.transformer_2d import Transformer2DModel +# from diffusers.models.transformers.transformer_temporal import ( +# TransformerSpatioTemporalModel, +# TransformerTemporalModel, +# ) +from animation.modules.transformer_temporal import TransformerTemporalModel, TransformerSpatioTemporalModel + +from diffusers.models.unets.unet_motion_model import ( + CrossAttnDownBlockMotion, + CrossAttnUpBlockMotion, + DownBlockMotion, + UNetMidBlockCrossAttnMotion, + UpBlockMotion, +) + + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +class DownBlockMotion(DownBlockMotion): + def __init__(self, *args, **kwargs): + deprecation_message = "Importing `DownBlockMotion` from `diffusers.models.unets.unet_3d_blocks` is deprecated and this will be removed in a future version. Please use `from diffusers.models.unets.unet_motion_model import DownBlockMotion` instead." + deprecate("DownBlockMotion", "1.0.0", deprecation_message) + super().__init__(*args, **kwargs) + + +class CrossAttnDownBlockMotion(CrossAttnDownBlockMotion): + def __init__(self, *args, **kwargs): + deprecation_message = "Importing `CrossAttnDownBlockMotion` from `diffusers.models.unets.unet_3d_blocks` is deprecated and this will be removed in a future version. Please use `from diffusers.models.unets.unet_motion_model import CrossAttnDownBlockMotion` instead." + deprecate("CrossAttnDownBlockMotion", "1.0.0", deprecation_message) + super().__init__(*args, **kwargs) + + +class UpBlockMotion(UpBlockMotion): + def __init__(self, *args, **kwargs): + deprecation_message = "Importing `UpBlockMotion` from `diffusers.models.unets.unet_3d_blocks` is deprecated and this will be removed in a future version. Please use `from diffusers.models.unets.unet_motion_model import UpBlockMotion` instead." + deprecate("UpBlockMotion", "1.0.0", deprecation_message) + super().__init__(*args, **kwargs) + + +class CrossAttnUpBlockMotion(CrossAttnUpBlockMotion): + def __init__(self, *args, **kwargs): + deprecation_message = "Importing `CrossAttnUpBlockMotion` from `diffusers.models.unets.unet_3d_blocks` is deprecated and this will be removed in a future version. Please use `from diffusers.models.unets.unet_motion_model import CrossAttnUpBlockMotion` instead." + deprecate("CrossAttnUpBlockMotion", "1.0.0", deprecation_message) + super().__init__(*args, **kwargs) + + +class UNetMidBlockCrossAttnMotion(UNetMidBlockCrossAttnMotion): + def __init__(self, *args, **kwargs): + deprecation_message = "Importing `UNetMidBlockCrossAttnMotion` from `diffusers.models.unets.unet_3d_blocks` is deprecated and this will be removed in a future version. Please use `from diffusers.models.unets.unet_motion_model import UNetMidBlockCrossAttnMotion` instead." + deprecate("UNetMidBlockCrossAttnMotion", "1.0.0", deprecation_message) + super().__init__(*args, **kwargs) + + +def get_down_block( + down_block_type: str, + num_layers: int, + in_channels: int, + out_channels: int, + temb_channels: int, + add_downsample: bool, + resnet_eps: float, + resnet_act_fn: str, + num_attention_heads: int, + resnet_groups: Optional[int] = None, + cross_attention_dim: Optional[int] = None, + downsample_padding: Optional[int] = None, + dual_cross_attention: bool = False, + use_linear_projection: bool = True, + only_cross_attention: bool = False, + upcast_attention: bool = False, + resnet_time_scale_shift: str = "default", + temporal_num_attention_heads: int = 8, + temporal_max_seq_length: int = 32, + transformer_layers_per_block: Union[int, Tuple[int]] = 1, + temporal_transformer_layers_per_block: Union[int, Tuple[int]] = 1, + dropout: float = 0.0, +) -> Union[ + "DownBlock3D", + "CrossAttnDownBlock3D", + "DownBlockSpatioTemporal", + "CrossAttnDownBlockSpatioTemporal", +]: + if down_block_type == "DownBlock3D": + return DownBlock3D( + num_layers=num_layers, + in_channels=in_channels, + out_channels=out_channels, + temb_channels=temb_channels, + add_downsample=add_downsample, + resnet_eps=resnet_eps, + resnet_act_fn=resnet_act_fn, + resnet_groups=resnet_groups, + downsample_padding=downsample_padding, + resnet_time_scale_shift=resnet_time_scale_shift, + dropout=dropout, + ) + elif down_block_type == "CrossAttnDownBlock3D": + if cross_attention_dim is None: + raise ValueError("cross_attention_dim must be specified for CrossAttnDownBlock3D") + return CrossAttnDownBlock3D( + num_layers=num_layers, + in_channels=in_channels, + out_channels=out_channels, + temb_channels=temb_channels, + add_downsample=add_downsample, + resnet_eps=resnet_eps, + resnet_act_fn=resnet_act_fn, + resnet_groups=resnet_groups, + downsample_padding=downsample_padding, + cross_attention_dim=cross_attention_dim, + num_attention_heads=num_attention_heads, + dual_cross_attention=dual_cross_attention, + use_linear_projection=use_linear_projection, + only_cross_attention=only_cross_attention, + upcast_attention=upcast_attention, + resnet_time_scale_shift=resnet_time_scale_shift, + dropout=dropout, + ) + elif down_block_type == "DownBlockSpatioTemporal": + # added for SDV + return DownBlockSpatioTemporal( + num_layers=num_layers, + in_channels=in_channels, + out_channels=out_channels, + temb_channels=temb_channels, + add_downsample=add_downsample, + ) + elif down_block_type == "CrossAttnDownBlockSpatioTemporal": + # added for SDV + if cross_attention_dim is None: + raise ValueError("cross_attention_dim must be specified for CrossAttnDownBlockSpatioTemporal") + return CrossAttnDownBlockSpatioTemporal( + in_channels=in_channels, + out_channels=out_channels, + temb_channels=temb_channels, + num_layers=num_layers, + transformer_layers_per_block=transformer_layers_per_block, + add_downsample=add_downsample, + cross_attention_dim=cross_attention_dim, + num_attention_heads=num_attention_heads, + ) + + raise ValueError(f"{down_block_type} does not exist.") + + +def get_up_block( + up_block_type: str, + num_layers: int, + in_channels: int, + out_channels: int, + prev_output_channel: int, + temb_channels: int, + add_upsample: bool, + resnet_eps: float, + resnet_act_fn: str, + num_attention_heads: int, + resolution_idx: Optional[int] = None, + resnet_groups: Optional[int] = None, + cross_attention_dim: Optional[int] = None, + dual_cross_attention: bool = False, + use_linear_projection: bool = True, + only_cross_attention: bool = False, + upcast_attention: bool = False, + resnet_time_scale_shift: str = "default", + temporal_num_attention_heads: int = 8, + temporal_cross_attention_dim: Optional[int] = None, + temporal_max_seq_length: int = 32, + transformer_layers_per_block: Union[int, Tuple[int]] = 1, + temporal_transformer_layers_per_block: Union[int, Tuple[int]] = 1, + dropout: float = 0.0, +) -> Union[ + "UpBlock3D", + "CrossAttnUpBlock3D", + "UpBlockSpatioTemporal", + "CrossAttnUpBlockSpatioTemporal", +]: + if up_block_type == "UpBlock3D": + return UpBlock3D( + num_layers=num_layers, + in_channels=in_channels, + out_channels=out_channels, + prev_output_channel=prev_output_channel, + temb_channels=temb_channels, + add_upsample=add_upsample, + resnet_eps=resnet_eps, + resnet_act_fn=resnet_act_fn, + resnet_groups=resnet_groups, + resnet_time_scale_shift=resnet_time_scale_shift, + resolution_idx=resolution_idx, + dropout=dropout, + ) + elif up_block_type == "CrossAttnUpBlock3D": + if cross_attention_dim is None: + raise ValueError("cross_attention_dim must be specified for CrossAttnUpBlock3D") + return CrossAttnUpBlock3D( + num_layers=num_layers, + in_channels=in_channels, + out_channels=out_channels, + prev_output_channel=prev_output_channel, + temb_channels=temb_channels, + add_upsample=add_upsample, + resnet_eps=resnet_eps, + resnet_act_fn=resnet_act_fn, + resnet_groups=resnet_groups, + cross_attention_dim=cross_attention_dim, + num_attention_heads=num_attention_heads, + dual_cross_attention=dual_cross_attention, + use_linear_projection=use_linear_projection, + only_cross_attention=only_cross_attention, + upcast_attention=upcast_attention, + resnet_time_scale_shift=resnet_time_scale_shift, + resolution_idx=resolution_idx, + dropout=dropout, + ) + elif up_block_type == "UpBlockSpatioTemporal": + # added for SDV + return UpBlockSpatioTemporal( + num_layers=num_layers, + in_channels=in_channels, + out_channels=out_channels, + prev_output_channel=prev_output_channel, + temb_channels=temb_channels, + resolution_idx=resolution_idx, + add_upsample=add_upsample, + ) + elif up_block_type == "CrossAttnUpBlockSpatioTemporal": + # added for SDV + if cross_attention_dim is None: + raise ValueError("cross_attention_dim must be specified for CrossAttnUpBlockSpatioTemporal") + return CrossAttnUpBlockSpatioTemporal( + in_channels=in_channels, + out_channels=out_channels, + prev_output_channel=prev_output_channel, + temb_channels=temb_channels, + num_layers=num_layers, + transformer_layers_per_block=transformer_layers_per_block, + add_upsample=add_upsample, + cross_attention_dim=cross_attention_dim, + num_attention_heads=num_attention_heads, + resolution_idx=resolution_idx, + ) + + raise ValueError(f"{up_block_type} does not exist.") + + +class UNetMidBlock3DCrossAttn(nn.Module): + def __init__( + self, + in_channels: int, + temb_channels: int, + dropout: float = 0.0, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_time_scale_shift: str = "default", + resnet_act_fn: str = "swish", + resnet_groups: int = 32, + resnet_pre_norm: bool = True, + num_attention_heads: int = 1, + output_scale_factor: float = 1.0, + cross_attention_dim: int = 1280, + dual_cross_attention: bool = False, + use_linear_projection: bool = True, + upcast_attention: bool = False, + ): + super().__init__() + + self.has_cross_attention = True + self.num_attention_heads = num_attention_heads + resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32) + + # there is always at least one resnet + resnets = [ + ResnetBlock2D( + in_channels=in_channels, + out_channels=in_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + dropout=dropout, + time_embedding_norm=resnet_time_scale_shift, + non_linearity=resnet_act_fn, + output_scale_factor=output_scale_factor, + pre_norm=resnet_pre_norm, + ) + ] + temp_convs = [ + TemporalConvLayer( + in_channels, + in_channels, + dropout=0.1, + norm_num_groups=resnet_groups, + ) + ] + attentions = [] + temp_attentions = [] + + for _ in range(num_layers): + attentions.append( + Transformer2DModel( + in_channels // num_attention_heads, + num_attention_heads, + in_channels=in_channels, + num_layers=1, + cross_attention_dim=cross_attention_dim, + norm_num_groups=resnet_groups, + use_linear_projection=use_linear_projection, + upcast_attention=upcast_attention, + ) + ) + temp_attentions.append( + TransformerTemporalModel( + in_channels // num_attention_heads, + num_attention_heads, + in_channels=in_channels, + num_layers=1, + cross_attention_dim=cross_attention_dim, + norm_num_groups=resnet_groups, + ) + ) + resnets.append( + ResnetBlock2D( + in_channels=in_channels, + out_channels=in_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + dropout=dropout, + time_embedding_norm=resnet_time_scale_shift, + non_linearity=resnet_act_fn, + output_scale_factor=output_scale_factor, + pre_norm=resnet_pre_norm, + ) + ) + temp_convs.append( + TemporalConvLayer( + in_channels, + in_channels, + dropout=0.1, + norm_num_groups=resnet_groups, + ) + ) + + self.resnets = nn.ModuleList(resnets) + self.temp_convs = nn.ModuleList(temp_convs) + self.attentions = nn.ModuleList(attentions) + self.temp_attentions = nn.ModuleList(temp_attentions) + + def forward( + self, + hidden_states: torch.Tensor, + temb: Optional[torch.Tensor] = None, + encoder_hidden_states: Optional[torch.Tensor] = None, + attention_mask: Optional[torch.Tensor] = None, + num_frames: int = 1, + cross_attention_kwargs: Optional[Dict[str, Any]] = None, + ) -> torch.Tensor: + hidden_states = self.resnets[0](hidden_states, temb) + hidden_states = self.temp_convs[0](hidden_states, num_frames=num_frames) + for attn, temp_attn, resnet, temp_conv in zip( + self.attentions, self.temp_attentions, self.resnets[1:], self.temp_convs[1:] + ): + hidden_states = attn( + hidden_states, + encoder_hidden_states=encoder_hidden_states, + cross_attention_kwargs=cross_attention_kwargs, + return_dict=False, + )[0] + hidden_states = temp_attn( + hidden_states, + num_frames=num_frames, + cross_attention_kwargs=cross_attention_kwargs, + return_dict=False, + )[0] + hidden_states = resnet(hidden_states, temb) + hidden_states = temp_conv(hidden_states, num_frames=num_frames) + + return hidden_states + + +class CrossAttnDownBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + temb_channels: int, + dropout: float = 0.0, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_time_scale_shift: str = "default", + resnet_act_fn: str = "swish", + resnet_groups: int = 32, + resnet_pre_norm: bool = True, + num_attention_heads: int = 1, + cross_attention_dim: int = 1280, + output_scale_factor: float = 1.0, + downsample_padding: int = 1, + add_downsample: bool = True, + dual_cross_attention: bool = False, + use_linear_projection: bool = False, + only_cross_attention: bool = False, + upcast_attention: bool = False, + ): + super().__init__() + resnets = [] + attentions = [] + temp_attentions = [] + temp_convs = [] + + self.has_cross_attention = True + self.num_attention_heads = num_attention_heads + + for i in range(num_layers): + in_channels = in_channels if i == 0 else out_channels + resnets.append( + ResnetBlock2D( + in_channels=in_channels, + out_channels=out_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + dropout=dropout, + time_embedding_norm=resnet_time_scale_shift, + non_linearity=resnet_act_fn, + output_scale_factor=output_scale_factor, + pre_norm=resnet_pre_norm, + ) + ) + temp_convs.append( + TemporalConvLayer( + out_channels, + out_channels, + dropout=0.1, + norm_num_groups=resnet_groups, + ) + ) + attentions.append( + Transformer2DModel( + out_channels // num_attention_heads, + num_attention_heads, + in_channels=out_channels, + num_layers=1, + cross_attention_dim=cross_attention_dim, + norm_num_groups=resnet_groups, + use_linear_projection=use_linear_projection, + only_cross_attention=only_cross_attention, + upcast_attention=upcast_attention, + ) + ) + temp_attentions.append( + TransformerTemporalModel( + out_channels // num_attention_heads, + num_attention_heads, + in_channels=out_channels, + num_layers=1, + cross_attention_dim=cross_attention_dim, + norm_num_groups=resnet_groups, + ) + ) + self.resnets = nn.ModuleList(resnets) + self.temp_convs = nn.ModuleList(temp_convs) + self.attentions = nn.ModuleList(attentions) + self.temp_attentions = nn.ModuleList(temp_attentions) + + if add_downsample: + self.downsamplers = nn.ModuleList( + [ + Downsample2D( + out_channels, + use_conv=True, + out_channels=out_channels, + padding=downsample_padding, + name="op", + ) + ] + ) + else: + self.downsamplers = None + + self.gradient_checkpointing = False + + def forward( + self, + hidden_states: torch.Tensor, + temb: Optional[torch.Tensor] = None, + encoder_hidden_states: Optional[torch.Tensor] = None, + attention_mask: Optional[torch.Tensor] = None, + num_frames: int = 1, + cross_attention_kwargs: Dict[str, Any] = None, + ) -> Union[torch.Tensor, Tuple[torch.Tensor, ...]]: + # TODO(Patrick, William) - attention mask is not used + output_states = () + + for resnet, temp_conv, attn, temp_attn in zip( + self.resnets, self.temp_convs, self.attentions, self.temp_attentions + ): + hidden_states = resnet(hidden_states, temb) + hidden_states = temp_conv(hidden_states, num_frames=num_frames) + hidden_states = attn( + hidden_states, + encoder_hidden_states=encoder_hidden_states, + cross_attention_kwargs=cross_attention_kwargs, + return_dict=False, + )[0] + hidden_states = temp_attn( + hidden_states, + num_frames=num_frames, + cross_attention_kwargs=cross_attention_kwargs, + return_dict=False, + )[0] + + output_states += (hidden_states,) + + if self.downsamplers is not None: + for downsampler in self.downsamplers: + hidden_states = downsampler(hidden_states) + + output_states += (hidden_states,) + + return hidden_states, output_states + + +class DownBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + temb_channels: int, + dropout: float = 0.0, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_time_scale_shift: str = "default", + resnet_act_fn: str = "swish", + resnet_groups: int = 32, + resnet_pre_norm: bool = True, + output_scale_factor: float = 1.0, + add_downsample: bool = True, + downsample_padding: int = 1, + ): + super().__init__() + resnets = [] + temp_convs = [] + + for i in range(num_layers): + in_channels = in_channels if i == 0 else out_channels + resnets.append( + ResnetBlock2D( + in_channels=in_channels, + out_channels=out_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + dropout=dropout, + time_embedding_norm=resnet_time_scale_shift, + non_linearity=resnet_act_fn, + output_scale_factor=output_scale_factor, + pre_norm=resnet_pre_norm, + ) + ) + temp_convs.append( + TemporalConvLayer( + out_channels, + out_channels, + dropout=0.1, + norm_num_groups=resnet_groups, + ) + ) + + self.resnets = nn.ModuleList(resnets) + self.temp_convs = nn.ModuleList(temp_convs) + + if add_downsample: + self.downsamplers = nn.ModuleList( + [ + Downsample2D( + out_channels, + use_conv=True, + out_channels=out_channels, + padding=downsample_padding, + name="op", + ) + ] + ) + else: + self.downsamplers = None + + self.gradient_checkpointing = False + + def forward( + self, + hidden_states: torch.Tensor, + temb: Optional[torch.Tensor] = None, + num_frames: int = 1, + ) -> Union[torch.Tensor, Tuple[torch.Tensor, ...]]: + output_states = () + + for resnet, temp_conv in zip(self.resnets, self.temp_convs): + hidden_states = resnet(hidden_states, temb) + hidden_states = temp_conv(hidden_states, num_frames=num_frames) + + output_states += (hidden_states,) + + if self.downsamplers is not None: + for downsampler in self.downsamplers: + hidden_states = downsampler(hidden_states) + + output_states += (hidden_states,) + + return hidden_states, output_states + + +class CrossAttnUpBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + prev_output_channel: int, + temb_channels: int, + dropout: float = 0.0, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_time_scale_shift: str = "default", + resnet_act_fn: str = "swish", + resnet_groups: int = 32, + resnet_pre_norm: bool = True, + num_attention_heads: int = 1, + cross_attention_dim: int = 1280, + output_scale_factor: float = 1.0, + add_upsample: bool = True, + dual_cross_attention: bool = False, + use_linear_projection: bool = False, + only_cross_attention: bool = False, + upcast_attention: bool = False, + resolution_idx: Optional[int] = None, + ): + super().__init__() + resnets = [] + temp_convs = [] + attentions = [] + temp_attentions = [] + + self.has_cross_attention = True + self.num_attention_heads = num_attention_heads + + for i in range(num_layers): + res_skip_channels = in_channels if (i == num_layers - 1) else out_channels + resnet_in_channels = prev_output_channel if i == 0 else out_channels + + resnets.append( + ResnetBlock2D( + in_channels=resnet_in_channels + res_skip_channels, + out_channels=out_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + dropout=dropout, + time_embedding_norm=resnet_time_scale_shift, + non_linearity=resnet_act_fn, + output_scale_factor=output_scale_factor, + pre_norm=resnet_pre_norm, + ) + ) + temp_convs.append( + TemporalConvLayer( + out_channels, + out_channels, + dropout=0.1, + norm_num_groups=resnet_groups, + ) + ) + attentions.append( + Transformer2DModel( + out_channels // num_attention_heads, + num_attention_heads, + in_channels=out_channels, + num_layers=1, + cross_attention_dim=cross_attention_dim, + norm_num_groups=resnet_groups, + use_linear_projection=use_linear_projection, + only_cross_attention=only_cross_attention, + upcast_attention=upcast_attention, + ) + ) + temp_attentions.append( + TransformerTemporalModel( + out_channels // num_attention_heads, + num_attention_heads, + in_channels=out_channels, + num_layers=1, + cross_attention_dim=cross_attention_dim, + norm_num_groups=resnet_groups, + ) + ) + self.resnets = nn.ModuleList(resnets) + self.temp_convs = nn.ModuleList(temp_convs) + self.attentions = nn.ModuleList(attentions) + self.temp_attentions = nn.ModuleList(temp_attentions) + + if add_upsample: + self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) + else: + self.upsamplers = None + + self.gradient_checkpointing = False + self.resolution_idx = resolution_idx + + def forward( + self, + hidden_states: torch.Tensor, + res_hidden_states_tuple: Tuple[torch.Tensor, ...], + temb: Optional[torch.Tensor] = None, + encoder_hidden_states: Optional[torch.Tensor] = None, + upsample_size: Optional[int] = None, + attention_mask: Optional[torch.Tensor] = None, + num_frames: int = 1, + cross_attention_kwargs: Dict[str, Any] = None, + ) -> torch.Tensor: + is_freeu_enabled = ( + getattr(self, "s1", None) + and getattr(self, "s2", None) + and getattr(self, "b1", None) + and getattr(self, "b2", None) + ) + + # TODO(Patrick, William) - attention mask is not used + for resnet, temp_conv, attn, temp_attn in zip( + self.resnets, self.temp_convs, self.attentions, self.temp_attentions + ): + # pop res hidden states + res_hidden_states = res_hidden_states_tuple[-1] + res_hidden_states_tuple = res_hidden_states_tuple[:-1] + + # FreeU: Only operate on the first two stages + if is_freeu_enabled: + hidden_states, res_hidden_states = apply_freeu( + self.resolution_idx, + hidden_states, + res_hidden_states, + s1=self.s1, + s2=self.s2, + b1=self.b1, + b2=self.b2, + ) + + hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) + + hidden_states = resnet(hidden_states, temb) + hidden_states = temp_conv(hidden_states, num_frames=num_frames) + hidden_states = attn( + hidden_states, + encoder_hidden_states=encoder_hidden_states, + cross_attention_kwargs=cross_attention_kwargs, + return_dict=False, + )[0] + hidden_states = temp_attn( + hidden_states, + num_frames=num_frames, + cross_attention_kwargs=cross_attention_kwargs, + return_dict=False, + )[0] + + if self.upsamplers is not None: + for upsampler in self.upsamplers: + hidden_states = upsampler(hidden_states, upsample_size) + + return hidden_states + + +class UpBlock3D(nn.Module): + def __init__( + self, + in_channels: int, + prev_output_channel: int, + out_channels: int, + temb_channels: int, + dropout: float = 0.0, + num_layers: int = 1, + resnet_eps: float = 1e-6, + resnet_time_scale_shift: str = "default", + resnet_act_fn: str = "swish", + resnet_groups: int = 32, + resnet_pre_norm: bool = True, + output_scale_factor: float = 1.0, + add_upsample: bool = True, + resolution_idx: Optional[int] = None, + ): + super().__init__() + resnets = [] + temp_convs = [] + + for i in range(num_layers): + res_skip_channels = in_channels if (i == num_layers - 1) else out_channels + resnet_in_channels = prev_output_channel if i == 0 else out_channels + + resnets.append( + ResnetBlock2D( + in_channels=resnet_in_channels + res_skip_channels, + out_channels=out_channels, + temb_channels=temb_channels, + eps=resnet_eps, + groups=resnet_groups, + dropout=dropout, + time_embedding_norm=resnet_time_scale_shift, + non_linearity=resnet_act_fn, + output_scale_factor=output_scale_factor, + pre_norm=resnet_pre_norm, + ) + ) + temp_convs.append( + TemporalConvLayer( + out_channels, + out_channels, + dropout=0.1, + norm_num_groups=resnet_groups, + ) + ) + + self.resnets = nn.ModuleList(resnets) + self.temp_convs = nn.ModuleList(temp_convs) + + if add_upsample: + self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) + else: + self.upsamplers = None + + self.gradient_checkpointing = False + self.resolution_idx = resolution_idx + + def forward( + self, + hidden_states: torch.Tensor, + res_hidden_states_tuple: Tuple[torch.Tensor, ...], + temb: Optional[torch.Tensor] = None, + upsample_size: Optional[int] = None, + num_frames: int = 1, + ) -> torch.Tensor: + is_freeu_enabled = ( + getattr(self, "s1", None) + and getattr(self, "s2", None) + and getattr(self, "b1", None) + and getattr(self, "b2", None) + ) + for resnet, temp_conv in zip(self.resnets, self.temp_convs): + # pop res hidden states + res_hidden_states = res_hidden_states_tuple[-1] + res_hidden_states_tuple = res_hidden_states_tuple[:-1] + + # FreeU: Only operate on the first two stages + if is_freeu_enabled: + hidden_states, res_hidden_states = apply_freeu( + self.resolution_idx, + hidden_states, + res_hidden_states, + s1=self.s1, + s2=self.s2, + b1=self.b1, + b2=self.b2, + ) + + hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) + + hidden_states = resnet(hidden_states, temb) + hidden_states = temp_conv(hidden_states, num_frames=num_frames) + + if self.upsamplers is not None: + for upsampler in self.upsamplers: + hidden_states = upsampler(hidden_states, upsample_size) + + return hidden_states + + +class MidBlockTemporalDecoder(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + attention_head_dim: int = 512, + num_layers: int = 1, + upcast_attention: bool = False, + ): + super().__init__() + + resnets = [] + attentions = [] + for i in range(num_layers): + input_channels = in_channels if i == 0 else out_channels + resnets.append( + SpatioTemporalResBlock( + in_channels=input_channels, + out_channels=out_channels, + temb_channels=None, + eps=1e-6, + temporal_eps=1e-5, + merge_factor=0.0, + merge_strategy="learned", + switch_spatial_to_temporal_mix=True, + ) + ) + + attentions.append( + Attention( + query_dim=in_channels, + heads=in_channels // attention_head_dim, + dim_head=attention_head_dim, + eps=1e-6, + upcast_attention=upcast_attention, + norm_num_groups=32, + bias=True, + residual_connection=True, + ) + ) + + self.attentions = nn.ModuleList(attentions) + self.resnets = nn.ModuleList(resnets) + + def forward( + self, + hidden_states: torch.Tensor, + image_only_indicator: torch.Tensor, + ): + hidden_states = self.resnets[0]( + hidden_states, + image_only_indicator=image_only_indicator, + ) + for resnet, attn in zip(self.resnets[1:], self.attentions): + hidden_states = attn(hidden_states) + hidden_states = resnet( + hidden_states, + image_only_indicator=image_only_indicator, + ) + + return hidden_states + + +class UpBlockTemporalDecoder(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + num_layers: int = 1, + add_upsample: bool = True, + ): + super().__init__() + resnets = [] + for i in range(num_layers): + input_channels = in_channels if i == 0 else out_channels + + resnets.append( + SpatioTemporalResBlock( + in_channels=input_channels, + out_channels=out_channels, + temb_channels=None, + eps=1e-6, + temporal_eps=1e-5, + merge_factor=0.0, + merge_strategy="learned", + switch_spatial_to_temporal_mix=True, + ) + ) + self.resnets = nn.ModuleList(resnets) + + if add_upsample: + self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) + else: + self.upsamplers = None + + def forward( + self, + hidden_states: torch.Tensor, + image_only_indicator: torch.Tensor, + ) -> torch.Tensor: + for resnet in self.resnets: + hidden_states = resnet( + hidden_states, + image_only_indicator=image_only_indicator, + ) + + if self.upsamplers is not None: + for upsampler in self.upsamplers: + hidden_states = upsampler(hidden_states) + + return hidden_states + + +class UNetMidBlockSpatioTemporal(nn.Module): + def __init__( + self, + in_channels: int, + temb_channels: int, + num_layers: int = 1, + transformer_layers_per_block: Union[int, Tuple[int]] = 1, + num_attention_heads: int = 1, + cross_attention_dim: int = 1280, + ): + super().__init__() + + self.has_cross_attention = True + self.num_attention_heads = num_attention_heads + + # support for variable transformer layers per block + if isinstance(transformer_layers_per_block, int): + transformer_layers_per_block = [transformer_layers_per_block] * num_layers + + # there is always at least one resnet + resnets = [ + SpatioTemporalResBlock( + in_channels=in_channels, + out_channels=in_channels, + temb_channels=temb_channels, + eps=1e-5, + ) + ] + attentions = [] + + for i in range(num_layers): + attentions.append( + TransformerSpatioTemporalModel( + num_attention_heads, + in_channels // num_attention_heads, + in_channels=in_channels, + num_layers=transformer_layers_per_block[i], + cross_attention_dim=cross_attention_dim, + ) + ) + + resnets.append( + SpatioTemporalResBlock( + in_channels=in_channels, + out_channels=in_channels, + temb_channels=temb_channels, + eps=1e-5, + ) + ) + + self.attentions = nn.ModuleList(attentions) + self.resnets = nn.ModuleList(resnets) + + self.gradient_checkpointing = False + + def forward( + self, + hidden_states: torch.Tensor, + temb: Optional[torch.Tensor] = None, + encoder_hidden_states: Optional[torch.Tensor] = None, + image_only_indicator: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + hidden_states = self.resnets[0]( + hidden_states, + temb, + image_only_indicator=image_only_indicator, + ) + + for attn, resnet in zip(self.attentions, self.resnets[1:]): + if self.training and self.gradient_checkpointing: # TODO + + def create_custom_forward(module, return_dict=None): + def custom_forward(*inputs): + if return_dict is not None: + return module(*inputs, return_dict=return_dict) + else: + return module(*inputs) + + return custom_forward + + ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + hidden_states = attn( + hidden_states, + encoder_hidden_states=encoder_hidden_states, + image_only_indicator=image_only_indicator, + return_dict=False, + )[0] + hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(resnet), + hidden_states, + temb, + image_only_indicator, + **ckpt_kwargs, + ) + else: + hidden_states = attn( + hidden_states, + encoder_hidden_states=encoder_hidden_states, + image_only_indicator=image_only_indicator, + return_dict=False, + )[0] + hidden_states = resnet( + hidden_states, + temb, + image_only_indicator=image_only_indicator, + ) + + return hidden_states + + +class DownBlockSpatioTemporal(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + temb_channels: int, + num_layers: int = 1, + add_downsample: bool = True, + ): + super().__init__() + resnets = [] + + for i in range(num_layers): + in_channels = in_channels if i == 0 else out_channels + resnets.append( + SpatioTemporalResBlock( + in_channels=in_channels, + out_channels=out_channels, + temb_channels=temb_channels, + eps=1e-5, + ) + ) + + self.resnets = nn.ModuleList(resnets) + + if add_downsample: + self.downsamplers = nn.ModuleList( + [ + Downsample2D( + out_channels, + use_conv=True, + out_channels=out_channels, + name="op", + ) + ] + ) + else: + self.downsamplers = None + + self.gradient_checkpointing = False + + def forward( + self, + hidden_states: torch.Tensor, + temb: Optional[torch.Tensor] = None, + image_only_indicator: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, Tuple[torch.Tensor, ...]]: + output_states = () + for resnet in self.resnets: + if self.training and 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"): + hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(resnet), + hidden_states, + temb, + image_only_indicator, + use_reentrant=False, + ) + else: + hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(resnet), + hidden_states, + temb, + image_only_indicator, + ) + else: + hidden_states = resnet( + hidden_states, + temb, + image_only_indicator=image_only_indicator, + ) + + output_states = output_states + (hidden_states,) + + if self.downsamplers is not None: + for downsampler in self.downsamplers: + hidden_states = downsampler(hidden_states) + + output_states = output_states + (hidden_states,) + + return hidden_states, output_states + + +class CrossAttnDownBlockSpatioTemporal(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + temb_channels: int, + num_layers: int = 1, + transformer_layers_per_block: Union[int, Tuple[int]] = 1, + num_attention_heads: int = 1, + cross_attention_dim: int = 1280, + add_downsample: bool = True, + ): + super().__init__() + resnets = [] + attentions = [] + + self.has_cross_attention = True + self.num_attention_heads = num_attention_heads + if isinstance(transformer_layers_per_block, int): + transformer_layers_per_block = [transformer_layers_per_block] * num_layers + + for i in range(num_layers): + in_channels = in_channels if i == 0 else out_channels + resnets.append( + SpatioTemporalResBlock( + in_channels=in_channels, + out_channels=out_channels, + temb_channels=temb_channels, + eps=1e-6, + ) + ) + attentions.append( + TransformerSpatioTemporalModel( + num_attention_heads, + out_channels // num_attention_heads, + in_channels=out_channels, + num_layers=transformer_layers_per_block[i], + cross_attention_dim=cross_attention_dim, + ) + ) + + self.attentions = nn.ModuleList(attentions) + self.resnets = nn.ModuleList(resnets) + + if add_downsample: + self.downsamplers = nn.ModuleList( + [ + Downsample2D( + out_channels, + use_conv=True, + out_channels=out_channels, + padding=1, + name="op", + ) + ] + ) + else: + self.downsamplers = None + + self.gradient_checkpointing = False + + def forward( + self, + hidden_states: torch.Tensor, + temb: Optional[torch.Tensor] = None, + encoder_hidden_states: Optional[torch.Tensor] = None, + image_only_indicator: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, Tuple[torch.Tensor, ...]]: + output_states = () + + blocks = list(zip(self.resnets, self.attentions)) + for resnet, attn in blocks: + if self.training and self.gradient_checkpointing: # TODO + + def create_custom_forward(module, return_dict=None): + def custom_forward(*inputs): + if return_dict is not None: + return module(*inputs, return_dict=return_dict) + else: + return module(*inputs) + + return custom_forward + + ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(resnet), + hidden_states, + temb, + image_only_indicator, + **ckpt_kwargs, + ) + + hidden_states = attn( + hidden_states, + encoder_hidden_states=encoder_hidden_states, + image_only_indicator=image_only_indicator, + return_dict=False, + )[0] + else: + hidden_states = resnet( + hidden_states, + temb, + image_only_indicator=image_only_indicator, + ) + hidden_states = attn( + hidden_states, + encoder_hidden_states=encoder_hidden_states, + image_only_indicator=image_only_indicator, + return_dict=False, + )[0] + + output_states = output_states + (hidden_states,) + + if self.downsamplers is not None: + for downsampler in self.downsamplers: + hidden_states = downsampler(hidden_states) + + output_states = output_states + (hidden_states,) + + return hidden_states, output_states + + +class UpBlockSpatioTemporal(nn.Module): + def __init__( + self, + in_channels: int, + prev_output_channel: int, + out_channels: int, + temb_channels: int, + resolution_idx: Optional[int] = None, + num_layers: int = 1, + resnet_eps: float = 1e-6, + add_upsample: bool = True, + ): + super().__init__() + resnets = [] + + for i in range(num_layers): + res_skip_channels = in_channels if (i == num_layers - 1) else out_channels + resnet_in_channels = prev_output_channel if i == 0 else out_channels + + resnets.append( + SpatioTemporalResBlock( + in_channels=resnet_in_channels + res_skip_channels, + out_channels=out_channels, + temb_channels=temb_channels, + eps=resnet_eps, + ) + ) + + self.resnets = nn.ModuleList(resnets) + + if add_upsample: + self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) + else: + self.upsamplers = None + + self.gradient_checkpointing = False + self.resolution_idx = resolution_idx + + def forward( + self, + hidden_states: torch.Tensor, + res_hidden_states_tuple: Tuple[torch.Tensor, ...], + temb: Optional[torch.Tensor] = None, + image_only_indicator: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + for resnet in self.resnets: + # pop res hidden states + res_hidden_states = res_hidden_states_tuple[-1] + res_hidden_states_tuple = res_hidden_states_tuple[:-1] + + hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) + + if self.training and 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"): + hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(resnet), + hidden_states, + temb, + image_only_indicator, + use_reentrant=False, + ) + else: + hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(resnet), + hidden_states, + temb, + image_only_indicator, + ) + else: + hidden_states = resnet( + hidden_states, + temb, + image_only_indicator=image_only_indicator, + ) + + if self.upsamplers is not None: + for upsampler in self.upsamplers: + hidden_states = upsampler(hidden_states) + + return hidden_states + + +class CrossAttnUpBlockSpatioTemporal(nn.Module): + def __init__( + self, + in_channels: int, + out_channels: int, + prev_output_channel: int, + temb_channels: int, + resolution_idx: Optional[int] = None, + num_layers: int = 1, + transformer_layers_per_block: Union[int, Tuple[int]] = 1, + resnet_eps: float = 1e-6, + num_attention_heads: int = 1, + cross_attention_dim: int = 1280, + add_upsample: bool = True, + ): + super().__init__() + resnets = [] + attentions = [] + + self.has_cross_attention = True + self.num_attention_heads = num_attention_heads + + if isinstance(transformer_layers_per_block, int): + transformer_layers_per_block = [transformer_layers_per_block] * num_layers + + for i in range(num_layers): + res_skip_channels = in_channels if (i == num_layers - 1) else out_channels + resnet_in_channels = prev_output_channel if i == 0 else out_channels + + resnets.append( + SpatioTemporalResBlock( + in_channels=resnet_in_channels + res_skip_channels, + out_channels=out_channels, + temb_channels=temb_channels, + eps=resnet_eps, + ) + ) + attentions.append( + TransformerSpatioTemporalModel( + num_attention_heads, + out_channels // num_attention_heads, + in_channels=out_channels, + num_layers=transformer_layers_per_block[i], + cross_attention_dim=cross_attention_dim, + ) + ) + + self.attentions = nn.ModuleList(attentions) + self.resnets = nn.ModuleList(resnets) + + if add_upsample: + self.upsamplers = nn.ModuleList([Upsample2D(out_channels, use_conv=True, out_channels=out_channels)]) + else: + self.upsamplers = None + + self.gradient_checkpointing = False + self.resolution_idx = resolution_idx + + def forward( + self, + hidden_states: torch.Tensor, + res_hidden_states_tuple: Tuple[torch.Tensor, ...], + temb: Optional[torch.Tensor] = None, + encoder_hidden_states: Optional[torch.Tensor] = None, + image_only_indicator: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + for resnet, attn in zip(self.resnets, self.attentions): + # pop res hidden states + res_hidden_states = res_hidden_states_tuple[-1] + res_hidden_states_tuple = res_hidden_states_tuple[:-1] + + # print("---------------------------------") + # print(len(self.resnets)) + # print(len(self.attentions)) + # print(len(res_hidden_states_tuple)) + # if len(res_hidden_states_tuple) > 0: + # print(res_hidden_states_tuple[0].size()) + # print(hidden_states.size()) # + # print(res_hidden_states.size()) # + # print("---------------------------------") + + hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) + + if self.training and self.gradient_checkpointing: # TODO + + def create_custom_forward(module, return_dict=None): + def custom_forward(*inputs): + if return_dict is not None: + return module(*inputs, return_dict=return_dict) + else: + return module(*inputs) + + return custom_forward + + ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + hidden_states = torch.utils.checkpoint.checkpoint( + create_custom_forward(resnet), + hidden_states, + temb, + image_only_indicator, + **ckpt_kwargs, + ) + hidden_states = attn( + hidden_states, + encoder_hidden_states=encoder_hidden_states, + image_only_indicator=image_only_indicator, + return_dict=False, + )[0] + else: + hidden_states = resnet( + hidden_states, + temb, + image_only_indicator=image_only_indicator, + ) + hidden_states = attn( + hidden_states, + encoder_hidden_states=encoder_hidden_states, + image_only_indicator=image_only_indicator, + return_dict=False, + )[0] + + if self.upsamplers is not None: + for upsampler in self.upsamplers: + hidden_states = upsampler(hidden_states) + + return hidden_states diff --git a/animation/StableAnimator/animation/pipelines/inference_pipeline_animation.py b/animation/StableAnimator/animation/pipelines/inference_pipeline_animation.py new file mode 100644 index 0000000..7d6c412 --- /dev/null +++ b/animation/StableAnimator/animation/pipelines/inference_pipeline_animation.py @@ -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) diff --git a/animation/StableAnimator/animation/pipelines/pipeline_animation.py b/animation/StableAnimator/animation/pipelines/pipeline_animation.py new file mode 100644 index 0000000..7aa438a --- /dev/null +++ b/animation/StableAnimator/animation/pipelines/pipeline_animation.py @@ -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) diff --git a/animation/StableAnimator/animation/pipelines/validation_pipeline_animation.py b/animation/StableAnimator/animation/pipelines/validation_pipeline_animation.py new file mode 100644 index 0000000..86d9332 --- /dev/null +++ b/animation/StableAnimator/animation/pipelines/validation_pipeline_animation.py @@ -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) diff --git a/animation/StableAnimator/animation/utils/__init__.py b/animation/StableAnimator/animation/utils/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/animation/StableAnimator/animation/utils/geglu_patch.py b/animation/StableAnimator/animation/utils/geglu_patch.py new file mode 100644 index 0000000..de0861f --- /dev/null +++ b/animation/StableAnimator/animation/utils/geglu_patch.py @@ -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 diff --git a/animation/StableAnimator/animation/utils/loader.py b/animation/StableAnimator/animation/utils/loader.py new file mode 100644 index 0000000..80e696a --- /dev/null +++ b/animation/StableAnimator/animation/utils/loader.py @@ -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 + diff --git a/animation/StableAnimator/animation/utils/utils.py b/animation/StableAnimator/animation/utils/utils.py new file mode 100644 index 0000000..85db8dc --- /dev/null +++ b/animation/StableAnimator/animation/utils/utils.py @@ -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 \ No newline at end of file diff --git a/animation/StableAnimator/app.py b/animation/StableAnimator/app.py new file mode 100644 index 0000000..8901fde --- /dev/null +++ b/animation/StableAnimator/app.py @@ -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(""" +
+

StableAnimator

+
+
+ 🌐 Github | + πŸ“œ arXiv +
+
+ ⚠️ This demo is for academic research and experiential use only. +
+ """) + 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) diff --git a/animation/StableAnimator/assets/figures/case-17.gif b/animation/StableAnimator/assets/figures/case-17.gif new file mode 100644 index 0000000..e5a632f Binary files /dev/null and b/animation/StableAnimator/assets/figures/case-17.gif differ diff --git a/animation/StableAnimator/assets/figures/case-18.gif b/animation/StableAnimator/assets/figures/case-18.gif new file mode 100644 index 0000000..69777c0 Binary files /dev/null and b/animation/StableAnimator/assets/figures/case-18.gif differ diff --git a/animation/StableAnimator/assets/figures/case-24.gif b/animation/StableAnimator/assets/figures/case-24.gif new file mode 100644 index 0000000..e27342c Binary files /dev/null and b/animation/StableAnimator/assets/figures/case-24.gif differ diff --git a/animation/StableAnimator/assets/figures/case-35.gif b/animation/StableAnimator/assets/figures/case-35.gif new file mode 100644 index 0000000..bdab604 Binary files /dev/null and b/animation/StableAnimator/assets/figures/case-35.gif differ diff --git a/animation/StableAnimator/assets/figures/case-42.gif b/animation/StableAnimator/assets/figures/case-42.gif new file mode 100644 index 0000000..76db98f Binary files /dev/null and b/animation/StableAnimator/assets/figures/case-42.gif differ diff --git a/animation/StableAnimator/assets/figures/case-45.gif b/animation/StableAnimator/assets/figures/case-45.gif new file mode 100644 index 0000000..66db02f Binary files /dev/null and b/animation/StableAnimator/assets/figures/case-45.gif differ diff --git a/animation/StableAnimator/assets/figures/case-46.gif b/animation/StableAnimator/assets/figures/case-46.gif new file mode 100644 index 0000000..faa7d55 Binary files /dev/null and b/animation/StableAnimator/assets/figures/case-46.gif differ diff --git a/animation/StableAnimator/assets/figures/case-47.gif b/animation/StableAnimator/assets/figures/case-47.gif new file mode 100644 index 0000000..8f55bdd Binary files /dev/null and b/animation/StableAnimator/assets/figures/case-47.gif differ diff --git a/animation/StableAnimator/assets/figures/case-5.gif b/animation/StableAnimator/assets/figures/case-5.gif new file mode 100644 index 0000000..060f872 Binary files /dev/null and b/animation/StableAnimator/assets/figures/case-5.gif differ diff --git a/animation/StableAnimator/assets/figures/case-61.gif b/animation/StableAnimator/assets/figures/case-61.gif new file mode 100644 index 0000000..11f3d2b Binary files /dev/null and b/animation/StableAnimator/assets/figures/case-61.gif differ diff --git a/animation/StableAnimator/assets/figures/framework.jpg b/animation/StableAnimator/assets/figures/framework.jpg new file mode 100644 index 0000000..ed44806 Binary files /dev/null and b/animation/StableAnimator/assets/figures/framework.jpg differ diff --git a/animation/StableAnimator/command_basic_infer.sh b/animation/StableAnimator/command_basic_infer.sh new file mode 100644 index 0000000..03bf6f2 --- /dev/null +++ b/animation/StableAnimator/command_basic_infer.sh @@ -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 \ No newline at end of file diff --git a/animation/StableAnimator/command_finetune.sh b/animation/StableAnimator/command_finetune.sh new file mode 100644 index 0000000..c7076ac --- /dev/null +++ b/animation/StableAnimator/command_finetune.sh @@ -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" diff --git a/animation/StableAnimator/command_train.sh b/animation/StableAnimator/command_train.sh new file mode 100644 index 0000000..42eff7c --- /dev/null +++ b/animation/StableAnimator/command_train.sh @@ -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" \ No newline at end of file diff --git a/animation/StableAnimator/command_train_single.sh b/animation/StableAnimator/command_train_single.sh new file mode 100644 index 0000000..87e650f --- /dev/null +++ b/animation/StableAnimator/command_train_single.sh @@ -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" \ No newline at end of file diff --git a/animation/StableAnimator/face_mask_extraction.py b/animation/StableAnimator/face_mask_extraction.py new file mode 100644 index 0000000..3565e27 --- /dev/null +++ b/animation/StableAnimator/face_mask_extraction.py @@ -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}") \ No newline at end of file diff --git a/animation/StableAnimator/inference_basic.py b/animation/StableAnimator/inference_basic.py new file mode 100644 index 0000000..fe3317c --- /dev/null +++ b/animation/StableAnimator/inference_basic.py @@ -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 diff --git a/animation/StableAnimator/requirements.txt b/animation/StableAnimator/requirements.txt new file mode 100644 index 0000000..a2f03af --- /dev/null +++ b/animation/StableAnimator/requirements.txt @@ -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 diff --git a/animation/StableAnimator/train.py b/animation/StableAnimator/train.py new file mode 100644 index 0000000..2bb137b --- /dev/null +++ b/animation/StableAnimator/train.py @@ -0,0 +1,1695 @@ +import argparse +import random +import logging +import math +import os + +import cv2 +import shutil +from pathlib import Path +from urllib.parse import urlparse +import numpy as np +import PIL +from PIL import Image, ImageDraw +import torch +import torch.nn.functional as F +import torch.utils.checkpoint +from diffusers.models.attention_processor import XFormersAttnProcessor + +from animation.dataset.animation_dataset import LargeScaleAnimationVideos +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.validation_pipeline_animation import ValidationAnimationPipeline +import transformers +from accelerate import Accelerator, DistributedType +from accelerate.logging import get_logger +from accelerate.utils import ProjectConfiguration, set_seed +from huggingface_hub import create_repo, upload_folder +from packaging import version +from tqdm.auto import tqdm +from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection +from einops import rearrange + +import datetime +import diffusers +from diffusers import AutoencoderKLTemporalDecoder, EulerDiscreteScheduler +from diffusers.image_processor import VaeImageProcessor +from diffusers.optimization import get_scheduler +from diffusers.training_utils import EMAModel +from diffusers.utils import check_min_version, deprecate, is_wandb_available, load_image +from diffusers.utils.import_utils import is_xformers_available +import warnings +import torch.nn as nn +from diffusers.utils.torch_utils import randn_tensor + +# Will error if the minimal version of diffusers is not installed. Remove at your own risks. +check_min_version("0.24.0.dev0") + +logger = get_logger(__name__, log_level="INFO") + +#i should make a utility function file +def validate_and_convert_image(image, target_size=(256, 256)): + if image is None: + print("Encountered a None image") + return None + + if isinstance(image, torch.Tensor): + # Convert PyTorch tensor to PIL Image + if image.ndim == 3 and image.shape[0] in [1, 3]: # Check for CxHxW format + if image.shape[0] == 1: # Convert single-channel grayscale to RGB + image = image.repeat(3, 1, 1) + image = image.mul(255).clamp(0, 255).byte().permute(1, 2, 0).cpu().numpy() + image = Image.fromarray(image) + else: + print(f"Invalid image tensor shape: {image.shape}") + return None + elif isinstance(image, Image.Image): + # Resize PIL Image + image = image.resize(target_size) + else: + print("Image is not a PIL Image or a PyTorch tensor") + return None + + return image + +def create_image_grid(images, rows, cols, target_size=(256, 256)): + valid_images = [validate_and_convert_image(img, target_size) for img in images] + valid_images = [img for img in valid_images if img is not None] + + if not valid_images: + print("No valid images to create a grid") + return None + + w, h = target_size + grid = Image.new('RGB', size=(cols * w, rows * h)) + + for i, image in enumerate(valid_images): + grid.paste(image, box=((i % cols) * w, (i // cols) * h)) + + return grid + +def save_combined_frames(batch_output, validation_images, validation_control_images,output_folder): + # Flatten batch_output, which is a list of lists of PIL Images + flattened_batch_output = [img for sublist in batch_output for img in sublist] + + # Combine frames into a list without converting (since they are already PIL Images) + combined_frames = validation_images + validation_control_images + flattened_batch_output + + # Calculate rows and columns for the grid + num_images = len(combined_frames) + cols = 3 # adjust number of columns as needed + rows = (num_images + cols - 1) // cols + timestamp = datetime.datetime.now().strftime("%Y%m%d-%H%M%S") + + filename = f"combined_frames_{timestamp}.png" + # Create and save the grid image + grid = create_image_grid(combined_frames, rows, cols) + output_folder = os.path.join(output_folder, "validation_images") + os.makedirs(output_folder, exist_ok=True) + + # Now define the full path for the file + timestamp = datetime.datetime.now().strftime("%Y%m%d-%H%M%S") + filename = f"combined_frames_{timestamp}.png" + output_loc = os.path.join(output_folder, filename) + + if grid is not None: + grid.save(output_loc) + else: + print("Failed to create image grid") + + + +# def load_images_from_folder(folder): +# images = [] +# valid_extensions = {".jpg", ".jpeg", ".png", ".bmp", ".gif", ".tiff"} # Add or remove extensions as needed +# +# # Function to extract frame number from the filename +# def frame_number(filename): +# # First, try the pattern 'frame_x_7fps' +# new_pattern_match = re.search(r'frame_(\d+)_7fps', filename) +# if new_pattern_match: +# return int(new_pattern_match.group(1)) +# # If the new pattern is not found, use the original digit extraction method +# matches = re.findall(r'\d+', filename) +# if matches: +# if matches[-1] == '0000' and len(matches) > 1: +# return int(matches[-2]) # Return the second-to-last sequence if the last is '0000' +# return int(matches[-1]) # Otherwise, return the last sequence +# return float('inf') # Return 'inf' +# +# # Sorting files based on frame number +# sorted_files = sorted(os.listdir(folder), key=frame_number) +# +# # Load images in sorted order +# for filename in sorted_files: +# ext = os.path.splitext(filename)[1].lower() +# if ext in valid_extensions: +# img = Image.open(os.path.join(folder, filename)).convert('RGB') +# images.append(img) +# +# return images + +def load_images_from_folder(folder): + 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') + images.append(img) + + return images + + + +# copy from https://github.com/crowsonkb/k-diffusion.git +def stratified_uniform(shape, group=0, groups=1, dtype=None, device=None): + """Draws stratified samples from a uniform distribution.""" + if groups <= 0: + raise ValueError(f"groups must be positive, got {groups}") + if group < 0 or group >= groups: + raise ValueError(f"group must be in [0, {groups})") + n = shape[-1] * groups + offsets = torch.arange(group, n, groups, dtype=dtype, device=device) + u = torch.rand(shape, dtype=dtype, device=device) + return (offsets + u) / n + + +def rand_cosine_interpolated(shape, image_d, noise_d_low, noise_d_high, sigma_data=1., min_value=1e-3, max_value=1e3, device='cpu', dtype=torch.float32): + """Draws samples from an interpolated cosine timestep distribution (from simple diffusion).""" + + def logsnr_schedule_cosine(t, logsnr_min, logsnr_max): + t_min = math.atan(math.exp(-0.5 * logsnr_max)) + t_max = math.atan(math.exp(-0.5 * logsnr_min)) + return -2 * torch.log(torch.tan(t_min + t * (t_max - t_min))) + + def logsnr_schedule_cosine_shifted(t, image_d, noise_d, logsnr_min, logsnr_max): + shift = 2 * math.log(noise_d / image_d) + return logsnr_schedule_cosine(t, logsnr_min - shift, logsnr_max - shift) + shift + + def logsnr_schedule_cosine_interpolated(t, image_d, noise_d_low, noise_d_high, logsnr_min, logsnr_max): + logsnr_low = logsnr_schedule_cosine_shifted( + t, image_d, noise_d_low, logsnr_min, logsnr_max) + logsnr_high = logsnr_schedule_cosine_shifted( + t, image_d, noise_d_high, logsnr_min, logsnr_max) + return torch.lerp(logsnr_low, logsnr_high, t) + + logsnr_min = -2 * math.log(min_value / sigma_data) + logsnr_max = -2 * math.log(max_value / sigma_data) + u = stratified_uniform( + shape, group=0, groups=1, dtype=dtype, device=device + ) + logsnr = logsnr_schedule_cosine_interpolated( + u, image_d, noise_d_low, noise_d_high, logsnr_min, logsnr_max) + return torch.exp(-logsnr / 2) * sigma_data + +def rand_log_normal(shape, loc=0., scale=1., device='cpu', dtype=torch.float32): + """Draws samples from an lognormal distribution.""" + u = torch.rand(shape, dtype=dtype, device=device) * (1 - 2e-7) + 1e-7 + return torch.distributions.Normal(loc, scale).icdf(u).exp() + +min_value = 0.002 +max_value = 700 +image_d = 64 +noise_d_low = 32 +noise_d_high = 64 +sigma_data = 0.5 + + +def _resize_with_antialiasing(input, size, interpolation="bicubic", align_corners=True): + h, w = input.shape[-2:] + factors = (h / size[0], w / size[1]) + + # First, we have to determine sigma + # Taken from skimage: https://github.com/scikit-image/scikit-image/blob/v0.19.2/skimage/transform/_warps.py#L171 + sigmas = ( + max((factors[0] - 1.0) / 2.0, 0.001), + max((factors[1] - 1.0) / 2.0, 0.001), + ) + + # Now kernel size. Good results are for 3 sigma, but that is kind of slow. Pillow uses 1 sigma + # https://github.com/python-pillow/Pillow/blob/master/src/libImaging/Resample.c#L206 + # But they do it in the 2 passes, which gives better results. Let's try 2 sigmas for now + ks = int(max(2.0 * 2 * sigmas[0], 3)), int(max(2.0 * 2 * sigmas[1], 3)) + + # Make sure it is odd + if (ks[0] % 2) == 0: + ks = ks[0] + 1, ks[1] + + if (ks[1] % 2) == 0: + ks = ks[0], ks[1] + 1 + + input = _gaussian_blur2d(input, ks, sigmas) + + output = torch.nn.functional.interpolate( + input, size=size, mode=interpolation, align_corners=align_corners) + return output + + +def _compute_padding(kernel_size): + """Compute padding tuple.""" + # 4 or 6 ints: (padding_left, padding_right,padding_top,padding_bottom) + # https://pytorch.org/docs/stable/nn.html#torch.nn.functional.pad + if len(kernel_size) < 2: + raise AssertionError(kernel_size) + computed = [k - 1 for k in kernel_size] + + # for even kernels we need to do asymmetric padding :( + out_padding = 2 * len(kernel_size) * [0] + + for i in range(len(kernel_size)): + computed_tmp = computed[-(i + 1)] + + pad_front = computed_tmp // 2 + pad_rear = computed_tmp - pad_front + + out_padding[2 * i + 0] = pad_front + out_padding[2 * i + 1] = pad_rear + + return out_padding + + +def _filter2d(input, kernel): + # prepare kernel + b, c, h, w = input.shape + tmp_kernel = kernel[:, None, ...].to( + device=input.device, dtype=input.dtype) + + tmp_kernel = tmp_kernel.expand(-1, c, -1, -1) + + height, width = tmp_kernel.shape[-2:] + + padding_shape: list[int] = _compute_padding([height, width]) + input = torch.nn.functional.pad(input, padding_shape, mode="reflect") + + # kernel and input tensor reshape to align element-wise or batch-wise params + tmp_kernel = tmp_kernel.reshape(-1, 1, height, width) + input = input.view(-1, tmp_kernel.size(0), input.size(-2), input.size(-1)) + + # convolve the tensor with the kernel. + output = torch.nn.functional.conv2d( + input, tmp_kernel, groups=tmp_kernel.size(0), padding=0, stride=1) + + out = output.view(b, c, h, w) + return out + + +def _gaussian(window_size: int, sigma): + if isinstance(sigma, float): + sigma = torch.tensor([[sigma]]) + + batch_size = sigma.shape[0] + + x = (torch.arange(window_size, device=sigma.device, + dtype=sigma.dtype) - window_size // 2).expand(batch_size, -1) + + if window_size % 2 == 0: + x = x + 0.5 + + gauss = torch.exp(-x.pow(2.0) / (2 * sigma.pow(2.0))) + + return gauss / gauss.sum(-1, keepdim=True) + + +def _gaussian_blur2d(input, kernel_size, sigma): + if isinstance(sigma, tuple): + sigma = torch.tensor([sigma], dtype=input.dtype) + else: + sigma = sigma.to(dtype=input.dtype) + + ky, kx = int(kernel_size[0]), int(kernel_size[1]) + bs = sigma.shape[0] + kernel_x = _gaussian(kx, sigma[:, 1].view(bs, 1)) + kernel_y = _gaussian(ky, sigma[:, 0].view(bs, 1)) + out_x = _filter2d(input, kernel_x[..., None, :]) + out = _filter2d(out_x, kernel_y[..., None]) + + return out + + +def export_to_video(video_frames, output_video_path, fps): + fourcc = cv2.VideoWriter_fourcc(*"mp4v") + h, w, _ = video_frames[0].shape + video_writer = cv2.VideoWriter( + output_video_path, fourcc, fps=fps, frameSize=(w, h)) + for i in range(len(video_frames)): + img = cv2.cvtColor(video_frames[i], cv2.COLOR_RGB2BGR) + video_writer.write(img) + + +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 tensor_to_vae_latent(t, vae, scale=True): + t = t.to(vae.dtype) + if len(t.shape) == 5: + video_length = t.shape[1] + + t = rearrange(t, "b f c h w -> (b f) c h w") + latents = vae.encode(t).latent_dist.sample() + latents = rearrange(latents, "(b f) c h w -> b f c h w", f=video_length) + elif len(t.shape) == 4: + latents = vae.encode(t).latent_dist.sample() + if scale: + latents = latents * vae.config.scaling_factor + return latents + + +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, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + + parser.add_argument( + "--num_frames", + type=int, + default=14, + ) + parser.add_argument( + "--dataset_type", + type=str, + default='ubc', + ) + parser.add_argument( + "--num_validation_images", + type=int, + default=1, + help="Number of images that should be generated during validation with `validation_prompt`.", + ) + parser.add_argument( + "--validation_steps", + type=int, + default=500, + help=( + "Run fine-tuning validation every X epochs. The validation process consists of running the text/image prompt" + " multiple times: `args.num_validation_images`." + ), + ) + parser.add_argument( + "--output_dir", + type=str, + default="./outputs", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument( + "--seed", type=int, default=None, help="A seed for reproducible training." + ) + parser.add_argument( + "--per_gpu_batch_size", + type=int, + default=1, + help="Batch size (per device) for the training dataloader.", + ) + parser.add_argument("--num_train_epochs", type=int, default=100) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + 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( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", + type=int, + default=500, + help="Number of steps for the warmup in the lr scheduler.", + ) + parser.add_argument( + "--conditioning_dropout_prob", + type=float, + default=0.1, + help="Conditioning dropout probability. Drops out the conditionings (image and edit prompt) used in training InstructPix2Pix. See section 3.2.1 in the paper: https://arxiv.org/abs/2211.09800.", + ) + parser.add_argument( + "--use_8bit_adam", + action="store_true", + help="Whether or not to use 8-bit Adam from bitsandbytes.", + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices" + ), + ) + parser.add_argument( + "--use_ema", action="store_true", help="Whether to use EMA model." + ) + parser.add_argument( + "--non_ema_revision", + type=str, + default=None, + required=False, + help=( + "Revision of pretrained non-ema model identifier. Must be a branch, tag or git identifier of the local or" + " remote repository specified with --pretrained_model_name_or_path." + ), + ) + parser.add_argument( + "--num_workers", + type=int, + default=8, + help=( + "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process." + ), + ) + parser.add_argument( + "--adam_beta1", + type=float, + default=0.9, + help="The beta1 parameter for the Adam optimizer.", + ) + parser.add_argument( + "--adam_beta2", + type=float, + default=0.999, + help="The beta2 parameter for the Adam optimizer.", + ) + parser.add_argument( + "--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use." + ) + parser.add_argument( + "--adam_epsilon", + type=float, + default=1e-08, + help="Epsilon value for the Adam optimizer", + ) + parser.add_argument( + "--max_grad_norm", default=1.0, type=float, help="Max gradient norm." + ) + parser.add_argument( + "--push_to_hub", + action="store_true", + help="Whether or not to push the model to the Hub.", + ) + parser.add_argument( + "--hub_token", + type=str, + default=None, + help="The token to use to push to the Model Hub.", + ) + parser.add_argument( + "--hub_model_id", + type=str, + default=None, + help="The name of the repository to keep in sync with the local `output_dir`.", + ) + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the" + " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config." + ), + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument( + "--local_rank", + type=int, + default=-1, + help="For distributed training: local_rank", + ) + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints are only suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=1, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument( + "--enable_xformers_memory_efficient_attention", + action="store_true", + help="Whether or not to use xformers.", + ) + parser.add_argument( + "--log_trainable_parameters", + action="store_true", + help="Whether to write the trainable parameters.", + ) + parser.add_argument( + "--pretrain_unet", + type=str, + default=None, + help="use weight for unet block", + ) + parser.add_argument( + "--rank", + type=int, + default=128, + help=("The dimension of the LoRA update matrices."), + ) + parser.add_argument( + "--csv_path", + type=str, + default=None, + help=( + "path to the dataset csv" + ), + ) + parser.add_argument( + "--video_folder", + type=str, + default=None, + help=( + "path to the video folder" + ), + ) + parser.add_argument( + "--condition_folder", + type=str, + default=None, + help=( + "path to the depth folder" + ), + ) + parser.add_argument( + "--motion_folder", + type=str, + default=None, + help=( + "path to the depth folder" + ), + ) + parser.add_argument( + "--validation_prompt", + type=str, + default=None, + help=( + "A set of prompts evaluated every `--validation_steps` and logged to `--report_to`." + " Provide either a matching number of `--validation_image`s, a single `--validation_image`" + " to be used with all prompts, or a single prompt that will be used with all `--validation_image`s." + ), + ) + parser.add_argument( + "--validation_image_folder", + 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_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( + "--sample_n_frames", + type=int, + default=14, + help=( + "the sample_n_frames" + ), + ) + + parser.add_argument( + "--ref_augment", + action="store_true", + help=( + "use augment for the reference image" + ), + ) + parser.add_argument( + "--train_stage", + type=int, + default=2, + help=( + "the training stage" + ), + ) + + 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( + "--data_root_path", + type=str, + default=None, + help="Path to the data root path", + ) + parser.add_argument( + "--rec_data_path", + type=str, + default=None, + help="Path to the rec data path", + ) + parser.add_argument( + "--vec_data_path", + type=str, + default=None, + help="Path to the vec data path", + ) + + parser.add_argument( + "--finetune_mode", + type=bool, + default=False, + help="Enable or disable the finetune mode (True/False).", + ) + parser.add_argument( + "--posenet_model_finetune_path", + type=str, + default=None, + help="Path to the pretrained posenet model", + ) + parser.add_argument( + "--face_encoder_finetune_path", + type=str, + default=None, + help="Path to the pretrained face encoder", + ) + parser.add_argument( + "--unet_model_finetune_path", + type=str, + default=None, + help="Path to the pretrained unet model", + ) + + args = parser.parse_args() + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + if env_local_rank != -1 and env_local_rank != args.local_rank: + args.local_rank = env_local_rank + + # default to using the same revision for the non-ema model if not specified + if args.non_ema_revision is None: + args.non_ema_revision = args.revision + + return args + + +def download_image(url): + original_image = ( + lambda image_url_or_path: load_image(image_url_or_path) + if urlparse(image_url_or_path).scheme + else PIL.Image.open(image_url_or_path).convert("RGB") + )(url) + return original_image + + +# This is for training using deepspeed. +# Since now the DeepSpeed only supports trainging with only one model +# So we create a virtual wrapper to contail all the models + +class DeepSpeedWrapperModel(nn.Module): + def __init__(self, **kwargs): + super().__init__() + for name, value in kwargs.items(): + assert isinstance(value, nn.Module) + self.register_module(name, value) + + +def main(): + + warnings.filterwarnings('ignore', category=DeprecationWarning) + warnings.filterwarnings('ignore', category=FutureWarning) + torch.multiprocessing.set_start_method('spawn') + + args = parse_args() + + if args.non_ema_revision is not None: + deprecate( + "non_ema_revision!=None", + "0.15.0", + message=( + "Downloading 'non_ema' weights from revision branches of the Hub is deprecated. Please make sure to" + " use `--variant=non_ema` instead." + ), + ) + logging_dir = os.path.join(args.output_dir, args.logging_dir) + accelerator_project_config = ProjectConfiguration( + project_dir=args.output_dir, logging_dir=logging_dir) + # ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True) + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + project_config=accelerator_project_config, + ) + + generator = torch.Generator( + device=accelerator.device).manual_seed(23123134) + + if args.report_to == "wandb": + if not is_wandb_available(): + raise ImportError( + "Make sure to install wandb if you want to use it for logging during training.") + import wandb + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed) + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + if args.push_to_hub: + repo_id = create_repo( + repo_id=args.hub_model_id or Path(args.output_dir).name, exist_ok=True, token=args.hub_token + ).repo_id + + # Load scheduler, tokenizer and models. + print(args.pretrained_model_name_or_path) + 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, variant="fp16") + unet = UNetSpatioTemporalConditionModel.from_pretrained( + args.pretrained_model_name_or_path if args.pretrain_unet is None else args.pretrain_unet, + subfolder="unet", + low_cpu_mem_usage=True, + variant="fp16" + ) + 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=1024, + num_tokens=4,) + face_model = FaceModel() + + # init adapter modules + 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 AnimationIDAttnNormalizedProcessor: {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) + + # triggering the finetune mode + if args.finetune_mode is True and args.posenet_model_finetune_path is not None and args.face_encoder_finetune_path is not None and args.unet_model_finetune_path is not None: + print("Loading existing posenet weights, face_encoder weights and unet weights.") + if args.posenet_model_finetune_path.endswith(".pth"): + pose_net_state_dict = torch.load(args.posenet_model_finetune_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_finetune_path.endswith(".pth"): + face_encoder_state_dict = torch.load(args.face_encoder_finetune_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_finetune_path.endswith(".pth"): + unet_state_dict = torch.load(args.unet_model_finetune_path, map_location="cpu") + unet.load_state_dict(unet_state_dict, strict=True) + else: + print("unet weights loading fail") + print(1/0) + + + vae_scale_factor = 2 ** (len(vae.config.block_out_channels) - 1) + image_processor = VaeImageProcessor(vae_scale_factor=vae_scale_factor) + + # Freeze vae and image_encoder + vae.requires_grad_(False) + image_encoder.requires_grad_(False) + unet.requires_grad_(False) + pose_net.requires_grad_(False) + face_encoder.requires_grad_(False) + + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + + image_encoder.to(accelerator.device, dtype=weight_dtype) + vae.to(accelerator.device, dtype=weight_dtype) + + if args.use_ema: + ema_unet = EMAModel(unet.parameters( + ), model_cls=UNetSpatioTemporalConditionModel, model_config=unet.config) + + if args.enable_xformers_memory_efficient_attention: + if is_xformers_available(): + import xformers + xformers_version = version.parse(xformers.__version__) + if xformers_version == version.parse("0.0.16"): + logger.warn( + "xFormers 0.0.16 cannot be used for training in some GPUs. If you observe problems during training, please update xFormers to at least 0.0.17. See https://huggingface.co/docs/diffusers/main/en/optimization/xformers for more details." + ) + unet.enable_xformers_memory_efficient_attention() + else: + raise ValueError( + "xformers is not available. Make sure it is installed correctly") + + + if args.gradient_checkpointing: + unet.enable_gradient_checkpointing() + + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * + args.per_gpu_batch_size * accelerator.num_processes + ) + + # Initialize the optimizer + if args.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError( + "Please install bitsandbytes to use 8-bit Adam. You can do so by running `pip install bitsandbytes`" + ) + + optimizer_cls = bnb.optim.AdamW8bit + else: + optimizer_cls = torch.optim.AdamW + + # if accelerator.distributed_type == DistributedType.DEEPSPEED: + # ds_wrapper = DeepSpeedWrapperModel( + # unet=unet, + # controlnext=controlnext + # ) + # unet = ds_wrapper.unet + # controlnext = ds_wrapper.controlnext + + + pose_net.requires_grad_(True) + face_encoder.requires_grad_(True) + + parameters_list = [] + + for name, para in pose_net.named_parameters(): + para.requires_grad = True + parameters_list.append({"params": para, "lr": args.learning_rate } ) + + for name, para in face_encoder.named_parameters(): + para.requires_grad = True + parameters_list.append({"params": para, "lr": args.learning_rate } ) + + + """ + For more details, please refer to: https://github.com/dvlab-research/ControlNeXt/issues/14#issuecomment-2290450333 + This is the selective parameters part. + As presented in our paper, we only select a small subset of parameters, which is fully adapted to the SD1.5 and SDXL backbones. By training fewer than 100 million parameters, we still achieve excellent performance. But this is is not suitable for the SD3 and SVD training. This is because, after SDXL, Stability faced significant legal risks due to the generation of highly realistic human images. After that, they stopped refining their models on human-related data, such as SVD and SD3, to avoid potential risks. + To achieve optimal performance, it's necessary to first continue training SVD and SD3 on human-related data to develop a robust backbone before fine-tuning. Of course, you can also combine the continual pretraining and finetuning. So you can find that we direct provide the full SVD parameters. + We have experimented with two approaches: 1.Directly training the model from scratch on human dancing data. 2. Continual training using a pre-trained human generation backbone, followed by fine-tuning a selective small subset of parameters. Interestingly, we observed no significant difference in performance between these two methods. + """ + + for name, para in unet.named_parameters(): + if "attentions" in name: + para.requires_grad = True + parameters_list.append({"params": para}) + else: + para.requires_grad = False + + optimizer = optimizer_cls( + parameters_list, + lr=args.learning_rate, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + ) + + # check para + if accelerator.is_main_process and args.log_trainable_parameters: + rec_txt1 = open('rec_para.txt', 'w') + rec_txt2 = open('rec_para_train.txt', 'w') + for name, para in unet.named_parameters(): + if para.requires_grad is False: + rec_txt1.write(f'{name}\n') + else: + rec_txt2.write(f'{name}\n') + rec_txt1.close() + rec_txt2.close() + # DataLoaders creation: + args.global_batch_size = args.per_gpu_batch_size * accelerator.num_processes + + root_path = args.data_root_path + txt_path_1 = args.rec_data_path + txt_path_2 = args.vec_data_path + train_dataset_1 = LargeScaleAnimationVideos( + root_path=root_path, + txt_path=txt_path_1, + width=512, + height=512, + n_sample_frames=args.sample_n_frames, + sample_frame_rate=4, + app=face_model.app, + handler_ante=face_model.handler_ante, + face_helper=face_model.face_helper + ) + train_dataloader_1 = torch.utils.data.DataLoader( + train_dataset_1, + batch_size=args.per_gpu_batch_size, + num_workers=args.num_workers, + shuffle=True, + ) + train_dataset_2 = LargeScaleAnimationVideos( + root_path=root_path, + txt_path=txt_path_2, + width=576, + height=1024, + n_sample_frames=args.sample_n_frames, + sample_frame_rate=4, + app=face_model.app, + handler_ante=face_model.handler_ante, + face_helper=face_model.face_helper + ) + train_dataloader_2 = torch.utils.data.DataLoader( + train_dataset_2, + batch_size=args.per_gpu_batch_size, + num_workers=args.num_workers, + shuffle=True + ) + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil((len(train_dataloader_1) + len(train_dataloader_2)) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + ) + + unet, pose_net, face_encoder, optimizer, lr_scheduler, train_dataloader_1, train_dataloader_2 = accelerator.prepare( + unet, pose_net, face_encoder, optimizer, lr_scheduler, train_dataloader_1, train_dataloader_2 + ) + + if args.use_ema: + ema_unet.to(accelerator.device) + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil((len(train_dataloader_1) + len(train_dataloader_2)) / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil( + args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + accelerator.init_trackers("StableAnimator", config=vars(args)) + + # Train! + total_batch_size = args.per_gpu_batch_size * \ + accelerator.num_processes * args.gradient_accumulation_steps + + len_zeros = len(train_dataloader_1) + len_ones = len(train_dataloader_2) + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset_1)+len(train_dataset_2)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info( + f" Instantaneous batch size per device = {args.per_gpu_batch_size}") + logger.info( + f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info( + f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + def encode_image(pixel_values): + pixel_values = _resize_with_antialiasing(pixel_values, (224, 224)) + pixel_values = (pixel_values + 1.0) / 2.0 + + pixel_values = pixel_values.to(torch.float32) + # Normalize the image with for CLIP input + pixel_values = feature_extractor( + images=pixel_values, + do_normalize=True, + do_center_crop=False, + do_resize=False, + do_rescale=False, + return_tensors="pt", + ).pixel_values + + pixel_values = pixel_values.to( + device=accelerator.device, dtype=image_encoder.dtype) + image_embeddings = image_encoder(pixel_values).image_embeds + image_embeddings= image_embeddings.unsqueeze(1) + return image_embeddings + + + def _get_add_time_ids( + fps, + motion_bucket_id, + noise_aug_strength, + dtype, + batch_size, + unet=None, + device=None + ): + add_time_ids = [fps, motion_bucket_id, noise_aug_strength] + + + add_time_ids = torch.tensor([add_time_ids], dtype=dtype, device=device) + add_time_ids = add_time_ids.repeat(batch_size, 1) + return add_time_ids + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the most recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + else: + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + global_step = int(path.split("-")[1]) + + resume_global_step = global_step * args.gradient_accumulation_steps + first_epoch = global_step // num_update_steps_per_epoch + resume_step = resume_global_step % ( + num_update_steps_per_epoch * args.gradient_accumulation_steps) + + # Only show the progress bar once on each machine. + progress_bar = tqdm(range(global_step, args.max_train_steps), + disable=not accelerator.is_local_main_process) + progress_bar.set_description("Steps") + + for epoch in range(first_epoch, args.num_train_epochs): + pose_net.train() + face_encoder.train() + unet.train() + train_loss = 0.0 + + iter1 = iter(train_dataloader_1) + iter2 = iter(train_dataloader_2) + list_0_and_1 = [0] * len_zeros + [1] * len_ones + random.shuffle(list_0_and_1) + for step in range(0, len(list_0_and_1)): + current_idx = list_0_and_1[step] + if current_idx == 0: + try: + batch = next(iter1) + except StopIteration: + iter1 = iter(train_dataloader_1) + batch = next(iter1) + elif current_idx == 1: + try: + batch = next(iter2) + except StopIteration: + iter2 = iter(train_dataloader_2) + batch = next(iter2) + + # Skip steps until we reach the resumed step + if args.resume_from_checkpoint and epoch == first_epoch and step < resume_step: + if step % args.gradient_accumulation_steps == 0: + progress_bar.update(1) + continue + + with accelerator.accumulate(pose_net, face_encoder, unet): + with accelerator.autocast(): + pixel_values = batch["pixel_values"].to(weight_dtype).to( + accelerator.device, non_blocking=True + ) + conditional_pixel_values = batch["reference_image"].to(weight_dtype).to( + accelerator.device, non_blocking=True + ) + + latents = tensor_to_vae_latent(pixel_values, vae).to(dtype=weight_dtype) + + # Get the text embedding for conditioning. + encoder_hidden_states = encode_image(conditional_pixel_values).to(dtype=weight_dtype) + image_embed = encoder_hidden_states.clone() + + train_noise_aug = 0.02 + conditional_pixel_values = conditional_pixel_values + train_noise_aug * \ + randn_tensor(conditional_pixel_values.shape, generator=generator, device=conditional_pixel_values.device, dtype=conditional_pixel_values.dtype) + conditional_latents = tensor_to_vae_latent(conditional_pixel_values, vae, scale=False) + + # Sample noise that we'll add to the latents + noise = torch.randn_like(latents) + bsz = latents.shape[0] + # Sample a random timestep for each image + sigmas = rand_cosine_interpolated(shape=[bsz,], image_d=image_d, noise_d_low=noise_d_low, noise_d_high=noise_d_high, sigma_data=sigma_data, min_value=min_value, max_value=max_value).to(latents.device, dtype=weight_dtype) + + # sigmas = rand_log_normal(shape=[bsz,], loc=0.7, scale=1.6).to(latents) + # Add noise to the latents according to the noise magnitude at each timestep + # (this is the forward diffusion process) + sigmas_reshaped = sigmas.clone() + while len(sigmas_reshaped.shape) < len(latents.shape): + sigmas_reshaped = sigmas_reshaped.unsqueeze(-1) + + + noisy_latents = latents + noise * sigmas_reshaped + + timesteps = torch.Tensor([0.25 * sigma.log() for sigma in sigmas]).to(latents.device, dtype=weight_dtype) + + + inp_noisy_latents = noisy_latents / ((sigmas_reshaped**2 + 1) ** 0.5) + + added_time_ids = _get_add_time_ids( + fps=6, + motion_bucket_id=127.0, + noise_aug_strength=train_noise_aug, # noise_aug_strength == 0.0 + dtype=encoder_hidden_states.dtype, + batch_size=bsz, + unet=unet, + device=latents.device + ) + + added_time_ids = added_time_ids.to(latents.device) + + # Conditioning dropout to support classifier-free guidance during inference. For more details + # check out the section 3.2.1 of the original paper https://arxiv.org/abs/2211.09800. + if args.conditioning_dropout_prob is not None: + random_p = torch.rand( + bsz, device=latents.device, generator=generator) + # Sample masks for the edit prompts. + prompt_mask = random_p < 2 * args.conditioning_dropout_prob + prompt_mask = prompt_mask.reshape(bsz, 1, 1) + # Final text conditioning. + null_conditioning = torch.zeros_like(encoder_hidden_states) + encoder_hidden_states = torch.where( + prompt_mask, null_conditioning, encoder_hidden_states) + + # Sample masks for the original images. + image_mask_dtype = conditional_latents.dtype + image_mask = 1 - ( + (random_p >= args.conditioning_dropout_prob).to( + image_mask_dtype) + * (random_p < 3 * args.conditioning_dropout_prob).to(image_mask_dtype) + ) + image_mask = image_mask.reshape(bsz, 1, 1, 1) + # Final image conditioning. + conditional_latents = image_mask * conditional_latents + + # Concatenate the `conditional_latents` with the `noisy_latents`. + conditional_latents = conditional_latents.unsqueeze( + 1).repeat(1, noisy_latents.shape[1], 1, 1, 1) + + pose_pixels = batch["pose_pixels"].to( + dtype=weight_dtype, device=accelerator.device, non_blocking=True + ) + faceid_embeds = batch["faceid_embeds"].to( + dtype=weight_dtype, device=accelerator.device, non_blocking=True + ) + pose_latents = pose_net(pose_pixels) + + # print("This is faceid_latents calculation") + # print(faceid_embeds.size()) # [1, 512] + # print(image_embed.size()) # [1, 1, 1024] + + faceid_latents = face_encoder(faceid_embeds, image_embed) + + + inp_noisy_latents = torch.cat( + [inp_noisy_latents, conditional_latents], dim=2) + target = latents + + # print(f"the size of encoder_hidden_states: {encoder_hidden_states.size()}") # [1, 1, 1024] + # print(f"the size of face latents: {faceid_latents.size()}") # [1, 4, 1024] + encoder_hidden_states = torch.cat([encoder_hidden_states, faceid_latents], dim=1) + + encoder_hidden_states = encoder_hidden_states.to(latents.dtype) + inp_noisy_latents = inp_noisy_latents.to(latents.dtype) + pose_latents = pose_latents.to(latents.dtype) + + # Predict the noise residual + model_pred = unet( + inp_noisy_latents, timesteps, encoder_hidden_states, + added_time_ids=added_time_ids, + pose_latents=pose_latents, + ).sample + + + sigmas = sigmas_reshaped + # Denoise the latents + c_out = -sigmas / ((sigmas**2 + 1)**0.5) + c_skip = 1 / (sigmas**2 + 1) + denoised_latents = model_pred * c_out + c_skip * noisy_latents + weighing = (1 + sigmas ** 2) * (sigmas**-2.0) + + tgt_face_masks = batch["tgt_face_masks"].to( + dtype=weight_dtype, device=accelerator.device, non_blocking=True + ) + tgt_face_masks = rearrange(tgt_face_masks, "b f c h w -> (b f) c h w") + tgt_face_masks = F.interpolate(tgt_face_masks, size=(target.size()[-2], target.size()[-1]), mode='nearest') + tgt_face_masks = rearrange(tgt_face_masks, "(b f) c h w -> b f c h w", f=args.sample_n_frames) + + # MSE loss + loss = torch.mean( + (weighing.float() * (denoised_latents.float() - + target.float()) ** 2 * (1 + tgt_face_masks)).reshape(target.shape[0], -1), + dim=1, + ) + loss = loss.mean() + + # Gather the losses across all processes for logging (if we use distributed training). + avg_loss = accelerator.gather( + loss.repeat(args.per_gpu_batch_size)).mean() + train_loss += avg_loss.item() / args.gradient_accumulation_steps + + # Backpropagate + accelerator.backward(loss) + # if accelerator.sync_gradients: + # accelerator.clip_grad_norm_(unet.parameters(), args.max_grad_norm) + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad() + + with torch.cuda.device(latents.device): + torch.cuda.empty_cache() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + if args.use_ema: + ema_unet.step(unet.parameters()) + progress_bar.update(1) + global_step += 1 + accelerator.log({"train_loss": train_loss}, step=global_step) + train_loss = 0.0 + + # save checkpoints! + # if global_step % args.checkpointing_steps == 0 and (accelerator.is_main_process or accelerator.distributed_type == DistributedType.DEEPSPEED): + if global_step % args.checkpointing_steps == 0 and accelerator.is_main_process: + + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None and accelerator.is_main_process: + checkpoints = os.listdir(args.output_dir) + checkpoints = [ + d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted( + checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len( + checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info( + f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join( + args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + save_path = os.path.join( + args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + unwrap_unet = accelerator.unwrap_model(unet) + unwrap_pose_net = accelerator.unwrap_model(pose_net) + unwrap_face_encoder = accelerator.unwrap_model(face_encoder) + unwrap_unet_state_dict = unwrap_unet.state_dict() + torch.save(unwrap_unet_state_dict, os.path.join(args.output_dir, f"checkpoint-{global_step}", f"unet-{global_step}.pth")) + unwrap_pose_net_state_dict = unwrap_pose_net.state_dict() + torch.save(unwrap_pose_net_state_dict, os.path.join(args.output_dir, f"checkpoint-{global_step}", f"pose_net-{global_step}.pth")) + unwrap_face_encoder_state_dict = unwrap_face_encoder.state_dict() + torch.save(unwrap_face_encoder_state_dict, os.path.join(args.output_dir, f"checkpoint-{global_step}", f"face_encoder-{global_step}.pth")) + logger.info(f"Saved state to {save_path}") + + if accelerator.is_main_process: + # sample images! + if global_step % args.validation_steps == 0: + logger.info( + f"Running validation... \n Generating {args.num_validation_images} videos." + ) + # create pipeline + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_unet.store(unet.parameters()) + ema_unet.copy_to(unet.parameters()) + + log_validation( + vae=vae, + image_encoder=image_encoder, + unet=unet, + pose_net=pose_net, + face_encoder=face_encoder, + app=face_model.app, + face_helper=face_model.face_helper, + handler_ante=face_model.handler_ante, + scheduler=noise_scheduler, + accelerator=accelerator, + feature_extractor=feature_extractor, + width=512, + height=512, + torch_dtype=weight_dtype, + validation_image_folder=args.validation_image_folder, + validation_image=args.validation_image, + validation_control_folder=args.validation_control_folder, + output_dir=args.output_dir, + generator=generator, + global_step=global_step, + num_validation_cases=1, + ) + + if args.use_ema: + # Switch back to the original UNet parameters. + ema_unet.restore(unet.parameters()) + + with torch.cuda.device(latents.device): + torch.cuda.empty_cache() + + logs = {"step_loss": loss.detach().item( + ), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if global_step >= args.max_train_steps: + break + + + # save checkpoints! + # if accelerator.is_main_process or accelerator.distributed_type == DistributedType.DEEPSPEED: + if accelerator.is_main_process: + save_path = os.path.join( + args.output_dir, f"checkpoint-last") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + +def log_validation( + vae, + image_encoder, + unet, + pose_net, + face_encoder, + app, + face_helper, + handler_ante, + scheduler, + accelerator, + feature_extractor, + width, + height, + torch_dtype, + validation_image_folder, + validation_image, + validation_control_folder, + output_dir, + generator, + global_step, + num_validation_cases=1, +): + logger.info("Running validation... ") + validation_unet = accelerator.unwrap_model(unet) + validation_image_encoder = accelerator.unwrap_model(image_encoder) + validation_vae = accelerator.unwrap_model(vae) + validation_pose_net = accelerator.unwrap_model(pose_net) + validation_face_encoder = accelerator.unwrap_model(face_encoder) + + pipeline = ValidationAnimationPipeline( + vae=validation_vae, + image_encoder=validation_image_encoder, + unet=validation_unet, + scheduler=scheduler, + feature_extractor=feature_extractor, + pose_net=validation_pose_net, + face_encoder=validation_face_encoder, + ) + pipeline = pipeline.to(accelerator.device) + validation_images = load_images_from_folder(validation_image_folder) + validation_image_path = validation_image + if validation_image is None: + validation_image = validation_images[0] + else: + validation_image = Image.open(validation_image).convert('RGB') + validation_control_images = load_images_from_folder(validation_control_folder) + + val_save_dir = os.path.join(output_dir, "validation_images") + if not os.path.exists(val_save_dir): + os.makedirs(val_save_dir) + + with accelerator.autocast(): + for val_img_idx in range(num_validation_cases): + # num_frames = args.num_frames + num_frames = len(validation_control_images) + + 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 = app.get(validation_face) + 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_helper.read_image(validation_face) + face_helper.get_face_landmarks_5(only_center_face=True) + face_helper.align_warp_face() + + if len(face_helper.cropped_faces) == 0: + validation_image_id_ante_embedding = np.zeros((512,)) + else: + validation_image_align_face = face_helper.cropped_faces[0] + print('fail to detect face using insightface, extract embedding on align face') + validation_image_id_ante_embedding = handler_ante.get_feat(validation_image_align_face) + + video_frames = pipeline( + image=validation_image, + image_pose=validation_control_images, + height=height, + width=width, + num_frames=num_frames, + tile_size=num_frames, + tile_overlap=4, + decode_chunk_size=4, + motion_bucket_id=127., + fps=7, + min_guidance_scale=3, + max_guidance_scale=3, + noise_aug_strength=0.02, + num_inference_steps=25, + generator=generator, + output_type="pil", + validation_image_id_ante_embedding=validation_image_id_ante_embedding, + ).frames[0] + # save_combined_frames(video_frames, validation_images, validation_control_images, val_save_dir) + + out_file = os.path.join( + val_save_dir, + f"step_{global_step}_val_img_{val_img_idx}.mp4", + ) + # print(video_frames.size()) # [16, 3, 512, 512] + for i in range(num_frames): + img = video_frames[i] + video_frames[i] = np.array(img) + export_to_gif(video_frames, out_file, 8) + + del pipeline + torch.cuda.empty_cache() + + + +if __name__ == "__main__": + main() diff --git a/animation/StableAnimator/train_single.py b/animation/StableAnimator/train_single.py new file mode 100644 index 0000000..880e8be --- /dev/null +++ b/animation/StableAnimator/train_single.py @@ -0,0 +1,1670 @@ +import argparse +import random +import logging +import math +import os + +import cv2 +import shutil +from pathlib import Path +from urllib.parse import urlparse +import numpy as np +import PIL +from PIL import Image, ImageDraw +import torch +import torch.nn.functional as F +import torch.utils.checkpoint +from diffusers.models.attention_processor import XFormersAttnProcessor + +from animation.dataset.animation_dataset import LargeScaleAnimationVideos +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.validation_pipeline_animation import ValidationAnimationPipeline +import transformers +from accelerate import Accelerator, DistributedType +from accelerate.logging import get_logger +from accelerate.utils import ProjectConfiguration, set_seed +from huggingface_hub import create_repo, upload_folder +from packaging import version +from tqdm.auto import tqdm +from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection +from einops import rearrange + +import datetime +import diffusers +from diffusers import AutoencoderKLTemporalDecoder, EulerDiscreteScheduler +from diffusers.image_processor import VaeImageProcessor +from diffusers.optimization import get_scheduler +from diffusers.training_utils import EMAModel +from diffusers.utils import check_min_version, deprecate, is_wandb_available, load_image +from diffusers.utils.import_utils import is_xformers_available +import warnings +import torch.nn as nn +from diffusers.utils.torch_utils import randn_tensor + +# Will error if the minimal version of diffusers is not installed. Remove at your own risks. +check_min_version("0.24.0.dev0") + +logger = get_logger(__name__, log_level="INFO") + + +# i should make a utility function file +def validate_and_convert_image(image, target_size=(256, 256)): + if image is None: + print("Encountered a None image") + return None + + if isinstance(image, torch.Tensor): + # Convert PyTorch tensor to PIL Image + if image.ndim == 3 and image.shape[0] in [1, 3]: # Check for CxHxW format + if image.shape[0] == 1: # Convert single-channel grayscale to RGB + image = image.repeat(3, 1, 1) + image = image.mul(255).clamp(0, 255).byte().permute(1, 2, 0).cpu().numpy() + image = Image.fromarray(image) + else: + print(f"Invalid image tensor shape: {image.shape}") + return None + elif isinstance(image, Image.Image): + # Resize PIL Image + image = image.resize(target_size) + else: + print("Image is not a PIL Image or a PyTorch tensor") + return None + + return image + + +def create_image_grid(images, rows, cols, target_size=(256, 256)): + valid_images = [validate_and_convert_image(img, target_size) for img in images] + valid_images = [img for img in valid_images if img is not None] + + if not valid_images: + print("No valid images to create a grid") + return None + + w, h = target_size + grid = Image.new('RGB', size=(cols * w, rows * h)) + + for i, image in enumerate(valid_images): + grid.paste(image, box=((i % cols) * w, (i // cols) * h)) + + return grid + + +def save_combined_frames(batch_output, validation_images, validation_control_images, output_folder): + # Flatten batch_output, which is a list of lists of PIL Images + flattened_batch_output = [img for sublist in batch_output for img in sublist] + + # Combine frames into a list without converting (since they are already PIL Images) + combined_frames = validation_images + validation_control_images + flattened_batch_output + + # Calculate rows and columns for the grid + num_images = len(combined_frames) + cols = 3 # adjust number of columns as needed + rows = (num_images + cols - 1) // cols + timestamp = datetime.datetime.now().strftime("%Y%m%d-%H%M%S") + + filename = f"combined_frames_{timestamp}.png" + # Create and save the grid image + grid = create_image_grid(combined_frames, rows, cols) + output_folder = os.path.join(output_folder, "validation_images") + os.makedirs(output_folder, exist_ok=True) + + # Now define the full path for the file + timestamp = datetime.datetime.now().strftime("%Y%m%d-%H%M%S") + filename = f"combined_frames_{timestamp}.png" + output_loc = os.path.join(output_folder, filename) + + if grid is not None: + grid.save(output_loc) + else: + print("Failed to create image grid") + + +# def load_images_from_folder(folder): +# images = [] +# valid_extensions = {".jpg", ".jpeg", ".png", ".bmp", ".gif", ".tiff"} # Add or remove extensions as needed +# +# # Function to extract frame number from the filename +# def frame_number(filename): +# # First, try the pattern 'frame_x_7fps' +# new_pattern_match = re.search(r'frame_(\d+)_7fps', filename) +# if new_pattern_match: +# return int(new_pattern_match.group(1)) +# # If the new pattern is not found, use the original digit extraction method +# matches = re.findall(r'\d+', filename) +# if matches: +# if matches[-1] == '0000' and len(matches) > 1: +# return int(matches[-2]) # Return the second-to-last sequence if the last is '0000' +# return int(matches[-1]) # Otherwise, return the last sequence +# return float('inf') # Return 'inf' +# +# # Sorting files based on frame number +# sorted_files = sorted(os.listdir(folder), key=frame_number) +# +# # Load images in sorted order +# for filename in sorted_files: +# ext = os.path.splitext(filename)[1].lower() +# if ext in valid_extensions: +# img = Image.open(os.path.join(folder, filename)).convert('RGB') +# images.append(img) +# +# return images + +def load_images_from_folder(folder): + 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') + images.append(img) + + return images + + +# copy from https://github.com/crowsonkb/k-diffusion.git +def stratified_uniform(shape, group=0, groups=1, dtype=None, device=None): + """Draws stratified samples from a uniform distribution.""" + if groups <= 0: + raise ValueError(f"groups must be positive, got {groups}") + if group < 0 or group >= groups: + raise ValueError(f"group must be in [0, {groups})") + n = shape[-1] * groups + offsets = torch.arange(group, n, groups, dtype=dtype, device=device) + u = torch.rand(shape, dtype=dtype, device=device) + return (offsets + u) / n + + +def rand_cosine_interpolated(shape, image_d, noise_d_low, noise_d_high, sigma_data=1., min_value=1e-3, max_value=1e3, + device='cpu', dtype=torch.float32): + """Draws samples from an interpolated cosine timestep distribution (from simple diffusion).""" + + def logsnr_schedule_cosine(t, logsnr_min, logsnr_max): + t_min = math.atan(math.exp(-0.5 * logsnr_max)) + t_max = math.atan(math.exp(-0.5 * logsnr_min)) + return -2 * torch.log(torch.tan(t_min + t * (t_max - t_min))) + + def logsnr_schedule_cosine_shifted(t, image_d, noise_d, logsnr_min, logsnr_max): + shift = 2 * math.log(noise_d / image_d) + return logsnr_schedule_cosine(t, logsnr_min - shift, logsnr_max - shift) + shift + + def logsnr_schedule_cosine_interpolated(t, image_d, noise_d_low, noise_d_high, logsnr_min, logsnr_max): + logsnr_low = logsnr_schedule_cosine_shifted( + t, image_d, noise_d_low, logsnr_min, logsnr_max) + logsnr_high = logsnr_schedule_cosine_shifted( + t, image_d, noise_d_high, logsnr_min, logsnr_max) + return torch.lerp(logsnr_low, logsnr_high, t) + + logsnr_min = -2 * math.log(min_value / sigma_data) + logsnr_max = -2 * math.log(max_value / sigma_data) + u = stratified_uniform( + shape, group=0, groups=1, dtype=dtype, device=device + ) + logsnr = logsnr_schedule_cosine_interpolated( + u, image_d, noise_d_low, noise_d_high, logsnr_min, logsnr_max) + return torch.exp(-logsnr / 2) * sigma_data + + +def rand_log_normal(shape, loc=0., scale=1., device='cpu', dtype=torch.float32): + """Draws samples from an lognormal distribution.""" + u = torch.rand(shape, dtype=dtype, device=device) * (1 - 2e-7) + 1e-7 + return torch.distributions.Normal(loc, scale).icdf(u).exp() + + +min_value = 0.002 +max_value = 700 +image_d = 64 +noise_d_low = 32 +noise_d_high = 64 +sigma_data = 0.5 + + +def _resize_with_antialiasing(input, size, interpolation="bicubic", align_corners=True): + h, w = input.shape[-2:] + factors = (h / size[0], w / size[1]) + + # First, we have to determine sigma + # Taken from skimage: https://github.com/scikit-image/scikit-image/blob/v0.19.2/skimage/transform/_warps.py#L171 + sigmas = ( + max((factors[0] - 1.0) / 2.0, 0.001), + max((factors[1] - 1.0) / 2.0, 0.001), + ) + + # Now kernel size. Good results are for 3 sigma, but that is kind of slow. Pillow uses 1 sigma + # https://github.com/python-pillow/Pillow/blob/master/src/libImaging/Resample.c#L206 + # But they do it in the 2 passes, which gives better results. Let's try 2 sigmas for now + ks = int(max(2.0 * 2 * sigmas[0], 3)), int(max(2.0 * 2 * sigmas[1], 3)) + + # Make sure it is odd + if (ks[0] % 2) == 0: + ks = ks[0] + 1, ks[1] + + if (ks[1] % 2) == 0: + ks = ks[0], ks[1] + 1 + + input = _gaussian_blur2d(input, ks, sigmas) + + output = torch.nn.functional.interpolate( + input, size=size, mode=interpolation, align_corners=align_corners) + return output + + +def _compute_padding(kernel_size): + """Compute padding tuple.""" + # 4 or 6 ints: (padding_left, padding_right,padding_top,padding_bottom) + # https://pytorch.org/docs/stable/nn.html#torch.nn.functional.pad + if len(kernel_size) < 2: + raise AssertionError(kernel_size) + computed = [k - 1 for k in kernel_size] + + # for even kernels we need to do asymmetric padding :( + out_padding = 2 * len(kernel_size) * [0] + + for i in range(len(kernel_size)): + computed_tmp = computed[-(i + 1)] + + pad_front = computed_tmp // 2 + pad_rear = computed_tmp - pad_front + + out_padding[2 * i + 0] = pad_front + out_padding[2 * i + 1] = pad_rear + + return out_padding + + +def _filter2d(input, kernel): + # prepare kernel + b, c, h, w = input.shape + tmp_kernel = kernel[:, None, ...].to( + device=input.device, dtype=input.dtype) + + tmp_kernel = tmp_kernel.expand(-1, c, -1, -1) + + height, width = tmp_kernel.shape[-2:] + + padding_shape: list[int] = _compute_padding([height, width]) + input = torch.nn.functional.pad(input, padding_shape, mode="reflect") + + # kernel and input tensor reshape to align element-wise or batch-wise params + tmp_kernel = tmp_kernel.reshape(-1, 1, height, width) + input = input.view(-1, tmp_kernel.size(0), input.size(-2), input.size(-1)) + + # convolve the tensor with the kernel. + output = torch.nn.functional.conv2d( + input, tmp_kernel, groups=tmp_kernel.size(0), padding=0, stride=1) + + out = output.view(b, c, h, w) + return out + + +def _gaussian(window_size: int, sigma): + if isinstance(sigma, float): + sigma = torch.tensor([[sigma]]) + + batch_size = sigma.shape[0] + + x = (torch.arange(window_size, device=sigma.device, + dtype=sigma.dtype) - window_size // 2).expand(batch_size, -1) + + if window_size % 2 == 0: + x = x + 0.5 + + gauss = torch.exp(-x.pow(2.0) / (2 * sigma.pow(2.0))) + + return gauss / gauss.sum(-1, keepdim=True) + + +def _gaussian_blur2d(input, kernel_size, sigma): + if isinstance(sigma, tuple): + sigma = torch.tensor([sigma], dtype=input.dtype) + else: + sigma = sigma.to(dtype=input.dtype) + + ky, kx = int(kernel_size[0]), int(kernel_size[1]) + bs = sigma.shape[0] + kernel_x = _gaussian(kx, sigma[:, 1].view(bs, 1)) + kernel_y = _gaussian(ky, sigma[:, 0].view(bs, 1)) + out_x = _filter2d(input, kernel_x[..., None, :]) + out = _filter2d(out_x, kernel_y[..., None]) + + return out + + +def export_to_video(video_frames, output_video_path, fps): + fourcc = cv2.VideoWriter_fourcc(*"mp4v") + h, w, _ = video_frames[0].shape + video_writer = cv2.VideoWriter( + output_video_path, fourcc, fps=fps, frameSize=(w, h)) + for i in range(len(video_frames)): + img = cv2.cvtColor(video_frames[i], cv2.COLOR_RGB2BGR) + video_writer.write(img) + + +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 tensor_to_vae_latent(t, vae, scale=True): + t = t.to(vae.dtype) + if len(t.shape) == 5: + video_length = t.shape[1] + + t = rearrange(t, "b f c h w -> (b f) c h w") + latents = vae.encode(t).latent_dist.sample() + latents = rearrange(latents, "(b f) c h w -> b f c h w", f=video_length) + elif len(t.shape) == 4: + latents = vae.encode(t).latent_dist.sample() + if scale: + latents = latents * vae.config.scaling_factor + return latents + + +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, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + + parser.add_argument( + "--num_frames", + type=int, + default=14, + ) + parser.add_argument( + "--dataset_type", + type=str, + default='ubc', + ) + parser.add_argument( + "--num_validation_images", + type=int, + default=1, + help="Number of images that should be generated during validation with `validation_prompt`.", + ) + parser.add_argument( + "--validation_steps", + type=int, + default=500, + help=( + "Run fine-tuning validation every X epochs. The validation process consists of running the text/image prompt" + " multiple times: `args.num_validation_images`." + ), + ) + parser.add_argument( + "--output_dir", + type=str, + default="./outputs", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument( + "--seed", type=int, default=None, help="A seed for reproducible training." + ) + parser.add_argument( + "--per_gpu_batch_size", + type=int, + default=1, + help="Batch size (per device) for the training dataloader.", + ) + parser.add_argument("--num_train_epochs", type=int, default=100) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + 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( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", + type=int, + default=500, + help="Number of steps for the warmup in the lr scheduler.", + ) + parser.add_argument( + "--conditioning_dropout_prob", + type=float, + default=0.1, + help="Conditioning dropout probability. Drops out the conditionings (image and edit prompt) used in training InstructPix2Pix. See section 3.2.1 in the paper: https://arxiv.org/abs/2211.09800.", + ) + parser.add_argument( + "--use_8bit_adam", + action="store_true", + help="Whether or not to use 8-bit Adam from bitsandbytes.", + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices" + ), + ) + parser.add_argument( + "--use_ema", action="store_true", help="Whether to use EMA model." + ) + parser.add_argument( + "--non_ema_revision", + type=str, + default=None, + required=False, + help=( + "Revision of pretrained non-ema model identifier. Must be a branch, tag or git identifier of the local or" + " remote repository specified with --pretrained_model_name_or_path." + ), + ) + parser.add_argument( + "--num_workers", + type=int, + default=8, + help=( + "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process." + ), + ) + parser.add_argument( + "--adam_beta1", + type=float, + default=0.9, + help="The beta1 parameter for the Adam optimizer.", + ) + parser.add_argument( + "--adam_beta2", + type=float, + default=0.999, + help="The beta2 parameter for the Adam optimizer.", + ) + parser.add_argument( + "--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use." + ) + parser.add_argument( + "--adam_epsilon", + type=float, + default=1e-08, + help="Epsilon value for the Adam optimizer", + ) + parser.add_argument( + "--max_grad_norm", default=1.0, type=float, help="Max gradient norm." + ) + parser.add_argument( + "--push_to_hub", + action="store_true", + help="Whether or not to push the model to the Hub.", + ) + parser.add_argument( + "--hub_token", + type=str, + default=None, + help="The token to use to push to the Model Hub.", + ) + parser.add_argument( + "--hub_model_id", + type=str, + default=None, + help="The name of the repository to keep in sync with the local `output_dir`.", + ) + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the" + " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config." + ), + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument( + "--local_rank", + type=int, + default=-1, + help="For distributed training: local_rank", + ) + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints are only suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=1, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument( + "--enable_xformers_memory_efficient_attention", + action="store_true", + help="Whether or not to use xformers.", + ) + parser.add_argument( + "--log_trainable_parameters", + action="store_true", + help="Whether to write the trainable parameters.", + ) + parser.add_argument( + "--pretrain_unet", + type=str, + default=None, + help="use weight for unet block", + ) + parser.add_argument( + "--rank", + type=int, + default=128, + help=("The dimension of the LoRA update matrices."), + ) + parser.add_argument( + "--csv_path", + type=str, + default=None, + help=( + "path to the dataset csv" + ), + ) + parser.add_argument( + "--video_folder", + type=str, + default=None, + help=( + "path to the video folder" + ), + ) + parser.add_argument( + "--condition_folder", + type=str, + default=None, + help=( + "path to the depth folder" + ), + ) + parser.add_argument( + "--motion_folder", + type=str, + default=None, + help=( + "path to the depth folder" + ), + ) + parser.add_argument( + "--validation_prompt", + type=str, + default=None, + help=( + "A set of prompts evaluated every `--validation_steps` and logged to `--report_to`." + " Provide either a matching number of `--validation_image`s, a single `--validation_image`" + " to be used with all prompts, or a single prompt that will be used with all `--validation_image`s." + ), + ) + parser.add_argument( + "--validation_image_folder", + 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_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( + "--sample_n_frames", + type=int, + default=14, + help=( + "the sample_n_frames" + ), + ) + + parser.add_argument( + "--ref_augment", + action="store_true", + help=( + "use augment for the reference image" + ), + ) + parser.add_argument( + "--train_stage", + type=int, + default=2, + help=( + "the training stage" + ), + ) + + 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( + "--data_root_path", + type=str, + default=None, + help="Path to the data root path", + ) + parser.add_argument( + "--data_path", + type=str, + default=None, + help="Path to the data path", + ) + + parser.add_argument( + "--finetune_mode", + type=bool, + default=False, + help="Enable or disable the finetune mode (True/False).", + ) + parser.add_argument( + "--posenet_model_finetune_path", + type=str, + default=None, + help="Path to the pretrained posenet model", + ) + parser.add_argument( + "--face_encoder_finetune_path", + type=str, + default=None, + help="Path to the pretrained face encoder", + ) + parser.add_argument( + "--unet_model_finetune_path", + type=str, + default=None, + help="Path to the pretrained unet model", + ) + + parser.add_argument( + "--dataset_width", + type=int, + default=512, + help="video dataset width", + ) + parser.add_argument( + "--dataset_height", + type=int, + default=512, + help="video dataset height", + ) + + args = parser.parse_args() + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + if env_local_rank != -1 and env_local_rank != args.local_rank: + args.local_rank = env_local_rank + + # default to using the same revision for the non-ema model if not specified + if args.non_ema_revision is None: + args.non_ema_revision = args.revision + + return args + + +def download_image(url): + original_image = ( + lambda image_url_or_path: load_image(image_url_or_path) + if urlparse(image_url_or_path).scheme + else PIL.Image.open(image_url_or_path).convert("RGB") + )(url) + return original_image + + +# This is for training using deepspeed. +# Since now the DeepSpeed only supports trainging with only one model +# So we create a virtual wrapper to contail all the models + +class DeepSpeedWrapperModel(nn.Module): + def __init__(self, **kwargs): + super().__init__() + for name, value in kwargs.items(): + assert isinstance(value, nn.Module) + self.register_module(name, value) + + +def main(): + warnings.filterwarnings('ignore', category=DeprecationWarning) + warnings.filterwarnings('ignore', category=FutureWarning) + torch.multiprocessing.set_start_method('spawn') + + args = parse_args() + + if args.non_ema_revision is not None: + deprecate( + "non_ema_revision!=None", + "0.15.0", + message=( + "Downloading 'non_ema' weights from revision branches of the Hub is deprecated. Please make sure to" + " use `--variant=non_ema` instead." + ), + ) + logging_dir = os.path.join(args.output_dir, args.logging_dir) + accelerator_project_config = ProjectConfiguration( + project_dir=args.output_dir, logging_dir=logging_dir) + # ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True) + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + project_config=accelerator_project_config, + ) + + generator = torch.Generator( + device=accelerator.device).manual_seed(23123134) + + if args.report_to == "wandb": + if not is_wandb_available(): + raise ImportError( + "Make sure to install wandb if you want to use it for logging during training.") + import wandb + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed) + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + if args.push_to_hub: + repo_id = create_repo( + repo_id=args.hub_model_id or Path(args.output_dir).name, exist_ok=True, token=args.hub_token + ).repo_id + + # Load scheduler, tokenizer and models. + print(args.pretrained_model_name_or_path) + 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, variant="fp16") + unet = UNetSpatioTemporalConditionModel.from_pretrained( + args.pretrained_model_name_or_path if args.pretrain_unet is None else args.pretrain_unet, + subfolder="unet", + low_cpu_mem_usage=True, + variant="fp16" + ) + 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=1024, + num_tokens=4, ) + face_model = FaceModel() + + # init adapter modules + 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 AnimationIDAttnNormalizedProcessor: {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) + + # triggering the finetune mode + if args.finetune_mode is True and args.posenet_model_finetune_path is not None and args.face_encoder_finetune_path is not None and args.unet_model_finetune_path is not None: + print("Loading existing posenet weights, face_encoder weights and unet weights.") + if args.posenet_model_finetune_path.endswith(".pth"): + pose_net_state_dict = torch.load(args.posenet_model_finetune_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_finetune_path.endswith(".pth"): + face_encoder_state_dict = torch.load(args.face_encoder_finetune_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_finetune_path.endswith(".pth"): + unet_state_dict = torch.load(args.unet_model_finetune_path, map_location="cpu") + unet.load_state_dict(unet_state_dict, strict=True) + else: + print("unet weights loading fail") + print(1 / 0) + + vae_scale_factor = 2 ** (len(vae.config.block_out_channels) - 1) + image_processor = VaeImageProcessor(vae_scale_factor=vae_scale_factor) + + # Freeze vae and image_encoder + vae.requires_grad_(False) + image_encoder.requires_grad_(False) + unet.requires_grad_(False) + pose_net.requires_grad_(False) + face_encoder.requires_grad_(False) + + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + + image_encoder.to(accelerator.device, dtype=weight_dtype) + vae.to(accelerator.device, dtype=weight_dtype) + + if args.use_ema: + ema_unet = EMAModel(unet.parameters( + ), model_cls=UNetSpatioTemporalConditionModel, model_config=unet.config) + + if args.enable_xformers_memory_efficient_attention: + if is_xformers_available(): + import xformers + xformers_version = version.parse(xformers.__version__) + if xformers_version == version.parse("0.0.16"): + logger.warn( + "xFormers 0.0.16 cannot be used for training in some GPUs. If you observe problems during training, please update xFormers to at least 0.0.17. See https://huggingface.co/docs/diffusers/main/en/optimization/xformers for more details." + ) + unet.enable_xformers_memory_efficient_attention() + else: + raise ValueError( + "xformers is not available. Make sure it is installed correctly") + + if args.gradient_checkpointing: + unet.enable_gradient_checkpointing() + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * + args.per_gpu_batch_size * accelerator.num_processes + ) + + # Initialize the optimizer + if args.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError( + "Please install bitsandbytes to use 8-bit Adam. You can do so by running `pip install bitsandbytes`" + ) + + optimizer_cls = bnb.optim.AdamW8bit + else: + optimizer_cls = torch.optim.AdamW + + # if accelerator.distributed_type == DistributedType.DEEPSPEED: + # ds_wrapper = DeepSpeedWrapperModel( + # unet=unet, + # controlnext=controlnext + # ) + # unet = ds_wrapper.unet + # controlnext = ds_wrapper.controlnext + + pose_net.requires_grad_(True) + face_encoder.requires_grad_(True) + + parameters_list = [] + + for name, para in pose_net.named_parameters(): + para.requires_grad = True + parameters_list.append({"params": para, "lr": args.learning_rate}) + + for name, para in face_encoder.named_parameters(): + para.requires_grad = True + parameters_list.append({"params": para, "lr": args.learning_rate}) + + """ + For more details, please refer to: https://github.com/dvlab-research/ControlNeXt/issues/14#issuecomment-2290450333 + This is the selective parameters part. + As presented in our paper, we only select a small subset of parameters, which is fully adapted to the SD1.5 and SDXL backbones. By training fewer than 100 million parameters, we still achieve excellent performance. But this is is not suitable for the SD3 and SVD training. This is because, after SDXL, Stability faced significant legal risks due to the generation of highly realistic human images. After that, they stopped refining their models on human-related data, such as SVD and SD3, to avoid potential risks. + To achieve optimal performance, it's necessary to first continue training SVD and SD3 on human-related data to develop a robust backbone before fine-tuning. Of course, you can also combine the continual pretraining and finetuning. So you can find that we direct provide the full SVD parameters. + We have experimented with two approaches: 1.Directly training the model from scratch on human dancing data. 2. Continual training using a pre-trained human generation backbone, followed by fine-tuning a selective small subset of parameters. Interestingly, we observed no significant difference in performance between these two methods. + """ + + for name, para in unet.named_parameters(): + if "attentions" in name: + para.requires_grad = True + parameters_list.append({"params": para}) + else: + para.requires_grad = False + + optimizer = optimizer_cls( + parameters_list, + lr=args.learning_rate, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + ) + + # check para + if accelerator.is_main_process and args.log_trainable_parameters: + rec_txt1 = open('rec_para.txt', 'w') + rec_txt2 = open('rec_para_train.txt', 'w') + for name, para in unet.named_parameters(): + if para.requires_grad is False: + rec_txt1.write(f'{name}\n') + else: + rec_txt2.write(f'{name}\n') + rec_txt1.close() + rec_txt2.close() + # DataLoaders creation: + args.global_batch_size = args.per_gpu_batch_size * accelerator.num_processes + + root_path = args.data_root_path + txt_path = args.data_path + train_dataset = LargeScaleAnimationVideos( + root_path=root_path, + txt_path=txt_path, + width=args.dataset_width, + height=args.dataset_height, + n_sample_frames=args.sample_n_frames, + sample_frame_rate=4, + app=face_model.app, + handler_ante=face_model.handler_ante, + face_helper=face_model.face_helper + ) + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_size=args.per_gpu_batch_size, + num_workers=args.num_workers, + shuffle=True, + ) + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + ) + + unet, pose_net, face_encoder, optimizer, lr_scheduler, train_dataloader = accelerator.prepare( + unet, pose_net, face_encoder, optimizer, lr_scheduler, train_dataloader + ) + + if args.use_ema: + ema_unet.to(accelerator.device) + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil( + args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + accelerator.init_trackers("StableAnimator", config=vars(args)) + + # Train! + total_batch_size = args.per_gpu_batch_size * \ + accelerator.num_processes * args.gradient_accumulation_steps + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info( + f" Instantaneous batch size per device = {args.per_gpu_batch_size}") + logger.info( + f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info( + f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + def encode_image(pixel_values): + pixel_values = _resize_with_antialiasing(pixel_values, (224, 224)) + pixel_values = (pixel_values + 1.0) / 2.0 + + pixel_values = pixel_values.to(torch.float32) + # Normalize the image with for CLIP input + pixel_values = feature_extractor( + images=pixel_values, + do_normalize=True, + do_center_crop=False, + do_resize=False, + do_rescale=False, + return_tensors="pt", + ).pixel_values + + pixel_values = pixel_values.to( + device=accelerator.device, dtype=image_encoder.dtype) + image_embeddings = image_encoder(pixel_values).image_embeds + image_embeddings = image_embeddings.unsqueeze(1) + return image_embeddings + + def _get_add_time_ids( + fps, + motion_bucket_id, + noise_aug_strength, + dtype, + batch_size, + unet=None, + device=None + ): + add_time_ids = [fps, motion_bucket_id, noise_aug_strength] + + add_time_ids = torch.tensor([add_time_ids], dtype=dtype, device=device) + add_time_ids = add_time_ids.repeat(batch_size, 1) + return add_time_ids + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the most recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + else: + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + global_step = int(path.split("-")[1]) + + resume_global_step = global_step * args.gradient_accumulation_steps + first_epoch = global_step // num_update_steps_per_epoch + resume_step = resume_global_step % ( + num_update_steps_per_epoch * args.gradient_accumulation_steps) + + # Only show the progress bar once on each machine. + progress_bar = tqdm(range(global_step, args.max_train_steps), + disable=not accelerator.is_local_main_process) + progress_bar.set_description("Steps") + + for epoch in range(first_epoch, args.num_train_epochs): + pose_net.train() + face_encoder.train() + unet.train() + train_loss = 0.0 + + for step, batch in enumerate(train_dataloader): + # Skip steps until we reach the resumed step + if args.resume_from_checkpoint and epoch == first_epoch and step < resume_step: + if step % args.gradient_accumulation_steps == 0: + progress_bar.update(1) + continue + + with accelerator.accumulate(pose_net, face_encoder, unet): + with accelerator.autocast(): + pixel_values = batch["pixel_values"].to(weight_dtype).to( + accelerator.device, non_blocking=True + ) + conditional_pixel_values = batch["reference_image"].to(weight_dtype).to( + accelerator.device, non_blocking=True + ) + + latents = tensor_to_vae_latent(pixel_values, vae).to(dtype=weight_dtype) + + # Get the text embedding for conditioning. + encoder_hidden_states = encode_image(conditional_pixel_values).to(dtype=weight_dtype) + image_embed = encoder_hidden_states.clone() + + train_noise_aug = 0.02 + conditional_pixel_values = conditional_pixel_values + train_noise_aug * \ + randn_tensor(conditional_pixel_values.shape, generator=generator, + device=conditional_pixel_values.device, + dtype=conditional_pixel_values.dtype) + conditional_latents = tensor_to_vae_latent(conditional_pixel_values, vae, scale=False) + + # Sample noise that we'll add to the latents + noise = torch.randn_like(latents) + bsz = latents.shape[0] + # Sample a random timestep for each image + sigmas = rand_cosine_interpolated(shape=[bsz, ], image_d=image_d, noise_d_low=noise_d_low, + noise_d_high=noise_d_high, sigma_data=sigma_data, + min_value=min_value, max_value=max_value).to(latents.device, + dtype=weight_dtype) + + # sigmas = rand_log_normal(shape=[bsz,], loc=0.7, scale=1.6).to(latents) + # Add noise to the latents according to the noise magnitude at each timestep + # (this is the forward diffusion process) + sigmas_reshaped = sigmas.clone() + while len(sigmas_reshaped.shape) < len(latents.shape): + sigmas_reshaped = sigmas_reshaped.unsqueeze(-1) + + noisy_latents = latents + noise * sigmas_reshaped + + timesteps = torch.Tensor([0.25 * sigma.log() for sigma in sigmas]).to(latents.device, + dtype=weight_dtype) + + inp_noisy_latents = noisy_latents / ((sigmas_reshaped ** 2 + 1) ** 0.5) + + added_time_ids = _get_add_time_ids( + fps=6, + motion_bucket_id=127.0, + noise_aug_strength=train_noise_aug, # noise_aug_strength == 0.0 + dtype=encoder_hidden_states.dtype, + batch_size=bsz, + unet=unet, + device=latents.device + ) + + added_time_ids = added_time_ids.to(latents.device) + + # Conditioning dropout to support classifier-free guidance during inference. For more details + # check out the section 3.2.1 of the original paper https://arxiv.org/abs/2211.09800. + if args.conditioning_dropout_prob is not None: + random_p = torch.rand( + bsz, device=latents.device, generator=generator) + # Sample masks for the edit prompts. + prompt_mask = random_p < 2 * args.conditioning_dropout_prob + prompt_mask = prompt_mask.reshape(bsz, 1, 1) + # Final text conditioning. + null_conditioning = torch.zeros_like(encoder_hidden_states) + encoder_hidden_states = torch.where( + prompt_mask, null_conditioning, encoder_hidden_states) + + # Sample masks for the original images. + image_mask_dtype = conditional_latents.dtype + image_mask = 1 - ( + (random_p >= args.conditioning_dropout_prob).to( + image_mask_dtype) + * (random_p < 3 * args.conditioning_dropout_prob).to(image_mask_dtype) + ) + image_mask = image_mask.reshape(bsz, 1, 1, 1) + # Final image conditioning. + conditional_latents = image_mask * conditional_latents + + # Concatenate the `conditional_latents` with the `noisy_latents`. + conditional_latents = conditional_latents.unsqueeze( + 1).repeat(1, noisy_latents.shape[1], 1, 1, 1) + + pose_pixels = batch["pose_pixels"].to( + dtype=weight_dtype, device=accelerator.device, non_blocking=True + ) + faceid_embeds = batch["faceid_embeds"].to( + dtype=weight_dtype, device=accelerator.device, non_blocking=True + ) + pose_latents = pose_net(pose_pixels) + + # print("This is faceid_latents calculation") + # print(faceid_embeds.size()) # [1, 512] + # print(image_embed.size()) # [1, 1, 1024] + + faceid_latents = face_encoder(faceid_embeds, image_embed) + + inp_noisy_latents = torch.cat( + [inp_noisy_latents, conditional_latents], dim=2) + target = latents + + # print(f"the size of encoder_hidden_states: {encoder_hidden_states.size()}") # [1, 1, 1024] + # print(f"the size of face latents: {faceid_latents.size()}") # [1, 4, 1024] + encoder_hidden_states = torch.cat([encoder_hidden_states, faceid_latents], dim=1) + + encoder_hidden_states = encoder_hidden_states.to(latents.dtype) + inp_noisy_latents = inp_noisy_latents.to(latents.dtype) + pose_latents = pose_latents.to(latents.dtype) + + # Predict the noise residual + model_pred = unet( + inp_noisy_latents, timesteps, encoder_hidden_states, + added_time_ids=added_time_ids, + pose_latents=pose_latents, + ).sample + + sigmas = sigmas_reshaped + # Denoise the latents + c_out = -sigmas / ((sigmas ** 2 + 1) ** 0.5) + c_skip = 1 / (sigmas ** 2 + 1) + denoised_latents = model_pred * c_out + c_skip * noisy_latents + weighing = (1 + sigmas ** 2) * (sigmas ** -2.0) + + tgt_face_masks = batch["tgt_face_masks"].to( + dtype=weight_dtype, device=accelerator.device, non_blocking=True + ) + tgt_face_masks = rearrange(tgt_face_masks, "b f c h w -> (b f) c h w") + tgt_face_masks = F.interpolate(tgt_face_masks, size=(target.size()[-2], target.size()[-1]), + mode='nearest') + tgt_face_masks = rearrange(tgt_face_masks, "(b f) c h w -> b f c h w", f=args.sample_n_frames) + + # MSE loss + loss = torch.mean( + (weighing.float() * (denoised_latents.float() - + target.float()) ** 2 * (1 + tgt_face_masks)).reshape(target.shape[0], -1), + dim=1, + ) + loss = loss.mean() + + # Gather the losses across all processes for logging (if we use distributed training). + avg_loss = accelerator.gather( + loss.repeat(args.per_gpu_batch_size)).mean() + train_loss += avg_loss.item() / args.gradient_accumulation_steps + + # Backpropagate + accelerator.backward(loss) + # if accelerator.sync_gradients: + # accelerator.clip_grad_norm_(unet.parameters(), args.max_grad_norm) + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad() + + with torch.cuda.device(latents.device): + torch.cuda.empty_cache() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + if args.use_ema: + ema_unet.step(unet.parameters()) + progress_bar.update(1) + global_step += 1 + accelerator.log({"train_loss": train_loss}, step=global_step) + train_loss = 0.0 + + # save checkpoints! + # if global_step % args.checkpointing_steps == 0 and (accelerator.is_main_process or accelerator.distributed_type == DistributedType.DEEPSPEED): + if global_step % args.checkpointing_steps == 0 and accelerator.is_main_process: + + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None and accelerator.is_main_process: + checkpoints = os.listdir(args.output_dir) + checkpoints = [ + d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted( + checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len( + checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info( + f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join( + args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + save_path = os.path.join( + args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + unwrap_unet = accelerator.unwrap_model(unet) + unwrap_pose_net = accelerator.unwrap_model(pose_net) + unwrap_face_encoder = accelerator.unwrap_model(face_encoder) + unwrap_unet_state_dict = unwrap_unet.state_dict() + torch.save(unwrap_unet_state_dict, + os.path.join(args.output_dir, f"checkpoint-{global_step}", f"unet-{global_step}.pth")) + unwrap_pose_net_state_dict = unwrap_pose_net.state_dict() + torch.save(unwrap_pose_net_state_dict, os.path.join(args.output_dir, f"checkpoint-{global_step}", + f"pose_net-{global_step}.pth")) + unwrap_face_encoder_state_dict = unwrap_face_encoder.state_dict() + torch.save(unwrap_face_encoder_state_dict, + os.path.join(args.output_dir, f"checkpoint-{global_step}", + f"face_encoder-{global_step}.pth")) + logger.info(f"Saved state to {save_path}") + + if accelerator.is_main_process: + # sample images! + if global_step % args.validation_steps == 0: + logger.info( + f"Running validation... \n Generating {args.num_validation_images} videos." + ) + # create pipeline + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_unet.store(unet.parameters()) + ema_unet.copy_to(unet.parameters()) + + log_validation( + vae=vae, + image_encoder=image_encoder, + unet=unet, + pose_net=pose_net, + face_encoder=face_encoder, + app=face_model.app, + face_helper=face_model.face_helper, + handler_ante=face_model.handler_ante, + scheduler=noise_scheduler, + accelerator=accelerator, + feature_extractor=feature_extractor, + width=512, + height=512, + torch_dtype=weight_dtype, + validation_image_folder=args.validation_image_folder, + validation_image=args.validation_image, + validation_control_folder=args.validation_control_folder, + output_dir=args.output_dir, + generator=generator, + global_step=global_step, + num_validation_cases=1, + ) + + if args.use_ema: + # Switch back to the original UNet parameters. + ema_unet.restore(unet.parameters()) + + with torch.cuda.device(latents.device): + torch.cuda.empty_cache() + + logs = {"step_loss": loss.detach().item( + ), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if global_step >= args.max_train_steps: + break + + # save checkpoints! + # if accelerator.is_main_process or accelerator.distributed_type == DistributedType.DEEPSPEED: + if accelerator.is_main_process: + save_path = os.path.join( + args.output_dir, f"checkpoint-last") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + +def log_validation( + vae, + image_encoder, + unet, + pose_net, + face_encoder, + app, + face_helper, + handler_ante, + scheduler, + accelerator, + feature_extractor, + width, + height, + torch_dtype, + validation_image_folder, + validation_image, + validation_control_folder, + output_dir, + generator, + global_step, + num_validation_cases=1, +): + logger.info("Running validation... ") + validation_unet = accelerator.unwrap_model(unet) + validation_image_encoder = accelerator.unwrap_model(image_encoder) + validation_vae = accelerator.unwrap_model(vae) + validation_pose_net = accelerator.unwrap_model(pose_net) + validation_face_encoder = accelerator.unwrap_model(face_encoder) + + pipeline = ValidationAnimationPipeline( + vae=validation_vae, + image_encoder=validation_image_encoder, + unet=validation_unet, + scheduler=scheduler, + feature_extractor=feature_extractor, + pose_net=validation_pose_net, + face_encoder=validation_face_encoder, + ) + pipeline = pipeline.to(accelerator.device) + validation_images = load_images_from_folder(validation_image_folder) + validation_image_path = validation_image + if validation_image is None: + validation_image = validation_images[0] + else: + validation_image = Image.open(validation_image).convert('RGB') + validation_control_images = load_images_from_folder(validation_control_folder) + + val_save_dir = os.path.join(output_dir, "validation_images") + if not os.path.exists(val_save_dir): + os.makedirs(val_save_dir) + + with accelerator.autocast(): + for val_img_idx in range(num_validation_cases): + # num_frames = args.num_frames + num_frames = len(validation_control_images) + + 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 = app.get(validation_face) + 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_helper.read_image(validation_face) + face_helper.get_face_landmarks_5(only_center_face=True) + face_helper.align_warp_face() + + if len(face_helper.cropped_faces) == 0: + validation_image_id_ante_embedding = np.zeros((512,)) + else: + validation_image_align_face = face_helper.cropped_faces[0] + print('fail to detect face using insightface, extract embedding on align face') + validation_image_id_ante_embedding = handler_ante.get_feat(validation_image_align_face) + + video_frames = pipeline( + image=validation_image, + image_pose=validation_control_images, + height=height, + width=width, + num_frames=num_frames, + tile_size=num_frames, + tile_overlap=4, + decode_chunk_size=4, + motion_bucket_id=127., + fps=7, + min_guidance_scale=3, + max_guidance_scale=3, + noise_aug_strength=0.02, + num_inference_steps=25, + generator=generator, + output_type="pil", + validation_image_id_ante_embedding=validation_image_id_ante_embedding, + ).frames[0] + # save_combined_frames(video_frames, validation_images, validation_control_images, val_save_dir) + + out_file = os.path.join( + val_save_dir, + f"step_{global_step}_val_img_{val_img_idx}.mp4", + ) + # print(video_frames.size()) # [16, 3, 512, 512] + for i in range(num_frames): + img = video_frames[i] + video_frames[i] = np.array(img) + export_to_gif(video_frames, out_file, 8) + + del pipeline + torch.cuda.empty_cache() + + +if __name__ == "__main__": + main()