Initial commit
This commit is contained in:
@@ -0,0 +1,2 @@
|
||||
.DS_Store
|
||||
.DS_Store ?
|
||||
@@ -0,0 +1,191 @@
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
|
||||
1. Definitions.
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to the Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no theory of
|
||||
liability, whether in contract, strict liability, or tort
|
||||
(including negligence or otherwise) arising in any way out of
|
||||
the use or inability to use the Work (even if such Holder or other
|
||||
party has been advised of the possibility of such damages), shall
|
||||
any Contributor be liable to You for damages, including any direct,
|
||||
indirect, special, incidental, or consequential damages of any
|
||||
character arising as a result of this License or out of the use or
|
||||
inability to use the Work (including but not limited to damages for
|
||||
loss of goodwill, work stoppage, computer failure or malfunction, or
|
||||
any and all other commercial damages or losses), even if such
|
||||
Contributor has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
Copyright 2025 FASHN AI
|
||||
|
||||
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.
|
||||
@@ -0,0 +1,74 @@
|
||||
# ComfyUI FASHN VTON v1.5 Custom Nodes
|
||||
|
||||
This custom node set implements the [FASHN VTON v1.5](https://github.com/fashn-AI/fashn-vton-1.5) model for virtual try-on in ComfyUI.
|
||||
|
||||
## Installation
|
||||
|
||||
1. Copy the `ComfyUI-FASHN-VTON` folder into your ComfyUI `custom_nodes` directory.
|
||||
2. Install the required dependencies:
|
||||
```bash
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
Note: If you are using a portable version of ComfyUI, use the corresponding python executable.
|
||||
|
||||
## How to Use
|
||||
|
||||
### 1. (Down)load FASHN VTON
|
||||
|
||||
Use the **(Down)load FASHN VTON** node.
|
||||
|
||||
The model **downloads automatically on first use**. In most cases, **no manual download is required**.
|
||||
|
||||
When this node runs, it will:
|
||||
- Download the **FASHN VTON v1.5** model weights
|
||||
- Download the required **DWPose** models
|
||||
- Store everything under `ComfyUI/models/fashn-vton/`
|
||||
- Load the pipeline automatically
|
||||
|
||||
If the files already exist locally, the download step is skipped.
|
||||
|
||||
#### Automatic Downloads
|
||||
|
||||
The following files are fetched automatically from Hugging Face:
|
||||
|
||||
- **FASHN VTON model**
|
||||
- `model.safetensors`
|
||||
- Source: https://huggingface.co/fashn-ai/fashn-vton-1.5
|
||||
|
||||
- **DWPose models**
|
||||
- `yolox_l.onnx`
|
||||
- `dw-ll_ucoco_384.onnx`
|
||||
- Source: https://huggingface.co/fashn-ai/DWPose
|
||||
|
||||
#### Manual Download (Optional)
|
||||
|
||||
If you prefer to download the models manually (e.g. for offline use), place the files in the following directory structure:
|
||||
|
||||
```
|
||||
ComfyUI/
|
||||
└── models/
|
||||
└── fashn-vton/
|
||||
├── model.safetensors
|
||||
└── dwpose/
|
||||
├── yolox_l.onnx
|
||||
└── dw-ll_ucoco_384.onnx
|
||||
```
|
||||
|
||||
### 2. Inference
|
||||
|
||||
Use the **FASHN VTON Inference** node:
|
||||
- **pipeline**: Connect from the Loader node.
|
||||
- **person_image**: The image of the person.
|
||||
- **garment_image**: The image of the garment.
|
||||
- **category**: `tops`, `bottoms`, or `one-pieces`.
|
||||
- **num_timesteps**: Recommended 30–50.
|
||||
- **guidance_scale**: Recommended 1.5–3.0.
|
||||
- **keep_model_loaded**: If set to `false`, the model will be moved to CPU after each inference to save VRAM.
|
||||
|
||||
## Progress Bar
|
||||
|
||||
The inference node supports the ComfyUI native progress bar to show the status of the sampling process.
|
||||
|
||||
## Credits
|
||||
|
||||
Model by [FASHN AI](https://fashn.ai/). Implementation based on their open-source repository.
|
||||
@@ -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,13 @@
|
||||
"""FASHN VTON v1.5"""
|
||||
|
||||
__version__ = "1.5.0"
|
||||
|
||||
from .pipeline import PipelineOutput, TryOnPipeline
|
||||
from .tryon_mmdit import TryOnModel
|
||||
|
||||
__all__ = [
|
||||
"TryOnPipeline",
|
||||
"PipelineOutput",
|
||||
"TryOnModel",
|
||||
"__version__",
|
||||
]
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,16 @@
|
||||
"""
|
||||
DWPose - Effective Whole-body Pose Estimation
|
||||
|
||||
This module is adapted from the IDEA-Research/DWPose repository (onnx branch):
|
||||
https://github.com/IDEA-Research/DWPose/tree/onnx/ControlNet-v1-1-nightly/annotator/dwpose
|
||||
|
||||
Original paper:
|
||||
"Effective Whole-body Pose Estimation with Two-stages Distillation"
|
||||
Zhendong Yang, Ailing Zeng, Chun Yuan, Yu Li
|
||||
ICCV 2023, CV4Metaverse Workshop
|
||||
https://arxiv.org/abs/2307.15880
|
||||
|
||||
License: Apache-2.0
|
||||
"""
|
||||
|
||||
from .dwpose import DWposeDetector, draw_pose
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,131 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from .utils import (
|
||||
draw_bodypose,
|
||||
draw_bodypose_gray,
|
||||
draw_facepose,
|
||||
draw_facepose_gray,
|
||||
draw_handpose,
|
||||
draw_handpose_gray,
|
||||
)
|
||||
from .wholebody import Wholebody
|
||||
|
||||
__all__ = ["DWposeDetector", "draw_pose"]
|
||||
|
||||
# Minimum confidence threshold for keypoint visibility
|
||||
KEYPOINT_VISIBILITY_THRESHOLD = 0.3
|
||||
|
||||
def draw_pose(pose, H, W, canvas_value: int = 0, grayscale: bool = False):
|
||||
bodies = pose["bodies"]
|
||||
candidate = bodies["candidate"]
|
||||
subset = bodies["subset"]
|
||||
|
||||
if grayscale:
|
||||
draw_bodypose_fn = draw_bodypose_gray
|
||||
draw_handpose_fn = draw_handpose_gray
|
||||
draw_facepose_fn = draw_facepose_gray
|
||||
canvas = np.full((H, W), canvas_value, dtype=np.uint8)
|
||||
else:
|
||||
draw_bodypose_fn = draw_bodypose
|
||||
draw_handpose_fn = draw_handpose
|
||||
draw_facepose_fn = draw_facepose
|
||||
canvas_value = int(canvas_value / 0.6) if canvas_value > 0 else 0
|
||||
canvas = np.full((H, W, 3), canvas_value, dtype=np.uint8)
|
||||
|
||||
canvas = draw_bodypose_fn(canvas, candidate, subset)
|
||||
|
||||
if "hands" in pose:
|
||||
canvas = draw_handpose_fn(canvas, pose.get("hands"))
|
||||
if "faces" in pose:
|
||||
canvas = draw_facepose_fn(canvas, pose.get("faces"))
|
||||
|
||||
return canvas
|
||||
|
||||
|
||||
class DWposeDetector:
|
||||
def __init__(
|
||||
self,
|
||||
checkpoints_dir,
|
||||
device="cuda:0",
|
||||
):
|
||||
|
||||
self.pose_estimation = Wholebody(checkpoints_dir=checkpoints_dir, device=device)
|
||||
|
||||
def _find_best_candidate(self, subset, candidate, score_threshold=KEYPOINT_VISIBILITY_THRESHOLD):
|
||||
# Apply score threshold to subset keypoints
|
||||
valid_keypoints = subset[:, 1:14] > score_threshold
|
||||
|
||||
# Calculate scores for each candidate, only counting valid keypoints
|
||||
headless_scores = np.sum(subset[:, 1:14] * valid_keypoints, axis=1)
|
||||
|
||||
# Extract keypoints for each candidate, excluding the head
|
||||
headless_keypoints = candidate[:, 1:14]
|
||||
|
||||
def compute_area(keypoints):
|
||||
# Filter keypoints based on the valid_keypoints mask
|
||||
valid_kp = keypoints[
|
||||
valid_keypoints[0]
|
||||
] # Assuming all candidates have the same validity mask for simplicity
|
||||
valid_x, valid_y = valid_kp[:, 0][valid_kp[:, 0] > 0], valid_kp[:, 1][valid_kp[:, 1] > 0]
|
||||
if not len(valid_x) or not len(valid_y):
|
||||
return 0
|
||||
return (np.max(valid_x) - np.min(valid_x)) * (np.max(valid_y) - np.min(valid_y))
|
||||
|
||||
areas = [compute_area(kp) for kp in headless_keypoints]
|
||||
|
||||
# Here, we multiply scores by areas, but we need to handle division by zero or invalid calculations
|
||||
with np.errstate(divide="ignore", invalid="ignore"):
|
||||
scores_times_areas = headless_scores * np.array(areas)
|
||||
|
||||
# Replace NaN or inf with 0 for np.nanargmax to work correctly
|
||||
scores_times_areas[np.isnan(scores_times_areas) | np.isinf(scores_times_areas)] = 0
|
||||
|
||||
# If all scores are zero (or invalid), we might want to handle this case differently
|
||||
if np.all(scores_times_areas == 0):
|
||||
best_candidate_idx = np.argmax(headless_scores)
|
||||
else:
|
||||
best_candidate_idx = np.nanargmax(scores_times_areas)
|
||||
|
||||
return (
|
||||
candidate[best_candidate_idx : best_candidate_idx + 1],
|
||||
subset[best_candidate_idx : best_candidate_idx + 1],
|
||||
)
|
||||
|
||||
@torch.inference_mode()
|
||||
def __call__(self, oriImg: np.array, single: bool = True) -> dict:
|
||||
oriImg = oriImg.copy()
|
||||
H, W, C = oriImg.shape
|
||||
|
||||
candidate, subset = self.pose_estimation(oriImg)
|
||||
nums, keys, locs = candidate.shape
|
||||
|
||||
if single and nums > 1:
|
||||
candidate, subset = self._find_best_candidate(subset, candidate)
|
||||
nums = 1 # Now we only have one candidate
|
||||
|
||||
candidate[..., 0] /= float(W)
|
||||
candidate[..., 1] /= float(H)
|
||||
|
||||
body = candidate[:, :18].copy()
|
||||
body = body.reshape(nums * 18, locs)
|
||||
score = subset[:, :18]
|
||||
for i in range(len(score)):
|
||||
for j in range(len(score[i])):
|
||||
if score[i][j] > KEYPOINT_VISIBILITY_THRESHOLD:
|
||||
score[i][j] = int(18 * i + j)
|
||||
else:
|
||||
score[i][j] = -1
|
||||
|
||||
un_visible = subset < KEYPOINT_VISIBILITY_THRESHOLD
|
||||
candidate[un_visible] = -1
|
||||
|
||||
foot = candidate[:, 18:24]
|
||||
faces = candidate[:, 24:92]
|
||||
hands = candidate[:, 92:113]
|
||||
hands = np.vstack([hands, candidate[:, 113:]])
|
||||
|
||||
bodies = dict(candidate=body, subset=score)
|
||||
pose = dict(bodies=bodies, hands=hands, faces=faces)
|
||||
|
||||
return pose
|
||||
@@ -0,0 +1,131 @@
|
||||
"""
|
||||
Detection utilities adapted from YOLOX (Apache-2.0):
|
||||
https://github.com/Megvii-BaseDetection/YOLOX
|
||||
"""
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
def nms(boxes, scores, nms_thr):
|
||||
"""Single class NMS implemented in Numpy."""
|
||||
x1 = boxes[:, 0]
|
||||
y1 = boxes[:, 1]
|
||||
x2 = boxes[:, 2]
|
||||
y2 = boxes[:, 3]
|
||||
|
||||
areas = (x2 - x1 + 1) * (y2 - y1 + 1)
|
||||
order = scores.argsort()[::-1]
|
||||
|
||||
keep = []
|
||||
while order.size > 0:
|
||||
i = order[0]
|
||||
keep.append(i)
|
||||
xx1 = np.maximum(x1[i], x1[order[1:]])
|
||||
yy1 = np.maximum(y1[i], y1[order[1:]])
|
||||
xx2 = np.minimum(x2[i], x2[order[1:]])
|
||||
yy2 = np.minimum(y2[i], y2[order[1:]])
|
||||
|
||||
w = np.maximum(0.0, xx2 - xx1 + 1)
|
||||
h = np.maximum(0.0, yy2 - yy1 + 1)
|
||||
inter = w * h
|
||||
ovr = inter / (areas[i] + areas[order[1:]] - inter)
|
||||
|
||||
inds = np.where(ovr <= nms_thr)[0]
|
||||
order = order[inds + 1]
|
||||
|
||||
return keep
|
||||
|
||||
|
||||
def multiclass_nms(boxes, scores, nms_thr, score_thr):
|
||||
"""Multiclass NMS implemented in Numpy. Class-aware version."""
|
||||
final_dets = []
|
||||
num_classes = scores.shape[1]
|
||||
for cls_ind in range(num_classes):
|
||||
cls_scores = scores[:, cls_ind]
|
||||
valid_score_mask = cls_scores > score_thr
|
||||
if valid_score_mask.sum() == 0:
|
||||
continue
|
||||
else:
|
||||
valid_scores = cls_scores[valid_score_mask]
|
||||
valid_boxes = boxes[valid_score_mask]
|
||||
keep = nms(valid_boxes, valid_scores, nms_thr)
|
||||
if len(keep) > 0:
|
||||
cls_inds = np.ones((len(keep), 1)) * cls_ind
|
||||
dets = np.concatenate([valid_boxes[keep], valid_scores[keep, None], cls_inds], 1)
|
||||
final_dets.append(dets)
|
||||
if len(final_dets) == 0:
|
||||
return None
|
||||
return np.concatenate(final_dets, 0)
|
||||
|
||||
|
||||
def demo_postprocess(outputs, img_size, p6=False):
|
||||
grids = []
|
||||
expanded_strides = []
|
||||
strides = [8, 16, 32] if not p6 else [8, 16, 32, 64]
|
||||
|
||||
hsizes = [img_size[0] // stride for stride in strides]
|
||||
wsizes = [img_size[1] // stride for stride in strides]
|
||||
|
||||
for hsize, wsize, stride in zip(hsizes, wsizes, strides):
|
||||
xv, yv = np.meshgrid(np.arange(wsize), np.arange(hsize))
|
||||
grid = np.stack((xv, yv), 2).reshape(1, -1, 2)
|
||||
grids.append(grid)
|
||||
shape = grid.shape[:2]
|
||||
expanded_strides.append(np.full((*shape, 1), stride))
|
||||
|
||||
grids = np.concatenate(grids, 1)
|
||||
expanded_strides = np.concatenate(expanded_strides, 1)
|
||||
outputs[..., :2] = (outputs[..., :2] + grids) * expanded_strides
|
||||
outputs[..., 2:4] = np.exp(outputs[..., 2:4]) * expanded_strides
|
||||
|
||||
return outputs
|
||||
|
||||
|
||||
def preprocess(img, input_size, swap=(2, 0, 1)):
|
||||
if len(img.shape) == 3:
|
||||
padded_img = np.ones((input_size[0], input_size[1], 3), dtype=np.uint8) * 114
|
||||
else:
|
||||
padded_img = np.ones(input_size, dtype=np.uint8) * 114
|
||||
|
||||
r = min(input_size[0] / img.shape[0], input_size[1] / img.shape[1])
|
||||
resized_img = cv2.resize(
|
||||
img,
|
||||
(int(img.shape[1] * r), int(img.shape[0] * r)),
|
||||
interpolation=cv2.INTER_LINEAR,
|
||||
).astype(np.uint8)
|
||||
padded_img[: int(img.shape[0] * r), : int(img.shape[1] * r)] = resized_img
|
||||
|
||||
padded_img = padded_img.transpose(swap)
|
||||
padded_img = np.ascontiguousarray(padded_img, dtype=np.float32)
|
||||
return padded_img, r
|
||||
|
||||
|
||||
def inference_detector(session, oriImg):
|
||||
input_shape = (640, 640)
|
||||
img, ratio = preprocess(oriImg, input_shape)
|
||||
|
||||
ort_inputs = {session.get_inputs()[0].name: img[None, :, :, :]}
|
||||
output = session.run(None, ort_inputs)
|
||||
predictions = demo_postprocess(output[0], input_shape)[0]
|
||||
|
||||
boxes = predictions[:, :4]
|
||||
scores = predictions[:, 4:5] * predictions[:, 5:]
|
||||
|
||||
boxes_xyxy = np.ones_like(boxes)
|
||||
boxes_xyxy[:, 0] = boxes[:, 0] - boxes[:, 2] / 2.0
|
||||
boxes_xyxy[:, 1] = boxes[:, 1] - boxes[:, 3] / 2.0
|
||||
boxes_xyxy[:, 2] = boxes[:, 0] + boxes[:, 2] / 2.0
|
||||
boxes_xyxy[:, 3] = boxes[:, 1] + boxes[:, 3] / 2.0
|
||||
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,364 @@
|
||||
"""
|
||||
Pose estimation adapted from DWPose/MMPose (Apache-2.0):
|
||||
https://github.com/IDEA-Research/DWPose
|
||||
"""
|
||||
|
||||
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.0) -> 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, 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.0, src_w * -0.5]), rot_rad)
|
||||
dst_dir = np.array([0.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.0] = -1
|
||||
|
||||
# reshape
|
||||
locs = locs.reshape(N, K, 2)
|
||||
vals = vals.reshape(N, K)
|
||||
|
||||
return locs, vals
|
||||
|
||||
|
||||
def decode(simcc_x: np.ndarray, simcc_y: np.ndarray, simcc_split_ratio) -> Tuple[np.ndarray, np.ndarray]:
|
||||
"""Modulate simcc distribution with Gaussian.
|
||||
|
||||
Args:
|
||||
simcc_x (np.ndarray[K, Wx]): model predicted simcc in x.
|
||||
simcc_y (np.ndarray[K, Wy]): model predicted simcc in y.
|
||||
simcc_split_ratio (int): The split ratio of simcc.
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing center and scale.
|
||||
- np.ndarray[float32]: keypoints in shape (K, 2) or (n, K, 2)
|
||||
- np.ndarray[float32]: scores in shape (K,) or (n, K)
|
||||
"""
|
||||
keypoints, scores = get_simcc_maximum(simcc_x, simcc_y)
|
||||
keypoints /= simcc_split_ratio
|
||||
|
||||
return keypoints, scores
|
||||
|
||||
|
||||
def inference_pose(session, out_bbox, oriImg):
|
||||
h, w = session.get_inputs()[0].shape[2:]
|
||||
model_input_size = (w, h)
|
||||
resized_img, center, scale = preprocess(oriImg, out_bbox, model_input_size)
|
||||
outputs = inference(session, resized_img)
|
||||
keypoints, scores = postprocess(outputs, model_input_size, center, scale)
|
||||
|
||||
return keypoints, scores
|
||||
@@ -0,0 +1,227 @@
|
||||
"""
|
||||
Drawing utilities adapted from DWPose (Apache-2.0):
|
||||
https://github.com/IDEA-Research/DWPose
|
||||
"""
|
||||
|
||||
import math
|
||||
|
||||
import cv2
|
||||
import matplotlib
|
||||
import numpy as np
|
||||
|
||||
eps = 0.01
|
||||
|
||||
|
||||
def draw_bodypose_gray(canvas, candidate, subset):
|
||||
H, W = canvas.shape
|
||||
candidate = np.array(candidate)
|
||||
subset = np.array(subset)
|
||||
|
||||
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],
|
||||
]
|
||||
|
||||
base_values = np.linspace(20, 240, 18).astype(np.uint8)
|
||||
stickwidth = 4
|
||||
limb_canvas = np.zeros_like(canvas)
|
||||
|
||||
for i in range(17):
|
||||
limb_gray = int(base_values[i])
|
||||
for n in range(len(subset)):
|
||||
indexA = int(subset[n][limbSeq[i][0] - 1])
|
||||
indexB = int(subset[n][limbSeq[i][1] - 1])
|
||||
if indexA == -1 or indexB == -1:
|
||||
continue
|
||||
|
||||
xA = int(candidate[indexA][0] * W)
|
||||
yA = int(candidate[indexA][1] * H)
|
||||
xB = int(candidate[indexB][0] * W)
|
||||
yB = int(candidate[indexB][1] * H)
|
||||
|
||||
mX = (xA + xB) // 2
|
||||
mY = (yA + yB) // 2
|
||||
length = int(math.hypot(xA - xB, yA - yB))
|
||||
angle = int(math.degrees(math.atan2(yA - yB, xA - xB)))
|
||||
|
||||
polygon = cv2.ellipse2Poly((mX, mY), (length // 2, stickwidth), angle, 0, 360, 1)
|
||||
cv2.fillConvexPoly(limb_canvas, polygon, limb_gray)
|
||||
|
||||
canvas = np.maximum(canvas, (limb_canvas * 0.6).astype(np.uint8))
|
||||
|
||||
keypoint_to_limb_map = {}
|
||||
for i, (a, b) in enumerate(limbSeq):
|
||||
if a not in keypoint_to_limb_map:
|
||||
keypoint_to_limb_map[a] = []
|
||||
if b not in keypoint_to_limb_map:
|
||||
keypoint_to_limb_map[b] = []
|
||||
keypoint_to_limb_map[a].append(i)
|
||||
keypoint_to_limb_map[b].append(i)
|
||||
|
||||
for i in range(1, 19):
|
||||
if i not in keypoint_to_limb_map:
|
||||
continue
|
||||
|
||||
connected_limbs = keypoint_to_limb_map[i]
|
||||
if connected_limbs:
|
||||
point_gray = min(255, int(base_values[connected_limbs[0]] * 1.3))
|
||||
else:
|
||||
point_gray = 200
|
||||
|
||||
for n in range(len(subset)):
|
||||
index = int(subset[n][i - 1])
|
||||
if index == -1:
|
||||
continue
|
||||
|
||||
x = int(candidate[index][0] * W)
|
||||
y = int(candidate[index][1] * H)
|
||||
cv2.circle(canvas, (x, y), 4, point_gray, thickness=-1)
|
||||
|
||||
return canvas
|
||||
|
||||
|
||||
def draw_bodypose(canvas, candidate, subset):
|
||||
H, W, C = canvas.shape
|
||||
candidate = np.array(candidate)
|
||||
subset = np.array(subset)
|
||||
|
||||
stickwidth = 4
|
||||
|
||||
limbSeq = [
|
||||
[2, 3], [2, 6], [3, 4], [4, 5], [6, 7], [7, 8], [2, 9], [9, 10],
|
||||
[10, 11], [2, 12], [12, 13], [13, 14], [2, 1], [1, 15], [15, 17],
|
||||
[1, 16], [16, 18], [3, 17], [6, 18],
|
||||
]
|
||||
|
||||
colors = [
|
||||
[255, 0, 0], [255, 85, 0], [255, 170, 0], [255, 255, 0],
|
||||
[170, 255, 0], [85, 255, 0], [0, 255, 0], [0, 255, 85],
|
||||
[0, 255, 170], [0, 255, 255], [0, 170, 255], [0, 85, 255],
|
||||
[0, 0, 255], [85, 0, 255], [170, 0, 255], [255, 0, 255],
|
||||
[255, 0, 170], [255, 0, 85],
|
||||
]
|
||||
|
||||
for i in range(17):
|
||||
for n in range(len(subset)):
|
||||
index = subset[n][np.array(limbSeq[i]) - 1]
|
||||
if -1 in index:
|
||||
continue
|
||||
Y = candidate[index.astype(int), 0] * float(W)
|
||||
X = candidate[index.astype(int), 1] * float(H)
|
||||
mX = np.mean(X)
|
||||
mY = np.mean(Y)
|
||||
length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5
|
||||
angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1]))
|
||||
polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stickwidth), int(angle), 0, 360, 1)
|
||||
cv2.fillConvexPoly(canvas, polygon, colors[i])
|
||||
|
||||
canvas = (canvas * 0.6).astype(np.uint8)
|
||||
|
||||
for i in range(18):
|
||||
for n in range(len(subset)):
|
||||
index = int(subset[n][i])
|
||||
if index == -1:
|
||||
continue
|
||||
x, y = candidate[index][0:2]
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
cv2.circle(canvas, (int(x), int(y)), 4, colors[i], thickness=-1)
|
||||
|
||||
return canvas
|
||||
|
||||
|
||||
def draw_handpose(canvas, all_hand_peaks):
|
||||
H, W, C = canvas.shape
|
||||
|
||||
edges = [
|
||||
[0, 1], [1, 2], [2, 3], [3, 4], [0, 5], [5, 6], [6, 7], [7, 8],
|
||||
[0, 9], [9, 10], [10, 11], [11, 12], [0, 13], [13, 14], [14, 15],
|
||||
[15, 16], [0, 17], [17, 18], [18, 19], [19, 20],
|
||||
]
|
||||
|
||||
for peaks in all_hand_peaks:
|
||||
peaks = np.array(peaks)
|
||||
|
||||
for ie, e in enumerate(edges):
|
||||
x1, y1 = peaks[e[0]]
|
||||
x2, y2 = peaks[e[1]]
|
||||
x1 = int(x1 * W)
|
||||
y1 = int(y1 * H)
|
||||
x2 = int(x2 * W)
|
||||
y2 = int(y2 * H)
|
||||
if x1 > eps and y1 > eps and x2 > eps and y2 > eps:
|
||||
cv2.line(
|
||||
canvas,
|
||||
(x1, y1),
|
||||
(x2, y2),
|
||||
matplotlib.colors.hsv_to_rgb([ie / float(len(edges)), 1.0, 1.0]) * 255,
|
||||
thickness=2,
|
||||
)
|
||||
|
||||
for i, keypoint in enumerate(peaks):
|
||||
x, y = keypoint
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
if x > eps and y > eps:
|
||||
cv2.circle(canvas, (x, y), 4, (0, 0, 255), thickness=-1)
|
||||
return canvas
|
||||
|
||||
|
||||
def draw_facepose(canvas, all_lmks):
|
||||
H, W, C = canvas.shape
|
||||
for lmks in all_lmks:
|
||||
lmks = np.array(lmks)
|
||||
for lmk in lmks:
|
||||
x, y = lmk
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
if x > eps and y > eps:
|
||||
cv2.circle(canvas, (x, y), 3, (255, 255, 255), thickness=-1)
|
||||
return canvas
|
||||
|
||||
|
||||
def draw_handpose_gray(canvas, all_hand_peaks):
|
||||
H, W = 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],
|
||||
]
|
||||
|
||||
edge_values = np.linspace(160, 220, len(edges)).astype(np.uint8)
|
||||
|
||||
for peaks in all_hand_peaks:
|
||||
peaks = np.array(peaks)
|
||||
|
||||
for ie, e in enumerate(edges):
|
||||
x1, y1 = peaks[e[0]]
|
||||
x2, y2 = peaks[e[1]]
|
||||
x1 = int(x1 * W)
|
||||
y1 = int(y1 * H)
|
||||
x2 = int(x2 * W)
|
||||
y2 = int(y2 * H)
|
||||
if x1 > eps and y1 > eps and x2 > eps and y2 > eps:
|
||||
cv2.line(canvas, (x1, y1), (x2, y2), int(edge_values[ie]), thickness=2)
|
||||
|
||||
for i, keypoint in enumerate(peaks):
|
||||
x, y = keypoint
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
if x > eps and y > eps:
|
||||
cv2.circle(canvas, (x, y), 4, 240, thickness=-1)
|
||||
return canvas
|
||||
|
||||
|
||||
def draw_facepose_gray(canvas, all_lmks):
|
||||
H, W = canvas.shape
|
||||
for lmks in all_lmks:
|
||||
lmks = np.array(lmks)
|
||||
for lmk in lmks:
|
||||
x, y = lmk
|
||||
x = int(x * W)
|
||||
y = int(y * H)
|
||||
if x > eps and y > eps:
|
||||
cv2.circle(canvas, (x, y), 3, 200, thickness=-1)
|
||||
return canvas
|
||||
@@ -0,0 +1,51 @@
|
||||
"""
|
||||
Wholebody pose detection adapted from DWPose (Apache-2.0):
|
||||
https://github.com/IDEA-Research/DWPose
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import onnxruntime as ort
|
||||
|
||||
from .onnxdet import inference_detector
|
||||
from .onnxpose import inference_pose
|
||||
|
||||
|
||||
class Wholebody:
|
||||
def __init__(self, checkpoints_dir, device="cuda:0"):
|
||||
if device.startswith("cuda"):
|
||||
device_id = int(device.split(":")[-1])
|
||||
provider_options = [{"device_id": str(device_id)}]
|
||||
providers = ["CUDAExecutionProvider"]
|
||||
else:
|
||||
providers = ["CPUExecutionProvider"]
|
||||
provider_options = None
|
||||
onnx_det = os.path.join(checkpoints_dir, "yolox_l.onnx")
|
||||
onnx_pose = os.path.join(checkpoints_dir, "dw-ll_ucoco_384.onnx")
|
||||
|
||||
self.session_det = ort.InferenceSession(
|
||||
path_or_bytes=onnx_det, providers=providers, provider_options=provider_options
|
||||
)
|
||||
self.session_pose = ort.InferenceSession(
|
||||
path_or_bytes=onnx_pose, providers=providers, provider_options=provider_options
|
||||
)
|
||||
|
||||
def __call__(self, oriImg):
|
||||
det_result = inference_detector(self.session_det, oriImg)
|
||||
keypoints, scores = inference_pose(self.session_pose, det_result, oriImg)
|
||||
|
||||
keypoints_info = np.concatenate((keypoints, scores[..., None]), axis=-1)
|
||||
# compute neck joint
|
||||
neck = np.mean(keypoints_info[:, [5, 6]], axis=1)
|
||||
# neck score when visualizing pred
|
||||
neck[:, 2:4] = np.logical_and(keypoints_info[:, 5, 2:4] > 0.3, keypoints_info[:, 6, 2:4] > 0.3).astype(int)
|
||||
new_keypoints_info = np.insert(keypoints_info, 17, neck, axis=1)
|
||||
mmpose_idx = [17, 6, 8, 10, 7, 9, 12, 14, 16, 13, 15, 2, 1, 4, 3]
|
||||
openpose_idx = [1, 2, 3, 4, 6, 7, 8, 9, 10, 12, 13, 14, 15, 16, 17]
|
||||
new_keypoints_info[:, openpose_idx] = new_keypoints_info[:, mmpose_idx]
|
||||
keypoints_info = new_keypoints_info
|
||||
|
||||
keypoints, scores = keypoints_info[..., :2], keypoints_info[..., 2]
|
||||
|
||||
return keypoints, scores
|
||||
@@ -0,0 +1,343 @@
|
||||
"""TryOn Pipeline."""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Literal, Optional
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
from fashn_human_parser import CATEGORY_TO_BODY_COVERAGE, FashnHumanParser
|
||||
from PIL import Image
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from .dwpose import DWposeDetector, draw_pose
|
||||
from .preprocessing import (
|
||||
BODY_COVERAGE_TO_FASHN_LABELS,
|
||||
FASHN_LABELS_TO_IDS,
|
||||
AspectPreserveResize,
|
||||
ResizePad,
|
||||
create_clothing_agnostic_image,
|
||||
create_garment_image,
|
||||
)
|
||||
from .tryon_mmdit import TryOnModel
|
||||
from .utils import (
|
||||
get_dummy_dw_keypoints,
|
||||
get_rf_schedule,
|
||||
load_checkpoint,
|
||||
normalize_uint8_to_neg1_1,
|
||||
numpy_to_torch,
|
||||
setup_logger,
|
||||
tensor_to_pil,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class PipelineOutput:
|
||||
"""Pipeline output container."""
|
||||
|
||||
images: List[Image.Image]
|
||||
|
||||
|
||||
class TryOnPipeline:
|
||||
"""
|
||||
TryOn inference pipeline.
|
||||
|
||||
Args:
|
||||
weights_dir: Directory containing model weights (model.safetensors, dwpose/)
|
||||
device: Device to run on ('cuda', 'cpu', or None for auto-detect)
|
||||
logger: Optional logger instance
|
||||
|
||||
Example:
|
||||
pipeline = TryOnPipeline(weights_dir="./weights")
|
||||
result = pipeline(person_image, garment_image, category="tops")
|
||||
"""
|
||||
|
||||
CATEGORY_TO_LABEL = {"tops": 1, "bottoms": 2, "one-pieces": 3}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
weights_dir: str,
|
||||
device: Optional[str] = None,
|
||||
logger: Optional[logging.Logger] = None,
|
||||
):
|
||||
self.weights_dir = os.path.abspath(weights_dir)
|
||||
self.logger = logger or setup_logger("TryOnPipeline", level=logging.INFO)
|
||||
|
||||
# Setup device
|
||||
self.device = torch.device(device if device else ("cuda" if torch.cuda.is_available() else "cpu"))
|
||||
self.logger.info(f"Using device: {self.device}")
|
||||
|
||||
# Setup inference dtype
|
||||
self.inference_dtype = torch.float32
|
||||
if self.device.type == "cuda" and torch.cuda.is_bf16_supported():
|
||||
self.inference_dtype = torch.bfloat16
|
||||
self.logger.info(f"Using dtype: {self.inference_dtype}")
|
||||
|
||||
# Validate weights exist
|
||||
self._validate_weights()
|
||||
|
||||
# Load models
|
||||
self._setup_tryon_model()
|
||||
self._setup_pose_model()
|
||||
self._setup_hp_model()
|
||||
|
||||
# Setup transforms (derived from model input shape)
|
||||
h, w = self.tryon_model.input_shape
|
||||
max_dim = max(h, w)
|
||||
self.pre_resize = AspectPreserveResize(target_size=(max_dim, max_dim), mode="fit", backend="pil")
|
||||
self.resize_pad_fn = ResizePad((w, h), backend="opencv")
|
||||
|
||||
def _validate_weights(self):
|
||||
"""Check that required weight files exist."""
|
||||
tryon_path = os.path.join(self.weights_dir, "model.safetensors")
|
||||
dwpose_dir = os.path.join(self.weights_dir, "dwpose")
|
||||
yolox_path = os.path.join(dwpose_dir, "yolox_l.onnx")
|
||||
dwpose_path = os.path.join(dwpose_dir, "dw-ll_ucoco_384.onnx")
|
||||
|
||||
missing = []
|
||||
if not os.path.exists(tryon_path):
|
||||
missing.append(tryon_path)
|
||||
if not os.path.exists(yolox_path):
|
||||
missing.append(yolox_path)
|
||||
if not os.path.exists(dwpose_path):
|
||||
missing.append(dwpose_path)
|
||||
|
||||
if missing:
|
||||
raise FileNotFoundError(
|
||||
"Missing model weights:\n"
|
||||
+ "\n".join(f" - {p}" for p in missing)
|
||||
+ f"\n\nPlease run:\n python scripts/download_weights.py --weights-dir {self.weights_dir}"
|
||||
)
|
||||
|
||||
def _setup_tryon_model(self):
|
||||
"""Load the TryOn model."""
|
||||
model_path = os.path.join(self.weights_dir, "model.safetensors")
|
||||
self.logger.info(f"Loading TryOnModel from {model_path}")
|
||||
|
||||
self.tryon_model = TryOnModel()
|
||||
state_dict = load_checkpoint(model_path, device=str(self.device))
|
||||
self.tryon_model.load_state_dict(state_dict)
|
||||
self.tryon_model.to(self.device, dtype=self.inference_dtype).eval()
|
||||
|
||||
self.logger.info("TryOnModel loaded")
|
||||
|
||||
def _setup_pose_model(self):
|
||||
"""Load DWPose model."""
|
||||
dwpose_dir = os.path.join(self.weights_dir, "dwpose")
|
||||
self.logger.info(f"Loading DWPose from {dwpose_dir}")
|
||||
|
||||
dwpose_device = f"cuda:{self.device.index or 0}" if self.device.type == "cuda" else "cpu"
|
||||
self.pose_model = DWposeDetector(checkpoints_dir=dwpose_dir, device=dwpose_device)
|
||||
|
||||
self.logger.info("DWPose loaded")
|
||||
|
||||
def _setup_hp_model(self):
|
||||
"""Load human parsing model."""
|
||||
self.logger.info("Loading FashnHumanParser")
|
||||
|
||||
hp_device = "cuda" if self.device.type == "cuda" else "cpu"
|
||||
self.hp_model = FashnHumanParser(device=hp_device)
|
||||
|
||||
self.logger.info("FashnHumanParser loaded")
|
||||
|
||||
@torch.inference_mode()
|
||||
def _sample(
|
||||
self,
|
||||
*,
|
||||
ca_images: torch.Tensor,
|
||||
garment_images: torch.Tensor,
|
||||
person_poses: torch.Tensor,
|
||||
garment_poses: torch.Tensor,
|
||||
garment_categories: torch.Tensor,
|
||||
num_timesteps: int = 30,
|
||||
time_shift_mu: float = 1.5,
|
||||
guidance_scale: float = 1.5,
|
||||
skip_cfg_last_n_steps: int = 1,
|
||||
use_tqdm: bool = True,
|
||||
callback: Optional[callable] = None,
|
||||
) -> List[Image.Image]:
|
||||
"""Euler sampling with CFG."""
|
||||
device, dtype = ca_images.device, ca_images.dtype
|
||||
batch_size = ca_images.shape[0]
|
||||
|
||||
# Init noisy images
|
||||
c, h, w = self.tryon_model.channels_in, *self.tryon_model.input_shape
|
||||
images = torch.randn((batch_size, c, h, w), dtype=dtype, device=device)
|
||||
|
||||
# Time schedule (from 0 -> 1)
|
||||
timesteps = get_rf_schedule(num_steps=num_timesteps, mu=time_shift_mu)
|
||||
|
||||
model_kwargs = {
|
||||
"person_poses": person_poses,
|
||||
"garment_poses": garment_poses,
|
||||
"ca_images": ca_images,
|
||||
"garment_images": garment_images,
|
||||
"garment_categories": garment_categories,
|
||||
}
|
||||
|
||||
# Euler sampling loop
|
||||
total_steps = len(timesteps) - 1
|
||||
for step_idx, (t_curr, t_prev) in enumerate(
|
||||
tqdm(
|
||||
zip(timesteps[:-1], timesteps[1:]),
|
||||
desc="Sampling",
|
||||
total=total_steps,
|
||||
disable=not use_tqdm,
|
||||
)
|
||||
):
|
||||
if callback:
|
||||
callback(step_idx, total_steps)
|
||||
dt = t_prev - t_curr
|
||||
t_vec = torch.full((batch_size,), t_curr, dtype=dtype, device=device)
|
||||
|
||||
pred = self.tryon_model.forward_for_cfg(images, t_vec, **model_kwargs)
|
||||
v_c, v_u = pred["v_c"], pred["v_u"]
|
||||
|
||||
# Skip CFG at final steps to prevent color saturation
|
||||
if skip_cfg_last_n_steps > 0 and step_idx >= num_timesteps - skip_cfg_last_n_steps:
|
||||
v_guided = v_c
|
||||
else:
|
||||
v_guided = v_u + guidance_scale * (v_c - v_u)
|
||||
|
||||
images = images + dt * v_guided
|
||||
|
||||
images = images.to(dtype=torch.float).clamp_(-1.0, 1.0)
|
||||
return [tensor_to_pil(img, unnormalize=True) for img in images]
|
||||
|
||||
@torch.inference_mode()
|
||||
def __call__(
|
||||
self,
|
||||
person_image: Image.Image,
|
||||
garment_image: Image.Image,
|
||||
category: Literal["tops", "bottoms", "one-pieces"],
|
||||
garment_photo_type: Literal["model", "flat-lay"] = "model",
|
||||
num_samples: int = 1,
|
||||
num_timesteps: int = 30,
|
||||
guidance_scale: float = 1.5,
|
||||
skip_cfg_last_n_steps: int = 1,
|
||||
seed: int = 42,
|
||||
segmentation_free: bool = True,
|
||||
callback: Optional[callable] = None,
|
||||
) -> PipelineOutput:
|
||||
"""
|
||||
Run virtual try-on inference.
|
||||
|
||||
Args:
|
||||
person_image: RGB image of the person to dress.
|
||||
garment_image: RGB image of the garment (model photo or flat-lay).
|
||||
category: Garment category - "tops", "bottoms", or "one-pieces".
|
||||
garment_photo_type: "model" if garment is worn by a person,
|
||||
"flat-lay" for product shots on plain backgrounds.
|
||||
num_samples: Number of output images to generate (1-4).
|
||||
num_timesteps: Diffusion sampling steps. Higher = better quality, slower.
|
||||
Recommended: 20 (fast), 30 (balanced), 50 (quality).
|
||||
guidance_scale: Classifier-free guidance strength.
|
||||
skip_cfg_last_n_steps: Skip CFG for final N steps to prevent color saturation.
|
||||
seed: Random seed for reproducibility.
|
||||
segmentation_free: If True, generate without masking the person image.
|
||||
Recommended for better body preservation and unconstrained garment volume
|
||||
(allows garments to expand beyond the original outfit's boundaries).
|
||||
|
||||
Returns:
|
||||
PipelineOutput with `images` list containing generated PIL Images.
|
||||
"""
|
||||
# Set seed
|
||||
torch.manual_seed(seed)
|
||||
if self.device.type == "cuda":
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
np.random.seed(seed)
|
||||
|
||||
# Pre-resize for pose detection quality
|
||||
person_image = self.pre_resize(person_image, allow_upsampling=False)
|
||||
garment_image = self.pre_resize(garment_image, allow_upsampling=False)
|
||||
|
||||
person_image_np = np.array(person_image)
|
||||
garment_image_np = np.array(garment_image)
|
||||
|
||||
# Pose detection (DWPose expects BGR)
|
||||
person_pose = self.pose_model(person_image_np[..., ::-1])
|
||||
garment_pose = (
|
||||
get_dummy_dw_keypoints()
|
||||
if garment_photo_type == "flat-lay"
|
||||
else self.pose_model(garment_image_np[..., ::-1])
|
||||
)
|
||||
|
||||
person_pose_img = draw_pose(person_pose, person_image_np.shape[0], person_image_np.shape[1], grayscale=True)
|
||||
garment_pose_img = draw_pose(garment_pose, garment_image_np.shape[0], garment_image_np.shape[1], grayscale=True)
|
||||
|
||||
# Human parsing
|
||||
person_seg_pred = self.hp_model.predict(person_image_np)
|
||||
garment_seg_pred = self.hp_model.predict(garment_image_np)
|
||||
|
||||
# Get labels to segment based on category
|
||||
body_coverage = CATEGORY_TO_BODY_COVERAGE.get(category)
|
||||
labels_to_segment = BODY_COVERAGE_TO_FASHN_LABELS.get(body_coverage)
|
||||
labels_to_segment_indices = [FASHN_LABELS_TO_IDS[label] for label in labels_to_segment]
|
||||
|
||||
# Create clothing-agnostic and garment images
|
||||
ca_image = create_clothing_agnostic_image(
|
||||
img_np=person_image_np.copy(),
|
||||
seg_pred=person_seg_pred.copy(),
|
||||
labels_to_segment_indices=labels_to_segment_indices.copy(),
|
||||
body_coverage=body_coverage,
|
||||
disable_masking=segmentation_free,
|
||||
logger=self.logger,
|
||||
)
|
||||
|
||||
garment_image_processed = create_garment_image(
|
||||
img_np=garment_image_np,
|
||||
seg_pred=garment_seg_pred,
|
||||
labels_to_segment_indices=labels_to_segment_indices.copy(),
|
||||
disable_masking=garment_photo_type == "flat-lay",
|
||||
)
|
||||
|
||||
# Resize/pad for model input
|
||||
ca_image = self.resize_pad_fn(ca_image, mem_padding=True)
|
||||
garment_image_processed = self.resize_pad_fn(garment_image_processed)
|
||||
person_pose_img = self.resize_pad_fn(person_pose_img, interpolation=cv2.INTER_NEAREST_EXACT)
|
||||
garment_pose_img = self.resize_pad_fn(garment_pose_img, interpolation=cv2.INTER_NEAREST_EXACT)
|
||||
|
||||
# Prepare tensors
|
||||
def prepare_tensor(img: np.ndarray) -> torch.Tensor:
|
||||
t = numpy_to_torch(img).unsqueeze(0)
|
||||
t = normalize_uint8_to_neg1_1(t)
|
||||
t = t.to(self.device).repeat(num_samples, 1, 1, 1)
|
||||
return t
|
||||
|
||||
ca_tensor = prepare_tensor(ca_image)
|
||||
garment_tensor = prepare_tensor(garment_image_processed)
|
||||
person_pose_tensor = prepare_tensor(person_pose_img)
|
||||
garment_pose_tensor = prepare_tensor(garment_pose_img)
|
||||
|
||||
garment_categories = (
|
||||
torch.tensor(self.CATEGORY_TO_LABEL[category]).unsqueeze(0).repeat(num_samples).to(self.device)
|
||||
)
|
||||
|
||||
# Cast to inference dtype
|
||||
ca_tensor = ca_tensor.to(dtype=self.inference_dtype)
|
||||
garment_tensor = garment_tensor.to(dtype=self.inference_dtype)
|
||||
person_pose_tensor = person_pose_tensor.to(dtype=self.inference_dtype)
|
||||
garment_pose_tensor = garment_pose_tensor.to(dtype=self.inference_dtype)
|
||||
|
||||
# Run sampling
|
||||
self.logger.info(f"Running inference with {num_timesteps} timesteps...")
|
||||
images = self._sample(
|
||||
ca_images=ca_tensor,
|
||||
garment_images=garment_tensor,
|
||||
person_poses=person_pose_tensor,
|
||||
garment_poses=garment_pose_tensor,
|
||||
garment_categories=garment_categories,
|
||||
num_timesteps=num_timesteps,
|
||||
guidance_scale=guidance_scale,
|
||||
skip_cfg_last_n_steps=skip_cfg_last_n_steps,
|
||||
callback=callback,
|
||||
)
|
||||
|
||||
# Unpad outputs
|
||||
images = [self.resize_pad_fn.unpad(img) for img in images]
|
||||
|
||||
self.logger.info(f"Generated {len(images)} images")
|
||||
|
||||
return PipelineOutput(images=images)
|
||||
@@ -0,0 +1,22 @@
|
||||
"""Preprocessing utilities."""
|
||||
|
||||
from .agnostic import (
|
||||
BODY_COVERAGE_TO_FASHN_LABELS,
|
||||
FASHN_LABELS_TO_IDS,
|
||||
create_clothing_agnostic_image,
|
||||
create_garment_image,
|
||||
)
|
||||
from .transforms import AspectPreserveResize, PadToShape, ResizePad
|
||||
|
||||
__all__ = [
|
||||
# Clothing-agnostic creation
|
||||
"create_clothing_agnostic_image",
|
||||
"create_garment_image",
|
||||
# Constants
|
||||
"FASHN_LABELS_TO_IDS",
|
||||
"BODY_COVERAGE_TO_FASHN_LABELS",
|
||||
# Transforms
|
||||
"AspectPreserveResize",
|
||||
"ResizePad",
|
||||
"PadToShape",
|
||||
]
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,212 @@
|
||||
"""Clothing-agnostic image creation."""
|
||||
|
||||
import logging
|
||||
from typing import List, Optional
|
||||
|
||||
import numpy as np
|
||||
from fashn_human_parser import BODY_COVERAGE_TO_LABELS, IDENTITY_LABELS, LABELS_TO_IDS
|
||||
|
||||
from ..utils import setup_logger
|
||||
from .masks import asymmetric_dilate_mask, create_bounded_mask, create_contour_following_mask, dilate_mask
|
||||
|
||||
# Re-export constants from fashn_human_parser for convenience
|
||||
FASHN_LABELS_TO_IDS = LABELS_TO_IDS
|
||||
BODY_COVERAGE_TO_FASHN_LABELS = BODY_COVERAGE_TO_LABELS
|
||||
IDENTITY_FASHN_LABELS = tuple(IDENTITY_LABELS)
|
||||
|
||||
|
||||
def _default(val, default_val):
|
||||
"""Return val if not None, else default_val (or call it if callable)."""
|
||||
if val is not None:
|
||||
return val
|
||||
return default_val() if callable(default_val) else default_val
|
||||
|
||||
|
||||
def _create_hybrid_contour_bounded_mask(
|
||||
contour_mask: np.ndarray,
|
||||
bounded_mask: np.ndarray,
|
||||
min_distance_threshold: float = 100.0,
|
||||
logger: Optional[logging.Logger] = None,
|
||||
baseline_height: float = 864.0,
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
Create hybrid mask by removing over-aggressive bounded expansions.
|
||||
|
||||
Combines contour-following and bounding-box masks, removing pixels from
|
||||
the bounded mask that are too far from the contour mask.
|
||||
|
||||
Args:
|
||||
contour_mask: Precise contour-following mask
|
||||
bounded_mask: More aggressive bounded box mask
|
||||
min_distance_threshold: Max distance from contour for bounded pixels (at baseline height)
|
||||
logger: Optional logger instance
|
||||
baseline_height: Reference height for scaling threshold
|
||||
|
||||
Returns:
|
||||
Hybrid mask with over-aggressive bounded pixels removed
|
||||
"""
|
||||
import cv2
|
||||
|
||||
logger = _default(logger, lambda: setup_logger("hybrid_mask"))
|
||||
|
||||
# Scale threshold based on image height
|
||||
height_scale = contour_mask.shape[0] / baseline_height
|
||||
scaled_threshold = min_distance_threshold * height_scale
|
||||
|
||||
if scaled_threshold <= 0:
|
||||
logger.debug("scaled_threshold<=0, returning pure contour mask")
|
||||
return contour_mask
|
||||
|
||||
hybrid_mask = bounded_mask.copy()
|
||||
|
||||
# Find pixels in bounded but not in contour (potential over-expansion)
|
||||
bounded_extra = bounded_mask & ~contour_mask
|
||||
|
||||
if not np.any(bounded_extra):
|
||||
logger.debug("No extra pixels in bounded mask, returning bounded mask")
|
||||
return bounded_mask
|
||||
|
||||
# Compute distance from bounded extra pixels to nearest contour mask pixel
|
||||
distance_from_contour = cv2.distanceTransform(
|
||||
(~contour_mask).astype(np.uint8), cv2.DIST_L2, 5
|
||||
)
|
||||
|
||||
# Remove pixels too far from contour
|
||||
bounded_extra_coords = np.where(bounded_extra)
|
||||
extra_distances = distance_from_contour[bounded_extra]
|
||||
remove_mask = extra_distances > scaled_threshold
|
||||
remove_coords = (bounded_extra_coords[0][remove_mask], bounded_extra_coords[1][remove_mask])
|
||||
|
||||
hybrid_mask[remove_coords] = False
|
||||
|
||||
return hybrid_mask
|
||||
|
||||
|
||||
def create_garment_image(
|
||||
img_np: np.ndarray,
|
||||
seg_pred: np.ndarray,
|
||||
labels_to_segment_indices: List[int],
|
||||
mask_value: int = 127,
|
||||
disable_masking: bool = False,
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
Create garment image with optional masking.
|
||||
|
||||
Masks out regions not belonging to the specified garment labels.
|
||||
|
||||
Args:
|
||||
img_np: Input image array (will be modified in-place)
|
||||
seg_pred: Segmentation prediction array
|
||||
labels_to_segment_indices: List of label indices to keep
|
||||
mask_value: Value to fill masked regions (default: 127 gray)
|
||||
disable_masking: If True, return image unchanged
|
||||
|
||||
Returns:
|
||||
Processed garment image array
|
||||
"""
|
||||
if not disable_masking:
|
||||
selected_labels_mask = np.isin(seg_pred, labels_to_segment_indices)
|
||||
img_np[~selected_labels_mask] = mask_value
|
||||
|
||||
return img_np
|
||||
|
||||
|
||||
def create_clothing_agnostic_image(
|
||||
img_np: np.ndarray,
|
||||
seg_pred: np.ndarray,
|
||||
labels_to_segment_indices: List[int],
|
||||
body_coverage: str,
|
||||
mask_value: int = 127,
|
||||
disable_masking: bool = False,
|
||||
min_distance_threshold: float = 100.0,
|
||||
baseline_height: float = 864.0,
|
||||
mask_limbs: bool = True,
|
||||
logger: Optional[logging.Logger] = None,
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
Create clothing-agnostic image.
|
||||
|
||||
Masks garments and body parts based on the target category.
|
||||
|
||||
Args:
|
||||
img_np: Input image array (will be modified in-place)
|
||||
seg_pred: Segmentation prediction array
|
||||
labels_to_segment_indices: List of label indices to mask
|
||||
body_coverage: Coverage type ("full", "upper", or "lower")
|
||||
mask_value: Value to fill masked regions (default: 127 gray)
|
||||
disable_masking: If True, return image unchanged
|
||||
min_distance_threshold: Distance threshold for hybrid mask (at baseline height)
|
||||
baseline_height: Reference height for parameter scaling
|
||||
mask_limbs: If True, also mask arms/legs based on body_coverage
|
||||
logger: Optional logger instance
|
||||
|
||||
Returns:
|
||||
Clothing-agnostic image array
|
||||
"""
|
||||
logger = _default(logger, lambda: setup_logger("clothing_agnostic"))
|
||||
|
||||
if disable_masking:
|
||||
return img_np
|
||||
|
||||
# Scale parameters based on image height
|
||||
height_scale = seg_pred.shape[0] / baseline_height
|
||||
logger.debug(f"Height scale factor: {height_scale:.3f} (height: {seg_pred.shape[0]})")
|
||||
|
||||
# Add body parts to mask based on body coverage
|
||||
labels_ids_dict = FASHN_LABELS_TO_IDS.copy()
|
||||
if mask_limbs:
|
||||
if body_coverage in ("full", "upper"):
|
||||
labels_to_segment_indices += [labels_ids_dict["arms"], labels_ids_dict["torso"]]
|
||||
if body_coverage in ("full", "lower"):
|
||||
labels_to_segment_indices += [labels_ids_dict["legs"]]
|
||||
|
||||
# Create base mask
|
||||
mask = np.isin(seg_pred, labels_to_segment_indices)
|
||||
|
||||
# Buffer mask to avoid leaks
|
||||
scaled_buffer_kernel = max(1, int(4 * height_scale))
|
||||
buffer_mask = dilate_mask(mask, kernel=(scaled_buffer_kernel, scaled_buffer_kernel))
|
||||
|
||||
# Create bounded mask
|
||||
bounded_mask = create_bounded_mask(mask)
|
||||
|
||||
# Create contour following mask
|
||||
scaled_brush_radius = max(1, int(18 * height_scale))
|
||||
contour_mask = create_contour_following_mask(mask, brush_radius=scaled_brush_radius)
|
||||
|
||||
# Create hybrid mask
|
||||
ca_mask = _create_hybrid_contour_bounded_mask(
|
||||
contour_mask, bounded_mask, logger=logger, min_distance_threshold=min_distance_threshold
|
||||
)
|
||||
|
||||
# Apply asymmetric dilation for inpainting workspace
|
||||
scaled_right = int(33 * height_scale)
|
||||
scaled_left = int(33 * height_scale)
|
||||
scaled_up = int(16 * height_scale)
|
||||
scaled_down = int(16 * height_scale)
|
||||
ca_mask = asymmetric_dilate_mask(ca_mask, right=scaled_right, left=scaled_left, up=scaled_up, down=scaled_down)
|
||||
|
||||
# Create exclusion mask (regions to preserve)
|
||||
identity_ids = [labels_ids_dict[label] for label in IDENTITY_FASHN_LABELS]
|
||||
|
||||
# Conditional identity based on coverage
|
||||
if body_coverage == "upper":
|
||||
identity_ids.append(labels_ids_dict["legs"])
|
||||
elif body_coverage == "lower":
|
||||
identity_ids.append(labels_ids_dict["arms"])
|
||||
|
||||
exclusion_mask = np.isin(seg_pred, identity_ids)
|
||||
|
||||
# Handle hands and feet
|
||||
if body_coverage in ("full", "upper"):
|
||||
hands_mask = seg_pred == labels_ids_dict["hands"]
|
||||
exclusion_mask = exclusion_mask | hands_mask
|
||||
|
||||
if body_coverage in ("full", "lower"):
|
||||
feet_mask = seg_pred == labels_ids_dict["feet"]
|
||||
exclusion_mask = exclusion_mask | feet_mask
|
||||
|
||||
final_mask = buffer_mask | (ca_mask & ~exclusion_mask)
|
||||
img_np[final_mask] = mask_value
|
||||
|
||||
return img_np
|
||||
@@ -0,0 +1,163 @@
|
||||
"""Mask processing utilities."""
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
def dilate_mask(mask: np.ndarray, kernel: tuple = (33, 33), iterations: int = 1) -> np.ndarray:
|
||||
"""
|
||||
Dilate the mask to create a buffer zone around the selected areas.
|
||||
|
||||
Args:
|
||||
mask: Input binary mask
|
||||
kernel: Dilation kernel size
|
||||
iterations: Number of dilation iterations
|
||||
|
||||
Returns:
|
||||
Dilated boolean mask
|
||||
"""
|
||||
kernel = np.ones(kernel, np.uint8)
|
||||
dilated_mask = cv2.dilate(mask.astype(np.uint8), kernel, iterations=iterations)
|
||||
return dilated_mask.astype(bool)
|
||||
|
||||
|
||||
def create_bounded_mask(mask: np.ndarray) -> np.ndarray:
|
||||
"""
|
||||
Create a mask that fills the bounding box of the input mask.
|
||||
|
||||
Args:
|
||||
mask: Input binary mask
|
||||
|
||||
Returns:
|
||||
Bounded mask filling the bounding rectangle
|
||||
"""
|
||||
bounded_mask = np.zeros_like(mask)
|
||||
x, y, w, h = cv2.boundingRect(mask.astype(np.uint8))
|
||||
bounded_mask[y : y + h, x : x + w] = 1
|
||||
return bounded_mask
|
||||
|
||||
|
||||
def asymmetric_dilate_mask(
|
||||
mask: np.ndarray, right: int, left: int, up: int, down: int
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
Dilate mask asymmetrically in different directions.
|
||||
|
||||
Args:
|
||||
mask: Input binary mask
|
||||
right: Dilation amount to the right
|
||||
left: Dilation amount to the left
|
||||
up: Dilation amount upward
|
||||
down: Dilation amount downward
|
||||
|
||||
Returns:
|
||||
Asymmetrically dilated boolean mask
|
||||
"""
|
||||
if mask.dtype == bool:
|
||||
mask = mask.astype(np.uint8) * 255
|
||||
|
||||
kernel_width = left + right + 1
|
||||
kernel_height = up + down + 1
|
||||
kernel = np.ones((kernel_height, kernel_width), np.uint8)
|
||||
|
||||
anchor_x = right
|
||||
anchor_y = down
|
||||
|
||||
mask = cv2.dilate(mask, kernel, anchor=(anchor_x, anchor_y))
|
||||
|
||||
return mask.astype(bool)
|
||||
|
||||
|
||||
def create_contour_following_mask(
|
||||
mask: np.ndarray,
|
||||
brush_radius: int = 36,
|
||||
smoothing_sigma: float | None = None,
|
||||
supersample: int = 1,
|
||||
keep_holes: bool = False,
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
Inflate mask so it looks like it was painted with a large soft brush.
|
||||
|
||||
Uses signed distance field for smooth contour following with optional
|
||||
supersampling for ultra-clean edges.
|
||||
|
||||
Args:
|
||||
mask: Input segmentation (foreground != 0)
|
||||
brush_radius: Extra pixels the virtual brush extends beyond the garment
|
||||
smoothing_sigma: Edge-smoothing sigma (Gaussian blur). If None, defaults to brush_radius / 2.5
|
||||
supersample: 1 = fastest; 2-4 for ultra-clean edges via upscale + max-pool downsample
|
||||
keep_holes: If True, preserve interior holes (e.g. neck opening)
|
||||
|
||||
Returns:
|
||||
Boolean mask, same H×W as input, guaranteed to contain the original mask
|
||||
"""
|
||||
if mask.dtype != np.bool_:
|
||||
mask = mask.astype(bool)
|
||||
|
||||
if smoothing_sigma is None:
|
||||
smoothing_sigma = brush_radius / 2.5
|
||||
|
||||
if supersample < 1 or not isinstance(supersample, int):
|
||||
raise ValueError("`supersample` must be a positive integer.")
|
||||
|
||||
# Optional super-sampling
|
||||
if supersample > 1:
|
||||
mask_work = cv2.resize(
|
||||
mask.astype(np.uint8),
|
||||
dsize=None,
|
||||
fx=supersample,
|
||||
fy=supersample,
|
||||
interpolation=cv2.INTER_NEAREST,
|
||||
).astype(bool)
|
||||
br = brush_radius * supersample
|
||||
sig = smoothing_sigma * supersample
|
||||
else:
|
||||
mask_work = mask.copy()
|
||||
br = brush_radius
|
||||
sig = smoothing_sigma
|
||||
|
||||
# Dilation ensures superset
|
||||
se = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (2 * br + 1, 2 * br + 1))
|
||||
dilated = cv2.dilate(mask_work.astype(np.uint8), se).astype(bool)
|
||||
|
||||
# Signed distance field
|
||||
dist_out = cv2.distanceTransform((~dilated).astype(np.uint8), cv2.DIST_L2, 5)
|
||||
dist_in = cv2.distanceTransform(dilated.astype(np.uint8), cv2.DIST_L2, 5)
|
||||
signed = dist_out - dist_in # <0 inside, >0 outside
|
||||
|
||||
# Smooth the level set
|
||||
signed_blur = cv2.GaussianBlur(signed, (0, 0), sig, borderType=cv2.BORDER_REPLICATE)
|
||||
smooth = signed_blur <= 0 # level-set 0
|
||||
|
||||
# Optional hole fill
|
||||
if not keep_holes:
|
||||
smooth = _fill_holes_cv(smooth)
|
||||
|
||||
# Containment: ensure original mask is included
|
||||
smooth |= mask_work
|
||||
|
||||
# Downsample if supersampled
|
||||
if supersample > 1:
|
||||
smooth = _max_pool_downsample(smooth, supersample)
|
||||
|
||||
return smooth.astype(bool)
|
||||
|
||||
|
||||
def _max_pool_downsample(arr: np.ndarray, factor: int) -> np.ndarray:
|
||||
"""Block-wise max-pooling downsample (factor must divide both axes)."""
|
||||
h, w = arr.shape
|
||||
if h % factor or w % factor:
|
||||
raise ValueError("Supersample factor must divide mask dimensions.")
|
||||
arr = arr.reshape(h // factor, factor, w // factor, factor)
|
||||
return arr.any(axis=(1, 3))
|
||||
|
||||
|
||||
def _fill_holes_cv(binary: np.ndarray) -> np.ndarray:
|
||||
"""Fill interior holes of a binary mask using flood-fill."""
|
||||
h, w = binary.shape
|
||||
flood_mask = np.zeros((h + 2, w + 2), np.uint8)
|
||||
inv = (~binary).astype(np.uint8) # holes & background = 1
|
||||
flood = inv.copy()
|
||||
cv2.floodFill(flood, flood_mask, (0, 0), 0) # erase the true background
|
||||
holes = flood == 1
|
||||
return binary | holes
|
||||
@@ -0,0 +1,222 @@
|
||||
"""Image transforms for preprocessing."""
|
||||
|
||||
from typing import Literal, Optional, Tuple, Union
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
from PIL import Image, ImageOps
|
||||
|
||||
|
||||
def _default(val, default_val):
|
||||
"""Return val if not None, else default_val (or call it if callable)."""
|
||||
if val is not None:
|
||||
return val
|
||||
return default_val() if callable(default_val) else default_val
|
||||
|
||||
|
||||
class AspectPreserveResize:
|
||||
"""
|
||||
Resize images while preserving aspect ratio.
|
||||
|
||||
Args:
|
||||
target_size: Target (width, height)
|
||||
mode: Resize mode
|
||||
- "fit": Scale to fit within target (may be smaller)
|
||||
- "exceed": Scale to exceed target (may be larger)
|
||||
- "short": Scale based on shorter dimension
|
||||
- "long": Scale based on longer dimension
|
||||
backend: "pil" or "opencv"
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
target_size: Tuple[int, int],
|
||||
mode: Literal["short", "long", "fit", "exceed"] = "fit",
|
||||
backend: Literal["pil", "opencv"] = "pil",
|
||||
):
|
||||
self.target_size = target_size
|
||||
self.mode = mode
|
||||
self.backend = backend
|
||||
|
||||
def _get_or_infer_scale_factor(
|
||||
self, width: int, height: int, allow_upsampling: bool = True
|
||||
) -> float:
|
||||
target_width, target_height = self.target_size
|
||||
scale_factor_width = target_width / width
|
||||
scale_factor_height = target_height / height
|
||||
|
||||
if self.mode == "long":
|
||||
scale_factor = min(scale_factor_width, scale_factor_height)
|
||||
elif self.mode == "short":
|
||||
scale_factor = max(scale_factor_width, scale_factor_height)
|
||||
elif self.mode == "fit":
|
||||
scale_factor = min(scale_factor_width, scale_factor_height)
|
||||
elif self.mode == "exceed":
|
||||
scale_factor = max(scale_factor_width, scale_factor_height)
|
||||
else:
|
||||
raise ValueError("Invalid mode. It should be 'short', 'long', 'fit', or 'exceed'.")
|
||||
|
||||
if not allow_upsampling and scale_factor > 1.0:
|
||||
return 1.0
|
||||
|
||||
return scale_factor
|
||||
|
||||
def _resize_image_pil(
|
||||
self, img: Image.Image, scale_factor: float, interpolation: Optional[int] = None
|
||||
) -> Image.Image:
|
||||
if scale_factor == 1.0:
|
||||
return img
|
||||
|
||||
width, height = img.size
|
||||
new_width = int(scale_factor * width)
|
||||
new_height = int(scale_factor * height)
|
||||
|
||||
interpolation = _default(interpolation, Image.LANCZOS)
|
||||
return img.resize((new_width, new_height), interpolation)
|
||||
|
||||
def _resize_image_opencv(
|
||||
self, img: np.ndarray, scale_factor: float, interpolation: Optional[int] = None
|
||||
) -> np.ndarray:
|
||||
if scale_factor == 1.0:
|
||||
return img
|
||||
|
||||
height, width = img.shape[:2]
|
||||
new_width = int(scale_factor * width)
|
||||
new_height = int(scale_factor * height)
|
||||
|
||||
interpolation = _default(
|
||||
interpolation, cv2.INTER_LANCZOS4 if scale_factor > 1 else cv2.INTER_AREA
|
||||
)
|
||||
|
||||
return cv2.resize(img, (new_width, new_height), interpolation=interpolation)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
img: Union[Image.Image, np.ndarray],
|
||||
allow_upsampling: bool = True,
|
||||
interpolation: Optional[int] = None,
|
||||
) -> Union[Image.Image, np.ndarray]:
|
||||
if self.backend == "pil":
|
||||
width, height = img.size
|
||||
elif self.backend == "opencv":
|
||||
height, width = img.shape[:2]
|
||||
|
||||
scale_factor = self._get_or_infer_scale_factor(width, height, allow_upsampling)
|
||||
|
||||
if self.backend == "pil":
|
||||
return self._resize_image_pil(img, scale_factor, interpolation=interpolation)
|
||||
elif self.backend == "opencv":
|
||||
return self._resize_image_opencv(img, scale_factor, interpolation=interpolation)
|
||||
|
||||
|
||||
class PadToShape:
|
||||
"""
|
||||
Pad images to a target shape with symmetric padding.
|
||||
|
||||
Args:
|
||||
target_size: Target (width, height)
|
||||
fill_value: Padding color (int for grayscale, tuple for RGB)
|
||||
backend: "pil" or "opencv"
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
target_size: Tuple[int, int],
|
||||
fill_value: Union[int, tuple] = 0,
|
||||
backend: Literal["pil", "opencv"] = "opencv",
|
||||
) -> None:
|
||||
self.target_width, self.target_height = target_size
|
||||
self.backend = backend
|
||||
if isinstance(fill_value, int):
|
||||
self.fill_value = (fill_value,) * 3
|
||||
else:
|
||||
self.fill_value = fill_value
|
||||
self.padding_mem: Optional[Tuple[int, int, int, int]] = None
|
||||
|
||||
@staticmethod
|
||||
def _calculate_needed_padding(
|
||||
width: int, height: int, target_width: int, target_height: int
|
||||
) -> Tuple[int, int, int, int]:
|
||||
total_width_padding = max(target_width - width, 0)
|
||||
total_height_padding = max(target_height - height, 0)
|
||||
|
||||
pad_left = total_width_padding // 2
|
||||
pad_top = total_height_padding // 2
|
||||
pad_right = total_width_padding - pad_left
|
||||
pad_bottom = total_height_padding - pad_top
|
||||
|
||||
return pad_left, pad_top, pad_right, pad_bottom
|
||||
|
||||
def _pad_image_pil(self, img: Image.Image, padding: Tuple[int, int, int, int]) -> Image.Image:
|
||||
return ImageOps.expand(img, border=padding, fill=self.fill_value)
|
||||
|
||||
def _pad_image_opencv(self, img: np.ndarray, padding: Tuple[int, int, int, int]) -> np.ndarray:
|
||||
pad_left, pad_top, pad_right, pad_bottom = padding
|
||||
return cv2.copyMakeBorder(
|
||||
img, pad_top, pad_bottom, pad_left, pad_right, cv2.BORDER_CONSTANT, value=self.fill_value
|
||||
)
|
||||
|
||||
def unpad(self, img: Union[Image.Image, np.ndarray]) -> Union[Image.Image, np.ndarray]:
|
||||
"""Remove padding using stored padding dimensions."""
|
||||
if self.padding_mem is None:
|
||||
raise ValueError("Padding memory is not set.")
|
||||
|
||||
pad_left, pad_top, pad_right, pad_bottom = self.padding_mem
|
||||
|
||||
if isinstance(img, Image.Image):
|
||||
return img.crop((pad_left, pad_top, img.width - pad_right, img.height - pad_bottom))
|
||||
return img[pad_top : img.shape[0] - pad_bottom, pad_left : img.shape[1] - pad_right]
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
img: Union[Image.Image, np.ndarray],
|
||||
mem_padding: bool = False,
|
||||
) -> Union[Image.Image, np.ndarray]:
|
||||
if self.backend == "pil":
|
||||
width, height = img.size
|
||||
else:
|
||||
height, width = img.shape[:2]
|
||||
|
||||
padding = self._calculate_needed_padding(width, height, self.target_width, self.target_height)
|
||||
|
||||
if mem_padding:
|
||||
self.padding_mem = padding
|
||||
|
||||
if self.backend == "pil":
|
||||
return self._pad_image_pil(img, padding)
|
||||
return self._pad_image_opencv(img, padding)
|
||||
|
||||
|
||||
class ResizePad:
|
||||
"""
|
||||
Aspect-preserving resize followed by symmetric padding.
|
||||
|
||||
Combines AspectPreserveResize and PadToShape to resize images to fit
|
||||
within target dimensions while preserving aspect ratio, then pads
|
||||
to reach exact target size.
|
||||
|
||||
Args:
|
||||
target_image_size: Target (width, height)
|
||||
backend: "pil" or "opencv"
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
target_image_size: Tuple[int, int],
|
||||
backend: Literal["pil", "opencv"] = "opencv",
|
||||
) -> None:
|
||||
self.resize_fn = AspectPreserveResize(target_size=target_image_size, mode="fit", backend=backend)
|
||||
self.pad_fn = PadToShape(target_image_size, backend=backend)
|
||||
|
||||
def unpad(self, img: Union[Image.Image, np.ndarray]) -> Union[Image.Image, np.ndarray]:
|
||||
"""Remove padding to restore original dimensions."""
|
||||
return self.pad_fn.unpad(img)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
img: Union[Image.Image, np.ndarray],
|
||||
mem_padding: bool = False,
|
||||
interpolation: Optional[int] = None,
|
||||
) -> Union[Image.Image, np.ndarray]:
|
||||
img = self.resize_fn(img, interpolation=interpolation)
|
||||
return self.pad_fn(img, mem_padding=mem_padding)
|
||||
@@ -0,0 +1,563 @@
|
||||
"""
|
||||
TryOn Model.
|
||||
|
||||
Contains components adapted from FLUX.1 by Black Forest Labs (Apache-2.0):
|
||||
https://github.com/black-forest-labs/flux
|
||||
"""
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
from einops import rearrange, repeat
|
||||
from torch import Tensor, nn
|
||||
|
||||
from .utils import cast_tuple, compact, exists, unpack_images
|
||||
|
||||
|
||||
# Use PyTorch's native scaled dot product attention (SDPA)
|
||||
def _attn_processor(q: Tensor, k: Tensor, v: Tensor) -> Tensor:
|
||||
"""Scaled dot product attention using PyTorch native implementation."""
|
||||
return torch.nn.functional.scaled_dot_product_attention(q, k, v)
|
||||
|
||||
|
||||
def attention(q: Tensor, k: Tensor, v: Tensor, pe: Tensor) -> Tensor:
|
||||
q, k = apply_rope(q, k, pe)
|
||||
x = _attn_processor(q, k, v)
|
||||
x = rearrange(x, "B H L D -> B L (H D)")
|
||||
|
||||
return x
|
||||
|
||||
|
||||
def rope(pos: Tensor, dim: int, theta: int) -> Tensor:
|
||||
assert dim % 2 == 0
|
||||
scale = torch.arange(0, dim, 2, dtype=torch.float64, device=pos.device) / dim
|
||||
omega = 1.0 / (theta**scale)
|
||||
out = torch.einsum("...n,d->...nd", pos, omega)
|
||||
out = torch.stack([torch.cos(out), -torch.sin(out), torch.sin(out), torch.cos(out)], dim=-1)
|
||||
out = rearrange(out, "b n d (i j) -> b n d i j", i=2, j=2)
|
||||
return out.float()
|
||||
|
||||
|
||||
def apply_rope(xq: Tensor, xk: Tensor, freqs_cis: Tensor) -> tuple[Tensor, Tensor]:
|
||||
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
|
||||
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
|
||||
xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
|
||||
xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
|
||||
return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)
|
||||
|
||||
|
||||
class EmbedND(nn.Module):
|
||||
def __init__(self, dim: int, theta: int, axes_dim: list[int]):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.theta = theta
|
||||
self.axes_dim = axes_dim
|
||||
|
||||
def forward(self, ids: Tensor) -> Tensor:
|
||||
n_axes = ids.shape[-1]
|
||||
emb = torch.cat(
|
||||
[rope(ids[..., i], self.axes_dim[i], self.theta) for i in range(n_axes)],
|
||||
dim=-3,
|
||||
)
|
||||
|
||||
return emb.unsqueeze(1)
|
||||
|
||||
class RMSNorm(torch.nn.Module):
|
||||
def __init__(self, dim: int):
|
||||
super().__init__()
|
||||
self.scale = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x: Tensor):
|
||||
x_dtype = x.dtype
|
||||
x = x.float()
|
||||
rrms = torch.rsqrt(torch.mean(x**2, dim=-1, keepdim=True) + 1e-6)
|
||||
return (x * rrms).to(dtype=x_dtype) * self.scale
|
||||
|
||||
|
||||
class QKNorm(torch.nn.Module):
|
||||
def __init__(self, dim: int):
|
||||
super().__init__()
|
||||
self.query_norm = RMSNorm(dim)
|
||||
self.key_norm = RMSNorm(dim)
|
||||
|
||||
def forward(self, q: Tensor, k: Tensor, v: Tensor) -> tuple[Tensor, Tensor]:
|
||||
q = self.query_norm(q)
|
||||
k = self.key_norm(k)
|
||||
return q.to(v), k.to(v)
|
||||
|
||||
|
||||
class SelfAttention(nn.Module):
|
||||
def __init__(self, dim: int, num_heads: int = 8, qkv_bias: bool = False):
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim // num_heads
|
||||
|
||||
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
||||
self.norm = QKNorm(head_dim)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
|
||||
def forward(self, x: Tensor, pe: Tensor) -> Tensor:
|
||||
qkv = self.qkv(x)
|
||||
q, k, v = rearrange(qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
|
||||
q, k = self.norm(q, k, v)
|
||||
x = attention(q, k, v, pe=pe)
|
||||
x = self.proj(x)
|
||||
return x
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModulationOut:
|
||||
shift: Tensor
|
||||
scale: Tensor
|
||||
gate: Tensor
|
||||
|
||||
|
||||
class Modulation(nn.Module):
|
||||
def __init__(self, dim: int, double: bool):
|
||||
super().__init__()
|
||||
self.is_double = double
|
||||
self.multiplier = 6 if double else 3
|
||||
self.lin = nn.Linear(dim, self.multiplier * dim, bias=True)
|
||||
|
||||
def forward(self, vec: Tensor) -> tuple[ModulationOut, ModulationOut | None]:
|
||||
out = self.lin(nn.functional.silu(vec))[:, None, :].chunk(self.multiplier, dim=-1)
|
||||
|
||||
return (
|
||||
ModulationOut(*out[:3]),
|
||||
ModulationOut(*out[3:]) if self.is_double else None,
|
||||
)
|
||||
|
||||
|
||||
class DoubleStreamBlock(nn.Module):
|
||||
def __init__(self, hidden_size: int, num_heads: int, mlp_ratio: float, qkv_bias: bool = False):
|
||||
super().__init__()
|
||||
|
||||
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
self.num_heads = num_heads
|
||||
self.hidden_size = hidden_size
|
||||
self.img_mod = Modulation(hidden_size, double=True)
|
||||
self.img_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.img_attn = SelfAttention(dim=hidden_size, num_heads=num_heads, qkv_bias=qkv_bias)
|
||||
|
||||
self.img_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.img_mlp = nn.Sequential(
|
||||
nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
|
||||
nn.GELU(approximate="tanh"),
|
||||
nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
|
||||
)
|
||||
|
||||
self.txt_mod = Modulation(hidden_size, double=True)
|
||||
self.txt_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.txt_attn = SelfAttention(dim=hidden_size, num_heads=num_heads, qkv_bias=qkv_bias)
|
||||
|
||||
self.txt_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.txt_mlp = nn.Sequential(
|
||||
nn.Linear(hidden_size, mlp_hidden_dim, bias=True),
|
||||
nn.GELU(approximate="tanh"),
|
||||
nn.Linear(mlp_hidden_dim, hidden_size, bias=True),
|
||||
)
|
||||
|
||||
def forward(self, img: Tensor, txt: Tensor, vec: Tensor, pe: Tensor) -> tuple[Tensor, Tensor]:
|
||||
img_mod1, img_mod2 = self.img_mod(vec)
|
||||
txt_mod1, txt_mod2 = self.txt_mod(vec)
|
||||
|
||||
# prepare image for attention
|
||||
img_modulated = self.img_norm1(img)
|
||||
img_modulated = (1 + img_mod1.scale) * img_modulated + img_mod1.shift
|
||||
img_qkv = self.img_attn.qkv(img_modulated)
|
||||
img_q, img_k, img_v = rearrange(img_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
|
||||
img_q, img_k = self.img_attn.norm(img_q, img_k, img_v)
|
||||
|
||||
# prepare txt for attention
|
||||
txt_modulated = self.txt_norm1(txt)
|
||||
txt_modulated = (1 + txt_mod1.scale) * txt_modulated + txt_mod1.shift
|
||||
txt_qkv = self.txt_attn.qkv(txt_modulated)
|
||||
txt_q, txt_k, txt_v = rearrange(txt_qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
|
||||
txt_q, txt_k = self.txt_attn.norm(txt_q, txt_k, txt_v)
|
||||
|
||||
# run actual attention
|
||||
q = torch.cat((txt_q, img_q), dim=2)
|
||||
k = torch.cat((txt_k, img_k), dim=2)
|
||||
v = torch.cat((txt_v, img_v), dim=2)
|
||||
|
||||
attn = attention(q, k, v, pe=pe)
|
||||
txt_attn, img_attn = attn[:, : txt.shape[1]], attn[:, txt.shape[1] :]
|
||||
|
||||
# calculate the img blocks
|
||||
img = img + img_mod1.gate * self.img_attn.proj(img_attn)
|
||||
img = img + img_mod2.gate * self.img_mlp((1 + img_mod2.scale) * self.img_norm2(img) + img_mod2.shift)
|
||||
|
||||
# calculate the txt blocks
|
||||
txt = txt + txt_mod1.gate * self.txt_attn.proj(txt_attn)
|
||||
txt = txt + txt_mod2.gate * self.txt_mlp((1 + txt_mod2.scale) * self.txt_norm2(txt) + txt_mod2.shift)
|
||||
return img, txt
|
||||
|
||||
|
||||
class SingleStreamBlock(nn.Module):
|
||||
"""
|
||||
A DiT block with parallel linear layers as described in
|
||||
https://arxiv.org/abs/2302.05442 and adapted modulation interface.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
head_dim = hidden_size // num_heads
|
||||
|
||||
self.mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
||||
# qkv and mlp_in
|
||||
self.linear1 = nn.Linear(hidden_size, hidden_size * 3 + self.mlp_hidden_dim)
|
||||
# proj and mlp_out
|
||||
self.linear2 = nn.Linear(hidden_size + self.mlp_hidden_dim, hidden_size)
|
||||
|
||||
self.norm = QKNorm(head_dim)
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.pre_norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
|
||||
self.mlp_act = nn.GELU(approximate="tanh")
|
||||
self.modulation = Modulation(hidden_size, double=False)
|
||||
|
||||
def forward(self, x: Tensor, vec: Tensor, pe: Tensor) -> Tensor:
|
||||
mod, _ = self.modulation(vec)
|
||||
x_mod = (1 + mod.scale) * self.pre_norm(x) + mod.shift
|
||||
qkv, mlp = torch.split(self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1)
|
||||
|
||||
q, k, v = rearrange(qkv, "B L (K H D) -> K B H L D", K=3, H=self.num_heads)
|
||||
q, k = self.norm(q, k, v)
|
||||
|
||||
# compute attention
|
||||
attn = attention(q, k, v, pe=pe)
|
||||
# compute activation in mlp stream, cat again and run second linear layer
|
||||
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
|
||||
return x + mod.gate * output
|
||||
|
||||
|
||||
class LastLayer(nn.Module):
|
||||
def __init__(self, hidden_size: int, out_channels: int):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.linear = nn.Linear(hidden_size, out_channels, bias=True)
|
||||
self.adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
|
||||
|
||||
def forward(self, x: Tensor, vec: Tensor) -> Tensor:
|
||||
shift, scale = self.adaLN_modulation(vec).chunk(2, dim=1)
|
||||
x = (1 + scale[:, None, :]) * self.norm_final(x) + shift[:, None, :]
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
def prepare(img: Tensor, patch_size: int = 1) -> dict[str, Tensor]:
|
||||
bs, c, h, w = img.shape
|
||||
|
||||
# Rearrange the image into patches based on the given patch size
|
||||
img = rearrange(img, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch_size, pw=patch_size)
|
||||
|
||||
# Ensure all images in the batch are processed if the input batch size was 1
|
||||
if img.shape[0] == 1 and bs > 1:
|
||||
img = repeat(img, "1 ... -> bs ...", bs=bs)
|
||||
|
||||
# Create image ids for positional encoding
|
||||
img_ids = torch.zeros(h // patch_size, w // patch_size, 3)
|
||||
img_ids[..., 1] = img_ids[..., 1] + torch.arange(h // patch_size)[:, None]
|
||||
img_ids[..., 2] = img_ids[..., 2] + torch.arange(w // patch_size)[None, :]
|
||||
img_ids = repeat(img_ids, "h w c -> b (h w) c", b=bs)
|
||||
|
||||
return img, img_ids.to(img.device)
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
"""2D Image to Patch Embedding"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
img_size=224,
|
||||
patch_size=16,
|
||||
in_chans=3,
|
||||
embed_dim=768,
|
||||
norm_layer=None,
|
||||
flatten=True,
|
||||
bias=True,
|
||||
):
|
||||
super().__init__()
|
||||
img_size = cast_tuple(img_size, 2)
|
||||
patch_size = cast_tuple(patch_size, 2)
|
||||
self.img_size = img_size
|
||||
self.patch_size = patch_size
|
||||
self.grid_size = (img_size[0] // patch_size[0], img_size[1] // patch_size[1])
|
||||
self.num_patches = self.grid_size[0] * self.grid_size[1]
|
||||
self.flatten = flatten
|
||||
|
||||
self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size, bias=bias)
|
||||
self.norm = norm_layer(embed_dim) if exists(norm_layer) else nn.Identity()
|
||||
|
||||
# Initialize patch_embed like nn.Linear (instead of nn.Conv2d):
|
||||
w = self.proj.weight.data
|
||||
nn.init.xavier_uniform_(w.view([w.shape[0], -1]))
|
||||
nn.init.constant_(self.proj.bias, 0)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.proj(x)
|
||||
if self.flatten:
|
||||
x = x.flatten(2).transpose(1, 2) # BCHW -> BNC
|
||||
x = self.norm(x)
|
||||
return x
|
||||
|
||||
class MLPEmbedder(nn.Module):
|
||||
def __init__(self, in_dim: int, hidden_dim: int):
|
||||
super().__init__()
|
||||
self.in_layer = nn.Linear(in_dim, hidden_dim, bias=True)
|
||||
self.silu = nn.SiLU()
|
||||
self.out_layer = nn.Linear(hidden_dim, hidden_dim, bias=True)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return self.out_layer(self.silu(self.in_layer(x)))
|
||||
|
||||
def timestep_embedding(t: Tensor, dim, max_period=10000, time_factor: float = 1000.0):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
:param t: a 1-D Tensor of N indices, one per batch element.
|
||||
These may be fractional.
|
||||
:param dim: the dimension of the output.
|
||||
:param max_period: controls the minimum frequency of the embeddings.
|
||||
:return: an (N, D) Tensor of positional embeddings.
|
||||
"""
|
||||
t = time_factor * t
|
||||
half = dim // 2
|
||||
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half).to(t.device)
|
||||
|
||||
args = t[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
if torch.is_floating_point(t):
|
||||
embedding = embedding.to(t)
|
||||
return embedding
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
def __init__(self, hidden_size, frequency_embedding_size=256):
|
||||
super().__init__()
|
||||
self.mlp = MLPEmbedder(frequency_embedding_size, hidden_size)
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
|
||||
def forward(self, t: Tensor) -> Tensor:
|
||||
return self.mlp(timestep_embedding(t, self.frequency_embedding_size))
|
||||
|
||||
def apply_conditional_dropout(tensor, mask, null_tensor=None):
|
||||
device, dtype = tensor.device, tensor.dtype
|
||||
mask_shape = [mask.shape[0]] + [1] * (tensor.dim() - 1)
|
||||
keep_mask = mask.view(*mask_shape)
|
||||
|
||||
if exists(null_tensor):
|
||||
null_tensor = null_tensor.to(device=device, dtype=dtype)
|
||||
null_tensor = null_tensor.expand_as(tensor)
|
||||
else:
|
||||
null_tensor = torch.zeros_like(tensor)
|
||||
|
||||
return torch.where(keep_mask, tensor, null_tensor)
|
||||
|
||||
|
||||
class TryOnModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_shape: Tuple[int] = (864, 576),
|
||||
hidden_size: int = 1280,
|
||||
n_heads=10,
|
||||
double_blocks_depth: int = 8,
|
||||
single_blocks_depth: int = 16,
|
||||
mlp_ratio: int = 4,
|
||||
channels_in: int = 3,
|
||||
patch_size: int = 12,
|
||||
theta: int = 10000,
|
||||
axes_dim: Tuple[int] = (16, 56, 56),
|
||||
qkv_bias: bool = True,
|
||||
guidance_embed: bool = False,
|
||||
n_classes: int = 3,
|
||||
use_patch_mixer: bool = True,
|
||||
patch_mixer_depth: int = 4,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# time
|
||||
self.t_embedder = TimestepEmbedder(hidden_size=hidden_size)
|
||||
|
||||
# category labels (tops, bottoms, one-pieces)
|
||||
self.y_embedder = nn.Embedding(n_classes + 1, hidden_size) if n_classes > 0 else None # +1 for null class
|
||||
|
||||
# guidance embeddings for guidance distillation
|
||||
self.guidance_embedder = TimestepEmbedder(hidden_size=hidden_size) if guidance_embed else None
|
||||
|
||||
# positional embeddings
|
||||
pe_dim = hidden_size // n_heads
|
||||
if sum(axes_dim) != pe_dim:
|
||||
raise ValueError(f"Got {axes_dim} but expected positional dim {pe_dim}")
|
||||
self.pe_embedder = EmbedND(dim=pe_dim, theta=theta, axes_dim=axes_dim)
|
||||
|
||||
# images
|
||||
self.input_shape = input_shape
|
||||
self.patch_size = patch_size
|
||||
self.channels_in = channels_in
|
||||
self.x_embedder = PatchEmbed(
|
||||
img_size=input_shape, patch_size=patch_size, in_chans=channels_in * 2 + 1, embed_dim=hidden_size, flatten=False
|
||||
)
|
||||
self.garment_embedder = PatchEmbed(
|
||||
img_size=input_shape, patch_size=patch_size, in_chans=channels_in + 1, embed_dim=hidden_size, flatten=False
|
||||
)
|
||||
|
||||
# patch mixer
|
||||
self.use_patch_mixer = use_patch_mixer
|
||||
if use_patch_mixer:
|
||||
self.x_patch_mixer = nn.ModuleList(
|
||||
[SingleStreamBlock(hidden_size, n_heads, mlp_ratio=mlp_ratio) for _ in range(patch_mixer_depth)]
|
||||
)
|
||||
# Buffer kept for checkpoint compatibility (not used at inference)
|
||||
self.register_buffer("patch_mixer_token", torch.zeros(1, 1, channels_in * self.patch_size**2))
|
||||
|
||||
# core MMDiT
|
||||
self.double_blocks = nn.ModuleList(
|
||||
[
|
||||
DoubleStreamBlock(
|
||||
hidden_size,
|
||||
n_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
)
|
||||
for _ in range(double_blocks_depth)
|
||||
]
|
||||
)
|
||||
|
||||
self.single_blocks = nn.ModuleList(
|
||||
[SingleStreamBlock(hidden_size, n_heads, mlp_ratio=mlp_ratio) for _ in range(single_blocks_depth)]
|
||||
)
|
||||
|
||||
self.final_layer = LastLayer(hidden_size, out_channels=channels_in * self.patch_size**2)
|
||||
|
||||
# forward with classifier free guidance
|
||||
|
||||
def forward_for_cfg(self, *args, **kwargs):
|
||||
# cleanup kwargs
|
||||
kwargs = compact(kwargs)
|
||||
|
||||
# infer batch size from the first tensor argument
|
||||
noisy_images = args[0]
|
||||
batch_size = noisy_images.shape[0]
|
||||
|
||||
# duplicate all tensor arguments and keyword arguments
|
||||
duplicated_args = [torch.cat([arg, arg], dim=0) if isinstance(arg, torch.Tensor) else arg for arg in args]
|
||||
duplicated_kwargs = {
|
||||
k: torch.cat([v, v], dim=0) if isinstance(v, torch.Tensor) else v for k, v in kwargs.items()
|
||||
}
|
||||
|
||||
# prepare cond drop masks for the duplicated inputs
|
||||
mask = torch.cat(
|
||||
[
|
||||
torch.ones(batch_size, device=noisy_images.device, dtype=torch.bool),
|
||||
torch.zeros(batch_size, device=noisy_images.device, dtype=torch.bool),
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
|
||||
# add cond_drop_probs to duplicated_kwargs
|
||||
duplicated_kwargs["mask"] = mask
|
||||
|
||||
# execute the forward pass with duplicated inputs
|
||||
all_logits = self.forward(*duplicated_args, **duplicated_kwargs)["x"]
|
||||
|
||||
# split the logits into original and null versions
|
||||
logits, null_logits = all_logits.split(batch_size)
|
||||
|
||||
return {"v_c": logits, "v_u": null_logits}
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
times,
|
||||
ca_images,
|
||||
garment_images,
|
||||
person_poses,
|
||||
garment_poses,
|
||||
mask: Optional[torch.Tensor] = None,
|
||||
guidance: Optional[torch.Tensor] = None,
|
||||
garment_categories: Optional[torch.Tensor] = None,
|
||||
):
|
||||
###################### CLASSIFIER FREE GUIDANCE ######################
|
||||
|
||||
batch_size, device = x.shape[0], x.device
|
||||
|
||||
# if mask is not provided, create a boolean mask of all true
|
||||
if not exists(mask):
|
||||
mask = torch.ones(batch_size, device=device, dtype=torch.bool)
|
||||
|
||||
####################### 2D IMAGES TO SEQUENCE ########################
|
||||
|
||||
ca_images = apply_conditional_dropout(ca_images, mask)
|
||||
person_poses = apply_conditional_dropout(person_poses, mask)
|
||||
x = torch.cat([x, ca_images, person_poses], dim=1)
|
||||
x = self.x_embedder(x)
|
||||
x, x_ids = prepare(x)
|
||||
|
||||
garment_poses = apply_conditional_dropout(garment_poses, mask)
|
||||
garment_images = apply_conditional_dropout(garment_images, mask)
|
||||
garment_images = torch.cat([garment_images, garment_poses], dim=1)
|
||||
garment_images = self.garment_embedder(garment_images)
|
||||
garment_images, garment_ids = prepare(garment_images)
|
||||
|
||||
###################### TIME & MODULATION ######################
|
||||
|
||||
t = self.t_embedder(times)
|
||||
|
||||
if exists(self.guidance_embedder):
|
||||
assert exists(guidance), "Guidance scale required for guidance distilled model"
|
||||
t = t + self.guidance_embedder(guidance)
|
||||
|
||||
if exists(self.y_embedder):
|
||||
assert exists(garment_categories), "Category labels required for y_embedder"
|
||||
y = apply_conditional_dropout(garment_categories, mask)
|
||||
t = t + self.y_embedder(y)
|
||||
|
||||
###################### POSITIONAL EMBEDDINGS ######################
|
||||
|
||||
img, txt, vec = x, garment_images, t # name change for consistency with the original code
|
||||
|
||||
x_pe = self.pe_embedder(x_ids)
|
||||
g_pe = self.pe_embedder(garment_ids)
|
||||
|
||||
###################### PATCH MIXER ######################
|
||||
|
||||
if self.use_patch_mixer:
|
||||
for block in self.x_patch_mixer:
|
||||
img = block(img, vec=vec, pe=x_pe)
|
||||
|
||||
###################### CORE MMDiT ########################
|
||||
|
||||
pe = torch.cat([x_pe, g_pe], dim=2)
|
||||
|
||||
for block in self.double_blocks:
|
||||
img, txt = block(img=img, txt=txt, vec=vec, pe=pe)
|
||||
|
||||
img = torch.cat((txt, img), 1)
|
||||
for block in self.single_blocks:
|
||||
img = block(img, vec=vec, pe=pe)
|
||||
img = img[:, txt.shape[1] :, ...]
|
||||
|
||||
x = self.final_layer(img, vec)
|
||||
|
||||
###################### SEQUENCE TO 2D IMAGES ########################
|
||||
|
||||
x = rearrange(
|
||||
x,
|
||||
"b (h w) c -> b c h w",
|
||||
h=self.input_shape[0] // self.patch_size,
|
||||
w=self.input_shape[1] // self.patch_size,
|
||||
)
|
||||
if self.patch_size > 1:
|
||||
x = unpack_images(x, self.patch_size)
|
||||
|
||||
return {"x": x}
|
||||
@@ -0,0 +1,40 @@
|
||||
"""
|
||||
Utility functions for FASHN VTON.
|
||||
|
||||
This package provides common utilities:
|
||||
- Common Python helpers (exists, default, cast_tuple, compact)
|
||||
- Model checkpoint loading
|
||||
- Tensor operations and conversions
|
||||
- Sampling schedules for Rectified Flow
|
||||
- Pose keypoint utilities
|
||||
- Logging setup
|
||||
"""
|
||||
|
||||
from .checkpoint import load_checkpoint
|
||||
from .common import cast_tuple, compact, default, exists
|
||||
from .keypoints import get_dummy_dw_keypoints
|
||||
from .logger import setup_logger
|
||||
from .sampling import get_rf_schedule, time_shift
|
||||
from .tensor import normalize_uint8_to_neg1_1, numpy_to_torch, tensor_to_pil, unpack_images
|
||||
|
||||
__all__ = [
|
||||
# Common helpers
|
||||
"exists",
|
||||
"default",
|
||||
"cast_tuple",
|
||||
"compact",
|
||||
# Checkpoint loading
|
||||
"load_checkpoint",
|
||||
# Tensor operations
|
||||
"numpy_to_torch",
|
||||
"unpack_images",
|
||||
"normalize_uint8_to_neg1_1",
|
||||
"tensor_to_pil",
|
||||
# Sampling
|
||||
"time_shift",
|
||||
"get_rf_schedule",
|
||||
# Pose
|
||||
"get_dummy_dw_keypoints",
|
||||
# Logging
|
||||
"setup_logger",
|
||||
]
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,44 @@
|
||||
"""Model checkpoint loading utilities."""
|
||||
|
||||
import os
|
||||
|
||||
import torch
|
||||
from safetensors.torch import load_file
|
||||
|
||||
|
||||
def load_checkpoint(checkpoint_path: str, device: str = "cpu") -> dict:
|
||||
"""
|
||||
Load model checkpoint from local file or HuggingFace Hub.
|
||||
|
||||
Supports:
|
||||
- Local .pt, .pth files (PyTorch format)
|
||||
- Local .safetensors files
|
||||
- HuggingFace repo IDs (e.g., "fashn-ai/fashn-vton-1.5")
|
||||
|
||||
Args:
|
||||
checkpoint_path: Local file path or HuggingFace repo ID
|
||||
device: Device to load the checkpoint to
|
||||
|
||||
Returns:
|
||||
The loaded state dictionary
|
||||
"""
|
||||
# Check if it's a local file
|
||||
if os.path.isfile(checkpoint_path):
|
||||
if checkpoint_path.endswith(".pt") or checkpoint_path.endswith(".pth"):
|
||||
return torch.load(checkpoint_path, map_location=device, weights_only=False)
|
||||
elif checkpoint_path.endswith(".safetensors"):
|
||||
return load_file(checkpoint_path, device=device)
|
||||
else:
|
||||
raise ValueError(f"Unknown checkpoint file format: {checkpoint_path}")
|
||||
|
||||
# Check if it looks like a HuggingFace repo ID
|
||||
if "/" in checkpoint_path and not checkpoint_path.endswith((".pt", ".pth", ".safetensors")):
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
local_path = hf_hub_download(
|
||||
repo_id=checkpoint_path,
|
||||
filename="model.safetensors",
|
||||
)
|
||||
return load_file(local_path, device=device)
|
||||
|
||||
raise ValueError(f"Checkpoint not found: {checkpoint_path}")
|
||||
@@ -0,0 +1,33 @@
|
||||
"""Common Python utility functions."""
|
||||
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
|
||||
|
||||
def exists(val: Any) -> bool:
|
||||
"""Check if value is not None."""
|
||||
return val is not None
|
||||
|
||||
|
||||
def default(val: Any, d: Any) -> Any:
|
||||
"""Return val if not None, else default (or call it if callable)."""
|
||||
if exists(val):
|
||||
return val
|
||||
return d() if callable(d) else d
|
||||
|
||||
|
||||
def cast_tuple(val: Any, length: Optional[int] = None) -> Tuple:
|
||||
"""Convert value to tuple with optional length validation."""
|
||||
if isinstance(val, list):
|
||||
val = tuple(val)
|
||||
|
||||
output = val if isinstance(val, tuple) else ((val,) * default(length, 1))
|
||||
|
||||
if exists(length):
|
||||
assert len(output) == length
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def compact(input_dict: Dict) -> Dict:
|
||||
"""Filter None values from dictionary."""
|
||||
return {key: value for key, value in input_dict.items() if exists(value)}
|
||||
@@ -0,0 +1,18 @@
|
||||
"""Keypoint utilities for pose detection."""
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def get_dummy_dw_keypoints() -> dict:
|
||||
"""
|
||||
Get dummy DWPose keypoints dictionary for flat-lay garments.
|
||||
|
||||
Returns a pose dictionary with all keypoints set to -1 to indicate
|
||||
no person is present (used for flat-lay garment images).
|
||||
|
||||
Returns:
|
||||
Dictionary with 'bodies' key containing dummy keypoints
|
||||
"""
|
||||
pose = {}
|
||||
pose["bodies"] = {"candidate": (-1) * np.ones((18, 2)), "subset": -1 * np.ones((1, 18))}
|
||||
return pose
|
||||
@@ -0,0 +1,55 @@
|
||||
"""Logging utilities for FASHN VTON with colored console output."""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
from .common import exists
|
||||
|
||||
|
||||
class CustomFormatter(logging.Formatter):
|
||||
COLORS = {
|
||||
"DEBUG": "\033[94m",
|
||||
"INFO": "\033[0m",
|
||||
"WARNING": "\033[93m",
|
||||
"ERROR": "\033[91m",
|
||||
"CRITICAL": "\033[1;91m",
|
||||
}
|
||||
RESET = "\033[0m"
|
||||
|
||||
def __init__(self, timestamp: bool = False, datefmt: str = "%Y-%m-%d %H:%M:%S"):
|
||||
fmt = "%(name)s - %(levelname)s - %(message)s"
|
||||
if timestamp:
|
||||
fmt = "%(asctime)s - " + fmt
|
||||
super().__init__(fmt, datefmt)
|
||||
|
||||
def format(self, record: logging.LogRecord) -> str:
|
||||
original_msg = record.msg
|
||||
if isinstance(original_msg, dict):
|
||||
record.msg = json.dumps(original_msg, indent=4, sort_keys=True)
|
||||
else:
|
||||
record.msg = original_msg
|
||||
|
||||
formatted_msg = super().format(record)
|
||||
levelname = record.levelname
|
||||
color_prefix = self.COLORS.get(levelname, self.COLORS["INFO"])
|
||||
return color_prefix + formatted_msg + self.RESET
|
||||
|
||||
|
||||
def setup_logger(
|
||||
name: str, timestamp: bool = False, level: Optional[int] = None
|
||||
) -> logging.Logger:
|
||||
logger = logging.getLogger(name)
|
||||
|
||||
if exists(level):
|
||||
logger.setLevel(level)
|
||||
|
||||
if not logger.handlers:
|
||||
handler = logging.StreamHandler()
|
||||
formatter = CustomFormatter(timestamp=timestamp)
|
||||
handler.setFormatter(formatter)
|
||||
logger.addHandler(handler)
|
||||
|
||||
logger.propagate = False
|
||||
|
||||
return logger
|
||||
@@ -0,0 +1,43 @@
|
||||
"""Sampling utilities for Rectified Flow inference."""
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def time_shift(mu: float, sigma: float, t: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Apply time shift to timesteps for flow matching schedule.
|
||||
|
||||
Args:
|
||||
mu: Time shift parameter (controls schedule steepness)
|
||||
sigma: Sigma parameter (typically 1.0)
|
||||
t: Timestep tensor with values in (0, 1]
|
||||
|
||||
Returns:
|
||||
Shifted timesteps
|
||||
"""
|
||||
return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma)
|
||||
|
||||
|
||||
def get_rf_schedule(num_steps: int, mu: float = 1.5, reverse: bool = True) -> list[float]:
|
||||
"""
|
||||
Generate timestep schedule for Rectified Flow sampling.
|
||||
|
||||
Creates a shifted linear schedule that provides better sample quality
|
||||
by spending more time at higher noise levels.
|
||||
|
||||
Args:
|
||||
num_steps: Number of sampling steps
|
||||
mu: Time shift parameter (higher = more time at high noise)
|
||||
reverse: If True, returns schedule from t=0 to t=1 (for denoising)
|
||||
|
||||
Returns:
|
||||
List of timesteps of length num_steps + 1
|
||||
"""
|
||||
if reverse:
|
||||
mu = -mu
|
||||
timesteps = torch.linspace(1, 0, num_steps + 1)
|
||||
timesteps = time_shift(mu, 1.0, timesteps)
|
||||
timesteps = timesteps.tolist()
|
||||
return timesteps[::-1] if reverse else timesteps
|
||||
@@ -0,0 +1,77 @@
|
||||
"""PyTorch tensor utilities for image processing."""
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
from torchvision.transforms.functional import to_pil_image
|
||||
|
||||
|
||||
def numpy_to_torch(img: np.ndarray) -> torch.Tensor:
|
||||
"""
|
||||
Convert numpy image to torch tensor.
|
||||
|
||||
For 3D arrays (H, W, C), permutes to (C, H, W).
|
||||
For 2D arrays (H, W), passes through unchanged.
|
||||
|
||||
Args:
|
||||
img: Input numpy array of shape (H, W, C) or (H, W)
|
||||
|
||||
Returns:
|
||||
Torch tensor of shape (C, H, W) or (H, W)
|
||||
"""
|
||||
t = torch.from_numpy(img)
|
||||
if t.ndim == 3:
|
||||
t = t.permute(2, 0, 1)
|
||||
return t
|
||||
|
||||
|
||||
def normalize_uint8_to_neg1_1(x: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Normalize uint8 image tensor from [0, 255] to [-1, 1] range.
|
||||
|
||||
Args:
|
||||
x: Input tensor with values in [0, 255]
|
||||
|
||||
Returns:
|
||||
Normalized tensor with values in [-1, 1]
|
||||
"""
|
||||
return x / 127.5 - 1.0
|
||||
|
||||
|
||||
def _neg1_1_to_0_1(normed_img: torch.Tensor) -> torch.Tensor:
|
||||
"""Convert [-1, 1] normalized tensor to [0, 1] range."""
|
||||
return (normed_img + 1) * 0.5
|
||||
|
||||
|
||||
def tensor_to_pil(img: torch.Tensor, unnormalize: bool = False) -> Image.Image:
|
||||
"""
|
||||
Convert PyTorch tensor to PIL Image.
|
||||
|
||||
Args:
|
||||
img: Input tensor of shape (C, H, W)
|
||||
unnormalize: If True, convert from [-1, 1] to [0, 1] range first
|
||||
|
||||
Returns:
|
||||
PIL Image
|
||||
"""
|
||||
if unnormalize:
|
||||
img = _neg1_1_to_0_1(img)
|
||||
return to_pil_image(img)
|
||||
|
||||
|
||||
def unpack_images(x: torch.Tensor, patch_size: int = 2) -> torch.Tensor:
|
||||
"""
|
||||
Unpack image patches back to full images.
|
||||
|
||||
Used after transformer processing to convert patch representations
|
||||
back to spatial images.
|
||||
|
||||
Args:
|
||||
x: Tensor of shape (batch_size, channels * patch_size^2, h, w)
|
||||
patch_size: Size of patches used during packing
|
||||
|
||||
Returns:
|
||||
Tensor of shape (batch_size, channels, h * patch_size, w * patch_size)
|
||||
"""
|
||||
return rearrange(x, "b (c p1 p2) h w -> b c (h p1) (w p2)", p1=patch_size, p2=patch_size)
|
||||
@@ -0,0 +1,153 @@
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
import folder_paths
|
||||
import comfy.utils
|
||||
import comfy.model_management
|
||||
from .fashn_vton import TryOnPipeline
|
||||
|
||||
model_list = [
|
||||
'fashn-ai/fashn-vton-1.5'
|
||||
]
|
||||
|
||||
class FashnVtonLoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": (model_list, {"default": 'fashn-ai/fashn-vton-1.5'})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("FASHN_VTON_PIPELINE",)
|
||||
RETURN_NAMES = ("pipeline",)
|
||||
FUNCTION = "load_pipeline"
|
||||
CATEGORY = "FashnAI"
|
||||
|
||||
def load_pipeline(self, model):
|
||||
weights_name = "fashn-vton"
|
||||
|
||||
base_weights_dir = os.path.join(folder_paths.models_dir, weights_name)
|
||||
|
||||
os.makedirs(base_weights_dir, exist_ok=True)
|
||||
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
# Download TryOnModel
|
||||
tryon_path = os.path.join(base_weights_dir, "model.safetensors")
|
||||
if not os.path.exists(tryon_path):
|
||||
print(f"FashnVTON: Downloading TryOnModel weights to {tryon_path}...")
|
||||
hf_hub_download(
|
||||
repo_id=model,
|
||||
filename="model.safetensors",
|
||||
local_dir=base_weights_dir,
|
||||
)
|
||||
|
||||
# Download DWPose
|
||||
dwpose_dir = os.path.join(base_weights_dir, "dwpose")
|
||||
os.makedirs(dwpose_dir, exist_ok=True)
|
||||
for filename in ["yolox_l.onnx", "dw-ll_ucoco_384.onnx"]:
|
||||
if not os.path.exists(os.path.join(dwpose_dir, filename)):
|
||||
print(f"FashnVTON: Downloading DWPose/{filename} to {dwpose_dir}...")
|
||||
hf_hub_download(
|
||||
repo_id="fashn-ai/DWPose",
|
||||
filename=filename,
|
||||
local_dir=dwpose_dir,
|
||||
)
|
||||
|
||||
# Initialize Pipeline
|
||||
print(f"FashnVTON: Loading pipeline from {base_weights_dir}...")
|
||||
pipeline = TryOnPipeline(weights_dir=base_weights_dir)
|
||||
|
||||
return (pipeline,)
|
||||
|
||||
class FashnVtonInference:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"pipeline": ("FASHN_VTON_PIPELINE",),
|
||||
"person_image": ("IMAGE",),
|
||||
"garment_image": ("IMAGE",),
|
||||
"category": (["tops", "bottoms", "one-pieces"], {"default": "tops"}),
|
||||
"num_timesteps": ("INT", {"default": 30, "min": 1, "max": 100, "step": 1}),
|
||||
"guidance_scale": ("FLOAT", {"default": 2.0, "min": 1.0, "max": 10.0, "step": 0.1}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"keep_model_loaded": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "FashnAI"
|
||||
|
||||
def process(self, pipeline, person_image, garment_image, category, num_timesteps, guidance_scale, seed, keep_model_loaded):
|
||||
|
||||
device = comfy.model_management.get_torch_device()
|
||||
|
||||
print(f"FashnVTON: Moving models to {device}...")
|
||||
if hasattr(pipeline, "tryon_model"):
|
||||
pipeline.tryon_model.to(device)
|
||||
if hasattr(pipeline, "hp_model"):
|
||||
if hasattr(pipeline.hp_model, "model"):
|
||||
pipeline.hp_model.model.to(device)
|
||||
|
||||
pbar = comfy.utils.ProgressBar(num_timesteps)
|
||||
|
||||
def progress_callback(step, total_steps):
|
||||
pbar.update_absolute(step + 1, total_steps)
|
||||
|
||||
# ComfyUI images are (B, H, W, C) tensors in [0, 1]
|
||||
def tensor_to_pil(tensor):
|
||||
img = tensor[0].cpu().numpy()
|
||||
img = (img * 255).astype(np.uint8)
|
||||
return Image.fromarray(img)
|
||||
|
||||
person_pil = tensor_to_pil(person_image)
|
||||
garment_pil = tensor_to_pil(garment_image)
|
||||
|
||||
seed = seed % (2**32)
|
||||
|
||||
try:
|
||||
result = pipeline(
|
||||
person_image=person_pil,
|
||||
garment_image=garment_pil,
|
||||
category=category,
|
||||
num_timesteps=num_timesteps,
|
||||
guidance_scale=guidance_scale,
|
||||
seed=seed,
|
||||
callback=progress_callback,
|
||||
)
|
||||
finally:
|
||||
# Handle Offloading
|
||||
if not keep_model_loaded:
|
||||
print("FashnVTON: Unloading models from VRAM...")
|
||||
|
||||
if hasattr(pipeline, "tryon_model"):
|
||||
pipeline.tryon_model.to("cpu")
|
||||
if hasattr(pipeline, "hp_model"):
|
||||
if hasattr(pipeline.hp_model, "model"):
|
||||
pipeline.hp_model.model.to("cpu")
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.ipc_collect()
|
||||
|
||||
comfy.model_management.soft_empty_cache()
|
||||
|
||||
# Convert back to ComfyUI format (B, H, W, C)
|
||||
output_img = np.array(result.images[0]).astype(np.float32) / 255.0
|
||||
output_tensor = torch.from_numpy(output_img).unsqueeze(0)
|
||||
|
||||
return (output_tensor,)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"FashnVtonLoader": FashnVtonLoader,
|
||||
"FashnVtonInference": FashnVtonInference,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FashnVtonLoader": "(Down)load Fashn VTON",
|
||||
"FashnVtonInference": "Fashn VTON Inference",
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
[project]
|
||||
name = "comfyui-fashn-vton"
|
||||
version = "1.0.0"
|
||||
description = "Implements the FASHN VTON v1.5 model for virtual try-on in ComfyUI"
|
||||
readme = "README.md"
|
||||
license = {file = "LICENSE"}
|
||||
classifiers = []
|
||||
dependencies = []
|
||||
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/drphero/ComfyUI-FASHN-VTON"
|
||||
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "drphero"
|
||||
DisplayName = "ComfyUI-FASHN-VTON"
|
||||
@@ -0,0 +1,10 @@
|
||||
safetensors>=0.3.0
|
||||
huggingface_hub>=0.20.0
|
||||
pillow>=9.0.0
|
||||
numpy>=1.21.0
|
||||
opencv-python>=4.5.0
|
||||
tqdm>=4.65.0
|
||||
einops>=0.6.0
|
||||
onnxruntime-gpu
|
||||
matplotlib>=3.5.0
|
||||
fashn-human-parser>=0.1.1
|
||||
Reference in New Issue
Block a user