diff --git a/.gitattributes b/.gitattributes index a6344aac8c09253b3b630fb776ae94478aa0275b..8af773f90dfc08bd6279a864836e9aacfaf319ed 100644 --- a/.gitattributes +++ b/.gitattributes @@ -33,3 +33,9 @@ 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 +assets/figure_edit_results.png filter=lfs diff=lfs merge=lfs -text +assets/logo.png filter=lfs diff=lfs merge=lfs -text +assets/table_editscore_qwen3_vl.png filter=lfs diff=lfs merge=lfs -text +assets/table_reward_model_results.png filter=lfs diff=lfs merge=lfs -text +example_images/input.png filter=lfs diff=lfs merge=lfs -text +example_images/output.png filter=lfs diff=lfs merge=lfs -text diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000000000000000000000000000000000000..8e9bed2007c6bfd4cd97f3b25ea3eea712b579a0 --- /dev/null +++ b/.gitignore @@ -0,0 +1,239 @@ +# Created by https://www.toptal.com/developers/gitignore/api/macos,python +# Edit at https://www.toptal.com/developers/gitignore?templates=macos,python + +### macOS ### +# General +.DS_Store +.AppleDouble +.LSOverride + +# Icon must end with two \r +Icon + + +# Thumbnails +._* + +# Files that might appear in the root of a volume +.DocumentRevisions-V100 +.fseventsd +.Spotlight-V100 +.TemporaryItems +.Trashes +.VolumeIcon.icns +.com.apple.timemachine.donotpresent + +# Directories potentially created on remote AFP share +.AppleDB +.AppleDesktop +Network Trash Folder +Temporary Items +.apdisk + +### macOS Patch ### +# iCloud generated files +*.icloud + +### Python ### +# Byte-compiled / optimized / DLL files +__pycache__/ +*.py[cod] +*$py.class + +# C extensions +*.so + +# Distribution / packaging +.Python +build/ +develop-eggs/ +dist/ +downloads/ +eggs/ +.eggs/ +lib/ +lib64/ +parts/ +sdist/ +var/ +wheels/ +share/python-wheels/ +*.egg-info/ +.installed.cfg +*.egg +MANIFEST + +# PyInstaller +# Usually these files are written by a python script from a template +# before PyInstaller builds the exe, so as to inject date/other infos into it. +*.manifest +*.spec + +# Installer logs +pip-log.txt +pip-delete-this-directory.txt + +# Unit test / coverage reports +htmlcov/ +.tox/ +.nox/ +.coverage +.coverage.* +.cache +nosetests.xml +coverage.xml +*.cover +*.py,cover +.hypothesis/ +.pytest_cache/ +cover/ + +# Translations +*.mo +*.pot + +# Django stuff: +*.log +local_settings.py +db.sqlite3 +db.sqlite3-journal + +# Flask stuff: +instance/ +.webassets-cache + +# Scrapy stuff: +.scrapy + +# Sphinx documentation +docs/_build/ + +# PyBuilder +.pybuilder/ +target/ + +# Jupyter Notebook +.ipynb_checkpoints + +# IPython +profile_default/ +ipython_config.py + +# pyenv +# For a library or package, you might want to ignore these files since the code is +# intended to run in multiple environments; otherwise, check them in: +# .python-version + +# pipenv +# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control. +# However, in case of collaboration, if having platform-specific dependencies or dependencies +# having no cross-platform support, pipenv may install dependencies that don't work, or not +# install all needed dependencies. +#Pipfile.lock + +# poetry +# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control. +# This is especially recommended for binary packages to ensure reproducibility, and is more +# commonly ignored for libraries. +# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control +#poetry.lock + +# pdm +# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control. +#pdm.lock +# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it +# in version control. +# https://pdm.fming.dev/#use-with-ide +.pdm.toml + +# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm +__pypackages__/ + +# Celery stuff +celerybeat-schedule +celerybeat.pid + +# SageMath parsed files +*.sage.py + +# Environments +.env +.venv +env/ +venv/ +ENV/ +env.bak/ +venv.bak/ + +# Spyder project settings +.spyderproject +.spyproject + +# Rope project settings +.ropeproject + +# mkdocs documentation +/site + +# mypy +.mypy_cache/ +.dmypy.json +dmypy.json + +# Pyre type checker +.pyre/ + +# pytype static type analyzer +.pytype/ + +# Cython debug symbols +cython_debug/ + +# PyCharm +# JetBrains specific template is maintained in a separate JetBrains.gitignore that can +# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore +# and can be added to the global gitignore or merged into this file. For a more nuclear +# option (not recommended) you can uncomment the following to ignore the entire idea folder. +#.idea/ + +### Python Patch ### +# Poetry local configuration file - https://python-poetry.org/docs/configuration/#local-configuration +poetry.toml + +# ruff +.ruff_cache/ + +# LSP config files +pyrightconfig.json + +# End of https://www.toptal.com/developers/gitignore/api/macos,python + +local_scripts/ + +omnigen2/utils/vpn_utils.py + +test_tokenizer.py +save_pipeline.py +app.sh +logs/ +results/ +test_jsonl* +pbs_files/ +convert_ckpt_to_pipeline.py +inference_test_efficiency.py +upload_pipeline* +example_images_resized/ +example_t2i_test_efficiency*.sh +example_edit_test_efficiency*.sh +example_in_context_generation_test_efficiency*.sh +intro* +resize_example_images.py +save_pipeline.py +outputs_gradio/* +test.py + +upload_to_huggingface.py +upload_to_modelscope.py + +editscore.egg-info/ +dist/ \ No newline at end of file diff --git a/FORK_PROVENANCE.md b/FORK_PROVENANCE.md new file mode 100644 index 0000000000000000000000000000000000000000..cf78fec05fb4a78987f189a51ea54567ac60229d --- /dev/null +++ b/FORK_PROVENANCE.md @@ -0,0 +1,11 @@ +# Fork provenance + +- **Upstream:** `VectorSpaceLab/EditScore (github)` +- **Upstream commit / SHA:** `4609c5d2ebb62fdebf665d3c924686d896ef1f74` +- **License:** Apache-2.0 +- **Kind:** Code mirror +- **Workspace consumer:** EditScore-7B (training + inference code) +- **Forked to:** `chibifire/EditScore-code` +- **Forked on:** 2026-09-05 +- **Author:** Ernest Lee +- **Reason:** shipping-surface completeness under our own HF org so a base/source we depend on cannot disappear or relicense out from under us diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..f49a4e16e68b128803cc2dcea614603632b04eac --- /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 [yyyy] [name of copyright owner] + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. \ No newline at end of file diff --git a/README.md b/README.md new file mode 100644 index 0000000000000000000000000000000000000000..e3c6f0a2434014af0242793c6f5bca67a2824397 --- /dev/null +++ b/README.md @@ -0,0 +1,215 @@ +

+ +

+ +

+ project page + arxiv + model + dataset + dataset + dataset +

+ +

+

+ News | + Quick Start | + Benchmark Usage | + Citation +

+

+ +**EditScore** is a series of state-of-the-art open-source reward models (7Bโ€“72B) designed to evaluate and enhance instruction-guided image editing. +## โœจ Highlights +- **State-of-the-Art Performance**: Effectively matches the performance of leading proprietary VLMs. With a self-ensembling strategy, **our largest model surpasses even GPT-5** on our comprehensive benchmark, **EditReward-Bench**. +- **A Reliable Evaluation Standard**: We introduce **EditReward-Bench**, the first public benchmark specifically designed for evaluating reward models in image editing, featuring 13 subtasks, 11 state-of-the-art editing models (*including proprietary models*) and expert human annotations. +- **Simple and Easy-to-Use**: Get an accurate quality score for your image edits with just a few lines of code. +- **Versatile Applications**: Ready to use as a best-in-class reranker to improve editing outputs, or as a high-fidelity reward signal for **stable and effective Reinforcement Learning (RL) fine-tuning**. + +## ๐Ÿ”ฅ News +- **2026-02-01**: Our work has been accepted to ICLR 2026 ๐ŸŽ‰ +- **2025-11-21**: We're excited to release the training configs for EditScore reward models! ๐ŸŽฏ Built on LLaMA-Factory, we trained EditScore ranging from 4B to 72B parameters with simple YAML configurations. Check out the [EditScore-train guide](examples/EditScore-train/README.md) to get started with training your own reward models. +- **2025-10-31**: Weโ€™re thrilled to announce the **Qwen3-VL** variants of **EditScore**! ๐Ÿš€ Powered by Qwen3-VL, the new 4B and 8B models achieve outstanding efficiency and performance. Impressively, the 4B model already matches the performance of the original 32B version, while the 8B model delivers results comparable to the original 72B model. The models are now available on [huggingface](https://huggingface.co/EditScore/models), see [Usage Example](#-usage-example) for how to use. Detailed comparisons with Qwen2.5-VL variants are in the [performance table](https://raw.githubusercontent.com/VectorSpaceLab/EditScore/refs/heads/main/assets/table_editscore_qwen3_vl.png). +- **2025-10-27**: Released [OmniGen2-EditScore7B-v1.1](https://huggingface.co/OmniGen2/OmniGen2-EditScore7B-v1.1), achieving a **7.01 (+0.73) GEdit score** within **700 steps**, by incorporating the **reweighting strategy** from [TempFlow](https://arxiv.org/abs/2508.04324). Additionally, the **JSON repair process** has been enhanced using [json_repair](https://github.com/mangiucugna/json_repair), improving **EditScoreโ€™s stability** under various conditions. Upgrade via `pip install -U editscore`. +- **2025-10-22**: **Introducing Our Reinforcement Learning Training Framework!** + We're excited to release our complete RL pipeline, the result of a massive effort to simplify fine-tuning for image editing models. Key features include: + - **Ready-to-Use RL Dataset**: Includes the complete dataset used in the EditScore project, along with clear usage guidelines and preparation scripts. + - **An Easy-to-Use Reward Model**: Seamlessly integrate **EditScore** as a reward signal. + - **A Scalable Reward Server**: Built with native multi-node support for high-throughput training. + - **Flexible Training Code**: Supports distributed training, variable image resolutions and mixed tasks (t2i, edit, in-context generation) out-of-the-box. + Dive into our comprehensive guide on [RL Fine-Tuning](examples/OmniGen2-RL#application-2-reinforcement-fine-tuning) to get started. + +- 2025-10-16: Training datasets [EditScore-Reward-Data](https://huggingface.co/datasets/EditScore/EditScore-Reward-Data) and [EditScore-RL-Data](https://huggingface.co/datasets/EditScore/EditScore-RL-Data) are available. +- 2025-10-15: **EditScore** is now available on PyPI โ€” install it easily with `pip install editscore`. +- 2025-10-15: Best-of-N inference scripts for OmniGen2, Flux-dev-Kontext, and Qwen-Image-Edit are now available! See [this](#apply-editscore-to-image-editing) for details. +- 2025-09-30: We release **OmniGen2-EditScore7B**, unlocking online RL For Image Editing via high-fidelity EditScore. LoRA weights are available at [Hugging Face](https://huggingface.co/OmniGen2/OmniGen2-EditScore7B) and [ModelScope](https://www.modelscope.cn/models/OmniGen2/OmniGen2-EditScore7B). +- 2025-09-30: We are excited to release **EditScore** and **EditReward-Bench**! Model weights and the benchmark dataset are now publicly available. You can access them on Hugging Face: [Models Collection](https://huggingface.co/collections/EditScore/editscore-68d8e27ee676981221db3cfe) and [Benchmark Dataset](https://huggingface.co/datasets/EditScore/EditReward-Bench), and on ModelScope: [Models Collection](https://www.modelscope.cn/collections/EditScore-8b0d53aa945d4e) and [Benchmark Dataset](https://www.modelscope.cn/datasets/EditScore/EditReward-Bench). + +## ๐Ÿ“– Introduction +While Reinforcement Learning (RL) holds immense potential for this domain, its progress has been severely hindered by the absence of a high-fidelity, efficient reward signal. + +To overcome this barrier, we provide a systematic, two-part solution: + +- **A Rigorous Evaluation Standard**: We first introduce **EditReward-Bench**, a new public benchmark for the direct and reliable evaluation of reward models. It features 13 diverse subtasks and expert human annotations, establishing a gold standard for measuring reward signal quality. + +- **A Powerful & Versatile Tool**: Guided by our benchmark, we developed the **EditScore** model series. Through meticulous data curation and an effective self-ensembling strategy, EditScore sets a new state of the art for open-source reward models, even surpassing the accuracy of leading proprietary VLMs. + +

+ +
+ Benchmark results on EditReward-Bench. +

+ +We demonstrate the practical utility of EditScore through two key applications: + +- **As a State-of-the-Art Reranker**: Use EditScore to perform Best-of-*N* selection and instantly improve the output quality of diverse editing models. +- **As a High-Fidelity Reward for RL**: Use EditScore as a robust reward signal to fine-tune models via RL, enabling stable training and unlocking significant performance gains where general-purpose VLMs fail. + +This repository releases both the **EditScore** models and the **EditReward-Bench** dataset to facilitate future research in reward modeling, policy optimization, and AI-driven model improvement. + +

+ +
+ EditScore as a superior reward signal for image editing. +

+ + +## ๐Ÿ“Œ TODO +We are actively working on improving EditScore and expanding its capabilities. Here's what's next: + +- [x] Qwen3-VL variants of EditScore. +- [x] Release training data for reward model and online RL. +- [x] Release RL training code applying EditScore to OmniGen2. +- [x] Provide Best-of-N inference scripts for OmniGen2, Flux-dev-Kontext, and Qwen-Image-Edit. + +## ๐Ÿš€ Quick Start + +### ๐Ÿ› ๏ธ Environment Setup +We offer two ways to install EditScore. Choose the one that best fits your needs. +**Method 1: Install from PyPI (Recommended for Users)**: If you want to use EditScore as a library in your own project. +**Method 2: Install from Source (For Developers)**: If you plan to contribute to the code, modify it, or run the examples in this repository + +#### Prerequisites: Installing PyTorch +Both installation methods require PyTorch to be installed first, as its version is dependent on your system's CUDA setup. +```bash +# (Optional) Create a clean Python environment +conda create -n editscore python=3.12 +conda activate editscore + +# Choose the command that matches your CUDA version. +# This example is for CUDA 12.6. +pip install torch==2.7.1 torchvision --extra-index-url https://download.pytorch.org/whl/cu126 +```` + +
+๐ŸŒ For users in Mainland China +```bash +# Install PyTorch from a domestic mirror +pip install torch==2.7.1 torchvision --index-url https://mirror.sjtu.edu.cn/pytorch-wheels/cu126 +``` +
+ +#### Method 1: Install from PyPI (Recommended for Users) +```bash +pip install -U editscore +``` + +#### Method 2: Install from Source (For Developers) +This method gives you a local, editable version of the project. +1. Clone the repository +```bash +git clone https://github.com/VectorSpaceLab/EditScore.git +cd EditScore +``` + +2. Install EditScore in editable mode +```bash +pip install -e . +``` + +#### โœ… (Recommended) Install Optional High-Performance Dependencies +For the best performance, especially during inference, we highly recommend installing vllm. +```bash +pip install -U vllm +``` + +--- + +### ๐Ÿงช Usage Example +Using EditScore is straightforward. The model will be automatically downloaded from the Hugging Face Hub on its first run. +```python +from PIL import Image +from editscore import EditScore + +# Load the EditScore model. It will be downloaded automatically. +# Replace with the specific model version you want to use. +model_path = "Qwen/Qwen3-VL-4B-Instruct" +lora_path = "EditScore/EditScore-Qwen3-VL-4B-Instruct" + +scorer = EditScore( + backbone="qwen3vl", # set to "qwen3vl_vllm" for faster inference + model_name_or_path=model_path, + lora_path=lora_path, + score_range=25, + num_pass=1, # Increase for better performance via self-ensembling +) + +# Below is Qwen2.5-VL version + +# model_path = "Qwen/Qwen2.5-VL-7B-Instruct" +# lora_path = "EditScore/EditScore-7B" + +# scorer = EditScore( +# backbone="qwen25vl", # set to "qwen25vl_vllm" for faster inference +# model_name_or_path=model_path, +# lora_path=lora_path, +# score_range=25, +# num_pass=1, # Increase for better performance via self-ensembling +# ) + +input_image = Image.open("example_images/input.png") +output_image = Image.open("example_images/output.png") +instruction = "Adjust the background to a glass wall." + +result = scorer.evaluate([input_image, output_image], instruction) +print(f"Edit Score: {result['overall']}") +# Expected output: A dictionary containing the final score and other details. +``` + +--- + +## ๐Ÿ“Š Benchmark Your Image-Editing Reward Model +#### Install benchmark dependencies +To use example code for benchmark, run following +```bash +pip install -r requirements.txt +``` + +We provide an evaluation script to benchmark reward models on **EditReward-Bench**. To evaluate your own custom reward model, simply create a scorer class with a similar interface and update the script. +```bash +# This script will evaluate the default EditScore model on the benchmark +bash evaluate.sh + +# Or speed up inference with VLLM +bash evaluate_vllm.sh +``` + +## Apply EditScore to Image Editing +We offer two example use cases for your exploration: +- **Best-of-N selection**: Use EditScore to automatically pick the most preferred image among multiple candidates. +- **Reinforcement fine-tuning**: Use EditScore as a reward model to guide RL-based optimization. + +For detailed instructions and examples, please refer to the [documentation](examples/OmniGen2-RL/README.md). + +## โค๏ธ Citing Us +If you find this repository or our work useful, please consider giving a star โญ and citation ๐Ÿฆ–, which would be greatly appreciated: + +```bibtex +@article{luo2025editscore, + title={EditScore: Unlocking Online RL for Image Editing via High-Fidelity Reward Modeling}, + author={Xin Luo and Jiahao Wang and Chenyuan Wu and Shitao Xiao and Xiyan Jiang and Defu Lian and Jiajun Zhang and Dong Liu and Zheng Liu}, + journal={arXiv preprint arXiv:2509.23909}, + year={2025} +} +``` diff --git a/assets/figure_edit_results.png b/assets/figure_edit_results.png new file mode 100644 index 0000000000000000000000000000000000000000..7cb2519d83906ddb03e9a2c3e388e56b6b2504bc --- /dev/null +++ b/assets/figure_edit_results.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a1f8ff4ad98c12fb3b7762e2787118ad6abd29fe3a22ac6ffacd4c6904f7788b +size 268064 diff --git a/assets/logo.png b/assets/logo.png new file mode 100644 index 0000000000000000000000000000000000000000..69f1dae18f73804f94308d550cd82e43ac6e4dec --- /dev/null +++ b/assets/logo.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:170b16d70efc7058027b8c019779202950a48c500b191956d0fd829639adf5e0 +size 195304 diff --git a/assets/table_editscore_qwen3_vl.png b/assets/table_editscore_qwen3_vl.png new file mode 100644 index 0000000000000000000000000000000000000000..84e5d80994d5b05873d6e19bffcbc348d5821e38 --- /dev/null +++ b/assets/table_editscore_qwen3_vl.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8210bbca111e39756c22400e75fbd2360456b1e84584a95b9d35de9e911a84c1 +size 300312 diff --git a/assets/table_reward_model_results.png b/assets/table_reward_model_results.png new file mode 100644 index 0000000000000000000000000000000000000000..065ce0bb4f02010265bcdecf17c50a969b1ad804 --- /dev/null +++ b/assets/table_reward_model_results.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6ce31391f25947c3863e2cc57cc7475a704b5b12754274d1f4549d603b0f98d0 +size 343358 diff --git a/calculate_statistics.py b/calculate_statistics.py new file mode 100644 index 0000000000000000000000000000000000000000..8e7213457dec71681696fd14dcd7defe96b5224a --- /dev/null +++ b/calculate_statistics.py @@ -0,0 +1,124 @@ +import os +import glob +import json +import numpy as np + +import argparse + +PROMPT_FOLLOWING = "prompt_following" +CONSISTENCY = "consistency" +OVERALL = "overall" +SCORE_CATEGORIES = [PROMPT_FOLLOWING, CONSISTENCY, OVERALL] + +def parse_args(): + parser = argparse.ArgumentParser() + parser.add_argument("--result_dir", type=str, required=True) + parser.add_argument("--backbone", type=str, default="qwen25vl", choices=["qwen25vl", "openai", "internvl3_5"]) + return parser.parse_args() + +def main(args): + task_types = sorted(os.listdir(args.result_dir)) + + print(task_types) + + prompt_following_results = dict() + consistency_results = dict() + overall_results = dict() + + all_prompt_following_scores = [] + all_consistency_scores = [] + all_overall_scores = [] + + for task_type in task_types: + task_type_dir = os.path.join(args.result_dir, task_type) + prompt_following_json_file = os.path.join(task_type_dir, f"{PROMPT_FOLLOWING}.jsonl") + consistency_json_file = os.path.join(task_type_dir, f"{CONSISTENCY}.jsonl") + overall_json_file = os.path.join(task_type_dir, f"{OVERALL}.jsonl") + + total_num = 0 + correct_num = 0 + with open(prompt_following_json_file, 'r') as f: + for line in f: + json_line = json.loads(line) + if json_line['score'][0] > json_line['score'][1]: + correct_num += 1 + total_num += 1 + all_prompt_following_scores.append(json_line['score'][0]) + all_prompt_following_scores.append(json_line['score'][1]) + prompt_following_results[task_type] = correct_num / total_num + + total_num = 0 + correct_num = 0 + with open(consistency_json_file, 'r') as f: + for line in f: + json_line = json.loads(line) + if json_line['score'][0] > json_line['score'][1]: + correct_num += 1 + total_num += 1 + all_consistency_scores.append(json_line['score'][0]) + all_consistency_scores.append(json_line['score'][1]) + consistency_results[task_type] = correct_num / total_num + + total_num = 0 + correct_num = 0 + with open(overall_json_file, 'r') as f: + for line in f: + json_line = json.loads(line) + if json_line['score'][0] > json_line['score'][1]: + correct_num += 1 + total_num += 1 + all_overall_scores.append(json_line['score'][0]) + all_overall_scores.append(json_line['score'][1]) + overall_results[task_type] = correct_num / total_num + + prompt_following_results['average'] = sum(prompt_following_results.values()) / len(prompt_following_results) + consistency_results['average'] = sum(consistency_results.values()) / len(consistency_results) + overall_results['average'] = sum(overall_results.values()) / len(overall_results) + + print(overall_results.keys()) + + task_types = [ + 'background_change', 'color_alter', 'style_change', 'subject-add', 'subject-remove', 'subject-replace', 'material_alter', + 'motion_change', 'ps_human', 'text_change', 'tone_transfer', 'extract', 'compose', 'average' + ] + + print(" & ".join(task_types)) + print("Prompt Following: " + " & ".join([f"{prompt_following_results[task_type]:.3f}" for task_type in task_types])) + print("Consistency: " + " & ".join([f"{consistency_results[task_type]:.3f}" for task_type in task_types])) + print("Overall: " + " & ".join([f"{overall_results[task_type]:.3f}" for task_type in task_types])) + + groups = { + 'object': ['subject-add', 'subject-remove', 'subject-replace'], + 'appearance': ['color_alter', 'material_alter', 'style_change', 'tone_transfer'], + 'scene': ['background_change', 'extract'], + 'advanced': ['ps_human', 'text_change', 'motion_change', 'compose'], + } + + print("--------------------------------") + print("--------------------------------") + + for group_name, group_task_types in groups.items(): + print(group_name + ":") + print("Prompt Following & Consistency & Overall") + prompt_following_mean = np.mean([prompt_following_results[task_type] for task_type in group_task_types]) + consistency_mean = np.mean([consistency_results[task_type] for task_type in group_task_types]) + overall_mean = np.mean([overall_results[task_type] for task_type in group_task_types]) + print(f"{prompt_following_mean:.3f} & {consistency_mean:.3f} & {overall_mean:.3f}") + + print("Average:") + print("Prompt Following & Consistency & Overall") + print(f"{prompt_following_results['average']:.3f} & {consistency_results['average']:.3f} & {overall_results['average']:.3f}") + + print("Prompt Following Scores:") + print("Min & Max & Mean & Std") + print(f"{np.min(all_prompt_following_scores):.3f} & {np.max(all_prompt_following_scores):.3f} & {np.mean(all_prompt_following_scores):.3f} & {np.std(all_prompt_following_scores):.3f}") + print("Consistency Scores:") + print("Min & Max & Mean & Std") + print(f"{np.min(all_consistency_scores):.3f} & {np.max(all_consistency_scores):.3f} & {np.mean(all_consistency_scores):.3f} & {np.std(all_consistency_scores):.3f}") + print("Overall Scores:") + print("Min & Max & Mean & Std") + print(f"{np.min(all_overall_scores):.3f} & {np.max(all_overall_scores):.3f} & {np.mean(all_overall_scores):.3f} & {np.std(all_overall_scores):.3f}") + +if __name__ == "__main__": + args = parse_args() + main(args) \ No newline at end of file diff --git a/editscore/__init__.py b/editscore/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..f1082e2e4d7f4e3879dbaac3f8caf6c23a0437ff --- /dev/null +++ b/editscore/__init__.py @@ -0,0 +1,217 @@ +import sys +sys.path.insert(0, 'editscore') + +from typing import Optional +from .utils import ( + mllm_output_to_dict +) +import math +from . import vie_prompts +import numpy as np +from .json_parser import parse_vlm_output_to_dict + +class EditScore: + def __init__( + self, + backbone="gpt-4.1", + openai_url="https://api.openai.com/v1/chat/completions", + key=None, + model_name_or_path="", + score_range: int=25, + temperature: float=0.7, + tensor_parallel_size: int=1, + max_model_len: int=1536, + max_num_batched_tokens: int=1536, + max_num_seqs: int=32, + num_pass: int=1, + reduction: str="average_last", + seed: int=42, + lora_path: Optional[str]=None, + cache_dir: Optional[str]=None, + ) -> None: + self.backbone = backbone + self.score_range = score_range + self.reduction = reduction + self.seed = seed + self.num_pass = num_pass + + if self.backbone == 'openai': + from .mllm_tools.openai import GPT4o + self.model = GPT4o(key, model_name=model_name_or_path, url=openai_url) + elif self.backbone == "qwen25vl": + from .mllm_tools.qwen25vl import Qwen25VL + self.model = Qwen25VL( + vlm_model=model_name_or_path, + temperature=temperature, + seed=seed, + lora_path=lora_path, + ) + elif self.backbone == "qwen25vl_vllm": + from .mllm_tools.qwen25vl_vllm import Qwen25VL + self.model = Qwen25VL( + vlm_model=model_name_or_path, + tensor_parallel_size=tensor_parallel_size, + max_model_len=max_model_len, + max_num_seqs=max_num_seqs, + max_num_batched_tokens=max_num_batched_tokens, + temperature=temperature, + seed=seed, + lora_path=lora_path, + cache_dir=cache_dir, + ) + elif self.backbone == "qwen3vl": + from .mllm_tools.qwen3vl import Qwen3VL + self.model = Qwen3VL( + vlm_model=model_name_or_path, + temperature=temperature, + seed=seed, + lora_path=lora_path, + ) + elif self.backbone == "qwen3vl_vllm": + from .mllm_tools.qwen3vl_vllm import Qwen3VL + self.model = Qwen3VL( + vlm_model=model_name_or_path, + tensor_parallel_size=tensor_parallel_size, + max_model_len=max_model_len, + max_num_seqs=max_num_seqs, + max_num_batched_tokens=max_num_batched_tokens, + temperature=temperature, + seed=seed, + lora_path=lora_path, + cache_dir=cache_dir, + ) + elif self.backbone == "internvl3_5": + from .mllm_tools.internvl35_lmdeploy import InternVL35 + self.model = InternVL35(model=model_name_or_path, tensor_parallel_size=tensor_parallel_size) + + self.context = vie_prompts._context_no_delimit_reasoning_first + + self.SC_prompt = "\n".join([self.context, vie_prompts._prompts_0shot_two_image_edit_rule, vie_prompts._prompts_0shot_tie_rule_SC.replace('10', str(self.score_range))]) + self.PQ_prompt = "\n".join([self.context, vie_prompts._prompts_0shot_rule_PQ.replace('10', str(self.score_range))]) + + def evaluate(self, image_prompts, text_prompt): + if not isinstance(image_prompts, list): + image_prompts = [image_prompts] + + if self.backbone in ['openai']: + self.model.use_encode = False if isinstance(image_prompts[0], str) else True + + _SC_prompt = self.SC_prompt.replace("", text_prompt) + + SC_prompt_final = self.model.prepare_input(image_prompts, _SC_prompt) + PQ_prompt_final = self.model.prepare_input(image_prompts[-1], self.PQ_prompt) # assume the last image is the edited image + + outputs_multi_pass = [] + + for i in range(self.num_pass): + SC_dict = False + PQ_dict = False + tries = 0 + max_tries = 2 + while SC_dict is False or PQ_dict is False: + tries += 1 + give_up_parsing = True if tries > max_tries else False + + result_SC = self.model.inference(SC_prompt_final, seed=self.seed + i) + result_PQ = self.model.inference(PQ_prompt_final, seed=self.seed + i) + + if result_SC in ["I'm sorry, but I can't assist with that request."] or result_PQ in ["I'm sorry, but I can't assist with that request."]: + give_up_parsing = True + + SC_dict = mllm_output_to_dict(result_SC, give_up_parsing=give_up_parsing, text_prompt=text_prompt, score_range=self.score_range) + PQ_dict = mllm_output_to_dict(result_PQ, give_up_parsing=give_up_parsing, text_prompt=text_prompt, score_range=self.score_range) + + if SC_dict == "rate_limit_exceeded" or PQ_dict == "rate_limit_exceeded": + print("rate_limit_exceeded") + raise ValueError("rate_limit_exceeded") + + try: + SC_score = min(SC_dict['score']) / (self.score_range / 10) + PQ_score = min(PQ_dict['score']) / (self.score_range / 10) + O_score = math.sqrt(SC_score * PQ_score) + except Exception as e: + print(f"{e=} {SC_dict['score']=} {PQ_dict['score']=}") + raise e + + try: + outputs_multi_pass.append({ + 'prompt_following': SC_dict['score'][0] / (self.score_range / 10), + 'consistency': SC_dict['score'][1] / (self.score_range / 10), + 'perceptual_quality': PQ_score, + 'overall': O_score, + }) + except Exception as e: + print(f"{e=} {SC_dict['score']=} {PQ_dict['score']=}") + raise e + + output = { + "prompt_following": np.mean([output_per_pass["prompt_following"] for output_per_pass in outputs_multi_pass]), + "consistency": np.mean([output_per_pass["consistency"] for output_per_pass in outputs_multi_pass]), + "perceptual_quality": np.mean([output_per_pass["perceptual_quality"] for output_per_pass in outputs_multi_pass]), + "overall": np.mean([output_per_pass["overall"] for output_per_pass in outputs_multi_pass]), + "SC_reasoning": SC_dict["reasoning"], + "PQ_reasoning": PQ_dict["reasoning"], + } + if self.reduction == "average_first": + output["overall"] = math.sqrt(output["prompt_following"] * output["perceptual_quality"]) + return output + + + def batch_evaluate(self, image_prompts, text_prompt): + SC_prompt = [self.SC_prompt.replace("", _text_prompt) for _text_prompt in text_prompt] + + SC_prompt = [self.model.prepare_input(image_prompt, _SC_prompt) for image_prompt, _SC_prompt in zip(image_prompts, SC_prompt)] + PQ_prompt = [self.model.prepare_input(image_prompt, self.PQ_prompt) for image_prompt in image_prompts] + + outputs_multi_pass = [[] for _ in range(len(image_prompts))] + for i in range(self.num_pass): + results = self.model.batch_inference(SC_prompt + PQ_prompt, seed=self.seed + i) + + SC_evaluations = [parse_vlm_output_to_dict(results[i]) for i in range(len(results) // 2)] + PQ_evaluations = [parse_vlm_output_to_dict(results[i]) for i in range(len(results) // 2, len(results))] + + for idx, (SC_evaluation, PQ_evaluation) in enumerate(zip(SC_evaluations, PQ_evaluations)): + SC_scores = SC_evaluation["score"] + PQ_scores = PQ_evaluation["score"] + + if len(SC_scores) == 0: + SC_scores = [self.score_range / 2] + if len(PQ_scores) == 0: + PQ_scores = [self.score_range / 2] + + SC_score = min(SC_scores) / (self.score_range / 10) + PQ_score = min(PQ_scores) / (self.score_range / 10) + if SC_score < 0 or SC_score > 10: + SC_score = self.score_range / 2 + if PQ_score < 0 or PQ_score > 10: + PQ_score = self.score_range / 2 + O_score = math.sqrt(SC_score * PQ_score) + + outputs_multi_pass[idx].append( + { + "SC_score": SC_score, + "PQ_score": PQ_score, + "O_score": O_score, + "SC_score_reasoning": SC_evaluation["reasoning"], + "PQ_score_reasoning": PQ_evaluation["reasoning"], + "SC_raw_output": results[idx], + "PQ_raw_output": results[len(results) // 2 + idx], + } + ) + + outputs = [] + for idx, outputs_per_prompt in enumerate(outputs_multi_pass): + outputs.append( + { + "SC_score": np.mean([output_per_pass["SC_score"] for output_per_pass in outputs_per_prompt]), + "PQ_score": np.mean([output_per_pass["PQ_score"] for output_per_pass in outputs_per_prompt]), + "O_score": np.mean([output_per_pass["O_score"] for output_per_pass in outputs_per_prompt]), + "SC_score_reasoning": outputs_per_prompt[0]["SC_score_reasoning"], + "PQ_score_reasoning": outputs_per_prompt[0]["PQ_score_reasoning"], + "SC_raw_output": outputs_per_prompt[0]["SC_raw_output"], + "PQ_raw_output": outputs_per_prompt[0]["PQ_raw_output"], + } + ) + if self.reduction == "average_first": + outputs[-1]["O_score"] = math.sqrt(outputs[-1]["SC_score"] * outputs[-1]["PQ_score"]) + return outputs \ No newline at end of file diff --git a/editscore/json_parser.py b/editscore/json_parser.py new file mode 100644 index 0000000000000000000000000000000000000000..1cf0488f3512bfce87556b5020d609b83c5209c6 --- /dev/null +++ b/editscore/json_parser.py @@ -0,0 +1,173 @@ +import json +import re +from typing import Dict, Any, List, Optional + +# ============================================================================== +# HELPER FUNCTIONS (Based on your provided robust fixers) +# ============================================================================== +# For clarity, these are named as internal functions (prefixed with an underscore). + +def _fix_json_quotes(s: str) -> str: + """First-stage repair: handle incorrect quotes and basic structure.""" + # Replace Python-style booleans/None with JSON standard + s = re.sub(r'\bTrue\b', 'true', s) + s = re.sub(r'\bFalse\b', 'false', s) + s = re.sub(r'\bNone\b', 'null', s) + + # Attempt to replace single quotes with double quotes (a common VLM error) + # This is a high-risk operation that might break the reasoning content, + # but it's worth trying early on. + try: + temp_s = s.replace("'", '"') + json.loads(temp_s) + return temp_s + except json.JSONDecodeError: + # If it's still invalid after replacement, return the original string for the next repair step. + pass + + # Add double quotes to keys (e.g., {reasoning: ...} -> {"reasoning": ...}) + s = re.sub(r'([\{\s,])(\w+)\s*:', r'\1"\2":', s) + return s + +def _repair_reasoning_field_robust(json_str: str) -> str: + """Second-stage repair: specifically fix unescaped double quotes inside the 'reasoning' field.""" + pattern = re.compile( + r'("reasoning"\s*:\s*")' # --- Group 1: "reasoning" : " + r'(.*?)' # --- Group 2: The content (non-greedy) + r'(?="\s*[,}])', # --- Lookahead: find " followed by , or } + re.DOTALL + ) + + def replacer(match): + prefix = match.group(1) + content = match.group(2) + # In the content, replace all unescaped " with \" + fixed_content = content.replace('"', '\\"') + return prefix + fixed_content + + return pattern.sub(replacer, json_str) + +def _fallback_extract_and_rebuild(input_str: str) -> str: + """Final fallback strategy: abandon repair, directly extract information, and rebuild a valid JSON.""" + # 1. Extract reasoning + # Find all content between "reasoning": and ,"score": + reasoning_text = "" + reason_match = re.search(r'["\']reasoning["\']\s*:\s*["\']?(.*?)["\']?\s*,\s*["\']score["\']', input_str, re.DOTALL | re.IGNORECASE) + if reason_match: + reasoning_text = reason_match.group(1).strip() + # Clean up any potentially remaining escape characters + reasoning_text = reasoning_text.replace('\\"', '"') + else: + # If not found, assume all text besides the score part is the reasoning. + # First, remove the score part. + score_part_match = re.search(r'["\']score["\']\s*:.*', input_str, re.IGNORECASE) + if score_part_match: + reasoning_text = input_str[:score_part_match.start()].strip() + else: + # If even 'score' cannot be found, assume the entire string is the reasoning. + reasoning_text = input_str + + # 2. Extract scores + scores = [] + # Prioritize searching after the 'score' keyword. + score_match = re.search(r'["\']score["\']\s*:\s*(.*)', input_str, re.DOTALL | re.IGNORECASE) + search_area = score_match.group(1) if score_match else input_str + + # Find all integers or floats. + numbers = re.findall(r'[-+]?\d*\.?\d+', search_area) + if numbers: + scores = [float(num) for num in numbers] + + # 3. Rebuild into a standard dictionary and return a JSON string. + rebuilt_data = { + "reasoning": reasoning_text, + "score": scores + } + return json.dumps(rebuilt_data, ensure_ascii=False) + + +def _format_and_validate_dict(data: Dict[str, Any]) -> Optional[Dict[str, Any]]: + """Validate and format the parsed dictionary to ensure it meets the final output standard.""" + if not isinstance(data, dict): + return None + + # Extract reasoning, tolerating case and spelling variations. + reasoning = "" + for key in ["reasoning", "reason", "rationale"]: + if key in data and isinstance(data[key], str): + reasoning = data[key] + break + + # Extract score, and ensure it is a list of floats. + scores = [] + if 'score' in data: + score_val = data['score'] + if isinstance(score_val, list): + scores = [float(s) for s in score_val if isinstance(s, (int, float, str))] + elif isinstance(score_val, (int, float)): + scores = [float(score_val)] + + # If any field was found, consider it a success. + if reasoning or scores: + return {"score": scores, "reasoning": reasoning} + + return None + +# ============================================================================== +# MAIN PARSING FUNCTION +# ============================================================================== + +def parse_vlm_output_to_dict(input_string: str) -> Dict[str, Any]: + """ + A highly robust function to parse a VLM's output string into a dictionary + containing 'score' and 'reasoning'. + + It uses a multi-stage repair pipeline, progressively degrading from standard + JSON parsing to a final information extraction fallback. + """ + # --- 0. Preprocessing --- + if not input_string or not input_string.strip(): + return {"score": [], "reasoning": "Input was empty."} + + # Find the substring enclosed by `{}`, which is often the core of the VLM output. + json_match = re.search(r'\{.*\}', input_string, re.DOTALL) + target_str = json_match.group(0) if json_match else input_string.strip() + + # --- Repair Pipeline --- + # Apply fixers in order, attempting to parse after each one. + + fixer_pipeline = [ + lambda s: s, # 1. Try the original string. + _fix_json_quotes, # 2. Fix basic quotes and keywords. + _repair_reasoning_field_robust, # 3. Fix internal quotes in the reasoning field. + ] + + for fixer in fixer_pipeline: + try: + fixed_str = fixer(target_str) + data = json.loads(fixed_str) + validated_data = _format_and_validate_dict(data) + if validated_data is not None: + return validated_data + except (json.JSONDecodeError, TypeError): + # If it fails, continue to the next fixer. + continue + + # --- Final Fallback Strategy --- + # If all repair and parsing attempts fail, activate the information extraction mode. + try: + fallback_str = _fallback_extract_and_rebuild(target_str) + # This function guarantees a valid JSON string, so we can load it directly. + data = json.loads(fallback_str) + # Still run it through the validator to standardize the format. + validated_data = _format_and_validate_dict(data) + if validated_data: + return validated_data + except Exception: + # If even the final fallback strategy fails, return an error message. + pass + + return { + "score": [], + "reasoning": f"Failed to parse after all strategies. Original output: '{input_string}'" + } diff --git a/editscore/mllm_tools/__init__.py b/editscore/mllm_tools/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/editscore/mllm_tools/internvl35_lmdeploy.py b/editscore/mllm_tools/internvl35_lmdeploy.py new file mode 100644 index 0000000000000000000000000000000000000000..5a9535243184addc456477a152310876601422f0 --- /dev/null +++ b/editscore/mllm_tools/internvl35_lmdeploy.py @@ -0,0 +1,73 @@ +from typing import List + +import random +# import magic +# import megfile + +import numpy as np +import torch + +from lmdeploy import pipeline, PytorchEngineConfig +from lmdeploy.vl import load_image +from lmdeploy.vl.constants import IMAGE_TOKEN + + +def set_seed(seed: int): + """ + Args: + Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`. + seed (`int`): The seed to set. + """ + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + + +def apply_chat_template(prompt, num_images: int = 2): + """ + This is used since the bug of transformers which do not support vision id https://github.com/QwenLM/Qwen2.5-VL/issues/716#issuecomment-2723316100 + """ + template = "\n".join([f"Image-{i}: {IMAGE_TOKEN}" for i in range(1, num_images + 1)]) + template += f"\n{prompt}" + return template + +class InternVL35(): + def __init__(self, model, max_model_len: int = 16384, tensor_parallel_size=1, max_num_seqs=32) -> None: + # attn_implementation = "flash_attention_2" if is_flash_attn_2_available() else None + self.model = pipeline(model, backend_config=PytorchEngineConfig(session_len=max_model_len, tp=tensor_parallel_size)) + + def prepare_input(self, images: List = [], text_prompt: str = ""): + if not isinstance(images, list): + images = [images] + messages = (apply_chat_template(text_prompt, num_images=len(images)), images) + return messages + + def inference(self, messages): + set_seed(42) + # Prepare the inputs + + response = self.model(messages) + print(f"{response.text=}", flush=True) + return response.text + +if __name__ == "__main__": + model = InternVL35( + vlm_model="OpenGVLab/InternVL3_5-8B", + max_model_len=16384, + tensor_parallel_size=1, + max_num_seqs=32 + ) + + from PIL import Image + prompt = model.prepare_input( + [Image.open("https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-VL/assets/demo.jpeg")], + 'Describe the image in detail.' + ) + + prompt2 = model.prepare_input( + [Image.open("https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-VL/assets/demo.jpeg")], + 'How well it looks? Give a score between 0 and 100.' + ) + res = model.inference([prompt, prompt2]) + print("result : \n", res) \ No newline at end of file diff --git a/editscore/mllm_tools/openai.py b/editscore/mllm_tools/openai.py new file mode 100644 index 0000000000000000000000000000000000000000..055904c40141090f886dd0d88398ef14a3218790 --- /dev/null +++ b/editscore/mllm_tools/openai.py @@ -0,0 +1,180 @@ +import base64 +import requests +from io import BytesIO, StringIO +from typing import Union, Optional, Tuple, List +from PIL import Image, ImageOps +import os + +def get_api_key(file_path): + # Read the API key from the first line of the file + with open(file_path, 'r') as file: + return file.readline().strip() + +# Function to encode the image +def encode_image(image_path): + with open(image_path, "rb") as image_file: + return base64.b64encode(image_file.read()).decode('utf-8') + +def pick_next_item(current_item, item_list): + if current_item not in item_list: + raise ValueError("Current item is not in the list") + current_index = item_list.index(current_item) + next_index = (current_index + 1) % len(item_list) + + return item_list[next_index] + +# Function to encode a PIL image +def encode_pil_image(pil_image): + # Create an in-memory binary stream + image_stream = BytesIO() + + # Save the PIL image to the binary stream in JPEG format (you can change the format if needed) + pil_image.save(image_stream, format='JPEG') + + # Get the binary data from the stream and encode it as base64 + image_data = image_stream.getvalue() + base64_image = base64.b64encode(image_data).decode('utf-8') + + return base64_image + + +def load_image(image: Union[str, Image.Image], format: str = "RGB", size: Optional[Tuple] = None) -> Image.Image: + """ + Load an image from a given path or URL and convert it to a PIL Image. + + Args: + image (Union[str, Image.Image]): The image path, URL, or a PIL Image object to be loaded. + format (str, optional): Desired color format of the resulting image. Defaults to "RGB". + size (Optional[Tuple], optional): Desired size for resizing the image. Defaults to None. + + Returns: + Image.Image: A PIL Image in the specified format and size. + + Raises: + ValueError: If the provided image format is not recognized. + """ + if isinstance(image, str): + if image.startswith("http://") or image.startswith("https://"): + image = Image.open(requests.get(image, stream=True).raw) + elif os.path.isfile(image): + image = Image.open(image) + else: + raise ValueError( + f"Incorrect path or url, URLs must start with `http://` or `https://`, and {image} is not a valid path" + ) + elif isinstance(image, Image.Image): + image = image + else: + raise ValueError( + "Incorrect format used for image. Should be an url linking to an image, a local path, or a PIL image." + ) + image = ImageOps.exif_transpose(image) + image = image.convert(format) + if (size != None): + image = image.resize(size, Image.LANCZOS) + return image + +class GPT4v(): + def __init__(self, key, url="https://api.openai.com/v1/chat/completions", are_images_encoded=False, model_name="gpt-4-vision-preview"): + """OpenAI GPT-4-vision model wrapper + Args: + api_key_path (str): Path to the API key file. Defaults to 'keys/secret.env'. + are_images_encoded (bool): Whether the images are encoded in base64. Defaults to False. + """ + self.multiple_api_keys = False + self.current_key_file = None + self.api_key = key + + self.url = url + self.model_name = model_name + self.use_encode = are_images_encoded + + def prepare_input(self, image_links: List = [], text_prompt: str = ""): + prompt_content = [] + text_dict = { + "type": "text", + "text": text_prompt + } + prompt_content.append(text_dict) + + if not isinstance(image_links, list): + image_links = [image_links] + + for image_link in image_links: + image = load_image(image_link) + if self.use_encode: + visual_dict = { + "type": "image_url", + "image_url": {"url": f"data:image/jpeg;base64,{encode_pil_image(image)}"} + } + else: + visual_dict = { + "type": "image_url", + "image_url": {"url": image_link} + } + prompt_content.append(visual_dict) + return prompt_content + + def inference(self, prompt, seed: Optional[int] = None): + payload = { + "model": self.model_name, + "messages": [ + { + "role": "user", + "content": prompt + } + ], + # "max_tokens": 1400 + } + headers = { + "Content-Type": "application/json", + "Authorization": f"Bearer {self.api_key}" + } + # try: + response = requests.post(self.url, json=payload, headers=headers, timeout=180) # Set timeout to 5 minutes (300 seconds) + # except Exception as e: + # print(f"Error: {e}") + # return "" + #return response.text + return self.extract_response(response) + + def extract_response(self, response): + try: + response = response.json() + out = response['choices'][0]['message']['content'] + return out + except Exception as e: + print(f"Error: {e}") + if response['error']['code'] == 'content_policy_violation': + print("Code is content_policy_violation") + elif response['error']['code'] in ['rate_limit_exceeded', 'insufficient_quota', 'insufficient_user_quota']: + print(f"Code is {response['error']['code']}", flush=True) + print(response['error']['message'], flush=True) + return "rate_limit_exceeded" + if self.multiple_api_keys == True: + new_key = pick_next_item(self.current_key_file, self.key_lists) + self.update_key(new_key) + self.current_key_file = new_key #override key + print("New key is from the file: ", new_key) + else: + print("Code is different") + print(response) + print(f"{response['error']['code']=}") + return "" + + def update_key(self, key, load_from_file=True): + if load_from_file: + self.api_key = get_api_key(key) + else: + self.api_key = key + +class GPT4o(GPT4v): + def __init__(self, key, url="https://api.openai.com/v1/chat/completions", are_images_encoded=False, model_name="gpt-4o-2024-05-13"): + super().__init__(key, url, are_images_encoded, model_name) + +if __name__ == "__main__": + model = GPT4o('sk-cB6h7HcCSDIp71gs6lFLZxKE0dOYOnJbxzES6kWXe1Wb2VHS', model_name="gpt-4.1") + prompt = model.prepare_prompt(['https://chromaica.github.io/Museum/ImagenHub_Text-Guided_IE/DiffEdit/sample_34_1.jpg', 'https://chromaica.github.io/Museum/ImagenHub_Text-Guided_IE/input/sample_34_1.jpg'], 'What is difference between two images?') + print("prompt : \n", prompt) + res = model.get_parsed_output(prompt) + print("result : \n", res) \ No newline at end of file diff --git a/editscore/mllm_tools/qwen25vl.py b/editscore/mllm_tools/qwen25vl.py new file mode 100644 index 0000000000000000000000000000000000000000..80859f893e7988eacf8ac2f30cbfcbd4e7c8b0bf --- /dev/null +++ b/editscore/mllm_tools/qwen25vl.py @@ -0,0 +1,97 @@ +from typing import Optional +import random +import numpy as np +import torch + +from transformers import Qwen2_5_VLForConditionalGeneration, AutoProcessor +from qwen_vl_utils import process_vision_info +from peft import PeftModel + + +def set_seed(seed: int): + """ + Args: + Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`. + seed (`int`): The seed to set. + """ + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + + +def apply_chat_template(prompt, num_images: int = 2): + """ + This is used since the bug of transformers which do not support vision id https://github.com/QwenLM/Qwen2.5-VL/issues/716#issuecomment-2723316100 + """ + template = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n" + template += "".join([f": <|vision_start|><|image_pad|><|vision_end|>" for i in range(1, num_images + 1)]) + template += f"{prompt}<|im_end|>\n<|im_start|>assistant\n" + return template + + +class Qwen25VL(): + def __init__( + self, + vlm_model, + temperature: float = 0.7, + seed: Optional[int] = None, + lora_path: Optional[str] = None, + ) -> None: + self.model = Qwen2_5_VLForConditionalGeneration.from_pretrained( + vlm_model, torch_dtype=torch.bfloat16, device_map="auto" + ) + if lora_path: + self.model = PeftModel.from_pretrained(self.model, lora_path) + self.model = self.model.merge_and_unload() + + self.processor = AutoProcessor.from_pretrained(vlm_model) + self.temperature = temperature + self.seed = seed + + def prepare_input(self, images, text_prompt: str = ""): + if not isinstance(images, list): + images = [images] + + messages = [ + { + "role": "user", + "content": [{"type": "image", "image": image} for image in images] + + [{"type": "text", "text": text_prompt}], + } + ] + text = apply_chat_template(text_prompt, num_images=len(images)) + image_inputs, video_inputs = process_vision_info(messages) + + inputs = self.processor( + text=[text], + images=image_inputs, + videos=video_inputs, + padding=True, + return_tensors="pt", + ) + inputs = inputs.to("cuda") + + return inputs + + def inference(self, inputs, seed: Optional[int] = None): + seed = self.seed if seed is None else seed + + set_seed(seed) + generated_ids = self.model.generate( + **inputs, + max_new_tokens=512, + do_sample=True, + temperature=self.temperature, + top_p=0.9, + top_k=20, + ) + generated_ids_trimmed = [ + out_ids[len(in_ids) :] for in_ids, out_ids in zip(inputs.input_ids, generated_ids) + ] + outputs = self.processor.batch_decode( + generated_ids_trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False + ) + + outputs = [output.strip() for output in outputs] + return outputs[0] \ No newline at end of file diff --git a/editscore/mllm_tools/qwen25vl_vllm.py b/editscore/mllm_tools/qwen25vl_vllm.py new file mode 100644 index 0000000000000000000000000000000000000000..76d0dfa089523b59e34f449945289db2de28a9c7 --- /dev/null +++ b/editscore/mllm_tools/qwen25vl_vllm.py @@ -0,0 +1,138 @@ +from typing import Optional + +import os +import hashlib +import random +import time +import numpy as np +import torch + +from vllm import LLM +from vllm.sampling_params import SamplingParams + +from transformers import Qwen2_5_VLForConditionalGeneration, AutoProcessor +from peft import PeftModel + +from qwen_vl_utils import process_vision_info + + +def set_seed(seed: int): + """ + Args: + Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`. + seed (`int`): The seed to set. + """ + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + + +def apply_chat_template(prompt, num_images: int = 2): + """ + This is used since the bug of transformers which do not support vision id https://github.com/QwenLM/Qwen2.5-VL/issues/716#issuecomment-2723316100 + """ + template = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n" + template += "".join([f": <|vision_start|><|image_pad|><|vision_end|>" for i in range(1, num_images + 1)]) + template += f"{prompt}<|im_end|>\n<|im_start|>assistant\n" + return template + + +class Qwen25VL(): + def __init__( + self, + vlm_model, + max_model_len: int = 1536, + tensor_parallel_size=1, + max_num_seqs=32, + max_num_batched_tokens=1536, + temperature: float = 0.7, + seed: Optional[int] = None, + lora_path: Optional[str] = None, + cache_dir: Optional[str] = None, + ) -> None: + if lora_path: + if cache_dir is None: + root_dir = torch.hub.get_dir() # default: ~/.cache/torch/hub + + lora_filename = os.path.splitext(os.path.basename(lora_path))[0] + lora_hash = hashlib.md5(lora_path.encode()).hexdigest()[:8] + lora_identifier = f"{lora_filename}_{lora_hash}" + + cache_dir = os.path.join(root_dir, "EditScore", f"{os.path.basename(vlm_model)}_merged_lora_{lora_identifier}") + + if not os.path.exists(cache_dir): + print(f"Merging LORA to {vlm_model} and saving to {cache_dir}", flush=True) + start_time = time.time() + model = Qwen2_5_VLForConditionalGeneration.from_pretrained( + vlm_model, torch_dtype=torch.bfloat16, device_map="cpu" + ) + model = PeftModel.from_pretrained(model, lora_path) + model = model.merge_and_unload() + model.save_pretrained(cache_dir) + + processor = AutoProcessor.from_pretrained(vlm_model) + processor.save_pretrained(cache_dir) + + print(f"Merging LORA to {vlm_model} and saving to {cache_dir} took {time.time() - start_time} seconds", flush=True) + else: + print(f"Skipping merging LORA, as merged model already exists in {cache_dir}", flush=True) + + vlm_model = cache_dir + + self.model = LLM( + model=vlm_model, + max_model_len=max_model_len, + tensor_parallel_size=tensor_parallel_size, + max_num_seqs=max_num_seqs, + max_num_batched_tokens=max_num_batched_tokens, + limit_mm_per_prompt={"image": 2}, + enable_prefix_caching=True, + ) + self.temperature = temperature + self.seed = seed + + def prepare_input(self, images, text_prompt: str = ""): + if not isinstance(images, list): + images = [images] + + messages = [ + { + "role": "user", + "content": [{"type": "image", "image": image} for image in images] + + [{"type": "text", "text": text_prompt}], + } + ] + text = apply_chat_template(text_prompt, num_images=len(images)) + image_inputs, _ = process_vision_info(messages) + + messages = { + "prompt": text, + "multi_modal_data": {"image": image_inputs}, + } + return messages + + def inference(self, messages, seed: Optional[int] = None): + seed = self.seed if seed is None else seed + sampling_params = SamplingParams(max_tokens=512, temperature=self.temperature, top_p=0.9, top_k=20, seed=seed) + outputs = self.model.generate(messages, sampling_params, use_tqdm=False) + + responses = [] + for output in outputs: + instruction = output.outputs[0].text.strip() + responses.append(instruction) + + return responses[0] + + + def batch_inference(self, messages, seed: Optional[int] = None): + seed = self.seed if seed is None else seed + sampling_params = SamplingParams(max_tokens=512, temperature=self.temperature, top_p=0.9, top_k=20, seed=seed) + outputs = self.model.generate(messages, sampling_params, use_tqdm=False) + + responses = [] + for output in outputs: + instruction = output.outputs[0].text.strip() + responses.append(instruction) + + return responses \ No newline at end of file diff --git a/editscore/mllm_tools/qwen3vl.py b/editscore/mllm_tools/qwen3vl.py new file mode 100644 index 0000000000000000000000000000000000000000..0fb330f65297a256d5ac3820c49d914988768e92 --- /dev/null +++ b/editscore/mllm_tools/qwen3vl.py @@ -0,0 +1,95 @@ +from typing import Optional +import random +import numpy as np +import torch + +from transformers import Qwen3VLForConditionalGeneration, AutoProcessor +from peft import PeftModel + + +def set_seed(seed: int): + """ + Args: + Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`. + seed (`int`): The seed to set. + """ + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + + +def apply_chat_template(prompt, num_images: int = 2): + """ + This is used since the bug of transformers which do not support vision id https://github.com/QwenLM/Qwen2.5-VL/issues/716#issuecomment-2723316100 + """ + template = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n" + template += "".join([f": <|vision_start|><|image_pad|><|vision_end|>" for i in range(1, num_images + 1)]) + template += f"{prompt}<|im_end|>\n<|im_start|>assistant\n" + return template + + +class Qwen3VL(): + def __init__( + self, + vlm_model, + temperature: float = 0.7, + seed: Optional[int] = None, + lora_path: Optional[str] = None, + ) -> None: + self.model = Qwen3VLForConditionalGeneration.from_pretrained( + vlm_model, torch_dtype=torch.bfloat16, device_map="auto" + ) + if lora_path: + self.model = PeftModel.from_pretrained(self.model, lora_path) + self.model = self.model.merge_and_unload() + + self.processor = AutoProcessor.from_pretrained(vlm_model) + self.temperature = temperature + self.seed = seed + + def prepare_input(self, images, text_prompt: str = ""): + if not isinstance(images, list): + images = [images] + + messages = [ + { + "role": "user", + "content": [{"type": "image", "image": image} for image in images] + + [{"type": "text", "text": text_prompt}], + } + ] + + inputs = self.processor.apply_chat_template( + messages, + tokenize=True, + add_generation_prompt=True, + return_dict=True, + return_tensors="pt" + ) + + inputs = inputs.to("cuda") + + return inputs + + def inference(self, inputs, seed: Optional[int] = None): + seed = self.seed if seed is None else seed + + set_seed(seed) + generated_ids = self.model.generate( + **inputs, + max_new_tokens=512, + do_sample=True, + temperature=self.temperature, + top_p=0.9, + top_k=20, + ) + generated_ids_trimmed = [ + out_ids[len(in_ids) :] for in_ids, out_ids in zip(inputs.input_ids, generated_ids) + ] + outputs = self.processor.batch_decode( + generated_ids_trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False + ) + + outputs = [output.strip() for output in outputs] + return outputs[0] \ No newline at end of file diff --git a/editscore/mllm_tools/qwen3vl_vllm.py b/editscore/mllm_tools/qwen3vl_vllm.py new file mode 100644 index 0000000000000000000000000000000000000000..2a4cdfa7db41bbb6180aca6013b4ff6219c281f7 --- /dev/null +++ b/editscore/mllm_tools/qwen3vl_vllm.py @@ -0,0 +1,141 @@ +from typing import Optional + +import os +import hashlib +import random +import time +import numpy as np +import torch + +from vllm import LLM +from vllm.sampling_params import SamplingParams + +from transformers import Qwen3VLForConditionalGeneration, AutoProcessor +from peft import PeftModel + +from qwen_vl_utils import process_vision_info + + +def set_seed(seed: int): + """ + Args: + Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`. + seed (`int`): The seed to set. + """ + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + + +def apply_chat_template(prompt, num_images: int = 2): + """ + This is used since the bug of transformers which do not support vision id https://github.com/QwenLM/Qwen2.5-VL/issues/716#issuecomment-2723316100 + """ + template = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n" + template += "".join([f": <|vision_start|><|image_pad|><|vision_end|>" for i in range(1, num_images + 1)]) + template += f"{prompt}<|im_end|>\n<|im_start|>assistant\n" + return template + + +class Qwen3VL(): + def __init__( + self, + vlm_model, + max_model_len: int = 1536, + tensor_parallel_size=1, + max_num_seqs=32, + max_num_batched_tokens=1536, + temperature: float = 0.7, + seed: Optional[int] = None, + lora_path: Optional[str] = None, + cache_dir: Optional[str] = None, + ) -> None: + if lora_path: + if cache_dir is None: + root_dir = torch.hub.get_dir() # default: ~/.cache/torch/hub + + lora_filename = os.path.splitext(os.path.basename(lora_path))[0] + lora_hash = hashlib.md5(lora_path.encode()).hexdigest()[:8] + lora_identifier = f"{lora_filename}_{lora_hash}" + + cache_dir = os.path.join(root_dir, "EditScore", f"{os.path.basename(vlm_model)}_merged_lora_{lora_identifier}") + + if not os.path.exists(cache_dir): + print(f"Merging LORA to {vlm_model} and saving to {cache_dir}", flush=True) + start_time = time.time() + model = Qwen3VLForConditionalGeneration.from_pretrained( + vlm_model, torch_dtype=torch.bfloat16, device_map="cpu" + ) + model = PeftModel.from_pretrained(model, lora_path) + model = model.merge_and_unload() + model.save_pretrained(cache_dir) + + processor = AutoProcessor.from_pretrained(vlm_model) + processor.save_pretrained(cache_dir) + + print(f"Merging LORA to {vlm_model} and saving to {cache_dir} took {time.time() - start_time} seconds", flush=True) + else: + print(f"Skipping merging LORA, as merged model already exists in {cache_dir}", flush=True) + + vlm_model = cache_dir + + self.model = LLM( + model=vlm_model, + max_model_len=max_model_len, + tensor_parallel_size=tensor_parallel_size, + max_num_seqs=max_num_seqs, + max_num_batched_tokens=max_num_batched_tokens, + limit_mm_per_prompt={"image": 2}, + enable_prefix_caching=True, + ) + + self.processor = AutoProcessor.from_pretrained(vlm_model) + self.temperature = temperature + self.seed = seed + + def prepare_input(self, images, text_prompt: str = ""): + if not isinstance(images, list): + images = [images] + + messages = [ + { + "role": "user", + "content": [{"type": "image", "image": image} for image in images] + + [{"type": "text", "text": text_prompt}], + } + ] + # text = apply_chat_template(text_prompt, num_images=len(images)) + text = self.processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) + image_inputs, _ = process_vision_info(messages) + + messages = { + "prompt": text, + "multi_modal_data": {"image": image_inputs}, + } + return messages + + def inference(self, messages, seed: Optional[int] = None): + seed = self.seed if seed is None else seed + sampling_params = SamplingParams(max_tokens=512, temperature=self.temperature, top_p=0.9, top_k=20, seed=seed) + outputs = self.model.generate(messages, sampling_params, use_tqdm=False) + + responses = [] + for output in outputs: + instruction = output.outputs[0].text.strip() + responses.append(instruction) + + return responses[0] + + + def batch_inference(self, messages, seed: Optional[int] = None): + seed = self.seed if seed is None else seed + sampling_params = SamplingParams(max_tokens=512, temperature=self.temperature, top_p=0.9, top_k=20, seed=seed) + outputs = self.model.generate(messages, sampling_params, use_tqdm=False) + + responses = [] + for output in outputs: + instruction = output.outputs[0].text.strip() + responses.append(instruction) + + return responses \ No newline at end of file diff --git a/editscore/mllm_tools/utils.py b/editscore/mllm_tools/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..3a85f3ae0c39245c5aa87a60e9bfcd0d361e5982 --- /dev/null +++ b/editscore/mllm_tools/utils.py @@ -0,0 +1,65 @@ +from typing import List +import base64 +from io import BytesIO +from PIL import Image +import requests + +def pil_image_to_base64(pil_image, format="PNG"): + buffered = BytesIO() + pil_image.save(buffered, format=format) # Save image to the buffer in the specified format + img_str = base64.b64encode(buffered.getvalue()).decode('utf-8') # Encode the buffer's content to base64 + return img_str + +def load_image(image_file): + if image_file.startswith("http"): + response = requests.get(image_file) + image = Image.open(BytesIO(response.content)).convert("RGB") + else: + import os + image = Image.open(image_file).convert("RGB") + return image + + +def load_images(image_files): + out = [] + for image_file in image_files: + image = load_image(image_file) + out.append(image) + return out + +def merge_images(image_links: List = []): + """Merge multiple images into one image + + Args: + image_links (List, optional): List of image links. Defaults to []. + + Returns: + [type]: [description] + """ + if len(image_links) == 0: + return None + images = load_images(image_links) + if len(images) == 1: + return images[0] + widths, heights = zip(*(i.size for i in images)) + average_height = sum(heights) // len(heights) + for i, im in enumerate(images): + # scale in proportion + images[i] = im.resize((int(im.size[0] * average_height / im.size[1]), average_height)) + widths, heights = zip(*(i.size for i in images)) + total_width = sum(widths) + max_height = max(heights) + new_im = Image.new("RGB", (total_width + 10 * (len(images) - 1), max_height)) + x_offset = 0 + for i, im in enumerate(images): + if i > 0: + # past a column of 1 pixel starting from x_offset width being black, 8 pixels being white, and 1 pixel being black + new_im.paste(Image.new("RGB", (1, max_height), (0, 0, 0)), (x_offset, 0)) + x_offset += 1 + new_im.paste(Image.new("RGB", (8, max_height), (255, 255, 255)), (x_offset, 0)) + x_offset += 8 + new_im.paste(Image.new("RGB", (1, max_height), (0, 0, 0)), (x_offset, 0)) + x_offset += 1 + new_im.paste(im, (x_offset, 0)) + x_offset += im.size[0] + return new_im \ No newline at end of file diff --git a/editscore/utils.py b/editscore/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..0e42e44b6a2cd38401301b179448892725055dd6 --- /dev/null +++ b/editscore/utils.py @@ -0,0 +1,534 @@ +import os +from typing import Union, List, Optional +import json +import regex as re +import ast +import random +import json_repair + +def fix_json(input_str): + # Add double quotes around keys using regex + fixed_str = re.sub(r'(\w+):', r'"\1":', input_str) + + # Add double quotes around string values if necessary and wrap int/float values in [] + def format_value(match): + key, value, comma = match.groups() + value = value.strip() + # Check if value is an integer or float + if re.match(r'^-?\d+(\.\d+)?$', value): + value = f'[{value}]' + # Check if value is a boolean or null + elif re.match(r'^(true|false|null)$', value, re.IGNORECASE): + pass # leave as is + else: + # Add quotes around string values + value = f'"{value}"' + return f'{key}: {value}{comma}' + + fixed_str = re.sub(r'(".*?"):(.*?)(,|})', format_value, fixed_str) + + return fixed_str + +def repair_reasoning_field_robust(json_str: str) -> str: + """ + Robustly repair unescaped double quotes inside the "reasoning" field of a JSON string. + This function uses regular expressions and a lookahead assertion to locate + the end of the "reasoning" value, even if it is not the last field in the JSON. + + Args: + json_str (str): A possibly malformed JSON string that may contain + unescaped quotes within the "reasoning" field. + + Returns: + str: A repaired JSON string with properly escaped quotes inside "reasoning". + """ + # 1. Define a regex pattern that locates the reasoning value using a lookahead. + # The re.DOTALL flag allows '.' to match newline characters. + pattern = re.compile( + # --- Group 1: prefix part including the "reasoning" key and opening quote --- + r'("reasoning"\s*:\s*")' + + # --- Group 2: content inside the reasoning string (non-greedy) --- + r'(.*?)' + + # --- Lookahead assertion --- + # Match the ending quote of the "reasoning" value, + # but only if it is followed by a comma or closing brace. + r'(?="\s*[,}])', + + re.DOTALL + ) + + # 2. Define a replacement function to escape quotes inside the "reasoning" content. + def replacer(match): + prefix = match.group(1) # e.g., '"reasoning": "' + content = match.group(2) # e.g., 'Overall building...' + + # Escape all unescaped double quotes inside the reasoning text. + fixed_content = content.replace('"', '\\"') + + # Reassemble the full matched segment. The suffix is not consumed by the pattern, + # so we just return the prefix + repaired content. + return prefix + fixed_content + + # 3. Apply the regex substitution across the entire JSON string. + repaired_str = pattern.sub(replacer, json_str) + + return repaired_str + +def fallback_repair_json(input_str: str) -> str: + """ + Last-resort JSON repair that tries to preserve the 'reasoning' text + even when it contains unescaped quotes or other corruption. + + Target output: + {"reasoning": "", "score": [float, float]} + + Approach: + 1. Locate 'reasoning' key position and 'score' key position. + 2. Extract the raw substring between them (reasoning_raw). + 3. Clean only the outer noise (leading/trailing quotes, commas, braces), + but preserve internal punctuation. + 4. Unescape common escape sequences and normalize quotes. + 5. Extract numeric scores robustly. + 6. Return a valid JSON string. + """ + + s = input_str + + # Normalize whitespace for easier searching (but keep original for slicing) + lowered = s.lower() + + # 1) find the start of reasoning key (case-insensitive) + m_reason = re.search(r'"?reasoning"?\s*[:๏ผš]', lowered) + m_score = re.search(r'"?score"?\s*[:๏ผš]', lowered) + + reasoning_text = "" + scores = [] + + if m_reason and m_score: + # compute the real indices in the original string + start_idx = m_reason.end() # right after colon in 'reasoning:' + score_start_idx = m_score.start() + + # 2) slice the original string between reasoning value start and score key start + reasoning_raw = s[start_idx:score_start_idx] + + # 3) clean outer noise but preserve inner content: + # - strip whitespace and outer commas/braces + reasoning_raw = reasoning_raw.strip() + # remove leading commas/braces/colons + reasoning_raw = re.sub(r'^[\s,{\[]+', '', reasoning_raw) + # remove trailing commas/braces/colons (but keep inner punctuation) + reasoning_raw = re.sub(r'[\s,}\]]+$', '', reasoning_raw) + + # If the reasoning starts with a quote char, drop it (we'll re-escape later). + if reasoning_raw.startswith(("'", '"')): + reasoning_raw = reasoning_raw[1:] + # If it ends with a quote char (common), drop it. + if reasoning_raw.endswith(("'", '"')): + reasoning_raw = reasoning_raw[:-1] + + # 4) normalize escapes: + # Replace common escaped sequences (\" -> "), but avoid creating unbalanced quotes. + reasoning_raw = reasoning_raw.replace('\\"', '"').replace("\\'", "'") + # Replace fancy quotes with straight quotes (optional) + reasoning_raw = re.sub(r'[โ€œโ€]', '"', reasoning_raw) + reasoning_raw = re.sub(r"[โ€˜โ€™]", "'", reasoning_raw) + + # Trim again + reasoning_text = reasoning_raw.strip() + else: + # If we couldn't find both keys, try a looser regex capturing 'reasoning' value + m_loose = re.search(r'"?reasoning"?\s*[:๏ผš]\s*["\']?(.*?)["\']?\s*(,|$)', s, re.DOTALL | re.IGNORECASE) + if m_loose: + reasoning_text = m_loose.group(1).strip() + # normalize escapes as above + reasoning_text = reasoning_text.replace('\\"', '"').replace("\\'", "'") + reasoning_text = re.sub(r'[โ€œโ€]', '"', reasoning_text) + reasoning_text = re.sub(r"[โ€˜โ€™]", "'", reasoning_text) + + # 5) Extract two numeric scores anywhere after the 'score' key (robust) + if m_score: + # slice from score key to the end + score_slice = s[m_score.end():] + # find numbers (integers or floats) + nums = re.findall(r'-?\d+(?:\.\d+)?', score_slice) + try: + scores = [float(n) for n in nums[:2]] + except Exception: + scores = [] + else: + # fallback: try to find any two numbers in the whole string + nums = re.findall(r'-?\d+(?:\.\d+)?', s) + try: + scores = [float(n) for n in nums[:2]] + except Exception: + scores = [] + + # Ensure we always return two floats + if len(scores) < 2: + scores += [0.0] * (2 - len(scores)) + + # 6) Construct final object. Let json.dumps handle escaping inside the reasoning. + repaired_obj = { + "reasoning": reasoning_text, + "score": scores + } + + return json.dumps(repaired_obj, ensure_ascii=False) + +def robust_json_fix(s: str): + try: + return json_repair.loads(s) + except Exception: + pass + + for fixer in [fix_json, repair_reasoning_field_robust]: + s = fixer(s) + try: + return json_repair.loads(s) + except Exception: + print(f"Error: Cannot fix {fixer.__name__} {s=}") + continue + + try: + repaired_str = fallback_repair_json(s) + return json_repair.loads(repaired_str) + except Exception as e: + print(f"Error: Cannot fix fallback_repair_json {s=} {e=}") + return False + +def read_file_to_string(file_path): + """ + Reads the contents of a text file and returns it as a string. + + :param file_path: The path to the text file. + :return: A string containing the contents of the file. + """ + try: + with open(file_path, 'r', encoding='utf-8') as file: + return file.read() + except FileNotFoundError: + print(f"The file {file_path} was not found.") + return None + except Exception as e: + print(f"An error occurred: {e}") + return None + +def read_files_to_string(file_paths): + """ + Reads the contents of multiple text files and returns them as a single string, + with each file's contents separated by a newline. + + :param file_paths: A list of paths to text files. + :return: A string containing the concatenated contents of the files. + """ + all_contents = [] # List to hold the contents of each file + + for file_path in file_paths: + try: + with open(file_path, 'r', encoding='utf-8') as file: + all_contents.append(file.read()) + except FileNotFoundError: + print(f"The file {file_path} was not found.") + except Exception as e: + print(f"An error occurred while reading {file_path}: {e}") + + # Join all the contents with a newline character + return "\n".join(all_contents) + +def get_file_path(filename: Union[str, os.PathLike], search_from: Union[str, os.PathLike] = "."): + """ + Search for a file across a directory and return its absolute path. + + Args: + filename (Union[str, os.PathLike]): The name of the file to search for. + search_from (Union[str, os.PathLike], optional): The directory from which to start the search. Defaults to ".". + + Returns: + str: Absolute path to the found file. + + Raises: + FileNotFoundError: If the file is not found. + """ + for root, dirs, files in os.walk(search_from): + for name in files: + if name == filename: + return os.path.abspath(os.path.join(root, name)) + raise FileNotFoundError(filename, "not found.") + + + +#+========================================================================================= +def verify(s, target_sequence): + # Count the occurrences of the target sequence + count = s.count(target_sequence) + + # Check if the target sequence appears exactly twice + return count == 2 + + +def is_int_between_0_and_10(s): + try: + num = int(s) + return 0 <= num <= 10 + except ValueError: + return False + +def is_str_a_list_of_ints_0_to_10(s): + try: + # Attempt to parse the string as a Python literal (list, dict, etc.) + parsed = ast.literal_eval(s) + + # Check if the parsed object is a list + if not isinstance(parsed, list): + return False + + # Check if all elements are integers and between 0 to 10 + return all(isinstance(item, int) and 0 <= item <= 10 for item in parsed) + + except (ValueError, SyntaxError): + # If parsing fails or any other error occurs + return False + +def is_str_valid_score_format_brackets(s): + try: + # Removing brackets and splitting the string by commas + content = s.strip("[]").split(',') + + length = len(content) + + # Parsing each element and checking the format and range + scores = {} + for item in content: + key, value = item.split(':') + key = key.strip() + value = int(value.strip()) + + # Check if the key starts with 'score' and the value is in the correct range + if not key.startswith("score") or not 0 <= value <= 10: + return False + + scores[key] = value + + fetch_words = [f"score{i+1}" for i in range(length)] + # Check if at least 'score1' and 'score2' are present + return all(key in scores for key in fetch_words) + + except (ValueError, SyntaxError): + # If any parsing error occurs + return False + +def normalize_quotes(s: str) -> str: + """ + Replace curly/smart quotes with normal ASCII quotes. + """ + # ๅธธ่ง็š„ๅ‡ ็งๆ™บ่ƒฝๅผ•ๅท U+201C U+201D U+2018 U+2019 + return s.replace("โ€œ", '"').replace("โ€", '"').replace("โ€˜", "'").replace("โ€™", "'") + + +#+========================================================================================= +def mllm_output_to_dict(input_string, give_up_parsing=False, text_prompt=None, score_range: int = 10): + """ + Args: + input_string (str): actually the output of the mllm model to be parsed + output_file_name (str): The name of the output file. + """ + # Catch for gpt4v rate_limit_exceeded error + if input_string == "rate_limit_exceeded": + return "rate_limit_exceeded" + + if give_up_parsing: + guessed_value = random.randint(0, score_range) + json_content = {'score': [guessed_value, guessed_value], "reasoning": f"guess_if_cannot_parse | {input_string}"} + return json_content + + # Define the delimiters + delimiter = '||V^=^V||' + + if input_string.count(delimiter) == 2: + if not verify(input_string, delimiter): + print("The required delimiters were not found correctly in the string.", flush=True) + return False + # Extract the content between the delimiters + start_index = input_string.find(delimiter) + len(delimiter) + end_index = input_string.rfind(delimiter) + else: + # find the json mannually + # some mllm tends not to output the delimiters, but it does output the json contents + # so we will find the json content mannually + start_index = input_string.find('{') + end_index = input_string.rfind('}') + 1 + if start_index == -1 or end_index == 0: + # json not found + # some mllm tends to output only a list of scores like [6, 0], + # this time we will just get the scores and ignore the reasoning (other part of the json) + start_index = input_string.find('[') + end_index = input_string.rfind(']') + 1 + if re.match(r'^\[\d+, ?\d+\]$', input_string[start_index:end_index]): + scores = json.loads(input_string[start_index:end_index]) + if not isinstance(scores, list): + scores = [scores] + json_content = {'score': scores, "reasoning": "System: output is simply a list of scores"} + json_str = json.dumps(json_content) + input_string = json_str + start_index = 0 + end_index = len(json_str) + elif is_int_between_0_and_10(input_string): # if output is simply a number + scores = [int(input_string)] + json_content = {'score': scores, "reasoning": "System: output is simply a number"} + json_str = json.dumps(json_content) + input_string = json_str + start_index = 0 + end_index = len(json_str) + else: + print(f"22 222 Failed to find the json content in the string. {text_prompt=} {input_string=}", flush=True) + return False + + # Check if we found two delimiters + if start_index != -1 and end_index != -1 and start_index != end_index: + # Extract the JSON string + json_str = input_string[start_index:end_index].strip() + json_str = json_str.replace("\n", "") + # Parse the JSON string into a dictionary + try: + json_str = normalize_quotes(json_str) + new_data = json.loads(json_str) + if not isinstance(new_data['score'], list): + new_data['score'] = [new_data['score']] + except Exception as e1: + print(f"Now fixing: {e1=} {json_str=}") + + new_data = robust_json_fix(json_str) + return new_data + else: + print("The required delimiters were not found correctly in the string.") + return False + +def write_entry_to_json_file(input_string, uid, prompt_input, vision_input, output_file_name, give_up_parsing=False): + """ + Args: + input_string (str): actually the output of the mllm model to be parsed + uid (str): The unique identifier for the each item in the test data + prompt_input (str): The prompt input for the entry. text prompt. + vision_input (str): The vision input for the entry. image links. + output_file_name (str): The name of the output file. + """ + # Catch for gpt4v rate_limit_exceeded error + if input_string == "rate_limit_exceeded": + return "rate_limit_exceeded" + + # Define the delimiters + delimiter = '||V^=^V||' + + if input_string.count(delimiter) == 2: + if not verify(input_string, delimiter): + print("The required delimiters were not found correctly in the string.") + return False + # Extract the content between the delimiters + start_index = input_string.find(delimiter) + len(delimiter) + end_index = input_string.rfind(delimiter) + else: + # find the json mannually + # some mllm tends not to output the delimiters, but it does output the json contents + # so we will find the json content mannually + start_index = input_string.find('{') + end_index = input_string.rfind('}') + 1 + if start_index == -1 or end_index == 0: + # json not found + # some mllm tends to output only a list of scores like [6, 0], + # this time we will just get the scores and ignore the reasoning (other part of the json) + start_index = input_string.find('[') + end_index = input_string.rfind(']') + 1 + if give_up_parsing: # if we want to give up parsing + guessed_value = random.randint(0, 10) + print(f"Failed to find the json content in the string. Guess a value : {guessed_value}.") + json_content = {'score': [guessed_value], "reasoning": f"guess_if_cannot_parse | {input_string}"} + json_str = json.dumps(json_content) + input_string = json_str + start_index = 0 + end_index = len(json_str) + elif re.match(r'^\[\d+, ?\d+\]$', input_string[start_index:end_index]): + scores = json.loads(input_string[start_index:end_index]) + json_content = {'score': scores, "reasoning": None} + json_str = json.dumps(json_content) + input_string = json_str + start_index = 0 + end_index = len(json_str) + elif is_int_between_0_and_10(input_string): # if output is simply a number + scores = [int(input_string)] + json_content = {'score': scores, "reasoning": None} + json_str = json.dumps(json_content) + input_string = json_str + start_index = 0 + end_index = len(json_str) + else: + print("Failed to find the json content in the string.") + return False + + # Check if we found two delimiters + if start_index != -1 and end_index != -1 and start_index != end_index: + # Extract the JSON string + json_str = input_string[start_index:end_index].strip() + json_str = json_str.replace("\n", "") + try: + # Parse the JSON string into a dictionary + new_data = json.loads(json_str) + + # Ensure the directory exists + os.makedirs(os.path.dirname(output_file_name), exist_ok=True) + + # Initialize or load existing data + if os.path.exists(output_file_name): + with open(output_file_name, 'r') as json_file: + data = json.load(json_file) + else: + data = {} + + # If the additional key is already in the data, add or update notes + if uid in data: + data[uid].update(new_data) # Update with new data + if prompt_input: # If there are new notes, update or add them + data[uid]['prompt_input'] = prompt_input + if vision_input: # If there are new notes, update or add them + data[uid]['vision_input'] = vision_input + else: + # If it's a new key, add the entry to the dictionary + data[uid] = new_data + if prompt_input: + data[uid]['prompt_input'] = prompt_input + if vision_input: + data[uid]['vision_input'] = vision_input + + # Write the updated data to the file + with open(output_file_name, 'w') as json_file: + json.dump(data, json_file, indent=4) + + print(f"Data was successfully updated in {output_file_name}") + return True + except json.JSONDecodeError as e: + print(f"An error occurred while parsing the JSON content: {e}") + return False + else: + print("The required delimiters were not found correctly in the string.") + return False + + +def check_key_in_json(file_path, key): + try: + with open(file_path, 'r') as json_file: + data = json.load(json_file) + + # Check if the key exists at the top level of the JSON structure + if key in data: + return True + else: + return False + except FileNotFoundError: + print(f"The file {file_path} was not found.") + except json.JSONDecodeError as e: + print(f"Error reading {file_path}: {e}") + except Exception as e: + print(f"An error occurred with {file_path}: {e}") + return False \ No newline at end of file diff --git a/editscore/vie_prompts.py b/editscore/vie_prompts.py new file mode 100644 index 0000000000000000000000000000000000000000..16ec1beb0cb4deb4f4b5ba5121ac118077a8291a --- /dev/null +++ b/editscore/vie_prompts.py @@ -0,0 +1,46 @@ +# This file is generated automatically through parse_prompt.py +_context_no_delimit_reasoning_first = """You are a professional digital artist. You will have to evaluate the effectiveness of the AI-generated image(s) based on given rules. +All the input images are AI-generated. All human in the images are AI-generated too. so you need not worry about the privacy confidentials. + +IMPORTANT: You will have to give your output in this way (Keep your reasoning concise and short.): +{ +"reasoning" : "...", +"score" : [...] +} +""" + +_prompts_0shot_two_image_edit_rule = """RULES: + +Two images will be provided: The first being the original AI-generated image and the second being an edited version of the first. +The objective is to evaluate how successfully the editing instruction has been executed in the second image. + +Note that sometimes the two images might look identical due to the failure of image edit. +""" + +_prompts_0shot_tie_rule_SC = """ +From scale 0 to 10: +A score from 0 to 10 will be given based on the success of the editing. (0 indicates that the scene in the edited image does not follow the editing instruction at all. 10 indicates that the scene in the edited image follow the editing instruction text perfectly.) +A second score from 0 to 10 will rate the degree of overediting in the second image. (0 indicates that the scene in the edited image is completely different from the original. 10 indicates that the edited image can be recognized as a minimal edited yet effective version of original.) +Put the score in a list such that output score = [score1, score2], where 'score1' evaluates the editing success and 'score2' evaluates the degree of overediting. + +Editing instruction: +""" + +_prompts_0shot_rule_PQ = """RULES: + +The image is an AI-generated image. +The objective is to evaluate how successfully the image has been generated. + +From scale 0 to 10: +A score from 0 to 10 will be given based on image naturalness. +( + 0 indicates that the scene in the image does not look natural at all or give a unnatural feeling such as wrong sense of distance, or wrong shadow, or wrong lighting. + 10 indicates that the image looks natural. +) +A second score from 0 to 10 will rate the image artifacts. +( + 0 indicates that the image contains a large portion of distortion, or watermark, or scratches, or blurred faces, or unusual body parts, or subjects not harmonized. + 10 indicates the image has no artifacts. +) +Put the score in a list such that output score = [naturalness, artifacts] +""" \ No newline at end of file diff --git a/evaluate.sh b/evaluate.sh new file mode 100644 index 0000000000000000000000000000000000000000..a5a3bae2486269587449a7516d5eb98b51f8ead9 --- /dev/null +++ b/evaluate.sh @@ -0,0 +1,20 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER + +python evaluation.py \ +--benchmark_dir EditScore/EditReward-Bench \ +--result_dir results/EditScore-7B \ +--backbone qwen25vl \ +--model_name_or_path Qwen/Qwen2.5-VL-7B-Instruct \ +--lora_path EditScore/EditScore-7B \ +--score_range 25 \ +--max_workers 1 \ +--max_model_len 4096 \ +--max_num_seqs 1 \ +--max_num_batched_tokens 4096 \ +--tensor_parallel_size 1 \ +--num_pass 1 + +python calculate_statistics.py \ +--result_dir results/EditScore-7B/qwen25vl \ No newline at end of file diff --git a/evaluate_72B_vllm.sh b/evaluate_72B_vllm.sh new file mode 100644 index 0000000000000000000000000000000000000000..ce7874648145b93edaabac577e834c595fa21358 --- /dev/null +++ b/evaluate_72B_vllm.sh @@ -0,0 +1,20 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER + +python evaluation.py \ +--benchmark_dir EditScore/EditReward-Bench \ +--result_dir results/EditScore-72B \ +--backbone qwen25vl_vllm \ +--model_name_or_path Qwen/Qwen2.5-VL-72B-Instruct \ +--lora_path EditScore/EditScore-72B \ +--score_range 25 \ +--max_workers 1 \ +--max_model_len 4096 \ +--max_num_seqs 1 \ +--max_num_batched_tokens 4096 \ +--tensor_parallel_size 4 \ +--num_pass 1 + +python calculate_statistics.py \ +--result_dir results/EditScore-72B/qwen25vl_vllm \ No newline at end of file diff --git a/evaluate_qwen3_vl_32B.sh b/evaluate_qwen3_vl_32B.sh new file mode 100644 index 0000000000000000000000000000000000000000..128e574aca3b542dc7c5a2978ff1a5e43b8f0500 --- /dev/null +++ b/evaluate_qwen3_vl_32B.sh @@ -0,0 +1,20 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER + +python evaluation.py \ +--benchmark_dir EditScore/EditReward-Bench \ +--result_dir results/EditScore-Qwen3-VL-32B \ +--backbone qwen3vl \ +--model_name_or_path Qwen/Qwen3-VL-32B-Instruct \ +--lora_path EditScore/EditScore-Qwen3-VL-32B-Instruct \ +--score_range 25 \ +--max_workers 1 \ +--max_model_len 4096 \ +--max_num_seqs 1 \ +--max_num_batched_tokens 4096 \ +--tensor_parallel_size 1 \ +--num_pass 1 + +python calculate_statistics.py \ +--result_dir results/EditScore-Qwen3-VL-32B/qwen3vl \ No newline at end of file diff --git a/evaluate_qwen3_vl_32B_avg4.sh b/evaluate_qwen3_vl_32B_avg4.sh new file mode 100644 index 0000000000000000000000000000000000000000..12ccf69065fe5b6408c1a8269b49bdac2f73843a --- /dev/null +++ b/evaluate_qwen3_vl_32B_avg4.sh @@ -0,0 +1,20 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER + +python evaluation.py \ +--benchmark_dir EditScore/EditReward-Bench \ +--result_dir results/EditScore-Qwen3-VL-32B-avg4 \ +--backbone qwen3vl \ +--model_name_or_path Qwen/Qwen3-VL-32B-Instruct \ +--lora_path EditScore/EditScore-Qwen3-VL-32B-Instruct \ +--score_range 25 \ +--max_workers 1 \ +--max_model_len 4096 \ +--max_num_seqs 1 \ +--max_num_batched_tokens 4096 \ +--tensor_parallel_size 1 \ +--num_pass 4 + +python calculate_statistics.py \ +--result_dir results/EditScore-Qwen3-VL-32B-avg4/qwen3vl \ No newline at end of file diff --git a/evaluate_qwen3_vl_4B.sh b/evaluate_qwen3_vl_4B.sh new file mode 100644 index 0000000000000000000000000000000000000000..6f40f84447d4d021168e6d0ceb76d06dec5b8d85 --- /dev/null +++ b/evaluate_qwen3_vl_4B.sh @@ -0,0 +1,20 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER + +python evaluation.py \ +--benchmark_dir EditScore/EditReward-Bench \ +--result_dir results/EditScore-Qwen3-VL-4B \ +--backbone qwen3vl \ +--model_name_or_path Qwen/Qwen3-VL-4B-Instruct \ +--lora_path EditScore/EditScore-Qwen3-VL-4B-Instruct \ +--score_range 25 \ +--max_workers 1 \ +--max_model_len 4096 \ +--max_num_seqs 1 \ +--max_num_batched_tokens 4096 \ +--tensor_parallel_size 1 \ +--num_pass 1 + +python calculate_statistics.py \ +--result_dir results/EditScore-Qwen3-VL-4B/qwen3vl \ No newline at end of file diff --git a/evaluate_qwen3_vl_4B_avg4.sh b/evaluate_qwen3_vl_4B_avg4.sh new file mode 100644 index 0000000000000000000000000000000000000000..d0c1d16373df3398e9a0c08c0c7f3cc2b3d5f478 --- /dev/null +++ b/evaluate_qwen3_vl_4B_avg4.sh @@ -0,0 +1,20 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER + +python evaluation.py \ +--benchmark_dir EditScore/EditReward-Bench \ +--result_dir results/EditScore-Qwen3-VL-4B-avg4 \ +--backbone qwen3vl \ +--model_name_or_path Qwen/Qwen3-VL-4B-Instruct \ +--lora_path EditScore/EditScore-Qwen3-VL-4B-Instruct \ +--score_range 25 \ +--max_workers 1 \ +--max_model_len 4096 \ +--max_num_seqs 1 \ +--max_num_batched_tokens 4096 \ +--tensor_parallel_size 1 \ +--num_pass 4 + +python calculate_statistics.py \ +--result_dir results/EditScore-Qwen3-VL-4B-avg4/qwen3vl \ No newline at end of file diff --git a/evaluate_qwen3_vl_4B_vllm.sh b/evaluate_qwen3_vl_4B_vllm.sh new file mode 100644 index 0000000000000000000000000000000000000000..095a52b689cd20c6257e1032b2ea51fb5aaf1cb0 --- /dev/null +++ b/evaluate_qwen3_vl_4B_vllm.sh @@ -0,0 +1,20 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER + +python evaluation.py \ +--benchmark_dir EditScore/EditReward-Bench \ +--result_dir results/EditScore-Qwen3-VL-4B \ +--backbone qwen3vl_vllm \ +--model_name_or_path Qwen/Qwen3-VL-4B-Instruct \ +--lora_path EditScore/EditScore-Qwen3-VL-4B-Instruct \ +--score_range 25 \ +--max_workers 1 \ +--max_model_len 4096 \ +--max_num_seqs 1 \ +--max_num_batched_tokens 4096 \ +--tensor_parallel_size 1 \ +--num_pass 1 + +python calculate_statistics.py \ +--result_dir results/EditScore-Qwen3-VL-4B/qwen3vl_vllm \ No newline at end of file diff --git a/evaluate_qwen3_vl_8B.sh b/evaluate_qwen3_vl_8B.sh new file mode 100644 index 0000000000000000000000000000000000000000..52531c13db95f5dbfdb2f8085623006f2eca5478 --- /dev/null +++ b/evaluate_qwen3_vl_8B.sh @@ -0,0 +1,20 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER + +python evaluation.py \ +--benchmark_dir EditScore/EditReward-Bench \ +--result_dir results/EditScore-Qwen3-VL-8B \ +--backbone qwen3vl \ +--model_name_or_path Qwen/Qwen3-VL-8B-Instruct \ +--lora_path EditScore/EditScore-Qwen3-VL-8B-Instruct \ +--score_range 25 \ +--max_workers 1 \ +--max_model_len 4096 \ +--max_num_seqs 1 \ +--max_num_batched_tokens 4096 \ +--tensor_parallel_size 1 \ +--num_pass 1 + +python calculate_statistics.py \ +--result_dir results/EditScore-Qwen3-VL-8B/qwen3vl \ No newline at end of file diff --git a/evaluate_qwen3_vl_8B_avg4.sh b/evaluate_qwen3_vl_8B_avg4.sh new file mode 100644 index 0000000000000000000000000000000000000000..824c91589418b093349fa6059853702e827ce390 --- /dev/null +++ b/evaluate_qwen3_vl_8B_avg4.sh @@ -0,0 +1,20 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER + +python evaluation.py \ +--benchmark_dir EditScore/EditReward-Bench \ +--result_dir results/EditScore-Qwen3-VL-8B-avg4 \ +--backbone qwen3vl \ +--model_name_or_path Qwen/Qwen3-VL-8B-Instruct \ +--lora_path EditScore/EditScore-Qwen3-VL-8B-Instruct \ +--score_range 25 \ +--max_workers 1 \ +--max_model_len 4096 \ +--max_num_seqs 1 \ +--max_num_batched_tokens 4096 \ +--tensor_parallel_size 1 \ +--num_pass 4 + +python calculate_statistics.py \ +--result_dir results/EditScore-Qwen3-VL-8B-avg4/qwen3vl \ No newline at end of file diff --git a/evaluate_vllm.sh b/evaluate_vllm.sh new file mode 100644 index 0000000000000000000000000000000000000000..fe0a3048f81a13970f2f4a67ee176e42c50a9195 --- /dev/null +++ b/evaluate_vllm.sh @@ -0,0 +1,20 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER + +python evaluation.py \ +--benchmark_dir EditScore/EditReward-Bench \ +--result_dir results/EditScore-7B \ +--backbone qwen25vl_vllm \ +--model_name_or_path Qwen/Qwen2.5-VL-7B-Instruct \ +--lora_path EditScore/EditScore-7B \ +--score_range 25 \ +--max_workers 1 \ +--max_model_len 4096 \ +--max_num_seqs 1 \ +--max_num_batched_tokens 4096 \ +--tensor_parallel_size 1 \ +--num_pass 1 + +python calculate_statistics.py \ +--result_dir results/EditScore-7B/qwen25vl_vllm \ No newline at end of file diff --git a/evaluation.py b/evaluation.py new file mode 100644 index 0000000000000000000000000000000000000000..59a2ffb9a36c34ecbe787139be95bf8c85ae07d8 --- /dev/null +++ b/evaluation.py @@ -0,0 +1,272 @@ +import dotenv + +dotenv.load_dotenv(override=True) + +import argparse +import glob +import hashlib +import json +import logging +import os +import time +import threading +from concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor, as_completed +from typing import List, Tuple, Dict, Any, Optional + +import dotenv +from PIL import Image +from tqdm import tqdm +from datasets import Dataset, load_dataset + +from editscore import EditScore + +PROMPT_FOLLOWING = "prompt_following" +CONSISTENCY = "consistency" +OVERALL = "overall" +SCORE_CATEGORIES = [PROMPT_FOLLOWING, CONSISTENCY, OVERALL] + + +class CacheManager: + def __init__(self, cache_file: str): + self.cache_file = cache_file + self.lock = threading.Lock() + self.cache = self._load() + + def _load(self) -> Dict[str, Any]: + cache = {} + if not os.path.exists(self.cache_file): + print( + f"Cache file not found at {self.cache_file}. A new one will be created." + ) + return cache + + with open(self.cache_file, "r", encoding="utf-8") as f: + for i, line in enumerate(f): + try: + data = json.loads(line) + cache[data["key"]] = data["result"] + except json.JSONDecodeError: + logging.warning( + f"Skipping corrupted line {i + 1} in cache file: {line.strip()}" + ) + print(f"Loaded {len(cache)} items from {self.cache_file}.") + return cache + + def get(self, key: str) -> Optional[Any]: + return self.cache.get(key) + + def append(self, key: str, result: Any): + with self.lock: + self.cache[key] = result + with open(self.cache_file, "a", encoding="utf-8") as f: + f.write( + json.dumps({"key": key, "result": result}, ensure_ascii=False) + + "\n" + ) + +def generate_cache_key(pair_key): + return hashlib.sha256(pair_key.encode("utf-8")).hexdigest() + +def load_pairs_dataset(dataset: Dataset) -> Dict[str, Tuple[str, Image.Image, Image.Image]]: + pairs = {} + for data in dataset: + key1, key2 = data["key"] + instruction = data["instruction"] + input_image = data["input_image"].convert("RGB") + + pairs[key1] = (instruction, input_image, data["output_images"][0].convert("RGB")) + pairs[key2] = (instruction, input_image, data["output_images"][1].convert("RGB")) + return pairs + +def _load_item(data: dict) -> list[tuple[str, tuple[str, Image.Image, Image.Image]]]: + key1, key2 = data["key"] + instruction = data["instruction"] + + input_image = data["input_image"].convert("RGB") + output_image1 = data["output_images"][0].convert("RGB") + output_image2 = data["output_images"][1].convert("RGB") + + return [ + (key1, (instruction, input_image, output_image1)), + (key2, (instruction, input_image, output_image2)), + ] + +def load_pairs_dataset_multithreaded(dataset: Dataset, max_workers: int = None) -> Dict[str, Tuple[str, Image.Image, Image.Image]]: + if max_workers is None: + # max_workers = min(32, (os.cpu_count() or 1) * 5) + max_workers = os.cpu_count() or 1 + + pairs = {} + + print(f"Processing dataset (length: {len(dataset)}) with {max_workers} threads", flush=True) + + with ProcessPoolExecutor(max_workers=max_workers) as executor: + results_iterator = tqdm( + executor.map(_load_item, dataset), + total=len(dataset), + desc="Processing dataset with multiple threads" + ) + + for result_pairs in results_iterator: + pairs.update(result_pairs) + + return pairs + + +def process_single_item(key, item, scorer): + instruction = item[0] + input_image = item[1] + output_image = item[2] + + output_image = output_image.resize((input_image.size[0], input_image.size[1])) + + score = scorer.evaluate([input_image, output_image], instruction) + return key, score + + +def parse_args(): + parser = argparse.ArgumentParser() + parser.add_argument( + "--benchmark_dir", type=str, default="EditScore/EditReward-Bench" + ) + parser.add_argument("--result_dir", type=str, required=True) + parser.add_argument( + "--backbone", + type=str, + default="openai", + choices=["openai", "qwen25vl", "qwen25vl_vllm", "internvl3_5", "qwen3vl", "qwen3vl_vllm"], + ) + parser.add_argument("--model_name_or_path", type=str, default="gpt-4.1") + parser.add_argument( + "--openai_url", type=str, default="https://api.openai.com/v1/chat/completions" + ) + parser.add_argument("--key", type=str, default="PUT YOUR API KEY HERE") + parser.add_argument("--num_pass", type=int, default=1) + parser.add_argument("--temperature", type=float, default=0.7) + parser.add_argument("--max_workers", type=int, default=20) + parser.add_argument("--score_range", type=int, default=25) + parser.add_argument("--tensor_parallel_size", type=int, default=1) + parser.add_argument("--max_model_len", type=int, default=1536) + parser.add_argument("--max_num_seqs", type=int, default=32) + parser.add_argument("--max_num_batched_tokens", type=int, default=1536) + parser.add_argument("--lora_path", type=str, default="EditScore/EditScore-7B") + parser.add_argument("--cache_dir", type=str, default=None) + return parser.parse_args() + + +def main(args): + start_time = time.time() + scorer = EditScore( + backbone=args.backbone, + key=args.key, + openai_url=args.openai_url, + model_name_or_path=args.model_name_or_path, + score_range=args.score_range, + temperature=args.temperature, + tensor_parallel_size=args.tensor_parallel_size, + max_model_len=args.max_model_len, + max_num_seqs=args.max_num_seqs, + max_num_batched_tokens=args.max_num_batched_tokens, + num_pass=args.num_pass, + lora_path=args.lora_path, + cache_dir=args.cache_dir, + ) + print(f"Scorer initialized in {time.time() - start_time} seconds", flush=True) + + cache_dir = os.path.join(args.result_dir, ".cache") + os.makedirs(cache_dir, exist_ok=True) + cache_file = os.path.join( + cache_dir, f"{args.backbone}_{args.model_name_or_path.replace('/', '_')}.jsonl" + ) + cache_manager = CacheManager(cache_file) + + start_time = time.time() + dataset = load_dataset(args.benchmark_dir, split="train") + print(f"Dataset loaded in {time.time() - start_time} seconds", flush=True) + + start_time = time.time() + unique_pairs = load_pairs_dataset_multithreaded(dataset) + print(f"Pairs loaded in {time.time() - start_time} seconds", flush=True) + + all_scores = {} + pairs_to_process = [ + pair_key + for pair_key in unique_pairs.keys() + if cache_manager.get(generate_cache_key(pair_key)) is None + ] + + for pair_key in unique_pairs.keys(): + if pair_key not in pairs_to_process: + all_scores[pair_key] = cache_manager.get(generate_cache_key(pair_key)) + + print( + f"{len(unique_pairs) - len(pairs_to_process)} pairs found in cache. Processing {len(pairs_to_process)} new pairs.", + flush=True + ) + + if pairs_to_process: + with ThreadPoolExecutor(max_workers=args.max_workers) as executor: + futures = [ + executor.submit(process_single_item, pair_key, unique_pairs[pair_key], scorer) + for pair_key in pairs_to_process + ] + + for future in tqdm( + as_completed(futures), + total=len(futures), + unit="pair", + desc="Processing", + ): + pair_key, result = future.result() + if result: + all_scores[pair_key] = result + cache_manager.append(generate_cache_key(pair_key), result) + + print("Writing results...", flush=True) + + start_time = time.time() + # dataset = dataset.remove_columns(["input_image", "output_images"]) + for idx, data in enumerate(dataset): + key1, key2 = data["key"] + task_type = data["task_type"] + dimension = data["dimension"] + + score1 = all_scores[key1][dimension] + score2 = all_scores[key2][dimension] + data["score"] = [score1, score2] + + input_image_path = os.path.join(args.result_dir, "images", f"{key1}_input.png") + output_image_path1 = os.path.join(args.result_dir, "images", f"{key1}.png") + output_image_path2 = os.path.join(args.result_dir, "images", f"{key2}.png") + + os.makedirs(os.path.dirname(input_image_path), exist_ok=True) + + data['input_image'].save(input_image_path) + data['output_images'][0].save(output_image_path1) + data['output_images'][1].save(output_image_path2) + + json_line = { + "key": (key1, key2), + "idx": idx, + "score": [score1, score2], + "SC_reasoning": [all_scores[key1]["SC_reasoning"], all_scores[key2]["SC_reasoning"]], + "PQ_reasoning": [all_scores[key1]["PQ_reasoning"], all_scores[key2]["PQ_reasoning"]], + "input_image": input_image_path, + "output_images": [output_image_path1, output_image_path2], + } + + save_file = os.path.join( + args.result_dir, args.backbone, task_type, f"{dimension}.jsonl" + ) + os.makedirs(os.path.dirname(save_file), exist_ok=True) + + with open(save_file, "a", encoding="utf-8") as f: + f.write(json.dumps(json_line, ensure_ascii=False) + "\n") + + print(f"Results written in {time.time() - start_time} seconds", flush=True) + print("--- Completed! ---", flush=True) + + +if __name__ == "__main__": + args = parse_args() + main(args) diff --git a/example_images/input.png b/example_images/input.png new file mode 100644 index 0000000000000000000000000000000000000000..761420bf28c5021344bd8adf143325ec0bdbfb8c --- /dev/null +++ b/example_images/input.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:62e32498503e4a45eb5d65c9f739e8933c58321c4036e86cc548561e6b258882 +size 422003 diff --git a/example_images/output.png b/example_images/output.png new file mode 100644 index 0000000000000000000000000000000000000000..0ae760267d00a5acb1bfa7301ed7f3f6419c6e83 --- /dev/null +++ b/example_images/output.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fc2f58112622ce178016b7dc3b48d0ae329e05ec612d2beb0b075e63f379369b +size 384464 diff --git a/examples/EditScore-train/README.md b/examples/EditScore-train/README.md new file mode 100644 index 0000000000000000000000000000000000000000..9062be03693a144f6b7f2043e1f948d23d43021f --- /dev/null +++ b/examples/EditScore-train/README.md @@ -0,0 +1,101 @@ +# EditScore Reward Model Training Guide + +This guide explains how to train EditScore reward models using LLaMA-Factory. + +## 1. Environment Setup + +### Clone LLaMA-Factory and Configure Virtual Environment + +```bash +git clone --depth 1 https://github.com/hiyouga/LLaMA-Factory.git +cd LLaMA-Factory +conda create -n llama-factory python=3.10 +conda activate llama-factory +pip install -e ".[torch,metrics]" --no-build-isolation +``` + +## 2. Directory Structure Configuration + +Create necessary folders and files in the LLaMA-Factory root directory: + +```bash +# Create log and output directories +mkdir -p logs +mkdir -p output + +# Create training configuration directory +mkdir -p examples/train_editscore + +# Copy training configuration files +cp EditScore/examples/EditScore-train/config/*.yaml examples/train_editscore/ + +# Copy training script +cp EditScore/examples/EditScore-train/train.sh . +``` + +## 3. Dataset Registration + +Register the EditScore-Reward-Data dataset in `LLaMA-Factory/data/dataset_info.json`: + +```json +"EditScore-Reward-Data": { + "file_name": "/path/to/your/reward.json", + "formatting": "sharegpt", + "columns": { + "messages": "conversations", + "images": "images" + } +} +``` + +## 4. Training Configuration Description + +### Single-Machine Training Configuration +- `editscore_7B.yaml` - Train EditScore-7B model (single machine) +- `editscore_qwen3_vl_4B_instruct.yaml` - Train EditScore_Qwen3_Vl_4B_Instruct model (single machine) +- `editscore_qwen3_vl_8B_instruct.yaml` - Train EditScore_Qwen3_Vl_8B_Instruct model (single machine) + +### Multi-Machine Training Configuration +- `editscore_32B.yaml` - Train EditScore-32B model (two machines) +- `editscore_72B.yaml` - Train EditScore-72B model (two machines) + +## 5. Start Training + +### Single-Machine Training + +```bash +# Modify experiment_name in train.sh to the corresponding configuration file name +# For example: name=editscore_7B +bash train.sh +``` + +### Multi-Machine Training + +**Master node (rank=0):** +```bash +bash train.sh --rank=0 --world_size=2 --master_addr=MASTER_NODE_IP --master_port=29500 +``` + +**Worker node (rank=1):** +```bash +bash train.sh --rank=1 --world_size=2 --master_addr=MASTER_NODE_IP --master_port=29500 +``` + +## 6. Parameter Configuration + +Users can modify the following parameters in the YAML configuration files as needed: + +- `per_device_train_batch_size`: Batch size per device +- `gradient_accumulation_steps`: Gradient accumulation steps +- `learning_rate`: Learning rate +- `num_train_epochs`: Number of training epochs +- `max_samples`: Maximum number of samples +- `output_dir`: Output directory + +## 7. Output Files + +After training completion, model files will be saved in the corresponding output directories: +- Single-machine training: `LLaMA-Factory/output/model_name/` +- Log files: `LLaMA-Factory/logs/experiment_name_rank.log` + + diff --git a/examples/EditScore-train/config/editscore_32B.yaml b/examples/EditScore-train/config/editscore_32B.yaml new file mode 100644 index 0000000000000000000000000000000000000000..4e12bc35dedca1d585268509b5fe41e4e07b0565 --- /dev/null +++ b/examples/EditScore-train/config/editscore_32B.yaml @@ -0,0 +1,42 @@ +model_name_or_path: Qwen/Qwen2.5-VL-32B-Instruct +image_max_pixels: 262144 +video_max_pixels: 16384 +trust_remote_code: true + +### method +stage: sft +do_train: true +finetuning_type: lora +lora_rank: 32 +lora_target: all + +deepspeed: LLaMA-Factory/examples/deepspeed/ds_z2_config.json +### dataset +dataset: EditScore-Reward-Data +template: qwen2_vl +cutoff_len: 8192 +max_samples: 100000 +overwrite_cache: true +preprocessing_num_workers: 16 +dataloader_num_workers: 4 + +### output +output_dir: LLaMA-Factory/output/editscore_32B/ +logging_steps: 1 +save_steps: 250 +plot_loss: true +overwrite_output_dir: true +save_only_model: false +report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow] + +### train +per_device_train_batch_size: 1 +gradient_accumulation_steps: 8 +learning_rate: 1.0e-4 +num_train_epochs: 3.0 +lr_scheduler_type: cosine +warmup_ratio: 0.1 +bf16: true +ddp_timeout: 180000000 +resume_from_checkpoint: null + diff --git a/examples/EditScore-train/config/editscore_72B.yaml b/examples/EditScore-train/config/editscore_72B.yaml new file mode 100644 index 0000000000000000000000000000000000000000..3d5c5feb64e6aa10249c2e295d83b361941c52f1 --- /dev/null +++ b/examples/EditScore-train/config/editscore_72B.yaml @@ -0,0 +1,42 @@ +model_name_or_path: Qwen/Qwen2.5-VL-72B-Instruct +image_max_pixels: 262144 +video_max_pixels: 16384 +trust_remote_code: true + +### method +stage: sft +do_train: true +finetuning_type: lora +lora_rank: 32 +lora_target: all + +deepspeed: LLaMA-Factory/examples/deepspeed/ds_z3_config.json +### dataset +dataset: EditScore-Reward-Data +template: qwen2_vl +cutoff_len: 8192 +max_samples: 100000 +overwrite_cache: true +preprocessing_num_workers: 16 +dataloader_num_workers: 4 + +### output +output_dir: LLaMA-Factory/output/editscore_72B/ +logging_steps: 1 +save_steps: 250 +plot_loss: true +overwrite_output_dir: true +save_only_model: false +report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow] + +### train +per_device_train_batch_size: 1 +gradient_accumulation_steps: 8 +learning_rate: 1.0e-4 +num_train_epochs: 3.0 +lr_scheduler_type: cosine +warmup_ratio: 0.1 +bf16: true +ddp_timeout: 180000000 +resume_from_checkpoint: null + diff --git a/examples/EditScore-train/config/editscore_7B.yaml b/examples/EditScore-train/config/editscore_7B.yaml new file mode 100644 index 0000000000000000000000000000000000000000..80e188ad71f57c77f11c5bab722ae5fdbcfa4f07 --- /dev/null +++ b/examples/EditScore-train/config/editscore_7B.yaml @@ -0,0 +1,41 @@ +model_name_or_path: Qwen/Qwen2.5-VL-7B-Instruct +image_max_pixels: 262144 +video_max_pixels: 16384 +trust_remote_code: true + +### method +stage: sft +do_train: true +finetuning_type: lora +lora_rank: 32 +lora_target: all + +### dataset +dataset: EditScore-Reward-Data +template: qwen2_vl +cutoff_len: 8192 +max_samples: 100000 +overwrite_cache: true +preprocessing_num_workers: 16 +dataloader_num_workers: 4 + +### output +output_dir: LLaMA-Factory/output/editscore_7B/ +logging_steps: 1 +save_steps: 250 +plot_loss: true +overwrite_output_dir: true +save_only_model: false +report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow] + +### train +per_device_train_batch_size: 1 +gradient_accumulation_steps: 16 +learning_rate: 1.0e-4 +num_train_epochs: 3.0 +lr_scheduler_type: cosine +warmup_ratio: 0.1 +bf16: true +ddp_timeout: 180000000 +resume_from_checkpoint: null + diff --git a/examples/EditScore-train/config/editscore_qwen3_vl_4B_instruct.yaml b/examples/EditScore-train/config/editscore_qwen3_vl_4B_instruct.yaml new file mode 100644 index 0000000000000000000000000000000000000000..a38173841dab6e076656c4d1d340dd33a1f5d477 --- /dev/null +++ b/examples/EditScore-train/config/editscore_qwen3_vl_4B_instruct.yaml @@ -0,0 +1,41 @@ +model_name_or_path: Qwen/Qwen3-VL-4B-Instruct +image_max_pixels: 262144 +video_max_pixels: 16384 +trust_remote_code: true + +### method +stage: sft +do_train: true +finetuning_type: lora +lora_rank: 32 +lora_target: all + +### dataset +dataset: EditScore-Reward-Data +template: qwen3_vl +cutoff_len: 8192 +max_samples: 500000 +overwrite_cache: true +preprocessing_num_workers: 16 +dataloader_num_workers: 4 + +### output +output_dir: LLaMA-Factory/output/editscore_qwen3_vl_4B_instruct/ +logging_steps: 1 +save_steps: 250 +plot_loss: true +overwrite_output_dir: true +save_only_model: false +report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow] + +### train +per_device_train_batch_size: 1 +gradient_accumulation_steps: 16 +learning_rate: 1.0e-4 +num_train_epochs: 3.0 +lr_scheduler_type: cosine +warmup_ratio: 0.1 +bf16: true +ddp_timeout: 180000000 +resume_from_checkpoint: null + diff --git a/examples/EditScore-train/config/editscore_qwen3_vl_8B_instruct.yaml b/examples/EditScore-train/config/editscore_qwen3_vl_8B_instruct.yaml new file mode 100644 index 0000000000000000000000000000000000000000..66050d729d9885e10e162025ff58db570865e3c1 --- /dev/null +++ b/examples/EditScore-train/config/editscore_qwen3_vl_8B_instruct.yaml @@ -0,0 +1,41 @@ +model_name_or_path: Qwen/Qwen3-VL-8B-Instruct +image_max_pixels: 262144 +video_max_pixels: 16384 +trust_remote_code: true + +### method +stage: sft +do_train: true +finetuning_type: lora +lora_rank: 32 +lora_target: all + +### dataset +dataset: EditScore-Reward-Data +template: qwen3_vl +cutoff_len: 8192 +max_samples: 500000 +overwrite_cache: true +preprocessing_num_workers: 16 +dataloader_num_workers: 4 + +### output +output_dir: LLaMA-Factory/output/editscore_qwen3_vl_8B_instruct/ +logging_steps: 1 +save_steps: 250 +plot_loss: true +overwrite_output_dir: true +save_only_model: false +report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow] + +### train +per_device_train_batch_size: 1 +gradient_accumulation_steps: 16 +learning_rate: 1.0e-4 +num_train_epochs: 3.0 +lr_scheduler_type: cosine +warmup_ratio: 0.1 +bf16: true +ddp_timeout: 180000000 +resume_from_checkpoint: null + diff --git a/examples/EditScore-train/train.sh b/examples/EditScore-train/train.sh new file mode 100644 index 0000000000000000000000000000000000000000..0f416024504700a92b6c17af22987c0b22412e6b --- /dev/null +++ b/examples/EditScore-train/train.sh @@ -0,0 +1,52 @@ +#!/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER + +# Activate conda environment of LLaMA-Factory +conda activate llama-factory + + +RANK=0 +WORLD_SIZE=1 +MASTER_ADDR="localhost" +MASTER_PORT=29500 + + +while [[ $# -gt 0 ]]; do + case "$1" in + --rank=*) + RANK="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + --master_addr=*) + MASTER_ADDR="${1#*=}" + shift + ;; + --master_port=*) + MASTER_PORT="${1#*=}" + shift + ;; + *) + echo "Unknown parameter: $1" + exit 1 + ;; + esac +done + +name=experiment_name +log_dir="LLaMA-Factory/logs" + +log_file="${log_dir}/${name}_${RANK}.log" + + +CONFIG_YAML="LLaMA-Factory/examples/train_editscore/${name}.yaml" +FORCE_TORCHRUN=1 \ +NNODES=${WORLD_SIZE} \ +NODE_RANK=${RANK} \ +MASTER_ADDR=${MASTER_ADDR} \ +MASTER_PORT=${MASTER_PORT} \ +llamafactory-cli train ${CONFIG_YAML} 2>&1 | tee ${log_file} diff --git a/examples/OmniGen2-RL/.gitignore b/examples/OmniGen2-RL/.gitignore new file mode 100644 index 0000000000000000000000000000000000000000..0e44547238a96c7a5e9ccdccccd324b6e4b17ed9 --- /dev/null +++ b/examples/OmniGen2-RL/.gitignore @@ -0,0 +1,233 @@ +# Created by https://www.toptal.com/developers/gitignore/api/macos,python +# Edit at https://www.toptal.com/developers/gitignore?templates=macos,python + +### macOS ### +# General +.DS_Store +.AppleDouble +.LSOverride + +# Icon must end with two \r +Icon + + +# Thumbnails +._* + +# Files that might appear in the root of a volume +.DocumentRevisions-V100 +.fseventsd +.Spotlight-V100 +.TemporaryItems +.Trashes +.VolumeIcon.icns +.com.apple.timemachine.donotpresent + +# Directories potentially created on remote AFP share +.AppleDB +.AppleDesktop +Network Trash Folder +Temporary Items +.apdisk + +### macOS Patch ### +# iCloud generated files +*.icloud + +### Python ### +# Byte-compiled / optimized / DLL files +__pycache__/ +*.py[cod] +*$py.class + +# C extensions +*.so + +# Distribution / packaging +.Python +build/ +develop-eggs/ +dist/ +downloads/ +eggs/ +.eggs/ +lib/ +lib64/ +parts/ +sdist/ +var/ +wheels/ +share/python-wheels/ +*.egg-info/ +.installed.cfg +*.egg +MANIFEST + +# PyInstaller +# Usually these files are written by a python script from a template +# before PyInstaller builds the exe, so as to inject date/other infos into it. +*.manifest +*.spec + +# Installer logs +pip-log.txt +pip-delete-this-directory.txt + +# Unit test / coverage reports +htmlcov/ +.tox/ +.nox/ +.coverage +.coverage.* +.cache +nosetests.xml +coverage.xml +*.cover +*.py,cover +.hypothesis/ +.pytest_cache/ +cover/ + +# Translations +*.mo +*.pot + +# Django stuff: +*.log +local_settings.py +db.sqlite3 +db.sqlite3-journal + +# Flask stuff: +instance/ +.webassets-cache + +# Scrapy stuff: +.scrapy + +# Sphinx documentation +docs/_build/ + +# PyBuilder +.pybuilder/ +target/ + +# Jupyter Notebook +.ipynb_checkpoints + +# IPython +profile_default/ +ipython_config.py + +# pyenv +# For a library or package, you might want to ignore these files since the code is +# intended to run in multiple environments; otherwise, check them in: +# .python-version + +# pipenv +# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control. +# However, in case of collaboration, if having platform-specific dependencies or dependencies +# having no cross-platform support, pipenv may install dependencies that don't work, or not +# install all needed dependencies. +#Pipfile.lock + +# poetry +# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control. +# This is especially recommended for binary packages to ensure reproducibility, and is more +# commonly ignored for libraries. +# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control +#poetry.lock + +# pdm +# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control. +#pdm.lock +# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it +# in version control. +# https://pdm.fming.dev/#use-with-ide +.pdm.toml + +# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm +__pypackages__/ + +# Celery stuff +celerybeat-schedule +celerybeat.pid + +# SageMath parsed files +*.sage.py + +# Environments +.env +.venv +env/ +venv/ +ENV/ +env.bak/ +venv.bak/ + +# Spyder project settings +.spyderproject +.spyproject + +# Rope project settings +.ropeproject + +# mkdocs documentation +/site + +# mypy +.mypy_cache/ +.dmypy.json +dmypy.json + +# Pyre type checker +.pyre/ + +# pytype static type analyzer +.pytype/ + +# Cython debug symbols +cython_debug/ + +# PyCharm +# JetBrains specific template is maintained in a separate JetBrains.gitignore that can +# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore +# and can be added to the global gitignore or merged into this file. For a more nuclear +# option (not recommended) you can uncomment the following to ignore the entire idea folder. +#.idea/ + +### Python Patch ### +# Poetry local configuration file - https://python-poetry.org/docs/configuration/#local-configuration +poetry.toml + +# ruff +.ruff_cache/ + +# LSP config files +pyrightconfig.json + +# End of https://www.toptal.com/developers/gitignore/api/macos,python + +local_scripts/ + +omnigen2/utils/vpn_utils.py + +test_tokenizer.py +save_pipeline.py +app.sh +logs/ +results/ +test_jsonl* +pbs_files/ +convert_ckpt_to_pipeline.py +inference_test_efficiency.py +upload_pipeline* +example_images_resized/ +example_t2i_test_efficiency*.sh +example_edit_test_efficiency*.sh +example_in_context_generation_test_efficiency*.sh +intro* +resize_example_images.py +save_pipeline.py +outputs_gradio/* +test.py \ No newline at end of file diff --git a/examples/OmniGen2-RL/LICENSE b/examples/OmniGen2-RL/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..f49a4e16e68b128803cc2dcea614603632b04eac --- /dev/null +++ b/examples/OmniGen2-RL/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 [yyyy] [name of copyright owner] + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. \ No newline at end of file diff --git a/examples/OmniGen2-RL/README.md b/examples/OmniGen2-RL/README.md new file mode 100644 index 0000000000000000000000000000000000000000..1e5e1d128cbce15982d35b6b95bfba7c4ea3c2b4 --- /dev/null +++ b/examples/OmniGen2-RL/README.md @@ -0,0 +1,189 @@ +# ๐Ÿš€ Advanced Applications of EditScore + +Welcome to the examples directory! Here, we demonstrate how to leverage **EditScore** not just as an evaluation metric, but as a powerful component to actively *improve image editing models*. + +This guide covers two primary downstream applications: +1. **Best-of-N selection**: A simple, training-free method to instantly boost the output quality of any image editing model. +2. **Reinforcement Learning (RL) Fine-Tuning**: Using EditScore as a high-fidelity reward signal to train models for significantly better performance. + +## ๐Ÿ› ๏ธ Setup for Examples +The examples require libraries for RL, data handling, and potentially experiment tracking. +```bash +# Navigate to this directory if you are in the root +cd examples/OmniGen2-RL + +# Install the required packages +pip install -r requirements.txt +``` + +## Application 1: Best-of-N for Superior Outputs +Best-of-N is an elegant and powerful technique. Instead of generating a single output for a given instruction, you generate multiple (N) candidates and then use a highly accurate evaluatorโ€”**EditScore**โ€”to select the best one. + +This acts as a powerful "reranker" that filters out suboptimal results, significantly improving the perceived quality of the model without any extra training. + +### How to Use +We provide ready-to-use scripts to perform a full Best-of-N workflow on the GEdit-Bench benchmark. The following instructions use **OmniGen2** as the base model, but we provide similar scripts for **FLUX-Kontext** and **Qwen-Image-Edit** in the `evaluation/GEdit-Bench/` directory. + +**1. Generate Candidates** +```bash +bash evaluation/GEdit-Bench/omnigen2_16samples.sh # default using 8 GPUs +``` + +> **โš ๏ธ Important Note on Resource Usage** +> +> This process is computationally expensive and slow due to the large number of generations (16 samples per instruction). For reference, completing this step for OmniGen2 takes approximately **3 hours using 64 H100 GPUs**. + +
+๐Ÿ‘‰ Click here for tips on the usage of the script + +- **Distributed Inference**: Our scripts natively support multi-machine and multi-GPU execution. To run inference across 4 machines, for example, execute the following commands on each respective machine: +```bash +# On the first machine (rank 0) +bash evaluation/GEdit-Bench/omnigen2_16samples.sh --world_size 4 --rank 0 + +# On the second machine (rank 1) +bash evaluation/GEdit-Bench/omnigen2_16samples.sh --world_size 4 --rank 1 + +# ...and so on for ranks 2 and 3. +``` + +- **Monitoring Progress**: The scripts utilize nohup for background execution. We recommend monitoring the file (specified in the script file) to track the status and progress of the generation process. +
+ +**2. Score and Select** +Next, use EditScore to evaluate all N candidates and identify the one with the highest score. + +```bash +bash evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass1.sh # EditScore-7B, single pass +bash evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass4.sh # EditScore-7B, Avg@4 +``` + +**3. Evaluate the Final Selections** +Finally, evaluate the performance of the images selected by EditScore on GEdit-Bench to quantify the improvement. +```bash +bash evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass1_eval.sh +bash evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass4_eval.sh +``` + +By comparing these results to the baseline performance of the original model, you will see the benefits of applying EditScore as a reranker. + +## Application 2: Reinforcement Fine-Tuning +Beyond evaluation, **EditScore** can be used as a high-quality reward signal to fine-tune your image editing models using Reinforcement Learning (RL), leading to significantly improved performance. + +We employ the **FlowGRPO** algorithm, combining its strengths with EditScore's accurate, real-time feedback to create a powerful end-to-end fine-tuning pipeline. This process effectively guides the model toward generating better edits. + +### 1. Prepare Training Data +First, set up the dataset for RL fine-tuning. +1. Download the Data +Downlaod the official RL training data from [EditScore-RL-Data](https://huggingface.co/datasets/EditScore/EditScore-RL-Data). +2. Create Meta File +The uploaded dataset uses relative image paths. Run the following script to convert them to absolute paths based on your local environment: +```bash +# Then +python scripts/data/process_jsonl.py --input /path/to/EditScore-RL-Data/rl.jsonl --output /path/to/EditScore-RL-Data/rl_abs.jsonl --base-path /path/to/EditScore-RL-Data + +# Due to the limitation of base model (OmniGen2), we discard text change and portrait beautification, as these tasks harm RL training. +python scripts/data/extract_9_tasks.py --input_path /path/to/EditScore-RL-Data/rl_abs.jsonl --output_path /path/to/EditScore-RL-Data/rl_abs_9tasks.jsonl +``` +3. Configure the Data Path +Specify the path to your processed `.jsonl` file in the data configuration located at `data_configs/train/example/edit/all.yml`. +For example: +```yaml +ratio_type: inside_ratio + +data: + - + path: '/path/to/EditScore-RL-Data/rl_abs_9tasks.jsonl' # <-- Ensure this path is correct + type: 'edit' + ratio: !!float 1 +``` + +### 2. Prepare the Base Model (OmniGen2) +```bash +python scripts/misc/extract_bin_from_pipe.py +``` + +### 3. Launch the Reward Server +RL training requires a live reward signal. Before starting the training process, you must launch the **EditScore Reward Server**. This server will provide real-time scores for the generated images during training. + +Our reward server is built with two components: a **proxy** and one or more **reward servers**. The proxy receives requests from the training node, distributes them to the individual reward servers for computation, and then collects the results to send back. This architecture allows for easy scaling across multiple machines. + +We provide a convenient script to launch the entire server stack across multiple machines, assuming you have `ssh` access to all reward server nodes. + +```bash +# Launch EditScore-7B Reward Server +bash reward_server/start_multi_machines.sh --model_name=editscore_7B --config_path=reward_server/server_configs/editscore_7B.yml + +# Launch EditScore-7B (Avg@4) Reward Server +bash reward_server/start_multi_machines.sh --model_name=editscore_7B_pass4 --config_path=reward_server/server_configs/editscore_7B_pass4.yml + +# Launch EditScore-72B Reward Server +bash reward_server/start_multi_machines.sh --model_name=editscore_72B --config_path=reward_server/server_configs/editscore_72B.yml +``` + +> **โš ๏ธ Important Notes** +> +> * Before running the script, you **must** specify the IP addresses of your reward server machines in the corresponding `.yml` configuration file. +> * If you cannot use `ssh` to control the nodes, please refer to the logic in `reward_server/start_multi_machines.sh` to manually start the proxy and server processes on each machine. +> * You can monitor the status of the proxy and servers by checking the log files in the `reward_server/logs/` directory. + +## 3.5 (Optional) Reward Server Sanity Check +To ensure the reward server is configured correctly and running as expected, we provide a sanity check script. +```bash +python reward_server/scripts/utils/reward_server_sanity_check.py --config_path=reward_server/server_configs/editscore_7B.yml +``` +Once these steps are complete, your environment is ready to begin the reinforcement learning fine-tuning process. + +### 4. Start RL Fine-Tuning + +**Configure Training Parameters** +Before launching, you may need to adjust key parameters in the configuration file: `options/omnigen2_edit_rl_4machine_editscore7b_avg4.yml`. + +Here are some important settings: +- `train.global_batch_size`: The total number of images generated across all GPUs in a single sampling phase before the policy is updated. It is calculated as `num_unique_prompts_per_sampling * num_images_per_prompt`. +- `train.batch_size`: Batch size per GPU (`batch_size_per_forward * gradient_accumulation_steps * num_update_steps_per_sampling`) +- `train.rl.num_images_per_prompt`: The number of candidate images to generate for each unique prompt. +- `train.rl.num_unique_prompts_per_sampling`: The number of unique prompts in a global batch +- `train.rl.num_update_steps_per_sampling`: The number of gradient updates to perform in each sampling phase. Set this to `> 1` to enable off-policy RL, which improves sample efficiency. +- `train.rl.batch_size_per_forward`: Batch size for each forward pass. Together with `num_update_steps_per_sampling`, it defines the total number of samples processed per policy update. + +**Launch Distributed Training** +We provide scripts for both single and multi-machine distributed training based on **FSDP**. +```bash +# Single-machine training (8 GPUs) using EditScore-7B as the reward model +bash scripts/train/omnigen2_edit_rl_single_machine_editscore7b.sh + +# Multi-machine training (e.g., 4 machines with 8 GPUs each) using EditScore-7B (Avg@4) +bash scripts/train/omnigen2_edit_rl_4machine_editscore7b_avg4.sh +``` + +### 4. Training Outputs and Monitoring +All training artifacts, including logs and model checkpoints, are saved to the `experiments/` directory. +For transparency and to ensure full reproducibility, we provide the training curves of **OmniGen2-EditScore7B-v1.1** on [Weights & Biases (wandb)](https://wandb.ai/omnigen-rl/OmniGen2-RL/reports/Training-Curves-of-OmniGen2-EditScore7B-v1-1---VmlldzoxNDg2MTI2NQ?accessToken=467cbc4jupu1maluk00pan7z611m6xxmnwviwk8nblml3ydoy3j9fu92el6c0s8i). + + +### 5. Evaluate your RL Fine-Tuned Model +After training, you must convert the FSDP-saved checkpoint (`.bin`) into the standard Hugging Face format before you can use it for inference. + +#### Step 1: Convert the Checkpoint +We provide a script to automatically handle the conversion from the distributed FSDP format to the standard Hugging Face format (`.bin`). +Run the following command, replacing the arguments with your experiment's details: + +```shell +bash scripts/misc/convert_dist_ckpt_to_hf_format.sh [EXPERIMENT_NAME] [STEP_NUMBER] +``` +- [EXPERIMENT_NAME]: The name of your training experiment (e.g., omnigen2_edit_rl_single_machine_editscore7b). +- [STEP_NUMBER]: The specific training step of the checkpoint you wish to evaluate (e.g., 500). + +This will create a new directory containing the converted model weights in the standard format, ready for inference. + +#### Step 2: Run Evaluation on GEdit-Bench +Once the checkpoint is converted, you can benchmark its performance. We provide evaluation scripts tailored for GEdit-Bench. + +You can use our example scripts as a template. Simply copy one and modify the internal paths to point to your newly converted model checkpoint. +```shell +# Run evaluation for the converted model from step 500 +bash evaluation/GEdit-Bench/omnigen2.sh --experiment_name=omnigen2_edit_rl_4machine_editscore7b_avg4 --step=500 +bash evaluation/GEdit-Bench/omnigen2_eval.sh --experiment_name=omnigen2_edit_rl_4machine_editscore7b_avg4 --step=500 +``` +By comparing the results to the baseline model's performance, you can quantify the improvements achieved through RL fine-tuning with EditScore. \ No newline at end of file diff --git a/examples/OmniGen2-RL/data_configs/train/example/edit/all.yml b/examples/OmniGen2-RL/data_configs/train/example/edit/all.yml new file mode 100644 index 0000000000000000000000000000000000000000..f7d0a04c73167ee30efc893dd3d9b40063f6dde1 --- /dev/null +++ b/examples/OmniGen2-RL/data_configs/train/example/edit/all.yml @@ -0,0 +1,7 @@ +ratio_type: inside_ratio + +data: + - + path: '/path/to/EditScore-RL-Data/rl_abs_9tasks.jsonl' + type: 'edit' + ratio: !!float 1 \ No newline at end of file diff --git a/examples/OmniGen2-RL/data_configs/train/example/train.yml b/examples/OmniGen2-RL/data_configs/train/example/train.yml new file mode 100644 index 0000000000000000000000000000000000000000..d1340104c1d76e0a4017ba276db9ef99c1224110 --- /dev/null +++ b/examples/OmniGen2-RL/data_configs/train/example/train.yml @@ -0,0 +1,5 @@ +data: + - + path: 'data_configs/train/example/edit/all.yml' + type: 'edit' + ratio: !!float 1 \ No newline at end of file diff --git a/examples/OmniGen2-RL/docs/README.md b/examples/OmniGen2-RL/docs/README.md new file mode 100644 index 0000000000000000000000000000000000000000..ebb337145b54a8b77632e7ea523201c252fc33b9 --- /dev/null +++ b/examples/OmniGen2-RL/docs/README.md @@ -0,0 +1,2 @@ +## Apply EditScore to Image Editing +### Best-of-N selection** \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/calculate_statistics.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/calculate_statistics.py new file mode 100644 index 0000000000000000000000000000000000000000..21bddda0330f1204ea5e2bcf9e0f32048d14494c --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/calculate_statistics.py @@ -0,0 +1,223 @@ +import os +import pandas as pd +from collections import defaultdict +import sys +import numpy as np +import math + +GROUPS = [ + "background_change", + "color_alter", + "style_change", + "subject-add", + "subject-remove", + "subject-replace", + "material_alter", + "motion_change", + "ps_human", + "text_change", + "tone_transfer", +] + +GROUPS2 = [ + "background_change", + "color_alter", + "material_alter", + "motion_change", + "ps_human", + "style_change", + "subject-add", + "subject-remove", + "subject-replace", + "text_change", + "tone_transfer", +] + +def analyze_scores(result_dir, language, num_samples): # ่ฟ™ไบ› group_scores ๅญ—ๅ…ธ็”จไบŽๅญ˜ๅ‚จๆฏไธช group ็š„ๆœ€็ปˆๅนณๅ‡ๅˆ† + group_scores_semantics = {} + group_scores_quality = {} + group_scores_overall = {} + group_scores_semantics_intersection = {} + group_scores_quality_intersection = {} + group_scores_overall_intersection = {} + + # ๅค–้ƒจๅพช็Žฏ๏ผŒๅค„็†ๆฏไธ€ไธช group + for group_name in GROUPS: + data_point_samples = defaultdict(list) + + # ๅพช็Žฏ่ฏปๅ– num_samples ไธช่ฏ„ๅˆ†ๆ–‡ไปถ + for turn in range(num_samples): + csv_path = os.path.join(result_dir, f"{group_name}_gpt_score{'_sample' + str(turn) if turn > 0 else ''}.csv") + if not os.path.exists(csv_path): + print(f"Warning: File not found, skipping: {csv_path}") + continue + + with open(csv_path, 'r') as f: + df = pd.read_csv(f) + + for _, row in df.iterrows(): + # ่ฟ‡ๆปค่ฏญ่จ€ + if row['instruction_language'] != language: + continue + + # ๅฎšไน‰ๅ”ฏไธ€ๆ ‡่ฏ†็ฌฆ + unique_key = os.path.basename(row['source_image']).split('_SRCIMG')[0] + + # ่ฎก็ฎ— overall_score + semantics_score = row['sementics_score'] + quality_score = row['quality_score'] + overall_score = math.sqrt(semantics_score * quality_score) + + # ๅฐ†ๅฝ“ๅ‰ๆ ทๆœฌ็š„ๅˆ†ๆ•ฐไฟกๆฏๅญ˜ๅ…ฅๅญ—ๅ…ธ + sample_data = { + 'semantics_score': semantics_score, + 'quality_score': quality_score, + 'overall_score': overall_score, + 'intersection_exist': row['intersection_exist'] + } + + # ๆŒ‰ๅ”ฏไธ€ๆ ‡่ฏ†็ฌฆ่šๅˆๆ‰€ๆœ‰ๆ ทๆœฌ + data_point_samples[unique_key].append(sample_data) + + # --- ๆ ธๅฟƒๆ”นๅŠจ้ƒจๅˆ†๏ผš็ฌฌไบŒ้˜ถๆฎต - ็ญ›้€‰ไธŽ่ฎก็ฎ— --- + # ็Žฐๅœจ data_point_samples ๅทฒ็ปๆ”ถ้›†ไบ†ๆ‰€ๆœ‰ๆต‹่ฏ•้กน็š„ๆ‰€ๆœ‰ๆ ทๆœฌๆ•ฐๆฎใ€‚ + # ๆˆ‘ไปฌ้œ€่ฆ้ๅކๅฎƒ๏ผŒไธบๆฏไธชๆต‹่ฏ•้กนๆ‰พๅˆฐๆœ€ไฝณๆ ทๆœฌ๏ผŒ็„ถๅŽๅฐ†ๆœ€ไฝณๅˆ†ๆ•ฐๅญ˜ๅ…ฅๆœ€็ปˆๅˆ—่กจใ€‚ + + best_semantics_scores = [] + best_quality_scores = [] + best_overall_scores = [] + + for unique_key, samples in data_point_samples.items(): + if not samples: + continue + + # ไปŽๅฝ“ๅ‰ๆต‹่ฏ•้กน็š„ๆ‰€ๆœ‰ๆ ทๆœฌไธญ๏ผŒๆ‰พๅˆฐ overall_score ๆœ€้ซ˜็š„้‚ฃไธช + # max() ๅ‡ฝๆ•ฐ็š„ key ๅ‚ๆ•ฐๅฏไปฅ่ฎฉๆˆ‘ไปฌๆŒ‡ๅฎšๆŒ‰ๅญ—ๅ…ธไธญ็š„ๅ“ชไธชๅ€ผๆฅๆฏ”่พƒ + best_sample = max(samples, key=lambda s: s['overall_score']) + + # ๅฐ†่ฟ™ไธชๆœ€ไฝณๆ ทๆœฌ็š„ๅˆ†ๆ•ฐๆทปๅŠ ๅˆฐๆœ€็ปˆๅˆ—่กจไธญ + best_semantics_scores.append(best_sample['semantics_score']) + best_quality_scores.append(best_sample['quality_score']) + best_overall_scores.append(best_sample['overall_score']) + + + group_scores_semantics[group_name] = np.mean(best_semantics_scores) + group_scores_quality[group_name] = np.mean(best_quality_scores) + group_scores_overall[group_name] = np.mean(best_overall_scores) + + print("\n--- Overall Model Averages ---") + + print("\nSemantics:") + model_scores = [group_scores_semantics[group] for group in GROUPS] + model_avg = np.mean(model_scores) + group_scores_semantics["avg_semantics"] = model_avg + + # print("\nSemantics Valid Num:") + # model_scores = [group_scores_semantics_valid_num[group] for group in GROUPS] + # model_avg = np.mean(model_scores) + # group_scores_semantics_valid_num["avg_semantics_valid_num"] = model_avg + + # print("\nSemantics Intersection:") + # model_scores = [group_scores_semantics_intersection[group] for group in GROUPS] + # model_avg = np.mean(model_scores) + # group_scores_semantics_intersection["avg_semantics"] = model_avg + + # print("\nSemantics Valid Num Intersection:") + # model_scores = [group_scores_semantics_valid_num_intersection[group] for group in GROUPS] + # model_avg = np.mean(model_scores) + # group_scores_semantics_valid_num_intersection["avg_semantics_valid_num"] = model_avg + + print("\nQuality:") + model_scores = [group_scores_quality[group] for group in GROUPS] + model_avg = np.mean(model_scores) + group_scores_quality["avg_quality"] = model_avg + + # print("\nQuality Valid Num:") + # model_scores = [group_scores_quality_valid_num[group] for group in GROUPS] + # model_avg = np.mean(model_scores) + # group_scores_quality_valid_num["avg_quality_valid_num"] = model_avg + + # print("\nQuality Intersection:") + # model_scores = [group_scores_quality_intersection[group] for group in GROUPS] + # model_avg = np.mean(model_scores) + # group_scores_quality_intersection["avg_quality"] = model_avg + + # print("\nQuality Valid Num Intersection:") + # model_scores = [group_scores_quality_valid_num_intersection[group] for group in GROUPS] + # model_avg = np.mean(model_scores) + # group_scores_quality_valid_num_intersection["avg_quality_valid_num"] = model_avg + + print("\nOverall:") + model_scores = [group_scores_overall[group] for group in GROUPS] + model_avg = np.mean(model_scores) + group_scores_overall["avg_overall"] = model_avg + + # print("\nOverall Valid Num:") + # model_scores = [group_scores_overall_valid_num[group] for group in GROUPS] + # model_avg = np.mean(model_scores) + # group_scores_overall_valid_num["avg_overall_valid_num"] = model_avg + + + return ( + group_scores_semantics, + group_scores_quality, + group_scores_overall, + # group_scores_semantics_valid_num, + # group_scores_quality_valid_num, + # group_scores_overall_valid_num + ) + +if __name__ == "__main__": + import argparse + parser = argparse.ArgumentParser() + parser.add_argument("--result_dir", type=str, default="/results/") + parser.add_argument("--language", type=str, default="en", choices=["en", "cn"]) + parser.add_argument("--num_samples", type=int, default=1) + parser.add_argument("--groups", type=str, default="GROUPS", choices=["GROUPS", "GROUPS2"]) + args = parser.parse_args() + result_dir = args.result_dir + + # result_dir = os.path.join(result_dir, "viescore") + + print("\nOverall:") + + ( + group_scores_semantics, + group_scores_quality, + group_scores_overall, + # group_scores_semantics_valid_num, + # group_scores_quality_valid_num, + # group_scores_overall_valid_num + ) = analyze_scores(result_dir, language=args.language, num_samples=args.num_samples) + + if args.groups == "GROUPS": + groups = GROUPS + else: + groups = GROUPS2 + + for group_name in groups: + print(f"{group_name}: {group_scores_semantics[group_name]:.2f}, {group_scores_quality[group_name]:.2f}, {group_scores_overall[group_name]:.2f}") + + print(f"Average: {group_scores_semantics['avg_semantics']:.2f}, {group_scores_quality['avg_quality']:.2f}, {group_scores_overall['avg_overall']:.2f}") + + print("Semantics: " + " & ".join([f"{group_scores_semantics[group_name]:.2f}" for group_name in groups] + [f"{group_scores_semantics['avg_semantics']:.2f}"])) + print("Quality: " + " & ".join([f"{group_scores_quality[group_name]:.2f}" for group_name in groups] + [f"{group_scores_quality['avg_quality']:.2f}"])) + print("Overall: " + " & ".join([f"{group_scores_overall[group_name]:.2f}" for group_name in groups] + [f"{group_scores_overall['avg_overall']:.2f}"])) + + # print("\nValid Num:") + # for group_name in GROUPS: + # print(f"{group_name}: {group_scores_semantics_valid_num[group_name]:.2f}, {group_scores_quality_valid_num[group_name]:.2f}, {group_scores_overall_valid_num[group_name]:.2f}") + + # print(f"Average Valid Num: {group_scores_semantics_valid_num['avg_semantics_valid_num']:.2f}, {group_scores_quality_valid_num['avg_quality_valid_num']:.2f}, {group_scores_overall_valid_num['avg_overall_valid_num']:.2f}") + + # print("\nIntersection:") + # for group_name in GROUPS: + # print(f"{group_name}: {group_scores_semantics_intersection[group_name]:.2f}, {group_scores_quality_intersection[group_name]:.2f}, {group_scores_overall_intersection[group_name]:.2f}") + + # print(f"Average Intersection: {group_scores_semantics_intersection['avg_semantics']:.2f}, {group_scores_quality_intersection['avg_quality']:.2f}, {group_scores_overall_intersection['avg_overall']:.2f}") + + # print("\nValid Num Intersection:") + # for group_name in GROUPS: + # print(f"{group_name}: {group_scores_semantics_valid_num_intersection[group_name]:.2f}, {group_scores_quality_valid_num_intersection[group_name]:.2f}, {group_scores_overall_valid_num_intersection[group_name]:.2f}") + + # print(f"Average Valid Num Intersection: {group_scores_semantics_valid_num_intersection['avg_semantics_valid_num']:.2f}, {group_scores_quality_valid_num_intersection['avg_quality_valid_num']:.2f}, {group_scores_overall_valid_num_intersection['avg_overall_valid_num']:.2f}") diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples.sh new file mode 100644 index 0000000000000000000000000000000000000000..1040d21b411da6481fa6aca09af0a9a68b78bfed --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples.sh @@ -0,0 +1,83 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +RANK=0 +MASTER_ADDR=1 +MASTER_PORT=29500 +WORLD_SIZE=1 + +# ๅค„็†ๅ‘ฝๅๅ‚ๆ•ฐ +while [[ $# -gt 0 ]]; do + case "$1" in + --rank=*) + RANK="${1#*=}" + shift + ;; + --master_addr=*) + MASTER_ADDR="${1#*=}" + shift + ;; + --master_port=*) + MASTER_PORT="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + *) + echo "ๆœช็Ÿฅๅ‚ๆ•ฐ: $1" + shift + ;; + esac +done + +# ่พ“ๅ‡บ้…็ฝฎ +echo "RANK: $RANK" +echo "MASTER_ADDR: $MASTER_ADDR" +echo "MASTER_PORT: $MASTER_PORT" +echo "WORLD_SIZE: $WORLD_SIZE" + +global_shift_index=0 +total_num_images=606 + +num_gpus_per_machine=$(python -c "import torch; print(torch.cuda.device_count())") +# Calculate images per machine, rounding up to ensure all data is covered +num_images_per_machine=$(( (total_num_images + WORLD_SIZE - 1) / WORLD_SIZE )) +shift_index=$((RANK * num_images_per_machine)) + +if [ $((total_num_images - shift_index)) -lt $num_images_per_machine ]; then + num_images_per_machine=$((total_num_images - shift_index)) +fi + +# Calculate base number of images per GPU (for first 7 GPUs) +num_images_per_gpu=$(( (num_images_per_machine + num_gpus_per_machine - 1) / num_gpus_per_machine )) + +guidance_scale=2.5 + +for ((i=0; i logs/gedit_FLUX-Kontext-dev_gs${guidance_scale}_16samples_${start_idx}_${end_idx}.log 2>&1 & +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples_select_best_editscore_pass1.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples_select_best_editscore_pass1.sh new file mode 100644 index 0000000000000000000000000000000000000000..30933ac44d600f8479fe426e7a93e5e38c94e034 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples_select_best_editscore_pass1.sh @@ -0,0 +1,89 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +RANK=0 +MASTER_ADDR=1 +MASTER_PORT=29500 +WORLD_SIZE=1 + +# ๅค„็†ๅ‘ฝๅๅ‚ๆ•ฐ +while [[ $# -gt 0 ]]; do + case "$1" in + --rank=*) + RANK="${1#*=}" + shift + ;; + --master_addr=*) + MASTER_ADDR="${1#*=}" + shift + ;; + --master_port=*) + MASTER_PORT="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + *) + echo "ๆœช็Ÿฅๅ‚ๆ•ฐ: $1" + shift + ;; + esac +done + +# ่พ“ๅ‡บ้…็ฝฎ +echo "RANK: $RANK" +echo "MASTER_ADDR: $MASTER_ADDR" +echo "MASTER_PORT: $MASTER_PORT" +echo "WORLD_SIZE: $WORLD_SIZE" + +global_shift_index=0 +total_num_images=606 + +num_gpus_per_machine=$(python -c "import torch; print(torch.cuda.device_count())") +# Calculate images per machine, rounding up to ensure all data is covered +num_images_per_machine=$(( (total_num_images + WORLD_SIZE - 1) / WORLD_SIZE )) +shift_index=$((RANK * num_images_per_machine)) + +if [ $((total_num_images - shift_index)) -lt $num_images_per_machine ]; then + num_images_per_machine=$((total_num_images - shift_index)) +fi + +# Calculate base number of images per GPU (for first 7 GPUs) +num_images_per_gpu=$(( (num_images_per_machine + num_gpus_per_machine - 1) / num_gpus_per_machine )) + +guidance_scale=2.5 + +for ((i=0; i logs/gedit_FLUX-Kontext-dev_gs${guidance_scale}_16samples_select_best_pass1_${start_idx}_${end_idx}.log 2>&1 & +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples_select_best_editscore_pass1_eval.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples_select_best_editscore_pass1_eval.sh new file mode 100644 index 0000000000000000000000000000000000000000..62fb31dd74c988608eb1a9598db3dde7dc14f283 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples_select_best_editscore_pass1_eval.sh @@ -0,0 +1,28 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +guidance_scale=2.5 + +best=( +1 +2 +4 +8 +16 +) + +for b in "${best[@]}" +do + accelerate launch --num_processes 1 evaluation/GEdit-Bench/test_gedit_score.py \ + --result_dir evaluation/GEdit-Bench/results/FLUX-Kontext-dev/results_gs${guidance_scale}_16samples_pass1_best${b} \ + --backbone gpt-4.1 \ + --openai_url https://api.openai.com/v1/chat/completions \ + --max_workers 30 \ + --key PUT-YOUR-KEY-HERE + + python evaluation/GEdit-Bench/calculate_statistics.py \ + --result_dir evaluation/GEdit-Bench/results/FLUX-Kontext-dev/results_gs${guidance_scale}_16samples_pass1_best${b}/viescore_gpt-4.1 \ + --language en +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples_select_best_editscore_pass4.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples_select_best_editscore_pass4.sh new file mode 100644 index 0000000000000000000000000000000000000000..72754904bd8bbc6592e478b10c88d4acfa95e124 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples_select_best_editscore_pass4.sh @@ -0,0 +1,89 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +RANK=0 +MASTER_ADDR=1 +MASTER_PORT=29500 +WORLD_SIZE=1 + +# ๅค„็†ๅ‘ฝๅๅ‚ๆ•ฐ +while [[ $# -gt 0 ]]; do + case "$1" in + --rank=*) + RANK="${1#*=}" + shift + ;; + --master_addr=*) + MASTER_ADDR="${1#*=}" + shift + ;; + --master_port=*) + MASTER_PORT="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + *) + echo "ๆœช็Ÿฅๅ‚ๆ•ฐ: $1" + shift + ;; + esac +done + +# ่พ“ๅ‡บ้…็ฝฎ +echo "RANK: $RANK" +echo "MASTER_ADDR: $MASTER_ADDR" +echo "MASTER_PORT: $MASTER_PORT" +echo "WORLD_SIZE: $WORLD_SIZE" + +global_shift_index=0 +total_num_images=606 + +num_gpus_per_machine=$(python -c "import torch; print(torch.cuda.device_count())") +# Calculate images per machine, rounding up to ensure all data is covered +num_images_per_machine=$(( (total_num_images + WORLD_SIZE - 1) / WORLD_SIZE )) +shift_index=$((RANK * num_images_per_machine)) + +if [ $((total_num_images - shift_index)) -lt $num_images_per_machine ]; then + num_images_per_machine=$((total_num_images - shift_index)) +fi + +# Calculate base number of images per GPU (for first 7 GPUs) +num_images_per_gpu=$(( (num_images_per_machine + num_gpus_per_machine - 1) / num_gpus_per_machine )) + +guidance_scale=2.5 + +for ((i=0; i logs/gedit_FLUX-Kontext-dev_gs${guidance_scale}_16samples_select_best_pass4_${start_idx}_${end_idx}.log 2>&1 & +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples_select_best_editscore_pass4_eval.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples_select_best_editscore_pass4_eval.sh new file mode 100644 index 0000000000000000000000000000000000000000..2d0d9819586ee257060e7e2445175b81a8831948 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples_select_best_editscore_pass4_eval.sh @@ -0,0 +1,28 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +guidance_scale=2.5 + +best=( +1 +2 +4 +8 +16 +) + +for b in "${best[@]}" +do + accelerate launch --num_processes 1 evaluation/GEdit-Bench/test_gedit_score.py \ + --result_dir evaluation/GEdit-Bench/results/FLUX-Kontext-dev/results_gs${guidance_scale}_16samples_pass4_best${b} \ + --backbone gpt-4.1 \ + --openai_url https://api.openai.com/v1/chat/completions \ + --max_workers 30 \ + --key PUT-YOUR-KEY-HERE + + python evaluation/GEdit-Bench/calculate_statistics.py \ + --result_dir evaluation/GEdit-Bench/results/FLUX-Kontext-dev/results_gs${guidance_scale}_16samples_pass4_best${b}/viescore_gpt-4.1 \ + --language en +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/inference.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..d26cc18efb9ab4376d56388c0b6844b08698cd14 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/inference.py @@ -0,0 +1,426 @@ +import dotenv + +dotenv.load_dotenv(override=True) + +import argparse +import os +import sys +from typing import List, Tuple + +from PIL import Image, ImageOps + +from omegaconf import OmegaConf +from tqdm import tqdm + +import torch +from torchvision.transforms.functional import to_pil_image, to_tensor + +from accelerate import Accelerator +from accelerate import init_empty_weights + +from datasets import load_dataset + +from transformers import AutoProcessor, AutoModelForVision2Seq + +from diffusers.models.autoencoders.autoencoder_kl import AutoencoderKL +from diffusers.hooks import apply_group_offloading + +sys.path.append(os.path.join(os.path.dirname(__file__), os.path.pardir, os.path.pardir)) + +from omnigen2.pipelines.omnigen2.pipeline_omnigen2 import OmniGen2Pipeline +from omnigen2.models.transformers.transformer_omnigen2 import OmniGen2Transformer2DModel +from omnigen2.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler + + +def parse_args(root_dir: str) -> argparse.Namespace: + """Parse command line arguments.""" + parser = argparse.ArgumentParser(description="OmniGen2 image generation script.") + parser.add_argument( + "--load_from_pipeline", + action="store_true", + help="Load from pipeline.", + ) + parser.add_argument( + "--pipeline_path", + type=str, + default=None, + ) + parser.add_argument( + "--experiment_name", + type=str, + default=None, + help="Name of experiment.", + ) + parser.add_argument( + "--model_path", + type=str, + default=None, + help="Path to model checkpoint.", + ) + parser.add_argument( + "--transformer_lora_path", + type=str, + default=None, + help="Path to transformer LoRA weights.", + ) + parser.add_argument( + "--scheduler", + type=str, + default="euler", + choices=["euler", "euler_maruyama", "dpmsolver++"], + help="Scheduler to use.", + ) + parser.add_argument( + "--num_inference_step", + type=int, + default=50, + help="Number of inference steps." + ) + parser.add_argument( + "--seed", + type=int, + default=0, + help="Random seed for generation." + ) + parser.add_argument( + "--height", + type=int, + default=1024, + help="Output image height." + ) + parser.add_argument( + "--width", + type=int, + default=1024, + help="Output image width." + ) + parser.add_argument( + "--max_input_image_pixels", + type=int, + default=1048576, + help="Maximum number of pixels for each input image." + ) + parser.add_argument( + "--dtype", + type=str, + default='bf16', + choices=['fp32', 'fp16', 'bf16'], + help="Data type for model weights." + ) + parser.add_argument( + "--text_guidance_scale", + type=float, + default=5.0, + help="Text guidance scale." + ) + parser.add_argument( + "--image_guidance_scale", + type=float, + default=2.0, + help="Image guidance scale." + ) + parser.add_argument( + "--cfg_range_start", + type=float, + default=0.0, + help="Start of the CFG range." + ) + parser.add_argument( + "--cfg_range_end", + type=float, + default=1.0, + help="End of the CFG range." + ) + parser.add_argument( + "--negative_prompt", + type=str, + default="(((deformed))), blurry, over saturation, bad anatomy, disfigured, poorly drawn face, mutation, mutated, (extra_limb), (ugly), (poorly drawn hands), fused fingers, messy drawing, broken legs censor, censored, censor_bar", + help="Negative prompt for generation." + ) + parser.add_argument( + "--result_dir", + type=str, + required=True, + help="Path to save results." + ) + parser.add_argument( + "--enable_model_cpu_offload", + action="store_true", + help="Enable model CPU offload." + ) + parser.add_argument( + "--enable_sequential_cpu_offload", + action="store_true", + help="Enable sequential CPU offload." + ) + parser.add_argument( + "--enable_group_offload", + action="store_true", + help="Enable group offload." + ) + parser.add_argument( + "--time_shift_base_res", + type=int, + default=320, + help="Time shift base resolution." + ) + parser.add_argument( + "--start_index", + type=int, + default=0, + ) + parser.add_argument( + "--end_index", + type=int, + default=1212, + ) + parser.add_argument( + "--use_ori_neg_prompt_template", + action="store_true", + help="Use original negative prompt template." + ) + parser.add_argument( + "--num_samples", + type=int, + default=1, + help="Number of samples to generate." + ) + parser.add_argument( + "--root_dir", + type=str, + default=None + ) + args = parser.parse_args() + + if args.root_dir is None: + args.root_dir = root_dir + return args + +def load_pipeline(args: argparse.Namespace, accelerator: Accelerator, weight_dtype: torch.dtype) -> OmniGen2Pipeline: + if args.load_from_pipeline: + pipeline = OmniGen2Pipeline.from_pretrained( + args.pipeline_path, + torch_dtype=weight_dtype, + trust_remote_code=True, + local_files_only=True, + ) + pipeline.transformer = OmniGen2Transformer2DModel.from_pretrained( + args.pipeline_path, + subfolder="transformer", + torch_dtype=weight_dtype, + local_files_only=True, + ) + else: + experiment_name = args.experiment_name + experiment_dir = os.path.join(args.root_dir, 'experiments', experiment_name) + + conf = OmegaConf.load(os.path.join(experiment_dir, f"{experiment_name}.yml")) + + with init_empty_weights(): + transformer = OmniGen2Transformer2DModel(**conf.model.arch_opt) + + state_dict = torch.load(os.path.join(experiment_dir, args.model_path), mmap=True, weights_only=True) + state_dict = torch.load(os.path.join(experiment_dir, args.model_path), mmap=True, weights_only=True) + missing, unexpect = transformer.load_state_dict(state_dict, assign=True, strict=False) + if len(missing) > 0 or len(unexpect) > 0: + print(f"missed parameters: {missing}") + print(f"unexpected parameters: {unexpect}") + + vae = AutoencoderKL.from_pretrained("black-forest-labs/FLUX.1-dev", subfolder="vae") + + mllm = AutoModelForVision2Seq.from_pretrained("Qwen/Qwen2.5-VL-3B-Instruct") + processor = AutoProcessor.from_pretrained("Qwen/Qwen2.5-VL-3B-Instruct") + + pipeline = OmniGen2Pipeline( + transformer=transformer, + vae=vae, + mllm=mllm, + processor=processor, + scheduler=FlowMatchEulerDiscreteScheduler(), + ) + + if args.transformer_lora_path: + print(f"LoRA weights loaded from {args.transformer_lora_path}") + pipeline.load_lora_weights(args.transformer_lora_path, + weight_name="pytorch_lora_weights.safetensors", + local_files_only=True) + + if args.scheduler == "dpmsolver++": + from omnigen2.schedulers.scheduling_dpmsolver_multistep import DPMSolverMultistepScheduler + scheduler = DPMSolverMultistepScheduler( + algorithm_type="dpmsolver++", + solver_type="midpoint", + solver_order=2, + prediction_type="flow_prediction", + ) + pipeline.scheduler = scheduler + elif args.scheduler == "euler_maruyama": + from omnigen2.schedulers.scheduling_flow_match_euler_maruyama_discrete import FlowMatchEulerMaruyamaDiscreteScheduler + scheduler = FlowMatchEulerMaruyamaDiscreteScheduler( + num_train_timesteps=args.num_inference_step, + sigma_schedule="v3" + ) + pipeline.scheduler = scheduler + elif args.scheduler == "euler": + scheduler = FlowMatchEulerDiscreteScheduler( + num_train_timesteps=args.num_inference_step, + time_shift_base_res=args.time_shift_base_res + ) + pipeline.scheduler = scheduler + + if args.enable_sequential_cpu_offload: + pipeline.enable_sequential_cpu_offload() + elif args.enable_model_cpu_offload: + pipeline.enable_model_cpu_offload() + elif args.enable_group_offload: + apply_group_offloading(pipeline.transformer, onload_device=accelerator.device, offload_type="block_level", num_blocks_per_group=2, use_stream=True) + apply_group_offloading(pipeline.mllm, onload_device=accelerator.device, offload_type="block_level", num_blocks_per_group=2, use_stream=True) + apply_group_offloading(pipeline.vae, onload_device=accelerator.device, offload_type="block_level", num_blocks_per_group=2, use_stream=True) + else: + pipeline = pipeline.to(device=accelerator.device) + pipeline = pipeline.to(dtype=weight_dtype) + return pipeline + + +def run(args: argparse.Namespace, + accelerator: Accelerator, + pipeline: OmniGen2Pipeline, + instruction: str, + negative_prompt: str, + input_images: List[Image.Image], + target_img_size: Tuple[int, int], + seed: int) -> Image.Image: + """Run the image generation pipeline with the given parameters.""" + generator = torch.Generator(device=accelerator.device).manual_seed(seed) + + if args.use_ori_neg_prompt_template: + negative_prompt = [ + { + "role": "system", + "content": "You are a helpful assistant.", + }, + {"role": "user", "content": negative_prompt}, + ] + negative_prompt = pipeline.processor.tokenizer.apply_chat_template( + negative_prompt, tokenize=False, add_generation_prompt=False + ) + + negative_prompt_embeds, negative_prompt_attention_mask = pipeline._get_qwen2_prompt_embeds( + prompt=negative_prompt, device=accelerator.device, max_sequence_length=1024 + ) + + results = pipeline( + prompt=[instruction], + input_images=[input_images], + size=[(target_img_size[0], target_img_size[1])], + num_inference_steps=args.num_inference_step, + max_sequence_length=1024, + text_guidance_scale=args.text_guidance_scale, + image_guidance_scale=args.image_guidance_scale, + cfg_range=(args.cfg_range_start, args.cfg_range_end), + negative_prompt_embeds=negative_prompt_embeds, + negative_prompt_attention_mask=negative_prompt_attention_mask, + num_images_per_prompt=1, + generator=generator, + output_type="pil", + ) + else: + results = pipeline( + prompt=[instruction], + input_images=[input_images], + size=[(target_img_size[0], target_img_size[1])], + num_inference_steps=args.num_inference_step, + max_sequence_length=1024, + text_guidance_scale=args.text_guidance_scale, + image_guidance_scale=args.image_guidance_scale, + cfg_range=(args.cfg_range_start, args.cfg_range_end), + negative_prompt=negative_prompt, + num_images_per_prompt=1, + generator=generator, + output_type="pil", + ) + return results + +def create_collage(images: List[torch.Tensor]) -> Image.Image: + """Create a horizontal collage from a list of images.""" + max_height = max(img.shape[-2] for img in images) + total_width = sum(img.shape[-1] for img in images) + canvas = torch.zeros((3, max_height, total_width), device=images[0].device) + + current_x = 0 + for img in images: + h, w = img.shape[-2:] + canvas[:, :h, current_x:current_x+w] = img * 0.5 + 0.5 + current_x += w + + return to_pil_image(canvas) + +def main(args: argparse.Namespace, root_dir: str) -> None: + """Main function to run the image generation process.""" + # Initialize accelerator + accelerator = Accelerator(mixed_precision=args.dtype if args.dtype != 'fp32' else 'no') + + # Set weight dtype + weight_dtype = torch.float32 + if args.dtype == 'fp16': + weight_dtype = torch.float16 + elif args.dtype == 'bf16': + weight_dtype = torch.bfloat16 + + # Load pipeline and process inputs + pipeline = load_pipeline(args, accelerator, weight_dtype) + pipeline.set_progress_bar_config(disable=True) + + test_dataset = load_dataset("stepfun-ai/GEdit-Bench", split='train') + + filtered_test_dataset = [ + item for item in test_dataset + if item['instruction_language'] != 'cn' + ] + test_dataset = filtered_test_dataset + + process_index = Accelerator().process_index + + data_index = list(range(args.start_index, args.end_index)) + + with tqdm( + total=len(data_index), + desc=f"process_index {process_index}: Processing {len(data_index)}/{len(test_dataset)}", + unit="image", + disable=not accelerator.is_main_process, + ) as pbar: + + for idx in data_index: + data_item = test_dataset[idx] + + task_type = data_item['task_type'] + instruction_language = data_item['instruction_language'] + + key = data_item['key'] + instruction = data_item['instruction'] + input_image = data_item['input_image'] + + ori_img_size = input_image.size + new_img_size = (ori_img_size[0] // 16 * 16, ori_img_size[1] // 16 * 16) + input_images = [input_image.resize(new_img_size)] + + for turn in range(args.num_samples): + results = run(args, accelerator, pipeline, instruction, args.negative_prompt, input_images, new_img_size, args.seed + idx + turn * len(test_dataset)) + output_image = results.images[0] + output_image = output_image.resize(ori_img_size) + + sub_dir = os.path.join(args.result_dir, "fullset", task_type, instruction_language) + os.makedirs(sub_dir, exist_ok=True) + + if turn > 0: + output_image.save(os.path.join(sub_dir, f"{key}_sample{turn}.png")) + else: + input_image.save(os.path.join(sub_dir, f"{key}_SRCIMG.png")) + output_image.save(os.path.join(sub_dir, f"{key}.png")) + + pbar.update(1) + +if __name__ == "__main__": + root_dir = os.path.abspath(os.path.join(__file__, os.path.pardir, os.path.pardir, os.path.pardir)) + args = parse_args(root_dir) + main(args, root_dir) \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/inference_flux_kontext_dev.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/inference_flux_kontext_dev.py new file mode 100644 index 0000000000000000000000000000000000000000..629fe3f9ce45aab665e100a87fe6ee15b853c652 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/inference_flux_kontext_dev.py @@ -0,0 +1,247 @@ +import dotenv + +dotenv.load_dotenv(override=True) + +import argparse +import os +import sys +from typing import List, Tuple + +from PIL import Image, ImageOps + +from omegaconf import OmegaConf +from tqdm import tqdm + +import torch +from torchvision.transforms.functional import to_pil_image, to_tensor + +from accelerate import Accelerator +from accelerate import init_empty_weights + +from datasets import load_from_disk, load_dataset + +from diffusers import FluxKontextPipeline + +sys.path.append(os.path.join(os.path.dirname(__file__), os.path.pardir, os.path.pardir)) + +def parse_args(root_dir: str) -> argparse.Namespace: + """Parse command line arguments.""" + parser = argparse.ArgumentParser(description="OmniGen2 image generation script.") + parser.add_argument( + "--load_from_pipeline", + action="store_true", + help="Load from pipeline.", + ) + parser.add_argument( + "--pipeline_path", + type=str, + default=None, + ) + parser.add_argument( + "--experiment_name", + type=str, + default=None, + help="Name of experiment.", + ) + parser.add_argument( + "--model_path", + type=str, + default=None, + help="Path to model checkpoint.", + ) + parser.add_argument( + "--transformer_lora_path", + type=str, + default=None, + help="Path to transformer LoRA weights.", + ) + parser.add_argument( + "--scheduler", + type=str, + default="euler", + choices=["euler", "euler_maruyama", "dpmsolver++"], + help="Scheduler to use.", + ) + parser.add_argument( + "--num_inference_step", + type=int, + default=28, + help="Number of inference steps." + ) + parser.add_argument( + "--seed", + type=int, + default=0, + help="Random seed for generation." + ) + parser.add_argument( + "--height", + type=int, + default=1024, + help="Output image height." + ) + parser.add_argument( + "--width", + type=int, + default=1024, + help="Output image width." + ) + parser.add_argument( + "--max_input_image_pixels", + type=int, + default=1048576, + help="Maximum number of pixels for each input image." + ) + parser.add_argument( + "--dtype", + type=str, + default='bf16', + choices=['fp32', 'fp16', 'bf16'], + help="Data type for model weights." + ) + parser.add_argument( + "--guidance_scale", + type=float, + default=2.5, + ) + parser.add_argument( + "--negative_prompt", + type=str, + default="(((deformed))), blurry, over saturation, bad anatomy, disfigured, poorly drawn face, mutation, mutated, (extra_limb), (ugly), (poorly drawn hands), fused fingers, messy drawing, broken legs censor, censored, censor_bar", + help="Negative prompt for generation." + ) + parser.add_argument( + "--result_dir", + type=str, + required=True, + help="Path to save results." + ) + parser.add_argument( + "--start_index", + type=int, + default=0, + ) + parser.add_argument( + "--end_index", + type=int, + default=1212, + ) + parser.add_argument( + "--num_samples", type=int, default=1 + ) + args = parser.parse_args() + args.root_dir = root_dir + return args + +def load_pipeline(args: argparse.Namespace, accelerator: Accelerator, weight_dtype: torch.dtype) -> FluxKontextPipeline: + pipeline = FluxKontextPipeline.from_pretrained(args.pipeline_path, torch_dtype=weight_dtype) + + pipeline = pipeline.to(device=accelerator.device) + pipeline = pipeline.to(dtype=weight_dtype) + return pipeline + +def run(args: argparse.Namespace, + accelerator: Accelerator, + pipeline: FluxKontextPipeline, + instruction: str, + negative_prompt: str, + input_image: Image.Image, + target_img_size: Tuple[int, int], + seed: int) -> Image.Image: + """Run the image generation pipeline with the given parameters.""" + generator = torch.Generator(device=accelerator.device).manual_seed(seed) + + results = pipeline( + prompt=instruction, + image=input_image, + width=input_image[0].size[0], + height=input_image[0].size[1], + guidance_scale=args.guidance_scale, + num_inference_steps=args.num_inference_step, + generator=generator, + output_type="pil", + ) + return results + +def create_collage(images: List[torch.Tensor]) -> Image.Image: + """Create a horizontal collage from a list of images.""" + max_height = max(img.shape[-2] for img in images) + total_width = sum(img.shape[-1] for img in images) + canvas = torch.zeros((3, max_height, total_width), device=images[0].device) + + current_x = 0 + for img in images: + h, w = img.shape[-2:] + canvas[:, :h, current_x:current_x+w] = img * 0.5 + 0.5 + current_x += w + + return to_pil_image(canvas) + +def main(args: argparse.Namespace, root_dir: str) -> None: + """Main function to run the image generation process.""" + # Initialize accelerator + accelerator = Accelerator(mixed_precision=args.dtype if args.dtype != 'fp32' else 'no') + + # Set weight dtype + weight_dtype = torch.float32 + if args.dtype == 'fp16': + weight_dtype = torch.float16 + elif args.dtype == 'bf16': + weight_dtype = torch.bfloat16 + + # Load pipeline and process inputs + pipeline = load_pipeline(args, accelerator, weight_dtype) + pipeline.set_progress_bar_config(disable=True) + + test_dataset = load_dataset("stepfun-ai/GEdit-Bench", split='train') + + filtered_test_dataset = [ + item for item in test_dataset + if item['instruction_language'] != 'cn' + ] + test_dataset = filtered_test_dataset + + process_index = Accelerator().process_index + + data_index = list(range(args.start_index, args.end_index)) + + with tqdm( + total=len(data_index), + desc=f"process_index {process_index}: Processing {len(data_index)}/{len(test_dataset)}", + unit="image", + disable=not accelerator.is_main_process, + ) as pbar: + for idx in data_index: + data_item = test_dataset[idx] + + task_type = data_item['task_type'] + instruction_language = data_item['instruction_language'] + + key = data_item['key'] + instruction = data_item['instruction'] + input_image = data_item['input_image'] + + ori_img_size = input_image.size + new_img_size = (ori_img_size[0] // 16 * 16, ori_img_size[1] // 16 * 16) + input_images = [input_image.resize(new_img_size)] + + for turn in range(args.num_samples): + results = run(args, accelerator, pipeline, instruction, args.negative_prompt, input_images, new_img_size, args.seed + idx + turn * len(test_dataset)) + output_image = results.images[0] + output_image = output_image.resize(ori_img_size) + + sub_dir = os.path.join(args.result_dir, "fullset", task_type, instruction_language) + os.makedirs(sub_dir, exist_ok=True) + + if turn > 0: + output_image.save(os.path.join(sub_dir, f"{key}_sample{turn}.png")) + else: + input_image.save(os.path.join(sub_dir, f"{key}_SRCIMG.png")) + output_image.save(os.path.join(sub_dir, f"{key}.png")) + + pbar.update(1) + +if __name__ == "__main__": + root_dir = os.path.abspath(os.path.join(__file__, os.path.pardir, os.path.pardir, os.path.pardir)) + args = parse_args(root_dir) + main(args, root_dir) \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/inference_qwen_image.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/inference_qwen_image.py new file mode 100644 index 0000000000000000000000000000000000000000..a8f6390442cf0ee79484db78769e3967f54f759c --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/inference_qwen_image.py @@ -0,0 +1,248 @@ +import dotenv + +dotenv.load_dotenv(override=True) + +import argparse +import os +import sys +from typing import List, Tuple + +from PIL import Image, ImageOps + +from omegaconf import OmegaConf +from tqdm import tqdm + +import torch +from torchvision.transforms.functional import to_pil_image, to_tensor + +from accelerate import Accelerator +from accelerate import init_empty_weights + +from datasets import load_dataset + +from diffusers import QwenImageEditPipeline + +sys.path.append(os.path.join(os.path.dirname(__file__), os.path.pardir, os.path.pardir)) + + +def parse_args(root_dir: str) -> argparse.Namespace: + """Parse command line arguments.""" + parser = argparse.ArgumentParser(description="OmniGen2 image generation script.") + parser.add_argument( + "--load_from_pipeline", + action="store_true", + help="Load from pipeline.", + ) + parser.add_argument( + "--pipeline_path", + type=str, + default=None, + ) + parser.add_argument( + "--experiment_name", + type=str, + default=None, + help="Name of experiment.", + ) + parser.add_argument( + "--model_path", + type=str, + default=None, + help="Path to model checkpoint.", + ) + parser.add_argument( + "--transformer_lora_path", + type=str, + default=None, + help="Path to transformer LoRA weights.", + ) + parser.add_argument( + "--scheduler", + type=str, + default="euler", + choices=["euler", "euler_maruyama", "dpmsolver++"], + help="Scheduler to use.", + ) + parser.add_argument( + "--num_inference_step", + type=int, + default=50, + help="Number of inference steps." + ) + parser.add_argument( + "--seed", + type=int, + default=0, + help="Random seed for generation." + ) + parser.add_argument( + "--height", + type=int, + default=1024, + help="Output image height." + ) + parser.add_argument( + "--width", + type=int, + default=1024, + help="Output image width." + ) + parser.add_argument( + "--max_input_image_pixels", + type=int, + default=1048576, + help="Maximum number of pixels for each input image." + ) + parser.add_argument( + "--dtype", + type=str, + default='bf16', + choices=['fp32', 'fp16', 'bf16'], + help="Data type for model weights." + ) + parser.add_argument( + "--guidance_scale", + type=float, + default=4.0, + ) + parser.add_argument( + "--negative_prompt", + type=str, + default="(((deformed))), blurry, over saturation, bad anatomy, disfigured, poorly drawn face, mutation, mutated, (extra_limb), (ugly), (poorly drawn hands), fused fingers, messy drawing, broken legs censor, censored, censor_bar", + help="Negative prompt for generation." + ) + parser.add_argument( + "--result_dir", + type=str, + required=True, + help="Path to save results." + ) + parser.add_argument( + "--start_index", + type=int, + default=0, + ) + parser.add_argument( + "--end_index", + type=int, + default=1212, + ) + parser.add_argument( + "--num_samples", type=int, default=1 + ) + args = parser.parse_args() + args.root_dir = root_dir + return args + +def load_pipeline(args: argparse.Namespace, accelerator: Accelerator, weight_dtype: torch.dtype) -> QwenImageEditPipeline: + pipeline = QwenImageEditPipeline.from_pretrained(args.pipeline_path, torch_dtype=weight_dtype) + + pipeline = pipeline.to(device=accelerator.device) + pipeline = pipeline.to(dtype=weight_dtype) + return pipeline + + +def run(args: argparse.Namespace, + accelerator: Accelerator, + pipeline: QwenImageEditPipeline, + instruction: str, + negative_prompt: str, + input_image: Image.Image, + target_img_size: Tuple[int, int], + seed: int) -> Image.Image: + """Run the image generation pipeline with the given parameters.""" + generator = torch.Generator(device=accelerator.device).manual_seed(seed) + + results = pipeline( + prompt=instruction, + image=input_image[0], + true_cfg_scale=args.guidance_scale, + negative_prompt=" ", + num_inference_steps=args.num_inference_step, + generator=generator, + output_type="pil", + ) + return results + +def create_collage(images: List[torch.Tensor]) -> Image.Image: + """Create a horizontal collage from a list of images.""" + max_height = max(img.shape[-2] for img in images) + total_width = sum(img.shape[-1] for img in images) + canvas = torch.zeros((3, max_height, total_width), device=images[0].device) + + current_x = 0 + for img in images: + h, w = img.shape[-2:] + canvas[:, :h, current_x:current_x+w] = img * 0.5 + 0.5 + current_x += w + + return to_pil_image(canvas) + +def main(args: argparse.Namespace, root_dir: str) -> None: + """Main function to run the image generation process.""" + # Initialize accelerator + accelerator = Accelerator(mixed_precision=args.dtype if args.dtype != 'fp32' else 'no') + + # Set weight dtype + weight_dtype = torch.float32 + if args.dtype == 'fp16': + weight_dtype = torch.float16 + elif args.dtype == 'bf16': + weight_dtype = torch.bfloat16 + + # Load pipeline and process inputs + pipeline = load_pipeline(args, accelerator, weight_dtype) + pipeline.set_progress_bar_config(disable=True) + + test_dataset = load_dataset("stepfun-ai/GEdit-Bench", split='train') + + filtered_test_dataset = [ + item for item in test_dataset + if item['instruction_language'] != 'cn' + ] + test_dataset = filtered_test_dataset + + process_index = Accelerator().process_index + + data_index = list(range(args.start_index, args.end_index)) + + with tqdm( + total=len(data_index), + desc=f"process_index {process_index}: Processing {len(data_index)}/{len(test_dataset)}", + unit="image", + disable=not accelerator.is_main_process, + ) as pbar: + for idx in data_index: + data_item = test_dataset[idx] + + task_type = data_item['task_type'] + instruction_language = data_item['instruction_language'] + + key = data_item['key'] + instruction = data_item['instruction'] + input_image = data_item['input_image'] + + ori_img_size = input_image.size + new_img_size = (ori_img_size[0] // 16 * 16, ori_img_size[1] // 16 * 16) + input_images = [input_image.resize(new_img_size)] + + for turn in range(args.num_samples): + results = run(args, accelerator, pipeline, instruction, args.negative_prompt, input_images, new_img_size, args.seed + idx + turn * len(test_dataset)) + output_image = results.images[0] + output_image = output_image.resize(ori_img_size) + + sub_dir = os.path.join(args.result_dir, "fullset", task_type, instruction_language) + os.makedirs(sub_dir, exist_ok=True) + + if turn > 0: + output_image.save(os.path.join(sub_dir, f"{key}_sample{turn}.png")) + else: + input_image.save(os.path.join(sub_dir, f"{key}_SRCIMG.png")) + output_image.save(os.path.join(sub_dir, f"{key}.png")) + + pbar.update(1) + +if __name__ == "__main__": + root_dir = os.path.abspath(os.path.join(__file__, os.path.pardir, os.path.pardir, os.path.pardir)) + args = parse_args(root_dir) + main(args, root_dir) \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2.sh new file mode 100644 index 0000000000000000000000000000000000000000..1af417b3875e7033179e1ba306f77bdd6b30f7b4 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2.sh @@ -0,0 +1,87 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +experiment_name=omnigen2_edit_rl_4machine_editscore7b_avg8 +step=700 +RANK=0 +WORLD_SIZE=1 + +while [[ $# -gt 0 ]]; do + case "$1" in + --experiment_name=*) + experiment_name="${1#*=}" + shift + ;; + --step=*) + step="${1#*=}" + shift + ;; + --rank=*) + RANK="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + *) + echo "ๆœช็Ÿฅๅ‚ๆ•ฐ: $1" + shift + ;; + esac +done + +# ่พ“ๅ‡บ้…็ฝฎ +echo "RANK: $RANK" +echo "WORLD_SIZE: $WORLD_SIZE" + +global_shift_index=0 +total_num_images=606 + +num_gpus_per_machine=$(python -c "import torch; print(torch.cuda.device_count())") +# Calculate images per machine, rounding up to ensure all data is covered +num_images_per_machine=$(( (total_num_images + WORLD_SIZE - 1) / WORLD_SIZE )) +shift_index=$((RANK * num_images_per_machine)) + +if [ $((total_num_images - shift_index)) -lt $num_images_per_machine ]; then + num_images_per_machine=$((total_num_images - shift_index)) +fi + +# Calculate base number of images per GPU (for first 7 GPUs) +num_images_per_gpu=$(( (num_images_per_machine + num_gpus_per_machine - 1) / num_gpus_per_machine )) + +text_guidance_scale=5.0 +image_guidance_scale=1.5 + +for ((i=0; i logs/gedit_${experiment_name}_step${step}_ts${text_guidance_scale}_ig${image_guidance_scale}_${start_idx}_${end_idx}.log 2>&1 & +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples.sh new file mode 100644 index 0000000000000000000000000000000000000000..3ad585e688af1c98d7478227c700ca51784e4ab9 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples.sh @@ -0,0 +1,90 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +RANK=0 +MASTER_ADDR=1 +MASTER_PORT=29500 +WORLD_SIZE=1 + +# ๅค„็†ๅ‘ฝๅๅ‚ๆ•ฐ +while [[ $# -gt 0 ]]; do + case "$1" in + --rank=*) + RANK="${1#*=}" + shift + ;; + --master_addr=*) + MASTER_ADDR="${1#*=}" + shift + ;; + --master_port=*) + MASTER_PORT="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + *) + echo "ๆœช็Ÿฅๅ‚ๆ•ฐ: $1" + shift + ;; + esac +done + +# ่พ“ๅ‡บ้…็ฝฎ +echo "RANK: $RANK" +echo "MASTER_ADDR: $MASTER_ADDR" +echo "MASTER_PORT: $MASTER_PORT" +echo "WORLD_SIZE: $WORLD_SIZE" + +global_shift_index=0 +total_num_images=606 + +num_gpus_per_machine=$(python -c "import torch; print(torch.cuda.device_count())") +# Calculate images per machine, rounding up to ensure all data is covered +num_images_per_machine=$(( (total_num_images + WORLD_SIZE - 1) / WORLD_SIZE )) +shift_index=$((RANK * num_images_per_machine)) + +if [ $((total_num_images - shift_index)) -lt $num_images_per_machine ]; then + num_images_per_machine=$((total_num_images - shift_index)) +fi + +# Calculate base number of images per GPU (for first 7 GPUs) +num_images_per_gpu=$(( (num_images_per_machine + num_gpus_per_machine - 1) / num_gpus_per_machine )) + +text_guidance_scale=5.0 +image_guidance_scale=1.5 + +for ((i=0; i logs/gedit_OmniGen2_ts${text_guidance_scale}_ig${image_guidance_scale}_16samples_${start_idx}_${end_idx}.log 2>&1 & +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass1.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass1.sh new file mode 100644 index 0000000000000000000000000000000000000000..57ba4e44d9a4e8173dc4b34e770e65df79163d36 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass1.sh @@ -0,0 +1,90 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +RANK=0 +MASTER_ADDR=1 +MASTER_PORT=29500 +WORLD_SIZE=1 + +# ๅค„็†ๅ‘ฝๅๅ‚ๆ•ฐ +while [[ $# -gt 0 ]]; do + case "$1" in + --rank=*) + RANK="${1#*=}" + shift + ;; + --master_addr=*) + MASTER_ADDR="${1#*=}" + shift + ;; + --master_port=*) + MASTER_PORT="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + *) + echo "ๆœช็Ÿฅๅ‚ๆ•ฐ: $1" + shift + ;; + esac +done + +# ่พ“ๅ‡บ้…็ฝฎ +echo "RANK: $RANK" +echo "MASTER_ADDR: $MASTER_ADDR" +echo "MASTER_PORT: $MASTER_PORT" +echo "WORLD_SIZE: $WORLD_SIZE" + +global_shift_index=0 +total_num_images=606 + +num_gpus_per_machine=$(python -c "import torch; print(torch.cuda.device_count())") +# Calculate images per machine, rounding up to ensure all data is covered +num_images_per_machine=$(( (total_num_images + WORLD_SIZE - 1) / WORLD_SIZE )) +shift_index=$((RANK * num_images_per_machine)) + +if [ $((total_num_images - shift_index)) -lt $num_images_per_machine ]; then + num_images_per_machine=$((total_num_images - shift_index)) +fi + +# Calculate base number of images per GPU (for first 7 GPUs) +num_images_per_gpu=$(( (num_images_per_machine + num_gpus_per_machine - 1) / num_gpus_per_machine )) + +text_guidance_scale=5.0 +image_guidance_scale=1.5 + +for ((i=0; i logs/gedit_OmniGen2_ts${text_guidance_scale}_ig${image_guidance_scale}_16samples_select_best_pass1_${start_idx}_${end_idx}.log 2>&1 & +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass1_eval.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass1_eval.sh new file mode 100644 index 0000000000000000000000000000000000000000..e4e0377cb564a76730a365514626c625f74c5ae8 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass1_eval.sh @@ -0,0 +1,26 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +best=( +1 +2 +4 +8 +16 +) + +for b in "${best[@]}" +do + accelerate launch --num_processes 1 evaluation/GEdit-Bench/test_gedit_score.py \ + --result_dir evaluation/GEdit-Bench/results/OmniGen2/results_ts${text_guidance_scale}_ig${image_guidance_scale}_16samples_pass1_best${b} \ + --backbone gpt-4.1 \ + --openai_url https://api.openai.com/v1/chat/completions \ + --max_workers 30 \ + --key PUT-YOUR-KEY-HERE + + python evaluation/GEdit-Bench/calculate_statistics.py \ + --result_dir evaluation/GEdit-Bench/results/OmniGen2/results_ts${text_guidance_scale}_ig${image_guidance_scale}_16samples_pass1_best${b}/viescore_gpt-4.1 \ + --language en +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass4.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass4.sh new file mode 100644 index 0000000000000000000000000000000000000000..04fdf7ab6a5adb84a90440502755e1d086c2f859 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass4.sh @@ -0,0 +1,90 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +RANK=0 +MASTER_ADDR=1 +MASTER_PORT=29500 +WORLD_SIZE=1 + +# ๅค„็†ๅ‘ฝๅๅ‚ๆ•ฐ +while [[ $# -gt 0 ]]; do + case "$1" in + --rank=*) + RANK="${1#*=}" + shift + ;; + --master_addr=*) + MASTER_ADDR="${1#*=}" + shift + ;; + --master_port=*) + MASTER_PORT="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + *) + echo "ๆœช็Ÿฅๅ‚ๆ•ฐ: $1" + shift + ;; + esac +done + +# ่พ“ๅ‡บ้…็ฝฎ +echo "RANK: $RANK" +echo "MASTER_ADDR: $MASTER_ADDR" +echo "MASTER_PORT: $MASTER_PORT" +echo "WORLD_SIZE: $WORLD_SIZE" + +global_shift_index=0 +total_num_images=606 + +num_gpus_per_machine=$(python -c "import torch; print(torch.cuda.device_count())") +# Calculate images per machine, rounding up to ensure all data is covered +num_images_per_machine=$(( (total_num_images + WORLD_SIZE - 1) / WORLD_SIZE )) +shift_index=$((RANK * num_images_per_machine)) + +if [ $((total_num_images - shift_index)) -lt $num_images_per_machine ]; then + num_images_per_machine=$((total_num_images - shift_index)) +fi + +# Calculate base number of images per GPU (for first 7 GPUs) +num_images_per_gpu=$(( (num_images_per_machine + num_gpus_per_machine - 1) / num_gpus_per_machine )) + +text_guidance_scale=5.0 +image_guidance_scale=1.5 + +for ((i=0; i logs/gedit_OmniGen2_ts${text_guidance_scale}_ig${image_guidance_scale}_16samples_select_best_pass4_${start_idx}_${end_idx}.log 2>&1 & +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass4_eval.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass4_eval.sh new file mode 100644 index 0000000000000000000000000000000000000000..c5b461a3f87f9474855db732cee2f876b63e54ef --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass4_eval.sh @@ -0,0 +1,26 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +best=( +1 +2 +4 +8 +16 +) + +for b in "${best[@]}" +do + accelerate launch --num_processes 1 evaluation/GEdit-Bench/test_gedit_score.py \ + --result_dir evaluation/GEdit-Bench/results/OmniGen2/results_ts${text_guidance_scale}_ig${image_guidance_scale}_16samples_pass4_best${b} \ + --backbone gpt-4.1 \ + --openai_url https://api.openai.com/v1/chat/completions \ + --max_workers 30 \ + --key PUT-YOUR-KEY-HERE + + python evaluation/GEdit-Bench/calculate_statistics.py \ + --result_dir evaluation/GEdit-Bench/results/OmniGen2/results_ts${text_guidance_scale}_ig${image_guidance_scale}_16samples_pass4_best${b}/viescore_gpt-4.1 \ + --language en +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_edit_rl_single_machine_editscore7b_step500.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_edit_rl_single_machine_editscore7b_step500.sh new file mode 100644 index 0000000000000000000000000000000000000000..621afd0db0a2b05fd67819103478aad4a08a2309 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_edit_rl_single_machine_editscore7b_step500.sh @@ -0,0 +1,91 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +RANK=0 +MASTER_ADDR=1 +MASTER_PORT=29500 +WORLD_SIZE=1 + +# ๅค„็†ๅ‘ฝๅๅ‚ๆ•ฐ +while [[ $# -gt 0 ]]; do + case "$1" in + --rank=*) + RANK="${1#*=}" + shift + ;; + --master_addr=*) + MASTER_ADDR="${1#*=}" + shift + ;; + --master_port=*) + MASTER_PORT="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + *) + echo "ๆœช็Ÿฅๅ‚ๆ•ฐ: $1" + shift + ;; + esac +done + +# ่พ“ๅ‡บ้…็ฝฎ +echo "RANK: $RANK" +echo "MASTER_ADDR: $MASTER_ADDR" +echo "MASTER_PORT: $MASTER_PORT" +echo "WORLD_SIZE: $WORLD_SIZE" + +global_shift_index=0 +total_num_images=606 + +num_gpus_per_machine=$(python -c "import torch; print(torch.cuda.device_count())") +# Calculate images per machine, rounding up to ensure all data is covered +num_images_per_machine=$(( (total_num_images + WORLD_SIZE - 1) / WORLD_SIZE )) +shift_index=$((RANK * num_images_per_machine)) + +if [ $((total_num_images - shift_index)) -lt $num_images_per_machine ]; then + num_images_per_machine=$((total_num_images - shift_index)) +fi + +# Calculate base number of images per GPU (for first 7 GPUs) +num_images_per_gpu=$(( (num_images_per_machine + num_gpus_per_machine - 1) / num_gpus_per_machine )) + +text_guidance_scale=5.0 +image_guidance_scale=1.5 + +for ((i=0; i logs/gedit_OmniGen2_ts${text_guidance_scale}_ig${image_guidance_scale}_16samples_${start_idx}_${end_idx}.log 2>&1 & +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_eval.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_eval.sh new file mode 100644 index 0000000000000000000000000000000000000000..1b2b1835f83e291aa7a964a4400e3a004b083c16 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/omnigen2_eval.sh @@ -0,0 +1,42 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +source "$(dirname $(which conda))/../etc/profile.d/conda.sh" +conda activate py3.12+pytorch2.7.1+cu126 + +experiment_name=omnigen2_edit_rl_4machine_editscore7b_avg8 +step=700 + +# ๅค„็†ๅ‘ฝๅๅ‚ๆ•ฐ +while [[ $# -gt 0 ]]; do + case "$1" in + --experiment_name=*) + experiment_name="${1#*=}" + shift + ;; + --step=*) + step="${1#*=}" + shift + ;; + *) + echo "ๆœช็Ÿฅๅ‚ๆ•ฐ: $1" + shift + ;; + esac +done + +text_guidance_scale=5.0 +image_guidance_scale=1.5 + +accelerate launch --num_processes 1 evaluation/GEdit-Bench/test_gedit_score.py \ +--result_dir evaluation/GEdit-Bench/results/${experiment_name}/results_step${step}_ts${text_guidance_scale}_ig${image_guidance_scale} \ +--backbone gpt-4.1 \ +--openai_url https://api.openai.com/v1/chat/completions \ +--max_workers 30 \ +--key PUT-YOUR-KEY-HERE + +python evaluation/GEdit-Bench/calculate_statistics.py \ +--result_dir evaluation/GEdit-Bench/results/${experiment_name}/results_step${step}_ts${text_guidance_scale}_ig${image_guidance_scale}/viescore_gpt-4.1 \ +--language en \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples.sh new file mode 100644 index 0000000000000000000000000000000000000000..81080fed4d242190dd72579a04d5021a3605a21e --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples.sh @@ -0,0 +1,84 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +RANK=0 +MASTER_ADDR=1 +MASTER_PORT=29500 +WORLD_SIZE=1 + +# ๅค„็†ๅ‘ฝๅๅ‚ๆ•ฐ +while [[ $# -gt 0 ]]; do + case "$1" in + --rank=*) + RANK="${1#*=}" + shift + ;; + --master_addr=*) + MASTER_ADDR="${1#*=}" + shift + ;; + --master_port=*) + MASTER_PORT="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + *) + echo "ๆœช็Ÿฅๅ‚ๆ•ฐ: $1" + shift + ;; + esac +done + +# ่พ“ๅ‡บ้…็ฝฎ +echo "RANK: $RANK" +echo "MASTER_ADDR: $MASTER_ADDR" +echo "MASTER_PORT: $MASTER_PORT" +echo "WORLD_SIZE: $WORLD_SIZE" + +global_shift_index=0 +total_num_images=606 + +num_gpus_per_machine=$(python -c "import torch; print(torch.cuda.device_count())") +# Calculate images per machine, rounding up to ensure all data is covered +num_images_per_machine=$(( (total_num_images + WORLD_SIZE - 1) / WORLD_SIZE )) +shift_index=$((RANK * num_images_per_machine)) + +if [ $((total_num_images - shift_index)) -lt $num_images_per_machine ]; then + num_images_per_machine=$((total_num_images - shift_index)) +fi + +# Calculate base number of images per GPU (for first 7 GPUs) +num_images_per_gpu=$(( (num_images_per_machine + num_gpus_per_machine - 1) / num_gpus_per_machine )) + +guidance_scale=4.0 + +for ((i=0; i logs/gedit_Qwen-Image-Edit_gs${guidance_scale}_16samples_${start_idx}_${end_idx}.log 2>&1 & +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples_select_best_editscore_pass1.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples_select_best_editscore_pass1.sh new file mode 100644 index 0000000000000000000000000000000000000000..6a801862887f93a147df0b683119cb813192712a --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples_select_best_editscore_pass1.sh @@ -0,0 +1,89 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +RANK=0 +MASTER_ADDR=1 +MASTER_PORT=29500 +WORLD_SIZE=1 + +# ๅค„็†ๅ‘ฝๅๅ‚ๆ•ฐ +while [[ $# -gt 0 ]]; do + case "$1" in + --rank=*) + RANK="${1#*=}" + shift + ;; + --master_addr=*) + MASTER_ADDR="${1#*=}" + shift + ;; + --master_port=*) + MASTER_PORT="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + *) + echo "ๆœช็Ÿฅๅ‚ๆ•ฐ: $1" + shift + ;; + esac +done + +# ่พ“ๅ‡บ้…็ฝฎ +echo "RANK: $RANK" +echo "MASTER_ADDR: $MASTER_ADDR" +echo "MASTER_PORT: $MASTER_PORT" +echo "WORLD_SIZE: $WORLD_SIZE" + +global_shift_index=0 +total_num_images=606 + +num_gpus_per_machine=$(python -c "import torch; print(torch.cuda.device_count())") +# Calculate images per machine, rounding up to ensure all data is covered +num_images_per_machine=$(( (total_num_images + WORLD_SIZE - 1) / WORLD_SIZE )) +shift_index=$((RANK * num_images_per_machine)) + +if [ $((total_num_images - shift_index)) -lt $num_images_per_machine ]; then + num_images_per_machine=$((total_num_images - shift_index)) +fi + +# Calculate base number of images per GPU (for first 7 GPUs) +num_images_per_gpu=$(( (num_images_per_machine + num_gpus_per_machine - 1) / num_gpus_per_machine )) + +guidance_scale=4.0 + +for ((i=0; i logs/gedit_Qwen-Image-Edit_gs${guidance_scale}_16samples_select_best_pass1_${start_idx}_${end_idx}.log 2>&1 & +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples_select_best_editscore_pass1_eval.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples_select_best_editscore_pass1_eval.sh new file mode 100644 index 0000000000000000000000000000000000000000..f6976cfbdefb2115556c7e64164f3bc37055e723 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples_select_best_editscore_pass1_eval.sh @@ -0,0 +1,26 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +best=( +1 +2 +4 +8 +16 +) + +for b in "${best[@]}" +do + accelerate launch --num_processes 1 evaluation/GEdit-Bench/test_gedit_score.py \ + --result_dir evaluation/GEdit-Bench/results/Qwen-Image-Edit/results_gs${guidance_scale}_16samples_pass1_best${b} \ + --backbone gpt-4.1 \ + --openai_url https://api.openai.com/v1/chat/completions \ + --max_workers 30 \ + --key PUT-YOUR-KEY-HERE + + python evaluation/GEdit-Bench/calculate_statistics.py \ + --result_dir evaluation/GEdit-Bench/results/Qwen-Image-Edit/results_gs${guidance_scale}_16samples_pass1_best${b}/viescore_gpt-4.1 \ + --language en +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples_select_best_editscore_pass4.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples_select_best_editscore_pass4.sh new file mode 100644 index 0000000000000000000000000000000000000000..dfa647ca434ec6b02df620129a1163dc8b0eda5d --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples_select_best_editscore_pass4.sh @@ -0,0 +1,89 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +RANK=0 +MASTER_ADDR=1 +MASTER_PORT=29500 +WORLD_SIZE=1 + +# ๅค„็†ๅ‘ฝๅๅ‚ๆ•ฐ +while [[ $# -gt 0 ]]; do + case "$1" in + --rank=*) + RANK="${1#*=}" + shift + ;; + --master_addr=*) + MASTER_ADDR="${1#*=}" + shift + ;; + --master_port=*) + MASTER_PORT="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + *) + echo "ๆœช็Ÿฅๅ‚ๆ•ฐ: $1" + shift + ;; + esac +done + +# ่พ“ๅ‡บ้…็ฝฎ +echo "RANK: $RANK" +echo "MASTER_ADDR: $MASTER_ADDR" +echo "MASTER_PORT: $MASTER_PORT" +echo "WORLD_SIZE: $WORLD_SIZE" + +global_shift_index=0 +total_num_images=606 + +num_gpus_per_machine=$(python -c "import torch; print(torch.cuda.device_count())") +# Calculate images per machine, rounding up to ensure all data is covered +num_images_per_machine=$(( (total_num_images + WORLD_SIZE - 1) / WORLD_SIZE )) +shift_index=$((RANK * num_images_per_machine)) + +if [ $((total_num_images - shift_index)) -lt $num_images_per_machine ]; then + num_images_per_machine=$((total_num_images - shift_index)) +fi + +# Calculate base number of images per GPU (for first 7 GPUs) +num_images_per_gpu=$(( (num_images_per_machine + num_gpus_per_machine - 1) / num_gpus_per_machine )) + +guidance_scale=4.0 + +for ((i=0; i logs/gedit_Qwen-Image-Edit_gs${guidance_scale}_16samples_select_best_pass4_${start_idx}_${end_idx}.log 2>&1 & +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples_select_best_editscore_pass4_eval.sh b/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples_select_best_editscore_pass4_eval.sh new file mode 100644 index 0000000000000000000000000000000000000000..d18eef94366ec1adb8ff63b4b3c90ad5ebcfa814 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/qwen_image_16samples_select_best_editscore_pass4_eval.sh @@ -0,0 +1,26 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +best=( +1 +2 +4 +8 +16 +) + +for b in "${best[@]}" +do + accelerate launch --num_processes 1 evaluation/GEdit-Bench/test_gedit_score.py \ + --result_dir evaluation/GEdit-Bench/results/Qwen-Image-Edit/results_gs${guidance_scale}_16samples_pass4_best${b} \ + --backbone gpt-4.1 \ + --openai_url https://api.openai.com/v1/chat/completions \ + --max_workers 30 \ + --key PUT-YOUR-KEY-HERE + + python evaluation/GEdit-Bench/calculate_statistics.py \ + --result_dir evaluation/GEdit-Bench/results/Qwen-Image-Edit/results_gs${guidance_scale}_16samples_pass4_best${b}/viescore_gpt-4.1 \ + --language en +done \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/select_best.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/select_best.py new file mode 100644 index 0000000000000000000000000000000000000000..4906ca0f8897497802869c9ada9f808c9183d446 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/select_best.py @@ -0,0 +1,181 @@ +import dotenv + +dotenv.load_dotenv(override=True) + +from editscore import EditScore +import PIL +import os +import glob +import json + +from PIL import Image +import argparse +from datasets import load_dataset +from shutil import copyfile + +import hashlib # ๆ–ฐๅขžๅฏผๅ…ฅ +import threading # ๆ–ฐๅขžๅฏผๅ…ฅ +from tqdm import tqdm + +def generate_cache_key(pair): + """ไธบๆฏไธชๆ ทๆœฌ็”Ÿๆˆไธ€ไธชๅ”ฏไธ€็š„SHA256ๅ“ˆๅธŒ้”ฎใ€‚""" + instruction, input_image, output_image = pair + # ๅฐ†ไธ‰ไธช็ป„ไปถ็”จ็‰นๆฎŠๅˆ†้š”็ฌฆ่ฟžๆŽฅ๏ผŒ็กฎไฟไธไผšๅ› ๅ†…ๅฎนๆœฌ่บซๅŒ…ๅซๅˆ†้š”็ฌฆ่€Œๆททๆท† + key_string = f"{instruction}|||{input_image}|||{output_image}" + return hashlib.sha256(key_string.encode('utf-8')).hexdigest() + +def load_cache(cache_file): + """ไปŽJSONLๆ–‡ไปถๅŠ ่ฝฝ็ผ“ๅญ˜ๅˆฐๅญ—ๅ…ธใ€‚""" + cache = {} + if not os.path.exists(cache_file): + return cache + with open(cache_file, 'r', encoding='utf-8') as f: + for line in f: + try: + data = json.loads(line) + cache[data['key']] = data['result'] + except json.JSONDecodeError: + print(f"Warning: Skipping corrupted line in cache file: {line.strip()}") + return cache + +def append_to_cache(cache_file, key, result, lock): + """ๅฐ†ๆ–ฐ็š„่ฎก็ฎ—็ป“ๆžœ็บฟ็จ‹ๅฎ‰ๅ…จๅœฐ่ฟฝๅŠ ๅˆฐ็ผ“ๅญ˜ๆ–‡ไปถใ€‚""" + with lock: + with open(cache_file, 'a', encoding='utf-8') as f: + f.write(json.dumps({'key': key, 'result': result}, ensure_ascii=False) + '\n') + +def process_single_item(item, vie_score, args, max_retries=10000): + instruction = item[0] + input_image = item[1] + output_image = item[2] + + for retry in range(max_retries): + # try: + pil_image_raw = Image.open(input_image).convert("RGB") + + pil_image_edited = Image.open(output_image).convert("RGB") + pil_image_edited = pil_image_edited.resize((pil_image_raw.size[0], pil_image_raw.size[1])) + + text_prompt = instruction + score = vie_score.evaluate( + [pil_image_raw, pil_image_edited], text_prompt, echo_output=False + ) + return item, score + +def parse_args(): + parser = argparse.ArgumentParser() + parser.add_argument("--result_dir", type=str, required=True) + parser.add_argument("--save_dir", type=str, required=True) + parser.add_argument("--scorer", type=str, default="viescore", choices=["viescore", "editscore"]) + parser.add_argument( + "--backbone", + type=str, + default="openai", + choices=["openai", "qwen25vl", "qwen25vl_vllm", "internvl3_5"], + ) + parser.add_argument("--model_name_or_path", type=str, default="gpt-4.1") + parser.add_argument( + "--openai_url", type=str, default="https://api.openai.com/v1/chat/completions" + ) + parser.add_argument( + "--key", type=str, default="sk-cB6h7HcCSDIp71gs6lFLZxKE0dOYOnJbxzES6kWXe1Wb2VHS" + ) + parser.add_argument( + "--context_version", type=str, default="v1", choices=["v1", "v2"] + ) + parser.add_argument( + "--prompt_version", type=str, default="default", choices=["default", "our", "editscore"] + ) + parser.add_argument( + "--start_index", + type=int, + default=0, + ) + parser.add_argument( + "--end_index", + type=int, + default=1212, + ) + parser.add_argument( + "--num_samples", type=int, default=1 + ) + parser.add_argument("--num_pass", type=int, default=1) + parser.add_argument("--temperature", type=float, default=0.7) + parser.add_argument("--max_workers", type=int, default=20) + parser.add_argument("--score_range", type=int, default=10) + parser.add_argument("--tensor_parallel_size", type=int, default=1) + parser.add_argument("--max_model_len", type=int, default=1536) + parser.add_argument("--max_num_seqs", type=int, default=32) + parser.add_argument("--max_num_batched_tokens", type=int, default=1536) + parser.add_argument("--enable_lora", action="store_true") + parser.add_argument("--lora_path", type=str, default="") + parser.add_argument("--cache_dir", type=str, default=None) + return parser.parse_args() + +def main(args): + scorer = EditScore( + backbone=args.backbone, + key=args.key, + openai_url=args.openai_url, + model_name_or_path=args.model_name_or_path, + score_range=args.score_range, + temperature=args.temperature, + tensor_parallel_size=args.tensor_parallel_size, + max_model_len=args.max_model_len, + max_num_seqs=args.max_num_seqs, + max_num_batched_tokens=args.max_num_batched_tokens, + num_pass=args.num_pass, + enable_lora=args.enable_lora, + lora_path=args.lora_path, + cache_dir=args.cache_dir, + ) + + dataset = load_dataset("stepfun-ai/GEdit-Bench", split='train') + dataset = dataset.remove_columns(["input_image", "input_image_raw"]) + dataset = dataset.filter(lambda x: x["instruction_language"] == "en", num_proc=4) + + data_index_list = list(range(args.start_index, args.end_index)) + with tqdm( + total=len(data_index_list), + desc=f"Processing {len(data_index_list)}/{len(dataset)}", + unit="image" + ) as pbar: + for idx in data_index_list: + data_item = dataset[idx] + + task_type = data_item['task_type'] + instruction_language = data_item['instruction_language'] + + key = data_item['key'] + instruction = data_item['instruction'] + input_image_path = f"{args.result_dir}/fullset/{task_type}/{instruction_language}/{key}_SRCIMG.png" + input_image = Image.open(input_image_path).convert('RGB') + + best_score = -float('inf') + best_output_image_path = None + + for turn in range(args.num_samples): + output_image_path = f"{args.result_dir}/fullset/{task_type}/{instruction_language}/{key}{'_sample' + str(turn) if turn > 0 else ''}.png" + output_image = Image.open(output_image_path).convert('RGB') + + score = scorer.evaluate( + [input_image, output_image], instruction, echo_output=False + )['overall'] + + if score > best_score: + best_score = score + best_output_image_path = output_image_path + + if (turn + 1) & (turn + 1 - 1) == 0: + save_input_image_path = f"{args.save_dir}_best{turn + 1}/fullset/{task_type}/{instruction_language}/{key}_SRCIMG.png" + save_output_image_path = f"{args.save_dir}_best{turn + 1}/fullset/{task_type}/{instruction_language}/{key}.png" + + os.makedirs(os.path.dirname(save_input_image_path), exist_ok=True) + if not os.path.exists(save_input_image_path): + copyfile(input_image_path, save_input_image_path) + copyfile(best_output_image_path, save_output_image_path) + pbar.update(1) + +if __name__ == "__main__": + args = parse_args() + main(args) \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/test_gedit_score.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/test_gedit_score.py new file mode 100644 index 0000000000000000000000000000000000000000..2e8e58e39dedb4f7e7fa0856be2dbacb868ceb3e --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/test_gedit_score.py @@ -0,0 +1,231 @@ +from viescore import VIEScore +import PIL +import os + +# import megfile +from PIL import Image +from tqdm import tqdm +from datasets import load_dataset, load_from_disk +import sys +import csv +import threading +import time +import argparse +from concurrent.futures import ThreadPoolExecutor, as_completed +from collections import defaultdict +import accelerate +from accelerate import Accelerator +from accelerate.state import AcceleratorState + +GROUPS = [ + "background_change", + "color_alter", + "material_alter", + "motion_change", + "ps_human", + "style_change", + "subject-add", + "subject-remove", + "subject-replace", + "text_change", + "tone_transfer", +] + +def process_single_item(item, vie_score, args, turn, max_retries=10000): + instruction = item["instruction"] + key = item["key"] + instruction_language = item["instruction_language"] + save_path_fullset_source_image = f"{args.result_dir}/fullset/{group_name}/{instruction_language}/{key}_SRCIMG.png" + save_path_fullset_result_image = ( + f"{args.result_dir}/fullset/{group_name}/{instruction_language}/{key}{'_sample' + str(turn) if turn > 0 else ''}.png" + ) + + src_image_path = save_path_fullset_source_image + save_path_item = save_path_fullset_result_image + + for retry in range(max_retries): + # try: + pil_image_raw = Image.open(open(src_image_path, "rb")).convert("RGB") + pil_image_edited = ( + Image.open(open(save_path_item, "rb")) + .convert("RGB") + .resize((pil_image_raw.size[0], pil_image_raw.size[1])) + ) + + text_prompt = instruction + score_list = vie_score.evaluate( + [pil_image_raw, pil_image_edited], text_prompt, echo_output=False + ) + sementics_score, quality_score, overall_score = score_list + + return { + "source_image": src_image_path, + "edited_image": save_path_item, + "instruction": instruction, + "sementics_score": sementics_score, + "quality_score": quality_score, + "intersection_exist": item["Intersection_exist"], + "instruction_language": item["instruction_language"], + } + +if __name__ == "__main__": + accelerator = Accelerator() + + parser = argparse.ArgumentParser() + parser.add_argument("--result_dir", type=str, default="/results/") + parser.add_argument("--csv_dir", type=str, default="viescore_gpt-4.1") + parser.add_argument( + "--backbone", type=str, default="gpt-4.1", choices=["gpt-4.1", "gpt-5"] + ) + parser.add_argument( + "--openai_url", type=str, default="https://api.openai.com/v1/chat/completions" + ) + parser.add_argument("--max_workers", type=int, default=20) + parser.add_argument( + "--key", type=str, required=True + ) + parser.add_argument("--num_samples", type=int, default=1) + + args = parser.parse_args() + + backbone = args.backbone + + cur_dir = os.path.dirname(os.path.abspath(__file__)) + vie_score = VIEScore( + backbone=backbone, task="tie", key=args.key, openai_url=args.openai_url + ) + max_workers = 20 + dataset = load_dataset("stepfun-ai/GEdit-Bench", split='train') + dataset = dataset.remove_columns(["input_image", "input_image_raw"]) + dataset = dataset.filter(lambda x: x["instruction_language"] == "en", num_proc=4) + + data_index_list = list( + range( + AcceleratorState().process_index, + len(dataset), + AcceleratorState().num_processes, + ) + ) + + all_csv_list = defaultdict(list) # Store all results for final combined CSV + for group_name in GROUPS: + for turn in range(args.num_samples): + group_csv_list = [] + group_dataset_list = dataset.filter( + lambda x: x["task_type"] == group_name, num_proc=4 + ) + + # Load existing group CSV if it exists + group_csv_path = os.path.join( + args.result_dir, f"viescore_{args.backbone}", f"{group_name}_gpt_score{'_sample' + str(turn) if turn > 0 else ''}.csv" + ) + + processed_samples = set() + + if os.path.exists(group_csv_path): + with open(group_csv_path, "r", newline="", encoding="utf-8-sig") as f: + reader = csv.DictReader(f) + group_results = list(reader) + group_csv_list.extend(group_results) + + for row in group_results: + sample_key = (row["source_image"], row["edited_image"]) + processed_samples.add(sample_key) + + print(f"Loaded existing results for {group_name}") + + print(f"Processing group: {group_name}") + + with ThreadPoolExecutor(max_workers=max_workers) as executor: + futures = [] + for item in group_dataset_list: + instruction = item["instruction"] + key = item["key"] + instruction_language = item["instruction_language"] + intersection_exist = item["Intersection_exist"] + sample_prefix = key + save_path_fullset_source_image = f"{args.result_dir}/fullset/{group_name}/{instruction_language}/{key}_SRCIMG.png" + save_path_fullset_result_image = f"{args.result_dir}/fullset/{group_name}/{instruction_language}/{key}{'_sample' + str(turn) if turn > 0 else ''}.png" + + if not os.path.exists( + save_path_fullset_result_image + ) or not os.path.exists(save_path_fullset_source_image): + print( + f"Skipping {sample_prefix}: Source or edited image does not exist {save_path_fullset_result_image=}" + ) + continue + + # Check if this sample has already been processed + sample_key = ( + save_path_fullset_source_image, + save_path_fullset_result_image, + ) + exists = sample_key in processed_samples + if exists: + print( + f"Skipping already processed sample: {sample_prefix}", + flush=True, + ) + continue + + future = executor.submit(process_single_item, item, vie_score, args, turn) + futures.append(future) + + for future in tqdm( + as_completed(futures), + total=len(futures), + unit="image", + desc=f"Processing {group_name} {turn}", + ): + result = future.result() + if result: + group_csv_list.append(result) + + from accelerate.utils import gather_object + + group_csv_list = gather_object(group_csv_list) + + if accelerator.is_main_process: + # Save group-specific CSV + group_csv_path = os.path.join( + args.result_dir, f"viescore_{args.backbone}", f"{group_name}_gpt_score{'_sample' + str(turn) if turn > 0 else ''}.csv" + ) + os.makedirs(os.path.dirname(group_csv_path), exist_ok=True) + with open(group_csv_path, "w", newline="", encoding="utf-8-sig") as f: + fieldnames = [ + "source_image", + "edited_image", + "instruction", + "sementics_score", + "quality_score", + "intersection_exist", + "instruction_language", + ] + writer = csv.DictWriter(f, fieldnames=fieldnames) + writer.writeheader() + for row in group_csv_list: + writer.writerow(row) + all_csv_list[turn].extend(group_csv_list) + + print(f"Saved group CSV for {group_name}, length๏ผš {len(group_csv_list)}, file: {group_csv_path}") + + if accelerator.is_main_process and all_csv_list[turn]: + # Save combined CSV + combined_csv_path = os.path.join( + args.result_dir, f"viescore_{args.backbone}", f"combined_gpt_score{'_sample' + str(turn) if turn > 0 else ''}.csv" + ) + os.makedirs(os.path.dirname(combined_csv_path), exist_ok=True) + with open(combined_csv_path, "w", newline="", encoding="utf-8-sig") as f: + fieldnames = [ + "source_image", + "edited_image", + "instruction", + "sementics_score", + "quality_score", + "intersection_exist", + "instruction_language", + ] + writer = csv.DictWriter(f, fieldnames=fieldnames) + writer.writeheader() + for row in all_csv_list[turn]: + writer.writerow(row) \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/__init__.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..28f46cfa1a57592d1109d7f79126cab659ef5480 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/__init__.py @@ -0,0 +1,137 @@ +import sys +sys.path.insert(0, 'viescore') + +from .utils import ( + mllm_output_to_dict +) +import math +from . import vie_prompts + +class VIEScore: + def __init__( + self, + backbone="gpt4o", + openai_url="https://api.openai.com/v1/chat/completions", + task="t2i", + key=None, + ) -> None: + self.task = task + self.backbone_name = backbone + + if self.task not in ["t2i", "tie", "t2v"]: + raise ValueError("task must be either 't2i' or 'tie'") + + if self.backbone_name in ["gpt4o", "gpt-4.1", "gpt-5"]: + from .mllm_tools.openai import GPT4o + self.model = GPT4o(key, model_name=self.backbone_name, url=openai_url) + elif self.backbone_name == "gpt4v": + from .mllm_tools.openai import GPT4v + self.model = GPT4v(key) + elif self.backbone_name == "gemini": + from .mllm_tools.gemini import Gemini + self.model = Gemini() + elif self.backbone_name == "idefics2": + from .mllm_tools.idefics2_eval import Idefics2 + self.model = Idefics2() + elif self.backbone_name == "mantis": + from .mllm_tools.mantis_idefics2_eval import Mantis + self.model = Mantis() + elif self.backbone_name == "minicpmv": + from .mllm_tools.minicpmv_eval import MiniCPMV + self.model = MiniCPMV() + elif self.backbone_name == "qwen25vl": + from .mllm_tools.qwen25vl_eval import Qwen25VL + self.model = Qwen25VL() + else: + raise NotImplementedError("backbone not supported") + + self.context = vie_prompts._context_no_delimit + if self.task == "t2i": + self.SC_prompt = "\n".join([self.context, vie_prompts._prompts_0shot_one_image_gen_rule, vie_prompts._prompts_0shot_t2i_rule_SC]) + self.PQ_prompt = "\n".join([self.context, vie_prompts._prompts_0shot_rule_PQ]) + elif self.task == "tie": + self.SC_prompt = "\n".join([self.context, vie_prompts._prompts_0shot_two_image_edit_rule, vie_prompts._prompts_0shot_tie_rule_SC]) + self.PQ_prompt = "\n".join([self.context, vie_prompts._prompts_0shot_rule_PQ]) + elif self.task == "t2v": + self.SC_prompt = "\n".join([self.context, vie_prompts._prompts_0shot_one_video_gen_rule, vie_prompts._prompts_0shot_t2v_rule_SC]) + self.PQ_prompt = "\n".join([self.context, vie_prompts._prompts_0shot_t2v_rule_PQ]) + + def evaluate(self, image_prompts, text_prompt, extract_overall_score_only=False, extract_all_score=True, echo_output=False): + if not isinstance(image_prompts, list): + image_prompts = [image_prompts] + if self.backbone_name in ['gpt4o', 'gpt4v', 'gpt-4.1', 'gpt-5']: + self.model.use_encode = False if isinstance(image_prompts[0], str) else True + + if self.task == "t2i": + _SC_prompt = self.SC_prompt.replace("", text_prompt) + elif self.task == "tie": + _SC_prompt = self.SC_prompt.replace("", text_prompt) + elif self.task == "t2v": + _SC_prompt = self.SC_prompt.replace("", text_prompt) + SC_prompt_final = self.model.prepare_prompt(image_prompts, _SC_prompt) + if self.task == "tie": + PQ_prompt_final = self.model.prepare_prompt(image_prompts[-1], self.PQ_prompt) + else: + PQ_prompt_final = self.model.prepare_prompt(image_prompts, self.PQ_prompt) + + results_dict = {} + + SC_dict = False + PQ_dict = False + tries = 0 + max_tries = 2 + while SC_dict is False or PQ_dict is False: + tries += 1 + guess_if_cannot_parse = True if tries > max_tries else False + result_SC = self.model.get_parsed_output(SC_prompt_final) + result_PQ = self.model.get_parsed_output(PQ_prompt_final) + + if result_SC in ["I'm sorry, but I can't assist with that request."] or result_PQ in ["I'm sorry, but I can't assist with that request."]: + guess_if_cannot_parse = True + + SC_dict = mllm_output_to_dict(result_SC, give_up_parsing=guess_if_cannot_parse, text_prompt=text_prompt) + PQ_dict = mllm_output_to_dict(result_PQ, give_up_parsing=guess_if_cannot_parse, text_prompt=text_prompt) + + if SC_dict == "rate_limit_exceeded" or PQ_dict == "rate_limit_exceeded": + print("rate_limit_exceeded") + raise ValueError("rate_limit_exceeded") + + if len(SC_dict['score']) == 0: + SC_dict['score'] = [5] + + if len(PQ_dict['score']) == 0: + PQ_dict['score'] = [5] + + results_dict['SC'] = SC_dict + results_dict['PQ'] = PQ_dict + if echo_output: + print("results_dict", results_dict, flush=True) + if extract_all_score: + try: + SC_score = min(results_dict['SC']['score']) + PQ_score = min(results_dict['PQ']['score']) + O_score = math.sqrt(SC_score * PQ_score) + return [SC_score, PQ_score, O_score] + except Exception as e: + print("results_dict", results_dict, flush=True) + raise ValueError("results_dict", e) + if extract_overall_score_only: + SC_scores = results_dict['SC']['score'] + PQ_scores = results_dict['PQ']['score'] + O_score = math.sqrt(min(SC_scores) * min(PQ_scores)) + return O_score + return results_dict + +if __name__ == "__main__": + model = VIEScore(backbone="gemini", task="t2i") + from datasets import load_dataset + dataset = load_dataset("TIGER-Lab/GenAI-Arena-Bench", "image_generation") + dataset = dataset["test"] + print("Now running the VIEScore model") + for idx in range(5): + left_image = dataset['left_image'][idx] + right_image = dataset['right_image'][idx] + prompt = dataset['prompt'][idx] + print(model.evaluate(left_image, prompt, extract_all_score=True)) + print(model.evaluate(right_image, prompt, extract_all_score=True)) + diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/__init__.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/gemini.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/gemini.py new file mode 100644 index 0000000000000000000000000000000000000000..85318c4a8f5ca33554485b881e4435d0aefdb9ce --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/gemini.py @@ -0,0 +1,147 @@ +""" +Install the Google AI Python SDK + +$ pip install google-generativeai + +See the getting started guide for more information: +https://ai.google.dev/gemini-api/docs/get-started/python +""" + +import requests +from PIL import Image +from io import BytesIO +import os +from typing import List +from urllib.parse import urlparse +import google.generativeai as genai +import tempfile + +genai.configure(api_key=os.environ["GEMINI_API_KEY"]) + +def upload_to_gemini(input, mime_type=None): + """Uploads the given file or PIL image to Gemini. + + See https://ai.google.dev/gemini-api/docs/prompting_with_media + """ + if isinstance(input, str): + # Input is a file path + file = genai.upload_file(input, mime_type=mime_type) + elif isinstance(input, Image.Image): + # Input is a PIL image + with tempfile.NamedTemporaryFile(suffix=".jpg", delete=False) as tmp_file: + input.save(tmp_file, format="JPEG") + tmp_file_path = tmp_file.name + file = genai.upload_file(tmp_file_path, mime_type=mime_type or "image/jpeg") + os.remove(tmp_file_path) + else: + raise ValueError("Unsupported input type. Must be a file path or PIL Image.") + + #print(f"Uploaded file '{file.display_name}' as: {file.uri}") + return file + +def save_image_from_url(url, base_save_directory='tmp', file_name=None): + # Parse the URL to create a directory path + parsed_url = urlparse(url) + url_path = os.path.join(parsed_url.netloc, parsed_url.path.lstrip('/')) + save_directory = os.path.join(base_save_directory, os.path.dirname(url_path)) + + # Create the directory if it doesn't exist + if not os.path.exists(save_directory): + os.makedirs(save_directory) + + # Get the image from the URL + response = requests.get(url) + if response.status_code == 200: + # Open the image + image = Image.open(BytesIO(response.content)) + + # Set the file name if not provided + if not file_name: + file_name = os.path.basename(parsed_url.path) + + # Save the image locally + file_path = os.path.join(save_directory, file_name) + image.save(file_path) + + return file_path + else: + raise Exception(f"Failed to retrieve image from URL. Status code: {response.status_code}") + +class Gemini(): + def __init__(self, model_name="gemini-1.5-pro-latest"): + # Create the model + # See https://ai.google.dev/api/python/google/generativeai/GenerativeModel + generation_config = { + "temperature": 1, + "top_p": 0.95, + "top_k": 64, + "max_output_tokens": 8192, + "response_mime_type": "text/plain", + } + safety_settings = [ + { + "category": "HARM_CATEGORY_HARASSMENT", + "threshold": "BLOCK_NONE", + }, + { + "category": "HARM_CATEGORY_HATE_SPEECH", + "threshold": "BLOCK_NONE", + }, + { + "category": "HARM_CATEGORY_SEXUALLY_EXPLICIT", + "threshold": "BLOCK_NONE", + }, + { + "category": "HARM_CATEGORY_DANGEROUS_CONTENT", + "threshold": "BLOCK_NONE", + }, + ] + self.model = genai.GenerativeModel( + model_name=model_name, + safety_settings=safety_settings, + generation_config=generation_config, + ) + + def prepare_prompt(self, image_links: List = [], text_prompt: str = ""): + if not isinstance(image_links, list): + image_links = [image_links] + + images_prompt = [] + for image_link in image_links: + if isinstance(image_link, str): + image = save_image_from_url(image_link) + else: + image = image_link + image = upload_to_gemini(image, mime_type="image/jpeg") + images_prompt.append(image) + + prompt_content = [images_prompt, text_prompt] + return prompt_content + + def get_parsed_output(self, prompt): + images_prompt = prompt[0] + text_prompt = prompt[1] + chat_session = self.model.start_chat( + history=[ + { + "role": "user", + "parts": images_prompt, + }, + ] + ) + try: + response = chat_session.send_message(text_prompt) + except: + return "Error in sending message to chat session." + return self.extract_response(response) + + def extract_response(self, response): + response = response.text + return response + +if __name__ == "__main__": + model = Gemini() + prompt = model.prepare_prompt(['https://chromaica.github.io/Museum/ImagenHub_Text-Guided_IE/DiffEdit/sample_34_1.jpg', 'https://chromaica.github.io/Museum/ImagenHub_Text-Guided_IE/input/sample_34_1.jpg'], 'What is difference between two images?') + print("prompt : \n", prompt) + res = model.get_parsed_output(prompt) + print("result : \n", res) \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/idefics2_eval.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/idefics2_eval.py new file mode 100644 index 0000000000000000000000000000000000000000..22cc8341411c51847ab9dab43c16db4d9d744954 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/idefics2_eval.py @@ -0,0 +1,43 @@ +import os +import torch +import time +from typing import List +from transformers import AutoProcessor, AutoModelForVision2Seq +from transformers.image_utils import load_image +from transformers.utils import is_flash_attn_2_available + + +class Idefics2(): + def __init__(self, model_path:str="HuggingFaceM4/idefics2-8b") -> None: + attn_implementation = "flash_attention_2" if is_flash_attn_2_available() else None + print(f"Using {attn_implementation} for attention implementation") + self.model = AutoModelForVision2Seq.from_pretrained(model_path, device_map="auto", torch_dtype=torch.float16, _attn_implementation=attn_implementation).eval() + self.processor = AutoProcessor.from_pretrained(model_path) + + def prepare_prompt(self, image_links: List = [], text_prompt: str = ""): + if not isinstance(image_links, list): + image_links = [image_links] + messages = [ + { + "role": "user", + "content": [ {"type": "image"}] * len(image_links) + [{"type": "text", "text": text_prompt}] + } + ] + prompt = self.processor.apply_chat_template(messages, add_generation_prompt=True) + images = [load_image(image_link) for image_link in image_links] #Support PIL images as well + inputs = self.processor(text=prompt, images=images, return_tensors="pt") + inputs = {k: v.to(self.model.device) for k, v in inputs.items()} + return inputs + + def get_parsed_output(self, inputs): + generate_ids = self.model.generate(**inputs, max_new_tokens=512, num_beams=1) + generated_text = self.processor.batch_decode(generate_ids[:, inputs['input_ids'].shape[1]:], skip_special_tokens=True, clean_up_tokenization_spaces=False)[0] + return generated_text + + +if __name__ == "__main__": + model = Idefics2() + prompt = model.prepare_prompt(['https://chromaica.github.io/Museum/ImagenHub_Text-Guided_IE/DiffEdit/sample_34_1.jpg', 'https://chromaica.github.io/Museum/ImagenHub_Text-Guided_IE/input/sample_34_1.jpg'], 'What is difference between two images?') + #print("prompt : \n", prompt) + res = model.get_parsed_output(prompt) + print("result : \n", res) \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/mantis_idefics2_eval.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/mantis_idefics2_eval.py new file mode 100644 index 0000000000000000000000000000000000000000..2c7711bff5146e9b0ea77d153cdf56d663e7d474 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/mantis_idefics2_eval.py @@ -0,0 +1,43 @@ +import os +import torch +import time +from typing import List +from transformers import AutoProcessor, AutoModelForVision2Seq +from transformers.image_utils import load_image +from transformers.utils import is_flash_attn_2_available + + +class Mantis(): + def __init__(self, model_path:str="TIGER-Lab/Mantis-8B-Idefics2") -> None: + attn_implementation = "flash_attention_2" if is_flash_attn_2_available() else None + print(f"Using {attn_implementation} for attention implementation") + self.model = AutoModelForVision2Seq.from_pretrained(model_path, device_map="auto", torch_dtype=torch.float16, _attn_implementation=attn_implementation).eval() + self.processor = AutoProcessor.from_pretrained(model_path) + + def prepare_prompt(self, image_links: List = [], text_prompt: str = ""): + if not isinstance(image_links, list): + image_links = [image_links] + messages = [ + { + "role": "user", + "content": [ {"type": "image"}] * len(image_links) + [{"type": "text", "text": text_prompt}] + } + ] + prompt = self.processor.apply_chat_template(messages, add_generation_prompt=True) + images = [load_image(image_link) for image_link in image_links] #Support PIL images as well + inputs = self.processor(text=prompt, images=images, return_tensors="pt") + inputs = {k: v.to(self.model.device) for k, v in inputs.items()} + return inputs + + def get_parsed_output(self, inputs): + generate_ids = self.model.generate(**inputs, max_new_tokens=512, num_beams=1) + generated_text = self.processor.batch_decode(generate_ids[:, inputs['input_ids'].shape[1]:], skip_special_tokens=True, clean_up_tokenization_spaces=False)[0] + return generated_text + + +if __name__ == "__main__": + model = Mantis() + prompt = model.prepare_prompt(['https://chromaica.github.io/Museum/ImagenHub_Text-Guided_IE/DiffEdit/sample_34_1.jpg', 'https://chromaica.github.io/Museum/ImagenHub_Text-Guided_IE/input/sample_34_1.jpg'], 'What is difference between two images?') + #print("prompt : \n", prompt) + res = model.get_parsed_output(prompt) + print("result : \n", res) \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/minicpmv_eval.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/minicpmv_eval.py new file mode 100644 index 0000000000000000000000000000000000000000..25732fdf723366735fd13bbed92077c152acc5f1 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/minicpmv_eval.py @@ -0,0 +1,42 @@ +import os +import torch +import time +from PIL import Image +from typing import List +from transformers import AutoModel, AutoTokenizer +from transformers.utils import is_flash_attn_2_available + +class MiniCPMV(): + def __init__(self) -> None: + attn_implementation = "flash_attention_2" if is_flash_attn_2_available() else None + self.model = AutoModel.from_pretrained('openbmb/MiniCPM-Llama3-V-2_5', trust_remote_code=True, torch_dtype=torch.float16, device_map='auto', _attn_implementation=attn_implementation).eval() + self.tokenizer = AutoTokenizer.from_pretrained('openbmb/MiniCPM-Llama3-V-2_5', trust_remote_code=True) + + print(f"Using {attn_implementation} for attention implementation") + + def prepare_prompt(self, image_links: List = [], text_prompt: str = ""): + if not isinstance(image_links, list): + image_links = [image_links] + messages = [ + { + "role": "user", + "content": [ {"type": "image"}] * len(image_links) + [{"type": "text", "text": text_prompt}] + } + ] + return messages + + def get_parsed_output(self, inputs): + res = self.model.chat( + image=None, + msgs=inputs, + tokenizer=self.tokenizer, + sampling=False, # if sampling=False, beam_search will be used by default + ) + return res + +if __name__ == "__main__": + model = MiniCPMV() + prompt = model.prepare_prompt(['https://chromaica.github.io/Museum/ImagenHub_Text-Guided_IE/DiffEdit/sample_34_1.jpg', 'https://chromaica.github.io/Museum/ImagenHub_Text-Guided_IE/input/sample_34_1.jpg'], 'What is difference between two images?') + #print("prompt : \n", prompt) + res = model.get_parsed_output(prompt) + print("result : \n", res) \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/openai.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/openai.py new file mode 100644 index 0000000000000000000000000000000000000000..81e5e44d390b2ec396b5f00411279493e239186d --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/openai.py @@ -0,0 +1,178 @@ +import base64 +import requests +from io import BytesIO, StringIO +from typing import Union, Optional, Tuple, List +from PIL import Image, ImageOps +import os + +def get_api_key(file_path): + # Read the API key from the first line of the file + with open(file_path, 'r') as file: + return file.readline().strip() + +# Function to encode the image +def encode_image(image_path): + with open(image_path, "rb") as image_file: + return base64.b64encode(image_file.read()).decode('utf-8') + +def pick_next_item(current_item, item_list): + if current_item not in item_list: + raise ValueError("Current item is not in the list") + current_index = item_list.index(current_item) + next_index = (current_index + 1) % len(item_list) + + return item_list[next_index] + +# Function to encode a PIL image +def encode_pil_image(pil_image): + # Create an in-memory binary stream + image_stream = BytesIO() + + # Save the PIL image to the binary stream in JPEG format (you can change the format if needed) + pil_image.save(image_stream, format='JPEG') + + # Get the binary data from the stream and encode it as base64 + image_data = image_stream.getvalue() + base64_image = base64.b64encode(image_data).decode('utf-8') + + return base64_image + + +def load_image(image: Union[str, Image.Image], format: str = "RGB", size: Optional[Tuple] = None) -> Image.Image: + """ + Load an image from a given path or URL and convert it to a PIL Image. + + Args: + image (Union[str, Image.Image]): The image path, URL, or a PIL Image object to be loaded. + format (str, optional): Desired color format of the resulting image. Defaults to "RGB". + size (Optional[Tuple], optional): Desired size for resizing the image. Defaults to None. + + Returns: + Image.Image: A PIL Image in the specified format and size. + + Raises: + ValueError: If the provided image format is not recognized. + """ + if isinstance(image, str): + if image.startswith("http://") or image.startswith("https://"): + image = Image.open(requests.get(image, stream=True).raw) + elif os.path.isfile(image): + image = Image.open(image) + else: + raise ValueError( + f"Incorrect path or url, URLs must start with `http://` or `https://`, and {image} is not a valid path" + ) + elif isinstance(image, Image.Image): + image = image + else: + raise ValueError( + "Incorrect format used for image. Should be an url linking to an image, a local path, or a PIL image." + ) + image = ImageOps.exif_transpose(image) + image = image.convert(format) + if (size != None): + image = image.resize(size, Image.LANCZOS) + return image + +class GPT4v(): + def __init__(self, key, url="https://api.openai.com/v1/chat/completions", are_images_encoded=False, model_name="gpt-4-vision-preview"): + """OpenAI GPT-4-vision model wrapper + Args: + api_key_path (str): Path to the API key file. Defaults to 'keys/secret.env'. + are_images_encoded (bool): Whether the images are encoded in base64. Defaults to False. + """ + self.multiple_api_keys = False + self.current_key_file = None + self.api_key = key + + self.url = url + self.model_name = model_name + self.use_encode = are_images_encoded + + def prepare_prompt(self, image_links: List = [], text_prompt: str = ""): + prompt_content = [] + text_dict = { + "type": "text", + "text": text_prompt + } + prompt_content.append(text_dict) + + if not isinstance(image_links, list): + image_links = [image_links] + for image_link in image_links: + image = load_image(image_link) + if self.use_encode: + visual_dict = { + "type": "image_url", + "image_url": {"url": f"data:image/jpeg;base64,{encode_pil_image(image)}"} + } + else: + visual_dict = { + "type": "image_url", + "image_url": {"url": image_link} + } + prompt_content.append(visual_dict) + return prompt_content + + def get_parsed_output(self, prompt): + payload = { + "model": self.model_name, + "messages": [ + { + "role": "user", + "content": prompt + } + ], + "max_tokens": 1400 + } + headers = { + "Content-Type": "application/json", + "Authorization": f"Bearer {self.api_key}" + } + try: + response = requests.post(self.url, json=payload, headers=headers, timeout=180) # Set timeout to 5 minutes (300 seconds) + except Exception as e: + print(f"Error: {e}") + return "" + #return response.text + return self.extract_response(response) + + def extract_response(self, response): + try: + response = response.json() + out = response['choices'][0]['message']['content'] + return out + except: + if response['error']['code'] == 'content_policy_violation': + print("Code is content_policy_violation") + elif response['error']['code'] in ['rate_limit_exceeded', 'insufficient_quota', 'insufficient_user_quota']: + print(f"Code is {response['error']['code']}", flush=True) + print(response['error']['message'], flush=True) + return "rate_limit_exceeded" + if self.multiple_api_keys == True: + new_key = pick_next_item(self.current_key_file, self.key_lists) + self.update_key(new_key) + self.current_key_file = new_key #override key + print("New key is from the file: ", new_key) + else: + print("Code is different") + print(response) + print(f"{response['error']['code']=}") + return "" + + def update_key(self, key, load_from_file=True): + if load_from_file: + self.api_key = get_api_key(key) + else: + self.api_key = key + +class GPT4o(GPT4v): + def __init__(self, key, url="https://api.openai.com/v1/chat/completions", are_images_encoded=False, model_name="gpt-4o-2024-05-13"): + super().__init__(key, url, are_images_encoded, model_name) + +if __name__ == "__main__": + model = GPT4o('sk-cB6h7HcCSDIp71gs6lFLZxKE0dOYOnJbxzES6kWXe1Wb2VHS', model_name="gpt-4.1") + prompt = model.prepare_prompt(['https://chromaica.github.io/Museum/ImagenHub_Text-Guided_IE/DiffEdit/sample_34_1.jpg', 'https://chromaica.github.io/Museum/ImagenHub_Text-Guided_IE/input/sample_34_1.jpg'], 'What is difference between two images?') + print("prompt : \n", prompt) + res = model.get_parsed_output(prompt) + print("result : \n", res) \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/qwen25_eval.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/qwen25_eval.py new file mode 100644 index 0000000000000000000000000000000000000000..d433036699fcbccce971b1d1c20b376fbf4ba47f --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/qwen25_eval.py @@ -0,0 +1,103 @@ +import os +import torch +import time +from PIL import Image +from typing import List +from transformers import AutoModel, AutoTokenizer +from transformers.utils import is_flash_attn_2_available +from transformers import AutoModelForCausalLM +from qwen_vl_utils import process_vision_info +from transformers import AutoTokenizer +import requests +from io import BytesIO +import random +import numpy as np +import base64 +import magic +import megfile + +def process_image(image): + img_byte_arr = BytesIO() + image.save(img_byte_arr, format='PNG') + img_byte_arr = img_byte_arr.getvalue() + return img_byte_arr + +def convert_image_to_base64(file_content): + mime_type = magic.from_buffer(file_content, mime=True) + base64_encoded_data = base64.b64encode(file_content).decode('utf-8') + return f"data:{mime_type};base64,{base64_encoded_data}" + + +def set_seed(seed: int): + """ + Args: + Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`. + seed (`int`): The seed to set. + """ + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + +class Qwen25(): + def __init__(self) -> None: + attn_implementation = "flash_attention_2" if is_flash_attn_2_available() else None + self.model = AutoModelForCausalLM.from_pretrained( + "/share_2/luoxin/modelscope/hub/models/Qwen/Qwen2.5-72B-Instruct", + torch_dtype="auto", + device_map="auto" + ) + self.tokenizer = AutoTokenizer.from_pretrained("/share_2/luoxin/modelscope/hub/models/Qwen/Qwen2.5-72B-Instruct") + + print(f"Using {attn_implementation} for attention implementation") + + def get_parsed_output(self, input_string): + set_seed(42) + # Prepare the inputs + messages = [ + {"role": "system", "content": "You are Qwen, created by Alibaba Cloud. You are a helpful assistant."}, + {"role": "user", "content": f"""Please rectify following json string to correct format. Expected format: + { + "score" : [...], + "reasoning" : "..." + } + Below is the input string: + {input_string}"""} + ] + + text = self.processor.apply_chat_template( + messages, tokenize=False, add_generation_prompt=True + ) + # Process inputs + inputs = self.tokenizer([text], return_tensors="pt").to(model.device) + inputs = inputs.to("cuda") + + generation_config = { + "max_new_tokens": 512, + "num_beams": 1, + "do_sample": True, + "temperature": 0.7, + "top_p": 0.8, + "top_k": 20, + } + + generated_ids = self.model.generate(**inputs, **generation_config) + generated_ids_trimmed = [ + out_ids[len(in_ids):] for in_ids, out_ids in zip(inputs.input_ids, generated_ids) + ] + output_text = self.processor.batch_decode( + generated_ids_trimmed, + skip_special_tokens=True, + clean_up_tokenization_spaces=False + ) + + return output_text[0] if output_text else "" + +if __name__ == "__main__": + model = Qwen25() + prompt = model.prepare_prompt( + ["https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-VL/assets/demo.jpeg"], + 'Describe the image in detail.' + ) + res = model.get_parsed_output(prompt) + print("result : \n", res) \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/qwen25vl_eval.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/qwen25vl_eval.py new file mode 100644 index 0000000000000000000000000000000000000000..846722c0bae52cc92e458abac3118c51050e54b7 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/qwen25vl_eval.py @@ -0,0 +1,132 @@ +import os +import torch +import time +from PIL import Image +from typing import List +from transformers import AutoModel, AutoTokenizer +from transformers.utils import is_flash_attn_2_available +from transformers import Qwen2_5_VLForConditionalGeneration +from qwen_vl_utils import process_vision_info +from transformers import AutoProcessor +import requests +from io import BytesIO +import random +import numpy as np +import base64 +import magic +import megfile + +def process_image(image): + img_byte_arr = BytesIO() + image.save(img_byte_arr, format='PNG') + img_byte_arr = img_byte_arr.getvalue() + return img_byte_arr + +def convert_image_to_base64(file_content): + mime_type = magic.from_buffer(file_content, mime=True) + base64_encoded_data = base64.b64encode(file_content).decode('utf-8') + return f"data:{mime_type};base64,{base64_encoded_data}" + + +def set_seed(seed: int): + """ + Args: + Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`. + seed (`int`): The seed to set. + """ + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + +class Qwen25VL(): + def __init__(self) -> None: + attn_implementation = "flash_attention_2" if is_flash_attn_2_available() else None + self.model = Qwen2_5_VLForConditionalGeneration.from_pretrained( + # "/share_2/luoxin/modelscope/hub/models/Qwen/Qwen2.5-VL-72B-Instruct-AWQ", + "/share/shared_models/Qwen2.5-VL-72B-Instruct/models--Qwen--Qwen2.5-VL-72B-Instruct/snapshots/5d8e171e5ee60e8ca4c6daa380bd29f78fe19021", + torch_dtype=torch.bfloat16, + device_map="auto" + ).eval() + # self.processor = AutoProcessor.from_pretrained("/share_2/luoxin/modelscope/hub/models/Qwen/Qwen2.5-VL-72B-Instruct-AWQ") + self.processor = AutoProcessor.from_pretrained("/share/shared_models/Qwen2.5-VL-72B-Instruct/models--Qwen--Qwen2.5-VL-72B-Instruct/snapshots/5d8e171e5ee60e8ca4c6daa380bd29f78fe19021") + + print(f"Using {attn_implementation} for attention implementation") + + def prepare_prompt(self, image_links: List = [], text_prompt: str = ""): + if not isinstance(image_links, list): + image_links = [image_links] + + image_links_base64 = [] + + for img_link in image_links: + if type(img_link) == str: + image_links_base64.append(convert_image_to_base64(process_image(megfile.smart_open(img_link, 'rb')))) + else: + image_links_base64.append(convert_image_to_base64(process_image(img_link))) + + messages = [ + { + "role": "user", + "content": [ + {"type": "image", "image": img_link} for img_link in image_links_base64 + ] + [{"type": "text", "text": text_prompt}] + } + ] + return messages + + def get_parsed_output(self, messages): + set_seed(42) + # Prepare the inputs + text = self.processor.apply_chat_template( + messages, tokenize=False, add_generation_prompt=True + ) + image_inputs, video_inputs = process_vision_info(messages) + + # Process inputs + inputs = self.processor( + text=[text], + images=image_inputs, + videos=video_inputs, + padding=True, + return_tensors="pt" + ) + inputs = inputs.to("cuda") + + # Generate output + # generation_config = { + # "max_new_tokens": 512, + # "num_beams": 1, + # "do_sample": False, + # "temperature": 0.1, + # "top_p": None, + # } + generation_config = { + "max_new_tokens": 512, + "num_beams": 1, + "do_sample": True, + "temperature": 0.7, + "top_p": 0.8, + "top_k": 20, + } + + generated_ids = self.model.generate(**inputs, **generation_config) + generated_ids_trimmed = [ + out_ids[len(in_ids):] for in_ids, out_ids in zip(inputs.input_ids, generated_ids) + ] + output_text = self.processor.batch_decode( + generated_ids_trimmed, + skip_special_tokens=True, + clean_up_tokenization_spaces=False + ) + + return output_text[0] if output_text else "" + +if __name__ == "__main__": + model = Qwen25VL() + prompt = model.prepare_prompt( + ["https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-VL/assets/demo.jpeg"], + 'Describe the image in detail.' + ) + res = model.get_parsed_output(prompt) + print("result : \n", res) \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/qwen25vl_vllm.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/qwen25vl_vllm.py new file mode 100644 index 0000000000000000000000000000000000000000..b9497a15d66e064f85d18a03df81d1422ed0e5c8 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/qwen25vl_vllm.py @@ -0,0 +1,113 @@ +from typing import List +from typing import Optional +import random +# import magic +# import megfile + +import numpy as np +import torch + +from vllm import LLM +from vllm.sampling_params import SamplingParams + +from qwen_vl_utils import process_vision_info + + +def set_seed(seed: int): + """ + Args: + Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`. + seed (`int`): The seed to set. + """ + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + + +def apply_chat_template(prompt, num_images: int = 2): + """ + This is used since the bug of transformers which do not support vision id https://github.com/QwenLM/Qwen2.5-VL/issues/716#issuecomment-2723316100 + """ + template = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n" + template += "".join([f": <|vision_start|><|image_pad|><|vision_end|>" for i in range(1, num_images + 1)]) + template += f"{prompt}<|im_end|>\n<|im_start|>assistant\n" + return template + + +class Qwen25VL(): + def __init__( + self, + vlm_model, + max_model_len: int = 1536, + tensor_parallel_size=1, + max_num_seqs=32, + max_num_batched_tokens=1536, + temperature: float = 0.7, + seed: Optional[int] = None, + ) -> None: + # attn_implementation = "flash_attention_2" if is_flash_attn_2_available() else None + self.model = LLM( + model=vlm_model, + max_model_len=max_model_len, + tensor_parallel_size=tensor_parallel_size, + max_num_seqs=max_num_seqs, + max_num_batched_tokens=max_num_batched_tokens, + limit_mm_per_prompt={"image": 2}, + enable_prefix_caching=True, + ) + self.temperature = temperature + self.seed = seed + + def prepare_input(self, images: List = [], text_prompt: str = ""): + if not isinstance(images, list): + images = [images] + + messages = [ + { + "role": "user", + "content": [{"type": "image", "image": image} for image in images] + + [{"type": "text", "text": text_prompt}], + } + ] + text = apply_chat_template(text_prompt, num_images=len(images)) + image_inputs, _ = process_vision_info(messages) + + messages = { + "prompt": text, + "multi_modal_data": {"image": image_inputs}, + } + return messages + + def inference(self, messages, seed: Optional[int] = None): + seed = self.seed if seed is None else seed + sampling_params = SamplingParams(max_tokens=512, temperature=self.temperature, top_p=0.9, top_k=20, seed=seed) + outputs = self.model.generate(messages, sampling_params, use_tqdm=False) + + responses = [] + for output in outputs: + instruction = output.outputs[0].text.strip() + responses.append(instruction) + + return responses[0] + +if __name__ == "__main__": + model = Qwen25VL( + vlm_model="Qwen/Qwen2.5-VL-7B-Instruct", + max_model_len=16384, + tensor_parallel_size=1, + max_num_seqs=32 + ) + + from PIL import Image + prompt = model.prepare_input( + [Image.open("https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-VL/assets/demo.jpeg")], + 'Describe the image in detail.' + ) + + prompt2 = model.prepare_input( + [Image.open("https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-VL/assets/demo.jpeg")], + 'How well it looks? Give a score between 0 and 100.' + ) + res = model.inference([prompt, prompt2]) + print("result : \n", res) \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/utils.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..3a85f3ae0c39245c5aa87a60e9bfcd0d361e5982 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/mllm_tools/utils.py @@ -0,0 +1,65 @@ +from typing import List +import base64 +from io import BytesIO +from PIL import Image +import requests + +def pil_image_to_base64(pil_image, format="PNG"): + buffered = BytesIO() + pil_image.save(buffered, format=format) # Save image to the buffer in the specified format + img_str = base64.b64encode(buffered.getvalue()).decode('utf-8') # Encode the buffer's content to base64 + return img_str + +def load_image(image_file): + if image_file.startswith("http"): + response = requests.get(image_file) + image = Image.open(BytesIO(response.content)).convert("RGB") + else: + import os + image = Image.open(image_file).convert("RGB") + return image + + +def load_images(image_files): + out = [] + for image_file in image_files: + image = load_image(image_file) + out.append(image) + return out + +def merge_images(image_links: List = []): + """Merge multiple images into one image + + Args: + image_links (List, optional): List of image links. Defaults to []. + + Returns: + [type]: [description] + """ + if len(image_links) == 0: + return None + images = load_images(image_links) + if len(images) == 1: + return images[0] + widths, heights = zip(*(i.size for i in images)) + average_height = sum(heights) // len(heights) + for i, im in enumerate(images): + # scale in proportion + images[i] = im.resize((int(im.size[0] * average_height / im.size[1]), average_height)) + widths, heights = zip(*(i.size for i in images)) + total_width = sum(widths) + max_height = max(heights) + new_im = Image.new("RGB", (total_width + 10 * (len(images) - 1), max_height)) + x_offset = 0 + for i, im in enumerate(images): + if i > 0: + # past a column of 1 pixel starting from x_offset width being black, 8 pixels being white, and 1 pixel being black + new_im.paste(Image.new("RGB", (1, max_height), (0, 0, 0)), (x_offset, 0)) + x_offset += 1 + new_im.paste(Image.new("RGB", (8, max_height), (255, 255, 255)), (x_offset, 0)) + x_offset += 8 + new_im.paste(Image.new("RGB", (1, max_height), (0, 0, 0)), (x_offset, 0)) + x_offset += 1 + new_im.paste(im, (x_offset, 0)) + x_offset += im.size[0] + return new_im \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/our_prompts.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/our_prompts.py new file mode 100644 index 0000000000000000000000000000000000000000..58533d528dbd5059636ae47c34be3cea90a2910c --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/our_prompts.py @@ -0,0 +1,29 @@ +_prompts_0shot_tie_rule_SC = """ +**Role:** You are an expert AI image quality evaluator. + +**Task:** You will be given two images: an **original** and an **edited version**. Your task is to evaluate the edit by comparing the edited image to the original based on the provided instruction. You will score it on two criteria. + +**Evaluation Criteria (Scale 0-10):** + +1. **Instruction Following**: How accurately was the instruction executed? + - 10: Perfectly executed. + - 5: Mostly executed, but with minor flaws or some effects not fully realized. + - 0: Completely ignored. + +2. **Image Consistency**: How well were unedited elements (background, subject identity, etc.) preserved? + - 10: No unnecessary changes; non-edited areas are identical to the original. + - 5: Noticeable but acceptable changes to unedited areas (e.g., slight background distortion). + - 0: The original image is unrecognizable due to massive, unintended changes. + +**Output Format:** +You MUST provide your output in a single JSON object. The `score` list must be in the order: `[Instruction Following score, Image Consistency score]`. Keep your reasoning concise. + +{ +"reasoning": "...", +"score": [int, int] +} + +**Note on Content:** All images are AI-generated. Do not comment on realism or privacy. + +**Instruction:** +""" diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/parse_prompt.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/parse_prompt.py new file mode 100644 index 0000000000000000000000000000000000000000..46a3b46fdaafbd46c2a9426bd1490c250fa2eefb --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/parse_prompt.py @@ -0,0 +1,20 @@ +import os + +def create_python_file_with_texts(folder_path, output_file): + with open(output_file, 'w', encoding='utf-8') as out_file: + out_file.write("# This file is generated automatically through parse_prompt.py\n\n") + for root, dirs, files in os.walk(folder_path): + for file in files: + if file.endswith(".txt"): + file_path = os.path.join(root, file) + var_name = "_" + file_path.replace(folder_path, "").replace(os.sep, "_").replace(".txt", "").strip("_") + with open(file_path, 'r', encoding='utf-8') as f: + content = f.read().replace('"""', '\"\"\"') + out_file.write(f'{var_name} = """{content}"""\n\n') + +# Example usage +current_file_path = os.path.abspath(__file__) +current_folder_path = os.path.dirname(current_file_path) +folder_path = os.path.join(current_folder_path, "prompts_raw") +output_file = os.path.join(current_folder_path, "vie_prompts.py") +create_python_file_with_texts(folder_path, output_file) diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/utils.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..206e2c099a9e7624a9daaeffba00a04c07812ef2 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/utils.py @@ -0,0 +1,418 @@ +import os +from typing import Union, List, Optional +import json +import regex as re +import ast +import random + +def fix_json(input_str): + # Add double quotes around keys using regex + fixed_str = re.sub(r'(\w+):', r'"\1":', input_str) + + # Add double quotes around string values if necessary and wrap int/float values in [] + def format_value(match): + key, value, comma = match.groups() + value = value.strip() + # Check if value is an integer or float + if re.match(r'^-?\d+(\.\d+)?$', value): + value = f'[{value}]' + # Check if value is a boolean or null + elif re.match(r'^(true|false|null)$', value, re.IGNORECASE): + pass # leave as is + else: + # Add quotes around string values + value = f'"{value}"' + return f'{key}: {value}{comma}' + + fixed_str = re.sub(r'(".*?"):(.*?)(,|})', format_value, fixed_str) + + return fixed_str + +def read_file_to_string(file_path): + """ + Reads the contents of a text file and returns it as a string. + + :param file_path: The path to the text file. + :return: A string containing the contents of the file. + """ + try: + with open(file_path, 'r', encoding='utf-8') as file: + return file.read() + except FileNotFoundError: + print(f"The file {file_path} was not found.") + return None + except Exception as e: + print(f"An error occurred: {e}") + return None + +def read_files_to_string(file_paths): + """ + Reads the contents of multiple text files and returns them as a single string, + with each file's contents separated by a newline. + + :param file_paths: A list of paths to text files. + :return: A string containing the concatenated contents of the files. + """ + all_contents = [] # List to hold the contents of each file + + for file_path in file_paths: + try: + with open(file_path, 'r', encoding='utf-8') as file: + all_contents.append(file.read()) + except FileNotFoundError: + print(f"The file {file_path} was not found.") + except Exception as e: + print(f"An error occurred while reading {file_path}: {e}") + + # Join all the contents with a newline character + return "\n".join(all_contents) + +def get_file_path(filename: Union[str, os.PathLike], search_from: Union[str, os.PathLike] = "."): + """ + Search for a file across a directory and return its absolute path. + + Args: + filename (Union[str, os.PathLike]): The name of the file to search for. + search_from (Union[str, os.PathLike], optional): The directory from which to start the search. Defaults to ".". + + Returns: + str: Absolute path to the found file. + + Raises: + FileNotFoundError: If the file is not found. + """ + for root, dirs, files in os.walk(search_from): + for name in files: + if name == filename: + return os.path.abspath(os.path.join(root, name)) + raise FileNotFoundError(filename, "not found.") + + + +#+========================================================================================= +def verify(s, target_sequence): + # Count the occurrences of the target sequence + count = s.count(target_sequence) + + # Check if the target sequence appears exactly twice + return count == 2 + + +def is_int_between_0_and_10(s): + try: + num = int(s) + return 0 <= num <= 10 + except ValueError: + return False + +def is_str_a_list_of_ints_0_to_10(s): + try: + # Attempt to parse the string as a Python literal (list, dict, etc.) + parsed = ast.literal_eval(s) + + # Check if the parsed object is a list + if not isinstance(parsed, list): + return False + + # Check if all elements are integers and between 0 to 10 + return all(isinstance(item, int) and 0 <= item <= 10 for item in parsed) + + except (ValueError, SyntaxError): + # If parsing fails or any other error occurs + return False + +def is_str_valid_score_format_brackets(s): + try: + # Removing brackets and splitting the string by commas + content = s.strip("[]").split(',') + + length = len(content) + + # Parsing each element and checking the format and range + scores = {} + for item in content: + key, value = item.split(':') + key = key.strip() + value = int(value.strip()) + + # Check if the key starts with 'score' and the value is in the correct range + if not key.startswith("score") or not 0 <= value <= 10: + return False + + scores[key] = value + + fetch_words = [f"score{i+1}" for i in range(length)] + # Check if at least 'score1' and 'score2' are present + return all(key in scores for key in fetch_words) + + except (ValueError, SyntaxError): + # If any parsing error occurs + return False + +def normalize_quotes(s: str) -> str: + """ + Replace curly/smart quotes with normal ASCII quotes. + """ + # ๅธธ่ง็š„ๅ‡ ็งๆ™บ่ƒฝๅผ•ๅท U+201C U+201D U+2018 U+2019 + return s.replace("โ€œ", '"').replace("โ€", '"').replace("โ€˜", "'").replace("โ€™", "'") + +def repair_reasoning_field_robust(json_str: str) -> str: + """ + ไฝฟ็”จๆญฃๅˆ™่กจ่พพๅผๅ’Œๅ…ˆ่กŒๆ–ญ่จ€๏ผŒๅฅๅฃฎๅœฐไฟฎๅค "reasoning" ๅญ—ๆฎตๅ†…้ƒจๆœช่ฝฌไน‰็š„ๅŒๅผ•ๅทใ€‚ + ๆญคๆ–นๆณ•ๅฏไปฅๅค„็† "reasoning" ๅญ—ๆฎตไธๆ˜ฏๆœ€ๅŽไธ€ไธชๅญ—ๆฎต็š„ๆƒ…ๅ†ตใ€‚ + + Args: + json_str: ๅฏ่ƒฝๅŒ…ๅซๆ ผๅผ้”™่ฏฏ็š„JSONๅญ—็ฌฆไธฒใ€‚ + + Returns: + ไฟฎๅคๅŽ็š„JSONๅญ—็ฌฆไธฒใ€‚ + """ + # 1. ๅฎšไน‰ๆ–ฐ็š„ๆญฃๅˆ™่กจ่พพๅผ๏ผŒไฝฟ็”จๆญฃๅ‘ๅ…ˆ่กŒๆ–ญ่จ€ๆฅๅฎšไฝ "reasoning" ๅ€ผ็š„็ป“ๆŸไฝ็ฝฎ + # re.DOTALL ๆ ‡ๅฟ—่ฎฉ '.' ๅฏไปฅๅŒน้…ๅŒ…ๆ‹ฌๆข่กŒ็ฌฆๅœจๅ†…็š„ไปปๆ„ๅญ—็ฌฆ + pattern = re.compile( + # --- ็ฌฌ1ไธชๆ•่Žท็ป„: reasoning ๅญ—ๆฎต็š„ "ๅ‰็ผ€" --- + r'("reasoning"\s*:\s*")' + + # --- ็ฌฌ2ไธชๆ•่Žท็ป„: reasoning ๅญ—ๆฎต็š„ "ๅ†…ๅฎน" --- + r'(.*?)' + + # --- ๆญฃๅ‘ๅ…ˆ่กŒๆ–ญ่จ€: ๅฏปๆ‰พๅ€ผ็š„็ป“ๆŸ่พน็•Œ๏ผŒไฝ†ไธๆถˆ่€—ๅฎƒ --- + # ๅŒน้…ๅˆฐ "reasoning" ๅ€ผ็š„็ป“ๆŸๅŒๅผ•ๅท๏ผŒ่ฟ™ไธชๅŒๅผ•ๅทๅŽ้ขๅฟ…้กป่ทŸ็€ไธ€ไธช้€—ๅทๆˆ–ไธ€ไธชๅณ่Šฑๆ‹ฌๅท + r'(?="\s*[,}])', + + re.DOTALL + ) + + # 2. ๅฎšไน‰ไธ€ไธชๆ›ด็ฎ€ๅ•็š„ๆ›ฟๆขๅ‡ฝๆ•ฐ + def replacer(match): + # ๆๅ–ๅ‡บไธคไธชๆ•่Žท็ป„ + prefix = match.group(1) # ไพ‹ๅฆ‚: '"reasoning" : "' + content = match.group(2) # ไพ‹ๅฆ‚: 'Overall building...' + + # ๅชๅœจ "ๅ†…ๅฎน" ้ƒจๅˆ†่ฟ›่กŒๆ›ฟๆข๏ผŒๅฐ†ๆ‰€ๆœ‰ๅŒๅผ•ๅท่ฝฌไน‰ + fixed_content = content.replace('"', '\\"') + + # ้‡ๆ–ฐ็ป„ๅˆใ€‚ๆณจๆ„๏ผšๆˆ‘ไปฌไธ้œ€่ฆๅค„็†ๅŽ็ผ€๏ผŒๅ› ไธบๅฎƒๆฒกๆœ‰่ขซๅŒน้…ๅ’Œๆถˆ่€—ๆމใ€‚ + return prefix + fixed_content + + # 3. ไฝฟ็”จ re.sub ๆ‰ง่กŒๆŸฅๆ‰พๅ’Œๆ›ฟๆข + repaired_str = pattern.sub(replacer, json_str) + + return repaired_str + +#+========================================================================================= +def mllm_output_to_dict(input_string, give_up_parsing=False, text_prompt=None): + """ + Args: + input_string (str): actually the output of the mllm model to be parsed + output_file_name (str): The name of the output file. + """ + # Catch for gpt4v rate_limit_exceeded error + if input_string == "rate_limit_exceeded": + return "rate_limit_exceeded" + + # Define the delimiters + delimiter = '||V^=^V||' + + if input_string.count(delimiter) == 2: + if not verify(input_string, delimiter): + print("The required delimiters were not found correctly in the string.", flush=True) + return False + # Extract the content between the delimiters + start_index = input_string.find(delimiter) + len(delimiter) + end_index = input_string.rfind(delimiter) + else: + # find the json mannually + # some mllm tends not to output the delimiters, but it does output the json contents + # so we will find the json content mannually + start_index = input_string.find('{') + end_index = input_string.rfind('}') + 1 + if start_index == -1 or end_index == 0: + # json not found + # some mllm tends to output only a list of scores like [6, 0], + # this time we will just get the scores and ignore the reasoning (other part of the json) + start_index = input_string.find('[') + end_index = input_string.rfind(']') + 1 + if give_up_parsing: # if we want to give up parsing + guessed_value = random.randint(0, 10) + print(f"1111 Failed to find the json content in the string. Guess a value : {text_prompt=} {input_string=} {guessed_value=}.", flush=True) + json_content = {'score': [guessed_value], "reasoning": f"guess_if_cannot_parse | {input_string}"} + json_str = json.dumps(json_content) + input_string = json_str + start_index = 0 + end_index = len(json_str) + elif re.match(r'^\[\d+, ?\d+\]$', input_string[start_index:end_index]): + scores = json.loads(input_string[start_index:end_index]) + if not isinstance(scores, list): + scores = [scores] + json_content = {'score': scores, "reasoning": "System: output is simply a list of scores"} + json_str = json.dumps(json_content) + input_string = json_str + start_index = 0 + end_index = len(json_str) + elif is_int_between_0_and_10(input_string): # if output is simply a number + scores = [int(input_string)] + json_content = {'score': scores, "reasoning": "System: output is simply a number"} + json_str = json.dumps(json_content) + input_string = json_str + start_index = 0 + end_index = len(json_str) + else: + print(f"22 222 Failed to find the json content in the string. {text_prompt=} {input_string=}", flush=True) + return False + + # Check if we found two delimiters + if start_index != -1 and end_index != -1 and start_index != end_index: + # Extract the JSON string + json_str = input_string[start_index:end_index].strip() + json_str = json_str.replace("\n", "") + # Parse the JSON string into a dictionary + try: + json_str = normalize_quotes(json_str) + new_data = json.loads(json_str) + if not isinstance(new_data['score'], list): + new_data['score'] = [new_data['score']] + except Exception as e1: + print(f"Now fixing: {e1=} {json_str=}") + try: + new_data = json.loads(fix_json(json_str)) + return new_data + except Exception as e2: + try: + print(f"Now fixing: {e2=} {fix_json(json_str)=}") + new_data = json.loads(repair_reasoning_field_robust(json_str)) + return new_data + except Exception as e3: + print(f"Error: Cannot fix {e3=} {repair_reasoning_field_robust(json_str)=}") + return False + return new_data + else: + print("The required delimiters were not found correctly in the string.") + return False + +def write_entry_to_json_file(input_string, uid, prompt_input, vision_input, output_file_name, give_up_parsing=False): + """ + Args: + input_string (str): actually the output of the mllm model to be parsed + uid (str): The unique identifier for the each item in the test data + prompt_input (str): The prompt input for the entry. text prompt. + vision_input (str): The vision input for the entry. image links. + output_file_name (str): The name of the output file. + """ + # Catch for gpt4v rate_limit_exceeded error + if input_string == "rate_limit_exceeded": + return "rate_limit_exceeded" + + # Define the delimiters + delimiter = '||V^=^V||' + + if input_string.count(delimiter) == 2: + if not verify(input_string, delimiter): + print("The required delimiters were not found correctly in the string.") + return False + # Extract the content between the delimiters + start_index = input_string.find(delimiter) + len(delimiter) + end_index = input_string.rfind(delimiter) + else: + # find the json mannually + # some mllm tends not to output the delimiters, but it does output the json contents + # so we will find the json content mannually + start_index = input_string.find('{') + end_index = input_string.rfind('}') + 1 + if start_index == -1 or end_index == 0: + # json not found + # some mllm tends to output only a list of scores like [6, 0], + # this time we will just get the scores and ignore the reasoning (other part of the json) + start_index = input_string.find('[') + end_index = input_string.rfind(']') + 1 + if give_up_parsing: # if we want to give up parsing + guessed_value = random.randint(0, 10) + print(f"Failed to find the json content in the string. Guess a value : {guessed_value}.") + json_content = {'score': [guessed_value], "reasoning": f"guess_if_cannot_parse | {input_string}"} + json_str = json.dumps(json_content) + input_string = json_str + start_index = 0 + end_index = len(json_str) + elif re.match(r'^\[\d+, ?\d+\]$', input_string[start_index:end_index]): + scores = json.loads(input_string[start_index:end_index]) + json_content = {'score': scores, "reasoning": None} + json_str = json.dumps(json_content) + input_string = json_str + start_index = 0 + end_index = len(json_str) + elif is_int_between_0_and_10(input_string): # if output is simply a number + scores = [int(input_string)] + json_content = {'score': scores, "reasoning": None} + json_str = json.dumps(json_content) + input_string = json_str + start_index = 0 + end_index = len(json_str) + else: + print("Failed to find the json content in the string.") + return False + + # Check if we found two delimiters + if start_index != -1 and end_index != -1 and start_index != end_index: + # Extract the JSON string + json_str = input_string[start_index:end_index].strip() + json_str = json_str.replace("\n", "") + try: + # Parse the JSON string into a dictionary + new_data = json.loads(json_str) + + # Ensure the directory exists + os.makedirs(os.path.dirname(output_file_name), exist_ok=True) + + # Initialize or load existing data + if os.path.exists(output_file_name): + with open(output_file_name, 'r') as json_file: + data = json.load(json_file) + else: + data = {} + + # If the additional key is already in the data, add or update notes + if uid in data: + data[uid].update(new_data) # Update with new data + if prompt_input: # If there are new notes, update or add them + data[uid]['prompt_input'] = prompt_input + if vision_input: # If there are new notes, update or add them + data[uid]['vision_input'] = vision_input + else: + # If it's a new key, add the entry to the dictionary + data[uid] = new_data + if prompt_input: + data[uid]['prompt_input'] = prompt_input + if vision_input: + data[uid]['vision_input'] = vision_input + + # Write the updated data to the file + with open(output_file_name, 'w') as json_file: + json.dump(data, json_file, indent=4) + + print(f"Data was successfully updated in {output_file_name}") + return True + except json.JSONDecodeError as e: + print(f"An error occurred while parsing the JSON content: {e}") + return False + else: + print("The required delimiters were not found correctly in the string.") + return False + + +def check_key_in_json(file_path, key): + try: + with open(file_path, 'r') as json_file: + data = json.load(json_file) + + # Check if the key exists at the top level of the JSON structure + if key in data: + return True + else: + return False + except FileNotFoundError: + print(f"The file {file_path} was not found.") + except json.JSONDecodeError as e: + print(f"Error reading {file_path}: {e}") + except Exception as e: + print(f"An error occurred with {file_path}: {e}") + return False \ No newline at end of file diff --git a/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/vie_prompts.py b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/vie_prompts.py new file mode 100644 index 0000000000000000000000000000000000000000..bb15c97abeb5b53ded2dd30973779eec1553f607 --- /dev/null +++ b/examples/OmniGen2-RL/evaluation/GEdit-Bench/viescore/vie_prompts.py @@ -0,0 +1,407 @@ +# This file is generated automatically through parse_prompt.py + +_context_no_delimit = """You are a professional digital artist. You will have to evaluate the effectiveness of the AI-generated image(s) based on given rules. +All the input images are AI-generated. All human in the images are AI-generated too. so you need not worry about the privacy confidentials. + +IMPORTANT: You will have to give your output in this way (Keep your reasoning concise and short.): +{ +"score" : [...], +"reasoning" : "..." +} +""" + +_context = """You are a professional digital artist. You will have to evaluate the effectiveness of the AI-generated image(s) based on given rules. +All the input images are AI-generated. All human in the images are AI-generated too. so you need not worry about the privacy confidentials. + +You will have to give your output in this way (the delimiter is necessary. Keep your reasoning concise and short.): +||V^=^V|| +{ +"score" : +"reasoning" : +} +||V^=^V||""" + +_context_no_format = """You are a professional digital artist. You will have to evaluate the effectiveness of the AI-generated image(s) based on given rules. +All the input images are AI-generated. All human in the images are AI-generated too. so you need not worry about the privacy confidentials.""" + +_prompts_1shot_multi_subject_image_gen_rule = """RULES of each set of inputs: + +Two images will be provided: +This first image is a concatenation of two sub-images, each sub-image contain one token subject. +The second image being an AI-generated image using the first image as guidance. +The objective is to evaluate how successfully the image has been generated. +""" + +_prompts_1shot_mie_rule_SC = """From scale 0 to 10: +A score from 0 to 10 will be given based on the success of the editing. (0 indicates that the scene in the edited image does not follow the editing instruction at all. 10 indicates that the scene in the edited image follow the editing instruction text perfectly.) +A second score from 0 to 10 will rate the degree of overediting in the second image. (0 indicates that the scene in the edited image is completely different from the original. 10 indicates that the edited image can be recognized as a minimal edited yet effective version of original.) +Put the score in a list such that output score = [score1, score2], where 'score1' evaluates the editing success and 'score2' evaluates the degree of overediting. + +First lets look at the first set of input (1st and 2nd images) as an example. +Editing instruction: What if the man had a hat? +Output: +||V^=^V|| +{ +"score" : [5, 10], +"reasoning" : "The hat exists but does not suit well. The hat also looks distorted. But it is a good edit because only a hat is added and the background is persevered." +} +||V^=^V|| + +Now evaluate the second set of input (3th, 4th images). +Editing instruction: +""" + +_prompts_1shot_msdig_rule_SC = """From scale 0 to 10: +A score from 0 to 10 will be given based on the success in following the prompt. +(0 indicates that the second image does not follow the prompt at all. 10 indicates the second image follows the prompt perfectly.) +A second score from 0 to 10 will rate how well the subject in the generated image resemble to the token subject in the first sub-image. +(0 indicates that the subject in the second image does not look like the token subject in the first sub-image at all. 10 indicates the subject in the second image look exactly alike the token subject in the first sub-image.) +A third score from 0 to 10 will rate how well the subject in the generated image resemble to the token subject in the second sub-image. +(0 indicates that the subject in the second image does not look like the token subject in the second sub-image at all. 10 indicates the subject in the second image look exactly alike the token subject in the second sub-image.) +Put the score in a list such that output score = [score1, score2, score3], where 'score1' evaluates the prompt and 'score2' evaluates the resemblance for the first sub-image, and 'score3' evaluates the resemblance for the second sub-image. + +First lets look at the first set of input (1st and 2nd images) as an example. +Text Prompt: A digital illustration of a cat beside a wooden pot +Output: +||V^=^V|| +{ +"score" : [5, 5, 10], +"reasoning" : "The cat is not beside the wooden pot. The pot looks partially resemble to the subject pot. The cat looks highly resemble to the subject cat." +} +||V^=^V|| + +Now evaluate the second set of input (3th, 4th images). +Text Prompt: """ + +_prompts_1shot_t2i_rule_SC = """From scale 0 to 10: +A score from 0 to 10 will be given based on the success in following the prompt. +(0 indicates that the AI generated image does not follow the prompt at all. 10 indicates the AI generated image follows the prompt perfectly.) + +Put the score in a list such that output score = [score]. + +First lets look at the first set of input (1st image) as an example. +Text Prompt: A pink and a white frisbee are on the ground. +Output: +||V^=^V|| +{ +"score" : [5], +"reasoning" : "White frisbee not present in the image." +} +||V^=^V|| + +Now evaluate the second set of input (2nd image). +Text Prompt: +""" + +_prompts_1shot_tie_rule_SC = """From scale 0 to 10: +A score from 0 to 10 will be given based on the success of the editing. (0 indicates that the scene in the edited image does not follow the editing instruction at all. 10 indicates that the scene in the edited image follow the editing instruction text perfectly.) +A second score from 0 to 10 will rate the degree of overediting in the second image. (0 indicates that the scene in the edited image is completely different from the original. 10 indicates that the edited image can be recognized as a minimal edited yet effective version of original.) +Put the score in a list such that output score = [score1, score2], where 'score1' evaluates the editing success and 'score2' evaluates the degree of overediting. + +First lets look at the first set of input (1st and 2nd images) as an example. +Editing instruction: What if the man had a hat? +Output: +||V^=^V|| +{ +"score" : [5, 10], +"reasoning" : "The hat exists but does not suit well. The hat also looks distorted. But it is a good edit because only a hat is added and the background is persevered." +} +||V^=^V|| + +Now evaluate the second set of input (3th, 4th images). +Editing instruction: +""" + +_prompts_1shot_sdie_rule_SC = """From scale 0 to 10: +A score from 0 to 10 will rate how well the subject in the generated image resemble to the token subject in the second image. +(0 indicates that the subject in the third image does not look like the token subject at all. 10 indicates the subject in the third image look exactly alike the token subject.) +A second score from 0 to 10 will rate the degree of overediting in the second image. +(0 indicates that the scene in the edited image is completely different from the first image. 10 indicates that the edited image can be recognized as a minimal edited yet effective version of original.) +Put the score in a list such that output score = [score1, score2], where 'score1' evaluates the resemblance and 'score2' evaluates the degree of overediting. + +First lets look at the first set of input (1st, 2nd and 3rd images) as an example. +Subject: +Output: +||V^=^V|| +{ +"score" : [5, 10], +"reasoning" : "The monster toy looks partially resemble to the token subject. The edit is minimal." +} +||V^=^V|| + +Now evaluate the second set of input (4th, 5th, and 6th images). +Subject: +""" + +_prompts_1shot_one_image_gen_rule = """RULES of each set of inputs: + +One image will be provided; The image is an AI-generated image. +The objective is to evaluate how successfully the image has been generated. +""" + +_prompts_1shot_sdig_rule_SC = """From scale 0 to 10: +A score from 0 to 10 will be given based on the success in following the prompt. +(0 indicates that the second image does not follow the prompt at all. 10 indicates the second image follows the prompt perfectly.) +A second score from 0 to 10 will rate how well the subject in the generated image resemble to the token subject in the first image. +(0 indicates that the subject in the second image does not look like the token subject at all. 10 indicates the subject in the second image look exactly alike the token subject.) +Put the score in a list such that output score = [score1, score2], where 'score1' evaluates the prompt and 'score2' evaluates the resemblance. + +First lets look at the first set of input (1st and 2nd images) as an example. +Text Prompt: a red cartoon figure eating a banana +Output: +||V^=^V|| +{ +"score" : [10, 5], +"reasoning" : "The red cartoon figure is eating a banana. The red cartoon figure looks partially resemble to the subject." +} +||V^=^V|| + +Now evaluate the second set of input (3th, 4th images). +Text Prompt: +""" + +_prompts_1shot_rule_PQ = """RULES of each set of inputs: + +One image will be provided; The image is an AI-generated image. +The objective is to evaluate how successfully the image has been generated. + +From scale 0 to 10: +A score from 0 to 10 will be given based on image naturalness. +( + 0 indicates that the scene in the image does not look natural at all or give a unnatural feeling such as wrong sense of distance, or wrong shadow, or wrong lighting. + 10 indicates that the image looks natural. +) +A second score from 0 to 10 will rate the image artifacts. +( + 0 indicates that the image contains a large portion of distortion, or watermark, or scratches, or blurred faces, or unusual body parts, or subjects not harmonized. + 10 indicates the image has no artifacts. +) +Put the score in a list such that output score = [naturalness, artifacts] + + +First lets look at the first set of input (1st image) as an example. +Output: +||V^=^V|| +{ +"score" : [5, 5], +"reasoning" : "The image gives an unnatural feeling on hands of the girl. There is also minor distortion on the eyes of the girl." +} +||V^=^V|| + +Now evaluate the second set of input (2nd image). + +""" + +_prompts_1shot_subject_image_gen_rule = """RULES of each set of inputs: + +Two images will be provided: The first being a token subject image and the second being an AI-generated image using the first image as guidance. +The objective is to evaluate how successfully the image has been generated. +""" + +_prompts_1shot_cig_rule_SC = """ +From scale 0 to 10: +A score from 0 to 10 will be given based on the success in following the prompt. +(0 indicates that the second image does not follow the prompt at all. 10 indicates the second image follows the prompt perfectly.) +A second score from 0 to 10 will rate how well the generated image is following the guidance image. +(0 indicates that the second image is not following the guidance at all. 10 indicates that second image is following the guidance image.) +Put the score in a list such that output score = [score1, score2], where 'score1' evaluates the prompt and 'score2' evaluates the guidance. + +First lets look at the first set of input (1st and 2nd images) as an example. +Text Prompt: the bridge is red, Golden Gate Bridge in San Francisco, USA +Output: +||V^=^V|| +{ +"score" : [5, 5], +"reasoning" : "The bridge is red. But half of the bridge is gone." +} +||V^=^V|| + +Now evaluate the second set of input (3th, 4th images). +Text Prompt: +""" + +_prompts_1shot_two_image_edit_rule = """RULES of each set of inputs: + +Two images will be provided: The first being the original AI-generated image and the second being an edited version of the first. +The objective is to evaluate how successfully the editing instruction has been executed in the second image. + +Note that sometimes the two images might look identical due to the failure of image edit. +""" + +_prompts_1shot_subject_image_edit_rule = """RULES of each set of inputs: + +Three images will be provided: +The first image is a input image to be edited. +The second image is a token subject image. +The third image is an AI-edited image from the first image. it should contain a subject that looks alike the subject in second image. +The objective is to evaluate how successfully the image has been edited. +""" + +_prompts_1shot_control_image_gen_rule = """RULES of each set of inputs: + +Two images will be provided: The first being a processed image (e.g. Canny edges, openpose, grayscale etc.) and the second being an AI-generated image using the first image as guidance. +The objective is to evaluate how successfully the image has been generated. +""" + +_prompts_0shot_two_image_edit_rule = """RULES: + +Two images will be provided: The first being the original AI-generated image and the second being an edited version of the first. +The objective is to evaluate how successfully the editing instruction has been executed in the second image. + +Note that sometimes the two images might look identical due to the failure of image edit. +""" + +_prompts_0shot_one_video_gen_rule = """RULES: + +The images are extracted from a AI-generated video according to the text prompt. +The objective is to evaluate how successfully the video has been generated. +""" + +_prompts_0shot_t2v_rule_PQ = """RULES: + +The image frames are AI-generated. +The objective is to evaluate how successfully the image frames has been generated. + +From scale 0 to 10: +A score from 0 to 10 will be given based on the image frames naturalness. +( + 0 indicates that the scene in the image frames does not look natural at all or give a unnatural feeling such as wrong sense of distance, or wrong shadow, or wrong lighting. + 10 indicates that the image frames looks natural. +) +A second score from 0 to 10 will rate the image frames artifacts. +( + 0 indicates that the image frames contains a large portion of distortion, or watermark, or scratches, or blurred faces, or unusual body parts, or subjects not harmonized. + 10 indicates the image frames has no artifacts. +) +Put the score in a list such that output score = [naturalness, artifacts] +""" + +_prompts_0shot_msdig_rule_SC = """From scale 0 to 10: +A score from 0 to 10 will be given based on the success in following the prompt. +(0 indicates that the second image does not follow the prompt at all. 10 indicates the second image follows the prompt perfectly.) +A second score from 0 to 10 will rate how well the subject in the generated image resemble to the token subject in the first sub-image. +(0 indicates that the subject in the second image does not look like the token subject in the first sub-image at all. 10 indicates the subject in the second image look exactly alike the token subject in the first sub-image.) +A third score from 0 to 10 will rate how well the subject in the generated image resemble to the token subject in the second sub-image. +(0 indicates that the subject in the second image does not look like the token subject in the second sub-image at all. 10 indicates the subject in the second image look exactly alike the token subject in the second sub-image.) +Put the score in a list such that output score = [score1, score2, score3], where 'score1' evaluates the prompt and 'score2' evaluates the resemblance for the first sub-image, and 'score3' evaluates the resemblance for the second sub-image. + +Text Prompt: +""" + +_prompts_0shot_sdie_rule_SC = """From scale 0 to 10: +A score from 0 to 10 will rate how well the subject in the generated image resemble to the token subject in the second image. +(0 indicates that the subject in the third image does not look like the token subject at all. 10 indicates the subject in the third image look exactly alike the token subject.) +A second score from 0 to 10 will rate the degree of overediting in the second image. +(0 indicates that the scene in the edited image is completely different from the first image. 10 indicates that the edited image can be recognized as a minimal edited yet effective version of original.) +Put the score in a list such that output score = [score1, score2], where 'score1' evaluates the resemblance and 'score2' evaluates the degree of overediting. + +Subject: """ + +_prompts_0shot_subject_image_edit_rule = """RULES: + +Three images will be provided: +The first image is a input image to be edited. +The second image is a token subject image. +The third image is an AI-edited image from the first image. it should contain a subject that looks alike the subject in second image. +The objective is to evaluate how successfully the image has been edited. +""" + +_prompts_0shot_mie_rule_SC = """From scale 0 to 10: +A score from 0 to 10 will be given based on the success of the editing. (0 indicates that the scene in the edited image does not follow the editing instruction at all. 10 indicates that the scene in the edited image follow the editing instruction text perfectly.) +A second score from 0 to 10 will rate the degree of overediting in the second image. (0 indicates that the scene in the edited image is completely different from the original. 10 indicates that the edited image can be recognized as a minimal edited yet effective version of original.) +Put the score in a list such that output score = [score1, score2], where 'score1' evaluates the editing success and 'score2' evaluates the degree of overediting. + +Editing instruction: +""" + +_prompts_0shot_sdig_rule_SC = """From scale 0 to 10: +A score from 0 to 10 will be given based on the success in following the prompt. +(0 indicates that the second image does not follow the prompt at all. 10 indicates the second image follows the prompt perfectly.) +A second score from 0 to 10 will rate how well the subject in the generated image resemble to the token subject in the first image. +(0 indicates that the subject in the second image does not look like the token subject at all. 10 indicates the subject in the second image look exactly alike the token subject.) +Put the score in a list such that output score = [score1, score2], where 'score1' evaluates the prompt and 'score2' evaluates the resemblance. + +Text Prompt: +""" + +_prompts_0shot_tie_rule_SC = """ +From scale 0 to 10: +A score from 0 to 10 will be given based on the success of the editing. (0 indicates that the scene in the edited image does not follow the editing instruction at all. 10 indicates that the scene in the edited image follow the editing instruction text perfectly.) +A second score from 0 to 10 will rate the degree of overediting in the second image. (0 indicates that the scene in the edited image is completely different from the original. 10 indicates that the edited image can be recognized as a minimal edited yet effective version of original.) +Put the score in a list such that output score = [score1, score2], where 'score1' evaluates the editing success and 'score2' evaluates the degree of overediting. + +Editing instruction: +""" + +_prompts_0shot_t2i_rule_SC = """From scale 0 to 10: +A score from 0 to 10 will be given based on the success in following the prompt. +(0 indicates that the AI generated image does not follow the prompt at all. 10 indicates the AI generated image follows the prompt perfectly.) + +Put the score in a list such that output score = [score]. + +Text Prompt: +""" + +_prompts_0shot_cig_rule_SC = """From scale 0 to 10: +A score from 0 to 10 will be given based on the success in following the prompt. +(0 indicates that the second image does not follow the prompt at all. 10 indicates the second image follows the prompt perfectly.) +A second score from 0 to 10 will rate how well the generated image is following the guidance image. +(0 indicates that the second image is not following the guidance at all. 10 indicates that second image is following the guidance image.) +Put the score in a list such that output score = [score1, score2], where 'score1' evaluates the prompt and 'score2' evaluates the guidance. + +Text Prompt: """ + +_prompts_0shot_control_image_gen_rule = """RULES: + +Two images will be provided: The first being a processed image (e.g. Canny edges, openpose, grayscale etc.) and the second being an AI-generated image using the first image as guidance. +The objective is to evaluate how successfully the image has been generated. +""" + +_prompts_0shot_rule_PQ = """RULES: + +The image is an AI-generated image. +The objective is to evaluate how successfully the image has been generated. + +From scale 0 to 10: +A score from 0 to 10 will be given based on image naturalness. +( + 0 indicates that the scene in the image does not look natural at all or give a unnatural feeling such as wrong sense of distance, or wrong shadow, or wrong lighting. + 10 indicates that the image looks natural. +) +A second score from 0 to 10 will rate the image artifacts. +( + 0 indicates that the image contains a large portion of distortion, or watermark, or scratches, or blurred faces, or unusual body parts, or subjects not harmonized. + 10 indicates the image has no artifacts. +) +Put the score in a list such that output score = [naturalness, artifacts] +""" + +_prompts_0shot_t2v_rule_SC = """From scale 0 to 10: +A score from 0 to 10 will be given based on the success in following the prompt. +(0 indicates that the image frames does not follow the prompt at all. 10 indicates the image frames follows the prompt perfectly.) + +Put the score in a list such that output score = [score]. + +Text Prompt: +""" + +_prompts_0shot_multi_subject_image_gen_rule = """RULES: + +Two images will be provided: +This first image is a concatenation of two sub-images, each sub-image contain one token subject. +The second image being an AI-generated image using the first image as guidance. +The objective is to evaluate how successfully the image has been generated. +""" + +_prompts_0shot_subject_image_gen_rule = """RULES: + +Two images will be provided: The first being a token subject image and the second being an AI-generated image using the first image as guidance. +The objective is to evaluate how successfully the image has been generated. +""" + +_prompts_0shot_one_image_gen_rule = """RULES: + +The image is an AI-generated image according to the text prompt. +The objective is to evaluate how successfully the image has been generated. +""" + diff --git a/examples/OmniGen2-RL/nccl_logs/.gitkeep b/examples/OmniGen2-RL/nccl_logs/.gitkeep new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/examples/OmniGen2-RL/omnigen2/__init__.py b/examples/OmniGen2-RL/omnigen2/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/examples/OmniGen2-RL/omnigen2/dataset/__init__.py b/examples/OmniGen2-RL/omnigen2/dataset/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/examples/OmniGen2-RL/omnigen2/dataset/omnigen2_train_dataset.py b/examples/OmniGen2-RL/omnigen2/dataset/omnigen2_train_dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..ac9c94b9dfcaa361ca560c89aaf6c51dd127f14c --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/dataset/omnigen2_train_dataset.py @@ -0,0 +1,394 @@ +from typing import Optional, Union, List + +import os +import random +import time +import math +import re +import yaml +import glob +from PIL import Image + +import torch +from torchvision import transforms + +from datasets import load_dataset, concatenate_datasets + +from ..pipelines.omnigen2.pipeline_omnigen2 import OmniGen2ImageProcessor + +from accelerate.logging import get_logger +logger = get_logger(__name__) + + +def normalize_whitespace(input_string): + # ๆ›ฟๆข่ฟž็ปญ็š„็ฉบๆ ผไธบไธ€ไธช็ฉบๆ ผ๏ผŒๅนถๅค„็†ไธŽๅˆถ่กจ็ฌฆๆˆ–ๆข่กŒ็ฌฆ็›ธ้‚ป็š„ๆƒ…ๅ†ต + input_string = re.sub(r' +\t+', '\n', input_string) + input_string = re.sub(r' +\n+', '\n', input_string) + input_string = re.sub(r'\t+ +', '\n', input_string) + input_string = re.sub(r'\n+ +', '\n', input_string) + + # ๆ›ฟๆข่ฟž็ปญ็š„็ฉบๆ ผไธบไธ€ไธช็ฉบๆ ผ + input_string = re.sub(r' +', ' ', input_string) + + # ๆ›ฟๆข่ฟž็ปญ็š„ๅˆถ่กจ็ฌฆไธบไธ€ไธชๅˆถ่กจ็ฌฆ + input_string = re.sub(r'\t+', '\t', input_string) + + # ๆ›ฟๆข่ฟž็ปญ็š„ๆข่กŒ็ฌฆไธบไธ€ไธชๆข่กŒ็ฌฆ + input_string = re.sub(r'\n+', '\n', input_string) + + return input_string.strip() + + +class OmniGen2TrainDataset(torch.utils.data.Dataset): + SYSTEM_PROMPT = "You are a helpful assistant that generates high-quality images based on user instructions." + SYSTEM_PROMPT_DROP = "You are a helpful assistant that generates images." + + def __init__( + self, + config_path: str, + tokenizer, + num_workers: int, + use_chat_template: bool, + max_input_pixels: Optional[Union[int, List[int]]] = None, + max_output_pixels: Optional[int] = None, + max_side_length: Optional[int] = None, + img_scale_num: int = 16, + prompt_dropout_prob: float = 0.0, + ref_img_dropout_prob: float = 0.0, + ): + self.max_input_pixels = max_input_pixels + self.max_output_pixels = max_output_pixels + + self.max_side_length = max_side_length + self.img_scale_num = img_scale_num + self.prompt_dropout_prob = prompt_dropout_prob + self.ref_img_dropout_prob = ref_img_dropout_prob + + with open(config_path, "r") as f: + self.config = yaml.load(f, Loader=yaml.FullLoader) + + self.num_workers = num_workers + + self.use_chat_template = use_chat_template + self.image_processor = OmniGen2ImageProcessor(vae_scale_factor=img_scale_num, do_resize=True) + + data = self._collect_annotations(self.config) + + self.data = data + self.tokenizer = tokenizer + + def _collect_annotations(self, config): + total_samples = 0 + total_ratio = 0 + json_datasets = [] + + ratio_type = config.get('ratio_type', 'outside_ratio') + + for data in config['data']: + data_path, data_type = data['path'], data.get("type", "default") + if os.path.isdir(data_path): + jsonl_files = list(glob.glob(os.path.join(data_path, "**/*.jsonl"), recursive=True)) + list(glob.glob(os.path.join(data_path, "**/*.json"), recursive=True)) + json_dataset = load_dataset('json', data_files=jsonl_files, cache_dir=None, num_proc=self.num_workers)['train'] + logger.info(f"Loaded {len(json_dataset)} samples from {data_path}", main_process_only=False) + else: + data_ext = os.path.splitext(data_path)[-1] + if data_ext in [".json", ".jsonl"]: + json_dataset = load_dataset('json', data_files=data_path, cache_dir=None, num_proc=self.num_workers)['train'] + logger.info(f"Loaded {len(json_dataset)} samples from {data_path}", main_process_only=False) + elif data_ext in [".yml", ".yaml"]: + with open(data_path, "r") as f: + sub_config = yaml.load(f, Loader=yaml.FullLoader) + json_dataset = self._collect_annotations(sub_config) + else: + raise NotImplementedError( + f'Unknown data file extension: "{data_ext}". ' + f"Currently, .json, .jsonl .yml .yaml are supported. " + "If you are using a supported format, please set the file extension so that the proper parsing " + "routine can be called." + ) + total_ratio += data['ratio'] + total_samples += len(json_dataset) + json_datasets.append(json_dataset) + + start_time = time.time() + for data, json_dataset in zip(config['data'], json_datasets): + if ratio_type == 'inside_ratio': + ratio = data['ratio'] + else: + ratio = data['ratio'] / total_ratio + if ratio != 1: + target_size = int(len(json_dataset) * ratio) # normalize the ratio + if target_size <= len(json_dataset): + # Random selection without replacement + indices = random.sample(range(len(json_dataset)), target_size) + else: + # Oversample with replacement + indices = random.choices(range(len(json_dataset)), k=target_size) + json_dataset = json_dataset.select(indices) + logger.info(f"Time taken to select samples: {time.time() - start_time} seconds", main_process_only=False) + + start_time = time.time() + json_dataset = concatenate_datasets(json_datasets) + logger.info(f"Time taken to concatenate datasets: {time.time() - start_time} seconds") + return json_dataset + + def clean_data_item(self, data_item): + data_item['task_type'] = data_item['task_type'] if 'task_type' in data_item and data_item['task_type'] else "" + + task_type = data_item['task_type'] + instruction = data_item['instruction'] + + if 'instruction_long' in data_item and data_item['instruction_long'] is not None: + if random.random() < 0.8 and len(data_item['instruction_long']) > 10: + instruction = data_item['instruction_long'] + + if 'instruction_short' in data_item and data_item['instruction_short'] is not None: + if random.random() < 0.3 and len(data_item['instruction_short']) > 10: + instruction = data_item['instruction_short'] + + if any([t2i_task_type in task_type for t2i_task_type in ['text_to_image', 't2i']]): + if all(key in data_item and data_item[key] is not None for key in ['instruction', 'instruction_zh', 'instruction_short', 'instruction_short_zh']): + instruction = random.choice( + [ + data_item[key] + for key in [ + "instruction", + "instruction_zh", + "instruction_short", + "instruction_short_zh", + ] + if len(data_item[key]) > 10 + ] + ) + elif 'instruction_zh' in data_item: + if "instruction_tag" in data_item and data_item['instruction_tag'] is not None and len(data_item['instruction_tag']) > 3 and random.random() < 0.08: + instruction = data_item['instruction_tag'] + elif random.random() < 0.5 and data_item['instruction_zh'] is not None and len(data_item['instruction_zh']) > 10: + instruction = data_item['instruction_zh'] + else: + instruction = data_item['instruction'] + + if "caption_qwenvl2_5" in data_item and data_item['caption_qwenvl2_5'] is not None: + randn_num = random.random() + if len(data_item["caption_qwenvl2_5"]) < 1000 and randn_num < 0.4 and len(data_item['caption_qwenvl2_5']) > 10: + instruction = data_item['caption_qwenvl2_5'] + elif len(data_item['caption_qwenvl2_5_zh']) < 1000 and randn_num < 0.99 and len(data_item['caption_qwenvl2_5_zh']) > 10: + instruction = data_item['caption_qwenvl2_5_zh'] + else: + instruction = data_item['instruction'] + + if "Hyper-Realistic photo. Photo of " in instruction: + instruction = instruction.replace("Hyper-Realistic photo. Photo of ", "") + if "Hyper-Realistic photo. " in instruction: + instruction = instruction.replace("Hyper-Realistic photo. ", "") + + prefixs = ["The image portrays ", "The image depicts ", "The image captures ", "The image highlights ", "The image shows ", "่ฟ™ๅผ ๅ›พ็‰‡ๅฑ•็คบไบ†"] + if random.random() < 0.5: + for p in prefixs: + if p in data_item['instruction']: + data_item['instruction'] = data_item['instruction'].replace(p, "") + break + + if "Hyper-Realistic photo. Photo of " in data_item['instruction']: + data_item['instruction'] = data_item['instruction'].replace("Hyper-Realistic photo. Photo of ", "") + if "Hyper-Realistic photo. " in data_item['instruction']: + data_item['instruction'] = data_item['instruction'].replace("Hyper-Realistic photo. ", "") + + unnecessary_words = ["", "", "<|image_1|>", "<|image_2|>", "<|image_3|>", "<|image_4|>", "<|image_5|>", + "<|img_1|>", "<|img_2|>", "<|img_3|>", "<|img_4|>", "<|img_5|>"] + for word in unnecessary_words: + data_item['instruction'] = data_item['instruction'].replace(word, '') + + instruction = normalize_whitespace(instruction) + + return data_item + + def apply_chat_template(self, instruction, system_prompt): + if self.use_chat_template: + prompt = [ + { + "role": "system", + "content": system_prompt, + }, + {"role": "user", "content": instruction}, + ] + instruction = self.tokenizer.apply_chat_template(prompt, tokenize=False, add_generation_prompt=False) + return instruction + + def process_item(self, data_item): + assert data_item['instruction'] is not None + data_item = self.clean_data_item(data_item) + + drop_prompt = random.random() < self.prompt_dropout_prob + drop_ref_img = drop_prompt and random.random() < self.ref_img_dropout_prob + + if drop_prompt: + instruction = self.apply_chat_template("", self.SYSTEM_PROMPT_DROP) + else: + instruction = self.apply_chat_template(data_item['instruction'], self.SYSTEM_PROMPT) + + if not drop_ref_img and 'input_images' in data_item and data_item['input_images'] is not None: + input_images_path = data_item['input_images'] + input_images = [] + input_images_pil = [] + + max_input_pixels = self.max_input_pixels[len(input_images_path) - 1] if isinstance(self.max_input_pixels, list) else self.max_input_pixels + + for input_image_path in input_images_path: + input_image_pil = Image.open(input_image_path).convert("RGB") + input_image = self.image_processor.preprocess(input_image_pil, max_pixels=max_input_pixels, max_side_length=self.max_side_length) + input_image_pil_resized = self.image_processor.postprocess(input_image, output_type="pil")[0] + input_images.append(input_image) + input_images_pil.append(input_image_pil_resized) + else: + input_images_path, input_images, input_images_pil = None, None, None + + target_img_size = data_item.get('target_img_size', (512, 512)) + + if input_images_pil is not None and len(input_images_pil) == 1: + target_img_size = (input_images_pil[0].width, input_images_pil[0].height) + + w, h = target_img_size + cur_pixels = w * h + ratio = min(1, (self.max_output_pixels / cur_pixels) ** 0.5) + + target_img_size = (int(w * ratio) // self.img_scale_num * self.img_scale_num, int(h * ratio) // self.img_scale_num * self.img_scale_num) + + # output_image_path = data_item['output_image'] + # output_image = Image.open(output_image_path).convert("RGB") + # output_image = self.image_processor.preprocess(output_image, max_pixels=self.max_output_pixels, max_side_length=self.max_side_length) + + data = { + 'task_type': data_item['task_type'], + 'instruction': instruction, + 'input_images_path': input_images_path, + 'input_images': input_images, + 'input_images_pil': input_images_pil, + 'target_img_size': target_img_size, + 'meta_data': data_item['meta_data'], + # 'output_image': output_image, + # 'output_image_path': output_image_path, + } + return data + + def __getitem__(self, index): + max_retries = 100 + + current_index = index + for attempt in range(max_retries): + try: + data_item = self.data[current_index] + return self.process_item(data_item) + except Exception as e: + print("error when loading data: ", e) + if attempt == max_retries - 1: + raise e + else: + # Try a different index for the next attempt + current_index = random.randint(0, len(self.data) - 1) + continue + + def __len__(self): + return len(self.data) + +class OmniGen2Collator(): + def __init__(self, tokenizer, max_token_len): + self.tokenizer = tokenizer + self.max_token_len = max_token_len + + def __call__(self, batch): + task_type = [data['task_type'] for data in batch] + instruction = [data['instruction'] for data in batch] + input_images_path = [data['input_images_path'] for data in batch] + input_images = [data['input_images'] for data in batch] + input_images_pil = [data['input_images_pil'] for data in batch] + target_img_size = [data['target_img_size'] for data in batch] + meta_data = [data['meta_data'] for data in batch] + # output_image = [data['output_image'] for data in batch] + # output_image_path = [data['output_image_path'] for data in batch] + + text_inputs = self.tokenizer( + instruction, + padding="longest", + max_length=self.max_token_len, + truncation=True, + return_tensors="pt", + ) + + data = { + "task_type": task_type, + "instruction": instruction, + "text_ids": text_inputs.input_ids, + "text_mask": text_inputs.attention_mask, + "input_images": input_images, + "input_images_path": input_images_path, + "input_images_pil": input_images_pil, + "target_img_size": target_img_size, + "meta_data": meta_data, + # "target_img_size": target_img_size, + # "output_image": output_image, + # "output_image_path": output_image_path, + } + return data + + +class RepeatedDistributedBatchSampler(torch.utils.data.Sampler): + def __init__( + self, + dataset, + batch_size: int, + num_repeats: int, + num_replicas: int, + rank: int, + shuffle: bool = True, + seed: int = 0, + drop_last: bool = False, + ): + self.dataset = dataset + self.num_repeats = num_repeats + self.num_replicas = num_replicas + self.rank = rank + self.shuffle = shuffle + self.seed = seed + self.drop_last = drop_last + + self.samples_per_iter = self.num_replicas * batch_size + assert self.samples_per_iter % self.num_repeats == 0, f"k can not div n*b, k{num_repeats}-num_replicas{num_replicas}-batch_size{batch_size}" + self.unique_samples_per_iter = self.samples_per_iter // self.num_repeats + + if self.drop_last and len(self.dataset) % self.unique_samples_per_iter != 0: # type: ignore[arg-type] + self.num_batches = len(self.dataset) // self.unique_samples_per_iter # type: ignore[arg-type] + else: + self.num_batches = math.ceil(len(self.dataset) / self.unique_samples_per_iter) # type: ignore[arg-type] + + self.total_size = self.num_batches * self.unique_samples_per_iter + self.batch_size = self.unique_samples_per_iter * self.num_repeats + self.epoch=0 + + def __iter__(self): + g = torch.Generator() + g.manual_seed(self.seed + self.epoch) + if self.shuffle: + # deterministically shuffle based on epoch and seed + indices = torch.randperm(len(self.dataset), generator=g).tolist() # type: ignore[arg-type] + else: + indices = list(range(len(self.dataset))) # type: ignore[arg-type] + indices = indices[: self.total_size] + + for i in range(self.num_batches): + start = i * self.unique_samples_per_iter + end = start + self.unique_samples_per_iter + batch_indices = indices[start:end] + + batch_indices = batch_indices * self.num_repeats + + shuffled_indices = torch.randperm(len(batch_indices), generator=g).tolist() + shuffled_samples = [batch_indices[j] for j in shuffled_indices] + + yield shuffled_samples + + def __len__(self): + return self.num_batches + + def set_epoch(self, epoch): + self.epoch = epoch \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/grpo/__init__.py b/examples/OmniGen2-RL/omnigen2/grpo/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/examples/OmniGen2-RL/omnigen2/grpo/reward_client_edit.py b/examples/OmniGen2-RL/omnigen2/grpo/reward_client_edit.py new file mode 100644 index 0000000000000000000000000000000000000000..965fd9097fb46bdaa142a5190c77cec4a6f6f9b6 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/grpo/reward_client_edit.py @@ -0,0 +1,187 @@ +#!/usr/bin/env python3 +""" +Pure Reward Client - Only responsible for data transmission +""" + +import base64 +from io import BytesIO +import requests +import time +import logging +from typing import List, Dict, Any, Optional, Tuple +from PIL import Image + +logger = logging.getLogger(__name__) + +class RewardClient: + """ + Pure Reward Client - Only responsible for communicating with proxy server + """ + + def __init__(self, proxy_host: str = "127.0.0.1", proxy_port: int = 23456, + timeout: int = 300, max_retries: int = 3): + """ + Initialize client + + Args: + proxy_host: Proxy server host address + proxy_port: Proxy server port + timeout: Request timeout in seconds + max_retries: Maximum number of retries + """ + self.proxy_url = f"http://{proxy_host}:{proxy_port}" + self.timeout = timeout + self.max_retries = max_retries + + logger.info(f"Initialize Reward client: {self.proxy_url}") + + @staticmethod + def _encode_image(image_obj: Any) -> str: + """Encode PIL image/bytes to base64 PNG string for safe transport.""" + if isinstance(image_obj, str): + return image_obj + if isinstance(image_obj, (bytes, bytearray)): + return base64.b64encode(image_obj).decode("ascii") + if not isinstance(image_obj, Image.Image): + raise TypeError(f"Unsupported image type: {type(image_obj)}") + + buffer = BytesIO() + image_obj.save(buffer, format="PNG") + return base64.b64encode(buffer.getvalue()).decode("ascii") + + def _build_safe_request_payload( + self, + input_images: List[Any], + output_image: List[Any], + meta_datas: List[Dict[str, Any]], + server_type: str, + ) -> Dict[str, Any]: + encoded_input_images = [] + for sample in input_images: + if not isinstance(sample, list): + raise TypeError(f"Each input_images item must be a list, got: {type(sample)}") + encoded_input_images.append([self._encode_image(img) for img in sample]) + + encoded_output_images = [self._encode_image(img) for img in output_image] + return { + "input_images": encoded_input_images, + "output_image": encoded_output_images, + "meta_datas": meta_datas, + "server_type": server_type, + } + + def evaluate(self, input_images: List[bytes], output_image: List[bytes], meta_datas: List[Dict[str, Any]], + server_type: str = 'geneval') -> Optional[Tuple[List[float], List[float], List[str], List[Dict]]]: + """ + Evaluate images and return rewards + + Args: + input_images: List of input image byte data + output_image: List of output image byte data + meta_datas: List of metadata + server_type: Server type ('geneval', 'ocr', etc.) + + Returns: + tuple: (scores, rewards, reasoning, meta_data) + - scores: List of scores + - rewards: List of rewards + - reasoning: List of reasoning results + - meta_data: List of metadata + """ + if not output_image: + return [], [], [], [] + + # Prepare request data + request_data = self._build_safe_request_payload( + input_images=input_images, + output_image=output_image, + meta_datas=meta_datas, + server_type=server_type, + ) + + # Retry logic + last_exception = None + for attempt in range(self.max_retries): + try: + response = requests.post( + self.proxy_url, + json=request_data, + timeout=self.timeout + ) + + if response.status_code == 200: + # Parse results + result = response.json() + scores = result.get('scores', []) + rewards = result.get('rewards', []) + reasoning = result.get('reasoning', []) + meta_data = result.get('meta_data', []) + + # Basic validation + if len(scores) != len(output_image) or len(rewards) != len(output_image): + logger.warning(f"Return data length mismatch: expected {len(output_image)}, got scores={len(scores)}, rewards={len(rewards)}") + + return scores, rewards, reasoning, meta_data + else: + logger.error(f"HTTP error: {response.status_code}") + last_exception = RuntimeError(f"HTTP {response.status_code}") + + except requests.exceptions.Timeout as e: + logger.error(f"Request timeout (attempt {attempt + 1}/{self.max_retries})") + last_exception = e + + except Exception as e: + logger.error(f"Request exception: {e} (attempt {attempt + 1}/{self.max_retries})") + last_exception = e + + # Wait before retry + if attempt < self.max_retries - 1: + time.sleep(2 ** attempt) + + logger.error(f"All retries failed, last exception: {last_exception}") + return None + + def ping(self) -> bool: + """Check if server is reachable""" + try: + response = requests.get(f"{self.proxy_url}/ping", timeout=5) + return response.status_code == 200 + except: + return False + +# Convenience function +def evaluate_images(input_images: List[bytes], output_image: List[bytes], meta_datas: List[Dict[str, Any]], + proxy_host: str = "127.0.0.1", proxy_port: int = 23456, + server_type: str = 'vlm') -> Optional[Tuple[List[float], List[float], List[str], List[Dict]]]: + """ + Convenience function: directly evaluate images + """ + client = RewardClient(proxy_host, proxy_port, timeout=600, max_retries=1) + return client.evaluate(input_images, output_image, meta_datas, server_type) + +# Usage example +if __name__ == "__main__": + # Create client + client = RewardClient() + + # Check connection + if not client.ping(): + print("โŒ Server unreachable") + exit(1) + + # Mock data + input_images = [b"fake_input_image"] # Input images + output_images = [b"fake_output_image"] # Output images + meta_datas = [{"tag": "test", "prompt": "a simple test"}] # Metadata + + # Evaluate images + print("๐Ÿ”ฅ Evaluating images:") + result = client.evaluate(input_images, output_images, meta_datas, server_type='vlm') + if result: + scores, rewards, reasoning, meta_data = result + print(f"Scores: {scores}") + print(f"Rewards: {rewards}") + print(f"Reasoning: {reasoning}") + print(f"Meta data: {meta_data}") + else: + print("Evaluation failed") diff --git a/examples/OmniGen2-RL/omnigen2/grpo/utils.py b/examples/OmniGen2-RL/omnigen2/grpo/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..e5fdaa1a05bbd0956a0445f3f7073f0efb1dc83e --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/grpo/utils.py @@ -0,0 +1,206 @@ +from typing import List, Dict, Any, Tuple +from collections import defaultdict + +import math +import numpy as np +import torch + +from accelerate.utils import gather_object + +def expand_as(tensor, other): + """ + Expands a tensor to match the dimensions of another tensor. + + If tensor has shape [b] and other has shape [b, c, h, w], + this function will reshape tensor to [b, 1, 1, 1] to enable broadcasting. + + Args: + tensor (`torch.FloatTensor`): The tensor to expand + other (`torch.FloatTensor`): The tensor whose shape will be matched + + Returns: + `torch.FloatTensor`: The expanded tensor + """ + for _ in range(other.ndim - tensor.ndim): + tensor = tensor.unsqueeze(-1) + return tensor + + +def process_grpo_rewards( + rewards: torch.Tensor, + prompts: List[str], + accelerator, + std_level: str = 'group' +) -> Tuple[torch.Tensor, Dict]: # Modified return type + """ + Process GRPO rewards and compute advantages + Only use group statistics from the same prompt in the current batch + + Args: + rewards: rewards of current batch [batch_size] + prompts: prompts of current batch [batch_size] + accelerator: distributed training accelerator + std_level: 'group' or 'batch' + + Returns: + advantages: computed advantages [batch_size] + prompt_stats: statistical information for each prompt + """ + # 1. Gather rewards, text_ids and mask from all processes + gathered_rewards = accelerator.gather(rewards) # [world_size * batch_size] + gathered_prompts = gather_object(prompts) # [world_size * batch_size] + + assert len(gathered_rewards) == len(gathered_prompts), f"{len(gathered_rewards)=} {len(gathered_prompts)=}" + + # 3. Group rewards by prompt + prompt_to_rewards = defaultdict(list) + for prompt, reward in zip(gathered_prompts, gathered_rewards): + prompt_to_rewards[prompt].append(reward.item()) + + # 4. Pre-compute statistical information for each prompt group + prompt_stats = {} + for prompt, group_rewards in prompt_to_rewards.items(): + prompt_stats[prompt] = { + "min": np.min(group_rewards), + "max": np.max(group_rewards), + "mean": np.mean(group_rewards), + "std": np.std(group_rewards), + } + + if std_level == 'batch': + batch_std = np.std([reward.item() for reward in gathered_rewards]) + + assert set(prompts).issubset(set(prompt_stats.keys())), f"{set(prompts)=} {set(prompt_stats.keys())=}" + + advantages = torch.zeros_like(rewards, device=rewards.device) + for i, prompt in enumerate(prompts): + stats = prompt_stats[prompt] + if std_level == 'group': + advantage = (rewards[i].item() - stats['mean']) / (stats['std'] + 1e-8) + elif std_level == 'batch': + advantage = (rewards[i].item() - stats['mean']) / (batch_std + 1e-8) + advantages[i] = torch.tensor(advantage, device=rewards.device) + + return advantages, prompt_stats + + +def forward_logprob( + latents: List[torch.Tensor], + latents_next: List[torch.Tensor], + t: torch.Tensor, + t_next: torch.Tensor, + step_index: int, + img_mask, + model, + model_kwargs: Dict[str, Any], + model_pred_kwargs: Dict[str, Any], + # model_pred_ref_kwargs: Dict[str, Any], + # model_pred_uncond_kwargs: Dict[str, Any], + scheduler, + apply_cfg: bool = True, + text_guidance_scale: float = 1.0, + image_guidance_scale: float = 1.0, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + # cfg + + model_pred = model(**model_kwargs, **model_pred_kwargs) + if apply_cfg: + if text_guidance_scale > 1.0 and image_guidance_scale > 1.0: + model_pred, model_pred_ref, model_pred_uncond = model_pred.chunk(3) + model_pred = ( + model_pred_uncond + + image_guidance_scale * (model_pred_ref - model_pred_uncond) + + text_guidance_scale * (model_pred - model_pred_ref) + ) + elif text_guidance_scale > 1.0: + model_pred, model_pred_uncond = model_pred.chunk(2) + model_pred = model_pred_uncond + text_guidance_scale * (model_pred - model_pred_uncond) + + + sigma_t = scheduler.get_sigma_t(t, t_next if step_index == 0 else None) # [batch_size] + sigma_t = expand_as(sigma_t.unsqueeze(1), latents) # [batch_size, max_img_len, dim] + t = expand_as(t.unsqueeze(1), latents) # [batch_size, max_img_len, dim] + t_next = expand_as(t_next.unsqueeze(1), latents) # [batch_size, max_img_len, dim] + dt = t_next - t + + sigma_t = sigma_t.to(dtype=torch.float32) + t = t.to(dtype=torch.float32) + t_next = t_next.to(dtype=torch.float32) + dt = dt.to(dtype=torch.float32) + + prev_sample_mean = ( + latents.to(dtype=torch.float32) * (1 - sigma_t**2 / (2 * (1 - t)) * dt) + + model_pred * (1 + sigma_t**2 * t / (2 * (1 - t))) * dt + ) + + log_prob = ( + -((latents_next.to(dtype=torch.float32).detach() - prev_sample_mean) ** 2) + / (2 * (sigma_t**2 * dt)) # Fix: denominator is 2 * ฯƒยฒ * dt + - torch.log(sigma_t * torch.sqrt(dt)) # Fix: log(ฯƒ * โˆšdt) + - 0.5 + * torch.log( + 2 * torch.as_tensor(math.pi, device=latents.device) + ) # Fix: 0.5 coefficient + ) + + img_mask = expand_as(img_mask, latents).expand(latents.shape) + log_prob = (log_prob * img_mask.detach()).sum( + dim=tuple(range(-log_prob.ndim + 1, 0)), dtype=torch.float32 + ) / img_mask.detach().sum( + dim=tuple(range(-log_prob.ndim + 1, 0)), dtype=torch.float32 + ) + + return log_prob, prev_sample_mean, sigma_t**2 + + +def compute_single_step_ppo_loss( + step_log_probs: torch.Tensor, # [batch_size] log probability of current time step + old_step_log_probs: torch.Tensor, # [batch_size] log probability of old policy + advantages: torch.Tensor, # [batch_size] advantages + clip_range: Tuple[float, float] = (1e-4, 1e-4), # PPO clipping range + adv_clip_max: float = 5 +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """ + Compute PPO loss for a single time step + + Args: + step_log_probs: log probability of current policy at this time step [batch_size] + old_step_log_probs: log probability of old policy at this time step [batch_size] + advantages: advantages [batch_size] + clip_range: PPO clipping range + + Returns: + step_pg_loss: policy loss at this time step (scalar) + pg_clipfrac_step: fraction of clipped samples (scalar) + approx_kl_step: approximate KL divergence (scalar) + """ + # Calculate ratio + log_ratio = step_log_probs - old_step_log_probs.detach() + approx_kl_step = (-log_ratio).mean() + + # Calculate two types of loss: original and clipped + ratio = torch.exp(log_ratio) + advantages = torch.clamp(advantages, -adv_clip_max, adv_clip_max) + unclipped_loss = -advantages.detach() * ratio # [batch_size] + clipped_loss = -advantages.detach() * torch.clamp( + ratio, 1.0 - clip_range[0], 1.0 + clip_range[1] + ) # [batch_size] + + # Take maximum value (more conservative loss) + step_pg_loss = torch.max(unclipped_loss, clipped_loss).mean() + + + # Calculate the fraction of clipped samples + pg_clipfrac_step = (clipped_loss > unclipped_loss).float().mean() + + num_positive = torch.where(advantages > 0, torch.ones_like(ratio), torch.zeros_like(ratio)).sum() + num_negative = torch.where(advantages < 0, torch.ones_like(ratio), torch.zeros_like(ratio)).sum() + ratio_positive = torch.where(advantages > 0, ratio, torch.zeros_like(ratio)).sum() / num_positive if num_positive > 0 else torch.tensor(0.0, device=ratio.device) + ratio_negative = torch.where(advantages < 0, ratio, torch.zeros_like(ratio)).sum() / num_negative if num_negative > 0 else torch.tensor(0.0, device=ratio.device) + + ratio_large_than_1 = torch.where(ratio > 1, torch.ones_like(ratio), torch.zeros_like(ratio)).sum() / len(ratio) + ratio_small_than_1 = torch.where(ratio < 1, torch.ones_like(ratio), torch.zeros_like(ratio)).sum() / len(ratio) + + return step_pg_loss, pg_clipfrac_step, approx_kl_step, unclipped_loss, clipped_loss, ratio, ratio_positive, ratio_negative, num_positive, num_negative, ratio_large_than_1, ratio_small_than_1 + + diff --git a/examples/OmniGen2-RL/omnigen2/models/__init__.py b/examples/OmniGen2-RL/omnigen2/models/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/examples/OmniGen2-RL/omnigen2/models/attention_processor.py b/examples/OmniGen2-RL/omnigen2/models/attention_processor.py new file mode 100644 index 0000000000000000000000000000000000000000..1f713c75f1f8a164440f45974f9540393a483a79 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/models/attention_processor.py @@ -0,0 +1,357 @@ +""" +OmniGen2 Attention Processor Module + +Copyright 2025 BAAI, The OmniGen2 Team and The HuggingFace Team. All rights reserved. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +import warnings +import math +from typing import Optional, Tuple, Dict, Any + +import torch +import torch.nn.functional as F +from einops import repeat + +from ..utils.import_utils import is_flash_attn_available + +if is_flash_attn_available(): + from flash_attn import flash_attn_varlen_func + from flash_attn.bert_padding import index_first_axis, pad_input, unpad_input +else: + warnings.warn("Cannot import flash_attn, install flash_attn to use Flash2Varlen attention for better performance") + + +from diffusers.models.attention_processor import Attention +from .embeddings import apply_rotary_emb + + +class OmniGen2AttnProcessorFlash2Varlen: + """ + Processor for implementing scaled dot-product attention with flash attention and variable length sequences. + + This processor implements: + - Flash attention with variable length sequences + - Rotary position embeddings (RoPE) + - Query-Key normalization + - Proportional attention scaling + + Args: + None + """ + + def __init__(self) -> None: + """Initialize the attention processor.""" + if not is_flash_attn_available(): + raise ImportError( + "OmniGen2AttnProcessorFlash2Varlen requires flash_attn. " + "Please install flash_attn." + ) + + def _upad_input( + self, + query_layer: torch.Tensor, + key_layer: torch.Tensor, + value_layer: torch.Tensor, + attention_mask: torch.Tensor, + query_length: int, + num_heads: int, + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, Tuple[torch.Tensor, torch.Tensor], Tuple[int, int]]: + """ + Unpad the input tensors for flash attention. + + Args: + query_layer: Query tensor of shape (batch_size, seq_len, num_heads, head_dim) + key_layer: Key tensor of shape (batch_size, seq_len, num_kv_heads, head_dim) + value_layer: Value tensor of shape (batch_size, seq_len, num_kv_heads, head_dim) + attention_mask: Attention mask tensor of shape (batch_size, seq_len) + query_length: Length of the query sequence + num_heads: Number of attention heads + + Returns: + Tuple containing: + - Unpadded query tensor + - Unpadded key tensor + - Unpadded value tensor + - Query indices + - Tuple of cumulative sequence lengths for query and key + - Tuple of maximum sequence lengths for query and key + """ + def _get_unpad_data(attention_mask: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, int]: + """Helper function to get unpadding data from attention mask.""" + seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32) + indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten() + max_seqlen_in_batch = seqlens_in_batch.max().item() + cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0)) + return indices, cu_seqlens, max_seqlen_in_batch + + indices_k, cu_seqlens_k, max_seqlen_in_batch_k = _get_unpad_data(attention_mask) + batch_size, kv_seq_len, num_key_value_heads, head_dim = key_layer.shape + + # Unpad key and value layers + key_layer = index_first_axis( + key_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), + indices_k, + ) + value_layer = index_first_axis( + value_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), + indices_k, + ) + + # Handle different query length cases + if query_length == kv_seq_len: + query_layer = index_first_axis( + query_layer.reshape(batch_size * kv_seq_len, num_heads, head_dim), + indices_k, + ) + cu_seqlens_q = cu_seqlens_k + max_seqlen_in_batch_q = max_seqlen_in_batch_k + indices_q = indices_k + elif query_length == 1: + max_seqlen_in_batch_q = 1 + cu_seqlens_q = torch.arange( + batch_size + 1, dtype=torch.int32, device=query_layer.device + ) + indices_q = cu_seqlens_q[:-1] + query_layer = query_layer.squeeze(1) + else: + attention_mask = attention_mask[:, -query_length:] + query_layer, indices_q, cu_seqlens_q, max_seqlen_in_batch_q = unpad_input(query_layer, attention_mask) + + return ( + query_layer, + key_layer, + value_layer, + indices_q, + (cu_seqlens_q, cu_seqlens_k), + (max_seqlen_in_batch_q, max_seqlen_in_batch_k), + ) + + def __call__( + self, + attn: Attention, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + attention_mask: Optional[torch.Tensor] = None, + image_rotary_emb: Optional[torch.Tensor] = None, + base_sequence_length: Optional[int] = None, + ) -> torch.Tensor: + """ + Process attention computation with flash attention. + + Args: + attn: Attention module + hidden_states: Hidden states tensor of shape (batch_size, seq_len, hidden_dim) + encoder_hidden_states: Encoder hidden states tensor + attention_mask: Optional attention mask tensor + image_rotary_emb: Optional rotary embeddings for image tokens + base_sequence_length: Optional base sequence length for proportional attention + + Returns: + torch.Tensor: Processed hidden states after attention computation + """ + batch_size, sequence_length, _ = hidden_states.shape + + # Get Query-Key-Value Pair + query = attn.to_q(hidden_states) + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + query_dim = query.shape[-1] + inner_dim = key.shape[-1] + head_dim = query_dim // attn.heads + dtype = query.dtype + + # Get key-value heads + kv_heads = inner_dim // head_dim + + # Reshape tensors for attention computation + query = query.view(batch_size, -1, attn.heads, head_dim) + key = key.view(batch_size, -1, kv_heads, head_dim) + value = value.view(batch_size, -1, kv_heads, head_dim) + + # Apply Query-Key normalization + if attn.norm_q is not None: + query = attn.norm_q(query) + if attn.norm_k is not None: + key = attn.norm_k(key) + + # Apply Rotary Position Embeddings + if image_rotary_emb is not None: + query = apply_rotary_emb(query, image_rotary_emb, use_real=False) + key = apply_rotary_emb(key, image_rotary_emb, use_real=False) + + query, key = query.to(dtype), key.to(dtype) + + # Calculate attention scale + if base_sequence_length is not None: + softmax_scale = math.sqrt(math.log(sequence_length, base_sequence_length)) * attn.scale + else: + softmax_scale = attn.scale + + # Unpad input for flash attention + ( + query_states, + key_states, + value_states, + indices_q, + cu_seq_lens, + max_seq_lens, + ) = self._upad_input(query, key, value, attention_mask, sequence_length, attn.heads) + + cu_seqlens_q, cu_seqlens_k = cu_seq_lens + max_seqlen_in_batch_q, max_seqlen_in_batch_k = max_seq_lens + + # Handle different number of heads + if kv_heads < attn.heads: + key_states = repeat(key_states, "l h c -> l (h k) c", k=attn.heads // kv_heads) + value_states = repeat(value_states, "l h c -> l (h k) c", k=attn.heads // kv_heads) + + # Apply flash attention + attn_output_unpad = flash_attn_varlen_func( + query_states, + key_states, + value_states, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_k=cu_seqlens_k, + max_seqlen_q=max_seqlen_in_batch_q, + max_seqlen_k=max_seqlen_in_batch_k, + dropout_p=0.0, + causal=False, + softmax_scale=softmax_scale, + ) + + # Pad output and apply final transformations + hidden_states = pad_input(attn_output_unpad, indices_q, batch_size, sequence_length) + hidden_states = hidden_states.flatten(-2) + hidden_states = hidden_states.type_as(query) + + # Apply output projection + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + + return hidden_states + + +class OmniGen2AttnProcessor: + """ + Processor for implementing scaled dot-product attention with flash attention and variable length sequences. + + This processor is optimized for PyTorch 2.0 and implements: + - Flash attention with variable length sequences + - Rotary position embeddings (RoPE) + - Query-Key normalization + - Proportional attention scaling + + Args: + None + + Raises: + ImportError: If PyTorch version is less than 2.0 + """ + + def __init__(self) -> None: + """Initialize the attention processor.""" + if not hasattr(F, "scaled_dot_product_attention"): + raise ImportError( + "OmniGen2AttnProcessorFlash2Varlen requires PyTorch 2.0. " + "Please upgrade PyTorch to version 2.0 or later." + ) + + def __call__( + self, + attn: Attention, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + attention_mask: Optional[torch.Tensor] = None, + image_rotary_emb: Optional[torch.Tensor] = None, + base_sequence_length: Optional[int] = None, + ) -> torch.Tensor: + """ + Process attention computation with flash attention. + + Args: + attn: Attention module + hidden_states: Hidden states tensor of shape (batch_size, seq_len, hidden_dim) + encoder_hidden_states: Encoder hidden states tensor + attention_mask: Optional attention mask tensor + image_rotary_emb: Optional rotary embeddings for image tokens + base_sequence_length: Optional base sequence length for proportional attention + + Returns: + torch.Tensor: Processed hidden states after attention computation + """ + batch_size, sequence_length, _ = hidden_states.shape + + # Get Query-Key-Value Pair + query = attn.to_q(hidden_states) + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + query_dim = query.shape[-1] + inner_dim = key.shape[-1] + head_dim = query_dim // attn.heads + dtype = query.dtype + + # Get key-value heads + kv_heads = inner_dim // head_dim + + # Reshape tensors for attention computation + query = query.view(batch_size, -1, attn.heads, head_dim) + key = key.view(batch_size, -1, kv_heads, head_dim) + value = value.view(batch_size, -1, kv_heads, head_dim) + + # Apply Query-Key normalization + if attn.norm_q is not None: + query = attn.norm_q(query) + if attn.norm_k is not None: + key = attn.norm_k(key) + + # Apply Rotary Position Embeddings + if image_rotary_emb is not None: + query = apply_rotary_emb(query, image_rotary_emb, use_real=False) + key = apply_rotary_emb(key, image_rotary_emb, use_real=False) + + query, key = query.to(dtype), key.to(dtype) + + # Calculate attention scale + if base_sequence_length is not None: + softmax_scale = math.sqrt(math.log(sequence_length, base_sequence_length)) * attn.scale + else: + softmax_scale = attn.scale + + # scaled_dot_product_attention expects attention_mask shape to be + # (batch, heads, source_length, target_length) + if attention_mask is not None: + attention_mask = attention_mask.bool().view(batch_size, 1, 1, -1) + + query = query.transpose(1, 2) + key = key.transpose(1, 2) + value = value.transpose(1, 2) + + # explicitly repeat key and value to match query length, otherwise using enable_gqa=True results in MATH backend of sdpa in our test of pytorch2.6 + key = key.repeat_interleave(query.size(-3) // key.size(-3), -3) + value = value.repeat_interleave(query.size(-3) // value.size(-3), -3) + + hidden_states = F.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, scale=softmax_scale + ) + hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + hidden_states = hidden_states.type_as(query) + + # Apply output projection + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + + return hidden_states \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/models/embeddings.py b/examples/OmniGen2-RL/omnigen2/models/embeddings.py new file mode 100644 index 0000000000000000000000000000000000000000..5282f2defef551b70276a24ae16995cfe515679c --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/models/embeddings.py @@ -0,0 +1,126 @@ +# Copyright 2024 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +from typing import List, Optional, Tuple, Union + +import torch +from torch import nn + + +from diffusers.models.activations import get_activation + + +class TimestepEmbedding(nn.Module): + def __init__( + self, + in_channels: int, + time_embed_dim: int, + act_fn: str = "silu", + out_dim: int = None, + post_act_fn: Optional[str] = None, + cond_proj_dim=None, + sample_proj_bias=True, + ): + super().__init__() + + self.linear_1 = nn.Linear(in_channels, time_embed_dim, sample_proj_bias) + + if cond_proj_dim is not None: + self.cond_proj = nn.Linear(cond_proj_dim, in_channels, bias=False) + else: + self.cond_proj = None + + self.act = get_activation(act_fn) + + if out_dim is not None: + time_embed_dim_out = out_dim + else: + time_embed_dim_out = time_embed_dim + self.linear_2 = nn.Linear(time_embed_dim, time_embed_dim_out, sample_proj_bias) + + if post_act_fn is None: + self.post_act = None + else: + self.post_act = get_activation(post_act_fn) + + self.initialize_weights() + + def initialize_weights(self): + nn.init.normal_(self.linear_1.weight, std=0.02) + nn.init.zeros_(self.linear_1.bias) + nn.init.normal_(self.linear_2.weight, std=0.02) + nn.init.zeros_(self.linear_2.bias) + + def forward(self, sample, condition=None): + if condition is not None: + sample = sample + self.cond_proj(condition) + sample = self.linear_1(sample) + + if self.act is not None: + sample = self.act(sample) + + sample = self.linear_2(sample) + + if self.post_act is not None: + sample = self.post_act(sample) + return sample + + +def apply_rotary_emb( + x: torch.Tensor, + freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]], + use_real: bool = True, + use_real_unbind_dim: int = -1, +) -> Tuple[torch.Tensor, torch.Tensor]: + """ + Apply rotary embeddings to input tensors using the given frequency tensor. This function applies rotary embeddings + to the given query or key 'x' tensors using the provided frequency tensor 'freqs_cis'. The input tensors are + reshaped as complex numbers, and the frequency tensor is reshaped for broadcasting compatibility. The resulting + tensors contain rotary embeddings and are returned as real tensors. + + Args: + x (`torch.Tensor`): + Query or key tensor to apply rotary embeddings. [B, H, S, D] xk (torch.Tensor): Key tensor to apply + freqs_cis (`Tuple[torch.Tensor]`): Precomputed frequency tensor for complex exponentials. ([S, D], [S, D],) + + Returns: + Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings. + """ + if use_real: + cos, sin = freqs_cis # [S, D] + cos = cos[None, None] + sin = sin[None, None] + cos, sin = cos.to(x.device), sin.to(x.device) + + if use_real_unbind_dim == -1: + # Used for flux, cogvideox, hunyuan-dit + x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2] + x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3) + elif use_real_unbind_dim == -2: + # Used for Stable Audio, OmniGen and CogView4 + x_real, x_imag = x.reshape(*x.shape[:-1], 2, -1).unbind(-2) # [B, S, H, D//2] + x_rotated = torch.cat([-x_imag, x_real], dim=-1) + else: + raise ValueError(f"`use_real_unbind_dim={use_real_unbind_dim}` but should be -1 or -2.") + + out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype) + + return out + else: + # used for lumina + # x_rotated = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2)) + x_rotated = torch.view_as_complex(x.float().reshape(*x.shape[:-1], x.shape[-1] // 2, 2)) + freqs_cis = freqs_cis.unsqueeze(2) + x_out = torch.view_as_real(x_rotated * freqs_cis).flatten(3) + + return x_out.type_as(x) \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/models/transformers/__init__.py b/examples/OmniGen2-RL/omnigen2/models/transformers/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..157de1d42af9ddb8162077aeae0ef52acdd792f8 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/models/transformers/__init__.py @@ -0,0 +1,3 @@ +from .transformer_omnigen2 import OmniGen2Transformer2DModel + +__all__ = ["OmniGen2Transformer2DModel"] \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/models/transformers/block_lumina2.py b/examples/OmniGen2-RL/omnigen2/models/transformers/block_lumina2.py new file mode 100644 index 0000000000000000000000000000000000000000..13739d3a596d196411b85af19a55b34eb566b0cf --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/models/transformers/block_lumina2.py @@ -0,0 +1,218 @@ + +# Copyright 2024 Alpha-VLLM Authors and The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import warnings +from typing import Optional, Tuple + +import torch +import torch.nn as nn + +from diffusers.models.embeddings import Timesteps +from ..embeddings import TimestepEmbedding + +from ...utils.import_utils import is_flash_attn_available, is_triton_available + +if is_triton_available(): + from ...ops.triton.layer_norm import RMSNorm +else: + from torch.nn import RMSNorm + warnings.warn("Cannot import triton, install triton to use fused RMSNorm for better performance") + +if is_flash_attn_available(): + from flash_attn.ops.activations import swiglu +else: + from .components import swiglu + warnings.warn("Cannot import flash_attn, install flash_attn to use fused SwiGLU for better performance") + +# try: +# from flash_attn.ops.activations import swiglu as fused_swiglu +# FUSEDSWIGLU_AVALIBLE = True +# except ImportError: + +# FUSEDSWIGLU_AVALIBLE = False +# warnings.warn("Cannot import apex RMSNorm, switch to vanilla implementation") + +class LuminaRMSNormZero(nn.Module): + """ + Norm layer adaptive RMS normalization zero. + + Parameters: + embedding_dim (`int`): The size of each embedding vector. + """ + + def __init__( + self, + embedding_dim: int, + norm_eps: float, + norm_elementwise_affine: bool, + ): + super().__init__() + self.silu = nn.SiLU() + self.linear = nn.Linear( + min(embedding_dim, 1024), + 4 * embedding_dim, + bias=True, + ) + + self.norm = RMSNorm(embedding_dim, eps=norm_eps) + + def forward( + self, + x: torch.Tensor, + emb: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + emb = self.linear(self.silu(emb)) + scale_msa, gate_msa, scale_mlp, gate_mlp = emb.chunk(4, dim=1) + x = self.norm(x) * (1 + scale_msa[:, None]) + return x, gate_msa, scale_mlp, gate_mlp + + +class LuminaLayerNormContinuous(nn.Module): + def __init__( + self, + embedding_dim: int, + conditioning_embedding_dim: int, + # NOTE: It is a bit weird that the norm layer can be configured to have scale and shift parameters + # because the output is immediately scaled and shifted by the projected conditioning embeddings. + # Note that AdaLayerNorm does not let the norm layer have scale and shift parameters. + # However, this is how it was implemented in the original code, and it's rather likely you should + # set `elementwise_affine` to False. + elementwise_affine=True, + eps=1e-5, + bias=True, + norm_type="layer_norm", + out_dim: Optional[int] = None, + ): + super().__init__() + + # AdaLN + self.silu = nn.SiLU() + self.linear_1 = nn.Linear(conditioning_embedding_dim, embedding_dim, bias=bias) + + if norm_type == "layer_norm": + self.norm = nn.LayerNorm(embedding_dim, eps, elementwise_affine, bias) + elif norm_type == "rms_norm": + self.norm = RMSNorm(embedding_dim, eps=eps, elementwise_affine=elementwise_affine) + else: + raise ValueError(f"unknown norm_type {norm_type}") + + self.linear_2 = None + if out_dim is not None: + self.linear_2 = nn.Linear(embedding_dim, out_dim, bias=bias) + + def forward( + self, + x: torch.Tensor, + conditioning_embedding: torch.Tensor, + ) -> torch.Tensor: + # convert back to the original dtype in case `conditioning_embedding`` is upcasted to float32 (needed for hunyuanDiT) + emb = self.linear_1(self.silu(conditioning_embedding).to(x.dtype)) + scale = emb + x = self.norm(x) * (1 + scale)[:, None, :] + + if self.linear_2 is not None: + x = self.linear_2(x) + + return x + + +class LuminaFeedForward(nn.Module): + r""" + A feed-forward layer. + + Parameters: + hidden_size (`int`): + The dimensionality of the hidden layers in the model. This parameter determines the width of the model's + hidden representations. + intermediate_size (`int`): The intermediate dimension of the feedforward layer. + multiple_of (`int`, *optional*): Value to ensure hidden dimension is a multiple + of this value. + ffn_dim_multiplier (float, *optional*): Custom multiplier for hidden + dimension. Defaults to None. + """ + + def __init__( + self, + dim: int, + inner_dim: int, + multiple_of: Optional[int] = 256, + ffn_dim_multiplier: Optional[float] = None, + ): + super().__init__() + self.swiglu = swiglu + + # custom hidden_size factor multiplier + if ffn_dim_multiplier is not None: + inner_dim = int(ffn_dim_multiplier * inner_dim) + inner_dim = multiple_of * ((inner_dim + multiple_of - 1) // multiple_of) + + self.linear_1 = nn.Linear( + dim, + inner_dim, + bias=False, + ) + self.linear_2 = nn.Linear( + inner_dim, + dim, + bias=False, + ) + self.linear_3 = nn.Linear( + dim, + inner_dim, + bias=False, + ) + + def forward(self, x): + h1, h2 = self.linear_1(x), self.linear_3(x) + return self.linear_2(self.swiglu(h1, h2)) + + +class Lumina2CombinedTimestepCaptionEmbedding(nn.Module): + def __init__( + self, + hidden_size: int = 4096, + text_feat_dim: int = 2048, + frequency_embedding_size: int = 256, + norm_eps: float = 1e-5, + timestep_scale: float = 1.0, + ) -> None: + super().__init__() + + self.time_proj = Timesteps( + num_channels=frequency_embedding_size, flip_sin_to_cos=True, downscale_freq_shift=0.0, scale=timestep_scale + ) + + self.timestep_embedder = TimestepEmbedding( + in_channels=frequency_embedding_size, time_embed_dim=min(hidden_size, 1024) + ) + + self.caption_embedder = nn.Sequential( + RMSNorm(text_feat_dim, eps=norm_eps), + nn.Linear(text_feat_dim, hidden_size, bias=True), + ) + + self._initialize_weights() + + def _initialize_weights(self): + nn.init.trunc_normal_(self.caption_embedder[1].weight, std=0.02) + nn.init.zeros_(self.caption_embedder[1].bias) + + def forward( + self, timestep: torch.Tensor, text_hidden_states: torch.Tensor, dtype: torch.dtype + ) -> Tuple[torch.Tensor, torch.Tensor]: + timestep_proj = self.time_proj(timestep).to(dtype=dtype) + time_embed = self.timestep_embedder(timestep_proj) + caption_embed = self.caption_embedder(text_hidden_states) + return time_embed, caption_embed \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/models/transformers/components.py b/examples/OmniGen2-RL/omnigen2/models/transformers/components.py new file mode 100644 index 0000000000000000000000000000000000000000..5e654b8c4d7228609817cea7c25036728d6f588d --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/models/transformers/components.py @@ -0,0 +1,4 @@ +import torch.nn.functional as F + +def swiglu(x, y): + return F.silu(x.float(), inplace=False).to(x.dtype) * y \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/models/transformers/repo.py b/examples/OmniGen2-RL/omnigen2/models/transformers/repo.py new file mode 100644 index 0000000000000000000000000000000000000000..08f3c784abc7423364a5c5ee65ad3731be78b04e --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/models/transformers/repo.py @@ -0,0 +1,129 @@ +from typing import List, Tuple + +import torch +import torch.nn as nn + +from einops import repeat +from diffusers.models.embeddings import get_1d_rotary_pos_embed + +class OmniGen2RotaryPosEmbed(nn.Module): + def __init__(self, theta: int, + axes_dim: Tuple[int, int, int], + axes_lens: Tuple[int, int, int] = (300, 512, 512), + patch_size: int = 2): + super().__init__() + self.theta = theta + self.axes_dim = axes_dim + self.axes_lens = axes_lens + self.patch_size = patch_size + + @staticmethod + def get_freqs_cis(axes_dim: Tuple[int, int, int], + axes_lens: Tuple[int, int, int], + theta: int) -> List[torch.Tensor]: + freqs_cis = [] + freqs_dtype = torch.float32 if torch.backends.mps.is_available() else torch.float64 + for i, (d, e) in enumerate(zip(axes_dim, axes_lens)): + emb = get_1d_rotary_pos_embed(d, e, theta=theta, freqs_dtype=freqs_dtype) + freqs_cis.append(emb) + return freqs_cis + + def _get_freqs_cis(self, freqs_cis, ids: torch.Tensor) -> torch.Tensor: + device = ids.device + if ids.device.type == "mps": + ids = ids.to("cpu") + + result = [] + for i in range(len(self.axes_dim)): + freqs = freqs_cis[i].to(ids.device) + index = ids[:, :, i : i + 1].repeat(1, 1, freqs.shape[-1]).to(torch.int64) + result.append(torch.gather(freqs.unsqueeze(0).repeat(index.shape[0], 1, 1), dim=1, index=index)) + return torch.cat(result, dim=-1).to(device) + + def forward( + self, + freqs_cis, + attention_mask, + l_effective_ref_img_len, + l_effective_img_len, + ref_img_sizes, + img_sizes, + device + ): + batch_size = len(attention_mask) + p = self.patch_size + + encoder_seq_len = attention_mask.shape[1] + l_effective_cap_len = attention_mask.sum(dim=1).tolist() + + seq_lengths = [cap_len + sum(ref_img_len) + img_len for cap_len, ref_img_len, img_len in zip(l_effective_cap_len, l_effective_ref_img_len, l_effective_img_len)] + + max_seq_len = max(seq_lengths) + max_ref_img_len = max([sum(ref_img_len) for ref_img_len in l_effective_ref_img_len]) + max_img_len = max(l_effective_img_len) + + # Create position IDs + position_ids = torch.zeros(batch_size, max_seq_len, 3, dtype=torch.int32, device=device) + + for i, (cap_seq_len, seq_len) in enumerate(zip(l_effective_cap_len, seq_lengths)): + # add text position ids + position_ids[i, :cap_seq_len] = repeat(torch.arange(cap_seq_len, dtype=torch.int32, device=device), "l -> l 3") + + pe_shift = cap_seq_len + pe_shift_len = cap_seq_len + + if ref_img_sizes[i] is not None: + for ref_img_size, ref_img_len in zip(ref_img_sizes[i], l_effective_ref_img_len[i]): + H, W = ref_img_size + ref_H_tokens, ref_W_tokens = H // p, W // p + assert ref_H_tokens * ref_W_tokens == ref_img_len + # add image position ids + + row_ids = repeat(torch.arange(ref_H_tokens, dtype=torch.int32, device=device), "h -> h w", w=ref_W_tokens).flatten() + col_ids = repeat(torch.arange(ref_W_tokens, dtype=torch.int32, device=device), "w -> h w", h=ref_H_tokens).flatten() + position_ids[i, pe_shift_len:pe_shift_len + ref_img_len, 0] = pe_shift + position_ids[i, pe_shift_len:pe_shift_len + ref_img_len, 1] = row_ids + position_ids[i, pe_shift_len:pe_shift_len + ref_img_len, 2] = col_ids + + pe_shift += max(ref_H_tokens, ref_W_tokens) + pe_shift_len += ref_img_len + + H, W = img_sizes[i] + H_tokens, W_tokens = H // p, W // p + assert H_tokens * W_tokens == l_effective_img_len[i] + + row_ids = repeat(torch.arange(H_tokens, dtype=torch.int32, device=device), "h -> h w", w=W_tokens).flatten() + col_ids = repeat(torch.arange(W_tokens, dtype=torch.int32, device=device), "w -> h w", h=H_tokens).flatten() + + assert pe_shift_len + l_effective_img_len[i] == seq_len + position_ids[i, pe_shift_len: seq_len, 0] = pe_shift + position_ids[i, pe_shift_len: seq_len, 1] = row_ids + position_ids[i, pe_shift_len: seq_len, 2] = col_ids + + # Get combined rotary embeddings + freqs_cis = self._get_freqs_cis(freqs_cis, position_ids) + + # create separate rotary embeddings for captions and images + cap_freqs_cis = torch.zeros( + batch_size, encoder_seq_len, freqs_cis.shape[-1], device=device, dtype=freqs_cis.dtype + ) + ref_img_freqs_cis = torch.zeros( + batch_size, max_ref_img_len, freqs_cis.shape[-1], device=device, dtype=freqs_cis.dtype + ) + img_freqs_cis = torch.zeros( + batch_size, max_img_len, freqs_cis.shape[-1], device=device, dtype=freqs_cis.dtype + ) + + for i, (cap_seq_len, ref_img_len, img_len, seq_len) in enumerate(zip(l_effective_cap_len, l_effective_ref_img_len, l_effective_img_len, seq_lengths)): + cap_freqs_cis[i, :cap_seq_len] = freqs_cis[i, :cap_seq_len] + ref_img_freqs_cis[i, :sum(ref_img_len)] = freqs_cis[i, cap_seq_len:cap_seq_len + sum(ref_img_len)] + img_freqs_cis[i, :img_len] = freqs_cis[i, cap_seq_len + sum(ref_img_len):cap_seq_len + sum(ref_img_len) + img_len] + + return ( + cap_freqs_cis, + ref_img_freqs_cis, + img_freqs_cis, + freqs_cis, + l_effective_cap_len, + seq_lengths, + ) \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/models/transformers/transformer_omnigen2.py b/examples/OmniGen2-RL/omnigen2/models/transformers/transformer_omnigen2.py new file mode 100644 index 0000000000000000000000000000000000000000..46771497d8f28f60d9237a5be45b72230c4eb4fe --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/models/transformers/transformer_omnigen2.py @@ -0,0 +1,642 @@ +import warnings +import itertools +from typing import Any, Dict, List, Optional, Tuple, Union + +import torch +import torch.nn as nn + +from einops import rearrange + +from diffusers.configuration_utils import ConfigMixin, register_to_config +from diffusers.loaders import PeftAdapterMixin +from diffusers.loaders.single_file_model import FromOriginalModelMixin +from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscale_lora_layers +from diffusers.models.attention_processor import Attention +from diffusers.models.modeling_outputs import Transformer2DModelOutput +from diffusers.models.modeling_utils import ModelMixin + +from ..attention_processor import OmniGen2AttnProcessorFlash2Varlen, OmniGen2AttnProcessor +from .repo import OmniGen2RotaryPosEmbed +from .block_lumina2 import LuminaLayerNormContinuous, LuminaRMSNormZero, LuminaFeedForward, Lumina2CombinedTimestepCaptionEmbedding + +from ...utils.import_utils import is_triton_available, is_flash_attn_available + +if is_triton_available(): + from ...ops.triton.layer_norm import RMSNorm +else: + from torch.nn import RMSNorm + +logger = logging.get_logger(__name__) + + +class OmniGen2TransformerBlock(nn.Module): + """ + Transformer block for OmniGen2 model. + + This block implements a transformer layer with: + - Multi-head attention with flash attention + - Feed-forward network with SwiGLU activation + - RMS normalization + - Optional modulation for conditional generation + + Args: + dim: Dimension of the input and output tensors + num_attention_heads: Number of attention heads + num_kv_heads: Number of key-value heads + multiple_of: Multiple of which the hidden dimension should be + ffn_dim_multiplier: Multiplier for the feed-forward network dimension + norm_eps: Epsilon value for normalization layers + modulation: Whether to use modulation for conditional generation + use_fused_rms_norm: Whether to use fused RMS normalization + use_fused_swiglu: Whether to use fused SwiGLU activation + """ + + def __init__( + self, + dim: int, + num_attention_heads: int, + num_kv_heads: int, + multiple_of: int, + ffn_dim_multiplier: float, + norm_eps: float, + modulation: bool = True, + ) -> None: + """Initialize the transformer block.""" + super().__init__() + self.head_dim = dim // num_attention_heads + self.modulation = modulation + + try: + processor = OmniGen2AttnProcessorFlash2Varlen() + except ImportError: + processor = OmniGen2AttnProcessor() + + # Initialize attention layer + self.attn = Attention( + query_dim=dim, + cross_attention_dim=None, + dim_head=dim // num_attention_heads, + qk_norm="rms_norm", + heads=num_attention_heads, + kv_heads=num_kv_heads, + eps=1e-5, + bias=False, + out_bias=False, + processor=processor, + ) + + # Initialize feed-forward network + self.feed_forward = LuminaFeedForward( + dim=dim, + inner_dim=4 * dim, + multiple_of=multiple_of, + ffn_dim_multiplier=ffn_dim_multiplier + ) + + # Initialize normalization layers + if modulation: + self.norm1 = LuminaRMSNormZero( + embedding_dim=dim, + norm_eps=norm_eps, + norm_elementwise_affine=True + ) + else: + self.norm1 = RMSNorm(dim, eps=norm_eps) + + self.ffn_norm1 = RMSNorm(dim, eps=norm_eps) + self.norm2 = RMSNorm(dim, eps=norm_eps) + self.ffn_norm2 = RMSNorm(dim, eps=norm_eps) + + self.initialize_weights() + + def initialize_weights(self) -> None: + """ + Initialize the weights of the transformer block. + + Uses Xavier uniform initialization for linear layers and zero initialization for biases. + """ + nn.init.xavier_uniform_(self.attn.to_q.weight) + nn.init.xavier_uniform_(self.attn.to_k.weight) + nn.init.xavier_uniform_(self.attn.to_v.weight) + nn.init.xavier_uniform_(self.attn.to_out[0].weight) + + nn.init.xavier_uniform_(self.feed_forward.linear_1.weight) + nn.init.xavier_uniform_(self.feed_forward.linear_2.weight) + nn.init.xavier_uniform_(self.feed_forward.linear_3.weight) + + if self.modulation: + nn.init.zeros_(self.norm1.linear.weight) + nn.init.zeros_(self.norm1.linear.bias) + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor, + image_rotary_emb: torch.Tensor, + temb: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """ + Forward pass of the transformer block. + + Args: + hidden_states: Input hidden states tensor + attention_mask: Attention mask tensor + image_rotary_emb: Rotary embeddings for image tokens + temb: Optional timestep embedding tensor + + Returns: + torch.Tensor: Output hidden states after transformer block processing + """ + import time + if self.modulation: + if temb is None: + raise ValueError("temb must be provided when modulation is enabled") + + norm_hidden_states, gate_msa, scale_mlp, gate_mlp = self.norm1(hidden_states, temb) + attn_output = self.attn( + hidden_states=norm_hidden_states, + encoder_hidden_states=norm_hidden_states, + attention_mask=attention_mask, + image_rotary_emb=image_rotary_emb, + ) + hidden_states = hidden_states + gate_msa.unsqueeze(1).tanh() * self.norm2(attn_output) + mlp_output = self.feed_forward(self.ffn_norm1(hidden_states) * (1 + scale_mlp.unsqueeze(1))) + hidden_states = hidden_states + gate_mlp.unsqueeze(1).tanh() * self.ffn_norm2(mlp_output) + else: + norm_hidden_states = self.norm1(hidden_states) + attn_output = self.attn( + hidden_states=norm_hidden_states, + encoder_hidden_states=norm_hidden_states, + attention_mask=attention_mask, + image_rotary_emb=image_rotary_emb, + ) + hidden_states = hidden_states + self.norm2(attn_output) + mlp_output = self.feed_forward(self.ffn_norm1(hidden_states)) + hidden_states = hidden_states + self.ffn_norm2(mlp_output) + + return hidden_states + + +class OmniGen2Transformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): + """ + OmniGen2 Transformer 2D Model. + + A transformer-based diffusion model for image generation with: + - Patch-based image processing + - Rotary position embeddings + - Multi-head attention + - Conditional generation support + + Args: + patch_size: Size of image patches + in_channels: Number of input channels + out_channels: Number of output channels (defaults to in_channels) + hidden_size: Size of hidden layers + num_layers: Number of transformer layers + num_refiner_layers: Number of refiner layers + num_attention_heads: Number of attention heads + num_kv_heads: Number of key-value heads + multiple_of: Multiple of which the hidden dimension should be + ffn_dim_multiplier: Multiplier for feed-forward network dimension + norm_eps: Epsilon value for normalization layers + axes_dim_rope: Dimensions for rotary position embeddings + axes_lens: Lengths for rotary position embeddings + text_feat_dim: Dimension of text features + timestep_scale: Scale factor for timestep embeddings + use_fused_rms_norm: Whether to use fused RMS normalization + use_fused_swiglu: Whether to use fused SwiGLU activation + """ + + _supports_gradient_checkpointing = True + _no_split_modules = ["Omnigen2TransformerBlock"] + _skip_layerwise_casting_patterns = ["x_embedder", "norm"] + + @register_to_config + def __init__( + self, + patch_size: int = 2, + in_channels: int = 16, + out_channels: Optional[int] = None, + hidden_size: int = 2304, + num_layers: int = 26, + num_refiner_layers: int = 2, + num_attention_heads: int = 24, + num_kv_heads: int = 8, + multiple_of: int = 256, + ffn_dim_multiplier: Optional[float] = None, + norm_eps: float = 1e-5, + axes_dim_rope: Tuple[int, int, int] = (32, 32, 32), + axes_lens: Tuple[int, int, int] = (300, 512, 512), + text_feat_dim: int = 1024, + timestep_scale: float = 1.0 + ) -> None: + """Initialize the OmniGen2 transformer model.""" + super().__init__() + + # Validate configuration + if (hidden_size // num_attention_heads) != sum(axes_dim_rope): + raise ValueError( + f"hidden_size // num_attention_heads ({hidden_size // num_attention_heads}) " + f"must equal sum(axes_dim_rope) ({sum(axes_dim_rope)})" + ) + + self.out_channels = out_channels or in_channels + + # Initialize embeddings + self.rope_embedder = OmniGen2RotaryPosEmbed( + theta=10000, + axes_dim=axes_dim_rope, + axes_lens=axes_lens, + patch_size=patch_size, + ) + + self.x_embedder = nn.Linear( + in_features=patch_size * patch_size * in_channels, + out_features=hidden_size, + ) + + self.ref_image_patch_embedder = nn.Linear( + in_features=patch_size * patch_size * in_channels, + out_features=hidden_size, + ) + + self.time_caption_embed = Lumina2CombinedTimestepCaptionEmbedding( + hidden_size=hidden_size, + text_feat_dim=text_feat_dim, + norm_eps=norm_eps, + timestep_scale=timestep_scale + ) + + # Initialize transformer blocks + self.noise_refiner = nn.ModuleList([ + OmniGen2TransformerBlock( + hidden_size, + num_attention_heads, + num_kv_heads, + multiple_of, + ffn_dim_multiplier, + norm_eps, + modulation=True + ) + for _ in range(num_refiner_layers) + ]) + + self.ref_image_refiner = nn.ModuleList([ + OmniGen2TransformerBlock( + hidden_size, + num_attention_heads, + num_kv_heads, + multiple_of, + ffn_dim_multiplier, + norm_eps, + modulation=True + ) + for _ in range(num_refiner_layers) + ]) + + self.context_refiner = nn.ModuleList( + [ + OmniGen2TransformerBlock( + hidden_size, + num_attention_heads, + num_kv_heads, + multiple_of, + ffn_dim_multiplier, + norm_eps, + modulation=False + ) + for _ in range(num_refiner_layers) + ] + ) + + # 3. Transformer blocks + self.layers = nn.ModuleList( + [ + OmniGen2TransformerBlock( + hidden_size, + num_attention_heads, + num_kv_heads, + multiple_of, + ffn_dim_multiplier, + norm_eps, + modulation=True + ) + for _ in range(num_layers) + ] + ) + + # 4. Output norm & projection + self.norm_out = LuminaLayerNormContinuous( + embedding_dim=hidden_size, + conditioning_embedding_dim=min(hidden_size, 1024), + elementwise_affine=False, + eps=1e-6, + bias=True, + out_dim=patch_size * patch_size * self.out_channels + ) + + # Add learnable embeddings to distinguish different images + self.image_index_embedding = nn.Parameter(torch.randn(5, hidden_size)) # support max 5 ref images + + self.gradient_checkpointing = False + + self.initialize_weights() + + def initialize_weights(self) -> None: + """ + Initialize the weights of the model. + + Uses Xavier uniform initialization for linear layers. + """ + nn.init.xavier_uniform_(self.x_embedder.weight) + nn.init.constant_(self.x_embedder.bias, 0.0) + + nn.init.xavier_uniform_(self.ref_image_patch_embedder.weight) + nn.init.constant_(self.ref_image_patch_embedder.bias, 0.0) + + nn.init.zeros_(self.norm_out.linear_1.weight) + nn.init.zeros_(self.norm_out.linear_1.bias) + nn.init.zeros_(self.norm_out.linear_2.weight) + nn.init.zeros_(self.norm_out.linear_2.bias) + + nn.init.normal_(self.image_index_embedding, std=0.02) + + def img_patch_embed_and_refine( + self, + hidden_states, + ref_image_hidden_states, + padded_img_mask, + padded_ref_img_mask, + noise_rotary_emb, + ref_img_rotary_emb, + l_effective_ref_img_len, + l_effective_img_len, + temb + ): + batch_size = len(hidden_states) + max_combined_img_len = max([img_len + sum(ref_img_len) for img_len, ref_img_len in zip(l_effective_img_len, l_effective_ref_img_len)]) + + hidden_states = self.x_embedder(hidden_states) + ref_image_hidden_states = self.ref_image_patch_embedder(ref_image_hidden_states) + + for i in range(batch_size): + shift = 0 + for j, ref_img_len in enumerate(l_effective_ref_img_len[i]): + ref_image_hidden_states[i, shift:shift + ref_img_len, :] = ref_image_hidden_states[i, shift:shift + ref_img_len, :] + self.image_index_embedding[j] + shift += ref_img_len + + for layer in self.noise_refiner: + hidden_states = layer(hidden_states, padded_img_mask, noise_rotary_emb, temb) + + flat_l_effective_ref_img_len = list(itertools.chain(*l_effective_ref_img_len)) + num_ref_images = len(flat_l_effective_ref_img_len) + max_ref_img_len = max(flat_l_effective_ref_img_len) + + batch_ref_img_mask = ref_image_hidden_states.new_zeros(num_ref_images, max_ref_img_len, dtype=torch.bool) + batch_ref_image_hidden_states = ref_image_hidden_states.new_zeros(num_ref_images, max_ref_img_len, self.config.hidden_size) + batch_ref_img_rotary_emb = hidden_states.new_zeros(num_ref_images, max_ref_img_len, ref_img_rotary_emb.shape[-1], dtype=ref_img_rotary_emb.dtype) + batch_temb = temb.new_zeros(num_ref_images, *temb.shape[1:], dtype=temb.dtype) + + # sequence of ref imgs to batch + idx = 0 + for i in range(batch_size): + shift = 0 + for ref_img_len in l_effective_ref_img_len[i]: + batch_ref_img_mask[idx, :ref_img_len] = True + batch_ref_image_hidden_states[idx, :ref_img_len] = ref_image_hidden_states[i, shift:shift + ref_img_len] + batch_ref_img_rotary_emb[idx, :ref_img_len] = ref_img_rotary_emb[i, shift:shift + ref_img_len] + batch_temb[idx] = temb[i] + shift += ref_img_len + idx += 1 + + # refine ref imgs separately + for layer in self.ref_image_refiner: + batch_ref_image_hidden_states = layer(batch_ref_image_hidden_states, batch_ref_img_mask, batch_ref_img_rotary_emb, batch_temb) + + # batch of ref imgs to sequence + idx = 0 + for i in range(batch_size): + shift = 0 + for ref_img_len in l_effective_ref_img_len[i]: + ref_image_hidden_states[i, shift:shift + ref_img_len] = batch_ref_image_hidden_states[idx, :ref_img_len] + shift += ref_img_len + idx += 1 + + combined_img_hidden_states = hidden_states.new_zeros(batch_size, max_combined_img_len, self.config.hidden_size) + for i, (ref_img_len, img_len) in enumerate(zip(l_effective_ref_img_len, l_effective_img_len)): + combined_img_hidden_states[i, :sum(ref_img_len)] = ref_image_hidden_states[i, :sum(ref_img_len)] + combined_img_hidden_states[i, sum(ref_img_len):sum(ref_img_len) + img_len] = hidden_states[i, :img_len] + + return combined_img_hidden_states + + def flat_and_pad_to_seq_ref_img(self, ref_image_hidden_states, batch_size, dtype, device): + p = self.config.patch_size + if ref_image_hidden_states is not None: + ref_img_sizes = [[(img.size(1), img.size(2)) for img in imgs] if imgs is not None else None for imgs in ref_image_hidden_states] + l_effective_ref_img_len = [[(ref_img_size[0] // p) * (ref_img_size[1] // p) for ref_img_size in _ref_img_sizes] if _ref_img_sizes is not None else [0] for _ref_img_sizes in ref_img_sizes] + else: + ref_img_sizes = [None for _ in range(batch_size)] + l_effective_ref_img_len = [[0] for _ in range(batch_size)] + + max_ref_img_len = max([sum(ref_img_len) for ref_img_len in l_effective_ref_img_len]) + + # print(f"{len(ref_image_hidden_states)=} {ref_img_sizes=} {l_effective_ref_img_len=} {max_ref_img_len=}") + # ref image patch embeddings + flat_ref_img_hidden_states = [] + for i in range(batch_size): + if ref_img_sizes[i] is not None: + imgs = [] + for ref_img in ref_image_hidden_states[i]: + C, H, W = ref_img.size() + ref_img = rearrange(ref_img, 'c (h p1) (w p2) -> (h w) (p1 p2 c)', p1=p, p2=p) + imgs.append(ref_img) + + img = torch.cat(imgs, dim=0) + flat_ref_img_hidden_states.append(img) + else: + flat_ref_img_hidden_states.append(None) + + padded_ref_img_hidden_states = torch.zeros(batch_size, max_ref_img_len, self.config.in_channels * p * p, device=device, dtype=dtype) + padded_ref_img_mask = torch.zeros(batch_size, max_ref_img_len, dtype=torch.bool, device=device) + for i in range(batch_size): + if ref_img_sizes[i] is not None: + padded_ref_img_hidden_states[i, :sum(l_effective_ref_img_len[i])] = flat_ref_img_hidden_states[i] + padded_ref_img_mask[i, :sum(l_effective_ref_img_len[i])] = True + + return ( + padded_ref_img_hidden_states, + padded_ref_img_mask, + l_effective_ref_img_len, + ref_img_sizes, + ) + + def flat_and_pad_to_seq(self, hidden_states, batch_size, device): + p = self.config.patch_size + + img_sizes = [(img.size(1), img.size(2)) for img in hidden_states] + l_effective_img_len = [(H // p) * (W // p) for (H, W) in img_sizes] + + max_img_len = max(l_effective_img_len) + + # image patch embeddings + flat_hidden_states = [] + for i in range(batch_size): + img = hidden_states[i] + C, H, W = img.size() + + img = rearrange(img, 'c (h p1) (w p2) -> (h w) (p1 p2 c)', p1=p, p2=p) + flat_hidden_states.append(img) + + padded_hidden_states = torch.zeros(batch_size, max_img_len, flat_hidden_states[0].shape[-1], device=device, dtype=flat_hidden_states[0].dtype) + padded_img_mask = torch.zeros(batch_size, max_img_len, dtype=torch.bool, device=device) + for i in range(batch_size): + padded_hidden_states[i, :l_effective_img_len[i]] = flat_hidden_states[i] + padded_img_mask[i, :l_effective_img_len[i]] = True + + return ( + padded_hidden_states, + padded_img_mask, + l_effective_img_len, + img_sizes, + ) + + def forward( + self, + hidden_states: Union[torch.Tensor, List[torch.Tensor]], + timestep: torch.Tensor, + text_hidden_states: torch.Tensor, + text_attention_mask: torch.Tensor, + freqs_cis: torch.Tensor, + ref_image_hidden_states: Optional[List[List[torch.Tensor]]] = None, + attention_kwargs: Optional[Dict[str, Any]] = None, + return_dict: bool = False, + flat_and_pad: bool = True, + img_mask: Optional[torch.Tensor] = None, + ref_img_mask: Optional[torch.Tensor] = None, + l_effective_ref_img_len: Optional[List[List[int]]] = None, + l_effective_img_len: Optional[List[int]] = None, + ref_img_sizes: Optional[List[List[Tuple[int, int]]]] = None, + img_sizes: Optional[List[Tuple[int, int]]] = None, + ) -> Union[torch.Tensor, Transformer2DModelOutput]: + if attention_kwargs is not None: + attention_kwargs = attention_kwargs.copy() + lora_scale = attention_kwargs.pop("scale", 1.0) + else: + lora_scale = 1.0 + + if USE_PEFT_BACKEND: + # weight the lora layers by setting `lora_scale` for each PEFT layer + scale_lora_layers(self, lora_scale) + else: + if attention_kwargs is not None and attention_kwargs.get("scale", None) is not None: + logger.warning( + "Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective." + ) + + # 1. Condition, positional & patch embedding + batch_size = len(hidden_states) + is_hidden_states_tensor = isinstance(hidden_states, torch.Tensor) + if is_hidden_states_tensor: + dtype = hidden_states.dtype + device = hidden_states.device + else: + dtype = hidden_states[0].dtype + device = hidden_states[0].device + + temb, text_hidden_states = self.time_caption_embed(timestep, text_hidden_states, hidden_states[0].dtype) + + if flat_and_pad: + if is_hidden_states_tensor: + hidden_states = [_hidden_states for _hidden_states in hidden_states] + + ( + hidden_states, + img_mask, + l_effective_img_len, + img_sizes, + ) = self.flat_and_pad_to_seq(hidden_states, batch_size, device) + + ( + ref_image_hidden_states, + ref_img_mask, + l_effective_ref_img_len, + ref_img_sizes, + ) = self.flat_and_pad_to_seq_ref_img(ref_image_hidden_states, batch_size, dtype, device) + + ( + context_rotary_emb, + ref_img_rotary_emb, + noise_rotary_emb, + rotary_emb, + encoder_seq_lengths, + seq_lengths, + ) = self.rope_embedder( + freqs_cis, + text_attention_mask, + l_effective_ref_img_len, + l_effective_img_len, + ref_img_sizes, + img_sizes, + device, + ) + + # 2. Context refinement + for layer in self.context_refiner: + text_hidden_states = layer(text_hidden_states, text_attention_mask, context_rotary_emb) + + combined_img_hidden_states = self.img_patch_embed_and_refine( + hidden_states, + ref_image_hidden_states, + img_mask, + ref_img_mask, + noise_rotary_emb, + ref_img_rotary_emb, + l_effective_ref_img_len, + l_effective_img_len, + temb, + ) + + # 3. Joint Transformer blocks + max_seq_len = max(seq_lengths) + + attention_mask = hidden_states.new_zeros(batch_size, max_seq_len, dtype=torch.bool) + joint_hidden_states = hidden_states.new_zeros(batch_size, max_seq_len, self.config.hidden_size) + for i, (encoder_seq_len, seq_len) in enumerate(zip(encoder_seq_lengths, seq_lengths)): + attention_mask[i, :seq_len] = True + joint_hidden_states[i, :encoder_seq_len] = text_hidden_states[i, :encoder_seq_len] + joint_hidden_states[i, encoder_seq_len:seq_len] = combined_img_hidden_states[i, :seq_len - encoder_seq_len] + + hidden_states = joint_hidden_states + + for layer_idx, layer in enumerate(self.layers): + if torch.is_grad_enabled() and self.gradient_checkpointing: + hidden_states = self._gradient_checkpointing_func( + layer, hidden_states, attention_mask, rotary_emb, temb + ) + else: + hidden_states = layer(hidden_states, attention_mask, rotary_emb, temb) + + # 4. Output norm & projection + hidden_states = self.norm_out(hidden_states, temb) + + if flat_and_pad: + p = self.config.patch_size + output = [] + for i, (img_size, img_len, seq_len) in enumerate(zip(img_sizes, l_effective_img_len, seq_lengths)): + height, width = img_size + output.append(rearrange(hidden_states[i][seq_len - img_len:seq_len], '(h w) (p1 p2 c) -> c (h p1) (w p2)', h=height // p, w=width // p, p1=p, p2=p)) + if is_hidden_states_tensor: + output = torch.stack(output, dim=0) + else: + max_img_len = max(l_effective_img_len) + output = torch.zeros(batch_size, max_img_len, hidden_states[0].shape[-1], device=device, dtype=hidden_states[0].dtype) + for i, (img_size, img_len, seq_len) in enumerate(zip(img_sizes, l_effective_img_len, seq_lengths)): + output[i][:img_len] = hidden_states[i][seq_len - img_len:seq_len] + + if USE_PEFT_BACKEND: + # remove `lora_scale` from each PEFT layer + unscale_lora_layers(self, lora_scale) + + if not return_dict: + return output + return Transformer2DModelOutput(sample=output) diff --git a/examples/OmniGen2-RL/omnigen2/models/transformers/transformer_omnigen2_test.py b/examples/OmniGen2-RL/omnigen2/models/transformers/transformer_omnigen2_test.py new file mode 100644 index 0000000000000000000000000000000000000000..46771497d8f28f60d9237a5be45b72230c4eb4fe --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/models/transformers/transformer_omnigen2_test.py @@ -0,0 +1,642 @@ +import warnings +import itertools +from typing import Any, Dict, List, Optional, Tuple, Union + +import torch +import torch.nn as nn + +from einops import rearrange + +from diffusers.configuration_utils import ConfigMixin, register_to_config +from diffusers.loaders import PeftAdapterMixin +from diffusers.loaders.single_file_model import FromOriginalModelMixin +from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscale_lora_layers +from diffusers.models.attention_processor import Attention +from diffusers.models.modeling_outputs import Transformer2DModelOutput +from diffusers.models.modeling_utils import ModelMixin + +from ..attention_processor import OmniGen2AttnProcessorFlash2Varlen, OmniGen2AttnProcessor +from .repo import OmniGen2RotaryPosEmbed +from .block_lumina2 import LuminaLayerNormContinuous, LuminaRMSNormZero, LuminaFeedForward, Lumina2CombinedTimestepCaptionEmbedding + +from ...utils.import_utils import is_triton_available, is_flash_attn_available + +if is_triton_available(): + from ...ops.triton.layer_norm import RMSNorm +else: + from torch.nn import RMSNorm + +logger = logging.get_logger(__name__) + + +class OmniGen2TransformerBlock(nn.Module): + """ + Transformer block for OmniGen2 model. + + This block implements a transformer layer with: + - Multi-head attention with flash attention + - Feed-forward network with SwiGLU activation + - RMS normalization + - Optional modulation for conditional generation + + Args: + dim: Dimension of the input and output tensors + num_attention_heads: Number of attention heads + num_kv_heads: Number of key-value heads + multiple_of: Multiple of which the hidden dimension should be + ffn_dim_multiplier: Multiplier for the feed-forward network dimension + norm_eps: Epsilon value for normalization layers + modulation: Whether to use modulation for conditional generation + use_fused_rms_norm: Whether to use fused RMS normalization + use_fused_swiglu: Whether to use fused SwiGLU activation + """ + + def __init__( + self, + dim: int, + num_attention_heads: int, + num_kv_heads: int, + multiple_of: int, + ffn_dim_multiplier: float, + norm_eps: float, + modulation: bool = True, + ) -> None: + """Initialize the transformer block.""" + super().__init__() + self.head_dim = dim // num_attention_heads + self.modulation = modulation + + try: + processor = OmniGen2AttnProcessorFlash2Varlen() + except ImportError: + processor = OmniGen2AttnProcessor() + + # Initialize attention layer + self.attn = Attention( + query_dim=dim, + cross_attention_dim=None, + dim_head=dim // num_attention_heads, + qk_norm="rms_norm", + heads=num_attention_heads, + kv_heads=num_kv_heads, + eps=1e-5, + bias=False, + out_bias=False, + processor=processor, + ) + + # Initialize feed-forward network + self.feed_forward = LuminaFeedForward( + dim=dim, + inner_dim=4 * dim, + multiple_of=multiple_of, + ffn_dim_multiplier=ffn_dim_multiplier + ) + + # Initialize normalization layers + if modulation: + self.norm1 = LuminaRMSNormZero( + embedding_dim=dim, + norm_eps=norm_eps, + norm_elementwise_affine=True + ) + else: + self.norm1 = RMSNorm(dim, eps=norm_eps) + + self.ffn_norm1 = RMSNorm(dim, eps=norm_eps) + self.norm2 = RMSNorm(dim, eps=norm_eps) + self.ffn_norm2 = RMSNorm(dim, eps=norm_eps) + + self.initialize_weights() + + def initialize_weights(self) -> None: + """ + Initialize the weights of the transformer block. + + Uses Xavier uniform initialization for linear layers and zero initialization for biases. + """ + nn.init.xavier_uniform_(self.attn.to_q.weight) + nn.init.xavier_uniform_(self.attn.to_k.weight) + nn.init.xavier_uniform_(self.attn.to_v.weight) + nn.init.xavier_uniform_(self.attn.to_out[0].weight) + + nn.init.xavier_uniform_(self.feed_forward.linear_1.weight) + nn.init.xavier_uniform_(self.feed_forward.linear_2.weight) + nn.init.xavier_uniform_(self.feed_forward.linear_3.weight) + + if self.modulation: + nn.init.zeros_(self.norm1.linear.weight) + nn.init.zeros_(self.norm1.linear.bias) + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor, + image_rotary_emb: torch.Tensor, + temb: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """ + Forward pass of the transformer block. + + Args: + hidden_states: Input hidden states tensor + attention_mask: Attention mask tensor + image_rotary_emb: Rotary embeddings for image tokens + temb: Optional timestep embedding tensor + + Returns: + torch.Tensor: Output hidden states after transformer block processing + """ + import time + if self.modulation: + if temb is None: + raise ValueError("temb must be provided when modulation is enabled") + + norm_hidden_states, gate_msa, scale_mlp, gate_mlp = self.norm1(hidden_states, temb) + attn_output = self.attn( + hidden_states=norm_hidden_states, + encoder_hidden_states=norm_hidden_states, + attention_mask=attention_mask, + image_rotary_emb=image_rotary_emb, + ) + hidden_states = hidden_states + gate_msa.unsqueeze(1).tanh() * self.norm2(attn_output) + mlp_output = self.feed_forward(self.ffn_norm1(hidden_states) * (1 + scale_mlp.unsqueeze(1))) + hidden_states = hidden_states + gate_mlp.unsqueeze(1).tanh() * self.ffn_norm2(mlp_output) + else: + norm_hidden_states = self.norm1(hidden_states) + attn_output = self.attn( + hidden_states=norm_hidden_states, + encoder_hidden_states=norm_hidden_states, + attention_mask=attention_mask, + image_rotary_emb=image_rotary_emb, + ) + hidden_states = hidden_states + self.norm2(attn_output) + mlp_output = self.feed_forward(self.ffn_norm1(hidden_states)) + hidden_states = hidden_states + self.ffn_norm2(mlp_output) + + return hidden_states + + +class OmniGen2Transformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): + """ + OmniGen2 Transformer 2D Model. + + A transformer-based diffusion model for image generation with: + - Patch-based image processing + - Rotary position embeddings + - Multi-head attention + - Conditional generation support + + Args: + patch_size: Size of image patches + in_channels: Number of input channels + out_channels: Number of output channels (defaults to in_channels) + hidden_size: Size of hidden layers + num_layers: Number of transformer layers + num_refiner_layers: Number of refiner layers + num_attention_heads: Number of attention heads + num_kv_heads: Number of key-value heads + multiple_of: Multiple of which the hidden dimension should be + ffn_dim_multiplier: Multiplier for feed-forward network dimension + norm_eps: Epsilon value for normalization layers + axes_dim_rope: Dimensions for rotary position embeddings + axes_lens: Lengths for rotary position embeddings + text_feat_dim: Dimension of text features + timestep_scale: Scale factor for timestep embeddings + use_fused_rms_norm: Whether to use fused RMS normalization + use_fused_swiglu: Whether to use fused SwiGLU activation + """ + + _supports_gradient_checkpointing = True + _no_split_modules = ["Omnigen2TransformerBlock"] + _skip_layerwise_casting_patterns = ["x_embedder", "norm"] + + @register_to_config + def __init__( + self, + patch_size: int = 2, + in_channels: int = 16, + out_channels: Optional[int] = None, + hidden_size: int = 2304, + num_layers: int = 26, + num_refiner_layers: int = 2, + num_attention_heads: int = 24, + num_kv_heads: int = 8, + multiple_of: int = 256, + ffn_dim_multiplier: Optional[float] = None, + norm_eps: float = 1e-5, + axes_dim_rope: Tuple[int, int, int] = (32, 32, 32), + axes_lens: Tuple[int, int, int] = (300, 512, 512), + text_feat_dim: int = 1024, + timestep_scale: float = 1.0 + ) -> None: + """Initialize the OmniGen2 transformer model.""" + super().__init__() + + # Validate configuration + if (hidden_size // num_attention_heads) != sum(axes_dim_rope): + raise ValueError( + f"hidden_size // num_attention_heads ({hidden_size // num_attention_heads}) " + f"must equal sum(axes_dim_rope) ({sum(axes_dim_rope)})" + ) + + self.out_channels = out_channels or in_channels + + # Initialize embeddings + self.rope_embedder = OmniGen2RotaryPosEmbed( + theta=10000, + axes_dim=axes_dim_rope, + axes_lens=axes_lens, + patch_size=patch_size, + ) + + self.x_embedder = nn.Linear( + in_features=patch_size * patch_size * in_channels, + out_features=hidden_size, + ) + + self.ref_image_patch_embedder = nn.Linear( + in_features=patch_size * patch_size * in_channels, + out_features=hidden_size, + ) + + self.time_caption_embed = Lumina2CombinedTimestepCaptionEmbedding( + hidden_size=hidden_size, + text_feat_dim=text_feat_dim, + norm_eps=norm_eps, + timestep_scale=timestep_scale + ) + + # Initialize transformer blocks + self.noise_refiner = nn.ModuleList([ + OmniGen2TransformerBlock( + hidden_size, + num_attention_heads, + num_kv_heads, + multiple_of, + ffn_dim_multiplier, + norm_eps, + modulation=True + ) + for _ in range(num_refiner_layers) + ]) + + self.ref_image_refiner = nn.ModuleList([ + OmniGen2TransformerBlock( + hidden_size, + num_attention_heads, + num_kv_heads, + multiple_of, + ffn_dim_multiplier, + norm_eps, + modulation=True + ) + for _ in range(num_refiner_layers) + ]) + + self.context_refiner = nn.ModuleList( + [ + OmniGen2TransformerBlock( + hidden_size, + num_attention_heads, + num_kv_heads, + multiple_of, + ffn_dim_multiplier, + norm_eps, + modulation=False + ) + for _ in range(num_refiner_layers) + ] + ) + + # 3. Transformer blocks + self.layers = nn.ModuleList( + [ + OmniGen2TransformerBlock( + hidden_size, + num_attention_heads, + num_kv_heads, + multiple_of, + ffn_dim_multiplier, + norm_eps, + modulation=True + ) + for _ in range(num_layers) + ] + ) + + # 4. Output norm & projection + self.norm_out = LuminaLayerNormContinuous( + embedding_dim=hidden_size, + conditioning_embedding_dim=min(hidden_size, 1024), + elementwise_affine=False, + eps=1e-6, + bias=True, + out_dim=patch_size * patch_size * self.out_channels + ) + + # Add learnable embeddings to distinguish different images + self.image_index_embedding = nn.Parameter(torch.randn(5, hidden_size)) # support max 5 ref images + + self.gradient_checkpointing = False + + self.initialize_weights() + + def initialize_weights(self) -> None: + """ + Initialize the weights of the model. + + Uses Xavier uniform initialization for linear layers. + """ + nn.init.xavier_uniform_(self.x_embedder.weight) + nn.init.constant_(self.x_embedder.bias, 0.0) + + nn.init.xavier_uniform_(self.ref_image_patch_embedder.weight) + nn.init.constant_(self.ref_image_patch_embedder.bias, 0.0) + + nn.init.zeros_(self.norm_out.linear_1.weight) + nn.init.zeros_(self.norm_out.linear_1.bias) + nn.init.zeros_(self.norm_out.linear_2.weight) + nn.init.zeros_(self.norm_out.linear_2.bias) + + nn.init.normal_(self.image_index_embedding, std=0.02) + + def img_patch_embed_and_refine( + self, + hidden_states, + ref_image_hidden_states, + padded_img_mask, + padded_ref_img_mask, + noise_rotary_emb, + ref_img_rotary_emb, + l_effective_ref_img_len, + l_effective_img_len, + temb + ): + batch_size = len(hidden_states) + max_combined_img_len = max([img_len + sum(ref_img_len) for img_len, ref_img_len in zip(l_effective_img_len, l_effective_ref_img_len)]) + + hidden_states = self.x_embedder(hidden_states) + ref_image_hidden_states = self.ref_image_patch_embedder(ref_image_hidden_states) + + for i in range(batch_size): + shift = 0 + for j, ref_img_len in enumerate(l_effective_ref_img_len[i]): + ref_image_hidden_states[i, shift:shift + ref_img_len, :] = ref_image_hidden_states[i, shift:shift + ref_img_len, :] + self.image_index_embedding[j] + shift += ref_img_len + + for layer in self.noise_refiner: + hidden_states = layer(hidden_states, padded_img_mask, noise_rotary_emb, temb) + + flat_l_effective_ref_img_len = list(itertools.chain(*l_effective_ref_img_len)) + num_ref_images = len(flat_l_effective_ref_img_len) + max_ref_img_len = max(flat_l_effective_ref_img_len) + + batch_ref_img_mask = ref_image_hidden_states.new_zeros(num_ref_images, max_ref_img_len, dtype=torch.bool) + batch_ref_image_hidden_states = ref_image_hidden_states.new_zeros(num_ref_images, max_ref_img_len, self.config.hidden_size) + batch_ref_img_rotary_emb = hidden_states.new_zeros(num_ref_images, max_ref_img_len, ref_img_rotary_emb.shape[-1], dtype=ref_img_rotary_emb.dtype) + batch_temb = temb.new_zeros(num_ref_images, *temb.shape[1:], dtype=temb.dtype) + + # sequence of ref imgs to batch + idx = 0 + for i in range(batch_size): + shift = 0 + for ref_img_len in l_effective_ref_img_len[i]: + batch_ref_img_mask[idx, :ref_img_len] = True + batch_ref_image_hidden_states[idx, :ref_img_len] = ref_image_hidden_states[i, shift:shift + ref_img_len] + batch_ref_img_rotary_emb[idx, :ref_img_len] = ref_img_rotary_emb[i, shift:shift + ref_img_len] + batch_temb[idx] = temb[i] + shift += ref_img_len + idx += 1 + + # refine ref imgs separately + for layer in self.ref_image_refiner: + batch_ref_image_hidden_states = layer(batch_ref_image_hidden_states, batch_ref_img_mask, batch_ref_img_rotary_emb, batch_temb) + + # batch of ref imgs to sequence + idx = 0 + for i in range(batch_size): + shift = 0 + for ref_img_len in l_effective_ref_img_len[i]: + ref_image_hidden_states[i, shift:shift + ref_img_len] = batch_ref_image_hidden_states[idx, :ref_img_len] + shift += ref_img_len + idx += 1 + + combined_img_hidden_states = hidden_states.new_zeros(batch_size, max_combined_img_len, self.config.hidden_size) + for i, (ref_img_len, img_len) in enumerate(zip(l_effective_ref_img_len, l_effective_img_len)): + combined_img_hidden_states[i, :sum(ref_img_len)] = ref_image_hidden_states[i, :sum(ref_img_len)] + combined_img_hidden_states[i, sum(ref_img_len):sum(ref_img_len) + img_len] = hidden_states[i, :img_len] + + return combined_img_hidden_states + + def flat_and_pad_to_seq_ref_img(self, ref_image_hidden_states, batch_size, dtype, device): + p = self.config.patch_size + if ref_image_hidden_states is not None: + ref_img_sizes = [[(img.size(1), img.size(2)) for img in imgs] if imgs is not None else None for imgs in ref_image_hidden_states] + l_effective_ref_img_len = [[(ref_img_size[0] // p) * (ref_img_size[1] // p) for ref_img_size in _ref_img_sizes] if _ref_img_sizes is not None else [0] for _ref_img_sizes in ref_img_sizes] + else: + ref_img_sizes = [None for _ in range(batch_size)] + l_effective_ref_img_len = [[0] for _ in range(batch_size)] + + max_ref_img_len = max([sum(ref_img_len) for ref_img_len in l_effective_ref_img_len]) + + # print(f"{len(ref_image_hidden_states)=} {ref_img_sizes=} {l_effective_ref_img_len=} {max_ref_img_len=}") + # ref image patch embeddings + flat_ref_img_hidden_states = [] + for i in range(batch_size): + if ref_img_sizes[i] is not None: + imgs = [] + for ref_img in ref_image_hidden_states[i]: + C, H, W = ref_img.size() + ref_img = rearrange(ref_img, 'c (h p1) (w p2) -> (h w) (p1 p2 c)', p1=p, p2=p) + imgs.append(ref_img) + + img = torch.cat(imgs, dim=0) + flat_ref_img_hidden_states.append(img) + else: + flat_ref_img_hidden_states.append(None) + + padded_ref_img_hidden_states = torch.zeros(batch_size, max_ref_img_len, self.config.in_channels * p * p, device=device, dtype=dtype) + padded_ref_img_mask = torch.zeros(batch_size, max_ref_img_len, dtype=torch.bool, device=device) + for i in range(batch_size): + if ref_img_sizes[i] is not None: + padded_ref_img_hidden_states[i, :sum(l_effective_ref_img_len[i])] = flat_ref_img_hidden_states[i] + padded_ref_img_mask[i, :sum(l_effective_ref_img_len[i])] = True + + return ( + padded_ref_img_hidden_states, + padded_ref_img_mask, + l_effective_ref_img_len, + ref_img_sizes, + ) + + def flat_and_pad_to_seq(self, hidden_states, batch_size, device): + p = self.config.patch_size + + img_sizes = [(img.size(1), img.size(2)) for img in hidden_states] + l_effective_img_len = [(H // p) * (W // p) for (H, W) in img_sizes] + + max_img_len = max(l_effective_img_len) + + # image patch embeddings + flat_hidden_states = [] + for i in range(batch_size): + img = hidden_states[i] + C, H, W = img.size() + + img = rearrange(img, 'c (h p1) (w p2) -> (h w) (p1 p2 c)', p1=p, p2=p) + flat_hidden_states.append(img) + + padded_hidden_states = torch.zeros(batch_size, max_img_len, flat_hidden_states[0].shape[-1], device=device, dtype=flat_hidden_states[0].dtype) + padded_img_mask = torch.zeros(batch_size, max_img_len, dtype=torch.bool, device=device) + for i in range(batch_size): + padded_hidden_states[i, :l_effective_img_len[i]] = flat_hidden_states[i] + padded_img_mask[i, :l_effective_img_len[i]] = True + + return ( + padded_hidden_states, + padded_img_mask, + l_effective_img_len, + img_sizes, + ) + + def forward( + self, + hidden_states: Union[torch.Tensor, List[torch.Tensor]], + timestep: torch.Tensor, + text_hidden_states: torch.Tensor, + text_attention_mask: torch.Tensor, + freqs_cis: torch.Tensor, + ref_image_hidden_states: Optional[List[List[torch.Tensor]]] = None, + attention_kwargs: Optional[Dict[str, Any]] = None, + return_dict: bool = False, + flat_and_pad: bool = True, + img_mask: Optional[torch.Tensor] = None, + ref_img_mask: Optional[torch.Tensor] = None, + l_effective_ref_img_len: Optional[List[List[int]]] = None, + l_effective_img_len: Optional[List[int]] = None, + ref_img_sizes: Optional[List[List[Tuple[int, int]]]] = None, + img_sizes: Optional[List[Tuple[int, int]]] = None, + ) -> Union[torch.Tensor, Transformer2DModelOutput]: + if attention_kwargs is not None: + attention_kwargs = attention_kwargs.copy() + lora_scale = attention_kwargs.pop("scale", 1.0) + else: + lora_scale = 1.0 + + if USE_PEFT_BACKEND: + # weight the lora layers by setting `lora_scale` for each PEFT layer + scale_lora_layers(self, lora_scale) + else: + if attention_kwargs is not None and attention_kwargs.get("scale", None) is not None: + logger.warning( + "Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective." + ) + + # 1. Condition, positional & patch embedding + batch_size = len(hidden_states) + is_hidden_states_tensor = isinstance(hidden_states, torch.Tensor) + if is_hidden_states_tensor: + dtype = hidden_states.dtype + device = hidden_states.device + else: + dtype = hidden_states[0].dtype + device = hidden_states[0].device + + temb, text_hidden_states = self.time_caption_embed(timestep, text_hidden_states, hidden_states[0].dtype) + + if flat_and_pad: + if is_hidden_states_tensor: + hidden_states = [_hidden_states for _hidden_states in hidden_states] + + ( + hidden_states, + img_mask, + l_effective_img_len, + img_sizes, + ) = self.flat_and_pad_to_seq(hidden_states, batch_size, device) + + ( + ref_image_hidden_states, + ref_img_mask, + l_effective_ref_img_len, + ref_img_sizes, + ) = self.flat_and_pad_to_seq_ref_img(ref_image_hidden_states, batch_size, dtype, device) + + ( + context_rotary_emb, + ref_img_rotary_emb, + noise_rotary_emb, + rotary_emb, + encoder_seq_lengths, + seq_lengths, + ) = self.rope_embedder( + freqs_cis, + text_attention_mask, + l_effective_ref_img_len, + l_effective_img_len, + ref_img_sizes, + img_sizes, + device, + ) + + # 2. Context refinement + for layer in self.context_refiner: + text_hidden_states = layer(text_hidden_states, text_attention_mask, context_rotary_emb) + + combined_img_hidden_states = self.img_patch_embed_and_refine( + hidden_states, + ref_image_hidden_states, + img_mask, + ref_img_mask, + noise_rotary_emb, + ref_img_rotary_emb, + l_effective_ref_img_len, + l_effective_img_len, + temb, + ) + + # 3. Joint Transformer blocks + max_seq_len = max(seq_lengths) + + attention_mask = hidden_states.new_zeros(batch_size, max_seq_len, dtype=torch.bool) + joint_hidden_states = hidden_states.new_zeros(batch_size, max_seq_len, self.config.hidden_size) + for i, (encoder_seq_len, seq_len) in enumerate(zip(encoder_seq_lengths, seq_lengths)): + attention_mask[i, :seq_len] = True + joint_hidden_states[i, :encoder_seq_len] = text_hidden_states[i, :encoder_seq_len] + joint_hidden_states[i, encoder_seq_len:seq_len] = combined_img_hidden_states[i, :seq_len - encoder_seq_len] + + hidden_states = joint_hidden_states + + for layer_idx, layer in enumerate(self.layers): + if torch.is_grad_enabled() and self.gradient_checkpointing: + hidden_states = self._gradient_checkpointing_func( + layer, hidden_states, attention_mask, rotary_emb, temb + ) + else: + hidden_states = layer(hidden_states, attention_mask, rotary_emb, temb) + + # 4. Output norm & projection + hidden_states = self.norm_out(hidden_states, temb) + + if flat_and_pad: + p = self.config.patch_size + output = [] + for i, (img_size, img_len, seq_len) in enumerate(zip(img_sizes, l_effective_img_len, seq_lengths)): + height, width = img_size + output.append(rearrange(hidden_states[i][seq_len - img_len:seq_len], '(h w) (p1 p2 c) -> c (h p1) (w p2)', h=height // p, w=width // p, p1=p, p2=p)) + if is_hidden_states_tensor: + output = torch.stack(output, dim=0) + else: + max_img_len = max(l_effective_img_len) + output = torch.zeros(batch_size, max_img_len, hidden_states[0].shape[-1], device=device, dtype=hidden_states[0].dtype) + for i, (img_size, img_len, seq_len) in enumerate(zip(img_sizes, l_effective_img_len, seq_lengths)): + output[i][:img_len] = hidden_states[i][seq_len - img_len:seq_len] + + if USE_PEFT_BACKEND: + # remove `lora_scale` from each PEFT layer + unscale_lora_layers(self, lora_scale) + + if not return_dict: + return output + return Transformer2DModelOutput(sample=output) diff --git a/examples/OmniGen2-RL/omnigen2/ops/triton/__init__.py b/examples/OmniGen2-RL/omnigen2/ops/triton/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/examples/OmniGen2-RL/omnigen2/ops/triton/layer_norm.py b/examples/OmniGen2-RL/omnigen2/ops/triton/layer_norm.py new file mode 100644 index 0000000000000000000000000000000000000000..b7d123304559857a3b9ab903cd4faab6029e436c --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/ops/triton/layer_norm.py @@ -0,0 +1,1257 @@ +# Copyright (c) 2024, Tri Dao. +# Implement dropout + residual + layer_norm / rms_norm. + +# Based on the Triton LayerNorm tutorial: https://triton-lang.org/main/getting-started/tutorials/05-layer-norm.html +# For the backward pass, we keep weight_grad and bias_grad in registers and accumulate. +# This is faster for dimensions up to 8k, but after that it's much slower due to register spilling. +# The models we train have hidden dim up to 8k anyway (e.g. Llama 70B), so this is fine. + +import math + +import torch +import torch.nn.functional as F + +import triton +import triton.language as tl + + +from typing import Callable + + +def custom_amp_decorator(dec: Callable, cuda_amp_deprecated: bool): + def decorator(*args, **kwargs): + if cuda_amp_deprecated: + kwargs["device_type"] = "cuda" + return dec(*args, **kwargs) + return decorator + + +if hasattr(torch.amp, "custom_fwd"): # type: ignore[attr-defined] + deprecated = True + from torch.amp import custom_fwd, custom_bwd # type: ignore[attr-defined] +else: + deprecated = False + from torch.cuda.amp import custom_fwd, custom_bwd + +custom_fwd = custom_amp_decorator(custom_fwd, deprecated) +custom_bwd = custom_amp_decorator(custom_bwd, deprecated) + + +def triton_autotune_configs(): + # Return configs with a valid warp count for the current device + configs=[] + # Maximum threads per block is architecture-dependent in theory, but in reality all are 1024 + max_threads_per_block=1024 + # Default to warp size 32 if not defined by device + warp_size=getattr(torch.cuda.get_device_properties(torch.cuda.current_device()), "warp_size", 32) + # Autotune for warp counts which are powers of 2 and do not exceed thread per block limit + warp_count=1 + while warp_count*warp_size <= max_threads_per_block: + configs.append(triton.Config({}, num_warps=warp_count)) + warp_count*=2 + return configs + +def layer_norm_ref( + x, + weight, + bias, + residual=None, + x1=None, + weight1=None, + bias1=None, + eps=1e-6, + dropout_p=0.0, + rowscale=None, + prenorm=False, + zero_centered_weight=False, + dropout_mask=None, + dropout_mask1=None, + upcast=False, +): + dtype = x.dtype + if upcast: + x = x.float() + weight = weight.float() + bias = bias.float() if bias is not None else None + residual = residual.float() if residual is not None else residual + x1 = x1.float() if x1 is not None else None + weight1 = weight1.float() if weight1 is not None else None + bias1 = bias1.float() if bias1 is not None else None + if zero_centered_weight: + weight = weight + 1.0 + if weight1 is not None: + weight1 = weight1 + 1.0 + if x1 is not None: + assert rowscale is None, "rowscale is not supported with parallel LayerNorm" + if rowscale is not None: + x = x * rowscale[..., None] + if dropout_p > 0.0: + if dropout_mask is not None: + x = x.masked_fill(~dropout_mask, 0.0) / (1.0 - dropout_p) + else: + x = F.dropout(x, p=dropout_p) + if x1 is not None: + if dropout_mask1 is not None: + x1 = x1.masked_fill(~dropout_mask1, 0.0) / (1.0 - dropout_p) + else: + x1 = F.dropout(x1, p=dropout_p) + if x1 is not None: + x = x + x1 + if residual is not None: + x = (x + residual).to(x.dtype) + out = F.layer_norm(x.to(weight.dtype), x.shape[-1:], weight=weight, bias=bias, eps=eps).to( + dtype + ) + if weight1 is None: + return out if not prenorm else (out, x) + else: + out1 = F.layer_norm( + x.to(weight1.dtype), x.shape[-1:], weight=weight1, bias=bias1, eps=eps + ).to(dtype) + return (out, out1) if not prenorm else (out, out1, x) + + +def rms_norm_ref( + x, + weight, + bias, + residual=None, + x1=None, + weight1=None, + bias1=None, + eps=1e-6, + dropout_p=0.0, + rowscale=None, + prenorm=False, + zero_centered_weight=False, + dropout_mask=None, + dropout_mask1=None, + upcast=False, +): + dtype = x.dtype + if upcast: + x = x.float() + weight = weight.float() + bias = bias.float() if bias is not None else None + residual = residual.float() if residual is not None else residual + x1 = x1.float() if x1 is not None else None + weight1 = weight1.float() if weight1 is not None else None + bias1 = bias1.float() if bias1 is not None else None + if zero_centered_weight: + weight = weight + 1.0 + if weight1 is not None: + weight1 = weight1 + 1.0 + if x1 is not None: + assert rowscale is None, "rowscale is not supported with parallel LayerNorm" + if rowscale is not None: + x = x * rowscale[..., None] + if dropout_p > 0.0: + if dropout_mask is not None: + x = x.masked_fill(~dropout_mask, 0.0) / (1.0 - dropout_p) + else: + x = F.dropout(x, p=dropout_p) + if x1 is not None: + if dropout_mask1 is not None: + x1 = x1.masked_fill(~dropout_mask1, 0.0) / (1.0 - dropout_p) + else: + x1 = F.dropout(x1, p=dropout_p) + if x1 is not None: + x = x + x1 + if residual is not None: + x = (x + residual).to(x.dtype) + rstd = 1 / torch.sqrt((x.square()).mean(dim=-1, keepdim=True) + eps) + out = ((x * rstd * weight) + bias if bias is not None else (x * rstd * weight)).to(dtype) + if weight1 is None: + return out if not prenorm else (out, x) + else: + out1 = ((x * rstd * weight1) + bias1 if bias1 is not None else (x * rstd * weight1)).to( + dtype + ) + return (out, out1) if not prenorm else (out, out1, x) + + +@triton.autotune( + configs=triton_autotune_configs(), + key=["N", "HAS_RESIDUAL", "STORE_RESIDUAL_OUT", "IS_RMS_NORM", "HAS_BIAS"], +) +# @triton.heuristics({"HAS_BIAS": lambda args: args["B"] is not None}) +# @triton.heuristics({"HAS_RESIDUAL": lambda args: args["RESIDUAL"] is not None}) +@triton.heuristics({"HAS_X1": lambda args: args["X1"] is not None}) +@triton.heuristics({"HAS_W1": lambda args: args["W1"] is not None}) +@triton.heuristics({"HAS_B1": lambda args: args["B1"] is not None}) +@triton.jit +def _layer_norm_fwd_1pass_kernel( + X, # pointer to the input + Y, # pointer to the output + W, # pointer to the weights + B, # pointer to the biases + RESIDUAL, # pointer to the residual + X1, + W1, + B1, + Y1, + RESIDUAL_OUT, # pointer to the residual + ROWSCALE, + SEEDS, # Dropout seeds for each row + DROPOUT_MASK, + Mean, # pointer to the mean + Rstd, # pointer to the 1/std + stride_x_row, # how much to increase the pointer when moving by 1 row + stride_y_row, + stride_res_row, + stride_res_out_row, + stride_x1_row, + stride_y1_row, + M, # number of rows in X + N, # number of columns in X + eps, # epsilon to avoid division by zero + dropout_p, # Dropout probability + zero_centered_weight, # If true, add 1.0 to the weight + IS_RMS_NORM: tl.constexpr, + BLOCK_N: tl.constexpr, + HAS_RESIDUAL: tl.constexpr, + STORE_RESIDUAL_OUT: tl.constexpr, + HAS_BIAS: tl.constexpr, + HAS_DROPOUT: tl.constexpr, + STORE_DROPOUT_MASK: tl.constexpr, + HAS_ROWSCALE: tl.constexpr, + HAS_X1: tl.constexpr, + HAS_W1: tl.constexpr, + HAS_B1: tl.constexpr, +): + # Map the program id to the row of X and Y it should compute. + row = tl.program_id(0) + X += row * stride_x_row + Y += row * stride_y_row + if HAS_RESIDUAL: + RESIDUAL += row * stride_res_row + if STORE_RESIDUAL_OUT: + RESIDUAL_OUT += row * stride_res_out_row + if HAS_X1: + X1 += row * stride_x1_row + if HAS_W1: + Y1 += row * stride_y1_row + # Compute mean and variance + cols = tl.arange(0, BLOCK_N) + x = tl.load(X + cols, mask=cols < N, other=0.0).to(tl.float32) + if HAS_ROWSCALE: + rowscale = tl.load(ROWSCALE + row).to(tl.float32) + x *= rowscale + if HAS_DROPOUT: + # Compute dropout mask + # 7 rounds is good enough, and reduces register pressure + keep_mask = tl.rand(tl.load(SEEDS + row).to(tl.uint32), cols, n_rounds=7) > dropout_p + x = tl.where(keep_mask, x / (1.0 - dropout_p), 0.0) + if STORE_DROPOUT_MASK: + tl.store(DROPOUT_MASK + row * N + cols, keep_mask, mask=cols < N) + if HAS_X1: + x1 = tl.load(X1 + cols, mask=cols < N, other=0.0).to(tl.float32) + if HAS_ROWSCALE: + rowscale = tl.load(ROWSCALE + M + row).to(tl.float32) + x1 *= rowscale + if HAS_DROPOUT: + # Compute dropout mask + # 7 rounds is good enough, and reduces register pressure + keep_mask = ( + tl.rand(tl.load(SEEDS + M + row).to(tl.uint32), cols, n_rounds=7) > dropout_p + ) + x1 = tl.where(keep_mask, x1 / (1.0 - dropout_p), 0.0) + if STORE_DROPOUT_MASK: + tl.store(DROPOUT_MASK + (M + row) * N + cols, keep_mask, mask=cols < N) + x += x1 + if HAS_RESIDUAL: + residual = tl.load(RESIDUAL + cols, mask=cols < N, other=0.0).to(tl.float32) + x += residual + if STORE_RESIDUAL_OUT: + tl.store(RESIDUAL_OUT + cols, x, mask=cols < N) + if not IS_RMS_NORM: + mean = tl.sum(x, axis=0) / N + tl.store(Mean + row, mean) + xbar = tl.where(cols < N, x - mean, 0.0) + var = tl.sum(xbar * xbar, axis=0) / N + else: + xbar = tl.where(cols < N, x, 0.0) + var = tl.sum(xbar * xbar, axis=0) / N + rstd = 1 / tl.sqrt(var + eps) + tl.store(Rstd + row, rstd) + # Normalize and apply linear transformation + mask = cols < N + w = tl.load(W + cols, mask=mask).to(tl.float32) + if zero_centered_weight: + w += 1.0 + if HAS_BIAS: + b = tl.load(B + cols, mask=mask).to(tl.float32) + x_hat = (x - mean) * rstd if not IS_RMS_NORM else x * rstd + y = x_hat * w + b if HAS_BIAS else x_hat * w + # Write output + tl.store(Y + cols, y, mask=mask) + if HAS_W1: + w1 = tl.load(W1 + cols, mask=mask).to(tl.float32) + if zero_centered_weight: + w1 += 1.0 + if HAS_B1: + b1 = tl.load(B1 + cols, mask=mask).to(tl.float32) + y1 = x_hat * w1 + b1 if HAS_B1 else x_hat * w1 + tl.store(Y1 + cols, y1, mask=mask) + + +def _layer_norm_fwd( + x, + weight, + bias, + eps, + residual=None, + x1=None, + weight1=None, + bias1=None, + dropout_p=0.0, + rowscale=None, + out_dtype=None, + residual_dtype=None, + zero_centered_weight=False, + is_rms_norm=False, + return_dropout_mask=False, + out=None, + residual_out=None +): + if residual is not None: + residual_dtype = residual.dtype + M, N = x.shape + assert x.stride(-1) == 1 + if residual is not None: + assert residual.stride(-1) == 1 + assert residual.shape == (M, N) + assert weight.shape == (N,) + assert weight.stride(-1) == 1 + if bias is not None: + assert bias.stride(-1) == 1 + assert bias.shape == (N,) + if x1 is not None: + assert x1.shape == x.shape + assert rowscale is None + assert x1.stride(-1) == 1 + if weight1 is not None: + assert weight1.shape == (N,) + assert weight1.stride(-1) == 1 + if bias1 is not None: + assert bias1.shape == (N,) + assert bias1.stride(-1) == 1 + if rowscale is not None: + assert rowscale.is_contiguous() + assert rowscale.shape == (M,) + # allocate output + if out is None: + out = torch.empty_like(x, dtype=x.dtype if out_dtype is None else out_dtype) + else: + assert out.shape == x.shape + assert out.stride(-1) == 1 + if weight1 is not None: + y1 = torch.empty_like(out) + assert y1.stride(-1) == 1 + else: + y1 = None + if ( + residual is not None + or (residual_dtype is not None and residual_dtype != x.dtype) + or dropout_p > 0.0 + or rowscale is not None + or x1 is not None + ): + if residual_out is None: + residual_out = torch.empty( + M, N, device=x.device, dtype=residual_dtype if residual_dtype is not None else x.dtype + ) + else: + assert residual_out.shape == x.shape + assert residual_out.stride(-1) == 1 + else: + residual_out = None + mean = torch.empty((M,), dtype=torch.float32, device=x.device) if not is_rms_norm else None + rstd = torch.empty((M,), dtype=torch.float32, device=x.device) + if dropout_p > 0.0: + seeds = torch.randint( + 2**32, (M if x1 is None else 2 * M,), device=x.device, dtype=torch.int64 + ) + else: + seeds = None + if return_dropout_mask and dropout_p > 0.0: + dropout_mask = torch.empty(M if x1 is None else 2 * M, N, device=x.device, dtype=torch.bool) + else: + dropout_mask = None + # Less than 64KB per feature: enqueue fused kernel + MAX_FUSED_SIZE = 65536 // x.element_size() + BLOCK_N = min(MAX_FUSED_SIZE, triton.next_power_of_2(N)) + if N > BLOCK_N: + raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.") + with torch.cuda.device(x.device.index): + _layer_norm_fwd_1pass_kernel[(M,)]( + x, + out, + weight, + bias, + residual, + x1, + weight1, + bias1, + y1, + residual_out, + rowscale, + seeds, + dropout_mask, + mean, + rstd, + x.stride(0), + out.stride(0), + residual.stride(0) if residual is not None else 0, + residual_out.stride(0) if residual_out is not None else 0, + x1.stride(0) if x1 is not None else 0, + y1.stride(0) if y1 is not None else 0, + M, + N, + eps, + dropout_p, + zero_centered_weight, + is_rms_norm, + BLOCK_N, + residual is not None, + residual_out is not None, + bias is not None, + dropout_p > 0.0, + dropout_mask is not None, + rowscale is not None, + ) + # residual_out is None if residual is None and residual_dtype == input_dtype and dropout_p == 0.0 + if dropout_mask is not None and x1 is not None: + dropout_mask, dropout_mask1 = dropout_mask.tensor_split(2, dim=0) + else: + dropout_mask1 = None + return ( + out, + y1, + mean, + rstd, + residual_out if residual_out is not None else x, + seeds, + dropout_mask, + dropout_mask1, + ) + + +@triton.autotune( + configs=triton_autotune_configs(), + key=["N", "HAS_DRESIDUAL", "STORE_DRESIDUAL", "IS_RMS_NORM", "HAS_BIAS", "HAS_DROPOUT"], +) +# @triton.heuristics({"HAS_BIAS": lambda args: args["B"] is not None}) +# @triton.heuristics({"HAS_DRESIDUAL": lambda args: args["DRESIDUAL"] is not None}) +# @triton.heuristics({"STORE_DRESIDUAL": lambda args: args["DRESIDUAL_IN"] is not None}) +@triton.heuristics({"HAS_ROWSCALE": lambda args: args["ROWSCALE"] is not None}) +@triton.heuristics({"HAS_DY1": lambda args: args["DY1"] is not None}) +@triton.heuristics({"HAS_DX1": lambda args: args["DX1"] is not None}) +@triton.heuristics({"HAS_B1": lambda args: args["DB1"] is not None}) +@triton.heuristics({"RECOMPUTE_OUTPUT": lambda args: args["Y"] is not None}) +@triton.jit +def _layer_norm_bwd_kernel( + X, # pointer to the input + W, # pointer to the weights + B, # pointer to the biases + Y, # pointer to the output to be recomputed + DY, # pointer to the output gradient + DX, # pointer to the input gradient + DW, # pointer to the partial sum of weights gradient + DB, # pointer to the partial sum of biases gradient + DRESIDUAL, + W1, + DY1, + DX1, + DW1, + DB1, + DRESIDUAL_IN, + ROWSCALE, + SEEDS, + Mean, # pointer to the mean + Rstd, # pointer to the 1/std + stride_x_row, # how much to increase the pointer when moving by 1 row + stride_y_row, + stride_dy_row, + stride_dx_row, + stride_dres_row, + stride_dy1_row, + stride_dx1_row, + stride_dres_in_row, + M, # number of rows in X + N, # number of columns in X + eps, # epsilon to avoid division by zero + dropout_p, + zero_centered_weight, + rows_per_program, + IS_RMS_NORM: tl.constexpr, + BLOCK_N: tl.constexpr, + HAS_DRESIDUAL: tl.constexpr, + STORE_DRESIDUAL: tl.constexpr, + HAS_BIAS: tl.constexpr, + HAS_DROPOUT: tl.constexpr, + HAS_ROWSCALE: tl.constexpr, + HAS_DY1: tl.constexpr, + HAS_DX1: tl.constexpr, + HAS_B1: tl.constexpr, + RECOMPUTE_OUTPUT: tl.constexpr, +): + # Map the program id to the elements of X, DX, and DY it should compute. + row_block_id = tl.program_id(0) + row_start = row_block_id * rows_per_program + # Do not early exit if row_start >= M, because we need to write DW and DB + cols = tl.arange(0, BLOCK_N) + mask = cols < N + X += row_start * stride_x_row + if HAS_DRESIDUAL: + DRESIDUAL += row_start * stride_dres_row + if STORE_DRESIDUAL: + DRESIDUAL_IN += row_start * stride_dres_in_row + DY += row_start * stride_dy_row + DX += row_start * stride_dx_row + if HAS_DY1: + DY1 += row_start * stride_dy1_row + if HAS_DX1: + DX1 += row_start * stride_dx1_row + if RECOMPUTE_OUTPUT: + Y += row_start * stride_y_row + w = tl.load(W + cols, mask=mask).to(tl.float32) + if zero_centered_weight: + w += 1.0 + if RECOMPUTE_OUTPUT and HAS_BIAS: + b = tl.load(B + cols, mask=mask, other=0.0).to(tl.float32) + if HAS_DY1: + w1 = tl.load(W1 + cols, mask=mask).to(tl.float32) + if zero_centered_weight: + w1 += 1.0 + dw = tl.zeros((BLOCK_N,), dtype=tl.float32) + if HAS_BIAS: + db = tl.zeros((BLOCK_N,), dtype=tl.float32) + if HAS_DY1: + dw1 = tl.zeros((BLOCK_N,), dtype=tl.float32) + if HAS_B1: + db1 = tl.zeros((BLOCK_N,), dtype=tl.float32) + row_end = min((row_block_id + 1) * rows_per_program, M) + for row in range(row_start, row_end): + # Load data to SRAM + x = tl.load(X + cols, mask=mask, other=0).to(tl.float32) + dy = tl.load(DY + cols, mask=mask, other=0).to(tl.float32) + if HAS_DY1: + dy1 = tl.load(DY1 + cols, mask=mask, other=0).to(tl.float32) + if not IS_RMS_NORM: + mean = tl.load(Mean + row) + rstd = tl.load(Rstd + row) + # Compute dx + xhat = (x - mean) * rstd if not IS_RMS_NORM else x * rstd + xhat = tl.where(mask, xhat, 0.0) + if RECOMPUTE_OUTPUT: + y = xhat * w + b if HAS_BIAS else xhat * w + tl.store(Y + cols, y, mask=mask) + wdy = w * dy + dw += dy * xhat + if HAS_BIAS: + db += dy + if HAS_DY1: + wdy += w1 * dy1 + dw1 += dy1 * xhat + if HAS_B1: + db1 += dy1 + if not IS_RMS_NORM: + c1 = tl.sum(xhat * wdy, axis=0) / N + c2 = tl.sum(wdy, axis=0) / N + dx = (wdy - (xhat * c1 + c2)) * rstd + else: + c1 = tl.sum(xhat * wdy, axis=0) / N + dx = (wdy - xhat * c1) * rstd + if HAS_DRESIDUAL: + dres = tl.load(DRESIDUAL + cols, mask=mask, other=0).to(tl.float32) + dx += dres + # Write dx + if STORE_DRESIDUAL: + tl.store(DRESIDUAL_IN + cols, dx, mask=mask) + if HAS_DX1: + if HAS_DROPOUT: + keep_mask = ( + tl.rand(tl.load(SEEDS + M + row).to(tl.uint32), cols, n_rounds=7) > dropout_p + ) + dx1 = tl.where(keep_mask, dx / (1.0 - dropout_p), 0.0) + else: + dx1 = dx + tl.store(DX1 + cols, dx1, mask=mask) + if HAS_DROPOUT: + keep_mask = tl.rand(tl.load(SEEDS + row).to(tl.uint32), cols, n_rounds=7) > dropout_p + dx = tl.where(keep_mask, dx / (1.0 - dropout_p), 0.0) + if HAS_ROWSCALE: + rowscale = tl.load(ROWSCALE + row).to(tl.float32) + dx *= rowscale + tl.store(DX + cols, dx, mask=mask) + + X += stride_x_row + if HAS_DRESIDUAL: + DRESIDUAL += stride_dres_row + if STORE_DRESIDUAL: + DRESIDUAL_IN += stride_dres_in_row + if RECOMPUTE_OUTPUT: + Y += stride_y_row + DY += stride_dy_row + DX += stride_dx_row + if HAS_DY1: + DY1 += stride_dy1_row + if HAS_DX1: + DX1 += stride_dx1_row + tl.store(DW + row_block_id * N + cols, dw, mask=mask) + if HAS_BIAS: + tl.store(DB + row_block_id * N + cols, db, mask=mask) + if HAS_DY1: + tl.store(DW1 + row_block_id * N + cols, dw1, mask=mask) + if HAS_B1: + tl.store(DB1 + row_block_id * N + cols, db1, mask=mask) + + +def _layer_norm_bwd( + dy, + x, + weight, + bias, + eps, + mean, + rstd, + dresidual=None, + dy1=None, + weight1=None, + bias1=None, + seeds=None, + dropout_p=0.0, + rowscale=None, + has_residual=False, + has_x1=False, + zero_centered_weight=False, + is_rms_norm=False, + x_dtype=None, + recompute_output=False, +): + M, N = x.shape + assert x.stride(-1) == 1 + assert dy.stride(-1) == 1 + assert dy.shape == (M, N) + if dresidual is not None: + assert dresidual.stride(-1) == 1 + assert dresidual.shape == (M, N) + assert weight.shape == (N,) + assert weight.stride(-1) == 1 + if bias is not None: + assert bias.stride(-1) == 1 + assert bias.shape == (N,) + if dy1 is not None: + assert weight1 is not None + assert dy1.shape == dy.shape + assert dy1.stride(-1) == 1 + if weight1 is not None: + assert weight1.shape == (N,) + assert weight1.stride(-1) == 1 + if bias1 is not None: + assert bias1.shape == (N,) + assert bias1.stride(-1) == 1 + if seeds is not None: + assert seeds.is_contiguous() + assert seeds.shape == (M if not has_x1 else M * 2,) + if rowscale is not None: + assert rowscale.is_contiguous() + assert rowscale.shape == (M,) + # allocate output + dx = ( + torch.empty_like(x) + if x_dtype is None + else torch.empty(M, N, dtype=x_dtype, device=x.device) + ) + dresidual_in = ( + torch.empty_like(x) + if has_residual + and (dx.dtype != x.dtype or dropout_p > 0.0 or rowscale is not None or has_x1) + else None + ) + dx1 = torch.empty_like(dx) if (has_x1 and dropout_p > 0.0) else None + y = torch.empty(M, N, dtype=dy.dtype, device=dy.device) if recompute_output else None + if recompute_output: + assert weight1 is None, "recompute_output is not supported with parallel LayerNorm" + + # Less than 64KB per feature: enqueue fused kernel + MAX_FUSED_SIZE = 65536 // x.element_size() + BLOCK_N = min(MAX_FUSED_SIZE, triton.next_power_of_2(N)) + if N > BLOCK_N: + raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.") + # Increasing the multiple (e.g. 8) will allow more thread blocks to be launched and hide the + # latency of the gmem reads/writes, but will increase the time of summing up dw / db. + sm_count = torch.cuda.get_device_properties(x.device).multi_processor_count * 8 + _dw = torch.empty((sm_count, N), dtype=torch.float32, device=weight.device) + _db = ( + torch.empty((sm_count, N), dtype=torch.float32, device=bias.device) + if bias is not None + else None + ) + _dw1 = torch.empty_like(_dw) if weight1 is not None else None + _db1 = torch.empty_like(_db) if bias1 is not None else None + rows_per_program = math.ceil(M / sm_count) + grid = (sm_count,) + with torch.cuda.device(x.device.index): + _layer_norm_bwd_kernel[grid]( + x, + weight, + bias, + y, + dy, + dx, + _dw, + _db, + dresidual, + weight1, + dy1, + dx1, + _dw1, + _db1, + dresidual_in, + rowscale, + seeds, + mean, + rstd, + x.stride(0), + 0 if not recompute_output else y.stride(0), + dy.stride(0), + dx.stride(0), + dresidual.stride(0) if dresidual is not None else 0, + dy1.stride(0) if dy1 is not None else 0, + dx1.stride(0) if dx1 is not None else 0, + dresidual_in.stride(0) if dresidual_in is not None else 0, + M, + N, + eps, + dropout_p, + zero_centered_weight, + rows_per_program, + is_rms_norm, + BLOCK_N, + dresidual is not None, + dresidual_in is not None, + bias is not None, + dropout_p > 0.0, + ) + dw = _dw.sum(0).to(weight.dtype) + db = _db.sum(0).to(bias.dtype) if bias is not None else None + dw1 = _dw1.sum(0).to(weight1.dtype) if weight1 is not None else None + db1 = _db1.sum(0).to(bias1.dtype) if bias1 is not None else None + # Don't need to compute dresidual_in separately in this case + if has_residual and dx.dtype == x.dtype and dropout_p == 0.0 and rowscale is None: + dresidual_in = dx + if has_x1 and dropout_p == 0.0: + dx1 = dx + return ( + (dx, dw, db, dresidual_in, dx1, dw1, db1) + if not recompute_output + else (dx, dw, db, dresidual_in, dx1, dw1, db1, y) + ) + + +class LayerNormFn(torch.autograd.Function): + @staticmethod + def forward( + ctx, + x, + weight, + bias, + residual=None, + x1=None, + weight1=None, + bias1=None, + eps=1e-6, + dropout_p=0.0, + rowscale=None, + prenorm=False, + residual_in_fp32=False, + zero_centered_weight=False, + is_rms_norm=False, + return_dropout_mask=False, + out=None, + residual_out=None + ): + x_shape_og = x.shape + # Check for zero sequence length + if x.numel() == 0: + ctx.zero_seq_length = True + # Only save minimal required tensors for backward + # ctx.save_for_backward(weight, bias, weight1, bias1) + ctx.x_shape_og = x_shape_og + ctx.weight_shape = weight.shape + ctx.weight_dtype = weight.dtype + ctx.weight_device = weight.device + + ctx.has_bias = bias is not None + ctx.bias_shape = bias.shape if bias is not None else None + ctx.bias_dtype = bias.dtype if bias is not None else None + ctx.bias_device = bias.device if bias is not None else None + + ctx.has_weight1 = weight1 is not None + ctx.weight1_shape = weight1.shape if weight1 is not None else None + ctx.weight1_dtype = weight1.dtype if weight1 is not None else None + ctx.weight1_device = weight1.device if weight1 is not None else None + + ctx.has_bias1 = bias1 is not None + ctx.bias1_shape = bias1.shape if bias1 is not None else None + ctx.bias1_dtype = bias1.dtype if bias1 is not None else None + ctx.bias1_device = bias1.device if bias1 is not None else None + + ctx.has_residual = residual is not None + ctx.has_x1 = x1 is not None + ctx.dropout_p = dropout_p + + # Handle output tensors with correct dtype + y = x # Preserve input tensor properties + y1 = torch.empty_like(x) if x1 is not None else None + + # Only create residual_out if prenorm is True + residual_out = torch.empty(x.shape, + dtype=torch.float32 if residual_in_fp32 else x.dtype, + device=x.device) if prenorm else None + + # Handle dropout masks + dropout_mask = None + dropout_mask1 = None + if return_dropout_mask: + dropout_mask = torch.empty_like(x, dtype=torch.uint8) + if x1 is not None: + dropout_mask1 = torch.empty_like(x, dtype=torch.uint8) + + # Return based on configuration + if not return_dropout_mask: + if weight1 is None: + return y if not prenorm else (y, residual_out) + else: + return (y, y1) if not prenorm else (y, y1, residual_out) + else: + if weight1 is None: + return ((y, dropout_mask, dropout_mask1) if not prenorm + else (y, residual_out, dropout_mask, dropout_mask1)) + else: + return ((y, y1, dropout_mask, dropout_mask1) if not prenorm + else (y, y1, residual_out, dropout_mask, dropout_mask1)) + + ctx.zero_seq_length = False + # reshape input data into 2D tensor + x = x.reshape(-1, x.shape[-1]) + if x.stride(-1) != 1: + x = x.contiguous() + if residual is not None: + assert residual.shape == x_shape_og + residual = residual.reshape(-1, residual.shape[-1]) + if residual.stride(-1) != 1: + residual = residual.contiguous() + if x1 is not None: + assert x1.shape == x_shape_og + assert rowscale is None, "rowscale is not supported with parallel LayerNorm" + x1 = x1.reshape(-1, x1.shape[-1]) + if x1.stride(-1) != 1: + x1 = x1.contiguous() + weight = weight.contiguous() + if bias is not None: + bias = bias.contiguous() + if weight1 is not None: + weight1 = weight1.contiguous() + if bias1 is not None: + bias1 = bias1.contiguous() + if rowscale is not None: + rowscale = rowscale.reshape(-1).contiguous() + residual_dtype = ( + residual.dtype + if residual is not None + else (torch.float32 if residual_in_fp32 else None) + ) + if out is not None: + out = out.reshape(-1, out.shape[-1]) + if residual_out is not None: + residual_out = residual_out.reshape(-1, residual_out.shape[-1]) + y, y1, mean, rstd, residual_out, seeds, dropout_mask, dropout_mask1 = _layer_norm_fwd( + x, + weight, + bias, + eps, + residual, + x1, + weight1, + bias1, + dropout_p=dropout_p, + rowscale=rowscale, + residual_dtype=residual_dtype, + zero_centered_weight=zero_centered_weight, + is_rms_norm=is_rms_norm, + return_dropout_mask=return_dropout_mask, + out=out, + residual_out=residual_out + ) + ctx.save_for_backward( + residual_out, weight, bias, weight1, bias1, rowscale, seeds, mean, rstd + ) + ctx.x_shape_og = x_shape_og + ctx.eps = eps + ctx.dropout_p = dropout_p + ctx.is_rms_norm = is_rms_norm + ctx.has_residual = residual is not None + ctx.has_x1 = x1 is not None + ctx.prenorm = prenorm + ctx.x_dtype = x.dtype + ctx.zero_centered_weight = zero_centered_weight + y = y.reshape(x_shape_og) + y1 = y1.reshape(x_shape_og) if y1 is not None else None + residual_out = residual_out.reshape(x_shape_og) if residual_out is not None else None + dropout_mask = dropout_mask.reshape(x_shape_og) if dropout_mask is not None else None + dropout_mask1 = dropout_mask1.reshape(x_shape_og) if dropout_mask1 is not None else None + if not return_dropout_mask: + if weight1 is None: + return y if not prenorm else (y, residual_out) + else: + return (y, y1) if not prenorm else (y, y1, residual_out) + else: + if weight1 is None: + return ( + (y, dropout_mask, dropout_mask1) + if not prenorm + else (y, residual_out, dropout_mask, dropout_mask1) + ) + else: + return ( + (y, y1, dropout_mask, dropout_mask1) + if not prenorm + else (y, y1, residual_out, dropout_mask, dropout_mask1) + ) + + @staticmethod + def backward(ctx, dy, *args): + if ctx.zero_seq_length: + return ( + torch.zeros(ctx.x_shape_og, dtype=dy.dtype, device=dy.device), + torch.zeros(ctx.weight_shape, dtype=ctx.weight_dtype, device=ctx.weight_device), + torch.zeros(ctx.bias_shape, dtype=ctx.bias_dtype, device=ctx.bias_device) if ctx.has_bias else None, + torch.zeros(ctx.x_shape_og, dtype=dy.dtype, device=dy.device) if ctx.has_residual else None, + torch.zeros(ctx.x_shape_og, dtype=dy.dtype, device=dy.device) if ctx.has_x1 and ctx.dropout_p > 0.0 else None, + torch.zeros(ctx.weight1_shape, dtype=ctx.weight1_dtype, device=ctx.weight1_device) if ctx.has_weight1 else None, + torch.zeros(ctx.bias1_shape, dtype=ctx.bias1_dtype, device=ctx.bias1_device) if ctx.has_bias1 else None, + None, + None, + None, + None, + None, + None, + None, + None, + None, + None, + ) + + x, weight, bias, weight1, bias1, rowscale, seeds, mean, rstd = ctx.saved_tensors + dy = dy.reshape(-1, dy.shape[-1]) + if dy.stride(-1) != 1: + dy = dy.contiguous() + assert dy.shape == x.shape + if weight1 is not None: + dy1, args = args[0], args[1:] + dy1 = dy1.reshape(-1, dy1.shape[-1]) + if dy1.stride(-1) != 1: + dy1 = dy1.contiguous() + assert dy1.shape == x.shape + else: + dy1 = None + if ctx.prenorm: + dresidual = args[0] + dresidual = dresidual.reshape(-1, dresidual.shape[-1]) + if dresidual.stride(-1) != 1: + dresidual = dresidual.contiguous() + assert dresidual.shape == x.shape + else: + dresidual = None + + dx, dw, db, dresidual_in, dx1, dw1, db1 = _layer_norm_bwd( + dy, + x, + weight, + bias, + ctx.eps, + mean, + rstd, + dresidual, + dy1, + weight1, + bias1, + seeds, + ctx.dropout_p, + rowscale, + ctx.has_residual, + ctx.has_x1, + ctx.zero_centered_weight, + ctx.is_rms_norm, + x_dtype=ctx.x_dtype, + ) + return ( + dx.reshape(ctx.x_shape_og), + dw, + db, + dresidual_in.reshape(ctx.x_shape_og) if ctx.has_residual else None, + dx1.reshape(ctx.x_shape_og) if dx1 is not None else None, + dw1, + db1, + None, + None, + None, + None, + None, + None, + None, + None, + None, + None, + ) + + +def layer_norm_fn( + x, + weight, + bias, + residual=None, + x1=None, + weight1=None, + bias1=None, + eps=1e-6, + dropout_p=0.0, + rowscale=None, + prenorm=False, + residual_in_fp32=False, + zero_centered_weight=False, + is_rms_norm=False, + return_dropout_mask=False, + out=None, + residual_out=None +): + return LayerNormFn.apply( + x, + weight, + bias, + residual, + x1, + weight1, + bias1, + eps, + dropout_p, + rowscale, + prenorm, + residual_in_fp32, + zero_centered_weight, + is_rms_norm, + return_dropout_mask, + out, + residual_out + ) + + +def rms_norm_fn( + x, + weight, + bias, + residual=None, + x1=None, + weight1=None, + bias1=None, + eps=1e-6, + dropout_p=0.0, + rowscale=None, + prenorm=False, + residual_in_fp32=False, + zero_centered_weight=False, + return_dropout_mask=False, + out=None, + residual_out=None +): + return LayerNormFn.apply( + x, + weight, + bias, + residual, + x1, + weight1, + bias1, + eps, + dropout_p, + rowscale, + prenorm, + residual_in_fp32, + zero_centered_weight, + True, + return_dropout_mask, + out, + residual_out + ) + + +class RMSNorm(torch.nn.Module): + + def __init__(self, hidden_size, eps=1e-5, dropout_p=0.0, zero_centered_weight=False, + device=None, dtype=None): + factory_kwargs = {"device": device, "dtype": dtype} + super().__init__() + self.eps = eps + if dropout_p > 0.0: + self.drop = torch.nn.Dropout(dropout_p) + else: + self.drop = None + self.zero_centered_weight = zero_centered_weight + self.weight = torch.nn.Parameter(torch.empty(hidden_size, **factory_kwargs)) + self.register_parameter("bias", None) + self.reset_parameters() + + def reset_parameters(self): + if not self.zero_centered_weight: + torch.nn.init.ones_(self.weight) + else: + torch.nn.init.zeros_(self.weight) + + def forward(self, x, residual=None, prenorm=False, residual_in_fp32=False): + return rms_norm_fn( + x, + self.weight, + self.bias, + residual=residual, + eps=self.eps, + dropout_p=self.drop.p if self.drop is not None and self.training else 0.0, + prenorm=prenorm, + residual_in_fp32=residual_in_fp32, + zero_centered_weight=self.zero_centered_weight, + ) + + +class LayerNormLinearFn(torch.autograd.Function): + @staticmethod + @custom_fwd + def forward( + ctx, + x, + norm_weight, + norm_bias, + linear_weight, + linear_bias, + residual=None, + eps=1e-6, + prenorm=False, + residual_in_fp32=False, + is_rms_norm=False, + ): + x_shape_og = x.shape + # reshape input data into 2D tensor + x = x.reshape(-1, x.shape[-1]) + if x.stride(-1) != 1: + x = x.contiguous() + if residual is not None: + assert residual.shape == x_shape_og + residual = residual.reshape(-1, residual.shape[-1]) + if residual.stride(-1) != 1: + residual = residual.contiguous() + norm_weight = norm_weight.contiguous() + if norm_bias is not None: + norm_bias = norm_bias.contiguous() + residual_dtype = ( + residual.dtype + if residual is not None + else (torch.float32 if residual_in_fp32 else None) + ) + y, _, mean, rstd, residual_out, *rest = _layer_norm_fwd( + x, + norm_weight, + norm_bias, + eps, + residual, + out_dtype=None if not torch.is_autocast_enabled() else torch.get_autocast_dtype("cuda"), + residual_dtype=residual_dtype, + is_rms_norm=is_rms_norm, + ) + y = y.reshape(x_shape_og) + dtype = torch.get_autocast_dtype("cuda") if torch.is_autocast_enabled() else y.dtype + linear_weight = linear_weight.to(dtype) + linear_bias = linear_bias.to(dtype) if linear_bias is not None else None + out = F.linear(y.to(linear_weight.dtype), linear_weight, linear_bias) + # We don't store y, will be recomputed in the backward pass to save memory + ctx.save_for_backward(residual_out, norm_weight, norm_bias, linear_weight, mean, rstd) + ctx.x_shape_og = x_shape_og + ctx.eps = eps + ctx.is_rms_norm = is_rms_norm + ctx.has_residual = residual is not None + ctx.prenorm = prenorm + ctx.x_dtype = x.dtype + ctx.linear_bias_is_none = linear_bias is None + return out if not prenorm else (out, residual_out.reshape(x_shape_og)) + + @staticmethod + @custom_bwd + def backward(ctx, dout, *args): + x, norm_weight, norm_bias, linear_weight, mean, rstd = ctx.saved_tensors + dout = dout.reshape(-1, dout.shape[-1]) + dy = F.linear(dout, linear_weight.t()) + dlinear_bias = None if ctx.linear_bias_is_none else dout.sum(0) + if dy.stride(-1) != 1: + dy = dy.contiguous() + assert dy.shape == x.shape + if ctx.prenorm: + dresidual = args[0] + dresidual = dresidual.reshape(-1, dresidual.shape[-1]) + if dresidual.stride(-1) != 1: + dresidual = dresidual.contiguous() + assert dresidual.shape == x.shape + else: + dresidual = None + dx, dnorm_weight, dnorm_bias, dresidual_in, _, _, _, y = _layer_norm_bwd( + dy, + x, + norm_weight, + norm_bias, + ctx.eps, + mean, + rstd, + dresidual=dresidual, + has_residual=ctx.has_residual, + is_rms_norm=ctx.is_rms_norm, + x_dtype=ctx.x_dtype, + recompute_output=True, + ) + dlinear_weight = torch.einsum("bo,bi->oi", dout, y) + return ( + dx.reshape(ctx.x_shape_og), + dnorm_weight, + dnorm_bias, + dlinear_weight, + dlinear_bias, + dresidual_in.reshape(ctx.x_shape_og) if ctx.has_residual else None, + None, + None, + None, + None, + ) + + +def layer_norm_linear_fn( + x, + norm_weight, + norm_bias, + linear_weight, + linear_bias, + residual=None, + eps=1e-6, + prenorm=False, + residual_in_fp32=False, + is_rms_norm=False, +): + return LayerNormLinearFn.apply( + x, + norm_weight, + norm_bias, + linear_weight, + linear_bias, + residual, + eps, + prenorm, + residual_in_fp32, + is_rms_norm, + ) diff --git a/examples/OmniGen2-RL/omnigen2/optim/__init__.py b/examples/OmniGen2-RL/omnigen2/optim/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/examples/OmniGen2-RL/omnigen2/optim/scheduler/__init__.py b/examples/OmniGen2-RL/omnigen2/optim/scheduler/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/examples/OmniGen2-RL/omnigen2/optim/scheduler/cosine_lr.py b/examples/OmniGen2-RL/omnigen2/optim/scheduler/cosine_lr.py new file mode 100644 index 0000000000000000000000000000000000000000..419f53475fd345e7c923e7287f3012a64eb37a63 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/optim/scheduler/cosine_lr.py @@ -0,0 +1,118 @@ +""" Cosine Scheduler + +Cosine LR schedule with warmup, cycle/restarts, noise, k-decay. + +Hacked together by / Copyright 2021 Ross Wightman +""" +import logging +import math +import torch +from typing import List + +from .scheduler import Scheduler + + +_logger = logging.getLogger(__name__) + + +class CosineLRScheduler(Scheduler): + """ + Cosine decay with restarts. + This is described in the paper https://arxiv.org/abs/1608.03983. + + Inspiration from + https://github.com/allenai/allennlp/blob/master/allennlp/training/learning_rate_schedulers/cosine.py + + k-decay option based on `k-decay: A New Method For Learning Rate Schedule` - https://arxiv.org/abs/2004.05909 + """ + + def __init__( + self, + optimizer: torch.optim.Optimizer, + t_initial: int, + lr_min: float = 0., + cycle_mul: float = 1., + cycle_decay: float = 1., + cycle_limit: int = 1, + warmup_t=0, + warmup_lr_init=0, + warmup_prefix=False, + t_in_epochs=True, + noise_range_t=None, + noise_pct=0.67, + noise_std=1.0, + noise_seed=42, + k_decay=1.0, + initialize=True, + ) -> None: + super().__init__( + optimizer, + param_group_field="lr", + t_in_epochs=t_in_epochs, + noise_range_t=noise_range_t, + noise_pct=noise_pct, + noise_std=noise_std, + noise_seed=noise_seed, + initialize=initialize, + ) + + assert t_initial > 0 + assert lr_min >= 0 + if t_initial == 1 and cycle_mul == 1 and cycle_decay == 1: + _logger.warning( + "Cosine annealing scheduler will have no effect on the learning " + "rate since t_initial = t_mul = eta_mul = 1.") + self.t_initial = t_initial + self.lr_min = lr_min + self.cycle_mul = cycle_mul + self.cycle_decay = cycle_decay + self.cycle_limit = cycle_limit + self.warmup_t = warmup_t + self.warmup_lr_init = warmup_lr_init + self.warmup_prefix = warmup_prefix + self.k_decay = k_decay + if self.warmup_t: + self.warmup_steps = [(v - warmup_lr_init) / self.warmup_t for v in self.base_values] + super().update_groups(self.warmup_lr_init) + else: + self.warmup_steps = [1 for _ in self.base_values] + + self._step_count = 0 # no use + + def _get_lr(self, t: int) -> List[float]: + + if t < self.warmup_t: + lrs = [self.warmup_lr_init + t * s for s in self.warmup_steps] + else: + if self.warmup_prefix: + t = t - self.warmup_t + + if self.cycle_mul != 1: + i = math.floor(math.log(1 - t / self.t_initial * (1 - self.cycle_mul), self.cycle_mul)) + t_i = self.cycle_mul ** i * self.t_initial + t_curr = t - (1 - self.cycle_mul ** i) / (1 - self.cycle_mul) * self.t_initial + else: + i = t // self.t_initial + t_i = self.t_initial + t_curr = t - (self.t_initial * i) + + gamma = self.cycle_decay ** i + lr_max_values = [v * gamma for v in self.base_values] + k = self.k_decay + + if i < self.cycle_limit: + lrs = [ + self.lr_min + 0.5 * (lr_max - self.lr_min) * (1 + math.cos(math.pi * t_curr ** k / t_i ** k)) + for lr_max in lr_max_values + ] + else: + lrs = [self.lr_min for _ in self.base_values] + + return lrs + + def get_cycle_length(self, cycles=0): + cycles = max(1, cycles or self.cycle_limit) + if self.cycle_mul == 1.0: + return self.t_initial * cycles + else: + return int(math.floor(-self.t_initial * (self.cycle_mul ** cycles - 1) / (1 - self.cycle_mul))) diff --git a/examples/OmniGen2-RL/omnigen2/optim/scheduler/scheduler.py b/examples/OmniGen2-RL/omnigen2/optim/scheduler/scheduler.py new file mode 100644 index 0000000000000000000000000000000000000000..359d83cba58f17c220c98387e50ca0f287bdbeb7 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/optim/scheduler/scheduler.py @@ -0,0 +1,131 @@ +import abc +from abc import ABC +from typing import Any, Dict, List, Optional + +import torch + + +class Scheduler(ABC): + """ Parameter Scheduler Base Class + A scheduler base class that can be used to schedule any optimizer parameter groups. + + Unlike the builtin PyTorch schedulers, this is intended to be consistently called + * At the END of each epoch, before incrementing the epoch count, to calculate next epoch's value + * At the END of each optimizer update, after incrementing the update count, to calculate next update's value + + The schedulers built on this should try to remain as stateless as possible (for simplicity). + + This family of schedulers is attempting to avoid the confusion of the meaning of 'last_epoch' + and -1 values for special behaviour. All epoch and update counts must be tracked in the training + code and explicitly passed in to the schedulers on the corresponding step or step_update call. + + Based on ideas from: + * https://github.com/pytorch/fairseq/tree/master/fairseq/optim/lr_scheduler + * https://github.com/allenai/allennlp/tree/master/allennlp/training/learning_rate_schedulers + """ + + def __init__( + self, + optimizer: torch.optim.Optimizer, + param_group_field: str, + t_in_epochs: bool = True, + noise_range_t=None, + noise_type='normal', + noise_pct=0.67, + noise_std=1.0, + noise_seed=None, + initialize: bool = True, + ) -> None: + self.optimizer = optimizer + self.param_group_field = param_group_field + self._initial_param_group_field = f"initial_{param_group_field}" + if initialize: + for i, group in enumerate(self.optimizer.param_groups): + if param_group_field not in group: + raise KeyError(f"{param_group_field} missing from param_groups[{i}]") + group.setdefault(self._initial_param_group_field, group[param_group_field]) + else: + for i, group in enumerate(self.optimizer.param_groups): + if self._initial_param_group_field not in group: + raise KeyError(f"{self._initial_param_group_field} missing from param_groups[{i}]") + self.base_values = [group[self._initial_param_group_field] for group in self.optimizer.param_groups] + self.metric = None # any point to having this for all? + self.t_in_epochs = t_in_epochs + self.noise_range_t = noise_range_t + self.noise_pct = noise_pct + self.noise_type = noise_type + self.noise_std = noise_std + self.noise_seed = noise_seed if noise_seed is not None else 42 + self.update_groups(self.base_values) + + def state_dict(self) -> Dict[str, Any]: + return {key: value for key, value in self.__dict__.items() if key != 'optimizer'} + + def load_state_dict(self, state_dict: Dict[str, Any]) -> None: + self.__dict__.update(state_dict) + + def get_last_lr(self): + """ Return last computed learning rate by current scheduler. + """ + return self._last_lr + + @abc.abstractmethod + def _get_lr(self, t: int) -> List[float]: + pass + + def _get_values(self, t: int, on_epoch: bool = True) -> Optional[List[float]]: + return self._get_lr(t) + + def step(self, epoch: int, metric: float = None) -> None: + self.metric = metric + values = self._get_values(epoch, on_epoch=True) + if values is not None: + values = self._add_noise(values, epoch) + self.update_groups(values) + + # def step_update(self, num_updates: int, metric: float = None): + # self.metric = metric + # values = self._get_values(num_updates, on_epoch=False) + # if values is not None: + # values = self._add_noise(values, num_updates) + # self.update_groups(values) + + def update_groups(self, values): + if not isinstance(values, (list, tuple)): + values = [values] * len(self.optimizer.param_groups) + for param_group, value in zip(self.optimizer.param_groups, values): + if 'lr_scale' in param_group: + param_group[self.param_group_field] = value * param_group['lr_scale'] + else: + param_group[self.param_group_field] = value + + self._last_lr = [group[self.param_group_field] for group in self.optimizer.param_groups] + + def _add_noise(self, lrs, t): + if self._is_apply_noise(t): + noise = self._calculate_noise(t) + lrs = [v + v * noise for v in lrs] + return lrs + + def _is_apply_noise(self, t) -> bool: + """Return True if scheduler in noise range.""" + apply_noise = False + if self.noise_range_t is not None: + if isinstance(self.noise_range_t, (list, tuple)): + apply_noise = self.noise_range_t[0] <= t < self.noise_range_t[1] + else: + apply_noise = t >= self.noise_range_t + return apply_noise + + def _calculate_noise(self, t) -> float: + g = torch.Generator() + g.manual_seed(self.noise_seed + t) + if self.noise_type == 'normal': + while True: + # resample if noise out of percent limit, brute force but shouldn't spin much + noise = torch.randn(1, generator=g).item() + if abs(noise) < self.noise_pct: + return noise + else: + noise = 2 * (torch.rand(1, generator=g).item() - 0.5) * self.noise_pct + return noise diff --git a/examples/OmniGen2-RL/omnigen2/optim/scheduler/step_lr.py b/examples/OmniGen2-RL/omnigen2/optim/scheduler/step_lr.py new file mode 100644 index 0000000000000000000000000000000000000000..5d2516526f0aa335cc16e7de9a43bc04b0c13d58 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/optim/scheduler/step_lr.py @@ -0,0 +1,63 @@ +""" Step Scheduler + +Basic step LR schedule with warmup, noise. + +Hacked together by / Copyright 2020 Ross Wightman +""" +import math +import torch +from typing import List + + +from .scheduler import Scheduler + + +class StepLRScheduler(Scheduler): + """ + """ + + def __init__( + self, + optimizer: torch.optim.Optimizer, + decay_t: float, + decay_rate: float = 1., + warmup_t=0, + warmup_lr_init=0, + warmup_prefix=True, + t_in_epochs=True, + noise_range_t=None, + noise_pct=0.67, + noise_std=1.0, + noise_seed=42, + initialize=True, + ) -> None: + super().__init__( + optimizer, + param_group_field="lr", + t_in_epochs=t_in_epochs, + noise_range_t=noise_range_t, + noise_pct=noise_pct, + noise_std=noise_std, + noise_seed=noise_seed, + initialize=initialize, + ) + + self.decay_t = decay_t + self.decay_rate = decay_rate + self.warmup_t = warmup_t + self.warmup_lr_init = warmup_lr_init + self.warmup_prefix = warmup_prefix + if self.warmup_t: + self.warmup_steps = [(v - warmup_lr_init) / self.warmup_t for v in self.base_values] + super().update_groups(self.warmup_lr_init) + else: + self.warmup_steps = [1 for _ in self.base_values] + + def _get_lr(self, t: int) -> List[float]: + if t < self.warmup_t: + lrs = [self.warmup_lr_init + t * s for s in self.warmup_steps] + else: + if self.warmup_prefix: + t = t - self.warmup_t + lrs = [v * (self.decay_rate ** (t // self.decay_t)) for v in self.base_values] + return lrs \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/pipelines/__init__.py b/examples/OmniGen2-RL/omnigen2/pipelines/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/examples/OmniGen2-RL/omnigen2/pipelines/image_processor.py b/examples/OmniGen2-RL/omnigen2/pipelines/image_processor.py new file mode 100644 index 0000000000000000000000000000000000000000..b1d699ca57091e6691393217c83eb1534c64a46c --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/pipelines/image_processor.py @@ -0,0 +1,267 @@ +# Copyright 2024 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import math +import warnings +from typing import List, Optional, Tuple, Union + +import numpy as np +import PIL.Image +import torch + +from diffusers.image_processor import PipelineImageInput, VaeImageProcessor, is_valid_image_imagelist +from diffusers.configuration_utils import register_to_config + +class OmniGen2ImageProcessor(VaeImageProcessor): + """ + Image processor for PixArt image resize and crop. + + Args: + do_resize (`bool`, *optional*, defaults to `True`): + Whether to downscale the image's (height, width) dimensions to multiples of `vae_scale_factor`. Can accept + `height` and `width` arguments from [`image_processor.VaeImageProcessor.preprocess`] method. + vae_scale_factor (`int`, *optional*, defaults to `8`): + VAE scale factor. If `do_resize` is `True`, the image is automatically resized to multiples of this factor. + resample (`str`, *optional*, defaults to `lanczos`): + Resampling filter to use when resizing the image. + do_normalize (`bool`, *optional*, defaults to `True`): + Whether to normalize the image to [-1,1]. + do_binarize (`bool`, *optional*, defaults to `False`): + Whether to binarize the image to 0/1. + do_convert_rgb (`bool`, *optional*, defaults to be `False`): + Whether to convert the images to RGB format. + do_convert_grayscale (`bool`, *optional*, defaults to be `False`): + Whether to convert the images to grayscale format. + """ + + @register_to_config + def __init__( + self, + do_resize: bool = True, + vae_scale_factor: int = 16, + resample: str = "lanczos", + max_pixels: Optional[int] = None, + max_side_length: Optional[int] = None, + do_normalize: bool = True, + do_binarize: bool = False, + do_convert_grayscale: bool = False, + ): + super().__init__( + do_resize=do_resize, + vae_scale_factor=vae_scale_factor, + resample=resample, + do_normalize=do_normalize, + do_binarize=do_binarize, + do_convert_grayscale=do_convert_grayscale, + ) + + self.max_pixels = max_pixels + self.max_side_length = max_side_length + + def get_new_height_width( + self, + image: Union[PIL.Image.Image, np.ndarray, torch.Tensor], + height: Optional[int] = None, + width: Optional[int] = None, + max_pixels: Optional[int] = None, + max_side_length: Optional[int] = None, + ) -> Tuple[int, int]: + r""" + Returns the height and width of the image, downscaled to the next integer multiple of `vae_scale_factor`. + + Args: + image (`Union[PIL.Image.Image, np.ndarray, torch.Tensor]`): + The image input, which can be a PIL image, NumPy array, or PyTorch tensor. If it is a NumPy array, it + should have shape `[batch, height, width]` or `[batch, height, width, channels]`. If it is a PyTorch + tensor, it should have shape `[batch, channels, height, width]`. + height (`Optional[int]`, *optional*, defaults to `None`): + The height of the preprocessed image. If `None`, the height of the `image` input will be used. + width (`Optional[int]`, *optional*, defaults to `None`): + The width of the preprocessed image. If `None`, the width of the `image` input will be used. + + Returns: + `Tuple[int, int]`: + A tuple containing the height and width, both resized to the nearest integer multiple of + `vae_scale_factor`. + """ + + if height is None: + if isinstance(image, PIL.Image.Image): + height = image.height + elif isinstance(image, torch.Tensor): + height = image.shape[2] + else: + height = image.shape[1] + + if width is None: + if isinstance(image, PIL.Image.Image): + width = image.width + elif isinstance(image, torch.Tensor): + width = image.shape[3] + else: + width = image.shape[2] + + if max_side_length is None: + max_side_length = self.max_side_length + + if max_pixels is None: + max_pixels = self.max_pixels + + ratio = 1.0 + if max_side_length is not None: + if height > width: + max_side_length_ratio = max_side_length / height + else: + max_side_length_ratio = max_side_length / width + + cur_pixels = height * width + max_pixels_ratio = (max_pixels / cur_pixels) ** 0.5 + ratio = min(max_pixels_ratio, max_side_length_ratio, 1.0) # do not upscale input image + + new_height, new_width = int(height * ratio) // self.config.vae_scale_factor * self.config.vae_scale_factor, int(width * ratio) // self.config.vae_scale_factor * self.config.vae_scale_factor + return new_height, new_width + + def preprocess( + self, + image: PipelineImageInput, + height: Optional[int] = None, + width: Optional[int] = None, + max_pixels: Optional[int] = None, + max_side_length: Optional[int] = None, + resize_mode: str = "default", # "default", "fill", "crop" + crops_coords: Optional[Tuple[int, int, int, int]] = None, + do_normalize: Optional[bool] = None, + ) -> torch.Tensor: + """ + Preprocess the image input. + + Args: + image (`PipelineImageInput`): + The image input, accepted formats are PIL images, NumPy arrays, PyTorch tensors; Also accept list of + supported formats. + height (`int`, *optional*): + The height in preprocessed image. If `None`, will use the `get_default_height_width()` to get default + height. + width (`int`, *optional*): + The width in preprocessed. If `None`, will use get_default_height_width()` to get the default width. + resize_mode (`str`, *optional*, defaults to `default`): + The resize mode, can be one of `default` or `fill`. If `default`, will resize the image to fit within + the specified width and height, and it may not maintaining the original aspect ratio. If `fill`, will + resize the image to fit within the specified width and height, maintaining the aspect ratio, and then + center the image within the dimensions, filling empty with data from image. If `crop`, will resize the + image to fit within the specified width and height, maintaining the aspect ratio, and then center the + image within the dimensions, cropping the excess. Note that resize_mode `fill` and `crop` are only + supported for PIL image input. + crops_coords (`List[Tuple[int, int, int, int]]`, *optional*, defaults to `None`): + The crop coordinates for each image in the batch. If `None`, will not crop the image. + + Returns: + `torch.Tensor`: + The preprocessed image. + """ + supported_formats = (PIL.Image.Image, np.ndarray, torch.Tensor) + + # Expand the missing dimension for 3-dimensional pytorch tensor or numpy array that represents grayscale image + if self.config.do_convert_grayscale and isinstance(image, (torch.Tensor, np.ndarray)) and image.ndim == 3: + if isinstance(image, torch.Tensor): + # if image is a pytorch tensor could have 2 possible shapes: + # 1. batch x height x width: we should insert the channel dimension at position 1 + # 2. channel x height x width: we should insert batch dimension at position 0, + # however, since both channel and batch dimension has same size 1, it is same to insert at position 1 + # for simplicity, we insert a dimension of size 1 at position 1 for both cases + image = image.unsqueeze(1) + else: + # if it is a numpy array, it could have 2 possible shapes: + # 1. batch x height x width: insert channel dimension on last position + # 2. height x width x channel: insert batch dimension on first position + if image.shape[-1] == 1: + image = np.expand_dims(image, axis=0) + else: + image = np.expand_dims(image, axis=-1) + + if isinstance(image, list) and isinstance(image[0], np.ndarray) and image[0].ndim == 4: + warnings.warn( + "Passing `image` as a list of 4d np.ndarray is deprecated." + "Please concatenate the list along the batch dimension and pass it as a single 4d np.ndarray", + FutureWarning, + ) + image = np.concatenate(image, axis=0) + if isinstance(image, list) and isinstance(image[0], torch.Tensor) and image[0].ndim == 4: + warnings.warn( + "Passing `image` as a list of 4d torch.Tensor is deprecated." + "Please concatenate the list along the batch dimension and pass it as a single 4d torch.Tensor", + FutureWarning, + ) + image = torch.cat(image, axis=0) + + if not is_valid_image_imagelist(image): + raise ValueError( + f"Input is in incorrect format. Currently, we only support {', '.join(str(x) for x in supported_formats)}" + ) + if not isinstance(image, list): + image = [image] + + if isinstance(image[0], PIL.Image.Image): + if crops_coords is not None: + image = [i.crop(crops_coords) for i in image] + if self.config.do_resize: + height, width = self.get_new_height_width(image[0], height, width, max_pixels, max_side_length) + image = [self.resize(i, height, width, resize_mode=resize_mode) for i in image] + if self.config.do_convert_rgb: + image = [self.convert_to_rgb(i) for i in image] + elif self.config.do_convert_grayscale: + image = [self.convert_to_grayscale(i) for i in image] + image = self.pil_to_numpy(image) # to np + image = self.numpy_to_pt(image) # to pt + + elif isinstance(image[0], np.ndarray): + image = np.concatenate(image, axis=0) if image[0].ndim == 4 else np.stack(image, axis=0) + + image = self.numpy_to_pt(image) + + height, width = self.get_new_height_width(image, height, width, max_pixels, max_side_length) + if self.config.do_resize: + image = self.resize(image, height, width) + + elif isinstance(image[0], torch.Tensor): + image = torch.cat(image, axis=0) if image[0].ndim == 4 else torch.stack(image, axis=0) + + if self.config.do_convert_grayscale and image.ndim == 3: + image = image.unsqueeze(1) + + channel = image.shape[1] + # don't need any preprocess if the image is latents + if channel == self.config.vae_latent_channels: + return image + + height, width = self.get_new_height_width(image, height, width, max_pixels, max_side_length) + if self.config.do_resize: + image = self.resize(image, height, width) + + # expected range [0,1], normalize to [-1,1] + do_normalize = do_normalize if do_normalize is not None else self.config.do_normalize + if do_normalize and image.min() < 0: + warnings.warn( + "Passing `image` as torch tensor with value range in [-1,1] is deprecated. The expected value range for image tensor is [0,1] " + f"when passing as pytorch tensor or numpy Array. You passed `image` with value range [{image.min()},{image.max()}]", + FutureWarning, + ) + do_normalize = False + if do_normalize: + image = self.normalize(image) + + if self.config.do_binarize: + image = self.binarize(image) + + return image \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/pipelines/lora_pipeline.py b/examples/OmniGen2-RL/omnigen2/pipelines/lora_pipeline.py new file mode 100644 index 0000000000000000000000000000000000000000..60dac28f674187b49deae601fbfd4e87261a8a19 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/pipelines/lora_pipeline.py @@ -0,0 +1,414 @@ +# Copyright 2024 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import os +from typing import Callable, Dict, List, Optional, Union + +import torch +from huggingface_hub.utils import validate_hf_hub_args + +from diffusers.utils import ( + USE_PEFT_BACKEND, + is_peft_available, + is_peft_version, + is_torch_version, + is_transformers_available, + is_transformers_version, + logging, +) +from diffusers.loaders.lora_base import ( # noqa + LoraBaseMixin, + _fetch_state_dict, + _pack_dict_with_prefix +) +from diffusers.loaders.lora_conversion_utils import ( + _convert_non_diffusers_lumina2_lora_to_diffusers, +) + + +_LOW_CPU_MEM_USAGE_DEFAULT_LORA = False +if is_torch_version(">=", "1.9.0"): + if ( + is_peft_available() + and is_peft_version(">=", "0.13.1") + and is_transformers_available() + and is_transformers_version(">", "4.45.2") + ): + _LOW_CPU_MEM_USAGE_DEFAULT_LORA = True + + +logger = logging.get_logger(__name__) + +TRANSFORMER_NAME = "transformer" + +class OmniGen2LoraLoaderMixin(LoraBaseMixin): + r""" + Load LoRA layers into [`OmniGen2Transformer2DModel`]. Specific to [`OmniGen2Pipeline`]. + """ + + _lora_loadable_modules = ["transformer"] + transformer_name = TRANSFORMER_NAME + + @classmethod + @validate_hf_hub_args + def lora_state_dict( + cls, + pretrained_model_name_or_path_or_dict: Union[str, Dict[str, torch.Tensor]], + **kwargs, + ): + r""" + Return state dict for lora weights and the network alphas. + + + + We support loading A1111 formatted LoRA checkpoints in a limited capacity. + + This function is experimental and might change in the future. + + + + Parameters: + pretrained_model_name_or_path_or_dict (`str` or `os.PathLike` or `dict`): + Can be either: + + - A string, the *model id* (for example `google/ddpm-celebahq-256`) of a pretrained model hosted on + the Hub. + - A path to a *directory* (for example `./my_model_directory`) containing the model weights saved + with [`ModelMixin.save_pretrained`]. + - A [torch state + dict](https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict). + + cache_dir (`Union[str, os.PathLike]`, *optional*): + Path to a directory where a downloaded pretrained model configuration is cached if the standard cache + is not used. + force_download (`bool`, *optional*, defaults to `False`): + Whether or not to force the (re-)download of the model weights and configuration files, overriding the + cached versions if they exist. + + proxies (`Dict[str, str]`, *optional*): + A dictionary of proxy servers to use by protocol or endpoint, for example, `{'http': 'foo.bar:3128', + 'http://hostname': 'foo.bar:4012'}`. The proxies are used on each request. + local_files_only (`bool`, *optional*, defaults to `False`): + Whether to only load local model weights and configuration files or not. If set to `True`, the model + won't be downloaded from the Hub. + token (`str` or *bool*, *optional*): + The token to use as HTTP bearer authorization for remote files. If `True`, the token generated from + `diffusers-cli login` (stored in `~/.huggingface`) is used. + revision (`str`, *optional*, defaults to `"main"`): + The specific model version to use. It can be a branch name, a tag name, a commit id, or any identifier + allowed by Git. + subfolder (`str`, *optional*, defaults to `""`): + The subfolder location of a model file within a larger model repository on the Hub or locally. + + """ + # Load the main state dict first which has the LoRA layers for either of + # transformer and text encoder or both. + cache_dir = kwargs.pop("cache_dir", None) + force_download = kwargs.pop("force_download", False) + proxies = kwargs.pop("proxies", None) + local_files_only = kwargs.pop("local_files_only", None) + token = kwargs.pop("token", None) + revision = kwargs.pop("revision", None) + subfolder = kwargs.pop("subfolder", None) + weight_name = kwargs.pop("weight_name", None) + use_safetensors = kwargs.pop("use_safetensors", None) + return_lora_metadata = kwargs.pop("return_lora_metadata", False) + + allow_pickle = False + if use_safetensors is None: + use_safetensors = True + allow_pickle = True + + user_agent = { + "file_type": "attn_procs_weights", + "framework": "pytorch", + } + + state_dict, metadata = _fetch_state_dict( + pretrained_model_name_or_path_or_dict=pretrained_model_name_or_path_or_dict, + weight_name=weight_name, + use_safetensors=use_safetensors, + local_files_only=local_files_only, + cache_dir=cache_dir, + force_download=force_download, + proxies=proxies, + token=token, + revision=revision, + subfolder=subfolder, + user_agent=user_agent, + allow_pickle=allow_pickle, + ) + + is_dora_scale_present = any("dora_scale" in k for k in state_dict) + if is_dora_scale_present: + warn_msg = "It seems like you are using a DoRA checkpoint that is not compatible in Diffusers at the moment. So, we are going to filter out the keys associated to 'dora_scale` from the state dict. If you think this is a mistake please open an issue https://github.com/huggingface/diffusers/issues/new." + logger.warning(warn_msg) + state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k} + + # conversion. + non_diffusers = any(k.startswith("diffusion_model.") for k in state_dict) + if non_diffusers: + state_dict = _convert_non_diffusers_lumina2_lora_to_diffusers(state_dict) + + out = (state_dict, metadata) if return_lora_metadata else state_dict + return out + + # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights + def load_lora_weights( + self, + pretrained_model_name_or_path_or_dict: Union[str, Dict[str, torch.Tensor]], + adapter_name=None, + hotswap: bool = False, + **kwargs, + ): + """ + Load LoRA weights specified in `pretrained_model_name_or_path_or_dict` into `self.transformer` and + `self.text_encoder`. All kwargs are forwarded to `self.lora_state_dict`. See + [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`] for more details on how the state dict is loaded. + See [`~loaders.StableDiffusionLoraLoaderMixin.load_lora_into_transformer`] for more details on how the state + dict is loaded into `self.transformer`. + + Parameters: + pretrained_model_name_or_path_or_dict (`str` or `os.PathLike` or `dict`): + See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`]. + adapter_name (`str`, *optional*): + Adapter name to be used for referencing the loaded adapter model. If not specified, it will use + `default_{i}` where i is the total number of adapters being loaded. + low_cpu_mem_usage (`bool`, *optional*): + Speed up model loading by only loading the pretrained LoRA weights and not initializing the random + weights. + kwargs (`dict`, *optional*): + See [`~loaders.StableDiffusionLoraLoaderMixin.lora_state_dict`]. + """ + if not USE_PEFT_BACKEND: + raise ValueError("PEFT backend is required for this method.") + + low_cpu_mem_usage = kwargs.pop("low_cpu_mem_usage", _LOW_CPU_MEM_USAGE_DEFAULT_LORA) + if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): + raise ValueError( + "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." + ) + + # if a dict is passed, copy it instead of modifying it inplace + if isinstance(pretrained_model_name_or_path_or_dict, dict): + pretrained_model_name_or_path_or_dict = pretrained_model_name_or_path_or_dict.copy() + + # First, ensure that the checkpoint is a compatible one and can be successfully loaded. + kwargs["return_lora_metadata"] = True + state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs) + # print(f"{state_dict=}") + + is_correct_format = all("lora" in key for key in state_dict.keys()) + if not is_correct_format: + raise ValueError("Invalid LoRA checkpoint.") + + self.load_lora_into_transformer( + state_dict, + transformer=getattr(self, self.transformer_name) if not hasattr(self, "transformer") else self.transformer, + adapter_name=adapter_name, + metadata=metadata, + _pipeline=self, + low_cpu_mem_usage=low_cpu_mem_usage, + hotswap=hotswap, + ) + + @classmethod + # Copied from diffusers.loaders.lora_pipeline.SD3LoraLoaderMixin.load_lora_into_transformer with SD3Transformer2DModel->Lumina2Transformer2DModel + def load_lora_into_transformer( + cls, + state_dict, + transformer, + adapter_name=None, + _pipeline=None, + low_cpu_mem_usage=False, + hotswap: bool = False, + metadata=None, + ): + """ + This will load the LoRA layers specified in `state_dict` into `transformer`. + + Parameters: + state_dict (`dict`): + A standard state dict containing the lora layer parameters. The keys can either be indexed directly + into the unet or prefixed with an additional `unet` which can be used to distinguish between text + encoder lora layers. + transformer (`Lumina2Transformer2DModel`): + The Transformer model to load the LoRA layers into. + adapter_name (`str`, *optional*): + Adapter name to be used for referencing the loaded adapter model. If not specified, it will use + `default_{i}` where i is the total number of adapters being loaded. + low_cpu_mem_usage (`bool`, *optional*): + Speed up model loading by only loading the pretrained LoRA weights and not initializing the random + weights. + hotswap : (`bool`, *optional*) + Defaults to `False`. Whether to substitute an existing (LoRA) adapter with the newly loaded adapter + in-place. This means that, instead of loading an additional adapter, this will take the existing + adapter weights and replace them with the weights of the new adapter. This can be faster and more + memory efficient. However, the main advantage of hotswapping is that when the model is compiled with + torch.compile, loading the new adapter does not require recompilation of the model. When using + hotswapping, the passed `adapter_name` should be the name of an already loaded adapter. + + If the new adapter and the old adapter have different ranks and/or LoRA alphas (i.e. scaling), you need + to call an additional method before loading the adapter: + + ```py + pipeline = ... # load diffusers pipeline + max_rank = ... # the highest rank among all LoRAs that you want to load + # call *before* compiling and loading the LoRA adapter + pipeline.enable_lora_hotswap(target_rank=max_rank) + pipeline.load_lora_weights(file_name) + # optionally compile the model now + ``` + + Note that hotswapping adapters of the text encoder is not yet supported. There are some further + limitations to this technique, which are documented here: + https://huggingface.co/docs/peft/main/en/package_reference/hotswap + """ + if low_cpu_mem_usage and is_peft_version("<", "0.13.0"): + raise ValueError( + "`low_cpu_mem_usage=True` is not compatible with this `peft` version. Please update it with `pip install -U peft`." + ) + + # Load the layers corresponding to transformer. + logger.info(f"Loading {cls.transformer_name}.") + transformer.load_lora_adapter( + state_dict, + network_alphas=None, + adapter_name=adapter_name, + metadata=metadata, + _pipeline=_pipeline, + low_cpu_mem_usage=low_cpu_mem_usage, + hotswap=hotswap, + ) + + @classmethod + # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.save_lora_weights + def save_lora_weights( + cls, + save_directory: Union[str, os.PathLike], + transformer_lora_layers: Dict[str, Union[torch.nn.Module, torch.Tensor]] = None, + is_main_process: bool = True, + weight_name: str = None, + save_function: Callable = None, + safe_serialization: bool = True, + transformer_lora_adapter_metadata: Optional[dict] = None, + ): + r""" + Save the LoRA parameters corresponding to the UNet and text encoder. + + Arguments: + save_directory (`str` or `os.PathLike`): + Directory to save LoRA parameters to. Will be created if it doesn't exist. + transformer_lora_layers (`Dict[str, torch.nn.Module]` or `Dict[str, torch.Tensor]`): + State dict of the LoRA layers corresponding to the `transformer`. + is_main_process (`bool`, *optional*, defaults to `True`): + Whether the process calling this is the main process or not. Useful during distributed training and you + need to call this function on all processes. In this case, set `is_main_process=True` only on the main + process to avoid race conditions. + save_function (`Callable`): + The function to use to save the state dictionary. Useful during distributed training when you need to + replace `torch.save` with another method. Can be configured with the environment variable + `DIFFUSERS_SAVE_MODE`. + safe_serialization (`bool`, *optional*, defaults to `True`): + Whether to save the model using `safetensors` or the traditional PyTorch way with `pickle`. + """ + state_dict = {} + lora_adapter_metadata = {} + + if not transformer_lora_layers: + raise ValueError("You must pass `transformer_lora_layers`.") + + state_dict.update(cls.pack_weights(transformer_lora_layers, cls.transformer_name)) + + if transformer_lora_adapter_metadata is not None: + lora_adapter_metadata.update( + _pack_dict_with_prefix(transformer_lora_adapter_metadata, cls.transformer_name) + ) + + # Save the model + cls.write_lora_layers( + state_dict=state_dict, + save_directory=save_directory, + is_main_process=is_main_process, + weight_name=weight_name, + save_function=save_function, + safe_serialization=safe_serialization, + lora_adapter_metadata=lora_adapter_metadata, + ) + + # Copied from diffusers.loaders.lora_pipeline.SanaLoraLoaderMixin.fuse_lora + def fuse_lora( + self, + components: List[str] = ["transformer"], + lora_scale: float = 1.0, + safe_fusing: bool = False, + adapter_names: Optional[List[str]] = None, + **kwargs, + ): + r""" + Fuses the LoRA parameters into the original parameters of the corresponding blocks. + + + + This is an experimental API. + + + + Args: + components: (`List[str]`): List of LoRA-injectable components to fuse the LoRAs into. + lora_scale (`float`, defaults to 1.0): + Controls how much to influence the outputs with the LoRA parameters. + safe_fusing (`bool`, defaults to `False`): + Whether to check fused weights for NaN values before fusing and if values are NaN not fusing them. + adapter_names (`List[str]`, *optional*): + Adapter names to be used for fusing. If nothing is passed, all active adapters will be fused. + + Example: + + ```py + from diffusers import DiffusionPipeline + import torch + + pipeline = DiffusionPipeline.from_pretrained( + "stabilityai/stable-diffusion-xl-base-1.0", torch_dtype=torch.float16 + ).to("cuda") + pipeline.load_lora_weights("nerijs/pixel-art-xl", weight_name="pixel-art-xl.safetensors", adapter_name="pixel") + pipeline.fuse_lora(lora_scale=0.7) + ``` + """ + super().fuse_lora( + components=components, + lora_scale=lora_scale, + safe_fusing=safe_fusing, + adapter_names=adapter_names, + **kwargs, + ) + + # Copied from diffusers.loaders.lora_pipeline.SanaLoraLoaderMixin.unfuse_lora + def unfuse_lora(self, components: List[str] = ["transformer"], **kwargs): + r""" + Reverses the effect of + [`pipe.fuse_lora()`](https://huggingface.co/docs/diffusers/main/en/api/loaders#diffusers.loaders.LoraBaseMixin.fuse_lora). + + + + This is an experimental API. + + + + Args: + components (`List[str]`): List of LoRA-injectable components to unfuse LoRA from. + unfuse_transformer (`bool`, defaults to `True`): Whether to unfuse the UNet LoRA parameters. + """ + super().unfuse_lora(components=components, **kwargs) \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/pipelines/omnigen2/pipeline_omnigen2.py b/examples/OmniGen2-RL/omnigen2/pipelines/omnigen2/pipeline_omnigen2.py new file mode 100644 index 0000000000000000000000000000000000000000..1aae7803afa500799c2c966e2e2e153dc8d389b5 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/pipelines/omnigen2/pipeline_omnigen2.py @@ -0,0 +1,1032 @@ +""" +OmniGen2 Diffusion Pipeline + +Copyright 2025 BAAI, The OmniGen2 Team and The HuggingFace Team. All rights reserved. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +import inspect +from typing import Any, Callable, Dict, List, Optional, Tuple, Union + +import math + +from PIL import Image +import numpy as np +import torch +import torch.nn.functional as F + +from transformers import Qwen2_5_VLForConditionalGeneration + +from diffusers.models.autoencoders import AutoencoderKL +from ...models.transformers import OmniGen2Transformer2DModel +from ...models.transformers.repo import OmniGen2RotaryPosEmbed +from diffusers.schedulers import FlowMatchEulerDiscreteScheduler +from diffusers.utils import ( + is_torch_xla_available, + logging, +) +from diffusers.utils.torch_utils import randn_tensor +from diffusers.pipelines.pipeline_utils import DiffusionPipeline + +from dataclasses import dataclass + +from einops import rearrange + +import PIL.Image + +from diffusers.utils import BaseOutput + +from omnigen2.pipelines.image_processor import OmniGen2ImageProcessor +from ..lora_pipeline import OmniGen2LoraLoaderMixin +from ...utils.tensor_util import pad_to_length + +if is_torch_xla_available(): + import torch_xla.core.xla_model as xm + + XLA_AVAILABLE = True +else: + XLA_AVAILABLE = False + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + +@dataclass +class FMPipelineOutput(BaseOutput): + """ + Output class for OmniGen2 pipeline. + + Args: + images (Union[List[PIL.Image.Image], np.ndarray]): + List of denoised PIL images of length `batch_size` or numpy array of shape + `(batch_size, height, width, num_channels)`. Contains the generated images. + """ + images: Union[List[PIL.Image.Image], np.ndarray] + middle_latents: Optional[List[torch.FloatTensor]] = None + log_probs: Optional[List[torch.FloatTensor]] = None + img_mask: Optional[torch.FloatTensor] = None + l_effective_img_len: Optional[List[int]] = None + img_sizes: Optional[List[Tuple[int, int]]] = None + ref_latents: Optional[torch.FloatTensor] = None + ref_img_mask: Optional[torch.FloatTensor] = None + l_effective_ref_img_len: Optional[List[int]] = None + ref_img_sizes: Optional[List[Tuple[int, int]]] = None + +# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps +def retrieve_timesteps( + scheduler, + num_inference_steps: Optional[int] = None, + device: Optional[Union[str, torch.device]] = None, + timesteps: Optional[List[int]] = None, + **kwargs, +): + """ + Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles + custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`. + + Args: + scheduler (`SchedulerMixin`): + The scheduler to get timesteps from. + num_inference_steps (`int`): + The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps` + must be `None`. + device (`str` or `torch.device`, *optional*): + The device to which the timesteps should be moved to. If `None`, the timesteps are not moved. + timesteps (`List[int]`, *optional*): + Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed, + `num_inference_steps` and `sigmas` must be `None`. + sigmas (`List[float]`, *optional*): + Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed, + `num_inference_steps` and `timesteps` must be `None`. + + Returns: + `Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the + second element is the number of inference steps. + """ + if timesteps is not None: + accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) + if not accepts_timesteps: + raise ValueError( + f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" + f" timestep schedules. Please check whether you are using the correct scheduler." + ) + scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs) + timesteps = scheduler.timesteps + num_inference_steps = len(timesteps) + else: + scheduler.set_timesteps(num_inference_steps, device=device, **kwargs) + timesteps = scheduler.timesteps + return timesteps, num_inference_steps + + +class OmniGen2Pipeline(DiffusionPipeline, OmniGen2LoraLoaderMixin): + """ + Pipeline for text-to-image generation using OmniGen2. + + This pipeline implements a text-to-image generation model that uses: + - Qwen2.5-VL for text encoding + - A custom transformer architecture for image generation + - VAE for image encoding/decoding + - FlowMatchEulerDiscreteScheduler for noise scheduling + + Args: + transformer (OmniGen2Transformer2DModel): The transformer model for image generation. + vae (AutoencoderKL): The VAE model for image encoding/decoding. + scheduler (FlowMatchEulerDiscreteScheduler): The scheduler for noise scheduling. + text_encoder (Qwen2_5_VLModel): The text encoder model. + tokenizer (Union[Qwen2Tokenizer, Qwen2TokenizerFast]): The tokenizer for text processing. + """ + + model_cpu_offload_seq = "mllm->transformer->vae" + + def __init__( + self, + transformer: OmniGen2Transformer2DModel, + vae: AutoencoderKL, + scheduler: FlowMatchEulerDiscreteScheduler, + mllm: Qwen2_5_VLForConditionalGeneration, + processor, + ) -> None: + """ + Initialize the OmniGen2 pipeline. + + Args: + transformer: The transformer model for image generation. + vae: The VAE model for image encoding/decoding. + scheduler: The scheduler for noise scheduling. + text_encoder: The text encoder model. + tokenizer: The tokenizer for text processing. + """ + super().__init__() + + self.register_modules( + transformer=transformer, + vae=vae, + scheduler=scheduler, + mllm=mllm, + processor=processor + ) + self.vae_scale_factor = ( + 2 ** (len(self.vae.config.block_out_channels) - 1) if hasattr(self, "vae") and self.vae is not None else 8 + ) + self.image_processor = OmniGen2ImageProcessor(vae_scale_factor=self.vae_scale_factor * 2, do_resize=True) + self.default_sample_size = 128 + + def prepare_latents( + self, + # batch_size: int, + num_channels_latents: int, + size: List[Tuple[int, int]], + num_images_per_prompt: int, + dtype: torch.dtype, + device: torch.device, + generator: Optional[torch.Generator], + latents: Optional[torch.FloatTensor] = None, + ) -> torch.FloatTensor: + """ + Prepare the initial latents for the diffusion process. + + Args: + batch_size: The number of images to generate. + num_channels_latents: The number of channels in the latent space. + height: The height of the generated image. + width: The width of the generated image. + dtype: The data type of the latents. + device: The device to place the latents on. + generator: The random number generator to use. + latents: Optional pre-computed latents to use instead of random initialization. + + Returns: + torch.FloatTensor: The prepared latents tensor. + """ + if latents is None: + latents = [] + for _size in size: + width = int(_size[0]) // self.vae_scale_factor + height = int(_size[1]) // self.vae_scale_factor + + shape = (num_images_per_prompt, num_channels_latents, height, width) + latent = randn_tensor(shape, generator=generator, device=device, dtype=dtype) + for i in range(num_images_per_prompt): + latents.append(latent[i]) + else: + # for i in range(num_images_per_prompt): + for i in range(len(latents)): + latents[i] = latents[i].to(device) + return latents + + def encode_vae(self, img: torch.FloatTensor) -> torch.FloatTensor: + """ + Encode an image into the VAE latent space. + + Args: + img: The input image tensor to encode. + + Returns: + torch.FloatTensor: The encoded latent representation. + """ + z0 = self.vae.encode(img.to(dtype=self.vae.dtype)).latent_dist.sample() + if self.vae.config.shift_factor is not None: + z0 = z0 - self.vae.config.shift_factor + if self.vae.config.scaling_factor is not None: + z0 = z0 * self.vae.config.scaling_factor + z0 = z0.to(dtype=self.vae.dtype) + return z0 + + def prepare_image( + self, + images: Union[List[PIL.Image.Image], PIL.Image.Image], + batch_size: int, + num_images_per_prompt: int, + max_pixels: int, + max_side_length: int, + do_normalize: bool, + device: torch.device, + dtype: torch.dtype, + ) -> List[Optional[torch.FloatTensor]]: + """ + Prepare input images for processing by encoding them into the VAE latent space. + + Args: + images: Single image or list of images to process. + batch_size: The number of images to generate per prompt. + num_images_per_prompt: The number of images to generate for each prompt. + device: The device to place the encoded latents on. + dtype: The data type of the encoded latents. + + Returns: + List[Optional[torch.FloatTensor]]: List of encoded latent representations for each image. + """ + latents = [] + for i, img in enumerate(images): + if img is not None and len(img) > 0: + ref_latents = [] + for j, img_j in enumerate(img): + img_j = self.image_processor.preprocess(img_j, max_pixels=max_pixels, max_side_length=max_side_length, do_normalize=do_normalize) + ref_latents.append(self.encode_vae(img_j.to(device=device)).squeeze(0)) + else: + ref_latents = None + latents.append(ref_latents) + + latents = [latent for latent in latents for _ in range(num_images_per_prompt)] + return latents + + def _get_qwen2_prompt_embeds( + self, + prompt: Union[str, List[str]], + device: Optional[torch.device] = None, + max_sequence_length: int = 256, + ) -> Tuple[torch.Tensor, torch.Tensor]: + """ + Get prompt embeddings from the Qwen2 text encoder. + + Args: + prompt: The prompt or list of prompts to encode. + device: The device to place the embeddings on. If None, uses the pipeline's device. + max_sequence_length: Maximum sequence length for tokenization. + + Returns: + Tuple[torch.Tensor, torch.Tensor]: A tuple containing: + - The prompt embeddings tensor + - The attention mask tensor + + Raises: + Warning: If the input text is truncated due to sequence length limitations. + """ + device = device or self._execution_device + prompt = [prompt] if isinstance(prompt, str) else prompt + + text_inputs = self.processor.tokenizer( + prompt, + padding="longest", + max_length=max_sequence_length, + truncation=True, + return_tensors="pt", + ) + + text_input_ids = text_inputs.input_ids.to(device) + untruncated_ids = self.processor.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids.to(device) + + if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids): + removed_text = self.processor.tokenizer.batch_decode(untruncated_ids[:, max_sequence_length - 1 : -1]) + logger.warning( + "The following part of your input was truncated because Gemma can only handle sequences up to" + f" {max_sequence_length} tokens: {removed_text}" + ) + + prompt_attention_mask = text_inputs.attention_mask.to(device) + prompt_embeds = self.mllm( + text_input_ids, + attention_mask=prompt_attention_mask, + output_hidden_states=True, + ).hidden_states[-1] + + if self.mllm is not None: + dtype = self.mllm.dtype + elif self.transformer is not None: + dtype = self.transformer.dtype + else: + dtype = None + + prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) + + return prompt_embeds, prompt_attention_mask + + def _apply_chat_template(self, prompt: str): + prompt = [ + { + "role": "system", + "content": "You are a helpful assistant that generates high-quality images based on user instructions.", + }, + {"role": "user", "content": prompt}, + ] + prompt = self.processor.tokenizer.apply_chat_template(prompt, tokenize=False, add_generation_prompt=False) + return prompt + + def encode_prompt( + self, + prompt: Union[str, List[str]], + do_classifier_free_guidance: bool = True, + negative_prompt: Optional[Union[str, List[str]]] = None, + num_images_per_prompt: int = 1, + device: Optional[torch.device] = None, + prompt_embeds: Optional[torch.Tensor] = None, + negative_prompt_embeds: Optional[torch.Tensor] = None, + prompt_attention_mask: Optional[torch.Tensor] = None, + negative_prompt_attention_mask: Optional[torch.Tensor] = None, + max_sequence_length: int = 256, + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + r""" + Encodes the prompt into text encoder hidden states. + + Args: + prompt (`str` or `List[str]`, *optional*): + prompt to be encoded + negative_prompt (`str` or `List[str]`, *optional*): + The prompt not to guide the image generation. If not defined, one has to pass `negative_prompt_embeds` + instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is less than `1`). For + Lumina-T2I, this should be "". + do_classifier_free_guidance (`bool`, *optional*, defaults to `True`): + whether to use classifier free guidance or not + num_images_per_prompt (`int`, *optional*, defaults to 1): + number of images that should be generated per prompt + device: (`torch.device`, *optional*): + torch device to place the resulting embeddings on + prompt_embeds (`torch.Tensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not + provided, text embeddings will be generated from `prompt` input argument. + negative_prompt_embeds (`torch.Tensor`, *optional*): + Pre-generated negative text embeddings. For Lumina-T2I, it's should be the embeddings of the "" string. + max_sequence_length (`int`, defaults to `256`): + Maximum sequence length to use for the prompt. + """ + device = device or self._execution_device + + if prompt is not None: + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + if prompt_embeds is None: + prompt = [prompt] if isinstance(prompt, str) else prompt + prompt = [self._apply_chat_template(_prompt) for _prompt in prompt] + + prompt_embeds, prompt_attention_mask = self._get_qwen2_prompt_embeds( + prompt=prompt, + device=device, + max_sequence_length=max_sequence_length + ) + + batch_size, seq_len, _ = prompt_embeds.shape + # duplicate text embeddings and attention mask for each generation per prompt, using mps friendly method + prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1) + prompt_embeds = prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1) + prompt_attention_mask = prompt_attention_mask.repeat(num_images_per_prompt, 1) + prompt_attention_mask = prompt_attention_mask.view(batch_size * num_images_per_prompt, -1) + + # Get negative embeddings for classifier free guidance + if do_classifier_free_guidance: + if negative_prompt_embeds is None: + negative_prompt = negative_prompt if negative_prompt is not None else "" + + # Normalize str to list + negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt + negative_prompt = [self._apply_chat_template(_negative_prompt) for _negative_prompt in negative_prompt] + + if prompt is not None and type(prompt) is not type(negative_prompt): + raise TypeError( + f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !=" + f" {type(prompt)}." + ) + elif isinstance(negative_prompt, str): + negative_prompt = [negative_prompt] + elif batch_size != len(negative_prompt): + raise ValueError( + f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:" + f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches" + " the batch size of `prompt`." + ) + negative_prompt_embeds, negative_prompt_attention_mask = self._get_qwen2_prompt_embeds( + prompt=negative_prompt, + device=device, + max_sequence_length=max_sequence_length, + ) + + batch_size, seq_len, _ = negative_prompt_embeds.shape + # duplicate text embeddings and attention mask for each generation per prompt, using mps friendly method + negative_prompt_embeds = negative_prompt_embeds.repeat(1, num_images_per_prompt, 1) + negative_prompt_embeds = negative_prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1) + negative_prompt_attention_mask = negative_prompt_attention_mask.repeat(num_images_per_prompt, 1) + negative_prompt_attention_mask = negative_prompt_attention_mask.view( + batch_size * num_images_per_prompt, -1 + ) + + return prompt_embeds, prompt_attention_mask, negative_prompt_embeds, negative_prompt_attention_mask + + @property + def num_timesteps(self): + return self._num_timesteps + + @property + def text_guidance_scale(self): + return self._text_guidance_scale + + @property + def image_guidance_scale(self): + return self._image_guidance_scale + + @property + def cfg_range(self): + return self._cfg_range + + @property + def enable_parallel_cfg(self): + return self._enable_parallel_cfg + + @property + def mixed_precision(self): + return self._mixed_precision + + @torch.no_grad() + def __call__( + self, + prompt: Optional[Union[str, List[str]]] = None, + negative_prompt: Optional[Union[str, List[str]]] = None, + prompt_embeds: Optional[torch.FloatTensor] = None, + negative_prompt_embeds: Optional[torch.FloatTensor] = None, + prompt_attention_mask: Optional[torch.LongTensor] = None, + negative_prompt_attention_mask: Optional[torch.LongTensor] = None, + max_sequence_length: Optional[int] = None, + callback_on_step_end_tensor_inputs: Optional[List[str]] = None, + input_images: Optional[List[PIL.Image.Image]] = None, + num_images_per_prompt: int = 1, + size: Optional[Union[Tuple[int, int], List[Tuple[int, int]]]] = None, + max_pixels: int = 1024 * 1024, + max_input_image_side_length: int = 1024, + align_res: bool = True, + num_inference_steps: int = 28, + text_guidance_scale: float = 4.0, + image_guidance_scale: float = 1.0, + cfg_range: Tuple[float, float] = (0.0, 1.0), + enable_parallel_cfg: bool = False, + attention_kwargs: Optional[Dict[str, Any]] = None, + timesteps: List[int] = None, + generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, + latents: Optional[torch.FloatTensor] = None, + output_type: Optional[str] = "pil", + return_middle_statistics: bool = False, + return_dict: bool = True, + verbose: bool = False, + step_func=None, + mixed_precision: bool = False, + do_normalize: Optional[bool] = None + ): + + size = size or self.default_sample_size * self.vae_scale_factor + + self._text_guidance_scale = text_guidance_scale + self._image_guidance_scale = image_guidance_scale + self._cfg_range = cfg_range + self._enable_parallel_cfg = enable_parallel_cfg + self._attention_kwargs = attention_kwargs + self._mixed_precision = mixed_precision + # 2. Define call parameters + if prompt is not None and isinstance(prompt, str): + batch_size = 1 + elif prompt is not None and isinstance(prompt, list): + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + if not isinstance(size, list): + size = [size] * batch_size + + device = self._execution_device + + # 3. Encode input prompt + ( + prompt_embeds, + prompt_attention_mask, + negative_prompt_embeds, + negative_prompt_attention_mask, + ) = self.encode_prompt( + prompt, + self.text_guidance_scale > 1.0, + negative_prompt=negative_prompt, + num_images_per_prompt=num_images_per_prompt, + device=device, + prompt_embeds=prompt_embeds, + negative_prompt_embeds=negative_prompt_embeds, + prompt_attention_mask=prompt_attention_mask, + negative_prompt_attention_mask=negative_prompt_attention_mask, + max_sequence_length=max_sequence_length, + ) + + dtype = self.vae.dtype + # 3. Prepare control image + ref_latents = self.prepare_image( + images=input_images, + batch_size=batch_size, + num_images_per_prompt=num_images_per_prompt, + max_pixels=max_pixels, + max_side_length=max_input_image_side_length, + do_normalize=do_normalize, + device=device, + dtype=dtype, + ) + + for i, _input_images in enumerate(input_images): + if _input_images is None: + input_images[i] = [] + + ori_size = [] + for i, (_input_images, _size) in enumerate(zip(input_images, size)): + if len(_input_images) == 1 and align_res: + size[i] = (ref_latents[i][0].shape[-1] * self.vae_scale_factor, ref_latents[i][0].shape[-2] * self.vae_scale_factor) + ori_size.append(size[i]) + else: + ori_size.append(_size) + + cur_pixels = _size[0] * _size[1] + ratio = (max_pixels / cur_pixels) ** 0.5 + ratio = min(ratio, 1.0) + + width, height = int(_size[0] * ratio) // 16 * 16, int(_size[1] * ratio) // 16 * 16 + size[i] = (width, height) + + # 4. Prepare latents. + latent_channels = self.transformer.config.in_channels + latents = self.prepare_latents( + latent_channels, + size, + num_images_per_prompt, + prompt_embeds.dtype, + device, + generator, + latents, + ) + + freqs_cis = OmniGen2RotaryPosEmbed.get_freqs_cis( + self.transformer.config.axes_dim_rope, + self.transformer.config.axes_lens, + theta=10000, + ) + + image = self.processing( + latents=latents, + ref_latents=ref_latents, + prompt_embeds=prompt_embeds, + freqs_cis=freqs_cis, + negative_prompt_embeds=negative_prompt_embeds, + prompt_attention_mask=prompt_attention_mask, + negative_prompt_attention_mask=negative_prompt_attention_mask, + num_inference_steps=num_inference_steps, + timesteps=timesteps, + generator=generator, + device=device, + dtype=dtype, + verbose=verbose, + step_func=step_func, + return_middle_statistics=return_middle_statistics, + ) + + if return_middle_statistics: + image, middle_latents, log_probs, img_mask, l_effective_img_len, img_sizes, ref_latents, ref_img_mask, l_effective_ref_img_len, ref_img_sizes = image + + ori_size = [_ori_size for _ori_size in ori_size for _ in range(num_images_per_prompt)] + + postprocessed_image = [] + for i, (_image, _ori_size) in enumerate(zip(image, ori_size)): + width, height = _ori_size + resized_image = F.interpolate(_image.unsqueeze(0), size=(height, width), mode='bilinear') + postprocessed_image.append(self.image_processor.postprocess(resized_image, output_type=output_type)[0]) + + image = postprocessed_image + + # Offload all models + self.maybe_free_model_hooks() + + if return_middle_statistics: + if not return_dict: + return image, middle_latents, log_probs, img_mask, l_effective_img_len, img_sizes, ref_latents, ref_img_mask, l_effective_ref_img_len, ref_img_sizes + else: + return FMPipelineOutput(images=image, middle_latents=middle_latents, log_probs=log_probs, img_mask=img_mask, l_effective_img_len=l_effective_img_len, img_sizes=img_sizes, ref_latents=ref_latents, ref_img_mask=ref_img_mask, l_effective_ref_img_len=l_effective_ref_img_len, ref_img_sizes=ref_img_sizes) + else: + if not return_dict: + return image + else: + return FMPipelineOutput(images=image) + + def _cfg_predict_sequential(self, + t, + prompt_embeds, + prompt_attention_mask, + negative_prompt_embeds, + negative_prompt_attention_mask, + freqs_cis, + latents, + img_mask, + l_effective_img_len, + img_sizes, + ref_latents, + ref_img_mask, + l_effective_ref_img_len, + ref_img_sizes, + ref_latents_N, + ref_img_mask_N, + l_effective_ref_img_len_N, + ref_img_sizes_N, + text_guidance_scale, + image_guidance_scale, + ): + model_kwargs = dict( + freqs_cis=freqs_cis, + flat_and_pad=False, + img_mask=img_mask, + l_effective_img_len=l_effective_img_len, + img_sizes=img_sizes, + ) + model_pred_kwargs = dict( + text_hidden_states=prompt_embeds, + text_attention_mask=prompt_attention_mask, + ref_image_hidden_states=ref_latents, + ref_img_mask=ref_img_mask, + l_effective_ref_img_len=l_effective_ref_img_len, + ref_img_sizes=ref_img_sizes, + ) + model_pred_ref_kwargs = dict( + text_hidden_states=negative_prompt_embeds, + text_attention_mask=negative_prompt_attention_mask, + ref_image_hidden_states=ref_latents, + ref_img_mask=ref_img_mask, + l_effective_ref_img_len=l_effective_ref_img_len, + ref_img_sizes=ref_img_sizes, + ) + model_pred_uncond_kwargs = dict( + text_hidden_states=negative_prompt_embeds, + text_attention_mask=negative_prompt_attention_mask, + ref_image_hidden_states=ref_latents_N, + ref_img_mask=ref_img_mask_N, + l_effective_ref_img_len=l_effective_ref_img_len_N, + ref_img_sizes=ref_img_sizes_N, + ) + + model_pred = self.predict( + t=t, + latents=latents, + **model_kwargs, + **model_pred_kwargs, + ) + if text_guidance_scale > 1.0 and image_guidance_scale > 1.0: + model_pred_ref = self.predict( + t=t, + latents=latents, + **model_kwargs, + **model_pred_ref_kwargs, + ) + model_pred_uncond = self.predict( + t=t, + latents=latents, + **model_kwargs, + **model_pred_uncond_kwargs, + ) + model_pred = model_pred_uncond + image_guidance_scale * (model_pred_ref - model_pred_uncond) + \ + text_guidance_scale * (model_pred - model_pred_ref) + + elif text_guidance_scale > 1.0: + model_pred_uncond = self.predict( + t=t, + latents=latents, + **model_kwargs, + **model_pred_uncond_kwargs, + ) + model_pred = model_pred_uncond + text_guidance_scale * (model_pred - model_pred_uncond) + return model_pred + + def _cfg_predict_parallel(self, + t, + prompt_embeds, + prompt_attention_mask, + negative_prompt_embeds, + negative_prompt_attention_mask, + freqs_cis, + latents, + img_mask, + l_effective_img_len, + img_sizes, + ref_latents, + ref_img_mask, + l_effective_ref_img_len, + ref_img_sizes, + ref_latents_N, + ref_img_mask_N, + l_effective_ref_img_len_N, + ref_img_sizes_N, + text_guidance_scale, + image_guidance_scale, + ): + model_kwargs = dict( + freqs_cis=freqs_cis, + flat_and_pad=False, + img_mask=img_mask, + l_effective_img_len=l_effective_img_len, + img_sizes=img_sizes, + ) + model_pred_kwargs = dict( + text_hidden_states=prompt_embeds, + text_attention_mask=prompt_attention_mask, + ref_image_hidden_states=ref_latents, + ref_img_mask=ref_img_mask, + l_effective_ref_img_len=l_effective_ref_img_len, + ref_img_sizes=ref_img_sizes, + ) + + if text_guidance_scale > 1.0 and image_guidance_scale > 1.0: + latents = torch.cat([latents, latents, latents], dim=0) + t = torch.cat([t, t, t], dim=0) + model_kwargs['img_mask'] = torch.cat([img_mask, img_mask, img_mask], dim=0) + model_kwargs['l_effective_img_len'] = l_effective_img_len * 3 + model_kwargs['img_sizes'] = img_sizes * 3 + + model_pred_kwargs['text_hidden_states'] = torch.cat([prompt_embeds, pad_to_length(negative_prompt_embeds, len=prompt_embeds.shape[1]), pad_to_length(negative_prompt_embeds, len=prompt_embeds.shape[1])], dim=0) + model_pred_kwargs['text_attention_mask'] = torch.cat([prompt_attention_mask, pad_to_length(negative_prompt_attention_mask, len=prompt_attention_mask.shape[1]), pad_to_length(negative_prompt_attention_mask, len=prompt_attention_mask.shape[1])], dim=0) + model_pred_kwargs['ref_image_hidden_states'] = torch.cat([ref_latents, ref_latents, pad_to_length(ref_latents_N, len=ref_latents.shape[1])], dim=0) + model_pred_kwargs['ref_img_mask'] = torch.cat([ref_img_mask, ref_img_mask, pad_to_length(ref_img_mask_N, len=ref_img_mask.shape[1])], dim=0) + model_pred_kwargs['l_effective_ref_img_len'] = l_effective_ref_img_len * 2 + l_effective_ref_img_len_N + model_pred_kwargs['ref_img_sizes'] = ref_img_sizes * 2 + ref_img_sizes_N + + elif text_guidance_scale > 1.0: + latents = torch.cat([latents, latents], dim=0) + t = torch.cat([t, t], dim=0) + model_kwargs['img_mask'] = torch.cat([img_mask, img_mask], dim=0) + model_kwargs['l_effective_img_len'] = l_effective_img_len * 2 + model_kwargs['img_sizes'] = img_sizes * 2 + + model_pred_kwargs['text_hidden_states'] = torch.cat([prompt_embeds, pad_to_length(negative_prompt_embeds, len=prompt_embeds.shape[1])], dim=0) + model_pred_kwargs['text_attention_mask'] = torch.cat([prompt_attention_mask, pad_to_length(negative_prompt_attention_mask, len=prompt_attention_mask.shape[1])], dim=0) + model_pred_kwargs['ref_image_hidden_states'] = torch.cat([ref_latents, pad_to_length(ref_latents_N, len=ref_latents.shape[1])], dim=0) + model_pred_kwargs['ref_img_mask'] = torch.cat([ref_img_mask, pad_to_length(ref_img_mask_N, len=ref_img_mask.shape[1])], dim=0) + model_pred_kwargs['l_effective_ref_img_len'] = l_effective_ref_img_len + l_effective_ref_img_len_N + model_pred_kwargs['ref_img_sizes'] = ref_img_sizes + ref_img_sizes_N + + model_pred = self.predict( + t=t, + latents=latents, + **model_kwargs, + **model_pred_kwargs, + ) + + if text_guidance_scale > 1.0 and image_guidance_scale > 1.0: + model_pred, model_pred_ref, model_pred_uncond = model_pred.chunk(3) + model_pred = model_pred_uncond + image_guidance_scale * (model_pred_ref - model_pred_uncond) + \ + text_guidance_scale * (model_pred - model_pred_ref) + elif text_guidance_scale > 1.0: + model_pred, model_pred_uncond = model_pred.chunk(2) + model_pred = model_pred_uncond + text_guidance_scale * (model_pred - model_pred_uncond) + + return model_pred + + def cfg_predict(self, + t, + prompt_embeds, + prompt_attention_mask, + negative_prompt_embeds, + negative_prompt_attention_mask, + freqs_cis, + latents, + img_mask, + l_effective_img_len, + img_sizes, + ref_latents, + ref_img_mask, + l_effective_ref_img_len, + ref_img_sizes, + ref_latents_N, + ref_img_mask_N, + l_effective_ref_img_len_N, + ref_img_sizes_N, + text_guidance_scale, + image_guidance_scale,): + if self.enable_parallel_cfg: + return self._cfg_predict_parallel( + t=t, + prompt_embeds=prompt_embeds, + prompt_attention_mask=prompt_attention_mask, + negative_prompt_embeds=negative_prompt_embeds, + negative_prompt_attention_mask=negative_prompt_attention_mask, + freqs_cis=freqs_cis, + latents=latents, + img_mask=img_mask, + l_effective_img_len=l_effective_img_len, + img_sizes=img_sizes, + ref_latents=ref_latents, + ref_img_mask=ref_img_mask, + l_effective_ref_img_len=l_effective_ref_img_len, + ref_img_sizes=ref_img_sizes, + ref_latents_N=ref_latents_N, + ref_img_mask_N=ref_img_mask_N, + l_effective_ref_img_len_N=l_effective_ref_img_len_N, + ref_img_sizes_N=ref_img_sizes_N, + text_guidance_scale=text_guidance_scale, + image_guidance_scale=image_guidance_scale, + ) + else: + return self._cfg_predict_sequential( + t=t, + prompt_embeds=prompt_embeds, + prompt_attention_mask=prompt_attention_mask, + negative_prompt_embeds=negative_prompt_embeds, + negative_prompt_attention_mask=negative_prompt_attention_mask, + freqs_cis=freqs_cis, + latents=latents, + img_mask=img_mask, + l_effective_img_len=l_effective_img_len, + img_sizes=img_sizes, + ref_latents=ref_latents, + ref_img_mask=ref_img_mask, + l_effective_ref_img_len=l_effective_ref_img_len, + ref_img_sizes=ref_img_sizes, + ref_latents_N=ref_latents_N, + ref_img_mask_N=ref_img_mask_N, + l_effective_ref_img_len_N=l_effective_ref_img_len_N, + ref_img_sizes_N=ref_img_sizes_N, + text_guidance_scale=text_guidance_scale, + image_guidance_scale=image_guidance_scale, + ) + + def processing( + self, + latents, + ref_latents, + prompt_embeds, + freqs_cis, + negative_prompt_embeds, + prompt_attention_mask, + negative_prompt_attention_mask, + num_inference_steps, + timesteps, + generator, + device, + dtype, + verbose, + step_func=None, + return_middle_statistics=False, + ): + timesteps, num_inference_steps = retrieve_timesteps( + self.scheduler, + num_inference_steps, + device, + timesteps, + num_tokens=[latent.shape[-2] * latent.shape[-1] for latent in latents] + ) + num_warmup_steps = max(len(timesteps[0]) - num_inference_steps * self.scheduler.order, 0) + self._num_timesteps = len(timesteps[0]) + + batch_size = len(latents) + ( + latents, + img_mask, + l_effective_img_len, + img_sizes, + ) = self.transformer.flat_and_pad_to_seq(latents, batch_size, device) + + ( + ref_latents, + ref_img_mask, + l_effective_ref_img_len, + ref_img_sizes, + ) = self.transformer.flat_and_pad_to_seq_ref_img(ref_latents, batch_size, dtype, device) + ( + ref_latents_N, + ref_img_mask_N, + l_effective_ref_img_len_N, + ref_img_sizes_N, + ) = self.transformer.flat_and_pad_to_seq_ref_img(None, batch_size, dtype, device) + + if return_middle_statistics: + latents_list = [latents] + log_probs_list = [] + + with self.progress_bar(total=num_inference_steps) as progress_bar: + for i in range(len(timesteps[0])): + t = timesteps[:, i] + text_guidance_scale = self.text_guidance_scale if self.cfg_range[0] <= i / len(timesteps[0]) <= self.cfg_range[1] else 1.0 + image_guidance_scale = self.image_guidance_scale if self.cfg_range[0] <= i / len(timesteps[0]) <= self.cfg_range[1] else 1.0 + + model_pred = self.cfg_predict( + t=t, + prompt_embeds=prompt_embeds, + prompt_attention_mask=prompt_attention_mask, + negative_prompt_embeds=negative_prompt_embeds, + negative_prompt_attention_mask=negative_prompt_attention_mask, + freqs_cis=freqs_cis, + latents=latents, + img_mask=img_mask, + l_effective_img_len=l_effective_img_len, + img_sizes=img_sizes, + ref_latents=ref_latents, + ref_img_mask=ref_img_mask, + l_effective_ref_img_len=l_effective_ref_img_len, + ref_img_sizes=ref_img_sizes, + ref_latents_N=ref_latents_N, + ref_img_mask_N=ref_img_mask_N, + l_effective_ref_img_len_N=l_effective_ref_img_len_N, + ref_img_sizes_N=ref_img_sizes_N, + text_guidance_scale=text_guidance_scale, + image_guidance_scale=image_guidance_scale + ) + latents_dtype = latents.dtype + if return_middle_statistics: + latents, log_probs = self.scheduler.step(model_pred.to(dtype=torch.float32), t, latents.to(dtype=torch.float32), generator=generator, return_log_prob=True, img_mask=img_mask, mixed_precision=self.mixed_precision, return_dict=False) + else: + latents = self.scheduler.step(model_pred, t, latents, generator=generator, mixed_precision=self.mixed_precision, return_dict=False)[0] + + if latents.dtype != latents_dtype and not self.mixed_precision: + if torch.backends.mps.is_available(): + # some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272 + latents = latents.to(latents_dtype) + + if return_middle_statistics: + latents_list.append(latents.clone()) + log_probs_list.append(log_probs) + + if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): + progress_bar.update(1) + + image = [] + for latent, img_len, img_size in zip(latents, l_effective_img_len, img_sizes): + height, width = img_size + p = self.transformer.config.patch_size + latent = rearrange(latent[:img_len], '(h w) (p1 p2 c) -> c (h p1) (w p2)', h=height // p, w=width // p, p1=p, p2=p) + latent = latent.to(dtype=dtype) + if self.vae.config.scaling_factor is not None: + latent = latent / self.vae.config.scaling_factor + if self.vae.config.shift_factor is not None: + latent = latent + self.vae.config.shift_factor + image.append(self.vae.decode(latent.unsqueeze(0), return_dict=False)[0].squeeze(0)) + + if return_middle_statistics: + return image, latents_list, log_probs_list, img_mask, l_effective_img_len, img_sizes, ref_latents, ref_img_mask, l_effective_ref_img_len, ref_img_sizes + else: + return image + + def predict( + self, + t, + latents, + text_hidden_states, + text_attention_mask, + ref_image_hidden_states, + freqs_cis, + **model_kwargs, + ): + # broadcast to batch dimension in a way that's compatible with ONNX/Core ML + if not self.mixed_precision: + timestep = t.expand(latents.shape[0]).to(latents.dtype) + else: + timestep = t + + # batch_size, num_channels_latents, height, width = latents.shape + from accelerate.utils import extract_model_from_parallel + if 'ref_image_hidden_states' in set(inspect.signature(extract_model_from_parallel(self.transformer).forward).parameters.keys()): + model_kwargs['ref_image_hidden_states'] = ref_image_hidden_states + + model_pred = self.transformer( + latents, + timestep, + text_hidden_states, + text_attention_mask, + freqs_cis, + **model_kwargs, + ) + return model_pred \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/pipelines/pipeline_utils.py b/examples/OmniGen2-RL/omnigen2/pipelines/pipeline_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..de31ff4e8627a93377e7c3c071162f6b395da688 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/pipelines/pipeline_utils.py @@ -0,0 +1,62 @@ +import torch + + +def get_pipeline_embeds(pipeline, prompt, negative_prompt, device): + """ Get pipeline embeds for prompts bigger than the maxlength of the pipe + :param pipeline: + :param prompt: + :param negative_prompt: + :param device: + :return: + """ + max_length = pipeline.tokenizer.model_max_length + + # simple way to determine length of tokens + # count_prompt = len(prompt.split(" ")) + # count_negative_prompt = len(negative_prompt.split(" ")) + + # create the tensor based on which prompt is longer + # if count_prompt >= count_negative_prompt: + input_ids = pipeline.tokenizer(prompt, return_tensors="pt", truncation=False, padding='longest').input_ids.to(device) + # input_ids = pipeline.tokenizer(prompt, padding="max_length", + # max_length=pipeline.tokenizer.model_max_length, + # truncation=True, + # return_tensors="pt",).input_ids.to(device) + shape_max_length = input_ids.shape[-1] + + if negative_prompt is not None: + negative_ids = pipeline.tokenizer(negative_prompt, truncation=True, padding="max_length", + max_length=shape_max_length, return_tensors="pt").input_ids.to(device) + + # else: + # negative_ids = pipeline.tokenizer(negative_prompt, return_tensors="pt", truncation=False).input_ids.to(device) + # shape_max_length = negative_ids.shape[-1] + # input_ids = pipeline.tokenizer(prompt, return_tensors="pt", truncation=False, padding="max_length", + # max_length=shape_max_length).input_ids.to(device) + + concat_embeds = [] + neg_embeds = [] + for i in range(0, shape_max_length, max_length): + if hasattr(pipeline.text_encoder.config, "use_attention_mask") and pipeline.text_encoder.config.use_attention_mask: + attention_mask = input_ids[:, i: i + max_length].attention_mask.to(device) + else: + attention_mask = None + concat_embeds.append(pipeline.text_encoder(input_ids[:, i: i + max_length], + attention_mask=attention_mask)[0]) + + if negative_prompt is not None: + if hasattr(pipeline.text_encoder.config, "use_attention_mask") and pipeline.text_encoder.config.use_attention_mask: + attention_mask = negative_ids[:, i: i + max_length].attention_mask.to(device) + else: + attention_mask = None + neg_embeds.append(pipeline.text_encoder(negative_ids[:, i: i + max_length], + attention_mask=attention_mask)[0]) + + concat_embeds = torch.cat(concat_embeds, dim=1) + + if negative_prompt is not None: + neg_embeds = torch.cat(neg_embeds, dim=1) + else: + neg_embeds = None + + return concat_embeds, neg_embeds diff --git a/examples/OmniGen2-RL/omnigen2/schedulers/__init__.py b/examples/OmniGen2-RL/omnigen2/schedulers/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/examples/OmniGen2-RL/omnigen2/schedulers/scheduling_dpmsolver_multistep.py b/examples/OmniGen2-RL/omnigen2/schedulers/scheduling_dpmsolver_multistep.py new file mode 100644 index 0000000000000000000000000000000000000000..f2dcced85ac72e1f1f4fbc7f2038b22c8e16341f --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/schedulers/scheduling_dpmsolver_multistep.py @@ -0,0 +1,1053 @@ +# Copyright 2024 TSAIL Team and The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# DISCLAIMER: This file is strongly influenced by https://github.com/LuChengTHU/dpm-solver + +import math +from typing import List, Optional, Tuple, Union + +import numpy as np +import torch + +from diffusers.configuration_utils import ConfigMixin, register_to_config +from diffusers.utils import deprecate, is_scipy_available +from diffusers.utils.torch_utils import randn_tensor +from diffusers.schedulers.scheduling_utils import KarrasDiffusionSchedulers, SchedulerMixin, SchedulerOutput + + +if is_scipy_available(): + import scipy.stats + + +# Copied from diffusers.schedulers.scheduling_ddpm.betas_for_alpha_bar +def betas_for_alpha_bar( + num_diffusion_timesteps, + max_beta=0.999, + alpha_transform_type="cosine", +): + """ + Create a beta schedule that discretizes the given alpha_t_bar function, which defines the cumulative product of + (1-beta) over time from t = [0,1]. + + Contains a function alpha_bar that takes an argument t and transforms it to the cumulative product of (1-beta) up + to that part of the diffusion process. + + + Args: + num_diffusion_timesteps (`int`): the number of betas to produce. + max_beta (`float`): the maximum beta to use; use values lower than 1 to + prevent singularities. + alpha_transform_type (`str`, *optional*, default to `cosine`): the type of noise schedule for alpha_bar. + Choose from `cosine` or `exp` + + Returns: + betas (`np.ndarray`): the betas used by the scheduler to step the model outputs + """ + if alpha_transform_type == "cosine": + + def alpha_bar_fn(t): + return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2 + + elif alpha_transform_type == "exp": + + def alpha_bar_fn(t): + return math.exp(t * -12.0) + + else: + raise ValueError(f"Unsupported alpha_transform_type: {alpha_transform_type}") + + betas = [] + for i in range(num_diffusion_timesteps): + t1 = i / num_diffusion_timesteps + t2 = (i + 1) / num_diffusion_timesteps + betas.append(min(1 - alpha_bar_fn(t2) / alpha_bar_fn(t1), max_beta)) + return torch.tensor(betas, dtype=torch.float32) + + +# Copied from diffusers.schedulers.scheduling_ddim.rescale_zero_terminal_snr +def rescale_zero_terminal_snr(betas): + """ + Rescales betas to have zero terminal SNR Based on https://arxiv.org/pdf/2305.08891.pdf (Algorithm 1) + + + Args: + betas (`torch.Tensor`): + the betas that the scheduler is being initialized with. + + Returns: + `torch.Tensor`: rescaled betas with zero terminal SNR + """ + # Convert betas to alphas_bar_sqrt + alphas = 1.0 - betas + alphas_cumprod = torch.cumprod(alphas, dim=0) + alphas_bar_sqrt = alphas_cumprod.sqrt() + + # Store old values. + alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone() + alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone() + + # Shift so the last timestep is zero. + alphas_bar_sqrt -= alphas_bar_sqrt_T + + # Scale so the first timestep is back to the old value. + alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T) + + # Convert alphas_bar_sqrt to betas + alphas_bar = alphas_bar_sqrt**2 # Revert sqrt + alphas = alphas_bar[1:] / alphas_bar[:-1] # Revert cumprod + alphas = torch.cat([alphas_bar[0:1], alphas]) + betas = 1 - alphas + + return betas + + +class DPMSolverMultistepScheduler(SchedulerMixin, ConfigMixin): + """ + `DPMSolverMultistepScheduler` is a fast dedicated high-order solver for diffusion ODEs. + + This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic + methods the library implements for all schedulers such as loading and saving. + + Args: + num_train_timesteps (`int`, defaults to 1000): + The number of diffusion steps to train the model. + beta_start (`float`, defaults to 0.0001): + The starting `beta` value of inference. + beta_end (`float`, defaults to 0.02): + The final `beta` value. + beta_schedule (`str`, defaults to `"linear"`): + The beta schedule, a mapping from a beta range to a sequence of betas for stepping the model. Choose from + `linear`, `scaled_linear`, or `squaredcos_cap_v2`. + trained_betas (`np.ndarray`, *optional*): + Pass an array of betas directly to the constructor to bypass `beta_start` and `beta_end`. + solver_order (`int`, defaults to 2): + The DPMSolver order which can be `1` or `2` or `3`. It is recommended to use `solver_order=2` for guided + sampling, and `solver_order=3` for unconditional sampling. + prediction_type (`str`, defaults to `epsilon`, *optional*): + Prediction type of the scheduler function; can be `epsilon` (predicts the noise of the diffusion process), + `sample` (directly predicts the noisy sample), `v_prediction` (see section 2.4 of [Imagen + Video](https://imagen.research.google/video/paper.pdf) paper), or `flow_prediction`. + thresholding (`bool`, defaults to `False`): + Whether to use the "dynamic thresholding" method. This is unsuitable for latent-space diffusion models such + as Stable Diffusion. + dynamic_thresholding_ratio (`float`, defaults to 0.995): + The ratio for the dynamic thresholding method. Valid only when `thresholding=True`. + sample_max_value (`float`, defaults to 1.0): + The threshold value for dynamic thresholding. Valid only when `thresholding=True` and + `algorithm_type="dpmsolver++"`. + algorithm_type (`str`, defaults to `dpmsolver++`): + Algorithm type for the solver; can be `dpmsolver`, `dpmsolver++`, `sde-dpmsolver` or `sde-dpmsolver++`. The + `dpmsolver` type implements the algorithms in the [DPMSolver](https://huggingface.co/papers/2206.00927) + paper, and the `dpmsolver++` type implements the algorithms in the + [DPMSolver++](https://huggingface.co/papers/2211.01095) paper. It is recommended to use `dpmsolver++` or + `sde-dpmsolver++` with `solver_order=2` for guided sampling like in Stable Diffusion. + solver_type (`str`, defaults to `midpoint`): + Solver type for the second-order solver; can be `midpoint` or `heun`. The solver type slightly affects the + sample quality, especially for a small number of steps. It is recommended to use `midpoint` solvers. + lower_order_final (`bool`, defaults to `True`): + Whether to use lower-order solvers in the final steps. Only valid for < 15 inference steps. This can + stabilize the sampling of DPMSolver for steps < 15, especially for steps <= 10. + euler_at_final (`bool`, defaults to `False`): + Whether to use Euler's method in the final step. It is a trade-off between numerical stability and detail + richness. This can stabilize the sampling of the SDE variant of DPMSolver for small number of inference + steps, but sometimes may result in blurring. + use_karras_sigmas (`bool`, *optional*, defaults to `False`): + Whether to use Karras sigmas for step sizes in the noise schedule during the sampling process. If `True`, + the sigmas are determined according to a sequence of noise levels {ฯƒi}. + use_exponential_sigmas (`bool`, *optional*, defaults to `False`): + Whether to use exponential sigmas for step sizes in the noise schedule during the sampling process. + use_beta_sigmas (`bool`, *optional*, defaults to `False`): + Whether to use beta sigmas for step sizes in the noise schedule during the sampling process. Refer to [Beta + Sampling is All You Need](https://huggingface.co/papers/2407.12173) for more information. + use_lu_lambdas (`bool`, *optional*, defaults to `False`): + Whether to use the uniform-logSNR for step sizes proposed by Lu's DPM-Solver in the noise schedule during + the sampling process. If `True`, the sigmas and time steps are determined according to a sequence of + `lambda(t)`. + use_flow_sigmas (`bool`, *optional*, defaults to `False`): + Whether to use flow sigmas for step sizes in the noise schedule during the sampling process. + flow_shift (`float`, *optional*, defaults to 1.0): + The shift value for the timestep schedule for flow matching. + final_sigmas_type (`str`, defaults to `"zero"`): + The final `sigma` value for the noise schedule during the sampling process. If `"sigma_min"`, the final + sigma is the same as the last sigma in the training schedule. If `zero`, the final sigma is set to 0. + lambda_min_clipped (`float`, defaults to `-inf`): + Clipping threshold for the minimum value of `lambda(t)` for numerical stability. This is critical for the + cosine (`squaredcos_cap_v2`) noise schedule. + variance_type (`str`, *optional*): + Set to "learned" or "learned_range" for diffusion models that predict variance. If set, the model's output + contains the predicted Gaussian variance. + timestep_spacing (`str`, defaults to `"linspace"`): + The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and + Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information. + steps_offset (`int`, defaults to 0): + An offset added to the inference steps, as required by some model families. + rescale_betas_zero_snr (`bool`, defaults to `False`): + Whether to rescale the betas to have zero terminal SNR. This enables the model to generate very bright and + dark samples instead of limiting it to samples with medium brightness. Loosely related to + [`--offset_noise`](https://github.com/huggingface/diffusers/blob/74fd735eb073eb1d774b1ab4154a0876eb82f055/examples/dreambooth/train_dreambooth.py#L506). + """ + + _compatibles = [e.name for e in KarrasDiffusionSchedulers] + order = 1 + + @register_to_config + def __init__( + self, + num_train_timesteps: int = 1000, + beta_start: float = 0.0001, + beta_end: float = 0.02, + beta_schedule: str = "linear", + trained_betas: Optional[Union[np.ndarray, List[float]]] = None, + solver_order: int = 2, + prediction_type: str = "epsilon", + thresholding: bool = False, + dynamic_thresholding_ratio: float = 0.995, + sample_max_value: float = 1.0, + algorithm_type: str = "dpmsolver++", + solver_type: str = "midpoint", + lower_order_final: bool = True, + euler_at_final: bool = False, + final_sigmas_type: str = 'zero', + dynamic_time_shift: bool = True + ): + if algorithm_type in ["dpmsolver", "sde-dpmsolver"]: + deprecation_message = f"algorithm_type {algorithm_type} is deprecated and will be removed in a future version. Choose from `dpmsolver++` or `sde-dpmsolver++` instead" + deprecate("algorithm_types dpmsolver and sde-dpmsolver", "1.0.0", deprecation_message) + + if trained_betas is not None: + self.betas = torch.tensor(trained_betas, dtype=torch.float32) + elif beta_schedule == "linear": + self.betas = torch.linspace(beta_start, beta_end, num_train_timesteps, dtype=torch.float32) + elif beta_schedule == "scaled_linear": + # this schedule is very specific to the latent diffusion model. + self.betas = torch.linspace(beta_start**0.5, beta_end**0.5, num_train_timesteps, dtype=torch.float32) ** 2 + elif beta_schedule == "squaredcos_cap_v2": + # Glide cosine schedule + self.betas = betas_for_alpha_bar(num_train_timesteps) + else: + raise NotImplementedError(f"{beta_schedule} is not implemented for {self.__class__}") + self.alphas = 1.0 - self.betas + self.alphas_cumprod = torch.cumprod(self.alphas, dim=0) + + # Currently we only support VP-type noise schedule + self.alpha_t = torch.sqrt(self.alphas_cumprod) + self.sigma_t = torch.sqrt(1 - self.alphas_cumprod) + self.lambda_t = torch.log(self.alpha_t) - torch.log(self.sigma_t) + self.sigmas = ((1 - self.alphas_cumprod) / self.alphas_cumprod) ** 0.5 + + # standard deviation of the initial noise distribution + self.init_noise_sigma = 1.0 + + # settings for DPM-Solver + if algorithm_type not in ["dpmsolver", "dpmsolver++", "sde-dpmsolver", "sde-dpmsolver++"]: + if algorithm_type == "deis": + self.register_to_config(algorithm_type="dpmsolver++") + else: + raise NotImplementedError(f"{algorithm_type} is not implemented for {self.__class__}") + + if solver_type not in ["midpoint", "heun"]: + if solver_type in ["logrho", "bh1", "bh2"]: + self.register_to_config(solver_type="midpoint") + else: + raise NotImplementedError(f"{solver_type} is not implemented for {self.__class__}") + + # if algorithm_type not in ["dpmsolver++", "sde-dpmsolver++"] and final_sigmas_type == "zero": + # raise ValueError( + # f"`final_sigmas_type` {final_sigmas_type} is not supported for `algorithm_type` {algorithm_type}. Please choose `sigma_min` instead." + # ) + + # setable values + self.num_inference_steps = None + timesteps = np.linspace(0, num_train_timesteps - 1, num_train_timesteps, dtype=np.float32)[::-1].copy() + self.timesteps = torch.from_numpy(timesteps) + self.model_outputs = [None] * solver_order + self.lower_order_nums = 0 + self._step_index = None + self._begin_index = None + self.sigmas = self.sigmas.to("cpu") # to avoid too much CPU/GPU communication + + @property + def step_index(self): + """ + The index counter for current timestep. It will increase 1 after each scheduler step. + """ + return self._step_index + + @property + def begin_index(self): + """ + The index for the first timestep. It should be set from pipeline with `set_begin_index` method. + """ + return self._begin_index + + def set_begin_index(self, begin_index: int = 0): + """ + Sets the begin index for the scheduler. This function should be run from pipeline before the inference. + + Args: + begin_index (`int`): + The begin index for the scheduler. + """ + self._begin_index = begin_index + + def set_timesteps( + self, + num_inference_steps: int = None, + device: Union[str, torch.device] = None, + timesteps: Optional[List[int]] = None, + num_tokens: Optional[int] = None + ): + if timesteps is None: + self.num_inference_steps = num_inference_steps + timesteps = np.linspace(0, 1, num_inference_steps + 1, dtype=np.float32)[:-1] + if self.config.dynamic_time_shift and num_tokens is not None: + m = np.sqrt(num_tokens) / 40 # when input resolution is 320 * 320, m = 1, when input resolution is 1024 * 1024, m = 3.2 + timesteps = timesteps / (m - m * timesteps + timesteps) + + timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32, device=device) + sigmas = torch.cat([1 - timesteps, torch.zeros(1, device=timesteps.device)]) + + self.sigmas = sigmas + self.timesteps = timesteps + + self.num_inference_steps = len(timesteps) + + self.model_outputs = [ + None, + ] * self.config.solver_order + self.lower_order_nums = 0 + + # add an index counter for schedulers that allow duplicated timesteps + self._step_index = None + self._begin_index = None + self.sigmas = self.sigmas.to("cpu") # to avoid too much CPU/GPU communication + + # Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler._threshold_sample + def _threshold_sample(self, sample: torch.Tensor) -> torch.Tensor: + """ + "Dynamic thresholding: At each sampling step we set s to a certain percentile absolute pixel value in xt0 (the + prediction of x_0 at timestep t), and if s > 1, then we threshold xt0 to the range [-s, s] and then divide by + s. Dynamic thresholding pushes saturated pixels (those near -1 and 1) inwards, thereby actively preventing + pixels from saturation at each step. We find that dynamic thresholding results in significantly better + photorealism as well as better image-text alignment, especially when using very large guidance weights." + + https://arxiv.org/abs/2205.11487 + """ + dtype = sample.dtype + batch_size, channels, *remaining_dims = sample.shape + + if dtype not in (torch.float32, torch.float64): + sample = sample.float() # upcast for quantile calculation, and clamp not implemented for cpu half + + # Flatten sample for doing quantile calculation along each image + sample = sample.reshape(batch_size, channels * np.prod(remaining_dims)) + + abs_sample = sample.abs() # "a certain percentile absolute pixel value" + + s = torch.quantile(abs_sample, self.config.dynamic_thresholding_ratio, dim=1) + s = torch.clamp( + s, min=1, max=self.config.sample_max_value + ) # When clamped to min=1, equivalent to standard clipping to [-1, 1] + s = s.unsqueeze(1) # (batch_size, 1) because clamp will broadcast along dim=0 + sample = torch.clamp(sample, -s, s) / s # "we threshold xt0 to the range [-s, s] and then divide by s" + + sample = sample.reshape(batch_size, channels, *remaining_dims) + sample = sample.to(dtype) + + return sample + + # Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._sigma_to_t + def _sigma_to_t(self, sigma, log_sigmas): + # get log sigma + log_sigma = np.log(np.maximum(sigma, 1e-10)) + + # get distribution + dists = log_sigma - log_sigmas[:, np.newaxis] + + # get sigmas range + low_idx = np.cumsum((dists >= 0), axis=0).argmax(axis=0).clip(max=log_sigmas.shape[0] - 2) + high_idx = low_idx + 1 + + low = log_sigmas[low_idx] + high = log_sigmas[high_idx] + + # interpolate sigmas + w = (low - log_sigma) / (low - high) + w = np.clip(w, 0, 1) + + # transform interpolation to time range + t = (1 - w) * low_idx + w * high_idx + t = t.reshape(sigma.shape) + return t + + def _sigma_to_alpha_sigma_t(self, sigma): + alpha_t = 1 - sigma + sigma_t = sigma + + return alpha_t, sigma_t + + # Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_karras + def _convert_to_karras(self, in_sigmas: torch.Tensor, num_inference_steps) -> torch.Tensor: + """Constructs the noise schedule of Karras et al. (2022).""" + + # Hack to make sure that other schedulers which copy this function don't break + # TODO: Add this logic to the other schedulers + if hasattr(self.config, "sigma_min"): + sigma_min = self.config.sigma_min + else: + sigma_min = None + + if hasattr(self.config, "sigma_max"): + sigma_max = self.config.sigma_max + else: + sigma_max = None + + sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item() + sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item() + + rho = 7.0 # 7.0 is the value used in the paper + ramp = np.linspace(0, 1, num_inference_steps) + min_inv_rho = sigma_min ** (1 / rho) + max_inv_rho = sigma_max ** (1 / rho) + sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho + return sigmas + + def _convert_to_lu(self, in_lambdas: torch.Tensor, num_inference_steps) -> torch.Tensor: + """Constructs the noise schedule of Lu et al. (2022).""" + + lambda_min: float = in_lambdas[-1].item() + lambda_max: float = in_lambdas[0].item() + + rho = 1.0 # 1.0 is the value used in the paper + ramp = np.linspace(0, 1, num_inference_steps) + min_inv_rho = lambda_min ** (1 / rho) + max_inv_rho = lambda_max ** (1 / rho) + lambdas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho + return lambdas + + # Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_exponential + def _convert_to_exponential(self, in_sigmas: torch.Tensor, num_inference_steps: int) -> torch.Tensor: + """Constructs an exponential noise schedule.""" + + # Hack to make sure that other schedulers which copy this function don't break + # TODO: Add this logic to the other schedulers + if hasattr(self.config, "sigma_min"): + sigma_min = self.config.sigma_min + else: + sigma_min = None + + if hasattr(self.config, "sigma_max"): + sigma_max = self.config.sigma_max + else: + sigma_max = None + + sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item() + sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item() + + sigmas = np.exp(np.linspace(math.log(sigma_max), math.log(sigma_min), num_inference_steps)) + return sigmas + + # Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_beta + def _convert_to_beta( + self, in_sigmas: torch.Tensor, num_inference_steps: int, alpha: float = 0.6, beta: float = 0.6 + ) -> torch.Tensor: + """From "Beta Sampling is All You Need" [arXiv:2407.12173] (Lee et. al, 2024)""" + + # Hack to make sure that other schedulers which copy this function don't break + # TODO: Add this logic to the other schedulers + if hasattr(self.config, "sigma_min"): + sigma_min = self.config.sigma_min + else: + sigma_min = None + + if hasattr(self.config, "sigma_max"): + sigma_max = self.config.sigma_max + else: + sigma_max = None + + sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item() + sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item() + + sigmas = np.array( + [ + sigma_min + (ppf * (sigma_max - sigma_min)) + for ppf in [ + scipy.stats.beta.ppf(timestep, alpha, beta) + for timestep in 1 - np.linspace(0, 1, num_inference_steps) + ] + ] + ) + return sigmas + + def convert_model_output( + self, + model_output: torch.Tensor, + *args, + sample: torch.Tensor = None, + **kwargs, + ) -> torch.Tensor: + """ + Convert the model output to the corresponding type the DPMSolver/DPMSolver++ algorithm needs. DPM-Solver is + designed to discretize an integral of the noise prediction model, and DPM-Solver++ is designed to discretize an + integral of the data prediction model. + + + + The algorithm and model type are decoupled. You can use either DPMSolver or DPMSolver++ for both noise + prediction and data prediction models. + + + + Args: + model_output (`torch.Tensor`): + The direct output from the learned diffusion model. + sample (`torch.Tensor`): + A current instance of a sample created by the diffusion process. + + Returns: + `torch.Tensor`: + The converted model output. + """ + timestep = args[0] if len(args) > 0 else kwargs.pop("timestep", None) + if sample is None: + if len(args) > 1: + sample = args[1] + else: + raise ValueError("missing `sample` as a required keyward argument") + if timestep is not None: + deprecate( + "timesteps", + "1.0.0", + "Passing `timesteps` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`", + ) + + # DPM-Solver++ needs to solve an integral of the data prediction model. + if self.config.algorithm_type in ["dpmsolver++", "sde-dpmsolver++"]: + if self.config.prediction_type == "epsilon": + # DPM-Solver and DPM-Solver++ only need the "mean" output. + if self.config.variance_type in ["learned", "learned_range"]: + model_output = model_output[:, :3] + sigma = self.sigmas[self.step_index] + alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma) + x0_pred = (sample - sigma_t * model_output) / alpha_t + elif self.config.prediction_type == "sample": + x0_pred = model_output + elif self.config.prediction_type == "v_prediction": + sigma = self.sigmas[self.step_index] + alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma) + x0_pred = alpha_t * sample - sigma_t * model_output + elif self.config.prediction_type == "flow_prediction": + sigma_t = self.sigmas[self.step_index] + x0_pred = sample + sigma_t * model_output + else: + raise ValueError( + f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample`, " + "`v_prediction`, or `flow_prediction` for the DPMSolverMultistepScheduler." + ) + + if self.config.thresholding: + x0_pred = self._threshold_sample(x0_pred) + + return x0_pred + + # DPM-Solver needs to solve an integral of the noise prediction model. + elif self.config.algorithm_type in ["dpmsolver", "sde-dpmsolver"]: + if self.config.prediction_type == "epsilon": + # DPM-Solver and DPM-Solver++ only need the "mean" output. + if self.config.variance_type in ["learned", "learned_range"]: + epsilon = model_output[:, :3] + else: + epsilon = model_output + elif self.config.prediction_type == "sample": + sigma = self.sigmas[self.step_index] + alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma) + epsilon = (sample - alpha_t * model_output) / sigma_t + elif self.config.prediction_type == "v_prediction": + sigma = self.sigmas[self.step_index] + alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma) + epsilon = alpha_t * model_output + sigma_t * sample + else: + raise ValueError( + f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample`, or" + " `v_prediction` for the DPMSolverMultistepScheduler." + ) + + if self.config.thresholding: + sigma = self.sigmas[self.step_index] + alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma) + x0_pred = (sample - sigma_t * epsilon) / alpha_t + x0_pred = self._threshold_sample(x0_pred) + epsilon = (sample - alpha_t * x0_pred) / sigma_t + + return epsilon + + def dpm_solver_first_order_update( + self, + model_output: torch.Tensor, + *args, + sample: torch.Tensor = None, + noise: Optional[torch.Tensor] = None, + **kwargs, + ) -> torch.Tensor: + """ + One step for the first-order DPMSolver (equivalent to DDIM). + + Args: + model_output (`torch.Tensor`): + The direct output from the learned diffusion model. + sample (`torch.Tensor`): + A current instance of a sample created by the diffusion process. + + Returns: + `torch.Tensor`: + The sample tensor at the previous timestep. + """ + timestep = args[0] if len(args) > 0 else kwargs.pop("timestep", None) + prev_timestep = args[1] if len(args) > 1 else kwargs.pop("prev_timestep", None) + if sample is None: + if len(args) > 2: + sample = args[2] + else: + raise ValueError(" missing `sample` as a required keyward argument") + if timestep is not None: + deprecate( + "timesteps", + "1.0.0", + "Passing `timesteps` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`", + ) + + if prev_timestep is not None: + deprecate( + "prev_timestep", + "1.0.0", + "Passing `prev_timestep` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`", + ) + + sigma_t, sigma_s = self.sigmas[self.step_index + 1], self.sigmas[self.step_index] + alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t) + alpha_s, sigma_s = self._sigma_to_alpha_sigma_t(sigma_s) + lambda_t = torch.log(alpha_t) - torch.log(sigma_t) + lambda_s = torch.log(alpha_s) - torch.log(sigma_s) + + h = lambda_t - lambda_s + if self.config.algorithm_type == "dpmsolver++": + x_t = (sigma_t / sigma_s) * sample - (alpha_t * (torch.exp(-h) - 1.0)) * model_output + elif self.config.algorithm_type == "dpmsolver": + x_t = (alpha_t / alpha_s) * sample - (sigma_t * (torch.exp(h) - 1.0)) * model_output + elif self.config.algorithm_type == "sde-dpmsolver++": + print(sigma_t * torch.sqrt(1.0 - torch.exp(-2 * h))) + assert noise is not None + x_t = ( + (sigma_t / sigma_s * torch.exp(-h)) * sample + + (alpha_t * (1 - torch.exp(-2.0 * h))) * model_output + + sigma_t * torch.sqrt(1.0 - torch.exp(-2 * h)) * noise + ) + elif self.config.algorithm_type == "sde-dpmsolver": + assert noise is not None + x_t = ( + (alpha_t / alpha_s) * sample + - 2.0 * (sigma_t * (torch.exp(h) - 1.0)) * model_output + + sigma_t * torch.sqrt(torch.exp(2 * h) - 1.0) * noise + ) + return x_t + + def multistep_dpm_solver_second_order_update( + self, + model_output_list: List[torch.Tensor], + *args, + sample: torch.Tensor = None, + noise: Optional[torch.Tensor] = None, + **kwargs, + ) -> torch.Tensor: + """ + One step for the second-order multistep DPMSolver. + + Args: + model_output_list (`List[torch.Tensor]`): + The direct outputs from learned diffusion model at current and latter timesteps. + sample (`torch.Tensor`): + A current instance of a sample created by the diffusion process. + + Returns: + `torch.Tensor`: + The sample tensor at the previous timestep. + """ + timestep_list = args[0] if len(args) > 0 else kwargs.pop("timestep_list", None) + prev_timestep = args[1] if len(args) > 1 else kwargs.pop("prev_timestep", None) + if sample is None: + if len(args) > 2: + sample = args[2] + else: + raise ValueError(" missing `sample` as a required keyward argument") + if timestep_list is not None: + deprecate( + "timestep_list", + "1.0.0", + "Passing `timestep_list` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`", + ) + + if prev_timestep is not None: + deprecate( + "prev_timestep", + "1.0.0", + "Passing `prev_timestep` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`", + ) + + sigma_t, sigma_s0, sigma_s1 = ( + self.sigmas[self.step_index + 1], + self.sigmas[self.step_index], + self.sigmas[self.step_index - 1], + ) + + alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t) + alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0) + alpha_s1, sigma_s1 = self._sigma_to_alpha_sigma_t(sigma_s1) + + lambda_t = torch.log(alpha_t) - torch.log(sigma_t) + lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0) + lambda_s1 = torch.log(alpha_s1) - torch.log(sigma_s1) + + m0, m1 = model_output_list[-1], model_output_list[-2] + + h, h_0 = lambda_t - lambda_s0, lambda_s0 - lambda_s1 + r0 = h_0 / h + D0, D1 = m0, (1.0 / r0) * (m0 - m1) + if self.config.algorithm_type == "dpmsolver++": + # See https://arxiv.org/abs/2211.01095 for detailed derivations + if self.config.solver_type == "midpoint": + x_t = ( + (sigma_t / sigma_s0) * sample + - (alpha_t * (torch.exp(-h) - 1.0)) * D0 + - 0.5 * (alpha_t * (torch.exp(-h) - 1.0)) * D1 + ) + elif self.config.solver_type == "heun": + x_t = ( + (sigma_t / sigma_s0) * sample + - (alpha_t * (torch.exp(-h) - 1.0)) * D0 + + (alpha_t * ((torch.exp(-h) - 1.0) / h + 1.0)) * D1 + ) + elif self.config.algorithm_type == "dpmsolver": + # See https://arxiv.org/abs/2206.00927 for detailed derivations + if self.config.solver_type == "midpoint": + x_t = ( + (alpha_t / alpha_s0) * sample + - (sigma_t * (torch.exp(h) - 1.0)) * D0 + - 0.5 * (sigma_t * (torch.exp(h) - 1.0)) * D1 + ) + elif self.config.solver_type == "heun": + x_t = ( + (alpha_t / alpha_s0) * sample + - (sigma_t * (torch.exp(h) - 1.0)) * D0 + - (sigma_t * ((torch.exp(h) - 1.0) / h - 1.0)) * D1 + ) + elif self.config.algorithm_type == "sde-dpmsolver++": + assert noise is not None + if self.config.solver_type == "midpoint": + x_t = ( + (sigma_t / sigma_s0 * torch.exp(-h)) * sample + + (alpha_t * (1 - torch.exp(-2.0 * h))) * D0 + + 0.5 * (alpha_t * (1 - torch.exp(-2.0 * h))) * D1 + + sigma_t * torch.sqrt(1.0 - torch.exp(-2 * h)) * noise + ) + elif self.config.solver_type == "heun": + x_t = ( + (sigma_t / sigma_s0 * torch.exp(-h)) * sample + + (alpha_t * (1 - torch.exp(-2.0 * h))) * D0 + + (alpha_t * ((1.0 - torch.exp(-2.0 * h)) / (-2.0 * h) + 1.0)) * D1 + + sigma_t * torch.sqrt(1.0 - torch.exp(-2 * h)) * noise + ) + elif self.config.algorithm_type == "sde-dpmsolver": + assert noise is not None + if self.config.solver_type == "midpoint": + x_t = ( + (alpha_t / alpha_s0) * sample + - 2.0 * (sigma_t * (torch.exp(h) - 1.0)) * D0 + - (sigma_t * (torch.exp(h) - 1.0)) * D1 + + sigma_t * torch.sqrt(torch.exp(2 * h) - 1.0) * noise + ) + elif self.config.solver_type == "heun": + x_t = ( + (alpha_t / alpha_s0) * sample + - 2.0 * (sigma_t * (torch.exp(h) - 1.0)) * D0 + - 2.0 * (sigma_t * ((torch.exp(h) - 1.0) / h - 1.0)) * D1 + + sigma_t * torch.sqrt(torch.exp(2 * h) - 1.0) * noise + ) + return x_t + + def multistep_dpm_solver_third_order_update( + self, + model_output_list: List[torch.Tensor], + *args, + sample: torch.Tensor = None, + noise: Optional[torch.Tensor] = None, + **kwargs, + ) -> torch.Tensor: + """ + One step for the third-order multistep DPMSolver. + + Args: + model_output_list (`List[torch.Tensor]`): + The direct outputs from learned diffusion model at current and latter timesteps. + sample (`torch.Tensor`): + A current instance of a sample created by diffusion process. + + Returns: + `torch.Tensor`: + The sample tensor at the previous timestep. + """ + + timestep_list = args[0] if len(args) > 0 else kwargs.pop("timestep_list", None) + prev_timestep = args[1] if len(args) > 1 else kwargs.pop("prev_timestep", None) + if sample is None: + if len(args) > 2: + sample = args[2] + else: + raise ValueError(" missing`sample` as a required keyward argument") + if timestep_list is not None: + deprecate( + "timestep_list", + "1.0.0", + "Passing `timestep_list` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`", + ) + + if prev_timestep is not None: + deprecate( + "prev_timestep", + "1.0.0", + "Passing `prev_timestep` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`", + ) + + sigma_t, sigma_s0, sigma_s1, sigma_s2 = ( + self.sigmas[self.step_index + 1], + self.sigmas[self.step_index], + self.sigmas[self.step_index - 1], + self.sigmas[self.step_index - 2], + ) + + alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma_t) + alpha_s0, sigma_s0 = self._sigma_to_alpha_sigma_t(sigma_s0) + alpha_s1, sigma_s1 = self._sigma_to_alpha_sigma_t(sigma_s1) + alpha_s2, sigma_s2 = self._sigma_to_alpha_sigma_t(sigma_s2) + + lambda_t = torch.log(alpha_t) - torch.log(sigma_t) + lambda_s0 = torch.log(alpha_s0) - torch.log(sigma_s0) + lambda_s1 = torch.log(alpha_s1) - torch.log(sigma_s1) + lambda_s2 = torch.log(alpha_s2) - torch.log(sigma_s2) + + m0, m1, m2 = model_output_list[-1], model_output_list[-2], model_output_list[-3] + + h, h_0, h_1 = lambda_t - lambda_s0, lambda_s0 - lambda_s1, lambda_s1 - lambda_s2 + r0, r1 = h_0 / h, h_1 / h + D0 = m0 + D1_0, D1_1 = (1.0 / r0) * (m0 - m1), (1.0 / r1) * (m1 - m2) + D1 = D1_0 + (r0 / (r0 + r1)) * (D1_0 - D1_1) + D2 = (1.0 / (r0 + r1)) * (D1_0 - D1_1) + if self.config.algorithm_type == "dpmsolver++": + # See https://arxiv.org/abs/2206.00927 for detailed derivations + x_t = ( + (sigma_t / sigma_s0) * sample + - (alpha_t * (torch.exp(-h) - 1.0)) * D0 + + (alpha_t * ((torch.exp(-h) - 1.0) / h + 1.0)) * D1 + - (alpha_t * ((torch.exp(-h) - 1.0 + h) / h**2 - 0.5)) * D2 + ) + elif self.config.algorithm_type == "dpmsolver": + # See https://arxiv.org/abs/2206.00927 for detailed derivations + x_t = ( + (alpha_t / alpha_s0) * sample + - (sigma_t * (torch.exp(h) - 1.0)) * D0 + - (sigma_t * ((torch.exp(h) - 1.0) / h - 1.0)) * D1 + - (sigma_t * ((torch.exp(h) - 1.0 - h) / h**2 - 0.5)) * D2 + ) + elif self.config.algorithm_type == "sde-dpmsolver++": + assert noise is not None + x_t = ( + (sigma_t / sigma_s0 * torch.exp(-h)) * sample + + (alpha_t * (1.0 - torch.exp(-2.0 * h))) * D0 + + (alpha_t * ((1.0 - torch.exp(-2.0 * h)) / (-2.0 * h) + 1.0)) * D1 + + (alpha_t * ((1.0 - torch.exp(-2.0 * h) - 2.0 * h) / (2.0 * h) ** 2 - 0.5)) * D2 + + sigma_t * torch.sqrt(1.0 - torch.exp(-2 * h)) * noise + ) + return x_t + + def index_for_timestep(self, timestep, schedule_timesteps=None): + if schedule_timesteps is None: + schedule_timesteps = self.timesteps + + index_candidates = (schedule_timesteps == timestep).nonzero() + + if len(index_candidates) == 0: + step_index = len(self.timesteps) - 1 + # The sigma index that is taken for the **very** first `step` + # is always the second index (or the last index if there is only 1) + # This way we can ensure we don't accidentally skip a sigma in + # case we start in the middle of the denoising schedule (e.g. for image-to-image) + elif len(index_candidates) > 1: + step_index = index_candidates[1].item() + else: + step_index = index_candidates[0].item() + + return step_index + + def _init_step_index(self, timestep): + """ + Initialize the step_index counter for the scheduler. + """ + + if self.begin_index is None: + if isinstance(timestep, torch.Tensor): + timestep = timestep.to(self.timesteps.device) + self._step_index = self.index_for_timestep(timestep) + else: + self._step_index = self._begin_index + + def step( + self, + model_output: torch.Tensor, + timestep: Union[int, torch.Tensor], + sample: torch.Tensor, + generator=None, + variance_noise: Optional[torch.Tensor] = None, + return_dict: bool = True, + ) -> Union[SchedulerOutput, Tuple]: + """ + Predict the sample from the previous timestep by reversing the SDE. This function propagates the sample with + the multistep DPMSolver. + + Args: + model_output (`torch.Tensor`): + The direct output from learned diffusion model. + timestep (`int`): + The current discrete timestep in the diffusion chain. + sample (`torch.Tensor`): + A current instance of a sample created by the diffusion process. + generator (`torch.Generator`, *optional*): + A random number generator. + variance_noise (`torch.Tensor`): + Alternative to generating noise with `generator` by directly providing the noise for the variance + itself. Useful for methods such as [`LEdits++`]. + return_dict (`bool`): + Whether or not to return a [`~schedulers.scheduling_utils.SchedulerOutput`] or `tuple`. + + Returns: + [`~schedulers.scheduling_utils.SchedulerOutput`] or `tuple`: + If return_dict is `True`, [`~schedulers.scheduling_utils.SchedulerOutput`] is returned, otherwise a + tuple is returned where the first element is the sample tensor. + + """ + if self.num_inference_steps is None: + raise ValueError( + "Number of inference steps is 'None', you need to run 'set_timesteps' after creating the scheduler" + ) + + if self.step_index is None: + self._init_step_index(timestep) + + # Improve numerical stability for small number of steps + lower_order_final = (self.step_index == len(self.timesteps) - 1) and ( + self.config.euler_at_final + or (self.config.lower_order_final and len(self.timesteps) < 15) + or self.config.final_sigmas_type == "zero" + ) + lower_order_second = ( + (self.step_index == len(self.timesteps) - 2) and self.config.lower_order_final and len(self.timesteps) < 15 + ) + + model_output = self.convert_model_output(model_output, sample=sample) + for i in range(self.config.solver_order - 1): + self.model_outputs[i] = self.model_outputs[i + 1] + self.model_outputs[-1] = model_output + + # Upcast to avoid precision issues when computing prev_sample + sample = sample.to(torch.float32) + if self.config.algorithm_type in ["sde-dpmsolver", "sde-dpmsolver++"] and variance_noise is None: + noise = randn_tensor( + model_output.shape, generator=generator, device=model_output.device, dtype=torch.float32 + ) + elif self.config.algorithm_type in ["sde-dpmsolver", "sde-dpmsolver++"]: + noise = variance_noise.to(device=model_output.device, dtype=torch.float32) + else: + noise = None + + if self.config.solver_order == 1 or self.lower_order_nums < 1 or lower_order_final: + prev_sample = self.dpm_solver_first_order_update(model_output, sample=sample, noise=noise) + elif self.config.solver_order == 2 or self.lower_order_nums < 2 or lower_order_second: + prev_sample = self.multistep_dpm_solver_second_order_update(self.model_outputs, sample=sample, noise=noise) + else: + prev_sample = self.multistep_dpm_solver_third_order_update(self.model_outputs, sample=sample, noise=noise) + + if self.lower_order_nums < self.config.solver_order: + self.lower_order_nums += 1 + + # Cast sample back to expected dtype + prev_sample = prev_sample.to(model_output.dtype) + + # upon completion increase step index by one + self._step_index += 1 + + if not return_dict: + return (prev_sample,) + + return SchedulerOutput(prev_sample=prev_sample) + + def scale_model_input(self, sample: torch.Tensor, *args, **kwargs) -> torch.Tensor: + """ + Ensures interchangeability with schedulers that need to scale the denoising model input depending on the + current timestep. + + Args: + sample (`torch.Tensor`): + The input sample. + + Returns: + `torch.Tensor`: + A scaled input sample. + """ + return sample + + def add_noise( + self, + original_samples: torch.Tensor, + noise: torch.Tensor, + timesteps: torch.IntTensor, + ) -> torch.Tensor: + # Make sure sigmas and timesteps have the same device and dtype as original_samples + sigmas = self.sigmas.to(device=original_samples.device, dtype=original_samples.dtype) + if original_samples.device.type == "mps" and torch.is_floating_point(timesteps): + # mps does not support float64 + schedule_timesteps = self.timesteps.to(original_samples.device, dtype=torch.float32) + timesteps = timesteps.to(original_samples.device, dtype=torch.float32) + else: + schedule_timesteps = self.timesteps.to(original_samples.device) + timesteps = timesteps.to(original_samples.device) + + # begin_index is None when the scheduler is used for training or pipeline does not implement set_begin_index + if self.begin_index is None: + step_indices = [self.index_for_timestep(t, schedule_timesteps) for t in timesteps] + elif self.step_index is not None: + # add_noise is called after first denoising step (for inpainting) + step_indices = [self.step_index] * timesteps.shape[0] + else: + # add noise is called before first denoising step to create initial latent(img2img) + step_indices = [self.begin_index] * timesteps.shape[0] + + sigma = sigmas[step_indices].flatten() + while len(sigma.shape) < len(original_samples.shape): + sigma = sigma.unsqueeze(-1) + + alpha_t, sigma_t = self._sigma_to_alpha_sigma_t(sigma) + noisy_samples = alpha_t * original_samples + sigma_t * noise + return noisy_samples + + def __len__(self): + return self.config.num_train_timesteps \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/schedulers/scheduling_flow_match_euler_discrete.py b/examples/OmniGen2-RL/omnigen2/schedulers/scheduling_flow_match_euler_discrete.py new file mode 100644 index 0000000000000000000000000000000000000000..b00680580776a093b81cf4256f815bb6841bfd95 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/schedulers/scheduling_flow_match_euler_discrete.py @@ -0,0 +1,261 @@ +# Copyright 2024 Stability AI, Katherine Crowson and The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import math +from dataclasses import dataclass +from typing import List, Optional, Tuple, Union + +import numpy as np +import torch + +from diffusers.configuration_utils import ConfigMixin, register_to_config +from diffusers.utils import BaseOutput, logging +from diffusers.schedulers.scheduling_utils import SchedulerMixin + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +def expand_as(tensor, other): + for _ in range(other.ndim - tensor.ndim): + tensor = tensor.unsqueeze(-1) + return tensor + + +@dataclass +class FlowMatchEulerDiscreteSchedulerOutput(BaseOutput): + """ + Output class for the scheduler's `step` function output. + + Args: + prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images): + Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the + denoising loop. + """ + + prev_sample: torch.FloatTensor + + +class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin): + """ + Euler scheduler. + + This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic + methods the library implements for all schedulers such as loading and saving. + + Args: + num_train_timesteps (`int`, defaults to 1000): + The number of diffusion steps to train the model. + timestep_spacing (`str`, defaults to `"linspace"`): + The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and + Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information. + shift (`float`, defaults to 1.0): + The shift value for the timestep schedule. + """ + + _compatibles = [] + order = 1 + + @register_to_config + def __init__( + self, + num_train_timesteps: int = 1000, + dynamic_time_shift: bool = True, + time_shift_base_res: int = 320 + ): + timesteps = torch.linspace(0, 1, num_train_timesteps + 1, dtype=torch.float32)[:-1] + + self.timesteps = timesteps + + self._step_index = None + self._begin_index = None + + @property + def step_index(self): + """ + The index counter for current timestep. It will increase 1 after each scheduler step. + """ + return self._step_index + + @property + def begin_index(self): + """ + The index for the first timestep. It should be set from pipeline with `set_begin_index` method. + """ + return self._begin_index + + # Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index + def set_begin_index(self, begin_index: int = 0): + """ + Sets the begin index for the scheduler. This function should be run from pipeline before the inference. + + Args: + begin_index (`int`): + The begin index for the scheduler. + """ + self._begin_index = begin_index + + def index_for_timestep(self, timestep, schedule_timesteps=None): + if schedule_timesteps is None: + schedule_timesteps = self._timesteps + + indices = (schedule_timesteps == timestep).nonzero() + + # The sigma index that is taken for the **very** first `step` + # is always the second index (or the last index if there is only 1) + # This way we can ensure we don't accidentally skip a sigma in + # case we start in the middle of the denoising schedule (e.g. for image-to-image) + pos = 1 if len(indices) > 1 else 0 + + return indices[pos].item() + + # def time_shift(self, mu: float, sigma: float, t: torch.Tensor): + # return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma) + + def set_timesteps( + self, + num_inference_steps: int = None, + device: Union[str, torch.device] = None, + timesteps: Optional[List[float]] = None, + num_tokens: Optional[int] = None + ): + """ + Sets the discrete timesteps used for the diffusion chain (to be run before inference). + + Args: + num_inference_steps (`int`): + The number of diffusion steps used when generating samples with a pre-trained model. + device (`str` or `torch.device`, *optional*): + The device to which the timesteps should be moved to. If `None`, the timesteps are not moved. + """ + + if timesteps is None: + self.num_inference_steps = num_inference_steps + if self.config.dynamic_time_shift and num_tokens is not None: + if isinstance(num_tokens, list): + timesteps = [] + for i in range(len(num_tokens)): + _timesteps = np.linspace(0, 1, num_inference_steps + 1, dtype=np.float32)[:-1] + # m = np.sqrt(num_tokens[i]) / 40 # when input resolution is 320 * 320, m = 1, when input resolution is 1024 * 1024, m = 3.2 + m = np.sqrt(num_tokens[i]) / (self.config.time_shift_base_res / 8) # when input resolution is 256 * 256, m = 1, when input resolution is 1024 * 1024, m = 3.2 + _timesteps = _timesteps / (m - m * _timesteps + _timesteps) + + timesteps.append(_timesteps) + timesteps = np.stack(timesteps, axis=0) + else: + timesteps = np.linspace(0, 1, num_inference_steps + 1, dtype=np.float32)[:-1] + m = np.sqrt(num_tokens) / (self.config.time_shift_base_res / 8) # when input resolution is 320 * 320, m = 1, when input resolution is 1024 * 1024, m = 3.2 + timesteps = timesteps / (m - m * timesteps + timesteps) + + timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32, device=device) + if self.config.dynamic_time_shift and num_tokens is not None and isinstance(num_tokens, list): + _timesteps = torch.cat([timesteps, torch.ones(len(num_tokens), 1, device=timesteps.device)], dim=1) + else: + _timesteps = torch.cat([timesteps, torch.ones(1, device=timesteps.device)]) + + self.timesteps = timesteps + self._timesteps = _timesteps + self._step_index = None + self._begin_index = 0 + + def _init_step_index(self, timestep): + if self.begin_index is None: + if isinstance(timestep, torch.Tensor): + timestep = timestep.to(self.timesteps.device) + self._step_index = self.index_for_timestep(timestep) + else: + self._step_index = self._begin_index + + def step( + self, + model_output: torch.FloatTensor, + timestep: Union[float, torch.FloatTensor], + sample: torch.FloatTensor, + generator: Optional[torch.Generator] = None, + mixed_precision: bool = False, + return_dict: bool = True, + ) -> Union[FlowMatchEulerDiscreteSchedulerOutput, Tuple]: + """ + Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion + process from the learned model outputs (most often the predicted noise). + + Args: + model_output (`torch.FloatTensor`): + The direct output from learned diffusion model. + timestep (`float`): + The current discrete timestep in the diffusion chain. + sample (`torch.FloatTensor`): + A current instance of a sample created by the diffusion process. + s_churn (`float`): + s_tmin (`float`): + s_tmax (`float`): + s_noise (`float`, defaults to 1.0): + Scaling factor for noise added to the sample. + generator (`torch.Generator`, *optional*): + A random number generator. + return_dict (`bool`): + Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or + tuple. + + Returns: + [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`: + If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is + returned, otherwise a tuple is returned where the first element is the sample tensor. + """ + + if ( + isinstance(timestep, int) + or isinstance(timestep, torch.IntTensor) + or isinstance(timestep, torch.LongTensor) + ): + raise ValueError( + ( + "Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to" + " `EulerDiscreteScheduler.step()` is not supported. Make sure to pass" + " one of the `scheduler.timesteps` as a timestep." + ), + ) + + if self.step_index is None: + self._init_step_index(timestep) + # Upcast to avoid precision issues when computing prev_sample + sample = sample.to(torch.float32) + if self._timesteps.dim() == 1: + t = self._timesteps[self.step_index] + t_next = self._timesteps[self.step_index + 1] + else: + t = self._timesteps[:, self.step_index] + t_next = self._timesteps[:, self.step_index + 1] + + t = expand_as(t, sample) + t_next = expand_as(t_next, sample) + + prev_sample = sample + (t_next - t) * model_output + + # Cast sample back to model compatible dtype + prev_sample = prev_sample.to(model_output.dtype) + + # upon completion increase step index by one + self._step_index += 1 + + if not return_dict: + return (prev_sample,) + + return FlowMatchEulerDiscreteSchedulerOutput(prev_sample=prev_sample) + + def __len__(self): + return self.config.num_train_timesteps \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/schedulers/scheduling_flow_match_euler_maruyama_discrete.py b/examples/OmniGen2-RL/omnigen2/schedulers/scheduling_flow_match_euler_maruyama_discrete.py new file mode 100644 index 0000000000000000000000000000000000000000..8b9b98c044d62be9eb0020c22b05e6fcc892b0d3 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/schedulers/scheduling_flow_match_euler_maruyama_discrete.py @@ -0,0 +1,298 @@ +# Copyright 2024 Stability AI, Katherine Crowson and The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import math +from dataclasses import dataclass +from typing import List, Optional, Tuple, Union + +import numpy as np +import torch + +from diffusers.configuration_utils import ConfigMixin, register_to_config +from diffusers.utils import BaseOutput, logging +from diffusers.utils.torch_utils import randn_tensor +from diffusers.schedulers.scheduling_utils import SchedulerMixin + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +def expand_as(tensor, other): + """ + Expands a tensor to match the dimensions of another tensor. + + If tensor has shape [b] and other has shape [b, c, h, w], + this function will reshape tensor to [b, 1, 1, 1] to enable broadcasting. + + Args: + tensor (`torch.FloatTensor`): The tensor to expand + other (`torch.FloatTensor`): The tensor whose shape will be matched + + Returns: + `torch.FloatTensor`: The expanded tensor + """ + for _ in range(other.ndim - tensor.ndim): + tensor = tensor.unsqueeze(-1) + return tensor + +@dataclass +class FlowMatchEulerMaruyamaDiscreteSchedulerOutput(BaseOutput): + """ + Output class for the scheduler's `step` function output. + + Args: + prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images): + Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the + denoising loop. + """ + + prev_sample: torch.FloatTensor + log_prob: Optional[torch.FloatTensor] = None + + +class FlowMatchEulerMaruyamaDiscreteScheduler(SchedulerMixin, ConfigMixin): + """ + Euler scheduler. + + This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic + methods the library implements for all schedulers such as loading and saving. + + Args: + num_train_timesteps (`int`, defaults to 1000): + The number of diffusion steps to train the model. + timestep_spacing (`str`, defaults to `"linspace"`): + The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and + Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information. + shift (`float`, defaults to 1.0): + The shift value for the timestep schedule. + """ + + _compatibles = [] + order = 1 + + @register_to_config + def __init__( + self, + num_train_timesteps: int = 1000, + dynamic_time_shift: bool = True, + sigma_schedule: str = "v1", + sigma_coef: float = 0.7, + time_shift_base_res: int = 320, + ): + timesteps = torch.linspace(0, 1, num_train_timesteps + 1, dtype=torch.float32)[:-1] + self.time_shift_base_res = time_shift_base_res + self.timesteps = timesteps + + self._step_index = None + self._begin_index = None + + @property + def step_index(self): + """ + The index counter for current timestep. It will increase 1 after each scheduler step. + """ + return self._step_index + + @property + def begin_index(self): + """ + The index for the first timestep. It should be set from pipeline with `set_begin_index` method. + """ + return self._begin_index + + # Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index + def set_begin_index(self, begin_index: int = 0): + """ + Sets the begin index for the scheduler. This function should be run from pipeline before the inference. + + Args: + begin_index (`int`): + The begin index for the scheduler. + """ + self._begin_index = begin_index + + def index_for_timestep(self, timestep, schedule_timesteps=None): + if schedule_timesteps is None: + schedule_timesteps = self._timesteps + + indices = (schedule_timesteps == timestep).nonzero() + + # The sigma index that is taken for the **very** first `step` + # is always the second index (or the last index if there is only 1) + # This way we can ensure we don't accidentally skip a sigma in + # case we start in the middle of the denoising schedule (e.g. for image-to-image) + pos = 1 if len(indices) > 1 else 0 + + return indices[pos].item() + + def set_timesteps( + self, + num_inference_steps: int = None, + device: Union[str, torch.device] = None, + timesteps: Optional[List[float]] = None, + num_tokens: Optional[Union[int, List[int]]] = None + ): + """ + Sets the discrete timesteps used for the diffusion chain (to be run before inference). + + Args: + num_inference_steps (`int`): + The number of diffusion steps used when generating samples with a pre-trained model. + device (`str` or `torch.device`, *optional*): + The device to which the timesteps should be moved to. If `None`, the timesteps are not moved. + """ + + if timesteps is None: + self.num_inference_steps = num_inference_steps + if self.config.dynamic_time_shift and num_tokens is not None: + if isinstance(num_tokens, list): + timesteps = [] + for i in range(len(num_tokens)): + _timesteps = np.linspace(0, 1, num_inference_steps + 1, dtype=np.float32)[:-1] + m = np.sqrt(num_tokens[i]) / (self.config.time_shift_base_res / 8) # when input resolution is 320 * 320, m = 1, when input resolution is 1024 * 1024, m = 3.2 + _timesteps = _timesteps / (m - m * _timesteps + _timesteps) + + timesteps.append(_timesteps) + timesteps = np.stack(timesteps, axis=0) + else: + timesteps = np.linspace(0, 1, num_inference_steps + 1, dtype=np.float32)[:-1] + m = np.sqrt(num_tokens) / (self.config.time_shift_base_res / 8) # when input resolution is 320 * 320, m = 1, when input resolution is 1024 * 1024, m = 3.2 + timesteps = timesteps / (m - m * timesteps + timesteps) + + timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32, device=device) + if self.config.dynamic_time_shift and num_tokens is not None and isinstance(num_tokens, list): + _timesteps = torch.cat([timesteps, torch.ones(len(num_tokens), 1, device=timesteps.device)], dim=1) + else: + _timesteps = torch.cat([timesteps, torch.ones(1, device=timesteps.device)]) + + self.timesteps = timesteps + self._timesteps = _timesteps + self._step_index = None + self._begin_index = 0 + + def _init_step_index(self, timestep): + if self.begin_index is None: + if isinstance(timestep, torch.Tensor): + timestep = timestep.to(self.timesteps.device) + self._step_index = self.index_for_timestep(timestep) + else: + self._step_index = self._begin_index + + def get_sigma_t(self, t, t_next=None): + if t_next is None: + t_next = t + def _get_sigma_t(t, t_next): + return self.config.sigma_coef * ((1 - t) / (t_next)) ** 0.5 + if t.ndim > 0: + return torch.stack([_get_sigma_t(_t, _t_next) for _t, _t_next in zip(t, t_next)]) + else: + return _get_sigma_t(t, t_next) + + def step( + self, + model_output: torch.FloatTensor, + timestep: Union[float, torch.FloatTensor], + sample: torch.FloatTensor, + generator: Optional[torch.Generator] = None, + return_log_prob: bool = False, + img_mask: Optional[torch.Tensor] = None, + mixed_precision: bool = False, + return_dict: bool = True, + ) -> Union[FlowMatchEulerMaruyamaDiscreteSchedulerOutput, Tuple]: + if ( + isinstance(timestep, int) + or isinstance(timestep, torch.IntTensor) + or isinstance(timestep, torch.LongTensor) + ): + raise ValueError( + ( + "Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to" + " `EulerDiscreteScheduler.step()` is not supported. Make sure to pass" + " one of the `scheduler.timesteps` as a timestep." + ), + ) + + if self.step_index is None: + self._init_step_index(timestep) + # Upcast to avoid precision issues when computing prev_sample + sample = sample.to(torch.float32) + t = self._timesteps[:, self.step_index] + t_next = self._timesteps[:, self.step_index + 1] + + sigma_t = self.get_sigma_t(t, t_next if self.step_index == 0 else None) + + sigma_t = expand_as(sigma_t, sample) + t = expand_as(t, sample) + t_next = expand_as(t_next, sample) + + dt = t_next - t + + sigma_t = sigma_t.to(dtype=torch.float32) + t = t.to(dtype=torch.float32) + t_next = t_next.to(dtype=torch.float32) + dt = dt.to(dtype=torch.float32) + + prev_sample_mean = ( + sample.to(dtype=torch.float32) * (1 - sigma_t**2 / (2 * (1 - t)) * dt) + + model_output * (1 + sigma_t**2 * t / (2 * (1 - t))) * dt + ) + variance_noise = randn_tensor( + model_output.shape, + generator=generator, + device=sample.device, + dtype=sample.dtype, + ) + prev_sample = ( + prev_sample_mean + sigma_t * torch.sqrt(dt) * variance_noise + ) + + if img_mask is not None: + img_mask = expand_as(img_mask, sample).expand(sample.shape) + prev_sample = prev_sample * img_mask + + log_prob = None + if return_log_prob: + log_prob = ( + -((prev_sample.detach().to(dtype=torch.float32) - prev_sample_mean) ** 2) + / (2 * ((sigma_t ** 2) * dt)) + - torch.log(sigma_t * torch.sqrt(dt)) + - 0.5 * torch.log(2 * torch.as_tensor(math.pi, device=sample.device)) + ) + + log_prob = (log_prob * img_mask).sum( + dim=tuple(range(-log_prob.ndim + 1, 0)), dtype=torch.float32 + ) / img_mask.sum( + dim=tuple(range(-log_prob.ndim + 1, 0)), dtype=torch.float32 + ) + + # Cast sample back to model compatible dtype + if not mixed_precision: + prev_sample = prev_sample.to(model_output.dtype) + + # upon completion increase step index by one + self._step_index += 1 + + if return_log_prob: + if not return_dict: + return (prev_sample, log_prob) + + return FlowMatchEulerMaruyamaDiscreteSchedulerOutput(prev_sample=prev_sample, log_prob=log_prob) + else: + if not return_dict: + return (prev_sample,) + + return FlowMatchEulerMaruyamaDiscreteSchedulerOutput(prev_sample=prev_sample) + + def __len__(self): + return self.config.num_train_timesteps \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/training_utils.py b/examples/OmniGen2-RL/omnigen2/training_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..a59ec9b5370698dd89229720fa5abe7aa241c7c4 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/training_utils.py @@ -0,0 +1,645 @@ +import contextlib +import copy +import gc +import math +import random +from typing import Any, Dict, Iterable, List, Optional, Tuple, Union + +import numpy as np +import torch + +from diffusers.models import UNet2DConditionModel +from diffusers.schedulers import SchedulerMixin +from diffusers.utils import ( + convert_state_dict_to_diffusers, + convert_state_dict_to_peft, + deprecate, + is_peft_available, + is_torch_npu_available, + is_torchvision_available, + is_transformers_available, +) + + +if is_transformers_available(): + import transformers + + if transformers.integrations.deepspeed.is_deepspeed_zero3_enabled(): + import deepspeed + +if is_peft_available(): + from peft import set_peft_model_state_dict + +if is_torchvision_available(): + from torchvision import transforms + +if is_torch_npu_available(): + import torch_npu # noqa: F401 + + +def set_seed(seed: int): + """ + Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`. + + Args: + seed (`int`): The seed to set. + + Returns: + `None` + """ + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + if is_torch_npu_available(): + torch.npu.manual_seed_all(seed) + else: + torch.cuda.manual_seed_all(seed) + # ^^ safe to call this function even if cuda is not available + + +def compute_snr(noise_scheduler, timesteps): + """ + Computes SNR as per + https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L847-L849 + for the given timesteps using the provided noise scheduler. + + Args: + noise_scheduler (`NoiseScheduler`): + An object containing the noise schedule parameters, specifically `alphas_cumprod`, which is used to compute + the SNR values. + timesteps (`torch.Tensor`): + A tensor of timesteps for which the SNR is computed. + + Returns: + `torch.Tensor`: A tensor containing the computed SNR values for each timestep. + """ + alphas_cumprod = noise_scheduler.alphas_cumprod + sqrt_alphas_cumprod = alphas_cumprod**0.5 + sqrt_one_minus_alphas_cumprod = (1.0 - alphas_cumprod) ** 0.5 + + # Expand the tensors. + # Adapted from https://github.com/TiankaiHang/Min-SNR-Diffusion-Training/blob/521b624bd70c67cee4bdf49225915f5945a872e3/guided_diffusion/gaussian_diffusion.py#L1026 + sqrt_alphas_cumprod = sqrt_alphas_cumprod.to(device=timesteps.device)[timesteps].float() + while len(sqrt_alphas_cumprod.shape) < len(timesteps.shape): + sqrt_alphas_cumprod = sqrt_alphas_cumprod[..., None] + alpha = sqrt_alphas_cumprod.expand(timesteps.shape) + + sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod.to(device=timesteps.device)[timesteps].float() + while len(sqrt_one_minus_alphas_cumprod.shape) < len(timesteps.shape): + sqrt_one_minus_alphas_cumprod = sqrt_one_minus_alphas_cumprod[..., None] + sigma = sqrt_one_minus_alphas_cumprod.expand(timesteps.shape) + + # Compute SNR. + snr = (alpha / sigma) ** 2 + return snr + + +def resolve_interpolation_mode(interpolation_type: str): + """ + Maps a string describing an interpolation function to the corresponding torchvision `InterpolationMode` enum. The + full list of supported enums is documented at + https://pytorch.org/vision/0.9/transforms.html#torchvision.transforms.functional.InterpolationMode. + + Args: + interpolation_type (`str`): + A string describing an interpolation method. Currently, `bilinear`, `bicubic`, `box`, `nearest`, + `nearest_exact`, `hamming`, and `lanczos` are supported, corresponding to the supported interpolation modes + in torchvision. + + Returns: + `torchvision.transforms.InterpolationMode`: an `InterpolationMode` enum used by torchvision's `resize` + transform. + """ + if not is_torchvision_available(): + raise ImportError( + "Please make sure to install `torchvision` to be able to use the `resolve_interpolation_mode()` function." + ) + + if interpolation_type == "bilinear": + interpolation_mode = transforms.InterpolationMode.BILINEAR + elif interpolation_type == "bicubic": + interpolation_mode = transforms.InterpolationMode.BICUBIC + elif interpolation_type == "box": + interpolation_mode = transforms.InterpolationMode.BOX + elif interpolation_type == "nearest": + interpolation_mode = transforms.InterpolationMode.NEAREST + elif interpolation_type == "nearest_exact": + interpolation_mode = transforms.InterpolationMode.NEAREST_EXACT + elif interpolation_type == "hamming": + interpolation_mode = transforms.InterpolationMode.HAMMING + elif interpolation_type == "lanczos": + interpolation_mode = transforms.InterpolationMode.LANCZOS + else: + raise ValueError( + f"The given interpolation mode {interpolation_type} is not supported. Currently supported interpolation" + f" modes are `bilinear`, `bicubic`, `box`, `nearest`, `nearest_exact`, `hamming`, and `lanczos`." + ) + + return interpolation_mode + + +def compute_dream_and_update_latents( + unet: UNet2DConditionModel, + noise_scheduler: SchedulerMixin, + timesteps: torch.Tensor, + noise: torch.Tensor, + noisy_latents: torch.Tensor, + target: torch.Tensor, + encoder_hidden_states: torch.Tensor, + dream_detail_preservation: float = 1.0, +) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor]]: + """ + Implements "DREAM (Diffusion Rectification and Estimation-Adaptive Models)" from http://arxiv.org/abs/2312.00210. + DREAM helps align training with sampling to help training be more efficient and accurate at the cost of an extra + forward step without gradients. + + Args: + `unet`: The state unet to use to make a prediction. + `noise_scheduler`: The noise scheduler used to add noise for the given timestep. + `timesteps`: The timesteps for the noise_scheduler to user. + `noise`: A tensor of noise in the shape of noisy_latents. + `noisy_latents`: Previously noise latents from the training loop. + `target`: The ground-truth tensor to predict after eps is removed. + `encoder_hidden_states`: Text embeddings from the text model. + `dream_detail_preservation`: A float value that indicates detail preservation level. + See reference. + + Returns: + `tuple[torch.Tensor, torch.Tensor]`: Adjusted noisy_latents and target. + """ + alphas_cumprod = noise_scheduler.alphas_cumprod.to(timesteps.device)[timesteps, None, None, None] + sqrt_one_minus_alphas_cumprod = (1.0 - alphas_cumprod) ** 0.5 + + # The paper uses lambda = sqrt(1 - alpha) ** p, with p = 1 in their experiments. + dream_lambda = sqrt_one_minus_alphas_cumprod**dream_detail_preservation + + pred = None + with torch.no_grad(): + pred = unet(noisy_latents, timesteps, encoder_hidden_states).sample + + _noisy_latents, _target = (None, None) + if noise_scheduler.config.prediction_type == "epsilon": + predicted_noise = pred + delta_noise = (noise - predicted_noise).detach() + delta_noise.mul_(dream_lambda) + _noisy_latents = noisy_latents.add(sqrt_one_minus_alphas_cumprod * delta_noise) + _target = target.add(delta_noise) + elif noise_scheduler.config.prediction_type == "v_prediction": + raise NotImplementedError("DREAM has not been implemented for v-prediction") + else: + raise ValueError(f"Unknown prediction type {noise_scheduler.config.prediction_type}") + + return _noisy_latents, _target + + +def unet_lora_state_dict(unet: UNet2DConditionModel) -> Dict[str, torch.Tensor]: + r""" + Returns: + A state dict containing just the LoRA parameters. + """ + lora_state_dict = {} + + for name, module in unet.named_modules(): + if hasattr(module, "set_lora_layer"): + lora_layer = getattr(module, "lora_layer") + if lora_layer is not None: + current_lora_layer_sd = lora_layer.state_dict() + for lora_layer_matrix_name, lora_param in current_lora_layer_sd.items(): + # The matrix name can either be "down" or "up". + lora_state_dict[f"{name}.lora.{lora_layer_matrix_name}"] = lora_param + + return lora_state_dict + + +def cast_training_params(model: Union[torch.nn.Module, List[torch.nn.Module]], dtype=torch.float32): + """ + Casts the training parameters of the model to the specified data type. + + Args: + model: The PyTorch model whose parameters will be cast. + dtype: The data type to which the model parameters will be cast. + """ + if not isinstance(model, list): + model = [model] + for m in model: + for param in m.parameters(): + # only upcast trainable parameters into fp32 + if param.requires_grad: + param.data = param.to(dtype) + + +def _set_state_dict_into_text_encoder( + lora_state_dict: Dict[str, torch.Tensor], prefix: str, text_encoder: torch.nn.Module +): + """ + Sets the `lora_state_dict` into `text_encoder` coming from `transformers`. + + Args: + lora_state_dict: The state dictionary to be set. + prefix: String identifier to retrieve the portion of the state dict that belongs to `text_encoder`. + text_encoder: Where the `lora_state_dict` is to be set. + """ + + text_encoder_state_dict = { + f"{k.replace(prefix, '')}": v for k, v in lora_state_dict.items() if k.startswith(prefix) + } + text_encoder_state_dict = convert_state_dict_to_peft(convert_state_dict_to_diffusers(text_encoder_state_dict)) + set_peft_model_state_dict(text_encoder, text_encoder_state_dict, adapter_name="default") + + +def compute_density_for_timestep_sampling( + weighting_scheme: str, + batch_size: int, + logit_mean: float = None, + logit_std: float = None, + mode_scale: float = None, + device: Union[torch.device, str] = "cpu", + generator: Optional[torch.Generator] = None, +): + """ + Compute the density for sampling the timesteps when doing SD3 training. + + Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528. + + SD3 paper reference: https://arxiv.org/abs/2403.03206v1. + """ + if weighting_scheme == "logit_normal": + u = torch.normal(mean=logit_mean, std=logit_std, size=(batch_size,), device=device, generator=generator) + u = torch.nn.functional.sigmoid(u) + elif weighting_scheme == "mode": + u = torch.rand(size=(batch_size,), device=device, generator=generator) + u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2) ** 2 - 1 + u) + else: + u = torch.rand(size=(batch_size,), device=device, generator=generator) + return u + + +def compute_loss_weighting_for_sd3(weighting_scheme: str, sigmas=None): + """ + Computes loss weighting scheme for SD3 training. + + Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528. + + SD3 paper reference: https://arxiv.org/abs/2403.03206v1. + """ + if weighting_scheme == "sigma_sqrt": + weighting = (sigmas**-2.0).float() + elif weighting_scheme == "cosmap": + bot = 1 - 2 * sigmas + 2 * sigmas**2 + weighting = 2 / (math.pi * bot) + else: + weighting = torch.ones_like(sigmas) + return weighting + + +def free_memory(): + """ + Runs garbage collection. Then clears the cache of the available accelerator. + """ + gc.collect() + + if torch.cuda.is_available(): + torch.cuda.empty_cache() + elif torch.backends.mps.is_available(): + torch.mps.empty_cache() + elif is_torch_npu_available(): + torch_npu.npu.empty_cache() + elif hasattr(torch, "xpu") and torch.xpu.is_available(): + torch.xpu.empty_cache() + + +# Adapted from torch-ema https://github.com/fadel/pytorch_ema/blob/master/torch_ema/ema.py#L14 +class EMAModel: + """ + Exponential Moving Average of models weights + """ + + def __init__( + self, + parameters: Iterable[torch.nn.Parameter], + decay: float = 0.9999, + min_decay: float = 0.0, + update_after_step: int = 0, + use_ema_warmup: bool = False, + inv_gamma: Union[float, int] = 1.0, + power: Union[float, int] = 2 / 3, + foreach: bool = False, + model_cls: Optional[Any] = None, + model_config: Dict[str, Any] = None, + **kwargs, + ): + """ + Args: + parameters (Iterable[torch.nn.Parameter]): The parameters to track. + decay (float): The decay factor for the exponential moving average. + min_decay (float): The minimum decay factor for the exponential moving average. + update_after_step (int): The number of steps to wait before starting to update the EMA weights. + use_ema_warmup (bool): Whether to use EMA warmup. + inv_gamma (float): + Inverse multiplicative factor of EMA warmup. Default: 1. Only used if `use_ema_warmup` is True. + power (float): Exponential factor of EMA warmup. Default: 2/3. Only used if `use_ema_warmup` is True. + foreach (bool): Use torch._foreach functions for updating shadow parameters. Should be faster. + device (Optional[Union[str, torch.device]]): The device to store the EMA weights on. If None, the EMA + weights will be stored on CPU. + + @crowsonkb's notes on EMA Warmup: + If gamma=1 and power=1, implements a simple average. gamma=1, power=2/3 are good values for models you plan + to train for a million or more steps (reaches decay factor 0.999 at 31.6K steps, 0.9999 at 1M steps), + gamma=1, power=3/4 for models you plan to train for less (reaches decay factor 0.999 at 10K steps, 0.9999 + at 215.4k steps). + """ + + if isinstance(parameters, torch.nn.Module): + deprecation_message = ( + "Passing a `torch.nn.Module` to `ExponentialMovingAverage` is deprecated. " + "Please pass the parameters of the module instead." + ) + deprecate( + "passing a `torch.nn.Module` to `ExponentialMovingAverage`", + "1.0.0", + deprecation_message, + standard_warn=False, + ) + parameters = parameters.parameters() + + # set use_ema_warmup to True if a torch.nn.Module is passed for backwards compatibility + use_ema_warmup = True + + if kwargs.get("max_value", None) is not None: + deprecation_message = "The `max_value` argument is deprecated. Please use `decay` instead." + deprecate("max_value", "1.0.0", deprecation_message, standard_warn=False) + decay = kwargs["max_value"] + + if kwargs.get("min_value", None) is not None: + deprecation_message = "The `min_value` argument is deprecated. Please use `min_decay` instead." + deprecate("min_value", "1.0.0", deprecation_message, standard_warn=False) + min_decay = kwargs["min_value"] + + parameters = list(parameters) + self.shadow_params = [p.clone().detach() for p in parameters] + + if kwargs.get("device", None) is not None: + deprecation_message = "The `device` argument is deprecated. Please use `to` instead." + deprecate("device", "1.0.0", deprecation_message, standard_warn=False) + self.to(device=kwargs["device"]) + + self.temp_stored_params = None + + self.decay = decay + self.min_decay = min_decay + self.update_after_step = update_after_step + self.use_ema_warmup = use_ema_warmup + self.inv_gamma = inv_gamma + self.power = power + self.optimization_step = 0 + self.cur_decay_value = None # set in `step()` + self.foreach = foreach + + self.model_cls = model_cls + self.model_config = model_config + + @classmethod + def from_pretrained(cls, path, model_cls, foreach=False) -> "EMAModel": + _, ema_kwargs = model_cls.from_config(path, return_unused_kwargs=True) + model = model_cls.from_pretrained(path) + + ema_model = cls(model.parameters(), model_cls=model_cls, model_config=model.config, foreach=foreach) + + ema_model.load_state_dict(ema_kwargs) + return ema_model + + def save_pretrained(self, path): + if self.model_cls is None: + raise ValueError("`save_pretrained` can only be used if `model_cls` was defined at __init__.") + + if self.model_config is None: + raise ValueError("`save_pretrained` can only be used if `model_config` was defined at __init__.") + + model = self.model_cls.from_config(self.model_config) + state_dict = self.state_dict() + state_dict.pop("shadow_params", None) + + model.register_to_config(**state_dict) + self.copy_to(model.parameters()) + model.save_pretrained(path) + + def get_decay(self, optimization_step: int) -> float: + """ + Compute the decay factor for the exponential moving average. + """ + step = max(0, optimization_step - self.update_after_step - 1) + + if step <= 0: + return 0.0 + + if self.use_ema_warmup: + cur_decay_value = 1 - (1 + step / self.inv_gamma) ** -self.power + else: + cur_decay_value = (1 + step) / (10 + step) + + cur_decay_value = min(cur_decay_value, self.decay) + # make sure decay is not smaller than min_decay + cur_decay_value = max(cur_decay_value, self.min_decay) + return cur_decay_value + + @torch.no_grad() + def step(self, parameters: Iterable[torch.nn.Parameter]): + if isinstance(parameters, torch.nn.Module): + deprecation_message = ( + "Passing a `torch.nn.Module` to `ExponentialMovingAverage.step` is deprecated. " + "Please pass the parameters of the module instead." + ) + deprecate( + "passing a `torch.nn.Module` to `ExponentialMovingAverage.step`", + "1.0.0", + deprecation_message, + standard_warn=False, + ) + parameters = parameters.parameters() + + parameters = list(parameters) + + self.optimization_step += 1 + + # Compute the decay factor for the exponential moving average. + decay = self.get_decay(self.optimization_step) + self.cur_decay_value = decay + one_minus_decay = 1 - decay + + context_manager = contextlib.nullcontext() + + if self.foreach: + if is_transformers_available() and transformers.integrations.deepspeed.is_deepspeed_zero3_enabled(): + context_manager = deepspeed.zero.GatheredParameters(parameters, modifier_rank=None) + + with context_manager: + params_grad = [param for param in parameters if param.requires_grad] + s_params_grad = [ + s_param for s_param, param in zip(self.shadow_params, parameters) if param.requires_grad + ] + + if len(params_grad) < len(parameters): + torch._foreach_copy_( + [s_param for s_param, param in zip(self.shadow_params, parameters) if not param.requires_grad], + [param for param in parameters if not param.requires_grad], + non_blocking=True, + ) + + torch._foreach_sub_( + s_params_grad, torch._foreach_sub(s_params_grad, params_grad), alpha=one_minus_decay + ) + + else: + for s_param, param in zip(self.shadow_params, parameters): + if is_transformers_available() and transformers.integrations.deepspeed.is_deepspeed_zero3_enabled(): + context_manager = deepspeed.zero.GatheredParameters(param, modifier_rank=None) + + with context_manager: + if param.requires_grad: + # print(f"{s_param.shape=} {param.shape=}") + s_param.sub_(one_minus_decay * (s_param - param)) + else: + s_param.copy_(param) + + def copy_to(self, parameters: Iterable[torch.nn.Parameter]) -> None: + """ + Copy current averaged parameters into given collection of parameters. + + Args: + parameters: Iterable of `torch.nn.Parameter`; the parameters to be + updated with the stored moving averages. If `None`, the parameters with which this + `ExponentialMovingAverage` was initialized will be used. + """ + parameters = list(parameters) + if self.foreach: + torch._foreach_copy_( + [param.data for param in parameters], + [s_param.to(param.device).data for s_param, param in zip(self.shadow_params, parameters)], + ) + else: + for s_param, param in zip(self.shadow_params, parameters): + param.data.copy_(s_param.to(param.device).data) + + def pin_memory(self) -> None: + r""" + Move internal buffers of the ExponentialMovingAverage to pinned memory. Useful for non-blocking transfers for + offloading EMA params to the host. + """ + + self.shadow_params = [p.pin_memory() for p in self.shadow_params] + + def to(self, device=None, dtype=None, non_blocking=False) -> None: + r""" + Move internal buffers of the ExponentialMovingAverage to `device`. + + Args: + device: like `device` argument to `torch.Tensor.to` + """ + # .to() on the tensors handles None correctly + self.shadow_params = [ + p.to(device=device, dtype=dtype, non_blocking=non_blocking) + if p.is_floating_point() + else p.to(device=device, non_blocking=non_blocking) + for p in self.shadow_params + ] + + def state_dict(self) -> dict: + r""" + Returns the state of the ExponentialMovingAverage as a dict. This method is used by accelerate during + checkpointing to save the ema state dict. + """ + # Following PyTorch conventions, references to tensors are returned: + # "returns a reference to the state and not its copy!" - + # https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict + return { + "decay": self.decay, + "min_decay": self.min_decay, + "optimization_step": self.optimization_step, + "update_after_step": self.update_after_step, + "use_ema_warmup": self.use_ema_warmup, + "inv_gamma": self.inv_gamma, + "power": self.power, + "shadow_params": self.shadow_params, + } + + def store(self, parameters: Iterable[torch.nn.Parameter]) -> None: + r""" + Saves the current parameters for restoring later. + + Args: + parameters: Iterable of `torch.nn.Parameter`. The parameters to be temporarily stored. + """ + self.temp_stored_params = [param.detach().cpu().clone() for param in parameters] + + def restore(self, parameters: Iterable[torch.nn.Parameter]) -> None: + r""" + Restore the parameters stored with the `store` method. Useful to validate the model with EMA parameters + without: affecting the original optimization process. Store the parameters before the `copy_to()` method. After + validation (or model saving), use this to restore the former parameters. + + Args: + parameters: Iterable of `torch.nn.Parameter`; the parameters to be + updated with the stored parameters. If `None`, the parameters with which this + `ExponentialMovingAverage` was initialized will be used. + """ + + if self.temp_stored_params is None: + raise RuntimeError("This ExponentialMovingAverage has no `store()`ed weights to `restore()`") + if self.foreach: + torch._foreach_copy_( + [param.data for param in parameters], [c_param.data for c_param in self.temp_stored_params] + ) + else: + for c_param, param in zip(self.temp_stored_params, parameters): + param.data.copy_(c_param.data) + + # Better memory-wise. + self.temp_stored_params = None + + def load_state_dict(self, state_dict: dict) -> None: + r""" + Loads the ExponentialMovingAverage state. This method is used by accelerate during checkpointing to save the + ema state dict. + + Args: + state_dict (dict): EMA state. Should be an object returned + from a call to :meth:`state_dict`. + """ + # deepcopy, to be consistent with module API + state_dict = copy.deepcopy(state_dict) + + self.decay = state_dict.get("decay", self.decay) + if self.decay < 0.0 or self.decay > 1.0: + raise ValueError("Decay must be between 0 and 1") + + self.min_decay = state_dict.get("min_decay", self.min_decay) + if not isinstance(self.min_decay, float): + raise ValueError("Invalid min_decay") + + self.optimization_step = state_dict.get("optimization_step", self.optimization_step) + if not isinstance(self.optimization_step, int): + raise ValueError("Invalid optimization_step") + + self.update_after_step = state_dict.get("update_after_step", self.update_after_step) + if not isinstance(self.update_after_step, int): + raise ValueError("Invalid update_after_step") + + self.use_ema_warmup = state_dict.get("use_ema_warmup", self.use_ema_warmup) + if not isinstance(self.use_ema_warmup, bool): + raise ValueError("Invalid use_ema_warmup") + + self.inv_gamma = state_dict.get("inv_gamma", self.inv_gamma) + if not isinstance(self.inv_gamma, (float, int)): + raise ValueError("Invalid inv_gamma") + + self.power = state_dict.get("power", self.power) + if not isinstance(self.power, (float, int)): + raise ValueError("Invalid power") + + shadow_params = state_dict.get("shadow_params", None) + if shadow_params is not None: + self.shadow_params = shadow_params + if not isinstance(self.shadow_params, list): + raise ValueError("shadow_params must be a list") + if not all(isinstance(p, torch.Tensor) for p in self.shadow_params): + raise ValueError("shadow_params must all be Tensors") diff --git a/examples/OmniGen2-RL/omnigen2/transport/__init__.py b/examples/OmniGen2-RL/omnigen2/transport/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..afc28718f9646356015abc81b97c3ad2bb54fe71 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/transport/__init__.py @@ -0,0 +1,74 @@ +from .transport import ModelType, PathType, Sampler, Transport, WeightType + + +def create_transport( + path_type="Linear", + prediction="velocity", + loss_weight=None, + train_eps=None, + sample_eps=None, + snr_type="uniform", + do_shift=True, + seq_len=1024, # corresponding to 512x512 + dynamic_time_shift: bool = False, + time_shift_version: str = "v1", +): + """function for creating Transport object + **Note**: model prediction defaults to velocity + Args: + - path_type: type of path to use; default to linear + - learn_score: set model prediction to score + - learn_noise: set model prediction to noise + - velocity_weighted: weight loss by velocity weight + - likelihood_weighted: weight loss by likelihood weight + - train_eps: small epsilon for avoiding instability during training + - sample_eps: small epsilon for avoiding instability during sampling + """ + + if prediction == "noise": + model_type = ModelType.NOISE + elif prediction == "score": + model_type = ModelType.SCORE + else: + model_type = ModelType.VELOCITY + + if loss_weight == "velocity": + loss_type = WeightType.VELOCITY + elif loss_weight == "likelihood": + loss_type = WeightType.LIKELIHOOD + else: + loss_type = WeightType.NONE + + path_choice = { + "Linear": PathType.LINEAR, + "GVP": PathType.GVP, + "VP": PathType.VP, + } + + path_type = path_choice[path_type] + + if path_type in [PathType.VP]: + train_eps = 1e-5 if train_eps is None else train_eps + sample_eps = 1e-3 if train_eps is None else sample_eps + elif path_type in [PathType.GVP, PathType.LINEAR] and model_type != ModelType.VELOCITY: + train_eps = 1e-3 if train_eps is None else train_eps + sample_eps = 1e-3 if train_eps is None else sample_eps + else: # velocity & [GVP, LINEAR] is stable everywhere + train_eps = 0 + sample_eps = 0 + + # create flow state + state = Transport( + model_type=model_type, + path_type=path_type, + loss_type=loss_type, + train_eps=train_eps, + sample_eps=sample_eps, + snr_type=snr_type, + do_shift=do_shift, + seq_len=seq_len, + dynamic_time_shift=dynamic_time_shift, + time_shift_version=time_shift_version, + ) + + return state diff --git a/examples/OmniGen2-RL/omnigen2/transport/dpm_solver.py b/examples/OmniGen2-RL/omnigen2/transport/dpm_solver.py new file mode 100644 index 0000000000000000000000000000000000000000..6af4b10446cd6e1f1c54bc71d369524f5fc8091f --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/transport/dpm_solver.py @@ -0,0 +1,1386 @@ +# Copyright 2024 NVIDIA CORPORATION & AFFILIATES +# +# 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. +# +# SPDX-License-Identifier: Apache-2.0 + +# This file is modified from https://github.com/PixArt-alpha/PixArt-sigma +import os + +import torch +from tqdm import tqdm + + +class NoiseScheduleFlow: + def __init__( + self, + schedule="discrete_flow", + ): + """Create a wrapper class for the forward SDE (EDM type).""" + self.T = 1 + self.t0 = 0.001 + self.schedule = schedule # ['continuous', 'discrete_flow'] + self.total_N = 1000 + + def marginal_log_mean_coeff(self, t): + """ + Compute log(alpha_t) of a given continuous-time label t in [0, T]. + """ + return torch.log(self.marginal_alpha(t)) + + def marginal_alpha(self, t): + """ + Compute alpha_t of a given continuous-time label t in [0, T]. + """ + return 1 - t + + @staticmethod + def marginal_std(t): + """ + Compute sigma_t of a given continuous-time label t in [0, T]. + """ + return t + + def marginal_lambda(self, t): + """ + Compute lambda_t = log(alpha_t) - log(sigma_t) of a given continuous-time label t in [0, T]. + """ + log_mean_coeff = self.marginal_log_mean_coeff(t) + log_std = torch.log(self.marginal_std(t)) + return log_mean_coeff - log_std + + @staticmethod + def inverse_lambda(lamb): + """ + Compute the continuous-time label t in [0, T] of a given half-logSNR lambda_t. + """ + return torch.exp(-lamb) + + +def model_wrapper( + model, + noise_schedule, + model_type="noise", + model_kwargs={}, + guidance_type="uncond", + condition=None, + unconditional_condition=None, + guidance_scale=1.0, + interval_guidance=[0, 1.0], + classifier_fn=None, + classifier_kwargs={}, +): + """Create a wrapper function for the noise prediction model. + + DPM-Solver needs to solve the continuous-time diffusion ODEs. For DPMs trained on discrete-time labels, we need to + firstly wrap the model function to a noise prediction model that accepts the continuous time as the input. + + We support four types of the diffusion model by setting `model_type`: + + 1. "noise": noise prediction model. (Trained by predicting noise). + + 2. "x_start": data prediction model. (Trained by predicting the data x_0 at time 0). + + 3. "v": velocity prediction model. (Trained by predicting the velocity). + The "v" prediction is derivation detailed in Appendix D of [1], and is used in Imagen-Video [2]. + + [1] Salimans, Tim, and Jonathan Ho. "Progressive distillation for fast sampling of diffusion models." + arXiv preprint arXiv:2202.00512 (2022). + [2] Ho, Jonathan, et al. "Imagen Video: High Definition Video Generation with Diffusion Models." + arXiv preprint arXiv:2210.02303 (2022). + + 4. "score": marginal score function. (Trained by denoising score matching). + Note that the score function and the noise prediction model follows a simple relationship: + ``` + noise(x_t, t) = -sigma_t * score(x_t, t) + ``` + + We support three types of guided sampling by DPMs by setting `guidance_type`: + 1. "uncond": unconditional sampling by DPMs. + The input `model` has the following format: + `` + model(x, t_input, **model_kwargs) -> noise | x_start | v | score + `` + + 2. "classifier": classifier guidance sampling [3] by DPMs and another classifier. + The input `model` has the following format: + `` + model(x, t_input, **model_kwargs) -> noise | x_start | v | score + `` + + The input `classifier_fn` has the following format: + `` + classifier_fn(x, t_input, cond, **classifier_kwargs) -> logits(x, t_input, cond) + `` + + [3] P. Dhariwal and A. Q. Nichol, "Diffusion models beat GANs on image synthesis," + in Advances in Neural Information Processing Systems, vol. 34, 2021, pp. 8780-8794. + + 3. "classifier-free": classifier-free guidance sampling by conditional DPMs. + The input `model` has the following format: + `` + model(x, t_input, cond, **model_kwargs) -> noise | x_start | v | score + `` + And if cond == `unconditional_condition`, the model output is the unconditional DPM output. + + [4] Ho, Jonathan, and Tim Salimans. "Classifier-free diffusion guidance." + arXiv preprint arXiv:2207.12598 (2022). + + + The `t_input` is the time label of the model, which may be discrete-time labels (i.e. 0 to 999) + or continuous-time labels (i.e. epsilon to T). + + We wrap the model function to accept only `x` and `t_continuous` as inputs, and outputs the predicted noise: + `` + def model_fn(x, t_continuous) -> noise: + t_input = get_model_input_time(t_continuous) + return noise_pred(model, x, t_input, **model_kwargs) + `` + where `t_continuous` is the continuous time labels (i.e. epsilon to T). And we use `model_fn` for DPM-Solver. + + =============================================================== + + Args: + model: A diffusion model with the corresponding format described above. + noise_schedule: A noise schedule object, such as NoiseScheduleVP. + model_type: A `str`. The parameterization type of the diffusion model. + "noise" or "x_start" or "v" or "score". + model_kwargs: A `dict`. A dict for the other inputs of the model function. + guidance_type: A `str`. The type of the guidance for sampling. + "uncond" or "classifier" or "classifier-free". + condition: A pytorch tensor. The condition for the guided sampling. + Only used for "classifier" or "classifier-free" guidance type. + unconditional_condition: A pytorch tensor. The condition for the unconditional sampling. + Only used for "classifier-free" guidance type. + guidance_scale: A `float`. The scale for the guided sampling. + classifier_fn: A classifier function. Only used for the classifier guidance. + classifier_kwargs: A `dict`. A dict for the other inputs of the classifier function. + Returns: + A noise prediction model that accepts the noised data and the continuous time as the inputs. + """ + + def get_model_input_time(t_continuous): + """ + Convert the continuous-time `t_continuous` (in [epsilon, T]) to the model input time. + For discrete-time DPMs, we convert `t_continuous` in [1 / N, 1] to `t_input` in [0, 1000 * (N - 1) / N]. + For continuous-time DPMs, we just use `t_continuous`. + """ + if noise_schedule.schedule == "discrete": + return (t_continuous - 1.0 / noise_schedule.total_N) * noise_schedule.total_N + elif noise_schedule.schedule == "discrete_flow": + return t_continuous * noise_schedule.total_N + else: + return t_continuous + + def noise_pred_fn(x, t_continuous, cond=None): + t_input = get_model_input_time(t_continuous) + if cond is None: + output = model(x, t_input, **model_kwargs) + else: + output = model(x, t_input, cond, **model_kwargs) + if model_type == "noise": + return output + elif model_type == "x_start": + alpha_t, sigma_t = noise_schedule.marginal_alpha(t_continuous), noise_schedule.marginal_std(t_continuous) + return (x - expand_dims(alpha_t, x.dim()) * output) / expand_dims(sigma_t, x.dim()) + elif model_type == "v": + alpha_t, sigma_t = noise_schedule.marginal_alpha(t_continuous), noise_schedule.marginal_std(t_continuous) + return expand_dims(alpha_t, x.dim()) * output + expand_dims(sigma_t, x.dim()) * x + elif model_type == "score": + sigma_t = noise_schedule.marginal_std(t_continuous) + return -expand_dims(sigma_t, x.dim()) * output + elif model_type == "flow": + _, sigma_t = noise_schedule.marginal_alpha(t_continuous), noise_schedule.marginal_std(t_continuous) + try: + noise = (1 - expand_dims(sigma_t, x.dim()).to(x)) * output + x + except: + noise = (1 - expand_dims(sigma_t, x.dim()).to(x)) * output[0] + x + return noise + + def cond_grad_fn(x, t_input): + """ + Compute the gradient of the classifier, i.e. nabla_{x} log p_t(cond | x_t). + """ + with torch.enable_grad(): + x_in = x.detach().requires_grad_(True) + log_prob = classifier_fn(x_in, t_input, condition, **classifier_kwargs) + return torch.autograd.grad(log_prob.sum(), x_in)[0] + + def model_fn(x, t_continuous): + """ + The noise predicition model function that is used for DPM-Solver. + """ + guidance_tp = guidance_type + if guidance_tp == "uncond": + return noise_pred_fn(x, t_continuous) + elif guidance_tp == "classifier": + assert classifier_fn is not None + t_input = get_model_input_time(t_continuous) + cond_grad = cond_grad_fn(x, t_input) + sigma_t = noise_schedule.marginal_std(t_continuous) + noise = noise_pred_fn(x, t_continuous) + return noise - guidance_scale * expand_dims(sigma_t, x.dim()) * cond_grad + elif guidance_tp == "classifier-free": + if ( + guidance_scale == 1.0 + or unconditional_condition is None + or not (interval_guidance[0] < t_continuous[0] < interval_guidance[1]) + ): + return noise_pred_fn(x, t_continuous, cond=condition) + else: + x_in = torch.cat([x] * 2) + t_in = torch.cat([t_continuous] * 2) + c_in = torch.cat([unconditional_condition, condition]) + try: + noise_uncond, noise = noise_pred_fn(x_in, t_in, cond=c_in).chunk(2) + except: + noise_uncond, noise = noise_pred_fn(x_in, t_in, cond=c_in)[0].chunk(2) + return noise_uncond + guidance_scale * (noise - noise_uncond) + + assert model_type in ["noise", "x_start", "v", "score", "flow"] + assert guidance_type in [ + "uncond", + "classifier", + "classifier-free", + ] + return model_fn + + +class DPM_Solver: + def __init__( + self, + model_fn, + noise_schedule, + algorithm_type="dpmsolver++", + correcting_x0_fn=None, + correcting_xt_fn=None, + thresholding_max_val=1.0, + dynamic_thresholding_ratio=0.995, + ): + """Construct a DPM-Solver. + + We support both DPM-Solver (`algorithm_type="dpmsolver"`) and DPM-Solver++ (`algorithm_type="dpmsolver++"`). + + We also support the "dynamic thresholding" method in Imagen[1]. For pixel-space diffusion models, you + can set both `algorithm_type="dpmsolver++"` and `correcting_x0_fn="dynamic_thresholding"` to use the + dynamic thresholding. The "dynamic thresholding" can greatly improve the sample quality for pixel-space + DPMs with large guidance scales. Note that the thresholding method is **unsuitable** for latent-space + DPMs (such as stable-diffusion). + + To support advanced algorithms in image-to-image applications, we also support corrector functions for + both x0 and xt. + + Args: + model_fn: A noise prediction model function which accepts the continuous-time input (t in [epsilon, T]): + `` + def model_fn(x, t_continuous): + return noise + `` + The shape of `x` is `(batch_size, **shape)`, and the shape of `t_continuous` is `(batch_size,)`. + noise_schedule: A noise schedule object, such as NoiseScheduleVP. + algorithm_type: A `str`. Either "dpmsolver" or "dpmsolver++". + correcting_x0_fn: A `str` or a function with the following format: + ``` + def correcting_x0_fn(x0, t): + x0_new = ... + return x0_new + ``` + This function is to correct the outputs of the data prediction model at each sampling step. e.g., + ``` + x0_pred = data_pred_model(xt, t) + if correcting_x0_fn is not None: + x0_pred = correcting_x0_fn(x0_pred, t) + xt_1 = update(x0_pred, xt, t) + ``` + If `correcting_x0_fn="dynamic_thresholding"`, we use the dynamic thresholding proposed in Imagen[1]. + correcting_xt_fn: A function with the following format: + ``` + def correcting_xt_fn(xt, t, step): + x_new = ... + return x_new + ``` + This function is to correct the intermediate samples xt at each sampling step. e.g., + ``` + xt = ... + xt = correcting_xt_fn(xt, t, step) + ``` + thresholding_max_val: A `float`. The max value for thresholding. + Valid only when use `dpmsolver++` and `correcting_x0_fn="dynamic_thresholding"`. + dynamic_thresholding_ratio: A `float`. The ratio for dynamic thresholding (see Imagen[1] for details). + Valid only when use `dpmsolver++` and `correcting_x0_fn="dynamic_thresholding"`. + + [1] Chitwan Saharia, William Chan, Saurabh Saxena, Lala Li, Jay Whang, Emily Denton, Seyed Kamyar Seyed Ghasemipour, + Burcu Karagol Ayan, S Sara Mahdavi, Rapha Gontijo Lopes, et al. Photorealistic text-to-image diffusion models + with deep language understanding. arXiv preprint arXiv:2205.11487, 2022b. + """ + self.model = lambda x, t: model_fn(x, t.expand(x.shape[0])) + self.noise_schedule = noise_schedule + assert algorithm_type in ["dpmsolver", "dpmsolver++"] + self.algorithm_type = algorithm_type + if correcting_x0_fn == "dynamic_thresholding": + self.correcting_x0_fn = self.dynamic_thresholding_fn + else: + self.correcting_x0_fn = correcting_x0_fn + self.correcting_xt_fn = correcting_xt_fn + self.dynamic_thresholding_ratio = dynamic_thresholding_ratio + self.thresholding_max_val = thresholding_max_val + self.register_progress_bar() + + def register_progress_bar(self, progress_fn=None): + """ + Register a progress bar callback function + + Args: + progress_fn: Callback function that takes current step and total steps as parameters + """ + self.progress_fn = progress_fn if progress_fn is not None else lambda step, total: None + + def update_progress(self, step, total_steps): + """ + Update sampling progress + + Args: + step: Current step number + total_steps: Total number of steps + """ + if hasattr(self, "progress_fn"): + try: + self.progress_fn(step / total_steps, desc=f"Generating {step}/{total_steps}") + except: + self.progress_fn(step, total_steps) + + else: + # If no progress_fn registered, use default empty function + pass + + def dynamic_thresholding_fn(self, x0, t): + """ + The dynamic thresholding method. + """ + dims = x0.dim() + p = self.dynamic_thresholding_ratio + s = torch.quantile(torch.abs(x0).reshape((x0.shape[0], -1)), p, dim=1) + s = expand_dims(torch.maximum(s, self.thresholding_max_val * torch.ones_like(s).to(s.device)), dims) + x0 = torch.clamp(x0, -s, s) / s + return x0 + + def noise_prediction_fn(self, x, t): + """ + Return the noise prediction model. + """ + return self.model(x, t) + + def data_prediction_fn(self, x, t): + """ + Return the data prediction model (with corrector). + """ + noise = self.noise_prediction_fn(x, t) + alpha_t, sigma_t = self.noise_schedule.marginal_alpha(t), self.noise_schedule.marginal_std(t) + x0 = (x - sigma_t * noise) / alpha_t + if self.correcting_x0_fn is not None: + x0 = self.correcting_x0_fn(x0, t) + return x0 + + def model_fn(self, x, t): + """ + Convert the model to the noise prediction model or the data prediction model. + """ + if self.algorithm_type == "dpmsolver++": + return self.data_prediction_fn(x, t) + else: + return self.noise_prediction_fn(x, t) + + def get_time_steps(self, skip_type, t_T, t_0, N, device, shift=1.0): + """Compute the intermediate time steps for sampling. + + Args: + skip_type: A `str`. The type for the spacing of the time steps. We support three types: + - 'logSNR': uniform logSNR for the time steps. + - 'time_uniform': uniform time for the time steps. (**Recommended for high-resolutional data**.) + - 'time_quadratic': quadratic time for the time steps. (Used in DDIM for low-resolutional data.) + t_T: A `float`. The starting time of the sampling (default is T). + t_0: A `float`. The ending time of the sampling (default is epsilon). + N: A `int`. The total number of the spacing of the time steps. + device: A torch device. + Returns: + A pytorch tensor of the time steps, with the shape (N + 1,). + """ + if skip_type == "logSNR": + lambda_T = self.noise_schedule.marginal_lambda(torch.tensor(t_T).to(device)) + lambda_0 = self.noise_schedule.marginal_lambda(torch.tensor(t_0).to(device)) + logSNR_steps = torch.linspace(lambda_T.cpu().item(), lambda_0.cpu().item(), N + 1).to(device) + return self.noise_schedule.inverse_lambda(logSNR_steps) + elif skip_type == "time_uniform": + return torch.linspace(t_T, t_0, N + 1).to(device) + elif skip_type == "time_quadratic": + t_order = 2 + t = torch.linspace(t_T ** (1.0 / t_order), t_0 ** (1.0 / t_order), N + 1).pow(t_order).to(device) + return t + elif skip_type == "time_uniform_flow": + betas = torch.linspace(t_T, t_0, N + 1).to(device) + sigmas = 1.0 - betas + sigmas = (shift * sigmas / (1 + (shift - 1) * sigmas)).flip(dims=[0]) + return sigmas + else: + raise ValueError( + f"Unsupported skip_type {skip_type}, need to be 'logSNR' or 'time_uniform' or 'time_quadratic'" + ) + + def get_orders_and_timesteps_for_singlestep_solver(self, steps, order, skip_type, t_T, t_0, device): + """ + Get the order of each step for sampling by the singlestep DPM-Solver. + + We combine both DPM-Solver-1,2,3 to use all the function evaluations, which is named as "DPM-Solver-fast". + Given a fixed number of function evaluations by `steps`, the sampling procedure by DPM-Solver-fast is: + - If order == 1: + We take `steps` of DPM-Solver-1 (i.e. DDIM). + - If order == 2: + - Denote K = (steps // 2). We take K or (K + 1) intermediate time steps for sampling. + - If steps % 2 == 0, we use K steps of DPM-Solver-2. + - If steps % 2 == 1, we use K steps of DPM-Solver-2 and 1 step of DPM-Solver-1. + - If order == 3: + - Denote K = (steps // 3 + 1). We take K intermediate time steps for sampling. + - If steps % 3 == 0, we use (K - 2) steps of DPM-Solver-3, and 1 step of DPM-Solver-2 and 1 step of DPM-Solver-1. + - If steps % 3 == 1, we use (K - 1) steps of DPM-Solver-3 and 1 step of DPM-Solver-1. + - If steps % 3 == 2, we use (K - 1) steps of DPM-Solver-3 and 1 step of DPM-Solver-2. + + ============================================ + Args: + order: A `int`. The max order for the solver (2 or 3). + steps: A `int`. The total number of function evaluations (NFE). + skip_type: A `str`. The type for the spacing of the time steps. We support three types: + - 'logSNR': uniform logSNR for the time steps. + - 'time_uniform': uniform time for the time steps. (**Recommended for high-resolutional data**.) + - 'time_quadratic': quadratic time for the time steps. (Used in DDIM for low-resolutional data.) + t_T: A `float`. The starting time of the sampling (default is T). + t_0: A `float`. The ending time of the sampling (default is epsilon). + device: A torch device. + Returns: + orders: A list of the solver order of each step. + """ + if order == 3: + K = steps // 3 + 1 + if steps % 3 == 0: + orders = [3,] * ( + K - 2 + ) + [2, 1] + elif steps % 3 == 1: + orders = [3,] * ( + K - 1 + ) + [1] + else: + orders = [3,] * ( + K - 1 + ) + [2] + elif order == 2: + if steps % 2 == 0: + K = steps // 2 + orders = [ + 2, + ] * K + else: + K = steps // 2 + 1 + orders = [2,] * ( + K - 1 + ) + [1] + elif order == 1: + K = 1 + orders = [ + 1, + ] * steps + else: + raise ValueError("'order' must be '1' or '2' or '3'.") + if skip_type == "logSNR": + # To reproduce the results in DPM-Solver paper + timesteps_outer = self.get_time_steps(skip_type, t_T, t_0, K, device) + else: + timesteps_outer = self.get_time_steps(skip_type, t_T, t_0, steps, device)[ + torch.cumsum( + torch.tensor( + [ + 0, + ] + + orders + ), + 0, + ).to(device) + ] + return timesteps_outer, orders + + def denoise_to_zero_fn(self, x, s): + """ + Denoise at the final step, which is equivalent to solve the ODE from lambda_s to infty by first-order discretization. + """ + return self.data_prediction_fn(x, s) + + def dpm_solver_first_update(self, x, s, t, model_s=None, return_intermediate=False): + """ + DPM-Solver-1 (equivalent to DDIM) from time `s` to time `t`. + + Args: + x: A pytorch tensor. The initial value at time `s`. + s: A pytorch tensor. The starting time, with the shape (1,). + t: A pytorch tensor. The ending time, with the shape (1,). + model_s: A pytorch tensor. The model function evaluated at time `s`. + If `model_s` is None, we evaluate the model by `x` and `s`; otherwise we directly use it. + return_intermediate: A `bool`. If true, also return the model value at time `s`. + Returns: + x_t: A pytorch tensor. The approximated solution at time `t`. + """ + ns = self.noise_schedule + dims = x.dim() + lambda_s, lambda_t = ns.marginal_lambda(s), ns.marginal_lambda(t) + h = lambda_t - lambda_s + log_alpha_s, log_alpha_t = ns.marginal_log_mean_coeff(s), ns.marginal_log_mean_coeff(t) + sigma_s, sigma_t = ns.marginal_std(s), ns.marginal_std(t) + alpha_t = torch.exp(log_alpha_t) + + if self.algorithm_type == "dpmsolver++": + phi_1 = torch.expm1(-h) + if model_s is None: + model_s = self.model_fn(x, s) + x_t = sigma_t / sigma_s * x - alpha_t * phi_1 * model_s + if return_intermediate: + return x_t, {"model_s": model_s} + else: + return x_t + else: + phi_1 = torch.expm1(h) + if model_s is None: + model_s = self.model_fn(x, s) + x_t = torch.exp(log_alpha_t - log_alpha_s) * x - (sigma_t * phi_1) * model_s + if return_intermediate: + return x_t, {"model_s": model_s} + else: + return x_t + + def singlestep_dpm_solver_second_update( + self, x, s, t, r1=0.5, model_s=None, return_intermediate=False, solver_type="dpmsolver" + ): + """ + Singlestep solver DPM-Solver-2 from time `s` to time `t`. + + Args: + x: A pytorch tensor. The initial value at time `s`. + s: A pytorch tensor. The starting time, with the shape (1,). + t: A pytorch tensor. The ending time, with the shape (1,). + r1: A `float`. The hyperparameter of the second-order solver. + model_s: A pytorch tensor. The model function evaluated at time `s`. + If `model_s` is None, we evaluate the model by `x` and `s`; otherwise we directly use it. + return_intermediate: A `bool`. If true, also return the model value at time `s` and `s1` (the intermediate time). + solver_type: either 'dpmsolver' or 'taylor'. The type for the high-order solvers. + The type slightly impacts the performance. We recommend to use 'dpmsolver' type. + Returns: + x_t: A pytorch tensor. The approximated solution at time `t`. + """ + if solver_type not in ["dpmsolver", "taylor"]: + raise ValueError(f"'solver_type' must be either 'dpmsolver' or 'taylor', got {solver_type}") + if r1 is None: + r1 = 0.5 + ns = self.noise_schedule + lambda_s, lambda_t = ns.marginal_lambda(s), ns.marginal_lambda(t) + h = lambda_t - lambda_s + lambda_s1 = lambda_s + r1 * h + s1 = ns.inverse_lambda(lambda_s1) + log_alpha_s, log_alpha_s1, log_alpha_t = ( + ns.marginal_log_mean_coeff(s), + ns.marginal_log_mean_coeff(s1), + ns.marginal_log_mean_coeff(t), + ) + sigma_s, sigma_s1, sigma_t = ns.marginal_std(s), ns.marginal_std(s1), ns.marginal_std(t) + alpha_s1, alpha_t = torch.exp(log_alpha_s1), torch.exp(log_alpha_t) + + if self.algorithm_type == "dpmsolver++": + phi_11 = torch.expm1(-r1 * h) + phi_1 = torch.expm1(-h) + + if model_s is None: + model_s = self.model_fn(x, s) + x_s1 = (sigma_s1 / sigma_s) * x - (alpha_s1 * phi_11) * model_s + model_s1 = self.model_fn(x_s1, s1) + if solver_type == "dpmsolver": + x_t = ( + (sigma_t / sigma_s) * x + - (alpha_t * phi_1) * model_s + - (0.5 / r1) * (alpha_t * phi_1) * (model_s1 - model_s) + ) + elif solver_type == "taylor": + x_t = ( + (sigma_t / sigma_s) * x + - (alpha_t * phi_1) * model_s + + (1.0 / r1) * (alpha_t * (phi_1 / h + 1.0)) * (model_s1 - model_s) + ) + else: + phi_11 = torch.expm1(r1 * h) + phi_1 = torch.expm1(h) + + if model_s is None: + model_s = self.model_fn(x, s) + x_s1 = torch.exp(log_alpha_s1 - log_alpha_s) * x - (sigma_s1 * phi_11) * model_s + model_s1 = self.model_fn(x_s1, s1) + if solver_type == "dpmsolver": + x_t = ( + torch.exp(log_alpha_t - log_alpha_s) * x + - (sigma_t * phi_1) * model_s + - (0.5 / r1) * (sigma_t * phi_1) * (model_s1 - model_s) + ) + elif solver_type == "taylor": + x_t = ( + torch.exp(log_alpha_t - log_alpha_s) * x + - (sigma_t * phi_1) * model_s + - (1.0 / r1) * (sigma_t * (phi_1 / h - 1.0)) * (model_s1 - model_s) + ) + if return_intermediate: + return x_t, {"model_s": model_s, "model_s1": model_s1} + else: + return x_t + + def singlestep_dpm_solver_third_update( + self, + x, + s, + t, + r1=1.0 / 3.0, + r2=2.0 / 3.0, + model_s=None, + model_s1=None, + return_intermediate=False, + solver_type="dpmsolver", + ): + """ + Singlestep solver DPM-Solver-3 from time `s` to time `t`. + + Args: + x: A pytorch tensor. The initial value at time `s`. + s: A pytorch tensor. The starting time, with the shape (1,). + t: A pytorch tensor. The ending time, with the shape (1,). + r1: A `float`. The hyperparameter of the third-order solver. + r2: A `float`. The hyperparameter of the third-order solver. + model_s: A pytorch tensor. The model function evaluated at time `s`. + If `model_s` is None, we evaluate the model by `x` and `s`; otherwise we directly use it. + model_s1: A pytorch tensor. The model function evaluated at time `s1` (the intermediate time given by `r1`). + If `model_s1` is None, we evaluate the model at `s1`; otherwise we directly use it. + return_intermediate: A `bool`. If true, also return the model value at time `s`, `s1` and `s2` (the intermediate times). + solver_type: either 'dpmsolver' or 'taylor'. The type for the high-order solvers. + The type slightly impacts the performance. We recommend to use 'dpmsolver' type. + Returns: + x_t: A pytorch tensor. The approximated solution at time `t`. + """ + if solver_type not in ["dpmsolver", "taylor"]: + raise ValueError(f"'solver_type' must be either 'dpmsolver' or 'taylor', got {solver_type}") + if r1 is None: + r1 = 1.0 / 3.0 + if r2 is None: + r2 = 2.0 / 3.0 + ns = self.noise_schedule + lambda_s, lambda_t = ns.marginal_lambda(s), ns.marginal_lambda(t) + h = lambda_t - lambda_s + lambda_s1 = lambda_s + r1 * h + lambda_s2 = lambda_s + r2 * h + s1 = ns.inverse_lambda(lambda_s1) + s2 = ns.inverse_lambda(lambda_s2) + log_alpha_s, log_alpha_s1, log_alpha_s2, log_alpha_t = ( + ns.marginal_log_mean_coeff(s), + ns.marginal_log_mean_coeff(s1), + ns.marginal_log_mean_coeff(s2), + ns.marginal_log_mean_coeff(t), + ) + sigma_s, sigma_s1, sigma_s2, sigma_t = ( + ns.marginal_std(s), + ns.marginal_std(s1), + ns.marginal_std(s2), + ns.marginal_std(t), + ) + alpha_s1, alpha_s2, alpha_t = torch.exp(log_alpha_s1), torch.exp(log_alpha_s2), torch.exp(log_alpha_t) + + if self.algorithm_type == "dpmsolver++": + phi_11 = torch.expm1(-r1 * h) + phi_12 = torch.expm1(-r2 * h) + phi_1 = torch.expm1(-h) + phi_22 = torch.expm1(-r2 * h) / (r2 * h) + 1.0 + phi_2 = phi_1 / h + 1.0 + phi_3 = phi_2 / h - 0.5 + + if model_s is None: + model_s = self.model_fn(x, s) + if model_s1 is None: + x_s1 = (sigma_s1 / sigma_s) * x - (alpha_s1 * phi_11) * model_s + model_s1 = self.model_fn(x_s1, s1) + x_s2 = ( + (sigma_s2 / sigma_s) * x + - (alpha_s2 * phi_12) * model_s + + r2 / r1 * (alpha_s2 * phi_22) * (model_s1 - model_s) + ) + model_s2 = self.model_fn(x_s2, s2) + if solver_type == "dpmsolver": + x_t = ( + (sigma_t / sigma_s) * x + - (alpha_t * phi_1) * model_s + + (1.0 / r2) * (alpha_t * phi_2) * (model_s2 - model_s) + ) + elif solver_type == "taylor": + D1_0 = (1.0 / r1) * (model_s1 - model_s) + D1_1 = (1.0 / r2) * (model_s2 - model_s) + D1 = (r2 * D1_0 - r1 * D1_1) / (r2 - r1) + D2 = 2.0 * (D1_1 - D1_0) / (r2 - r1) + x_t = ( + (sigma_t / sigma_s) * x + - (alpha_t * phi_1) * model_s + + (alpha_t * phi_2) * D1 + - (alpha_t * phi_3) * D2 + ) + else: + phi_11 = torch.expm1(r1 * h) + phi_12 = torch.expm1(r2 * h) + phi_1 = torch.expm1(h) + phi_22 = torch.expm1(r2 * h) / (r2 * h) - 1.0 + phi_2 = phi_1 / h - 1.0 + phi_3 = phi_2 / h - 0.5 + + if model_s is None: + model_s = self.model_fn(x, s) + if model_s1 is None: + x_s1 = (torch.exp(log_alpha_s1 - log_alpha_s)) * x - (sigma_s1 * phi_11) * model_s + model_s1 = self.model_fn(x_s1, s1) + x_s2 = ( + (torch.exp(log_alpha_s2 - log_alpha_s)) * x + - (sigma_s2 * phi_12) * model_s + - r2 / r1 * (sigma_s2 * phi_22) * (model_s1 - model_s) + ) + model_s2 = self.model_fn(x_s2, s2) + if solver_type == "dpmsolver": + x_t = ( + (torch.exp(log_alpha_t - log_alpha_s)) * x + - (sigma_t * phi_1) * model_s + - (1.0 / r2) * (sigma_t * phi_2) * (model_s2 - model_s) + ) + elif solver_type == "taylor": + D1_0 = (1.0 / r1) * (model_s1 - model_s) + D1_1 = (1.0 / r2) * (model_s2 - model_s) + D1 = (r2 * D1_0 - r1 * D1_1) / (r2 - r1) + D2 = 2.0 * (D1_1 - D1_0) / (r2 - r1) + x_t = ( + (torch.exp(log_alpha_t - log_alpha_s)) * x + - (sigma_t * phi_1) * model_s + - (sigma_t * phi_2) * D1 + - (sigma_t * phi_3) * D2 + ) + + if return_intermediate: + return x_t, {"model_s": model_s, "model_s1": model_s1, "model_s2": model_s2} + else: + return x_t + + def multistep_dpm_solver_second_update(self, x, model_prev_list, t_prev_list, t, solver_type="dpmsolver"): + """ + Multistep solver DPM-Solver-2 from time `t_prev_list[-1]` to time `t`. + + Args: + x: A pytorch tensor. The initial value at time `s`. + model_prev_list: A list of pytorch tensor. The previous computed model values. + t_prev_list: A list of pytorch tensor. The previous times, each time has the shape (1,) + t: A pytorch tensor. The ending time, with the shape (1,). + solver_type: either 'dpmsolver' or 'taylor'. The type for the high-order solvers. + The type slightly impacts the performance. We recommend to use 'dpmsolver' type. + Returns: + x_t: A pytorch tensor. The approximated solution at time `t`. + """ + if solver_type not in ["dpmsolver", "taylor"]: + raise ValueError(f"'solver_type' must be either 'dpmsolver' or 'taylor', got {solver_type}") + ns = self.noise_schedule + model_prev_1, model_prev_0 = model_prev_list[-2], model_prev_list[-1] + t_prev_1, t_prev_0 = t_prev_list[-2], t_prev_list[-1] + lambda_prev_1, lambda_prev_0, lambda_t = ( + ns.marginal_lambda(t_prev_1), + ns.marginal_lambda(t_prev_0), + ns.marginal_lambda(t), + ) + log_alpha_prev_0, log_alpha_t = ns.marginal_log_mean_coeff(t_prev_0), ns.marginal_log_mean_coeff(t) + sigma_prev_0, sigma_t = ns.marginal_std(t_prev_0), ns.marginal_std(t) + alpha_t = torch.exp(log_alpha_t) + + h_0 = lambda_prev_0 - lambda_prev_1 + h = lambda_t - lambda_prev_0 + r0 = h_0 / h + D1_0 = (1.0 / r0) * (model_prev_0 - model_prev_1) + if self.algorithm_type == "dpmsolver++": + phi_1 = torch.expm1(-h) + if solver_type == "dpmsolver": + x_t = (sigma_t / sigma_prev_0) * x - (alpha_t * phi_1) * model_prev_0 - 0.5 * (alpha_t * phi_1) * D1_0 + elif solver_type == "taylor": + x_t = ( + (sigma_t / sigma_prev_0) * x + - (alpha_t * phi_1) * model_prev_0 + + (alpha_t * (phi_1 / h + 1.0)) * D1_0 + ) + else: + phi_1 = torch.expm1(h) + if solver_type == "dpmsolver": + x_t = ( + (torch.exp(log_alpha_t - log_alpha_prev_0)) * x + - (sigma_t * phi_1) * model_prev_0 + - 0.5 * (sigma_t * phi_1) * D1_0 + ) + elif solver_type == "taylor": + x_t = ( + (torch.exp(log_alpha_t - log_alpha_prev_0)) * x + - (sigma_t * phi_1) * model_prev_0 + - (sigma_t * (phi_1 / h - 1.0)) * D1_0 + ) + return x_t + + def multistep_dpm_solver_third_update(self, x, model_prev_list, t_prev_list, t, solver_type="dpmsolver"): + """ + Multistep solver DPM-Solver-3 from time `t_prev_list[-1]` to time `t`. + + Args: + x: A pytorch tensor. The initial value at time `s`. + model_prev_list: A list of pytorch tensor. The previous computed model values. + t_prev_list: A list of pytorch tensor. The previous times, each time has the shape (1,) + t: A pytorch tensor. The ending time, with the shape (1,). + solver_type: either 'dpmsolver' or 'taylor'. The type for the high-order solvers. + The type slightly impacts the performance. We recommend to use 'dpmsolver' type. + Returns: + x_t: A pytorch tensor. The approximated solution at time `t`. + """ + ns = self.noise_schedule + model_prev_2, model_prev_1, model_prev_0 = model_prev_list + t_prev_2, t_prev_1, t_prev_0 = t_prev_list + lambda_prev_2, lambda_prev_1, lambda_prev_0, lambda_t = ( + ns.marginal_lambda(t_prev_2), + ns.marginal_lambda(t_prev_1), + ns.marginal_lambda(t_prev_0), + ns.marginal_lambda(t), + ) + log_alpha_prev_0, log_alpha_t = ns.marginal_log_mean_coeff(t_prev_0), ns.marginal_log_mean_coeff(t) + sigma_prev_0, sigma_t = ns.marginal_std(t_prev_0), ns.marginal_std(t) + alpha_t = torch.exp(log_alpha_t) + + h_1 = lambda_prev_1 - lambda_prev_2 + h_0 = lambda_prev_0 - lambda_prev_1 + h = lambda_t - lambda_prev_0 + r0, r1 = h_0 / h, h_1 / h + D1_0 = (1.0 / r0) * (model_prev_0 - model_prev_1) + D1_1 = (1.0 / r1) * (model_prev_1 - model_prev_2) + D1 = D1_0 + (r0 / (r0 + r1)) * (D1_0 - D1_1) + D2 = (1.0 / (r0 + r1)) * (D1_0 - D1_1) + if self.algorithm_type == "dpmsolver++": + phi_1 = torch.expm1(-h) + phi_2 = phi_1 / h + 1.0 + phi_3 = phi_2 / h - 0.5 + x_t = ( + (sigma_t / sigma_prev_0) * x + - (alpha_t * phi_1) * model_prev_0 + + (alpha_t * phi_2) * D1 + - (alpha_t * phi_3) * D2 + ) + else: + phi_1 = torch.expm1(h) + phi_2 = phi_1 / h - 1.0 + phi_3 = phi_2 / h - 0.5 + x_t = ( + (torch.exp(log_alpha_t - log_alpha_prev_0)) * x + - (sigma_t * phi_1) * model_prev_0 + - (sigma_t * phi_2) * D1 + - (sigma_t * phi_3) * D2 + ) + return x_t + + def singlestep_dpm_solver_update( + self, x, s, t, order, return_intermediate=False, solver_type="dpmsolver", r1=None, r2=None + ): + """ + Singlestep DPM-Solver with the order `order` from time `s` to time `t`. + + Args: + x: A pytorch tensor. The initial value at time `s`. + s: A pytorch tensor. The starting time, with the shape (1,). + t: A pytorch tensor. The ending time, with the shape (1,). + order: A `int`. The order of DPM-Solver. We only support order == 1 or 2 or 3. + return_intermediate: A `bool`. If true, also return the model value at time `s`, `s1` and `s2` (the intermediate times). + solver_type: either 'dpmsolver' or 'taylor'. The type for the high-order solvers. + The type slightly impacts the performance. We recommend to use 'dpmsolver' type. + r1: A `float`. The hyperparameter of the second-order or third-order solver. + r2: A `float`. The hyperparameter of the third-order solver. + Returns: + x_t: A pytorch tensor. The approximated solution at time `t`. + """ + if order == 1: + return self.dpm_solver_first_update(x, s, t, return_intermediate=return_intermediate) + elif order == 2: + return self.singlestep_dpm_solver_second_update( + x, s, t, return_intermediate=return_intermediate, solver_type=solver_type, r1=r1 + ) + elif order == 3: + return self.singlestep_dpm_solver_third_update( + x, s, t, return_intermediate=return_intermediate, solver_type=solver_type, r1=r1, r2=r2 + ) + else: + raise ValueError(f"Solver order must be 1 or 2 or 3, got {order}") + + def multistep_dpm_solver_update(self, x, model_prev_list, t_prev_list, t, order, solver_type="dpmsolver"): + """ + Multistep DPM-Solver with the order `order` from time `t_prev_list[-1]` to time `t`. + + Args: + x: A pytorch tensor. The initial value at time `s`. + model_prev_list: A list of pytorch tensor. The previous computed model values. + t_prev_list: A list of pytorch tensor. The previous times, each time has the shape (1,) + t: A pytorch tensor. The ending time, with the shape (1,). + order: A `int`. The order of DPM-Solver. We only support order == 1 or 2 or 3. + solver_type: either 'dpmsolver' or 'taylor'. The type for the high-order solvers. + The type slightly impacts the performance. We recommend to use 'dpmsolver' type. + Returns: + x_t: A pytorch tensor. The approximated solution at time `t`. + """ + if order == 1: + return self.dpm_solver_first_update(x, t_prev_list[-1], t, model_s=model_prev_list[-1]) + elif order == 2: + return self.multistep_dpm_solver_second_update(x, model_prev_list, t_prev_list, t, solver_type=solver_type) + elif order == 3: + return self.multistep_dpm_solver_third_update(x, model_prev_list, t_prev_list, t, solver_type=solver_type) + else: + raise ValueError(f"Solver order must be 1 or 2 or 3, got {order}") + + def dpm_solver_adaptive( + self, x, order, t_T, t_0, h_init=0.05, atol=0.0078, rtol=0.05, theta=0.9, t_err=1e-5, solver_type="dpmsolver" + ): + """ + The adaptive step size solver based on singlestep DPM-Solver. + + Args: + x: A pytorch tensor. The initial value at time `t_T`. + order: A `int`. The (higher) order of the solver. We only support order == 2 or 3. + t_T: A `float`. The starting time of the sampling (default is T). + t_0: A `float`. The ending time of the sampling (default is epsilon). + h_init: A `float`. The initial step size (for logSNR). + atol: A `float`. The absolute tolerance of the solver. For image data, the default setting is 0.0078, followed [1]. + rtol: A `float`. The relative tolerance of the solver. The default setting is 0.05. + theta: A `float`. The safety hyperparameter for adapting the step size. The default setting is 0.9, followed [1]. + t_err: A `float`. The tolerance for the time. We solve the diffusion ODE until the absolute error between the + current time and `t_0` is less than `t_err`. The default setting is 1e-5. + solver_type: either 'dpmsolver' or 'taylor'. The type for the high-order solvers. + The type slightly impacts the performance. We recommend to use 'dpmsolver' type. + Returns: + x_0: A pytorch tensor. The approximated solution at time `t_0`. + + [1] A. Jolicoeur-Martineau, K. Li, R. Pichรฉ-Taillefer, T. Kachman, and I. Mitliagkas, "Gotta go fast when generating data with score-based models," arXiv preprint arXiv:2105.14080, 2021. + """ + ns = self.noise_schedule + s = t_T * torch.ones((1,)).to(x) + lambda_s = ns.marginal_lambda(s) + lambda_0 = ns.marginal_lambda(t_0 * torch.ones_like(s).to(x)) + h = h_init * torch.ones_like(s).to(x) + x_prev = x + nfe = 0 + if order == 2: + r1 = 0.5 + lower_update = lambda x, s, t: self.dpm_solver_first_update(x, s, t, return_intermediate=True) + higher_update = lambda x, s, t, **kwargs: self.singlestep_dpm_solver_second_update( + x, s, t, r1=r1, solver_type=solver_type, **kwargs + ) + elif order == 3: + r1, r2 = 1.0 / 3.0, 2.0 / 3.0 + lower_update = lambda x, s, t: self.singlestep_dpm_solver_second_update( + x, s, t, r1=r1, return_intermediate=True, solver_type=solver_type + ) + higher_update = lambda x, s, t, **kwargs: self.singlestep_dpm_solver_third_update( + x, s, t, r1=r1, r2=r2, solver_type=solver_type, **kwargs + ) + else: + raise ValueError(f"For adaptive step size solver, order must be 2 or 3, got {order}") + while torch.abs(s - t_0).mean() > t_err: + t = ns.inverse_lambda(lambda_s + h) + x_lower, lower_noise_kwargs = lower_update(x, s, t) + x_higher = higher_update(x, s, t, **lower_noise_kwargs) + delta = torch.max(torch.ones_like(x).to(x) * atol, rtol * torch.max(torch.abs(x_lower), torch.abs(x_prev))) + norm_fn = lambda v: torch.sqrt(torch.square(v.reshape((v.shape[0], -1))).mean(dim=-1, keepdim=True)) + E = norm_fn((x_higher - x_lower) / delta).max() + if torch.all(E <= 1.0): + x = x_higher + s = t + x_prev = x_lower + lambda_s = ns.marginal_lambda(s) + h = torch.min(theta * h * torch.float_power(E, -1.0 / order).float(), lambda_0 - lambda_s) + nfe += order + print("adaptive solver nfe", nfe) + return x + + def add_noise(self, x, t, noise=None): + """ + Compute the noised input xt = alpha_t * x + sigma_t * noise. + + Args: + x: A `torch.Tensor` with shape `(batch_size, *shape)`. + t: A `torch.Tensor` with shape `(t_size,)`. + Returns: + xt with shape `(t_size, batch_size, *shape)`. + """ + alpha_t, sigma_t = self.noise_schedule.marginal_alpha(t), self.noise_schedule.marginal_std(t) + if noise is None: + noise = torch.randn((t.shape[0], *x.shape), device=x.device) + x = x.reshape((-1, *x.shape)) + xt = expand_dims(alpha_t, x.dim()) * x + expand_dims(sigma_t, x.dim()) * noise + if t.shape[0] == 1: + return xt.squeeze(0) + else: + return xt + + def inverse( + self, + x, + steps=20, + t_start=None, + t_end=None, + order=2, + skip_type="time_uniform", + method="multistep", + lower_order_final=True, + denoise_to_zero=False, + solver_type="dpmsolver", + atol=0.0078, + rtol=0.05, + return_intermediate=False, + ): + """ + Inverse the sample `x` from time `t_start` to `t_end` by DPM-Solver. + For discrete-time DPMs, we use `t_start=1/N`, where `N` is the total time steps during training. + """ + t_0 = 1.0 / self.noise_schedule.total_N if t_start is None else t_start + t_T = self.noise_schedule.T if t_end is None else t_end + assert ( + t_0 > 0 and t_T > 0 + ), "Time range needs to be greater than 0. For discrete-time DPMs, it needs to be in [1 / N, 1], where N is the length of betas array" + return self.sample( + x, + steps=steps, + t_start=t_0, + t_end=t_T, + order=order, + skip_type=skip_type, + method=method, + lower_order_final=lower_order_final, + denoise_to_zero=denoise_to_zero, + solver_type=solver_type, + atol=atol, + rtol=rtol, + return_intermediate=return_intermediate, + ) + + def sample( + self, + x, + steps=20, + t_start=None, + t_end=None, + order=2, + skip_type="time_uniform", + method="multistep", + lower_order_final=True, + denoise_to_zero=False, + solver_type="dpmsolver", + atol=0.0078, + rtol=0.05, + return_intermediate=False, + flow_shift=1.0, + ): + """ + Compute the sample at time `t_end` by DPM-Solver, given the initial `x` at time `t_start`. + + ===================================================== + + We support the following algorithms for both noise prediction model and data prediction model: + - 'singlestep': + Singlestep DPM-Solver (i.e. "DPM-Solver-fast" in the paper), which combines different orders of singlestep DPM-Solver. + We combine all the singlestep solvers with order <= `order` to use up all the function evaluations (steps). + The total number of function evaluations (NFE) == `steps`. + Given a fixed NFE == `steps`, the sampling procedure is: + - If `order` == 1: + - Denote K = steps. We use K steps of DPM-Solver-1 (i.e. DDIM). + - If `order` == 2: + - Denote K = (steps // 2) + (steps % 2). We take K intermediate time steps for sampling. + - If steps % 2 == 0, we use K steps of singlestep DPM-Solver-2. + - If steps % 2 == 1, we use (K - 1) steps of singlestep DPM-Solver-2 and 1 step of DPM-Solver-1. + - If `order` == 3: + - Denote K = (steps // 3 + 1). We take K intermediate time steps for sampling. + - If steps % 3 == 0, we use (K - 2) steps of singlestep DPM-Solver-3, and 1 step of singlestep DPM-Solver-2 and 1 step of DPM-Solver-1. + - If steps % 3 == 1, we use (K - 1) steps of singlestep DPM-Solver-3 and 1 step of DPM-Solver-1. + - If steps % 3 == 2, we use (K - 1) steps of singlestep DPM-Solver-3 and 1 step of singlestep DPM-Solver-2. + - 'multistep': + Multistep DPM-Solver with the order of `order`. The total number of function evaluations (NFE) == `steps`. + We initialize the first `order` values by lower order multistep solvers. + Given a fixed NFE == `steps`, the sampling procedure is: + Denote K = steps. + - If `order` == 1: + - We use K steps of DPM-Solver-1 (i.e. DDIM). + - If `order` == 2: + - We firstly use 1 step of DPM-Solver-1, then use (K - 1) step of multistep DPM-Solver-2. + - If `order` == 3: + - We firstly use 1 step of DPM-Solver-1, then 1 step of multistep DPM-Solver-2, then (K - 2) step of multistep DPM-Solver-3. + - 'singlestep_fixed': + Fixed order singlestep DPM-Solver (i.e. DPM-Solver-1 or singlestep DPM-Solver-2 or singlestep DPM-Solver-3). + We use singlestep DPM-Solver-`order` for `order`=1 or 2 or 3, with total [`steps` // `order`] * `order` NFE. + - 'adaptive': + Adaptive step size DPM-Solver (i.e. "DPM-Solver-12" and "DPM-Solver-23" in the paper). + We ignore `steps` and use adaptive step size DPM-Solver with a higher order of `order`. + You can adjust the absolute tolerance `atol` and the relative tolerance `rtol` to balance the computatation costs + (NFE) and the sample quality. + - If `order` == 2, we use DPM-Solver-12 which combines DPM-Solver-1 and singlestep DPM-Solver-2. + - If `order` == 3, we use DPM-Solver-23 which combines singlestep DPM-Solver-2 and singlestep DPM-Solver-3. + + ===================================================== + + Some advices for choosing the algorithm: + - For **unconditional sampling** or **guided sampling with small guidance scale** by DPMs: + Use singlestep DPM-Solver or DPM-Solver++ ("DPM-Solver-fast" in the paper) with `order = 3`. + e.g., DPM-Solver: + >>> dpm_solver = DPM_Solver(model_fn, noise_schedule, algorithm_type="dpmsolver") + >>> x_sample = dpm_solver.sample(x, steps=steps, t_start=t_start, t_end=t_end, order=3, + skip_type='time_uniform', method='singlestep') + e.g., DPM-Solver++: + >>> dpm_solver = DPM_Solver(model_fn, noise_schedule, algorithm_type="dpmsolver++") + >>> x_sample = dpm_solver.sample(x, steps=steps, t_start=t_start, t_end=t_end, order=3, + skip_type='time_uniform', method='singlestep') + - For **guided sampling with large guidance scale** by DPMs: + Use multistep DPM-Solver with `algorithm_type="dpmsolver++"` and `order = 2`. + e.g. + >>> dpm_solver = DPM_Solver(model_fn, noise_schedule, algorithm_type="dpmsolver++") + >>> x_sample = dpm_solver.sample(x, steps=steps, t_start=t_start, t_end=t_end, order=2, + skip_type='time_uniform', method='multistep') + + We support three types of `skip_type`: + - 'logSNR': uniform logSNR for the time steps. **Recommended for low-resolutional images** + - 'time_uniform': uniform time for the time steps. **Recommended for high-resolutional images**. + - 'time_quadratic': quadratic time for the time steps. + + ===================================================== + Args: + x: A pytorch tensor. The initial value at time `t_start` + e.g. if `t_start` == T, then `x` is a sample from the standard normal distribution. + steps: A `int`. The total number of function evaluations (NFE). + t_start: A `float`. The starting time of the sampling. + If `T` is None, we use self.noise_schedule.T (default is 1.0). + t_end: A `float`. The ending time of the sampling. + If `t_end` is None, we use 1. / self.noise_schedule.total_N. + e.g. if total_N == 1000, we have `t_end` == 1e-3. + For discrete-time DPMs: + - We recommend `t_end` == 1. / self.noise_schedule.total_N. + For continuous-time DPMs: + - We recommend `t_end` == 1e-3 when `steps` <= 15; and `t_end` == 1e-4 when `steps` > 15. + order: A `int`. The order of DPM-Solver. + skip_type: A `str`. The type for the spacing of the time steps. 'time_uniform' or 'logSNR' or 'time_quadratic'. + method: A `str`. The method for sampling. 'singlestep' or 'multistep' or 'singlestep_fixed' or 'adaptive'. + denoise_to_zero: A `bool`. Whether to denoise to time 0 at the final step. + Default is `False`. If `denoise_to_zero` is `True`, the total NFE is (`steps` + 1). + + This trick is firstly proposed by DDPM (https://arxiv.org/abs/2006.11239) and + score_sde (https://arxiv.org/abs/2011.13456). Such trick can improve the FID + for diffusion models sampling by diffusion SDEs for low-resolutional images + (such as CIFAR-10). However, we observed that such trick does not matter for + high-resolutional images. As it needs an additional NFE, we do not recommend + it for high-resolutional images. + lower_order_final: A `bool`. Whether to use lower order solvers at the final steps. + Only valid for `method=multistep` and `steps < 15`. We empirically find that + this trick is a key to stabilizing the sampling by DPM-Solver with very few steps + (especially for steps <= 10). So we recommend to set it to be `True`. + solver_type: A `str`. The taylor expansion type for the solver. `dpmsolver` or `taylor`. We recommend `dpmsolver`. + atol: A `float`. The absolute tolerance of the adaptive step size solver. Valid when `method` == 'adaptive'. + rtol: A `float`. The relative tolerance of the adaptive step size solver. Valid when `method` == 'adaptive'. + return_intermediate: A `bool`. Whether to save the xt at each step. + When set to `True`, method returns a tuple (x0, intermediates); when set to False, method returns only x0. + Returns: + x_end: A pytorch tensor. The approximated solution at time `t_end`. + + """ + t_0 = 1.0 / self.noise_schedule.total_N if t_end is None else t_end + t_T = self.noise_schedule.T if t_start is None else t_start + assert ( + t_0 > 0 and t_T > 0 + ), "Time range needs to be greater than 0. For discrete-time DPMs, it needs to be in [1 / N, 1], where N is the length of betas array" + if return_intermediate: + assert method in [ + "multistep", + "singlestep", + "singlestep_fixed", + ], "Cannot use adaptive solver when saving intermediate values" + if self.correcting_xt_fn is not None: + assert method in [ + "multistep", + "singlestep", + "singlestep_fixed", + ], "Cannot use adaptive solver when correcting_xt_fn is not None" + device = x.device + intermediates = [] + with torch.no_grad(): + if method == "adaptive": + x = self.dpm_solver_adaptive( + x, order=order, t_T=t_T, t_0=t_0, atol=atol, rtol=rtol, solver_type=solver_type + ) + elif method == "multistep": + assert steps >= order + timesteps = self.get_time_steps( + skip_type=skip_type, t_T=t_T, t_0=t_0, N=steps, device=device, shift=flow_shift + ) + assert timesteps.shape[0] - 1 == steps + # Init the initial values. + step = 0 + t = timesteps[step] + t_prev_list = [t] + model_prev_list = [self.model_fn(x, t)] + if self.correcting_xt_fn is not None: + x = self.correcting_xt_fn(x, t, step) + if return_intermediate: + intermediates.append(x) + self.update_progress(step + 1, len(timesteps)) + # Init the first `order` values by lower order multistep DPM-Solver. + for step in range(1, order): + t = timesteps[step] + x = self.multistep_dpm_solver_update( + x, model_prev_list, t_prev_list, t, step, solver_type=solver_type + ) + if self.correcting_xt_fn is not None: + x = self.correcting_xt_fn(x, t, step) + if return_intermediate: + intermediates.append(x) + t_prev_list.append(t) + model_prev_list.append(self.model_fn(x, t)) + # update progress bar + self.update_progress(step + 1, len(timesteps)) + # Compute the remaining values by `order`-th order multistep DPM-Solver. + for step in tqdm(range(order, steps + 1), disable=os.getenv("DPM_TQDM", "False") == "True"): + t = timesteps[step] + # We only use lower order for steps < 10 + # if lower_order_final and steps < 10: + if lower_order_final: # recommended by Shuchen Xue + step_order = min(order, steps + 1 - step) + else: + step_order = order + x = self.multistep_dpm_solver_update( + x, model_prev_list, t_prev_list, t, step_order, solver_type=solver_type + ) + if self.correcting_xt_fn is not None: + x = self.correcting_xt_fn(x, t, step) + if return_intermediate: + intermediates.append(x) + for i in range(order - 1): + t_prev_list[i] = t_prev_list[i + 1] + model_prev_list[i] = model_prev_list[i + 1] + t_prev_list[-1] = t + # We do not need to evaluate the final model value. + if step < steps: + model_prev_list[-1] = self.model_fn(x, t) + # update progress bar + self.update_progress(step + 1, len(timesteps)) + elif method in ["singlestep", "singlestep_fixed"]: + if method == "singlestep": + timesteps_outer, orders = self.get_orders_and_timesteps_for_singlestep_solver( + steps=steps, order=order, skip_type=skip_type, t_T=t_T, t_0=t_0, device=device + ) + elif method == "singlestep_fixed": + K = steps // order + orders = [ + order, + ] * K + timesteps_outer = self.get_time_steps(skip_type=skip_type, t_T=t_T, t_0=t_0, N=K, device=device) + for step, order in enumerate(orders): + s, t = timesteps_outer[step], timesteps_outer[step + 1] + timesteps_inner = self.get_time_steps( + skip_type=skip_type, t_T=s.item(), t_0=t.item(), N=order, device=device + ) + lambda_inner = self.noise_schedule.marginal_lambda(timesteps_inner) + h = lambda_inner[-1] - lambda_inner[0] + r1 = None if order <= 1 else (lambda_inner[1] - lambda_inner[0]) / h + r2 = None if order <= 2 else (lambda_inner[2] - lambda_inner[0]) / h + x = self.singlestep_dpm_solver_update(x, s, t, order, solver_type=solver_type, r1=r1, r2=r2) + if self.correcting_xt_fn is not None: + x = self.correcting_xt_fn(x, t, step) + if return_intermediate: + intermediates.append(x) + self.update_progress(step + 1, len(timesteps_outer)) + else: + raise ValueError(f"Got wrong method {method}") + if denoise_to_zero: + t = torch.ones((1,)).to(device) * t_0 + x = self.denoise_to_zero_fn(x, t) + if self.correcting_xt_fn is not None: + x = self.correcting_xt_fn(x, t, step + 1) + if return_intermediate: + intermediates.append(x) + if return_intermediate: + return x, intermediates + else: + return x + + +############################################################# +# other utility functions +############################################################# + + +def interpolate_fn(x, xp, yp): + """ + A piecewise linear function y = f(x), using xp and yp as keypoints. + We implement f(x) in a differentiable way (i.e. applicable for autograd). + The function f(x) is well-defined for all x-axis. (For x beyond the bounds of xp, we use the outmost points of xp to define the linear function.) + + Args: + x: PyTorch tensor with shape [N, C], where N is the batch size, C is the number of channels (we use C = 1 for DPM-Solver). + xp: PyTorch tensor with shape [C, K], where K is the number of keypoints. + yp: PyTorch tensor with shape [C, K]. + Returns: + The function values f(x), with shape [N, C]. + """ + N, K = x.shape[0], xp.shape[1] + all_x = torch.cat([x.unsqueeze(2), xp.unsqueeze(0).repeat((N, 1, 1))], dim=2) + sorted_all_x, x_indices = torch.sort(all_x, dim=2) + x_idx = torch.argmin(x_indices, dim=2) + cand_start_idx = x_idx - 1 + start_idx = torch.where( + torch.eq(x_idx, 0), + torch.tensor(1, device=x.device), + torch.where( + torch.eq(x_idx, K), + torch.tensor(K - 2, device=x.device), + cand_start_idx, + ), + ) + end_idx = torch.where(torch.eq(start_idx, cand_start_idx), start_idx + 2, start_idx + 1) + start_x = torch.gather(sorted_all_x, dim=2, index=start_idx.unsqueeze(2)).squeeze(2) + end_x = torch.gather(sorted_all_x, dim=2, index=end_idx.unsqueeze(2)).squeeze(2) + start_idx2 = torch.where( + torch.eq(x_idx, 0), + torch.tensor(0, device=x.device), + torch.where( + torch.eq(x_idx, K), + torch.tensor(K - 2, device=x.device), + cand_start_idx, + ), + ) + y_positions_expanded = yp.unsqueeze(0).expand(N, -1, -1) + start_y = torch.gather(y_positions_expanded, dim=2, index=start_idx2.unsqueeze(2)).squeeze(2) + end_y = torch.gather(y_positions_expanded, dim=2, index=(start_idx2 + 1).unsqueeze(2)).squeeze(2) + cand = start_y + (x - start_x) * (end_y - start_y) / (end_x - start_x) + return cand + + +def expand_dims(v, dims): + """ + Expand the tensor `v` to the dim `dims`. + + Args: + `v`: a PyTorch tensor with shape [N]. + `dim`: a `int`. + Returns: + a PyTorch tensor with shape [N, 1, 1, ..., 1] and the total dimension is `dims`. + """ + return v[(...,) + (None,) * (dims - 1)] \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/transport/integrators.py b/examples/OmniGen2-RL/omnigen2/transport/integrators.py new file mode 100644 index 0000000000000000000000000000000000000000..29ac24f310e68b7435771a252c301d1cd7ea1113 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/transport/integrators.py @@ -0,0 +1,122 @@ +import torch as th +from torchdiffeq import odeint +from .utils import time_shift, get_lin_function + +class sde: + """SDE solver class""" + + def __init__( + self, + drift, + diffusion, + *, + t0, + t1, + num_steps, + sampler_type, + ): + assert t0 < t1, "SDE sampler has to be in forward time" + + self.num_timesteps = num_steps + self.t = th.linspace(t0, t1, num_steps) + self.dt = self.t[1] - self.t[0] + self.drift = drift + self.diffusion = diffusion + self.sampler_type = sampler_type + + def __Euler_Maruyama_step(self, x, mean_x, t, model, **model_kwargs): + w_cur = th.randn(x.size()).to(x) + t = th.ones(x.size(0)).to(x) * t + dw = w_cur * th.sqrt(self.dt) + drift = self.drift(x, t, model, **model_kwargs) + diffusion = self.diffusion(x, t) + mean_x = x + drift * self.dt + x = mean_x + th.sqrt(2 * diffusion) * dw + return x, mean_x + + def __Heun_step(self, x, _, t, model, **model_kwargs): + w_cur = th.randn(x.size()).to(x) + dw = w_cur * th.sqrt(self.dt) + t_cur = th.ones(x.size(0)).to(x) * t + diffusion = self.diffusion(x, t_cur) + xhat = x + th.sqrt(2 * diffusion) * dw + K1 = self.drift(xhat, t_cur, model, **model_kwargs) + xp = xhat + self.dt * K1 + K2 = self.drift(xp, t_cur + self.dt, model, **model_kwargs) + return ( + xhat + 0.5 * self.dt * (K1 + K2), + xhat, + ) # at last time point we do not perform the heun step + + def __forward_fn(self): + """TODO: generalize here by adding all private functions ending with steps to it""" + sampler_dict = { + "Euler": self.__Euler_Maruyama_step, + "Heun": self.__Heun_step, + } + + try: + sampler = sampler_dict[self.sampler_type] + except: + raise NotImplementedError("Smapler type not implemented.") + + return sampler + + def sample(self, init, model, **model_kwargs): + """forward loop of sde""" + x = init + mean_x = init + samples = [] + sampler = self.__forward_fn() + for ti in self.t[:-1]: + with th.no_grad(): + x, mean_x = sampler(x, mean_x, ti, model, **model_kwargs) + samples.append(x) + + return samples + + +class ode: + """ODE solver class""" + + def __init__( + self, + drift, + *, + t0, + t1, + sampler_type, + num_steps, + atol, + rtol, + do_shift=False, + time_shifting_factor=None, + ): + assert t0 < t1, "ODE sampler has to be in forward time" + + self.drift = drift + self.do_shift = do_shift + self.t = th.linspace(t0, t1, num_steps) + if time_shifting_factor: + self.t = self.t / (self.t + time_shifting_factor - time_shifting_factor * self.t) + self.atol = atol + self.rtol = rtol + self.sampler_type = sampler_type + + def sample(self, x, model, **model_kwargs): + x = x.float() + device = x[0].device if isinstance(x, tuple) else x.device + + def _fn(t, x): + t = th.ones(x[0].size(0)).to(device) * t if isinstance(x, tuple) else th.ones(x.size(0)).to(device) * t + model_output = self.drift(x, t, model, **model_kwargs).float() + return model_output + + t = self.t.to(device) + if self.do_shift: + mu = get_lin_function(y1=0.5, y2=1.15)(x.shape[1]) + t = time_shift(mu, 1.0, t) + atol = [self.atol] * len(x) if isinstance(x, tuple) else [self.atol] + rtol = [self.rtol] * len(x) if isinstance(x, tuple) else [self.rtol] + samples = odeint(_fn, x, t, method=self.sampler_type, atol=atol, rtol=rtol) + return samples diff --git a/examples/OmniGen2-RL/omnigen2/transport/path.py b/examples/OmniGen2-RL/omnigen2/transport/path.py new file mode 100644 index 0000000000000000000000000000000000000000..3a5b1ea132c03dd324ada858053d23aa76de9be6 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/transport/path.py @@ -0,0 +1,201 @@ +import numpy as np +import torch as th + + +def expand_t_like_x(t, x): + """Function to reshape time t to broadcastable dimension of x + Args: + t: [batch_dim,], time vector + x: [batch_dim,...], data point + """ + dims = [1] * len(x[0].size()) + t = t.view(t.size(0), *dims) + return t + + +#################### Coupling Plans #################### + + +class ICPlan: + """Linear Coupling Plan""" + + def __init__(self, sigma=0.0): + self.sigma = sigma + + def compute_alpha_t(self, t): + """Compute the data coefficient along the path""" + return t, 1 + + def compute_sigma_t(self, t): + """Compute the noise coefficient along the path""" + return 1 - t, -1 + + def compute_d_alpha_alpha_ratio_t(self, t): + """Compute the ratio between d_alpha and alpha""" + return 1 / t + + def compute_drift(self, x, t): + """We always output sde according to score parametrization;""" + t = expand_t_like_x(t, x) + alpha_ratio = self.compute_d_alpha_alpha_ratio_t(t) + sigma_t, d_sigma_t = self.compute_sigma_t(t) + drift = alpha_ratio * x + diffusion = alpha_ratio * (sigma_t**2) - sigma_t * d_sigma_t + + return -drift, diffusion + + def compute_diffusion(self, x, t, form="constant", norm=1.0): + """Compute the diffusion term of the SDE + Args: + x: [batch_dim, ...], data point + t: [batch_dim,], time vector + form: str, form of the diffusion term + norm: float, norm of the diffusion term + """ + t = expand_t_like_x(t, x) + choices = { + "constant": norm, + "SBDM": norm * self.compute_drift(x, t)[1], + "sigma": norm * self.compute_sigma_t(t)[0], + "linear": norm * (1 - t), + "decreasing": 0.25 * (norm * th.cos(np.pi * t) + 1) ** 2, + "inccreasing-decreasing": norm * th.sin(np.pi * t) ** 2, + } + + try: + diffusion = choices[form] + except KeyError: + raise NotImplementedError(f"Diffusion form {form} not implemented") + + return diffusion + + def get_score_from_velocity(self, velocity, x, t): + """Wrapper function: transfrom velocity prediction model to score + Args: + velocity: [batch_dim, ...] shaped tensor; velocity model output + x: [batch_dim, ...] shaped tensor; x_t data point + t: [batch_dim,] time tensor + """ + t = expand_t_like_x(t, x) + alpha_t, d_alpha_t = self.compute_alpha_t(t) + sigma_t, d_sigma_t = self.compute_sigma_t(t) + mean = x + reverse_alpha_ratio = alpha_t / d_alpha_t + var = sigma_t**2 - reverse_alpha_ratio * d_sigma_t * sigma_t + score = (reverse_alpha_ratio * velocity - mean) / var + return score + + def get_noise_from_velocity(self, velocity, x, t): + """Wrapper function: transfrom velocity prediction model to denoiser + Args: + velocity: [batch_dim, ...] shaped tensor; velocity model output + x: [batch_dim, ...] shaped tensor; x_t data point + t: [batch_dim,] time tensor + """ + t = expand_t_like_x(t, x) + alpha_t, d_alpha_t = self.compute_alpha_t(t) + sigma_t, d_sigma_t = self.compute_sigma_t(t) + mean = x + reverse_alpha_ratio = alpha_t / d_alpha_t + var = reverse_alpha_ratio * d_sigma_t - sigma_t + noise = (reverse_alpha_ratio * velocity - mean) / var + return noise + + def get_velocity_from_score(self, score, x, t): + """Wrapper function: transfrom score prediction model to velocity + Args: + score: [batch_dim, ...] shaped tensor; score model output + x: [batch_dim, ...] shaped tensor; x_t data point + t: [batch_dim,] time tensor + """ + t = expand_t_like_x(t, x) + drift, var = self.compute_drift(x, t) + velocity = var * score - drift + return velocity + + def compute_mu_t(self, t, x0, x1): + """Compute the mean of time-dependent density p_t""" + t = expand_t_like_x(t, x1) + alpha_t, _ = self.compute_alpha_t(t) + sigma_t, _ = self.compute_sigma_t(t) + if isinstance(x1, (list, tuple)): + return [alpha_t[i] * x1[i] + sigma_t[i] * x0[i] for i in range(len(x1))] + else: + return alpha_t * x1 + sigma_t * x0 + + def compute_xt(self, t, x0, x1): + """Sample xt from time-dependent density p_t; rng is required""" + xt = self.compute_mu_t(t, x0, x1) + return xt + + def compute_ut(self, t, x0, x1, xt): + """Compute the vector field corresponding to p_t""" + t = expand_t_like_x(t, x1) + _, d_alpha_t = self.compute_alpha_t(t) + _, d_sigma_t = self.compute_sigma_t(t) + if isinstance(x1, (list, tuple)): + return [d_alpha_t * x1[i] + d_sigma_t * x0[i] for i in range(len(x1))] + else: + return d_alpha_t * x1 + d_sigma_t * x0 + + def plan(self, t, x0, x1): + xt = self.compute_xt(t, x0, x1) + ut = self.compute_ut(t, x0, x1, xt) + return t, xt, ut + + +class VPCPlan(ICPlan): + """class for VP path flow matching""" + + def __init__(self, sigma_min=0.1, sigma_max=20.0): + self.sigma_min = sigma_min + self.sigma_max = sigma_max + self.log_mean_coeff = ( + lambda t: -0.25 * ((1 - t) ** 2) * (self.sigma_max - self.sigma_min) - 0.5 * (1 - t) * self.sigma_min + ) + self.d_log_mean_coeff = lambda t: 0.5 * (1 - t) * (self.sigma_max - self.sigma_min) + 0.5 * self.sigma_min + + def compute_alpha_t(self, t): + """Compute coefficient of x1""" + alpha_t = self.log_mean_coeff(t) + alpha_t = th.exp(alpha_t) + d_alpha_t = alpha_t * self.d_log_mean_coeff(t) + return alpha_t, d_alpha_t + + def compute_sigma_t(self, t): + """Compute coefficient of x0""" + p_sigma_t = 2 * self.log_mean_coeff(t) + sigma_t = th.sqrt(1 - th.exp(p_sigma_t)) + d_sigma_t = th.exp(p_sigma_t) * (2 * self.d_log_mean_coeff(t)) / (-2 * sigma_t) + return sigma_t, d_sigma_t + + def compute_d_alpha_alpha_ratio_t(self, t): + """Special purposed function for computing numerical stabled d_alpha_t / alpha_t""" + return self.d_log_mean_coeff(t) + + def compute_drift(self, x, t): + """Compute the drift term of the SDE""" + t = expand_t_like_x(t, x) + beta_t = self.sigma_min + (1 - t) * (self.sigma_max - self.sigma_min) + return -0.5 * beta_t * x, beta_t / 2 + + +class GVPCPlan(ICPlan): + def __init__(self, sigma=0.0): + super().__init__(sigma) + + def compute_alpha_t(self, t): + """Compute coefficient of x1""" + alpha_t = th.sin(t * np.pi / 2) + d_alpha_t = np.pi / 2 * th.cos(t * np.pi / 2) + return alpha_t, d_alpha_t + + def compute_sigma_t(self, t): + """Compute coefficient of x0""" + sigma_t = th.cos(t * np.pi / 2) + d_sigma_t = -np.pi / 2 * th.sin(t * np.pi / 2) + return sigma_t, d_sigma_t + + def compute_d_alpha_alpha_ratio_t(self, t): + """Special purposed function for computing numerical stabled d_alpha_t / alpha_t""" + return np.pi / (2 * th.tan(t * np.pi / 2)) diff --git a/examples/OmniGen2-RL/omnigen2/transport/transport.py b/examples/OmniGen2-RL/omnigen2/transport/transport.py new file mode 100644 index 0000000000000000000000000000000000000000..f0d3dcec89d181d1d5137bdea0c7b24491515c70 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/transport/transport.py @@ -0,0 +1,545 @@ +import enum +import math +from typing import Callable, Optional + +import numpy as np +import torch as th +import random + +from . import path +from .integrators import ode, sde +from .utils import mean_flat, expand_dims +from .dpm_solver import NoiseScheduleFlow, model_wrapper, DPM_Solver + + +class ModelType(enum.Enum): + """ + Which type of output the model predicts. + """ + + NOISE = enum.auto() # the model predicts epsilon + SCORE = enum.auto() # the model predicts \nabla \log p(x) + VELOCITY = enum.auto() # the model predicts v(x) + + +class PathType(enum.Enum): + """ + Which type of path to use. + """ + + LINEAR = enum.auto() + GVP = enum.auto() + VP = enum.auto() + + +class WeightType(enum.Enum): + """ + Which type of weighting to use. + """ + + NONE = enum.auto() + VELOCITY = enum.auto() + LIKELIHOOD = enum.auto() + + +class Transport: + def __init__(self, *, model_type, path_type, loss_type, train_eps, sample_eps, snr_type, do_shift, seq_len, + dynamic_time_shift: bool = False, + time_shift_version: str = "v1"): + path_options = { + PathType.LINEAR: path.ICPlan, + PathType.GVP: path.GVPCPlan, + PathType.VP: path.VPCPlan, + } + + self.loss_type = loss_type + self.model_type = model_type + self.path_sampler = path_options[path_type]() + self.train_eps = train_eps + self.sample_eps = sample_eps + + self.snr_type = snr_type + self.do_shift = do_shift + self.seq_len = seq_len + self.dynamic_time_shift = dynamic_time_shift + self.time_shift_version = time_shift_version + def prior_logp(self, z): + """ + Standard multivariate normal prior + Assume z is batched + """ + shape = th.tensor(z.size()) + N = th.prod(shape[1:]) + _fn = lambda x: -N / 2.0 * np.log(2 * np.pi) - th.sum(x**2) / 2.0 + return th.vmap(_fn)(z) + + def check_interval( + self, + train_eps, + sample_eps, + *, + diffusion_form="SBDM", + sde=False, + reverse=False, + eval=False, + last_step_size=0.0, + ): + t0 = 0 + t1 = 1 + eps = train_eps if not eval else sample_eps + if type(self.path_sampler) in [path.VPCPlan]: + t1 = 1 - eps if (not sde or last_step_size == 0) else 1 - last_step_size + + elif (type(self.path_sampler) in [path.ICPlan, path.GVPCPlan]) and ( + self.model_type != ModelType.VELOCITY or sde + ): # avoid numerical issue by taking a first semi-implicit step + t0 = eps if (diffusion_form == "SBDM" and sde) or self.model_type != ModelType.VELOCITY else 0 + t1 = 1 - eps if (not sde or last_step_size == 0) else 1 - last_step_size + + if reverse: + t0, t1 = 1 - t0, 1 - t1 + + return t0, t1 + + def sample(self, x1, process_index, num_processes): + """Sampling x0 & t based on shape of x1 (if needed) + Args: + x1 - data point; [batch, *dim] + """ + if isinstance(x1, (list, tuple)): + x0 = [th.randn_like(img_start) for img_start in x1] + else: + x0 = th.randn_like(x1) + t0, t1 = self.check_interval(self.train_eps, self.sample_eps) + + if self.snr_type.startswith("uniform"): + assert t0 == 0.0 and t1 == 1.0, "not implemented." + if "_" in self.snr_type: + _, t0, t1 = self.snr_type.split("_") + t0, t1 = float(t0), float(t1) + t = th.rand((len(x1),)) * (t1 - t0) + t0 + if self.snr_type == "stratified_uniform": + batch_size = len(x1) + n = batch_size * num_processes + offsets = th.arange(process_index, n, num_processes) + u = th.rand(size=(batch_size,)) + t = ((offsets + u) / n) + elif self.snr_type == "lognorm": + u = th.normal(mean=0.0, std=1.0, size=(len(x1),)) + t = 1 / (1 + th.exp(-u)) * (t1 - t0) + t0 + elif self.snr_type == "zero": + t = th.rand((len(x1),)) + for _ in range(len(x1)): + if random.random() < 1.0: + t[_] = 0.0 + # print(t) + else: + raise NotImplementedError("Not implemented snr_type %s" % self.snr_type) + + if self.do_shift: + if self.dynamic_time_shift: + if self.time_shift_version == "v1": + base_shift: float = 0.5 + max_shift: float = 1.15 + lin_func = self.get_lin_function(y1=base_shift, y2=max_shift) + + mu = th.tensor([lin_func((_x1.shape[-2] // 2) * (_x1.shape[-1] // 2)) for _x1 in x1], dtype=t.dtype, device=t.device).view_as(t) + t = self.time_shift(mu, 1.0, t) + elif self.time_shift_version == "v2": + tokens = th.tensor([(_x1.shape[-2] // 2) * (_x1.shape[-1] // 2) for _x1 in x1], dtype=t.dtype, device=t.device).view_as(t) + t = self.time_shift_v2(tokens, t) + else: + if self.time_shift_version == "v1": + base_shift: float = 0.5 + max_shift: float = 1.15 + mu = self.get_lin_function(y1=base_shift, y2=max_shift)(self.seq_len) + t = self.time_shift(mu, 1.0, t) + elif self.time_shift_version == "v2": + tokens = th.tensor([self.seq_len] * len(x1), dtype=t.dtype, device=t.device).view_as(t) + t = self.time_shift_v2(tokens, t) + t = t.to(x1[0]) + return t, x0, x1 + + def time_shift(self, mu: float, sigma: float, t: th.Tensor): + # the following implementation was original for t=0: clean / t=1: noise + # Since we adopt the reverse, the 1-t operations are needed + t = 1 - t + if isinstance(mu, th.Tensor): + t = th.exp(mu) / (th.exp(mu) + (1 / t - 1) ** sigma) + else: + t = math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma) + t = 1 - t + return t + + def time_shift_v2(self, tokens: th.Tensor, t: th.Tensor): + # t = th.exp(mu) / (th.exp(mu) + (1 / t - 1) ** sigma) + m = th.sqrt(tokens) / 20 + t = t / (m - m * t + t) + return t + + def get_lin_function( + self, x1: float = 256, y1: float = 0.5, x2: float = 4096, y2: float = 1.15 + ) -> Callable[[float], float]: + m = (y2 - y1) / (x2 - x1) + b = y1 - m * x1 + return lambda x: m * x + b + + def training_losses( + self, + model, + x1, + model_kwargs=None, + process_index: Optional[int] = None, + num_processes: Optional[int] = None, + reduction: str = 'mean', + ): + """Loss for training the score model + Args: + - model: backbone model; could be score, noise, or velocity + - x1: datapoint + - model_kwargs: additional arguments for the model + """ + + terms = {} + + if model_kwargs is None: + model_kwargs = {} + t, x0, x1 = self.sample(x1, process_index, num_processes) + t, xt, ut = self.path_sampler.plan(t, x0, x1) + + terms = {} + terms['t'] = t + terms['xt'] = xt + + if "cond" in model_kwargs: + conds = model_kwargs.pop("cond") + xt = [th.cat([x, cond], dim=0) if cond is not None else x for x, cond in zip(xt, conds)] + model_output = model(xt, t, **model_kwargs) + B = len(x0) + + terms['pred'] = model_output + if self.model_type == ModelType.VELOCITY: + if isinstance(x1, (list, tuple)): + assert len(model_output) == len(ut) == len(x1) + for i in range(B): + assert ( + model_output[i].shape == ut[i].shape == x1[i].shape + ), f"{model_output[i].shape} {ut[i].shape} {x1[i].shape}" + terms["task_loss"] = th.stack( + [th.nn.functional.mse_loss(ut[i].float(), model_output[i].float(), reduction=reduction) for i in range(B)], + dim=0, + ) + else: + terms["task_loss"] = mean_flat(((model_output - ut) ** 2)) + else: + raise NotImplementedError + + terms["loss"] = terms["task_loss"] + terms["t"] = t + return terms + + def get_drift(self): + """member function for obtaining the drift of the probability flow ODE""" + + def score_ode(x, t, model, **model_kwargs): + drift_mean, drift_var = self.path_sampler.compute_drift(x, t) + model_output = model(x, t, **model_kwargs) + return -drift_mean + drift_var * model_output # by change of variable + + def noise_ode(x, t, model, **model_kwargs): + drift_mean, drift_var = self.path_sampler.compute_drift(x, t) + sigma_t, _ = self.path_sampler.compute_sigma_t(path.expand_t_like_x(t, x)) + model_output = model(x, t, **model_kwargs) + score = model_output / -sigma_t + return -drift_mean + drift_var * score + + def velocity_ode(x, t, model, **model_kwargs): + model_output = model(x, t, **model_kwargs) + return model_output + + if self.model_type == ModelType.NOISE: + drift_fn = noise_ode + elif self.model_type == ModelType.SCORE: + drift_fn = score_ode + else: + drift_fn = velocity_ode + + def body_fn(x, t, model, **model_kwargs): + model_output = drift_fn(x, t, model, **model_kwargs) + assert model_output.shape == x.shape, "Output shape from ODE solver must match input shape" + return model_output + + return body_fn + + def get_score( + self, + ): + """member function for obtaining score of + x_t = alpha_t * x + sigma_t * eps""" + if self.model_type == ModelType.NOISE: + score_fn = ( + lambda x, t, model, **kwargs: model(x, t, **kwargs) + / -self.path_sampler.compute_sigma_t(path.expand_t_like_x(t, x))[0] + ) + elif self.model_type == ModelType.SCORE: + score_fn = lambda x, t, model, **kwagrs: model(x, t, **kwagrs) + elif self.model_type == ModelType.VELOCITY: + score_fn = lambda x, t, model, **kwargs: self.path_sampler.get_score_from_velocity( + model(x, t, **kwargs), x, t + ) + else: + raise NotImplementedError() + + return score_fn + + +class Sampler: + """Sampler class for the transport model""" + + def __init__( + self, + transport, + ): + """Constructor for a general sampler; supporting different sampling methods + Args: + - transport: an tranport object specify model prediction & interpolant type + """ + + self.transport = transport + self.drift = self.transport.get_drift() + self.score = self.transport.get_score() + + def __get_sde_diffusion_and_drift( + self, + *, + diffusion_form="SBDM", + diffusion_norm=1.0, + ): + def diffusion_fn(x, t): + diffusion = self.transport.path_sampler.compute_diffusion(x, t, form=diffusion_form, norm=diffusion_norm) + return diffusion + + sde_drift = lambda x, t, model, **kwargs: self.drift(x, t, model, **kwargs) + diffusion_fn(x, t) * self.score( + x, t, model, **kwargs + ) + + sde_diffusion = diffusion_fn + + return sde_drift, sde_diffusion + + def __get_last_step( + self, + sde_drift, + *, + last_step, + last_step_size, + ): + """Get the last step function of the SDE solver""" + + if last_step is None: + last_step_fn = lambda x, t, model, **model_kwargs: x + elif last_step == "Mean": + last_step_fn = ( + lambda x, t, model, **model_kwargs: x + sde_drift(x, t, model, **model_kwargs) * last_step_size + ) + elif last_step == "Tweedie": + alpha = self.transport.path_sampler.compute_alpha_t # simple aliasing; the original name was too long + sigma = self.transport.path_sampler.compute_sigma_t + last_step_fn = lambda x, t, model, **model_kwargs: x / alpha(t)[0][0] + (sigma(t)[0][0] ** 2) / alpha(t)[0][ + 0 + ] * self.score(x, t, model, **model_kwargs) + elif last_step == "Euler": + last_step_fn = ( + lambda x, t, model, **model_kwargs: x + self.drift(x, t, model, **model_kwargs) * last_step_size + ) + else: + raise NotImplementedError() + + return last_step_fn + + def sample_sde( + self, + *, + sampling_method="Euler", + diffusion_form="SBDM", + diffusion_norm=1.0, + last_step="Mean", + last_step_size=0.04, + num_steps=250, + ): + """returns a sampling function with given SDE settings + Args: + - sampling_method: type of sampler used in solving the SDE; default to be Euler-Maruyama + - diffusion_form: function form of diffusion coefficient; default to be matching SBDM + - diffusion_norm: function magnitude of diffusion coefficient; default to 1 + - last_step: type of the last step; default to identity + - last_step_size: size of the last step; default to match the stride of 250 steps over [0,1] + - num_steps: total integration step of SDE + """ + + if last_step is None: + last_step_size = 0.0 + + sde_drift, sde_diffusion = self.__get_sde_diffusion_and_drift( + diffusion_form=diffusion_form, + diffusion_norm=diffusion_norm, + ) + + t0, t1 = self.transport.check_interval( + self.transport.train_eps, + self.transport.sample_eps, + diffusion_form=diffusion_form, + sde=True, + eval=True, + reverse=False, + last_step_size=last_step_size, + ) + + _sde = sde( + sde_drift, + sde_diffusion, + t0=t0, + t1=t1, + num_steps=num_steps, + sampler_type=sampling_method, + ) + + last_step_fn = self.__get_last_step(sde_drift, last_step=last_step, last_step_size=last_step_size) + + def _sample(init, model, **model_kwargs): + xs = _sde.sample(init, model, **model_kwargs) + ts = th.ones(init.size(0), device=init.device) * t1 + x = last_step_fn(xs[-1], ts, model, **model_kwargs) + xs.append(x) + + assert len(xs) == num_steps, "Samples does not match the number of steps" + + return xs + + return _sample + + def sample_dpm( + self, + model, + model_kwargs=None, + ): + + noise_schedule = NoiseScheduleFlow(schedule="discrete_flow") + + def noise_pred_fn(x, t_continuous): + output = model(x, 1 - t_continuous, **model_kwargs) + _, sigma_t = noise_schedule.marginal_alpha(t_continuous), noise_schedule.marginal_std(t_continuous) + try: + noise = x - (1 - expand_dims(sigma_t, x.dim()).to(x)) * output + except: + noise = x - (1 - expand_dims(sigma_t, x.dim()).to(x)) * output[0] + return noise + + return DPM_Solver(noise_pred_fn, noise_schedule, algorithm_type="dpmsolver++").sample + + + def sample_ode( + self, + *, + sampling_method="dopri5", + num_steps=50, + atol=1e-6, + rtol=1e-3, + reverse=False, + do_shift=False, + time_shifting_factor=None, + ): + """returns a sampling function with given ODE settings + Args: + - sampling_method: type of sampler used in solving the ODE; default to be Dopri5 + - num_steps: + - fixed solver (Euler, Heun): the actual number of integration steps performed + - adaptive solver (Dopri5): the number of datapoints saved during integration; produced by interpolation + - atol: absolute error tolerance for the solver + - rtol: relative error tolerance for the solver + """ + + # for flux + drift = lambda x, t, model, **kwargs: self.drift(x, t, model, **kwargs) + + t0, t1 = self.transport.check_interval( + self.transport.train_eps, + self.transport.sample_eps, + sde=False, + eval=True, + reverse=reverse, + last_step_size=0.0, + ) + + _ode = ode( + drift=drift, + t0=t0, + t1=t1, + sampler_type=sampling_method, + num_steps=num_steps, + atol=atol, + rtol=rtol, + do_shift=do_shift, + time_shifting_factor=time_shifting_factor, + ) + + return _ode.sample + + def sample_ode_likelihood( + self, + *, + sampling_method="dopri5", + num_steps=50, + atol=1e-6, + rtol=1e-3, + ): + """returns a sampling function for calculating likelihood with given ODE settings + Args: + - sampling_method: type of sampler used in solving the ODE; default to be Dopri5 + - num_steps: + - fixed solver (Euler, Heun): the actual number of integration steps performed + - adaptive solver (Dopri5): the number of datapoints saved during integration; produced by interpolation + - atol: absolute error tolerance for the solver + - rtol: relative error tolerance for the solver + """ + + def _likelihood_drift(x, t, model, **model_kwargs): + x, _ = x + eps = th.randint(2, x.size(), dtype=th.float, device=x.device) * 2 - 1 + t = th.ones_like(t) * (1 - t) + with th.enable_grad(): + x.requires_grad = True + grad = th.autograd.grad(th.sum(self.drift(x, t, model, **model_kwargs) * eps), x)[0] + logp_grad = th.sum(grad * eps, dim=tuple(range(1, len(x.size())))) + drift = self.drift(x, t, model, **model_kwargs) + return (-drift, logp_grad) + + t0, t1 = self.transport.check_interval( + self.transport.train_eps, + self.transport.sample_eps, + sde=False, + eval=True, + reverse=False, + last_step_size=0.0, + ) + + _ode = ode( + drift=_likelihood_drift, + t0=t0, + t1=t1, + sampler_type=sampling_method, + num_steps=num_steps, + atol=atol, + rtol=rtol, + ) + + def _sample_fn(x, model, **model_kwargs): + init_logp = th.zeros(x.size(0)).to(x) + input = (x, init_logp) + drift, delta_logp = _ode.sample(input, model, **model_kwargs) + drift, delta_logp = drift[-1], delta_logp[-1] + prior_logp = self.transport.prior_logp(drift) + logp = prior_logp - delta_logp + return logp, drift + + return _sample_fn diff --git a/examples/OmniGen2-RL/omnigen2/transport/utils.py b/examples/OmniGen2-RL/omnigen2/transport/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..9af21dc828f7ed4df10c4e66dd39ca581f3cc1aa --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/transport/utils.py @@ -0,0 +1,56 @@ +import torch as th +import math + +class EasyDict: + def __init__(self, sub_dict): + for k, v in sub_dict.items(): + setattr(self, k, v) + + def __getitem__(self, key): + return getattr(self, key) + + +def mean_flat(x): + """ + Take the mean over all non-batch dimensions. + """ + return th.mean(x, dim=list(range(1, len(x.size())))) + + +def log_state(state): + result = [] + + sorted_state = dict(sorted(state.items())) + for key, value in sorted_state.items(): + # Check if the value is an instance of a class + if " Image.Image: + """Create a horizontal collage from a list of images.""" + max_height = max(img.shape[-2] for img in images) + total_width = sum(img.shape[-1] for img in images) + canvas = torch.zeros((3, max_height, total_width), device=images[0].device) + + current_x = 0 + for img in images: + h, w = img.shape[-2:] + canvas[:, :h, current_x:current_x+w] = img * 0.5 + 0.5 + current_x += w + + return to_pil_image(canvas) \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/utils/import_utils.py b/examples/OmniGen2-RL/omnigen2/utils/import_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..dc946d77782b66ccb8d6402b2e521a9e5de4e81c --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/utils/import_utils.py @@ -0,0 +1,46 @@ +# Copyright 2024 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +""" +Import utilities: Utilities related to imports and our lazy inits. +""" + +import importlib.util +import sys + +# The package importlib_metadata is in a different place, depending on the python version. +if sys.version_info < (3, 8): + import importlib_metadata +else: + import importlib.metadata as importlib_metadata + +def _is_package_available(pkg_name: str): + pkg_exists = importlib.util.find_spec(pkg_name) is not None + pkg_version = "N/A" + + if pkg_exists: + try: + pkg_version = importlib_metadata.version(pkg_name) + except (ImportError, importlib_metadata.PackageNotFoundError): + pkg_exists = False + + return pkg_exists, pkg_version + +_triton_available, _triton_version = _is_package_available("triton") +_flash_attn_available, _flash_attn_version = _is_package_available("flash_attn") + +def is_triton_available(): + return _triton_available + +def is_flash_attn_available(): + return _flash_attn_available \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/utils/logging_utils.py b/examples/OmniGen2-RL/omnigen2/utils/logging_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..0b942a74e1d4c43c3cd2c2df7e28849e69de87c4 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/utils/logging_utils.py @@ -0,0 +1,15 @@ +import logging + +class TqdmToLogger(object): + """File-like object to redirect tqdm output to a logger.""" + def __init__(self, logger, level=logging.INFO): + self.logger = logger + self.level = level + + def write(self, buf): + for line in buf.rstrip().splitlines(): + self.logger.log(self.level, line) + + def flush(self): + for handler in self.logger.logger.handlers: + handler.flush() \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/utils/reproducibility.py b/examples/OmniGen2-RL/omnigen2/utils/reproducibility.py new file mode 100644 index 0000000000000000000000000000000000000000..0e89c86ee290690d061dda76e178a03d3f4bc742 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/utils/reproducibility.py @@ -0,0 +1,22 @@ +import random +import numpy as np + +import torch + +from diffusers.utils import is_torch_npu_available + + +def worker_init_fn(worker_id, num_processes, num_workers, process_index, seed, same_seed_per_epoch=False): + if same_seed_per_epoch: + worker_seed = seed + num_processes + num_workers * process_index + worker_id + else: + worker_seed = torch.initial_seed() + + random.seed(worker_seed) + np.random.seed(worker_seed % 2**32) + torch.manual_seed(worker_seed) + + if is_torch_npu_available(): + torch.npu.manual_seed_all(seed) + else: + torch.cuda.manual_seed_all(seed) \ No newline at end of file diff --git a/examples/OmniGen2-RL/omnigen2/utils/tensor_util.py b/examples/OmniGen2-RL/omnigen2/utils/tensor_util.py new file mode 100644 index 0000000000000000000000000000000000000000..a2469f5822b70cee8cc08beb91af1f263ebec4a9 --- /dev/null +++ b/examples/OmniGen2-RL/omnigen2/utils/tensor_util.py @@ -0,0 +1,26 @@ +from typing import List + +from PIL import Image + +from torch.nn import functional as F + +def pad_to_length(tensor, len, dim=1): + return F.pad(tensor, [0] * ((tensor.dim() - 2) * 2) + [0, len - tensor.shape[dim]]) + +def expand_as(tensor, other): + """ + Expands a tensor to match the dimensions of another tensor. + + If tensor has shape [b] and other has shape [b, c, h, w], + this function will reshape tensor to [b, 1, 1, 1] to enable broadcasting. + + Args: + tensor (`torch.FloatTensor`): The tensor to expand + other (`torch.FloatTensor`): The tensor whose shape will be matched + + Returns: + `torch.FloatTensor`: The expanded tensor + """ + for _ in range(other.ndim - tensor.ndim): + tensor = tensor.unsqueeze(-1) + return tensor \ No newline at end of file diff --git a/examples/OmniGen2-RL/options/omnigen2_edit_rl_4machine_editscore7b_avg4.yml b/examples/OmniGen2-RL/options/omnigen2_edit_rl_4machine_editscore7b_avg4.yml new file mode 100644 index 0000000000000000000000000000000000000000..673325bebd1ee7c450d15eed244aa2233ce0149b --- /dev/null +++ b/examples/OmniGen2-RL/options/omnigen2_edit_rl_4machine_editscore7b_avg4.yml @@ -0,0 +1,122 @@ +name: omnigen2_edit_rl_4machine_editscore7b_avg4 + +seed: 2233 +device_specific_seed: true +workder_specific_seed: true + +reward_server_config: reward_server/server_configs/editscore_7B_avg4.yml + +data: + data_path: data_configs/train/example/train.yml + use_chat_template: true + maximum_text_tokens: 888 + prompt_dropout_prob: !!float 0.0 + ref_img_dropout_prob: !!float 0.0 + max_output_pixels: 262144 # 512 * 512 + max_input_pixels: [262144, 262144, 262144, 262144] # [512 * 512, 512 * 512, 512 * 512, 512 * 512] + max_side_length: 2048 + +model: + pretrained_vae_model_name_or_path: black-forest-labs/FLUX.1-dev + pretrained_text_encoder_model_name_or_path: Qwen/Qwen2.5-VL-3B-Instruct + pretrained_model_path: pretrained_models/OmniGen2/transformer/pytorch_model.bin + + arch_opt: + patch_size: 2 + in_channels: 16 + hidden_size: 2520 + num_layers: 32 + num_refiner_layers: 2 + num_attention_heads: 21 + num_kv_heads: 7 + multiple_of: 256 + norm_eps: !!float 1e-05 + axes_dim_rope: [40, 40, 40] + axes_lens: [10000, 10000, 10000] + text_feat_dim: 2048 + timestep_scale: !!float 1000 + +transport: + snr_type: lognorm + do_shift: true + dynamic_time_shift: true + +train: + global_batch_size: 576 + batch_size: 18 + gradient_accumulation_steps: 1 + + max_train_steps: 1000 + + dataloader_num_workers: 12 + + # Optimizer + learning_rate: !!float 4e-4 + scale_lr: false + lr_scheduler: timm_constant_with_warmup + warmup_t: 0 + warmup_lr_init: 1e-7 + warmup_prefix: true + t_in_epochs: false + + # resume_from_checkpoint: + + use_8bit_adam: false + adam_beta1: 0.9 + adam_beta2: 0.95 + adam_weight_decay: !!float 0.01 + adam_epsilon: !!float 1e-08 + max_grad_norm: 1 + + gradient_checkpointing: true + + set_grads_to_none: true + + # Misc + allow_tf32: false + mixed_precision: 'bf16' + + ema_decay: 0.0 + + lora_ft: true + lora_rank: 32 + lora_alpha: 64 + lora_dropout: 0 + + rl: + num_unique_prompts_per_sampling: 48 + num_update_steps_per_sampling: 2 + batch_size_per_forward: 9 + num_images_per_prompt: 12 + sigma_coef: 0.7 + negative_prompt: "" + num_inference_step: 20 + max_sequence_length: 1024 + text_guidance_scale: 4 + image_guidance_scale: 2 + cfg_range_start: 0.0 + cfg_range_end: 0.6 + train_timesteps_fraction: 0.6 + reuse_samples_nums: 1 + clip_range: [!!float 1e-4, !!float 5e-4] + adv_clip_max: !!float 5 + kl_loss_weight: !!float 0.04 + apply_cfg_in_training: true + server_type: vlm + use_ori_neg_prompt_template: true + time_shift_base_res: 168 + policy_loss_reweighting: true + +val: + train_visualization_interval: 5 + num_train_visualization_samples: 3 + +logger: + log_with: [wandb, tensorboard] + # log_with: ~ + + checkpointing_steps: 50 + checkpoints_total_limit: ~ + +cache_dir: +resume_from_checkpoint: latest \ No newline at end of file diff --git a/examples/OmniGen2-RL/options/omnigen2_edit_rl_4machine_editscore7b_avg8.yml b/examples/OmniGen2-RL/options/omnigen2_edit_rl_4machine_editscore7b_avg8.yml new file mode 100644 index 0000000000000000000000000000000000000000..4692b75bd0bb2b30cc04c3e06fefcef80222275f --- /dev/null +++ b/examples/OmniGen2-RL/options/omnigen2_edit_rl_4machine_editscore7b_avg8.yml @@ -0,0 +1,122 @@ +name: omnigen2_edit_rl_4machine_editscore7b_avg8 + +seed: 2233 +device_specific_seed: true +workder_specific_seed: true + +reward_server_config: reward_server/server_configs/editscore_7B_avg8.yml + +data: + data_path: data_configs/train/example/train.yml + use_chat_template: true + maximum_text_tokens: 888 + prompt_dropout_prob: !!float 0.0 + ref_img_dropout_prob: !!float 0.0 + max_output_pixels: 262144 # 512 * 512 + max_input_pixels: [262144, 262144, 262144, 262144] # [512 * 512, 512 * 512, 512 * 512, 512 * 512] + max_side_length: 2048 + +model: + pretrained_vae_model_name_or_path: black-forest-labs/FLUX.1-dev + pretrained_text_encoder_model_name_or_path: Qwen/Qwen2.5-VL-3B-Instruct + pretrained_model_path: pretrained_models/OmniGen2/transformer/pytorch_model.bin + + arch_opt: + patch_size: 2 + in_channels: 16 + hidden_size: 2520 + num_layers: 32 + num_refiner_layers: 2 + num_attention_heads: 21 + num_kv_heads: 7 + multiple_of: 256 + norm_eps: !!float 1e-05 + axes_dim_rope: [40, 40, 40] + axes_lens: [10000, 10000, 10000] + text_feat_dim: 2048 + timestep_scale: !!float 1000 + +transport: + snr_type: lognorm + do_shift: true + dynamic_time_shift: true + +train: + global_batch_size: 576 + batch_size: 18 + gradient_accumulation_steps: 1 + + max_train_steps: 1000 + + dataloader_num_workers: 12 + + # Optimizer + learning_rate: !!float 4e-4 + scale_lr: false + lr_scheduler: timm_constant_with_warmup + warmup_t: 0 + warmup_lr_init: 1e-7 + warmup_prefix: true + t_in_epochs: false + + # resume_from_checkpoint: + + use_8bit_adam: false + adam_beta1: 0.9 + adam_beta2: 0.95 + adam_weight_decay: !!float 0.01 + adam_epsilon: !!float 1e-08 + max_grad_norm: 1 + + gradient_checkpointing: true + + set_grads_to_none: true + + # Misc + allow_tf32: false + mixed_precision: 'bf16' + + ema_decay: 0.0 + + lora_ft: true + lora_rank: 32 + lora_alpha: 64 + lora_dropout: 0 + + rl: + num_unique_prompts_per_sampling: 48 + num_update_steps_per_sampling: 2 + batch_size_per_forward: 9 + num_images_per_prompt: 12 + sigma_coef: 0.7 + negative_prompt: "" + num_inference_step: 20 + max_sequence_length: 1024 + text_guidance_scale: 4 + image_guidance_scale: 2 + cfg_range_start: 0.0 + cfg_range_end: 0.6 + train_timesteps_fraction: 0.6 + reuse_samples_nums: 1 + clip_range: [!!float 1e-4, !!float 5e-4] + adv_clip_max: !!float 5 + kl_loss_weight: !!float 0.04 + apply_cfg_in_training: true + server_type: vlm + use_ori_neg_prompt_template: true + time_shift_base_res: 168 + policy_loss_reweighting: true + +val: + train_visualization_interval: 5 + num_train_visualization_samples: 3 + +logger: + log_with: [wandb, tensorboard] + # log_with: ~ + + checkpointing_steps: 50 + checkpoints_total_limit: ~ + +cache_dir: +resume_from_checkpoint: latest \ No newline at end of file diff --git a/examples/OmniGen2-RL/options/omnigen2_edit_rl_single_machine_editscore7b.yml b/examples/OmniGen2-RL/options/omnigen2_edit_rl_single_machine_editscore7b.yml new file mode 100644 index 0000000000000000000000000000000000000000..3890d0f590a754598d01c68a6b397929937fefa4 --- /dev/null +++ b/examples/OmniGen2-RL/options/omnigen2_edit_rl_single_machine_editscore7b.yml @@ -0,0 +1,122 @@ +name: omnigen2_edit_rl_single_machine_editscore7b + +seed: 2233 +device_specific_seed: true +workder_specific_seed: true + +reward_server_config: reward_server/server_configs/editscore_7B.yml + +data: + data_path: data_configs/train/example/train.yml + use_chat_template: true + maximum_text_tokens: 888 + prompt_dropout_prob: !!float 0.0 + ref_img_dropout_prob: !!float 0.0 + max_output_pixels: 262144 # 512 * 512 + max_input_pixels: [262144, 262144, 262144, 262144] # [512 * 512, 512 * 512, 512 * 512, 512 * 512] + max_side_length: 2048 + +model: + pretrained_vae_model_name_or_path: black-forest-labs/FLUX.1-dev + pretrained_text_encoder_model_name_or_path: Qwen/Qwen2.5-VL-3B-Instruct + pretrained_model_path: pretrained_models/OmniGen2/transformer/pytorch_model.bin + + arch_opt: + patch_size: 2 + in_channels: 16 + hidden_size: 2520 + num_layers: 32 + num_refiner_layers: 2 + num_attention_heads: 21 + num_kv_heads: 7 + multiple_of: 256 + norm_eps: !!float 1e-05 + axes_dim_rope: [40, 40, 40] + axes_lens: [10000, 10000, 10000] + text_feat_dim: 2048 + timestep_scale: !!float 1000 + +transport: + snr_type: lognorm + do_shift: true + dynamic_time_shift: true + +train: + global_batch_size: 144 + batch_size: 18 + gradient_accumulation_steps: 1 + + max_train_steps: 1000 + + dataloader_num_workers: 12 + + # Optimizer + learning_rate: !!float 4e-4 + scale_lr: false + lr_scheduler: timm_constant_with_warmup + warmup_t: 0 + warmup_lr_init: 1e-7 + warmup_prefix: true + t_in_epochs: false + + # resume_from_checkpoint: + + use_8bit_adam: false + adam_beta1: 0.9 + adam_beta2: 0.95 + adam_weight_decay: !!float 0.01 + adam_epsilon: !!float 1e-08 + max_grad_norm: 1 + + gradient_checkpointing: true + + set_grads_to_none: true + + # Misc + allow_tf32: false + mixed_precision: 'bf16' + + ema_decay: 0.0 + + lora_ft: true + lora_rank: 32 + lora_alpha: 64 + lora_dropout: 0 + + rl: + num_unique_prompts_per_sampling: 12 + num_update_steps_per_sampling: 2 + batch_size_per_forward: 9 + num_images_per_prompt: 12 + sigma_coef: 0.7 + negative_prompt: "" + num_inference_step: 20 + max_sequence_length: 1024 + text_guidance_scale: 4 + image_guidance_scale: 2 + cfg_range_start: 0.0 + cfg_range_end: 0.6 + train_timesteps_fraction: 0.6 + reuse_samples_nums: 1 + clip_range: [!!float 1e-4, !!float 5e-4] + adv_clip_max: !!float 5 + kl_loss_weight: !!float 0.04 + apply_cfg_in_training: true + server_type: vlm + use_ori_neg_prompt_template: true + time_shift_base_res: 168 + policy_loss_reweighting: true + +val: + train_visualization_interval: 5 + num_train_visualization_samples: 3 + +logger: + log_with: [wandb, tensorboard] + # log_with: ~ + + checkpointing_steps: 50 + checkpoints_total_limit: ~ + +cache_dir: +resume_from_checkpoint: latest \ No newline at end of file diff --git a/examples/OmniGen2-RL/pretrained_models/.gitkeep b/examples/OmniGen2-RL/pretrained_models/.gitkeep new file mode 100644 index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391 diff --git a/examples/OmniGen2-RL/requirements.txt b/examples/OmniGen2-RL/requirements.txt new file mode 100644 index 0000000000000000000000000000000000000000..a8d16b69963f971de2ccf4d6adf8c51df3777e9b --- /dev/null +++ b/examples/OmniGen2-RL/requirements.txt @@ -0,0 +1,14 @@ +editscore +datasets +tqdm +python-dotenv +wheel +omegaconf +diffusers +pandas +wandb +ninja +triton-windows; sys_platform == "win32" +matplotlib +flash_attn +flask \ No newline at end of file diff --git a/examples/OmniGen2-RL/reward_server/kill_multi_machines.sh b/examples/OmniGen2-RL/reward_server/kill_multi_machines.sh new file mode 100644 index 0000000000000000000000000000000000000000..235141072024f8cdca251f31b4759106eb6f0bae --- /dev/null +++ b/examples/OmniGen2-RL/reward_server/kill_multi_machines.sh @@ -0,0 +1,32 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) + +root_dir=$SHELL_FOLDER + +config_path=${root_dir}/server_configs/editscore_7B.yml + +while [[ $# -gt 0 ]]; do + case "$1" in + --config_path=*) + config_path="${1#*=}" + shift + ;; + *) + echo "Unknown parameter: $1" + shift + ;; + esac +done + +hosts_string=$(python ${root_dir}/scripts/misc/load_host.py --config_path=${config_path}) +hosts=(${hosts_string}) + +SCRIPT="pkill python" + +for i in "${!hosts[@]}"; do + echo "$i ${hosts[$i]}" + ssh -o StrictHostKeyChecking=no "${hosts[$i]}" "$SCRIPT" & +done + +wait +echo "All tasks in remote machines killed" \ No newline at end of file diff --git a/examples/OmniGen2-RL/reward_server/reward_proxy.py b/examples/OmniGen2-RL/reward_server/reward_proxy.py new file mode 100644 index 0000000000000000000000000000000000000000..6c15181fdbe2fc4ad0762b5e1252299b2fda97f5 --- /dev/null +++ b/examples/OmniGen2-RL/reward_server/reward_proxy.py @@ -0,0 +1,330 @@ +#!/usr/bin/env python3 + +from typing import List, Dict, Any, Tuple +import argparse +import requests +import json +import time +import logging +from flask import Flask, request, jsonify +from concurrent.futures import ThreadPoolExecutor +from collections import defaultdict +import math +import yaml + +logging.basicConfig( + level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s" +) +logger = logging.getLogger("RewardProxy") + +app = Flask(__name__) + + +def reorder_results( + merged_item_list: List[Dict[str, Any]], original_batch_size: int, app: Flask +) -> Dict[str, Any]: + """ + Reorder the merged result list according to the original indices. + + Args: + merged_item_list: A flattened list, each element is a dict containing a single image result. + e.g. [{'score': 0.8, 'meta_data': {'original_index': 5}}, ...] + original_batch_size: The original batch size of the request. + + Returns: + A dict containing sorted 'scores', 'rewards', 'meta_data', etc. + """ + if not merged_item_list: + logger.warning("Merged result list is empty, cannot reorder results.") + # Return an empty result in the expected format + return { + "scores": [0.0] * original_batch_size, + "rewards": [0.0] * original_batch_size, + "reasoning": [""] * original_batch_size, + "strict_rewards": [0.0] * original_batch_size, + "meta_data": [ + {"original_index": i, "error": "No result received"} + for i in range(original_batch_size) + ], + "group_rewards": {}, + "group_strict_rewards": {}, + } + + # 1. Create placeholder lists, pre-allocated to the correct size + ordered_scores = [0.0] * original_batch_size + ordered_rewards = [0.0] * original_batch_size + ordered_reasoning = [""] * original_batch_size + ordered_strict_rewards = [0.0] * original_batch_size + ordered_meta_datas = [ + {"original_index": i, "error": "Result missing from server response"} + for i in range(original_batch_size) + ] + + # 2. Iterate over the flattened result list and place each result in the correct position + found_count = 0 + for item in merged_item_list: + if not isinstance(item, dict): + logger.warning(f"Found non-dict result item, skipped: {item}") + continue + + meta = item.get("meta_data", {}) + original_index = meta.get("original_index") + + if original_index is not None and 0 <= original_index < original_batch_size: + ordered_scores[original_index] = item.get("score", 0.0) + ordered_rewards[original_index] = item.get("reward", 0.0) + ordered_reasoning[original_index] = item.get("reasoning", "") + ordered_strict_rewards[original_index] = item.get("strict_reward", 0.0) + ordered_meta_datas[original_index] = meta + found_count += 1 + else: + logger.warning(f"Found invalid or missing 'original_index' in result item: {item}") + + # 4. Logging + if found_count < original_batch_size: + logger.warning( + f"Result reordering incomplete: {found_count}/{original_batch_size} results found." + ) + else: + logger.info(f"Result reordering complete: {found_count}/{original_batch_size} matched successfully.") + + return { + "scores": ordered_scores, + "rewards": ordered_rewards, + "reasoning": ordered_reasoning, + "strict_rewards": ordered_strict_rewards, + "meta_data": ordered_meta_datas, + } + + +class RewardProxy: + def __init__(self, worker_configs: List[Dict[str, Any]]): + self.server_urls = self._build_server_urls(worker_configs) + print(f"{len(self.server_urls)=}, {worker_configs=}", flush=True) + self.executor = ThreadPoolExecutor(max_workers=len(self.server_urls)) + + logger.info("๐Ÿš€ Proxy initialized") + logger.info(f" -> servers {self.server_urls=} ...") + + @staticmethod + def _build_server_urls(worker_configs: List[Dict[str, Any]]) -> Dict[str, List[str]]: + server_urls = [] + for conf in worker_configs: + server_urls.extend([f"http://{conf['host']}:{conf['base_port'] + i}" for i in range(conf['num_servers'])]) + + return server_urls + + def _send_request_to_worker( + self, server_url: str, batch_data: Dict[str, Any] + ) -> Dict[str, Any]: + """Send request to a single worker server and return the result.""" + try: + response = requests.post( + server_url, + json=batch_data, + timeout=600, # 300 seconds timeout + ) + response.raise_for_status() # Raise exception for 4xx or 5xx status codes + return response.json() + except requests.exceptions.RequestException as e: + logger.error(f"Request to server {server_url} failed: {e}") + except ValueError as e: + logger.error(f"Failed to parse response from {server_url}: {e}") + return None # Return None to indicate failure + + def rebatch_with_instruction( + self, + input_images, + output_image: List, + meta_datas: List, + ): + input_images_group = defaultdict(list) + output_image_group = defaultdict(list) + meta_datas_group = defaultdict(list) + + original_index_group = defaultdict(list) + + for i in range(len(input_images)): + key = meta_datas[i]["instruction"] + input_images_group[key].append(input_images[i]) + output_image_group[key].append(output_image[i]) + meta_datas_group[key].append(meta_datas[i]) + + original_index_group[key].append(i) + + return input_images_group, output_image_group, meta_datas_group, original_index_group + + + def process_batch( + self, + input_images, + output_image: List, + meta_datas: List, + **kwargs, + ) -> Dict[str, List[Any]]: + """ + Dispatch batch tasks to the specified type of worker servers and merge results. + """ + + input_images_group, output_image_group, meta_datas_group, original_index_group = self.rebatch_with_instruction(input_images, output_image, meta_datas) + + num_workers = len(self.server_urls) + + original_index = [] + futures = [] + for i, key in enumerate(input_images_group.keys()): + server_url = self.server_urls[i % num_workers] + payload = { + "input_images": input_images_group[key], + "output_image": output_image_group[key], + "meta_data": meta_datas_group[key], + **kwargs, # Pass use_flowgrpo, debug, etc. + } + # print(f"{server_url=}, {start_idx + i * size_per_worker}:{min(start_idx + (i + 1) * size_per_worker, end_idx)}: {len(payload['input_images'])=}, {len(payload['output_image'])=}, {len(payload['meta_data'])=}", flush=True) + futures.append( + self.executor.submit(self._send_request_to_worker, server_url, payload) + ) + + original_index.extend(original_index_group[key]) + + inverse_original_index = {i: idx for idx, i in enumerate(original_index)} + + # Merge all successful results + merged_results = [] + for future in futures: + result = future.result() + if result is None: + continue + if isinstance(result, dict) and result.get("error"): + logger.error(f"Worker returned error: {result['error']}") + continue + if isinstance(result, list): + merged_results.extend(result) + continue + logger.error(f"Unexpected worker response type: {type(result)}") + + # reorder results by original index + merged_results = [merged_results[inverse_original_index[i]] for i in range(len(original_index))] + + return merged_results + + +def prepare_request_data(request_body: bytes) -> Tuple[List, List, str, Dict]: + """Parse request body and add original index to meta data.""" + data = json.loads(request_body) + if not isinstance(data, dict): + raise ValueError("Request body must be a JSON object") + + input_images = data["input_images"] + output_image = data["output_image"] + meta_datas = data["meta_datas"] + + if not isinstance(input_images, list) or not isinstance(output_image, list): + raise ValueError("'input_images' and 'output_image' must be lists") + if not isinstance(meta_datas, list): + raise ValueError("'meta_datas' must be a list") + + normalized_meta_datas = [] + for meta in meta_datas: + if isinstance(meta, str): + normalized_meta_datas.append(json.loads(meta)) + elif isinstance(meta, dict): + normalized_meta_datas.append(meta) + else: + raise ValueError("Each meta_data item must be a dict or JSON string") + meta_datas = normalized_meta_datas + + # Add original index to each meta_data for later sorting + for i, meta in enumerate(meta_datas): + meta["original_index"] = i + + server_type = data.get("server_type", "geneval") + return input_images, output_image, meta_datas, server_type + + +# Flask route +@app.route("/", methods=["POST"]) +def evaluate(): + try: + input_images, output_image, meta_datas, server_type = prepare_request_data( + request.data + ) + original_batch_size = len(output_image) + logger.info( + f"Received evaluation request: {original_batch_size} images, server type: {server_type}" + ) + except Exception as e: + logger.error(f"Failed to parse request: {e}", exc_info=True) + # Return a JSON error, more universal than pickle + return jsonify( + {"error": "Failed to parse request data", "details": str(e)} + ), 400 + + start_time = time.time() + + proxy = app.proxy + # Dispatch processing + merged_results = proxy.process_batch( + input_images, output_image, meta_datas + ) + + # Reorder results by index + ordered_result = reorder_results(merged_results, original_batch_size, app) + + total_time = time.time() - start_time + logger.info( + f"Evaluation complete! Total time: {total_time:.3f}s ({total_time / original_batch_size * 1000:.1f} ms/image)" + ) + + return jsonify(ordered_result) + + +def main(): + parser = argparse.ArgumentParser(description="Universal Reward Proxy Server") + parser.add_argument("--host", type=str, default="0.0.0.0", help="Server host address") + parser.add_argument( + "--config_path", + type=str, + default="server_configs/editscore_7B.yml", + help="Configuration file path", + ) + # parser.add_argument("--port", type=int, default=23456, help="Proxy server port") + + # parser.add_argument("--worker_host", type=str, default="127.0.0.1") + # parser.add_argument("--worker_base_port", type=int, default=18888) + # parser.add_argument("--worker_num_machines", type=int, default=1) + # parser.add_argument("--max_workers_per_machine", type=int, default=128) + # parser.add_argument("--batch_size", type=int, default=64) + args = parser.parse_args() + + # print( + # f"{args.worker_host=}, {args.worker_base_port=}, {args.worker_num_machines=}, {args.max_workers_per_machine=}, {args.batch_size=}", + # flush=True, + # ) + config = yaml.load(open(args.config_path, "r"), Loader=yaml.FullLoader) + proxy_port = config["server"]["proxy_port"] + + worker_configs = [] + + hosts = config["server"]["hosts"] + for i, host in enumerate(hosts): + worker_configs.append( + { + "host": host, + "base_port": config["server"]["worker_base_port"], + "num_servers": 8 // config["reward"]["tensor_parallel_size"], + } + ) + + proxy_instance = RewardProxy(worker_configs) + app.proxy = proxy_instance + + logger.info(f"Starting proxy server at {worker_configs=}") + + app.run( + host=args.host, port=proxy_port, debug=False, threaded=True, use_reloader=False + ) + + +if __name__ == "__main__": + main() diff --git a/examples/OmniGen2-RL/reward_server/reward_server.py b/examples/OmniGen2-RL/reward_server/reward_server.py new file mode 100644 index 0000000000000000000000000000000000000000..5587de2c67edd6d6efb3cca374f608f74a1549b9 --- /dev/null +++ b/examples/OmniGen2-RL/reward_server/reward_server.py @@ -0,0 +1,244 @@ +import dotenv + +dotenv.load_dotenv(override=True) + +from typing import List, Optional +import argparse +import json +import os +import warnings +import threading +from queue import Queue +from typing import Dict, Tuple +import uuid +import time +import base64 +from io import BytesIO + +from flask import Flask, request, jsonify +from PIL import Image + +from editscore import EditScore +import yaml + +warnings.filterwarnings("ignore") + +app = Flask(__name__) + +# --- Global queue and result storage --- +request_queue = Queue() +results = {} # Use a dict to store results, associated by unique ID + +def apply_chat_template(prompt, num_images: int = 2): + """ + This is used since the bug of transformers which do not support vision id https://github.com/QwenLM/Qwen2.5-VL/issues/716#issuecomment-2723316100 + """ + template = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n" + template += "".join([f": <|vision_start|><|image_pad|><|vision_end|>" for i in range(1, num_images + 1)]) + template += f"{prompt}<|im_end|>\n<|im_start|>assistant\n" + return template + +class VLMScorer: + """Encapsulates vLLM model and scoring logic.""" + def __init__(self, config: Dict[str, any]): + print("๐Ÿ”ง Initializing VLMScorer...") + self.scorer = EditScore( + backbone=config["backbone"], + model_name_or_path=config["model_name_or_path"], + score_range=config["score_range"], + temperature=config["temperature"], + tensor_parallel_size=config["tensor_parallel_size"], + max_model_len=config["max_model_len"], + max_num_seqs=config["max_num_seqs"], + max_num_batched_tokens=config["max_num_batched_tokens"], + num_pass=config["num_pass"], + lora_path=config["lora_path"], + seed=config["seed"], + ) + print("โœ… VLMScorer initialization complete.") + + def score(self, input_images: List[List[Image.Image]], output_image: List[Image.Image], metadata: Dict[str, any]) -> float: + """Score a batch of samples.""" + + image_prompts = [] + for input_image, _output_image in zip(input_images, output_image): + image_prompts.append(input_image + [_output_image]) + + results = self.scorer.batch_evaluate(image_prompts, [_metadata['instruction'] for _metadata in metadata]) + + outputs = [] + for result in results: + reward = result['O_score'] / 10 + reasoning = f"SC_score: {result['SC_score']}\n" + reasoning += f"SC_score_reasoning: {result['SC_score_reasoning']}\n" + reasoning += f"PQ_score: {result['PQ_score']}\n" + reasoning += f"PQ_score_reasoning: {result['PQ_score_reasoning']}\n" + reasoning += f"SC_raw_output: {result['SC_raw_output']}\n" + reasoning += f"PQ_raw_output: {result['PQ_raw_output']}\n" + outputs.append((reward, reasoning)) + return outputs + +def vlm_worker(scorer: VLMScorer): + """Background worker thread, continuously fetches and processes tasks from the queue.""" + print("๐Ÿš€ VLM background worker thread started, waiting for tasks...") + while True: + try: + task_id, input_images, output_image, meta_data = request_queue.get() + + # print(f"๐Ÿ”ฉ Start processing task {task_id[:8]}...") + outputs = scorer.score(input_images, output_image, meta_data) + result_payload = [] + for (reward, reasoning), _meta_data in zip(outputs, meta_data): + result_payload.append( + { + "score": 1.0 if reward >= 0.5 else 0.0, + "reward": reward, + "reasoning": reasoning, + "strict_reward": reward, + "meta_data": _meta_data, + "group_reward": {_meta_data.get("tag", "vlm"): reward}, + "group_strict_reward": {_meta_data.get("tag", "vlm"): reward}, + } + ) + results[task_id] = result_payload + + except Exception as e: + print(f"โŒ Worker thread error while processing task {task_id[:8]}: {e}") + import traceback + traceback.print_exc() + error_result = {"error": f"Internal server error: {e}"} + results[task_id] = error_result + finally: + request_queue.task_done() + +# --- Web layer (Flask App) --- + +def decode_base64_image(image_data: str) -> Image.Image: + """Decode base64 image bytes into a PIL image.""" + try: + raw_bytes = base64.b64decode(image_data, validate=True) + image = Image.open(BytesIO(raw_bytes)) + image.load() + return image + except Exception as e: + raise ValueError(f"Invalid base64 image data: {e}") from e + + +def parse_and_validate_request(raw_data: bytes) -> Tuple[List[Image.Image], Image.Image, Dict, str]: + """Parse request data, validate and convert to required format.""" + try: + data = json.loads(raw_data) + if not isinstance(data, dict): + raise ValueError("Request body must be a JSON object") + input_images_datas = data['input_images'] + output_image_datas = data['output_image'] + meta_data = data['meta_data'] + except Exception as e: + print(f"Failed to parse request data: {e}") + return None, None, None, f"Failed to parse request data: {e}" + + if not isinstance(input_images_datas, list) or not isinstance(output_image_datas, list): + return None, None, None, "'input_images' and 'output_image' must be lists" + if not isinstance(meta_data, list): + return None, None, None, "'meta_data' must be a list" + + try: + batch_output_image = [] + for output_image_data in output_image_datas: + if not isinstance(output_image_data, str): + return None, None, None, "Each output image must be a base64 string" + batch_output_image.append(decode_base64_image(output_image_data).convert('RGB')) + + batch_input_images = [] + for input_image_data in input_images_datas: + if not isinstance(input_image_data, list): + return None, None, None, "Each input_images item must be a list" + batch_input_images.append([]) + for _input_image_data in input_image_data: + if not isinstance(_input_image_data, str): + return None, None, None, "Each input image must be a base64 string" + batch_input_images[-1].append(decode_base64_image(_input_image_data).convert('RGB')) + except Exception as e: + return None, None, None, f"Invalid image payload: {e}" + + batch_meta_data = [] + for _meta_data in meta_data: + if isinstance(_meta_data, str): + try: + _meta_data = json.loads(_meta_data) + except json.JSONDecodeError: + _meta_data = {'prompt': _meta_data} + + if not isinstance(_meta_data, dict): + return None, None, None, f"Meta data must be a dict or JSON string" + batch_meta_data.append(_meta_data) + return batch_input_images, batch_output_image, batch_meta_data, None + +@app.route('/', methods=['POST']) +def evaluate_batch_samples(): + """Receive request, put it into the queue, and wait for the result to return.""" + + input_images, output_image, meta_data, error_msg = parse_and_validate_request(request.data) + if error_msg: + print(f"โŒ Request validation failed: {error_msg}") + return jsonify({"error": error_msg}), 400 + + task_id = str(uuid.uuid4()) + request_queue.put((task_id, input_images, output_image, meta_data)) + print(f"๐Ÿ“ฅ Task {task_id[:8]} enqueued, {len(input_images)=}, {len(output_image)=}, {len(meta_data)=}, current queue size: {request_queue.qsize()}", flush=True) + + timeout_seconds = 600 + start_time = time.time() + + while True: + if task_id in results: + result_data = results.pop(task_id) + print(f"๐Ÿ“ค Task {task_id[:8]} result returned. Time elapsed: {time.time() - start_time:.2f}s") + if isinstance(result_data, dict) and "error" in result_data: + return jsonify(result_data), 500 + return jsonify(result_data), 200 + + if time.time() - start_time > timeout_seconds: + print(f"โŒ›๏ธ Task {task_id[:8]} timed out waiting.") + return jsonify({"error": "Request timed out"}), 504 + + time.sleep(0.05) + + +def arg_parser(): + parser = argparse.ArgumentParser(description='VLM Reward Server - High concurrency optimized (Flask native server)') + parser.add_argument('--host', type=str, default='0.0.0.0', help='Server host (0.0.0.0 means listen on all interfaces)') + parser.add_argument('--port', type=int, default=18096, help='Server port') + parser.add_argument('--config_path', type=str, default='examples/OmniGen2-RL/reward_server/server_configs/editscore_7B.yml', help='Configuration file path') + args = parser.parse_args() + return args + +def main(args): + """Main function, loads model, starts background worker thread and web server.""" + + # 1. Load model + print("โšก Preloading VLM model...") + config = yaml.load(open(args.config_path, "r"), Loader=yaml.FullLoader) + scorer = VLMScorer(config["reward"]) + + # 2. Start background worker thread + worker_thread = threading.Thread(target=vlm_worker, args=(scorer,), daemon=True) + worker_thread.start() + + # 3. Start Flask web server + print(f"๐Ÿ”ฅ Starting VLM reward server at http://{args.host}:{args.port}") + print("๐Ÿš€ Mode: High concurrency single-sample requests (queue-based processing)") + + # Use Flask's built-in development server with threading enabled + try: + # threaded=True allows the server to handle multiple HTTP requests simultaneously + # use_reloader=False is necessary when using background threads to prevent the reloader from creating duplicate threads and model instances + app.run(host=args.host, port=args.port, debug=False, threaded=True, use_reloader=False) + except KeyboardInterrupt: + print("\n๐Ÿ‘‹ VLM server stopped.") + except Exception as e: + print(f"โŒ VLM server failed to start: {e}") + +if __name__ == '__main__': + args = arg_parser() + main(args) diff --git a/examples/OmniGen2-RL/reward_server/scripts/misc/load_host.py b/examples/OmniGen2-RL/reward_server/scripts/misc/load_host.py new file mode 100644 index 0000000000000000000000000000000000000000..1fb00e7d0da03b9677689fceeb68940e400b6ec2 --- /dev/null +++ b/examples/OmniGen2-RL/reward_server/scripts/misc/load_host.py @@ -0,0 +1,20 @@ +import argparse +import yaml + + +def parse_args(): + parser = argparse.ArgumentParser() + parser.add_argument("--config_path", type=str, required=True) + return parser.parse_args() + + +def main(args): + config = yaml.load(open(args.config_path, "r"), Loader=yaml.FullLoader) + hosts = config["server"]["hosts"] + output_string = " ".join(hosts) + print(output_string) + + +if __name__ == "__main__": + args = parse_args() + main(args) diff --git a/examples/OmniGen2-RL/reward_server/scripts/utils/reward_server_sanity_check.py b/examples/OmniGen2-RL/reward_server/scripts/utils/reward_server_sanity_check.py new file mode 100644 index 0000000000000000000000000000000000000000..3f27519fda777967fd10f0a21c1a58ae212854f7 --- /dev/null +++ b/examples/OmniGen2-RL/reward_server/scripts/utils/reward_server_sanity_check.py @@ -0,0 +1,65 @@ +import dotenv + +dotenv.load_dotenv(override=True) + +import os +import sys +import json +from PIL import Image +import argparse +import yaml + +sys.path.append(os.path.join(os.path.dirname(__file__), "..", "..", "..")) + +from omnigen2.grpo.reward_client_edit import evaluate_images + + +def parse_args(): + parser = argparse.ArgumentParser() + parser.add_argument("--config_path", type=str, required=True) + return parser.parse_args() + + +def main(args): + config = yaml.load(open(args.config_path, "r"), Loader=yaml.FullLoader) + proxy_host = config["server"]["hosts"][0] + proxy_port = config["server"]["proxy_port"] + + root_dir = os.path.join( + os.path.dirname(__file__), os.path.pardir, os.path.pardir, os.path.pardir + ) + + N = 48 + K = 12 + images = [ + Image.open(os.path.join(root_dir, "../../example_images/output.png")).resize( + (512, 512) + ) + ] * (N * K) + input_images = [ + [ + Image.open(os.path.join(root_dir, "../../example_images/input.png")).resize( + (512, 512) + ) + ] + ] * (N * K) + meta_datas = [ + json.dumps({"instruction": f"Adjust the background to a glass wall.{i}"}) + for i in range(K) + for j in range(N) + ] + + scores, rewards, reasoning, meta_data = evaluate_images( + input_images=input_images, + output_image=images, + meta_datas=meta_datas, + proxy_host=proxy_host, + proxy_port=proxy_port, + server_type="vlm", + ) + print(scores, rewards, reasoning, meta_data) + + +if __name__ == "__main__": + args = parse_args() + main(args) diff --git a/examples/OmniGen2-RL/reward_server/server_configs/editscore_72B.yml b/examples/OmniGen2-RL/reward_server/server_configs/editscore_72B.yml new file mode 100644 index 0000000000000000000000000000000000000000..75ddea21332a2d737c56c625779b631ecc8fe5dd --- /dev/null +++ b/examples/OmniGen2-RL/reward_server/server_configs/editscore_72B.yml @@ -0,0 +1,27 @@ +server: + hosts: + # Your machine IPs, here we use 4 machines for example + - 127.0.0.1 + - 127.0.0.2 + - 127.0.0.3 + - 127.0.0.4 + - 127.0.0.5 + - 127.0.0.6 + - 127.0.0.7 + - 127.0.0.8 + + worker_base_port: 18888 + proxy_port: 23456 + +reward: + backbone: qwen25vl_vllm + model_name_or_path: Qwen/Qwen2.5-VL-72B-Instruct + lora_path: EditScore/EditScore-72B + score_range: 25 + tensor_parallel_size: 4 + max_num_seqs: 64 + max_model_len: 1536 + max_num_batched_tokens: 98304 + num_pass: 1 + seed: 42 + temperature: 0.7 \ No newline at end of file diff --git a/examples/OmniGen2-RL/reward_server/server_configs/editscore_7B.yml b/examples/OmniGen2-RL/reward_server/server_configs/editscore_7B.yml new file mode 100644 index 0000000000000000000000000000000000000000..ad8395c9a66b264360afc48471f6a14988997169 --- /dev/null +++ b/examples/OmniGen2-RL/reward_server/server_configs/editscore_7B.yml @@ -0,0 +1,23 @@ +server: + hosts: + # Your machine IPs, here we use 4 machines for example + - 127.0.0.1 + - 127.0.0.2 + - 127.0.0.3 + - 127.0.0.4 + + worker_base_port: 18888 + proxy_port: 23456 + +reward: + backbone: qwen25vl_vllm + model_name_or_path: Qwen/Qwen2.5-VL-7B-Instruct + lora_path: EditScore/EditScore-7B + score_range: 25 + tensor_parallel_size: 1 + max_num_seqs: 64 + max_model_len: 1536 + max_num_batched_tokens: 98304 + num_pass: 1 + seed: 42 + temperature: 0.7 \ No newline at end of file diff --git a/examples/OmniGen2-RL/reward_server/server_configs/editscore_7B_avg4.yml b/examples/OmniGen2-RL/reward_server/server_configs/editscore_7B_avg4.yml new file mode 100644 index 0000000000000000000000000000000000000000..21e468879dce84c16f25a2b0fff249db7d30d967 --- /dev/null +++ b/examples/OmniGen2-RL/reward_server/server_configs/editscore_7B_avg4.yml @@ -0,0 +1,23 @@ +server: + hosts: + # Your machine IPs, here we use 4 machines for example + - 127.0.0.1 + - 127.0.0.2 + - 127.0.0.3 + - 127.0.0.4 + + worker_base_port: 18888 + proxy_port: 23456 + +reward: + backbone: qwen25vl_vllm + model_name_or_path: Qwen/Qwen2.5-VL-7B-Instruct + lora_path: EditScore/EditScore-7B + score_range: 25 + tensor_parallel_size: 1 + max_num_seqs: 64 + max_model_len: 1536 + max_num_batched_tokens: 98304 + num_pass: 4 + seed: 42 + temperature: 0.7 \ No newline at end of file diff --git a/examples/OmniGen2-RL/reward_server/server_configs/editscore_7B_avg8.yml b/examples/OmniGen2-RL/reward_server/server_configs/editscore_7B_avg8.yml new file mode 100644 index 0000000000000000000000000000000000000000..a88c9b0d1b09d8446cf20559c623afeb23adae28 --- /dev/null +++ b/examples/OmniGen2-RL/reward_server/server_configs/editscore_7B_avg8.yml @@ -0,0 +1,23 @@ +server: + hosts: + # Your machine IPs, here we use 4 machines for example + - 127.0.0.1 + - 127.0.0.2 + - 127.0.0.3 + - 127.0.0.4 + + worker_base_port: 18888 + proxy_port: 23456 + +reward: + backbone: qwen25vl_vllm + model_name_or_path: Qwen/Qwen2.5-VL-7B-Instruct + lora_path: EditScore/EditScore-7B + score_range: 25 + tensor_parallel_size: 1 + max_num_seqs: 64 + max_model_len: 1536 + max_num_batched_tokens: 98304 + num_pass: 8 + seed: 42 + temperature: 0.7 \ No newline at end of file diff --git a/examples/OmniGen2-RL/reward_server/start_multi_machines.sh b/examples/OmniGen2-RL/reward_server/start_multi_machines.sh new file mode 100644 index 0000000000000000000000000000000000000000..940c94561fefd1ebbc74e0f110275043bf272ed8 --- /dev/null +++ b/examples/OmniGen2-RL/reward_server/start_multi_machines.sh @@ -0,0 +1,45 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) + +root_dir=$SHELL_FOLDER +model_name=editscore_7B +config_path=${root_dir}/server_configs/editscore_7B.yml + +while [[ $# -gt 0 ]]; do + case "$1" in + --config_path=*) + config_path="${1#*=}" + shift + ;; + --model_name=*) + model_name="${1#*=}" + shift + ;; + *) + echo "Unknown parameter: $1" + shift + ;; + esac +done + +hosts_string=$(python ${root_dir}/scripts/misc/load_host.py --config_path=${config_path}) +hosts=(${hosts_string}) + +SCRIPT="bash ${root_dir}/step1.sh --config_path=${config_path} --model_name=${model_name}" + +for i in "${!hosts[@]}"; do + echo "$i ${hosts[$i]}" + ssh -o StrictHostKeyChecking=no "${hosts[$i]}" "tmux new-session -d -s reward_server '$SCRIPT --machine_id=$i'" & +done + +SCRIPT="bash ${root_dir}/step2.sh --config_path=${config_path} --model_name=${model_name}" + +for i in "${!hosts[@]}"; do + echo "$i ${hosts[$i]}" + ssh -o StrictHostKeyChecking=no "${hosts[i]}" "tmux new-session -d -s reward_proxy '$SCRIPT --machine_id=$i'" & +done + +echo "step2.sh started successfully" + +wait +echo "All tasks started successfully" \ No newline at end of file diff --git a/examples/OmniGen2-RL/reward_server/start_multi_servers.py b/examples/OmniGen2-RL/reward_server/start_multi_servers.py new file mode 100644 index 0000000000000000000000000000000000000000..e52c942eeb4ada8ce739dd3e0ab13ad2aeab1bd4 --- /dev/null +++ b/examples/OmniGen2-RL/reward_server/start_multi_servers.py @@ -0,0 +1,172 @@ +#!/usr/bin/env python3 +""" +Script to launch multi-GPU servers. +Start specified server instances on all GPUs. +""" + +import subprocess +import time +import os +import signal +import sys +import argparse +from multiprocessing import Process +import yaml + +# Store subprocesses +server_processes = [] + +def start_server(args, worker_idx, num_gpus_per_worker, port, server_script, unknown_args, log_name): + """Start server on the specified GPU(s)""" + cmd = [ + sys.executable, server_script, + "--port", str(port), + "--config_path", args.config_path, + *unknown_args, + ] + + env = os.environ.copy() + env['CUDA_VISIBLE_DEVICES'] = ','.join(str(i) for i in range(worker_idx * num_gpus_per_worker, (worker_idx + 1) * num_gpus_per_worker)) + + # Create log directory + os.makedirs("./logs", exist_ok=True) + log_file = f"./logs/{log_name}_worker_{worker_idx}.log" + + try: + with open(log_file, 'w') as f: + process = subprocess.Popen( + cmd, + env=env, + stdout=f, + stderr=subprocess.STDOUT, + universal_newlines=True + ) + print(f"๐Ÿ“ Worker {worker_idx} log file: {log_file}") + return process + except Exception as e: + print(f"โŒ Failed to start server for Worker {worker_idx}: {e}") + return None + +def signal_handler(signum, frame): + """Handle exit signals""" + print("\n๐Ÿ›‘ Received exit signal, shutting down all servers...") + for i, process in enumerate(server_processes): + if process and process.poll() is None: + print(f"Shutting down server Worker {i}") + process.terminate() + try: + process.wait(timeout=5) + except subprocess.TimeoutExpired: + process.kill() + sys.exit(0) + +def check_server_logs(args, worker_idx): + """Check server log file""" + log_file = f"./logs/{args.log_name}_worker_{worker_idx}.log" + if os.path.exists(log_file): + try: + with open(log_file, 'r') as f: + content = f.read() + if content: + print(f"\n๐Ÿ“„ Worker {worker_idx} error log:") + print(content[-1000:]) # Show last 1000 characters + except Exception as e: + print(f"Failed to read log file {log_file}: {e}") + +def main(): + # Argument parsing + parser = argparse.ArgumentParser(description='Launch multi-GPU servers') + parser.add_argument('server_script', help='Server script filename to launch') + parser.add_argument('--config_path', type=str, required=True, help='Config path') + parser.add_argument('--machine_id', type=int, default=0, help='Machine id') + parser.add_argument('--model_name', type=str, required=True, help='Model name') + args, unknown_args = parser.parse_known_args() + + log_name = f"reward_server_{args.model_name}_machine_{args.machine_id}" + + config = yaml.load(open(args.config_path, "r"), Loader=yaml.FullLoader) + hosts = config["server"]["hosts"] + worker_base_port = config["server"]["worker_base_port"] + num_gpus_per_worker = config["reward"]["tensor_parallel_size"] + + # Add .py suffix if missing + server_script = args.server_script + if not server_script.endswith('.py'): + server_script += '.py' + + # Check if script exists + if not os.path.exists(server_script): + print(f"โŒ Script file {server_script} does not exist") + sys.exit(1) + + # Register signal handlers + signal.signal(signal.SIGINT, signal_handler) + signal.signal(signal.SIGTERM, signal_handler) + + # Detect available GPUs + import torch + if torch.cuda.is_available(): + num_available_gpus = torch.cuda.device_count() + else: + num_available_gpus = 0 + + num_workers = num_available_gpus // num_gpus_per_worker + + if num_workers == 0: + print("โŒ No available GPUs, exiting") + sys.exit(1) + + print(f"๐Ÿ”ฅ Launching {num_workers} {server_script} server(s)") + print(f"๐Ÿ“ Base port: {worker_base_port}") + print(f"๐Ÿ–ฅ๏ธ Host: {hosts[args.machine_id]}") + print("-" * 50) + + # Start all servers + for worker_idx in range(num_workers): + port = worker_base_port + worker_idx + process = start_server(args, worker_idx, num_gpus_per_worker, port, server_script, unknown_args, log_name) + server_processes.append(process) + time.sleep(3) # Add delay to avoid resource contention + + print(f"\nโœ… Attempted to launch {num_workers} server(s)") + print("Server list:") + for worker_idx in range(num_workers): + port = worker_base_port + worker_idx + print(f" Server {worker_idx}: http://{hosts[args.machine_id]}:{port}") + + print("\nโณ Waiting for servers to start...") + time.sleep(10) # Wait for servers to start + + # Check server status + running_servers = 0 + for i, process in enumerate(server_processes): + if process and process.poll() is None: + running_servers += 1 + print(f"โœ… Server {i} is running") + else: + print(f"โŒ Server {i} failed to start") + check_server_logs(args, i) + + if running_servers == 0: + print("โŒ No servers started successfully") + sys.exit(1) + + print(f"\n๐ŸŽ‰ Successfully started {running_servers} server(s)") + print("โณ Servers are running... (Press Ctrl+C to exit)") + + # Monitor server status + try: + while True: + time.sleep(10) # Check interval + # Check for crashed servers + for i, process in enumerate(server_processes): + if process and process.poll() is not None: + print(f"โš ๏ธ Server {i} exited (return code: {process.returncode})") + check_server_logs(args, i) + # Set exited process to None to avoid duplicate reporting + server_processes[i] = None + except KeyboardInterrupt: + signal_handler(signal.SIGINT, None) + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/examples/OmniGen2-RL/reward_server/step1.sh b/examples/OmniGen2-RL/reward_server/step1.sh new file mode 100644 index 0000000000000000000000000000000000000000..b8fa6555da27fb7de74f195eef4e82954bdf1384 --- /dev/null +++ b/examples/OmniGen2-RL/reward_server/step1.sh @@ -0,0 +1,40 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER + +# uncomment this if you are using conda +# source "$(dirname $(which conda))/../etc/profile.d/conda.sh" +# conda activate editscore + +root_dir=$SHELL_FOLDER +machine_id=0 +model_name=editscore_7B +config_path=${root_dir}/server_configs/editscore_7B.yml + +# process named parameters +while [[ $# -gt 0 ]]; do + case "$1" in + --config_path=*) + config_path="${1#*=}" + shift + ;; + --machine_id=*) + machine_id="${1#*=}" + shift + ;; + --model_name=*) + model_name="${1#*=}" + shift + ;; + *) + echo "Unknown parameter: $1" + shift + ;; + esac +done + +export VLLM_LOGGING_LEVEL=DEBUG +export VLLM_LOG_BATCHSIZE_INTERVAL=60 + +VLLM_USE_V1=1 VLLM_FLASH_ATTN_VERSION=3 python start_multi_servers.py reward_server --config_path ${config_path} --model_name ${model_name} \ +--machine_id ${machine_id} \ No newline at end of file diff --git a/examples/OmniGen2-RL/reward_server/step2.sh b/examples/OmniGen2-RL/reward_server/step2.sh new file mode 100644 index 0000000000000000000000000000000000000000..0e9a83fa647b1c325868be94e0b13067cc698e60 --- /dev/null +++ b/examples/OmniGen2-RL/reward_server/step2.sh @@ -0,0 +1,35 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER + +# uncomment this if you are using conda +# source "$(dirname $(which conda))/../etc/profile.d/conda.sh" +# conda activate editscore + +machine_id=0 +config_path=examples/OmniGen2-RL/reward_server/server_configs/editscore_7B.yml +model_name=editscore_7B + +while [[ $# -gt 0 ]]; do + case "$1" in + --config_path=*) + config_path="${1#*=}" + shift + ;; + --machine_id=*) + machine_id="${1#*=}" + shift + ;; + --model_name=*) + model_name="${1#*=}" + shift + ;; + *) + echo "Unknown parameter: $1" + shift + ;; + esac +done + +python reward_proxy.py --config_path ${config_path} \ +>logs/reward_proxy_${model_name}_machine${machine_id}.log 2>&1 \ No newline at end of file diff --git a/examples/OmniGen2-RL/scripts/data/extract_9_tasks.py b/examples/OmniGen2-RL/scripts/data/extract_9_tasks.py new file mode 100644 index 0000000000000000000000000000000000000000..7ee25f8b2ab625e65ea67ec5cbfa8f1d1c9b77d4 --- /dev/null +++ b/examples/OmniGen2-RL/scripts/data/extract_9_tasks.py @@ -0,0 +1,42 @@ +import json +import os +import argparse + +DESIRED_TASKS = [ + "background", + "color_alter", + "material_alter", + "motion_change", + # "ps_human", + "style", + "subject_add", + "subject_remove", + "subject_replace", + "tone_transfer", + # "text_change" +] + + +def main(args): + filtered_json_lines = [] + with open(args.input_path, "r", encoding="utf-8") as f: + for line in f: + json_line = json.loads(line) + + if json_line["task_type"] in DESIRED_TASKS: + filtered_json_lines.append(json_line) + + os.makedirs(os.path.dirname(args.output_path), exist_ok=True) + with open(args.output_path, "w", encoding="utf-8") as f: + for json_line in filtered_json_lines: + f.write(json.dumps(json_line, ensure_ascii=False) + "\n") + +def parse_args(): + parser = argparse.ArgumentParser() + parser.add_argument("--input_path", type=str, required=True) + parser.add_argument("--output_path", type=str, required=True) + return parser.parse_args() + +if __name__ == "__main__": + args = parse_args() + main(args) \ No newline at end of file diff --git a/examples/OmniGen2-RL/scripts/data/process_jsonl.py b/examples/OmniGen2-RL/scripts/data/process_jsonl.py new file mode 100644 index 0000000000000000000000000000000000000000..bbfd13803121588ac6df8f916118565086174a00 --- /dev/null +++ b/examples/OmniGen2-RL/scripts/data/process_jsonl.py @@ -0,0 +1,91 @@ +import json +import os +import argparse + +def convert_jsonl_paths(input_file, output_file, base_path, max_records=100000): + """ + Convert JSONL file image paths from relative to absolute paths + + Args: + input_file: Path to input JSONL file + output_file: Path to output JSONL file + base_path: Base directory path for converting relative paths + max_records: Maximum number of records to process, default 10000 + """ + + with open(input_file, 'r', encoding='utf-8') as infile, \ + open(output_file, 'w', encoding='utf-8') as outfile: + + count = 0 + for line in infile: + if count >= max_records: + break + + # Parse JSON line + try: + data = json.loads(line.strip()) + + # Check if input_images field exists + if "input_images" in data and isinstance(data["input_images"], list): + # Convert paths + new_paths = [] + for path in data["input_images"]: + if isinstance(path, str) and path.startswith("images/"): + # Convert relative path to absolute path + new_path = path.replace("images/", f"{base_path}/images/") + new_paths.append(new_path) + else: + # Keep original path unchanged + new_paths.append(path) + + data["input_images"] = new_paths + + # Write converted data + outfile.write(json.dumps(data, ensure_ascii=False) + '\n') + count += 1 + + # Print progress every 1000 records + if count % 1000 == 0: + print(f"Processed {count} records") + + except json.JSONDecodeError as e: + print(f"Skipping invalid JSON line: {e}") + continue + + print(f"Conversion completed! Total processed records: {count}") + print(f"Output file: {output_file}") + +def main(): + parser = argparse.ArgumentParser(description="Convert JSONL file image paths from relative to absolute") + parser.add_argument("--input", "-i", required=True, help="Input JSONL file path") + parser.add_argument("--output", "-o", required=True, help="Output JSONL file path") + parser.add_argument("--base-path", "-b", required=True, help="Base directory path for converting relative paths") + parser.add_argument("--max-records", "-m", type=int, default=100000, help="Maximum number of records to process (default: 100000)") + + args = parser.parse_args() + + # Validate input file exists + if not os.path.exists(args.input): + print(f"Error: Input file '{args.input}' does not exist") + return + + # Create output directory if it doesn't exist + output_dir = os.path.dirname(args.output) + if output_dir and not os.path.exists(output_dir): + os.makedirs(output_dir) + print(f"Created output directory: {output_dir}") + + # Validate base path exists + if not os.path.exists(args.base_path): + print(f"Warning: Base path '{args.base_path}' does not exist") + + print(f"Input file: {args.input}") + print(f"Output file: {args.output}") + print(f"Base path: {args.base_path}") + print(f"Max records: {args.max_records}") + print("-" * 50) + + convert_jsonl_paths(args.input, args.output, args.base_path, args.max_records) + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/examples/OmniGen2-RL/scripts/misc/convert_ckpt_to_hf_format.py b/examples/OmniGen2-RL/scripts/misc/convert_ckpt_to_hf_format.py new file mode 100644 index 0000000000000000000000000000000000000000..d29fa3356d1221e92d7c0f9231cb45e5d17111fd --- /dev/null +++ b/examples/OmniGen2-RL/scripts/misc/convert_ckpt_to_hf_format.py @@ -0,0 +1,80 @@ +import dotenv + +dotenv.load_dotenv(override=True) + +import sys +import os +import argparse + +from omegaconf import OmegaConf + +import torch +from accelerate import init_empty_weights + +from peft import LoraConfig +from peft.utils import get_peft_model_state_dict + +sys.path.append(os.path.join(os.path.dirname(__file__), '..', '..')) + +from omnigen2.models.transformers.transformer_omnigen2 import OmniGen2Transformer2DModel +from omnigen2.pipelines.omnigen2.pipeline_omnigen2 import OmniGen2Pipeline + + +def main(args): + config_path = args.config_path + model_path = args.model_path + + conf = OmegaConf.load(config_path) + arch_opt = conf.model.arch_opt + + arch_opt = OmegaConf.to_object(arch_opt) + # Convert lists to tuples in conf.model.arch_opt + for key, value in arch_opt.items(): + if isinstance(value, list): + arch_opt[key] = tuple(value) + + with init_empty_weights(): + transformer = OmniGen2Transformer2DModel(**arch_opt) + + if conf.train.get('lora_ft', False): + target_modules = ["to_k", "to_q", "to_v", "to_out.0"] + + # now we will add new LoRA weights the transformer layers + lora_config = LoraConfig( + r=conf.train.lora_rank, + lora_alpha=conf.train.lora_rank, + lora_dropout=conf.train.lora_dropout, + init_lora_weights="gaussian", + target_modules=target_modules, + ) + transformer.add_adapter(lora_config) + + state_dict = torch.load(model_path, mmap=True, weights_only=True) + missing, unexpect = transformer.load_state_dict( + state_dict, assign=True, strict=False + ) + print(f"missed parameters: {missing}") + print(f"unexpected parameters: {unexpect}") + + save_path = args.save_path + if conf.train.get('lora_ft', False): + transformer_lora_layers = get_peft_model_state_dict(transformer) + OmniGen2Pipeline.save_lora_weights( + save_directory=save_path, + transformer_lora_layers=transformer_lora_layers, + ) + else: + transformer.save_pretrained(save_path) + + +def parse_args(): + parser = argparse.ArgumentParser() + parser.add_argument("--config_path", type=str, required=True) + parser.add_argument("--model_path", type=str, required=True) + parser.add_argument("--save_path", type=str, required=True) + return parser.parse_args() + + +if __name__ == "__main__": + args = parse_args() + main(args) diff --git a/examples/OmniGen2-RL/scripts/misc/convert_dist_ckpt_to_ckpt.py b/examples/OmniGen2-RL/scripts/misc/convert_dist_ckpt_to_ckpt.py new file mode 100644 index 0000000000000000000000000000000000000000..04d51af04c072f63d2cf87c623df8a9d2d08e3bd --- /dev/null +++ b/examples/OmniGen2-RL/scripts/misc/convert_dist_ckpt_to_ckpt.py @@ -0,0 +1,38 @@ +import dotenv + +dotenv.load_dotenv(override=True) + +import sys +import os +import argparse + +import torch +from torch.distributed.checkpoint.format_utils import dcp_to_torch_save + +sys.path.append(os.path.join(os.path.dirname(__file__), '..', '..')) + +from omnigen2.models.transformers.transformer_omnigen2 import OmniGen2Transformer2DModel +from omnigen2.pipelines.omnigen2.pipeline_omnigen2 import OmniGen2Pipeline + + +def main(args): + model_path = args.model_path + save_path = args.save_path + + dcp_to_torch_save(model_path, save_path) + + state_dict = torch.load(save_path, weights_only=True)['model'] + + torch.save(state_dict, save_path) + + +def parse_args(): + parser = argparse.ArgumentParser() + parser.add_argument("--model_path", type=str, required=True) + parser.add_argument("--save_path", type=str, required=True) + return parser.parse_args() + + +if __name__ == "__main__": + args = parse_args() + main(args) diff --git a/examples/OmniGen2-RL/scripts/misc/convert_dist_ckpt_to_hf_format.sh b/examples/OmniGen2-RL/scripts/misc/convert_dist_ckpt_to_hf_format.sh new file mode 100644 index 0000000000000000000000000000000000000000..cef5a77d10132387dc12dea3b8e31cb138262b3c --- /dev/null +++ b/examples/OmniGen2-RL/scripts/misc/convert_dist_ckpt_to_hf_format.sh @@ -0,0 +1,19 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER +cd ../../ + +experiment_name=$1 +step=$2 + +model_path=experiments/${experiment_name}/checkpoint-${step} +config_path=experiments/${experiment_name}/${experiment_name}.yml + +python scripts/misc/convert_dist_ckpt_to_ckpt.py \ +--model_path $model_path/pytorch_model_fsdp_0 \ +--save_path $model_path/pytorch_model_fsdp.bin + +python scripts/misc/convert_ckpt_to_hf_format.py \ +--config_path $config_path \ +--model_path $model_path/pytorch_model_fsdp.bin \ +--save_path $model_path/transformer_lora \ No newline at end of file diff --git a/examples/OmniGen2-RL/scripts/misc/extract_bin_from_pipe.py b/examples/OmniGen2-RL/scripts/misc/extract_bin_from_pipe.py new file mode 100644 index 0000000000000000000000000000000000000000..061cd020b6ee9cd980624e85f66e91846b8f6929 --- /dev/null +++ b/examples/OmniGen2-RL/scripts/misc/extract_bin_from_pipe.py @@ -0,0 +1,26 @@ +import dotenv + +dotenv.load_dotenv(override=True) + +import os +import sys + +import torch + +sys.path.append(os.path.join(os.path.dirname(__file__), '..', '..')) + +from omnigen2.models.transformers.transformer_omnigen2 import OmniGen2Transformer2DModel + + +def main(): + transformer = OmniGen2Transformer2DModel.from_pretrained("OmniGen2/OmniGen2", subfolder="transformer") + + state_dict = transformer.state_dict() + + save_path = os.path.join(os.path.dirname(__file__), os.path.pardir, os.path.pardir, "pretrained_models", "OmniGen2", "transformer/pytorch_model.bin") + os.makedirs(os.path.dirname(save_path), exist_ok=True) + + torch.save(state_dict, save_path) + +if __name__ == "__main__": + main() diff --git a/examples/OmniGen2-RL/scripts/misc/merge_ckpt.py b/examples/OmniGen2-RL/scripts/misc/merge_ckpt.py new file mode 100644 index 0000000000000000000000000000000000000000..daf6715bd4722a597c58a4dbdf21dcc91beb5d4a --- /dev/null +++ b/examples/OmniGen2-RL/scripts/misc/merge_ckpt.py @@ -0,0 +1,74 @@ +import dotenv + +dotenv.load_dotenv(override=True) + +import sys +import os +import argparse + +from omegaconf import OmegaConf + +import torch +from accelerate import init_empty_weights + +from peft import LoraConfig + +sys.path.append(os.path.join(os.path.dirname(__file__), '..', '..')) + +from omnigen2.models.transformers.transformer_omnigen2 import OmniGen2Transformer2DModel + + +def main(args): + config_path = args.config_path + model_path = args.model_path + + conf = OmegaConf.load(config_path) + arch_opt = conf.model.arch_opt + + arch_opt = OmegaConf.to_object(arch_opt) + # Convert lists to tuples in conf.model.arch_opt + for key, value in arch_opt.items(): + if isinstance(value, list): + arch_opt[key] = tuple(value) + + with init_empty_weights(): + transformer = OmniGen2Transformer2DModel(**arch_opt) + + if conf.train.get('lora_ft', False): + target_modules = ["to_k", "to_q", "to_v", "to_out.0"] + + # now we will add new LoRA weights the transformer layers + lora_config = LoraConfig( + r=conf.train.lora_rank, + lora_alpha=conf.train.lora_rank, + lora_dropout=conf.train.lora_dropout, + init_lora_weights="gaussian", + target_modules=target_modules, + ) + transformer.add_adapter(lora_config) + + state_dict = torch.load(model_path, mmap=True, weights_only=True) + missing, unexpect = transformer.load_state_dict( + state_dict, assign=True, strict=False + ) + print(f"missed parameters: {missing}") + print(f"unexpected parameters: {unexpect}") + + save_path = args.save_path + + transformer.fuse_lora() + transformer.unload_lora() + state_dict = transformer.state_dict() + torch.save(state_dict, save_path) + +def parse_args(): + parser = argparse.ArgumentParser() + parser.add_argument("--config_path", type=str, required=True) + parser.add_argument("--model_path", type=str, required=True) + parser.add_argument("--save_path", type=str, required=True) + return parser.parse_args() + + +if __name__ == "__main__": + args = parse_args() + main(args) diff --git a/examples/OmniGen2-RL/scripts/misc/merge_ckpt.sh b/examples/OmniGen2-RL/scripts/misc/merge_ckpt.sh new file mode 100644 index 0000000000000000000000000000000000000000..217d5a6e6039a55da41a70f7b37a08b2fce23113 --- /dev/null +++ b/examples/OmniGen2-RL/scripts/misc/merge_ckpt.sh @@ -0,0 +1,15 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $SHELL_FOLDER +cd ../../ + +experiment_name=$1 +step=$2 + +model_path=experiments/${experiment_name}/checkpoint-${step} +config_path=experiments/${experiment_name}/${experiment_name}.yml + +python scripts/misc/merge_ckpt.py \ +--config_path $config_path \ +--model_path $model_path/pytorch_model_fsdp.bin \ +--save_path $model_path/pytorch_model_fsdp_merged.bin \ No newline at end of file diff --git a/examples/OmniGen2-RL/scripts/train/omnigen2_edit_rl_4machine_editscore7b_avg4.sh b/examples/OmniGen2-RL/scripts/train/omnigen2_edit_rl_4machine_editscore7b_avg4.sh new file mode 100644 index 0000000000000000000000000000000000000000..b5535f31284230e02ba6ac96b0a0c878f068f05e --- /dev/null +++ b/examples/OmniGen2-RL/scripts/train/omnigen2_edit_rl_4machine_editscore7b_avg4.sh @@ -0,0 +1,68 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +debug=false +RANK=0 +MASTER_ADDR=1 +MASTER_PORT=29500 +WORLD_SIZE=1 + +# process named arguments +while [[ $# -gt 0 ]]; do + case "$1" in + --rank=*) + RANK="${1#*=}" + shift + ;; + --master_addr=*) + MASTER_ADDR="${1#*=}" + shift + ;; + --master_port=*) + MASTER_PORT="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + *) + echo "unknown argument: $1" + shift + ;; + esac +done + +echo "RANK: $RANK" +echo "MASTER_ADDR: $MASTER_ADDR" +echo "MASTER_PORT: $MASTER_PORT" +echo "WORLD_SIZE: $WORLD_SIZE" + +num_gpu_cards=$(nvidia-smi -L | wc -l) +num_processes=$(($WORLD_SIZE * $num_gpu_cards)) +# num_processes=2 + +echo "num_processes: $num_processes" +echo $NCCL_DEBUG_FILE + +experiment_name="omnigen2_edit_rl_4machine_editscore7b_avg4" + +accelerate launch \ +--machine_rank=$RANK \ +--main_process_ip=$MASTER_ADDR \ +--main_process_port=$MASTER_PORT \ +--num_machines=$WORLD_SIZE \ +--num_processes=$num_processes \ +--use_fsdp \ +--fsdp_offload_params false \ +--fsdp_sharding_strategy HYBRID_SHARD \ +--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP \ +--fsdp_transformer_layer_cls_to_wrap OmniGen2TransformerBlock \ +--fsdp_state_dict_type SHARDED_STATE_DICT \ +--fsdp_forward_prefetch false \ +--fsdp_use_orig_params True \ +--fsdp_cpu_ram_efficient_loading false \ +--fsdp_sync_module_states True \ +train.py --config options/${experiment_name}.yml \ No newline at end of file diff --git a/examples/OmniGen2-RL/scripts/train/omnigen2_edit_rl_4machine_editscore7b_avg8.sh b/examples/OmniGen2-RL/scripts/train/omnigen2_edit_rl_4machine_editscore7b_avg8.sh new file mode 100644 index 0000000000000000000000000000000000000000..052a7bcb1808f6d4f2ae4b00b5102c76d264c6d7 --- /dev/null +++ b/examples/OmniGen2-RL/scripts/train/omnigen2_edit_rl_4machine_editscore7b_avg8.sh @@ -0,0 +1,68 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +debug=false +RANK=0 +MASTER_ADDR=1 +MASTER_PORT=29500 +WORLD_SIZE=1 + +# process named arguments +while [[ $# -gt 0 ]]; do + case "$1" in + --rank=*) + RANK="${1#*=}" + shift + ;; + --master_addr=*) + MASTER_ADDR="${1#*=}" + shift + ;; + --master_port=*) + MASTER_PORT="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + *) + echo "unknown argument: $1" + shift + ;; + esac +done + +echo "RANK: $RANK" +echo "MASTER_ADDR: $MASTER_ADDR" +echo "MASTER_PORT: $MASTER_PORT" +echo "WORLD_SIZE: $WORLD_SIZE" + +num_gpu_cards=$(nvidia-smi -L | wc -l) +num_processes=$(($WORLD_SIZE * $num_gpu_cards)) +# num_processes=2 + +echo "num_gpu_cards: $num_gpu_cards, num_processes: $num_processes" +echo $NCCL_DEBUG_FILE + +experiment_name="omnigen2_edit_rl_4machine_editscore7b_avg8" + +accelerate launch \ +--machine_rank=$RANK \ +--main_process_ip=$MASTER_ADDR \ +--main_process_port=$MASTER_PORT \ +--num_machines=$WORLD_SIZE \ +--num_processes=$num_processes \ +--use_fsdp \ +--fsdp_offload_params false \ +--fsdp_sharding_strategy HYBRID_SHARD \ +--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP \ +--fsdp_transformer_layer_cls_to_wrap OmniGen2TransformerBlock \ +--fsdp_state_dict_type SHARDED_STATE_DICT \ +--fsdp_forward_prefetch false \ +--fsdp_use_orig_params True \ +--fsdp_cpu_ram_efficient_loading false \ +--fsdp_sync_module_states True \ +train.py --config options/${experiment_name}.yml \ No newline at end of file diff --git a/examples/OmniGen2-RL/scripts/train/omnigen2_edit_rl_single_machine_editscore7b.sh b/examples/OmniGen2-RL/scripts/train/omnigen2_edit_rl_single_machine_editscore7b.sh new file mode 100644 index 0000000000000000000000000000000000000000..4a4912ce7fca9c28d37294625cede21eccd80baa --- /dev/null +++ b/examples/OmniGen2-RL/scripts/train/omnigen2_edit_rl_single_machine_editscore7b.sh @@ -0,0 +1,68 @@ +# !/bin/bash +SHELL_FOLDER=$(cd "$(dirname "$0")";pwd) +cd $(dirname $SHELL_FOLDER) +cd ../ + +debug=false +RANK=0 +MASTER_ADDR=1 +MASTER_PORT=29500 +WORLD_SIZE=1 + +# process named arguments +while [[ $# -gt 0 ]]; do + case "$1" in + --rank=*) + RANK="${1#*=}" + shift + ;; + --master_addr=*) + MASTER_ADDR="${1#*=}" + shift + ;; + --master_port=*) + MASTER_PORT="${1#*=}" + shift + ;; + --world_size=*) + WORLD_SIZE="${1#*=}" + shift + ;; + *) + echo "unknown argument: $1" + shift + ;; + esac +done + +echo "RANK: $RANK" +echo "MASTER_ADDR: $MASTER_ADDR" +echo "MASTER_PORT: $MASTER_PORT" +echo "WORLD_SIZE: $WORLD_SIZE" + +num_gpu_cards=$(nvidia-smi -L | wc -l) +num_processes=$(($WORLD_SIZE * $num_gpu_cards)) +# num_processes=2 + +echo "num_processes: $num_processes" +echo $NCCL_DEBUG_FILE + +experiment_name="omnigen2_edit_rl_single_machine_editscore7b" + +accelerate launch \ +--machine_rank=$RANK \ +--main_process_ip=$MASTER_ADDR \ +--main_process_port=$MASTER_PORT \ +--num_machines=$WORLD_SIZE \ +--num_processes=$num_processes \ +--use_fsdp \ +--fsdp_offload_params false \ +--fsdp_sharding_strategy HYBRID_SHARD \ +--fsdp_auto_wrap_policy TRANSFORMER_BASED_WRAP \ +--fsdp_transformer_layer_cls_to_wrap OmniGen2TransformerBlock \ +--fsdp_state_dict_type SHARDED_STATE_DICT \ +--fsdp_forward_prefetch false \ +--fsdp_use_orig_params True \ +--fsdp_cpu_ram_efficient_loading false \ +--fsdp_sync_module_states True \ +train.py --config options/${experiment_name}.yml \ No newline at end of file diff --git a/examples/OmniGen2-RL/train.py b/examples/OmniGen2-RL/train.py new file mode 100644 index 0000000000000000000000000000000000000000..f009c08e10be9e5a28cd60372b903bf52b6c3c2b --- /dev/null +++ b/examples/OmniGen2-RL/train.py @@ -0,0 +1,1071 @@ +import dotenv + +dotenv.load_dotenv(override=True) + +from typing import Union, List, Optional, Tuple +import time +from contextlib import contextmanager +from copy import deepcopy +import argparse +from collections import defaultdict +import logging +import math +import os +import json +import random +import shutil +from functools import partial +from pathlib import Path +from omegaconf import OmegaConf +from tqdm.auto import tqdm + +import numpy as np + +import matplotlib.pyplot as plt + +import torch + +import torch.nn.functional as F +import torch.utils.checkpoint + +from torchvision.transforms.functional import crop, to_pil_image, to_tensor + +from einops import repeat, rearrange + +import accelerate +from accelerate import Accelerator +from accelerate.state import AcceleratorState +from accelerate.logging import get_logger +from accelerate.utils import ProjectConfiguration, set_seed, DataLoaderConfiguration +from accelerate import init_empty_weights +from accelerate.utils import gather_object + +import transformers +from transformers import AutoTokenizer, AutoProcessor +from transformers import Qwen2_5_VLModel as TextEncoder + +import diffusers +from diffusers.optimization import get_scheduler +from diffusers.utils.torch_utils import is_compiled_module +from diffusers.models.autoencoders.autoencoder_kl import AutoencoderKL + +from peft import LoraConfig + +from omnigen2.training_utils import EMAModel +from omnigen2.utils.logging_utils import TqdmToLogger +from omnigen2.utils.tensor_util import pad_to_length, expand_as +from omnigen2.dataset.omnigen2_train_dataset import OmniGen2TrainDataset, OmniGen2Collator, RepeatedDistributedBatchSampler +from omnigen2.models.transformers.transformer_omnigen2 import OmniGen2Transformer2DModel +from omnigen2.models.transformers.repo import OmniGen2RotaryPosEmbed +from omnigen2.grpo.reward_client_edit import evaluate_images +from omnigen2.grpo.utils import forward_logprob, process_grpo_rewards, compute_single_step_ppo_loss +from omnigen2.pipelines.omnigen2.pipeline_omnigen2 import FMPipelineOutput + + +logger = get_logger(__name__) + + +def parse_args(root_path) -> OmegaConf: + parser = argparse.ArgumentParser(description="OmniGen2 training script") + parser.add_argument( + "--config", + type=str, + required=True, + help="Path to configuration file (YAML format)", + ) + parser.add_argument( + "--global_batch_size", + type=int, + default=None, + help="Global batch size.", + ) + parser.add_argument( + "--data_path", + type=str, + default=None, + help="Data path.", + ) + args = parser.parse_args() + conf = OmegaConf.load(args.config) + + output_dir = os.path.join(root_path, 'experiments', conf.name) + conf.root_dir = root_path + conf.output_dir = output_dir + conf.config_file = args.config + + # Override config with command line arguments + if args.global_batch_size is not None: + conf.train.global_batch_size = args.global_batch_size + + if args.data_path is not None: + conf.data.data_path = args.data_path + return conf + +def setup_logging(args: OmegaConf, accelerator: Accelerator) -> None: + """ + Set up logging configuration for training. + + Args: + accelerator: Accelerator instance + args: Configuration object + logging_dir: Directory for log files + """ + + logging_dir = Path(args.output_dir, "logs") + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + shutil.copy(args.config_file, args.output_dir) + + # Create logging directory and file handler + os.makedirs(logging_dir, exist_ok=True) + log_file = Path(logging_dir, f'{time.strftime("%Y%m%d-%H%M%S")}.log') + + formatter = logging.Formatter('%(asctime)s - %(levelname)s - %(name)s - %(message)s') + file_handler = logging.FileHandler(log_file, 'w') + file_handler.setFormatter(formatter) + file_handler.setLevel(logging.INFO) + logger.logger.addHandler(file_handler) + + # Configure basic logging + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + + # Set verbosity for different processes + log_level = logging.INFO if accelerator.is_local_main_process else logging.ERROR + transformers.utils.logging.set_verbosity(log_level) + diffusers.utils.logging.set_verbosity(log_level) + + +def log_model_info(name: str, model: torch.nn.Module): + """Logs parameter counts for a given model.""" + total_params = sum(p.numel() for p in model.parameters()) + trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad) + logger.info(f"--- {name} ---") + logger.info(model) + logger.info(f"Total parameters (M): {total_params / 1e6:.2f}") + logger.info(f"Trainable parameters (M): {trainable_params / 1e6:.2f}") + + +def get_qwen2_prompt_embeds( + text_encoder, + tokenizer, + prompt: Union[str, List[str]], + device: Optional[torch.device] = None, + max_sequence_length: int = 256, +) -> Tuple[torch.Tensor, torch.Tensor]: + """ + Get prompt embeddings from the Qwen2 text encoder. + + Args: + prompt: The prompt or list of prompts to encode. + device: The device to place the embeddings on. If None, uses the pipeline's device. + max_sequence_length: Maximum sequence length for tokenization. + + Returns: + Tuple[torch.Tensor, torch.Tensor]: A tuple containing: + - The prompt embeddings tensor + - The attention mask tensor + + Raises: + Warning: If the input text is truncated due to sequence length limitations. + """ + prompt = [prompt] if isinstance(prompt, str) else prompt + + text_inputs = tokenizer( + prompt, + padding="longest", + max_length=max_sequence_length, + truncation=True, + return_tensors="pt", + ) + + text_input_ids = text_inputs.input_ids.to(device) + untruncated_ids = tokenizer(prompt, padding="longest", return_tensors="pt").input_ids.to(device) + + if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids): + removed_text = tokenizer.batch_decode(untruncated_ids[:, max_sequence_length - 1 : -1]) + logger.warning( + "The following part of your input was truncated because Gemma can only handle sequences up to" + f" {max_sequence_length} tokens: {removed_text}" + ) + + prompt_attention_mask = text_inputs.attention_mask.to(device) + + prompt_embeds = text_encoder( + input_ids=text_input_ids, + attention_mask=prompt_attention_mask, + output_hidden_states=False, + ).last_hidden_state + + return prompt_embeds, prompt_attention_mask + +@contextmanager +def disabled_adapters(model): + try: + # model.disable_adapters() + from peft.tuners.tuners_utils import BaseTunerLayer + + for _, module in model.named_modules(): + if isinstance(module, BaseTunerLayer): + if hasattr(module, "enable_adapters"): + module._disable_adapters = True + else: + module.disable_adapters = True + yield + finally: + # model.enable_adapters() + from peft.tuners.tuners_utils import BaseTunerLayer + + for _, module in model.named_modules(): + if isinstance(module, BaseTunerLayer): + if hasattr(module, "enable_adapters"): + module._disable_adapters = False + else: + module.disable_adapters = False + + +def main(args): + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=Path(args.output_dir, 'logs')) + + accelerator = Accelerator( + gradient_accumulation_steps=args.train.gradient_accumulation_steps + * math.ceil( + args.train.rl.num_inference_step + * args.train.rl.get("train_timesteps_fraction", 1.0) + ), + mixed_precision=args.train.mixed_precision, + log_with=OmegaConf.to_object(args.logger.log_with), + project_config=accelerator_project_config, + dataloader_config=DataLoaderConfiguration(split_batches=True), + ) + + setup_logging(args, accelerator) + + # Reproducibility + if args.seed is not None: + set_seed(args.seed, device_specific=args.get('device_specific_seed', False)) + + # Set performance flags + if args.train.allow_tf32: + torch.backends.cuda.matmul.allow_tf32 = True + if args.train.get('benchmark_cudnn', False): + torch.backends.cudnn.benchmark = True + + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + ema_decay = args.train.get('ema_decay', 0) + + if args.model.pretrained_model_path: + with init_empty_weights(): + model = OmniGen2Transformer2DModel(**args.model.arch_opt) + + state_dict = torch.load(args.model.pretrained_model_path, mmap=True, weights_only=True) + missing, unexpect = model.load_state_dict(state_dict, assign=True, strict=False) + else: + model = OmniGen2Transformer2DModel(**args.model.arch_opt) + model.train() + + freqs_cis = OmniGen2RotaryPosEmbed.get_freqs_cis( + model.config.axes_dim_rope, + model.config.axes_lens, + theta=10000, + ) + + if ema_decay != 0: + model_ema = deepcopy(model) + model_ema._requires_grad = False + + processor = AutoProcessor.from_pretrained(args.model.pretrained_text_encoder_model_name_or_path) + + text_tokenizer = processor.tokenizer + text_tokenizer.padding_side = "right" + + if accelerator.is_main_process: + text_tokenizer.save_pretrained(os.path.join(args.output_dir, 'tokenizer')) + + text_encoder = TextEncoder.from_pretrained( + args.model.pretrained_text_encoder_model_name_or_path, + torch_dtype=weight_dtype, + ) + if args.model.get('resize_token_embeddings', False): + text_encoder.resize_token_embeddings(len(text_tokenizer)) + + if accelerator.is_main_process: + text_encoder.save_pretrained(os.path.join(args.output_dir, 'text_encoder')) + + log_model_info("text_encoder", text_encoder) + + vae = AutoencoderKL.from_pretrained( + args.model.pretrained_vae_model_name_or_path, + subfolder=args.model.get("vae_subfolder", "vae"), + local_files_only=True, + ) + + logger.info(vae) + logger.info("***** Move vae, text_encoder to device and cast to weight_dtype *****") + # Move vae, unet, text_encoder and controlnet_ema to device and cast to weight_dtype + vae = vae.to(accelerator.device, dtype=weight_dtype) + text_encoder = text_encoder.to(accelerator.device, dtype=weight_dtype) + + args.train.lora_ft = args.train.get('lora_ft', False) + if args.train.lora_ft: + model.requires_grad_(False) + + target_modules = ["to_k", "to_q", "to_v", "to_out.0"] + + lora_config = LoraConfig( + r=args.train.lora_rank, + lora_alpha=args.train.lora_alpha, + lora_dropout=args.train.lora_dropout, + init_lora_weights="gaussian", + target_modules=target_modules, + ) + model.add_adapter(lora_config) + + if args.train.gradient_checkpointing: + model.enable_gradient_checkpointing() + + if args.train.scale_lr: + args.train.learning_rate = ( + args.train.learning_rate * args.train.gradient_accumulation_steps * args.train.batch_size * accelerator.num_processes + ) + + # Use 8-bit Adam for lower memory usage or to fine-tune the model in 16GB GPUs + if args.train.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError( + "To use 8-bit Adam, please install the bitsandbytes library: `pip install bitsandbytes`." + ) + + optimizer_class = bnb.optim.AdamW8bit + else: + optimizer_class = torch.optim.AdamW + + log_model_info("transformer", model) + + # Optimizer creation + trainable_params = list(filter(lambda p: p.requires_grad, model.parameters())) + + optimizer = optimizer_class( + trainable_params, + lr=args.train.learning_rate, + betas=(args.train.adam_beta1, args.train.adam_beta2), + weight_decay=args.train.adam_weight_decay, + eps=args.train.adam_epsilon, + ) + + logger.info("***** Prepare dataset *****") + + with accelerator.main_process_first(): + train_dataset = OmniGen2TrainDataset( + args.data.data_path, + tokenizer=text_tokenizer, + num_workers=args.train.dataloader_num_workers, + use_chat_template=args.data.use_chat_template, + prompt_dropout_prob=args.data.get('prompt_dropout_prob', 0.0), + ref_img_dropout_prob=args.data.get('ref_img_dropout_prob', 0.0), + max_input_pixels=OmegaConf.to_object(args.data.get('max_input_pixels', 1024 * 1024)), + max_output_pixels=args.data.get('max_output_pixels', 1024 * 1024), + max_side_length=args.data.get('max_side_length', 2048), + ) + + logger.info(f"Number of training samples: {len(train_dataset)}") + + if args.seed is not None and args.get("workder_specific_seed", False): + from omnigen2.utils.reproducibility import worker_init_fn + + worker_init_fn = partial( + worker_init_fn, + num_processes=AcceleratorState().num_processes, + num_workers=args.train.dataloader_num_workers, + process_index=AcceleratorState().process_index, + seed=args.seed, + same_seed_per_epoch=args.get("same_seed_per_epoch", False), + ) + else: + worker_init_fn = None + + logger.info("***** Prepare dataLoader *****") + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + num_workers=args.train.dataloader_num_workers, + batch_sampler=RepeatedDistributedBatchSampler( + dataset=train_dataset, + batch_size=args.train.batch_size, + num_repeats=args.train.rl.num_images_per_prompt, + num_replicas=AcceleratorState().num_processes, + rank=AcceleratorState().process_index, + shuffle=True, + seed=args.seed, + drop_last=True, + ), + worker_init_fn=worker_init_fn, + collate_fn=OmniGen2Collator(tokenizer=text_tokenizer, max_token_len=args.data.maximum_text_tokens) + ) + + logger.info(f"{args.train.batch_size=} {args.train.gradient_accumulation_steps=} {accelerator.num_processes=} {args.train.global_batch_size=}") + + + assert args.train.batch_size % (args.train.rl.batch_size_per_forward * args.train.gradient_accumulation_steps) == 0, f"{args.train.batch_size=} % ({args.train.rl.batch_size_per_forward=} * {args.train.rl.gradient_accumulation_steps=}) != 0" + assert args.train.batch_size // (args.train.rl.batch_size_per_forward * args.train.gradient_accumulation_steps) == args.train.rl.num_update_steps_per_sampling + assert args.train.global_batch_size // args.train.rl.num_images_per_prompt == args.train.rl.num_unique_prompts_per_sampling, f"{args.train.global_batch_size=} // {args.train.rl.num_images_per_prompt=} != {args.train.rl.num_unique_prompts_per_sampling=}" + + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataloader) * args.train.rl.num_update_steps_per_sampling) + if 'max_train_steps' not in args.train: + args.train.max_train_steps = args.train.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + if args.train.lr_scheduler == 'timm_cosine': + from omnigen2.optim.scheduler.cosine_lr import CosineLRScheduler + + lr_scheduler = CosineLRScheduler(optimizer=optimizer, + t_initial=args.train.t_initial, + lr_min=args.train.lr_min, + cycle_decay=args.train.cycle_decay, + warmup_t=args.train.warmup_t, + warmup_lr_init=args.train.warmup_lr_init, + warmup_prefix=args.train.warmup_prefix, + t_in_epochs=args.train.t_in_epochs) + elif args.train.lr_scheduler == 'timm_constant_with_warmup': + from omnigen2.optim.scheduler.step_lr import StepLRScheduler + + lr_scheduler = StepLRScheduler( + optimizer=optimizer, + decay_t=1, + decay_rate=1, + warmup_t=args.train.warmup_t, + warmup_lr_init=args.train.warmup_lr_init, + warmup_prefix=args.train.warmup_prefix, + t_in_epochs=args.train.t_in_epochs, + ) + else: + lr_scheduler = get_scheduler( + args.train.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.train.lr_warmup_steps, + num_training_steps=args.train.max_train_steps, + num_cycles=args.train.lr_num_cycles, + power=args.train.lr_power, + ) + + logger.info("***** Prepare everything with our accelerator *****") + + if args.train.ema_decay != 0: + model, model_ema, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + model, model_ema, optimizer, train_dataloader, lr_scheduler + ) + model_ema = EMAModel(model_ema.parameters(), decay=ema_decay, model_cls=type(unwrap_model(model)), model_config=model_ema.config) + else: + model, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + model, optimizer, train_dataloader, lr_scheduler + ) + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil(len(train_dataloader) * args.train.rl.num_update_steps_per_sampling) + if overrode_max_train_steps: + args.train.max_train_steps = args.train.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.train.num_train_epochs = math.ceil(args.train.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + accelerator.init_trackers("OmniGen2-RL", init_kwargs={"wandb": {"name": args.name}}) + + # Train! + total_batch_size = args.train.batch_size * accelerator.num_processes * args.train.gradient_accumulation_steps + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num batches each epoch = {len(train_dataloader)}") + logger.info(f" Num Epochs = {args.train.num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {args.train.batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Gradient Accumulation steps = {args.train.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.train.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the most recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + first_epoch = global_step // num_update_steps_per_epoch + else: + initial_global_step = 0 + + progress_bar = tqdm( + range(0, args.train.max_train_steps), + initial=initial_global_step, + desc="Steps", + # Only show the progress bar once on each machine. + disable=not accelerator.is_local_main_process, + file=TqdmToLogger(logger, level=logging.INFO) + ) + + if accelerator.is_main_process: + for tracker in accelerator.trackers: + if tracker.name == "wandb": + logger.info(f"***** Wandb log dir: {tracker.run.dir} *****") + + from omnigen2.pipelines.omnigen2.pipeline_omnigen2 import OmniGen2Pipeline + from omnigen2.schedulers.scheduling_flow_match_euler_maruyama_discrete import FlowMatchEulerMaruyamaDiscreteScheduler + pipeline = OmniGen2Pipeline( + transformer=model, + vae=vae, + scheduler=FlowMatchEulerMaruyamaDiscreteScheduler( + sigma_coef=args.train.rl.get('sigma_coef', 0.7), + time_shift_base_res=args.train.rl.get('time_shift_base_res', 320) + ), + mllm=None, + processor=processor, + ) + pipeline.set_progress_bar_config(disable=True) + + with torch.no_grad(): + + if args.train.rl.get('use_ori_neg_prompt_template', False): + negative_prompt = [ + { + "role": "system", + "content": "You are a helpful assistant.", + }, + {"role": "user", "content": args.train.rl.negative_prompt}, + ] + negative_prompt = pipeline.processor.tokenizer.apply_chat_template( + negative_prompt, tokenize=False, add_generation_prompt=False + ) + else: + negative_prompt = pipeline._apply_chat_template(args.train.rl.negative_prompt) + + negative_prompt_embeds, negative_prompt_attention_mask = get_qwen2_prompt_embeds( + text_encoder=text_encoder, + tokenizer=text_tokenizer, + prompt=negative_prompt, + device=accelerator.device, + max_sequence_length=1024, + ) + + seq_len = negative_prompt_embeds.shape[1] + + negative_prompt_embeds = negative_prompt_embeds.repeat(1, args.train.rl.batch_size_per_forward, 1) + negative_prompt_embeds = negative_prompt_embeds.view(args.train.rl.batch_size_per_forward, seq_len, -1) + negative_prompt_attention_mask = negative_prompt_attention_mask.repeat(args.train.rl.batch_size_per_forward, 1) + negative_prompt_attention_mask = negative_prompt_attention_mask.view( + args.train.rl.batch_size_per_forward, -1 + ) + + ref_latents_N, ref_img_mask_N, l_effective_ref_img_len_N, ref_img_sizes_N = pipeline.transformer.flat_and_pad_to_seq_ref_img(None, args.train.rl.batch_size_per_forward, weight_dtype, accelerator.device) + + reward_server_config = OmegaConf.load(args.reward_server_config) + + for epoch in range(first_epoch, args.train.num_train_epochs): + if 'max_train_steps' in args.train and global_step >= args.train.max_train_steps: + break + train_dataloader.batch_sampler.batch_sampler.set_epoch(epoch) + for step, batch in enumerate(train_dataloader): + instruction = batch['instruction'] + input_images = batch['input_images'] + input_images_pil = batch['input_images_pil'] + target_img_size = batch['target_img_size'] + text_mask = batch['text_mask'] + text_input_ids = batch['text_ids'] + + total_results = None + total_text_feats = None + + batch_size_per_forward = args.train.rl.batch_size_per_forward + for i in range(args.train.batch_size // batch_size_per_forward): + with torch.no_grad(): + text_feats = text_encoder( + input_ids=text_input_ids[i*batch_size_per_forward:(i+1)*batch_size_per_forward], + attention_mask=text_mask[i*batch_size_per_forward:(i+1)*batch_size_per_forward], + output_hidden_states=False, + ).last_hidden_state + + results = pipeline( + prompt_embeds=text_feats, + prompt_attention_mask=text_mask[i*batch_size_per_forward:(i+1)*batch_size_per_forward], + negative_prompt_embeds=negative_prompt_embeds, + negative_prompt_attention_mask=negative_prompt_attention_mask, + input_images=input_images[i*batch_size_per_forward:(i+1)*batch_size_per_forward], + size=target_img_size[i*batch_size_per_forward:(i+1)*batch_size_per_forward], + num_inference_steps=args.train.rl.num_inference_step, + max_sequence_length=1024, + text_guidance_scale=args.train.rl.text_guidance_scale, + image_guidance_scale=args.train.rl.image_guidance_scale, + cfg_range=(args.train.rl.cfg_range_start, args.train.rl.cfg_range_end), + num_images_per_prompt=1, + output_type="pil", + enable_parallel_cfg=True, + return_middle_statistics=True, + mixed_precision=True, + do_normalize=False + ) + + if i == 0: + total_text_feats = text_feats + total_results = results + for k in ['img_mask', 'ref_latents', 'ref_img_mask', 'middle_latents']: + total_results.__dict__[k] = [total_results.__dict__[k]] + else: + total_text_feats = torch.cat([total_text_feats, text_feats], dim=0) + + for k in ['images', 'l_effective_img_len', 'img_sizes', 'l_effective_ref_img_len', 'ref_img_sizes']: + total_results.__dict__[k].extend(results.__dict__[k]) + + for k in ['img_mask', 'ref_latents', 'ref_img_mask', 'middle_latents']: + total_results.__dict__[k].append(results.__dict__[k]) + + for i in range(len(results.log_probs)): + total_results.log_probs[i] = torch.cat([total_results.log_probs[i], results.log_probs[i]], dim=0) + + for i in range(len(batch['meta_data'])): + json_data = json.loads(batch['meta_data'][i]) + json_data['id'] = f"{global_step * args.train.global_batch_size + accelerator.process_index * args.train.batch_size + i}" + batch['meta_data'][i] = json.dumps(json_data) + + local_batch_size = len(input_images_pil) + gathered_input_images_pil = gather_object(input_images_pil) + gathered_output_images = gather_object(total_results.images) + gathered_meta_data = gather_object(batch['meta_data']) + + if accelerator.is_main_process: + scores, rewards, reasoning, meta_data = evaluate_images( + input_images=gathered_input_images_pil, + output_image=gathered_output_images, + meta_datas=gathered_meta_data, + proxy_host=reward_server_config.server.hosts[0], + proxy_port=reward_server_config.server.proxy_port, + server_type=args.train.rl.get('server_type', 'vlm') + ) + + rewards_to_scatter = [rewards[i:i + local_batch_size] for i in range(0, len(rewards), local_batch_size)] + reasoning_to_scatter = [reasoning[i:i + local_batch_size] for i in range(0, len(reasoning), local_batch_size)] + meta_data_to_scatter = [meta_data[i:i + local_batch_size] for i in range(0, len(meta_data), local_batch_size)] + else: + rewards_to_scatter = [None for _ in range(accelerator.num_processes)] + reasoning_to_scatter = [None for _ in range(accelerator.num_processes)] + meta_data_to_scatter = [None for _ in range(accelerator.num_processes)] + + accelerator.wait_for_everyone() + # Extract the current processโ€™s own rewards, reasoning, and meta_data. + rewards = [None] + reasoning = [None] + meta_data = [None] + torch.distributed.scatter_object_list(rewards, rewards_to_scatter) + torch.distributed.scatter_object_list(reasoning, reasoning_to_scatter) + torch.distributed.scatter_object_list(meta_data, meta_data_to_scatter) + rewards = rewards[0] + reasoning = reasoning[0] + meta_data = meta_data[0] + + assert len(rewards) == len(reasoning) == len(meta_data) == local_batch_size + + rewards = torch.tensor(rewards, dtype=torch.float32, device=accelerator.device) + + advantages, prompt_stats = process_grpo_rewards( + rewards=rewards, + prompts=instruction, + accelerator=accelerator, + std_level=args.train.rl.get('std_level', 'group'), + ) + + reuse_samples_nums = args.train.rl.reuse_samples_nums # reuse times of samples + clip_range = args.train.rl.clip_range # PPO clip range + + # prepare data for GRPO + timesteps = pipeline.scheduler._timesteps # [batch_size, num_timesteps+1] + + assert reuse_samples_nums == 1 + + for reuse_step in range(reuse_samples_nums): + + logs = defaultdict(list) + for forward_step in range(args.train.batch_size // batch_size_per_forward): + results = FMPipelineOutput( + images=[], + middle_latents=[], + log_probs=[], + img_mask=[], + l_effective_img_len=[], + img_sizes=[], + ref_latents=[], + ref_img_mask=[], + l_effective_ref_img_len=[], + ref_img_sizes=[], + ) + + for k in ['images', 'l_effective_img_len', 'img_sizes', 'l_effective_ref_img_len', 'ref_img_sizes']: + results.__dict__[k] = total_results.__dict__[k][forward_step*batch_size_per_forward:(forward_step+1)*batch_size_per_forward] + + for k in ['img_mask', 'ref_latents', 'ref_img_mask', 'middle_latents']: + results.__dict__[k] = total_results.__dict__[k][forward_step] + + results.log_probs = [total_results.log_probs[i][forward_step*batch_size_per_forward:(forward_step+1)*batch_size_per_forward] for i in range(len(total_results.log_probs))] + + text_feats = total_text_feats[forward_step*batch_size_per_forward:(forward_step+1)*batch_size_per_forward] + + old_log_probs = [total_results.log_probs[i][forward_step*batch_size_per_forward:(forward_step+1)*batch_size_per_forward] for i in range(len(total_results.log_probs))] + + train_timesteps = list(range(args.train.rl.num_inference_step)) + sample_steps = math.ceil(args.train.rl.num_inference_step * args.train.rl.get('train_timesteps_fraction', 1.0)) + train_timesteps = sorted(random.sample(train_timesteps, k=sample_steps)) + + if args.train.rl.policy_loss_reweighting: + sigma_ts = [] + for i in train_timesteps: + t = timesteps[0, i] + t_next = timesteps[0, i+1] + + sigma_t = pipeline.scheduler.get_sigma_t(t, t_next if i == 0 else None) # [batch_size] + dt = t_next - t + sigma_ts.append(sigma_t * math.sqrt(dt)) + + sigma_ts = torch.stack(sigma_ts) + normalize_factor = sigma_ts.mean() + + for idx, i in enumerate(train_timesteps): + with accelerator.accumulate(model): + text_guidance_scale = args.train.rl.text_guidance_scale if args.train.rl.cfg_range_start <= i / args.train.rl.num_inference_step <= args.train.rl.cfg_range_end else 1.0 + image_guidance_scale = args.train.rl.image_guidance_scale if args.train.rl.cfg_range_start <= i / args.train.rl.num_inference_step <= args.train.rl.cfg_range_end else 1.0 + + latents = results.middle_latents[i] + latents_next = results.middle_latents[i+1] + t = timesteps[:, i] + t_next = timesteps[:, i+1] + + model_kwargs = dict( + hidden_states=latents, + timestep=t, + freqs_cis=freqs_cis, + flat_and_pad=False, + img_mask=results.img_mask, + l_effective_img_len=results.l_effective_img_len, + img_sizes=results.img_sizes, + ) + model_pred_kwargs = dict( + text_hidden_states=text_feats, + text_attention_mask=text_mask[forward_step*batch_size_per_forward:(forward_step+1)*batch_size_per_forward], + ref_image_hidden_states=results.ref_latents, + ref_img_mask=results.ref_img_mask, + l_effective_ref_img_len=results.l_effective_ref_img_len, + ref_img_sizes=results.ref_img_sizes, + ) + + if image_guidance_scale > 1 and text_guidance_scale > 1: + model_kwargs['hidden_states'] = torch.cat([latents, latents, latents], dim=0) + model_kwargs['timestep'] = torch.cat([t, t, t], dim=0) + model_kwargs['img_mask'] = torch.cat([results.img_mask, results.img_mask, results.img_mask], dim=0) + model_kwargs['l_effective_img_len'] = results.l_effective_img_len * 3 + model_kwargs['img_sizes'] = results.img_sizes * 3 + + model_pred_kwargs['text_hidden_states'] = torch.cat([text_feats, pad_to_length(negative_prompt_embeds, len=text_feats.shape[1]), pad_to_length(negative_prompt_embeds, len=text_feats.shape[1])], dim=0) + model_pred_kwargs['text_attention_mask'] = torch.cat([text_mask[forward_step*batch_size_per_forward:(forward_step+1)*batch_size_per_forward], pad_to_length(negative_prompt_attention_mask, len=text_mask.shape[1]), pad_to_length(negative_prompt_attention_mask, len=text_mask.shape[1])], dim=0) + model_pred_kwargs['ref_image_hidden_states'] = torch.cat([results.ref_latents, results.ref_latents, pad_to_length(ref_latents_N, len=results.ref_latents.shape[1])], dim=0) + model_pred_kwargs['ref_img_mask'] = torch.cat([results.ref_img_mask, results.ref_img_mask, pad_to_length(ref_img_mask_N, len=results.ref_img_mask.shape[1])], dim=0) + model_pred_kwargs['l_effective_ref_img_len'] = results.l_effective_ref_img_len * 2 + l_effective_ref_img_len_N + model_pred_kwargs['ref_img_sizes'] = results.ref_img_sizes * 2 + ref_img_sizes_N + + elif text_guidance_scale > 1: + model_kwargs['hidden_states'] = torch.cat([latents, latents], dim=0) + model_kwargs['timestep'] = torch.cat([t, t], dim=0) + model_kwargs['img_mask'] = torch.cat([results.img_mask, results.img_mask], dim=0) + model_kwargs['l_effective_img_len'] = results.l_effective_img_len * 2 + model_kwargs['img_sizes'] = results.img_sizes * 2 + + model_pred_kwargs['text_hidden_states'] = torch.cat([text_feats, pad_to_length(negative_prompt_embeds, len=text_feats.shape[1])], dim=0) + model_pred_kwargs['text_attention_mask'] = torch.cat([text_mask[forward_step*batch_size_per_forward:(forward_step+1)*batch_size_per_forward], pad_to_length(negative_prompt_attention_mask, len=text_mask.shape[1])], dim=0) + model_pred_kwargs['ref_image_hidden_states'] = torch.cat([results.ref_latents, pad_to_length(ref_latents_N, len=results.ref_latents.shape[1])], dim=0) + model_pred_kwargs['ref_img_mask'] = torch.cat([results.ref_img_mask, pad_to_length(ref_img_mask_N, len=results.ref_img_mask.shape[1])], dim=0) + model_pred_kwargs['l_effective_ref_img_len'] = results.l_effective_ref_img_len + l_effective_ref_img_len_N + model_pred_kwargs['ref_img_sizes'] = results.ref_img_sizes + ref_img_sizes_N + + step_log_probs, mean_t, sigma_t = forward_logprob( + latents=latents, + latents_next=latents_next, + t=t, + t_next=t_next, + step_index=i, + img_mask=results.img_mask, + model=model, + model_kwargs=model_kwargs, + model_pred_kwargs=model_pred_kwargs, + scheduler=pipeline.scheduler, + apply_cfg=args.train.rl.apply_cfg_in_training, + text_guidance_scale=text_guidance_scale, + image_guidance_scale=image_guidance_scale, + ) + + if args.train.rl.kl_loss_weight > 0: + with torch.no_grad(): + with disabled_adapters(unwrap_model(model)): + _, mean_t_ref, _ = forward_logprob( + latents=latents, + latents_next=latents_next, + t=t, + t_next=t_next, + step_index=i, + img_mask=results.img_mask, + model=model, + model_kwargs=model_kwargs, + model_pred_kwargs=model_pred_kwargs, + scheduler=pipeline.scheduler, + apply_cfg=args.train.rl.apply_cfg_in_training, + text_guidance_scale=text_guidance_scale, + image_guidance_scale=image_guidance_scale, + ) + + loss = 0 + + ( + policy_loss, + policy_clip_frac, + approx_kl, + unclipped_loss, + clipped_loss, + ratio, + ratio_positive, + ratio_negative, + num_positive, + num_negative, + ratio_large_than_1, + ratio_small_than_1, + ) = compute_single_step_ppo_loss( + step_log_probs=step_log_probs, + old_step_log_probs=old_log_probs[i], + advantages=advantages[ + forward_step * batch_size_per_forward : ( + forward_step + 1 + ) + * batch_size_per_forward + ], + clip_range=clip_range, + adv_clip_max=args.train.rl.adv_clip_max, + ) + logs['policy_loss'].append(policy_loss.detach()) + logs['policy_clip_frac'].append(policy_clip_frac.detach()) + logs['approx_kl'].append(approx_kl.detach()) + logs['advantages'].append(advantages[forward_step*batch_size_per_forward:(forward_step+1)*batch_size_per_forward].mean().detach()) + logs['advantages_std'].append(advantages[forward_step*batch_size_per_forward:(forward_step+1)*batch_size_per_forward].std().detach()) + logs['advantages_min'].append(advantages[forward_step*batch_size_per_forward:(forward_step+1)*batch_size_per_forward].min().detach()) + logs['advantages_max'].append(advantages[forward_step*batch_size_per_forward:(forward_step+1)*batch_size_per_forward].max().detach()) + logs['policy_loss_unclipped'].append(unclipped_loss.mean().detach()) + logs['policy_loss_clipped'].append(clipped_loss.mean().detach()) + logs['ratio'].append(ratio.mean().detach()) + logs['ratio_positive'].append(ratio_positive.detach()) + logs['ratio_negative'].append(ratio_negative.detach()) + logs['num_positive'].append(num_positive.detach()) + logs['num_negative'].append(num_negative.detach()) + logs['ratio_large_than_1'].append(ratio_large_than_1.detach()) + logs['ratio_small_than_1'].append(ratio_small_than_1.detach()) + + if args.train.rl.policy_loss_reweighting: + policy_loss = policy_loss * sample_steps * (sigma_ts[idx] / normalize_factor) + + loss += policy_loss + + if args.train.rl.kl_loss_weight > 0: + + kl_loss = torch.mean( + ((mean_t - mean_t_ref.detach()) ** 2) + .flatten(start_dim=1) + .mean(dim=1) + / (2 * sigma_t) + ) + logs['kl_loss'].append(kl_loss.detach()) + + loss += args.train.rl.kl_loss_weight * kl_loss + + logs['loss'].append(loss.detach()) + accelerator.backward(loss) + if accelerator.sync_gradients: + grad_norm = accelerator.clip_grad_norm_(model.parameters(), args.train.max_grad_norm) + logs['grad_norm'].append(grad_norm.to(accelerator.device).detach()) + optimizer.step() + if 'timm' in args.train.lr_scheduler: + lr_scheduler.step(global_step) + else: + lr_scheduler.step() + optimizer.zero_grad(set_to_none=args.train.set_grads_to_none) + + # Checks if the accelerator has performed an optimization step behind the scenes + + if accelerator.sync_gradients: + + logs = {k: torch.mean(torch.stack(v)) for k, v in logs.items()} + logs = accelerator.reduce(logs, reduction="mean") + logs = {k: v.item() for k, v in logs.items()} + logs.update( + { + "lr": lr_scheduler.get_last_lr()[0], + "rewards_min": np.mean( + [v["min"] for k, v in prompt_stats.items()] + ), + "rewards_max": np.mean( + [v["max"] for k, v in prompt_stats.items()] + ), + "rewards_mean": np.mean( + [v["mean"] for k, v in prompt_stats.items()] + ), + "rewards_std": np.mean( + [v["std"] for k, v in prompt_stats.items()] + ), + "zero_std_ratio": np.mean( + [ + v["std"] == 0 for k, v in prompt_stats.items() + ] + ) + } + ) + + if ema_decay != 0: + model_ema.step(model.parameters()) + + global_step += 1 + + if global_step % args.logger.checkpointing_steps == 0: + if accelerator.is_main_process: + if args.logger.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + if len(checkpoints) >= args.logger.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.logger.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + accelerator.wait_for_everyone() + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + if 'train_visualization_interval' in args.val and (global_step - 1) % args.val.train_visualization_interval == 0: + num_samples = min(args.val.get('num_train_visualization_samples', 2), args.train.batch_size) + + if accelerator.is_main_process: + target_instruction = instruction[:num_samples] + else: + target_instruction = [None] * num_samples + torch.distributed.broadcast_object_list(target_instruction) + + rewards_per_instruction = [[] for _ in range(num_samples)] + for i in range(len(target_instruction)): + for j in range(len(instruction)): + if target_instruction[i] == instruction[j]: + rewards_per_instruction[i].append((rewards[j].item(), accelerator.process_index, j)) + gathered_rewards_per_instruction = [None for _ in range(accelerator.num_processes)] + torch.distributed.all_gather_object(gathered_rewards_per_instruction, rewards_per_instruction) + + gathered_rewards_per_instruction_flat = [[] for _ in range(num_samples)] + for i in range(num_samples): + for j in range(accelerator.num_processes): + gathered_rewards_per_instruction_flat[i].extend(gathered_rewards_per_instruction[j][i]) + + p = [{} for _ in range(num_samples)] + for i in range(len(target_instruction)): + gathered_rewards_per_instruction_flat[i].sort(key=lambda x: x[0]) + for j in range(len(gathered_rewards_per_instruction_flat[i])): + if gathered_rewards_per_instruction_flat[i][j][1] == accelerator.process_index: + p[i][gathered_rewards_per_instruction_flat[i][j][2]] = j + + with torch.no_grad(): + for i in range(len(target_instruction)): + cnt = 0 + for j in range(len(instruction)): + if instruction[j] == target_instruction[i]: + if cnt == 0: + for k in range(len(input_images_pil[j])): + input_images_pil[j][k].save(os.path.join(args.output_dir, f"input_visualization_{global_step}_{i}_input_{k}.png")) + + total_results.images[j].save(os.path.join(args.output_dir, f"input_visualization_{global_step}_{i}_{p[i][j]}.png")) + with open(os.path.join(args.output_dir, f"instruction_{global_step}_{i}_{p[i][j]}.txt"), "w", encoding='utf-8') as f: + f.write(f"instruction: {instruction[j]}\nreward: {rewards[j].item()}\nadvantages: {advantages[j].item()}\nreasoning: {reasoning[j]}\nmeta_data: {batch['meta_data'][j]}\ncur_receive_meta_data: {meta_data[j]}") + cnt += 1 + + progress_bar.set_postfix(**logs) + progress_bar.update(1) + + accelerator.log(logs, step=global_step) + logs = defaultdict(list) + + if 'max_train_steps' in args.train and global_step >= args.train.max_train_steps: + break + + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + if len(checkpoints) > 0 and int(checkpoints[-1].split("-")[1]) < global_step: + if accelerator.is_main_process: + if args.logger.checkpoints_total_limit is not None: + if len(checkpoints) >= args.logger.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.logger.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + accelerator.wait_for_everyone() + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + + accelerator.end_training() + +if __name__ == "__main__": + root_path = os.path.abspath(os.path.join(__file__, os.path.pardir)) + args = parse_args(root_path) + main(args) \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000000000000000000000000000000000000..e64d15e4e4a091d21ae6447b635288e11d979c66 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,39 @@ +[build-system] +requires = ["setuptools>=61.0"] +build-backend = "setuptools.build_meta" + +[project] +name = "editscore" +version = "0.2" +authors = [ + { name="Xin Luo", email="xinluo@mail.ustc.edu.cn" }, + { name="Jiahao Wang", email="jiahaowang0917@gmail.com" }, + { name="Chenyuan Wu", email="wuchenyuan@mail.ustc.edu.cn" }, +] +description = "A high-fidelity reward model for instruction-based image editing." +readme = "README.md" +requires-python = ">=3.8" +classifiers = [ + "Programming Language :: Python :: 3", + "License :: OSI Approved :: Apache Software License", # ๅ‡่ฎพๆ˜ฏ Apache 2.0 + "Operating System :: OS Independent", +] +# ่ฟ™้‡Œๅˆ—ๅ‡บ EditScore ๆ ธๅฟƒๅบ“็š„ไพ่ต– +dependencies = [ + "torch", + "torchvision", + "accelerate", + "transformers", + "qwen-vl-utils", + "peft", + "Pillow", + "json-repair" +] + +[project.urls] +Homepage = "https://github.com/VectorSpaceLab/EditScore" +Issues = "https://github.com/VectorSpaceLab/EditScore/issues" + +[tool.setuptools.packages.find] +# ่‡ชๅŠจๅ‘็Žฐๅไธบ 'editscore' ็š„ๅŒ…๏ผˆๅณ editscore/ ็›ฎๅฝ•๏ผ‰ +where = ["."] \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000000000000000000000000000000000000..486baa870b7fae9fc0f31ce5643baef8eeee7cc0 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,4 @@ +datasets +tqdm +python-dotenv +wheel \ No newline at end of file