init
This commit is contained in:
+160
@@ -0,0 +1,160 @@
|
|||||||
|
# Byte-compiled / optimized / DLL files
|
||||||
|
__pycache__/
|
||||||
|
*.py[cod]
|
||||||
|
*$py.class
|
||||||
|
|
||||||
|
# C extensions
|
||||||
|
*.so
|
||||||
|
|
||||||
|
# Distribution / packaging
|
||||||
|
.Python
|
||||||
|
build/
|
||||||
|
develop-eggs/
|
||||||
|
dist/
|
||||||
|
downloads/
|
||||||
|
eggs/
|
||||||
|
.eggs/
|
||||||
|
lib/
|
||||||
|
lib64/
|
||||||
|
parts/
|
||||||
|
sdist/
|
||||||
|
var/
|
||||||
|
wheels/
|
||||||
|
share/python-wheels/
|
||||||
|
*.egg-info/
|
||||||
|
.installed.cfg
|
||||||
|
*.egg
|
||||||
|
MANIFEST
|
||||||
|
|
||||||
|
# PyInstaller
|
||||||
|
# Usually these files are written by a python script from a template
|
||||||
|
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
||||||
|
*.manifest
|
||||||
|
*.spec
|
||||||
|
|
||||||
|
# Installer logs
|
||||||
|
pip-log.txt
|
||||||
|
pip-delete-this-directory.txt
|
||||||
|
|
||||||
|
# Unit test / coverage reports
|
||||||
|
htmlcov/
|
||||||
|
.tox/
|
||||||
|
.nox/
|
||||||
|
.coverage
|
||||||
|
.coverage.*
|
||||||
|
.cache
|
||||||
|
nosetests.xml
|
||||||
|
coverage.xml
|
||||||
|
*.cover
|
||||||
|
*.py,cover
|
||||||
|
.hypothesis/
|
||||||
|
.pytest_cache/
|
||||||
|
cover/
|
||||||
|
|
||||||
|
# Translations
|
||||||
|
*.mo
|
||||||
|
*.pot
|
||||||
|
|
||||||
|
# Django stuff:
|
||||||
|
*.log
|
||||||
|
local_settings.py
|
||||||
|
db.sqlite3
|
||||||
|
db.sqlite3-journal
|
||||||
|
|
||||||
|
# Flask stuff:
|
||||||
|
instance/
|
||||||
|
.webassets-cache
|
||||||
|
|
||||||
|
# Scrapy stuff:
|
||||||
|
.scrapy
|
||||||
|
|
||||||
|
# Sphinx documentation
|
||||||
|
docs/_build/
|
||||||
|
|
||||||
|
# PyBuilder
|
||||||
|
.pybuilder/
|
||||||
|
target/
|
||||||
|
|
||||||
|
# Jupyter Notebook
|
||||||
|
.ipynb_checkpoints
|
||||||
|
|
||||||
|
# IPython
|
||||||
|
profile_default/
|
||||||
|
ipython_config.py
|
||||||
|
|
||||||
|
# pyenv
|
||||||
|
# For a library or package, you might want to ignore these files since the code is
|
||||||
|
# intended to run in multiple environments; otherwise, check them in:
|
||||||
|
# .python-version
|
||||||
|
|
||||||
|
# pipenv
|
||||||
|
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
|
||||||
|
# However, in case of collaboration, if having platform-specific dependencies or dependencies
|
||||||
|
# having no cross-platform support, pipenv may install dependencies that don't work, or not
|
||||||
|
# install all needed dependencies.
|
||||||
|
#Pipfile.lock
|
||||||
|
|
||||||
|
# poetry
|
||||||
|
# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
|
||||||
|
# This is especially recommended for binary packages to ensure reproducibility, and is more
|
||||||
|
# commonly ignored for libraries.
|
||||||
|
# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
|
||||||
|
#poetry.lock
|
||||||
|
|
||||||
|
# PEP 582; used by e.g. github.com/David-OConnor/pyflow
|
||||||
|
__pypackages__/
|
||||||
|
|
||||||
|
# Celery stuff
|
||||||
|
celerybeat-schedule
|
||||||
|
celerybeat.pid
|
||||||
|
|
||||||
|
# SageMath parsed files
|
||||||
|
*.sage.py
|
||||||
|
|
||||||
|
# Environments
|
||||||
|
.env
|
||||||
|
.venv
|
||||||
|
env/
|
||||||
|
venv/
|
||||||
|
ENV/
|
||||||
|
env.bak/
|
||||||
|
venv.bak/
|
||||||
|
|
||||||
|
# Spyder project settings
|
||||||
|
.spyderproject
|
||||||
|
.spyproject
|
||||||
|
|
||||||
|
# Rope project settings
|
||||||
|
.ropeproject
|
||||||
|
|
||||||
|
# mkdocs documentation
|
||||||
|
#/site
|
||||||
|
|
||||||
|
# mypy
|
||||||
|
.mypy_cache/
|
||||||
|
.dmypy.json
|
||||||
|
dmypy.json
|
||||||
|
|
||||||
|
# Pyre type checker
|
||||||
|
.pyre/
|
||||||
|
|
||||||
|
# pytype static type analyzer
|
||||||
|
.pytype/
|
||||||
|
|
||||||
|
# Cython debug symbols
|
||||||
|
cython_debug/
|
||||||
|
|
||||||
|
# PyCharm
|
||||||
|
# JetBrains specific template is maintained in a separate JetBrains.gitignore that can
|
||||||
|
# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore
|
||||||
|
# and can be added to the global gitignore or merged into this file. For a more nuclear
|
||||||
|
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
|
||||||
|
.idea/
|
||||||
|
|
||||||
|
# custom ignores
|
||||||
|
.DS_Store
|
||||||
|
_.*
|
||||||
|
|
||||||
|
# models and outputs
|
||||||
|
models/dwpose
|
||||||
|
outputs/
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
# ComfyUI nodes for ControlNext-SVD v2
|
||||||
|
|
||||||
|
These nodes include my wrapper for the original diffusers pipeline, as well as work in progress native ComfyUI implementation.
|
||||||
|
|
||||||
|
Original repo:
|
||||||
|
|
||||||
|
https://github.com/dvlab-research/ControlNeXt/tree/main/ControlNeXt-SVD-v2
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||||
|
|
||||||
|
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
{
|
||||||
|
"_class_name": "UNetSpatioTemporalConditionModel",
|
||||||
|
"_diffusers_version": "0.24.0.dev0",
|
||||||
|
"_name_or_path": "/home/suraj_huggingface_co/.cache/huggingface/hub/models--diffusers--svd-xt/snapshots/9703ded20c957c340781ee710b75660826deb487/unet",
|
||||||
|
"addition_time_embed_dim": 256,
|
||||||
|
"block_out_channels": [
|
||||||
|
320,
|
||||||
|
640,
|
||||||
|
1280,
|
||||||
|
1280
|
||||||
|
],
|
||||||
|
"cross_attention_dim": 1024,
|
||||||
|
"down_block_types": [
|
||||||
|
"CrossAttnDownBlockSpatioTemporal",
|
||||||
|
"CrossAttnDownBlockSpatioTemporal",
|
||||||
|
"CrossAttnDownBlockSpatioTemporal",
|
||||||
|
"DownBlockSpatioTemporal"
|
||||||
|
],
|
||||||
|
"in_channels": 8,
|
||||||
|
"layers_per_block": 2,
|
||||||
|
"num_attention_heads": [
|
||||||
|
5,
|
||||||
|
10,
|
||||||
|
20,
|
||||||
|
20
|
||||||
|
],
|
||||||
|
"num_frames": 25,
|
||||||
|
"out_channels": 4,
|
||||||
|
"projection_class_embeddings_input_dim": 768,
|
||||||
|
"sample_size": 96,
|
||||||
|
"transformer_layers_per_block": 1,
|
||||||
|
"up_block_types": [
|
||||||
|
"UpBlockSpatioTemporal",
|
||||||
|
"CrossAttnUpBlockSpatioTemporal",
|
||||||
|
"CrossAttnUpBlockSpatioTemporal",
|
||||||
|
"CrossAttnUpBlockSpatioTemporal"
|
||||||
|
]
|
||||||
|
}
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
*.pyc
|
||||||
@@ -0,0 +1,64 @@
|
|||||||
|
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 DWposeDetector:
|
||||||
|
"""
|
||||||
|
A pose detect method for image-like data.
|
||||||
|
|
||||||
|
Parameters:
|
||||||
|
model_det: (str) serialized ONNX format model path,
|
||||||
|
such as https://huggingface.co/yzd-v/DWPose/blob/main/yolox_l.onnx
|
||||||
|
model_pose: (str) serialized ONNX format model path,
|
||||||
|
such as https://huggingface.co/yzd-v/DWPose/blob/main/dw-ll_ucoco_384.onnx
|
||||||
|
device: (str) 'cpu' or 'cuda:{device_id}'
|
||||||
|
"""
|
||||||
|
def __init__(self, model_det, model_pose, device='cpu'):
|
||||||
|
self.pose_estimation = Wholebody(model_det=model_det, model_pose=model_pose)
|
||||||
|
|
||||||
|
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 = DWposeDetector(
|
||||||
|
# model_det="models/DWPose/yolox_l.onnx",
|
||||||
|
# model_pose="models/DWPose/dw-ll_ucoco_384.onnx",
|
||||||
|
# device=device)
|
||||||
@@ -0,0 +1,125 @@
|
|||||||
|
import cv2
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
|
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(model, oriImg, detect_classes=[0]):
|
||||||
|
input_shape = (640,640)
|
||||||
|
img, ratio = preprocess(oriImg, input_shape)
|
||||||
|
|
||||||
|
device, dtype = next(model.parameters()).device, next(model.parameters()).dtype
|
||||||
|
input = img[None, :, :, :]
|
||||||
|
input = torch.from_numpy(input).to(device, dtype)
|
||||||
|
|
||||||
|
output = model(input).float().cpu().detach().numpy()
|
||||||
|
predictions = demo_postprocess(output[0], input_shape)
|
||||||
|
|
||||||
|
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 None:
|
||||||
|
return None
|
||||||
|
final_boxes, final_scores, final_cls_inds = dets[:, :4], dets[:, 4], dets[:, 5]
|
||||||
|
isscore = final_scores>0.3
|
||||||
|
iscat = np.isin(final_cls_inds, detect_classes)
|
||||||
|
isbbox = [ i and j for (i, j) in zip(isscore, iscat)]
|
||||||
|
final_boxes = final_boxes[isbbox]
|
||||||
|
return final_boxes
|
||||||
@@ -0,0 +1,363 @@
|
|||||||
|
from typing import List, Tuple
|
||||||
|
|
||||||
|
import cv2
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
|
def preprocess(
|
||||||
|
img: np.ndarray, out_bbox, input_size: Tuple[int, int] = (192, 256)
|
||||||
|
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
|
||||||
|
"""Do preprocessing for DWPose 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 out_bbox is None or 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(model, img, bs=5):
|
||||||
|
"""Inference DWPose model implemented in TorchScript.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model : TorchScript Model.
|
||||||
|
img : Input image in shape.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
outputs : Output of DWPose model.
|
||||||
|
"""
|
||||||
|
all_out = []
|
||||||
|
# build input
|
||||||
|
orig_img_count = len(img)
|
||||||
|
#Pad zeros to fit batch size
|
||||||
|
for _ in range(bs - (orig_img_count % bs)):
|
||||||
|
img.append(np.zeros_like(img[0]))
|
||||||
|
input = np.stack(img, axis=0).transpose(0, 3, 1, 2)
|
||||||
|
device, dtype = next(model.parameters()).device, next(model.parameters()).dtype
|
||||||
|
input = torch.from_numpy(input).to(device, dtype)
|
||||||
|
|
||||||
|
out1, out2 = [], []
|
||||||
|
for i in range(input.shape[0] // bs):
|
||||||
|
curr_batch_output = model(input[i*bs:(i+1)*bs])
|
||||||
|
out1.append(curr_batch_output[0].float())
|
||||||
|
out2.append(curr_batch_output[1].float())
|
||||||
|
out1, out2 = torch.cat(out1, dim=0)[:orig_img_count], torch.cat(out2, dim=0)[:orig_img_count]
|
||||||
|
out1, out2 = out1.float().cpu().detach().numpy(), out2.float().cpu().detach().numpy()
|
||||||
|
all_outputs = out1, out2
|
||||||
|
|
||||||
|
for batch_idx in range(len(all_outputs[0])):
|
||||||
|
outputs = [all_outputs[i][batch_idx:batch_idx+1,...] for i in range(len(all_outputs))]
|
||||||
|
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 DWPose 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(model, out_bbox, oriImg, model_input_size=(288, 384)):
|
||||||
|
resized_img, center, scale = preprocess(oriImg, out_bbox, model_input_size)
|
||||||
|
#outputs = inference(session, resized_img, dtype)
|
||||||
|
outputs = inference(model, resized_img)
|
||||||
|
|
||||||
|
keypoints, scores = postprocess(outputs, model_input_size, center, scale)
|
||||||
|
|
||||||
|
return keypoints, scores
|
||||||
@@ -0,0 +1,145 @@
|
|||||||
|
import cv2
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
|
||||||
|
def nms(boxes, scores, nms_thr):
|
||||||
|
"""Single class NMS implemented in Numpy.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
boxes (np.ndarray): shape=(N,4); N is number of boxes
|
||||||
|
scores (np.ndarray): the score of bboxes
|
||||||
|
nms_thr (float): the threshold in NMS
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List[int]: output bbox ids
|
||||||
|
"""
|
||||||
|
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.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
boxes (np.ndarray): shape=(N,4); N is number of boxes
|
||||||
|
scores (np.ndarray): the score of bboxes
|
||||||
|
nms_thr (float): the threshold in NMS
|
||||||
|
score_thr (float): the threshold of cls score
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
np.ndarray: outputs bboxes coordinate
|
||||||
|
"""
|
||||||
|
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):
|
||||||
|
"""run human detect
|
||||||
|
"""
|
||||||
|
input_shape = (640,640)
|
||||||
|
img, ratio = preprocess(oriImg, input_shape)
|
||||||
|
|
||||||
|
ort_inputs = {session.get_inputs()[0].name: img[None, :, :, :]}
|
||||||
|
output = session.run(None, ort_inputs)
|
||||||
|
predictions = demo_postprocess(output[0], input_shape)[0]
|
||||||
|
|
||||||
|
boxes = predictions[:, :4]
|
||||||
|
scores = predictions[:, 4:5] * predictions[:, 5:]
|
||||||
|
|
||||||
|
boxes_xyxy = np.ones_like(boxes)
|
||||||
|
boxes_xyxy[:, 0] = boxes[:, 0] - boxes[:, 2]/2.
|
||||||
|
boxes_xyxy[:, 1] = boxes[:, 1] - boxes[:, 3]/2.
|
||||||
|
boxes_xyxy[:, 2] = boxes[:, 0] + boxes[:, 2]/2.
|
||||||
|
boxes_xyxy[:, 3] = boxes[:, 1] + boxes[:, 3]/2.
|
||||||
|
boxes_xyxy /= ratio
|
||||||
|
dets = multiclass_nms(boxes_xyxy, scores, nms_thr=0.45, score_thr=0.1)
|
||||||
|
if dets is not None:
|
||||||
|
final_boxes, final_scores, final_cls_inds = dets[:, :4], dets[:, 4], dets[:, 5]
|
||||||
|
isscore = final_scores>0.3
|
||||||
|
iscat = final_cls_inds == 0
|
||||||
|
isbbox = [ i and j for (i, j) in zip(isscore, iscat)]
|
||||||
|
final_boxes = final_boxes[isbbox]
|
||||||
|
else:
|
||||||
|
final_boxes = np.array([])
|
||||||
|
|
||||||
|
return final_boxes
|
||||||
@@ -0,0 +1,375 @@
|
|||||||
|
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):
|
||||||
|
"""run pose detect
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session (ort.InferenceSession): ONNXRuntime session.
|
||||||
|
out_bbox (np.ndarray): bbox list
|
||||||
|
oriImg (np.ndarray): Input image in shape.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple:
|
||||||
|
- keypoints (np.ndarray): Rescaled keypoints.
|
||||||
|
- scores (np.ndarray): Model predict scores.
|
||||||
|
"""
|
||||||
|
h, w = session.get_inputs()[0].shape[2:]
|
||||||
|
model_input_size = (w, h)
|
||||||
|
# preprocess for rtm-pose model inference.
|
||||||
|
resized_img, center, scale = preprocess(oriImg, out_bbox, model_input_size)
|
||||||
|
# run pose estimation for processed img
|
||||||
|
outputs = inference(session, resized_img)
|
||||||
|
# postprocess for rtm-pose model output.
|
||||||
|
keypoints, scores = postprocess(outputs, model_input_size, center, scale)
|
||||||
|
|
||||||
|
return keypoints, scores
|
||||||
@@ -0,0 +1,70 @@
|
|||||||
|
import decord
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from .util import draw_pose
|
||||||
|
from .dwpose_detector import dwpose_detector as dwprocessor
|
||||||
|
|
||||||
|
|
||||||
|
def get_video_pose(
|
||||||
|
video_path: str,
|
||||||
|
ref_image: np.ndarray,
|
||||||
|
sample_stride: int=1):
|
||||||
|
"""preprocess ref image pose and video pose
|
||||||
|
|
||||||
|
Args:
|
||||||
|
video_path (str): video pose path
|
||||||
|
ref_image (np.ndarray): reference image
|
||||||
|
sample_stride (int, optional): Defaults to 1.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
np.ndarray: sequence of video pose
|
||||||
|
"""
|
||||||
|
# select ref-keypoint from reference pose for pose rescale
|
||||||
|
ref_pose = dwprocessor(ref_image)
|
||||||
|
ref_keypoint_id = [0, 1, 2, 5, 8, 11, 14, 15, 16, 17]
|
||||||
|
ref_keypoint_id = [i for i in ref_keypoint_id \
|
||||||
|
if ref_pose['bodies']['score'].shape[0] > 0 and ref_pose['bodies']['score'][0][i] > 0.3]
|
||||||
|
ref_body = ref_pose['bodies']['candidate'][ref_keypoint_id]
|
||||||
|
|
||||||
|
height, width, _ = ref_image.shape
|
||||||
|
|
||||||
|
# read input video
|
||||||
|
vr = decord.VideoReader(video_path, ctx=decord.cpu(0))
|
||||||
|
sample_stride *= max(1, int(vr.get_avg_fps() / 24))
|
||||||
|
|
||||||
|
detected_poses = [dwprocessor(frm) for frm in vr.get_batch(list(range(0, len(vr), sample_stride))).asnumpy()]
|
||||||
|
|
||||||
|
detected_bodies = np.stack(
|
||||||
|
[p['bodies']['candidate'] for p in detected_poses if p['bodies']['candidate'].shape[0] == 18])[:,
|
||||||
|
ref_keypoint_id]
|
||||||
|
# compute linear-rescale params
|
||||||
|
ay, by = np.polyfit(detected_bodies[:, :, 1].flatten(), np.tile(ref_body[:, 1], len(detected_bodies)), 1)
|
||||||
|
fh, fw, _ = vr[0].shape
|
||||||
|
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):
|
||||||
|
"""process image pose
|
||||||
|
|
||||||
|
Args:
|
||||||
|
ref_image (np.ndarray): reference image pixel value
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
np.ndarray: pose visual image in RGB-mode
|
||||||
|
"""
|
||||||
|
height, width, _ = ref_image.shape
|
||||||
|
ref_pose = dwprocessor(ref_image)
|
||||||
|
pose_img = draw_pose(ref_pose, height, width)
|
||||||
|
return np.array(pose_img)
|
||||||
+136
@@ -0,0 +1,136 @@
|
|||||||
|
import math
|
||||||
|
import numpy as np
|
||||||
|
import matplotlib
|
||||||
|
import cv2
|
||||||
|
|
||||||
|
|
||||||
|
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, include_body, include_hand, include_face, 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 #####################################################
|
||||||
|
if include_body:
|
||||||
|
canvas = draw_bodypose(canvas, candidate, subset, score=bodies['score'])
|
||||||
|
|
||||||
|
########################################### draw hand pose #####################################################
|
||||||
|
if include_hand:
|
||||||
|
canvas = draw_handpose(canvas, hands, pose['hands_score'])
|
||||||
|
|
||||||
|
########################################### draw face pose #####################################################
|
||||||
|
if include_face:
|
||||||
|
canvas = draw_facepose(canvas, faces, pose['faces_score'])
|
||||||
|
|
||||||
|
return cv2.cvtColor(cv2.resize(canvas, (W, H)), cv2.COLOR_BGR2RGB).transpose(2, 0, 1)
|
||||||
@@ -0,0 +1,65 @@
|
|||||||
|
import numpy as np
|
||||||
|
|
||||||
|
import comfy.model_management as mm
|
||||||
|
|
||||||
|
#import onnxruntime as ort
|
||||||
|
# from .onnxdet import inference_detector
|
||||||
|
# from .onnxpose import inference_pose
|
||||||
|
|
||||||
|
from .jit_det import inference_detector as inference_jit_yolox
|
||||||
|
from .jit_pose import inference_pose as inference_jit_pose
|
||||||
|
|
||||||
|
class Wholebody:
|
||||||
|
"""detect human pose by dwpose
|
||||||
|
"""
|
||||||
|
def __init__(self, model_det, model_pose):
|
||||||
|
#providers = ['CPUExecutionProvider'] if device == 'cpu' else ['CUDAExecutionProvider']
|
||||||
|
#provider_options = None if device == 'cpu' else [{'device_id': 0}]
|
||||||
|
|
||||||
|
# self.session_det = ort.InferenceSession(
|
||||||
|
# path_or_bytes=model_det, providers=providers, provider_options=provider_options
|
||||||
|
# )
|
||||||
|
# self.session_pose = ort.InferenceSession(
|
||||||
|
# path_or_bytes=model_pose, providers=providers, provider_options=provider_options
|
||||||
|
# )
|
||||||
|
|
||||||
|
self.det = model_det
|
||||||
|
self.pose = model_pose
|
||||||
|
|
||||||
|
def __call__(self, oriImg):
|
||||||
|
"""call to process dwpose-detect
|
||||||
|
|
||||||
|
Args:
|
||||||
|
oriImg (np.ndarray): detected image
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
det_result = inference_jit_yolox(self.det, oriImg, detect_classes=[0])
|
||||||
|
keypoints, scores = inference_jit_pose(self.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
|
||||||
|
|
||||||
|
|
||||||
Binary file not shown.
@@ -0,0 +1,176 @@
|
|||||||
|
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||||
|
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
|
||||||
|
from diffusers.models.modeling_utils import ModelMixin
|
||||||
|
from diffusers.models.resnet import Downsample2D, ResnetBlock2D
|
||||||
|
|
||||||
|
|
||||||
|
class ControlNeXtSDVModel(ModelMixin, ConfigMixin):
|
||||||
|
_supports_gradient_checkpointing = True
|
||||||
|
|
||||||
|
@register_to_config
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
time_embed_dim = 256,
|
||||||
|
in_channels = [128, 128],
|
||||||
|
out_channels = [128, 256],
|
||||||
|
groups = [4, 8]
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.time_proj = Timesteps(128, True, downscale_freq_shift=0)
|
||||||
|
self.time_embedding = TimestepEmbedding(128, time_embed_dim)
|
||||||
|
self.embedding = nn.Sequential(
|
||||||
|
nn.Conv2d(3, 64, kernel_size=3, stride=2, padding=1),
|
||||||
|
nn.GroupNorm(2, 64),
|
||||||
|
nn.ReLU(),
|
||||||
|
nn.Conv2d(64, 64, kernel_size=3),
|
||||||
|
nn.GroupNorm(2, 64),
|
||||||
|
nn.ReLU(),
|
||||||
|
nn.Conv2d(64, 128, kernel_size=3),
|
||||||
|
nn.GroupNorm(2, 128),
|
||||||
|
nn.ReLU(),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.down_res = nn.ModuleList()
|
||||||
|
self.down_sample = nn.ModuleList()
|
||||||
|
for i in range(len(in_channels)):
|
||||||
|
self.down_res.append(
|
||||||
|
ResnetBlock2D(
|
||||||
|
in_channels=in_channels[i],
|
||||||
|
out_channels=out_channels[i],
|
||||||
|
temb_channels=time_embed_dim,
|
||||||
|
groups=groups[i]
|
||||||
|
),
|
||||||
|
)
|
||||||
|
self.down_sample.append(
|
||||||
|
Downsample2D(
|
||||||
|
out_channels[i],
|
||||||
|
use_conv=True,
|
||||||
|
out_channels=out_channels[i],
|
||||||
|
padding=1,
|
||||||
|
name="op",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.mid_convs = nn.ModuleList()
|
||||||
|
self.mid_convs.append(nn.Sequential(
|
||||||
|
nn.Conv2d(
|
||||||
|
in_channels=out_channels[-1],
|
||||||
|
out_channels=out_channels[-1],
|
||||||
|
kernel_size=3,
|
||||||
|
stride=1,
|
||||||
|
padding=1
|
||||||
|
),
|
||||||
|
nn.ReLU(),
|
||||||
|
nn.GroupNorm(8, out_channels[-1]),
|
||||||
|
nn.Conv2d(
|
||||||
|
in_channels=out_channels[-1],
|
||||||
|
out_channels=out_channels[-1],
|
||||||
|
kernel_size=3,
|
||||||
|
stride=1,
|
||||||
|
padding=1
|
||||||
|
),
|
||||||
|
nn.GroupNorm(8, out_channels[-1]),
|
||||||
|
))
|
||||||
|
self.mid_convs.append(
|
||||||
|
nn.Conv2d(
|
||||||
|
in_channels=out_channels[-1],
|
||||||
|
out_channels=320,
|
||||||
|
kernel_size=1,
|
||||||
|
stride=1,
|
||||||
|
))
|
||||||
|
|
||||||
|
self.scale = 1.
|
||||||
|
|
||||||
|
def _set_gradient_checkpointing(self, module, value=False):
|
||||||
|
if hasattr(module, "gradient_checkpointing"):
|
||||||
|
module.gradient_checkpointing = value
|
||||||
|
|
||||||
|
# Copied from diffusers.models.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],
|
||||||
|
):
|
||||||
|
|
||||||
|
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_batch = self.time_embedding(t_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_batch.repeat_interleave(num_frames, dim=0)
|
||||||
|
|
||||||
|
sample = self.embedding(sample)
|
||||||
|
|
||||||
|
for res, downsample in zip(self.down_res, self.down_sample):
|
||||||
|
sample = res(sample, emb)
|
||||||
|
sample = downsample(sample, emb)
|
||||||
|
|
||||||
|
sample = self.mid_convs[0](sample) + sample
|
||||||
|
sample = self.mid_convs[1](sample)
|
||||||
|
|
||||||
|
return {
|
||||||
|
'output': sample,
|
||||||
|
'scale': self.scale,
|
||||||
|
}
|
||||||
|
|
||||||
@@ -0,0 +1,517 @@
|
|||||||
|
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.utils import BaseOutput, logging
|
||||||
|
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.models.unets.unet_3d_blocks import UNetMidBlockSpatioTemporal, get_down_block, get_up_block
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
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 UNetSpatioTemporalConditionControlNeXtModel(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,
|
||||||
|
upcast_attention: bool = False,
|
||||||
|
):
|
||||||
|
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`. `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`. `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`. `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`. `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`. `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",
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
)
|
||||||
|
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",
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
)
|
||||||
|
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` 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.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,
|
||||||
|
down_block_additional_residuals: Optional[Tuple[torch.Tensor]] = None,
|
||||||
|
mid_block_additional_residual: Optional[torch.Tensor] = None,
|
||||||
|
conditional_controls: Optional[torch.Tensor] = None,
|
||||||
|
return_dict: bool = True,
|
||||||
|
added_time_ids: torch.Tensor=None,
|
||||||
|
image_only_indicator: torch.Tensor=None,
|
||||||
|
) -> 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.
|
||||||
|
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 image_only_indicator is None:
|
||||||
|
image_only_indicator = torch.zeros(batch_size, num_frames, dtype=sample.dtype, device=sample.device)
|
||||||
|
|
||||||
|
down_block_res_samples = (sample,)
|
||||||
|
for idx,downsample_block in enumerate(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
|
||||||
|
|
||||||
|
if idx == 0 and conditional_controls is not None:
|
||||||
|
scale = conditional_controls['scale']
|
||||||
|
conditional_controls = conditional_controls['output']
|
||||||
|
mean_latents, std_latents = torch.mean(sample, dim=(1, 2, 3), keepdim=True), torch.std(sample, dim=(1, 2, 3), keepdim=True)
|
||||||
|
mean_control, std_control = torch.mean(conditional_controls, dim=(1, 2, 3), keepdim=True), torch.std(conditional_controls, dim=(1, 2, 3), keepdim=True)
|
||||||
|
conditional_controls = (conditional_controls - mean_control) * (std_latents / (std_control + 1e-5)) + mean_latents
|
||||||
|
conditional_controls = F.adaptive_avg_pool2d(conditional_controls, sample.shape[-2:])
|
||||||
|
|
||||||
|
sample = sample + conditional_controls * scale * 0.2
|
||||||
|
|
||||||
|
if down_block_additional_residuals is not None:
|
||||||
|
new_down_block_res_samples = ()
|
||||||
|
|
||||||
|
for down_block_res_sample, down_block_additional_residual in zip(
|
||||||
|
down_block_res_samples, down_block_additional_residuals
|
||||||
|
):
|
||||||
|
down_block_res_sample = down_block_res_sample + down_block_additional_residual
|
||||||
|
new_down_block_res_samples = new_down_block_res_samples + (down_block_res_sample,)
|
||||||
|
|
||||||
|
down_block_res_samples = new_down_block_res_samples
|
||||||
|
|
||||||
|
# 4. mid
|
||||||
|
sample = self.mid_block(
|
||||||
|
hidden_states=sample,
|
||||||
|
temb=emb,
|
||||||
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
|
image_only_indicator=image_only_indicator,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 5. up
|
||||||
|
for i, upsample_block in enumerate(self.up_blocks):
|
||||||
|
res_samples = down_block_res_samples[-len(upsample_block.resnets) :]
|
||||||
|
down_block_res_samples = down_block_res_samples[: -len(upsample_block.resnets)]
|
||||||
|
|
||||||
|
if hasattr(upsample_block, "has_cross_attention") and upsample_block.has_cross_attention:
|
||||||
|
sample = upsample_block(
|
||||||
|
hidden_states=sample,
|
||||||
|
temb=emb,
|
||||||
|
res_hidden_states_tuple=res_samples,
|
||||||
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
|
image_only_indicator=image_only_indicator,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
sample = upsample_block(
|
||||||
|
hidden_states=sample,
|
||||||
|
temb=emb,
|
||||||
|
res_hidden_states_tuple=res_samples,
|
||||||
|
image_only_indicator=image_only_indicator,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 6. post-process
|
||||||
|
sample = self.conv_norm_out(sample)
|
||||||
|
sample = self.conv_act(sample)
|
||||||
|
sample = self.conv_out(sample)
|
||||||
|
|
||||||
|
# 7. Reshape back to original shape
|
||||||
|
sample = sample.reshape(batch_size, num_frames, *sample.shape[1:])
|
||||||
|
|
||||||
|
if not return_dict:
|
||||||
|
return (sample,)
|
||||||
|
|
||||||
|
return UNetSpatioTemporalConditionOutput(sample=sample)
|
||||||
@@ -0,0 +1,561 @@
|
|||||||
|
import os
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
import numpy as np
|
||||||
|
import gc
|
||||||
|
|
||||||
|
import folder_paths
|
||||||
|
import comfy.model_management as mm
|
||||||
|
import comfy.utils
|
||||||
|
|
||||||
|
try:
|
||||||
|
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
|
||||||
|
except:
|
||||||
|
pass
|
||||||
|
|
||||||
|
from .pipeline.pipeline_stable_video_diffusion_controlnext import StableVideoDiffusionPipelineControlNeXt, tensor2vid
|
||||||
|
|
||||||
|
from .models.controlnext_vid_svd import ControlNeXtSDVModel
|
||||||
|
from .models.unet_spatio_temporal_condition_controlnext import UNetSpatioTemporalConditionControlNeXtModel
|
||||||
|
from .utils.scheduling_euler_discrete_karras_fix import EulerDiscreteScheduler as EulerDiscreteSchedulerKarras
|
||||||
|
from diffusers.schedulers import EulerDiscreteScheduler
|
||||||
|
|
||||||
|
from transformers import CLIPVisionModelWithProjection, CLIPImageProcessor
|
||||||
|
from diffusers import AutoencoderKLTemporalDecoder
|
||||||
|
|
||||||
|
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
|
||||||
|
|
||||||
|
from contextlib import nullcontext
|
||||||
|
try:
|
||||||
|
from accelerate import init_empty_weights
|
||||||
|
from accelerate.utils import set_module_tensor_to_device
|
||||||
|
is_accelerate_available = True
|
||||||
|
except:
|
||||||
|
is_accelerate_available = False
|
||||||
|
pass
|
||||||
|
|
||||||
|
import logging
|
||||||
|
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
||||||
|
log = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
def loglinear_interp(t_steps, num_steps):
|
||||||
|
"""
|
||||||
|
Performs log-linear interpolation of a given array of decreasing numbers.
|
||||||
|
"""
|
||||||
|
xs = np.linspace(0, 1, len(t_steps))
|
||||||
|
ys = np.log(t_steps[::-1])
|
||||||
|
|
||||||
|
new_xs = np.linspace(0, 1, num_steps)
|
||||||
|
new_ys = np.interp(new_xs, xs, ys)
|
||||||
|
|
||||||
|
interped_ys = np.exp(new_ys)[::-1].copy()
|
||||||
|
return interped_ys
|
||||||
|
|
||||||
|
class DownloadAndLoadControlNeXt:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
|
||||||
|
"precision": (
|
||||||
|
[
|
||||||
|
'fp32',
|
||||||
|
'fp16',
|
||||||
|
'bf16',
|
||||||
|
], {
|
||||||
|
"default": 'fp16'
|
||||||
|
}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("CONTROLNEXT_PIPE",)
|
||||||
|
RETURN_NAMES = ("controlnext_pipeline",)
|
||||||
|
FUNCTION = "loadmodel"
|
||||||
|
CATEGORY = "ControlNeXtWrapper"
|
||||||
|
|
||||||
|
def loadmodel(self, precision):
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||||
|
|
||||||
|
pbar = comfy.utils.ProgressBar(5)
|
||||||
|
|
||||||
|
download_path = os.path.join(folder_paths.models_dir, "diffusers", "controlnext")
|
||||||
|
unet_model_path = os.path.join(download_path, "controlnext-svd_v2-unet-fp16.safetensors")
|
||||||
|
contolnet_model_path = os.path.join(download_path, "controlnext-svd_v2-controlnet-fp16.safetensors")
|
||||||
|
|
||||||
|
if not os.path.exists(unet_model_path):
|
||||||
|
log.info(f"Downloading model to: {unet_model_path}")
|
||||||
|
from huggingface_hub import snapshot_download
|
||||||
|
snapshot_download(repo_id="Pbihao/ControlNeXt",
|
||||||
|
local_dir=download_path,
|
||||||
|
local_dir_use_symlinks=False)
|
||||||
|
|
||||||
|
log.info(f"Loading model from: {unet_model_path}")
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
if not os.path.exists(svd_path):
|
||||||
|
log.info(f"Downloading SVD model to: {svd_path}")
|
||||||
|
from huggingface_hub import snapshot_download
|
||||||
|
snapshot_download(repo_id="vdo/stable-video-diffusion-img2vid-xt-1-1",
|
||||||
|
allow_patterns=[f"*.json", "*fp16*"],
|
||||||
|
ignore_patterns=["*unet*"],
|
||||||
|
local_dir=svd_path,
|
||||||
|
local_dir_use_symlinks=False)
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
svd_path = os.path.join(folder_paths.models_dir, "diffusers", "stable-video-diffusion-img2vid-xt-1-1")
|
||||||
|
|
||||||
|
unet_config = UNetSpatioTemporalConditionControlNeXtModel.load_config(os.path.join(script_directory, "configs", "unet_config.json"))
|
||||||
|
log.info("Loading UNET")
|
||||||
|
with (init_empty_weights() if is_accelerate_available else nullcontext()):
|
||||||
|
self.unet = UNetSpatioTemporalConditionControlNeXtModel.from_config(unet_config)
|
||||||
|
sd = comfy.utils.load_torch_file(os.path.join(unet_model_path))
|
||||||
|
if is_accelerate_available:
|
||||||
|
for key in sd:
|
||||||
|
set_module_tensor_to_device(self.unet, key, dtype=dtype, device=device, value=sd[key])
|
||||||
|
else:
|
||||||
|
self.unet.load_state_dict(sd, strict=False)
|
||||||
|
del sd
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
log.info("Loading VAE")
|
||||||
|
self.vae = AutoencoderKLTemporalDecoder.from_pretrained(svd_path, subfolder="vae", variant="fp16", low_cpu_mem_usage=True).to(dtype).to(device).eval()
|
||||||
|
|
||||||
|
log.info("Loading IMAGE_ENCODER")
|
||||||
|
self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(svd_path, subfolder="image_encoder", variant="fp16", low_cpu_mem_usage=True).to(dtype).to(device).eval()
|
||||||
|
pbar.update(1)
|
||||||
|
self.noise_scheduler = EulerDiscreteScheduler.from_pretrained(svd_path, subfolder="scheduler")
|
||||||
|
self.feature_extractor = CLIPImageProcessor.from_pretrained(svd_path, subfolder="feature_extractor")
|
||||||
|
|
||||||
|
log.info("Loading ControlNeXt")
|
||||||
|
self.controlnext = ControlNeXtSDVModel()
|
||||||
|
self.controlnext.load_state_dict(comfy.utils.load_torch_file(os.path.join(contolnet_model_path)))
|
||||||
|
self.controlnext = self.controlnext.to(dtype).to(device).eval()
|
||||||
|
|
||||||
|
pipeline = StableVideoDiffusionPipelineControlNeXt(
|
||||||
|
vae = self.vae,
|
||||||
|
image_encoder = self.image_encoder,
|
||||||
|
unet = self.unet,
|
||||||
|
scheduler = self.noise_scheduler,
|
||||||
|
feature_extractor = self.feature_extractor,
|
||||||
|
controlnext=self.controlnext,
|
||||||
|
)
|
||||||
|
|
||||||
|
controlnextsvd_model = {
|
||||||
|
'pipeline': pipeline,
|
||||||
|
'dtype': dtype,
|
||||||
|
}
|
||||||
|
pbar.update(1)
|
||||||
|
return (controlnextsvd_model,)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
class ControlNextDiffusersScheduler:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"scheduler": (
|
||||||
|
[
|
||||||
|
'EulerDiscreteScheduler',
|
||||||
|
'EulerDiscreteSchedulerKarras',
|
||||||
|
'EulerDiscreteScheduler_AYS',
|
||||||
|
],
|
||||||
|
),
|
||||||
|
"sigma_min": ("FLOAT", {"default": 0.002, "min": 0.0, "max": 700.0, "step": 0.001}),
|
||||||
|
"sigma_max": ("FLOAT", {"default": 700.0, "min": 0.0, "max": 700.0, "step": 0.001}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("DIFFUSERS_SCHEDULER",)
|
||||||
|
RETURN_NAMES = ("scheduler",)
|
||||||
|
FUNCTION = "loadmodel"
|
||||||
|
CATEGORY = "ControlNeXtWrapper"
|
||||||
|
|
||||||
|
def loadmodel(self, scheduler, sigma_min, sigma_max):
|
||||||
|
|
||||||
|
scheduler_config = {
|
||||||
|
"beta_end": 0.012,
|
||||||
|
"beta_schedule": "scaled_linear",
|
||||||
|
"beta_start": 0.00085,
|
||||||
|
"clip_sample": False,
|
||||||
|
"interpolation_type": "linear",
|
||||||
|
"num_train_timesteps": 1000,
|
||||||
|
"prediction_type": "v_prediction",
|
||||||
|
"set_alpha_to_one": False,
|
||||||
|
"sigma_max": sigma_max,
|
||||||
|
"sigma_min": sigma_min,
|
||||||
|
"skip_prk_steps": True,
|
||||||
|
"steps_offset": 1,
|
||||||
|
"timestep_spacing": "leading",
|
||||||
|
"timestep_type": "continuous",
|
||||||
|
"trained_betas": None,
|
||||||
|
"use_karras_sigmas": False
|
||||||
|
}
|
||||||
|
if scheduler == 'EulerDiscreteScheduler':
|
||||||
|
noise_scheduler = EulerDiscreteScheduler.from_config(scheduler_config)
|
||||||
|
sigmas = None
|
||||||
|
elif scheduler == 'EulerDiscreteScheduler_AYS':
|
||||||
|
noise_scheduler = EulerDiscreteScheduler.from_config(scheduler_config)
|
||||||
|
sigmas = [700.00, 54.5, 15.886, 7.977, 4.248, 1.789, 0.981, 0.403, 0.173, 0.034, 0.002]
|
||||||
|
elif scheduler == 'EulerDiscreteSchedulerKarras':
|
||||||
|
scheduler_config['use_karras_sigmas'] = True
|
||||||
|
noise_scheduler = EulerDiscreteSchedulerKarras.from_config(scheduler_config)
|
||||||
|
sigmas = None
|
||||||
|
|
||||||
|
scheduler_options = {
|
||||||
|
"noise_scheduler": noise_scheduler,
|
||||||
|
"sigmas": sigmas,
|
||||||
|
}
|
||||||
|
|
||||||
|
return (scheduler_options,)
|
||||||
|
|
||||||
|
class ControlNextSampler:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"controlnext_pipeline": ("CONTROLNEXT_PIPE",),
|
||||||
|
"ref_image": ("IMAGE",),
|
||||||
|
"pose_images": ("IMAGE",),
|
||||||
|
"steps": ("INT", {"default": 25, "min": 1, "max": 200, "step": 1}),
|
||||||
|
"motion_bucket_id": ("INT", {"default": 127, "min": 0, "max": 1000, "step": 1}),
|
||||||
|
"cfg_min": ("FLOAT", {"default": 3.0, "min": 0.0, "max": 20.0, "step": 0.01}),
|
||||||
|
"cfg_max": ("FLOAT", {"default": 3.0, "min": 0.0, "max": 20.0, "step": 0.01}),
|
||||||
|
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||||
|
"fps": ("INT", {"default": 7, "min": 2, "max": 100, "step": 1}),
|
||||||
|
"controlnext_cond_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||||
|
"noise_aug_strength": ("FLOAT", {"default": 0.02, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||||
|
"context_size": ("INT", {"default": 24, "min": 1, "max": 128, "step": 1}),
|
||||||
|
"context_overlap": ("INT", {"default": 6, "min": 1, "max": 128, "step": 1}),
|
||||||
|
"keep_model_loaded": ("BOOLEAN", {"default": True}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"optional_scheduler": ("DIFFUSERS_SCHEDULER",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("LATENT",)
|
||||||
|
RETURN_NAMES = ("samples",)
|
||||||
|
FUNCTION = "process"
|
||||||
|
CATEGORY = "ControlNextWrapper"
|
||||||
|
|
||||||
|
def process(self, controlnext_pipeline, ref_image, pose_images, cfg_min, cfg_max, controlnext_cond_scale, motion_bucket_id, steps, seed, noise_aug_strength, fps, keep_model_loaded,
|
||||||
|
context_size, context_overlap, optional_scheduler=None):
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
offload_device = mm.unet_offload_device()
|
||||||
|
mm.unload_all_models()
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
dtype = controlnext_pipeline['dtype']
|
||||||
|
pipeline = controlnext_pipeline['pipeline']
|
||||||
|
|
||||||
|
original_scheduler = pipeline.scheduler
|
||||||
|
|
||||||
|
if optional_scheduler is not None:
|
||||||
|
log.info(f"Using optional scheduler: {optional_scheduler['noise_scheduler']}")
|
||||||
|
pipeline.scheduler = optional_scheduler['noise_scheduler']
|
||||||
|
sigmas = optional_scheduler['sigmas']
|
||||||
|
|
||||||
|
if sigmas is not None and (steps + 1) != len(sigmas):
|
||||||
|
sigmas = loglinear_interp(sigmas, steps + 1)
|
||||||
|
sigmas = sigmas[-(steps + 1):]
|
||||||
|
sigmas[-1] = 0
|
||||||
|
log.info(f"Using timesteps: {sigmas}")
|
||||||
|
else:
|
||||||
|
pipeline.scheduler = original_scheduler
|
||||||
|
sigmas = None
|
||||||
|
|
||||||
|
B, H, W, C = pose_images.shape
|
||||||
|
|
||||||
|
assert B >= context_size, "The number of poses must be greater than the context size"
|
||||||
|
|
||||||
|
ref_image = ref_image.permute(0, 3, 1, 2)
|
||||||
|
pose_images = pose_images.permute(0, 3, 1, 2)
|
||||||
|
pose_images = pose_images * 2 - 1
|
||||||
|
|
||||||
|
ref_image = ref_image.to(device).to(dtype)
|
||||||
|
pose_images = pose_images.to(device).to(dtype)
|
||||||
|
|
||||||
|
generator = torch.Generator(device=device)
|
||||||
|
generator.manual_seed(seed)
|
||||||
|
|
||||||
|
frames = pipeline(
|
||||||
|
ref_image,
|
||||||
|
pose_images,
|
||||||
|
num_frames=B,
|
||||||
|
frames_per_batch=context_size,
|
||||||
|
overlap=context_overlap,
|
||||||
|
motion_bucket_id=motion_bucket_id,
|
||||||
|
min_guidance_scale=cfg_min,
|
||||||
|
max_guidance_scale=cfg_max,
|
||||||
|
controlnext_cond_scale=controlnext_cond_scale,
|
||||||
|
height=H,
|
||||||
|
width=W,
|
||||||
|
fps=fps,
|
||||||
|
noise_aug_strength=noise_aug_strength,
|
||||||
|
num_inference_steps=steps,
|
||||||
|
generator=generator,
|
||||||
|
sigmas = sigmas,
|
||||||
|
decode_chunk_size=2,
|
||||||
|
output_type="latent",
|
||||||
|
return_dict="false",
|
||||||
|
#device=device,
|
||||||
|
).frames
|
||||||
|
|
||||||
|
if not keep_model_loaded:
|
||||||
|
pipeline.unet.to(offload_device)
|
||||||
|
pipeline.vae.to(offload_device)
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
gc.collect()
|
||||||
|
|
||||||
|
return {"samples": frames},
|
||||||
|
|
||||||
|
class ControlNextDecode:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"controlnext_pipeline": ("CONTROLNEXT_PIPE",),
|
||||||
|
"samples": ("LATENT",),
|
||||||
|
"decode_chunk_size": ("INT", {"default": 4, "min": 1, "max": 200, "step": 1})
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("IMAGE",)
|
||||||
|
RETURN_NAMES = ("images",)
|
||||||
|
FUNCTION = "process"
|
||||||
|
CATEGORY = "ControlNextWrapper"
|
||||||
|
|
||||||
|
def process(self, controlnext_pipeline, samples, decode_chunk_size):
|
||||||
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
|
pipeline = controlnext_pipeline['pipeline']
|
||||||
|
num_frames = samples['samples'].shape[0]
|
||||||
|
try:
|
||||||
|
frames = pipeline.decode_latents(samples['samples'], num_frames, decode_chunk_size)
|
||||||
|
except:
|
||||||
|
frames = pipeline.decode_latents(samples['samples'], num_frames, 1)
|
||||||
|
frames = tensor2vid(frames, pipeline.image_processor, output_type="pt")
|
||||||
|
|
||||||
|
frames = frames.squeeze(1)[1:].permute(0, 2, 3, 1).cpu().float()
|
||||||
|
|
||||||
|
return frames,
|
||||||
|
|
||||||
|
class ControlNextGetPoses:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"ref_image": ("IMAGE",),
|
||||||
|
"pose_images": ("IMAGE",),
|
||||||
|
"include_body": ("BOOLEAN", {"default": True}),
|
||||||
|
"include_hand": ("BOOLEAN", {"default": True}),
|
||||||
|
"include_face": ("BOOLEAN", {"default": True}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("IMAGE", "IMAGE",)
|
||||||
|
RETURN_NAMES = ("poses_with_ref", "pose_images")
|
||||||
|
FUNCTION = "process"
|
||||||
|
CATEGORY = "ControlNextWrapper"
|
||||||
|
|
||||||
|
def process(self, ref_image, pose_images, include_body, include_hand, include_face):
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
offload_device = mm.unet_offload_device()
|
||||||
|
from .dwpose.util import draw_pose
|
||||||
|
from .dwpose.dwpose_detector import DWposeDetector
|
||||||
|
|
||||||
|
assert ref_image.shape[1:3] == pose_images.shape[1:3], "ref_image and pose_images must have the same resolution"
|
||||||
|
|
||||||
|
#yolo_model = "yolox_l.onnx"
|
||||||
|
#dw_pose_model = "dw-ll_ucoco_384.onnx"
|
||||||
|
dw_pose_model = "dw-ll_ucoco_384_bs5.torchscript.pt"
|
||||||
|
yolo_model = "yolox_l.torchscript.pt"
|
||||||
|
|
||||||
|
model_base_path = os.path.join(script_directory, "models", "DWPose")
|
||||||
|
|
||||||
|
model_det=os.path.join(model_base_path, yolo_model)
|
||||||
|
model_pose=os.path.join(model_base_path, dw_pose_model)
|
||||||
|
|
||||||
|
if not os.path.exists(model_det):
|
||||||
|
log.info(f"Downloading yolo model to: {model_base_path}")
|
||||||
|
from huggingface_hub import snapshot_download
|
||||||
|
snapshot_download(repo_id="hr16/yolox-onnx",
|
||||||
|
allow_patterns=[f"*{yolo_model}*"],
|
||||||
|
local_dir=model_base_path,
|
||||||
|
local_dir_use_symlinks=False)
|
||||||
|
|
||||||
|
if not os.path.exists(model_pose):
|
||||||
|
log.info(f"Downloading dwpose model to: {model_base_path}")
|
||||||
|
from huggingface_hub import snapshot_download
|
||||||
|
snapshot_download(repo_id="hr16/DWPose-TorchScript-BatchSize5",
|
||||||
|
allow_patterns=[f"*{dw_pose_model}*"],
|
||||||
|
local_dir=model_base_path,
|
||||||
|
local_dir_use_symlinks=False)
|
||||||
|
|
||||||
|
model_det=os.path.join(model_base_path, yolo_model)
|
||||||
|
model_pose=os.path.join(model_base_path, dw_pose_model)
|
||||||
|
|
||||||
|
if not hasattr(self, "det") or not hasattr(self, "pose"):
|
||||||
|
self.det = torch.jit.load(model_det)
|
||||||
|
self.pose = torch.jit.load(model_pose)
|
||||||
|
|
||||||
|
self.dwprocessor = DWposeDetector(
|
||||||
|
model_det=self.det,
|
||||||
|
model_pose=self.pose)
|
||||||
|
|
||||||
|
ref_image = ref_image.squeeze(0).cpu().numpy() * 255
|
||||||
|
|
||||||
|
self.det = self.det.to(device)
|
||||||
|
self.pose = self.pose.to(device)
|
||||||
|
|
||||||
|
# select ref-keypoint from reference pose for pose rescale
|
||||||
|
ref_pose = self.dwprocessor(ref_image)
|
||||||
|
#ref_keypoint_id = [0, 1, 2, 5, 8, 11, 14, 15, 16, 17]
|
||||||
|
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 ref_pose['bodies']['score'].shape[0] > 0 and ref_pose['bodies']['score'][0][i] > 0.3]
|
||||||
|
if len(ref_pose['bodies']['subset']) > 0 and ref_pose['bodies']['subset'][0][i] >= .0]
|
||||||
|
ref_body = ref_pose['bodies']['candidate'][ref_keypoint_id]
|
||||||
|
|
||||||
|
height, width, _ = ref_image.shape
|
||||||
|
pose_images_np = pose_images.cpu().numpy() * 255
|
||||||
|
|
||||||
|
# read input video
|
||||||
|
pbar = comfy.utils.ProgressBar(len(pose_images_np))
|
||||||
|
detected_poses_np_list = []
|
||||||
|
for img_np in pose_images_np:
|
||||||
|
detected_poses_np_list.append(self.dwprocessor(img_np))
|
||||||
|
pbar.update(1)
|
||||||
|
|
||||||
|
self.det = self.det.to(offload_device)
|
||||||
|
self.pose = self.pose.to(offload_device)
|
||||||
|
|
||||||
|
detected_bodies = np.stack(
|
||||||
|
[p['bodies']['candidate'] for p in detected_poses_np_list if p['bodies']['candidate'].shape[0] == 18])[:,
|
||||||
|
ref_keypoint_id]
|
||||||
|
# compute linear-rescale params
|
||||||
|
ay, by = np.polyfit(detected_bodies[:, :, 1].flatten(), np.tile(ref_body[:, 1], len(detected_bodies)), 1)
|
||||||
|
fh, fw, _ = pose_images_np[0].shape
|
||||||
|
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_np_list:
|
||||||
|
if include_body:
|
||||||
|
detected_pose['bodies']['candidate'] = detected_pose['bodies']['candidate'] * a + b
|
||||||
|
if include_hand:
|
||||||
|
detected_pose['faces'] = detected_pose['faces'] * a + b
|
||||||
|
if include_face:
|
||||||
|
detected_pose['hands'] = detected_pose['hands'] * a + b
|
||||||
|
im = draw_pose(detected_pose, height, width, include_body=include_body, include_hand=include_hand, include_face=include_face)
|
||||||
|
output_pose.append(np.array(im))
|
||||||
|
|
||||||
|
output_pose_tensors = [torch.tensor(np.array(im)) for im in output_pose]
|
||||||
|
output_tensor = torch.stack(output_pose_tensors) / 255
|
||||||
|
|
||||||
|
ref_pose_img = draw_pose(ref_pose, height, width, include_body=include_body, include_hand=include_hand, include_face=include_face)
|
||||||
|
ref_pose_tensor = torch.tensor(np.array(ref_pose_img)) / 255
|
||||||
|
output_tensor = torch.cat((ref_pose_tensor.unsqueeze(0), output_tensor))
|
||||||
|
output_tensor = output_tensor.permute(0, 2, 3, 1).cpu().float()
|
||||||
|
|
||||||
|
return output_tensor, output_tensor[1:]
|
||||||
|
|
||||||
|
|
||||||
|
class ControlNextSVDApply:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"model": ("MODEL",),
|
||||||
|
"pose_images": ("IMAGE",),
|
||||||
|
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||||
|
"blocks": ("STRING",{"default": "3"}),
|
||||||
|
"input_block_patch_after_skip": ("BOOLEAN", {"default": True}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("MODEL",)
|
||||||
|
RETURN_NAMES = ("model", )
|
||||||
|
FUNCTION = "patch"
|
||||||
|
CATEGORY = "MimicMotionAdvanced"
|
||||||
|
|
||||||
|
def patch(self, model, pose_images, strength, blocks, input_block_patch_after_skip):
|
||||||
|
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
dtype = mm.unet_dtype()
|
||||||
|
|
||||||
|
B, H, W, C = pose_images.shape
|
||||||
|
|
||||||
|
pose_images = pose_images.clone()
|
||||||
|
pose_images = pose_images.permute(0, 3, 1, 2).unsqueeze(0)
|
||||||
|
#pose_images = pose_images * 2 - 1
|
||||||
|
#pose_images = pose_images.to(device).to(dtype)
|
||||||
|
|
||||||
|
if not hasattr(self, 'controlnext'):
|
||||||
|
self.controlnext = ControlNeXtSDVModel()
|
||||||
|
self.controlnext.load_state_dict(comfy.utils.load_torch_file(os.path.join(script_directory, 'models', 'controlnext-svd_v2-controlnet-fp16.safetensors')))
|
||||||
|
self.controlnext = self.controlnext.to(dtype).to(device).eval()
|
||||||
|
|
||||||
|
block_list = [int(x) for x in blocks.split(',')] #for testing, blocks 0-3 possible to apply to, 3 after skip so far best
|
||||||
|
|
||||||
|
def input_block_patch(h, transformer_options):
|
||||||
|
if transformer_options['block'][1] in block_list and 0 in transformer_options["cond_or_uncond"]:
|
||||||
|
|
||||||
|
sigma = transformer_options["sigmas"][0]
|
||||||
|
|
||||||
|
log_sigma = sigma.log()
|
||||||
|
min_log_sigma = torch.tensor(0.0002).log()
|
||||||
|
max_log_sigma = torch.tensor(700).log() #can I get these from the model?
|
||||||
|
normalized_log_sigma = (log_sigma - min_log_sigma) / (max_log_sigma - min_log_sigma)
|
||||||
|
|
||||||
|
#AnimateDiff-Evolved context windowing, is this method slower than it should be?
|
||||||
|
if "ad_params" in transformer_options and transformer_options["ad_params"]['sub_idxs'] is not None:
|
||||||
|
sub_idxs = transformer_options['ad_params']['sub_idxs']
|
||||||
|
controlnext_input = pose_images[:,sub_idxs].to(h.dtype).to(h.device).contiguous()
|
||||||
|
|
||||||
|
controlnext_input[:, 0, ...] = pose_images[:, 0, ...]
|
||||||
|
else:
|
||||||
|
controlnext_input = pose_images.to(h.dtype).to(h.device)
|
||||||
|
|
||||||
|
print("controlnext_input shape: ", controlnext_input.shape)
|
||||||
|
print("h shape: ", h.shape)
|
||||||
|
|
||||||
|
conditional_controls = self.controlnext(controlnext_input, normalized_log_sigma)['output']
|
||||||
|
|
||||||
|
mean_latents, std_latents = torch.mean(h, dim=(1, 2, 3), keepdim=True), torch.std(h, dim=(1, 2, 3), keepdim=True)
|
||||||
|
mean_control, std_control = torch.mean(conditional_controls, dim=(1, 2, 3), keepdim=True), torch.std(conditional_controls, dim=(1, 2, 3), keepdim=True)
|
||||||
|
conditional_controls = (conditional_controls - mean_control) * (std_latents / (std_control + 1e-5)) + mean_latents
|
||||||
|
conditional_controls = F.adaptive_avg_pool2d(conditional_controls, h.shape[-2:])
|
||||||
|
|
||||||
|
h = h + conditional_controls * 0.2 * strength
|
||||||
|
|
||||||
|
return h
|
||||||
|
model_clone = model.clone()
|
||||||
|
if not input_block_patch_after_skip:
|
||||||
|
model_clone.set_model_input_block_patch(input_block_patch)
|
||||||
|
else:
|
||||||
|
model_clone.set_model_input_block_patch_after_skip(input_block_patch)
|
||||||
|
|
||||||
|
return (model_clone, )
|
||||||
|
|
||||||
|
NODE_CLASS_MAPPINGS = {
|
||||||
|
"DownloadAndLoadControlNeXt": DownloadAndLoadControlNeXt,
|
||||||
|
"ControlNextSampler": ControlNextSampler,
|
||||||
|
"ControlNextDecode": ControlNextDecode,
|
||||||
|
"ControlNextGetPoses": ControlNextGetPoses,
|
||||||
|
"ControlNextDiffusersScheduler": ControlNextDiffusersScheduler,
|
||||||
|
"ControlNextSVDApply": ControlNextSVDApply
|
||||||
|
}
|
||||||
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
|
"DownloadAndLoadControlNeXt": "(Down)Load ControlNeXt",
|
||||||
|
"ControlNextSampler": "ControlNext Sampler",
|
||||||
|
"ControlNextDecode": "ControlNext Decode",
|
||||||
|
"ControlNextGetPoses": "ControlNext GetPoses",
|
||||||
|
"ControlNextDiffusersScheduler": "ControlNext Diffusers Scheduler",
|
||||||
|
"ControlNextSVDApply": "ControlNext SVD Apply"
|
||||||
|
}
|
||||||
@@ -0,0 +1,770 @@
|
|||||||
|
# Copyright 2023 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.
|
||||||
|
|
||||||
|
import inspect
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Callable, Dict, List, Optional, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import PIL.Image
|
||||||
|
import torch
|
||||||
|
from transformers import CLIPImageProcessor, CLIPVisionModelWithProjection
|
||||||
|
from ..models.controlnext_vid_svd import ControlNeXtSDVModel
|
||||||
|
|
||||||
|
from diffusers.image_processor import VaeImageProcessor
|
||||||
|
from diffusers.models import AutoencoderKLTemporalDecoder, UNetSpatioTemporalConditionModel
|
||||||
|
from diffusers.utils import BaseOutput, logging
|
||||||
|
from diffusers.utils.torch_utils import randn_tensor
|
||||||
|
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||||
|
from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion import retrieve_timesteps
|
||||||
|
from ..models.unet_spatio_temporal_condition_controlnext import UNetSpatioTemporalConditionControlNeXtModel
|
||||||
|
from ..utils.scheduling_euler_discrete_karras_fix import EulerDiscreteScheduler
|
||||||
|
#from diffusers.pipelines.utils import PIL_INTERPOLATION, BaseOutput, logging
|
||||||
|
from diffusers.pipelines.stable_video_diffusion.pipeline_stable_video_diffusion import StableVideoDiffusionPipeline
|
||||||
|
|
||||||
|
from comfy.utils import ProgressBar
|
||||||
|
import comfy.model_management as mm
|
||||||
|
from comfy.clip_vision import clip_preprocess
|
||||||
|
offload_device = mm.unet_offload_device()
|
||||||
|
|
||||||
|
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||||
|
|
||||||
|
def _get_add_time_ids(
|
||||||
|
noise_aug_strength,
|
||||||
|
dtype,
|
||||||
|
batch_size,
|
||||||
|
fps=4,
|
||||||
|
motion_bucket_id=128,
|
||||||
|
unet=None,
|
||||||
|
):
|
||||||
|
add_time_ids = [fps, motion_bucket_id, noise_aug_strength]
|
||||||
|
|
||||||
|
passed_add_embed_dim = unet.config.addition_time_embed_dim * len(add_time_ids)
|
||||||
|
expected_add_embed_dim = 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}, but a vector of {passed_add_embed_dim} was created. The model has an incorrect config. 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)
|
||||||
|
|
||||||
|
|
||||||
|
return add_time_ids
|
||||||
|
|
||||||
|
|
||||||
|
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, output_type="pt"):
|
||||||
|
# Based on:
|
||||||
|
# https://github.com/modelscope/modelscope/blob/1509fdb973e5871f37148a4b5e5964cafd43e64d/modelscope/pipelines/multi_modal/text_to_video_synthesis_pipeline.py#L78
|
||||||
|
|
||||||
|
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 == "pt":
|
||||||
|
outputs = torch.stack(outputs)
|
||||||
|
return outputs
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class StableVideoDiffusionPipelineOutput(BaseOutput):
|
||||||
|
r"""
|
||||||
|
Output class for zero-shot text-to-video pipeline.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
frames (`[List[PIL.Image.Image]`, `np.ndarray`]):
|
||||||
|
List of denoised PIL images of length `batch_size` or NumPy array of shape `(batch_size, height, width,
|
||||||
|
num_channels)`.
|
||||||
|
"""
|
||||||
|
|
||||||
|
frames: Union[List[PIL.Image.Image], np.ndarray]
|
||||||
|
|
||||||
|
|
||||||
|
class StableVideoDiffusionPipelineControlNeXt(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 ([`AutoencoderKL`]):
|
||||||
|
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.
|
||||||
|
"""
|
||||||
|
|
||||||
|
model_cpu_offload_seq = "image_encoder->unet->vae"
|
||||||
|
_callback_tensor_inputs = ["latents"]
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
vae: AutoencoderKLTemporalDecoder,
|
||||||
|
image_encoder: CLIPVisionModelWithProjection,
|
||||||
|
unet: UNetSpatioTemporalConditionControlNeXtModel,
|
||||||
|
controlnext: ControlNeXtSDVModel,
|
||||||
|
scheduler: EulerDiscreteScheduler,
|
||||||
|
feature_extractor: CLIPImageProcessor,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.register_modules(
|
||||||
|
vae=vae,
|
||||||
|
image_encoder=image_encoder,
|
||||||
|
controlnext=controlnext,
|
||||||
|
unet=unet,
|
||||||
|
scheduler=scheduler,
|
||||||
|
feature_extractor=feature_extractor,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
|
||||||
|
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
|
||||||
|
|
||||||
|
|
||||||
|
def _encode_image(self, image, device, num_videos_per_prompt, do_classifier_free_guidance):
|
||||||
|
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.permute(0, 2, 3, 1)
|
||||||
|
image = clip_preprocess(image.clone(), 224)
|
||||||
|
|
||||||
|
image = image.to(device=device, dtype=dtype)
|
||||||
|
self.image_encoder.to(device)
|
||||||
|
image_embeddings = self.image_encoder(image).image_embeds
|
||||||
|
image_embeddings = image_embeddings.unsqueeze(1)
|
||||||
|
self.image_encoder.to(offload_device)
|
||||||
|
|
||||||
|
# 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,
|
||||||
|
num_videos_per_prompt,
|
||||||
|
do_classifier_free_guidance,
|
||||||
|
):
|
||||||
|
image = image.to(device=device)
|
||||||
|
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,
|
||||||
|
motion_bucket_id,
|
||||||
|
noise_aug_strength,
|
||||||
|
dtype,
|
||||||
|
batch_size,
|
||||||
|
num_videos_per_prompt,
|
||||||
|
do_classifier_free_guidance,
|
||||||
|
):
|
||||||
|
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}, but a vector of {passed_add_embed_dim} was created. The model has an incorrect config. 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, num_frames, decode_chunk_size=14):
|
||||||
|
# [batch, frames, channels, height, width] -> [batch*frames, channels, height, width]
|
||||||
|
latents = latents.flatten(0, 1)
|
||||||
|
|
||||||
|
latents = 1 / self.vae.config.scaling_factor * latents
|
||||||
|
|
||||||
|
accepts_num_frames = "num_frames" in set(inspect.signature(self.vae.forward).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
|
||||||
|
decode_kwargs["num_frames"] = num_frames_in
|
||||||
|
frame = self.vae.decode(latents[i : i + decode_chunk_size], **decode_kwargs).sample
|
||||||
|
frames.append(frame)
|
||||||
|
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,
|
||||||
|
num_frames,
|
||||||
|
num_channels_latents,
|
||||||
|
height,
|
||||||
|
width,
|
||||||
|
dtype,
|
||||||
|
device,
|
||||||
|
generator,
|
||||||
|
latents=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):
|
||||||
|
return self._guidance_scale >= 1 and self.unet.config.time_cond_proj_dim is None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def num_timesteps(self):
|
||||||
|
return self._num_timesteps
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def __call__(
|
||||||
|
self,
|
||||||
|
image: Union[PIL.Image.Image, List[PIL.Image.Image], torch.FloatTensor],
|
||||||
|
controlnext_condition:Optional[torch.FloatTensor] = None,
|
||||||
|
height: int = 576,
|
||||||
|
width: int = 1024,
|
||||||
|
num_frames: Optional[int] = None,
|
||||||
|
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: int = 0.02,
|
||||||
|
sigmas: Optional[List[float]] = None,
|
||||||
|
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,
|
||||||
|
controlnext_cond_scale=1.0,
|
||||||
|
batch_size=1,
|
||||||
|
overlap=5,
|
||||||
|
frames_per_batch = 14,
|
||||||
|
):
|
||||||
|
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 (`int`, *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.
|
||||||
|
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 or not to return a [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] instead of a
|
||||||
|
plain tuple.
|
||||||
|
|
||||||
|
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
|
||||||
|
frames_per_batch = min(frames_per_batch, 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
|
||||||
|
|
||||||
|
# 3. Encode input image
|
||||||
|
image_embeddings = self._encode_image(image, device, num_videos_per_prompt, do_classifier_free_guidance)
|
||||||
|
|
||||||
|
# NOTE: Stable Diffusion Video was conditioned on fps - 1, which
|
||||||
|
# is why it is reduced here.
|
||||||
|
# See: https://github.com/Stability-AI/generative-models/blob/ed0997173f98eaf8f4edf7ba5fe8f15c6b877fd3/scripts/sampling/simple_video_sample.py#L188
|
||||||
|
fps = fps - 1
|
||||||
|
|
||||||
|
# 4. Encode input image using VAE
|
||||||
|
image = self.image_processor.preprocess(image, height=height, width=width)
|
||||||
|
noise = randn_tensor(image.shape, generator=generator, device=image.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, num_videos_per_prompt, do_classifier_free_guidance)
|
||||||
|
image_latents = image_latents.to(image_embeddings.dtype)
|
||||||
|
|
||||||
|
# cast back to fp16 if needed
|
||||||
|
# if needs_upcasting:
|
||||||
|
# self.vae.to(dtype=self_vae_dtype)
|
||||||
|
|
||||||
|
# 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)
|
||||||
|
#image_latents = torch.cat([image_latents] * 2) if do_classifier_free_guidance else image_latents
|
||||||
|
|
||||||
|
# 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,
|
||||||
|
do_classifier_free_guidance,
|
||||||
|
)
|
||||||
|
added_time_ids = added_time_ids.to(device)
|
||||||
|
|
||||||
|
# 4. Prepare timesteps
|
||||||
|
if sigmas is not None:
|
||||||
|
timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, None, sigmas)
|
||||||
|
else:
|
||||||
|
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,
|
||||||
|
num_frames,
|
||||||
|
num_channels_latents,
|
||||||
|
height,
|
||||||
|
width,
|
||||||
|
image_embeddings.dtype,
|
||||||
|
device,
|
||||||
|
generator,
|
||||||
|
latents,
|
||||||
|
)
|
||||||
|
#prepare controlnext condition
|
||||||
|
controlnext_condition = self.image_processor.preprocess(controlnext_condition, height=height, width=width)
|
||||||
|
controlnext_condition = (controlnext_condition + 1.0) / 2
|
||||||
|
controlnext_condition = controlnext_condition.unsqueeze(0)
|
||||||
|
if do_classifier_free_guidance:
|
||||||
|
controlnext_condition = torch.cat([controlnext_condition] * 2)
|
||||||
|
controlnext_condition = controlnext_condition.to(device, latents.dtype)
|
||||||
|
controlnext_condition_all = controlnext_condition * controlnext_cond_scale
|
||||||
|
latents_all = latents
|
||||||
|
|
||||||
|
# 7. Prepare guidance scale
|
||||||
|
guidance_scale = torch.linspace(min_guidance_scale, max_guidance_scale, frames_per_batch).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
|
||||||
|
|
||||||
|
noise_aug_strength = 0.02 #"¯\_(ツ)_/¯
|
||||||
|
added_time_ids = _get_add_time_ids(
|
||||||
|
noise_aug_strength,
|
||||||
|
image_embeddings.dtype,
|
||||||
|
batch_size,
|
||||||
|
6,
|
||||||
|
128,
|
||||||
|
unet=self.unet,
|
||||||
|
)
|
||||||
|
if do_classifier_free_guidance:
|
||||||
|
added_time_ids = torch.cat([added_time_ids] * 2)
|
||||||
|
added_time_ids = added_time_ids.to(latents.device)
|
||||||
|
|
||||||
|
|
||||||
|
# 8. Denoising loop
|
||||||
|
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||||
|
self._num_timesteps = len(timesteps)
|
||||||
|
comfy_pbar = ProgressBar(num_inference_steps)
|
||||||
|
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||||
|
for i, t in enumerate(timesteps):
|
||||||
|
pred_tmp = torch.zeros_like(latents_all)
|
||||||
|
counter = torch.zeros((latents.shape[0], num_frames, 1, 1, 1 )).to(device=latents.device)
|
||||||
|
for batch, ind_start_idx in enumerate(range(0, num_frames-overlap, frames_per_batch-overlap)):
|
||||||
|
self.scheduler._step_index = None
|
||||||
|
if ind_start_idx + frames_per_batch > num_frames:
|
||||||
|
ind_start = num_frames - frames_per_batch
|
||||||
|
else:
|
||||||
|
ind_start = ind_start_idx
|
||||||
|
latents = latents_all[:,ind_start:ind_start+frames_per_batch].contiguous()
|
||||||
|
controlnext_condition = controlnext_condition_all[:,ind_start:ind_start+frames_per_batch].contiguous()
|
||||||
|
|
||||||
|
controlnext_condition[:, 0, ...] = controlnext_condition_all[:, 0, ...]
|
||||||
|
|
||||||
|
# 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)
|
||||||
|
|
||||||
|
controlnext_output = self.controlnext(
|
||||||
|
controlnext_condition,
|
||||||
|
t,
|
||||||
|
)
|
||||||
|
if do_classifier_free_guidance:
|
||||||
|
N = controlnext_output['output'].shape[0]
|
||||||
|
controlnext_output['scale'] = torch.tensor(controlnext_output['scale']).to(latent_model_input).repeat(N)[:, None, None, None]
|
||||||
|
controlnext_output['scale'][:N // 2] *= 0
|
||||||
|
|
||||||
|
|
||||||
|
# Concatenate image_latents over channels dimention
|
||||||
|
latent_model_input = torch.cat([latent_model_input, image_latents[:,ind_start:ind_start+frames_per_batch].contiguous()], dim=2)
|
||||||
|
|
||||||
|
|
||||||
|
# predict the noise residual
|
||||||
|
noise_pred = self.unet(
|
||||||
|
latent_model_input,
|
||||||
|
t,
|
||||||
|
encoder_hidden_states=image_embeddings,
|
||||||
|
added_time_ids=added_time_ids,
|
||||||
|
conditional_controls=controlnext_output,
|
||||||
|
return_dict=False,
|
||||||
|
)[0]
|
||||||
|
|
||||||
|
# perform guidance
|
||||||
|
if 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).prev_sample
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
if ind_start == 0:
|
||||||
|
pred_tmp[:,ind_start:ind_start+frames_per_batch] += latents
|
||||||
|
counter[:,ind_start:ind_start+frames_per_batch] += 1
|
||||||
|
else:
|
||||||
|
pred_tmp[:,ind_start + 1 : ind_start+frames_per_batch] += latents[:, 1:, ...]
|
||||||
|
counter[:,ind_start + 1:ind_start+frames_per_batch] += 1
|
||||||
|
pred_tmp /= counter
|
||||||
|
latents_all = pred_tmp
|
||||||
|
|
||||||
|
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||||
|
progress_bar.update()
|
||||||
|
comfy_pbar.update(1)
|
||||||
|
latents = latents_all
|
||||||
|
|
||||||
|
if not output_type == "latent":
|
||||||
|
# cast back to fp16 if needed
|
||||||
|
#if needs_upcasting:
|
||||||
|
# self.vae.to(dtype=self_vae_dtype)
|
||||||
|
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 StableVideoDiffusionPipelineOutput(frames=frames)
|
||||||
|
|
||||||
|
|
||||||
|
# resizing utils
|
||||||
|
# TODO: clean up later
|
||||||
|
def _resize_with_antialiasing(input, size, interpolation="bicubic", align_corners=True):
|
||||||
|
|
||||||
|
if input.ndim == 3:
|
||||||
|
input = input.unsqueeze(0) # Add a batch dimension
|
||||||
|
|
||||||
|
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
|
||||||
@@ -0,0 +1,5 @@
|
|||||||
|
diffusers>=0.30.0
|
||||||
|
accelerate
|
||||||
|
huggingface_hub
|
||||||
|
transformers
|
||||||
|
opencv-python
|
||||||
@@ -0,0 +1,282 @@
|
|||||||
|
import os
|
||||||
|
import torch
|
||||||
|
import numpy as np
|
||||||
|
from PIL import Image
|
||||||
|
from pipeline.pipeline_stable_video_diffusion_controlnext import StableVideoDiffusionPipelineControlNeXt
|
||||||
|
from models.controlnext_vid_svd import ControlNeXtSDVModel
|
||||||
|
from models.unet_spatio_temporal_condition_controlnext import UNetSpatioTemporalConditionControlNeXtModel
|
||||||
|
from transformers import CLIPVisionModelWithProjection
|
||||||
|
import re
|
||||||
|
from diffusers import AutoencoderKLTemporalDecoder
|
||||||
|
from moviepy.editor import ImageSequenceClip
|
||||||
|
from decord import VideoReader
|
||||||
|
import argparse
|
||||||
|
from safetensors.torch import load_file
|
||||||
|
from utils.pre_process import preprocess
|
||||||
|
|
||||||
|
|
||||||
|
def write_mp4(video_path, samples, fps=14, audio_bitrate="192k"):
|
||||||
|
clip = ImageSequenceClip(samples, fps=fps)
|
||||||
|
clip.write_videofile(video_path, audio_codec="aac", audio_bitrate=audio_bitrate,
|
||||||
|
ffmpeg_params=["-crf", "18", "-preset", "slow"])
|
||||||
|
|
||||||
|
def save_vid_side_by_side(batch_output, validation_control_images, output_folder, fps):
|
||||||
|
# Helper function to convert tensors to PIL images and save as GIF
|
||||||
|
flattened_batch_output = [img for sublist in batch_output for img in sublist]
|
||||||
|
video_path = output_folder+'/test_1.mp4'
|
||||||
|
final_images = []
|
||||||
|
outputs = []
|
||||||
|
# Helper function to concatenate images horizontally
|
||||||
|
def get_concat_h(im1, im2):
|
||||||
|
dst = Image.new('RGB', (im1.width + im2.width, max(im1.height, im2.height)))
|
||||||
|
dst.paste(im1, (0, 0))
|
||||||
|
dst.paste(im2, (im1.width, 0))
|
||||||
|
return dst
|
||||||
|
for image_list in zip(validation_control_images, flattened_batch_output):
|
||||||
|
predict_img = image_list[1].resize(image_list[0].size)
|
||||||
|
result = get_concat_h(image_list[0], predict_img)
|
||||||
|
final_images.append(np.array(result))
|
||||||
|
outputs.append(np.array(predict_img))
|
||||||
|
write_mp4(video_path, final_images, fps=fps)
|
||||||
|
|
||||||
|
output_path = output_folder + "/output.mp4"
|
||||||
|
write_mp4(output_path, outputs, fps=fps)
|
||||||
|
|
||||||
|
|
||||||
|
def load_images_from_folder_to_pil(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_video_to_pil(video_path):
|
||||||
|
images = []
|
||||||
|
|
||||||
|
vr = VideoReader(video_path)
|
||||||
|
length = len(vr)
|
||||||
|
|
||||||
|
for idx in range(length):
|
||||||
|
frame = vr[idx].asnumpy()
|
||||||
|
images.append(Image.fromarray(frame))
|
||||||
|
return images
|
||||||
|
|
||||||
|
|
||||||
|
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_control_images_folder",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
required=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--validation_control_video_path",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
required=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
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.,
|
||||||
|
required=False
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--num_inference_steps",
|
||||||
|
type=int,
|
||||||
|
default=25,
|
||||||
|
required=False
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--controlnext_path",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
required=True
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--unet_path",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
required=True
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--max_frame_num",
|
||||||
|
type=int,
|
||||||
|
default=50,
|
||||||
|
required=False
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--ref_image_path",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
required=True
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--batch_frames",
|
||||||
|
type=int,
|
||||||
|
default=14,
|
||||||
|
required=False
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--overlap",
|
||||||
|
type=int,
|
||||||
|
default=4,
|
||||||
|
required=False
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--sample_stride",
|
||||||
|
type=int,
|
||||||
|
default=2,
|
||||||
|
required=False
|
||||||
|
)
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
return args
|
||||||
|
|
||||||
|
|
||||||
|
def load_tensor(tensor_path):
|
||||||
|
if os.path.splitext(tensor_path)[1] == '.bin':
|
||||||
|
return torch.load(tensor_path)
|
||||||
|
elif os.path.splitext(tensor_path)[1] == ".safetensors":
|
||||||
|
return load_file(tensor_path)
|
||||||
|
else:
|
||||||
|
print("without supported tensors")
|
||||||
|
os._exit()
|
||||||
|
|
||||||
|
|
||||||
|
# Main script
|
||||||
|
if __name__ == "__main__":
|
||||||
|
args = parse_args()
|
||||||
|
|
||||||
|
assert (args.validation_control_images_folder is None) ^ (args.validation_control_video_path is None), "must and only one of [validation_control_images_folder, validation_control_video_path] should be given"
|
||||||
|
|
||||||
|
unet = UNetSpatioTemporalConditionControlNeXtModel.from_pretrained(
|
||||||
|
args.pretrained_model_name_or_path,
|
||||||
|
subfolder="unet",
|
||||||
|
low_cpu_mem_usage=True,
|
||||||
|
)
|
||||||
|
controlnext = ControlNeXtSDVModel()
|
||||||
|
controlnext.load_state_dict(load_tensor(args.controlnext_path))
|
||||||
|
unet.load_state_dict(load_tensor(args.unet_path), strict=False)
|
||||||
|
|
||||||
|
image_encoder = CLIPVisionModelWithProjection.from_pretrained(
|
||||||
|
args.pretrained_model_name_or_path, subfolder="image_encoder")
|
||||||
|
vae = AutoencoderKLTemporalDecoder.from_pretrained(
|
||||||
|
args.pretrained_model_name_or_path, subfolder="vae")
|
||||||
|
|
||||||
|
pipeline = StableVideoDiffusionPipelineControlNeXt.from_pretrained(
|
||||||
|
args.pretrained_model_name_or_path,
|
||||||
|
controlnext=controlnext,
|
||||||
|
unet=unet,
|
||||||
|
vae=vae,
|
||||||
|
image_encoder=image_encoder)
|
||||||
|
# pipeline.to(dtype=torch.float16)
|
||||||
|
pipeline.enable_model_cpu_offload()
|
||||||
|
|
||||||
|
os.makedirs(args.output_dir, exist_ok=True)
|
||||||
|
|
||||||
|
# Inference and saving loop
|
||||||
|
# ref_image = Image.open(args.ref_image_path).convert('RGB')
|
||||||
|
# ref_image = ref_image.resize((args.width, args.height))
|
||||||
|
# validation_control_images = [img.resize((args.width, args.height)) for img in validation_control_images]
|
||||||
|
|
||||||
|
validation_control_images, ref_image = preprocess(args.validation_control_video_path, args.ref_image_path, width=args.width, height=args.height, max_frame_num=args.max_frame_num, sample_stride=args.sample_stride)
|
||||||
|
|
||||||
|
|
||||||
|
final_result = []
|
||||||
|
frames = args.batch_frames
|
||||||
|
num_frames = min(args.max_frame_num, len(validation_control_images))
|
||||||
|
|
||||||
|
for i in range(num_frames):
|
||||||
|
validation_control_images[i] = Image.fromarray(np.array(validation_control_images[i]))
|
||||||
|
|
||||||
|
video_frames = pipeline(
|
||||||
|
ref_image,
|
||||||
|
validation_control_images[:num_frames],
|
||||||
|
decode_chunk_size=2,
|
||||||
|
num_frames=num_frames,
|
||||||
|
motion_bucket_id=127.0,
|
||||||
|
fps=7,
|
||||||
|
controlnext_cond_scale=1.0,
|
||||||
|
width=args.width,
|
||||||
|
height=args.height,
|
||||||
|
min_guidance_scale=args.guidance_scale,
|
||||||
|
max_guidance_scale=args.guidance_scale,
|
||||||
|
frames_per_batch=frames,
|
||||||
|
num_inference_steps=args.num_inference_steps,
|
||||||
|
overlap=args.overlap).frames[0]
|
||||||
|
final_result.append(video_frames)
|
||||||
|
|
||||||
|
fps =VideoReader(args.validation_control_video_path).get_avg_fps() // args.sample_stride
|
||||||
|
|
||||||
|
save_vid_side_by_side(
|
||||||
|
final_result,
|
||||||
|
validation_control_images[:num_frames],
|
||||||
|
args.output_dir,
|
||||||
|
fps=fps)
|
||||||
@@ -0,0 +1,53 @@
|
|||||||
|
import os
|
||||||
|
import argparse
|
||||||
|
import logging
|
||||||
|
import math
|
||||||
|
from omegaconf import OmegaConf
|
||||||
|
from datetime import datetime
|
||||||
|
from pathlib import Path
|
||||||
|
from PIL import Image
|
||||||
|
import numpy as np
|
||||||
|
import torch.jit
|
||||||
|
from torchvision.datasets.folder import pil_loader
|
||||||
|
from torchvision.transforms.functional import pil_to_tensor, resize, center_crop
|
||||||
|
from torchvision.transforms.functional import to_pil_image
|
||||||
|
from dwpose.preprocess import get_image_pose, get_video_pose
|
||||||
|
|
||||||
|
ASPECT_RATIO = 9 / 16
|
||||||
|
|
||||||
|
def preprocess(video_path, image_path, width=576, height=1024, sample_stride=2, max_frame_num=None):
|
||||||
|
"""preprocess ref image pose and video pose
|
||||||
|
|
||||||
|
Args:
|
||||||
|
video_path (str): input video pose path
|
||||||
|
image_path (str): reference image path
|
||||||
|
resolution (int, optional): Defaults to 576.
|
||||||
|
sample_stride (int, optional): Defaults to 2.
|
||||||
|
"""
|
||||||
|
image_pixels = pil_loader(image_path)
|
||||||
|
image_pixels = pil_to_tensor(image_pixels) # (c, h, w)
|
||||||
|
h, w = image_pixels.shape[-2:]
|
||||||
|
############################ compute target h/w according to original aspect ratio ###############################
|
||||||
|
# if h>w:
|
||||||
|
# w_target, h_target = resolution, int(resolution / ASPECT_RATIO // 64) * 64
|
||||||
|
# else:
|
||||||
|
# w_target, h_target = int(resolution / ASPECT_RATIO // 64) * 64, resolution
|
||||||
|
w_target, h_target = width, height
|
||||||
|
h_w_ratio = float(h) / float(w)
|
||||||
|
if h_w_ratio < h_target / w_target:
|
||||||
|
h_resize, w_resize = h_target, math.ceil(h_target / h_w_ratio)
|
||||||
|
else:
|
||||||
|
h_resize, w_resize = math.ceil(w_target * h_w_ratio), w_target
|
||||||
|
image_pixels = resize(image_pixels, [h_resize, w_resize], antialias=None)
|
||||||
|
image_pixels = center_crop(image_pixels, [h_target, w_target])
|
||||||
|
image_pixels = image_pixels.permute((1, 2, 0)).numpy()
|
||||||
|
##################################### get image&video pose value #################################################
|
||||||
|
image_pose = get_image_pose(image_pixels)
|
||||||
|
video_pose = get_video_pose(video_path, image_pixels, sample_stride=sample_stride, max_frame_num=max_frame_num)
|
||||||
|
pose_pixels = np.concatenate([np.expand_dims(image_pose, 0), video_pose])
|
||||||
|
# image_pixels = np.transpose(np.expand_dims(image_pixels, 0), (0, 3, 1, 2))
|
||||||
|
image_pixels = Image.fromarray(image_pixels)
|
||||||
|
pose_pixels = [Image.fromarray(p.transpose((1,2,0))) for p in pose_pixels]
|
||||||
|
# return torch.from_numpy(pose_pixels.copy()) / 127.5 - 1, torch.from_numpy(image_pixels) / 127.5 - 1
|
||||||
|
return pose_pixels, image_pixels
|
||||||
|
|
||||||
@@ -0,0 +1,556 @@
|
|||||||
|
# Copyright 2023 Katherine Crowson and 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.
|
||||||
|
|
||||||
|
import math
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||||
|
from diffusers.utils import BaseOutput, logging
|
||||||
|
from diffusers.utils.torch_utils import randn_tensor
|
||||||
|
from diffusers.schedulers.scheduling_utils import KarrasDiffusionSchedulers, SchedulerMixin
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
|
||||||
|
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
# Copied from diffusers.schedulers.scheduling_ddpm.DDPMSchedulerOutput with DDPM->EulerDiscrete
|
||||||
|
class EulerDiscreteSchedulerOutput(BaseOutput):
|
||||||
|
"""
|
||||||
|
Output class for the scheduler's `step` function output.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
|
||||||
|
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
|
||||||
|
denoising loop.
|
||||||
|
pred_original_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
|
||||||
|
The predicted denoised sample `(x_{0})` based on the model output from the current timestep.
|
||||||
|
`pred_original_sample` can be used to preview progress or for guidance.
|
||||||
|
"""
|
||||||
|
|
||||||
|
prev_sample: torch.FloatTensor
|
||||||
|
pred_original_sample: Optional[torch.FloatTensor] = None
|
||||||
|
|
||||||
|
|
||||||
|
# Copied from diffusers.schedulers.scheduling_ddpm.betas_for_alpha_bar
|
||||||
|
def betas_for_alpha_bar(
|
||||||
|
num_diffusion_timesteps,
|
||||||
|
max_beta=0.999,
|
||||||
|
alpha_transform_type="cosine",
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Create a beta schedule that discretizes the given alpha_t_bar function, which defines the cumulative product of
|
||||||
|
(1-beta) over time from t = [0,1].
|
||||||
|
|
||||||
|
Contains a function alpha_bar that takes an argument t and transforms it to the cumulative product of (1-beta) up
|
||||||
|
to that part of the diffusion process.
|
||||||
|
|
||||||
|
|
||||||
|
Args:
|
||||||
|
num_diffusion_timesteps (`int`): the number of betas to produce.
|
||||||
|
max_beta (`float`): the maximum beta to use; use values lower than 1 to
|
||||||
|
prevent singularities.
|
||||||
|
alpha_transform_type (`str`, *optional*, default to `cosine`): the type of noise schedule for alpha_bar.
|
||||||
|
Choose from `cosine` or `exp`
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
betas (`np.ndarray`): the betas used by the scheduler to step the model outputs
|
||||||
|
"""
|
||||||
|
if alpha_transform_type == "cosine":
|
||||||
|
|
||||||
|
def alpha_bar_fn(t):
|
||||||
|
return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
|
||||||
|
|
||||||
|
elif alpha_transform_type == "exp":
|
||||||
|
|
||||||
|
def alpha_bar_fn(t):
|
||||||
|
return math.exp(t * -12.0)
|
||||||
|
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unsupported alpha_tranform_type: {alpha_transform_type}")
|
||||||
|
|
||||||
|
betas = []
|
||||||
|
for i in range(num_diffusion_timesteps):
|
||||||
|
t1 = i / num_diffusion_timesteps
|
||||||
|
t2 = (i + 1) / num_diffusion_timesteps
|
||||||
|
betas.append(min(1 - alpha_bar_fn(t2) / alpha_bar_fn(t1), max_beta))
|
||||||
|
return torch.tensor(betas, dtype=torch.float32)
|
||||||
|
|
||||||
|
|
||||||
|
# Copied from diffusers.schedulers.scheduling_ddim.rescale_zero_terminal_snr
|
||||||
|
def rescale_zero_terminal_snr(betas):
|
||||||
|
"""
|
||||||
|
Rescales betas to have zero terminal SNR Based on https://arxiv.org/pdf/2305.08891.pdf (Algorithm 1)
|
||||||
|
|
||||||
|
|
||||||
|
Args:
|
||||||
|
betas (`torch.FloatTensor`):
|
||||||
|
the betas that the scheduler is being initialized with.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
`torch.FloatTensor`: rescaled betas with zero terminal SNR
|
||||||
|
"""
|
||||||
|
# Convert betas to alphas_bar_sqrt
|
||||||
|
alphas = 1.0 - betas
|
||||||
|
alphas_cumprod = torch.cumprod(alphas, dim=0)
|
||||||
|
alphas_bar_sqrt = alphas_cumprod.sqrt()
|
||||||
|
|
||||||
|
# Store old values.
|
||||||
|
alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone()
|
||||||
|
alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone()
|
||||||
|
|
||||||
|
# Shift so the last timestep is zero.
|
||||||
|
alphas_bar_sqrt -= alphas_bar_sqrt_T
|
||||||
|
|
||||||
|
# Scale so the first timestep is back to the old value.
|
||||||
|
alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T)
|
||||||
|
|
||||||
|
# Convert alphas_bar_sqrt to betas
|
||||||
|
alphas_bar = alphas_bar_sqrt**2 # Revert sqrt
|
||||||
|
alphas = alphas_bar[1:] / alphas_bar[:-1] # Revert cumprod
|
||||||
|
alphas = torch.cat([alphas_bar[0:1], alphas])
|
||||||
|
betas = 1 - alphas
|
||||||
|
|
||||||
|
return betas
|
||||||
|
|
||||||
|
|
||||||
|
class EulerDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||||
|
"""
|
||||||
|
Euler scheduler.
|
||||||
|
|
||||||
|
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
|
||||||
|
methods the library implements for all schedulers such as loading and saving.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
num_train_timesteps (`int`, defaults to 1000):
|
||||||
|
The number of diffusion steps to train the model.
|
||||||
|
beta_start (`float`, defaults to 0.0001):
|
||||||
|
The starting `beta` value of inference.
|
||||||
|
beta_end (`float`, defaults to 0.02):
|
||||||
|
The final `beta` value.
|
||||||
|
beta_schedule (`str`, defaults to `"linear"`):
|
||||||
|
The beta schedule, a mapping from a beta range to a sequence of betas for stepping the model. Choose from
|
||||||
|
`linear` or `scaled_linear`.
|
||||||
|
trained_betas (`np.ndarray`, *optional*):
|
||||||
|
Pass an array of betas directly to the constructor to bypass `beta_start` and `beta_end`.
|
||||||
|
prediction_type (`str`, defaults to `epsilon`, *optional*):
|
||||||
|
Prediction type of the scheduler function; can be `epsilon` (predicts the noise of the diffusion process),
|
||||||
|
`sample` (directly predicts the noisy sample`) or `v_prediction` (see section 2.4 of [Imagen
|
||||||
|
Video](https://imagen.research.google/video/paper.pdf) paper).
|
||||||
|
interpolation_type(`str`, defaults to `"linear"`, *optional*):
|
||||||
|
The interpolation type to compute intermediate sigmas for the scheduler denoising steps. Should be on of
|
||||||
|
`"linear"` or `"log_linear"`.
|
||||||
|
use_karras_sigmas (`bool`, *optional*, defaults to `False`):
|
||||||
|
Whether to use Karras sigmas for step sizes in the noise schedule during the sampling process. If `True`,
|
||||||
|
the sigmas are determined according to a sequence of noise levels {σi}.
|
||||||
|
timestep_spacing (`str`, defaults to `"linspace"`):
|
||||||
|
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
|
||||||
|
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
|
||||||
|
steps_offset (`int`, defaults to 0):
|
||||||
|
An offset added to the inference steps. You can use a combination of `offset=1` and
|
||||||
|
`set_alpha_to_one=False` to make the last step use step 0 for the previous alpha product like in Stable
|
||||||
|
Diffusion.
|
||||||
|
rescale_betas_zero_snr (`bool`, defaults to `False`):
|
||||||
|
Whether to rescale the betas to have zero terminal SNR. This enables the model to generate very bright and
|
||||||
|
dark samples instead of limiting it to samples with medium brightness. Loosely related to
|
||||||
|
[`--offset_noise`](https://github.com/huggingface/diffusers/blob/74fd735eb073eb1d774b1ab4154a0876eb82f055/examples/dreambooth/train_dreambooth.py#L506).
|
||||||
|
"""
|
||||||
|
|
||||||
|
_compatibles = [e.name for e in KarrasDiffusionSchedulers]
|
||||||
|
order = 1
|
||||||
|
|
||||||
|
@register_to_config
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
num_train_timesteps: int = 1000,
|
||||||
|
beta_start: float = 0.0001,
|
||||||
|
beta_end: float = 0.02,
|
||||||
|
beta_schedule: str = "linear",
|
||||||
|
trained_betas: Optional[Union[np.ndarray, List[float]]] = None,
|
||||||
|
prediction_type: str = "epsilon",
|
||||||
|
interpolation_type: str = "linear",
|
||||||
|
use_karras_sigmas: Optional[bool] = False,
|
||||||
|
sigma_min: Optional[float] = None,
|
||||||
|
sigma_max: Optional[float] = None,
|
||||||
|
timestep_spacing: str = "linspace",
|
||||||
|
timestep_type: str = "discrete", # can be "discrete" or "continuous"
|
||||||
|
steps_offset: int = 0,
|
||||||
|
rescale_betas_zero_snr: bool = False,
|
||||||
|
):
|
||||||
|
if trained_betas is not None:
|
||||||
|
self.betas = torch.tensor(trained_betas, dtype=torch.float32)
|
||||||
|
elif beta_schedule == "linear":
|
||||||
|
self.betas = torch.linspace(beta_start, beta_end, num_train_timesteps, dtype=torch.float32)
|
||||||
|
elif beta_schedule == "scaled_linear":
|
||||||
|
# this schedule is very specific to the latent diffusion model.
|
||||||
|
self.betas = torch.linspace(beta_start**0.5, beta_end**0.5, num_train_timesteps, dtype=torch.float32) ** 2
|
||||||
|
elif beta_schedule == "squaredcos_cap_v2":
|
||||||
|
# Glide cosine schedule
|
||||||
|
self.betas = betas_for_alpha_bar(num_train_timesteps)
|
||||||
|
else:
|
||||||
|
raise NotImplementedError(f"{beta_schedule} does is not implemented for {self.__class__}")
|
||||||
|
|
||||||
|
if rescale_betas_zero_snr:
|
||||||
|
self.betas = rescale_zero_terminal_snr(self.betas)
|
||||||
|
|
||||||
|
self.alphas = 1.0 - self.betas
|
||||||
|
self.alphas_cumprod = torch.cumprod(self.alphas, dim=0)
|
||||||
|
|
||||||
|
if rescale_betas_zero_snr:
|
||||||
|
# Close to 0 without being 0 so first sigma is not inf
|
||||||
|
# FP16 smallest positive subnormal works well here
|
||||||
|
self.alphas_cumprod[-1] = 2**-24
|
||||||
|
|
||||||
|
sigmas = np.array(((1 - self.alphas_cumprod) / self.alphas_cumprod) ** 0.5)
|
||||||
|
timesteps = np.linspace(0, num_train_timesteps - 1, num_train_timesteps, dtype=float)[::-1].copy()
|
||||||
|
|
||||||
|
sigmas = sigmas[::-1].copy()
|
||||||
|
|
||||||
|
if self.use_karras_sigmas:
|
||||||
|
log_sigmas = np.log(sigmas)
|
||||||
|
sigmas = self._convert_to_karras(in_sigmas=sigmas, num_inference_steps=num_train_timesteps)
|
||||||
|
timesteps = np.array([self._sigma_to_t(sigma, log_sigmas) for sigma in sigmas])
|
||||||
|
|
||||||
|
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32)
|
||||||
|
|
||||||
|
# setable values
|
||||||
|
self.num_inference_steps = None
|
||||||
|
|
||||||
|
# TODO: Support the full EDM scalings for all prediction types and timestep types
|
||||||
|
if timestep_type == "continuous" and prediction_type == "v_prediction":
|
||||||
|
self.timesteps = torch.Tensor([0.25 * sigma.log() for sigma in sigmas])
|
||||||
|
else:
|
||||||
|
self.timesteps = torch.from_numpy(timesteps.astype(np.float32))
|
||||||
|
|
||||||
|
self.sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
|
||||||
|
|
||||||
|
self.is_scale_input_called = False
|
||||||
|
self.use_karras_sigmas = use_karras_sigmas
|
||||||
|
|
||||||
|
self._step_index = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def init_noise_sigma(self):
|
||||||
|
# standard deviation of the initial noise distribution
|
||||||
|
max_sigma = max(self.sigmas) if isinstance(self.sigmas, list) else self.sigmas.max()
|
||||||
|
if self.config.timestep_spacing in ["linspace", "trailing"]:
|
||||||
|
return max_sigma
|
||||||
|
|
||||||
|
return (max_sigma**2 + 1) ** 0.5
|
||||||
|
|
||||||
|
@property
|
||||||
|
def step_index(self):
|
||||||
|
"""
|
||||||
|
The index counter for current timestep. It will increae 1 after each scheduler step.
|
||||||
|
"""
|
||||||
|
return self._step_index
|
||||||
|
|
||||||
|
def scale_model_input(
|
||||||
|
self, sample: torch.FloatTensor, timestep: Union[float, torch.FloatTensor]
|
||||||
|
) -> torch.FloatTensor:
|
||||||
|
"""
|
||||||
|
Ensures interchangeability with schedulers that need to scale the denoising model input depending on the
|
||||||
|
current timestep. Scales the denoising model input by `(sigma**2 + 1) ** 0.5` to match the Euler algorithm.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
sample (`torch.FloatTensor`):
|
||||||
|
The input sample.
|
||||||
|
timestep (`int`, *optional*):
|
||||||
|
The current timestep in the diffusion chain.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
`torch.FloatTensor`:
|
||||||
|
A scaled input sample.
|
||||||
|
"""
|
||||||
|
if self.step_index is None:
|
||||||
|
self._init_step_index(timestep)
|
||||||
|
|
||||||
|
sigma = self.sigmas[self.step_index]
|
||||||
|
sample = sample / ((sigma**2 + 1) ** 0.5)
|
||||||
|
|
||||||
|
self.is_scale_input_called = True
|
||||||
|
return sample
|
||||||
|
|
||||||
|
def set_timesteps(self, num_inference_steps: int, device: Union[str, torch.device] = None):
|
||||||
|
"""
|
||||||
|
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
num_inference_steps (`int`):
|
||||||
|
The number of diffusion steps used when generating samples with a pre-trained model.
|
||||||
|
device (`str` or `torch.device`, *optional*):
|
||||||
|
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||||
|
"""
|
||||||
|
self.num_inference_steps = num_inference_steps
|
||||||
|
|
||||||
|
# "linspace", "leading", "trailing" corresponds to annotation of Table 2. of https://arxiv.org/abs/2305.08891
|
||||||
|
if self.config.timestep_spacing == "linspace":
|
||||||
|
timesteps = np.linspace(0, self.config.num_train_timesteps - 1, num_inference_steps, dtype=np.float32)[
|
||||||
|
::-1
|
||||||
|
].copy()
|
||||||
|
elif self.config.timestep_spacing == "leading":
|
||||||
|
step_ratio = self.config.num_train_timesteps // self.num_inference_steps
|
||||||
|
# creates integer timesteps by multiplying by ratio
|
||||||
|
# casting to int to avoid issues when num_inference_step is power of 3
|
||||||
|
timesteps = (np.arange(0, num_inference_steps) * step_ratio).round()[::-1].copy().astype(np.float32)
|
||||||
|
timesteps += self.config.steps_offset
|
||||||
|
elif self.config.timestep_spacing == "trailing":
|
||||||
|
step_ratio = self.config.num_train_timesteps / self.num_inference_steps
|
||||||
|
# creates integer timesteps by multiplying by ratio
|
||||||
|
# casting to int to avoid issues when num_inference_step is power of 3
|
||||||
|
timesteps = (np.arange(self.config.num_train_timesteps, 0, -step_ratio)).round().copy().astype(np.float32)
|
||||||
|
timesteps -= 1
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
f"{self.config.timestep_spacing} is not supported. Please make sure to choose one of 'linspace', 'leading' or 'trailing'."
|
||||||
|
)
|
||||||
|
|
||||||
|
sigmas = np.array(((1 - self.alphas_cumprod) / self.alphas_cumprod) ** 0.5)
|
||||||
|
log_sigmas = np.log(sigmas)
|
||||||
|
|
||||||
|
if self.config.interpolation_type == "linear":
|
||||||
|
sigmas = np.interp(timesteps, np.arange(0, len(sigmas)), sigmas)
|
||||||
|
elif self.config.interpolation_type == "log_linear":
|
||||||
|
sigmas = torch.linspace(np.log(sigmas[-1]), np.log(sigmas[0]), num_inference_steps + 1).exp().numpy()
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
f"{self.config.interpolation_type} is not implemented. Please specify interpolation_type to either"
|
||||||
|
" 'linear' or 'log_linear'"
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.use_karras_sigmas:
|
||||||
|
sigmas = self._convert_to_karras(in_sigmas=sigmas, num_inference_steps=self.num_inference_steps)
|
||||||
|
timesteps = np.array([self._sigma_to_t(sigma, log_sigmas) for sigma in sigmas])
|
||||||
|
|
||||||
|
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32, device=device)
|
||||||
|
|
||||||
|
# TODO: Support the full EDM scalings for all prediction types and timestep types
|
||||||
|
if self.config.timestep_type == "continuous" and self.config.prediction_type == "v_prediction":
|
||||||
|
self.timesteps = torch.Tensor([0.25 * sigma.log() for sigma in sigmas]).to(device=device)
|
||||||
|
else:
|
||||||
|
self.timesteps = torch.from_numpy(timesteps.astype(np.float32)).to(device=device)
|
||||||
|
|
||||||
|
self.sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
|
||||||
|
self._step_index = None
|
||||||
|
|
||||||
|
def _sigma_to_t(self, sigma, log_sigmas):
|
||||||
|
# get log sigma
|
||||||
|
log_sigma = np.log(np.maximum(sigma, 1e-10))
|
||||||
|
|
||||||
|
# get distribution
|
||||||
|
dists = log_sigma - log_sigmas[:, np.newaxis]
|
||||||
|
|
||||||
|
# get sigmas range
|
||||||
|
low_idx = np.cumsum((dists >= 0), axis=0).argmax(axis=0).clip(max=log_sigmas.shape[0] - 2)
|
||||||
|
high_idx = low_idx + 1
|
||||||
|
|
||||||
|
low = log_sigmas[low_idx]
|
||||||
|
high = log_sigmas[high_idx]
|
||||||
|
|
||||||
|
# interpolate sigmas
|
||||||
|
w = (low - log_sigma) / (low - high)
|
||||||
|
w = np.clip(w, 0, 1)
|
||||||
|
|
||||||
|
# transform interpolation to time range
|
||||||
|
t = (1 - w) * low_idx + w * high_idx
|
||||||
|
t = t.reshape(sigma.shape)
|
||||||
|
return t
|
||||||
|
|
||||||
|
# Copied from https://github.com/crowsonkb/k-diffusion/blob/686dbad0f39640ea25c8a8c6a6e56bb40eacefa2/k_diffusion/sampling.py#L17
|
||||||
|
def _convert_to_karras(self, in_sigmas: torch.FloatTensor, num_inference_steps) -> torch.FloatTensor:
|
||||||
|
"""Constructs the noise schedule of Karras et al. (2022)."""
|
||||||
|
|
||||||
|
# Hack to make sure that other schedulers which copy this function don't break
|
||||||
|
# TODO: Add this logic to the other schedulers
|
||||||
|
if hasattr(self.config, "sigma_min"):
|
||||||
|
sigma_min = self.config.sigma_min
|
||||||
|
else:
|
||||||
|
sigma_min = None
|
||||||
|
|
||||||
|
if hasattr(self.config, "sigma_max"):
|
||||||
|
sigma_max = self.config.sigma_max
|
||||||
|
else:
|
||||||
|
sigma_max = None
|
||||||
|
|
||||||
|
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
|
||||||
|
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
|
||||||
|
|
||||||
|
rho = 7.0 # 7.0 is the value used in the paper
|
||||||
|
ramp = np.linspace(0, 1, num_inference_steps)
|
||||||
|
min_inv_rho = sigma_min ** (1 / rho)
|
||||||
|
max_inv_rho = sigma_max ** (1 / rho)
|
||||||
|
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho
|
||||||
|
return sigmas
|
||||||
|
|
||||||
|
def _init_step_index(self, timestep):
|
||||||
|
if isinstance(timestep, torch.Tensor):
|
||||||
|
timestep = timestep.to(self.timesteps.device)
|
||||||
|
|
||||||
|
index_candidates = (self.timesteps == timestep).nonzero()
|
||||||
|
|
||||||
|
# The sigma index that is taken for the **very** first `step`
|
||||||
|
# is always the second index (or the last index if there is only 1)
|
||||||
|
# This way we can ensure we don't accidentally skip a sigma in
|
||||||
|
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
|
||||||
|
if len(index_candidates) > 1:
|
||||||
|
step_index = index_candidates[1]
|
||||||
|
else:
|
||||||
|
step_index = index_candidates[0]
|
||||||
|
|
||||||
|
self._step_index = step_index.item()
|
||||||
|
|
||||||
|
def step(
|
||||||
|
self,
|
||||||
|
model_output: torch.FloatTensor,
|
||||||
|
timestep: Union[float, torch.FloatTensor],
|
||||||
|
sample: torch.FloatTensor,
|
||||||
|
s_churn: float = 0.0,
|
||||||
|
s_tmin: float = 0.0,
|
||||||
|
s_tmax: float = float("inf"),
|
||||||
|
s_noise: float = 1.0,
|
||||||
|
generator: Optional[torch.Generator] = None,
|
||||||
|
return_dict: bool = True,
|
||||||
|
) -> Union[EulerDiscreteSchedulerOutput, Tuple]:
|
||||||
|
"""
|
||||||
|
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
|
||||||
|
process from the learned model outputs (most often the predicted noise).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_output (`torch.FloatTensor`):
|
||||||
|
The direct output from learned diffusion model.
|
||||||
|
timestep (`float`):
|
||||||
|
The current discrete timestep in the diffusion chain.
|
||||||
|
sample (`torch.FloatTensor`):
|
||||||
|
A current instance of a sample created by the diffusion process.
|
||||||
|
s_churn (`float`):
|
||||||
|
s_tmin (`float`):
|
||||||
|
s_tmax (`float`):
|
||||||
|
s_noise (`float`, defaults to 1.0):
|
||||||
|
Scaling factor for noise added to the sample.
|
||||||
|
generator (`torch.Generator`, *optional*):
|
||||||
|
A random number generator.
|
||||||
|
return_dict (`bool`):
|
||||||
|
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
|
||||||
|
tuple.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
|
||||||
|
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
|
||||||
|
returned, otherwise a tuple is returned where the first element is the sample tensor.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if (
|
||||||
|
isinstance(timestep, int)
|
||||||
|
or isinstance(timestep, torch.IntTensor)
|
||||||
|
or isinstance(timestep, torch.LongTensor)
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
(
|
||||||
|
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||||
|
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||||
|
" one of the `scheduler.timesteps` as a timestep."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
if not self.is_scale_input_called:
|
||||||
|
logger.warning(
|
||||||
|
"The `scale_model_input` function should be called before `step` to ensure correct denoising. "
|
||||||
|
"See `StableDiffusionPipeline` for a usage example."
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.step_index is None:
|
||||||
|
self._init_step_index(timestep)
|
||||||
|
|
||||||
|
# Upcast to avoid precision issues when computing prev_sample
|
||||||
|
sample = sample.to(torch.float32)
|
||||||
|
|
||||||
|
sigma = self.sigmas[self.step_index]
|
||||||
|
|
||||||
|
gamma = min(s_churn / (len(self.sigmas) - 1), 2**0.5 - 1) if s_tmin <= sigma <= s_tmax else 0.0
|
||||||
|
|
||||||
|
noise = randn_tensor(
|
||||||
|
model_output.shape, dtype=model_output.dtype, device=model_output.device, generator=generator
|
||||||
|
)
|
||||||
|
|
||||||
|
eps = noise * s_noise
|
||||||
|
sigma_hat = sigma * (gamma + 1)
|
||||||
|
|
||||||
|
if gamma > 0:
|
||||||
|
sample = sample + eps * (sigma_hat**2 - sigma**2) ** 0.5
|
||||||
|
|
||||||
|
# 1. compute predicted original sample (x_0) from sigma-scaled predicted noise
|
||||||
|
# NOTE: "original_sample" should not be an expected prediction_type but is left in for
|
||||||
|
# backwards compatibility
|
||||||
|
if self.config.prediction_type == "original_sample" or self.config.prediction_type == "sample":
|
||||||
|
pred_original_sample = model_output
|
||||||
|
elif self.config.prediction_type == "epsilon":
|
||||||
|
pred_original_sample = sample - sigma_hat * model_output
|
||||||
|
elif self.config.prediction_type == "v_prediction":
|
||||||
|
# denoised = model_output * c_out + input * c_skip
|
||||||
|
pred_original_sample = model_output * (-sigma / (sigma**2 + 1) ** 0.5) + (sample / (sigma**2 + 1))
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, or `v_prediction`"
|
||||||
|
)
|
||||||
|
|
||||||
|
# 2. Convert to an ODE derivative
|
||||||
|
derivative = (sample - pred_original_sample) / sigma_hat
|
||||||
|
|
||||||
|
dt = self.sigmas[self.step_index + 1] - sigma_hat
|
||||||
|
|
||||||
|
prev_sample = sample + derivative * dt
|
||||||
|
|
||||||
|
# Cast sample back to model compatible dtype
|
||||||
|
prev_sample = prev_sample.to(model_output.dtype)
|
||||||
|
|
||||||
|
# upon completion increase step index by one
|
||||||
|
self._step_index += 1
|
||||||
|
|
||||||
|
if not return_dict:
|
||||||
|
return (prev_sample,)
|
||||||
|
|
||||||
|
return EulerDiscreteSchedulerOutput(prev_sample=prev_sample, pred_original_sample=pred_original_sample)
|
||||||
|
|
||||||
|
def add_noise(
|
||||||
|
self,
|
||||||
|
original_samples: torch.FloatTensor,
|
||||||
|
noise: torch.FloatTensor,
|
||||||
|
timesteps: torch.FloatTensor,
|
||||||
|
) -> torch.FloatTensor:
|
||||||
|
# Make sure sigmas and timesteps have the same device and dtype as original_samples
|
||||||
|
sigmas = self.sigmas.to(device=original_samples.device, dtype=original_samples.dtype)
|
||||||
|
if original_samples.device.type == "mps" and torch.is_floating_point(timesteps):
|
||||||
|
# mps does not support float64
|
||||||
|
schedule_timesteps = self.timesteps.to(original_samples.device, dtype=torch.float32)
|
||||||
|
timesteps = timesteps.to(original_samples.device, dtype=torch.float32)
|
||||||
|
else:
|
||||||
|
schedule_timesteps = self.timesteps.to(original_samples.device)
|
||||||
|
timesteps = timesteps.to(original_samples.device)
|
||||||
|
|
||||||
|
step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps]
|
||||||
|
|
||||||
|
sigma = sigmas[step_indices].flatten()
|
||||||
|
while len(sigma.shape) < len(original_samples.shape):
|
||||||
|
sigma = sigma.unsqueeze(-1)
|
||||||
|
|
||||||
|
noisy_samples = original_samples + noise * sigma
|
||||||
|
return noisy_samples
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return self.config.num_train_timesteps
|
||||||
Reference in New Issue
Block a user