Initial commit

This commit is contained in:
drphero
2026-01-29 22:31:58 +01:00
commit 643c253a8f
47 changed files with 3218 additions and 0 deletions
+2
View File
@@ -0,0 +1,2 @@
.DS_Store
.DS_Store ?
+191
View File
@@ -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.
+74
View File
@@ -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.
+3
View File
@@ -0,0 +1,3 @@
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+13
View File
@@ -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.
+16
View File
@@ -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.
+131
View File
@@ -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
+131
View File
@@ -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
+364
View File
@@ -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
+227
View File
@@ -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
+51
View File
@@ -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
+343
View File
@@ -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)
+22
View File
@@ -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",
]
+212
View File
@@ -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
+163
View File
@@ -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
+222
View File
@@ -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)
+563
View File
@@ -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}
+40
View File
@@ -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.
+44
View File
@@ -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}")
+33
View File
@@ -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)}
+18
View File
@@ -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
+55
View File
@@ -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
+43
View File
@@ -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
+77
View File
@@ -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)
+153
View File
@@ -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",
}
+17
View File
@@ -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"
+10
View File
@@ -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