commit 643c253a8fce2900993842584daeb330ee96a464 Author: drphero <14174687+drphero@users.noreply.github.com> Date: Thu Jan 29 22:31:58 2026 +0100 Initial commit diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..9923533 --- /dev/null +++ b/.gitignore @@ -0,0 +1,2 @@ +.DS_Store +.DS_Store ? \ No newline at end of file diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..d095272 --- /dev/null +++ b/LICENSE @@ -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. \ No newline at end of file diff --git a/README.md b/README.md new file mode 100644 index 0000000..7e1db6b --- /dev/null +++ b/README.md @@ -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. diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..39a8c6b --- /dev/null +++ b/__init__.py @@ -0,0 +1,3 @@ +from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/fashn_vton/__init__.py b/fashn_vton/__init__.py new file mode 100644 index 0000000..a3204a1 --- /dev/null +++ b/fashn_vton/__init__.py @@ -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__", +] diff --git a/fashn_vton/__pycache__/__init__.cpython-312.pyc b/fashn_vton/__pycache__/__init__.cpython-312.pyc new file mode 100644 index 0000000..9d8a700 Binary files /dev/null and b/fashn_vton/__pycache__/__init__.cpython-312.pyc differ diff --git a/fashn_vton/__pycache__/pipeline.cpython-312.pyc b/fashn_vton/__pycache__/pipeline.cpython-312.pyc new file mode 100644 index 0000000..6dd8882 Binary files /dev/null and b/fashn_vton/__pycache__/pipeline.cpython-312.pyc differ diff --git a/fashn_vton/__pycache__/tryon_mmdit.cpython-312.pyc b/fashn_vton/__pycache__/tryon_mmdit.cpython-312.pyc new file mode 100644 index 0000000..5cdee3f Binary files /dev/null and b/fashn_vton/__pycache__/tryon_mmdit.cpython-312.pyc differ diff --git a/fashn_vton/dwpose/__init__.py b/fashn_vton/dwpose/__init__.py new file mode 100644 index 0000000..4566b7b --- /dev/null +++ b/fashn_vton/dwpose/__init__.py @@ -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 diff --git a/fashn_vton/dwpose/__pycache__/__init__.cpython-312.pyc b/fashn_vton/dwpose/__pycache__/__init__.cpython-312.pyc new file mode 100644 index 0000000..bc57a1e Binary files /dev/null and b/fashn_vton/dwpose/__pycache__/__init__.cpython-312.pyc differ diff --git a/fashn_vton/dwpose/__pycache__/dwpose.cpython-312.pyc b/fashn_vton/dwpose/__pycache__/dwpose.cpython-312.pyc new file mode 100644 index 0000000..3b29447 Binary files /dev/null and b/fashn_vton/dwpose/__pycache__/dwpose.cpython-312.pyc differ diff --git a/fashn_vton/dwpose/__pycache__/onnxdet.cpython-312.pyc b/fashn_vton/dwpose/__pycache__/onnxdet.cpython-312.pyc new file mode 100644 index 0000000..7969252 Binary files /dev/null and b/fashn_vton/dwpose/__pycache__/onnxdet.cpython-312.pyc differ diff --git a/fashn_vton/dwpose/__pycache__/onnxpose.cpython-312.pyc b/fashn_vton/dwpose/__pycache__/onnxpose.cpython-312.pyc new file mode 100644 index 0000000..ee4d321 Binary files /dev/null and b/fashn_vton/dwpose/__pycache__/onnxpose.cpython-312.pyc differ diff --git a/fashn_vton/dwpose/__pycache__/utils.cpython-312.pyc b/fashn_vton/dwpose/__pycache__/utils.cpython-312.pyc new file mode 100644 index 0000000..8b30ceb Binary files /dev/null and b/fashn_vton/dwpose/__pycache__/utils.cpython-312.pyc differ diff --git a/fashn_vton/dwpose/__pycache__/wholebody.cpython-312.pyc b/fashn_vton/dwpose/__pycache__/wholebody.cpython-312.pyc new file mode 100644 index 0000000..8cc2af0 Binary files /dev/null and b/fashn_vton/dwpose/__pycache__/wholebody.cpython-312.pyc differ diff --git a/fashn_vton/dwpose/dwpose.py b/fashn_vton/dwpose/dwpose.py new file mode 100644 index 0000000..1e2c21a --- /dev/null +++ b/fashn_vton/dwpose/dwpose.py @@ -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 diff --git a/fashn_vton/dwpose/onnxdet.py b/fashn_vton/dwpose/onnxdet.py new file mode 100644 index 0000000..e2b0abd --- /dev/null +++ b/fashn_vton/dwpose/onnxdet.py @@ -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 diff --git a/fashn_vton/dwpose/onnxpose.py b/fashn_vton/dwpose/onnxpose.py new file mode 100644 index 0000000..a5c1ff2 --- /dev/null +++ b/fashn_vton/dwpose/onnxpose.py @@ -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 diff --git a/fashn_vton/dwpose/utils.py b/fashn_vton/dwpose/utils.py new file mode 100644 index 0000000..31f3b2b --- /dev/null +++ b/fashn_vton/dwpose/utils.py @@ -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 diff --git a/fashn_vton/dwpose/wholebody.py b/fashn_vton/dwpose/wholebody.py new file mode 100644 index 0000000..0331a53 --- /dev/null +++ b/fashn_vton/dwpose/wholebody.py @@ -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 diff --git a/fashn_vton/pipeline.py b/fashn_vton/pipeline.py new file mode 100644 index 0000000..6a90ecb --- /dev/null +++ b/fashn_vton/pipeline.py @@ -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) diff --git a/fashn_vton/preprocessing/__init__.py b/fashn_vton/preprocessing/__init__.py new file mode 100644 index 0000000..635d4ec --- /dev/null +++ b/fashn_vton/preprocessing/__init__.py @@ -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", +] diff --git a/fashn_vton/preprocessing/__pycache__/__init__.cpython-312.pyc b/fashn_vton/preprocessing/__pycache__/__init__.cpython-312.pyc new file mode 100644 index 0000000..1dd7794 Binary files /dev/null and b/fashn_vton/preprocessing/__pycache__/__init__.cpython-312.pyc differ diff --git a/fashn_vton/preprocessing/__pycache__/agnostic.cpython-312.pyc b/fashn_vton/preprocessing/__pycache__/agnostic.cpython-312.pyc new file mode 100644 index 0000000..93ad4d8 Binary files /dev/null and b/fashn_vton/preprocessing/__pycache__/agnostic.cpython-312.pyc differ diff --git a/fashn_vton/preprocessing/__pycache__/masks.cpython-312.pyc b/fashn_vton/preprocessing/__pycache__/masks.cpython-312.pyc new file mode 100644 index 0000000..bafd562 Binary files /dev/null and b/fashn_vton/preprocessing/__pycache__/masks.cpython-312.pyc differ diff --git a/fashn_vton/preprocessing/__pycache__/transforms.cpython-312.pyc b/fashn_vton/preprocessing/__pycache__/transforms.cpython-312.pyc new file mode 100644 index 0000000..0fbd629 Binary files /dev/null and b/fashn_vton/preprocessing/__pycache__/transforms.cpython-312.pyc differ diff --git a/fashn_vton/preprocessing/agnostic.py b/fashn_vton/preprocessing/agnostic.py new file mode 100644 index 0000000..09f864f --- /dev/null +++ b/fashn_vton/preprocessing/agnostic.py @@ -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 diff --git a/fashn_vton/preprocessing/masks.py b/fashn_vton/preprocessing/masks.py new file mode 100644 index 0000000..43c4f6b --- /dev/null +++ b/fashn_vton/preprocessing/masks.py @@ -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 diff --git a/fashn_vton/preprocessing/transforms.py b/fashn_vton/preprocessing/transforms.py new file mode 100644 index 0000000..b03aae3 --- /dev/null +++ b/fashn_vton/preprocessing/transforms.py @@ -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) diff --git a/fashn_vton/tryon_mmdit.py b/fashn_vton/tryon_mmdit.py new file mode 100644 index 0000000..fab0ac3 --- /dev/null +++ b/fashn_vton/tryon_mmdit.py @@ -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} diff --git a/fashn_vton/utils/__init__.py b/fashn_vton/utils/__init__.py new file mode 100644 index 0000000..2757c66 --- /dev/null +++ b/fashn_vton/utils/__init__.py @@ -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", +] diff --git a/fashn_vton/utils/__pycache__/__init__.cpython-312.pyc b/fashn_vton/utils/__pycache__/__init__.cpython-312.pyc new file mode 100644 index 0000000..21fd330 Binary files /dev/null and b/fashn_vton/utils/__pycache__/__init__.cpython-312.pyc differ diff --git a/fashn_vton/utils/__pycache__/checkpoint.cpython-312.pyc b/fashn_vton/utils/__pycache__/checkpoint.cpython-312.pyc new file mode 100644 index 0000000..1f806d4 Binary files /dev/null and b/fashn_vton/utils/__pycache__/checkpoint.cpython-312.pyc differ diff --git a/fashn_vton/utils/__pycache__/common.cpython-312.pyc b/fashn_vton/utils/__pycache__/common.cpython-312.pyc new file mode 100644 index 0000000..a68ba61 Binary files /dev/null and b/fashn_vton/utils/__pycache__/common.cpython-312.pyc differ diff --git a/fashn_vton/utils/__pycache__/keypoints.cpython-312.pyc b/fashn_vton/utils/__pycache__/keypoints.cpython-312.pyc new file mode 100644 index 0000000..2f699c0 Binary files /dev/null and b/fashn_vton/utils/__pycache__/keypoints.cpython-312.pyc differ diff --git a/fashn_vton/utils/__pycache__/logger.cpython-312.pyc b/fashn_vton/utils/__pycache__/logger.cpython-312.pyc new file mode 100644 index 0000000..53f852b Binary files /dev/null and b/fashn_vton/utils/__pycache__/logger.cpython-312.pyc differ diff --git a/fashn_vton/utils/__pycache__/sampling.cpython-312.pyc b/fashn_vton/utils/__pycache__/sampling.cpython-312.pyc new file mode 100644 index 0000000..1f521bf Binary files /dev/null and b/fashn_vton/utils/__pycache__/sampling.cpython-312.pyc differ diff --git a/fashn_vton/utils/__pycache__/tensor.cpython-312.pyc b/fashn_vton/utils/__pycache__/tensor.cpython-312.pyc new file mode 100644 index 0000000..5e43df5 Binary files /dev/null and b/fashn_vton/utils/__pycache__/tensor.cpython-312.pyc differ diff --git a/fashn_vton/utils/checkpoint.py b/fashn_vton/utils/checkpoint.py new file mode 100644 index 0000000..7155575 --- /dev/null +++ b/fashn_vton/utils/checkpoint.py @@ -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}") diff --git a/fashn_vton/utils/common.py b/fashn_vton/utils/common.py new file mode 100644 index 0000000..06f3716 --- /dev/null +++ b/fashn_vton/utils/common.py @@ -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)} diff --git a/fashn_vton/utils/keypoints.py b/fashn_vton/utils/keypoints.py new file mode 100644 index 0000000..0cc6d69 --- /dev/null +++ b/fashn_vton/utils/keypoints.py @@ -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 diff --git a/fashn_vton/utils/logger.py b/fashn_vton/utils/logger.py new file mode 100644 index 0000000..bb6418d --- /dev/null +++ b/fashn_vton/utils/logger.py @@ -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 diff --git a/fashn_vton/utils/sampling.py b/fashn_vton/utils/sampling.py new file mode 100644 index 0000000..b42dee0 --- /dev/null +++ b/fashn_vton/utils/sampling.py @@ -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 diff --git a/fashn_vton/utils/tensor.py b/fashn_vton/utils/tensor.py new file mode 100644 index 0000000..046aab4 --- /dev/null +++ b/fashn_vton/utils/tensor.py @@ -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) diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..a5903fc --- /dev/null +++ b/nodes.py @@ -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", +} \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..4ecd92e --- /dev/null +++ b/pyproject.toml @@ -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" \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..503c51b --- /dev/null +++ b/requirements.txt @@ -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