diff --git a/.gitattributes b/.gitattributes index a6344aac8c09253b3b630fb776ae94478aa0275b..26ab6f9c1e9a522b6ece1fd16738e1f7e42f77a6 100644 --- a/.gitattributes +++ b/.gitattributes @@ -33,3 +33,30 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text *.zip filter=lfs diff=lfs merge=lfs -text *.zst filter=lfs diff=lfs merge=lfs -text *tfevents* filter=lfs diff=lfs merge=lfs -text +SCAIL-Pose/resources/animation.png filter=lfs diff=lfs merge=lfs -text +SCAIL-Pose/resources/data.png filter=lfs diff=lfs merge=lfs -text +SCAIL-Pose/resources/pose_comp.png filter=lfs diff=lfs merge=lfs -text +SCAIL-Pose/resources/pose_result.png filter=lfs diff=lfs merge=lfs -text +SCAIL-Pose/resources/pose_teaser.png filter=lfs diff=lfs merge=lfs -text +SCAIL-Pose/resources/preteaser.png filter=lfs diff=lfs merge=lfs -text +SCAIL-Pose/resources/teaser.png filter=lfs diff=lfs merge=lfs -text +examples/animation_001/combined.gif filter=lfs diff=lfs merge=lfs -text +examples/animation_001/driving.mp4 filter=lfs diff=lfs merge=lfs -text +examples/animation_001/ref.jpg filter=lfs diff=lfs merge=lfs -text +examples/animation_001/rendered_mask_v2.mp4 filter=lfs diff=lfs merge=lfs -text +examples/animation_001/rendered_v2.mp4 filter=lfs diff=lfs merge=lfs -text +examples/animation_001_posedriven/combined.gif filter=lfs diff=lfs merge=lfs -text +examples/animation_001_posedriven/driving.mp4 filter=lfs diff=lfs merge=lfs -text +examples/animation_001_posedriven/ref.jpg filter=lfs diff=lfs merge=lfs -text +examples/animation_001_posedriven/rendered_mask_v2.mp4 filter=lfs diff=lfs merge=lfs -text +examples/animation_001_posedriven/rendered_v2.mp4 filter=lfs diff=lfs merge=lfs -text +examples/animation_002/driving.mp4 filter=lfs diff=lfs merge=lfs -text +examples/animation_002/ref.jpg filter=lfs diff=lfs merge=lfs -text +examples/animation_002/rendered_mask_v2.mp4 filter=lfs diff=lfs merge=lfs -text +examples/animation_002/rendered_v2.mp4 filter=lfs diff=lfs merge=lfs -text +examples/animation_003_multi_ref/background.png filter=lfs diff=lfs merge=lfs -text +examples/animation_003_multi_ref/character_0.png filter=lfs diff=lfs merge=lfs -text +examples/animation_003_multi_ref/character_1.png filter=lfs diff=lfs merge=lfs -text +examples/animation_003_multi_ref/driving.mp4 filter=lfs diff=lfs merge=lfs -text +examples/animation_003_multi_ref/ref.png filter=lfs diff=lfs merge=lfs -text +examples/animation_003_multi_ref/ref_mask.jpg filter=lfs diff=lfs merge=lfs -text diff --git a/.gitmodules b/.gitmodules new file mode 100644 index 0000000000000000000000000000000000000000..53bcd98494cbbf6b5a6177be0a68cedcf443746f --- /dev/null +++ b/.gitmodules @@ -0,0 +1,3 @@ +[submodule "SCAIL-Pose"] + path = SCAIL-Pose + url = https://github.com/zai-org/SCAIL-Pose diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..fd957948ecf9f2f385c89bb58110bc6031ca5c39 --- /dev/null +++ b/LICENSE @@ -0,0 +1,201 @@ + 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 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 legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, 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 + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright 2026 Zhipu 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. diff --git a/ORIGINAL_README.md b/ORIGINAL_README.md new file mode 100644 index 0000000000000000000000000000000000000000..d258d117c727ae398b3a4d8fe1405370ee2ed8c9 --- /dev/null +++ b/ORIGINAL_README.md @@ -0,0 +1,434 @@ +

SCAIL-2: Unifying Controlled Character Animation with End-to-end In-Context Conditioning

+ + +
+ + arXiv + + + HuggingFace + + + Project Page + + + Datasets + +
+ + + +This repository contains the official implementation code of SCAIL-2: Unifying Controlled Character Animation with End-to-end In-Context Conditioning. The code is for the inference of SCAIL-2 Model, an open-source model to support **End-to-End** Character Animation. + + +

+ Teaser +

+ +## 🔎 Introduction +SCAIL-1 identifies the key bottlenecks that hinder character animation towards production level: how to represent the pose and how to inject the pose. However, the reliance on intermediate pose representation still hinders the model towards complex motion and generalizable identity. We define the issue as over reliance on intermediates. + +As intermediates, skeleton maps suffer from inherent ambiguity under complex scenarios. Further, it restricts the driving source to be exocentric human movements and thus cannot handle driving sources like animals. Character replacement and multi-character animation suffers from similar issues, where state-of-the-art methods use inpainting masks, but such masks are still a form of intermediates and limits the application and bounds the performance. + + +

+ preteaser +

+ + + +To bypass intermediate pose representation, we utilize several off-the-shelf models, including [SCAIL-Preview](https://github.com/zai-org/SCAIL), [Wan-Animate](https://github.com/Wan-Video/Wan2.2), [MoCha](https://github.com/Orange-3DV-Team/MoCha) to synthesize 60K motion pairs. By designing a Unified Motion Transfer Interface containing 2 type of masking channels and a dedicated RoPE design, we support training with all those data. We utilize **reserve driving**, so that the model can learn capabilities beyond those models. From the data composition and the training recipe, the final model yield emergent capabilities. For example, it supports cross-identity replacement, animal-driving scenarios, and support more advanced control intermediate like [SAM3D-Body](https://github.com/facebookresearch/sam-3d-body)'s mesh rendering in zero-shot manner. + +

+ pipeline +

+ +

+ Teaser +

+ +We model the bias of pose-driven generators as preference and introduce Bias-Aware DPO, a novel mechanisim to further improve details. The DPO LoRA is also released at HuggingFace and can be enabled in this repo as well as ComfyUI implementations. + + +## 🎨 Community Works + +❤️ We thank the community for sharing their amazing creations! Special thanks to Ablejones (Discord), 机智波, 肥猴, 绘篇AI手艺 (Bilibili), Fuzzy-Mastodon (Reddit). Audio comes from reference videos. + + + + + + + + + + + +
+ + +## 🚀 Getting Started + + +### Mask Semantics +The mask is a critical input to SCAIL-2 even in Animation Mode. To visualize the channels that the mask is actually for, we encode them with colors: + +- **Black** — tells the model the background at this location should *not* be visible. +- **White** — tells the model the background at this location *should* be visible. +- **Color** — encodes the correspondence between character regions and the driving motion. + +Animation mode (end-to-end) example (left: reference mask, right: driving mask): + +

+ animation mask example +

+ +

+ multi animation mask example +

+ + +Animation mode (pose-driven) example (left: reference mask, right: driving mask): + +

+ pose-driven animation mask example +

+ +Replacement mode example (left: reference mask, right: driving mask): + +

+ replacement mask example +

+ +Without a correct mask Animation mode collapse into Replacement-Mode behavior in certain inputs. + +The masks also enable zero-shot multi-reference generation, where additional visual inputs provide information that single reference may not cover, such as back view, close-up view and occluded background. According to the color assignment logic, in multi-reference the following inputs get the corresponding masks as shown below: + + + + + + + + + + + + + + +
+ + +### Checkpoints Download + +| ckpts | Download Link | Notes | +|--------------|------------------------------------------------------------------------------------------------------------------------------|-------------------------------| +| SCAIL-2 | [🤗 Hugging Face](https://huggingface.co/zai-org/SCAIL-2)
[🤖 ModelScope](https://modelscope.cn/models/ZhipuAI/SCAIL-2) | Trained with mixed resolutions and fps.
End-to-end driven supports both 512p and 704p.
Pose-driven performs better under 704p.
H and W should be both divisible by 32
(e.g. 704*1280) if using other resolutions. | + +Use the following commands to download the model weights +(We have integrated both Wan VAE and T5 modules into this checkpoint for convenience). + +```bash +hf download zai-org/SCAIL-2 +``` +The files should be organized like: +``` +SCAIL-2/ +├── Wan2.1_VAE.pth +├── model +│ ├── 1 +│ │ └── fsdp2_rank_0000_checkpoint.pt +│ └── latest +└── umt5-xxl + ├── ... +``` + +The model weights are intended for `sat` branch, for usage in `wan` branch, convert to `safetensors` format: +```bash +python convert.py --scail-dir /path/to/SCAIL-2 --save-path /path/to/SCAIL-2.safetensors +``` + +### Environment Setup +Please make sure your Python version is between 3.10 and 3.12, inclusive of both 3.10 and 3.12. +``` +pip install -r requirements.txt +``` + + + + +### Input Preparation + +`SCAIL-Pose` contains the preprocessing code used to prepare SCAIL-2 inputs, including pose extraction, pose rendering, reference masks, and driving-video masks. It can prepare both animation inputs and character replacement inputs. The submodule should live under the project root: + +``` +SCAIL-2/ +├── generate.py +├── examples/ +├── SCAIL-Pose/ +└── ... +``` + +After cloning this repository, initialize the submodule: + +```bash +git submodule update --init --recursive +``` + +Enter the submodule and follow its environment setup. `SCAIL-Pose` recommends an OpenMMLab/MMPose environment, then installing its own requirements: + +```bash +cd SCAIL-Pose +pip install -r requirements.txt +``` + +Download the pose-preprocessing weights inside `SCAIL-Pose/pretrained_weights`. The required layout is: + +``` +pretrained_weights/ +├── nlf_l_multi_0.3.2.torchscript +└── DWPose/ + ├── dw-ll_ucoco_384.onnx + └── yolox_l.onnx +``` + +For SCAIL-2 animation, `SCAIL-Pose` provides an all-in-one preprocessing entrypoint: + +```bash +# Recommended end-to-end mode: rendered_v2.mp4 is the driving video copy, +# and the mask video is generated from SAM3 masks. +python NLFPoseExtract/process_animation_aio.py --subdir /path/to/input --e2e_mode + +# Pose-driven mode: runs NLF + DWPose and writes a skeleton render. +python NLFPoseExtract/process_animation_aio.py --subdir /path/to/input +``` + +For character replacement, use: + +```bash +python NLFPoseExtract/process_replacement.py --subdir /path/to/input + +# If the driving video has multiple people and only one should be replaced: +python NLFPoseExtract/process_replacement.py --subdir /path/to/input --matchnearest +``` + +The preprocessing outputs are written back to the example folder and can be passed to `generate.py` as `--image`, `--mask_image`, `--pose`, and `--mask_video`. + + +## 🦾 Usage +### Generate Input Conditions + +`generate.py` runs one SCAIL-2 inference job from four local input files: + +``` +examples/001/ +├── ref.jpg # reference character image +├── ref_mask.jpg # foreground mask of the reference image +├── rendered_v2.mp4 # driving / pose video consumed by --pose +└── rendered_mask_v2.mp4 # per-frame driving mask consumed by --mask_video +``` + +The paths passed to `--image`, `--mask_image`, `--pose`, and `--mask_video` must exist. The script checks them before loading the image/video data. + +For animation mode, `--pose` can be an end-to-end driving video or a pose-rendered video, depending on how the sample was prepared. `--mask_video` should be the corresponding per-frame foreground/control mask. For replacement mode, pass `--replace_flag` and provide the replacement-region mask through `--mask_video`. + +### Prompt Semantics + +For both animation and character replacement, `--prompt` should describe the generated video itself. It should not be an instruction to the model. + +For replacement tasks, the prompt should describe the video after replacement has already happened. For better results, describe the replacement character's visible clothing and appearance, and include objects the character interacts with or stays close to in the video, such as tools, instruments, chairs, tables, vehicles, doors, or handheld items. + +### Character Replacement Prompt Enhancer + +We provide an optional Gemini-based helper, `prompt_enhancer.py`, to turn a short replacement instruction into a positive prompt for `generate.py`. The helper samples frames from the source video, reads the replacement reference image, uses few-shot examples from `prompt_examples.txt`, and outputs a long English description of the replaced video. + +`google-genai` is not installed by default in `requirements.txt`. Install it before using the enhancer: + +```bash +pip install google-genai +``` + +Set a Gemini API key before running. + +```bash +export GEMINI_API_KEY=your_api_key +``` + +Example: + +```bash +python prompt_enhancer.py \ + --video /path/to/driving.mp4 \ + --image /path/to/ref.png \ + --instruction "replace the man in the blue jacket in the video with the person in the image" \ + --examples prompt_examples.txt \ + --num_frames 8 \ + --output enhanced_prompt.txt \ + --caption_out source_caption.txt +``` + +The `--instruction` argument is only for Gemini, so it can say who should be replaced by whom. The file written to `--output` is the positive generated-video description that should be passed to `generate.py --prompt`; the enhancer is instructed to include useful SCAIL-2 prompt details such as the replacement character's clothing and objects the character interacts with. + +Use the enhanced prompt for replacement inference: + +```bash +python generate.py \ + --model SCAIL-14B \ + --ckpt_dir /path/to/SCAIL-2 \ + --scail_path /path/to/SCAIL-2.safetensors \ + --replace_flag \ + --target_w 896 --target_h 512 \ + --image /path/to/ref.png \ + --mask_image /path/to/ref_mask.png \ + --pose /path/to/driving.mp4 \ + --mask_video /path/to/replace_mask.mp4 \ + --prompt "$(cat enhanced_prompt.txt)" \ + --save_file replacement_output.mp4 +``` + +`prompt_examples.txt` is used as few-shot style guidance. Add more examples there if you want the enhanced prompts to follow a different level of detail or wording. + +### Single-GPU Inference + +Run inference directly with `generate.py`: + +Example for animation: + +```bash +python generate.py \ + --model SCAIL-14B \ + --ckpt_dir /path/to/SCAIL-2 \ + --scail_path /path/to/SCAIL-2.safetensors \ + --target_w 896 --target_h 512 \ + --image examples/001/ref.jpg \ + --mask_image examples/001/ref_mask.jpg \ + --pose examples/001/rendered_v2.mp4 \ + --mask_video examples/001/rendered_mask_v2.mp4 \ + --prompt "The girl is dancing" \ + --save_file output.mp4 +``` + +Example for replacement: + +```bash +python generate.py \ + --model SCAIL-14B \ + --ckpt_dir /path/to/SCAIL-2 \ + --scail_path /path/to/SCAIL-2.safetensors \ + --target_w 896 --target_h 512 \ + --image examples/replace_001/ref.png \ + --mask_image examples/replace_001/ref_mask.png \ + --pose examples/replace_001/rendered_v2.mp4 \ + --mask_video examples/replace_001/replace_mask.mp4 \ + --prompt "A blond white male wearing a black suit, trousers, and leather shoes is playing the violin on the street while pedestrians walk past him." \ + --save_file output.mp4 \ + --replace_flag +``` + +Useful sampling options: + +- `--sample_steps`: number of denoising steps. Defaults to `40`. +- `--sample_shift`: flow-matching scheduler shift. Defaults to `3.0` if not specified. +- `--sample_guide_scale`: classifier-free guidance scale. Defaults to `5.0`. +- `--sample_solver`: `unipc` or `dpm++`. Defaults to `unipc`. +- `--offload_model`: whether to offload model components between stages. For single-process inference, the default is `True`. + +Note that SCAIL-2 is trained with long, detailed prompts. Short prompts or an empty prompt can run, but detailed descriptions of the reference subject and motion usually produce better results. + + +### LoRA Integrations + +If you use a Lightx2v LoRA checkpoint, pass it with `--lora_path` and set its strength with `--lora_alpha`: + +```bash +python generate.py \ + --model SCAIL-14B \ + --ckpt_dir /path/to/SCAIL-2 \ + --scail_path /path/to/SCAIL-2.safetensors \ + --lora_path Lightx2v/lightx2v_I2V_14B_480p_cfg_step_distill_rank128_bf16.safetensors \ + --lora_alpha 1.0 \ + --sample_steps 8 \ + --sample_shift 1 \ + --sample_guide_scale 1.0 \ + --target_w 896 --target_h 512 \ + --image examples/001/ref.jpg \ + --mask_image examples/001/ref_mask.jpg \ + --pose examples/001/rendered_v2.mp4 \ + --mask_video examples/001/rendered_mask_v2.mp4 \ + --prompt "The girl is dancing" \ + --save_file output.mp4 +``` + +For the DPO LoRA, you can checkout the [`sat-scail2`](https://github.com/zai-org/SCAIL-2/tree/sat-scail2) branch to fully reproduce original results, or convert it into this branch after format matching. The DPO LoRA does not only alleviate hands distortion, but also improved the synchronization of the lips and eyes. If you use ComfyUI, here is a brief comparison showing the effect of the DPO LoRA in Kijai's ComfyUI Workflow: + + +

+ + + + + +### Experimental Functions: Multi-Reference + +SCAIL-2 supports zero-shot multi-reference inference though not optimized for it. Extra references are optional images that provide additional visual evidence, such as another view of the character, a close-up of clothing details, or a clean background reference. Pass them with `--additional_ref_image` and pass one mask for each image with `--additional_ref_mask_image`. The two lists must have the same length and are paired position by position. + +Choose each extra-reference mask according to the mask semantics described above: +- For a clean background reference whose visible background should be preserved, use a **white** mask over the valid background area. If the background is not occluded by the character, a full-white mask is usually appropriate. +- For extra character references where the background is different from the target scene, keep the character/control region in the semantic mask color and make the unrelated background **black**, so the model does not treat that background as visible target content. +- Use consistent mask colors for the same character or region across the main reference, extra references, and driving mask when you want them to refer to the same subject. +The following code provides a simple example of multi-reference inference: + + +```bash +python generate.py \ + --model SCAIL-14B \ + --ckpt_dir /path/to/SCAIL-2 \ + --scail_path /path/to/SCAIL-2.safetensors \ + --target_w 896 --target_h 512 \ + --image examples/animation_003_multi_ref/ref.png \ + --mask_image examples/animation_003_multi_ref/ref_mask.jpg \ + --pose examples/animation_003_multi_ref/rendered_v2.mp4 \ + --mask_video examples/animation_003_multi_ref/rendered_mask_v2.mp4 \ + --additional_ref_image \ + examples/animation_003_multi_ref/background.png \ + examples/animation_003_multi_ref/character_1.png \ + examples/animation_003_multi_ref/character_0.png \ + --additional_ref_mask_image \ + examples/animation_003_multi_ref/background_mask.png \ + examples/animation_003_multi_ref/character_1_mask.png \ + examples/animation_003_multi_ref/character_0_mask.png \ + --prompt "An anime style character with yellow hair, wearing a white and green sailor uniform and a green skirt, is dancing in a warm anime-style classroom." \ + --save_file output_multi_ref.mp4 +``` + +However, as the model is not optimized for such inputs, video qualities may degrade even though additional information do get referenced. To address this, mocking those reference images as videos reduce degradation and artifacts. We specially thanks [wuwukasi](https://github.com/wuwukaka) and [iceage](https://github.com/user2318) for the collaboration to provide empircal results and implementations to support the findings. Check their refined implementations here: [WanAnimatePlus](https://github.com/wuwukaka/ComfyUI-WanAnimatePlus) and [CustomNodeKit](https://github.com/user2318/ComfyUI-CustomNodeKit/), where they will provide their workflows for SCAIL-2's multi-ref mode. + + + + +## 🗃️ Datasets +We provide a large subset of the **MotionPair** dataset used to train SCAIL-2. To request access, please [fill out this form](https://docs.google.com/forms/d/e/1FAIpQLSfZjC0fZmiYFYHg90_79Yl45ipQLfR8ZhOAahOs19nO8nMvxA/viewform?usp=sharing&ouid=108574921907991336711) and agree to the terms of use. If you have not received a reply within a week after submitting the form, feel free to follow up at teal024@foxmail.com. + + +## ✨ Acknowledgements +Our implementation is built upon the foundation of [Wan 2.1](https://github.com/Wan-Video/Wan2.1) and the overall project architecture is inherited from [SCAIL](https://github.com/zai-org/SCAIL). We specially thank [Wan-Animate](https://github.com/Wan-Video/Wan2.2), [MoCha](https://github.com/Orange-3DV-Team/MoCha) as supplement data generators besides SCAIL to make MotionPair-60K. We also thank [HuMo Dataset](https://github.com/Phantom-video/HuMo) as the high-quality source video provider. + +## 📄 Citation + +If you find this work useful in your research, please cite: + +```bibtex +@misc{yan2026scail2, + title={SCAIL-2: Unifying Controlled Character Animation with End-to-end In-Context Conditioning}, + author={Wenhao Yan and Fengjia Guo and Zhuoyi Yang and Jie Tang}, + year={2026}, + eprint={2606.10804}, + archivePrefix={arXiv}, + primaryClass={cs.CV}, + url={https://arxiv.org/abs/2606.10804}, +} +``` + +## 🗝️ License +This project is licensed under the Apache License 2.0 - see the [LICENSE](LICENSE) file for details. diff --git a/SCAIL-Pose/.github/ISSUE_TEMPLATE/bug_report.yaml b/SCAIL-Pose/.github/ISSUE_TEMPLATE/bug_report.yaml new file mode 100644 index 0000000000000000000000000000000000000000..dde7c11197312fd98996d0a83bbed1adb661b61a --- /dev/null +++ b/SCAIL-Pose/.github/ISSUE_TEMPLATE/bug_report.yaml @@ -0,0 +1,72 @@ +name: "\U0001F41B Bug Report" +description: Submit a bug report to help us improve SCAIL-Pose / 提交一个 Bug 问题报告来帮助我们改进 SCAIL-Pose +body: + - type: textarea + id: system-info + attributes: + label: System Info / 系統信息 + description: Your operating environment / 您的运行环境信息 + placeholder: Includes Cuda version, Transformers version, Python version, operating system, hardware information (if you suspect a hardware problem)... / 包括Cuda版本,Transformers版本,Python版本,操作系统,硬件信息(如果您怀疑是硬件方面的问题)... + validations: + required: true + + - type: textarea + id: who-can-help + attributes: + label: Who can help? / 谁可以帮助到您? + description: | + Your issue will be replied to more quickly if you can figure out the right person to tag with @ + All issues are read by one of the maintainers, so if you don't know who to tag, just leave this blank and our maintainer will ping the right person. + + Please tag fewer than 3 people. + + 如果您能找到合适的标签 @,您的问题会更快得到回复。 + 所有问题都会由我们的维护者阅读,如果您不知道该标记谁,只需留空,我们的维护人员会找到合适的开发组成员来解决问题。 + + 标记的人数应该不超过 3 个人。 + + If it's not a bug in these three subsections, you may not specify the helper. Our maintainer will find the right person in the development group to solve the problem. + + 如果不是这三个子版块的bug,您可以不指明帮助者,我们的维护人员会找到合适的开发组成员来解决问题。 + + placeholder: "@Username ..." + + - type: checkboxes + id: information-scripts-examples + attributes: + label: Information / 问题信息 + description: 'The problem arises when using: / 问题出现在' + options: + - label: "The official example scripts / 官方的示例脚本" + - label: "My own modified scripts / 我自己修改的脚本和任务" + + - type: textarea + id: reproduction + validations: + required: true + attributes: + label: Reproduction / 复现过程 + description: | + Please provide a code example that reproduces the problem you encountered, preferably with a minimal reproduction unit. + If you have code snippets, error messages, stack traces, please provide them here as well. + Please format your code correctly using code tags. See https://help.github.com/en/github/writing-on-github/creating-and-highlighting-code-blocks#syntax-highlighting + Do not use screenshots, as they are difficult to read and (more importantly) do not allow others to copy and paste your code. + + 请提供能重现您遇到的问题的代码示例,最好是最小复现单元。 + 如果您有代码片段、错误信息、堆栈跟踪,也请在此提供。 + 请使用代码标签正确格式化您的代码。请参见 https://help.github.com/en/github/writing-on-github/creating-and-highlighting-code-blocks#syntax-highlighting + 请勿使用截图,因为截图难以阅读,而且(更重要的是)不允许他人复制粘贴您的代码。 + placeholder: | + Steps to reproduce the behavior/复现Bug的步骤: + + 1. + 2. + 3. + + - type: textarea + id: expected-behavior + validations: + required: true + attributes: + label: Expected behavior / 期待表现 + description: "A clear and concise description of what you would expect to happen. /简单描述您期望发生的事情。" diff --git a/SCAIL-Pose/.github/ISSUE_TEMPLATE/feature-request.yaml b/SCAIL-Pose/.github/ISSUE_TEMPLATE/feature-request.yaml new file mode 100644 index 0000000000000000000000000000000000000000..f1827d5f89e1c5df90ed6b2af8c7896cbf9bdeb7 --- /dev/null +++ b/SCAIL-Pose/.github/ISSUE_TEMPLATE/feature-request.yaml @@ -0,0 +1,34 @@ +name: "\U0001F680 Feature request" +description: Submit a request for a new SCAIL-Pose feature / 提交一个新的 SCAIL-Pose 的功能建议 +labels: [ "feature" ] +body: + - type: textarea + id: feature-request + validations: + required: true + attributes: + label: Feature request / 功能建议 + description: | + A brief description of the functional proposal. Links to corresponding papers and code are desirable. + 对功能建议的简述。最好提供对应的论文和代码链接 + + - type: textarea + id: motivation + validations: + required: true + attributes: + label: Motivation / 动机 + description: | + Your motivation for making the suggestion. If that motivation is related to another GitHub issue, link to it here. + 您提出建议的动机。如果该动机与另一个 GitHub 问题有关,请在此处提供对应的链接。 + + - type: textarea + id: contribution + validations: + required: true + attributes: + label: Your contribution / 您的贡献 + description: | + + Your PR link or any other link you can help with. + 您的PR链接或者其他您能提供帮助的链接。 diff --git a/SCAIL-Pose/DWPoseProcess/AAUtils.py b/SCAIL-Pose/DWPoseProcess/AAUtils.py new file mode 100644 index 0000000000000000000000000000000000000000..441250d681898b0854d4753489ddc4b0766d5c39 --- /dev/null +++ b/SCAIL-Pose/DWPoseProcess/AAUtils.py @@ -0,0 +1,305 @@ +# MooreAA 同样API +import importlib +import os +import os.path as osp +import shutil +import sys +from pathlib import Path +import decord +from decord import VideoReader, cpu, gpu +import av +import numpy as np +import torch +import torchvision +from einops import rearrange +from PIL import Image +import time +from fractions import Fraction +import cv2 +import jsonlines +import random +import io + + + +def seed_everything(seed): + import random + + import numpy as np + + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + np.random.seed(seed % (2**32)) + random.seed(seed) + + +def import_filename(filename): + spec = importlib.util.spec_from_file_location("mymodule", filename) + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module + + +def delete_additional_ckpt(base_path, num_keep): + dirs = [] + for d in os.listdir(base_path): + if d.startswith("checkpoint-"): + dirs.append(d) + num_tot = len(dirs) + if num_tot <= num_keep: + return + # ensure ckpt is sorted and delete the ealier! + del_dirs = sorted(dirs, key=lambda x: int(x.split("-")[-1]))[: num_tot - num_keep] + for d in del_dirs: + path_to_dir = osp.join(base_path, d) + if osp.exists(path_to_dir): + shutil.rmtree(path_to_dir) + + +# def save_videos_from_pil(pil_images, path, fps=8): +# if fps is None or fps <= 0 or fps > 240: +# print(f"Warning: Invalid FPS {fps}") +# return + +# save_fmt = Path(path).suffix +# os.makedirs(os.path.dirname(path), exist_ok=True) +# width, height = pil_images[0].size + +# if save_fmt == ".mp4": +# try: +# codec = "libx264" +# container = av.open(path, "w") +# stream = container.add_stream(codec, rate=fps) + +# stream.width = width +# stream.height = height + +# for pil_image in pil_images: +# # pil_image = Image.fromarray(image_arr).convert("RGB") +# av_frame = av.VideoFrame.from_image(pil_image) +# container.mux(stream.encode(av_frame)) +# container.mux(stream.encode()) +# container.close() +# except Exception as e: +# print(f"Unexpected error while saving video {path}: {e}") +# if os.path.exists(path): +# try: +# os.remove(path) +# print(f"Corrupted file {path} removed successfully.") +# except Exception as rm_e: +# print(f"Failed to remove corrupted file {path}: {rm_e}") + +# elif save_fmt == ".gif": +# pil_images[0].save( +# fp=path, +# format="GIF", +# append_images=pil_images[1:], +# save_all=True, +# duration=(1 / fps * 1000), +# loop=0, +# ) +# else: +# raise ValueError("Unsupported file type. Use .mp4 or .gif.") + + +def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=6, fps=8): + videos = rearrange(videos, "b c t h w -> t b c h w") + height, width = videos.shape[-2:] + outputs = [] + + for x in videos: # x: b c h w + if x.shape[0] != 1: + x = torchvision.utils.make_grid(x, nrow=n_rows) # (c h w) + else: + x = x.squeeze(0) + x = x.transpose(0, 1).transpose(1, 2).squeeze(-1) # (h w c) + if rescale: + x = (x + 1.0) / 2.0 # -1,1 -> 0,1 + x = (x * 255).numpy().astype(np.uint8) + x = Image.fromarray(x) + + outputs.append(x) + + os.makedirs(os.path.dirname(path), exist_ok=True) + + save_videos_from_pil(outputs, path, fps) + + +def read_frames(video_path): + try: + # 使用 decord 打开视频 + vr = VideoReader(video_path) + frames = [] + + # 逐帧解码 + for i in range(len(vr)): + frame = vr[i] # 获取帧,返回的是 mx.ndarray + image = Image.fromarray(frame.asnumpy()) # 转换为 PIL 格式 + frames.append(image) + + return frames + + except Exception as e: + print(f"Error reading frames from {video_path}: {e}") + return None # 返回 None 避免代码崩溃 + +def get_fps(video_path): + + # container = av.open(video_path) + # video_stream = next(s for s in container.streams if s.type == "video") + # fps = video_stream.average_rate + # container.close() + # print("pyav_fps") + # print(fps) + try: + vr = decord.VideoReader(video_path) + fps = vr.get_avg_fps() + return Fraction(fps).limit_denominator(1001) + except Exception as e: + print(f"Error reading FPS from {video_path}: {e}") + return None # 返回 None 避免代码崩溃 + + +def read_frames_and_fps(video_path): + try: + # 使用 decord 打开视频 + vr = VideoReader(video_path) + fps = vr.get_avg_fps() + frames = [] + + # 逐帧解码 + for i in range(len(vr)): + frame = vr[i] # 获取帧,返回的是 mx.ndarray + image = Image.fromarray(frame.asnumpy()) # 转换为 PIL 格式 + frames.append(image) + + processed_fps = Fraction(fps).limit_denominator(1001) + return frames, processed_fps + + except Exception as e: + print(f"Error reading frames from {video_path}: {e}") + return None, None # 返回 None 避免代码崩溃 + +def read_frames_and_fps_as_np(video_path): + try: + # 使用 decord 打开视频 + vr = VideoReader(video_path) + fps = vr.get_avg_fps() + frames = [] + + # 逐帧解码 + for i in range(len(vr)): + frame = vr[i] # 获取帧,返回的是 mx.ndarray + image = frame.asnumpy() # 转换为 PIL 格式 + frames.append(image) + processed_fps = Fraction(fps).limit_denominator(1001) + return frames, processed_fps + + except Exception as e: + print(f"Error reading frames from {video_path}: {e}") + return None, None # 返回 None 避免代码崩溃 + + + +def save_videos_from_pil(pil_images, path, fps=8): + if fps is None or fps <= 0 or fps > 240: + print(f"Warning: Invalid FPS {fps}") + return + + save_fmt = Path(path).suffix + os.makedirs(os.path.dirname(path), exist_ok=True) + width, height = pil_images[0].size + + if save_fmt == ".mp4": + try: + codec = "libx264" + container = av.open(path, "w") + stream = container.add_stream(codec, rate=fps) + + stream.width = width + stream.height = height + + for pil_image in pil_images: + # pil_image = Image.fromarray(image_arr).convert("RGB") + av_frame = av.VideoFrame.from_image(pil_image) + container.mux(stream.encode(av_frame)) + container.mux(stream.encode()) + container.close() + except Exception as e: + print(f"Unexpected error while saving video {path}: {e}") + if os.path.exists(path): + try: + os.remove(path) + print(f"Corrupted file {path} removed successfully.") + except Exception as rm_e: + print(f"Failed to remove corrupted file {path}: {rm_e}") + + elif save_fmt == ".gif": + pil_images[0].save( + fp=path, + format="GIF", + append_images=pil_images[1:], + save_all=True, + duration=(1 / fps * 1000), + loop=0, + ) + else: + raise ValueError("Unsupported file type. Use .mp4 or .gif.") + + +def load_video_with_pose_from_first_frame(video_data, pose_data, sampling="uniform", duration=None, num_frames=99, wanted_fps=None, actual_fps=None, + skip_frms_num=4., nb_read_frames=None): + decord.bridge.set_bridge("torch") + vr = VideoReader(uri=video_data, height=-1, width=-1) + vr_pose = VideoReader(uri=pose_data, height=-1, width=-1) + + start = 0 + end = int(start + num_frames / wanted_fps * actual_fps) + n_frms = num_frames + 1 # 要取到num_frames帧, +1把第一帧需要额外拿出来处理,让第一帧和后续帧不同 + + if sampling == "uniform": + indices = np.arange(start, end, (end - start) / n_frms).astype(int) + else: + raise NotImplementedError + + # get_batch -> T, H, W, C + + temp_frms = vr.get_batch(np.arange(start, end)) + temp_frms_pose = vr_pose.get_batch(np.arange(start, end)) + + assert temp_frms is not None + assert temp_frms_pose is not None + + tensor_frms = torch.from_numpy(temp_frms) if type(temp_frms) is not torch.Tensor else temp_frms + tensor_frms = tensor_frms[torch.tensor((indices - start).tolist())] + + tensor_frms_pose = torch.from_numpy(temp_frms_pose) if type(temp_frms_pose) is not torch.Tensor else temp_frms_pose + tensor_frms_pose = temp_frms_pose[torch.tensor((indices - start).tolist())] + + # print(f"n_frms: {n_frms}; tensor_frms.shape: {tensor_frms.shape} tensor_frms_pose.shape: {tensor_frms_pose.shape}") + return pad_last_frame(tensor_frms, n_frms), pad_last_frame(tensor_frms_pose, n_frms) + + +def pad_last_frame(tensor, sampling_frms_num): + # T, H, W, C + if tensor.shape[0] < sampling_frms_num: + # 复制最后一帧 + last_frame = tensor[-int(sampling_frms_num-tensor.shape[0]):] + # 将最后一帧添加到第二个维度 + padded_tensor = torch.cat([tensor, last_frame], dim=0) + return padded_tensor + else: + return tensor[:sampling_frms_num] + + +def load_video_sampling(video_data, pose_data, num_frames, wanted_fps): + decord.bridge.set_bridge("torch") + # 以video_data的为准 + vr = VideoReader(uri=video_data, height=-1, width=-1) + actual_fps = vr.get_avg_fps() + if video_data: + video, pose = load_video_with_pose_from_first_frame(video_data, pose_data, sampling="uniform", duration=100000, num_frames=num_frames, wanted_fps=wanted_fps, actual_fps=actual_fps, skip_frms_num=0, nb_read_frames=None) + return video, pose + else: + raise ValueError("mooreAA should have video data") \ No newline at end of file diff --git a/SCAIL-Pose/DWPoseProcess/checkUtils.py b/SCAIL-Pose/DWPoseProcess/checkUtils.py new file mode 100644 index 0000000000000000000000000000000000000000..a445a2619cffb13a0384a63734e6d9dfe72290cc --- /dev/null +++ b/SCAIL-Pose/DWPoseProcess/checkUtils.py @@ -0,0 +1,263 @@ +import numpy as np +from collections import deque +import numpy as np +from collections import deque +import math + +def get_bbox_area(bbox): + x1, y1, x2, y2 = bbox + return (x2 - x1) * (y2 - y1) + + +def check_consistant(boxA, boxB, scoreA_lst, scoreB_lst, beta, all_threshold): + """ + 计算两个锚框之间的 IoU(交并比)以及分数的变化比例,来判断连续性 + """ + # 计算交集框的坐标 + scoreA = scoreA_lst[0] + scoreB = scoreB_lst[0] + iou = get_IoU(boxA, boxB) + + reduction_ratio = (scoreA - scoreB) / (scoreA + scoreB) if scoreA > scoreB else 0 # 如果分数减少的很多,就更不连续 + + return iou - reduction_ratio * beta > all_threshold + + +def get_IoU(boxA, boxB): + """ + 计算两个锚框之间的 IoU(交并比)以及分数的变化比例,来判断连续性 + """ + # 计算交集框的坐标 + x1_int = max(boxA[0], boxB[0]) + y1_int = max(boxA[1], boxB[1]) + x2_int = min(boxA[2], boxB[2]) + y2_int = min(boxA[3], boxB[3]) + + # 计算交集的面积 + inter_width = max(0, x2_int - x1_int) + inter_height = max(0, y2_int - y1_int) + inter_area = inter_width * inter_height + + # 计算两个锚框的面积 + areaA = (boxA[2] - boxA[0]) * (boxA[3] - boxA[1]) + areaB = (boxB[2] - boxB[0]) * (boxB[3] - boxB[1]) + + # 计算并集的面积 + union_area = areaA + areaB - inter_area + + # 计算 IoU + iou = inter_area / union_area if union_area > 0 else 0.0 + + return iou + +def check_bbox_single_for_video(bbox, reference_width, reference_height, min_bbox_width=1/6, min_bbox_area=1/30): + """ + 每一帧都要检查,判断 bbox 是否符合视频格式下的要求 + """ + x1, y1, x2, y2 = bbox + bbox_width = x2 - x1 + bbox_height = y2 - y1 + + # 防止非法 bbox + if bbox_width <= 0 or bbox_height <= 0: + return False + + bbox_area = bbox_width * bbox_height + + if bbox_area < min_bbox_area: + # print("filtered: bbox too small or too large") + return False + + # 横屏额外筛 + if reference_width > reference_height: + if bbox_width < min_bbox_width: + # print("filtered: bbox too wide") + return False + return True + + + +##############################判断点是否满足################################ + +def part5_valid(valid_joints): + """ + 判断身体五块是不是都有点 + + 参数: + valid_joints:布尔数组 + + 返回: + bool: 如果满足要求返回True,否则返回False。 + """ + + ### 认为以下的视频是满足我们采样需要的: + ### A只有上半身手部动作的:手部有可能在挥舞过程中移出屏幕,但是上半身应该一直在屏幕内,此时1-0, 1-2, 1-5这三条骨骼应该都存在;14-17应该至少有一个点存在 + ### B全身动作:1-0, 1-2, 1-5, 1-8, 1-11这五条骨骼应该都存在 + + top_core_joints = valid_joints[[0,1,2,5]] + top_nece_joints = valid_joints[[14,15,16,17]] + if all(top_core_joints) and any(top_nece_joints): + return True + + wholebody_core_joints = valid_joints[[0,1,2,5,8,11]] + if all(wholebody_core_joints): + return True + return False + + +def check_valid_sequence(valid_keypoints, threshold=0.3): + valid_joints = np.zeros(24) + for valid_keypoint in valid_keypoints: + valid_joints += valid_keypoint + return part5_valid(valid_joints) + + +def get_valid_indice_from_keypoints(ref_part_poses, ref_part_indices): + # ref_part_poses: poses序列,ref_part_indices poses序列里面每个值对应的在整个序列里的index + # return: 每个pose序列里面,满足要求的pose的index + valid_indice = [] + for i, (keypoint_all, indice) in enumerate(zip(ref_part_poses, ref_part_indices)): + body_subset = keypoint_all["bodies"]["subset"][0] + valid_joints = body_subset > -1 # 得到一个布尔索引 + + if not part5_valid(valid_joints): + continue + + faces = keypoint_all["faces"][0] + + left_eye = faces[36:42] # 左眼关键点 5个关键点有4个认为有左眼 + right_eye = faces[42:48] # 右眼关键点 5个关键点有4个认为有右眼 + nose = faces[27:36] # 鼻子关键点 8个关键点有5个认为有鼻子 + mouth = faces[48:68] # 嘴巴关键点 21个关键点有15个认为有嘴巴 + + # 计算每个部位有效的关键点数 + left_eye_valid = sum(1 for point in left_eye if point[0] > 0 and point[1] > 0) + right_eye_valid = sum(1 for point in right_eye if point[0] > 0 and point[1] > 0) + nose_valid = sum(1 for point in nose if point[0] > 0 and point[1] > 0) + mouth_valid = sum(1 for point in mouth if point[0] > 0 and point[1] > 0) + + # 如果有两个或以上部位有效,则认为是正脸 + valid_face_parts = 0 + if left_eye_valid >= 4: + valid_face_parts += 1 + if right_eye_valid >= 4: + valid_face_parts += 1 + if nose_valid >= 5: + valid_face_parts += 1 + if mouth_valid >= 15: + valid_face_parts += 1 + + if valid_face_parts >= 2: + valid_indice.append(int(indice)) + + return valid_indice + +def check_from_keypoints_core_keypoints(keypoints, bboxs): + # 用于keypoints版本,根据每一帧的18个keypoints和bbox iou来判断是否满足要求 + valid_sequence = deque(maxlen=4) + for i, (keypoint_all, bbox_all) in enumerate(zip(keypoints, bboxs)): + body_subset = keypoint_all["bodies"]["subset"][0] + valid_joints = body_subset > -1 # 得到一个布尔索引 + if len(valid_sequence) == 4: + if not check_valid_sequence(valid_sequence): + # print("filtered: 骨骼不满足要求") + return False # 关键点异常 + valid_sequence.append(valid_joints) + return True + +def select_ref_from_keypoints_bbox_multi(ref_part_indices, ref_part_bboxes, bboxs): + for ref_index, ref_bbox in zip(ref_part_indices, ref_part_bboxes): + bbox_areas_ref = [get_bbox_area(bbox) for bbox in ref_bbox] + max_bbox_area_ref = max(bbox_areas_ref) + num_human_ref = sum(1 for bbox in ref_bbox if get_bbox_area(bbox) > max_bbox_area_ref * 0.5) + if num_human_ref < 2 or num_human_ref > 5: + continue + driving_bbox_ok = True + for i, bbox_all in enumerate(bboxs): + bbox_areas = [get_bbox_area(bbox) for bbox in bbox_all] + max_bbox_area = max(bbox_areas) + num_human = sum(1 for bbox in bbox_all if get_bbox_area(bbox) > max_bbox_area * 0.5) + if num_human != num_human_ref: + driving_bbox_ok = False + break + if driving_bbox_ok: + return int(ref_index) + else: + continue + return None + +def check_from_keypoints_bbox(keypoints, bboxs, IoU_thresthold, reference_width, reference_height, multi_person=False): + # 用于keypoints版本,根据每一帧的18个keypoints和bbox iou来判断是否满足要求 + last_bbox = None + for i, (keypoint_all, bbox_all) in enumerate(zip(keypoints, bboxs)): + if not len(bbox_all): + return False + else: + if multi_person: + for bbox in bbox_all: + if not check_bbox_single_for_video(bbox, reference_width, reference_height, min_bbox_width=1/6): + return False + else: + bbox = bbox_all[0] + if not check_bbox_single_for_video(bbox, reference_width, reference_height, min_bbox_width=1/7): + return False # bbox大小异常 + if last_bbox is not None: + if not get_IoU(bbox, last_bbox) > IoU_thresthold: + return False # IoU异常 + last_bbox = bbox + return True + + + +def check_from_keypoints_stick_movement(keypoints, angle_threshold): + # 骨骼选择:列表中每个元组表示由两个关节确定一条骨骼:格式 (joint_a, joint_b) + # bones = [(1, 0), (1, 2), (1, 5), (1, 8), (1, 11)] + bones = [(1, 0), (1, 2), (1, 5), (1, 8), (1, 11), (2, 3), (5, 6), (8, 9), (11, 12)] + max_delta_list = [] + # 遍历从第二帧开始,对比前一帧和当前帧 + human_num_list = [len(keypoints[idx]["bodies"]["candidate"]) for idx in range(0, len(keypoints))] + min_human_num = min(human_num_list) + for human_idx in range(min_human_num): + for i in range(1, len(keypoints)): + # 获取上一帧和当前帧的关键点数据(格式为 (18,3) 数组) + prev_frame_subset = keypoints[i-1]["bodies"]["subset"][human_idx] + curr_frame_subset = keypoints[i]["bodies"]["subset"][human_idx] + prev_frame_keypoints = keypoints[i-1]["bodies"]["candidate"][human_idx] + curr_frame_keypoints = keypoints[i]["bodies"]["candidate"][human_idx] + + + max_delta = 0 + for (j1, j2) in bones: + # 检查上一帧中两个关节是否有效(假设 x, y 坐标需大于 0 才认为有效) + if prev_frame_subset[j1] < 0 or prev_frame_subset[j2] < 0: + continue + if curr_frame_subset[j1] < 0 or curr_frame_subset[j2] < 0: + continue + + # 计算上一帧和当前帧中对应骨骼的向量(方向一致,均从 j1 指向 j2) + vec_prev = np.array([prev_frame_keypoints[j2][0] - prev_frame_keypoints[j1][0], + prev_frame_keypoints[j2][1] - prev_frame_keypoints[j1][1]]) + vec_curr = np.array([curr_frame_keypoints[j2][0] - curr_frame_keypoints[j1][0], + curr_frame_keypoints[j2][1] - curr_frame_keypoints[j1][1]]) + + # 如果向量模长为0,则无法计算角度,跳过 + if np.linalg.norm(vec_prev) == 0 or np.linalg.norm(vec_curr) == 0: + continue + # 计算向量对应的角度(弧度制) + angle_prev = math.atan2(vec_prev[1], vec_prev[0]) + angle_curr = math.atan2(vec_curr[1], vec_curr[0]) + + # 计算角度差,并规范到 [0, pi] 范围 + delta = abs(angle_curr - angle_prev) + if delta > math.pi: + delta = 2 * math.pi - delta + + max_delta = max(delta, max_delta) + max_delta_list.append(max_delta) + max_delta_list = sorted(max_delta_list) + max_delta_list = max_delta_list[len(max_delta_list)//8:-len(max_delta_list)//8] # 去掉两端8分之一的值 + avg_movement = sum(max_delta_list) / len(max_delta_list) + if avg_movement < angle_threshold: # 筛去过小的动作 + return False + return True + \ No newline at end of file diff --git a/SCAIL-Pose/DWPoseProcess/dwpose/__init__.py b/SCAIL-Pose/DWPoseProcess/dwpose/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..d753f7dc0633bfc46ddef7f4684c3a8463babe43 --- /dev/null +++ b/SCAIL-Pose/DWPoseProcess/dwpose/__init__.py @@ -0,0 +1,129 @@ +# https://github.com/IDEA-Research/DWPose +# Openpose +# Original from CMU https://github.com/CMU-Perceptual-Computing-Lab/openpose +# 2nd Edited by https://github.com/Hzzone/pytorch-openpose +# 3rd Edited by ControlNet +# 4th Edited by ControlNet (added face and correct hands) + +import copy +import os + +os.environ["KMP_DUPLICATE_LIB_OK"] = "TRUE" +import cv2 +import numpy as np +import torch +from controlnet_aux.util import HWC3, resize_image +from PIL import Image + +from . import util +from .wholebody import Wholebody + + +class DWposeDetector: + def __init__(self, use_batch=False): + self.use_batch = use_batch + pass + + def to(self, device): + self.pose_estimation = Wholebody(device, self.use_batch) + return self + + def _get_multi_result_from_est(self, candidate, score_result, det_result, H, W): + nums, keys, locs = candidate.shape # n 所有身体关键点数量,坐标 + candidate[..., 0] /= float(W) + candidate[..., 1] /= float(H) + subset_score = score_result[:, :24] # 按照24个骨骼关键点来区分可见位置 + face_score = score_result[:, 24:92] + hand_score = score_result[:, 92:113] + hand_score = np.vstack([hand_score, score_result[:, 113:]]) + + body_candidate = candidate[:, :24].copy() # body(n, 24, 2) + for i in range(len(subset_score)): # n 个 + for j in range(len(subset_score[i])): + if subset_score[i][j] > 0.3: + subset_score[i][j] = j # 标注序号,这样后续用的时候可以快速查出可用点 + else: + subset_score[i][j] = -1 # 躯干中去除掉不可见的骨骼 + + un_visible = score_result < 0.3 + candidate[un_visible] = -1 # 全部关键点中去掉不可见骨骼 + + faces = candidate[:, 24:92] + hands = candidate[:, 92:113] # hands(2*n, 21, 2) + hands = np.vstack([hands, candidate[:, 113:]]) + + bodies = dict(candidate=body_candidate, subset=subset_score) + pose = dict(bodies=bodies, hands=hands, faces=faces) + score = dict(body_score=subset_score, hand_score=hand_score, face_score=face_score) + + new_det_result = [] + for bbox in det_result: + x1, y1, x2, y2 = bbox + new_x1 = x1 / W + new_y1 = y1 / H + new_x2 = x2 / W + new_y2 = y2 / H + new_bbox = [new_x1, new_y1, new_x2, new_y2] + new_det_result.append(new_bbox) + + return pose, score, new_det_result # body_score是原始的躯干骨骼分数 + + # def _get_result_from_est(self, input_image, candidate, subset, det_result, image_resolution, output_type, H, W): + # nums, keys, locs = candidate.shape + # candidate[..., 0] /= float(W) + # candidate[..., 1] /= float(H) + # score = subset[:, :18] # 前18个是躯干骨骼 score(n, 18) + # max_ind = np.mean(score, axis=-1).argmax(axis=0) # 返回分数最高的锚框对应的骨骼 + # score = score[[max_ind]] + # body = candidate[:, :18].copy() + # body = body[[max_ind]] + # nums = 1 + # body = body.reshape(nums * 18, locs) # Moore-AA只有一个人体, 0-18表示body + # body_score = copy.deepcopy(score) # 已经去过max_ind + # for i in range(len(score)): + # for j in range(len(score[i])): + # if score[i][j] > 0.3: + # score[i][j] = int(18 * i + j) + # else: + # score[i][j] = -1 # 躯干中去除掉不可见的骨骼 + + # un_visible = subset < 0.3 + # candidate[un_visible] = -1 # 全部关键点中去掉不可见骨骼 + + # foot = candidate[:, 18:24] + + # faces = candidate[[max_ind], 24:92] + + # hands = candidate[[max_ind], 92:113] + # hands = np.vstack([hands, candidate[[max_ind], 113:]]) + + # bodies = dict(candidate=body, subset=score) + # pose = dict(bodies=bodies, hands=hands, faces=faces) + + # return pose, body_score, det_result # body_score是原始的躯干骨骼分数 + + def __call__( + self, + input, + **kwargs, + ): + if not self.use_batch: + # PIL要不要颜色反转? + input = cv2.cvtColor( + np.array(input, dtype=np.uint8), cv2.COLOR_RGB2BGR + ) + input = HWC3(input) + H, W, C = input.shape + + with torch.no_grad(): + candidate, subset, det_result = self.pose_estimation(input) # candidate (n, 134, 2) 候选点 / subset (n, 134) 得分 + return self._get_multi_result_from_est(candidate, subset, det_result, H, W) + else: + raise NotImplementedError("DWposeDetector does not support batch mode") + + + + + + + diff --git a/SCAIL-Pose/DWPoseProcess/dwpose/onnxdet.py b/SCAIL-Pose/DWPoseProcess/dwpose/onnxdet.py new file mode 100644 index 0000000000000000000000000000000000000000..2255e077accd40e88db6668b5c54686e77898d9d --- /dev/null +++ b/SCAIL-Pose/DWPoseProcess/dwpose/onnxdet.py @@ -0,0 +1,134 @@ +# https://github.com/IDEA-Research/DWPose +import cv2 +import numpy as np +import onnxruntime + + +def nms(boxes, scores, nms_thr): + """Single class NMS implemented in Numpy.""" + x1 = boxes[:, 0] + y1 = boxes[:, 1] + x2 = boxes[:, 2] + y2 = boxes[:, 3] + + areas = (x2 - x1 + 1) * (y2 - y1 + 1) + order = scores.argsort()[::-1] + + keep = [] + while order.size > 0: + i = order[0] + keep.append(i) + xx1 = np.maximum(x1[i], x1[order[1:]]) + yy1 = np.maximum(y1[i], y1[order[1:]]) + xx2 = np.minimum(x2[i], x2[order[1:]]) + yy2 = np.minimum(y2[i], y2[order[1:]]) + + w = np.maximum(0.0, xx2 - xx1 + 1) + h = np.maximum(0.0, yy2 - yy1 + 1) + inter = w * h + ovr = inter / (areas[i] + areas[order[1:]] - inter) + + inds = np.where(ovr <= nms_thr)[0] + order = order[inds + 1] + + return keep + + +def multiclass_nms(boxes, scores, nms_thr, score_thr): + """Multiclass NMS implemented in Numpy. Class-aware version.""" + final_dets = [] + num_classes = scores.shape[1] + for cls_ind in range(num_classes): + cls_scores = scores[:, cls_ind] + valid_score_mask = cls_scores > score_thr + if valid_score_mask.sum() == 0: + continue + else: + valid_scores = cls_scores[valid_score_mask] + valid_boxes = boxes[valid_score_mask] + keep = nms(valid_boxes, valid_scores, nms_thr) + if len(keep) > 0: + cls_inds = np.ones((len(keep), 1)) * cls_ind + dets = np.concatenate( + [valid_boxes[keep], valid_scores[keep, None], cls_inds], 1 + ) + final_dets.append(dets) + if len(final_dets) == 0: + return None + return np.concatenate(final_dets, 0) + + +def demo_postprocess(outputs, img_size, p6=False): + grids = [] + expanded_strides = [] + strides = [8, 16, 32] if not p6 else [8, 16, 32, 64] + + hsizes = [img_size[0] // stride for stride in strides] + wsizes = [img_size[1] // stride for stride in strides] + + for hsize, wsize, stride in zip(hsizes, wsizes, strides): + xv, yv = np.meshgrid(np.arange(wsize), np.arange(hsize)) + grid = np.stack((xv, yv), 2).reshape(1, -1, 2) + grids.append(grid) + shape = grid.shape[:2] + expanded_strides.append(np.full((*shape, 1), stride)) + + grids = np.concatenate(grids, 1) + expanded_strides = np.concatenate(expanded_strides, 1) + outputs[..., :2] = (outputs[..., :2] + grids) * expanded_strides + outputs[..., 2:4] = np.exp(outputs[..., 2:4]) * expanded_strides + + return outputs + + +def preprocess(img, input_size, swap=(2, 0, 1)): + if len(img.shape) == 3: + padded_img = np.ones((input_size[0], input_size[1], 3), dtype=np.uint8) * 114 + else: + padded_img = np.ones(input_size, dtype=np.uint8) * 114 + + r = min(input_size[0] / img.shape[0], input_size[1] / img.shape[1]) + resized_img = cv2.resize( + img, + (int(img.shape[1] * r), int(img.shape[0] * r)), + interpolation=cv2.INTER_LINEAR, + ).astype(np.uint8) + padded_img[: int(img.shape[0] * r), : int(img.shape[1] * r)] = resized_img + + padded_img = padded_img.transpose(swap) + padded_img = np.ascontiguousarray(padded_img, dtype=np.float32) + return padded_img, r + + +def inference_detector(session, oriImg): + input_shape = (640, 640) + img, ratio = preprocess(oriImg, input_shape) + + ort_inputs = {session.get_inputs()[0].name: img[None, :, :, :]} + output = session.run(None, ort_inputs) + predictions = demo_postprocess(output[0], input_shape)[0] + + boxes = predictions[:, :4] + scores = predictions[:, 4:5] * predictions[:, 5:] + + boxes_xyxy = np.ones_like(boxes) + boxes_xyxy[:, 0] = boxes[:, 0] - boxes[:, 2] / 2.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: + # print("no boxes detected") + return [] + + return final_boxes + + + diff --git a/SCAIL-Pose/DWPoseProcess/dwpose/onnxpose.py b/SCAIL-Pose/DWPoseProcess/dwpose/onnxpose.py new file mode 100644 index 0000000000000000000000000000000000000000..a8f181539823540d4b82dcec7963000d11f3221c --- /dev/null +++ b/SCAIL-Pose/DWPoseProcess/dwpose/onnxpose.py @@ -0,0 +1,487 @@ +# 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) + + outputs = sess.run(sess_output, sess_input) + # outputs 也是len为2的列表,outputs[0] outputs[1]第一维为1 + all_out.append(outputs) + + # breakpoint() + + return all_out + + +def postprocess( + outputs: List[np.ndarray], + model_input_size: Tuple[int, int], + center: Tuple[int, int], # 实际是List of tuple + 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): List of Center of bbox in shape (x, y). + scale (tuple): List of 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 + + +def inference_pose_batch(session, out_bbox, oriImg): + h, w = session.get_inputs()[0].shape[2:] + model_input_size = (w, h) + # breakpoint() + num_boxes = [box.shape[0] if len(box)>=1 else 1 for box in out_bbox] + resized_img, center, scale = preprocess_batch(oriImg, out_bbox, model_input_size) + outputs = inference_batch(session, resized_img) # nparray 第0维为(\sum_{batch_num} num_box) + keypoints, scores = postprocess_batch(outputs, num_boxes, model_input_size, center, scale) + # keypoints和scores在postprocess中从(\sum_{batch_num} num_box) 恢复为 a list of batch_num of nparray[num_box, ...] + + return keypoints, scores + + +def preprocess_batch( + img_list: List[np.ndarray], out_bbox_list, input_size: Tuple[int, int] = (192, 256) +) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: + # img_list shape: batch_num * nparray[C, H, W] + # out_bbox_list shape: batch_bum * nparray[num_box, 4] + resized_imgs = [] + centers = [] + scales = [] + + for img, out_bbox in zip(img_list, out_bbox_list): + resized_img, center, scale = preprocess(img, out_bbox, input_size) + # 有的out_bbox是两个 + # resized_img: list of array, num_box*nparray [C, H', W'] + resized_imgs.extend(resized_img) + centers.extend(center) + scales.extend(scale) + # resized_imgs 为 nparray 第0维为(\sum_{batch_num} num_box) + # centers和scales 为 list of nparray len为(\sum_{batch_num} num_box) + return np.stack(resized_imgs), centers, scales + + +def inference_batch(sess: ort.InferenceSession, batch_img): + """Inference RTMPose model in batch.""" + # build input + input = batch_img.transpose(0, 3, 1, 2).astype(np.float32) + + # build output + sess_input = {sess.get_inputs()[0].name: input} + sess_output = [] + for out in sess.get_outputs(): + sess_output.append(out.name) + + outputs = sess.run(sess_output, sess_input) + num = outputs[0].shape[0] + split_0 = np.split(outputs[0], num, axis=0) + split_1 = np.split(outputs[1], num, axis=0) + # 创建最终的输出列表,其中每个元素是一个包含两个 nparray 的列表 + final_output = [[split_0[i], split_1[i]] for i in range(num)] + return final_output # list of nparray 长度为(\sum_{batch_num} num_box) + +def postprocess_batch( + outputs: List[np.ndarray], + num_boxes: List[int], + model_input_size: Tuple[int, int], + center: List[Tuple[int, int]], + scale: List[Tuple[int, int]], + simcc_split_ratio: float = 2.0, +) -> Tuple[np.ndarray, np.ndarray]: + # keypoints和scores在postprocess中根据num_boxes从(\sum_{batch_num} num_box) 恢复为 a list of batch_num of nparray[num_box, ...] + all_key, all_score = postprocess(outputs, model_input_size, center, scale, simcc_split_ratio) + relist_all_key, relist_all_score = [], [] + start_idx = 0 + for num_box in num_boxes: + # 按照每个 batch 的 num_box 从 all_key 和 all_score 中恢复出对应的部分 + split_indices = np.cumsum(num_boxes) + split_indices = np.insert(split_indices, 0, 0) # 在开头插入0 -> 0, len1, len1+len2... + # 恢复每个batch的 keypoints 和 scores + relist_all_key = [all_key[split_indices[i]:split_indices[i + 1]] for i in range(len(split_indices) - 1)] + relist_all_score = [all_score[split_indices[i]:split_indices[i + 1]] for i in range(len(split_indices) - 1)] + return relist_all_key, relist_all_score # a list of batch_num of nparray[num_box, ...] + + +# Main function for testing inference_pose_batch +def main(): + # 创建一个简单的测试图像和假边界框 + img = np.random.rand(256, 192, 3).astype(np.float32) # 一张随机图像 + out_bbox = np.array([[50, 50, 150, 150]]) + + img_list = [np.random.rand(256, 192, 3).astype(np.float32) for _ in range(3)] # 3张随机图像 + out_bbox_list = [ + np.array([[50, 50, 150, 150]]), # 第一个批次: 1个框 + np.array([[30, 30, 130, 130], [50, 50, 200, 200]]), # 第二个批次: 2个框 + np.array([[10, 10, 100, 100]]), # 第三个批次: 1个框 + ] + + # 加载 ONNX 模型 + model_path = "/workspace/yanwenhao/Moore-AnimateAnyone/pretrained_weights/DWPose/dw-ll_ucoco_384.onnx" # 在此替换为实际的模型路径 + sess = ort.InferenceSession(model_path) + + # 调用 inference_pose 进行单张图像推理 + keypoints, scores = inference_pose(sess, out_bbox, img) + + # 输出结果 + print("single infer test") + print("Keypoints:", keypoints.shape) + print("Scores:", scores.shape) + + # 调用 inference_pose_batch 进行推理 + keypoints, scores = inference_pose_batch(sess, out_bbox_list, img_list) + + # 输出结果 + print("batch infer test") + for i, (keypoint, score) in enumerate(zip(keypoints, scores)): + print(f"Batch_index {i + 1}:") + print("Keypoints:", keypoint.shape) + print("Scores:", score.shape) + +# 调用 main 函数 +if __name__ == "__main__": + main() diff --git a/SCAIL-Pose/DWPoseProcess/dwpose/util.py b/SCAIL-Pose/DWPoseProcess/dwpose/util.py new file mode 100644 index 0000000000000000000000000000000000000000..0b8806bc950903e15741a79e05cf1ec729fdcc9f --- /dev/null +++ b/SCAIL-Pose/DWPoseProcess/dwpose/util.py @@ -0,0 +1,378 @@ +# https://github.com/IDEA-Research/DWPose +import math +import numpy as np +import matplotlib +import cv2 + + +eps = 0.01 + + +def smart_resize(x, s): + Ht, Wt = s + if x.ndim == 2: + Ho, Wo = x.shape + Co = 1 + else: + Ho, Wo, Co = x.shape + if Co == 3 or Co == 1: + k = float(Ht + Wt) / float(Ho + Wo) + return cv2.resize( + x, + (int(Wt), int(Ht)), + interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4, + ) + else: + return np.stack([smart_resize(x[:, :, i], s) for i in range(Co)], axis=2) + + +def smart_resize_k(x, fx, fy): + if x.ndim == 2: + Ho, Wo = x.shape + Co = 1 + else: + Ho, Wo, Co = x.shape + Ht, Wt = Ho * fy, Wo * fx + if Co == 3 or Co == 1: + k = float(Ht + Wt) / float(Ho + Wo) + return cv2.resize( + x, + (int(Wt), int(Ht)), + interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4, + ) + else: + return np.stack([smart_resize_k(x[:, :, i], fx, fy) for i in range(Co)], axis=2) + + +def padRightDownCorner(img, stride, padValue): + h = img.shape[0] + w = img.shape[1] + + pad = 4 * [None] + pad[0] = 0 # up + pad[1] = 0 # left + pad[2] = 0 if (h % stride == 0) else stride - (h % stride) # down + pad[3] = 0 if (w % stride == 0) else stride - (w % stride) # right + + img_padded = img + pad_up = np.tile(img_padded[0:1, :, :] * 0 + padValue, (pad[0], 1, 1)) + img_padded = np.concatenate((pad_up, img_padded), axis=0) + pad_left = np.tile(img_padded[:, 0:1, :] * 0 + padValue, (1, pad[1], 1)) + img_padded = np.concatenate((pad_left, img_padded), axis=1) + pad_down = np.tile(img_padded[-2:-1, :, :] * 0 + padValue, (pad[2], 1, 1)) + img_padded = np.concatenate((img_padded, pad_down), axis=0) + pad_right = np.tile(img_padded[:, -2:-1, :] * 0 + padValue, (1, pad[3], 1)) + img_padded = np.concatenate((img_padded, pad_right), axis=1) + + return img_padded, pad + + +def transfer(model, model_weights): + transfered_model_weights = {} + for weights_name in model.state_dict().keys(): + transfered_model_weights[weights_name] = model_weights[ + ".".join(weights_name.split(".")[1:]) + ] + return transfered_model_weights + + +def draw_bodypose(canvas, candidate, subset): + H, W, C = canvas.shape + candidate = np.array(candidate) + subset = np.array(subset) + + stickwidth = 4 + + limbSeq = [ + [2, 3], + [2, 6], + [3, 4], + [4, 5], + [6, 7], + [7, 8], + [2, 9], + [9, 10], + [10, 11], + [2, 12], + [12, 13], + [13, 14], + [2, 1], + [1, 15], + [15, 17], + [1, 16], + [16, 18], + [3, 17], + [6, 18], + ] + + colors = [ + [255, 0, 0], + [255, 85, 0], + [255, 170, 0], + [255, 255, 0], + [170, 255, 0], + [85, 255, 0], + [0, 255, 0], + [0, 255, 85], + [0, 255, 170], + [0, 255, 255], + [0, 170, 255], + [0, 85, 255], + [0, 0, 255], + [85, 0, 255], + [170, 0, 255], + [255, 0, 255], + [255, 0, 170], + [255, 0, 85], + ] + + for i in range(17): + for n in range(len(subset)): + index = subset[n][np.array(limbSeq[i]) - 1] + if -1 in index: + continue + Y = candidate[index.astype(int), 0] * float(W) + X = candidate[index.astype(int), 1] * float(H) + mX = np.mean(X) + mY = np.mean(Y) + length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5 + angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1])) + polygon = cv2.ellipse2Poly( + (int(mY), int(mX)), (int(length / 2), stickwidth), int(angle), 0, 360, 1 + ) + cv2.fillConvexPoly(canvas, polygon, colors[i]) + + canvas = (canvas * 0.6).astype(np.uint8) + + for i in range(18): + for n in range(len(subset)): + index = int(subset[n][i]) + if index == -1: + continue + x, y = candidate[index][0:2] + x = int(x * W) + y = int(y * H) + cv2.circle(canvas, (int(x), int(y)), 4, colors[i], thickness=-1) + + return canvas + + +def draw_handpose(canvas, all_hand_peaks): + H, W, C = canvas.shape + + edges = [ + [0, 1], + [1, 2], + [2, 3], + [3, 4], + [0, 5], + [5, 6], + [6, 7], + [7, 8], + [0, 9], + [9, 10], + [10, 11], + [11, 12], + [0, 13], + [13, 14], + [14, 15], + [15, 16], + [0, 17], + [17, 18], + [18, 19], + [19, 20], + ] + + for peaks in all_hand_peaks: + peaks = np.array(peaks) + + for ie, e in enumerate(edges): + x1, y1 = peaks[e[0]] + x2, y2 = peaks[e[1]] + x1 = int(x1 * W) + y1 = int(y1 * H) + x2 = int(x2 * W) + y2 = int(y2 * H) + if x1 > eps and y1 > eps and x2 > eps and y2 > eps: + cv2.line( + canvas, + (x1, y1), + (x2, y2), + matplotlib.colors.hsv_to_rgb([ie / float(len(edges)), 1.0, 1.0]) + * 255, + thickness=2, + ) + + for i, keyponit in enumerate(peaks): + x, y = keyponit + x = int(x * W) + y = int(y * H) + if x > eps and y > eps: + cv2.circle(canvas, (x, y), 2, (0, 0, 255), thickness=-1) + return canvas + + +def draw_facepose(canvas, all_lmks): + H, W, C = canvas.shape + for lmks in all_lmks: + lmks = np.array(lmks) + for lmk in lmks: + x, y = lmk + x = int(x * W) + y = int(y * H) + if x > eps and y > eps: + cv2.circle(canvas, (x, y), 3, (255, 255, 255), thickness=-1) + return canvas + + +# detect hand according to body pose keypoints +# please refer to https://github.com/CMU-Perceptual-Computing-Lab/openpose/blob/master/src/openpose/hand/handDetector.cpp +def handDetect(candidate, subset, oriImg): + # right hand: wrist 4, elbow 3, shoulder 2 + # left hand: wrist 7, elbow 6, shoulder 5 + ratioWristElbow = 0.33 + detect_result = [] + image_height, image_width = oriImg.shape[0:2] + for person in subset.astype(int): + # if any of three not detected + has_left = np.sum(person[[5, 6, 7]] == -1) == 0 + has_right = np.sum(person[[2, 3, 4]] == -1) == 0 + if not (has_left or has_right): + continue + hands = [] + # left hand + if has_left: + left_shoulder_index, left_elbow_index, left_wrist_index = person[[5, 6, 7]] + x1, y1 = candidate[left_shoulder_index][:2] + x2, y2 = candidate[left_elbow_index][:2] + x3, y3 = candidate[left_wrist_index][:2] + hands.append([x1, y1, x2, y2, x3, y3, True]) + # right hand + if has_right: + right_shoulder_index, right_elbow_index, right_wrist_index = person[ + [2, 3, 4] + ] + x1, y1 = candidate[right_shoulder_index][:2] + x2, y2 = candidate[right_elbow_index][:2] + x3, y3 = candidate[right_wrist_index][:2] + hands.append([x1, y1, x2, y2, x3, y3, False]) + + for x1, y1, x2, y2, x3, y3, is_left in hands: + # pos_hand = pos_wrist + ratio * (pos_wrist - pos_elbox) = (1 + ratio) * pos_wrist - ratio * pos_elbox + # handRectangle.x = posePtr[wrist*3] + ratioWristElbow * (posePtr[wrist*3] - posePtr[elbow*3]); + # handRectangle.y = posePtr[wrist*3+1] + ratioWristElbow * (posePtr[wrist*3+1] - posePtr[elbow*3+1]); + # const auto distanceWristElbow = getDistance(poseKeypoints, person, wrist, elbow); + # const auto distanceElbowShoulder = getDistance(poseKeypoints, person, elbow, shoulder); + # handRectangle.width = 1.5f * fastMax(distanceWristElbow, 0.9f * distanceElbowShoulder); + x = x3 + ratioWristElbow * (x3 - x2) + y = y3 + ratioWristElbow * (y3 - y2) + distanceWristElbow = math.sqrt((x3 - x2) ** 2 + (y3 - y2) ** 2) + distanceElbowShoulder = math.sqrt((x2 - x1) ** 2 + (y2 - y1) ** 2) + width = 1.5 * max(distanceWristElbow, 0.9 * distanceElbowShoulder) + # x-y refers to the center --> offset to topLeft point + # handRectangle.x -= handRectangle.width / 2.f; + # handRectangle.y -= handRectangle.height / 2.f; + x -= width / 2 + y -= width / 2 # width = height + # overflow the image + if x < 0: + x = 0 + if y < 0: + y = 0 + width1 = width + width2 = width + if x + width > image_width: + width1 = image_width - x + if y + width > image_height: + width2 = image_height - y + width = min(width1, width2) + # the max hand box value is 20 pixels + if width >= 20: + detect_result.append([int(x), int(y), int(width), is_left]) + + """ + return value: [[x, y, w, True if left hand else False]]. + width=height since the network require squared input. + x, y is the coordinate of top left + """ + return detect_result + + +# Written by Lvmin +def faceDetect(candidate, subset, oriImg): + # left right eye ear 14 15 16 17 + detect_result = [] + image_height, image_width = oriImg.shape[0:2] + for person in subset.astype(int): + has_head = person[0] > -1 + if not has_head: + continue + + has_left_eye = person[14] > -1 + has_right_eye = person[15] > -1 + has_left_ear = person[16] > -1 + has_right_ear = person[17] > -1 + + if not (has_left_eye or has_right_eye or has_left_ear or has_right_ear): + continue + + head, left_eye, right_eye, left_ear, right_ear = person[[0, 14, 15, 16, 17]] + + width = 0.0 + x0, y0 = candidate[head][:2] + + if has_left_eye: + x1, y1 = candidate[left_eye][:2] + d = max(abs(x0 - x1), abs(y0 - y1)) + width = max(width, d * 3.0) + + if has_right_eye: + x1, y1 = candidate[right_eye][:2] + d = max(abs(x0 - x1), abs(y0 - y1)) + width = max(width, d * 3.0) + + if has_left_ear: + x1, y1 = candidate[left_ear][:2] + d = max(abs(x0 - x1), abs(y0 - y1)) + width = max(width, d * 1.5) + + if has_right_ear: + x1, y1 = candidate[right_ear][:2] + d = max(abs(x0 - x1), abs(y0 - y1)) + width = max(width, d * 1.5) + + x, y = x0, y0 + + x -= width + y -= width + + if x < 0: + x = 0 + + if y < 0: + y = 0 + + width1 = width * 2 + width2 = width * 2 + + if x + width > image_width: + width1 = image_width - x + + if y + width > image_height: + width2 = image_height - y + + width = min(width1, width2) + + if width >= 20: + detect_result.append([int(x), int(y), int(width)]) + + return detect_result + + +# get max index of 2d array +def npmax(array): + arrayindex = array.argmax(1) + arrayvalue = array.max(1) + i = arrayvalue.argmax() + j = arrayindex[i] + return i, j diff --git a/SCAIL-Pose/DWPoseProcess/dwpose/wholebody.py b/SCAIL-Pose/DWPoseProcess/dwpose/wholebody.py new file mode 100644 index 0000000000000000000000000000000000000000..7fbdadfcbffef24f642258c083b1470c6fc7c43c --- /dev/null +++ b/SCAIL-Pose/DWPoseProcess/dwpose/wholebody.py @@ -0,0 +1,58 @@ +# https://github.com/IDEA-Research/DWPose +from pathlib import Path + +import cv2 +import numpy as np +import onnxruntime as ort +import time +from .onnxdet import inference_detector +from .onnxpose import inference_pose, inference_pose_batch + +ModelDataPathPrefix = Path("./pretrained_weights") + + +class Wholebody: + def __init__(self, device, use_batch=False): + providers = [("CUDAExecutionProvider", { + "device_id": device + })] + # providers = [("CPUExecutionProvider", {})] + onnx_det = ModelDataPathPrefix.joinpath("DWPose/yolox_l.onnx") + onnx_pose = ModelDataPathPrefix.joinpath("DWPose/dw-ll_ucoco_384.onnx") + + self.session_det = ort.InferenceSession( + path_or_bytes=onnx_det, providers=providers + ) + self.session_pose = ort.InferenceSession( + path_or_bytes=onnx_pose, providers=providers + ) + self.use_batch = use_batch + + def _get_result_from_det_pose(self, det_result, keypoints, scores): + keypoints_info = np.concatenate((keypoints, scores[..., None]), axis=-1) # (1, 133, 3) + # compute neck joint + neck = np.mean(keypoints_info[:, [5, 6]], axis=1) # (1, 3),对第五第六个点做平均 + # 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) # 从第二个开始切片,这里维度为3,只切一片 + new_keypoints_info = np.insert(keypoints_info, 17, neck, axis=1) # 在17索引处插入neck + # 调换骨骼索引 + 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] # openpose需要检测17点+1脖子关键点 + 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, det_result + + + def __call__(self, oriImg): + if not self.use_batch: + det_result = inference_detector(self.session_det, oriImg) + keypoints, scores = inference_pose(self.session_pose, det_result, oriImg) # keypoints: (n_bbox, 133, 2) scores: (n_bbox, 133) 不管输入是什么初步提取的关键点数量都是一致的 + return self._get_result_from_det_pose(det_result=det_result, keypoints=keypoints, scores=scores) + + else: + raise NotImplementedError("DWposeDetector does not support batch mode") \ No newline at end of file diff --git a/SCAIL-Pose/DWPoseProcess/extractUtils.py b/SCAIL-Pose/DWPoseProcess/extractUtils.py new file mode 100644 index 0000000000000000000000000000000000000000..591e6de860d499a29b2b55bdbe65d5d683f5b897 --- /dev/null +++ b/SCAIL-Pose/DWPoseProcess/extractUtils.py @@ -0,0 +1,75 @@ +import numpy as np +import copy + +def get_bbox_area(bbox): + x1, y1, x2, y2 = bbox + return (x2 - x1) * (y2 - y1) + +def check_single_human_requirements(det_result): + # filter results + if len(det_result) > 3 or len(det_result) == 0: + return False + elif len(det_result) == 1: + return True + elif len(det_result) > 1: # [2, 3] + bbox_areas = [get_bbox_area(bbox) for bbox in det_result] + # 获取最大 bbox 面积的索引 + max_ind = max(range(len(bbox_areas)), key=lambda i: bbox_areas[i]) + # 获取次大面积(需要排除 max_ind) + other_indices = [i for i in range(len(bbox_areas)) if i != max_ind] + second_max_area = max([bbox_areas[i] for i in other_indices]) + + max_area = bbox_areas[max_ind] + if max_area < 2 * second_max_area: + return False + else: + return True + +def human_select(poses, det_results, multi_person): + new_poses = [] + new_det_results = [] + for pose, det_result in zip(poses, det_results): + if multi_person: + new_pose, new_det_result = get_multi_human(pose, det_result) + else: + new_pose, new_det_result = get_single_human(pose, det_result) + new_poses.append(new_pose) + new_det_results.append(new_det_result) + return new_poses, new_det_results + + +def get_single_human(pose, det_result): + if len(det_result) <= 1: + return pose, det_result + else: + bbox_areas = [get_bbox_area(bbox) for bbox in det_result] + max_ind = max(range(len(bbox_areas)), key=lambda i: bbox_areas[i]) + pose_copy = copy.deepcopy(pose) + pose_copy['bodies']['candidate'] = pose_copy['bodies']['candidate'][max_ind:max_ind+1] + pose_copy['bodies']['subset'] = pose_copy['bodies']['subset'][max_ind:max_ind+1] + pose_copy['hands'] = pose_copy['hands'][2*max_ind:2*max_ind+2] + pose_copy['faces'] = pose_copy['faces'][max_ind:max_ind+1] + return pose_copy, det_result[max_ind:max_ind+1] + +def check_multi_human_requirements(det_result): + # filter results + if len(det_result) < 2 or len(det_result) > 4: # 2-4个人 + return False + else: # [3, 6] + bbox_areas = [get_bbox_area(bbox) for bbox in det_result] + # 获取最大 bbox 面积的索引 + max_ind = max(range(len(bbox_areas)), key=lambda i: bbox_areas[i]) + max_area = bbox_areas[max_ind] + + # 选择面积大于等于最大面积50%的bbox + selected_indices = [i for i in range(len(bbox_areas)) if bbox_areas[i] >= 0.5 * max_area] # 包含max_ind + + # 检查选中的bbox数量是否大于等于2 + if len(selected_indices) >= 2: + return True + else: + return False + +def get_multi_human(pose, det_result): + # 后续再筛比较好,后续从65帧里面筛的时候可以把背景里的人的筛掉 + return pose, det_result \ No newline at end of file diff --git a/SCAIL-Pose/LICENSE b/SCAIL-Pose/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..feb89430888d83ff5d976fed09d4b144e8d23ddd --- /dev/null +++ b/SCAIL-Pose/LICENSE @@ -0,0 +1,201 @@ + 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 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 legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, 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 + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright 2025 W Yan, S Ye, Z Yang, J Teng, ZH Dong, K Wen, X Gu, YJ Liu, J Tang + + 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. diff --git a/SCAIL-Pose/NLFPoseExtract/align3d.py b/SCAIL-Pose/NLFPoseExtract/align3d.py new file mode 100644 index 0000000000000000000000000000000000000000..f1323938244798c446ce2ece601004185bccec16 --- /dev/null +++ b/SCAIL-Pose/NLFPoseExtract/align3d.py @@ -0,0 +1,138 @@ +from sympy import beta +import torch +import numpy as np +from scipy.optimize import minimize + + +def solve_new_camera_params_central(three_d_points, focal_length, imshape, new_2d_points): + """ + 通过最小化原始2D投影点和新的2D投影点之间的误差,求解新的相机参数。 + + 参数: + three_d_points (torch.Tensor): N*3 的 3D 点 + focal_length (float): 原始相机的焦距 + imshape (tuple): 图像的尺寸,例如 [512, 896] + original_2d_points (torch.Tensor): N*2 的原始2D投影点 + new_2d_points (torch.Tensor): N*2 的新的2D投影点 + + 返回: + m, n, p, q: 新的相机内参中的参数 + """ + + # 原始相机内参矩阵 + K_orig = np.array([ + [focal_length, 0, imshape[1] / 2], + [0, focal_length, imshape[0] / 2], + [0, 0, 1] + ]) + + # 目标函数:最小化原始投影点和新的投影点之间的误差 + def objective(params): + m, s, p, q = params + # 构建新的相机内参矩阵 + K_new = np.array([ + [focal_length * m , 0, imshape[1] / 2 + p], + [0, focal_length * m * s, imshape[0] / 2 + q], + [0, 0, 1] + ]) + + # 计算新的2D投影点 + new_projections = [] + for point in three_d_points: + X, Y, Z = point + u = (K_new[0, 0] * X / Z) + K_new[0, 2] + v = (K_new[1, 1] * Y / Z) + K_new[1, 2] + new_projections.append([u, v]) + new_projections = np.array(new_projections) + + # 计算原始2D投影点和新的投影点之间的误差 + # 第0个投影点特殊处理 + error0 = np.sum((new_2d_points[:1] - new_projections[:1]) ** 2) + error = np.sum((new_2d_points[1:] - new_projections[1:]) ** 2) + return error0 * 8 + error + + # 初始化参数 m, beta, p, q + initial_params = [1.0, 1.0, 0.0, 0.0] # 初始值 + + # 使用最小二乘法求解 p, q) + result = minimize(objective, initial_params, bounds=[(0.7, 1.4), (0.8, 1.15), (-imshape[1], imshape[1]), (-imshape[0], imshape[0])]) + + # 输出求解结果 + m, s, p, q = result.x + print(f"debug: solved camera params m={m}, s={s}, p={p}, q={q}") + + K_final = np.array([ + [focal_length * m, 0, imshape[1] / 2 + p], + [0, focal_length * m * s, imshape[0] / 2 + q], + [0, 0, 1] + ]) + + + return K_final, m, s + + +def solve_new_camera_params_down(three_d_points, focal_length, imshape, new_2d_points): + """ + 通过最小化原始2D投影点和新的2D投影点之间的误差,求解新的相机参数。 + + 参数: + three_d_points (torch.Tensor): N*3 的 3D 点 + focal_length (float): 原始相机的焦距 + imshape (tuple): 图像的尺寸,例如 [512, 896] + original_2d_points (torch.Tensor): N*2 的原始2D投影点 + new_2d_points (torch.Tensor): N*2 的新的2D投影点 + + 返回: + m, n, p, q: 新的相机内参中的参数 + """ + + # 原始相机内参矩阵 + K_orig = np.array([ + [focal_length, 0, imshape[1] / 2], + [0, focal_length, imshape[0] / 2], + [0, 0, 1] + ]) + + # 目标函数:最小化原始投影点和新的投影点之间的误差 + def objective(params): + m, s, p, q = params + # 构建新的相机内参矩阵 + K_new = np.array([ + [focal_length * m , 0, imshape[1] / 2 + p], + [0, focal_length * m * s, imshape[0] / 2 + q], + [0, 0, 1] + ]) + + # 计算新的2D投影点 + new_projections = [] + for point in three_d_points: + X, Y, Z = point + u = (K_new[0, 0] * X / Z) + K_new[0, 2] + v = (K_new[1, 1] * Y / Z) + K_new[1, 2] + new_projections.append([u, v]) + new_projections = np.array(new_projections) + + # 计算原始2D投影点和新的投影点之间的误差 + # 第0个投影点特殊处理 + error0 = np.sum((new_2d_points[:1] - new_projections[:1]) ** 2) + error = np.sum((new_2d_points[1:] - new_projections[1:]) ** 2) + return error0 + error * 4 + + # 初始化参数 m, beta, p, q + initial_params = [1.0, 1.0, 0.0, 0.0] # 初始值 + + # 使用最小二乘法求解 p, q) + result = minimize(objective, initial_params, bounds=[(0.7, 1.4), (0.8, 1.15), (-imshape[1], imshape[1]), (-imshape[0], imshape[0])]) + + # 输出求解结果 + m, s, p, q = result.x + print(f"debug: solved camera params m={m}, s={s}, p={p}, q={q}") + + K_final = np.array([ + [focal_length * m, 0, imshape[1] / 2 + p], + [0, focal_length * m * s, imshape[0] / 2 + q], + [0, 0, 1] + ]) + + + return K_final, m, s \ No newline at end of file diff --git a/SCAIL-Pose/NLFPoseExtract/debug_nlf.py b/SCAIL-Pose/NLFPoseExtract/debug_nlf.py new file mode 100644 index 0000000000000000000000000000000000000000..a58a5beda6864721c68f663b9c41891830f9abd6 --- /dev/null +++ b/SCAIL-Pose/NLFPoseExtract/debug_nlf.py @@ -0,0 +1,10 @@ +smpl_path = "/workspace/ywh_data/DataProcessNew/dongman/smpl/ca8637a1f14435234a9f1f9bc31ca4bb.pkl" + +import pickle + +with open(smpl_path, 'rb') as f: + smpl_data = pickle.load(f) + data = smpl_data['pose']['joints3d_nonparam'] + + breakpoint() + print(data) \ No newline at end of file diff --git a/SCAIL-Pose/NLFPoseExtract/extract_nlfpose_batch.py b/SCAIL-Pose/NLFPoseExtract/extract_nlfpose_batch.py new file mode 100644 index 0000000000000000000000000000000000000000..f50e2f3538fffed8efac07211f226a17dcd3d939 --- /dev/null +++ b/SCAIL-Pose/NLFPoseExtract/extract_nlfpose_batch.py @@ -0,0 +1,365 @@ +import os +import sys + +# 动态添加项目根目录到 sys.path,这样就不需要 export PYTHONPATH +current_dir = os.path.dirname(os.path.abspath(__file__)) +project_root = os.path.dirname(current_dir) # SCAIL_Pose 目录 +if project_root not in sys.path: + sys.path.insert(0, project_root) + +import random +from pathlib import Path +import multiprocessing +import numpy as np +import time +from DWPoseProcess.checkUtils import * +from collections import deque +import shutil +import torch +import yaml +import webdataset as wds +from torch.utils.data import DataLoader +from tqdm import tqdm +from functools import partial +import threading +import time +from concurrent.futures import ThreadPoolExecutor, wait, FIRST_COMPLETED, ALL_COMPLETED, TimeoutError +from decord import VideoReader +from fractions import Fraction +import io +import gc +from PIL import Image +from multiprocessing import Process +import json +import jsonlines +from webdataset import TarWriter +import math +import glob +import pickle +import copy +import decord + + +def process_video_nlf(model, vr_frames, bboxes): + # Ensure output directory exists + # pose_results = { + # 'joints3d_nonparam': [], + # } + pose_meta_list = [] + vr_frames = vr_frames.cuda() + height, width = vr_frames.shape[1], vr_frames.shape[2] + result_list = [] + + batch_size = 64 + buffer = torch.zeros( + (batch_size, height, width, 3), + dtype=vr_frames.dtype, + device='cuda' + ) + buffer_count = 0 + with torch.inference_mode(), torch.device('cuda'): + for frame, bbox_list in zip(vr_frames, bboxes): + for bbox in bbox_list: + x1, y1, x2, y2 = bbox + x1_px = max(0, math.floor(x1 * width - width * 0.025)) + y1_px = max(0, math.floor(y1 * height - height * 0.05)) + x2_px = min(width, math.ceil(x2 * width + width * 0.025)) + y2_px = min(height, math.ceil(y2 * height + height * 0.05)) + + cropped_region = frame[y1_px:y2_px, x1_px:x2_px, :] + buffer[buffer_count, y1_px:y2_px, x1_px:x2_px, :] = cropped_region + buffer_count += 1 + + # 一旦 buffer 满了,推理并清空 + if buffer_count == batch_size: + frame_batch = buffer.permute(0, 3, 1, 2) + pred = model.detect_smpl_batched(frame_batch) + if 'joints3d_nonparam' in pred: + result_list.extend(pred['joints3d_nonparam']) + else: + result_list.extend([None] * buffer_count) + + buffer.zero_() + buffer_count = 0 + + # 处理最后不满一批的残余 + if buffer_count > 0: + frame_batch = buffer[:buffer_count].permute(0, 3, 1, 2) + pred = model.detect_smpl_batched(frame_batch) + if 'joints3d_nonparam' in pred: + result_list.extend(pred['joints3d_nonparam']) + else: + result_list.extend([None] * buffer_count) + + index = 0 + for bbox_list in bboxes: + n = len(bbox_list) + pose_meta_list.append({"video_height": height, "video_width": width, "bboxes": bbox_list, "nlfpose": result_list[index : index + n]}) + index += n + + del buffer # 删除 Python 引用 + torch.cuda.empty_cache() + return pose_meta_list + + +def process_video_multi_nlf(model, vr_frames_list): # vr_frames_list里支持1-3人 + # Ensure output directory exists + # pose_results = { + # 'joints3d_nonparam': [], + # } + pose_meta_list = [] + vr_frames_first = vr_frames_list[0].cuda() + # vr_frames_second = vr_frames_second.cuda() + height, width = vr_frames_first.shape[1], vr_frames_first.shape[2] + result_list = [] + + batch_size = 64 + buffer = torch.zeros( + (batch_size, height, width, 3), + dtype=vr_frames_first.dtype, + device='cuda' + ) + buffer_count = 0 + with torch.inference_mode(), torch.device('cuda'): + for frame_idx in range(len(vr_frames_first)): + for person_idx in range(len(vr_frames_list)): + buffer[buffer_count, :, :, :] = vr_frames_first[frame_idx] if person_idx == 0 else vr_frames_list[person_idx][frame_idx] + buffer_count += 1 + # 一旦 buffer 满了,推理并清空 + if buffer_count == batch_size: + frame_batch = buffer.permute(0, 3, 1, 2) + pred = model.detect_smpl_batched(frame_batch) + if 'joints3d_nonparam' in pred: + result_list.extend(pred['joints3d_nonparam']) + else: + result_list.extend([None] * buffer_count) + + buffer.zero_() + buffer_count = 0 + + # 处理最后不满一批的残余 + if buffer_count > 0: + frame_batch = buffer[:buffer_count].permute(0, 3, 1, 2) + pred = model.detect_smpl_batched(frame_batch) + if 'joints3d_nonparam' in pred: + result_list.extend(pred['joints3d_nonparam']) + else: + result_list.extend([None] * buffer_count) + + index = 0 + length_step = len(vr_frames_list) + for _ in range(len(vr_frames_first)): + pose_meta_list.append({"video_height": height, "video_width": width, "bboxes": None, "nlfpose": result_list[index : index + length_step]}) + index += length_step + + del buffer # 删除 Python 引用 + torch.cuda.empty_cache() + return pose_meta_list + + +def process_video_nlf_original(model, vr_frames): + # Ensure output directory exists + # pose_results = { + # 'joints3d_nonparam': [], + # } + pose_meta_list = [] + vr_frames = vr_frames.cuda() + height, width = vr_frames.shape[1], vr_frames.shape[2] + result_list = [] + people_count_list = [] + + batch_size = 64 + buffer = torch.zeros( + (batch_size, height, width, 3), + dtype=vr_frames.dtype, + device='cuda' + ) + buffer_count = 0 + with torch.inference_mode(), torch.device('cuda'): + for frame in vr_frames: + buffer[buffer_count] = frame + buffer_count += 1 + + # 一旦 buffer 满了,推理并清空 + if buffer_count == batch_size: + frame_batch = buffer.permute(0, 3, 1, 2) + pred = model.detect_smpl_batched(frame_batch) + if 'joints3d_nonparam' in pred: + result_list.extend(pred['joints3d_nonparam']) + else: + result_list.extend([None] * buffer_count) + + buffer.zero_() + buffer_count = 0 + + # 处理最后不满一批的残余 + if buffer_count > 0: + frame_batch = buffer[:buffer_count].permute(0, 3, 1, 2) + pred = model.detect_smpl_batched(frame_batch) + if 'joints3d_nonparam' in pred: + result_list.extend(pred['joints3d_nonparam']) + else: + result_list.extend([None] * buffer_count) + + index = 0 + for index in range(len(vr_frames)): + pose_meta_list.append({"video_height": height, "video_width": width, "bboxes": None, "nlfpose": result_list[index]}) + + del buffer # 删除 Python 引用 + torch.cuda.empty_cache() + return pose_meta_list + + +def process_fn_video(src, bbox_dir): + worker_info = torch.utils.data.get_worker_info() + for i, r in enumerate(src): + if worker_info is not None: + if i % worker_info.num_workers != worker_info.id: + continue + key = r['__key__'] + mp4_bytes = r.get("mp4", None) + + try: + decord.bridge.set_bridge("torch") + vr = VideoReader(io.BytesIO(mp4_bytes)) # 这里都是原视频,没有动的 + frames = vr.get_batch(range(len(vr))) + frames = torch.from_numpy(frames) if type(frames) is not torch.Tensor else frames + bbox_path = os.path.join(bbox_dir, key + '.pt') + if os.path.exists(bbox_path): + bboxes = torch.load(bbox_path) + else: + print('no bboxes file: ', key) + continue + except Exception as e: + print(e) + print('load video error: ', key) + continue + item = {'__key__': key, 'frames': frames, 'bboxes': bboxes} + yield item + + +def producer_worker_wds(tar_paths, save_dir_bbox, task_queue): + for tar_path in tar_paths: + produce_nlfpose(tar_path, save_dir_bbox, task_queue) + + +def produce_nlfpose(wds_path, save_dir_bbox, task_queue): + dataset = wds.DataPipeline( + wds.SimpleShardList(wds_path, seed=None), + wds.tarfile_to_samples(), + partial(process_fn_video, bbox_dir=save_dir_bbox), + ) + dataloader = DataLoader(dataset, batch_size=1, num_workers=4, shuffle=False, collate_fn=lambda x: x[0]) + for data in tqdm(dataloader): + task_queue.put(data) + +def gpu_worker(task_queue, save_dir_smpl): + model = torch.jit.load("/workspace/yanwenhao/dwpose_draw/NLFPoseExtract/nlf_l_multi_0.3.2.torchscript").cuda().eval() + while True: + item = task_queue.get() + if item is None: + break + try: + frames = item['frames'] + key = item['__key__'] + bboxes = item['bboxes'] + output_data = process_video_nlf(model, frames, bboxes) + + with open(os.path.join(save_dir_smpl, key + '.pkl'), 'wb') as f: + pickle.dump(output_data, f) + except Exception as e: + print(f"Task failed: {e}") + + +def load_config(config_path): + with open(config_path, 'r') as f: + config = yaml.safe_load(f) + return config + + +# def process_tar_debug(wds_path): +# model = torch.jit.load("/workspace/yanwenhao/dwpose_draw/NLFPoseExtract/nlf_l_multi_0.3.2.torchscript").cuda().eval() +# dataset = wds.DataPipeline( +# wds.SimpleShardList(wds_path, seed=None), +# wds.tarfile_to_samples(), +# partial(process_fn_video), +# ) +# dataloader = DataLoader(dataset, batch_size=1, num_workers=4, shuffle=False, collate_fn=lambda x: x[0]) +# for data in tqdm(dataloader): +# item = data +# if item is None: +# break +# try: +# frames = item['frames'] +# key = item['__key__'] +# bboxes = torch.load(os.path.join(save_dir_bbox, key + '.pt')) +# output_data = process_video_nlf(model, frames, bboxes) + +# with open(os.path.join(save_dir_smpl, key + '.pkl'), 'wb') as f: +# pickle.dump(output_data, f) + +# except Exception as e: +# print(f"Task failed: {e}") + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser() + parser.add_argument('--config', type=str, default='video_directories.yaml', + help='Path to YAML configuration file') + parser.add_argument('--input_root', type=str, default='/workspace/ywh_data/pose_pack_wds_0923add_step1', + help='Input root') + parser.add_argument('--local_rank', type=int, default=0, + help='Local rank') + parser.add_argument('--world_size', type=int, default=1, + help='World size') + + args = parser.parse_args() + config = load_config(args.config) + os.environ['CUDA_VISIBLE_DEVICES'] = str(args.local_rank) + + video_root = config.get('video_root', '') + + + save_dir_smpl = os.path.join(video_root, 'smpl') + save_dir_bbox = os.path.join(video_root, 'bboxes') + os.makedirs(save_dir_smpl, exist_ok=True) + + + processes = [] # 存储进程的列表 + max_queue_size = 32 + task_queue = multiprocessing.Queue(maxsize=max_queue_size) + # Split wds_list into chunks + input_root = os.path.join(args.input_root, os.path.basename(os.path.normpath(video_root))) + input_tar_paths = glob.glob(os.path.join(input_root, "**", "*.tar"), recursive=True) + input_tar_paths = sorted(input_tar_paths) + input_tar_paths_for_the_rank = input_tar_paths[args.local_rank::args.world_size] + + # 并行流程 + p = multiprocessing.Process(target=gpu_worker, args=(task_queue, save_dir_smpl)) + p.start() + + producer_worker_wds(input_tar_paths_for_the_rank, save_dir_bbox, task_queue) + for _ in range(max_queue_size): + task_queue.put(None) + + p.join(timeout=6000) + if p.is_alive(): + print("Warning: GPU worker process did not finish within the expected time") + p.terminate() + + # 串行debug + # for wds_path in input_tar_paths_for_the_rank: + # process_tar_debug(wds_path) + + + + + + + + + + + diff --git a/SCAIL-Pose/NLFPoseExtract/nlf_draw.py b/SCAIL-Pose/NLFPoseExtract/nlf_draw.py new file mode 100644 index 0000000000000000000000000000000000000000..594ea1e28463932d70b8e9828ebf7bd31b720757 --- /dev/null +++ b/SCAIL-Pose/NLFPoseExtract/nlf_draw.py @@ -0,0 +1,135 @@ +import cv2 +import numpy as np +import math +from PIL import Image +from DWPoseProcess.dwpose.util import draw_bodypose + + +def process_data_to_COCO_format(joints): + """Args: + joints: numpy array of shape (24, 2) or (24, 3) + Returns: + new_joints: numpy array of shape (17, 2) or (17, 3) + """ + if joints.ndim != 2: + raise ValueError(f"Expected shape (24,2) or (24,3), got {joints.shape}") + + dim = joints.shape[1] # 2D or 3D + + mapping = { + 15: 0, # head + 12: 1, # neck + 17: 2, # left shoulder + 16: 5, # right shoulder + 19: 3, # left elbow + 18: 6, # right elbow + 21: 4, # left hand + 20: 7, # right hand + 2: 8, # left pelvis + 1: 11, # right pelvis + 5: 9, # left knee + 4: 12, # right knee + 8: 10, # left feet + 7: 13, # right feet + } + + new_joints = np.zeros((18, dim), dtype=joints.dtype) + for src, dst in mapping.items(): + new_joints[dst] = joints[src] + + return new_joints + + +def intrinsic_matrix_from_field_of_view(imshape, fov_degrees:float =55): # nlf default fov_degrees 55 + imshape = np.array(imshape) + fov_radians = fov_degrees * np.array(np.pi / 180) + larger_side = np.max(imshape) + focal_length = larger_side / (np.tan(fov_radians / 2) * 2) + # intrinsic_matrix 3*3 + return np.array([ + [focal_length, 0, imshape[1] / 2], + [0, focal_length, imshape[0] / 2], + [0, 0, 1], + ]) + + +def p3d_to_p2d(point_3d, height, width): # point3d n*num_points*3 + camera_matrix = intrinsic_matrix_from_field_of_view((height,width)) + camera_matrix = np.expand_dims(camera_matrix, axis=0) + camera_matrix = np.expand_dims(camera_matrix, axis=0) # 1*1*3*3 + point_3d = np.expand_dims(point_3d,axis=-1) # n*num_points*3*1 + point_2d = (camera_matrix@point_3d).squeeze(-1) # n*num_points*3 + point_2d[:,:,:2] = point_2d[:,:,:2]/point_2d[:,:,2:3] # 相对位置 + return point_2d[:,:,:] # n*num_points*2 + +def preview_nlf_2d(data): + """ return a list of images """ + height, width = data['video_height'], data['video_width'] + offset = [height, width, 0] + np_images = [] + for image_result in data['pose']['joints3d_nonparam']: + final_canvas = np.zeros(shape=(offset[0], offset[1], 3), dtype=np.uint8) + for joints3d in image_result: # 每个人的pose + canvas = np.zeros(shape=(offset[0], offset[1], 3), dtype=np.uint8) + joints3d = joints3d.cpu().numpy() + joints2d = p3d_to_p2d(joints3d, offset[0], offset[1]) + joints2d = joints2d[0][:, :2] + joints2d[:, 0] = joints2d[:, 0] / offset[1] # x坐标归一化 + joints2d[:, 1] = joints2d[:, 1] / offset[0] # y坐标归一化 + joints2d = process_data_to_COCO_format(joints2d) + subset = np.expand_dims(np.concatenate([np.arange(14), [-1, -1, -1, -1]]), axis=0) + canvas = draw_bodypose(canvas, joints2d, subset) + final_canvas = final_canvas + canvas + np_images.append(final_canvas) + + return np_images + + + +def preview_nlf_2d_ori(nlf_results): + """ return a list of images """ + height, width = nlf_results[0]['video_height'], nlf_results[0]['video_width'] + offset = [height, width, 0] + np_images = [] + for single_result in nlf_results: + image_result = single_result['nlfpose'] + final_canvas = np.zeros(shape=(offset[0], offset[1], 3), dtype=np.uint8) + for joints3d in image_result: # 每个人的pose + canvas = np.zeros(shape=(offset[0], offset[1], 3), dtype=np.uint8) + joints3d = joints3d.cpu().numpy() + joints2d = p3d_to_p2d(joints3d, offset[0], offset[1]) + joints2d = joints2d[0][:, :2] + joints2d[:, 0] = joints2d[:, 0] / offset[1] # x坐标归一化 + joints2d[:, 1] = joints2d[:, 1] / offset[0] # y坐标归一化 + joints2d = process_data_to_COCO_format(joints2d) + subset = np.expand_dims(np.concatenate([np.arange(14), [-1, -1, -1, -1]]), axis=0) + canvas = draw_bodypose(canvas, joints2d, subset) + final_canvas = final_canvas + canvas + np_images.append(final_canvas) + + return np_images + + +def preview_nlf_2d_new(nlf_results): + """ return a list of images """ + height, width = nlf_results[0]['video_height'], nlf_results[0]['video_width'] + offset = [height, width, 0] + np_images = [] + for bbox_result in nlf_results: + final_canvas = np.zeros(shape=(offset[0], offset[1], 3), dtype=np.uint8) + bbox_result = bbox_result['nlfpose'] + for image_result in bbox_result: + for joints3d in image_result: # 每个人的pose + canvas = np.zeros(shape=(offset[0], offset[1], 3), dtype=np.uint8) + joints3d = joints3d.cpu().numpy() + joints2d = p3d_to_p2d(joints3d, offset[0], offset[1]) + joints2d = joints2d[0][:, :2] + joints2d[:, 0] = joints2d[:, 0] / offset[1] # x坐标归一化 + joints2d[:, 1] = joints2d[:, 1] / offset[0] # y坐标归一化 + joints2d = process_data_to_COCO_format(joints2d) + subset = np.expand_dims(np.concatenate([np.arange(14), [-1, -1, -1, -1]]), axis=0) + canvas = draw_bodypose(canvas, joints2d, subset) + final_canvas = final_canvas + canvas + np_images.append(final_canvas) + + return np_images \ No newline at end of file diff --git a/SCAIL-Pose/NLFPoseExtract/nlf_render.py b/SCAIL-Pose/NLFPoseExtract/nlf_render.py new file mode 100644 index 0000000000000000000000000000000000000000..3136751084642d9dd74ec6d07e2e731eff699efa --- /dev/null +++ b/SCAIL-Pose/NLFPoseExtract/nlf_render.py @@ -0,0 +1,658 @@ +import cv2 +import numpy as np +import math +from PIL import Image +from render_3d.taichi_cylinder import render_whole +from NLFPoseExtract.nlf_draw import intrinsic_matrix_from_field_of_view, process_data_to_COCO_format, preview_nlf_2d, p3d_to_p2d +from concurrent.futures import ProcessPoolExecutor, as_completed +from pose_draw.draw_pose_utils import draw_pose_to_canvas_np, scale_image_hw_keep_size +import pose_draw.draw_utils as draw_utils +import torch.multiprocessing as mp +import os +os.environ['PYOPENGL_PLATFORM'] = 'osmesa' +import copy +import random +import torch +try: + import moviepy.editor as mpy +except Exception: + import moviepy as mpy + +def p3d_single_p2d(points, intrinsic_matrix): + X, Y, Z = points[0], points[1], points[2] + u = (intrinsic_matrix[0, 0] * X / Z) + intrinsic_matrix[0, 2] + v = (intrinsic_matrix[1, 1] * Y / Z) + intrinsic_matrix[1, 2] + u_np = u.cpu().numpy() + v_np = v.cpu().numpy() + return np.array([u_np, v_np]) + +def scale_around_center(points, center, dim, scale=1.0): + return (points[:, dim] - center[dim]) * scale + center[dim] + +def shift_dwpose_according_to_nlf(smpl_poses, aligned_poses, ori_intrinstics, modified_intrinstics, height, width, scale_x = 1.0, scale_y = 1.0): + ########## warning: 会改变body; shift 之后 body是不准的 ########## + for i in range(len(smpl_poses)): + persons_joints_list = smpl_poses[i] + poses_list = aligned_poses[i] + # 对里面每一个人,取关节并进行变形;并且修改2d;如果3d不存在,把2d的手/脸也去掉 + for person_idx, person_joints in enumerate(persons_joints_list): + face = poses_list["faces"][person_idx] + right_hand = poses_list["hands"][2 * person_idx] + left_hand = poses_list["hands"][2 * person_idx + 1] + candidate = poses_list["bodies"]["candidate"][person_idx] + # 注意,这里不是coco format + person_joint_15_2d_shift = p3d_single_p2d(person_joints[15], modified_intrinstics) - p3d_single_p2d(person_joints[15], ori_intrinstics) if person_joints[15, 2] > 0.01 else np.array([0.0, 0.0]) # face + person_joint_20_2d_shift = p3d_single_p2d(person_joints[20], modified_intrinstics) - p3d_single_p2d(person_joints[20], ori_intrinstics) if person_joints[20, 2] > 0.01 else np.array([0.0, 0.0]) # right hand + person_joint_21_2d_shift = p3d_single_p2d(person_joints[21], modified_intrinstics) - p3d_single_p2d(person_joints[21], ori_intrinstics) if person_joints[21, 2] > 0.01 else np.array([0.0, 0.0]) # left hand + + face[:, 0] += person_joint_15_2d_shift[0] / width + face[:, 1] += person_joint_15_2d_shift[1] / height + right_hand[:, 0] += person_joint_20_2d_shift[0] / width + right_hand[:, 1] += person_joint_20_2d_shift[1] / height + left_hand[:, 0] += person_joint_21_2d_shift[0] / width + left_hand[:, 1] += person_joint_21_2d_shift[1] / height + candidate[:, 0] += person_joint_15_2d_shift[0] / width + candidate[:, 1] += person_joint_15_2d_shift[1] / height + + scales = [scale_x, scale_y] + # apply camera scale around wrist (hand[0]). + for dim in [0,1]: + right_hand[:, dim] = scale_around_center(right_hand, right_hand[0, :], dim=dim, scale=scales[dim]) + left_hand[:, dim] = scale_around_center(left_hand, left_hand[0, :], dim=dim, scale=scales[dim]) + +def get_single_pose_cylinder_specs(args): + """渲染单个pose的辅助函数,用于并行处理""" + idx, pose, focal, princpt, height, width, colors, limb_seq, draw_seq = args + cylinder_specs = [] + + for joints3d in pose: # 多人 + joints3d = joints3d.cpu().numpy() + joints3d = process_data_to_COCO_format(joints3d) + for line_idx in draw_seq: + line = limb_seq[line_idx] + start, end = line[0], line[1] + if np.sum(joints3d[start]) == 0 or np.sum(joints3d[end]) == 0: + continue + else: + cylinder_specs.append((joints3d[start], joints3d[end], colors[line_idx])) + return cylinder_specs + +def get_single_pose_cylinder_specs_mono(args): + """渲染单个pose的辅助函数,用于并行处理""" + idx, pose, ori_pose, binary_frame, intrinsic_matrix, height, width, limb_seq, draw_seq = args + cylinder_specs = [] + + for joints3d, ori_joint3d in zip(pose, ori_pose): # 多人 + joints3d = joints3d.cpu().numpy() + joints3d = process_data_to_COCO_format(joints3d) + ori_joint3d = ori_joint3d.cpu().numpy() + ori_joint3d = process_data_to_COCO_format(ori_joint3d) + specific_color = locate_binary_color(binary_frame, ori_joint3d, height, width) # 通过3D点的2D投影的像素位置,计算原本这个人对应的颜色 + for line_idx in draw_seq: + line = limb_seq[line_idx] + start, end = line[0], line[1] + if np.sum(joints3d[start]) == 0 or np.sum(joints3d[end]) == 0: + continue + else: + cylinder_specs.append((joints3d[start], joints3d[end], specific_color)) + return cylinder_specs + +def get_single_pose_cylinder_specs_colored(args): + """直接使用传入的每人颜色渲染,不做颜色查找。""" + idx, pose, person_colors, limb_seq, draw_seq = args + cylinder_specs = [] + for person_idx, joints3d in enumerate(pose): + joints3d = joints3d.cpu().numpy() + joints3d = process_data_to_COCO_format(joints3d) + color = person_colors[person_idx] if person_idx < len(person_colors) else [0, 0, 0, 1] + for line_idx in draw_seq: + line = limb_seq[line_idx] + start, end = line[0], line[1] + if np.sum(joints3d[start]) == 0 or np.sum(joints3d[end]) == 0: + continue + cylinder_specs.append((joints3d[start], joints3d[end], color)) + return cylinder_specs + + +def locate_binary_color(binary_frame, ori_joint3d, height, width): + """通过3D点的2D投影的像素位置,计算原本这个人对应的颜色。 + ori_joint3d: COCO format (18, 3) numpy array + binary_frame: H x W x 3,BGR,像素颜色只有6种纯色之一 + """ + key_joint_indices = [1, 2, 5, 8, 11] # neck, left shoulder, right shoulder, left pelvis, right pelvis + key_joints_3d = ori_joint3d[key_joint_indices] # (5, 3) + valid_flag = key_joints_3d[:, 2] > 0.0001 + + point_2d = p3d_to_p2d(key_joints_3d[np.newaxis], height, width)[0] # (5, 3) + + sampled = [] + for is_valid, p2d in zip(valid_flag, point_2d): + if not is_valid: + continue + u = int(round(p2d[0])) + v = int(round(p2d[1])) + if 0 <= u < width and 0 <= v < height: + sampled.append(binary_frame[v, u].astype(np.float32)) + + if len(sampled) == 0: + return [0, 0, 0, 1] + + avg_color = np.mean(sampled, axis=0) + binarized = (avg_color > 127).astype(np.float32) * 1.0 + + return [binarized[0], binarized[1], binarized[2], 1] + +def collect_smpl_poses(data): + uncollected_smpl_poses = [item['nlfpose'] for item in data] + smpl_poses = [[] for _ in range(len(uncollected_smpl_poses))] + for frame_idx in range(len(uncollected_smpl_poses)): + for person_idx in range(len(uncollected_smpl_poses[frame_idx])): # 每个人(每个bbox)只给出一个pose + if len(uncollected_smpl_poses[frame_idx][person_idx]) > 0: # 有返回的骨骼 + smpl_poses[frame_idx].append(uncollected_smpl_poses[frame_idx][person_idx][0]) + else: + smpl_poses[frame_idx].append(torch.zeros((24, 3), dtype=torch.float32)) # 没有检测到人,就放一个全0的 + + return smpl_poses + + + +def collect_smpl_poses_samurai(data): + uncollected_smpl_poses = [item['nlfpose'] for item in data] + smpl_poses_first = [[] for _ in range(len(uncollected_smpl_poses))] + smpl_poses_second = [[] for _ in range(len(uncollected_smpl_poses))] + + for frame_idx in range(len(uncollected_smpl_poses)): + for person_idx in range(len(uncollected_smpl_poses[frame_idx])): # 每个人(每个bbox)只给出一个pose + if len(uncollected_smpl_poses[frame_idx][person_idx]) > 0: # 有返回的骨骼 + if person_idx == 0: + smpl_poses_first[frame_idx].append(uncollected_smpl_poses[frame_idx][person_idx][0]) + elif person_idx == 1: + smpl_poses_second[frame_idx].append(uncollected_smpl_poses[frame_idx][person_idx][0]) + else: + if person_idx == 0: + smpl_poses_first[frame_idx].append(torch.zeros((24, 3), dtype=torch.float32)) # 没有检测到人,就放一个全0的 + elif person_idx == 1: + smpl_poses_second[frame_idx].append(torch.zeros((24, 3), dtype=torch.float32)) + + return smpl_poses_first, smpl_poses_second + + + + +def render_nlf_as_images(data, poses, reshape_pool=None, intrinsic_matrix=None, draw_2d=True, aug_2d=False, aug_cam=False, binary_mask=None, person_colors=None, palette_offset=0): + """ return a list of images """ + height, width = data[0]['video_height'], data[0]['video_width'] + video_length = len(data) + + base_colors_255_dict = { + # Warm Colors for Right Side (R.) - Red, Orange, Yellow + "Red": [255, 0, 0], + "Orange": [255, 85, 0], + "Golden Orange": [255, 170, 0], + "Yellow": [255, 240, 0], + "Yellow-Green": [180, 255, 0], + # Cool Colors for Left Side (L.) - Green, Blue, Purple + "Bright Green": [0, 255, 0], + "Light Green-Blue": [0, 255, 85], + "Aqua": [0, 255, 170], + "Cyan": [0, 255, 255], + "Sky Blue": [0, 170, 255], + "Medium Blue": [0, 85, 255], + "Pure Blue": [0, 0, 255], + "Purple-Blue": [85, 0, 255], + "Medium Purple": [170, 0, 255], + # Neutral/Central Colors (e.g., for Neck, Nose, Eyes, Ears) + "Grey": [150, 150, 150], + "Pink-Magenta": [255, 0, 170], + "Dark Pink": [255, 0, 85], + "Violet": [100, 0, 255], + "Dark Violet": [50, 0, 255], + } + + ordered_colors_255 = [ + base_colors_255_dict["Red"], # Neck -> R. Shoulder (Red) + base_colors_255_dict["Cyan"], # Neck -> L. Shoulder (Cyan) + base_colors_255_dict["Orange"], # R. Shoulder -> R. Elbow (Orange) + base_colors_255_dict["Golden Orange"], # R. Elbow -> R. Wrist (Golden Orange) + base_colors_255_dict["Sky Blue"], # L. Shoulder -> L. Elbow (Sky Blue) + base_colors_255_dict["Medium Blue"], # L. Elbow -> L. Wrist (Medium Blue) + base_colors_255_dict["Yellow-Green"], # Neck -> R. Hip ( Yellow-Green) + base_colors_255_dict["Bright Green"], # R. Hip -> R. Knee (Bright Green - transitioning warm to cool spectrum) + base_colors_255_dict["Light Green-Blue"], # R. Knee -> R. Ankle (Light Green-Blue - transitioning) + base_colors_255_dict["Pure Blue"], # Neck -> L. Hip (Pure Blue) + base_colors_255_dict["Purple-Blue"], # L. Hip -> L. Knee (Purple-Blue) + base_colors_255_dict["Medium Purple"], # L. Knee -> L. Ankle (Medium Purple) + base_colors_255_dict["Grey"], # Neck -> Nose (Grey) + base_colors_255_dict["Pink-Magenta"], # Nose -> R. Eye (Pink/Magenta) + base_colors_255_dict["Dark Violet"], # R. Eye -> R. Ear (Dark Pink) + base_colors_255_dict["Pink-Magenta"], # Nose -> L. Eye (Violet) + base_colors_255_dict["Dark Violet"], # L. Eye -> L. Ear (Dark Violet) + ] + + limb_seq = [ + [1, 2], # 0 Neck -> R. Shoulder + [1, 5], # 1 Neck -> L. Shoulder + [2, 3], # 2 R. Shoulder -> R. Elbow + [3, 4], # 3 R. Elbow -> R. Wrist + [5, 6], # 4 L. Shoulder -> L. Elbow + [6, 7], # 5 L. Elbow -> L. Wrist + [1, 8], # 6 Neck -> R. Hip + [8, 9], # 7 R. Hip -> R. Knee + [9, 10], # 8 R. Knee -> R. Ankle + [1, 11], # 9 Neck -> L. Hip + [11, 12], # 10 L. Hip -> L. Knee + [12, 13], # 11 L. Knee -> L. Ankle + [1, 0], # 12 Neck -> Nose + [0, 14], # 13 Nose -> R. Eye + [14, 16], # 14 R. Eye -> R. Ear + [0, 15], # 15 Nose -> L. Eye + [15, 17], # 16 L. Eye -> L. Ear + ] + + draw_seq = [0, 2, 3, # Neck -> R. Shoulder -> R. Elbow -> R. Wrist + 1, 4, 5, # Neck -> L. Shoulder -> L. Elbow -> L. Wrist + 6, 7, 8, # Neck -> R. Hip -> R. Knee -> R. Ankle + 9, 10, 11, # Neck -> L. Hip -> L. Knee -> L. Ankle + 12, # Neck -> Nose + 13, 14, # Nose -> R. Eye -> R. Ear + 15, 16, # Nose -> L. Eye -> L. Ear + ] # 从近心端往外扩展 + + colors = [[c / 300 + 0.15 for c in color_rgb] + [0.8] for color_rgb in ordered_colors_255] + + + + # smpl_poses 会在这里被修改 + if poses is not None or binary_mask is not None or person_colors is not None: + # 重新收集poses + smpl_poses = collect_smpl_poses(data) + if binary_mask is not None: + original_smpl_poses = copy.deepcopy(smpl_poses) + if poses is not None: + aligned_poses = copy.deepcopy(poses) # 2d poses + if reshape_pool is not None: + for i in range(video_length): + persons_joints_list = smpl_poses[i] + poses_list = aligned_poses[i] + # 对里面每一个人,取关节并进行变形;并且修改2d;如果3d不存在,把2d的手/脸也去掉 + for person_idx, person_joints in enumerate(persons_joints_list): + candidate = poses_list['bodies']['candidate'][person_idx] + subset = poses_list['bodies']['subset'][person_idx] + face = poses_list["faces"][person_idx] + right_hand = poses_list["hands"][2 * person_idx] + left_hand = poses_list["hands"][2 * person_idx + 1] + reshape_pool.apply_random_reshapes(person_joints, candidate, left_hand, right_hand, face, subset) + else: + smpl_poses = [item['nlfpose'] for item in data] # 主要为了兼容多人评测集;搭配process_video_nlf_original + + + if intrinsic_matrix is None: + intrinsic_matrix = intrinsic_matrix_from_field_of_view((height, width)) + focal_x = intrinsic_matrix[0,0] + focal_y = intrinsic_matrix[1,1] + princpt = (intrinsic_matrix[0,2], intrinsic_matrix[1,2]) # 主点 (cx, cy) + if aug_cam and random.random() < 0.3: + w_shift_factor = random.uniform(-0.04, 0.04) + h_shift_factor = random.uniform(-0.04, 0.04) + princpt = (princpt[0] - w_shift_factor * width, princpt[1] - h_shift_factor * height) # princpt变化和点的变化相反 + new_intrinsic_matrix = copy.deepcopy(intrinsic_matrix) + new_intrinsic_matrix[0,2] = princpt[0] + new_intrinsic_matrix[1,2] = princpt[1] + shift_dwpose_according_to_nlf(smpl_poses, aligned_poses, intrinsic_matrix, new_intrinsic_matrix, height, width) + + # person_colors 传入时,为每人生成独立肢体颜色方案(同 render_multi_nlf_as_images 的两套配色) + if person_colors is not None: + _palettes_255 = [ + # Person 0: 浅色调 + [[255,150,150],[180,230,240],[255,180,140],[255,215,150],[160,200,255],[100,120,255], + [200,255,100],[100,255,100],[140,255,180],[120,140,255],[180, 90,255],[190,120,255], + [210,210,210],[255,120,200],[130, 80,255],[255,120,200],[130, 80,255]], + # Person 1: 饱和色调 + [[255, 20, 20],[ 0,230,255],[255, 60, 0],[255,110, 0],[ 0,130,255],[ 0, 70,255], + [160,255, 40],[ 0,255, 50],[ 0,255,100],[ 0, 0,255],[ 80, 0,255],[160, 0,255], + [130,130,130],[255, 0,150],[ 60, 0,255],[255, 0,150],[ 60, 0,255]], + ] + colors_per_person = [ + [[c / 300 + 0.15 for c in rgb] + [0.8] + for rgb in _palettes_255[(p + palette_offset) % len(_palettes_255)]] + for p in range(len(person_colors)) + ] + + # 串行获取每一帧的cylinder_specs + cylinder_specs_list = [] + cylinder_specs_list_mono = [] + for i in range(video_length): + if person_colors is not None: + cylinder_specs = [] + for p_idx, person_pose in enumerate(smpl_poses[i]): + p_limb_colors = colors_per_person[p_idx] if p_idx < len(colors_per_person) else colors + cylinder_specs.extend(get_single_pose_cylinder_specs( + (i, [person_pose], None, None, None, None, p_limb_colors, limb_seq, draw_seq))) + else: + cylinder_specs = get_single_pose_cylinder_specs((i, smpl_poses[i], None, None, None, None, colors, limb_seq, draw_seq)) + cylinder_specs_list.append(cylinder_specs) + if person_colors is not None: + cylinder_specs_colored = get_single_pose_cylinder_specs_colored((i, smpl_poses[i], person_colors, limb_seq, draw_seq)) + cylinder_specs_list_mono.append(cylinder_specs_colored) + elif binary_mask is not None: + cylinder_specs_mono = get_single_pose_cylinder_specs_mono((i, smpl_poses[i], original_smpl_poses[i], binary_mask[i], intrinsic_matrix, height, width, limb_seq, draw_seq)) + cylinder_specs_list_mono.append(cylinder_specs_mono) + + + frames_np_rgba = render_whole(cylinder_specs_list, H=height, W=width, fx=focal_x, fy=focal_y, cx=princpt[0], cy=princpt[1]) + frames_np_rgba_mono = render_whole(cylinder_specs_list_mono, H=height, W=width, fx=focal_x, fy=focal_y, cx=princpt[0], cy=princpt[1], use_specular=False) if (binary_mask is not None or person_colors is not None) else None + + bg_color = np.array([0, 0, 0], dtype=np.uint8) + for frame in frames_np_rgba: + bg_mask = frame[:, :, 3] == 0 + frame[:, :, :3][bg_mask] = bg_color + + scale_h = random.uniform(0.85, 1.15) + scale_w = random.uniform(0.85, 1.15) + rescale_flag = random.random() < 0.4 if reshape_pool is not None else False + + if poses is not None and draw_2d: + canvas_2d = draw_pose_to_canvas_np(aligned_poses, pool=None, H=height, W=width, reshape_scale=0, show_feet_flag=False, show_body_flag=False, show_cheek_flag=True, dw_hand=True) + for i in range(len(frames_np_rgba)): + frame_img = frames_np_rgba[i] + canvas_img = canvas_2d[i] + mask = canvas_img != 0 + frame_img[:, :, :3][mask] = canvas_img[mask] + frames_np_rgba[i] = frame_img # no alpha blending + # 在 mono 版上用每人的颜色画 cheek/hand/face 2D 关键点 + if frames_np_rgba_mono is not None and person_colors is not None: + poses_list = aligned_poses[i] + n_draw = min(len(poses_list['bodies']['candidate']), len(person_colors)) + for p_idx in range(n_draw): + temp_canvas = np.zeros((height, width, 3), dtype=np.uint8) + p_candidate = poses_list['bodies']['candidate'][p_idx] + p_subset = poses_list['bodies']['subset'][p_idx:p_idx+1] + p_faces = poses_list['faces'][p_idx:p_idx+1] + p_hands = poses_list['hands'][2*p_idx:2*p_idx+2] + temp_canvas = draw_utils.draw_bodypose_augmentation(temp_canvas, p_candidate, p_subset, drop_aug=False, shift_aug=False, all_cheek_aug=True) + temp_canvas = draw_utils.draw_handpose(temp_canvas, p_hands) + temp_canvas = draw_utils.draw_facepose(temp_canvas, p_faces, optimized_face=True) + mask_2d = np.any(temp_canvas != 0, axis=-1) + mono_color = [int(c * 255) for c in person_colors[p_idx][:3]] + frames_np_rgba_mono[i][:, :, :3][mask_2d] = mono_color + if aug_2d: + if rescale_flag: + frames_np_rgba[i] = scale_image_hw_keep_size(frames_np_rgba[i], scale_h, scale_w) + border_mask = frames_np_rgba[i][:, :, 3] == 0 + frames_np_rgba[i][:, :, :3][border_mask] = bg_color + if reshape_pool is not None and random.random() < 0.04: + # 4%的概率完全消除某些帧,两组同步 + frames_np_rgba[i][:, :, :3] = bg_color + if frames_np_rgba_mono is not None: + frames_np_rgba_mono[i][:, :, 0:3] = 0 + if frames_np_rgba_mono is not None and rescale_flag: + frames_np_rgba_mono[i] = scale_image_hw_keep_size(frames_np_rgba_mono[i], scale_h, scale_w) + else: + for i in range(len(frames_np_rgba)): + if aug_2d: + if rescale_flag: + frames_np_rgba[i] = scale_image_hw_keep_size(frames_np_rgba[i], scale_h, scale_w) + border_mask = frames_np_rgba[i][:, :, 3] == 0 + frames_np_rgba[i][:, :, :3][border_mask] = bg_color + if reshape_pool is not None and random.random() < 0.04: + # 4%的概率完全消除某些帧,两组同步 + frames_np_rgba[i][:, :, :3] = bg_color + if frames_np_rgba_mono is not None: + frames_np_rgba_mono[i][:, :, 0:3] = 0 + if frames_np_rgba_mono is not None and rescale_flag: + frames_np_rgba_mono[i] = scale_image_hw_keep_size(frames_np_rgba_mono[i], scale_h, scale_w) + + if binary_mask is not None or person_colors is not None: + return frames_np_rgba, frames_np_rgba_mono + return frames_np_rgba + + + + + + + +def render_multi_nlf_as_images(data, poses, reshape_pool=None, intrinsic_matrix=None, draw_2d=True, aug_2d=False, aug_cam=False): + """ return a list of images """ + height, width = data[0]['video_height'], data[0]['video_width'] + video_length = len(data) + + second_person_base_colors_255_dict = { + # Warm Colors for Right Side (R.) - Red, Orange, Yellow + "Red": [255, 20, 20], + "Orange": [255, 60, 0], + "Golden Orange": [255, 110, 0], + "Yellow": [255, 200, 0], + "Yellow-Green": [160, 255, 40], + + # Cool Colors for Left Side (L.) - Green, Blue, Purple + "Bright Green": [0, 255, 50], + "Light Green-Blue": [0, 255, 100], + "Aqua": [0, 255, 200], + "Cyan": [0, 230, 255], + "Sky Blue": [0, 130, 255], + "Medium Blue": [0, 70, 255], + "Pure Blue": [0, 0, 255], + "Purple-Blue": [80, 0, 255], + "Medium Purple": [160, 0, 255], + + # Neutral/Central Colors (e.g., for Neck, Nose, Eyes, Ears) + "Grey": [130, 130, 130], + "Pink-Magenta": [255, 0, 150], + "Dark Pink": [255, 0, 100], + "Violet": [120, 0, 255], + "Dark Violet": [60, 0, 255], + } + + first_person_base_colors_255_dict = { + # Warm Colors for Right Side (R.) - Red, Orange, Yellow + "Red": [255, 150, 150], + "Orange": [255, 180, 140], + "Golden Orange": [255, 215, 150], + "Yellow": [255, 240, 170], + "Yellow-Green": [200, 255, 100], + + # Cool Colors for Left Side (L.) - Green, Blue, Purple + "Bright Green": [100, 255, 100], + "Light Green-Blue": [140, 255, 180], + "Aqua": [150, 240, 200], + "Cyan": [180, 230, 240], + "Sky Blue": [160, 200, 255], + "Medium Blue": [100, 120, 255], + "Pure Blue": [120, 140, 255], + "Purple-Blue": [180, 90, 255], + "Medium Purple": [190, 120, 255], + + # Neutral/Central Colors (e.g., for Neck, Nose, Eyes, Ears) + "Grey": [210, 210, 210], + "Pink-Magenta": [255, 120, 200], + "Dark Pink": [255, 150, 180], + "Violet": [200, 90, 255], + "Dark Violet": [130, 80, 255], + } + + base_colors_255_dict_list = [first_person_base_colors_255_dict, second_person_base_colors_255_dict] + ordered_colors_255_list = [[ + base_colors_255_dict["Red"], # Neck -> R. Shoulder (Red) + base_colors_255_dict["Cyan"], # Neck -> L. Shoulder (Cyan) + base_colors_255_dict["Orange"], # R. Shoulder -> R. Elbow (Orange) + base_colors_255_dict["Golden Orange"], # R. Elbow -> R. Wrist (Golden Orange) + base_colors_255_dict["Sky Blue"], # L. Shoulder -> L. Elbow (Sky Blue) + base_colors_255_dict["Medium Blue"], # L. Elbow -> L. Wrist (Medium Blue) + base_colors_255_dict["Yellow-Green"], # Neck -> R. Hip ( Yellow-Green) + base_colors_255_dict["Bright Green"], # R. Hip -> R. Knee (Bright Green - transitioning warm to cool spectrum) + base_colors_255_dict["Light Green-Blue"], # R. Knee -> R. Ankle (Light Green-Blue - transitioning) + base_colors_255_dict["Pure Blue"], # Neck -> L. Hip (Pure Blue) + base_colors_255_dict["Purple-Blue"], # L. Hip -> L. Knee (Purple-Blue) + base_colors_255_dict["Medium Purple"], # L. Knee -> L. Ankle (Medium Purple) + base_colors_255_dict["Grey"], # Neck -> Nose (Grey) + base_colors_255_dict["Pink-Magenta"], # Nose -> R. Eye (Pink/Magenta) + base_colors_255_dict["Dark Violet"], # R. Eye -> R. Ear (Dark Pink) + base_colors_255_dict["Pink-Magenta"], # Nose -> L. Eye (Violet) + base_colors_255_dict["Dark Violet"], # L. Eye -> L. Ear (Dark Violet) + ] for base_colors_255_dict in base_colors_255_dict_list] + + limb_seq = [ + [1, 2], # 0 Neck -> R. Shoulder + [1, 5], # 1 Neck -> L. Shoulder + [2, 3], # 2 R. Shoulder -> R. Elbow + [3, 4], # 3 R. Elbow -> R. Wrist + [5, 6], # 4 L. Shoulder -> L. Elbow + [6, 7], # 5 L. Elbow -> L. Wrist + [1, 8], # 6 Neck -> R. Hip + [8, 9], # 7 R. Hip -> R. Knee + [9, 10], # 8 R. Knee -> R. Ankle + [1, 11], # 9 Neck -> L. Hip + [11, 12], # 10 L. Hip -> L. Knee + [12, 13], # 11 L. Knee -> L. Ankle + [1, 0], # 12 Neck -> Nose + [0, 14], # 13 Nose -> R. Eye + [14, 16], # 14 R. Eye -> R. Ear + [0, 15], # 15 Nose -> L. Eye + [15, 17], # 16 L. Eye -> L. Ear + ] + + draw_seq = [0, 2, 3, # Neck -> R. Shoulder -> R. Elbow -> R. Wrist + 1, 4, 5, # Neck -> L. Shoulder -> L. Elbow -> L. Wrist + 6, 7, 8, # Neck -> R. Hip -> R. Knee -> R. Ankle + 9, 10, 11, # Neck -> L. Hip -> L. Knee -> L. Ankle + 12, # Neck -> Nose + 13, 14, # Nose -> R. Eye -> R. Ear + 15, 16, # Nose -> L. Eye -> L. Ear + ] # 从近心端往外扩展 + + colors_first = [[c / 300 + 0.15 for c in color_rgb] + [0.8] for color_rgb in ordered_colors_255_list[0]] + colors_second = [[c / 300 + 0.15 for c in color_rgb] + [0.8] for color_rgb in ordered_colors_255_list[1]] + + smpl_poses_first, smpl_poses_second = collect_smpl_poses_samurai(data) + + + if intrinsic_matrix is None: + intrinsic_matrix = intrinsic_matrix_from_field_of_view((height, width)) + focal_x = intrinsic_matrix[0,0] + focal_y = intrinsic_matrix[1,1] + princpt = (intrinsic_matrix[0,2], intrinsic_matrix[1,2]) # 主点 (cx, cy) + + # 串行获取每一帧的cylinder_specs + cylinder_specs_list = [] + for i in range(video_length): + cylinder_specs_first = get_single_pose_cylinder_specs((i, smpl_poses_first[i], None, None, None, None, colors_first, limb_seq, draw_seq)) + cylinder_specs_second = get_single_pose_cylinder_specs((i, smpl_poses_second[i], None, None, None, None, colors_second, limb_seq, draw_seq)) + cylinder_specs = cylinder_specs_first + cylinder_specs_second + cylinder_specs_list.append(cylinder_specs) + + + frames_np_rgba = render_whole(cylinder_specs_list, H=height, W=width, fx=focal_x, fy=focal_y, cx=princpt[0], cy=princpt[1]) + if poses is not None and draw_2d: + aligned_poses = copy.deepcopy(poses) + canvas_2d = draw_pose_to_canvas_np(aligned_poses, pool=None, H=height, W=width, reshape_scale=0, show_feet_flag=False, show_body_flag=False, show_cheek_flag=True, dw_hand=True) + for i in range(len(frames_np_rgba)): + frame_img = frames_np_rgba[i] + canvas_img = canvas_2d[i] + mask = canvas_img != 0 + frame_img[:, :, :3][mask] = canvas_img[mask] + frames_np_rgba[i] = frame_img + + return frames_np_rgba + + +def run_nlf_from_masks(video_frames, masks, colors, model_nlf, nlf_render_path, + nlf_render_mask_path, fps=16, detector=None): + """对每个人用墨绿色背景隔离后提取 NLF 姿态,再分别渲染普通和 mono 结果并保存为 MP4。 + + Args: + video_frames: (T, H, W, 3) uint8 numpy array, RGB + masks: list of (T, H, W) bool ndarray,每人一个 + colors: list of BGR color tuples,与 masks 一一对应 + model_nlf: TorchScript NLF 模型 + nlf_render_path: 普通渲染输出路径(含 2D 关键点叠加) + nlf_render_mask_path: mono 渲染输出路径 + fps: 输出帧率 + detector: DWposeDetector,对原始帧提取多人 2D 关键点 + """ + from NLFPoseExtract.extract_nlfpose_batch import process_video_multi_nlf + + if len(masks) == 0: + print("No masks provided, skipping.") + return + + T, H, W, C = video_frames.shape + dark_green = np.array([0, 100, 0], dtype=np.uint8) + + vr_frames_list = [] + for mask in masks: + person_frames = np.full((T, H, W, C), dark_green, dtype=np.uint8) + person_frames[mask] = video_frames[mask] + vr_frames_list.append(torch.from_numpy(person_frames)) + + nlf_results = process_video_multi_nlf(model_nlf, vr_frames_list) + + poses = None + if detector is not None: + # Per-person DWpose: run detector on each SAM3 person's dark_green-bg crop so the + # 2D keypoints (face/hands/body) align with SAM3 person order. Stack the per-person + # single-person dicts back into multi-person dicts per frame, in SAM3 order. + N = len(masks) + EMPTY_BODY = np.full((24, 2), -1.0, dtype=np.float32) + EMPTY_SUBSET = np.full((24,), -1.0, dtype=np.float32) + EMPTY_FACE = np.full((68, 2), -1.0, dtype=np.float32) + EMPTY_HAND = np.full((21, 2), -1.0, dtype=np.float32) + + per_person_per_frame = [[None] * T for _ in range(N)] + for p_idx in range(N): + person_frames_np = vr_frames_list[p_idx].numpy() # (T, H, W, 3) RGB, dark_green bg + for t in range(T): + pose_dict, _, _ = detector(Image.fromarray(person_frames_np[t])) + per_person_per_frame[p_idx][t] = pose_dict + + poses = [] + for t in range(T): + cand_rows, sub_rows, face_rows = [], [], [] + hand_rows = [] + for p_idx in range(N): + pd = per_person_per_frame[p_idx][t] + cands = pd['bodies']['candidate'] + if cands is not None and len(cands) > 0: + cand_rows.append(cands[0]) + sub_rows.append(pd['bodies']['subset'][0]) + face_rows.append(pd['faces'][0]) + hand_rows.append(pd['hands'][0]) + hand_rows.append(pd['hands'][1]) + else: + cand_rows.append(EMPTY_BODY) + sub_rows.append(EMPTY_SUBSET) + face_rows.append(EMPTY_FACE) + hand_rows.append(EMPTY_HAND) + hand_rows.append(EMPTY_HAND) + poses.append({ + 'bodies': { + 'candidate': np.stack(cand_rows, axis=0), + 'subset': np.stack(sub_rows, axis=0), + }, + 'faces': np.stack(face_rows, axis=0), + 'hands': np.stack(hand_rows, axis=0), + }) + + person_colors_rgba = [] + for bgr in colors: + b, g, r = bgr[0] / 255.0, bgr[1] / 255.0, bgr[2] / 255.0 + person_colors_rgba.append([r, g, b, 1.0]) + + palette_offset = 1 if len(masks) == 1 else 0 + frames_regular, frames_mono = render_nlf_as_images( + copy.deepcopy(nlf_results), poses=copy.deepcopy(poses), + reshape_pool=None, intrinsic_matrix=None, + draw_2d=True, aug_2d=False, aug_cam=False, + person_colors=person_colors_rgba, palette_offset=palette_offset, + ) + + for out_path in (nlf_render_path, nlf_render_mask_path): + out_dir = os.path.dirname(out_path) + if out_dir: + os.makedirs(out_dir, exist_ok=True) + + frames_regular_rgb = [f[:, :, :3] for f in frames_regular] + frames_mono_rgb = [f[:, :, :3] for f in frames_mono] + + mpy.ImageSequenceClip(frames_regular_rgb, fps=fps).write_videofile(nlf_render_path) + mpy.ImageSequenceClip(frames_mono_rgb, fps=fps).write_videofile(nlf_render_mask_path) \ No newline at end of file diff --git a/SCAIL-Pose/NLFPoseExtract/process_animation_aio.py b/SCAIL-Pose/NLFPoseExtract/process_animation_aio.py new file mode 100644 index 0000000000000000000000000000000000000000..181f2dfb8e3768e00517efd88b325c5cd0ebb594 --- /dev/null +++ b/SCAIL-Pose/NLFPoseExtract/process_animation_aio.py @@ -0,0 +1,236 @@ +import os +import sys + +# 动态添加项目根目录到 sys.path,这样就不需要 export PYTHONPATH +current_dir = os.path.dirname(os.path.abspath(__file__)) +project_root = os.path.dirname(current_dir) # SCAIL_Pose 目录 +if project_root not in sys.path: + sys.path.insert(0, project_root) + +import argparse +import glob +import shutil +import time +import traceback + +import numpy as np +import torch +from decord import VideoReader, cpu + +from NLFPoseExtract.v2_helper import ( + find_ref_image, + match_driving_to_ref_by_center, + save_colored_mask_image, + write_colored_mask_video, + write_kept_driving_video, +) + + +def process_one(subdir, video_name, e2e_mode, crop_kind, max_persons, text, + model_nlf, detector, predictor, image_predictor, + crop_margin=0.05): + from TrackSam3.track import get_mask_from_image_via_video, get_mask_from_video + + if crop_kind is not None and not e2e_mode: + raise ValueError("--crop_e2e_bbox / --crop_e2e_mask / --crop_e2e_steady_bbox require --e2e_mode") + + mp4_path = os.path.join(subdir, video_name) + if not os.path.exists(mp4_path): + raise FileNotFoundError(f"No {video_name} found in {subdir}") + + ref_image_path = find_ref_image(subdir) + + out_path_rendered = os.path.join(subdir, 'rendered_v2.mp4') + out_path_mask = os.path.join(subdir, 'rendered_mask_v2.mp4') + ref_mask_path = os.path.join(subdir, 'ref_mask.jpg') + + # 1) Driving is the canonical authority for person count (video tracker is reliable; + # SAM3 image mode tends to return duplicate masks for the same character). + print(f"Getting driving masks from {mp4_path} (text={text})...") + drv_masks, drv_colors = get_mask_from_video( + mp4_path, predictor, max_targets=max_persons, sort_by='x', fixed_colors=None, + text=text, + ) + if len(drv_masks) == 0: + raise RuntimeError(f"No valid persons detected in driving {mp4_path}") + N = len(drv_masks) + print(f"Driving defines {N} person(s) (capped at --max_persons={max_persons}); colors={drv_colors}") + + # 2) Ref provides per-person colors; cap at N and sort left-to-right to match driving. + # Route ref through the video predictor (single-frame mp4 wrapper) — image-mode SAM3 + # often misses small / distant subjects that the video pipeline picks up reliably. + print(f"Getting ref masks from {ref_image_path} (text={text})...") + ref_masks, ref_colors = get_mask_from_image_via_video( + ref_image_path, predictor, max_targets=N, sort_by='x', fixed_colors=None, + text=text, + ) + if len(ref_masks) < N: + print(f" Ref has {len(ref_masks)} person(s) but driving has {N}; " + f"matching driving subset to ref by normalized center ...") + drv_masks, drv_colors = match_driving_to_ref_by_center( + drv_masks, drv_colors, ref_masks) + if len(drv_masks) != len(ref_masks): + raise RuntimeError( + f"Center matching produced {len(drv_masks)} driving tracks but " + f"ref has {len(ref_masks)}; cannot align." + ) + N = len(drv_masks) + print(f" Reduced driving to {N} person(s).") + save_colored_mask_image(ref_masks, ref_colors, ref_mask_path, bg_color=(255, 255, 255)) + print(f" Ref mask saved: {ref_mask_path}") + + # 3) Read driving frames and fps once + vr = VideoReader(mp4_path, ctx=cpu(0)) + fps = vr.get_avg_fps() + fps_int = max(1, int(round(fps))) + video_frames_np = vr.get_batch(list(range(len(vr)))).asnumpy() # (T, H, W, 3) RGB + + # 4) Branch on e2e_mode (+ optional crop_kind for rendered_v2 only) + if e2e_mode: + if crop_kind is not None: + print(f"[e2e_mode+crop_e2e_{crop_kind}] writing kept driving as rendered_v2.mp4 " + f"(bbox_margin={crop_margin}) ...") + write_kept_driving_video(video_frames_np, drv_masks, out_path_rendered, + fps_int, crop_kind=crop_kind, bbox_margin=crop_margin) + else: + print("[e2e_mode] copying driving as rendered_v2.mp4 ...") + shutil.copyfile(mp4_path, out_path_rendered) + print("[e2e_mode] writing colored mask video as rendered_mask_v2.mp4 ...") + write_colored_mask_video(drv_masks, drv_colors, out_path_mask, fps_int) + else: + from NLFPoseExtract.nlf_render import run_nlf_from_masks + print("Running NLF and rendering skeletons ...") + run_nlf_from_masks( + video_frames=video_frames_np, + masks=drv_masks, + colors=drv_colors, + model_nlf=model_nlf, + nlf_render_path=out_path_rendered, + nlf_render_mask_path=out_path_mask, + fps=fps_int, + detector=detector, + ) + + print("Done!") + print(f" Rendered: {out_path_rendered}") + print(f" Rendered mask: {out_path_mask}") + print(f" Ref mask: {ref_mask_path}") + + +if __name__ == '__main__': + parser = argparse.ArgumentParser( + description='SCAIL pose AIO: SAM3 mask + (optional) NLF skeleton render, with ' + 'ref-aligned multi-person colors. Pass exactly one of --subdir ' + '(single example) OR --input_root (batch over a directory of examples).' + ) + src = parser.add_mutually_exclusive_group(required=True) + src.add_argument('--subdir', type=str, default=None, + help='Single-example mode: path to one subdir containing driving video + ref_image. ' + 'Mutually exclusive with --input_root.') + src.add_argument('--input_root', type=str, default=None, + help='Batch mode: directory whose immediate subdirs are each an example. ' + 'Models load once and every subdir is processed in one invocation. ' + 'Mutually exclusive with --subdir.') + parser.add_argument('--video_name', type=str, default='driving.mp4', + choices=['driving.mp4', 'GT.mp4', 'raw.mp4'], + help='Filename of the driving video inside each subdir.') + parser.add_argument('--e2e_mode', action='store_true', + help='If set, skip pose extraction: rendered_v2.mp4 is a copy of the driving video, ' + 'and rendered_mask_v2.mp4 is the colored SAM3 mask video. Recommanded setting to True as it\'s more accurate and easy-to-use than pose-driven for most cases.' + 'Default pose-driven provides more control and interpretability, you can adopt pose-driven for extremely challenging inputs.') + crop_group = parser.add_mutually_exclusive_group() + crop_group.add_argument('--crop_e2e_mask', action='store_true', + help='Sub-flag of --e2e_mode. rendered_v2.mp4 keeps only pixels inside ' + 'each person\'s mask silhouette; rest is black. The behaviour is between pose-driven and crop_e2e_bbox.' + 'For 512p e2e runs or portrait videos, as the main training is in under this resolution, typically the function is not needed.' + 'For 704p horizontal (especially multi-human scenarios as it\'s zero-shot), we notice the crop can be an inference optimization to reduce artifacts.' + 'For other cases the crop may also not be necessary, as the default full e2e is usually better especially in terms of human-object interactions.') + crop_group.add_argument('--crop_e2e_bbox', action='store_true', + help='Sub-flag of --e2e_mode. rendered_v2.mp4 keeps only pixels inside ' + 'each person\'s mask bbox (margin-padded, overlap-merged); rest is black. ') + crop_group.add_argument('--crop_e2e_steady_bbox', action='store_true', + help='Sub-flag of --e2e_mode. One static bbox from the union of all person masks ' + 'across all frames (with margin); the kept rectangle never moves. Usually not needed unless you specifically want behaviour between full e2e and crop_e2e_bbox. ') + # All the expected behaviour of the 4 cropping alternatives (including not cropping, i.e full e2e) are based on 50 steps with cfg 4.0, when using lightx2v the results may be different + # Anyway, the cropping is generally a harmless optimization to reduce artifacts, and as long as it can cover the main body movements and interactions with objects, it should be fine. + parser.add_argument('--crop_margin', type=float, default=0.05, + help='Fractional padding added to bboxes in --crop_e2e_bbox / ' + '--crop_e2e_steady_bbox (default 0.05).') + parser.add_argument('--max_persons', type=int, default=2, + help='Maximum number of persons to track from ref / driving (default 2).') + parser.add_argument('--text', type=str, nargs='+', + default=['human', 'character'], + help='Text prompts passed to SAM3 for both driving and ref. Add extras ' + 'like "robot arm" "gripper" if the subject is a non-human character ' + '(e.g. a robotic arm in egocentric/animation data).') + parser.add_argument('--skip_existing', action='store_true', + help='In --input_root mode, skip subdirs whose rendered_mask_v2.mp4 already exists.') + parser.add_argument('--model_path', type=str, + default='pretrained_weights/nlf_l_multi_0.3.2.torchscript', + help='Path to NLF TorchScript model (only used when --e2e_mode is not set).') + parser.add_argument('--sam3_model', type=str, + default='pretrained_weights/sam3.pt', + help='Path to SAM3 model weights.') + args = parser.parse_args() + + crop_kind = ('bbox' if args.crop_e2e_bbox else + 'mask' if args.crop_e2e_mask else + 'steady_bbox' if args.crop_e2e_steady_bbox else None) + if crop_kind is not None and not args.e2e_mode: + parser.error("--crop_e2e_bbox / --crop_e2e_mask / --crop_e2e_steady_bbox require --e2e_mode") + + from ultralytics.models.sam import SAM3SemanticPredictor, SAM3VideoSemanticPredictor + + print("Initializing SAM3 video predictor...") + overrides = dict( + conf=0.25, task="segment", mode="predict", imgsz=640, + model=args.sam3_model, half=True, save=False, verbose=False, + ) + predictor = SAM3VideoSemanticPredictor(overrides=overrides, new_det_thresh=1.0) + + print("Initializing SAM3 image predictor...") + image_predictor = SAM3SemanticPredictor(overrides=overrides) + + if args.e2e_mode: + model_nlf = None + detector = None + else: + from DWPoseProcess.dwpose import DWposeDetector + print("Loading NLF model...") + model_nlf = torch.jit.load(args.model_path).cuda().eval() + print("Loading DWpose detector...") + detector = DWposeDetector(use_batch=False).to(0) + print("All models loaded.") + + if args.subdir is not None: + subdirs = [args.subdir] + else: + subdirs = sorted(d for d in glob.glob(os.path.join(args.input_root, '*')) + if os.path.isdir(d)) + if not subdirs: + print(f"No subdirs found under {args.input_root}") + sys.exit(0) + + n_ok, n_skip, n_err = 0, 0, 0 + for i, subdir in enumerate(subdirs): + if args.skip_existing and os.path.exists(os.path.join(subdir, 'rendered_mask_v2.mp4')): + print(f"[{i+1}/{len(subdirs)}] skip (already done): {subdir}") + n_skip += 1 + continue + + print(f"\n{'='*60}") + print(f"[{i+1}/{len(subdirs)}] {subdir} (video_name={args.video_name}, e2e_mode={args.e2e_mode}, crop_kind={crop_kind})") + print(f"{'='*60}") + t0 = time.time() + try: + process_one(subdir, args.video_name, args.e2e_mode, crop_kind, + args.max_persons, args.text, model_nlf, detector, + predictor, image_predictor, crop_margin=args.crop_margin) + n_ok += 1 + print(f" -> ok ({time.time() - t0:.1f}s)") + except Exception as e: + n_err += 1 + print(f" -> FAILED: {e}") + traceback.print_exc() + + print(f"\nDone. ok={n_ok} skipped={n_skip} failed={n_err} total={len(subdirs)}") diff --git a/SCAIL-Pose/NLFPoseExtract/process_replacement.py b/SCAIL-Pose/NLFPoseExtract/process_replacement.py new file mode 100644 index 0000000000000000000000000000000000000000..402dc1f5bebd77fd0606186046586d3605dd9a6e --- /dev/null +++ b/SCAIL-Pose/NLFPoseExtract/process_replacement.py @@ -0,0 +1,235 @@ +import os +import sys + +current_dir = os.path.dirname(os.path.abspath(__file__)) +project_root = os.path.dirname(current_dir) +if project_root not in sys.path: + sys.path.insert(0, project_root) + +import argparse +import glob +import shutil +import time +import traceback + +import cv2 +import numpy as np + +from NLFPoseExtract.v2_helper import ( + find_ref_image, + save_colored_mask_image, + save_real_pixel_mask_image, + write_colored_mask_video, +) + + +def _select_closest_to_ref(drv_masks, drv_colors, ref_masks): + """matchnearest: among multiple driving tracks, pick the one whose first-frame + mask has the highest IoU with the ref mask, after resizing ref to driving + resolution. Returns ([selected_mask], [selected_color]). + """ + if len(drv_masks) <= 1: + return drv_masks, drv_colors + + H_drv, W_drv = drv_masks[0].shape[1:] + ref_u8 = ref_masks[0][0].astype(np.uint8) * 255 + ref_resized = cv2.resize(ref_u8, (W_drv, H_drv), interpolation=cv2.INTER_NEAREST) > 127 + + best_iou, best_idx = -1.0, 0 + for i, mask in enumerate(drv_masks): + drv_first = mask[0] + inter = int(np.logical_and(ref_resized, drv_first).sum()) + union = int(np.logical_or(ref_resized, drv_first).sum()) + iou = inter / max(union, 1) + print(f" matchnearest IoU track {i} (color={drv_colors[i]}): {iou:.4f}") + if iou > best_iou: + best_iou, best_idx = iou, i + + print(f" matchnearest selected track {best_idx} (IoU={best_iou:.4f})") + return [drv_masks[best_idx]], [drv_colors[best_idx]] + + +def _union_masks(masks, colors): + """Combine N masks into one via logical OR; reuse the first track's color. + Used for egocentric mode where left/right arms are detected as separate SAM3 + instances but should be treated as a single actor.""" + if len(masks) <= 1: + return masks, colors + combined = np.logical_or.reduce(masks) + print(f" egocentric: unioned {len(masks)} masks into 1 (color={colors[0]})") + return [combined], [colors[0]] + + +def process_one(subdir, video_name, test_mode, matchnearest, egocentric, + predictor, image_predictor, text): + from TrackSam3.track import get_mask_from_image, get_mask_from_video + + mp4_path = os.path.join(subdir, video_name) + if not os.path.exists(mp4_path): + raise FileNotFoundError(f"No {video_name} found in {subdir}") + + out_path_rendered = os.path.join(subdir, 'rendered_v2.mp4') + out_path_mask = os.path.join(subdir, 'replace_mask.mp4') + ref_image_out_path = os.path.join(subdir, 'ref_image.png') + ref_mask_path = os.path.join(subdir, 'ref_mask.png') + + # 1) Read fps + first frame via cv2 — decord VideoReader corrupts CUDA fds before/between + # SAM3 calls regardless of ordering; cv2 is CUDA-agnostic and safe at any point. + cap = cv2.VideoCapture(mp4_path) + fps = cap.get(cv2.CAP_PROP_FPS) + fps_int = max(1, int(round(fps))) + ret, first_frame_bgr = cap.read() + cap.release() + if not ret: + raise RuntimeError(f"Could not read first frame from {mp4_path}") + first_frame_rgb = first_frame_bgr[:, :, ::-1] + + # 2) Driving → full-video masks. matchnearest allows 2 tracks then picks via IoU; + # egocentric allows 2 tracks then unions them. + max_drv = 2 if (matchnearest or egocentric) else 1 + print(f"Getting driving masks from {mp4_path} (max_targets={max_drv}, text={text})...") + drv_masks, drv_colors = get_mask_from_video( + mp4_path, predictor, max_targets=max_drv, sort_by='x', fixed_colors=None, text=text, + ) + if len(drv_masks) == 0: + raise RuntimeError(f"No valid persons detected in driving {mp4_path}") + print(f"Driving detected: {len(drv_masks)} person(s); colors={drv_colors}") + + if egocentric: + drv_masks, drv_colors = _union_masks(drv_masks, drv_colors) + + # 3) Resolve ref_image_path: test_mode auto-generates from first driving frame + if test_mode: + save_real_pixel_mask_image([drv_masks[0][0:1]], first_frame_rgb, ref_image_out_path) + print(f"[test_mode] Ref image saved (real pixels, black bg): {ref_image_out_path}") + ref_image_path = ref_image_out_path + else: + ref_image_path = find_ref_image(subdir) + + # 4) Get ref masks from ref_image (same path for both modes) + max_ref = 2 if egocentric else 1 + print(f"Getting ref masks from {ref_image_path} (max_targets={max_ref})...") + ref_masks, ref_colors = get_mask_from_image( + ref_image_path, image_predictor, max_targets=max_ref, + sort_by='x', fixed_colors=None, text=text, + ) + if len(ref_masks) == 0: + raise RuntimeError(f"No qualifying person found in ref image {ref_image_path}") + + if egocentric: + ref_masks, ref_colors = _union_masks(ref_masks, ref_colors) + + # ref_mask.png: solid-color mask on black bg + save_colored_mask_image(ref_masks, ref_colors, ref_mask_path, bg_color=(0, 0, 0)) + print(f" Ref mask saved: {ref_mask_path}") + + # 4.5) matchnearest: pick the driving track closest to ref by IoU + if matchnearest: + drv_masks, drv_colors = _select_closest_to_ref(drv_masks, drv_colors, ref_masks) + + # 4) rendered_v2.mp4 is always a copy of driving + shutil.copyfile(mp4_path, out_path_rendered) + print(f" Copied driving → {out_path_rendered}") + + # 5) replace_mask.mp4 — white background + print("Writing replace_mask.mp4 (white bg)...") + write_colored_mask_video(drv_masks, drv_colors, out_path_mask, fps_int, + bg_color=(255, 255, 255)) + + print("Done!") + print(f" Rendered: {out_path_rendered}") + print(f" Replace mask: {out_path_mask}") + print(f" Ref mask: {ref_mask_path}") + + +if __name__ == '__main__': + parser = argparse.ArgumentParser( + description='SCAIL replacement pipeline: SAM3 mask extraction for character replacement. ' + 'Outputs rendered_v2.mp4 (driving copy), replace_mask.mp4 (white bg), ' + 'ref_mask.png (black bg). Pass exactly one of --subdir or --input_root.' + ) + src = parser.add_mutually_exclusive_group(required=True) + src.add_argument('--subdir', type=str, default=None, + help='Single-example mode: path to one subdir. ' + 'Mutually exclusive with --input_root.') + src.add_argument('--input_root', type=str, default=None, + help='Batch mode: directory whose immediate subdirs are each an example. ' + 'Mutually exclusive with --subdir.') + parser.add_argument('--video_name', type=str, default='driving.mp4', + choices=['driving.mp4', 'GT.mp4'], + help='Filename of the driving video inside each subdir.') + parser.add_argument('--test_mode', action='store_true', + help='Use driving first frame as ref: saves ref_image.png with real ' + 'pixels inside mask area (black outside). No ref_image file needed.') + parser.add_argument('--matchnearest', action='store_true', + help='Driving may contain 2 persons; ref has 1. Picks the driving ' + 'track whose first-frame mask has highest IoU with the ref mask ' + '(after resizing ref to driving resolution). Other tracks are dropped.') + parser.add_argument('--egocentric', action='store_true', + help='ONLY for egocentric/first-person data where the actor appears as ' + 'multiple disconnected parts (e.g. left + right arms or grippers). ' + 'Sets max_targets=2 for both driving and ref, then unions the ' + 'resulting masks into one (same color), treating both arms as a ' + 'single actor. Do NOT use on normal third-person data. ' + 'Mutually exclusive with --matchnearest.') + parser.add_argument('--text', type=str, nargs='+', + default=['human', 'character'], + help='Text prompts passed to SAM3 for both driving and ref. Add extras ' + 'like "bear" if the subject is a non-human character.') + parser.add_argument('--skip_existing', action='store_true', + help='In --input_root mode, skip subdirs whose replace_mask.mp4 already exists.') + parser.add_argument('--sam3_model', type=str, + default='pretrained_weights/sam3.pt', + help='Path to SAM3 model weights.') + args = parser.parse_args() + + if args.matchnearest and args.egocentric: + parser.error("--matchnearest and --egocentric are mutually exclusive: " + "the first picks one track out of many, the second unions multiple " + "tracks into one.") + + from ultralytics.models.sam import SAM3SemanticPredictor, SAM3VideoSemanticPredictor + + print("Initializing SAM3 video predictor...") + overrides = dict( + conf=0.25, task="segment", mode="predict", imgsz=640, + model=args.sam3_model, half=True, save=False, verbose=False, + ) + predictor = SAM3VideoSemanticPredictor(overrides=overrides, new_det_thresh=1.0) + + print("Initializing SAM3 image predictor...") + image_predictor = SAM3SemanticPredictor(overrides=overrides) + + print("All models loaded.") + + if args.subdir is not None: + subdirs = [args.subdir] + else: + subdirs = sorted(d for d in glob.glob(os.path.join(args.input_root, '*')) + if os.path.isdir(d)) + if not subdirs: + print(f"No subdirs found under {args.input_root}") + sys.exit(0) + + n_ok, n_skip, n_err = 0, 0, 0 + for i, subdir in enumerate(subdirs): + if args.skip_existing and os.path.exists(os.path.join(subdir, 'replace_mask.mp4')): + print(f"[{i+1}/{len(subdirs)}] skip (already done): {subdir}") + n_skip += 1 + continue + + print(f"\n{'='*60}") + print(f"[{i+1}/{len(subdirs)}] {subdir} (video_name={args.video_name}, test_mode={args.test_mode}, matchnearest={args.matchnearest}, egocentric={args.egocentric})") + print(f"{'='*60}") + t0 = time.time() + try: + process_one(subdir, args.video_name, args.test_mode, args.matchnearest, + args.egocentric, predictor, image_predictor, args.text) + n_ok += 1 + print(f" -> ok ({time.time() - t0:.1f}s)") + except Exception as e: + n_err += 1 + print(f" -> FAILED: {e}") + traceback.print_exc() + + print(f"\nDone. ok={n_ok} skipped={n_skip} failed={n_err} total={len(subdirs)}") diff --git a/SCAIL-Pose/NLFPoseExtract/reshape_utils_3d.py b/SCAIL-Pose/NLFPoseExtract/reshape_utils_3d.py new file mode 100644 index 0000000000000000000000000000000000000000..da44f4bc7f7dbff186af1b8e37d18d5e598c5cf6 --- /dev/null +++ b/SCAIL-Pose/NLFPoseExtract/reshape_utils_3d.py @@ -0,0 +1,304 @@ +import numpy as np +import random +from NLFPoseExtract.nlf_draw import intrinsic_matrix_from_field_of_view, process_data_to_COCO_format, p3d_to_p2d +import torch + + + +# reshapePool只负责形变,骨骼偏移、丢弃等得从draw层来做 +class reshapePool3d: + def __init__(self, reshape_type, height, width): # 对每个视频只初始化一次 + self.reshape_type = reshape_type + self.height = height + self.width = width + self.shoulder_alpha = 0 + self.upper_arm_alpha = 0 + self.forearm_alpha = 0 + self.body_alpha = 0 + self.thigh_alpha = 0 + self.calf_alpha = 0 + self.face_alpha = random.choices([-0.4, -0.2, 0, 0.2, 0.4], weights=[0.2, 0.15, 0.3, 0.15, 0.2], k=1)[0] + + + self.body_reshape_methods = [ + self.reshape_body, + self.reshape_arm, + self.reshape_leg, + self.reshape_shoulder, + ] + + self.face_reshape_methods = [ + self.reshape_face, + ] + + self.body_offset_selected_methods = [] + + options = ["normal_human", "dwarf", "slender", "elf", "random_long_arm_long_leg", "king-kong"] + if self.reshape_type == "low": + weights = [0.8, 0, 0, 0.1, 0.1, 0] + self.body_offset_selected_methods = [] + if self.reshape_type == "normal": + weights = [0.4, 0.1, 0.1, 0.1, 0.2, 0.1] + elif self.reshape_type == "high": + weights = [0.3, 0.2, 0.1, 0.1, 0.2, 0.1] + elif self.reshape_type == "dongman": + weights = [0.1, 0.2, 0.2, 0.2, 0.2, 0.1] + self.face_alpha = random.choices([-0.4, -0.2, 0.4], weights=[0.4, 0.2, 0.4], k=1)[0] + choice = random.choices(options, weights=weights, k=1)[0] + self.aug_init(choice) + + + def pose_reshape_2d_for_face(self, alpha, candidate, face, subset, body_anchor_point, body_affected_points): + anchor_x, anchor_y = candidate[body_anchor_point] + if subset[body_anchor_point] == -1: + return + + for face_point_idx in range(len(face)): + face_point_x, face_point_y = face[face_point_idx] + if face_point_x == -1 or face_point_y == -1: + continue + vector_x = face_point_x - anchor_x + vector_y = face_point_y - anchor_y + offset_x = vector_x * alpha + offset_y = vector_y * alpha + face[face_point_idx] = [face_point_x + offset_x, face_point_y + offset_y] + + for body_affected_point_idx in body_affected_points: + body_point_x, body_point_y = candidate[body_affected_point_idx] + if subset[body_affected_point_idx] == -1 or body_point_x == -1 or body_point_y == -1: + continue + vector_x = body_point_x - anchor_x + vector_y = body_point_y - anchor_y + offset_x = vector_x * alpha + offset_y = vector_y * alpha + candidate[body_affected_point_idx] = [body_point_x + offset_x, body_point_y + offset_y] + + def pose_reshape_3d(self, alpha, smpl_joints, candidate, subset, left_hand, right_hand, face, + anchor_point, end_point, + affected_body_points): + """ + # joints3d: n, 24, 3 + alpha: 变化比例 + anchor_point: 起始点 + anchor_part: 起始点所属的身体部分,0代表body candidate, 1代表face, 2代表hand + end_point: 中止点 + """ + if torch.sum(smpl_joints[anchor_point]) == 0: + right_hand[:] = -1 + left_hand[:] = -1 + face[:] = -1 + subset[:] = -1 + return + + + anchor_x, anchor_y, anchor_z = smpl_joints[anchor_point] + end_x, end_y, end_z = smpl_joints[end_point] + + vector_x = end_x - anchor_x + vector_y = end_y - anchor_y + vector_z = end_z - anchor_z + offset_x = (vector_x * alpha).item() + offset_y = (vector_y * alpha).item() + offset_z = (vector_z * alpha).item() + + map_to_2d = {} + map_to_2d[4] = 11 + map_to_2d[7] = 12 + map_to_2d[10] = 13 + map_to_2d[5] = 8 + map_to_2d[8] = 9 + map_to_2d[11] = 10 + map_to_2d[20] = 6 + map_to_2d[22] = 7 + map_to_2d[21] = 3 + map_to_2d[23] = 4 + map_to_2d[18] = 5 + map_to_2d[19] = 2 + + + for affected_body_point in affected_body_points: + if torch.sum(smpl_joints[affected_body_point]) == 0: + continue + new_smpl_joint = smpl_joints[affected_body_point] + torch.tensor([offset_x, offset_y, offset_z]).to(smpl_joints.device) + new_smpl_joint_2d_offset = p3d_to_p2d(new_smpl_joint.reshape(1,1,3).cpu().numpy(), self.height, self.width)[0][0] - p3d_to_p2d(smpl_joints[affected_body_point].reshape(1,1,3).cpu().numpy(), self.height, self.width)[0][0] + new_smpl_joint_2d_offset = np.array([new_smpl_joint_2d_offset[0] / self.width, new_smpl_joint_2d_offset[1] / self.height]) + smpl_joints[affected_body_point] = new_smpl_joint + if affected_body_point in map_to_2d.keys(): + affected_candidate_point_idx = map_to_2d[affected_body_point] + if subset[affected_candidate_point_idx] != -1 and candidate[affected_candidate_point_idx][0] != -1 and candidate[affected_candidate_point_idx][1] != -1: + candidate[affected_candidate_point_idx] = candidate[affected_candidate_point_idx] + new_smpl_joint_2d_offset # 2d的也移动这么多 + if affected_candidate_point_idx == 4: # dwpose 右手 (反的 + left_hand[:] = left_hand + new_smpl_joint_2d_offset + if affected_candidate_point_idx == 7: # dwpose 左手 (反的 + right_hand[:] = right_hand + new_smpl_joint_2d_offset + + + + def aug_init(self, body_type): + print(f"augmentation: using body_type: {body_type}") + self.shoulder_alpha = 0 + self.upper_arm_alpha = 0 + self.forearm_alpha = 0 + self.body_alpha = 0 + self.thigh_alpha = 0 + self.calf_alpha = 0 + if body_type == "normal_human": + self.body_reshape_selected_methods = [] + elif body_type == "dwarf": # body不动 + self.upper_arm_alpha = random.uniform(-0.3, -0.2) + self.forearm_alpha = self.upper_arm_alpha + self.shoulder_alpha = -0.2 + self.thigh_alpha = random.uniform(-0.3, -0.2) + self.calf_alpha = self.thigh_alpha + self.body_reshape_selected_methods = [self.reshape_shoulder, self.reshape_arm, self.reshape_leg] + self.face_alpha = 0.2 + elif body_type == "slender": + self.upper_arm_alpha = 0.3 + self.forearm_alpha = 0.2 + self.shoulder_alpha = 0.2 + self.thigh_alpha = 0.1 + self.calf_alpha = 0.1 + self.body_reshape_selected_methods = [self.reshape_shoulder, self.reshape_arm, self.reshape_leg] + self.face_alpha = -0.2 + elif body_type == "elf": + self.body_alpha = random.uniform(-0.2, 0.2) + self.shoulder_alpha = 0.1 + self.upper_arm_alpha = 0.1 + self.forearm_alpha = 0.1 + self.thigh_alpha = 0.25 + self.calf_alpha = 0.25 + self.body_reshape_selected_methods = [self.reshape_body, self.reshape_shoulder, self.reshape_arm, self.reshape_leg] + self.face_alpha = 0 + elif body_type == "king-kong": + self.body_alpha = 0.1 + self.thigh_alpha = -0.25 + self.calf_alpha = -0.25 + self.upper_arm_alpha = 0.2 + self.forearm_alpha = 0.2 + self.shoulder_alpha = 0.3 + self.body_reshape_selected_methods = [self.reshape_body, self.reshape_shoulder, self.reshape_arm, self.reshape_leg] + self.face_alpha = 0 + elif body_type == "random_long_arm_long_leg": + self.upper_arm_alpha = random.uniform(-0.2, 0.2) + self.forearm_alpha = random.uniform(-0.2, 0.2) + self.thigh_alpha = random.uniform(-0.2, 0.2) + self.calf_alpha = random.uniform(-0.2, 0.2) + self.body_alpha = random.uniform(-0.1, 0.1) + self.body_reshape_selected_methods = [self.reshape_body, self.reshape_arm, self.reshape_leg] + elif body_type == "test_case_1": + self.upper_arm_alpha = -0.4 + self.forearm_alpha = -0.4 + self.shoulder_alpha = -0.3 + self.thigh_alpha = 0.2 + self.calf_alpha = 0.2 + self.body_alpha = 0.1 + self.face_alpha = 0.4 + self.body_reshape_selected_methods = [self.reshape_body, self.reshape_shoulder, self.reshape_arm, self.reshape_leg] + elif body_type == "test_case_2": + self.upper_arm_alpha = 0.4 + self.forearm_alpha = 0.4 + self.shoulder_alpha = 0.3 + self.thigh_alpha = -0.2 + self.calf_alpha = -0.25 + self.body_alpha = -0.2 + self.face_alpha = -0.2 + self.body_reshape_selected_methods = [self.reshape_body, self.reshape_shoulder, self.reshape_arm, self.reshape_leg] + elif body_type == "test_case_3": + self.upper_arm_alpha = 0.2 + self.forearm_alpha = 0.2 + self.shoulder_alpha = 0.2 + self.thigh_alpha = 0.4 + self.calf_alpha = 0.4 + self.body_alpha = 0.3 + self.face_alpha = 0.15 + self.body_reshape_selected_methods = [self.reshape_body, self.reshape_shoulder, self.reshape_arm, self.reshape_leg] + elif body_type == "normal_human_test": + self.body_reshape_selected_methods = [] + self.face_alpha = 0 + + + + def apply_random_reshapes(self, smpl_joints_list, candidate, left_hand, right_hand, face, subset): + # Apply the two selected reshape methods + for method in self.body_reshape_selected_methods: + method(smpl_joints_list, candidate, subset, left_hand, right_hand, face) + + for method in self.body_offset_selected_methods: + method(smpl_joints_list, candidate, left_hand, right_hand, face) + + for method in self.face_reshape_methods: + method(candidate, face, subset) + + + + + def reshape_body(self, smpl_joints_list, candidate, subset, left_hand, right_hand, face): + self.pose_reshape_3d(self.body_alpha, smpl_joints_list, candidate, subset, left_hand, right_hand, face, + 12, 1, + [1, 4, 7, 10] + ) + self.pose_reshape_3d(self.body_alpha, smpl_joints_list, candidate, subset, left_hand, right_hand, face, + 12, 2, + [2, 5, 8, 11] + ) + + def reshape_arm(self, smpl_joints_list, candidate, subset, left_hand, right_hand, face): + self.pose_reshape_3d(self.upper_arm_alpha, smpl_joints_list, candidate, subset, left_hand, right_hand, face, + 16, 18, + [18, 20, 22] + ) + self.pose_reshape_3d(self.upper_arm_alpha, smpl_joints_list, candidate, subset, left_hand, right_hand, face, + 17, 19, + [19, 21, 23] + ) + self.pose_reshape_3d(self.forearm_alpha, smpl_joints_list, candidate, subset, left_hand, right_hand, face, + 18, 20, + [20, 22] + ) + self.pose_reshape_3d(self.forearm_alpha, smpl_joints_list, candidate, subset, left_hand, right_hand, face, + 19, 21, + [21, 23] + ) + + + def reshape_leg(self, smpl_joints_list, candidate, subset, left_hand, right_hand, face): + self.pose_reshape_3d(self.thigh_alpha, smpl_joints_list, candidate, subset, left_hand, right_hand, face, + 1, 4, + [4, 7, 10] + ) + self.pose_reshape_3d(self.thigh_alpha, smpl_joints_list, candidate, subset, left_hand, right_hand, face, + 2, 5, + [5, 8, 11] + ) + self.pose_reshape_3d(self.calf_alpha, smpl_joints_list, candidate, subset, left_hand, right_hand, face, + 4, 7, + [7, 10] + ) + self.pose_reshape_3d(self.calf_alpha, smpl_joints_list, candidate, subset, left_hand, right_hand, face, + 5, 8, + [8, 11] + ) + + + def reshape_shoulder(self, smpl_joints_list, candidate, subset, left_hand, right_hand, face): + self.pose_reshape_3d(self.shoulder_alpha, smpl_joints_list, candidate, subset, left_hand, right_hand, face, + 12, 16, + [16, 18, 20, 22] + ) + self.pose_reshape_3d(self.shoulder_alpha, smpl_joints_list, candidate, subset, left_hand, right_hand, face, + 12, 17, + [17, 19, 21, 23] + ) + + # def offset_3d_all(self, smpl_joints_list, candidate, left_hand, right_hand, face): + # smpl_joints_list = smpl_joints_list + torch.tensor([self.offset_3d_x, self.offset_3d_y, self.offset_3d_z]).to(smpl_joints_list.device) + # offset_2d = p3d_to_p2d(np.array([[[self.offset_3d_x, self.offset_3d_y, self.offset_3d_z]]]), self.height, self.width)[0][0] + # candidate = candidate + np.array([offset_2d[0], offset_2d[1]]) + # left_hand = left_hand + np.array([offset_2d[0], offset_2d[1]]) + # right_hand = right_hand + np.array([offset_2d[0], offset_2d[1]]) + # face = face + np.array([offset_2d[0], offset_2d[1]]) + + + def reshape_face(self, candidate, face, subset): + self.pose_reshape_2d_for_face(alpha=self.face_alpha, candidate=candidate, face=face, subset=subset, + body_anchor_point=0, body_affected_points=[14, 15, 16, 17]) diff --git a/SCAIL-Pose/NLFPoseExtract/v1_process_pose.py b/SCAIL-Pose/NLFPoseExtract/v1_process_pose.py new file mode 100644 index 0000000000000000000000000000000000000000..5478a8ab485c014977c90af2c8040ce499148d90 --- /dev/null +++ b/SCAIL-Pose/NLFPoseExtract/v1_process_pose.py @@ -0,0 +1,313 @@ +import os +import sys + +# 动态添加项目根目录到 sys.path,这样就不需要 export PYTHONPATH +current_dir = os.path.dirname(os.path.abspath(__file__)) +project_root = os.path.dirname(current_dir) # SCAIL_Pose 目录 +if project_root not in sys.path: + sys.path.insert(0, project_root) + +import cv2 +import torch +import pickle +import torchvision +import shutil +import glob +import random +from tqdm import tqdm +import decord +from decord import VideoReader, cpu, gpu +from torchvision.transforms import ToPILImage +from PIL import Image +import numpy as np +import argparse +from NLFPoseExtract.nlf_render import render_nlf_as_images, collect_smpl_poses, shift_dwpose_according_to_nlf, p3d_single_p2d +from NLFPoseExtract.nlf_draw import intrinsic_matrix_from_field_of_view, process_data_to_COCO_format, p3d_to_p2d +from DWPoseProcess.dwpose import DWposeDetector +from concurrent.futures import ProcessPoolExecutor, as_completed +import multiprocessing +import traceback +from NLFPoseExtract.extract_nlfpose_batch import process_video_nlf +from NLFPoseExtract.reshape_utils_3d import reshapePool3d +try: + import moviepy.editor as mpy +except: + import moviepy as mpy +from torchvision.transforms.functional import center_crop, resize +from torchvision.transforms import InterpolationMode +import torchvision.transforms as TT +import copy +from NLFPoseExtract.align3d import solve_new_camera_params_central, solve_new_camera_params_down + + +def recollect_nlf(data): + new_data = [] + for item in data: + new_item = item.copy() + if len(item['bboxes']) > 0: + new_item['bboxes'] = item['bboxes'][:1] + new_item['nlfpose'] = item['nlfpose'][:1] + new_data.append(new_item) + return new_data + +def recollect_dwposes(poses): + new_poses = [] + for pose in poses: + new_pose = pose.copy() + for i in range(1): + bodies = pose["bodies"] + faces = pose["faces"][i:i+1] + hands = pose["hands"][2*i:2*i+2] + candidate = bodies["candidate"][i:i+1] # candidate是所有点的坐标和置信度 + subset = bodies["subset"][i:i+1] # subset是认为的有效点 + new_pose = { + "bodies": { + "candidate": candidate, + "subset": subset + }, + "faces": faces, + "hands": hands + } + new_poses.append(new_pose) + return new_poses + + + +def resize_for_rectangle_crop(arr, image_size, reshape_mode='random'): + if arr.shape[3] / arr.shape[2] > image_size[1] / image_size[0]: + arr = resize(arr, size=[image_size[0], int(arr.shape[3] * image_size[0] / arr.shape[2])], interpolation=InterpolationMode.BICUBIC) + else: + arr = resize(arr, size=[int(arr.shape[2] * image_size[1] / arr.shape[3]), image_size[1]], interpolation=InterpolationMode.BICUBIC) + + h, w = arr.shape[2], arr.shape[3] + + delta_h = h - image_size[0] + delta_w = w - image_size[1] + + if reshape_mode == 'random' or reshape_mode == 'none': + top = np.random.randint(0, delta_h + 1) + left = np.random.randint(0, delta_w + 1) + elif reshape_mode == 'center': + top, left = delta_h // 2, delta_w // 2 + else: + raise NotImplementedError + arr = TT.functional.crop( + arr, top=top, left=left, height=image_size[0], width=image_size[1] + ) + return arr + +def scale_faces(poses, pose_2d_ref): + # 输入:两个list of dict,poses[0]['faces'].shape: 1, 68, 2 , poses_ref[0]['faces'].shape: 1, 68, 2 + # 根据脸部的中心点,对poses中的脸部关键点进行缩放 + # 也即:计算ref里面脸部中心点(idx: 30)到其他脸部关键点的中心距离, 计算poses里面脸部中心点到其他脸部关键点的中心距离,得到scale_n + # 对scale_n 取一下0.8-1.5的上下界,然后应用在poses上 + # 注意:需要inplace改变poses + + ref = pose_2d_ref[0] + pose_0 = poses[0] + + + face_0 = pose_0['faces'] # shape: (1, 68, 2) + face_ref = ref['faces'] + + # 提取 numpy 数组 + face_0 = np.array(face_0[0]) # (68, 2) + face_ref = np.array(face_ref[0]) + + # 中心点(鼻尖或面部中心) + center_idx = 30 + center_0 = face_0[center_idx] + center_ref = face_ref[center_idx] + + # 计算到中心点的距离 + dist = np.linalg.norm(face_0 - center_0, axis=1) + dist_ref = np.linalg.norm(face_ref - center_ref, axis=1) + + # 避免中心点自身的 0 距离影响 + dist = np.delete(dist, center_idx) + dist_ref = np.delete(dist_ref, center_idx) + + mean_dist = np.mean(dist) + mean_dist_ref = np.mean(dist_ref) + + if mean_dist < 1e-6: + scale_n = 1.0 + else: + scale_n = mean_dist_ref / mean_dist + + # 限制在 [0.8, 1.5] + scale_n = np.clip(scale_n, 0.8, 1.5) + + for i, pose in enumerate(poses): + face = pose['faces'] + # 提取 numpy 数组 + face = np.array(face[0]) # (68, 2) + center = face[center_idx] + scaled_face = (face - center) * scale_n + center + poses[i]['faces'][0] = scaled_face + + body = pose['bodies'] + candidate = body['candidate'] + candidate_np = np.array(candidate[0]) # (14, 2) + body_center = candidate_np[0] + scaled_candidate = (candidate_np - body_center) * scale_n + body_center + poses[i]['bodies']['candidate'][0] = scaled_candidate + + # inplace 修改 + pose['faces'][0] = scaled_face + + return scale_n + + +if __name__ == '__main__': + parser = argparse.ArgumentParser(description='Process video with NLF pose estimation') + parser.add_argument('--subdir', type=str, default="../examples/001", help='Path to the subdirectory to process') + parser.add_argument('--model_path', type=str, default='pretrained_weights/nlf_l_multi_0.3.2.torchscript', + help='Path to NLF model') + parser.add_argument('--use_align', action='store_true', help='Whether to use 2D keypoints from reference image for alignment') + parser.add_argument('--resolution', type=int, nargs=2, default=[512, 896], + metavar=('HEIGHT', 'WIDTH'), + help='Target resolution as [height, width], currently only [512, 896] are supported') + args = parser.parse_args() + + subdir = args.subdir + model_nlf = torch.jit.load(args.model_path).cuda().eval() + decord.bridge.set_bridge("torch") + + # 设置路径 + mp4_path = os.path.join(subdir, 'driving.mp4') + if not os.path.exists(mp4_path): + raise FileNotFoundError(f"No video file found in {subdir}") + + if args.use_align: + out_path_aligned = os.path.join(subdir, 'rendered_aligned.mp4') + else: + out_path_aligned = os.path.join(subdir, 'rendered.mp4') + + ref_image_path = os.path.join(subdir, 'ref_image.jpg') + if not os.path.exists(ref_image_path): + ref_image_path = os.path.join(subdir, 'ref_image.png') + if not os.path.exists(ref_image_path): + ref_image_path = os.path.join(subdir, 'ref.jpg') + if not os.path.exists(ref_image_path): + raise FileNotFoundError(f"No reference image found in {subdir}") + + print(f"Processing: {subdir}") + print(f"Video: {mp4_path}") + print(f"Reference: {ref_image_path}") + print(f"Resolution: {args.resolution}") + + # 读取视频 + vr = VideoReader(mp4_path) + vr_frames = vr.get_batch(list(range(len(vr)))) # T H W C + sampling_image_size = args.resolution + if vr_frames.shape[1] < vr_frames.shape[2]: + target_H, target_W = sampling_image_size + else: + target_W, target_H = sampling_image_size + vr_frames = resize_for_rectangle_crop(vr_frames.permute(0, 3, 1, 2), [target_H, target_W], reshape_mode='center').permute(0, 2, 3, 1) # T H W C ->T C H W -> T H W C + + # 读取参考图片 + img_ref = cv2.imread(ref_image_path) + img_ref = cv2.cvtColor(img_ref, cv2.COLOR_BGR2RGB) + vr_frames_ref = torch.from_numpy(img_ref).unsqueeze(0) + vr_frames_ref = resize_for_rectangle_crop(vr_frames_ref.permute(0, 3, 1, 2), [target_H, target_W], reshape_mode='center').permute(0, 2, 3, 1) # 1 H W C ->1 C H W -> 1 H W C + + # 初始化检测器 + detector = DWposeDetector(use_batch=False).to(0) + + # 处理Driving视频 + print("Processing driving video...") + detector_return_list = [] + pil_frames = [] + for i in tqdm(range(len(vr_frames)), desc="Detecting poses in video"): + pil_frame = Image.fromarray(vr_frames[i].numpy()) + pil_frames.append(pil_frame) + detector_result = detector(pil_frame) + detector_return_list.append(detector_result) + + W, H = pil_frames[0].size + poses, scores, det_results = zip(*detector_return_list) + + print("Running NLF on driving video...") + nlf_results = process_video_nlf(model_nlf, vr_frames, det_results) + + # 处理ref图片 + print("Processing reference image...") + detector_return_list_ref = [] + pil_frames_ref = [] + for i in range(len(vr_frames_ref)): + pil_frame = Image.fromarray(vr_frames_ref[i].numpy()) + pil_frames_ref.append(pil_frame) + detector_result = detector(pil_frame) + detector_return_list_ref.append(detector_result) + + poses_ref, scores_ref, det_results_ref = zip(*detector_return_list_ref) + + print("Running NLF on reference image...") + nlf_results_ref = process_video_nlf(model_nlf, vr_frames_ref, det_results_ref) + + # 进行对齐和渲染 + print("Aligning and rendering...") + ori_camera_pose = intrinsic_matrix_from_field_of_view([target_H, target_W]) + ori_focal = ori_camera_pose[0, 0] + pose_3d_first_driving_frame = nlf_results[0]['nlfpose'][0][0].cpu().numpy() # 3D点 frame-idx bbox-idx detect-idx + pose_3d_coco_first_driving_frame = process_data_to_COCO_format(pose_3d_first_driving_frame) + + if args.use_align: + poses_2d_ref = poses_ref[0]['bodies']['candidate'][0][:14] + poses_2d_ref[:, 0] = poses_2d_ref[:, 0] * target_W + poses_2d_ref[:, 1] = poses_2d_ref[:, 1] * target_H + + poses_2d_subset = poses_ref[0]['bodies']['subset'][0][:14] + pose_3d_coco_first_driving_frame = pose_3d_coco_first_driving_frame[:14] + + valid_indices = [] + valid_upper_indices = [] + valid_lower_indices = [] + upper_body_indices = [0, 2, 3, 5, 6] + lower_body_indices = [9, 10, 12, 13] + excluded_indices = [3, 4, 6, 7] # 去除手 + for i in range(len(poses_2d_subset)): + if poses_2d_subset[i] != -1.0 and np.sum(pose_3d_coco_first_driving_frame[i]) != 0: + if i in upper_body_indices: + valid_upper_indices.append(i) + if i in lower_body_indices: + valid_lower_indices.append(i) + + if len(valid_lower_indices) >= 4: + print("Align feet") + valid_indices = [1] + valid_lower_indices + else: + print("Align body") + valid_indices = [1] + valid_upper_indices + + pose_2d_ref = poses_2d_ref[valid_indices] + pose_3d_coco_first_driving_frame = pose_3d_coco_first_driving_frame[valid_indices] + + if len(valid_lower_indices) >= 4: + new_camera_intrinsics, scale_m, scale_s = solve_new_camera_params_down(pose_3d_coco_first_driving_frame, ori_focal, [target_H, target_W], pose_2d_ref) + else: + new_camera_intrinsics, scale_m, scale_s = solve_new_camera_params_central(pose_3d_coco_first_driving_frame, ori_focal, [target_H, target_W], pose_2d_ref) + + # m 代表缩放了多少 + scale_face = scale_faces(list(poses), list(poses_ref)) # poses[0]['faces'].shape: 1, 68, 2 , poses_ref[0]['faces'].shape: 1, 68, 2 + + print(f"Scale - m: {scale_m}, face: {scale_face}") + + nlf_results = recollect_nlf(nlf_results) + poses = recollect_dwposes(list(poses)) + shift_dwpose_according_to_nlf(collect_smpl_poses(nlf_results), poses, ori_camera_pose, new_camera_intrinsics, target_H, target_W, scale_x=scale_m, scale_y=scale_m*scale_s) + + print("Rendering final video...") + frames_np = render_nlf_as_images(nlf_results, poses, reshape_pool=None, intrinsic_matrix=new_camera_intrinsics) + + else: + nlf_results = recollect_nlf(nlf_results) + print("Rendering final video...") + frames_np = render_nlf_as_images(nlf_results, poses, reshape_pool=None, intrinsic_matrix=ori_camera_pose) + + mpy.ImageSequenceClip(frames_np, fps=16).write_videofile(out_path_aligned) + print(f"Done! Output saved to: {out_path_aligned}") + + diff --git a/SCAIL-Pose/NLFPoseExtract/v1_process_pose_multi.py b/SCAIL-Pose/NLFPoseExtract/v1_process_pose_multi.py new file mode 100644 index 0000000000000000000000000000000000000000..b4177fbf4d6924af1ea98429079d9114da5332d0 --- /dev/null +++ b/SCAIL-Pose/NLFPoseExtract/v1_process_pose_multi.py @@ -0,0 +1,243 @@ +import argparse +import os +import os.path as osp +import shutil +import numpy as np +import torch +import gc +import sys +import cv2 +from tqdm import tqdm +from PIL import Image +import decord +from decord import VideoReader, cpu +try: + import moviepy.editor as mpy +except ImportError: + import moviepy as mpy +import copy +import glob + +# Add project root to sys.path +current_dir = os.path.dirname(os.path.abspath(__file__)) +project_root = os.path.dirname(current_dir) +if project_root not in sys.path: + sys.path.insert(0, project_root) + +sys.path.append(os.path.join(project_root, "sam2")) +from sam2.build_sam import build_sam2_video_predictor + +from DWPoseProcess.dwpose import DWposeDetector +from NLFPoseExtract.extract_nlfpose_batch import process_video_multi_nlf +from NLFPoseExtract.nlf_render import render_multi_nlf_as_images + +def get_largest_bbox_indices(bboxes, num_bboxes=2): + # 计算每个bbox的面积 + def calculate_area(bbox): + x1, y1, x2, y2 = bbox + return (x2 - x1) * (y2 - y1) + + # 计算每个bbox的面积,并保留原索引 + bboxes_with_area = [(i, calculate_area(bbox)) for i, bbox in enumerate(bboxes)] + + # 根据面积从大到小排序 + bboxes_with_area.sort(key=lambda x: x[1], reverse=True) + + # 取出面积最大的 num_bboxes 个索引 + largest_indices = [idx for idx, _ in bboxes_with_area[:num_bboxes]] + + return largest_indices + +def change_poses_to_limit_num(poses, bboxes, num_bboxes=2): + bboxes = list(bboxes) # ✅ 转换为可变列表 + for idx, (pose, bbox) in enumerate(zip(poses, bboxes)): + if len(bbox) == 0: + continue + largest_indices = get_largest_bbox_indices(bbox, num_bboxes) + + # 过滤 subset、hands、faces + pose['bodies']['subset'] = pose['bodies']['subset'][largest_indices] + + new_hands = [] + for i in largest_indices: + if 2*i+1 < len(pose['hands']): + new_hands.append(pose['hands'][2*i]) + new_hands.append(pose['hands'][2*i+1]) + pose['hands'] = new_hands + + pose['faces'] = [pose['faces'][i] for i in largest_indices if i < len(pose['faces'])] + + bboxes[idx] = [bbox[i] for i in largest_indices] + + return poses, bboxes + +def get_samurai_crop_video(video_input_path, video_output_root, bboxes_0, final_keypoints_list, predictor=None, use_green_background=True): + decord.bridge.set_bridge("torch") + # 用 decord 读取视频帧 + if video_input_path.endswith(".mp4"): + vr = VideoReader(video_input_path) + loaded_frames = vr.get_batch(list(range(len(vr)))).numpy() + height, width = loaded_frames[0].shape[:2] + + # 每个人一个输出视频 + num_persons = len(final_keypoints_list) + print(f"Detected {num_persons} persons, will save {num_persons} videos.") + + prompts = {fid: ((x1, y1, x2, y2), 0) for fid, (x1, y1, x2, y2) in enumerate(bboxes_0)} + with torch.inference_mode(), torch.autocast("cuda", dtype=torch.float16): + for person_idx in range(num_persons): + print(f"Processing person {person_idx + 1}/{num_persons}...") + + state = predictor.init_state(video_input_path, offload_video_to_cpu=True) + bbox, track_label = prompts[person_idx] + bbox = (bbox[0] * width, bbox[1] * height, bbox[2] * width, bbox[3] * height) + points = copy.deepcopy(final_keypoints_list[person_idx]) + points[:, 0] *= width + points[:, 1] *= height + + _, _, masks = predictor.add_new_points_or_box(state, box=bbox, points=points, labels=np.ones(points.shape[0]), frame_idx=0, obj_id=0) + + output_frames = [] + output_mask_frames = [] + repeat_flag = False + + for frame_idx, object_ids, masks in predictor.propagate_in_video(state): + img = loaded_frames[frame_idx].copy() + for obj_id, mask in zip(object_ids, masks): + mask = mask[0].cpu().numpy() > 0.0 # 更新 mask + mask_log = np.zeros_like(img) + if use_green_background: + mask_img = np.full_like(img, (30, 60, 30)) + else: + mask_img = np.zeros_like(img) + mask_img[mask] = img[mask] + mask_log[mask] = 255 + output_frames.append(mask_img) # mask_img: array of [h, w, 3] + output_mask_frames.append(mask_log) # mask: array of [h, w] + + del state + gc.collect() + torch.cuda.empty_cache() + + # 用 moviepy 保存视频 + output_name = os.path.join(video_output_root, f"{person_idx+1}.mp4") + clip = mpy.ImageSequenceClip(output_frames, fps=16) + clip.write_videofile(output_name, codec="libx264", audio=False) + print(f"Saved {output_name}") + + # del predictor # Do not delete predictor here as it might be reused or managed outside + gc.collect() + torch.clear_autocast_cache() + torch.cuda.empty_cache() + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument('--subdir', type=str, required=True, help='Path to the subdirectory containing GT.mp4') + parser.add_argument('--model_path', type=str, default='pretrained_weights/nlf_l_multi_0.3.2.torchscript', + help='Path to NLF model') + parser.add_argument('--resolution', type=int, nargs=2, default=[512, 512], help='Resolution [H, W]') + args = parser.parse_args() + + subdir = args.subdir + model_path = args.model_path + resolution = args.resolution + + video_input_path = osp.join(subdir, "driving.mp4") + if not osp.exists(video_input_path): + print(f"Error: {video_input_path} does not exist.") + + # 1. Extract DWpose and BBoxes + print("Extracting DWpose and BBoxes...") + detector = DWposeDetector(use_batch=False).to(0) + + vr = VideoReader(video_input_path) + vr_frames = vr.get_batch(list(range(len(vr)))).asnumpy() # T H W C + + # Resize if needed? process_pose.py does resize. + # But run_samurai_mp4.py seems to use original video for samurai. + # process_multinlf_after_samurai.py uses samurai output which is same resolution as input? + # Let's stick to original resolution for extraction to match run_samurai_mp4 logic which uses original video. + + detector_return_list = [] + pil_frames = [] + for i in tqdm(range(len(vr_frames)), desc="Detecting poses"): + pil_frame = Image.fromarray(vr_frames[i]) + pil_frames.append(pil_frame) + detector_result = detector(pil_frame) + detector_return_list.append(detector_result) + + poses, scores, det_results = zip(*detector_return_list) + # poses is tuple of dicts, det_results is tuple of lists of bboxes + + # Save meta if needed, or just use in memory. + # run_samurai_mp4.py saves to meta/keypoints.pt and meta/bboxes.pt + meta_dir = osp.join(subdir, "meta") + os.makedirs(meta_dir, exist_ok=True) + torch.save(poses, osp.join(meta_dir, "keypoints.pt")) + torch.save(det_results, osp.join(meta_dir, "bboxes.pt")) + + # 2. Run Samurai Segmentation + print("Running Samurai Segmentation...") + samurai_output_root = osp.join(subdir, "samurai") + if osp.exists(samurai_output_root): + shutil.rmtree(samurai_output_root) + os.makedirs(samurai_output_root, exist_ok=True) + + device = "cuda:0" + predictor = build_sam2_video_predictor("configs/sam2.1/sam2.1_hiera_l.yaml", "sam2/checkpoints/sam2.1_hiera_large.pt", device=device) + + # Prepare inputs for samurai + bboxes_0 = det_results[0] + indices = get_largest_bbox_indices(bboxes_0) + bboxes_0 = [bboxes_0[index] for index in indices] + + keypoints_0 = poses[0]['bodies']['candidate'] + subset_0 = poses[0]['bodies']['subset'] + chosen_keypoints = keypoints_0[indices] + + final_keypoints_list = [] + for i in range(len(chosen_keypoints)): + keypoints_for_person = chosen_keypoints[i] + subset_for_person = subset_0[i] + considered_points = [0, 1, 14, 15] + # Create a copy to avoid modifying original if needed, though subset_for_person is from tensor + subset_for_person_mod = subset_for_person.copy() + for k in range(len(subset_for_person_mod)): + if k not in considered_points: + subset_for_person_mod[k] = -1 + new_keypoints = keypoints_for_person[subset_for_person_mod != -1] + final_keypoints_list.append(new_keypoints) + + get_samurai_crop_video(video_input_path, samurai_output_root, bboxes_0, final_keypoints_list, predictor=predictor) + + del predictor + gc.collect() + torch.cuda.empty_cache() + + # 3. Render Multi NLF + print("Rendering Multi NLF...") + model_nlf = torch.jit.load(model_path).cuda().eval() + decord.bridge.set_bridge("torch") + + vr_frames_list = [] + for samurai_mp4_path in sorted(glob.glob(osp.join(samurai_output_root, '*.mp4'))): + vr_tmp = VideoReader(samurai_mp4_path) + vr_frames_tmp = vr_tmp.get_batch(list(range(len(vr_tmp)))) + vr_frames_list.append(vr_frames_tmp) + + # Filter poses for rendering + # change_poses_to_limit_num modifies poses in-place or returns new ones? + # It returns poses, bboxes. And it modifies the lists passed to it? + # poses is a tuple from zip, convert to list + poses_list = list(poses) + det_results_list = list(det_results) + + poses_list, det_results_list = change_poses_to_limit_num(poses_list, det_results_list) + + nlf_results = process_video_multi_nlf(model_nlf, vr_frames_list) + frames_ori_np = render_multi_nlf_as_images(nlf_results, poses_list, reshape_pool=None) + + out_path = osp.join(subdir, 'rendered.mp4') + mpy.ImageSequenceClip(frames_ori_np, fps=16).write_videofile(out_path) + print(f"Done! Output saved to: {out_path}") + diff --git a/SCAIL-Pose/NLFPoseExtract/v2_helper.py b/SCAIL-Pose/NLFPoseExtract/v2_helper.py new file mode 100644 index 0000000000000000000000000000000000000000..925b0e60f9ea23f449b31ec3635018da210f004b --- /dev/null +++ b/SCAIL-Pose/NLFPoseExtract/v2_helper.py @@ -0,0 +1,233 @@ +"""Helpers shared by process_animation_aio.py and process_replacement.py. + +Small, side-effect-free utilities for finding ref images, rendering SAM3 masks +(image + video), and writing the e2e "kept-region" driving video used by +--crop_e2e_bbox / --crop_e2e_mask in process_animation_aio. +""" +import os + +import cv2 +import numpy as np + +try: + import moviepy.editor as mpy +except Exception: + import moviepy as mpy + + +def find_ref_image(subdir): + for name in ('ref_image.jpg', 'ref_image.png', 'ref.jpg', 'ref.png'): + p = os.path.join(subdir, name) + if os.path.exists(p): + return p + raise FileNotFoundError(f"No reference image (ref_image.jpg/png or ref.jpg) found in {subdir}") + + +def imread_bgr(image_path): + """Robust BGR uint8 image reader. Tries cv2.imread first, falls back to PIL + when cv2 returns None — handles cases where libpng's global state gets + corrupted by SAM3/torch deps and rejects PNGs that PIL still decodes fine. + + Raises FileNotFoundError if both readers fail.""" + img = cv2.imread(str(image_path)) + if img is not None: + return img + try: + from PIL import Image + pil = Image.open(image_path).convert('RGB') + print(f" [warn] cv2.imread failed for {image_path}; used PIL fallback") + return np.array(pil)[:, :, ::-1].copy() # RGB -> BGR + except Exception as e: + raise FileNotFoundError(f"Cannot read image: {image_path} ({e})") + + +def save_colored_mask_image(masks, colors, out_path, bg_color=(0, 0, 0)): + """Render the first frame of each person's mask onto a single BGR image. + bg_color is BGR; default black.""" + H, W = masks[0].shape[1:] + frame = np.full((H, W, 3), bg_color, dtype=np.uint8) + for mask_t, color in zip(masks, colors): + frame[mask_t[0]] = color + cv2.imwrite(out_path, frame) + + +def save_real_pixel_mask_image(masks, frame_rgb, out_path): + """Black bg; mask regions show real pixels from frame_rgb (H,W,3 RGB).""" + H, W = masks[0].shape[1:] + canvas = np.zeros((H, W, 3), dtype=np.uint8) + frame_bgr = frame_rgb[:, :, ::-1] + for mask_t in masks: + canvas[mask_t[0]] = frame_bgr[mask_t[0]] + cv2.imwrite(out_path, canvas) + + +def write_colored_mask_video(masks, colors, out_path, fps, bg_color=(0, 0, 0)): + """Per-frame colored mask mp4 (each person painted in their BGR color). + bg_color is BGR; default black.""" + T = masks[0].shape[0] + H, W = masks[0].shape[1:] + rgb_colors = [(int(c[2]), int(c[1]), int(c[0])) for c in colors] # BGR -> RGB for moviepy + bg_rgb = (int(bg_color[2]), int(bg_color[1]), int(bg_color[0])) + + frames = [] + for t in range(T): + frame = np.full((H, W, 3), bg_rgb, dtype=np.uint8) + for mask_t, rgb in zip(masks, rgb_colors): + frame[mask_t[t]] = rgb + frames.append(frame) + + out_dir = os.path.dirname(out_path) + if out_dir: + os.makedirs(out_dir, exist_ok=True) + mpy.ImageSequenceClip(frames, fps=fps).write_videofile(out_path) + + +def _mask_to_numpy_bool(m_t): + if hasattr(m_t, 'cpu'): + m_t = m_t.cpu().numpy() + return np.asarray(m_t).astype(bool) + + +def match_driving_to_ref_by_center(drv_masks, drv_colors, ref_masks): + """Greedy nearest-center matching for the one-to-many case where SAM3 finds + more driving persons than ref (often false positives like instruments/props + matching the 'character' prompt). For each ref person, pick the not-yet-used + driving person whose first-frame mask centroid (normalized to [0,1]^2) is + closest. Returns (matched_masks, matched_colors) in ref order, len == len(ref). + """ + def norm_center(mask2d, H, W): + mt = _mask_to_numpy_bool(mask2d) + ys, xs = np.where(mt) + if len(xs) == 0: + return None + return (float(xs.mean()) / W, float(ys.mean()) / H) + + H_d, W_d = drv_masks[0].shape[1:] + H_r, W_r = ref_masks[0].shape[1:] + drv_c = [norm_center(m[0], H_d, W_d) for m in drv_masks] + ref_c = [norm_center(m[0], H_r, W_r) for m in ref_masks] + + used = set() + out_masks, out_colors = [], [] + for ri, rc in enumerate(ref_c): + if rc is None: + continue + best_i, best_d = None, float('inf') + for di, dc in enumerate(drv_c): + if di in used or dc is None: + continue + d = (dc[0] - rc[0]) ** 2 + (dc[1] - rc[1]) ** 2 + if d < best_d: + best_d, best_i = d, di + if best_i is None: + continue + used.add(best_i) + out_masks.append(drv_masks[best_i]) + out_colors.append(drv_colors[best_i]) + print(f" match: ref[{ri}] center={rc} -> drv[{best_i}] " + f"center={drv_c[best_i]} (d^2={best_d:.4f})") + return out_masks, out_colors + + +def _merge_overlapping_bboxes(boxes): + """Iteratively merge any pair of overlapping bboxes into their bounding union + until no overlaps remain. O(N^2) — fine since N <= max_persons (typically 2).""" + def overlaps(a, b): + return not (a[2] <= b[0] or b[2] <= a[0] or a[3] <= b[1] or b[3] <= a[1]) + def union(a, b): + return (min(a[0], b[0]), min(a[1], b[1]), max(a[2], b[2]), max(a[3], b[3])) + boxes = list(boxes) + changed = True + while changed: + changed = False + for i in range(len(boxes)): + for j in range(i + 1, len(boxes)): + if overlaps(boxes[i], boxes[j]): + boxes[i] = union(boxes[i], boxes[j]) + boxes.pop(j) + changed = True + break + if changed: + break + return boxes + + +def _frame_keep_mask(masks_t, H, W, crop_kind, bbox_margin): + """(H, W) bool mask of pixels to keep this frame. + crop_kind='mask': union of per-person silhouettes. + crop_kind='bbox': union of per-person bboxes (margin-padded, overlap-merged).""" + if crop_kind == 'mask': + keep = np.zeros((H, W), bool) + for m_t in masks_t: + keep |= _mask_to_numpy_bool(m_t) + return keep + boxes = [] + for m_t in masks_t: + ys, xs = np.where(_mask_to_numpy_bool(m_t)) + if len(xs) == 0: + continue + x0, x1 = int(xs.min()), int(xs.max()) + 1 + y0, y1 = int(ys.min()), int(ys.max()) + 1 + mx = int(round((x1 - x0) * bbox_margin)) + my = int(round((y1 - y0) * bbox_margin)) + boxes.append((max(0, x0 - mx), max(0, y0 - my), + min(W, x1 + mx), min(H, y1 + my))) + keep = np.zeros((H, W), bool) + for x0, y0, x1, y1 in _merge_overlapping_bboxes(boxes): + keep[y0:y1, x0:x1] = True + return keep + + +def _compute_steady_bbox(masks, H, W, bbox_margin): + """Single bbox covering the union of every per-person mask across all frames, + with fractional margin. Returns (x0, y0, x1, y1) or None if all masks empty.""" + T = masks[0].shape[0] + union = np.zeros((H, W), bool) + for m in masks: + for t in range(T): + union |= _mask_to_numpy_bool(m[t]) + ys, xs = np.where(union) + if len(xs) == 0: + return None + x0, x1 = int(xs.min()), int(xs.max()) + 1 + y0, y1 = int(ys.min()), int(ys.max()) + 1 + mx = int(round((x1 - x0) * bbox_margin)) + my = int(round((y1 - y0) * bbox_margin)) + return (max(0, x0 - mx), max(0, y0 - my), + min(W, x1 + mx), min(H, y1 + my)) + + +def write_kept_driving_video(video_frames_rgb, masks, out_path, fps, + crop_kind, bbox_margin=0.05): + """Same dims as the driving video; per-frame keep only the region selected by + crop_kind ('bbox' | 'mask' | 'steady_bbox') and blacken the rest. + + 'steady_bbox' uses one static bbox = union of all masks across all frames + (with margin), so the kept rectangle never moves — useful when the driving + has camera motion and you want a stable crop window.""" + T = masks[0].shape[0] + H, W = video_frames_rgb.shape[1:3] + + static_keep = None + if crop_kind == 'steady_bbox': + bbox = _compute_steady_bbox(masks, H, W, bbox_margin) + static_keep = np.zeros((H, W), bool) + if bbox is not None: + x0, y0, x1, y1 = bbox + static_keep[y0:y1, x0:x1] = True + print(f" steady_bbox: {bbox} (W={x1-x0}, H={y1-y0}) in {W}x{H} frame") + + frames = [] + for t in range(T): + if static_keep is not None: + keep = static_keep + else: + keep = _frame_keep_mask([masks[i][t] for i in range(len(masks))], + H, W, crop_kind, bbox_margin) + out = np.zeros_like(video_frames_rgb[t]) + out[keep] = video_frames_rgb[t][keep] + frames.append(out) + out_dir = os.path.dirname(out_path) + if out_dir: + os.makedirs(out_dir, exist_ok=True) + mpy.ImageSequenceClip(frames, fps=fps).write_videofile(out_path) diff --git a/SCAIL-Pose/README.md b/SCAIL-Pose/README.md new file mode 100644 index 0000000000000000000000000000000000000000..7ef5e9599d268959e06967cff5a81975d2117288 --- /dev/null +++ b/SCAIL-Pose/README.md @@ -0,0 +1,197 @@ +

Official Code for Processing Driving Videos for SCAIL Series

+
+ + + + + +
+ + +This repository contains the code to process driving videos for **SCAIL**, a series of frameworks towards Studio-Grade Character Animation via In-Context Learning. The frameworks enable complex animation under diverse and challenging +conditions, including large motion variations and multi-character interactions. The main repo is at [zai-org/SCAIL](https://github.com/zai-org/SCAIL). +

+ teaser
+ SCAIL +

+ +

+ teaser
+ SCAIL-2 +

+ + +## 📋 Methods +**SCAIL** is a series of frameworks towards Studio-Grade Character Animation via In-Context Learning. The first open-source work of this series is SCAIL-Preview, a pose-driven animation framework. We develop a 3D skeleton for the pose representation to be fully identity agnostic and depth-aware. The representation can process multi-human interactions, yielding robust results from [NLFPose](https://github.com/isarandi/nlf)’s reliable depth estimation. + +

+ Teaser +

+ +Despite current progress, skeleton maps suffer from inherent ambiguity under complex scenarios. As intermediates, skeleton maps suffer from inherent ambiguity under complex scenarios. Further, it restricts the driving source to be exocentric human movements and thus cannot handle driving sources like animals. Character replacement and multi-character animation suffers from similar issues, where state-of-the-art methods use inpainting masks, but such masks are still a form of intermediates and limits the application and bounds the performance. + +

+ Preteaser +

+ +Our latest **SCAIL-2** is an end-to-end framework to bypass the pose estimation to obtain more reliable and expressive motion, utilizing the inherent in-context learning capability in the diffusion transformer. We adopt a unification design to support both Animation Mode and Replacement Mode, using + [SAM3](https://github.com/facebookresearch/sam3) to extract the explicit mask for both the reference image and the driving sequence to augment the conditioning. Benefiting from the end-to-end unification, SCAIL-2 supports diverse driving tasks. You can directly use the full driving video to drive the reference image, or use pose-driven just like SCAIL-Preview. We will elaborate different ways of driving in lateral usage instructions. + + +## 🚀 Getting Started + +Make sure you have already clone the main repo, this repo should be cloned under the main repo folder: +``` +SCAIL/ (or SCAIL-2/) +├── examples +├── sat +├── configs +├── ... +├── SCAIL-Pose +``` + +Change dir to this pose extraction & rendering folder: + +``` +cd SCAIL-Pose/ +``` + +### Environment Setup + +We recommend using [mmpose](https://github.com/open-mmlab) for the environment setup. You can refer to the official +mmpose [installation guide](https://mmpose.readthedocs.io/en/latest/installation.html). Note that the example in the guide uses python 3.8, however we recommend using python>=3.10 for better compatibility with SAM models. +The following commands are used to install the required packages once you have setup the environment. + +```bash +conda activate openmmlab +pip install -r requirements.txt + +# [Optional] SAM2 is only for multi-human extraction of SCAIL-Preview, for SCAIL-2 we use SAM3 +git clone https://github.com/facebookresearch/sam2.git && cd sam2 +pip install -e . +cd checkpoints && \ +./download_ckpts.sh && \ +cd ../.. +``` + + + +### Weights Download + +First, download pretrained weights for pose extraction & rendering. The script below +downloads [NLFPose](https://github.com/isarandi/nlf) (torchscript), [DWPose](https://github.com/IDEA-Research/DWPose) ( +onnx) and [YOLOX](https://github.com/Megvii-BaseDetection/YOLOX) (onnx) weights. You can also download the weights +manually and put them into the `pretrained_weights` folder. + +```bash +mkdir pretrained_weights && cd pretrained_weights +# download NLFPose Model Weights +wget https://github.com/isarandi/nlf/releases/download/v0.3.2/nlf_l_multi_0.3.2.torchscript +# download DWPose Model Weights & Detection Model Weights +mkdir DWPose +wget -O DWPose/dw-ll_ucoco_384.onnx \ + https://huggingface.co/yzd-v/DWPose/resolve/main/dw-ll_ucoco_384.onnx +wget -O DWPose/yolox_l.onnx \ + https://huggingface.co/yzd-v/DWPose/resolve/main/yolox_l.onnx +cd .. +``` + +For **SCAIL-2**, you additionally need the SAM3 weights. SAM3 is gated on HuggingFace, +so you must first request access at [facebook/sam3](https://huggingface.co/facebook/sam3) +and agree to Meta's license. Once approved, download `sam3.pt` into `pretrained_weights/`: + +```bash +# After being granted access on HuggingFace +huggingface-cli login +huggingface-cli download facebook/sam3 sam3.pt --local-dir pretrained_weights/ +``` + +The weights should be formatted as follows: + +``` +pretrained_weights/ +├── nlf_l_multi_0.3.2.torchscript +├── sam3.pt +└── DWPose/ + ├── dw-ll_ucoco_384.onnx + └── yolox_l.onnx +``` + + +## 🦾 Usage + +### SCAIL-Preview + +``` +# Single Character w/o 3D Retarget +python NLFPoseExtract/v1_process_pose.py --subdir --resolution [512, 896] + +# Single Character w/ 3D Retarget +python NLFPoseExtract/v1_process_pose.py --subdir --use_align --resolution [512, 896] + +# Multi-Human +python NLFPoseExtract/v1_process_pose_multi.py --subdir --resolution [512, 896] +``` + +### SCAIL-2 +For SCAIL-2, two entrypoints cover the two tasks: **Animation** (`process_animation_aio.py`) and **Replacement** (`process_replacement.py`). + +#### Animation Mode + +```bash +# (Recommended) End-to-end: rendered_v2.mp4 = driving copy, mask video is colored SAM3 masks. +# More accurate and easier than pose-driven for most cases. +python NLFPoseExtract/process_animation_aio.py --subdir --e2e_mode + +# Pose-driven (no --e2e_mode): runs NLF + DWpose, rendered_v2.mp4 is the skeleton render. +# More interpretable / controllable; use it for extremely challenging inputs. +python NLFPoseExtract/process_animation_aio.py --subdir + +## Following options allow behaviours between pose-driven and full-e2e. Useful for 704p horizontal / multi-human inputs where the zero-shot resolution gap causes artifacts + +# E2E + per-frame mask silhouette crop. +python NLFPoseExtract/process_animation_aio.py --subdir --e2e_mode --crop_e2e_mask + +# E2E + per-frame bbox crop. +python NLFPoseExtract/process_animation_aio.py --subdir --e2e_mode --crop_e2e_bbox + + +``` + +Other useful flags: `--max_persons N` (default 2), `--text human character ...` (extra SAM3 prompts, e.g. add `"robot arm" "gripper"` for egocentric/robotic subjects), `--sam3_model ` (override the default `pretrained_weights/sam3.pt` location). The same `--sam3_model` flag is also accepted by `process_replacement.py`. + +#### Replacement Mode + +```bash +# Standard: ref image in the subdir, driving has 1 actor matching the ref. +python NLFPoseExtract/process_replacement.py --subdir + +# Driving has 2 persons but you only want to replace 1: pick the driving track whose first-frame mask has +# highest IoU with the ref mask; drop the other. +python NLFPoseExtract/process_replacement.py --subdir --matchnearest +``` + +Examples are in the main repo folder; you can also use your own images or videos. After extraction the results live in the example folder and can be fed straight into the main repo to generate character animations. + +#### Notes for Animation +Although our model supports a variety of driving modalities, end-to-end driving typically achieves the best results, as the model has access to the complete visual information. This is especially evident in cases involving object interactions. + +

+ Preteaser +

+ + + + +## 📄 Citation + +If you find this work useful in your research, please cite: + +```bibtex +@article{yan2025scail, + title={SCAIL: Towards Studio-Grade Character Animation via In-Context Learning of 3D-Consistent Pose Representations}, + author={Yan, Wenhao and Ye, Sheng and Yang, Zhuoyi and Teng, Jiayan and Dong, ZhenHui and Wen, Kairui and Gu, Xiaotao and Liu, Yong-Jin and Tang, Jie}, + journal={arXiv preprint arXiv:2512.05905}, + year={2025} +} +``` diff --git a/SCAIL-Pose/TrackSam3/track.py b/SCAIL-Pose/TrackSam3/track.py new file mode 100644 index 0000000000000000000000000000000000000000..11f3b3425264eac9cc71da16ba89e3bd4ea07060 --- /dev/null +++ b/SCAIL-Pose/TrackSam3/track.py @@ -0,0 +1,355 @@ +import os +from random import shuffle + +import cv2 +import numpy as np +from decord import VideoReader + + +# Minimum mask ratio threshold (percentage of frame). Override via env var +# SCAIL_MIN_MASK_RATIO for small-subject scenes (e.g. paper figures where the +# subject occupies <1% of the frame) without editing this file. +MIN_MASK_RATIO = float(os.environ.get('SCAIL_MIN_MASK_RATIO', '1.0')) + +# Default cap on number of targets when caller does not override +DEFAULT_MAX_TARGETS = 4 + +# Deterministic BGR palette used when callers want stable colors across runs. +DEFAULT_PALETTE_BGR = [ + (255, 0, 0), # Blue + (0, 0, 255), # Red + (0, 255, 0), # Green + (255, 0, 255), # Magenta + (255, 255, 0), # Cyan + (0, 255, 255), # Yellow +] + + +def remove_small_tracks_from_predictor(predictor, invalid_track_ids): + """Remove invalid track IDs from predictor's internal tracker state.""" + if not invalid_track_ids: + return + + metadata = predictor.inference_state.get("tracker_metadata", {}) + if not metadata: + return + + obj_ids = metadata.get("obj_ids_all_gpu", np.array([])) + if len(obj_ids) == 0: + return + + keep_mask = np.array([int(oid) not in invalid_track_ids for oid in obj_ids]) + metadata["obj_ids_all_gpu"] = obj_ids[keep_mask] + + arrays_to_filter = [ + "obj_id_to_score", "obj_id_to_cls", "obj_id_to_tracker_score" + ] + for key in arrays_to_filter: + if key in metadata and isinstance(metadata[key], dict): + metadata[key] = {k: v for k, v in metadata[key].items() if int(k) not in invalid_track_ids} + + tracker_states = predictor.inference_state.get("tracker_inference_states", []) + if tracker_states: + for state in tracker_states: + if hasattr(state, 'obj_ids') and state.obj_ids is not None: + state_keep = np.array([int(oid) not in invalid_track_ids for oid in state.obj_ids]) + state.obj_ids = state.obj_ids[state_keep] + + print(f"Removed track IDs {invalid_track_ids} from tracker state") + + +def visualize_and_save_mask(results, width, height, predictor, new_indices, full_length, + max_targets=DEFAULT_MAX_TARGETS, shuffle_colors=True, + direct_return=False): + """Run through SAM3 streaming results and gather per-track binary masks. + + Returns (valid_track_ids_ordered, mask_arrays, track_colors) ordered by descending + mask area in the first frame; or None if no valid track is detected. + """ + colors = list(DEFAULT_PALETTE_BGR) + if shuffle_colors: + shuffle(colors) + + frame_idx = 0 + valid_track_ids = None + total_pixels = height * width + valid_track_ids_ordered = [] + mask_arrays = {} + track_colors = {} + _color_counter = 0 + + for result_idx, result in enumerate(results): + index_result = new_indices[result_idx] + if result.masks is not None: + masks = result.masks.data.cpu().numpy() # (N, H, W) + track_ids = result.boxes.id.cpu().numpy() if result.boxes.id is not None else np.arange(len(masks)) + + if frame_idx == 0: + valid_track_ids = set() + invalid_track_ids = set() + candidates = [] + for i, (mask, track_id) in enumerate(zip(masks, track_ids)): + if mask.shape[:2] != (height, width): + mask_resized = cv2.resize(mask.astype(np.float32), (width, height)) + else: + mask_resized = mask + mask_bool = mask_resized > 0.5 + mask_ratio = np.sum(mask_bool) / total_pixels * 100 + if mask_ratio >= MIN_MASK_RATIO: + candidates.append((int(track_id), mask_ratio)) + else: + invalid_track_ids.add(int(track_id)) + + candidates.sort(key=lambda x: x[1], reverse=True) + + if len(candidates) == 0 and direct_return: + print(f" No valid candidates (all < MIN_MASK_RATIO={MIN_MASK_RATIO}%) in first frame, return") + return + + if len(candidates) > max_targets: + if direct_return: + print(f" Found {len(candidates)} candidates, return") + return + print(f" Found {len(candidates)} candidates, limiting to top {max_targets}") + kept_candidates = candidates[:max_targets] + dropped_candidates = candidates[max_targets:] + for track_id, _ in kept_candidates: + valid_track_ids.add(track_id) + for track_id, _ in dropped_candidates: + invalid_track_ids.add(track_id) + else: + kept_candidates = candidates + for track_id, _ in candidates: + valid_track_ids.add(track_id) + + if kept_candidates and direct_return: + max_ratio = kept_candidates[0][1] + if max_ratio < 1.5 or max_ratio > 50: + print(f" Max mask ratio {max_ratio:.2f}% out of valid range [1.5, 50], return") + return + + if len(kept_candidates) >= 2 and direct_return: + max_ratio = kept_candidates[0][1] + min_ratio = kept_candidates[-1][1] + if min_ratio < max_ratio / 3: + print(f" Smallest person ({min_ratio:.2f}%) < 1/3 of largest ({max_ratio:.2f}%), return") + return + + valid_track_ids_ordered = [tid for tid, _ in kept_candidates] + mask_arrays = {tid: np.zeros((full_length, height, width), dtype=bool) + for tid in valid_track_ids_ordered} + + if invalid_track_ids: + remove_small_tracks_from_predictor(predictor, invalid_track_ids) + + for i, (mask, track_id) in enumerate(zip(masks, track_ids)): + if valid_track_ids is not None and int(track_id) not in valid_track_ids: + continue + + tid = int(track_id) + if tid not in track_colors: + track_colors[tid] = colors[_color_counter % len(colors)] + _color_counter += 1 + + if mask.shape[:2] != (height, width): + mask = cv2.resize(mask.astype(np.float32), (width, height)) + mask_bool = mask > 0.5 + + if tid in mask_arrays: + mask_arrays[tid][index_result] = mask_bool + + frame_idx += 1 + + if not valid_track_ids_ordered: + return None + return valid_track_ids_ordered, mask_arrays, track_colors + + +def _centroid_x(mask_2d): + """X-coordinate of the centroid of a 2D bool mask. Returns +inf if mask is empty.""" + cols = np.where(mask_2d.any(axis=0))[0] + if len(cols) == 0: + return float('inf') + rows = np.where(mask_2d.any(axis=1))[0] + # use bounding-box center (cheap and stable) + return 0.5 * (cols[0] + cols[-1]) + + +def _reorder_and_color(valid_track_ids_ordered, mask_arrays, sort_by, fixed_colors): + """Apply left-to-right sort and deterministic color assignment. + + Returns (masks, colors) where masks is a list of (T, H, W) bool ndarray and + colors is a list of BGR tuples, both in the chosen ordering. + """ + if sort_by == 'x': + ordered = sorted(valid_track_ids_ordered, + key=lambda tid: _centroid_x(mask_arrays[tid][0])) + elif sort_by == 'area': + ordered = list(valid_track_ids_ordered) + else: + raise ValueError(f"unknown sort_by: {sort_by}") + + n = len(ordered) + if fixed_colors is not None: + if len(fixed_colors) < n: + raise ValueError(f"fixed_colors has {len(fixed_colors)} entries but {n} tracks") + colors = [tuple(c) for c in fixed_colors[:n]] + else: + colors = [DEFAULT_PALETTE_BGR[i % len(DEFAULT_PALETTE_BGR)] for i in range(n)] + + masks = [mask_arrays[tid] for tid in ordered] + return masks, colors + + +def get_mask_from_video(video_path, predictor, max_targets=DEFAULT_MAX_TARGETS, + sort_by='area', fixed_colors=None, + text=("human", "character")): + """Run SAM3 tracking on a video file and return per-person binary masks and colors. + + Args: + video_path: path to input video (str or Path). + predictor: SAM3VideoSemanticPredictor instance (state will be reset). + max_targets: cap on the number of tracked persons (kept by descending area). + sort_by: 'area' (default, descending area) or 'x' (left-to-right by first-frame + centroid x). + fixed_colors: optional list of BGR tuples assigned to ordered tracks instead of + the default palette. Must have at least len(tracks) entries. + + Returns: + masks: list of (T, H, W) bool ndarray, one per tracked person. + colors: list of BGR color tuples corresponding to each person. + Both lists are empty if no valid persons are detected. + """ + video_path = str(video_path) + + predictor.inference_state = {} + if hasattr(predictor, 'dataset'): + predictor.dataset = None + + vr = VideoReader(video_path) + full_length = len(vr) + height, width = vr[0].asnumpy().shape[:2] + del vr + + results = predictor(source=video_path, text=list(text), stream=True) + ret = visualize_and_save_mask( + results, width, height, predictor, + new_indices=np.arange(full_length), full_length=full_length, + max_targets=max_targets, shuffle_colors=fixed_colors is None, + direct_return=False, + ) + if ret is None: + return [], [] + valid_track_ids_ordered, mask_arrays, _ = ret + return _reorder_and_color(valid_track_ids_ordered, mask_arrays, sort_by, fixed_colors) + + +def get_mask_from_image_via_video(image_path, video_predictor, max_targets=DEFAULT_MAX_TARGETS, + sort_by='x', fixed_colors=None, + text=("human", "character"), n_repeat=4, fps=8): + """Detect persons in a still image by wrapping it as a tiny mp4 and routing through + SAM3VideoSemanticPredictor. Workaround for image-mode SAM3 missing small / distant + subjects that the video pipeline picks up reliably. + + Returns (masks, colors) with each mask shaped (1, H, W) bool — only the first frame + of the synthetic clip is kept. + """ + import tempfile + from NLFPoseExtract.v2_helper import imread_bgr + image_path = str(image_path) + img = imread_bgr(image_path) + H, W = img.shape[:2] + + tmp_fd, tmp_path = tempfile.mkstemp(suffix='.mp4') + os.close(tmp_fd) + try: + fourcc = cv2.VideoWriter_fourcc(*'mp4v') + vw = cv2.VideoWriter(tmp_path, fourcc, float(fps), (W, H)) + if not vw.isOpened(): + raise RuntimeError(f"cv2.VideoWriter failed to open {tmp_path}") + for _ in range(n_repeat): + vw.write(img) + vw.release() + + masks, colors = get_mask_from_video( + tmp_path, video_predictor, + max_targets=max_targets, sort_by=sort_by, + fixed_colors=fixed_colors, text=text, + ) + finally: + try: + os.unlink(tmp_path) + except OSError: + pass + + masks = [m[:1] for m in masks] + return masks, colors + + +def get_mask_from_image(image_path, predictor, max_targets=DEFAULT_MAX_TARGETS, + sort_by='x', fixed_colors=None, + text=("human", "character")): + """Run SAM3SemanticPredictor (image variant) on a single image. + + Args: + image_path: path to input image (str or Path). + predictor: SAM3SemanticPredictor instance. + max_targets: cap on number of persons (kept by descending area). + sort_by: 'x' (default, left-to-right) or 'area'. + fixed_colors: optional list of BGR tuples assigned in order; otherwise the + deterministic palette is used. + + Returns: + masks: list of (1, H, W) bool ndarray, one per detected person. + colors: list of BGR color tuples corresponding to each person. + """ + image_path = str(image_path) + results = predictor(source=image_path, text=list(text)) + if not results: + return [], [] + result = results[0] + if result.masks is None or len(result.masks) == 0: + return [], [] + + masks_NHW = result.masks.data.cpu().numpy() # (N, H, W) + masks_NHW = masks_NHW > 0.5 + + return _filter_image_masks(masks_NHW, max_targets, sort_by, fixed_colors) + + +def _filter_image_masks(masks_NHW, max_targets, sort_by, fixed_colors): + """Apply MIN_MASK_RATIO + max_targets filter to per-image SAM masks, then order + and color them. Returns (masks_list, colors) where each mask is (1, H, W) bool. + """ + N, H, W = masks_NHW.shape + total_pixels = H * W + + candidates = [] # (idx, mask_ratio) + for i in range(N): + ratio = float(np.sum(masks_NHW[i])) / total_pixels * 100 + if ratio >= MIN_MASK_RATIO: + candidates.append((i, ratio)) + + candidates.sort(key=lambda x: x[1], reverse=True) + candidates = candidates[:max_targets] + if not candidates: + return [], [] + + kept = [masks_NHW[idx] for idx, _ in candidates] # list of (H, W) bool + + if sort_by == 'x': + order = sorted(range(len(kept)), key=lambda i: _centroid_x(kept[i])) + kept = [kept[i] for i in order] + elif sort_by != 'area': + raise ValueError(f"unknown sort_by: {sort_by}") + + n = len(kept) + if fixed_colors is not None: + if len(fixed_colors) < n: + raise ValueError(f"fixed_colors has {len(fixed_colors)} entries but {n} masks") + colors = [tuple(c) for c in fixed_colors[:n]] + else: + colors = [DEFAULT_PALETTE_BGR[i % len(DEFAULT_PALETTE_BGR)] for i in range(n)] + + masks_out = [m[None] for m in kept] # add T=1 axis + return masks_out, colors diff --git a/SCAIL-Pose/pose_draw/draw_3d_utils.py b/SCAIL-Pose/pose_draw/draw_3d_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..fbe1cb5f95029567aa165eea98752ddec9062d58 --- /dev/null +++ b/SCAIL-Pose/pose_draw/draw_3d_utils.py @@ -0,0 +1,220 @@ +import numpy as np + +def convert_3dpose_to_2dpose_body(body_keypoints, face_keypoints): + """ + 将20点的3D坐标映射到18点的2D坐标。 + :param poses: 输入的20点坐标列表,每个点为 [x, y, z] + :return: 映射得到的18点坐标列表,每个点为 [x, y] + """ + # 映射关系:索引位置 + body_mapping = { + 0: 1, 1: 2, 2: 3, 3: 4, 4: 5, 5: 6, 6: 7, 7: 8, 8: 8, 9: 9, 10: 10, 11: 23, 13: 22, 12: 21, + 14: 11, 15: 12, 16: 13, 17: 20, 18: 18, 19: 19 + } + face_mapping = { + 1: 16, 8: 14, 4: 0, 7: 15, 0: 17 + } + + # 初始化18点坐标列表,默认值为 [-1, -1] + result = [[-1, -1] for _ in range(24)] + + # 遍历映射关系,将对应的20点坐标映射到18点坐标 + for src_idx, dst_idx in body_mapping.items(): + if src_idx < len(body_keypoints): # 确保索引不越界 + result[dst_idx] = [body_keypoints[src_idx][1],body_keypoints[src_idx][0]] # 提取 x, y 坐标 + for src_idx, dst_idx in face_mapping.items(): + if src_idx < len(face_keypoints): + result[dst_idx] = [face_keypoints[src_idx][1], face_keypoints[src_idx][0]] + return result + +def convert_3dpose_to_2dpose_hand(left_hand_keypoints, right_hand_keypoints, body_keypoints): + """ + 将20点的3D坐标映射到18点的2D坐标。 + :param poses: 输入的20点坐标列表,每个点为 [x, y, z] + :return: 映射得到的18点坐标列表,每个点为 [x, y] + """ + # 映射关系:索引位置 + hand_mapping = { + 0: 1, 1: 2, 2: 3, 3: 4, 4: 5, 5: 6, 6: 7, 7: 8, 8: 9, 9: 10, + 10: 11, 11: 12, 12: 13, 13: 14, 14: 15, 15: 16, 16: 17, 17: 18, + 18: 19, 19: 20 + } + + body_mapping_left = {3: 0} + body_mapping_right = {6: 0} + + # 初始化18点坐标列表,默认值为 [-1, -1] + left_result = [[-1, -1] for _ in range(21)] + right_result = [[-1, -1] for _ in range(21)] + + # 遍历映射关系,将对应的20点坐标映射到18点坐标 + for src_idx, dst_idx in hand_mapping.items(): + if src_idx < len(left_hand_keypoints): # 确保索引不越界 + left_result[dst_idx] = [left_hand_keypoints[src_idx][1], left_hand_keypoints[src_idx][0]] # 提取 x, y 坐标 + right_result[dst_idx] = [right_hand_keypoints[src_idx][1], right_hand_keypoints[src_idx][0]] + + for src_idx, dst_idx in body_mapping_left.items(): + if src_idx < len(body_keypoints): + left_result[dst_idx] = [body_keypoints[src_idx][1], body_keypoints[src_idx][0]] + for src_idx, dst_idx in body_mapping_right.items(): + if src_idx < len(body_keypoints): + right_result[dst_idx] = [body_keypoints[src_idx][1], body_keypoints[src_idx][0]] + + return [left_result, right_result] + +def convert_3dpose_to_2dpose_face(face_keypoints): + result = [[-1, -1] if i in [0, 1, 4, 5, 6, 7, 8] else [pt[1], pt[0]] for i, pt in enumerate(face_keypoints)] + return result + +def correct_lift_end_kpt_by_phmr(start, end, dwpose_kpts, lift_start, lift_end, phmr_start, phmr_end): + ''' + 检查另一端是否符合要求, 符合要求则返回lift后结果,不然返回phmr结果 + ''' + if dwpose_kpts[start][0] == -1: + return + lift_vec = np.array(lift_end) - np.array(lift_start) + phmr_vec = np.array(phmr_end) - np.array(phmr_start) + start_distance = np.linalg.norm(np.array(lift_start) - np.array(phmr_start)) + end_distance = np.linalg.norm(np.array(lift_end) - np.array(phmr_end)) + lift_vec_len = np.linalg.norm(lift_vec) + phmr_vec_len = np.linalg.norm(phmr_vec) + if start_distance + end_distance > phmr_vec_len: + dwpose_kpts[end] = [-1, -1] + theta = np.arccos(np.dot(lift_vec, phmr_vec) / (lift_vec_len * phmr_vec_len)) + if lift_vec_len > phmr_vec_len * 1.65 or lift_vec_len < phmr_vec_len * 0.4 or theta > np.pi / 4: + dwpose_kpts[end] = [-1, -1] + return + + + +def mix_3d_poses(poses_dwpose, poses_3dpose): + ''' + 组合两种pose,用3dPose的身体,DWPose的face和hand + ''' + poses = [] + for pose_dwpose, pose_3dpose in zip(poses_dwpose, poses_3dpose): + pose = { + "bodies": { + "candidate": pose_3dpose["bodies"]["candidate"], + "subset": pose_dwpose["bodies"]["subset"] + }, + "faces": pose_dwpose["faces"], + "hands": pose_dwpose["hands"] + } + poses.append(pose) + return poses + +def correct_hand_from_3d(hand_keypoints_dwpose, hand_keypoints_3dpose): + ''' + 如果dwpose的手部关节点和3dpose的手部关节点相差过大,则去掉最远的那一端 + ''' + edges_palm = [ + [1, 2], [2, 3], [3, 4], + [5, 6], [6, 7], [7, 8], + [9, 10], [10, 11], [11, 12], + [13, 14], [14, 15], [15, 16], + [17, 18], [18, 19], [19, 20], + ] + edges_finger = [[0, 1], [0, 5], [0, 9], [0, 13], [0, 17]] + max_length_palm = 0 + max_length_finger = 0 + for edge in edges_palm: + limb_length_3dpose = np.linalg.norm(np.array(hand_keypoints_3dpose[edge[0]]) - np.array(hand_keypoints_3dpose[edge[1]])) + if limb_length_3dpose > max_length_palm: + max_length_palm = limb_length_3dpose + for edge in edges_finger: + limb_length_3dpose = np.linalg.norm(np.array(hand_keypoints_3dpose[edge[0]]) - np.array(hand_keypoints_3dpose[edge[1]])) + if limb_length_3dpose > max_length_finger: + max_length_finger = limb_length_3dpose + for edge in edges_palm: + limb_length_dwpose = np.linalg.norm(np.array(hand_keypoints_dwpose[edge[0]]) - np.array(hand_keypoints_dwpose[edge[1]])) + if limb_length_dwpose > max_length_palm * 1.5: + if -1 in hand_keypoints_dwpose[edge[0]] or -1 in hand_keypoints_dwpose[edge[1]] or -1 in hand_keypoints_3dpose[edge[0]] or -1 in hand_keypoints_3dpose[edge[1]]: + continue + distance_point_0 = np.linalg.norm(np.array(hand_keypoints_dwpose[edge[0]]) - np.array(hand_keypoints_3dpose[edge[0]])) + distance_point_1 = np.linalg.norm(np.array(hand_keypoints_dwpose[edge[1]]) - np.array(hand_keypoints_3dpose[edge[1]])) + if distance_point_0 > distance_point_1: + hand_keypoints_dwpose[edge[1]] = [-1, -1] + else: + hand_keypoints_dwpose[edge[0]] = [-1, -1] + for edge in edges_finger: + limb_length_dwpose = np.linalg.norm(np.array(hand_keypoints_dwpose[edge[0]]) - np.array(hand_keypoints_dwpose[edge[1]])) + if limb_length_dwpose > max_length_finger * 1.5: + if -1 in hand_keypoints_dwpose[edge[0]] or -1 in hand_keypoints_dwpose[edge[1]] or -1 in hand_keypoints_3dpose[edge[0]] or -1 in hand_keypoints_3dpose[edge[1]]: + continue + distance_point_0 = np.linalg.norm(np.array(hand_keypoints_dwpose[edge[0]]) - np.array(hand_keypoints_3dpose[edge[0]])) + distance_point_1 = np.linalg.norm(np.array(hand_keypoints_dwpose[edge[1]]) - np.array(hand_keypoints_3dpose[edge[1]])) + if distance_point_0 > distance_point_1: + hand_keypoints_dwpose[edge[1]] = [-1, -1] + else: + hand_keypoints_dwpose[edge[0]] = [-1, -1] + return hand_keypoints_dwpose + +def correct_body_from_3d(body_keypoints_dwpose, body_keypoints_3dpose, subset_dwpose, subset_3dpose): + ''' + 如果dwpose的骨骼长度和3dpose的骨骼长度相差过大,则去掉最远的那一端 + ''' + 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], + ] + + for ori_limb in limbSeq: + limb = [ori_limb[0] - 1, ori_limb[1] - 1] + limb_length_dwpose = np.linalg.norm(np.array(body_keypoints_dwpose[limb[0]]) - np.array(body_keypoints_dwpose[limb[1]])) + limb_length_3dpose = np.linalg.norm(np.array(body_keypoints_3dpose[limb[0]]) - np.array(body_keypoints_3dpose[limb[1]])) + if subset_dwpose[0][limb[0]] == -1 or subset_dwpose[0][limb[1]] == -1 or subset_3dpose[0][limb[0]] == -1 or subset_3dpose[0][limb[1]] == -1: + continue + if limb_length_dwpose > limb_length_3dpose * 2: + # 判断较远端 + distance_point_0 = np.linalg.norm(np.array(body_keypoints_dwpose[limb[0]]) - np.array(body_keypoints_3dpose[limb[0]])) + distance_point_1 = np.linalg.norm(np.array(body_keypoints_dwpose[limb[1]]) - np.array(body_keypoints_3dpose[limb[1]])) + if distance_point_0 > distance_point_1: + if limb[1] == 1: # 核心 + continue + body_keypoints_dwpose[limb[1]] = [-1, -1] + subset_dwpose[0][limb[1]] = -1 + else: + if limb[0] == 1: # 核心 + continue + body_keypoints_dwpose[limb[0]] = [-1, -1] + subset_dwpose[0][limb[0]] = -1 + return body_keypoints_dwpose, subset_dwpose + +def correct_full_pose_from_3d(poses_dwpose, poses_3dpose): + ''' + 如果dwpose的骨骼长度和3dpose的骨骼长度相差过大,则去掉离3d pose最远的那一端 + ''' + poses = [] + for pose_dwpose, pose_3dpose in zip(poses_dwpose, poses_3dpose): + new_candidate, new_subset = correct_body_from_3d(pose_dwpose["bodies"]["candidate"], pose_3dpose["bodies"]["candidate"], pose_dwpose["bodies"]["subset"], pose_3dpose["bodies"]["subset"]) + new_hands_0 = correct_hand_from_3d(pose_dwpose["hands"][0], pose_3dpose["hands"][0]) + new_hands_1 = correct_hand_from_3d(pose_dwpose["hands"][1], pose_3dpose["hands"][1]) + pose = { + "bodies": { + "candidate": new_candidate, + "subset": new_subset + }, + "faces": pose_dwpose["faces"], + "hands": [new_hands_0, new_hands_1] + } + poses.append(pose) + + return poses \ No newline at end of file diff --git a/SCAIL-Pose/pose_draw/draw_pose_utils.py b/SCAIL-Pose/pose_draw/draw_pose_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..d47cd3d2a9010eff39aae3d5c12e50e6b91e68ac --- /dev/null +++ b/SCAIL-Pose/pose_draw/draw_pose_utils.py @@ -0,0 +1,169 @@ +import cv2 +import numpy as np +from PIL import Image +import torch +import sys +import os +sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +import pose_draw.draw_utils as util +from pose_draw.draw_3d_utils import * +from pose_draw.reshape_utils import * +from DWPoseProcess.AAUtils import read_frames_and_fps_as_np, save_videos_from_pil +from DWPoseProcess.checkUtils import * +import random +import shutil +import argparse +import yaml # Add this import +import os +from tqdm import tqdm +from multiprocessing import Pool, cpu_count +from decord import VideoReader +import copy + + +def draw_pose(pose, H, W, show_feet=False, show_body=True, show_hand=True, show_face=True, show_cheek=False, dw_bgr=False, dw_hand=False, aug_body_draw=False, optimized_face=False): + final_canvas = np.zeros(shape=(H, W, 3), dtype=np.uint8) + for i in range(len(pose["bodies"]["candidate"])): + canvas = np.zeros(shape=(H, W, 3), dtype=np.uint8) + bodies = pose["bodies"] + faces = pose["faces"][i:i+1] + hands = pose["hands"][2*i:2*i+2] + candidate = bodies["candidate"][i] + subset = bodies["subset"][i:i+1] # subset是认为的有效点 + + if show_body: + if len(subset[0]) <= 18 or show_feet == False: + if aug_body_draw: + raise NotImplementedError("aug_body_draw is not implemented yet") + else: + canvas = util.draw_bodypose(canvas, candidate, subset) + else: + canvas = util.draw_bodypose_with_feet(canvas, candidate, subset) + if dw_bgr: + canvas = cv2.cvtColor(canvas, cv2.COLOR_BGR2RGB) + if show_cheek: + assert show_body == False, "show_cheek and show_body cannot be True at the same time" + canvas = util.draw_bodypose_augmentation(canvas, candidate, subset, drop_aug=True, shift_aug=False, all_cheek_aug=True) + if show_hand: + if not dw_hand: + canvas = util.draw_handpose_lr(canvas, hands) + else: + canvas = util.draw_handpose(canvas, hands) + if show_face: + canvas = util.draw_facepose(canvas, faces, optimized_face=optimized_face) + final_canvas = final_canvas + canvas + return final_canvas + + +def scale_image_hw_keep_size(img, scale_h, scale_w): + """分别按 scale_h, scale_w 缩放图像,保持输出尺寸不变。""" + H, W = img.shape[:2] + new_H, new_W = int(H * scale_h), int(W * scale_w) + scaled = cv2.resize(img, (new_W, new_H), interpolation=cv2.INTER_LINEAR) + + result = np.zeros_like(img) + + # 计算在目标图上的放置范围 + # --- Y方向 --- + if new_H >= H: + y_start_src = (new_H - H) // 2 + y_end_src = y_start_src + H + y_start_dst = 0 + y_end_dst = H + else: + y_start_src = 0 + y_end_src = new_H + y_start_dst = (H - new_H) // 2 + y_end_dst = y_start_dst + new_H + + # --- X方向 --- + if new_W >= W: + x_start_src = (new_W - W) // 2 + x_end_src = x_start_src + W + x_start_dst = 0 + x_end_dst = W + else: + x_start_src = 0 + x_end_src = new_W + x_start_dst = (W - new_W) // 2 + x_end_dst = x_start_dst + new_W + + # 将 scaled 映射到 result + result[y_start_dst:y_end_dst, x_start_dst:x_end_dst] = scaled[y_start_src:y_end_src, x_start_src:x_end_src] + + return result + +def draw_pose_to_canvas_np(poses, pool, H, W, reshape_scale, show_feet_flag=False, show_body_flag=True, show_hand_flag=True, show_face_flag=True, show_cheek_flag=False, dw_bgr=False, dw_hand=False, aug_body_draw=False): + canvas_np_lst = [] + for pose in poses: + if reshape_scale > 0: + pool.apply_random_reshapes(pose) + canvas = draw_pose(pose, H, W, show_feet_flag, show_body_flag, show_hand_flag, show_face_flag, show_cheek_flag, dw_bgr, dw_hand, aug_body_draw, optimized_face=True) + canvas_np_lst.append(canvas) + return canvas_np_lst + + +def draw_pose_to_canvas(poses, pool, H, W, reshape_scale, points_only_flag, show_feet_flag, show_body_flag=True, show_hand_flag=True, show_face_flag=True, show_cheek_flag=False, dw_bgr=False, dw_hand=False, aug_body_draw=False): + canvas_lst = [] + for pose in poses: + if reshape_scale > 0: + pool.apply_random_reshapes(pose) + canvas = draw_pose(pose, H, W, show_feet_flag, show_body_flag, show_hand_flag, show_face_flag, show_cheek_flag, dw_bgr, dw_hand, aug_body_draw, optimized_face=False) + canvas_img = Image.fromarray(canvas) + canvas_lst.append(canvas_img) + return canvas_lst + + +def get_mp4_filenames_from_directory(dwpose_keypoints_dir): + mp4_filenames_dwpose = [] + # 通过keypoints和mp4的交集取所有可用的mp4 + if dwpose_keypoints_dir: + for root, dirs, files in os.walk(dwpose_keypoints_dir): + for file in files: + if file.lower().endswith('.pt'): # 只查找 .mp4 文件 + mp4_filenames_dwpose.append(file.replace(".pt", ".mp4")) # 获取绝对路径 + return mp4_filenames_dwpose + +def project_dwpose_to_3d(dwpose_keypoint, original_threed_keypoint, focal, princpt, H, W): + # 相机内参 + # fx, fy = focal, focal + fx, fy = focal + cx, cy = princpt + + # 2D 关键点坐标 + x_2d, y_2d = dwpose_keypoint[0] * W, dwpose_keypoint[1] * H + + # 原始 3D 点(相机坐标系下) + ori_x, ori_y, ori_z = original_threed_keypoint + + # 使用新的 2D 点和原始深度反投影计算新的 3D 点 + # 公式: x = (u - cx) * z / fx + new_x = (x_2d - cx) * ori_z / fx + new_y = (y_2d - cy) * ori_z / fy + new_z = ori_z # 保持深度不变 + + return [new_x, new_y, new_z] + + + +def process_video(mp4_path, dwpose_keypoint_path, threed_keypoint_pair, reshape_scale, points_only_flag, show_feet_flag, wanted_fps=None, output_dirname=None, pose_type="dwpose"): + frames, fps = read_frames_and_fps_as_np(mp4_path) + initial_frame = frames[0] + output_path = os.path.join(output_dirname, os.path.basename(mp4_path)) + os.makedirs(output_dirname, exist_ok=True) + + if "3dpose" in pose_type: + raise NotImplementedError("3dpose is not implemented") + else: + poses = torch.load(dwpose_keypoint_path) + pool = reshapePool(alpha=reshape_scale) + canvas_lst = draw_pose_to_canvas(poses, pool, initial_frame.shape[0], initial_frame.shape[1], reshape_scale, points_only_flag, show_feet_flag, show_body_flag=True) + save_videos_from_pil(canvas_lst, output_path, wanted_fps) + + +def load_config(config_path): + with open(config_path, 'r') as f: + config = yaml.safe_load(f) + return config + + diff --git a/SCAIL-Pose/pose_draw/draw_utils.py b/SCAIL-Pose/pose_draw/draw_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..f7115e5d7518a75d57748e2030b9ff91aba040d9 --- /dev/null +++ b/SCAIL-Pose/pose_draw/draw_utils.py @@ -0,0 +1,658 @@ +# https://github.com/IDEA-Research/DWPose +import math +import numpy as np +import matplotlib +import cv2 +import random + +eps = 0.01 + + +def smart_resize(x, s): + Ht, Wt = s + if x.ndim == 2: + Ho, Wo = x.shape + Co = 1 + else: + Ho, Wo, Co = x.shape + if Co == 3 or Co == 1: + k = float(Ht + Wt) / float(Ho + Wo) + return cv2.resize( + x, + (int(Wt), int(Ht)), + interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4, + ) + else: + return np.stack([smart_resize(x[:, :, i], s) for i in range(Co)], axis=2) + + +def smart_resize_k(x, fx, fy): + if x.ndim == 2: + Ho, Wo = x.shape + Co = 1 + else: + Ho, Wo, Co = x.shape + Ht, Wt = Ho * fy, Wo * fx + if Co == 3 or Co == 1: + k = float(Ht + Wt) / float(Ho + Wo) + return cv2.resize( + x, + (int(Wt), int(Ht)), + interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4, + ) + else: + return np.stack([smart_resize_k(x[:, :, i], fx, fy) for i in range(Co)], axis=2) + + +def padRightDownCorner(img, stride, padValue): + h = img.shape[0] + w = img.shape[1] + + pad = 4 * [None] + pad[0] = 0 # up + pad[1] = 0 # left + pad[2] = 0 if (h % stride == 0) else stride - (h % stride) # down + pad[3] = 0 if (w % stride == 0) else stride - (w % stride) # right + + img_padded = img + pad_up = np.tile(img_padded[0:1, :, :] * 0 + padValue, (pad[0], 1, 1)) + img_padded = np.concatenate((pad_up, img_padded), axis=0) + pad_left = np.tile(img_padded[:, 0:1, :] * 0 + padValue, (1, pad[1], 1)) + img_padded = np.concatenate((pad_left, img_padded), axis=1) + pad_down = np.tile(img_padded[-2:-1, :, :] * 0 + padValue, (pad[2], 1, 1)) + img_padded = np.concatenate((img_padded, pad_down), axis=0) + pad_right = np.tile(img_padded[:, -2:-1, :] * 0 + padValue, (1, pad[3], 1)) + img_padded = np.concatenate((img_padded, pad_right), axis=1) + + return img_padded, pad + + +def transfer(model, model_weights): + transfered_model_weights = {} + for weights_name in model.state_dict().keys(): + transfered_model_weights[weights_name] = model_weights[ + ".".join(weights_name.split(".")[1:]) + ] + return transfered_model_weights + +def draw_bodypose_with_feet(canvas, candidate, subset): + H, W, C = canvas.shape + candidate = np.array(candidate) + subset = np.array(subset) + + stickwidth = 4 + + # 原始18个关节点的连接顺序(和 OpenPose 的 COCO 模型一致) + 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], + ] + + # 添加脚部连接线:10->18, 10->19, 10->20;13->21, 13->22, 13->23 + foot_limbSeq = [ + [14, 19], + [14, 20], + [14, 21], + [11, 22], + [11, 23], + [11, 24], + ] + + # 生成颜色(原始18条颜色 + 6条新颜色) + 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], + ] + + colors_feet = [ + [100, 0, 215], [80, 0, 235], [60, 0, 255], + [0, 235, 150], [0, 215, 170], [0, 195, 190], + ] + + colors = colors + colors_feet + + 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]) + + for i in range(6): + for n in range(len(subset)): + index = subset[n][np.array(foot_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_feet[i]) + + + canvas = (canvas * 0.6).astype(np.uint8) + + # 画关键点 + for i in range(24): + 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_bodypose_augmentation(canvas, candidate, subset, drop_aug=True, shift_aug=False, all_cheek_aug=False): + H, W, C = canvas.shape + candidate = np.array(candidate) + subset = np.array(subset) + + stickwidth = 4 + + limbSeq = [ + [2, 3], # 1->2 左肩 0 + [2, 6], # 1->5 右肩 1 + [3, 4], # 2->3 左臂 2 + [4, 5], # 3->4 左肘 3 + [6, 7], # 5->6 右臂 4 + [7, 8], # 6->7 右肘 5 + [2, 9], # 6 + [9, 10], # 7 + [10, 11], # 8 + [2, 12], # 9 + [12, 13], # 10 + [13, 14], # 11 + [2, 1], # 12 + [1, 15], # 13 cheek + [15, 17], # 14 cheek + [1, 16], # 15 cheek + [16, 18], # 16 cheek + [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], + ] + + # 随机选0-2根骨骼进行丢弃 + if drop_aug: + arr_drop = list(range(17)) + k_drop = random.choices([0, 1, 2], weights=[0.5, 0.3, 0.2])[0] + drop_indices = random.sample(arr_drop, k_drop) + else: + drop_indices = [] + if shift_aug: + shift_indices = random.sample(list(range(17)), 2) + else: + shift_indices = [] + if all_cheek_aug: + drop_indices = list(range(13)) # 0-12对应的骨骼都扔掉 + + 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) + + if i in drop_indices: + continue + + mX = np.mean(X) # 计算两个关节点之间的中点 + mY = np.mean(Y) + length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5 + if i in shift_indices: + mX = mX + random.uniform(-length/4, length/4) + mY = mY + random.uniform(-length/4, length/4) + 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): + if all_cheek_aug: + if not i in [0, 14, 15, 16, 17]: + continue + 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_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_lr(canvas, all_hand_peaks): + H, W, C = canvas.shape + + # 连接顺序:21个关键点的骨架连线 + 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], + ] + + all_num_hands = len(all_hand_peaks) + for peaks_idx, peaks in enumerate(all_hand_peaks): + left_or_right = not (peaks_idx >= all_num_hands / 2) + base_hue = 0 if left_or_right == 0 else 0.3 + 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: + if left_or_right == 0: + hsv_color = [ (base_hue + ie / float(len(edges)) * 0.8), 0.9, 0.9 ] + else: + hsv_color = [ (base_hue + ie / float(len(edges)) * 0.8), 0.8, 1 ] + rgb_color = matplotlib.colors.hsv_to_rgb(hsv_color) * 255 + cv2.line( + canvas, + (x1, y1), + (x2, y2), + rgb_color, + 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: + # 关键点也用淡色标注(左手蓝、右手红) + point_color = (245, 100, 100) if left_or_right == 0 else (100, 100, 255) + cv2.circle(canvas, (x, y), 4, point_color, thickness=-1) + + return canvas + +def draw_handpose(canvas, all_hand_peaks): + H, W, C = canvas.shape + stickwidth_thin = min(max(int(min(H, W) / 300), 1), 2) + + 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=stickwidth_thin, + ) + + for i, keyponit in enumerate(peaks): + x, y = keyponit + x = int(x * W) + y = int(y * H) + if x > eps and y > eps: + cv2.circle(canvas, (x, y), stickwidth_thin, (0, 0, 255), thickness=-1) + return canvas + + +def draw_facepose(canvas, all_lmks, optimized_face=True): + H, W, C = canvas.shape + stickwidth = min(max(int(min(H, W) / 200), 1), 3) + stickwidth_thin = min(max(int(min(H, W) / 300), 1), 2) + + for lmks in all_lmks: + lmks = np.array(lmks) + for lmk_idx, lmk in enumerate(lmks): + x, y = lmk + x = int(x * W) + y = int(y * H) + if x > eps and y > eps: + if optimized_face: + if lmk_idx in list(range(17, 27)) + list(range(36, 70)): + cv2.circle(canvas, (x, y), stickwidth_thin, (255, 255, 255), thickness=-1) + else: + cv2.circle(canvas, (x, y), stickwidth, (255, 255, 255), thickness=-1) + return canvas + + + + + +# detect hand according to body pose keypoints +# please refer to https://github.com/CMU-Perceptual-Computing-Lab/openpose/blob/master/src/openpose/hand/handDetector.cpp +def handDetect(candidate, subset, oriImg): + # right hand: wrist 4, elbow 3, shoulder 2 + # left hand: wrist 7, elbow 6, shoulder 5 + ratioWristElbow = 0.33 + detect_result = [] + image_height, image_width = oriImg.shape[0:2] + for person in subset.astype(int): + # if any of three not detected + has_left = np.sum(person[[5, 6, 7]] == -1) == 0 + has_right = np.sum(person[[2, 3, 4]] == -1) == 0 + if not (has_left or has_right): + continue + hands = [] + # left hand + if has_left: + left_shoulder_index, left_elbow_index, left_wrist_index = person[[5, 6, 7]] + x1, y1 = candidate[left_shoulder_index][:2] + x2, y2 = candidate[left_elbow_index][:2] + x3, y3 = candidate[left_wrist_index][:2] + hands.append([x1, y1, x2, y2, x3, y3, True]) + # right hand + if has_right: + right_shoulder_index, right_elbow_index, right_wrist_index = person[ + [2, 3, 4] + ] + x1, y1 = candidate[right_shoulder_index][:2] + x2, y2 = candidate[right_elbow_index][:2] + x3, y3 = candidate[right_wrist_index][:2] + hands.append([x1, y1, x2, y2, x3, y3, False]) + + for x1, y1, x2, y2, x3, y3, is_left in hands: + # pos_hand = pos_wrist + ratio * (pos_wrist - pos_elbox) = (1 + ratio) * pos_wrist - ratio * pos_elbox + # handRectangle.x = posePtr[wrist*3] + ratioWristElbow * (posePtr[wrist*3] - posePtr[elbow*3]); + # handRectangle.y = posePtr[wrist*3+1] + ratioWristElbow * (posePtr[wrist*3+1] - posePtr[elbow*3+1]); + # const auto distanceWristElbow = getDistance(poseKeypoints, person, wrist, elbow); + # const auto distanceElbowShoulder = getDistance(poseKeypoints, person, elbow, shoulder); + # handRectangle.width = 1.5f * fastMax(distanceWristElbow, 0.9f * distanceElbowShoulder); + x = x3 + ratioWristElbow * (x3 - x2) + y = y3 + ratioWristElbow * (y3 - y2) + distanceWristElbow = math.sqrt((x3 - x2) ** 2 + (y3 - y2) ** 2) + distanceElbowShoulder = math.sqrt((x2 - x1) ** 2 + (y2 - y1) ** 2) + width = 1.5 * max(distanceWristElbow, 0.9 * distanceElbowShoulder) + # x-y refers to the center --> offset to topLeft point + # handRectangle.x -= handRectangle.width / 2.f; + # handRectangle.y -= handRectangle.height / 2.f; + x -= width / 2 + y -= width / 2 # width = height + # overflow the image + if x < 0: + x = 0 + if y < 0: + y = 0 + width1 = width + width2 = width + if x + width > image_width: + width1 = image_width - x + if y + width > image_height: + width2 = image_height - y + width = min(width1, width2) + # the max hand box value is 20 pixels + if width >= 20: + detect_result.append([int(x), int(y), int(width), is_left]) + + """ + return value: [[x, y, w, True if left hand else False]]. + width=height since the network require squared input. + x, y is the coordinate of top left + """ + return detect_result + + +# Written by Lvmin +def faceDetect(candidate, subset, oriImg): + # left right eye ear 14 15 16 17 + detect_result = [] + image_height, image_width = oriImg.shape[0:2] + for person in subset.astype(int): + has_head = person[0] > -1 + if not has_head: + continue + + has_left_eye = person[14] > -1 + has_right_eye = person[15] > -1 + has_left_ear = person[16] > -1 + has_right_ear = person[17] > -1 + + if not (has_left_eye or has_right_eye or has_left_ear or has_right_ear): + continue + + head, left_eye, right_eye, left_ear, right_ear = person[[0, 14, 15, 16, 17]] + + width = 0.0 + x0, y0 = candidate[head][:2] + + if has_left_eye: + x1, y1 = candidate[left_eye][:2] + d = max(abs(x0 - x1), abs(y0 - y1)) + width = max(width, d * 3.0) + + if has_right_eye: + x1, y1 = candidate[right_eye][:2] + d = max(abs(x0 - x1), abs(y0 - y1)) + width = max(width, d * 3.0) + + if has_left_ear: + x1, y1 = candidate[left_ear][:2] + d = max(abs(x0 - x1), abs(y0 - y1)) + width = max(width, d * 1.5) + + if has_right_ear: + x1, y1 = candidate[right_ear][:2] + d = max(abs(x0 - x1), abs(y0 - y1)) + width = max(width, d * 1.5) + + x, y = x0, y0 + + x -= width + y -= width + + if x < 0: + x = 0 + + if y < 0: + y = 0 + + width1 = width * 2 + width2 = width * 2 + + if x + width > image_width: + width1 = image_width - x + + if y + width > image_height: + width2 = image_height - y + + width = min(width1, width2) + + if width >= 20: + detect_result.append([int(x), int(y), int(width)]) + + return detect_result + + +# get max index of 2d array +def npmax(array): + arrayindex = array.argmax(1) + arrayvalue = array.max(1) + i = arrayvalue.argmax() + j = arrayindex[i] + return i, j diff --git a/SCAIL-Pose/pose_draw/reshape_utils.py b/SCAIL-Pose/pose_draw/reshape_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..4dc8486564a6d68dc913c05f7650c98b0df22390 --- /dev/null +++ b/SCAIL-Pose/pose_draw/reshape_utils.py @@ -0,0 +1,444 @@ +import numpy as np +import random + +def pose_offset(offset_x, offset_y, pose, vitposes): + for i in range(len(pose["bodies"]["candidate"])): + bodies = pose["bodies"] + faces = pose["faces"][i] + right_hand = pose["hands"][2 * i] + left_hand = pose["hands"][2 * i + 1] + candidate = bodies["candidate"][i] + subset = bodies["subset"][i] + + for face_point in faces: + if face_point[0] == -1 and face_point[1] == -1: + continue + face_point[0] = face_point[0] + offset_x + face_point[1] = face_point[1] + offset_y + for left_hand_point in left_hand: + if left_hand_point[0] == -1 and left_hand_point[1] == -1: + continue + left_hand_point[0] = left_hand_point[0] + offset_x + left_hand_point[1] = left_hand_point[1] + offset_y + for right_hand_point in right_hand: + if right_hand_point[0] == -1 and right_hand_point[1] == -1: + continue + right_hand_point[0] = right_hand_point[0] + offset_x + right_hand_point[1] = right_hand_point[1] + offset_y + + assert len(candidate) == len(subset), f"candidate, length: {len(candidate)} and subset, length: {len(subset)} must have the same length" + for idx, body_point in enumerate(candidate): + if subset[idx] == -1: + continue + if body_point[0] == -1 and body_point[1] == -1: + continue + body_point[0] = body_point[0] + offset_x + body_point[1] = body_point[1] + offset_y + + if vitposes is not None: + for vitpose in vitposes: + vitpose['keypoints_body'] = vitpose['keypoints_body'] + np.array([offset_x, offset_y, 0]) + vitpose['keypoints_left_hand'] = vitpose['keypoints_left_hand'] + np.array([offset_x, offset_y, 0]) + vitpose['keypoints_right_hand'] = vitpose['keypoints_right_hand'] + np.array([offset_x, offset_y, 0]) + + +def pose_whole_scale(scale_x, scale_y, pose, vitposes): # 目前仅能用于单人 + # 获取最左端的x 最右端的x 最上端的y 最下端的y + min_x = float('inf') + max_x = float('-inf') + min_y = float('inf') + max_y = float('-inf') + if len(pose["bodies"]["candidate"]) > 1: + return + for i in range(len(pose["bodies"]["candidate"])): + bodies = pose["bodies"] + faces = pose["faces"][i] + right_hand = pose["hands"][2 * i] + left_hand = pose["hands"][2 * i + 1] + candidate = bodies["candidate"][i] + subset = bodies["subset"][i] + + for face_point in faces: + if face_point[0] == -1 and face_point[1] == -1: + continue + min_x = min(min_x, face_point[0]) + max_x = max(max_x, face_point[0]) + min_y = min(min_y, face_point[1]) + max_y = max(max_y, face_point[1]) + for left_hand_point in left_hand: + if left_hand_point[0] == -1 and left_hand_point[1] == -1: + continue + min_x = min(min_x, left_hand_point[0]) + max_x = max(max_x, left_hand_point[0]) + min_y = min(min_y, left_hand_point[1]) + max_y = max(max_y, left_hand_point[1]) + for right_hand_point in right_hand: + if right_hand_point[0] == -1 and right_hand_point[1] == -1: + continue + min_x = min(min_x, right_hand_point[0]) + max_x = max(max_x, right_hand_point[0]) + min_y = min(min_y, right_hand_point[1]) + max_y = max(max_y, right_hand_point[1]) + assert len(candidate) == len(subset), "candidate and subset must have the same length" + for idx, body_point in enumerate(candidate): + if subset[idx] == -1: + continue + if body_point[0] == -1 and body_point[1] == -1: + continue + min_x = min(min_x, body_point[0]) + max_x = max(max_x, body_point[0]) + min_y = min(min_y, body_point[1]) + max_y = max(max_y, body_point[1]) + + # 计算缩放比例 + # 根据bbox中心进行dilate, 倍数为x_scale, y_scale + bbox_center_x = (min_x + max_x) / 2 + bbox_center_y = (min_y + max_y) / 2 + for i in range(len(pose["bodies"]["candidate"])): + bodies = pose["bodies"] + candidate = bodies["candidate"][i] + for body_point in candidate: + if body_point[0] == -1 and body_point[1] == -1: + continue + body_point[0] = body_point[0] + (body_point[0] - bbox_center_x) * scale_x # scale越大,越远离中心,变化越大 + body_point[1] = body_point[1] + (body_point[1] - bbox_center_y) * scale_y + for face_point in faces: + if face_point[0] == -1 and face_point[1] == -1: + continue + face_point[0] = face_point[0] + (face_point[0] - bbox_center_x) * scale_x + face_point[1] = face_point[1] + (face_point[1] - bbox_center_y) * scale_y + for left_hand_point in left_hand: + if left_hand_point[0] == -1 and left_hand_point[1] == -1: + continue + left_hand_point[0] = left_hand_point[0] + (left_hand_point[0] - bbox_center_x) * scale_x + left_hand_point[1] = left_hand_point[1] + (left_hand_point[1] - bbox_center_y) * scale_y + for right_hand_point in right_hand: + if right_hand_point[0] == -1 and right_hand_point[1] == -1: + continue + right_hand_point[0] = right_hand_point[0] + (right_hand_point[0] - bbox_center_x) * scale_x + right_hand_point[1] = right_hand_point[1] + (right_hand_point[1] - bbox_center_y) * scale_y + + if vitposes is not None: + min_x = float('inf') + max_x = float('-inf') + min_y = float('inf') + max_y = float('-inf') + if len(vitposes) > 1: + return + for vitpose in vitposes: + for body_point in vitpose['keypoints_body']: + if body_point[2] < 0.4: + continue + min_x = min(min_x, body_point[0]) + max_x = max(max_x, body_point[0]) + min_y = min(min_y, body_point[1]) + max_y = max(max_y, body_point[1]) + for left_hand_point in vitpose['keypoints_left_hand']: + if left_hand_point[2] < 0.4: + continue + min_x = min(min_x, left_hand_point[0]) + max_x = max(max_x, left_hand_point[0]) + min_y = min(min_y, left_hand_point[1]) + max_y = max(max_y, left_hand_point[1]) + for right_hand_point in vitpose['keypoints_right_hand']: + if right_hand_point[2] < 0.4: + continue + min_x = min(min_x, right_hand_point[0]) + max_x = max(max_x, right_hand_point[0]) + min_y = min(min_y, right_hand_point[1]) + max_y = max(max_y, right_hand_point[1]) + bbox_center_x = (min_x + max_x) / 2 + bbox_center_y = (min_y + max_y) / 2 + for body_point in vitpose['keypoints_body']: + body_point[0] = body_point[0] + (body_point[0] - bbox_center_x) * scale_x + body_point[1] = body_point[1] + (body_point[1] - bbox_center_y) * scale_y + for left_hand_point in vitpose['keypoints_left_hand']: + left_hand_point[0] = left_hand_point[0] + (left_hand_point[0] - bbox_center_x) * scale_x + left_hand_point[1] = left_hand_point[1] + (left_hand_point[1] - bbox_center_y) * scale_y + for right_hand_point in vitpose['keypoints_right_hand']: + right_hand_point[0] = right_hand_point[0] + (right_hand_point[0] - bbox_center_x) * scale_x + right_hand_point[1] = right_hand_point[1] + (right_hand_point[1] - bbox_center_y) * scale_y + + + + +# 如果做多人增强,需要修改逻辑,每个人都有随机性,现在都是做单人增强,增强方法一致 +def pose_reshape(alpha, pose, vitposes, + anchor_point, anchor_part, end_point, end_part, + affected_body_points, affected_faces_points, affected_left_hands_points, affected_right_hands_points): + """ + 直接修改原始pose中的某些点的位置,通过根据一个不动点和端点之间的向量偏移被影响点的位置。 + 一次只能有一个anchor_point, 一个end_point,多个affected_points + alpha: 变化比例 + anchor_point: 起始点 + anchor_part: 起始点所属的身体部分,0代表body candidate, 1代表face, 2代表hand + end_point: 中止点 + end_part: 中止点所属的身体部分,0代表body candidate, 1代表face, 2代表hand + """ + for i in range(len(pose["bodies"]["candidate"])): + bodies = pose["bodies"] + # faces = pose["faces"][i:i+1] + faces = pose["faces"][i] + # hands = pose["hands"][2*i:2*i+2] + right_hand = pose["hands"][2 * i] + left_hand = pose["hands"][2 * i + 1] + candidate = bodies["candidate"][i] + # subset = bodies["subset"][i:i+1] + subset = bodies["subset"][i] + assert anchor_part == 0, "anchor part must belong to the body" + anchor_x, anchor_y = candidate[anchor_point] + if subset[anchor_point] == -1: + continue + + if end_part == 0: + end_x, end_y = candidate[end_point] + if subset[end_point] == -1: + continue + elif end_part == 1: + end_x, end_y = faces[end_point] + # 不考虑Hands + if end_x == -1 and end_y == -1: + continue + + vector_x = end_x - anchor_x + vector_y = end_y - anchor_y + offset_x = vector_x * alpha + offset_y = vector_y * alpha + + for affected_body_point in affected_body_points: + if subset[affected_body_point] == -1: + continue + affected_x, affected_y = candidate[affected_body_point] + candidate[affected_body_point] = [affected_x + offset_x, affected_y + offset_y] + for affected_faces_point in affected_faces_points: + affected_x, affected_y = faces[affected_faces_point] + if affected_x == -1 and affected_y == -1: + continue + faces[affected_faces_point] = [affected_x + offset_x, affected_y + offset_y] + for affected_hands_point in affected_left_hands_points: + affected_x, affected_y = left_hand[affected_hands_point] + if affected_x == -1 and affected_y == -1: + continue + left_hand[affected_hands_point] = [affected_x + offset_x, affected_y + offset_y] + for affected_hands_point in affected_right_hands_points: + affected_x, affected_y = right_hand[affected_hands_point] + if affected_x == -1 and affected_y == -1: + continue + right_hand[affected_hands_point] = [affected_x + offset_x, affected_y + offset_y] + + if vitposes is not None: + vit_visible_threshold = 0.5 + for vitpose in vitposes: + assert anchor_part == 0, "anchor part must belong to the body" + anchor_x, anchor_y, _ = vitpose['keypoints_body'][anchor_point] + if vitpose['keypoints_body'][anchor_point][2] < vit_visible_threshold: + continue + if end_part == 0: + end_x, end_y, _ = vitpose['keypoints_body'][end_point] + if vitpose['keypoints_body'][end_point][2] < vit_visible_threshold: + continue + # 不考虑hands + elif end_part == 1: + continue # ViTPose的脸不动 + + vector_x = end_x - anchor_x + vector_y = end_y - anchor_y + offset_x = vector_x * alpha + offset_y = vector_y * alpha + + for affected_body_point in affected_body_points: + if vitpose['keypoints_body'][affected_body_point][2] < vit_visible_threshold: + continue + vitpose['keypoints_body'][affected_body_point][0] = vitpose['keypoints_body'][affected_body_point][0] + offset_x + vitpose['keypoints_body'][affected_body_point][1] = vitpose['keypoints_body'][affected_body_point][1] + offset_y + # 左右标记相反,需要互换 + for affected_left_hand_point in affected_left_hands_points: + if vitpose['keypoints_right_hand'][affected_left_hand_point][2] < vit_visible_threshold: + continue + vitpose['keypoints_right_hand'][affected_left_hand_point][0] = vitpose['keypoints_right_hand'][affected_left_hand_point][0] + offset_x + vitpose['keypoints_right_hand'][affected_left_hand_point][1] = vitpose['keypoints_right_hand'][affected_left_hand_point][1] + offset_y + for affected_right_hand_point in affected_right_hands_points: + if vitpose['keypoints_left_hand'][affected_right_hand_point][2] < vit_visible_threshold: + continue + vitpose['keypoints_left_hand'][affected_right_hand_point][0] = vitpose['keypoints_left_hand'][affected_right_hand_point][0] + offset_x + vitpose['keypoints_left_hand'][affected_right_hand_point][1] = vitpose['keypoints_left_hand'][affected_right_hand_point][1] + offset_y + + + +# reshapePool只负责形变,骨骼偏移、丢弃等得从draw层来做 +class reshapePool: + def __init__(self, alpha): # 对每个视频只初始化一次 + self.faces_indices = np.arange(0, 68) + self.left_hands_indices = np.arange(0, 21) + self.right_hands_indices = np.arange(0, 21) + self.alpha = alpha # 0.1 + self.offset_x = random.uniform(-1/16, 1/16) + self.offset_y = random.uniform(-1/16, 1/16) + self.scale_x = random.uniform(-alpha/2, alpha/2) + self.scale_y = random.uniform(-alpha/2, alpha/2) + + self.body_reshape_methods = [ + self.extend_body, + self.extend_arm, + self.extend_leg, + self.shrink_body, + self.shrink_arm, + self.shrink_leg, + ] + self.scale_reshape_methods = [ + self.offset_wholebody, + # self.scale_wholebody, + self.dilate_face, + self.shrink_face, + ] + self.selected_methods = random.sample(self.body_reshape_methods, 2) + random.sample(self.scale_reshape_methods, random.choice([0, 1])) + + def apply_random_reshapes(self, pose, vitposes=None): + # Apply the two selected reshape methods + for method in self.selected_methods: + method(pose, vitposes) + + def offset_wholebody(self, pose, vitposes): + pose_offset(self.offset_x, self.offset_y, pose, vitposes) + + def scale_wholebody(self, pose, vitposes): + pose_whole_scale(self.scale_x, self.scale_y, pose, vitposes) + + + def extend_body(self, pose, vitposes): + pose_reshape(self.alpha, pose, vitposes, + 1, 0, 8, 0, + [8, 9, 10], [], [], [] + ) + pose_reshape(self.alpha, pose, vitposes, + 1, 0, 11, 0, + [11, 12, 13], [], [], [] + ) + + def shrink_body(self, pose, vitposes): + pose_reshape(-self.alpha, pose, vitposes, + 1, 0, 8, 0, + [8, 9, 10], [], [], [] + ) + pose_reshape(-self.alpha, pose, vitposes, + 1, 0, 11, 0, + [11, 12, 13], [], [], [] + ) + + def extend_arm(self, pose, vitposes): + pose_reshape(self.alpha, pose, vitposes, + 2, 0, 3, 0, + [3, 4], [], self.left_hands_indices, [] + ) + pose_reshape(self.alpha, pose, vitposes, + 5, 0, 6, 0, + [6, 7], [], [], self.right_hands_indices + ) + pose_reshape(self.alpha, pose, vitposes, + 3, 0, 4, 0, + [4], [], self.left_hands_indices, [] + ) + pose_reshape(self.alpha, pose, vitposes, + 6, 0, 7, 0, + [7], [], [], self.right_hands_indices + ) + + def shrink_arm(self, pose, vitposes): + pose_reshape(-self.alpha, pose, vitposes, + 2, 0, 3, 0, + [3, 4], [], self.left_hands_indices, [] + ) + pose_reshape(-self.alpha, pose, vitposes, + 5, 0, 6, 0, + [6, 7], [], [], self.right_hands_indices + ) + pose_reshape(-self.alpha, pose, vitposes, + 3, 0, 4, 0, + [4], [], self.left_hands_indices, [] + ) + pose_reshape(-self.alpha, pose, vitposes, + 6, 0, 7, 0, + [7], [], [], self.right_hands_indices + ) + + def extend_leg(self, pose, vitposes): + pose_reshape(self.alpha, pose, vitposes, + 8, 0, 9, 0, + [9, 10], [], [], [] + ) + pose_reshape(self.alpha, pose, vitposes, + 11, 0, 12, 0, + [12, 13], [], [], [] + ) + pose_reshape(self.alpha, pose, vitposes, + 9, 0, 10, 0, + [10], [], [], [] + ) + pose_reshape(self.alpha, pose, vitposes, + 12, 0, 13, 0, + [13], [], [], [] + ) + + def shrink_leg(self, pose, vitposes): + pose_reshape(-self.alpha, pose, vitposes, + 8, 0, 9, 0, + [9, 10], [], [], [] + ) + pose_reshape(-self.alpha, pose, vitposes, + 11, 0, 12, 0, + [12, 13], [], [], [] + ) + pose_reshape(-self.alpha, pose, vitposes, + 9, 0, 10, 0, + [10], [], [], [] + ) + pose_reshape(-self.alpha, pose, vitposes, + 12, 0, 13, 0, + [13], [], [], [] + ) + + def extend_shoulder(self, pose, vitposes): + pose_reshape(self.alpha, pose, vitposes, + 1, 0, 2, 0, + [2, 3, 4], [], self.left_hands_indices, [] + ) + pose_reshape(self.alpha, pose, vitposes, + 1, 0, 5, 0, + [5, 6, 7], [], [], self.right_hands_indices + ) + + def shrink_shoulder(self, pose, vitposes): + pose_reshape(-self.alpha, pose, vitposes, + 1, 0, 2, 0, + [2, 3, 4], [], self.left_hands_indices, [] + ) + pose_reshape(-self.alpha, pose, vitposes, + 1, 0, 5, 0, + [5, 6, 7], [], [], self.right_hands_indices + ) + + + def dilate_face(self, pose, vitposes): + for i in self.faces_indices: + pose_reshape(self.alpha, pose, vitposes, + 0, 0, i, 1, + [], [i], [], [] + ) + for i in [14, 15, 16, 17]: + pose_reshape(self.alpha, pose, vitposes, + 0, 0, i, 0, + [i], [], [], [] + ) + + def shrink_face(self, pose, vitposes): # 缩小脸会看不清 + for i in self.faces_indices: + pose_reshape(-self.alpha / 2, pose, vitposes, + 0, 0, i, 1, + [], [i], [], [] + ) + + for i in [14, 15, 16, 17]: + pose_reshape(-self.alpha / 2, pose, vitposes, + 0, 0, i, 0, + [i], [], [], [] + ) \ No newline at end of file diff --git a/SCAIL-Pose/render_3d/render_cylinder.py b/SCAIL-Pose/render_3d/render_cylinder.py new file mode 100644 index 0000000000000000000000000000000000000000..5b0e74d71ed3cebf66b5df741b97a7bc0bc0f277 --- /dev/null +++ b/SCAIL-Pose/render_3d/render_cylinder.py @@ -0,0 +1,129 @@ +import os +import torch +import matplotlib.pyplot as plt +import numpy as np +import cv2 +from PIL import Image + + +def render_colored_cylinders(cylinder_specs, focal, princpt, image_size=(1280, 1280), img=None): + os.environ['PYOPENGL_PLATFORM'] = 'osmesa' + import pyrender + import trimesh + H, W = image_size + if isinstance(focal, float) or isinstance(focal, int): + fx, fy = focal, focal + else: + fx, fy = focal[0], focal[1] + cx, cy = princpt + + # 初始化场景 + scene = pyrender.Scene(bg_color=[0, 0, 0, 0], ambient_light=[0.1, 0.1, 0.1]) + + # 设置相机 + camera = pyrender.IntrinsicsCamera(fx=fx, fy=fy, cx=cx, cy=cy, znear=0.5, zfar=10000) + pyrender2opencv = np.array([[1.0, 0, 0, 0], + [0, -1, 0, 0], + [0, 0, -1, 0], + [0, 0, 0, 1]]) + cam_pose = pyrender2opencv @ np.eye(4) + scene.add(camera, pose=cam_pose) + + # 添加光源 + light = pyrender.DirectionalLight(color=np.ones(3), intensity=3.0) + scene.add(light, pose=cam_pose) + + points_to_draw = [] + + for start, end, color in cylinder_specs: + start = np.array(start) + end = np.array(end) + vec = end - start + height = np.linalg.norm(vec) + if height == 0: + continue + + tm = trimesh.creation.cylinder(radius=12, height=height, sections=16) + + # 旋转对齐z轴 + z_axis = np.array([0, 0, 1]) + axis = np.cross(z_axis, vec) + if np.linalg.norm(axis) > 1e-6: + axis = axis / np.linalg.norm(axis) + angle = np.arccos(np.dot(z_axis, vec) / height) + rot = trimesh.transformations.rotation_matrix(angle, axis) + tm.apply_transform(rot) + + tm.apply_translation(start + vec / 2) + + # 材质颜色(支持 RGBA) + rgba = np.array(color) + material = pyrender.MetallicRoughnessMaterial( + metallicFactor=0.1, + roughnessFactor=0.5, + baseColorFactor=rgba + ) + + mesh = pyrender.Mesh.from_trimesh(tm, material=material) + scene.add(mesh) + + # 投影点用于可视化,检查投射是否正确 + x1 = fx * (start[0] / start[2]) + cx + y1 = fy * (start[1] / start[2]) + cy + x2 = fx * (end[0] / end[2]) + cx + y2 = fy * (end[1] / end[2]) + cy + points_to_draw.append((x1, y1)) + points_to_draw.append((x2, y2)) + + # 渲染 + r = pyrender.OffscreenRenderer(viewport_width=W, viewport_height=H, point_size=1.0) + color, _ = r.render(scene, flags=pyrender.RenderFlags.RGBA) + + # 后处理 + color = color.astype(np.float32) / 255.0 + # 转 uint8 + final_img = (color * 255).astype(np.uint8) + + # 画点,检查投射是否正确 + for (x, y) in points_to_draw: + print(f" debug point: {x}, {y}") + x_draw = int(x) + y_draw = int(y) + cv2.circle(final_img, (x_draw, y_draw), radius=4, color=(0, 255, 0), thickness=-1) + + return Image.fromarray(final_img) + + +# test +if __name__ == "__main__": + # 构造一个空白背景 + H, W = 480, 640 + img = np.zeros((H, W, 3), dtype=np.uint8) + 255 # 白色背景 + + # 构造几组3D点对和颜色 + cylinder_specs = [ + # 起点 (0,0,100), Y轴方向的红色圆柱终点 Y 调整为 40 + (np.array([0, 20, 120]), np.array([0, 40, 100]), [1.0, 0.0, 0.0, 1.0]), # 红色 + + # 起点 (0,0,100), X轴方向的绿色圆柱终点 X 调整为 60 + (np.array([0, 0, 100]), np.array([60, 40, 100]), [0.0, 1.0, 0.0, 1.0]), # 绿色 + + # Z轴方向的蓝色圆柱长度调整为50 (从100到150) + (np.array([0, 0, 100]), np.array([0, 0, 150]), [0.0, 0.0, 1.0, 1.0]), # 蓝色 + ] + + # 简单的相机参数 + fx, fy = 500, 500 + cx, cy = W // 2, H // 2 + + # 调用渲染函数 + img_pil = render_colored_cylinders( + cylinder_specs=cylinder_specs, + focal=(fx, fy), + princpt=(cx, cy), + image_size=(H, W), + img=img + ) + + # 显示或保存结果 + img_pil.save("test_render_cylinder.png") diff --git a/SCAIL-Pose/render_3d/taichi_cylinder.py b/SCAIL-Pose/render_3d/taichi_cylinder.py new file mode 100644 index 0000000000000000000000000000000000000000..f2f001ac4228e2af0adb99c202a9e5362d829772 --- /dev/null +++ b/SCAIL-Pose/render_3d/taichi_cylinder.py @@ -0,0 +1,217 @@ +import taichi as ti +import numpy as np +from PIL import Image +import random +import math +import time +import imageio +ti.init(arch=ti.cuda) + +def flatten_specs(specs_list): + """把 specs_list 拉平为 numpy 数组 + 索引表""" + starts, ends, colors = [], [], [] + frame_offset, frame_count = [], [] + offset = 0 + for specs in specs_list: + frame_offset.append(offset) + frame_count.append(len(specs)) + for (s, e, c) in specs: + starts.append(s) + ends.append(e) + colors.append(c) + offset += len(specs) + return ( + np.array(starts, dtype=np.float32), + np.array(ends, dtype=np.float32), + np.array(colors, dtype=np.float32), + np.array(frame_offset, dtype=np.int32), + np.array(frame_count, dtype=np.int32), + ) + +def render_whole(specs_list, H=480, W=640, fx=500, fy=500, cx=240, cy=320, radius=21.5, use_specular=True): + img = ti.Vector.field(4, dtype=ti.f32, shape=(H, W)) + starts, ends, colors, frame_offset, frame_count = flatten_specs(specs_list) + total_cyl = len(starts) + n_frames = len(specs_list) + z_min = min(starts[:, 2].min(), ends[:, 2].min()) + z_max = max(starts[:, 2].max(), ends[:, 2].max()) + + # ========= 相机内参 ========= + znear = 0.1 + zfar = max(min(z_max, 25000), 10000) + C = ti.Vector([0.0, 0.0, 0.0]) # 相机中心 + light_dir = ti.Vector([0.0, 0.0, 1.0]) + + c_start = ti.Vector.field(3, dtype=ti.f32, shape=total_cyl) + c_end = ti.Vector.field(3, dtype=ti.f32, shape=total_cyl) + c_rgba = ti.Vector.field(4, dtype=ti.f32, shape=total_cyl) + n_cyl = ti.field(dtype=ti.i32, shape=()) # 实际数量 + f_offset = ti.field(dtype=ti.i32, shape=n_frames) + f_count = ti.field(dtype=ti.i32, shape=n_frames) + frame_id = ti.field(dtype=ti.i32, shape=()) # 当前帧号 + z_min_field = ti.field(dtype=ti.f32, shape=()) + z_max_field = ti.field(dtype=ti.f32, shape=()) + use_spec_field = ti.field(dtype=ti.i32, shape=()) + + z_min_field[None] = z_min + z_max_field[None] = z_max + use_spec_field[None] = 1 if use_specular else 0 + + # # ====== 拷贝数据一次 ====== + c_start.from_numpy(starts) + c_end.from_numpy(ends) + c_rgba.from_numpy(colors) + f_offset.from_numpy(frame_offset) + f_count.from_numpy(frame_count) + + @ti.func + def sd_cylinder(p, a, b, r): + pa = p - a + ba = b - a + h = ba.norm() + eps = 1e-8 + res = 0.0 + if h < eps: + res = pa.norm() - r + else: + ba_n = ba / h + proj = pa.dot(ba_n) + proj_clamped = min(max(proj, 0.0), h) + res = (pa - proj_clamped * ba_n).norm() - r + return res + + @ti.func + def scene_sdf(p): + best_d = 1e6 + best_col = ti.Vector([0.0, 0.0, 0.0, 0.0]) + fid = frame_id[None] # 从 field 里读出来,变成一个普通 int + off = f_offset[fid] + cnt = f_count[fid] + for i in range(cnt): # 只遍历实际数量 + a = c_start[off + i] + b = c_end[off + i] + r = radius + col = c_rgba[off + i] + d = sd_cylinder(p, a, b, r) + if d < best_d: + best_d = d + best_col = col + return best_d, best_col + + @ti.func + def get_normal(p): + e = 1e-3 + dx = scene_sdf(p + ti.Vector([e, 0.0, 0.0]))[0] - scene_sdf(p - ti.Vector([e, 0.0, 0.0]))[0] + dy = scene_sdf(p + ti.Vector([0.0, e, 0.0]))[0] - scene_sdf(p - ti.Vector([0.0, e, 0.0]))[0] + dz = scene_sdf(p + ti.Vector([0.0, 0.0, e]))[0] - scene_sdf(p - ti.Vector([0.0, 0.0, e]))[0] + n = ti.Vector([dx, dy, dz]) + return n.normalized() + + @ti.func + def pixel_to_ray(xi, yi): + u = (xi - cx) / fx + v = (yi - cy) / fy + dir_cam = ti.Vector([u, v, 1.0]).normalized() + Rcw = ti.Matrix.identity(ti.f32, 3) + rd_world = Rcw @ dir_cam + ro_world = C + return ro_world, rd_world + + @ti.kernel + def render(): + depth_near, depth_far = ti.max(z_min_field[None], 0.1), ti.min(z_max_field[None] + 6000, 20000) # 能渲染出来的点,最大12000 + for y, x in img: + ro, rd = pixel_to_ray(x, y) + t = znear + col_out = ti.Vector([0.0, 0.0, 0.0, 0.0]) + for _ in range(300): + p = ro + rd * t + d, col = scene_sdf(p) + if d < 1e-3: + # n = get_normal(p) + # diff = max(n.dot(-light_dir), 0.0) + # lit = 0.3 + 0.7 * diff + # col_out = ti.Vector([col.x * lit, col.y * lit, col.z * lit, col.w]) + # break + + if use_spec_field[None] == 1: + n = get_normal(p) + diff = max(n.dot(-light_dir), 0.0) + + depth_factor = 1.0 - (p.z - depth_near) / (depth_far - znear) + depth_factor = ti.max(0.0, ti.min(1.0, depth_factor)) + + # diffuse/ambient 光照 + diffuse_term = 0.3 + 0.7 * diff + base = col.xyz * diffuse_term * depth_factor + + # === Blinn-Phong 镜面反射 === + view_dir = -rd.normalized() + half_dir = (view_dir + -light_dir).normalized() + spec = max(n.dot(half_dir), 0.0) ** 32 # shininess=32,越小越散,越大越锐 + highlight = ti.Vector([1.0, 1.0, 1.0]) * (0.5 * spec) * depth_factor + col_out = ti.Vector([base.x + highlight.x, + base.y + highlight.y, + base.z + highlight.z, + col.w]) + else: + # mono:无光照,纯平色 + col_out = ti.Vector([col.x, col.y, col.z, col.w]) + break + + if t > zfar: + break + t += max(d, 1e-4) + img[y, x] = col_out + + frames_np_rgba = [] + for f in range(len(specs_list)): + # start_time = time.time() + frame_id[None] = f + render() + arr = np.clip(img.to_numpy(), 0, 1) + # end_time = time.time() + # print(f"Frame {f} time: {end_time - start_time} seconds") + arr8 = (arr * 255).astype(np.uint8) + frames_np_rgba.append(arr8) + + return frames_np_rgba + + +def random_cylinder(): + """生成一根随机圆柱 (start, end, color)。""" + # 起点 [-200,200]^2, z 在 [-300,-100] + ax = random.uniform(-200, 200) + ay = random.uniform(-200, 200) + az = random.uniform(300, 400) + start = [ax, ay, az] + + # 随机方向和长度 + theta = random.uniform(0, 2*math.pi) + phi = random.uniform(-math.pi/4, math.pi/4) # 倾斜角 + L = 100 + dx = math.cos(phi) * math.cos(theta) + dy = math.cos(phi) * math.sin(theta) + dz = math.sin(phi) + end = [ax + dx * L, ay + dy * L, az + dz * L] + + # 随机颜色 (RGB + alpha=1) + color = [random.random(), random.random(), random.random(), 1.0] + + return (start, end, color) + +def generate_specs_list(num_frames=120, min_cyl=10, max_cyl=120): + """生成 specs_list,每帧有若干随机圆柱.""" + specs_list = [] + for _ in range(num_frames): + n_cyl = random.randint(min_cyl, max_cyl) + specs = [random_cylinder() for _ in range(n_cyl)] + specs_x_shift = [([spec[0][0] + 50, spec[0][1], spec[0][2]], [spec[1][0] + 50, spec[1][1], spec[1][2]], spec[2]) for spec in specs] + specs_list.append(specs) + specs_list.append(specs_x_shift) + return specs_list + + +if __name__ == "__main__": + specs_list = generate_specs_list(num_frames=24, min_cyl=10, max_cyl=120) + frames = render_whole(specs_list) diff --git a/SCAIL-Pose/requirements.txt b/SCAIL-Pose/requirements.txt new file mode 100644 index 0000000000000000000000000000000000000000..39b5e590239f351532e140ea04790e7ff3d8217c --- /dev/null +++ b/SCAIL-Pose/requirements.txt @@ -0,0 +1,23 @@ +onnxruntime-gpu>=1.23.2 +moviepy>=1.0.3 +taichi>=1.7.4 +openmim>=0.3.9 +torch +torchvision +mmcv>=2.2.0 +pre-commit>=4.5.0 +matplotlib==3.7 +tikzplotlib +jpeg4py +opencv-python +lmdb +pandas +scipy +loguru +moviepy +decord +controlnet_aux +webdataset +jsonlines +av +ultralytics \ No newline at end of file diff --git a/SCAIL-Pose/resources/animation.png b/SCAIL-Pose/resources/animation.png new file mode 100644 index 0000000000000000000000000000000000000000..e450144e4df9165ad947c86ddb54852c0fae3ce1 --- /dev/null +++ b/SCAIL-Pose/resources/animation.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d5eddd7caf6fbbbd40f56d861a7d00cc7f4a174df1e124fa825836c87ff721f5 +size 2284719 diff --git a/SCAIL-Pose/resources/data.png b/SCAIL-Pose/resources/data.png new file mode 100644 index 0000000000000000000000000000000000000000..f4e4272175d26bcf9ff44938faedd1cdeef41012 --- /dev/null +++ b/SCAIL-Pose/resources/data.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1235a1ba44c76988250536d83870a648d68646c2c8a2d4fcbaab462d1ac13b67 +size 643234 diff --git a/SCAIL-Pose/resources/pose_comp.png b/SCAIL-Pose/resources/pose_comp.png new file mode 100644 index 0000000000000000000000000000000000000000..0839657ec1a3440d336ba81a26d3f6620a9045ee --- /dev/null +++ b/SCAIL-Pose/resources/pose_comp.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:72efc184058c7dfedc5d986d6e9f927cb39bc7566bdd1f2b7a9f278107f51d78 +size 1171760 diff --git a/SCAIL-Pose/resources/pose_result.png b/SCAIL-Pose/resources/pose_result.png new file mode 100644 index 0000000000000000000000000000000000000000..42cbcbbc75b5932a3741993efc07378cb0aa8f68 --- /dev/null +++ b/SCAIL-Pose/resources/pose_result.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:56ae0cfd546b0797cef43e8006ab57f0baf83ea7a1c382eae737666223cae3f2 +size 705565 diff --git a/SCAIL-Pose/resources/pose_teaser.png b/SCAIL-Pose/resources/pose_teaser.png new file mode 100644 index 0000000000000000000000000000000000000000..710890b3ff3ae6073fc52494a303563472baf603 --- /dev/null +++ b/SCAIL-Pose/resources/pose_teaser.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2bf5e8413cbd7afad6d050099b7a55d349552c24df917e9b3135db550a2357f0 +size 3460705 diff --git a/SCAIL-Pose/resources/preteaser.png b/SCAIL-Pose/resources/preteaser.png new file mode 100644 index 0000000000000000000000000000000000000000..8ae5660e10c09bf28ba7a84016b0bb300dea80a7 --- /dev/null +++ b/SCAIL-Pose/resources/preteaser.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:04d1eaa51b780fbda7e2b097412eecf32d2c3b1b1deab573012034a0935d2225 +size 163401 diff --git a/SCAIL-Pose/resources/teaser.png b/SCAIL-Pose/resources/teaser.png new file mode 100644 index 0000000000000000000000000000000000000000..30b41a1f2905845e5b3314e7fd088a470b7f253e --- /dev/null +++ b/SCAIL-Pose/resources/teaser.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:56c3c1745976ad37bae932a77623482becc09936a548c4cad30295402baabb1e +size 4509665 diff --git a/configs/config-1.3b.json b/configs/config-1.3b.json new file mode 100644 index 0000000000000000000000000000000000000000..d2ac6dfc38382dec0e80a35d87a053593096a09f --- /dev/null +++ b/configs/config-1.3b.json @@ -0,0 +1,14 @@ +{ + "_class_name": "WanSCAILModel", + "_diffusers_version": "0.30.0", + "dim": 1536, + "eps": 1e-06, + "ffn_dim": 8960, + "freq_dim": 256, + "in_dim": 20, + "model_type": "i2v", + "num_heads": 12, + "num_layers": 30, + "out_dim": 16, + "text_len": 512 +} diff --git a/configs/config-14b.json b/configs/config-14b.json new file mode 100644 index 0000000000000000000000000000000000000000..a45f9d261c9ca4fe67396a7cc0f43e9ea732e377 --- /dev/null +++ b/configs/config-14b.json @@ -0,0 +1,15 @@ +{ + "_class_name": "WanSCAILModel", + "_diffusers_version": "0.30.0", + "dim": 5120, + "eps": 1e-06, + "ffn_dim": 13824, + "freq_dim": 256, + "in_dim": 20, + "mask_dim": 28, + "model_type": "i2v", + "num_heads": 40, + "num_layers": 40, + "out_dim": 16, + "text_len": 512 +} diff --git a/convert.py b/convert.py new file mode 100644 index 0000000000000000000000000000000000000000..b08ca6ca34de8837a3c0a3ddf9adcb43557779cf --- /dev/null +++ b/convert.py @@ -0,0 +1,199 @@ + +import torch +from safetensors.torch import save_file +from typing import Dict +import os + +class ModuleParser(): + def __init__(self, key: str): + self.key = key + self.modules = key.split('.') + self.idx = 0 + + def match(self, pattern: str) -> bool: + patterns = pattern.split('.') + for j, p in enumerate(patterns): + if self.idx + j < len(self.modules) and self.modules[self.idx+j] == p: + continue + else: + return False + self.idx += len(patterns) + return True + + def step(self) -> str: + m = self.modules[self.idx] + self.idx += 1 + return m + + def eof(self) -> bool: + return self.idx == len(self.modules) + +def get_new_mappings(key: str, param: torch.Tensor) -> Dict[str, torch.Tensor]: + modules = [] + parser = ModuleParser(key) + sat_prefix = "model.diffusion_model" + assert parser.match(sat_prefix), key + if parser.match("mixins"): + if parser.match("adaln_layer"): + if parser.match("adaLN_modulations"): + modules.append("blocks") + modules.append(parser.step()) + modules.append("modulation") + elif parser.match("query_layernorm_list"): + modules.append("blocks") + modules.append(parser.step()) + modules.append("self_attn.norm_q") + modules.append(parser.step()) + elif parser.match("key_layernorm_list"): + modules.append("blocks") + modules.append(parser.step()) + modules.append("self_attn.norm_k") + modules.append(parser.step()) + elif parser.match("cross_query_layernorm_list"): + modules.append("blocks") + modules.append(parser.step()) + modules.append("cross_attn.norm_q") + modules.append(parser.step()) + elif parser.match("cross_key_layernorm_list"): + modules.append("blocks") + modules.append(parser.step()) + modules.append("cross_attn.norm_k") + modules.append(parser.step()) + elif parser.match("clip_feature_key_layernorm_list"): + modules.append("blocks") + modules.append(parser.step()) + modules.append("cross_attn.norm_k_img") + modules.append(parser.step()) + elif parser.match("clip_feature_key_value_list"): + modules.append("blocks") + modules.append(parser.step()) + modules.append("cross_attn") + prefix = '.'.join(modules) + suffix = parser.step() # weight or bias + key_param, value_param = param.chunk(2, dim=0) + return { + ".".join([prefix, "k_img", suffix]): key_param, + ".".join([prefix, "v_img", suffix]): value_param, + } + else: + raise ValueError(key) + elif parser.match("final_layer"): + modules.append("head") + if parser.match("adaLN_modulation"): + modules.append("modulation") + elif parser.match("linear"): + modules.append("head") + modules.append(parser.step()) # weight or bias + elif parser.match("patch_embed"): + if parser.match("proj"): + modules.append("patch_embedding") + elif parser.match("proj_pose"): + modules.append("patch_embedding_pose") + elif parser.match("proj_mask"): + modules.append("patch_embedding_mask") + else: + raise ValueError(key) + modules.append(parser.step()) + else: + raise ValueError(key) + elif parser.match("transformer.layers"): + modules.append("blocks") + modules.append(parser.step()) + if parser.match("attention"): + modules.append("self_attn") + if parser.match("dense"): + modules.append("o") + modules.append(parser.step()) + elif parser.match("query_key_value"): + prefix = '.'.join(modules) + suffix = parser.step() + query_param, key_param, value_param = param.chunk(3, dim=0) + return { + ".".join([prefix, "q", suffix]): query_param, + ".".join([prefix, "k", suffix]): key_param, + ".".join([prefix, "v", suffix]): value_param, + } + else: + raise ValueError(key) + elif parser.match("cross_attention"): + modules.append("cross_attn") + if parser.match("dense"): + modules.append("o") + modules.append(parser.step()) + elif parser.match("query"): + modules.append("q") + modules.append(parser.step()) + elif parser.match("key_value"): + prefix = '.'.join(modules) + suffix = parser.step() + key_param, value_param = param.chunk(2, dim=0) + return { + ".".join([prefix, "k", suffix]): key_param, + ".".join([prefix, "v", suffix]): value_param, + } + else: + raise ValueError(key) + elif parser.match("post_cross_attention_layernorm"): + modules.append("norm3") + modules.append(parser.step()) # weight or bias + elif parser.match("mlp"): + modules.append("ffn") + if parser.match("dense_h_to_4h"): + modules.append("0") + elif parser.match("dense_4h_to_h"): + modules.append("2") + else: + raise ValueError(key) + modules.append(parser.step()) # weight or bias + else: + raise ValueError(key) + elif parser.match("time_embed"): + modules.append("time_embedding") + modules.append(parser.step()) + modules.append(parser.step()) + elif parser.match("adaln_projection"): + modules.append("time_projection") + modules.append(parser.step()) + modules.append(parser.step()) + elif parser.match("text_embedding"): + modules.append("text_embedding") + modules.append(parser.step()) + modules.append(parser.step()) + elif parser.match("clip_proj"): + assert parser.match("proj"), key + modules.append("img_emb.proj") + modules.append(parser.step()) + modules.append(parser.step()) + else: + raise ValueError(key) + assert parser.eof(), key + return {'.'.join(modules): param} + +def get_new_state_dict(old: Dict[str, torch.Tensor]): + new = dict() + for key, value in old.items(): + map = get_new_mappings(key, value) + for new_key, new_value in map.items(): + if new_key in new: + print(f"Warning: duplicate new key {new_key} converted from {key}!") + new[new_key] = new_value + return new + +def main(args): + pt_file_path = os.path.join(args.scail_dir, args.sat_model_path) + print(f"Loading from {pt_file_path}...") + checkpoint = torch.load(pt_file_path) + state_dict = checkpoint["module"] + new_state_dict = get_new_state_dict(state_dict) + print(f"Saving to {args.save_path}...") + save_file(new_state_dict, args.save_path) + print("Done.") + +import argparse +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--scail-dir", default="SCAIL-2/") + parser.add_argument("--sat-model-path", default="model/1/fsdp2_rank_0000_checkpoint.pt") + parser.add_argument("--save-path", default="SCAIL-2.safetensors") + args = parser.parse_args() + main(args) \ No newline at end of file diff --git a/examples/animation_001/combined.gif b/examples/animation_001/combined.gif new file mode 100644 index 0000000000000000000000000000000000000000..2afcef794af6e8a2d0f19cfec9f7612691aeb66c --- /dev/null +++ b/examples/animation_001/combined.gif @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:089acb223778fc3b26e1273abda72640ff9a310cdd0d46067ceb7dc15be91e74 +size 299420 diff --git a/examples/animation_001/driving.mp4 b/examples/animation_001/driving.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..90c576c5fbffa10f3fdcaebefb63430dae4dd1c1 --- /dev/null +++ b/examples/animation_001/driving.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a6395c6c0f927fa6d6f1f7e761aa902f31c90ae112f4ca38457baa2455108b6a +size 517884 diff --git a/examples/animation_001/ref.jpg b/examples/animation_001/ref.jpg new file mode 100644 index 0000000000000000000000000000000000000000..46453829bd18ed902a30f99249a5151da942df1c --- /dev/null +++ b/examples/animation_001/ref.jpg @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3485a50ddf5cf73d81da87ff7a9da4a16e4ad3e51cb7c8e1f2393d2dea79f781 +size 224421 diff --git a/examples/animation_001/ref_mask.jpg b/examples/animation_001/ref_mask.jpg new file mode 100644 index 0000000000000000000000000000000000000000..60bbc4f0567a3def18039f9bfb45a76fdb9a3391 Binary files /dev/null and b/examples/animation_001/ref_mask.jpg differ diff --git a/examples/animation_001/rendered_mask_v2.mp4 b/examples/animation_001/rendered_mask_v2.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..6b84a5023eb61307db2c0ed2044d20efc1eda7f8 --- /dev/null +++ b/examples/animation_001/rendered_mask_v2.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b8dca4f6148feb0ea4cd225ea5ab5d7500747f240e89312f1092ec88d68b5bba +size 206989 diff --git a/examples/animation_001/rendered_v2.mp4 b/examples/animation_001/rendered_v2.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..90c576c5fbffa10f3fdcaebefb63430dae4dd1c1 --- /dev/null +++ b/examples/animation_001/rendered_v2.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a6395c6c0f927fa6d6f1f7e761aa902f31c90ae112f4ca38457baa2455108b6a +size 517884 diff --git a/examples/animation_001_posedriven/combined.gif b/examples/animation_001_posedriven/combined.gif new file mode 100644 index 0000000000000000000000000000000000000000..5771b221ca8963001ec9369c130375bc9ddc1df1 --- /dev/null +++ b/examples/animation_001_posedriven/combined.gif @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6c64f4d51ba95ae30d80091db7a0448a8998dd0e80b5eb5f336d7947782bdaf7 +size 425579 diff --git a/examples/animation_001_posedriven/driving.mp4 b/examples/animation_001_posedriven/driving.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..90c576c5fbffa10f3fdcaebefb63430dae4dd1c1 --- /dev/null +++ b/examples/animation_001_posedriven/driving.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a6395c6c0f927fa6d6f1f7e761aa902f31c90ae112f4ca38457baa2455108b6a +size 517884 diff --git a/examples/animation_001_posedriven/ref.jpg b/examples/animation_001_posedriven/ref.jpg new file mode 100644 index 0000000000000000000000000000000000000000..46453829bd18ed902a30f99249a5151da942df1c --- /dev/null +++ b/examples/animation_001_posedriven/ref.jpg @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3485a50ddf5cf73d81da87ff7a9da4a16e4ad3e51cb7c8e1f2393d2dea79f781 +size 224421 diff --git a/examples/animation_001_posedriven/ref_mask.jpg b/examples/animation_001_posedriven/ref_mask.jpg new file mode 100644 index 0000000000000000000000000000000000000000..60bbc4f0567a3def18039f9bfb45a76fdb9a3391 Binary files /dev/null and b/examples/animation_001_posedriven/ref_mask.jpg differ diff --git a/examples/animation_001_posedriven/rendered_mask_v2.mp4 b/examples/animation_001_posedriven/rendered_mask_v2.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..cf05ea8418b8913a52ba9f02be7bdcfcf6580998 --- /dev/null +++ b/examples/animation_001_posedriven/rendered_mask_v2.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:226c2916348f76a003c16969e656bdce7d5dd3335e20aa69ee685b463ba19ec6 +size 293660 diff --git a/examples/animation_001_posedriven/rendered_v2.mp4 b/examples/animation_001_posedriven/rendered_v2.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..5fe9fdcc1a72f68bd6fd55472d6dd3b9b15b64b1 --- /dev/null +++ b/examples/animation_001_posedriven/rendered_v2.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:99d2f8128424ab88049c8f0c0d2febf6d13557be5903989fe1945491a8e7d388 +size 435034 diff --git a/examples/animation_002/driving.mp4 b/examples/animation_002/driving.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..f06582bd5aa446f7c802b1d5b2c126048b78118f --- /dev/null +++ b/examples/animation_002/driving.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0f20b6b4368ff99be96c7d87fd77d2ce27b05c1884c4a0a2002960527434ab84 +size 494414 diff --git a/examples/animation_002/ref.jpg b/examples/animation_002/ref.jpg new file mode 100644 index 0000000000000000000000000000000000000000..b1d00efbf59784dfc866e305dc46fc1f362e7c81 --- /dev/null +++ b/examples/animation_002/ref.jpg @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:526eeb944355c126ab7d41f2c54a1d8af7d9efe66edf758637eed99b14d5859b +size 103653 diff --git a/examples/animation_002/ref_mask.jpg b/examples/animation_002/ref_mask.jpg new file mode 100644 index 0000000000000000000000000000000000000000..2a4fea3c769aa5488bd00d463154ea3d265a210c Binary files /dev/null and b/examples/animation_002/ref_mask.jpg differ diff --git a/examples/animation_002/rendered_mask_v2.mp4 b/examples/animation_002/rendered_mask_v2.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..09266e37bc8de8d744610a89fa6971536b3d2067 --- /dev/null +++ b/examples/animation_002/rendered_mask_v2.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4f1032a1d9a885f88835a42dbfd6b2e1e383669fc34e2e6d8c6d77f254f50cec +size 176767 diff --git a/examples/animation_002/rendered_v2.mp4 b/examples/animation_002/rendered_v2.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..f06582bd5aa446f7c802b1d5b2c126048b78118f --- /dev/null +++ b/examples/animation_002/rendered_v2.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0f20b6b4368ff99be96c7d87fd77d2ce27b05c1884c4a0a2002960527434ab84 +size 494414 diff --git a/examples/animation_003_multi_ref/background.png b/examples/animation_003_multi_ref/background.png new file mode 100644 index 0000000000000000000000000000000000000000..4e29399404b2b2a34bab2a06bbc2262d9d86c7f9 --- /dev/null +++ b/examples/animation_003_multi_ref/background.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9164694778a010d91314f9c1da9f6678717139900ab2c59415a525faae25c08c +size 8675279 diff --git a/examples/animation_003_multi_ref/background_mask.png b/examples/animation_003_multi_ref/background_mask.png new file mode 100644 index 0000000000000000000000000000000000000000..ebde9ec4603cfb9966270a212a55757ba8a84a35 Binary files /dev/null and b/examples/animation_003_multi_ref/background_mask.png differ diff --git a/examples/animation_003_multi_ref/character_0.png b/examples/animation_003_multi_ref/character_0.png new file mode 100644 index 0000000000000000000000000000000000000000..70528ceb168f460c2324786c64d02819fc3a042a --- /dev/null +++ b/examples/animation_003_multi_ref/character_0.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:59d4fd62f23b709b2a479d8e5507528391ccd59f38430af89f5fdc2b95ea91fd +size 1353282 diff --git a/examples/animation_003_multi_ref/character_0_mask.png b/examples/animation_003_multi_ref/character_0_mask.png new file mode 100644 index 0000000000000000000000000000000000000000..e43bd54dcd8a0e3c42dcabc397bb37f1e9415182 Binary files /dev/null and b/examples/animation_003_multi_ref/character_0_mask.png differ diff --git a/examples/animation_003_multi_ref/character_1.png b/examples/animation_003_multi_ref/character_1.png new file mode 100644 index 0000000000000000000000000000000000000000..69226b3055d7b5ea6a6006dac72de59ff996420d --- /dev/null +++ b/examples/animation_003_multi_ref/character_1.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a638d034ea9e9730faaa57cd505480eb8fa13652156c0df4b39ec7d14e5b832d +size 1203064 diff --git a/examples/animation_003_multi_ref/character_1_mask.png b/examples/animation_003_multi_ref/character_1_mask.png new file mode 100644 index 0000000000000000000000000000000000000000..97c9cabc7e13f397ff81665800e7ed2c5b0a8a3c Binary files /dev/null and b/examples/animation_003_multi_ref/character_1_mask.png differ diff --git a/examples/animation_003_multi_ref/driving.mp4 b/examples/animation_003_multi_ref/driving.mp4 new file mode 100644 index 0000000000000000000000000000000000000000..dcb15800d29ca0f891771a289a3304266d288794 --- /dev/null +++ b/examples/animation_003_multi_ref/driving.mp4 @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:45d85e19aa1155e3a0e5d8134d81f01752f5a9eaa1588cfee18413ef3de1d5c1 +size 4269531 diff --git a/examples/animation_003_multi_ref/ref.png b/examples/animation_003_multi_ref/ref.png new file mode 100644 index 0000000000000000000000000000000000000000..93b7cfcd4f2e1da75d0e2e51c1f11f7c2c710158 --- /dev/null +++ b/examples/animation_003_multi_ref/ref.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e45def94e82ac70bb00c35701a4ad097a7d7089cd4c095e957a66a076439fe2b +size 8564724 diff --git a/examples/animation_003_multi_ref/ref_mask.jpg b/examples/animation_003_multi_ref/ref_mask.jpg new file mode 100644 index 0000000000000000000000000000000000000000..083f9708ec9100587f48bf446204944c1d8d5475 --- /dev/null +++ b/examples/animation_003_multi_ref/ref_mask.jpg @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6bdfeb772286d6aed94f769c1ca8c7a89c8ce5735abb389f32309c940edcd291 +size 122424 diff --git a/examples/input.txt b/examples/input.txt new file mode 100644 index 0000000000000000000000000000000000000000..91fc0c47d495a374f040d3be7504b1930ce84b76 --- /dev/null +++ b/examples/input.txt @@ -0,0 +1 @@ +the girl is dancing@@examples/001 \ No newline at end of file diff --git a/examples/prompt_examples.txt b/examples/prompt_examples.txt new file mode 100644 index 0000000000000000000000000000000000000000..be8b8cef05d25ab4d09d665290a2c6b5fc33ca36 --- /dev/null +++ b/examples/prompt_examples.txt @@ -0,0 +1,2 @@ +A middle-aged white man is working in a woodworking workshop. Positioned at the center of the frame, he wears a yellow and black checkered shirt with blue jeans and glasses, his graying hair neatly combed. He stands before a workbench cluttered with wooden boards and tools, holding a measuring tape in his left hand and a pen in his right, intently marking measurements on the wood. His expression is serious and focused, exuding professionalism and dedication. The background features various woodworking machinery and stacks of lumber, creating an atmosphere of diligence and precision. +A medical worker is washing his hands in the video. Positioned on the left side of the frame, he wears a light blue surgical gown, dark pants, and a blue surgical cap, with his back to the camera. He stands at a stainless steel handwashing station, facing the left side of the image, bending over the sink to wash his hands before straightening up and moving toward the center. The man first pulls the faucet lever with his right hand, then rubs his hands together while leaning slightly forward. Next, he presses the soap dispenser button with his right hand to lather his hands, ensuring thorough cleaning. After finishing, he moves toward the paper towel dispenser at the center of the frame, using it to dry his hands. In the background, white tiled walls are visible, along with a stainless steel soap dispenser and hand dryer mounted on the wall. A door can be seen in the distance on the right side of the hallway. The entire scene conveys a clean, professional, and hygienic atmosphere. \ No newline at end of file diff --git a/generate.py b/generate.py new file mode 100644 index 0000000000000000000000000000000000000000..9ed2bf4d5cf8bf4f3536ca1a2e1bf43285fd0b29 --- /dev/null +++ b/generate.py @@ -0,0 +1,455 @@ +# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved. +import argparse +import logging +import os +import sys +import warnings +from datetime import datetime + +warnings.filterwarnings('ignore') + +import random + +import torch +import torch.distributed as dist + +from einops import rearrange +from PIL import Image + +import wan +from wan.configs import SCAIL_CONFIGS, SCAIL_CONFIG_PATHS +from wan.utils.utils import cache_video, str2bool +from wan.utils.scail_utils import load_image_to_tensor_chw_normalized, load_video_for_pose_sample, resize_for_rectangle_crop, get_tasks_from_txt + + +def _validate_args(args): + assert args.ckpt_dir is not None, "Please specify the checkpoint directory." + if args.txt is None: + assert args.pose is not None, "Please specify the pose video." + assert args.image is not None, "Please specify the reference image." + assert str(args.model).upper() in SCAIL_CONFIGS + + args.model = str(args.model).upper() + + if args.scail_config_path is None: + args.scail_config_path = SCAIL_CONFIG_PATHS[args.model] + + if args.sample_steps is None: + args.sample_steps = 40 + + if args.sample_shift is None: + args.sample_shift = 3.0 + + if args.additional_ref_image is not None and args.additional_ref_mask_image is None: + raise ValueError("Please specify --additional_ref_mask_image when using --additional_ref_image.") + if args.additional_ref_image is None and args.additional_ref_mask_image is not None: + raise ValueError("--additional_ref_mask_image requires --additional_ref_image.") + if args.additional_ref_image is not None and len(args.additional_ref_image) != len(args.additional_ref_mask_image): + raise ValueError( + f"--additional_ref_image and --additional_ref_mask_image must have the same number of paths, " + f"got {len(args.additional_ref_image)} and {len(args.additional_ref_mask_image)}.") + + args.base_seed = args.base_seed if args.base_seed >= 0 else random.randint(0, sys.maxsize) + + +def _parse_args(): + parser = argparse.ArgumentParser() + parser.add_argument( + "--model", + type=str, + default="SCAIL-14B", + help="Type of SCAIL model. Choices: [SCAIL-14B, SCAIL-1.3B]") + parser.add_argument( + "--ckpt_dir", + type=str, + default="./SCAIL-Preview/", + help="The path to the checkpoint directory.") + parser.add_argument( + "--offload_model", + type=str2bool, + default=None, + help="Whether to offload the model to CPU after each model forward, reducing GPU memory usage." + ) + parser.add_argument( + "--ulysses_size", + type=int, + default=1, + help="The size of the ulysses parallelism in DiT.") + parser.add_argument( + "--ring_size", + type=int, + default=1, + help="The size of the ring attention parallelism in DiT.") + parser.add_argument( + "--t5_fsdp", + action="store_true", + default=False, + help="Whether to use FSDP for T5.") + parser.add_argument( + "--t5_cpu", + action="store_true", + default=False, + help="Whether to place T5 model on CPU.") + parser.add_argument( + "--dit_fsdp", + action="store_true", + default=False, + help="Whether to use FSDP for DiT.") + parser.add_argument( + "--save_dir", + type=str, + default="samples", + help="The directory to save the generated videos when --txt is not None.") + parser.add_argument( + "--save_file", + type=str, + default=None, + help="The file to save the generated video to.") + parser.add_argument( + "--prompt", + type=str, + default=None, + help="The prompt to generate the video from.") + parser.add_argument( + "--base_seed", + type=int, + default=-1, + help="The seed to use for generating the video.") + parser.add_argument( + "--txt", + type=str, + default=None, + help="Path to txt file. Default: None") + parser.add_argument( + "--image", + type=str, + default=None, + help="The reference image to generate the video from.") + parser.add_argument( + "--additional_ref_image", "--additional_image", + dest="additional_ref_image", + type=str, + nargs="+", + default=None, + help="Additional reference image paths (beta).") + parser.add_argument( + "--additional_ref_mask_image", "--additional_mask_image", + dest="additional_ref_mask_image", + type=str, + nargs="+", + default=None, + help="Mask image paths for the additional reference images (beta).") + parser.add_argument( + "--mask_image", + type=str, + default=None, + help="The mask of reference image.") + parser.add_argument( + "--pose", + type=str, + default=None, + help="The rendered pose video to generate the video from.") + parser.add_argument( + "--mask_video", + type=str, + default=None, + help="The mask of driving video.") + parser.add_argument( + "--replace_flag", + action="store_true", + default=False, + help="Pass --replace_flag to run in replacement mode. Default: False (animation mode).") + parser.add_argument( + "--target_h", + type=int, + default=512, + help="The target height of the generated video.") + parser.add_argument( + "--target_w", + type=int, + default=896, + help="The target width of the generated video.") + parser.add_argument( + "--scail_path", + type=str, + default=None, + help="Path to converted SCAIL.safetensors") + parser.add_argument( + "--scail_config_path", + type=str, + default=None, + help="Path to config.json of SCAIL") + parser.add_argument( + "--sample_solver", + type=str, + default='unipc', + choices=['unipc', 'dpm++'], + help="The solver used to sample.") + parser.add_argument( + "--sample_steps", + type=int, + default=None, + help="The sampling steps.") + parser.add_argument( + "--sample_shift", + type=float, + default=None, + help="Sampling shift factor for flow matching schedulers.") + parser.add_argument( + "--sample_guide_scale", + type=float, + default=5.0, + help="Classifier free guidance scale.") + parser.add_argument( + "--segment_len", + type=int, + default=81, + help="The number of pixel frames to sample per segment for long-video inference.") + parser.add_argument( + "--segment_overlap", + type=int, + default=5, + help="The number of pixel frames reused as clean history between adjacent segments.") + parser.add_argument( + "--lora_path", + type=str, + default=None, + help="Path to safetensors of LoRA." + ) + parser.add_argument( + "--lora_alpha", + type=float, + default=1.0, + help="Strength of LoRA. Default: 1.0" + ) + + args = parser.parse_args() + + _validate_args(args) + + return args + + +def _init_logging(rank): + # logging + if rank == 0: + # set format + logging.basicConfig( + level=logging.INFO, + format="[%(asctime)s] %(levelname)s: %(message)s", + handlers=[logging.StreamHandler(stream=sys.stdout)]) + else: + logging.basicConfig(level=logging.ERROR) + +def _check_input_path(path, name): + if path is None: + raise ValueError(f"Please specify {name}.") + if not os.path.exists(path): + raise FileNotFoundError(f"{name} does not exist: {path}") + if not os.path.isfile(path): + raise FileNotFoundError(f"{name} is not a file: {path}") + + +def generate_video(pipeline: wan.SCAIL2Pipeline, prompt: str, image_path: str, image_mask_path: str, pose_path: str, driving_mask_path: str, args, device, rank, cfg, input_idx, replace_flag, additional_task_input=None): + _check_input_path(image_path, "input image") + _check_input_path(image_mask_path, "input mask image") + _check_input_path(pose_path, "input pose video") + _check_input_path(driving_mask_path, "input mask video") + + additional_task_input = additional_task_input or {} + additional_input = {} + + logging.info(f"Input prompt: {prompt}") + logging.info(f"Input image: {image_path}") + img = Image.open(image_path).convert("RGB") + target_h = args.target_h + target_w = args.target_w + + img_uncropped = load_image_to_tensor_chw_normalized(img).to(device) # 1 c h w, -1 to 1 + _, _, h, w = img_uncropped.shape + if target_h is None or target_w is None: + target_h, target_w = h, w + if (h < w and target_h > target_w) or (h > w and target_h < target_w): + target_h, target_w = target_w, target_h + + logging.info(f"Input mask image: {image_mask_path}") + mask_img = Image.open(image_mask_path).convert("RGB") + mask_img_uncropped = load_image_to_tensor_chw_normalized(mask_img).to(device) + + if additional_task_input.get("additional_ref_image_paths", None) is not None: + additional_ref_image_paths = additional_task_input["additional_ref_image_paths"] + additional_ref_mask_image_paths = additional_task_input["additional_ref_mask_image_paths"] + additional_imgs = [] + additional_mask_imgs = [] + for idx, (additional_ref_image_path, additional_ref_mask_image_path) in enumerate( + zip(additional_ref_image_paths, additional_ref_mask_image_paths)): + _check_input_path(additional_ref_image_path, f"additional ref image {idx}") + _check_input_path(additional_ref_mask_image_path, f"additional ref mask image {idx}") + logging.info(f"Input additional reference image {idx}: {additional_ref_image_path}") + additional_img = Image.open(additional_ref_image_path).convert("RGB") + additional_img_uncropped = load_image_to_tensor_chw_normalized(additional_img).to(device) + additional_img = resize_for_rectangle_crop(additional_img_uncropped, (target_h, target_w), reshape_mode="center") + additional_imgs.append(additional_img.squeeze(0)) # c h w, -1, 1 + logging.info(f"Input additional reference mask image {idx}: {additional_ref_mask_image_path}") + additional_mask_img = Image.open(additional_ref_mask_image_path).convert("RGB") + additional_mask_img_uncropped = load_image_to_tensor_chw_normalized(additional_mask_img).to(device) + additional_mask_img = resize_for_rectangle_crop(additional_mask_img_uncropped, (target_h, target_w), reshape_mode="center") + additional_mask_imgs.append(additional_mask_img.squeeze(0)) # c h w, -1, 1 + additional_input["additional_ref_imgs"] = additional_imgs + additional_input["additional_ref_mask_imgs"] = additional_mask_imgs + + logging.info(f"Input pose video: {pose_path}") + pose_video = load_video_for_pose_sample(pose_path) # t h w c + pose_video = pose_video.permute(0, 3, 1, 2) # t c h w + pose_video = resize_for_rectangle_crop(pose_video, (target_h, target_w), reshape_mode="center") + pose_video = (pose_video - 127.5) / 127.5 # -1 1 + + logging.info(f"Input mask video: {driving_mask_path}") + driving_mask_video = load_video_for_pose_sample(driving_mask_path) # t h w c + driving_mask_video = driving_mask_video.permute(0, 3, 1, 2) # t c h w + driving_mask_video = resize_for_rectangle_crop(driving_mask_video, (target_h, target_w), reshape_mode="center") + driving_mask_video = (driving_mask_video - 127.5) / 127.5 # -1 1 + driving_mask_video = rearrange(driving_mask_video, 't c h w -> c t h w') + + img = resize_for_rectangle_crop(img_uncropped, (target_h, target_w), reshape_mode="center") + img = img.squeeze(0) # c h w, -1, 1 + + mask_img = resize_for_rectangle_crop(mask_img_uncropped, (target_h, target_w), reshape_mode="center") + mask_img = mask_img.squeeze(0) + + logging.info(f"Mode: {'Replacement' if replace_flag else 'Animation'}") + + logging.info("Generating video ...") + video = pipeline.generate( + prompt, + img, + ref_mask_img=mask_img, + pose_video=pose_video, + driving_mask_video=driving_mask_video, + replace_flag=replace_flag, + shift=args.sample_shift, + sample_solver=args.sample_solver, + segment_len=args.segment_len, + segment_overlap=args.segment_overlap, + sampling_steps=args.sample_steps, + guide_scale=args.sample_guide_scale, + seed=args.base_seed, + offload_model=args.offload_model, + **additional_input + ) + + if rank == 0: + if args.save_file is None: + formatted_time = datetime.now().strftime("%Y%m%d_%H%M%S") + formatted_prompt = args.prompt.replace(" ", "_").replace("/", + "_")[:50] + suffix = '.mp4' + args.save_file = f"SCAIL2_{args.target_w}{'x' if sys.platform=='win32' else '*'}{args.target_h}_{args.ring_size}_{formatted_prompt}_{formatted_time}" + suffix + save_file = args.save_file + if input_idx is not None: + save_dir = os.path.join(args.save_dir, f"{input_idx:07}") + os.makedirs(save_dir, exist_ok=True) + save_file = os.path.join(save_dir, args.save_file) + + logging.info(f"Saving generated video to {save_file}") + cache_video( + tensor=video[None], + save_file=save_file, + fps=cfg.sample_fps, + nrow=1, + normalize=True, + value_range=(-1, 1)) + +def generate(args): + rank = int(os.getenv("RANK", 0)) + world_size = int(os.getenv("WORLD_SIZE", 1)) + local_rank = int(os.getenv("LOCAL_RANK", 0)) + device = local_rank + _init_logging(rank) + + if args.offload_model is None: + args.offload_model = False if world_size > 1 else True + logging.info( + f"offload_model is not specified, set to {args.offload_model}.") + if world_size > 1: + torch.cuda.set_device(local_rank) + # dist.init_process_group( + # backend="nccl", + # init_method="env://", + # rank=rank, + # world_size=world_size) + else: + assert not ( + args.t5_fsdp or args.dit_fsdp + ), f"t5_fsdp and dit_fsdp are not supported in non-distributed environments." + assert not ( + args.ulysses_size > 1 or args.ring_size > 1 + ), f"context parallel are not supported in non-distributed environments." + + if args.ulysses_size > 1 or args.ring_size > 1: + assert args.ulysses_size * args.ring_size == world_size, f"The number of ulysses_size and ring_size should be equal to the world size." + from xfuser.core.distributed import ( + init_distributed_environment, + initialize_model_parallel, + ) + init_distributed_environment( + rank=dist.get_rank(), world_size=dist.get_world_size()) + + initialize_model_parallel( + sequence_parallel_degree=dist.get_world_size(), + ring_degree=args.ring_size, + ulysses_degree=args.ulysses_size, + ) + + cfg = SCAIL_CONFIGS[args.model] + if args.ulysses_size > 1: + assert cfg.num_heads % args.ulysses_size == 0, f"`{cfg.num_heads=}` cannot be divided evenly by `{args.ulysses_size=}`." + + logging.info(f"Generation job args: {args}") + + if dist.is_initialized(): + base_seed = [args.base_seed] if rank == 0 else [None] + dist.broadcast_object_list(base_seed, src=0) + args.base_seed = base_seed[0] + + if args.prompt is None: + args.prompt = "" + + additional_task_input = {} + if args.additional_ref_image is not None: + additional_task_input["additional_ref_image_paths"] = args.additional_ref_image + additional_task_input["additional_ref_mask_image_paths"] = args.additional_ref_mask_image + + if args.txt is not None: + raise NotImplementedError() + tasks = get_tasks_from_txt(args.txt) + logging.info(f"Total number of generation tasks: {len(tasks)}.") + tasks = tasks[rank::world_size] + else: + tasks = [(args.prompt, args.image, args.mask_image, args.pose, args.mask_video, None, additional_task_input)] + + logging.info("Creating SCAIL-2 pipeline.") + scail_pipeline = wan.SCAIL2Pipeline( + config=cfg, + checkpoint_dir=args.ckpt_dir, + scail_safetensors_path=args.scail_path, + scail_config_path=args.scail_config_path, + device_id=device, + rank=rank, + t5_fsdp=args.t5_fsdp, + dit_fsdp=args.dit_fsdp, + use_usp=(args.ulysses_size > 1 or args.ring_size > 1), + t5_cpu=args.t5_cpu, + lora_path=args.lora_path, + lora_alpha=args.lora_alpha, + ) + + for task in tasks: + prompt, image_path, image_mask_path, pose_path, driving_mask_path, input_idx, additional_task_input = task + generate_video(scail_pipeline, prompt, image_path, image_mask_path, pose_path, driving_mask_path, args, device, rank, cfg, input_idx, args.replace_flag, additional_task_input) + + logging.info("Finished.") + +if __name__ == "__main__": + args = _parse_args() + generate(args) diff --git a/prompt_enhancer.py b/prompt_enhancer.py new file mode 100644 index 0000000000000000000000000000000000000000..0b77e79131cfb153263be7ddc14dfda09add7ef6 --- /dev/null +++ b/prompt_enhancer.py @@ -0,0 +1,296 @@ +#!/usr/bin/env python3 +import argparse +import mimetypes +import tempfile +import os +from pathlib import Path + + +VIDEO_CAPTION_PROMPT = """You are captioning sampled frames from a source video for a character replacement video generation task. + +Describe the source video in one detailed English paragraph. Focus on: +- the scene, location, lighting, camera framing, and background; +- the action, motion, timing, and camera movement across the sampled frames; +- the clothing, pose, body motion, and nearby objects touched or interacted with by the person/character being replaced. + +If the user specifies who should be replaced, identify that source subject clearly in the caption. Pay special attention to the source subject's clothing and any objects they hold, touch, operate, sit on, stand near, or otherwise interact with, because those details help locate the replacement region. + +Do not mention the replacement target image. Do not invent an identity for the replacement target. +Output only the source-video caption. +""" + + +REPLACEMENT_PROMPT_TEMPLATE = """You are a prompt enhancer for SCAIL-2 character replacement. + +Your task is to write one detailed English description of the final replaced video. This is not an editing instruction. The output must describe the video after replacement has already happened: the replacement character from the reference image is performing the source subject's motion in the source scene. + +Replacement instruction from user: +{instruction} + +Source video caption: +{caption} + +Few-shot examples of the desired prompt style: +{examples} + +Rules: +1. Output a positive video-generation prompt describing the replaced video itself. Do not output wording like "replace X with Y", "swap", "edit", or "the task is". +2. Remove the original source subject's identity and appearance. Keep only the original subject's motion, pose, timing, spatial position, and interaction with the scene. +3. The final prompt for SCAIL-2 should describe the replacement character's visible clothing and appearance in enough detail, using the reference image as the source of identity and wardrobe details. +4. The final prompt should also describe important objects the character interacts with or stays close to in the source video, such as tools, instruments, furniture, vehicles, doors, tables, handheld items, or work surfaces. +5. Keep the original video environment, lighting, camera angle, shot scale, background objects, and motion trajectory. +6. If the source caption mentions the original subject's clothing only to locate body regions or interactions, translate those grounding details into the replacement character's final appearance instead of preserving the original identity. +7. Use natural video wording with concrete verbs. Avoid mentioning masks, segmentation, editing software, Gemini, or the prompt generation process. +8. Output only the final enhanced prompt, in one English paragraph, around 90-140 words. +""" + + +def _check_file(path: str, name: str) -> Path: + if path is None: + raise ValueError(f"Please specify {name}.") + file_path = Path(path) + if not file_path.exists(): + raise FileNotFoundError(f"{name} does not exist: {file_path}") + if not file_path.is_file(): + raise FileNotFoundError(f"{name} is not a file: {file_path}") + return file_path + + +def _read_examples(path: str | None, max_chars: int) -> str: + if path is None: + return "" + file_path = _check_file(path, "prompt examples") + text = file_path.read_text(encoding="utf-8").strip() + return text[:max_chars] + + +def _guess_mime(path: Path, fallback: str) -> str: + mime_type, _ = mimetypes.guess_type(str(path)) + return mime_type or fallback + + +def extract_video_frames(video_path: Path, num_frames: int, image_format: str = "jpg") -> list[bytes]: + import cv2 + + cap = cv2.VideoCapture(str(video_path)) + if not cap.isOpened(): + raise RuntimeError(f"Failed to open source video: {video_path}") + + frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT)) + if frame_count <= 0: + cap.release() + raise RuntimeError(f"Could not determine frame count for source video: {video_path}") + + num_frames = max(1, min(num_frames, frame_count)) + if num_frames == 1: + indices = [frame_count // 2] + else: + indices = [round(i * (frame_count - 1) / (num_frames - 1)) for i in range(num_frames)] + + ext = ".jpg" if image_format.lower() in ("jpg", "jpeg") else ".png" + encode_params = [int(cv2.IMWRITE_JPEG_QUALITY), 92] if ext == ".jpg" else [] + frames = [] + for idx in indices: + cap.set(cv2.CAP_PROP_POS_FRAMES, idx) + ok, frame = cap.read() + if not ok: + continue + ok, encoded = cv2.imencode(ext, frame, encode_params) + if ok: + frames.append(encoded.tobytes()) + cap.release() + + if not frames: + raise RuntimeError(f"Failed to extract frames from source video: {video_path}") + return frames + + +def _make_client(api_key: str | None): + try: + from google import genai + import google.genai.types as gtypes + except ImportError as exc: + raise ImportError( + "Missing Gemini SDK. Install it with: pip install google-genai" + ) from exc + + api_key = api_key or os.environ.get("GEMINI_API_KEY") or os.environ.get("GOOGLE_API_KEY") + if not api_key: + raise ValueError("Set GEMINI_API_KEY or pass --api_key.") + base_url = os.environ.get("GEMINI_BASE_URL") + http_opts = gtypes.HttpOptions(baseUrl=base_url) if base_url else None + return genai.Client(api_key=api_key, http_options=http_opts) + + +def _generate_text(client, model: str, contents, temperature: float) -> str: + from google.genai import types + + response = client.models.generate_content( + model=model, + contents=contents, + config=types.GenerateContentConfig(temperature=temperature), + ) + text = getattr(response, "text", None) + if not text: + raise RuntimeError(f"Gemini returned an empty response: {response}") + return text.strip() + + +def caption_video( + client, + model: str, + video_path: Path, + instruction: str, + temperature: float, + num_frames: int, +) -> str: + import google.genai.types as gtypes + + frame_bytes = extract_video_frames(video_path, num_frames=num_frames) + prompt = ( + f"{VIDEO_CAPTION_PROMPT}\n\n" + f"User replacement instruction: {instruction}\n" + f"The following images are {len(frame_bytes)} sampled frames in chronological order." + ) + parts = [ + gtypes.Part( + inline_data=gtypes.Blob(data=data, mime_type="image/jpeg") + ) + for data in frame_bytes + ] + parts.append(gtypes.Part(text=prompt)) + contents = gtypes.Content(parts=parts) + return _generate_text(client, model, contents, temperature) + + +def enhance_prompt( + client, + model: str, + image_path: Path, + instruction: str, + caption: str, + examples: str, + temperature: float, +) -> str: + import google.genai.types as gtypes + + image_bytes = image_path.read_bytes() + prompt = REPLACEMENT_PROMPT_TEMPLATE.format( + instruction=instruction.strip(), + caption=caption.strip(), + examples=examples.strip() or "(No examples provided.)", + ) + contents = gtypes.Content( + parts=[ + gtypes.Part( + inline_data=gtypes.Blob( + data=image_bytes, + mime_type=_guess_mime(image_path, "image/jpeg"), + ) + ), + gtypes.Part(text=prompt), + ] + ) + return _generate_text(client, model, contents, temperature) + + +def parse_args(): + parser = argparse.ArgumentParser( + description="Enhance a character replacement prompt with Gemini." + ) + parser.add_argument( + "--video", + required=True, + help="Source/original video for Gemini to caption.", + ) + parser.add_argument( + "--image", + required=True, + help="Reference image of the replacement character.", + ) + parser.add_argument( + "--instruction", + required=True, + help='User replacement instruction, e.g. "replace the man with the person in the image".', + ) + parser.add_argument( + "--examples", + default="prompt_examples.txt", + help="Few-shot prompt examples. Default: prompt_examples.txt", + ) + parser.add_argument( + "--model", + default="gemini-3-flash-preview", + help="Gemini model name. Default: gemini-3-flash-preview", + ) + parser.add_argument( + "--api_key", + default=None, + help="Gemini API key. Defaults to GEMINI_API_KEY or GOOGLE_API_KEY.", + ) + parser.add_argument( + "--temperature", + type=float, + default=0.4, + help="Gemini sampling temperature.", + ) + parser.add_argument( + "--max_example_chars", + type=int, + default=4000, + help="Maximum characters loaded from the few-shot examples file.", + ) + parser.add_argument( + "--num_frames", + type=int, + default=8, + help="Number of source-video frames sampled for Gemini captioning.", + ) + parser.add_argument( + "--caption_out", + default=None, + help="Optional path to save the intermediate source video caption.", + ) + parser.add_argument( + "--output", + default=None, + help="Optional path to save the enhanced prompt.", + ) + return parser.parse_args() + + +def main(): + args = parse_args() + video_path = _check_file(args.video, "source video") + image_path = _check_file(args.image, "replacement image") + examples = _read_examples(args.examples, args.max_example_chars) + + client = _make_client(args.api_key) + + caption = caption_video( + client=client, + model=args.model, + video_path=video_path, + instruction=args.instruction, + temperature=args.temperature, + num_frames=args.num_frames, + ) + if args.caption_out: + Path(args.caption_out).write_text(caption + "\n", encoding="utf-8") + + enhanced_prompt = enhance_prompt( + client=client, + model=args.model, + image_path=image_path, + instruction=args.instruction, + caption=caption, + examples=examples, + temperature=args.temperature, + ) + if args.output: + Path(args.output).write_text(enhanced_prompt + "\n", encoding="utf-8") + print(enhanced_prompt) + + +if __name__ == "__main__": + main() diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000000000000000000000000000000000000..d416e7ba5f627bb46ceec3a770345cd1198250e1 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,16 @@ +torch>=2.4.0 +torchvision>=0.19.0 +opencv-python>=4.9.0.80 +diffusers>=0.31.0 +transformers>=4.49.0 +tokenizers>=0.20.3 +accelerate>=1.1.1 +tqdm +imageio +easydict +ftfy +dashscope +imageio-ffmpeg +flash_attn +gradio>=5.0.0 +numpy>=1.23.5,<2