Initial mirror of VectorSpaceLab/EditScore@4609c5d
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +6 -0
- .gitignore +239 -0
- FORK_PROVENANCE.md +11 -0
- LICENSE +201 -0
- README.md +215 -0
- assets/figure_edit_results.png +3 -0
- assets/logo.png +3 -0
- assets/table_editscore_qwen3_vl.png +3 -0
- assets/table_reward_model_results.png +3 -0
- calculate_statistics.py +124 -0
- editscore/__init__.py +217 -0
- editscore/json_parser.py +173 -0
- editscore/mllm_tools/__init__.py +0 -0
- editscore/mllm_tools/internvl35_lmdeploy.py +73 -0
- editscore/mllm_tools/openai.py +180 -0
- editscore/mllm_tools/qwen25vl.py +97 -0
- editscore/mllm_tools/qwen25vl_vllm.py +138 -0
- editscore/mllm_tools/qwen3vl.py +95 -0
- editscore/mllm_tools/qwen3vl_vllm.py +141 -0
- editscore/mllm_tools/utils.py +65 -0
- editscore/utils.py +534 -0
- editscore/vie_prompts.py +46 -0
- evaluate.sh +20 -0
- evaluate_72B_vllm.sh +20 -0
- evaluate_qwen3_vl_32B.sh +20 -0
- evaluate_qwen3_vl_32B_avg4.sh +20 -0
- evaluate_qwen3_vl_4B.sh +20 -0
- evaluate_qwen3_vl_4B_avg4.sh +20 -0
- evaluate_qwen3_vl_4B_vllm.sh +20 -0
- evaluate_qwen3_vl_8B.sh +20 -0
- evaluate_qwen3_vl_8B_avg4.sh +20 -0
- evaluate_vllm.sh +20 -0
- evaluation.py +272 -0
- example_images/input.png +3 -0
- example_images/output.png +3 -0
- examples/EditScore-train/README.md +101 -0
- examples/EditScore-train/config/editscore_32B.yaml +42 -0
- examples/EditScore-train/config/editscore_72B.yaml +42 -0
- examples/EditScore-train/config/editscore_7B.yaml +41 -0
- examples/EditScore-train/config/editscore_qwen3_vl_4B_instruct.yaml +41 -0
- examples/EditScore-train/config/editscore_qwen3_vl_8B_instruct.yaml +41 -0
- examples/EditScore-train/train.sh +52 -0
- examples/OmniGen2-RL/.gitignore +233 -0
- examples/OmniGen2-RL/LICENSE +201 -0
- examples/OmniGen2-RL/README.md +189 -0
- examples/OmniGen2-RL/data_configs/train/example/edit/all.yml +7 -0
- examples/OmniGen2-RL/data_configs/train/example/train.yml +5 -0
- examples/OmniGen2-RL/docs/README.md +2 -0
- examples/OmniGen2-RL/evaluation/GEdit-Bench/calculate_statistics.py +223 -0
- examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples.sh +83 -0
.gitattributes
CHANGED
|
@@ -33,3 +33,9 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
| 36 |
+
assets/figure_edit_results.png filter=lfs diff=lfs merge=lfs -text
|
| 37 |
+
assets/logo.png filter=lfs diff=lfs merge=lfs -text
|
| 38 |
+
assets/table_editscore_qwen3_vl.png filter=lfs diff=lfs merge=lfs -text
|
| 39 |
+
assets/table_reward_model_results.png filter=lfs diff=lfs merge=lfs -text
|
| 40 |
+
example_images/input.png filter=lfs diff=lfs merge=lfs -text
|
| 41 |
+
example_images/output.png filter=lfs diff=lfs merge=lfs -text
|
.gitignore
ADDED
|
@@ -0,0 +1,239 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Created by https://www.toptal.com/developers/gitignore/api/macos,python
|
| 2 |
+
# Edit at https://www.toptal.com/developers/gitignore?templates=macos,python
|
| 3 |
+
|
| 4 |
+
### macOS ###
|
| 5 |
+
# General
|
| 6 |
+
.DS_Store
|
| 7 |
+
.AppleDouble
|
| 8 |
+
.LSOverride
|
| 9 |
+
|
| 10 |
+
# Icon must end with two \r
|
| 11 |
+
Icon
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
# Thumbnails
|
| 15 |
+
._*
|
| 16 |
+
|
| 17 |
+
# Files that might appear in the root of a volume
|
| 18 |
+
.DocumentRevisions-V100
|
| 19 |
+
.fseventsd
|
| 20 |
+
.Spotlight-V100
|
| 21 |
+
.TemporaryItems
|
| 22 |
+
.Trashes
|
| 23 |
+
.VolumeIcon.icns
|
| 24 |
+
.com.apple.timemachine.donotpresent
|
| 25 |
+
|
| 26 |
+
# Directories potentially created on remote AFP share
|
| 27 |
+
.AppleDB
|
| 28 |
+
.AppleDesktop
|
| 29 |
+
Network Trash Folder
|
| 30 |
+
Temporary Items
|
| 31 |
+
.apdisk
|
| 32 |
+
|
| 33 |
+
### macOS Patch ###
|
| 34 |
+
# iCloud generated files
|
| 35 |
+
*.icloud
|
| 36 |
+
|
| 37 |
+
### Python ###
|
| 38 |
+
# Byte-compiled / optimized / DLL files
|
| 39 |
+
__pycache__/
|
| 40 |
+
*.py[cod]
|
| 41 |
+
*$py.class
|
| 42 |
+
|
| 43 |
+
# C extensions
|
| 44 |
+
*.so
|
| 45 |
+
|
| 46 |
+
# Distribution / packaging
|
| 47 |
+
.Python
|
| 48 |
+
build/
|
| 49 |
+
develop-eggs/
|
| 50 |
+
dist/
|
| 51 |
+
downloads/
|
| 52 |
+
eggs/
|
| 53 |
+
.eggs/
|
| 54 |
+
lib/
|
| 55 |
+
lib64/
|
| 56 |
+
parts/
|
| 57 |
+
sdist/
|
| 58 |
+
var/
|
| 59 |
+
wheels/
|
| 60 |
+
share/python-wheels/
|
| 61 |
+
*.egg-info/
|
| 62 |
+
.installed.cfg
|
| 63 |
+
*.egg
|
| 64 |
+
MANIFEST
|
| 65 |
+
|
| 66 |
+
# PyInstaller
|
| 67 |
+
# Usually these files are written by a python script from a template
|
| 68 |
+
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
| 69 |
+
*.manifest
|
| 70 |
+
*.spec
|
| 71 |
+
|
| 72 |
+
# Installer logs
|
| 73 |
+
pip-log.txt
|
| 74 |
+
pip-delete-this-directory.txt
|
| 75 |
+
|
| 76 |
+
# Unit test / coverage reports
|
| 77 |
+
htmlcov/
|
| 78 |
+
.tox/
|
| 79 |
+
.nox/
|
| 80 |
+
.coverage
|
| 81 |
+
.coverage.*
|
| 82 |
+
.cache
|
| 83 |
+
nosetests.xml
|
| 84 |
+
coverage.xml
|
| 85 |
+
*.cover
|
| 86 |
+
*.py,cover
|
| 87 |
+
.hypothesis/
|
| 88 |
+
.pytest_cache/
|
| 89 |
+
cover/
|
| 90 |
+
|
| 91 |
+
# Translations
|
| 92 |
+
*.mo
|
| 93 |
+
*.pot
|
| 94 |
+
|
| 95 |
+
# Django stuff:
|
| 96 |
+
*.log
|
| 97 |
+
local_settings.py
|
| 98 |
+
db.sqlite3
|
| 99 |
+
db.sqlite3-journal
|
| 100 |
+
|
| 101 |
+
# Flask stuff:
|
| 102 |
+
instance/
|
| 103 |
+
.webassets-cache
|
| 104 |
+
|
| 105 |
+
# Scrapy stuff:
|
| 106 |
+
.scrapy
|
| 107 |
+
|
| 108 |
+
# Sphinx documentation
|
| 109 |
+
docs/_build/
|
| 110 |
+
|
| 111 |
+
# PyBuilder
|
| 112 |
+
.pybuilder/
|
| 113 |
+
target/
|
| 114 |
+
|
| 115 |
+
# Jupyter Notebook
|
| 116 |
+
.ipynb_checkpoints
|
| 117 |
+
|
| 118 |
+
# IPython
|
| 119 |
+
profile_default/
|
| 120 |
+
ipython_config.py
|
| 121 |
+
|
| 122 |
+
# pyenv
|
| 123 |
+
# For a library or package, you might want to ignore these files since the code is
|
| 124 |
+
# intended to run in multiple environments; otherwise, check them in:
|
| 125 |
+
# .python-version
|
| 126 |
+
|
| 127 |
+
# pipenv
|
| 128 |
+
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
|
| 129 |
+
# However, in case of collaboration, if having platform-specific dependencies or dependencies
|
| 130 |
+
# having no cross-platform support, pipenv may install dependencies that don't work, or not
|
| 131 |
+
# install all needed dependencies.
|
| 132 |
+
#Pipfile.lock
|
| 133 |
+
|
| 134 |
+
# poetry
|
| 135 |
+
# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
|
| 136 |
+
# This is especially recommended for binary packages to ensure reproducibility, and is more
|
| 137 |
+
# commonly ignored for libraries.
|
| 138 |
+
# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
|
| 139 |
+
#poetry.lock
|
| 140 |
+
|
| 141 |
+
# pdm
|
| 142 |
+
# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control.
|
| 143 |
+
#pdm.lock
|
| 144 |
+
# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it
|
| 145 |
+
# in version control.
|
| 146 |
+
# https://pdm.fming.dev/#use-with-ide
|
| 147 |
+
.pdm.toml
|
| 148 |
+
|
| 149 |
+
# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm
|
| 150 |
+
__pypackages__/
|
| 151 |
+
|
| 152 |
+
# Celery stuff
|
| 153 |
+
celerybeat-schedule
|
| 154 |
+
celerybeat.pid
|
| 155 |
+
|
| 156 |
+
# SageMath parsed files
|
| 157 |
+
*.sage.py
|
| 158 |
+
|
| 159 |
+
# Environments
|
| 160 |
+
.env
|
| 161 |
+
.venv
|
| 162 |
+
env/
|
| 163 |
+
venv/
|
| 164 |
+
ENV/
|
| 165 |
+
env.bak/
|
| 166 |
+
venv.bak/
|
| 167 |
+
|
| 168 |
+
# Spyder project settings
|
| 169 |
+
.spyderproject
|
| 170 |
+
.spyproject
|
| 171 |
+
|
| 172 |
+
# Rope project settings
|
| 173 |
+
.ropeproject
|
| 174 |
+
|
| 175 |
+
# mkdocs documentation
|
| 176 |
+
/site
|
| 177 |
+
|
| 178 |
+
# mypy
|
| 179 |
+
.mypy_cache/
|
| 180 |
+
.dmypy.json
|
| 181 |
+
dmypy.json
|
| 182 |
+
|
| 183 |
+
# Pyre type checker
|
| 184 |
+
.pyre/
|
| 185 |
+
|
| 186 |
+
# pytype static type analyzer
|
| 187 |
+
.pytype/
|
| 188 |
+
|
| 189 |
+
# Cython debug symbols
|
| 190 |
+
cython_debug/
|
| 191 |
+
|
| 192 |
+
# PyCharm
|
| 193 |
+
# JetBrains specific template is maintained in a separate JetBrains.gitignore that can
|
| 194 |
+
# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore
|
| 195 |
+
# and can be added to the global gitignore or merged into this file. For a more nuclear
|
| 196 |
+
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
|
| 197 |
+
#.idea/
|
| 198 |
+
|
| 199 |
+
### Python Patch ###
|
| 200 |
+
# Poetry local configuration file - https://python-poetry.org/docs/configuration/#local-configuration
|
| 201 |
+
poetry.toml
|
| 202 |
+
|
| 203 |
+
# ruff
|
| 204 |
+
.ruff_cache/
|
| 205 |
+
|
| 206 |
+
# LSP config files
|
| 207 |
+
pyrightconfig.json
|
| 208 |
+
|
| 209 |
+
# End of https://www.toptal.com/developers/gitignore/api/macos,python
|
| 210 |
+
|
| 211 |
+
local_scripts/
|
| 212 |
+
|
| 213 |
+
omnigen2/utils/vpn_utils.py
|
| 214 |
+
|
| 215 |
+
test_tokenizer.py
|
| 216 |
+
save_pipeline.py
|
| 217 |
+
app.sh
|
| 218 |
+
logs/
|
| 219 |
+
results/
|
| 220 |
+
test_jsonl*
|
| 221 |
+
pbs_files/
|
| 222 |
+
convert_ckpt_to_pipeline.py
|
| 223 |
+
inference_test_efficiency.py
|
| 224 |
+
upload_pipeline*
|
| 225 |
+
example_images_resized/
|
| 226 |
+
example_t2i_test_efficiency*.sh
|
| 227 |
+
example_edit_test_efficiency*.sh
|
| 228 |
+
example_in_context_generation_test_efficiency*.sh
|
| 229 |
+
intro*
|
| 230 |
+
resize_example_images.py
|
| 231 |
+
save_pipeline.py
|
| 232 |
+
outputs_gradio/*
|
| 233 |
+
test.py
|
| 234 |
+
|
| 235 |
+
upload_to_huggingface.py
|
| 236 |
+
upload_to_modelscope.py
|
| 237 |
+
|
| 238 |
+
editscore.egg-info/
|
| 239 |
+
dist/
|
FORK_PROVENANCE.md
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Fork provenance
|
| 2 |
+
|
| 3 |
+
- **Upstream:** `VectorSpaceLab/EditScore (github)`
|
| 4 |
+
- **Upstream commit / SHA:** `4609c5d2ebb62fdebf665d3c924686d896ef1f74`
|
| 5 |
+
- **License:** Apache-2.0
|
| 6 |
+
- **Kind:** Code mirror
|
| 7 |
+
- **Workspace consumer:** EditScore-7B (training + inference code)
|
| 8 |
+
- **Forked to:** `chibifire/EditScore-code`
|
| 9 |
+
- **Forked on:** 2026-09-05
|
| 10 |
+
- **Author:** Ernest Lee <ernest.lee@chibifire.com>
|
| 11 |
+
- **Reason:** shipping-surface completeness under our own HF org so a base/source we depend on cannot disappear or relicense out from under us
|
LICENSE
ADDED
|
@@ -0,0 +1,201 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Apache License
|
| 2 |
+
Version 2.0, January 2004
|
| 3 |
+
http://www.apache.org/licenses/
|
| 4 |
+
|
| 5 |
+
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
| 6 |
+
|
| 7 |
+
1. Definitions.
|
| 8 |
+
|
| 9 |
+
"License" shall mean the terms and conditions for use, reproduction,
|
| 10 |
+
and distribution as defined by Sections 1 through 9 of this document.
|
| 11 |
+
|
| 12 |
+
"Licensor" shall mean the copyright owner or entity authorized by
|
| 13 |
+
the copyright owner that is granting the License.
|
| 14 |
+
|
| 15 |
+
"Legal Entity" shall mean the union of the acting entity and all
|
| 16 |
+
other entities that control, are controlled by, or are under common
|
| 17 |
+
control with that entity. For the purposes of this definition,
|
| 18 |
+
"control" means (i) the power, direct or indirect, to cause the
|
| 19 |
+
direction or management of such entity, whether by contract or
|
| 20 |
+
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
| 21 |
+
outstanding shares, or (iii) beneficial ownership of such entity.
|
| 22 |
+
|
| 23 |
+
"You" (or "Your") shall mean an individual or Legal Entity
|
| 24 |
+
exercising permissions granted by this License.
|
| 25 |
+
|
| 26 |
+
"Source" form shall mean the preferred form for making modifications,
|
| 27 |
+
including but not limited to software source code, documentation
|
| 28 |
+
source, and configuration files.
|
| 29 |
+
|
| 30 |
+
"Object" form shall mean any form resulting from mechanical
|
| 31 |
+
transformation or translation of a Source form, including but
|
| 32 |
+
not limited to compiled object code, generated documentation,
|
| 33 |
+
and conversions to other media types.
|
| 34 |
+
|
| 35 |
+
"Work" shall mean the work of authorship, whether in Source or
|
| 36 |
+
Object form, made available under the License, as indicated by a
|
| 37 |
+
copyright notice that is included in or attached to the work
|
| 38 |
+
(an example is provided in the Appendix below).
|
| 39 |
+
|
| 40 |
+
"Derivative Works" shall mean any work, whether in Source or Object
|
| 41 |
+
form, that is based on (or derived from) the Work and for which the
|
| 42 |
+
editorial revisions, annotations, elaborations, or other modifications
|
| 43 |
+
represent, as a whole, an original work of authorship. For the purposes
|
| 44 |
+
of this License, Derivative Works shall not include works that remain
|
| 45 |
+
separable from, or merely link (or bind by name) to the interfaces of,
|
| 46 |
+
the Work and Derivative Works thereof.
|
| 47 |
+
|
| 48 |
+
"Contribution" shall mean any work of authorship, including
|
| 49 |
+
the original version of the Work and any modifications or additions
|
| 50 |
+
to that Work or Derivative Works thereof, that is intentionally
|
| 51 |
+
submitted to Licensor for inclusion in the Work by the copyright owner
|
| 52 |
+
or by an individual or Legal Entity authorized to submit on behalf of
|
| 53 |
+
the copyright owner. For the purposes of this definition, "submitted"
|
| 54 |
+
means any form of electronic, verbal, or written communication sent
|
| 55 |
+
to the Licensor or its representatives, including but not limited to
|
| 56 |
+
communication on electronic mailing lists, source code control systems,
|
| 57 |
+
and issue tracking systems that are managed by, or on behalf of, the
|
| 58 |
+
Licensor for the purpose of discussing and improving the Work, but
|
| 59 |
+
excluding communication that is conspicuously marked or otherwise
|
| 60 |
+
designated in writing by the copyright owner as "Not a Contribution."
|
| 61 |
+
|
| 62 |
+
"Contributor" shall mean Licensor and any individual or Legal Entity
|
| 63 |
+
on behalf of whom a Contribution has been received by Licensor and
|
| 64 |
+
subsequently incorporated within the Work.
|
| 65 |
+
|
| 66 |
+
2. Grant of Copyright License. Subject to the terms and conditions of
|
| 67 |
+
this License, each Contributor hereby grants to You a perpetual,
|
| 68 |
+
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 69 |
+
copyright license to reproduce, prepare Derivative Works of,
|
| 70 |
+
publicly display, publicly perform, sublicense, and distribute the
|
| 71 |
+
Work and such Derivative Works in Source or Object form.
|
| 72 |
+
|
| 73 |
+
3. Grant of Patent License. Subject to the terms and conditions of
|
| 74 |
+
this License, each Contributor hereby grants to You a perpetual,
|
| 75 |
+
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 76 |
+
(except as stated in this section) patent license to make, have made,
|
| 77 |
+
use, offer to sell, sell, import, and otherwise transfer the Work,
|
| 78 |
+
where such license applies only to those patent claims licensable
|
| 79 |
+
by such Contributor that are necessarily infringed by their
|
| 80 |
+
Contribution(s) alone or by combination of their Contribution(s)
|
| 81 |
+
with the Work to which such Contribution(s) was submitted. If You
|
| 82 |
+
institute patent litigation against any entity (including a
|
| 83 |
+
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
| 84 |
+
or a Contribution incorporated within the Work constitutes direct
|
| 85 |
+
or contributory patent infringement, then any patent licenses
|
| 86 |
+
granted to You under this License for that Work shall terminate
|
| 87 |
+
as of the date such litigation is filed.
|
| 88 |
+
|
| 89 |
+
4. Redistribution. You may reproduce and distribute copies of the
|
| 90 |
+
Work or Derivative Works thereof in any medium, with or without
|
| 91 |
+
modifications, and in Source or Object form, provided that You
|
| 92 |
+
meet the following conditions:
|
| 93 |
+
|
| 94 |
+
(a) You must give any other recipients of the Work or
|
| 95 |
+
Derivative Works a copy of this License; and
|
| 96 |
+
|
| 97 |
+
(b) You must cause any modified files to carry prominent notices
|
| 98 |
+
stating that You changed the files; and
|
| 99 |
+
|
| 100 |
+
(c) You must retain, in the Source form of any Derivative Works
|
| 101 |
+
that You distribute, all copyright, patent, trademark, and
|
| 102 |
+
attribution notices from the Source form of the Work,
|
| 103 |
+
excluding those notices that do not pertain to any part of
|
| 104 |
+
the Derivative Works; and
|
| 105 |
+
|
| 106 |
+
(d) If the Work includes a "NOTICE" text file as part of its
|
| 107 |
+
distribution, then any Derivative Works that You distribute must
|
| 108 |
+
include a readable copy of the attribution notices contained
|
| 109 |
+
within such NOTICE file, excluding those notices that do not
|
| 110 |
+
pertain to any part of the Derivative Works, in at least one
|
| 111 |
+
of the following places: within a NOTICE text file distributed
|
| 112 |
+
as part of the Derivative Works; within the Source form or
|
| 113 |
+
documentation, if provided along with the Derivative Works; or,
|
| 114 |
+
within a display generated by the Derivative Works, if and
|
| 115 |
+
wherever such third-party notices normally appear. The contents
|
| 116 |
+
of the NOTICE file are for informational purposes only and
|
| 117 |
+
do not modify the License. You may add Your own attribution
|
| 118 |
+
notices within Derivative Works that You distribute, alongside
|
| 119 |
+
or as an addendum to the NOTICE text from the Work, provided
|
| 120 |
+
that such additional attribution notices cannot be construed
|
| 121 |
+
as modifying the License.
|
| 122 |
+
|
| 123 |
+
You may add Your own copyright statement to Your modifications and
|
| 124 |
+
may provide additional or different license terms and conditions
|
| 125 |
+
for use, reproduction, or distribution of Your modifications, or
|
| 126 |
+
for any such Derivative Works as a whole, provided Your use,
|
| 127 |
+
reproduction, and distribution of the Work otherwise complies with
|
| 128 |
+
the conditions stated in this License.
|
| 129 |
+
|
| 130 |
+
5. Submission of Contributions. Unless You explicitly state otherwise,
|
| 131 |
+
any Contribution intentionally submitted for inclusion in the Work
|
| 132 |
+
by You to the Licensor shall be under the terms and conditions of
|
| 133 |
+
this License, without any additional terms or conditions.
|
| 134 |
+
Notwithstanding the above, nothing herein shall supersede or modify
|
| 135 |
+
the terms of any separate license agreement you may have executed
|
| 136 |
+
with Licensor regarding such Contributions.
|
| 137 |
+
|
| 138 |
+
6. Trademarks. This License does not grant permission to use the trade
|
| 139 |
+
names, trademarks, service marks, or product names of the Licensor,
|
| 140 |
+
except as required for reasonable and customary use in describing the
|
| 141 |
+
origin of the Work and reproducing the content of the NOTICE file.
|
| 142 |
+
|
| 143 |
+
7. Disclaimer of Warranty. Unless required by applicable law or
|
| 144 |
+
agreed to in writing, Licensor provides the Work (and each
|
| 145 |
+
Contributor provides its Contributions) on an "AS IS" BASIS,
|
| 146 |
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
| 147 |
+
implied, including, without limitation, any warranties or conditions
|
| 148 |
+
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
| 149 |
+
PARTICULAR PURPOSE. You are solely responsible for determining the
|
| 150 |
+
appropriateness of using or redistributing the Work and assume any
|
| 151 |
+
risks associated with Your exercise of permissions under this License.
|
| 152 |
+
|
| 153 |
+
8. Limitation of Liability. In no event and under no legal theory,
|
| 154 |
+
whether in tort (including negligence), contract, or otherwise,
|
| 155 |
+
unless required by applicable law (such as deliberate and grossly
|
| 156 |
+
negligent acts) or agreed to in writing, shall any Contributor be
|
| 157 |
+
liable to You for damages, including any direct, indirect, special,
|
| 158 |
+
incidental, or consequential damages of any character arising as a
|
| 159 |
+
result of this License or out of the use or inability to use the
|
| 160 |
+
Work (including but not limited to damages for loss of goodwill,
|
| 161 |
+
work stoppage, computer failure or malfunction, or any and all
|
| 162 |
+
other commercial damages or losses), even if such Contributor
|
| 163 |
+
has been advised of the possibility of such damages.
|
| 164 |
+
|
| 165 |
+
9. Accepting Warranty or Additional Liability. While redistributing
|
| 166 |
+
the Work or Derivative Works thereof, You may choose to offer,
|
| 167 |
+
and charge a fee for, acceptance of support, warranty, indemnity,
|
| 168 |
+
or other liability obligations and/or rights consistent with this
|
| 169 |
+
License. However, in accepting such obligations, You may act only
|
| 170 |
+
on Your own behalf and on Your sole responsibility, not on behalf
|
| 171 |
+
of any other Contributor, and only if You agree to indemnify,
|
| 172 |
+
defend, and hold each Contributor harmless for any liability
|
| 173 |
+
incurred by, or claims asserted against, such Contributor by reason
|
| 174 |
+
of your accepting any such warranty or additional liability.
|
| 175 |
+
|
| 176 |
+
END OF TERMS AND CONDITIONS
|
| 177 |
+
|
| 178 |
+
APPENDIX: How to apply the Apache License to your work.
|
| 179 |
+
|
| 180 |
+
To apply the Apache License to your work, attach the following
|
| 181 |
+
boilerplate notice, with the fields enclosed by brackets "[]"
|
| 182 |
+
replaced with your own identifying information. (Don't include
|
| 183 |
+
the brackets!) The text should be enclosed in the appropriate
|
| 184 |
+
comment syntax for the file format. We also recommend that a
|
| 185 |
+
file or class name and description of purpose be included on the
|
| 186 |
+
same "printed page" as the copyright notice for easier
|
| 187 |
+
identification within third-party archives.
|
| 188 |
+
|
| 189 |
+
Copyright [yyyy] [name of copyright owner]
|
| 190 |
+
|
| 191 |
+
Licensed under the Apache License, Version 2.0 (the "License");
|
| 192 |
+
you may not use this file except in compliance with the License.
|
| 193 |
+
You may obtain a copy of the License at
|
| 194 |
+
|
| 195 |
+
http://www.apache.org/licenses/LICENSE-2.0
|
| 196 |
+
|
| 197 |
+
Unless required by applicable law or agreed to in writing, software
|
| 198 |
+
distributed under the License is distributed on an "AS IS" BASIS,
|
| 199 |
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 200 |
+
See the License for the specific language governing permissions and
|
| 201 |
+
limitations under the License.
|
README.md
ADDED
|
@@ -0,0 +1,215 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
<p align="center">
|
| 2 |
+
<img src="https://raw.githubusercontent.com/VectorSpaceLab/EditScore/refs/heads/main/assets/logo.png" width="65%">
|
| 3 |
+
</p>
|
| 4 |
+
|
| 5 |
+
<p align="center">
|
| 6 |
+
<a href="https://vectorspacelab.github.io/EditScore"><img src="https://img.shields.io/badge/Project%20Page-EditScore-yellow" alt="project page"></a>
|
| 7 |
+
<a href="https://arxiv.org/abs/2509.23909"><img src="https://img.shields.io/badge/arXiv%20paper-2509.23909-b31b1b.svg" alt="arxiv"></a>
|
| 8 |
+
<a href="https://huggingface.co/collections/EditScore/editscore-68d8e27ee676981221db3cfe"><img src="https://img.shields.io/badge/EditScore-🤗-yellow" alt="model"></a>
|
| 9 |
+
<a href="https://huggingface.co/datasets/EditScore/EditReward-Bench"><img src="https://img.shields.io/badge/EditReward--Bench-🤗-yellow" alt="dataset"></a>
|
| 10 |
+
<a href="https://huggingface.co/datasets/EditScore/EditScore-Reward-Data"><img src="https://img.shields.io/badge/EditScore--Reward--Data-🤗-yellow" alt="dataset"></a>
|
| 11 |
+
<a href="https://huggingface.co/datasets/EditScore/EditScore-RL-Data"><img src="https://img.shields.io/badge/EditScore--RL--Data-🤗-yellow" alt="dataset"></a>
|
| 12 |
+
</p>
|
| 13 |
+
|
| 14 |
+
<h4 align="center">
|
| 15 |
+
<p>
|
| 16 |
+
<a href=#-news>News</a> |
|
| 17 |
+
<a href=#-quick-start>Quick Start</a> |
|
| 18 |
+
<a href=#-benchmark-your-image-editing-reward-model usage>Benchmark Usage</a> |
|
| 19 |
+
<a href=#%EF%B8%8F-citing-us>Citation</a>
|
| 20 |
+
<p>
|
| 21 |
+
</h4>
|
| 22 |
+
|
| 23 |
+
**EditScore** is a series of state-of-the-art open-source reward models (7B–72B) designed to evaluate and enhance instruction-guided image editing.
|
| 24 |
+
## ✨ Highlights
|
| 25 |
+
- **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**.
|
| 26 |
+
- **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.
|
| 27 |
+
- **Simple and Easy-to-Use**: Get an accurate quality score for your image edits with just a few lines of code.
|
| 28 |
+
- **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**.
|
| 29 |
+
|
| 30 |
+
## 🔥 News
|
| 31 |
+
- **2026-02-01**: Our work has been accepted to ICLR 2026 🎉
|
| 32 |
+
- **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.
|
| 33 |
+
- **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).
|
| 34 |
+
- **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`.
|
| 35 |
+
- **2025-10-22**: **Introducing Our Reinforcement Learning Training Framework!**
|
| 36 |
+
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:
|
| 37 |
+
- **Ready-to-Use RL Dataset**: Includes the complete dataset used in the EditScore project, along with clear usage guidelines and preparation scripts.
|
| 38 |
+
- **An Easy-to-Use Reward Model**: Seamlessly integrate **EditScore** as a reward signal.
|
| 39 |
+
- **A Scalable Reward Server**: Built with native multi-node support for high-throughput training.
|
| 40 |
+
- **Flexible Training Code**: Supports distributed training, variable image resolutions and mixed tasks (t2i, edit, in-context generation) out-of-the-box.
|
| 41 |
+
Dive into our comprehensive guide on [RL Fine-Tuning](examples/OmniGen2-RL#application-2-reinforcement-fine-tuning) to get started.
|
| 42 |
+
|
| 43 |
+
- 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.
|
| 44 |
+
- 2025-10-15: **EditScore** is now available on PyPI — install it easily with `pip install editscore`.
|
| 45 |
+
- 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.
|
| 46 |
+
- 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).
|
| 47 |
+
- 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).
|
| 48 |
+
|
| 49 |
+
## 📖 Introduction
|
| 50 |
+
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.
|
| 51 |
+
|
| 52 |
+
To overcome this barrier, we provide a systematic, two-part solution:
|
| 53 |
+
|
| 54 |
+
- **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.
|
| 55 |
+
|
| 56 |
+
- **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.
|
| 57 |
+
|
| 58 |
+
<p align="center">
|
| 59 |
+
<img src="https://raw.githubusercontent.com/VectorSpaceLab/EditScore/refs/heads/main/assets/table_reward_model_results.png" width="95%">
|
| 60 |
+
<br>
|
| 61 |
+
<em>Benchmark results on EditReward-Bench.</em>
|
| 62 |
+
</p>
|
| 63 |
+
|
| 64 |
+
We demonstrate the practical utility of EditScore through two key applications:
|
| 65 |
+
|
| 66 |
+
- **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.
|
| 67 |
+
- **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.
|
| 68 |
+
|
| 69 |
+
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.
|
| 70 |
+
|
| 71 |
+
<p align="center">
|
| 72 |
+
<img src="https://raw.githubusercontent.com/VectorSpaceLab/EditScore/refs/heads/main/assets/figure_edit_results.png" width="95%">
|
| 73 |
+
<br>
|
| 74 |
+
<em>EditScore as a superior reward signal for image editing.</em>
|
| 75 |
+
</p>
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
## 📌 TODO
|
| 79 |
+
We are actively working on improving EditScore and expanding its capabilities. Here's what's next:
|
| 80 |
+
|
| 81 |
+
- [x] Qwen3-VL variants of EditScore.
|
| 82 |
+
- [x] Release training data for reward model and online RL.
|
| 83 |
+
- [x] Release RL training code applying EditScore to OmniGen2.
|
| 84 |
+
- [x] Provide Best-of-N inference scripts for OmniGen2, Flux-dev-Kontext, and Qwen-Image-Edit.
|
| 85 |
+
|
| 86 |
+
## 🚀 Quick Start
|
| 87 |
+
|
| 88 |
+
### 🛠️ Environment Setup
|
| 89 |
+
We offer two ways to install EditScore. Choose the one that best fits your needs.
|
| 90 |
+
**Method 1: Install from PyPI (Recommended for Users)**: If you want to use EditScore as a library in your own project.
|
| 91 |
+
**Method 2: Install from Source (For Developers)**: If you plan to contribute to the code, modify it, or run the examples in this repository
|
| 92 |
+
|
| 93 |
+
#### Prerequisites: Installing PyTorch
|
| 94 |
+
Both installation methods require PyTorch to be installed first, as its version is dependent on your system's CUDA setup.
|
| 95 |
+
```bash
|
| 96 |
+
# (Optional) Create a clean Python environment
|
| 97 |
+
conda create -n editscore python=3.12
|
| 98 |
+
conda activate editscore
|
| 99 |
+
|
| 100 |
+
# Choose the command that matches your CUDA version.
|
| 101 |
+
# This example is for CUDA 12.6.
|
| 102 |
+
pip install torch==2.7.1 torchvision --extra-index-url https://download.pytorch.org/whl/cu126
|
| 103 |
+
````
|
| 104 |
+
|
| 105 |
+
<details>
|
| 106 |
+
<summary>🌏 For users in Mainland China</summary>
|
| 107 |
+
```bash
|
| 108 |
+
# Install PyTorch from a domestic mirror
|
| 109 |
+
pip install torch==2.7.1 torchvision --index-url https://mirror.sjtu.edu.cn/pytorch-wheels/cu126
|
| 110 |
+
```
|
| 111 |
+
</details>
|
| 112 |
+
|
| 113 |
+
#### Method 1: Install from PyPI (Recommended for Users)
|
| 114 |
+
```bash
|
| 115 |
+
pip install -U editscore
|
| 116 |
+
```
|
| 117 |
+
|
| 118 |
+
#### Method 2: Install from Source (For Developers)
|
| 119 |
+
This method gives you a local, editable version of the project.
|
| 120 |
+
1. Clone the repository
|
| 121 |
+
```bash
|
| 122 |
+
git clone https://github.com/VectorSpaceLab/EditScore.git
|
| 123 |
+
cd EditScore
|
| 124 |
+
```
|
| 125 |
+
|
| 126 |
+
2. Install EditScore in editable mode
|
| 127 |
+
```bash
|
| 128 |
+
pip install -e .
|
| 129 |
+
```
|
| 130 |
+
|
| 131 |
+
#### ✅ (Recommended) Install Optional High-Performance Dependencies
|
| 132 |
+
For the best performance, especially during inference, we highly recommend installing vllm.
|
| 133 |
+
```bash
|
| 134 |
+
pip install -U vllm
|
| 135 |
+
```
|
| 136 |
+
|
| 137 |
+
---
|
| 138 |
+
|
| 139 |
+
### 🧪 Usage Example
|
| 140 |
+
Using EditScore is straightforward. The model will be automatically downloaded from the Hugging Face Hub on its first run.
|
| 141 |
+
```python
|
| 142 |
+
from PIL import Image
|
| 143 |
+
from editscore import EditScore
|
| 144 |
+
|
| 145 |
+
# Load the EditScore model. It will be downloaded automatically.
|
| 146 |
+
# Replace with the specific model version you want to use.
|
| 147 |
+
model_path = "Qwen/Qwen3-VL-4B-Instruct"
|
| 148 |
+
lora_path = "EditScore/EditScore-Qwen3-VL-4B-Instruct"
|
| 149 |
+
|
| 150 |
+
scorer = EditScore(
|
| 151 |
+
backbone="qwen3vl", # set to "qwen3vl_vllm" for faster inference
|
| 152 |
+
model_name_or_path=model_path,
|
| 153 |
+
lora_path=lora_path,
|
| 154 |
+
score_range=25,
|
| 155 |
+
num_pass=1, # Increase for better performance via self-ensembling
|
| 156 |
+
)
|
| 157 |
+
|
| 158 |
+
# Below is Qwen2.5-VL version
|
| 159 |
+
|
| 160 |
+
# model_path = "Qwen/Qwen2.5-VL-7B-Instruct"
|
| 161 |
+
# lora_path = "EditScore/EditScore-7B"
|
| 162 |
+
|
| 163 |
+
# scorer = EditScore(
|
| 164 |
+
# backbone="qwen25vl", # set to "qwen25vl_vllm" for faster inference
|
| 165 |
+
# model_name_or_path=model_path,
|
| 166 |
+
# lora_path=lora_path,
|
| 167 |
+
# score_range=25,
|
| 168 |
+
# num_pass=1, # Increase for better performance via self-ensembling
|
| 169 |
+
# )
|
| 170 |
+
|
| 171 |
+
input_image = Image.open("example_images/input.png")
|
| 172 |
+
output_image = Image.open("example_images/output.png")
|
| 173 |
+
instruction = "Adjust the background to a glass wall."
|
| 174 |
+
|
| 175 |
+
result = scorer.evaluate([input_image, output_image], instruction)
|
| 176 |
+
print(f"Edit Score: {result['overall']}")
|
| 177 |
+
# Expected output: A dictionary containing the final score and other details.
|
| 178 |
+
```
|
| 179 |
+
|
| 180 |
+
---
|
| 181 |
+
|
| 182 |
+
## 📊 Benchmark Your Image-Editing Reward Model
|
| 183 |
+
#### Install benchmark dependencies
|
| 184 |
+
To use example code for benchmark, run following
|
| 185 |
+
```bash
|
| 186 |
+
pip install -r requirements.txt
|
| 187 |
+
```
|
| 188 |
+
|
| 189 |
+
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.
|
| 190 |
+
```bash
|
| 191 |
+
# This script will evaluate the default EditScore model on the benchmark
|
| 192 |
+
bash evaluate.sh
|
| 193 |
+
|
| 194 |
+
# Or speed up inference with VLLM
|
| 195 |
+
bash evaluate_vllm.sh
|
| 196 |
+
```
|
| 197 |
+
|
| 198 |
+
## Apply EditScore to Image Editing
|
| 199 |
+
We offer two example use cases for your exploration:
|
| 200 |
+
- **Best-of-N selection**: Use EditScore to automatically pick the most preferred image among multiple candidates.
|
| 201 |
+
- **Reinforcement fine-tuning**: Use EditScore as a reward model to guide RL-based optimization.
|
| 202 |
+
|
| 203 |
+
For detailed instructions and examples, please refer to the [documentation](examples/OmniGen2-RL/README.md).
|
| 204 |
+
|
| 205 |
+
## ❤️ Citing Us
|
| 206 |
+
If you find this repository or our work useful, please consider giving a star ⭐ and citation 🦖, which would be greatly appreciated:
|
| 207 |
+
|
| 208 |
+
```bibtex
|
| 209 |
+
@article{luo2025editscore,
|
| 210 |
+
title={EditScore: Unlocking Online RL for Image Editing via High-Fidelity Reward Modeling},
|
| 211 |
+
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},
|
| 212 |
+
journal={arXiv preprint arXiv:2509.23909},
|
| 213 |
+
year={2025}
|
| 214 |
+
}
|
| 215 |
+
```
|
assets/figure_edit_results.png
ADDED
|
Git LFS Details
|
assets/logo.png
ADDED
|
Git LFS Details
|
assets/table_editscore_qwen3_vl.png
ADDED
|
Git LFS Details
|
assets/table_reward_model_results.png
ADDED
|
Git LFS Details
|
calculate_statistics.py
ADDED
|
@@ -0,0 +1,124 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import glob
|
| 3 |
+
import json
|
| 4 |
+
import numpy as np
|
| 5 |
+
|
| 6 |
+
import argparse
|
| 7 |
+
|
| 8 |
+
PROMPT_FOLLOWING = "prompt_following"
|
| 9 |
+
CONSISTENCY = "consistency"
|
| 10 |
+
OVERALL = "overall"
|
| 11 |
+
SCORE_CATEGORIES = [PROMPT_FOLLOWING, CONSISTENCY, OVERALL]
|
| 12 |
+
|
| 13 |
+
def parse_args():
|
| 14 |
+
parser = argparse.ArgumentParser()
|
| 15 |
+
parser.add_argument("--result_dir", type=str, required=True)
|
| 16 |
+
parser.add_argument("--backbone", type=str, default="qwen25vl", choices=["qwen25vl", "openai", "internvl3_5"])
|
| 17 |
+
return parser.parse_args()
|
| 18 |
+
|
| 19 |
+
def main(args):
|
| 20 |
+
task_types = sorted(os.listdir(args.result_dir))
|
| 21 |
+
|
| 22 |
+
print(task_types)
|
| 23 |
+
|
| 24 |
+
prompt_following_results = dict()
|
| 25 |
+
consistency_results = dict()
|
| 26 |
+
overall_results = dict()
|
| 27 |
+
|
| 28 |
+
all_prompt_following_scores = []
|
| 29 |
+
all_consistency_scores = []
|
| 30 |
+
all_overall_scores = []
|
| 31 |
+
|
| 32 |
+
for task_type in task_types:
|
| 33 |
+
task_type_dir = os.path.join(args.result_dir, task_type)
|
| 34 |
+
prompt_following_json_file = os.path.join(task_type_dir, f"{PROMPT_FOLLOWING}.jsonl")
|
| 35 |
+
consistency_json_file = os.path.join(task_type_dir, f"{CONSISTENCY}.jsonl")
|
| 36 |
+
overall_json_file = os.path.join(task_type_dir, f"{OVERALL}.jsonl")
|
| 37 |
+
|
| 38 |
+
total_num = 0
|
| 39 |
+
correct_num = 0
|
| 40 |
+
with open(prompt_following_json_file, 'r') as f:
|
| 41 |
+
for line in f:
|
| 42 |
+
json_line = json.loads(line)
|
| 43 |
+
if json_line['score'][0] > json_line['score'][1]:
|
| 44 |
+
correct_num += 1
|
| 45 |
+
total_num += 1
|
| 46 |
+
all_prompt_following_scores.append(json_line['score'][0])
|
| 47 |
+
all_prompt_following_scores.append(json_line['score'][1])
|
| 48 |
+
prompt_following_results[task_type] = correct_num / total_num
|
| 49 |
+
|
| 50 |
+
total_num = 0
|
| 51 |
+
correct_num = 0
|
| 52 |
+
with open(consistency_json_file, 'r') as f:
|
| 53 |
+
for line in f:
|
| 54 |
+
json_line = json.loads(line)
|
| 55 |
+
if json_line['score'][0] > json_line['score'][1]:
|
| 56 |
+
correct_num += 1
|
| 57 |
+
total_num += 1
|
| 58 |
+
all_consistency_scores.append(json_line['score'][0])
|
| 59 |
+
all_consistency_scores.append(json_line['score'][1])
|
| 60 |
+
consistency_results[task_type] = correct_num / total_num
|
| 61 |
+
|
| 62 |
+
total_num = 0
|
| 63 |
+
correct_num = 0
|
| 64 |
+
with open(overall_json_file, 'r') as f:
|
| 65 |
+
for line in f:
|
| 66 |
+
json_line = json.loads(line)
|
| 67 |
+
if json_line['score'][0] > json_line['score'][1]:
|
| 68 |
+
correct_num += 1
|
| 69 |
+
total_num += 1
|
| 70 |
+
all_overall_scores.append(json_line['score'][0])
|
| 71 |
+
all_overall_scores.append(json_line['score'][1])
|
| 72 |
+
overall_results[task_type] = correct_num / total_num
|
| 73 |
+
|
| 74 |
+
prompt_following_results['average'] = sum(prompt_following_results.values()) / len(prompt_following_results)
|
| 75 |
+
consistency_results['average'] = sum(consistency_results.values()) / len(consistency_results)
|
| 76 |
+
overall_results['average'] = sum(overall_results.values()) / len(overall_results)
|
| 77 |
+
|
| 78 |
+
print(overall_results.keys())
|
| 79 |
+
|
| 80 |
+
task_types = [
|
| 81 |
+
'background_change', 'color_alter', 'style_change', 'subject-add', 'subject-remove', 'subject-replace', 'material_alter',
|
| 82 |
+
'motion_change', 'ps_human', 'text_change', 'tone_transfer', 'extract', 'compose', 'average'
|
| 83 |
+
]
|
| 84 |
+
|
| 85 |
+
print(" & ".join(task_types))
|
| 86 |
+
print("Prompt Following: " + " & ".join([f"{prompt_following_results[task_type]:.3f}" for task_type in task_types]))
|
| 87 |
+
print("Consistency: " + " & ".join([f"{consistency_results[task_type]:.3f}" for task_type in task_types]))
|
| 88 |
+
print("Overall: " + " & ".join([f"{overall_results[task_type]:.3f}" for task_type in task_types]))
|
| 89 |
+
|
| 90 |
+
groups = {
|
| 91 |
+
'object': ['subject-add', 'subject-remove', 'subject-replace'],
|
| 92 |
+
'appearance': ['color_alter', 'material_alter', 'style_change', 'tone_transfer'],
|
| 93 |
+
'scene': ['background_change', 'extract'],
|
| 94 |
+
'advanced': ['ps_human', 'text_change', 'motion_change', 'compose'],
|
| 95 |
+
}
|
| 96 |
+
|
| 97 |
+
print("--------------------------------")
|
| 98 |
+
print("--------------------------------")
|
| 99 |
+
|
| 100 |
+
for group_name, group_task_types in groups.items():
|
| 101 |
+
print(group_name + ":")
|
| 102 |
+
print("Prompt Following & Consistency & Overall")
|
| 103 |
+
prompt_following_mean = np.mean([prompt_following_results[task_type] for task_type in group_task_types])
|
| 104 |
+
consistency_mean = np.mean([consistency_results[task_type] for task_type in group_task_types])
|
| 105 |
+
overall_mean = np.mean([overall_results[task_type] for task_type in group_task_types])
|
| 106 |
+
print(f"{prompt_following_mean:.3f} & {consistency_mean:.3f} & {overall_mean:.3f}")
|
| 107 |
+
|
| 108 |
+
print("Average:")
|
| 109 |
+
print("Prompt Following & Consistency & Overall")
|
| 110 |
+
print(f"{prompt_following_results['average']:.3f} & {consistency_results['average']:.3f} & {overall_results['average']:.3f}")
|
| 111 |
+
|
| 112 |
+
print("Prompt Following Scores:")
|
| 113 |
+
print("Min & Max & Mean & Std")
|
| 114 |
+
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}")
|
| 115 |
+
print("Consistency Scores:")
|
| 116 |
+
print("Min & Max & Mean & Std")
|
| 117 |
+
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}")
|
| 118 |
+
print("Overall Scores:")
|
| 119 |
+
print("Min & Max & Mean & Std")
|
| 120 |
+
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}")
|
| 121 |
+
|
| 122 |
+
if __name__ == "__main__":
|
| 123 |
+
args = parse_args()
|
| 124 |
+
main(args)
|
editscore/__init__.py
ADDED
|
@@ -0,0 +1,217 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import sys
|
| 2 |
+
sys.path.insert(0, 'editscore')
|
| 3 |
+
|
| 4 |
+
from typing import Optional
|
| 5 |
+
from .utils import (
|
| 6 |
+
mllm_output_to_dict
|
| 7 |
+
)
|
| 8 |
+
import math
|
| 9 |
+
from . import vie_prompts
|
| 10 |
+
import numpy as np
|
| 11 |
+
from .json_parser import parse_vlm_output_to_dict
|
| 12 |
+
|
| 13 |
+
class EditScore:
|
| 14 |
+
def __init__(
|
| 15 |
+
self,
|
| 16 |
+
backbone="gpt-4.1",
|
| 17 |
+
openai_url="https://api.openai.com/v1/chat/completions",
|
| 18 |
+
key=None,
|
| 19 |
+
model_name_or_path="",
|
| 20 |
+
score_range: int=25,
|
| 21 |
+
temperature: float=0.7,
|
| 22 |
+
tensor_parallel_size: int=1,
|
| 23 |
+
max_model_len: int=1536,
|
| 24 |
+
max_num_batched_tokens: int=1536,
|
| 25 |
+
max_num_seqs: int=32,
|
| 26 |
+
num_pass: int=1,
|
| 27 |
+
reduction: str="average_last",
|
| 28 |
+
seed: int=42,
|
| 29 |
+
lora_path: Optional[str]=None,
|
| 30 |
+
cache_dir: Optional[str]=None,
|
| 31 |
+
) -> None:
|
| 32 |
+
self.backbone = backbone
|
| 33 |
+
self.score_range = score_range
|
| 34 |
+
self.reduction = reduction
|
| 35 |
+
self.seed = seed
|
| 36 |
+
self.num_pass = num_pass
|
| 37 |
+
|
| 38 |
+
if self.backbone == 'openai':
|
| 39 |
+
from .mllm_tools.openai import GPT4o
|
| 40 |
+
self.model = GPT4o(key, model_name=model_name_or_path, url=openai_url)
|
| 41 |
+
elif self.backbone == "qwen25vl":
|
| 42 |
+
from .mllm_tools.qwen25vl import Qwen25VL
|
| 43 |
+
self.model = Qwen25VL(
|
| 44 |
+
vlm_model=model_name_or_path,
|
| 45 |
+
temperature=temperature,
|
| 46 |
+
seed=seed,
|
| 47 |
+
lora_path=lora_path,
|
| 48 |
+
)
|
| 49 |
+
elif self.backbone == "qwen25vl_vllm":
|
| 50 |
+
from .mllm_tools.qwen25vl_vllm import Qwen25VL
|
| 51 |
+
self.model = Qwen25VL(
|
| 52 |
+
vlm_model=model_name_or_path,
|
| 53 |
+
tensor_parallel_size=tensor_parallel_size,
|
| 54 |
+
max_model_len=max_model_len,
|
| 55 |
+
max_num_seqs=max_num_seqs,
|
| 56 |
+
max_num_batched_tokens=max_num_batched_tokens,
|
| 57 |
+
temperature=temperature,
|
| 58 |
+
seed=seed,
|
| 59 |
+
lora_path=lora_path,
|
| 60 |
+
cache_dir=cache_dir,
|
| 61 |
+
)
|
| 62 |
+
elif self.backbone == "qwen3vl":
|
| 63 |
+
from .mllm_tools.qwen3vl import Qwen3VL
|
| 64 |
+
self.model = Qwen3VL(
|
| 65 |
+
vlm_model=model_name_or_path,
|
| 66 |
+
temperature=temperature,
|
| 67 |
+
seed=seed,
|
| 68 |
+
lora_path=lora_path,
|
| 69 |
+
)
|
| 70 |
+
elif self.backbone == "qwen3vl_vllm":
|
| 71 |
+
from .mllm_tools.qwen3vl_vllm import Qwen3VL
|
| 72 |
+
self.model = Qwen3VL(
|
| 73 |
+
vlm_model=model_name_or_path,
|
| 74 |
+
tensor_parallel_size=tensor_parallel_size,
|
| 75 |
+
max_model_len=max_model_len,
|
| 76 |
+
max_num_seqs=max_num_seqs,
|
| 77 |
+
max_num_batched_tokens=max_num_batched_tokens,
|
| 78 |
+
temperature=temperature,
|
| 79 |
+
seed=seed,
|
| 80 |
+
lora_path=lora_path,
|
| 81 |
+
cache_dir=cache_dir,
|
| 82 |
+
)
|
| 83 |
+
elif self.backbone == "internvl3_5":
|
| 84 |
+
from .mllm_tools.internvl35_lmdeploy import InternVL35
|
| 85 |
+
self.model = InternVL35(model=model_name_or_path, tensor_parallel_size=tensor_parallel_size)
|
| 86 |
+
|
| 87 |
+
self.context = vie_prompts._context_no_delimit_reasoning_first
|
| 88 |
+
|
| 89 |
+
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))])
|
| 90 |
+
self.PQ_prompt = "\n".join([self.context, vie_prompts._prompts_0shot_rule_PQ.replace('10', str(self.score_range))])
|
| 91 |
+
|
| 92 |
+
def evaluate(self, image_prompts, text_prompt):
|
| 93 |
+
if not isinstance(image_prompts, list):
|
| 94 |
+
image_prompts = [image_prompts]
|
| 95 |
+
|
| 96 |
+
if self.backbone in ['openai']:
|
| 97 |
+
self.model.use_encode = False if isinstance(image_prompts[0], str) else True
|
| 98 |
+
|
| 99 |
+
_SC_prompt = self.SC_prompt.replace("<instruction>", text_prompt)
|
| 100 |
+
|
| 101 |
+
SC_prompt_final = self.model.prepare_input(image_prompts, _SC_prompt)
|
| 102 |
+
PQ_prompt_final = self.model.prepare_input(image_prompts[-1], self.PQ_prompt) # assume the last image is the edited image
|
| 103 |
+
|
| 104 |
+
outputs_multi_pass = []
|
| 105 |
+
|
| 106 |
+
for i in range(self.num_pass):
|
| 107 |
+
SC_dict = False
|
| 108 |
+
PQ_dict = False
|
| 109 |
+
tries = 0
|
| 110 |
+
max_tries = 2
|
| 111 |
+
while SC_dict is False or PQ_dict is False:
|
| 112 |
+
tries += 1
|
| 113 |
+
give_up_parsing = True if tries > max_tries else False
|
| 114 |
+
|
| 115 |
+
result_SC = self.model.inference(SC_prompt_final, seed=self.seed + i)
|
| 116 |
+
result_PQ = self.model.inference(PQ_prompt_final, seed=self.seed + i)
|
| 117 |
+
|
| 118 |
+
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."]:
|
| 119 |
+
give_up_parsing = True
|
| 120 |
+
|
| 121 |
+
SC_dict = mllm_output_to_dict(result_SC, give_up_parsing=give_up_parsing, text_prompt=text_prompt, score_range=self.score_range)
|
| 122 |
+
PQ_dict = mllm_output_to_dict(result_PQ, give_up_parsing=give_up_parsing, text_prompt=text_prompt, score_range=self.score_range)
|
| 123 |
+
|
| 124 |
+
if SC_dict == "rate_limit_exceeded" or PQ_dict == "rate_limit_exceeded":
|
| 125 |
+
print("rate_limit_exceeded")
|
| 126 |
+
raise ValueError("rate_limit_exceeded")
|
| 127 |
+
|
| 128 |
+
try:
|
| 129 |
+
SC_score = min(SC_dict['score']) / (self.score_range / 10)
|
| 130 |
+
PQ_score = min(PQ_dict['score']) / (self.score_range / 10)
|
| 131 |
+
O_score = math.sqrt(SC_score * PQ_score)
|
| 132 |
+
except Exception as e:
|
| 133 |
+
print(f"{e=} {SC_dict['score']=} {PQ_dict['score']=}")
|
| 134 |
+
raise e
|
| 135 |
+
|
| 136 |
+
try:
|
| 137 |
+
outputs_multi_pass.append({
|
| 138 |
+
'prompt_following': SC_dict['score'][0] / (self.score_range / 10),
|
| 139 |
+
'consistency': SC_dict['score'][1] / (self.score_range / 10),
|
| 140 |
+
'perceptual_quality': PQ_score,
|
| 141 |
+
'overall': O_score,
|
| 142 |
+
})
|
| 143 |
+
except Exception as e:
|
| 144 |
+
print(f"{e=} {SC_dict['score']=} {PQ_dict['score']=}")
|
| 145 |
+
raise e
|
| 146 |
+
|
| 147 |
+
output = {
|
| 148 |
+
"prompt_following": np.mean([output_per_pass["prompt_following"] for output_per_pass in outputs_multi_pass]),
|
| 149 |
+
"consistency": np.mean([output_per_pass["consistency"] for output_per_pass in outputs_multi_pass]),
|
| 150 |
+
"perceptual_quality": np.mean([output_per_pass["perceptual_quality"] for output_per_pass in outputs_multi_pass]),
|
| 151 |
+
"overall": np.mean([output_per_pass["overall"] for output_per_pass in outputs_multi_pass]),
|
| 152 |
+
"SC_reasoning": SC_dict["reasoning"],
|
| 153 |
+
"PQ_reasoning": PQ_dict["reasoning"],
|
| 154 |
+
}
|
| 155 |
+
if self.reduction == "average_first":
|
| 156 |
+
output["overall"] = math.sqrt(output["prompt_following"] * output["perceptual_quality"])
|
| 157 |
+
return output
|
| 158 |
+
|
| 159 |
+
|
| 160 |
+
def batch_evaluate(self, image_prompts, text_prompt):
|
| 161 |
+
SC_prompt = [self.SC_prompt.replace("<instruction>", _text_prompt) for _text_prompt in text_prompt]
|
| 162 |
+
|
| 163 |
+
SC_prompt = [self.model.prepare_input(image_prompt, _SC_prompt) for image_prompt, _SC_prompt in zip(image_prompts, SC_prompt)]
|
| 164 |
+
PQ_prompt = [self.model.prepare_input(image_prompt, self.PQ_prompt) for image_prompt in image_prompts]
|
| 165 |
+
|
| 166 |
+
outputs_multi_pass = [[] for _ in range(len(image_prompts))]
|
| 167 |
+
for i in range(self.num_pass):
|
| 168 |
+
results = self.model.batch_inference(SC_prompt + PQ_prompt, seed=self.seed + i)
|
| 169 |
+
|
| 170 |
+
SC_evaluations = [parse_vlm_output_to_dict(results[i]) for i in range(len(results) // 2)]
|
| 171 |
+
PQ_evaluations = [parse_vlm_output_to_dict(results[i]) for i in range(len(results) // 2, len(results))]
|
| 172 |
+
|
| 173 |
+
for idx, (SC_evaluation, PQ_evaluation) in enumerate(zip(SC_evaluations, PQ_evaluations)):
|
| 174 |
+
SC_scores = SC_evaluation["score"]
|
| 175 |
+
PQ_scores = PQ_evaluation["score"]
|
| 176 |
+
|
| 177 |
+
if len(SC_scores) == 0:
|
| 178 |
+
SC_scores = [self.score_range / 2]
|
| 179 |
+
if len(PQ_scores) == 0:
|
| 180 |
+
PQ_scores = [self.score_range / 2]
|
| 181 |
+
|
| 182 |
+
SC_score = min(SC_scores) / (self.score_range / 10)
|
| 183 |
+
PQ_score = min(PQ_scores) / (self.score_range / 10)
|
| 184 |
+
if SC_score < 0 or SC_score > 10:
|
| 185 |
+
SC_score = self.score_range / 2
|
| 186 |
+
if PQ_score < 0 or PQ_score > 10:
|
| 187 |
+
PQ_score = self.score_range / 2
|
| 188 |
+
O_score = math.sqrt(SC_score * PQ_score)
|
| 189 |
+
|
| 190 |
+
outputs_multi_pass[idx].append(
|
| 191 |
+
{
|
| 192 |
+
"SC_score": SC_score,
|
| 193 |
+
"PQ_score": PQ_score,
|
| 194 |
+
"O_score": O_score,
|
| 195 |
+
"SC_score_reasoning": SC_evaluation["reasoning"],
|
| 196 |
+
"PQ_score_reasoning": PQ_evaluation["reasoning"],
|
| 197 |
+
"SC_raw_output": results[idx],
|
| 198 |
+
"PQ_raw_output": results[len(results) // 2 + idx],
|
| 199 |
+
}
|
| 200 |
+
)
|
| 201 |
+
|
| 202 |
+
outputs = []
|
| 203 |
+
for idx, outputs_per_prompt in enumerate(outputs_multi_pass):
|
| 204 |
+
outputs.append(
|
| 205 |
+
{
|
| 206 |
+
"SC_score": np.mean([output_per_pass["SC_score"] for output_per_pass in outputs_per_prompt]),
|
| 207 |
+
"PQ_score": np.mean([output_per_pass["PQ_score"] for output_per_pass in outputs_per_prompt]),
|
| 208 |
+
"O_score": np.mean([output_per_pass["O_score"] for output_per_pass in outputs_per_prompt]),
|
| 209 |
+
"SC_score_reasoning": outputs_per_prompt[0]["SC_score_reasoning"],
|
| 210 |
+
"PQ_score_reasoning": outputs_per_prompt[0]["PQ_score_reasoning"],
|
| 211 |
+
"SC_raw_output": outputs_per_prompt[0]["SC_raw_output"],
|
| 212 |
+
"PQ_raw_output": outputs_per_prompt[0]["PQ_raw_output"],
|
| 213 |
+
}
|
| 214 |
+
)
|
| 215 |
+
if self.reduction == "average_first":
|
| 216 |
+
outputs[-1]["O_score"] = math.sqrt(outputs[-1]["SC_score"] * outputs[-1]["PQ_score"])
|
| 217 |
+
return outputs
|
editscore/json_parser.py
ADDED
|
@@ -0,0 +1,173 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import json
|
| 2 |
+
import re
|
| 3 |
+
from typing import Dict, Any, List, Optional
|
| 4 |
+
|
| 5 |
+
# ==============================================================================
|
| 6 |
+
# HELPER FUNCTIONS (Based on your provided robust fixers)
|
| 7 |
+
# ==============================================================================
|
| 8 |
+
# For clarity, these are named as internal functions (prefixed with an underscore).
|
| 9 |
+
|
| 10 |
+
def _fix_json_quotes(s: str) -> str:
|
| 11 |
+
"""First-stage repair: handle incorrect quotes and basic structure."""
|
| 12 |
+
# Replace Python-style booleans/None with JSON standard
|
| 13 |
+
s = re.sub(r'\bTrue\b', 'true', s)
|
| 14 |
+
s = re.sub(r'\bFalse\b', 'false', s)
|
| 15 |
+
s = re.sub(r'\bNone\b', 'null', s)
|
| 16 |
+
|
| 17 |
+
# Attempt to replace single quotes with double quotes (a common VLM error)
|
| 18 |
+
# This is a high-risk operation that might break the reasoning content,
|
| 19 |
+
# but it's worth trying early on.
|
| 20 |
+
try:
|
| 21 |
+
temp_s = s.replace("'", '"')
|
| 22 |
+
json.loads(temp_s)
|
| 23 |
+
return temp_s
|
| 24 |
+
except json.JSONDecodeError:
|
| 25 |
+
# If it's still invalid after replacement, return the original string for the next repair step.
|
| 26 |
+
pass
|
| 27 |
+
|
| 28 |
+
# Add double quotes to keys (e.g., {reasoning: ...} -> {"reasoning": ...})
|
| 29 |
+
s = re.sub(r'([\{\s,])(\w+)\s*:', r'\1"\2":', s)
|
| 30 |
+
return s
|
| 31 |
+
|
| 32 |
+
def _repair_reasoning_field_robust(json_str: str) -> str:
|
| 33 |
+
"""Second-stage repair: specifically fix unescaped double quotes inside the 'reasoning' field."""
|
| 34 |
+
pattern = re.compile(
|
| 35 |
+
r'("reasoning"\s*:\s*")' # --- Group 1: "reasoning" : "
|
| 36 |
+
r'(.*?)' # --- Group 2: The content (non-greedy)
|
| 37 |
+
r'(?="\s*[,}])', # --- Lookahead: find " followed by , or }
|
| 38 |
+
re.DOTALL
|
| 39 |
+
)
|
| 40 |
+
|
| 41 |
+
def replacer(match):
|
| 42 |
+
prefix = match.group(1)
|
| 43 |
+
content = match.group(2)
|
| 44 |
+
# In the content, replace all unescaped " with \"
|
| 45 |
+
fixed_content = content.replace('"', '\\"')
|
| 46 |
+
return prefix + fixed_content
|
| 47 |
+
|
| 48 |
+
return pattern.sub(replacer, json_str)
|
| 49 |
+
|
| 50 |
+
def _fallback_extract_and_rebuild(input_str: str) -> str:
|
| 51 |
+
"""Final fallback strategy: abandon repair, directly extract information, and rebuild a valid JSON."""
|
| 52 |
+
# 1. Extract reasoning
|
| 53 |
+
# Find all content between "reasoning": and ,"score":
|
| 54 |
+
reasoning_text = ""
|
| 55 |
+
reason_match = re.search(r'["\']reasoning["\']\s*:\s*["\']?(.*?)["\']?\s*,\s*["\']score["\']', input_str, re.DOTALL | re.IGNORECASE)
|
| 56 |
+
if reason_match:
|
| 57 |
+
reasoning_text = reason_match.group(1).strip()
|
| 58 |
+
# Clean up any potentially remaining escape characters
|
| 59 |
+
reasoning_text = reasoning_text.replace('\\"', '"')
|
| 60 |
+
else:
|
| 61 |
+
# If not found, assume all text besides the score part is the reasoning.
|
| 62 |
+
# First, remove the score part.
|
| 63 |
+
score_part_match = re.search(r'["\']score["\']\s*:.*', input_str, re.IGNORECASE)
|
| 64 |
+
if score_part_match:
|
| 65 |
+
reasoning_text = input_str[:score_part_match.start()].strip()
|
| 66 |
+
else:
|
| 67 |
+
# If even 'score' cannot be found, assume the entire string is the reasoning.
|
| 68 |
+
reasoning_text = input_str
|
| 69 |
+
|
| 70 |
+
# 2. Extract scores
|
| 71 |
+
scores = []
|
| 72 |
+
# Prioritize searching after the 'score' keyword.
|
| 73 |
+
score_match = re.search(r'["\']score["\']\s*:\s*(.*)', input_str, re.DOTALL | re.IGNORECASE)
|
| 74 |
+
search_area = score_match.group(1) if score_match else input_str
|
| 75 |
+
|
| 76 |
+
# Find all integers or floats.
|
| 77 |
+
numbers = re.findall(r'[-+]?\d*\.?\d+', search_area)
|
| 78 |
+
if numbers:
|
| 79 |
+
scores = [float(num) for num in numbers]
|
| 80 |
+
|
| 81 |
+
# 3. Rebuild into a standard dictionary and return a JSON string.
|
| 82 |
+
rebuilt_data = {
|
| 83 |
+
"reasoning": reasoning_text,
|
| 84 |
+
"score": scores
|
| 85 |
+
}
|
| 86 |
+
return json.dumps(rebuilt_data, ensure_ascii=False)
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
def _format_and_validate_dict(data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
| 90 |
+
"""Validate and format the parsed dictionary to ensure it meets the final output standard."""
|
| 91 |
+
if not isinstance(data, dict):
|
| 92 |
+
return None
|
| 93 |
+
|
| 94 |
+
# Extract reasoning, tolerating case and spelling variations.
|
| 95 |
+
reasoning = ""
|
| 96 |
+
for key in ["reasoning", "reason", "rationale"]:
|
| 97 |
+
if key in data and isinstance(data[key], str):
|
| 98 |
+
reasoning = data[key]
|
| 99 |
+
break
|
| 100 |
+
|
| 101 |
+
# Extract score, and ensure it is a list of floats.
|
| 102 |
+
scores = []
|
| 103 |
+
if 'score' in data:
|
| 104 |
+
score_val = data['score']
|
| 105 |
+
if isinstance(score_val, list):
|
| 106 |
+
scores = [float(s) for s in score_val if isinstance(s, (int, float, str))]
|
| 107 |
+
elif isinstance(score_val, (int, float)):
|
| 108 |
+
scores = [float(score_val)]
|
| 109 |
+
|
| 110 |
+
# If any field was found, consider it a success.
|
| 111 |
+
if reasoning or scores:
|
| 112 |
+
return {"score": scores, "reasoning": reasoning}
|
| 113 |
+
|
| 114 |
+
return None
|
| 115 |
+
|
| 116 |
+
# ==============================================================================
|
| 117 |
+
# MAIN PARSING FUNCTION
|
| 118 |
+
# ==============================================================================
|
| 119 |
+
|
| 120 |
+
def parse_vlm_output_to_dict(input_string: str) -> Dict[str, Any]:
|
| 121 |
+
"""
|
| 122 |
+
A highly robust function to parse a VLM's output string into a dictionary
|
| 123 |
+
containing 'score' and 'reasoning'.
|
| 124 |
+
|
| 125 |
+
It uses a multi-stage repair pipeline, progressively degrading from standard
|
| 126 |
+
JSON parsing to a final information extraction fallback.
|
| 127 |
+
"""
|
| 128 |
+
# --- 0. Preprocessing ---
|
| 129 |
+
if not input_string or not input_string.strip():
|
| 130 |
+
return {"score": [], "reasoning": "Input was empty."}
|
| 131 |
+
|
| 132 |
+
# Find the substring enclosed by `{}`, which is often the core of the VLM output.
|
| 133 |
+
json_match = re.search(r'\{.*\}', input_string, re.DOTALL)
|
| 134 |
+
target_str = json_match.group(0) if json_match else input_string.strip()
|
| 135 |
+
|
| 136 |
+
# --- Repair Pipeline ---
|
| 137 |
+
# Apply fixers in order, attempting to parse after each one.
|
| 138 |
+
|
| 139 |
+
fixer_pipeline = [
|
| 140 |
+
lambda s: s, # 1. Try the original string.
|
| 141 |
+
_fix_json_quotes, # 2. Fix basic quotes and keywords.
|
| 142 |
+
_repair_reasoning_field_robust, # 3. Fix internal quotes in the reasoning field.
|
| 143 |
+
]
|
| 144 |
+
|
| 145 |
+
for fixer in fixer_pipeline:
|
| 146 |
+
try:
|
| 147 |
+
fixed_str = fixer(target_str)
|
| 148 |
+
data = json.loads(fixed_str)
|
| 149 |
+
validated_data = _format_and_validate_dict(data)
|
| 150 |
+
if validated_data is not None:
|
| 151 |
+
return validated_data
|
| 152 |
+
except (json.JSONDecodeError, TypeError):
|
| 153 |
+
# If it fails, continue to the next fixer.
|
| 154 |
+
continue
|
| 155 |
+
|
| 156 |
+
# --- Final Fallback Strategy ---
|
| 157 |
+
# If all repair and parsing attempts fail, activate the information extraction mode.
|
| 158 |
+
try:
|
| 159 |
+
fallback_str = _fallback_extract_and_rebuild(target_str)
|
| 160 |
+
# This function guarantees a valid JSON string, so we can load it directly.
|
| 161 |
+
data = json.loads(fallback_str)
|
| 162 |
+
# Still run it through the validator to standardize the format.
|
| 163 |
+
validated_data = _format_and_validate_dict(data)
|
| 164 |
+
if validated_data:
|
| 165 |
+
return validated_data
|
| 166 |
+
except Exception:
|
| 167 |
+
# If even the final fallback strategy fails, return an error message.
|
| 168 |
+
pass
|
| 169 |
+
|
| 170 |
+
return {
|
| 171 |
+
"score": [],
|
| 172 |
+
"reasoning": f"Failed to parse after all strategies. Original output: '{input_string}'"
|
| 173 |
+
}
|
editscore/mllm_tools/__init__.py
ADDED
|
File without changes
|
editscore/mllm_tools/internvl35_lmdeploy.py
ADDED
|
@@ -0,0 +1,73 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import List
|
| 2 |
+
|
| 3 |
+
import random
|
| 4 |
+
# import magic
|
| 5 |
+
# import megfile
|
| 6 |
+
|
| 7 |
+
import numpy as np
|
| 8 |
+
import torch
|
| 9 |
+
|
| 10 |
+
from lmdeploy import pipeline, PytorchEngineConfig
|
| 11 |
+
from lmdeploy.vl import load_image
|
| 12 |
+
from lmdeploy.vl.constants import IMAGE_TOKEN
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def set_seed(seed: int):
|
| 16 |
+
"""
|
| 17 |
+
Args:
|
| 18 |
+
Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`.
|
| 19 |
+
seed (`int`): The seed to set.
|
| 20 |
+
"""
|
| 21 |
+
random.seed(seed)
|
| 22 |
+
np.random.seed(seed)
|
| 23 |
+
torch.manual_seed(seed)
|
| 24 |
+
torch.cuda.manual_seed_all(seed)
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def apply_chat_template(prompt, num_images: int = 2):
|
| 28 |
+
"""
|
| 29 |
+
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
|
| 30 |
+
"""
|
| 31 |
+
template = "\n".join([f"Image-{i}: {IMAGE_TOKEN}" for i in range(1, num_images + 1)])
|
| 32 |
+
template += f"\n{prompt}"
|
| 33 |
+
return template
|
| 34 |
+
|
| 35 |
+
class InternVL35():
|
| 36 |
+
def __init__(self, model, max_model_len: int = 16384, tensor_parallel_size=1, max_num_seqs=32) -> None:
|
| 37 |
+
# attn_implementation = "flash_attention_2" if is_flash_attn_2_available() else None
|
| 38 |
+
self.model = pipeline(model, backend_config=PytorchEngineConfig(session_len=max_model_len, tp=tensor_parallel_size))
|
| 39 |
+
|
| 40 |
+
def prepare_input(self, images: List = [], text_prompt: str = ""):
|
| 41 |
+
if not isinstance(images, list):
|
| 42 |
+
images = [images]
|
| 43 |
+
messages = (apply_chat_template(text_prompt, num_images=len(images)), images)
|
| 44 |
+
return messages
|
| 45 |
+
|
| 46 |
+
def inference(self, messages):
|
| 47 |
+
set_seed(42)
|
| 48 |
+
# Prepare the inputs
|
| 49 |
+
|
| 50 |
+
response = self.model(messages)
|
| 51 |
+
print(f"{response.text=}", flush=True)
|
| 52 |
+
return response.text
|
| 53 |
+
|
| 54 |
+
if __name__ == "__main__":
|
| 55 |
+
model = InternVL35(
|
| 56 |
+
vlm_model="OpenGVLab/InternVL3_5-8B",
|
| 57 |
+
max_model_len=16384,
|
| 58 |
+
tensor_parallel_size=1,
|
| 59 |
+
max_num_seqs=32
|
| 60 |
+
)
|
| 61 |
+
|
| 62 |
+
from PIL import Image
|
| 63 |
+
prompt = model.prepare_input(
|
| 64 |
+
[Image.open("https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-VL/assets/demo.jpeg")],
|
| 65 |
+
'Describe the image in detail.'
|
| 66 |
+
)
|
| 67 |
+
|
| 68 |
+
prompt2 = model.prepare_input(
|
| 69 |
+
[Image.open("https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-VL/assets/demo.jpeg")],
|
| 70 |
+
'How well it looks? Give a score between 0 and 100.'
|
| 71 |
+
)
|
| 72 |
+
res = model.inference([prompt, prompt2])
|
| 73 |
+
print("result : \n", res)
|
editscore/mllm_tools/openai.py
ADDED
|
@@ -0,0 +1,180 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import base64
|
| 2 |
+
import requests
|
| 3 |
+
from io import BytesIO, StringIO
|
| 4 |
+
from typing import Union, Optional, Tuple, List
|
| 5 |
+
from PIL import Image, ImageOps
|
| 6 |
+
import os
|
| 7 |
+
|
| 8 |
+
def get_api_key(file_path):
|
| 9 |
+
# Read the API key from the first line of the file
|
| 10 |
+
with open(file_path, 'r') as file:
|
| 11 |
+
return file.readline().strip()
|
| 12 |
+
|
| 13 |
+
# Function to encode the image
|
| 14 |
+
def encode_image(image_path):
|
| 15 |
+
with open(image_path, "rb") as image_file:
|
| 16 |
+
return base64.b64encode(image_file.read()).decode('utf-8')
|
| 17 |
+
|
| 18 |
+
def pick_next_item(current_item, item_list):
|
| 19 |
+
if current_item not in item_list:
|
| 20 |
+
raise ValueError("Current item is not in the list")
|
| 21 |
+
current_index = item_list.index(current_item)
|
| 22 |
+
next_index = (current_index + 1) % len(item_list)
|
| 23 |
+
|
| 24 |
+
return item_list[next_index]
|
| 25 |
+
|
| 26 |
+
# Function to encode a PIL image
|
| 27 |
+
def encode_pil_image(pil_image):
|
| 28 |
+
# Create an in-memory binary stream
|
| 29 |
+
image_stream = BytesIO()
|
| 30 |
+
|
| 31 |
+
# Save the PIL image to the binary stream in JPEG format (you can change the format if needed)
|
| 32 |
+
pil_image.save(image_stream, format='JPEG')
|
| 33 |
+
|
| 34 |
+
# Get the binary data from the stream and encode it as base64
|
| 35 |
+
image_data = image_stream.getvalue()
|
| 36 |
+
base64_image = base64.b64encode(image_data).decode('utf-8')
|
| 37 |
+
|
| 38 |
+
return base64_image
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
def load_image(image: Union[str, Image.Image], format: str = "RGB", size: Optional[Tuple] = None) -> Image.Image:
|
| 42 |
+
"""
|
| 43 |
+
Load an image from a given path or URL and convert it to a PIL Image.
|
| 44 |
+
|
| 45 |
+
Args:
|
| 46 |
+
image (Union[str, Image.Image]): The image path, URL, or a PIL Image object to be loaded.
|
| 47 |
+
format (str, optional): Desired color format of the resulting image. Defaults to "RGB".
|
| 48 |
+
size (Optional[Tuple], optional): Desired size for resizing the image. Defaults to None.
|
| 49 |
+
|
| 50 |
+
Returns:
|
| 51 |
+
Image.Image: A PIL Image in the specified format and size.
|
| 52 |
+
|
| 53 |
+
Raises:
|
| 54 |
+
ValueError: If the provided image format is not recognized.
|
| 55 |
+
"""
|
| 56 |
+
if isinstance(image, str):
|
| 57 |
+
if image.startswith("http://") or image.startswith("https://"):
|
| 58 |
+
image = Image.open(requests.get(image, stream=True).raw)
|
| 59 |
+
elif os.path.isfile(image):
|
| 60 |
+
image = Image.open(image)
|
| 61 |
+
else:
|
| 62 |
+
raise ValueError(
|
| 63 |
+
f"Incorrect path or url, URLs must start with `http://` or `https://`, and {image} is not a valid path"
|
| 64 |
+
)
|
| 65 |
+
elif isinstance(image, Image.Image):
|
| 66 |
+
image = image
|
| 67 |
+
else:
|
| 68 |
+
raise ValueError(
|
| 69 |
+
"Incorrect format used for image. Should be an url linking to an image, a local path, or a PIL image."
|
| 70 |
+
)
|
| 71 |
+
image = ImageOps.exif_transpose(image)
|
| 72 |
+
image = image.convert(format)
|
| 73 |
+
if (size != None):
|
| 74 |
+
image = image.resize(size, Image.LANCZOS)
|
| 75 |
+
return image
|
| 76 |
+
|
| 77 |
+
class GPT4v():
|
| 78 |
+
def __init__(self, key, url="https://api.openai.com/v1/chat/completions", are_images_encoded=False, model_name="gpt-4-vision-preview"):
|
| 79 |
+
"""OpenAI GPT-4-vision model wrapper
|
| 80 |
+
Args:
|
| 81 |
+
api_key_path (str): Path to the API key file. Defaults to 'keys/secret.env'.
|
| 82 |
+
are_images_encoded (bool): Whether the images are encoded in base64. Defaults to False.
|
| 83 |
+
"""
|
| 84 |
+
self.multiple_api_keys = False
|
| 85 |
+
self.current_key_file = None
|
| 86 |
+
self.api_key = key
|
| 87 |
+
|
| 88 |
+
self.url = url
|
| 89 |
+
self.model_name = model_name
|
| 90 |
+
self.use_encode = are_images_encoded
|
| 91 |
+
|
| 92 |
+
def prepare_input(self, image_links: List = [], text_prompt: str = ""):
|
| 93 |
+
prompt_content = []
|
| 94 |
+
text_dict = {
|
| 95 |
+
"type": "text",
|
| 96 |
+
"text": text_prompt
|
| 97 |
+
}
|
| 98 |
+
prompt_content.append(text_dict)
|
| 99 |
+
|
| 100 |
+
if not isinstance(image_links, list):
|
| 101 |
+
image_links = [image_links]
|
| 102 |
+
|
| 103 |
+
for image_link in image_links:
|
| 104 |
+
image = load_image(image_link)
|
| 105 |
+
if self.use_encode:
|
| 106 |
+
visual_dict = {
|
| 107 |
+
"type": "image_url",
|
| 108 |
+
"image_url": {"url": f"data:image/jpeg;base64,{encode_pil_image(image)}"}
|
| 109 |
+
}
|
| 110 |
+
else:
|
| 111 |
+
visual_dict = {
|
| 112 |
+
"type": "image_url",
|
| 113 |
+
"image_url": {"url": image_link}
|
| 114 |
+
}
|
| 115 |
+
prompt_content.append(visual_dict)
|
| 116 |
+
return prompt_content
|
| 117 |
+
|
| 118 |
+
def inference(self, prompt, seed: Optional[int] = None):
|
| 119 |
+
payload = {
|
| 120 |
+
"model": self.model_name,
|
| 121 |
+
"messages": [
|
| 122 |
+
{
|
| 123 |
+
"role": "user",
|
| 124 |
+
"content": prompt
|
| 125 |
+
}
|
| 126 |
+
],
|
| 127 |
+
# "max_tokens": 1400
|
| 128 |
+
}
|
| 129 |
+
headers = {
|
| 130 |
+
"Content-Type": "application/json",
|
| 131 |
+
"Authorization": f"Bearer {self.api_key}"
|
| 132 |
+
}
|
| 133 |
+
# try:
|
| 134 |
+
response = requests.post(self.url, json=payload, headers=headers, timeout=180) # Set timeout to 5 minutes (300 seconds)
|
| 135 |
+
# except Exception as e:
|
| 136 |
+
# print(f"Error: {e}")
|
| 137 |
+
# return ""
|
| 138 |
+
#return response.text
|
| 139 |
+
return self.extract_response(response)
|
| 140 |
+
|
| 141 |
+
def extract_response(self, response):
|
| 142 |
+
try:
|
| 143 |
+
response = response.json()
|
| 144 |
+
out = response['choices'][0]['message']['content']
|
| 145 |
+
return out
|
| 146 |
+
except Exception as e:
|
| 147 |
+
print(f"Error: {e}")
|
| 148 |
+
if response['error']['code'] == 'content_policy_violation':
|
| 149 |
+
print("Code is content_policy_violation")
|
| 150 |
+
elif response['error']['code'] in ['rate_limit_exceeded', 'insufficient_quota', 'insufficient_user_quota']:
|
| 151 |
+
print(f"Code is {response['error']['code']}", flush=True)
|
| 152 |
+
print(response['error']['message'], flush=True)
|
| 153 |
+
return "rate_limit_exceeded"
|
| 154 |
+
if self.multiple_api_keys == True:
|
| 155 |
+
new_key = pick_next_item(self.current_key_file, self.key_lists)
|
| 156 |
+
self.update_key(new_key)
|
| 157 |
+
self.current_key_file = new_key #override key
|
| 158 |
+
print("New key is from the file: ", new_key)
|
| 159 |
+
else:
|
| 160 |
+
print("Code is different")
|
| 161 |
+
print(response)
|
| 162 |
+
print(f"{response['error']['code']=}")
|
| 163 |
+
return ""
|
| 164 |
+
|
| 165 |
+
def update_key(self, key, load_from_file=True):
|
| 166 |
+
if load_from_file:
|
| 167 |
+
self.api_key = get_api_key(key)
|
| 168 |
+
else:
|
| 169 |
+
self.api_key = key
|
| 170 |
+
|
| 171 |
+
class GPT4o(GPT4v):
|
| 172 |
+
def __init__(self, key, url="https://api.openai.com/v1/chat/completions", are_images_encoded=False, model_name="gpt-4o-2024-05-13"):
|
| 173 |
+
super().__init__(key, url, are_images_encoded, model_name)
|
| 174 |
+
|
| 175 |
+
if __name__ == "__main__":
|
| 176 |
+
model = GPT4o('sk-cB6h7HcCSDIp71gs6lFLZxKE0dOYOnJbxzES6kWXe1Wb2VHS', model_name="gpt-4.1")
|
| 177 |
+
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?')
|
| 178 |
+
print("prompt : \n", prompt)
|
| 179 |
+
res = model.get_parsed_output(prompt)
|
| 180 |
+
print("result : \n", res)
|
editscore/mllm_tools/qwen25vl.py
ADDED
|
@@ -0,0 +1,97 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import Optional
|
| 2 |
+
import random
|
| 3 |
+
import numpy as np
|
| 4 |
+
import torch
|
| 5 |
+
|
| 6 |
+
from transformers import Qwen2_5_VLForConditionalGeneration, AutoProcessor
|
| 7 |
+
from qwen_vl_utils import process_vision_info
|
| 8 |
+
from peft import PeftModel
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def set_seed(seed: int):
|
| 12 |
+
"""
|
| 13 |
+
Args:
|
| 14 |
+
Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`.
|
| 15 |
+
seed (`int`): The seed to set.
|
| 16 |
+
"""
|
| 17 |
+
random.seed(seed)
|
| 18 |
+
np.random.seed(seed)
|
| 19 |
+
torch.manual_seed(seed)
|
| 20 |
+
torch.cuda.manual_seed_all(seed)
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def apply_chat_template(prompt, num_images: int = 2):
|
| 24 |
+
"""
|
| 25 |
+
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
|
| 26 |
+
"""
|
| 27 |
+
template = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n"
|
| 28 |
+
template += "".join([f"<img{i}>: <|vision_start|><|image_pad|><|vision_end|>" for i in range(1, num_images + 1)])
|
| 29 |
+
template += f"{prompt}<|im_end|>\n<|im_start|>assistant\n"
|
| 30 |
+
return template
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
class Qwen25VL():
|
| 34 |
+
def __init__(
|
| 35 |
+
self,
|
| 36 |
+
vlm_model,
|
| 37 |
+
temperature: float = 0.7,
|
| 38 |
+
seed: Optional[int] = None,
|
| 39 |
+
lora_path: Optional[str] = None,
|
| 40 |
+
) -> None:
|
| 41 |
+
self.model = Qwen2_5_VLForConditionalGeneration.from_pretrained(
|
| 42 |
+
vlm_model, torch_dtype=torch.bfloat16, device_map="auto"
|
| 43 |
+
)
|
| 44 |
+
if lora_path:
|
| 45 |
+
self.model = PeftModel.from_pretrained(self.model, lora_path)
|
| 46 |
+
self.model = self.model.merge_and_unload()
|
| 47 |
+
|
| 48 |
+
self.processor = AutoProcessor.from_pretrained(vlm_model)
|
| 49 |
+
self.temperature = temperature
|
| 50 |
+
self.seed = seed
|
| 51 |
+
|
| 52 |
+
def prepare_input(self, images, text_prompt: str = ""):
|
| 53 |
+
if not isinstance(images, list):
|
| 54 |
+
images = [images]
|
| 55 |
+
|
| 56 |
+
messages = [
|
| 57 |
+
{
|
| 58 |
+
"role": "user",
|
| 59 |
+
"content": [{"type": "image", "image": image} for image in images]
|
| 60 |
+
+ [{"type": "text", "text": text_prompt}],
|
| 61 |
+
}
|
| 62 |
+
]
|
| 63 |
+
text = apply_chat_template(text_prompt, num_images=len(images))
|
| 64 |
+
image_inputs, video_inputs = process_vision_info(messages)
|
| 65 |
+
|
| 66 |
+
inputs = self.processor(
|
| 67 |
+
text=[text],
|
| 68 |
+
images=image_inputs,
|
| 69 |
+
videos=video_inputs,
|
| 70 |
+
padding=True,
|
| 71 |
+
return_tensors="pt",
|
| 72 |
+
)
|
| 73 |
+
inputs = inputs.to("cuda")
|
| 74 |
+
|
| 75 |
+
return inputs
|
| 76 |
+
|
| 77 |
+
def inference(self, inputs, seed: Optional[int] = None):
|
| 78 |
+
seed = self.seed if seed is None else seed
|
| 79 |
+
|
| 80 |
+
set_seed(seed)
|
| 81 |
+
generated_ids = self.model.generate(
|
| 82 |
+
**inputs,
|
| 83 |
+
max_new_tokens=512,
|
| 84 |
+
do_sample=True,
|
| 85 |
+
temperature=self.temperature,
|
| 86 |
+
top_p=0.9,
|
| 87 |
+
top_k=20,
|
| 88 |
+
)
|
| 89 |
+
generated_ids_trimmed = [
|
| 90 |
+
out_ids[len(in_ids) :] for in_ids, out_ids in zip(inputs.input_ids, generated_ids)
|
| 91 |
+
]
|
| 92 |
+
outputs = self.processor.batch_decode(
|
| 93 |
+
generated_ids_trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False
|
| 94 |
+
)
|
| 95 |
+
|
| 96 |
+
outputs = [output.strip() for output in outputs]
|
| 97 |
+
return outputs[0]
|
editscore/mllm_tools/qwen25vl_vllm.py
ADDED
|
@@ -0,0 +1,138 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import Optional
|
| 2 |
+
|
| 3 |
+
import os
|
| 4 |
+
import hashlib
|
| 5 |
+
import random
|
| 6 |
+
import time
|
| 7 |
+
import numpy as np
|
| 8 |
+
import torch
|
| 9 |
+
|
| 10 |
+
from vllm import LLM
|
| 11 |
+
from vllm.sampling_params import SamplingParams
|
| 12 |
+
|
| 13 |
+
from transformers import Qwen2_5_VLForConditionalGeneration, AutoProcessor
|
| 14 |
+
from peft import PeftModel
|
| 15 |
+
|
| 16 |
+
from qwen_vl_utils import process_vision_info
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def set_seed(seed: int):
|
| 20 |
+
"""
|
| 21 |
+
Args:
|
| 22 |
+
Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`.
|
| 23 |
+
seed (`int`): The seed to set.
|
| 24 |
+
"""
|
| 25 |
+
random.seed(seed)
|
| 26 |
+
np.random.seed(seed)
|
| 27 |
+
torch.manual_seed(seed)
|
| 28 |
+
torch.cuda.manual_seed_all(seed)
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def apply_chat_template(prompt, num_images: int = 2):
|
| 32 |
+
"""
|
| 33 |
+
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
|
| 34 |
+
"""
|
| 35 |
+
template = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n"
|
| 36 |
+
template += "".join([f"<img{i}>: <|vision_start|><|image_pad|><|vision_end|>" for i in range(1, num_images + 1)])
|
| 37 |
+
template += f"{prompt}<|im_end|>\n<|im_start|>assistant\n"
|
| 38 |
+
return template
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
class Qwen25VL():
|
| 42 |
+
def __init__(
|
| 43 |
+
self,
|
| 44 |
+
vlm_model,
|
| 45 |
+
max_model_len: int = 1536,
|
| 46 |
+
tensor_parallel_size=1,
|
| 47 |
+
max_num_seqs=32,
|
| 48 |
+
max_num_batched_tokens=1536,
|
| 49 |
+
temperature: float = 0.7,
|
| 50 |
+
seed: Optional[int] = None,
|
| 51 |
+
lora_path: Optional[str] = None,
|
| 52 |
+
cache_dir: Optional[str] = None,
|
| 53 |
+
) -> None:
|
| 54 |
+
if lora_path:
|
| 55 |
+
if cache_dir is None:
|
| 56 |
+
root_dir = torch.hub.get_dir() # default: ~/.cache/torch/hub
|
| 57 |
+
|
| 58 |
+
lora_filename = os.path.splitext(os.path.basename(lora_path))[0]
|
| 59 |
+
lora_hash = hashlib.md5(lora_path.encode()).hexdigest()[:8]
|
| 60 |
+
lora_identifier = f"{lora_filename}_{lora_hash}"
|
| 61 |
+
|
| 62 |
+
cache_dir = os.path.join(root_dir, "EditScore", f"{os.path.basename(vlm_model)}_merged_lora_{lora_identifier}")
|
| 63 |
+
|
| 64 |
+
if not os.path.exists(cache_dir):
|
| 65 |
+
print(f"Merging LORA to {vlm_model} and saving to {cache_dir}", flush=True)
|
| 66 |
+
start_time = time.time()
|
| 67 |
+
model = Qwen2_5_VLForConditionalGeneration.from_pretrained(
|
| 68 |
+
vlm_model, torch_dtype=torch.bfloat16, device_map="cpu"
|
| 69 |
+
)
|
| 70 |
+
model = PeftModel.from_pretrained(model, lora_path)
|
| 71 |
+
model = model.merge_and_unload()
|
| 72 |
+
model.save_pretrained(cache_dir)
|
| 73 |
+
|
| 74 |
+
processor = AutoProcessor.from_pretrained(vlm_model)
|
| 75 |
+
processor.save_pretrained(cache_dir)
|
| 76 |
+
|
| 77 |
+
print(f"Merging LORA to {vlm_model} and saving to {cache_dir} took {time.time() - start_time} seconds", flush=True)
|
| 78 |
+
else:
|
| 79 |
+
print(f"Skipping merging LORA, as merged model already exists in {cache_dir}", flush=True)
|
| 80 |
+
|
| 81 |
+
vlm_model = cache_dir
|
| 82 |
+
|
| 83 |
+
self.model = LLM(
|
| 84 |
+
model=vlm_model,
|
| 85 |
+
max_model_len=max_model_len,
|
| 86 |
+
tensor_parallel_size=tensor_parallel_size,
|
| 87 |
+
max_num_seqs=max_num_seqs,
|
| 88 |
+
max_num_batched_tokens=max_num_batched_tokens,
|
| 89 |
+
limit_mm_per_prompt={"image": 2},
|
| 90 |
+
enable_prefix_caching=True,
|
| 91 |
+
)
|
| 92 |
+
self.temperature = temperature
|
| 93 |
+
self.seed = seed
|
| 94 |
+
|
| 95 |
+
def prepare_input(self, images, text_prompt: str = ""):
|
| 96 |
+
if not isinstance(images, list):
|
| 97 |
+
images = [images]
|
| 98 |
+
|
| 99 |
+
messages = [
|
| 100 |
+
{
|
| 101 |
+
"role": "user",
|
| 102 |
+
"content": [{"type": "image", "image": image} for image in images]
|
| 103 |
+
+ [{"type": "text", "text": text_prompt}],
|
| 104 |
+
}
|
| 105 |
+
]
|
| 106 |
+
text = apply_chat_template(text_prompt, num_images=len(images))
|
| 107 |
+
image_inputs, _ = process_vision_info(messages)
|
| 108 |
+
|
| 109 |
+
messages = {
|
| 110 |
+
"prompt": text,
|
| 111 |
+
"multi_modal_data": {"image": image_inputs},
|
| 112 |
+
}
|
| 113 |
+
return messages
|
| 114 |
+
|
| 115 |
+
def inference(self, messages, seed: Optional[int] = None):
|
| 116 |
+
seed = self.seed if seed is None else seed
|
| 117 |
+
sampling_params = SamplingParams(max_tokens=512, temperature=self.temperature, top_p=0.9, top_k=20, seed=seed)
|
| 118 |
+
outputs = self.model.generate(messages, sampling_params, use_tqdm=False)
|
| 119 |
+
|
| 120 |
+
responses = []
|
| 121 |
+
for output in outputs:
|
| 122 |
+
instruction = output.outputs[0].text.strip()
|
| 123 |
+
responses.append(instruction)
|
| 124 |
+
|
| 125 |
+
return responses[0]
|
| 126 |
+
|
| 127 |
+
|
| 128 |
+
def batch_inference(self, messages, seed: Optional[int] = None):
|
| 129 |
+
seed = self.seed if seed is None else seed
|
| 130 |
+
sampling_params = SamplingParams(max_tokens=512, temperature=self.temperature, top_p=0.9, top_k=20, seed=seed)
|
| 131 |
+
outputs = self.model.generate(messages, sampling_params, use_tqdm=False)
|
| 132 |
+
|
| 133 |
+
responses = []
|
| 134 |
+
for output in outputs:
|
| 135 |
+
instruction = output.outputs[0].text.strip()
|
| 136 |
+
responses.append(instruction)
|
| 137 |
+
|
| 138 |
+
return responses
|
editscore/mllm_tools/qwen3vl.py
ADDED
|
@@ -0,0 +1,95 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import Optional
|
| 2 |
+
import random
|
| 3 |
+
import numpy as np
|
| 4 |
+
import torch
|
| 5 |
+
|
| 6 |
+
from transformers import Qwen3VLForConditionalGeneration, AutoProcessor
|
| 7 |
+
from peft import PeftModel
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
def set_seed(seed: int):
|
| 11 |
+
"""
|
| 12 |
+
Args:
|
| 13 |
+
Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`.
|
| 14 |
+
seed (`int`): The seed to set.
|
| 15 |
+
"""
|
| 16 |
+
random.seed(seed)
|
| 17 |
+
np.random.seed(seed)
|
| 18 |
+
torch.manual_seed(seed)
|
| 19 |
+
torch.cuda.manual_seed_all(seed)
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def apply_chat_template(prompt, num_images: int = 2):
|
| 23 |
+
"""
|
| 24 |
+
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
|
| 25 |
+
"""
|
| 26 |
+
template = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n"
|
| 27 |
+
template += "".join([f"<img{i}>: <|vision_start|><|image_pad|><|vision_end|>" for i in range(1, num_images + 1)])
|
| 28 |
+
template += f"{prompt}<|im_end|>\n<|im_start|>assistant\n"
|
| 29 |
+
return template
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
class Qwen3VL():
|
| 33 |
+
def __init__(
|
| 34 |
+
self,
|
| 35 |
+
vlm_model,
|
| 36 |
+
temperature: float = 0.7,
|
| 37 |
+
seed: Optional[int] = None,
|
| 38 |
+
lora_path: Optional[str] = None,
|
| 39 |
+
) -> None:
|
| 40 |
+
self.model = Qwen3VLForConditionalGeneration.from_pretrained(
|
| 41 |
+
vlm_model, torch_dtype=torch.bfloat16, device_map="auto"
|
| 42 |
+
)
|
| 43 |
+
if lora_path:
|
| 44 |
+
self.model = PeftModel.from_pretrained(self.model, lora_path)
|
| 45 |
+
self.model = self.model.merge_and_unload()
|
| 46 |
+
|
| 47 |
+
self.processor = AutoProcessor.from_pretrained(vlm_model)
|
| 48 |
+
self.temperature = temperature
|
| 49 |
+
self.seed = seed
|
| 50 |
+
|
| 51 |
+
def prepare_input(self, images, text_prompt: str = ""):
|
| 52 |
+
if not isinstance(images, list):
|
| 53 |
+
images = [images]
|
| 54 |
+
|
| 55 |
+
messages = [
|
| 56 |
+
{
|
| 57 |
+
"role": "user",
|
| 58 |
+
"content": [{"type": "image", "image": image} for image in images]
|
| 59 |
+
+ [{"type": "text", "text": text_prompt}],
|
| 60 |
+
}
|
| 61 |
+
]
|
| 62 |
+
|
| 63 |
+
inputs = self.processor.apply_chat_template(
|
| 64 |
+
messages,
|
| 65 |
+
tokenize=True,
|
| 66 |
+
add_generation_prompt=True,
|
| 67 |
+
return_dict=True,
|
| 68 |
+
return_tensors="pt"
|
| 69 |
+
)
|
| 70 |
+
|
| 71 |
+
inputs = inputs.to("cuda")
|
| 72 |
+
|
| 73 |
+
return inputs
|
| 74 |
+
|
| 75 |
+
def inference(self, inputs, seed: Optional[int] = None):
|
| 76 |
+
seed = self.seed if seed is None else seed
|
| 77 |
+
|
| 78 |
+
set_seed(seed)
|
| 79 |
+
generated_ids = self.model.generate(
|
| 80 |
+
**inputs,
|
| 81 |
+
max_new_tokens=512,
|
| 82 |
+
do_sample=True,
|
| 83 |
+
temperature=self.temperature,
|
| 84 |
+
top_p=0.9,
|
| 85 |
+
top_k=20,
|
| 86 |
+
)
|
| 87 |
+
generated_ids_trimmed = [
|
| 88 |
+
out_ids[len(in_ids) :] for in_ids, out_ids in zip(inputs.input_ids, generated_ids)
|
| 89 |
+
]
|
| 90 |
+
outputs = self.processor.batch_decode(
|
| 91 |
+
generated_ids_trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False
|
| 92 |
+
)
|
| 93 |
+
|
| 94 |
+
outputs = [output.strip() for output in outputs]
|
| 95 |
+
return outputs[0]
|
editscore/mllm_tools/qwen3vl_vllm.py
ADDED
|
@@ -0,0 +1,141 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import Optional
|
| 2 |
+
|
| 3 |
+
import os
|
| 4 |
+
import hashlib
|
| 5 |
+
import random
|
| 6 |
+
import time
|
| 7 |
+
import numpy as np
|
| 8 |
+
import torch
|
| 9 |
+
|
| 10 |
+
from vllm import LLM
|
| 11 |
+
from vllm.sampling_params import SamplingParams
|
| 12 |
+
|
| 13 |
+
from transformers import Qwen3VLForConditionalGeneration, AutoProcessor
|
| 14 |
+
from peft import PeftModel
|
| 15 |
+
|
| 16 |
+
from qwen_vl_utils import process_vision_info
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def set_seed(seed: int):
|
| 20 |
+
"""
|
| 21 |
+
Args:
|
| 22 |
+
Helper function for reproducible behavior to set the seed in `random`, `numpy`, `torch`.
|
| 23 |
+
seed (`int`): The seed to set.
|
| 24 |
+
"""
|
| 25 |
+
random.seed(seed)
|
| 26 |
+
np.random.seed(seed)
|
| 27 |
+
torch.manual_seed(seed)
|
| 28 |
+
torch.cuda.manual_seed_all(seed)
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def apply_chat_template(prompt, num_images: int = 2):
|
| 32 |
+
"""
|
| 33 |
+
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
|
| 34 |
+
"""
|
| 35 |
+
template = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n"
|
| 36 |
+
template += "".join([f"<img{i}>: <|vision_start|><|image_pad|><|vision_end|>" for i in range(1, num_images + 1)])
|
| 37 |
+
template += f"{prompt}<|im_end|>\n<|im_start|>assistant\n"
|
| 38 |
+
return template
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
class Qwen3VL():
|
| 42 |
+
def __init__(
|
| 43 |
+
self,
|
| 44 |
+
vlm_model,
|
| 45 |
+
max_model_len: int = 1536,
|
| 46 |
+
tensor_parallel_size=1,
|
| 47 |
+
max_num_seqs=32,
|
| 48 |
+
max_num_batched_tokens=1536,
|
| 49 |
+
temperature: float = 0.7,
|
| 50 |
+
seed: Optional[int] = None,
|
| 51 |
+
lora_path: Optional[str] = None,
|
| 52 |
+
cache_dir: Optional[str] = None,
|
| 53 |
+
) -> None:
|
| 54 |
+
if lora_path:
|
| 55 |
+
if cache_dir is None:
|
| 56 |
+
root_dir = torch.hub.get_dir() # default: ~/.cache/torch/hub
|
| 57 |
+
|
| 58 |
+
lora_filename = os.path.splitext(os.path.basename(lora_path))[0]
|
| 59 |
+
lora_hash = hashlib.md5(lora_path.encode()).hexdigest()[:8]
|
| 60 |
+
lora_identifier = f"{lora_filename}_{lora_hash}"
|
| 61 |
+
|
| 62 |
+
cache_dir = os.path.join(root_dir, "EditScore", f"{os.path.basename(vlm_model)}_merged_lora_{lora_identifier}")
|
| 63 |
+
|
| 64 |
+
if not os.path.exists(cache_dir):
|
| 65 |
+
print(f"Merging LORA to {vlm_model} and saving to {cache_dir}", flush=True)
|
| 66 |
+
start_time = time.time()
|
| 67 |
+
model = Qwen3VLForConditionalGeneration.from_pretrained(
|
| 68 |
+
vlm_model, torch_dtype=torch.bfloat16, device_map="cpu"
|
| 69 |
+
)
|
| 70 |
+
model = PeftModel.from_pretrained(model, lora_path)
|
| 71 |
+
model = model.merge_and_unload()
|
| 72 |
+
model.save_pretrained(cache_dir)
|
| 73 |
+
|
| 74 |
+
processor = AutoProcessor.from_pretrained(vlm_model)
|
| 75 |
+
processor.save_pretrained(cache_dir)
|
| 76 |
+
|
| 77 |
+
print(f"Merging LORA to {vlm_model} and saving to {cache_dir} took {time.time() - start_time} seconds", flush=True)
|
| 78 |
+
else:
|
| 79 |
+
print(f"Skipping merging LORA, as merged model already exists in {cache_dir}", flush=True)
|
| 80 |
+
|
| 81 |
+
vlm_model = cache_dir
|
| 82 |
+
|
| 83 |
+
self.model = LLM(
|
| 84 |
+
model=vlm_model,
|
| 85 |
+
max_model_len=max_model_len,
|
| 86 |
+
tensor_parallel_size=tensor_parallel_size,
|
| 87 |
+
max_num_seqs=max_num_seqs,
|
| 88 |
+
max_num_batched_tokens=max_num_batched_tokens,
|
| 89 |
+
limit_mm_per_prompt={"image": 2},
|
| 90 |
+
enable_prefix_caching=True,
|
| 91 |
+
)
|
| 92 |
+
|
| 93 |
+
self.processor = AutoProcessor.from_pretrained(vlm_model)
|
| 94 |
+
self.temperature = temperature
|
| 95 |
+
self.seed = seed
|
| 96 |
+
|
| 97 |
+
def prepare_input(self, images, text_prompt: str = ""):
|
| 98 |
+
if not isinstance(images, list):
|
| 99 |
+
images = [images]
|
| 100 |
+
|
| 101 |
+
messages = [
|
| 102 |
+
{
|
| 103 |
+
"role": "user",
|
| 104 |
+
"content": [{"type": "image", "image": image} for image in images]
|
| 105 |
+
+ [{"type": "text", "text": text_prompt}],
|
| 106 |
+
}
|
| 107 |
+
]
|
| 108 |
+
# text = apply_chat_template(text_prompt, num_images=len(images))
|
| 109 |
+
text = self.processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
|
| 110 |
+
image_inputs, _ = process_vision_info(messages)
|
| 111 |
+
|
| 112 |
+
messages = {
|
| 113 |
+
"prompt": text,
|
| 114 |
+
"multi_modal_data": {"image": image_inputs},
|
| 115 |
+
}
|
| 116 |
+
return messages
|
| 117 |
+
|
| 118 |
+
def inference(self, messages, seed: Optional[int] = None):
|
| 119 |
+
seed = self.seed if seed is None else seed
|
| 120 |
+
sampling_params = SamplingParams(max_tokens=512, temperature=self.temperature, top_p=0.9, top_k=20, seed=seed)
|
| 121 |
+
outputs = self.model.generate(messages, sampling_params, use_tqdm=False)
|
| 122 |
+
|
| 123 |
+
responses = []
|
| 124 |
+
for output in outputs:
|
| 125 |
+
instruction = output.outputs[0].text.strip()
|
| 126 |
+
responses.append(instruction)
|
| 127 |
+
|
| 128 |
+
return responses[0]
|
| 129 |
+
|
| 130 |
+
|
| 131 |
+
def batch_inference(self, messages, seed: Optional[int] = None):
|
| 132 |
+
seed = self.seed if seed is None else seed
|
| 133 |
+
sampling_params = SamplingParams(max_tokens=512, temperature=self.temperature, top_p=0.9, top_k=20, seed=seed)
|
| 134 |
+
outputs = self.model.generate(messages, sampling_params, use_tqdm=False)
|
| 135 |
+
|
| 136 |
+
responses = []
|
| 137 |
+
for output in outputs:
|
| 138 |
+
instruction = output.outputs[0].text.strip()
|
| 139 |
+
responses.append(instruction)
|
| 140 |
+
|
| 141 |
+
return responses
|
editscore/mllm_tools/utils.py
ADDED
|
@@ -0,0 +1,65 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import List
|
| 2 |
+
import base64
|
| 3 |
+
from io import BytesIO
|
| 4 |
+
from PIL import Image
|
| 5 |
+
import requests
|
| 6 |
+
|
| 7 |
+
def pil_image_to_base64(pil_image, format="PNG"):
|
| 8 |
+
buffered = BytesIO()
|
| 9 |
+
pil_image.save(buffered, format=format) # Save image to the buffer in the specified format
|
| 10 |
+
img_str = base64.b64encode(buffered.getvalue()).decode('utf-8') # Encode the buffer's content to base64
|
| 11 |
+
return img_str
|
| 12 |
+
|
| 13 |
+
def load_image(image_file):
|
| 14 |
+
if image_file.startswith("http"):
|
| 15 |
+
response = requests.get(image_file)
|
| 16 |
+
image = Image.open(BytesIO(response.content)).convert("RGB")
|
| 17 |
+
else:
|
| 18 |
+
import os
|
| 19 |
+
image = Image.open(image_file).convert("RGB")
|
| 20 |
+
return image
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def load_images(image_files):
|
| 24 |
+
out = []
|
| 25 |
+
for image_file in image_files:
|
| 26 |
+
image = load_image(image_file)
|
| 27 |
+
out.append(image)
|
| 28 |
+
return out
|
| 29 |
+
|
| 30 |
+
def merge_images(image_links: List = []):
|
| 31 |
+
"""Merge multiple images into one image
|
| 32 |
+
|
| 33 |
+
Args:
|
| 34 |
+
image_links (List, optional): List of image links. Defaults to [].
|
| 35 |
+
|
| 36 |
+
Returns:
|
| 37 |
+
[type]: [description]
|
| 38 |
+
"""
|
| 39 |
+
if len(image_links) == 0:
|
| 40 |
+
return None
|
| 41 |
+
images = load_images(image_links)
|
| 42 |
+
if len(images) == 1:
|
| 43 |
+
return images[0]
|
| 44 |
+
widths, heights = zip(*(i.size for i in images))
|
| 45 |
+
average_height = sum(heights) // len(heights)
|
| 46 |
+
for i, im in enumerate(images):
|
| 47 |
+
# scale in proportion
|
| 48 |
+
images[i] = im.resize((int(im.size[0] * average_height / im.size[1]), average_height))
|
| 49 |
+
widths, heights = zip(*(i.size for i in images))
|
| 50 |
+
total_width = sum(widths)
|
| 51 |
+
max_height = max(heights)
|
| 52 |
+
new_im = Image.new("RGB", (total_width + 10 * (len(images) - 1), max_height))
|
| 53 |
+
x_offset = 0
|
| 54 |
+
for i, im in enumerate(images):
|
| 55 |
+
if i > 0:
|
| 56 |
+
# past a column of 1 pixel starting from x_offset width being black, 8 pixels being white, and 1 pixel being black
|
| 57 |
+
new_im.paste(Image.new("RGB", (1, max_height), (0, 0, 0)), (x_offset, 0))
|
| 58 |
+
x_offset += 1
|
| 59 |
+
new_im.paste(Image.new("RGB", (8, max_height), (255, 255, 255)), (x_offset, 0))
|
| 60 |
+
x_offset += 8
|
| 61 |
+
new_im.paste(Image.new("RGB", (1, max_height), (0, 0, 0)), (x_offset, 0))
|
| 62 |
+
x_offset += 1
|
| 63 |
+
new_im.paste(im, (x_offset, 0))
|
| 64 |
+
x_offset += im.size[0]
|
| 65 |
+
return new_im
|
editscore/utils.py
ADDED
|
@@ -0,0 +1,534 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
from typing import Union, List, Optional
|
| 3 |
+
import json
|
| 4 |
+
import regex as re
|
| 5 |
+
import ast
|
| 6 |
+
import random
|
| 7 |
+
import json_repair
|
| 8 |
+
|
| 9 |
+
def fix_json(input_str):
|
| 10 |
+
# Add double quotes around keys using regex
|
| 11 |
+
fixed_str = re.sub(r'(\w+):', r'"\1":', input_str)
|
| 12 |
+
|
| 13 |
+
# Add double quotes around string values if necessary and wrap int/float values in []
|
| 14 |
+
def format_value(match):
|
| 15 |
+
key, value, comma = match.groups()
|
| 16 |
+
value = value.strip()
|
| 17 |
+
# Check if value is an integer or float
|
| 18 |
+
if re.match(r'^-?\d+(\.\d+)?$', value):
|
| 19 |
+
value = f'[{value}]'
|
| 20 |
+
# Check if value is a boolean or null
|
| 21 |
+
elif re.match(r'^(true|false|null)$', value, re.IGNORECASE):
|
| 22 |
+
pass # leave as is
|
| 23 |
+
else:
|
| 24 |
+
# Add quotes around string values
|
| 25 |
+
value = f'"{value}"'
|
| 26 |
+
return f'{key}: {value}{comma}'
|
| 27 |
+
|
| 28 |
+
fixed_str = re.sub(r'(".*?"):(.*?)(,|})', format_value, fixed_str)
|
| 29 |
+
|
| 30 |
+
return fixed_str
|
| 31 |
+
|
| 32 |
+
def repair_reasoning_field_robust(json_str: str) -> str:
|
| 33 |
+
"""
|
| 34 |
+
Robustly repair unescaped double quotes inside the "reasoning" field of a JSON string.
|
| 35 |
+
This function uses regular expressions and a lookahead assertion to locate
|
| 36 |
+
the end of the "reasoning" value, even if it is not the last field in the JSON.
|
| 37 |
+
|
| 38 |
+
Args:
|
| 39 |
+
json_str (str): A possibly malformed JSON string that may contain
|
| 40 |
+
unescaped quotes within the "reasoning" field.
|
| 41 |
+
|
| 42 |
+
Returns:
|
| 43 |
+
str: A repaired JSON string with properly escaped quotes inside "reasoning".
|
| 44 |
+
"""
|
| 45 |
+
# 1. Define a regex pattern that locates the reasoning value using a lookahead.
|
| 46 |
+
# The re.DOTALL flag allows '.' to match newline characters.
|
| 47 |
+
pattern = re.compile(
|
| 48 |
+
# --- Group 1: prefix part including the "reasoning" key and opening quote ---
|
| 49 |
+
r'("reasoning"\s*:\s*")'
|
| 50 |
+
|
| 51 |
+
# --- Group 2: content inside the reasoning string (non-greedy) ---
|
| 52 |
+
r'(.*?)'
|
| 53 |
+
|
| 54 |
+
# --- Lookahead assertion ---
|
| 55 |
+
# Match the ending quote of the "reasoning" value,
|
| 56 |
+
# but only if it is followed by a comma or closing brace.
|
| 57 |
+
r'(?="\s*[,}])',
|
| 58 |
+
|
| 59 |
+
re.DOTALL
|
| 60 |
+
)
|
| 61 |
+
|
| 62 |
+
# 2. Define a replacement function to escape quotes inside the "reasoning" content.
|
| 63 |
+
def replacer(match):
|
| 64 |
+
prefix = match.group(1) # e.g., '"reasoning": "'
|
| 65 |
+
content = match.group(2) # e.g., 'Overall building...'
|
| 66 |
+
|
| 67 |
+
# Escape all unescaped double quotes inside the reasoning text.
|
| 68 |
+
fixed_content = content.replace('"', '\\"')
|
| 69 |
+
|
| 70 |
+
# Reassemble the full matched segment. The suffix is not consumed by the pattern,
|
| 71 |
+
# so we just return the prefix + repaired content.
|
| 72 |
+
return prefix + fixed_content
|
| 73 |
+
|
| 74 |
+
# 3. Apply the regex substitution across the entire JSON string.
|
| 75 |
+
repaired_str = pattern.sub(replacer, json_str)
|
| 76 |
+
|
| 77 |
+
return repaired_str
|
| 78 |
+
|
| 79 |
+
def fallback_repair_json(input_str: str) -> str:
|
| 80 |
+
"""
|
| 81 |
+
Last-resort JSON repair that tries to preserve the 'reasoning' text
|
| 82 |
+
even when it contains unescaped quotes or other corruption.
|
| 83 |
+
|
| 84 |
+
Target output:
|
| 85 |
+
{"reasoning": "<text>", "score": [float, float]}
|
| 86 |
+
|
| 87 |
+
Approach:
|
| 88 |
+
1. Locate 'reasoning' key position and 'score' key position.
|
| 89 |
+
2. Extract the raw substring between them (reasoning_raw).
|
| 90 |
+
3. Clean only the outer noise (leading/trailing quotes, commas, braces),
|
| 91 |
+
but preserve internal punctuation.
|
| 92 |
+
4. Unescape common escape sequences and normalize quotes.
|
| 93 |
+
5. Extract numeric scores robustly.
|
| 94 |
+
6. Return a valid JSON string.
|
| 95 |
+
"""
|
| 96 |
+
|
| 97 |
+
s = input_str
|
| 98 |
+
|
| 99 |
+
# Normalize whitespace for easier searching (but keep original for slicing)
|
| 100 |
+
lowered = s.lower()
|
| 101 |
+
|
| 102 |
+
# 1) find the start of reasoning key (case-insensitive)
|
| 103 |
+
m_reason = re.search(r'"?reasoning"?\s*[::]', lowered)
|
| 104 |
+
m_score = re.search(r'"?score"?\s*[::]', lowered)
|
| 105 |
+
|
| 106 |
+
reasoning_text = ""
|
| 107 |
+
scores = []
|
| 108 |
+
|
| 109 |
+
if m_reason and m_score:
|
| 110 |
+
# compute the real indices in the original string
|
| 111 |
+
start_idx = m_reason.end() # right after colon in 'reasoning:'
|
| 112 |
+
score_start_idx = m_score.start()
|
| 113 |
+
|
| 114 |
+
# 2) slice the original string between reasoning value start and score key start
|
| 115 |
+
reasoning_raw = s[start_idx:score_start_idx]
|
| 116 |
+
|
| 117 |
+
# 3) clean outer noise but preserve inner content:
|
| 118 |
+
# - strip whitespace and outer commas/braces
|
| 119 |
+
reasoning_raw = reasoning_raw.strip()
|
| 120 |
+
# remove leading commas/braces/colons
|
| 121 |
+
reasoning_raw = re.sub(r'^[\s,{\[]+', '', reasoning_raw)
|
| 122 |
+
# remove trailing commas/braces/colons (but keep inner punctuation)
|
| 123 |
+
reasoning_raw = re.sub(r'[\s,}\]]+$', '', reasoning_raw)
|
| 124 |
+
|
| 125 |
+
# If the reasoning starts with a quote char, drop it (we'll re-escape later).
|
| 126 |
+
if reasoning_raw.startswith(("'", '"')):
|
| 127 |
+
reasoning_raw = reasoning_raw[1:]
|
| 128 |
+
# If it ends with a quote char (common), drop it.
|
| 129 |
+
if reasoning_raw.endswith(("'", '"')):
|
| 130 |
+
reasoning_raw = reasoning_raw[:-1]
|
| 131 |
+
|
| 132 |
+
# 4) normalize escapes:
|
| 133 |
+
# Replace common escaped sequences (\" -> "), but avoid creating unbalanced quotes.
|
| 134 |
+
reasoning_raw = reasoning_raw.replace('\\"', '"').replace("\\'", "'")
|
| 135 |
+
# Replace fancy quotes with straight quotes (optional)
|
| 136 |
+
reasoning_raw = re.sub(r'[“”]', '"', reasoning_raw)
|
| 137 |
+
reasoning_raw = re.sub(r"[‘’]", "'", reasoning_raw)
|
| 138 |
+
|
| 139 |
+
# Trim again
|
| 140 |
+
reasoning_text = reasoning_raw.strip()
|
| 141 |
+
else:
|
| 142 |
+
# If we couldn't find both keys, try a looser regex capturing 'reasoning' value
|
| 143 |
+
m_loose = re.search(r'"?reasoning"?\s*[::]\s*["\']?(.*?)["\']?\s*(,|$)', s, re.DOTALL | re.IGNORECASE)
|
| 144 |
+
if m_loose:
|
| 145 |
+
reasoning_text = m_loose.group(1).strip()
|
| 146 |
+
# normalize escapes as above
|
| 147 |
+
reasoning_text = reasoning_text.replace('\\"', '"').replace("\\'", "'")
|
| 148 |
+
reasoning_text = re.sub(r'[“”]', '"', reasoning_text)
|
| 149 |
+
reasoning_text = re.sub(r"[‘’]", "'", reasoning_text)
|
| 150 |
+
|
| 151 |
+
# 5) Extract two numeric scores anywhere after the 'score' key (robust)
|
| 152 |
+
if m_score:
|
| 153 |
+
# slice from score key to the end
|
| 154 |
+
score_slice = s[m_score.end():]
|
| 155 |
+
# find numbers (integers or floats)
|
| 156 |
+
nums = re.findall(r'-?\d+(?:\.\d+)?', score_slice)
|
| 157 |
+
try:
|
| 158 |
+
scores = [float(n) for n in nums[:2]]
|
| 159 |
+
except Exception:
|
| 160 |
+
scores = []
|
| 161 |
+
else:
|
| 162 |
+
# fallback: try to find any two numbers in the whole string
|
| 163 |
+
nums = re.findall(r'-?\d+(?:\.\d+)?', s)
|
| 164 |
+
try:
|
| 165 |
+
scores = [float(n) for n in nums[:2]]
|
| 166 |
+
except Exception:
|
| 167 |
+
scores = []
|
| 168 |
+
|
| 169 |
+
# Ensure we always return two floats
|
| 170 |
+
if len(scores) < 2:
|
| 171 |
+
scores += [0.0] * (2 - len(scores))
|
| 172 |
+
|
| 173 |
+
# 6) Construct final object. Let json.dumps handle escaping inside the reasoning.
|
| 174 |
+
repaired_obj = {
|
| 175 |
+
"reasoning": reasoning_text,
|
| 176 |
+
"score": scores
|
| 177 |
+
}
|
| 178 |
+
|
| 179 |
+
return json.dumps(repaired_obj, ensure_ascii=False)
|
| 180 |
+
|
| 181 |
+
def robust_json_fix(s: str):
|
| 182 |
+
try:
|
| 183 |
+
return json_repair.loads(s)
|
| 184 |
+
except Exception:
|
| 185 |
+
pass
|
| 186 |
+
|
| 187 |
+
for fixer in [fix_json, repair_reasoning_field_robust]:
|
| 188 |
+
s = fixer(s)
|
| 189 |
+
try:
|
| 190 |
+
return json_repair.loads(s)
|
| 191 |
+
except Exception:
|
| 192 |
+
print(f"Error: Cannot fix {fixer.__name__} {s=}")
|
| 193 |
+
continue
|
| 194 |
+
|
| 195 |
+
try:
|
| 196 |
+
repaired_str = fallback_repair_json(s)
|
| 197 |
+
return json_repair.loads(repaired_str)
|
| 198 |
+
except Exception as e:
|
| 199 |
+
print(f"Error: Cannot fix fallback_repair_json {s=} {e=}")
|
| 200 |
+
return False
|
| 201 |
+
|
| 202 |
+
def read_file_to_string(file_path):
|
| 203 |
+
"""
|
| 204 |
+
Reads the contents of a text file and returns it as a string.
|
| 205 |
+
|
| 206 |
+
:param file_path: The path to the text file.
|
| 207 |
+
:return: A string containing the contents of the file.
|
| 208 |
+
"""
|
| 209 |
+
try:
|
| 210 |
+
with open(file_path, 'r', encoding='utf-8') as file:
|
| 211 |
+
return file.read()
|
| 212 |
+
except FileNotFoundError:
|
| 213 |
+
print(f"The file {file_path} was not found.")
|
| 214 |
+
return None
|
| 215 |
+
except Exception as e:
|
| 216 |
+
print(f"An error occurred: {e}")
|
| 217 |
+
return None
|
| 218 |
+
|
| 219 |
+
def read_files_to_string(file_paths):
|
| 220 |
+
"""
|
| 221 |
+
Reads the contents of multiple text files and returns them as a single string,
|
| 222 |
+
with each file's contents separated by a newline.
|
| 223 |
+
|
| 224 |
+
:param file_paths: A list of paths to text files.
|
| 225 |
+
:return: A string containing the concatenated contents of the files.
|
| 226 |
+
"""
|
| 227 |
+
all_contents = [] # List to hold the contents of each file
|
| 228 |
+
|
| 229 |
+
for file_path in file_paths:
|
| 230 |
+
try:
|
| 231 |
+
with open(file_path, 'r', encoding='utf-8') as file:
|
| 232 |
+
all_contents.append(file.read())
|
| 233 |
+
except FileNotFoundError:
|
| 234 |
+
print(f"The file {file_path} was not found.")
|
| 235 |
+
except Exception as e:
|
| 236 |
+
print(f"An error occurred while reading {file_path}: {e}")
|
| 237 |
+
|
| 238 |
+
# Join all the contents with a newline character
|
| 239 |
+
return "\n".join(all_contents)
|
| 240 |
+
|
| 241 |
+
def get_file_path(filename: Union[str, os.PathLike], search_from: Union[str, os.PathLike] = "."):
|
| 242 |
+
"""
|
| 243 |
+
Search for a file across a directory and return its absolute path.
|
| 244 |
+
|
| 245 |
+
Args:
|
| 246 |
+
filename (Union[str, os.PathLike]): The name of the file to search for.
|
| 247 |
+
search_from (Union[str, os.PathLike], optional): The directory from which to start the search. Defaults to ".".
|
| 248 |
+
|
| 249 |
+
Returns:
|
| 250 |
+
str: Absolute path to the found file.
|
| 251 |
+
|
| 252 |
+
Raises:
|
| 253 |
+
FileNotFoundError: If the file is not found.
|
| 254 |
+
"""
|
| 255 |
+
for root, dirs, files in os.walk(search_from):
|
| 256 |
+
for name in files:
|
| 257 |
+
if name == filename:
|
| 258 |
+
return os.path.abspath(os.path.join(root, name))
|
| 259 |
+
raise FileNotFoundError(filename, "not found.")
|
| 260 |
+
|
| 261 |
+
|
| 262 |
+
|
| 263 |
+
#+=========================================================================================
|
| 264 |
+
def verify(s, target_sequence):
|
| 265 |
+
# Count the occurrences of the target sequence
|
| 266 |
+
count = s.count(target_sequence)
|
| 267 |
+
|
| 268 |
+
# Check if the target sequence appears exactly twice
|
| 269 |
+
return count == 2
|
| 270 |
+
|
| 271 |
+
|
| 272 |
+
def is_int_between_0_and_10(s):
|
| 273 |
+
try:
|
| 274 |
+
num = int(s)
|
| 275 |
+
return 0 <= num <= 10
|
| 276 |
+
except ValueError:
|
| 277 |
+
return False
|
| 278 |
+
|
| 279 |
+
def is_str_a_list_of_ints_0_to_10(s):
|
| 280 |
+
try:
|
| 281 |
+
# Attempt to parse the string as a Python literal (list, dict, etc.)
|
| 282 |
+
parsed = ast.literal_eval(s)
|
| 283 |
+
|
| 284 |
+
# Check if the parsed object is a list
|
| 285 |
+
if not isinstance(parsed, list):
|
| 286 |
+
return False
|
| 287 |
+
|
| 288 |
+
# Check if all elements are integers and between 0 to 10
|
| 289 |
+
return all(isinstance(item, int) and 0 <= item <= 10 for item in parsed)
|
| 290 |
+
|
| 291 |
+
except (ValueError, SyntaxError):
|
| 292 |
+
# If parsing fails or any other error occurs
|
| 293 |
+
return False
|
| 294 |
+
|
| 295 |
+
def is_str_valid_score_format_brackets(s):
|
| 296 |
+
try:
|
| 297 |
+
# Removing brackets and splitting the string by commas
|
| 298 |
+
content = s.strip("[]").split(',')
|
| 299 |
+
|
| 300 |
+
length = len(content)
|
| 301 |
+
|
| 302 |
+
# Parsing each element and checking the format and range
|
| 303 |
+
scores = {}
|
| 304 |
+
for item in content:
|
| 305 |
+
key, value = item.split(':')
|
| 306 |
+
key = key.strip()
|
| 307 |
+
value = int(value.strip())
|
| 308 |
+
|
| 309 |
+
# Check if the key starts with 'score' and the value is in the correct range
|
| 310 |
+
if not key.startswith("score") or not 0 <= value <= 10:
|
| 311 |
+
return False
|
| 312 |
+
|
| 313 |
+
scores[key] = value
|
| 314 |
+
|
| 315 |
+
fetch_words = [f"score{i+1}" for i in range(length)]
|
| 316 |
+
# Check if at least 'score1' and 'score2' are present
|
| 317 |
+
return all(key in scores for key in fetch_words)
|
| 318 |
+
|
| 319 |
+
except (ValueError, SyntaxError):
|
| 320 |
+
# If any parsing error occurs
|
| 321 |
+
return False
|
| 322 |
+
|
| 323 |
+
def normalize_quotes(s: str) -> str:
|
| 324 |
+
"""
|
| 325 |
+
Replace curly/smart quotes with normal ASCII quotes.
|
| 326 |
+
"""
|
| 327 |
+
# 常见的几种智能引号 U+201C U+201D U+2018 U+2019
|
| 328 |
+
return s.replace("“", '"').replace("”", '"').replace("‘", "'").replace("’", "'")
|
| 329 |
+
|
| 330 |
+
|
| 331 |
+
#+=========================================================================================
|
| 332 |
+
def mllm_output_to_dict(input_string, give_up_parsing=False, text_prompt=None, score_range: int = 10):
|
| 333 |
+
"""
|
| 334 |
+
Args:
|
| 335 |
+
input_string (str): actually the output of the mllm model to be parsed
|
| 336 |
+
output_file_name (str): The name of the output file.
|
| 337 |
+
"""
|
| 338 |
+
# Catch for gpt4v rate_limit_exceeded error
|
| 339 |
+
if input_string == "rate_limit_exceeded":
|
| 340 |
+
return "rate_limit_exceeded"
|
| 341 |
+
|
| 342 |
+
if give_up_parsing:
|
| 343 |
+
guessed_value = random.randint(0, score_range)
|
| 344 |
+
json_content = {'score': [guessed_value, guessed_value], "reasoning": f"guess_if_cannot_parse | {input_string}"}
|
| 345 |
+
return json_content
|
| 346 |
+
|
| 347 |
+
# Define the delimiters
|
| 348 |
+
delimiter = '||V^=^V||'
|
| 349 |
+
|
| 350 |
+
if input_string.count(delimiter) == 2:
|
| 351 |
+
if not verify(input_string, delimiter):
|
| 352 |
+
print("The required delimiters were not found correctly in the string.", flush=True)
|
| 353 |
+
return False
|
| 354 |
+
# Extract the content between the delimiters
|
| 355 |
+
start_index = input_string.find(delimiter) + len(delimiter)
|
| 356 |
+
end_index = input_string.rfind(delimiter)
|
| 357 |
+
else:
|
| 358 |
+
# find the json mannually
|
| 359 |
+
# some mllm tends not to output the delimiters, but it does output the json contents
|
| 360 |
+
# so we will find the json content mannually
|
| 361 |
+
start_index = input_string.find('{')
|
| 362 |
+
end_index = input_string.rfind('}') + 1
|
| 363 |
+
if start_index == -1 or end_index == 0:
|
| 364 |
+
# json not found
|
| 365 |
+
# some mllm tends to output only a list of scores like [6, 0],
|
| 366 |
+
# this time we will just get the scores and ignore the reasoning (other part of the json)
|
| 367 |
+
start_index = input_string.find('[')
|
| 368 |
+
end_index = input_string.rfind(']') + 1
|
| 369 |
+
if re.match(r'^\[\d+, ?\d+\]$', input_string[start_index:end_index]):
|
| 370 |
+
scores = json.loads(input_string[start_index:end_index])
|
| 371 |
+
if not isinstance(scores, list):
|
| 372 |
+
scores = [scores]
|
| 373 |
+
json_content = {'score': scores, "reasoning": "System: output is simply a list of scores"}
|
| 374 |
+
json_str = json.dumps(json_content)
|
| 375 |
+
input_string = json_str
|
| 376 |
+
start_index = 0
|
| 377 |
+
end_index = len(json_str)
|
| 378 |
+
elif is_int_between_0_and_10(input_string): # if output is simply a number
|
| 379 |
+
scores = [int(input_string)]
|
| 380 |
+
json_content = {'score': scores, "reasoning": "System: output is simply a number"}
|
| 381 |
+
json_str = json.dumps(json_content)
|
| 382 |
+
input_string = json_str
|
| 383 |
+
start_index = 0
|
| 384 |
+
end_index = len(json_str)
|
| 385 |
+
else:
|
| 386 |
+
print(f"22 222 Failed to find the json content in the string. {text_prompt=} {input_string=}", flush=True)
|
| 387 |
+
return False
|
| 388 |
+
|
| 389 |
+
# Check if we found two delimiters
|
| 390 |
+
if start_index != -1 and end_index != -1 and start_index != end_index:
|
| 391 |
+
# Extract the JSON string
|
| 392 |
+
json_str = input_string[start_index:end_index].strip()
|
| 393 |
+
json_str = json_str.replace("\n", "")
|
| 394 |
+
# Parse the JSON string into a dictionary
|
| 395 |
+
try:
|
| 396 |
+
json_str = normalize_quotes(json_str)
|
| 397 |
+
new_data = json.loads(json_str)
|
| 398 |
+
if not isinstance(new_data['score'], list):
|
| 399 |
+
new_data['score'] = [new_data['score']]
|
| 400 |
+
except Exception as e1:
|
| 401 |
+
print(f"Now fixing: {e1=} {json_str=}")
|
| 402 |
+
|
| 403 |
+
new_data = robust_json_fix(json_str)
|
| 404 |
+
return new_data
|
| 405 |
+
else:
|
| 406 |
+
print("The required delimiters were not found correctly in the string.")
|
| 407 |
+
return False
|
| 408 |
+
|
| 409 |
+
def write_entry_to_json_file(input_string, uid, prompt_input, vision_input, output_file_name, give_up_parsing=False):
|
| 410 |
+
"""
|
| 411 |
+
Args:
|
| 412 |
+
input_string (str): actually the output of the mllm model to be parsed
|
| 413 |
+
uid (str): The unique identifier for the each item in the test data
|
| 414 |
+
prompt_input (str): The prompt input for the entry. text prompt.
|
| 415 |
+
vision_input (str): The vision input for the entry. image links.
|
| 416 |
+
output_file_name (str): The name of the output file.
|
| 417 |
+
"""
|
| 418 |
+
# Catch for gpt4v rate_limit_exceeded error
|
| 419 |
+
if input_string == "rate_limit_exceeded":
|
| 420 |
+
return "rate_limit_exceeded"
|
| 421 |
+
|
| 422 |
+
# Define the delimiters
|
| 423 |
+
delimiter = '||V^=^V||'
|
| 424 |
+
|
| 425 |
+
if input_string.count(delimiter) == 2:
|
| 426 |
+
if not verify(input_string, delimiter):
|
| 427 |
+
print("The required delimiters were not found correctly in the string.")
|
| 428 |
+
return False
|
| 429 |
+
# Extract the content between the delimiters
|
| 430 |
+
start_index = input_string.find(delimiter) + len(delimiter)
|
| 431 |
+
end_index = input_string.rfind(delimiter)
|
| 432 |
+
else:
|
| 433 |
+
# find the json mannually
|
| 434 |
+
# some mllm tends not to output the delimiters, but it does output the json contents
|
| 435 |
+
# so we will find the json content mannually
|
| 436 |
+
start_index = input_string.find('{')
|
| 437 |
+
end_index = input_string.rfind('}') + 1
|
| 438 |
+
if start_index == -1 or end_index == 0:
|
| 439 |
+
# json not found
|
| 440 |
+
# some mllm tends to output only a list of scores like [6, 0],
|
| 441 |
+
# this time we will just get the scores and ignore the reasoning (other part of the json)
|
| 442 |
+
start_index = input_string.find('[')
|
| 443 |
+
end_index = input_string.rfind(']') + 1
|
| 444 |
+
if give_up_parsing: # if we want to give up parsing
|
| 445 |
+
guessed_value = random.randint(0, 10)
|
| 446 |
+
print(f"Failed to find the json content in the string. Guess a value : {guessed_value}.")
|
| 447 |
+
json_content = {'score': [guessed_value], "reasoning": f"guess_if_cannot_parse | {input_string}"}
|
| 448 |
+
json_str = json.dumps(json_content)
|
| 449 |
+
input_string = json_str
|
| 450 |
+
start_index = 0
|
| 451 |
+
end_index = len(json_str)
|
| 452 |
+
elif re.match(r'^\[\d+, ?\d+\]$', input_string[start_index:end_index]):
|
| 453 |
+
scores = json.loads(input_string[start_index:end_index])
|
| 454 |
+
json_content = {'score': scores, "reasoning": None}
|
| 455 |
+
json_str = json.dumps(json_content)
|
| 456 |
+
input_string = json_str
|
| 457 |
+
start_index = 0
|
| 458 |
+
end_index = len(json_str)
|
| 459 |
+
elif is_int_between_0_and_10(input_string): # if output is simply a number
|
| 460 |
+
scores = [int(input_string)]
|
| 461 |
+
json_content = {'score': scores, "reasoning": None}
|
| 462 |
+
json_str = json.dumps(json_content)
|
| 463 |
+
input_string = json_str
|
| 464 |
+
start_index = 0
|
| 465 |
+
end_index = len(json_str)
|
| 466 |
+
else:
|
| 467 |
+
print("Failed to find the json content in the string.")
|
| 468 |
+
return False
|
| 469 |
+
|
| 470 |
+
# Check if we found two delimiters
|
| 471 |
+
if start_index != -1 and end_index != -1 and start_index != end_index:
|
| 472 |
+
# Extract the JSON string
|
| 473 |
+
json_str = input_string[start_index:end_index].strip()
|
| 474 |
+
json_str = json_str.replace("\n", "")
|
| 475 |
+
try:
|
| 476 |
+
# Parse the JSON string into a dictionary
|
| 477 |
+
new_data = json.loads(json_str)
|
| 478 |
+
|
| 479 |
+
# Ensure the directory exists
|
| 480 |
+
os.makedirs(os.path.dirname(output_file_name), exist_ok=True)
|
| 481 |
+
|
| 482 |
+
# Initialize or load existing data
|
| 483 |
+
if os.path.exists(output_file_name):
|
| 484 |
+
with open(output_file_name, 'r') as json_file:
|
| 485 |
+
data = json.load(json_file)
|
| 486 |
+
else:
|
| 487 |
+
data = {}
|
| 488 |
+
|
| 489 |
+
# If the additional key is already in the data, add or update notes
|
| 490 |
+
if uid in data:
|
| 491 |
+
data[uid].update(new_data) # Update with new data
|
| 492 |
+
if prompt_input: # If there are new notes, update or add them
|
| 493 |
+
data[uid]['prompt_input'] = prompt_input
|
| 494 |
+
if vision_input: # If there are new notes, update or add them
|
| 495 |
+
data[uid]['vision_input'] = vision_input
|
| 496 |
+
else:
|
| 497 |
+
# If it's a new key, add the entry to the dictionary
|
| 498 |
+
data[uid] = new_data
|
| 499 |
+
if prompt_input:
|
| 500 |
+
data[uid]['prompt_input'] = prompt_input
|
| 501 |
+
if vision_input:
|
| 502 |
+
data[uid]['vision_input'] = vision_input
|
| 503 |
+
|
| 504 |
+
# Write the updated data to the file
|
| 505 |
+
with open(output_file_name, 'w') as json_file:
|
| 506 |
+
json.dump(data, json_file, indent=4)
|
| 507 |
+
|
| 508 |
+
print(f"Data was successfully updated in {output_file_name}")
|
| 509 |
+
return True
|
| 510 |
+
except json.JSONDecodeError as e:
|
| 511 |
+
print(f"An error occurred while parsing the JSON content: {e}")
|
| 512 |
+
return False
|
| 513 |
+
else:
|
| 514 |
+
print("The required delimiters were not found correctly in the string.")
|
| 515 |
+
return False
|
| 516 |
+
|
| 517 |
+
|
| 518 |
+
def check_key_in_json(file_path, key):
|
| 519 |
+
try:
|
| 520 |
+
with open(file_path, 'r') as json_file:
|
| 521 |
+
data = json.load(json_file)
|
| 522 |
+
|
| 523 |
+
# Check if the key exists at the top level of the JSON structure
|
| 524 |
+
if key in data:
|
| 525 |
+
return True
|
| 526 |
+
else:
|
| 527 |
+
return False
|
| 528 |
+
except FileNotFoundError:
|
| 529 |
+
print(f"The file {file_path} was not found.")
|
| 530 |
+
except json.JSONDecodeError as e:
|
| 531 |
+
print(f"Error reading {file_path}: {e}")
|
| 532 |
+
except Exception as e:
|
| 533 |
+
print(f"An error occurred with {file_path}: {e}")
|
| 534 |
+
return False
|
editscore/vie_prompts.py
ADDED
|
@@ -0,0 +1,46 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# This file is generated automatically through parse_prompt.py
|
| 2 |
+
_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.
|
| 3 |
+
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.
|
| 4 |
+
|
| 5 |
+
IMPORTANT: You will have to give your output in this way (Keep your reasoning concise and short.):
|
| 6 |
+
{
|
| 7 |
+
"reasoning" : "...",
|
| 8 |
+
"score" : [...]
|
| 9 |
+
}
|
| 10 |
+
"""
|
| 11 |
+
|
| 12 |
+
_prompts_0shot_two_image_edit_rule = """RULES:
|
| 13 |
+
|
| 14 |
+
Two images will be provided: The first being the original AI-generated image and the second being an edited version of the first.
|
| 15 |
+
The objective is to evaluate how successfully the editing instruction has been executed in the second image.
|
| 16 |
+
|
| 17 |
+
Note that sometimes the two images might look identical due to the failure of image edit.
|
| 18 |
+
"""
|
| 19 |
+
|
| 20 |
+
_prompts_0shot_tie_rule_SC = """
|
| 21 |
+
From scale 0 to 10:
|
| 22 |
+
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.)
|
| 23 |
+
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.)
|
| 24 |
+
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.
|
| 25 |
+
|
| 26 |
+
Editing instruction: <instruction>
|
| 27 |
+
"""
|
| 28 |
+
|
| 29 |
+
_prompts_0shot_rule_PQ = """RULES:
|
| 30 |
+
|
| 31 |
+
The image is an AI-generated image.
|
| 32 |
+
The objective is to evaluate how successfully the image has been generated.
|
| 33 |
+
|
| 34 |
+
From scale 0 to 10:
|
| 35 |
+
A score from 0 to 10 will be given based on image naturalness.
|
| 36 |
+
(
|
| 37 |
+
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.
|
| 38 |
+
10 indicates that the image looks natural.
|
| 39 |
+
)
|
| 40 |
+
A second score from 0 to 10 will rate the image artifacts.
|
| 41 |
+
(
|
| 42 |
+
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.
|
| 43 |
+
10 indicates the image has no artifacts.
|
| 44 |
+
)
|
| 45 |
+
Put the score in a list such that output score = [naturalness, artifacts]
|
| 46 |
+
"""
|
evaluate.sh
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# !/bin/bash
|
| 2 |
+
SHELL_FOLDER=$(cd "$(dirname "$0")";pwd)
|
| 3 |
+
cd $SHELL_FOLDER
|
| 4 |
+
|
| 5 |
+
python evaluation.py \
|
| 6 |
+
--benchmark_dir EditScore/EditReward-Bench \
|
| 7 |
+
--result_dir results/EditScore-7B \
|
| 8 |
+
--backbone qwen25vl \
|
| 9 |
+
--model_name_or_path Qwen/Qwen2.5-VL-7B-Instruct \
|
| 10 |
+
--lora_path EditScore/EditScore-7B \
|
| 11 |
+
--score_range 25 \
|
| 12 |
+
--max_workers 1 \
|
| 13 |
+
--max_model_len 4096 \
|
| 14 |
+
--max_num_seqs 1 \
|
| 15 |
+
--max_num_batched_tokens 4096 \
|
| 16 |
+
--tensor_parallel_size 1 \
|
| 17 |
+
--num_pass 1
|
| 18 |
+
|
| 19 |
+
python calculate_statistics.py \
|
| 20 |
+
--result_dir results/EditScore-7B/qwen25vl
|
evaluate_72B_vllm.sh
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# !/bin/bash
|
| 2 |
+
SHELL_FOLDER=$(cd "$(dirname "$0")";pwd)
|
| 3 |
+
cd $SHELL_FOLDER
|
| 4 |
+
|
| 5 |
+
python evaluation.py \
|
| 6 |
+
--benchmark_dir EditScore/EditReward-Bench \
|
| 7 |
+
--result_dir results/EditScore-72B \
|
| 8 |
+
--backbone qwen25vl_vllm \
|
| 9 |
+
--model_name_or_path Qwen/Qwen2.5-VL-72B-Instruct \
|
| 10 |
+
--lora_path EditScore/EditScore-72B \
|
| 11 |
+
--score_range 25 \
|
| 12 |
+
--max_workers 1 \
|
| 13 |
+
--max_model_len 4096 \
|
| 14 |
+
--max_num_seqs 1 \
|
| 15 |
+
--max_num_batched_tokens 4096 \
|
| 16 |
+
--tensor_parallel_size 4 \
|
| 17 |
+
--num_pass 1
|
| 18 |
+
|
| 19 |
+
python calculate_statistics.py \
|
| 20 |
+
--result_dir results/EditScore-72B/qwen25vl_vllm
|
evaluate_qwen3_vl_32B.sh
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# !/bin/bash
|
| 2 |
+
SHELL_FOLDER=$(cd "$(dirname "$0")";pwd)
|
| 3 |
+
cd $SHELL_FOLDER
|
| 4 |
+
|
| 5 |
+
python evaluation.py \
|
| 6 |
+
--benchmark_dir EditScore/EditReward-Bench \
|
| 7 |
+
--result_dir results/EditScore-Qwen3-VL-32B \
|
| 8 |
+
--backbone qwen3vl \
|
| 9 |
+
--model_name_or_path Qwen/Qwen3-VL-32B-Instruct \
|
| 10 |
+
--lora_path EditScore/EditScore-Qwen3-VL-32B-Instruct \
|
| 11 |
+
--score_range 25 \
|
| 12 |
+
--max_workers 1 \
|
| 13 |
+
--max_model_len 4096 \
|
| 14 |
+
--max_num_seqs 1 \
|
| 15 |
+
--max_num_batched_tokens 4096 \
|
| 16 |
+
--tensor_parallel_size 1 \
|
| 17 |
+
--num_pass 1
|
| 18 |
+
|
| 19 |
+
python calculate_statistics.py \
|
| 20 |
+
--result_dir results/EditScore-Qwen3-VL-32B/qwen3vl
|
evaluate_qwen3_vl_32B_avg4.sh
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# !/bin/bash
|
| 2 |
+
SHELL_FOLDER=$(cd "$(dirname "$0")";pwd)
|
| 3 |
+
cd $SHELL_FOLDER
|
| 4 |
+
|
| 5 |
+
python evaluation.py \
|
| 6 |
+
--benchmark_dir EditScore/EditReward-Bench \
|
| 7 |
+
--result_dir results/EditScore-Qwen3-VL-32B-avg4 \
|
| 8 |
+
--backbone qwen3vl \
|
| 9 |
+
--model_name_or_path Qwen/Qwen3-VL-32B-Instruct \
|
| 10 |
+
--lora_path EditScore/EditScore-Qwen3-VL-32B-Instruct \
|
| 11 |
+
--score_range 25 \
|
| 12 |
+
--max_workers 1 \
|
| 13 |
+
--max_model_len 4096 \
|
| 14 |
+
--max_num_seqs 1 \
|
| 15 |
+
--max_num_batched_tokens 4096 \
|
| 16 |
+
--tensor_parallel_size 1 \
|
| 17 |
+
--num_pass 4
|
| 18 |
+
|
| 19 |
+
python calculate_statistics.py \
|
| 20 |
+
--result_dir results/EditScore-Qwen3-VL-32B-avg4/qwen3vl
|
evaluate_qwen3_vl_4B.sh
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# !/bin/bash
|
| 2 |
+
SHELL_FOLDER=$(cd "$(dirname "$0")";pwd)
|
| 3 |
+
cd $SHELL_FOLDER
|
| 4 |
+
|
| 5 |
+
python evaluation.py \
|
| 6 |
+
--benchmark_dir EditScore/EditReward-Bench \
|
| 7 |
+
--result_dir results/EditScore-Qwen3-VL-4B \
|
| 8 |
+
--backbone qwen3vl \
|
| 9 |
+
--model_name_or_path Qwen/Qwen3-VL-4B-Instruct \
|
| 10 |
+
--lora_path EditScore/EditScore-Qwen3-VL-4B-Instruct \
|
| 11 |
+
--score_range 25 \
|
| 12 |
+
--max_workers 1 \
|
| 13 |
+
--max_model_len 4096 \
|
| 14 |
+
--max_num_seqs 1 \
|
| 15 |
+
--max_num_batched_tokens 4096 \
|
| 16 |
+
--tensor_parallel_size 1 \
|
| 17 |
+
--num_pass 1
|
| 18 |
+
|
| 19 |
+
python calculate_statistics.py \
|
| 20 |
+
--result_dir results/EditScore-Qwen3-VL-4B/qwen3vl
|
evaluate_qwen3_vl_4B_avg4.sh
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# !/bin/bash
|
| 2 |
+
SHELL_FOLDER=$(cd "$(dirname "$0")";pwd)
|
| 3 |
+
cd $SHELL_FOLDER
|
| 4 |
+
|
| 5 |
+
python evaluation.py \
|
| 6 |
+
--benchmark_dir EditScore/EditReward-Bench \
|
| 7 |
+
--result_dir results/EditScore-Qwen3-VL-4B-avg4 \
|
| 8 |
+
--backbone qwen3vl \
|
| 9 |
+
--model_name_or_path Qwen/Qwen3-VL-4B-Instruct \
|
| 10 |
+
--lora_path EditScore/EditScore-Qwen3-VL-4B-Instruct \
|
| 11 |
+
--score_range 25 \
|
| 12 |
+
--max_workers 1 \
|
| 13 |
+
--max_model_len 4096 \
|
| 14 |
+
--max_num_seqs 1 \
|
| 15 |
+
--max_num_batched_tokens 4096 \
|
| 16 |
+
--tensor_parallel_size 1 \
|
| 17 |
+
--num_pass 4
|
| 18 |
+
|
| 19 |
+
python calculate_statistics.py \
|
| 20 |
+
--result_dir results/EditScore-Qwen3-VL-4B-avg4/qwen3vl
|
evaluate_qwen3_vl_4B_vllm.sh
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# !/bin/bash
|
| 2 |
+
SHELL_FOLDER=$(cd "$(dirname "$0")";pwd)
|
| 3 |
+
cd $SHELL_FOLDER
|
| 4 |
+
|
| 5 |
+
python evaluation.py \
|
| 6 |
+
--benchmark_dir EditScore/EditReward-Bench \
|
| 7 |
+
--result_dir results/EditScore-Qwen3-VL-4B \
|
| 8 |
+
--backbone qwen3vl_vllm \
|
| 9 |
+
--model_name_or_path Qwen/Qwen3-VL-4B-Instruct \
|
| 10 |
+
--lora_path EditScore/EditScore-Qwen3-VL-4B-Instruct \
|
| 11 |
+
--score_range 25 \
|
| 12 |
+
--max_workers 1 \
|
| 13 |
+
--max_model_len 4096 \
|
| 14 |
+
--max_num_seqs 1 \
|
| 15 |
+
--max_num_batched_tokens 4096 \
|
| 16 |
+
--tensor_parallel_size 1 \
|
| 17 |
+
--num_pass 1
|
| 18 |
+
|
| 19 |
+
python calculate_statistics.py \
|
| 20 |
+
--result_dir results/EditScore-Qwen3-VL-4B/qwen3vl_vllm
|
evaluate_qwen3_vl_8B.sh
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# !/bin/bash
|
| 2 |
+
SHELL_FOLDER=$(cd "$(dirname "$0")";pwd)
|
| 3 |
+
cd $SHELL_FOLDER
|
| 4 |
+
|
| 5 |
+
python evaluation.py \
|
| 6 |
+
--benchmark_dir EditScore/EditReward-Bench \
|
| 7 |
+
--result_dir results/EditScore-Qwen3-VL-8B \
|
| 8 |
+
--backbone qwen3vl \
|
| 9 |
+
--model_name_or_path Qwen/Qwen3-VL-8B-Instruct \
|
| 10 |
+
--lora_path EditScore/EditScore-Qwen3-VL-8B-Instruct \
|
| 11 |
+
--score_range 25 \
|
| 12 |
+
--max_workers 1 \
|
| 13 |
+
--max_model_len 4096 \
|
| 14 |
+
--max_num_seqs 1 \
|
| 15 |
+
--max_num_batched_tokens 4096 \
|
| 16 |
+
--tensor_parallel_size 1 \
|
| 17 |
+
--num_pass 1
|
| 18 |
+
|
| 19 |
+
python calculate_statistics.py \
|
| 20 |
+
--result_dir results/EditScore-Qwen3-VL-8B/qwen3vl
|
evaluate_qwen3_vl_8B_avg4.sh
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# !/bin/bash
|
| 2 |
+
SHELL_FOLDER=$(cd "$(dirname "$0")";pwd)
|
| 3 |
+
cd $SHELL_FOLDER
|
| 4 |
+
|
| 5 |
+
python evaluation.py \
|
| 6 |
+
--benchmark_dir EditScore/EditReward-Bench \
|
| 7 |
+
--result_dir results/EditScore-Qwen3-VL-8B-avg4 \
|
| 8 |
+
--backbone qwen3vl \
|
| 9 |
+
--model_name_or_path Qwen/Qwen3-VL-8B-Instruct \
|
| 10 |
+
--lora_path EditScore/EditScore-Qwen3-VL-8B-Instruct \
|
| 11 |
+
--score_range 25 \
|
| 12 |
+
--max_workers 1 \
|
| 13 |
+
--max_model_len 4096 \
|
| 14 |
+
--max_num_seqs 1 \
|
| 15 |
+
--max_num_batched_tokens 4096 \
|
| 16 |
+
--tensor_parallel_size 1 \
|
| 17 |
+
--num_pass 4
|
| 18 |
+
|
| 19 |
+
python calculate_statistics.py \
|
| 20 |
+
--result_dir results/EditScore-Qwen3-VL-8B-avg4/qwen3vl
|
evaluate_vllm.sh
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# !/bin/bash
|
| 2 |
+
SHELL_FOLDER=$(cd "$(dirname "$0")";pwd)
|
| 3 |
+
cd $SHELL_FOLDER
|
| 4 |
+
|
| 5 |
+
python evaluation.py \
|
| 6 |
+
--benchmark_dir EditScore/EditReward-Bench \
|
| 7 |
+
--result_dir results/EditScore-7B \
|
| 8 |
+
--backbone qwen25vl_vllm \
|
| 9 |
+
--model_name_or_path Qwen/Qwen2.5-VL-7B-Instruct \
|
| 10 |
+
--lora_path EditScore/EditScore-7B \
|
| 11 |
+
--score_range 25 \
|
| 12 |
+
--max_workers 1 \
|
| 13 |
+
--max_model_len 4096 \
|
| 14 |
+
--max_num_seqs 1 \
|
| 15 |
+
--max_num_batched_tokens 4096 \
|
| 16 |
+
--tensor_parallel_size 1 \
|
| 17 |
+
--num_pass 1
|
| 18 |
+
|
| 19 |
+
python calculate_statistics.py \
|
| 20 |
+
--result_dir results/EditScore-7B/qwen25vl_vllm
|
evaluation.py
ADDED
|
@@ -0,0 +1,272 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import dotenv
|
| 2 |
+
|
| 3 |
+
dotenv.load_dotenv(override=True)
|
| 4 |
+
|
| 5 |
+
import argparse
|
| 6 |
+
import glob
|
| 7 |
+
import hashlib
|
| 8 |
+
import json
|
| 9 |
+
import logging
|
| 10 |
+
import os
|
| 11 |
+
import time
|
| 12 |
+
import threading
|
| 13 |
+
from concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor, as_completed
|
| 14 |
+
from typing import List, Tuple, Dict, Any, Optional
|
| 15 |
+
|
| 16 |
+
import dotenv
|
| 17 |
+
from PIL import Image
|
| 18 |
+
from tqdm import tqdm
|
| 19 |
+
from datasets import Dataset, load_dataset
|
| 20 |
+
|
| 21 |
+
from editscore import EditScore
|
| 22 |
+
|
| 23 |
+
PROMPT_FOLLOWING = "prompt_following"
|
| 24 |
+
CONSISTENCY = "consistency"
|
| 25 |
+
OVERALL = "overall"
|
| 26 |
+
SCORE_CATEGORIES = [PROMPT_FOLLOWING, CONSISTENCY, OVERALL]
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
class CacheManager:
|
| 30 |
+
def __init__(self, cache_file: str):
|
| 31 |
+
self.cache_file = cache_file
|
| 32 |
+
self.lock = threading.Lock()
|
| 33 |
+
self.cache = self._load()
|
| 34 |
+
|
| 35 |
+
def _load(self) -> Dict[str, Any]:
|
| 36 |
+
cache = {}
|
| 37 |
+
if not os.path.exists(self.cache_file):
|
| 38 |
+
print(
|
| 39 |
+
f"Cache file not found at {self.cache_file}. A new one will be created."
|
| 40 |
+
)
|
| 41 |
+
return cache
|
| 42 |
+
|
| 43 |
+
with open(self.cache_file, "r", encoding="utf-8") as f:
|
| 44 |
+
for i, line in enumerate(f):
|
| 45 |
+
try:
|
| 46 |
+
data = json.loads(line)
|
| 47 |
+
cache[data["key"]] = data["result"]
|
| 48 |
+
except json.JSONDecodeError:
|
| 49 |
+
logging.warning(
|
| 50 |
+
f"Skipping corrupted line {i + 1} in cache file: {line.strip()}"
|
| 51 |
+
)
|
| 52 |
+
print(f"Loaded {len(cache)} items from {self.cache_file}.")
|
| 53 |
+
return cache
|
| 54 |
+
|
| 55 |
+
def get(self, key: str) -> Optional[Any]:
|
| 56 |
+
return self.cache.get(key)
|
| 57 |
+
|
| 58 |
+
def append(self, key: str, result: Any):
|
| 59 |
+
with self.lock:
|
| 60 |
+
self.cache[key] = result
|
| 61 |
+
with open(self.cache_file, "a", encoding="utf-8") as f:
|
| 62 |
+
f.write(
|
| 63 |
+
json.dumps({"key": key, "result": result}, ensure_ascii=False)
|
| 64 |
+
+ "\n"
|
| 65 |
+
)
|
| 66 |
+
|
| 67 |
+
def generate_cache_key(pair_key):
|
| 68 |
+
return hashlib.sha256(pair_key.encode("utf-8")).hexdigest()
|
| 69 |
+
|
| 70 |
+
def load_pairs_dataset(dataset: Dataset) -> Dict[str, Tuple[str, Image.Image, Image.Image]]:
|
| 71 |
+
pairs = {}
|
| 72 |
+
for data in dataset:
|
| 73 |
+
key1, key2 = data["key"]
|
| 74 |
+
instruction = data["instruction"]
|
| 75 |
+
input_image = data["input_image"].convert("RGB")
|
| 76 |
+
|
| 77 |
+
pairs[key1] = (instruction, input_image, data["output_images"][0].convert("RGB"))
|
| 78 |
+
pairs[key2] = (instruction, input_image, data["output_images"][1].convert("RGB"))
|
| 79 |
+
return pairs
|
| 80 |
+
|
| 81 |
+
def _load_item(data: dict) -> list[tuple[str, tuple[str, Image.Image, Image.Image]]]:
|
| 82 |
+
key1, key2 = data["key"]
|
| 83 |
+
instruction = data["instruction"]
|
| 84 |
+
|
| 85 |
+
input_image = data["input_image"].convert("RGB")
|
| 86 |
+
output_image1 = data["output_images"][0].convert("RGB")
|
| 87 |
+
output_image2 = data["output_images"][1].convert("RGB")
|
| 88 |
+
|
| 89 |
+
return [
|
| 90 |
+
(key1, (instruction, input_image, output_image1)),
|
| 91 |
+
(key2, (instruction, input_image, output_image2)),
|
| 92 |
+
]
|
| 93 |
+
|
| 94 |
+
def load_pairs_dataset_multithreaded(dataset: Dataset, max_workers: int = None) -> Dict[str, Tuple[str, Image.Image, Image.Image]]:
|
| 95 |
+
if max_workers is None:
|
| 96 |
+
# max_workers = min(32, (os.cpu_count() or 1) * 5)
|
| 97 |
+
max_workers = os.cpu_count() or 1
|
| 98 |
+
|
| 99 |
+
pairs = {}
|
| 100 |
+
|
| 101 |
+
print(f"Processing dataset (length: {len(dataset)}) with {max_workers} threads", flush=True)
|
| 102 |
+
|
| 103 |
+
with ProcessPoolExecutor(max_workers=max_workers) as executor:
|
| 104 |
+
results_iterator = tqdm(
|
| 105 |
+
executor.map(_load_item, dataset),
|
| 106 |
+
total=len(dataset),
|
| 107 |
+
desc="Processing dataset with multiple threads"
|
| 108 |
+
)
|
| 109 |
+
|
| 110 |
+
for result_pairs in results_iterator:
|
| 111 |
+
pairs.update(result_pairs)
|
| 112 |
+
|
| 113 |
+
return pairs
|
| 114 |
+
|
| 115 |
+
|
| 116 |
+
def process_single_item(key, item, scorer):
|
| 117 |
+
instruction = item[0]
|
| 118 |
+
input_image = item[1]
|
| 119 |
+
output_image = item[2]
|
| 120 |
+
|
| 121 |
+
output_image = output_image.resize((input_image.size[0], input_image.size[1]))
|
| 122 |
+
|
| 123 |
+
score = scorer.evaluate([input_image, output_image], instruction)
|
| 124 |
+
return key, score
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
def parse_args():
|
| 128 |
+
parser = argparse.ArgumentParser()
|
| 129 |
+
parser.add_argument(
|
| 130 |
+
"--benchmark_dir", type=str, default="EditScore/EditReward-Bench"
|
| 131 |
+
)
|
| 132 |
+
parser.add_argument("--result_dir", type=str, required=True)
|
| 133 |
+
parser.add_argument(
|
| 134 |
+
"--backbone",
|
| 135 |
+
type=str,
|
| 136 |
+
default="openai",
|
| 137 |
+
choices=["openai", "qwen25vl", "qwen25vl_vllm", "internvl3_5", "qwen3vl", "qwen3vl_vllm"],
|
| 138 |
+
)
|
| 139 |
+
parser.add_argument("--model_name_or_path", type=str, default="gpt-4.1")
|
| 140 |
+
parser.add_argument(
|
| 141 |
+
"--openai_url", type=str, default="https://api.openai.com/v1/chat/completions"
|
| 142 |
+
)
|
| 143 |
+
parser.add_argument("--key", type=str, default="PUT YOUR API KEY HERE")
|
| 144 |
+
parser.add_argument("--num_pass", type=int, default=1)
|
| 145 |
+
parser.add_argument("--temperature", type=float, default=0.7)
|
| 146 |
+
parser.add_argument("--max_workers", type=int, default=20)
|
| 147 |
+
parser.add_argument("--score_range", type=int, default=25)
|
| 148 |
+
parser.add_argument("--tensor_parallel_size", type=int, default=1)
|
| 149 |
+
parser.add_argument("--max_model_len", type=int, default=1536)
|
| 150 |
+
parser.add_argument("--max_num_seqs", type=int, default=32)
|
| 151 |
+
parser.add_argument("--max_num_batched_tokens", type=int, default=1536)
|
| 152 |
+
parser.add_argument("--lora_path", type=str, default="EditScore/EditScore-7B")
|
| 153 |
+
parser.add_argument("--cache_dir", type=str, default=None)
|
| 154 |
+
return parser.parse_args()
|
| 155 |
+
|
| 156 |
+
|
| 157 |
+
def main(args):
|
| 158 |
+
start_time = time.time()
|
| 159 |
+
scorer = EditScore(
|
| 160 |
+
backbone=args.backbone,
|
| 161 |
+
key=args.key,
|
| 162 |
+
openai_url=args.openai_url,
|
| 163 |
+
model_name_or_path=args.model_name_or_path,
|
| 164 |
+
score_range=args.score_range,
|
| 165 |
+
temperature=args.temperature,
|
| 166 |
+
tensor_parallel_size=args.tensor_parallel_size,
|
| 167 |
+
max_model_len=args.max_model_len,
|
| 168 |
+
max_num_seqs=args.max_num_seqs,
|
| 169 |
+
max_num_batched_tokens=args.max_num_batched_tokens,
|
| 170 |
+
num_pass=args.num_pass,
|
| 171 |
+
lora_path=args.lora_path,
|
| 172 |
+
cache_dir=args.cache_dir,
|
| 173 |
+
)
|
| 174 |
+
print(f"Scorer initialized in {time.time() - start_time} seconds", flush=True)
|
| 175 |
+
|
| 176 |
+
cache_dir = os.path.join(args.result_dir, ".cache")
|
| 177 |
+
os.makedirs(cache_dir, exist_ok=True)
|
| 178 |
+
cache_file = os.path.join(
|
| 179 |
+
cache_dir, f"{args.backbone}_{args.model_name_or_path.replace('/', '_')}.jsonl"
|
| 180 |
+
)
|
| 181 |
+
cache_manager = CacheManager(cache_file)
|
| 182 |
+
|
| 183 |
+
start_time = time.time()
|
| 184 |
+
dataset = load_dataset(args.benchmark_dir, split="train")
|
| 185 |
+
print(f"Dataset loaded in {time.time() - start_time} seconds", flush=True)
|
| 186 |
+
|
| 187 |
+
start_time = time.time()
|
| 188 |
+
unique_pairs = load_pairs_dataset_multithreaded(dataset)
|
| 189 |
+
print(f"Pairs loaded in {time.time() - start_time} seconds", flush=True)
|
| 190 |
+
|
| 191 |
+
all_scores = {}
|
| 192 |
+
pairs_to_process = [
|
| 193 |
+
pair_key
|
| 194 |
+
for pair_key in unique_pairs.keys()
|
| 195 |
+
if cache_manager.get(generate_cache_key(pair_key)) is None
|
| 196 |
+
]
|
| 197 |
+
|
| 198 |
+
for pair_key in unique_pairs.keys():
|
| 199 |
+
if pair_key not in pairs_to_process:
|
| 200 |
+
all_scores[pair_key] = cache_manager.get(generate_cache_key(pair_key))
|
| 201 |
+
|
| 202 |
+
print(
|
| 203 |
+
f"{len(unique_pairs) - len(pairs_to_process)} pairs found in cache. Processing {len(pairs_to_process)} new pairs.",
|
| 204 |
+
flush=True
|
| 205 |
+
)
|
| 206 |
+
|
| 207 |
+
if pairs_to_process:
|
| 208 |
+
with ThreadPoolExecutor(max_workers=args.max_workers) as executor:
|
| 209 |
+
futures = [
|
| 210 |
+
executor.submit(process_single_item, pair_key, unique_pairs[pair_key], scorer)
|
| 211 |
+
for pair_key in pairs_to_process
|
| 212 |
+
]
|
| 213 |
+
|
| 214 |
+
for future in tqdm(
|
| 215 |
+
as_completed(futures),
|
| 216 |
+
total=len(futures),
|
| 217 |
+
unit="pair",
|
| 218 |
+
desc="Processing",
|
| 219 |
+
):
|
| 220 |
+
pair_key, result = future.result()
|
| 221 |
+
if result:
|
| 222 |
+
all_scores[pair_key] = result
|
| 223 |
+
cache_manager.append(generate_cache_key(pair_key), result)
|
| 224 |
+
|
| 225 |
+
print("Writing results...", flush=True)
|
| 226 |
+
|
| 227 |
+
start_time = time.time()
|
| 228 |
+
# dataset = dataset.remove_columns(["input_image", "output_images"])
|
| 229 |
+
for idx, data in enumerate(dataset):
|
| 230 |
+
key1, key2 = data["key"]
|
| 231 |
+
task_type = data["task_type"]
|
| 232 |
+
dimension = data["dimension"]
|
| 233 |
+
|
| 234 |
+
score1 = all_scores[key1][dimension]
|
| 235 |
+
score2 = all_scores[key2][dimension]
|
| 236 |
+
data["score"] = [score1, score2]
|
| 237 |
+
|
| 238 |
+
input_image_path = os.path.join(args.result_dir, "images", f"{key1}_input.png")
|
| 239 |
+
output_image_path1 = os.path.join(args.result_dir, "images", f"{key1}.png")
|
| 240 |
+
output_image_path2 = os.path.join(args.result_dir, "images", f"{key2}.png")
|
| 241 |
+
|
| 242 |
+
os.makedirs(os.path.dirname(input_image_path), exist_ok=True)
|
| 243 |
+
|
| 244 |
+
data['input_image'].save(input_image_path)
|
| 245 |
+
data['output_images'][0].save(output_image_path1)
|
| 246 |
+
data['output_images'][1].save(output_image_path2)
|
| 247 |
+
|
| 248 |
+
json_line = {
|
| 249 |
+
"key": (key1, key2),
|
| 250 |
+
"idx": idx,
|
| 251 |
+
"score": [score1, score2],
|
| 252 |
+
"SC_reasoning": [all_scores[key1]["SC_reasoning"], all_scores[key2]["SC_reasoning"]],
|
| 253 |
+
"PQ_reasoning": [all_scores[key1]["PQ_reasoning"], all_scores[key2]["PQ_reasoning"]],
|
| 254 |
+
"input_image": input_image_path,
|
| 255 |
+
"output_images": [output_image_path1, output_image_path2],
|
| 256 |
+
}
|
| 257 |
+
|
| 258 |
+
save_file = os.path.join(
|
| 259 |
+
args.result_dir, args.backbone, task_type, f"{dimension}.jsonl"
|
| 260 |
+
)
|
| 261 |
+
os.makedirs(os.path.dirname(save_file), exist_ok=True)
|
| 262 |
+
|
| 263 |
+
with open(save_file, "a", encoding="utf-8") as f:
|
| 264 |
+
f.write(json.dumps(json_line, ensure_ascii=False) + "\n")
|
| 265 |
+
|
| 266 |
+
print(f"Results written in {time.time() - start_time} seconds", flush=True)
|
| 267 |
+
print("--- Completed! ---", flush=True)
|
| 268 |
+
|
| 269 |
+
|
| 270 |
+
if __name__ == "__main__":
|
| 271 |
+
args = parse_args()
|
| 272 |
+
main(args)
|
example_images/input.png
ADDED
|
Git LFS Details
|
example_images/output.png
ADDED
|
Git LFS Details
|
examples/EditScore-train/README.md
ADDED
|
@@ -0,0 +1,101 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# EditScore Reward Model Training Guide
|
| 2 |
+
|
| 3 |
+
This guide explains how to train EditScore reward models using LLaMA-Factory.
|
| 4 |
+
|
| 5 |
+
## 1. Environment Setup
|
| 6 |
+
|
| 7 |
+
### Clone LLaMA-Factory and Configure Virtual Environment
|
| 8 |
+
|
| 9 |
+
```bash
|
| 10 |
+
git clone --depth 1 https://github.com/hiyouga/LLaMA-Factory.git
|
| 11 |
+
cd LLaMA-Factory
|
| 12 |
+
conda create -n llama-factory python=3.10
|
| 13 |
+
conda activate llama-factory
|
| 14 |
+
pip install -e ".[torch,metrics]" --no-build-isolation
|
| 15 |
+
```
|
| 16 |
+
|
| 17 |
+
## 2. Directory Structure Configuration
|
| 18 |
+
|
| 19 |
+
Create necessary folders and files in the LLaMA-Factory root directory:
|
| 20 |
+
|
| 21 |
+
```bash
|
| 22 |
+
# Create log and output directories
|
| 23 |
+
mkdir -p logs
|
| 24 |
+
mkdir -p output
|
| 25 |
+
|
| 26 |
+
# Create training configuration directory
|
| 27 |
+
mkdir -p examples/train_editscore
|
| 28 |
+
|
| 29 |
+
# Copy training configuration files
|
| 30 |
+
cp EditScore/examples/EditScore-train/config/*.yaml examples/train_editscore/
|
| 31 |
+
|
| 32 |
+
# Copy training script
|
| 33 |
+
cp EditScore/examples/EditScore-train/train.sh .
|
| 34 |
+
```
|
| 35 |
+
|
| 36 |
+
## 3. Dataset Registration
|
| 37 |
+
|
| 38 |
+
Register the EditScore-Reward-Data dataset in `LLaMA-Factory/data/dataset_info.json`:
|
| 39 |
+
|
| 40 |
+
```json
|
| 41 |
+
"EditScore-Reward-Data": {
|
| 42 |
+
"file_name": "/path/to/your/reward.json",
|
| 43 |
+
"formatting": "sharegpt",
|
| 44 |
+
"columns": {
|
| 45 |
+
"messages": "conversations",
|
| 46 |
+
"images": "images"
|
| 47 |
+
}
|
| 48 |
+
}
|
| 49 |
+
```
|
| 50 |
+
|
| 51 |
+
## 4. Training Configuration Description
|
| 52 |
+
|
| 53 |
+
### Single-Machine Training Configuration
|
| 54 |
+
- `editscore_7B.yaml` - Train EditScore-7B model (single machine)
|
| 55 |
+
- `editscore_qwen3_vl_4B_instruct.yaml` - Train EditScore_Qwen3_Vl_4B_Instruct model (single machine)
|
| 56 |
+
- `editscore_qwen3_vl_8B_instruct.yaml` - Train EditScore_Qwen3_Vl_8B_Instruct model (single machine)
|
| 57 |
+
|
| 58 |
+
### Multi-Machine Training Configuration
|
| 59 |
+
- `editscore_32B.yaml` - Train EditScore-32B model (two machines)
|
| 60 |
+
- `editscore_72B.yaml` - Train EditScore-72B model (two machines)
|
| 61 |
+
|
| 62 |
+
## 5. Start Training
|
| 63 |
+
|
| 64 |
+
### Single-Machine Training
|
| 65 |
+
|
| 66 |
+
```bash
|
| 67 |
+
# Modify experiment_name in train.sh to the corresponding configuration file name
|
| 68 |
+
# For example: name=editscore_7B
|
| 69 |
+
bash train.sh
|
| 70 |
+
```
|
| 71 |
+
|
| 72 |
+
### Multi-Machine Training
|
| 73 |
+
|
| 74 |
+
**Master node (rank=0):**
|
| 75 |
+
```bash
|
| 76 |
+
bash train.sh --rank=0 --world_size=2 --master_addr=MASTER_NODE_IP --master_port=29500
|
| 77 |
+
```
|
| 78 |
+
|
| 79 |
+
**Worker node (rank=1):**
|
| 80 |
+
```bash
|
| 81 |
+
bash train.sh --rank=1 --world_size=2 --master_addr=MASTER_NODE_IP --master_port=29500
|
| 82 |
+
```
|
| 83 |
+
|
| 84 |
+
## 6. Parameter Configuration
|
| 85 |
+
|
| 86 |
+
Users can modify the following parameters in the YAML configuration files as needed:
|
| 87 |
+
|
| 88 |
+
- `per_device_train_batch_size`: Batch size per device
|
| 89 |
+
- `gradient_accumulation_steps`: Gradient accumulation steps
|
| 90 |
+
- `learning_rate`: Learning rate
|
| 91 |
+
- `num_train_epochs`: Number of training epochs
|
| 92 |
+
- `max_samples`: Maximum number of samples
|
| 93 |
+
- `output_dir`: Output directory
|
| 94 |
+
|
| 95 |
+
## 7. Output Files
|
| 96 |
+
|
| 97 |
+
After training completion, model files will be saved in the corresponding output directories:
|
| 98 |
+
- Single-machine training: `LLaMA-Factory/output/model_name/`
|
| 99 |
+
- Log files: `LLaMA-Factory/logs/experiment_name_rank.log`
|
| 100 |
+
|
| 101 |
+
|
examples/EditScore-train/config/editscore_32B.yaml
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
model_name_or_path: Qwen/Qwen2.5-VL-32B-Instruct
|
| 2 |
+
image_max_pixels: 262144
|
| 3 |
+
video_max_pixels: 16384
|
| 4 |
+
trust_remote_code: true
|
| 5 |
+
|
| 6 |
+
### method
|
| 7 |
+
stage: sft
|
| 8 |
+
do_train: true
|
| 9 |
+
finetuning_type: lora
|
| 10 |
+
lora_rank: 32
|
| 11 |
+
lora_target: all
|
| 12 |
+
|
| 13 |
+
deepspeed: LLaMA-Factory/examples/deepspeed/ds_z2_config.json
|
| 14 |
+
### dataset
|
| 15 |
+
dataset: EditScore-Reward-Data
|
| 16 |
+
template: qwen2_vl
|
| 17 |
+
cutoff_len: 8192
|
| 18 |
+
max_samples: 100000
|
| 19 |
+
overwrite_cache: true
|
| 20 |
+
preprocessing_num_workers: 16
|
| 21 |
+
dataloader_num_workers: 4
|
| 22 |
+
|
| 23 |
+
### output
|
| 24 |
+
output_dir: LLaMA-Factory/output/editscore_32B/
|
| 25 |
+
logging_steps: 1
|
| 26 |
+
save_steps: 250
|
| 27 |
+
plot_loss: true
|
| 28 |
+
overwrite_output_dir: true
|
| 29 |
+
save_only_model: false
|
| 30 |
+
report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow]
|
| 31 |
+
|
| 32 |
+
### train
|
| 33 |
+
per_device_train_batch_size: 1
|
| 34 |
+
gradient_accumulation_steps: 8
|
| 35 |
+
learning_rate: 1.0e-4
|
| 36 |
+
num_train_epochs: 3.0
|
| 37 |
+
lr_scheduler_type: cosine
|
| 38 |
+
warmup_ratio: 0.1
|
| 39 |
+
bf16: true
|
| 40 |
+
ddp_timeout: 180000000
|
| 41 |
+
resume_from_checkpoint: null
|
| 42 |
+
|
examples/EditScore-train/config/editscore_72B.yaml
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
model_name_or_path: Qwen/Qwen2.5-VL-72B-Instruct
|
| 2 |
+
image_max_pixels: 262144
|
| 3 |
+
video_max_pixels: 16384
|
| 4 |
+
trust_remote_code: true
|
| 5 |
+
|
| 6 |
+
### method
|
| 7 |
+
stage: sft
|
| 8 |
+
do_train: true
|
| 9 |
+
finetuning_type: lora
|
| 10 |
+
lora_rank: 32
|
| 11 |
+
lora_target: all
|
| 12 |
+
|
| 13 |
+
deepspeed: LLaMA-Factory/examples/deepspeed/ds_z3_config.json
|
| 14 |
+
### dataset
|
| 15 |
+
dataset: EditScore-Reward-Data
|
| 16 |
+
template: qwen2_vl
|
| 17 |
+
cutoff_len: 8192
|
| 18 |
+
max_samples: 100000
|
| 19 |
+
overwrite_cache: true
|
| 20 |
+
preprocessing_num_workers: 16
|
| 21 |
+
dataloader_num_workers: 4
|
| 22 |
+
|
| 23 |
+
### output
|
| 24 |
+
output_dir: LLaMA-Factory/output/editscore_72B/
|
| 25 |
+
logging_steps: 1
|
| 26 |
+
save_steps: 250
|
| 27 |
+
plot_loss: true
|
| 28 |
+
overwrite_output_dir: true
|
| 29 |
+
save_only_model: false
|
| 30 |
+
report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow]
|
| 31 |
+
|
| 32 |
+
### train
|
| 33 |
+
per_device_train_batch_size: 1
|
| 34 |
+
gradient_accumulation_steps: 8
|
| 35 |
+
learning_rate: 1.0e-4
|
| 36 |
+
num_train_epochs: 3.0
|
| 37 |
+
lr_scheduler_type: cosine
|
| 38 |
+
warmup_ratio: 0.1
|
| 39 |
+
bf16: true
|
| 40 |
+
ddp_timeout: 180000000
|
| 41 |
+
resume_from_checkpoint: null
|
| 42 |
+
|
examples/EditScore-train/config/editscore_7B.yaml
ADDED
|
@@ -0,0 +1,41 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
model_name_or_path: Qwen/Qwen2.5-VL-7B-Instruct
|
| 2 |
+
image_max_pixels: 262144
|
| 3 |
+
video_max_pixels: 16384
|
| 4 |
+
trust_remote_code: true
|
| 5 |
+
|
| 6 |
+
### method
|
| 7 |
+
stage: sft
|
| 8 |
+
do_train: true
|
| 9 |
+
finetuning_type: lora
|
| 10 |
+
lora_rank: 32
|
| 11 |
+
lora_target: all
|
| 12 |
+
|
| 13 |
+
### dataset
|
| 14 |
+
dataset: EditScore-Reward-Data
|
| 15 |
+
template: qwen2_vl
|
| 16 |
+
cutoff_len: 8192
|
| 17 |
+
max_samples: 100000
|
| 18 |
+
overwrite_cache: true
|
| 19 |
+
preprocessing_num_workers: 16
|
| 20 |
+
dataloader_num_workers: 4
|
| 21 |
+
|
| 22 |
+
### output
|
| 23 |
+
output_dir: LLaMA-Factory/output/editscore_7B/
|
| 24 |
+
logging_steps: 1
|
| 25 |
+
save_steps: 250
|
| 26 |
+
plot_loss: true
|
| 27 |
+
overwrite_output_dir: true
|
| 28 |
+
save_only_model: false
|
| 29 |
+
report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow]
|
| 30 |
+
|
| 31 |
+
### train
|
| 32 |
+
per_device_train_batch_size: 1
|
| 33 |
+
gradient_accumulation_steps: 16
|
| 34 |
+
learning_rate: 1.0e-4
|
| 35 |
+
num_train_epochs: 3.0
|
| 36 |
+
lr_scheduler_type: cosine
|
| 37 |
+
warmup_ratio: 0.1
|
| 38 |
+
bf16: true
|
| 39 |
+
ddp_timeout: 180000000
|
| 40 |
+
resume_from_checkpoint: null
|
| 41 |
+
|
examples/EditScore-train/config/editscore_qwen3_vl_4B_instruct.yaml
ADDED
|
@@ -0,0 +1,41 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
model_name_or_path: Qwen/Qwen3-VL-4B-Instruct
|
| 2 |
+
image_max_pixels: 262144
|
| 3 |
+
video_max_pixels: 16384
|
| 4 |
+
trust_remote_code: true
|
| 5 |
+
|
| 6 |
+
### method
|
| 7 |
+
stage: sft
|
| 8 |
+
do_train: true
|
| 9 |
+
finetuning_type: lora
|
| 10 |
+
lora_rank: 32
|
| 11 |
+
lora_target: all
|
| 12 |
+
|
| 13 |
+
### dataset
|
| 14 |
+
dataset: EditScore-Reward-Data
|
| 15 |
+
template: qwen3_vl
|
| 16 |
+
cutoff_len: 8192
|
| 17 |
+
max_samples: 500000
|
| 18 |
+
overwrite_cache: true
|
| 19 |
+
preprocessing_num_workers: 16
|
| 20 |
+
dataloader_num_workers: 4
|
| 21 |
+
|
| 22 |
+
### output
|
| 23 |
+
output_dir: LLaMA-Factory/output/editscore_qwen3_vl_4B_instruct/
|
| 24 |
+
logging_steps: 1
|
| 25 |
+
save_steps: 250
|
| 26 |
+
plot_loss: true
|
| 27 |
+
overwrite_output_dir: true
|
| 28 |
+
save_only_model: false
|
| 29 |
+
report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow]
|
| 30 |
+
|
| 31 |
+
### train
|
| 32 |
+
per_device_train_batch_size: 1
|
| 33 |
+
gradient_accumulation_steps: 16
|
| 34 |
+
learning_rate: 1.0e-4
|
| 35 |
+
num_train_epochs: 3.0
|
| 36 |
+
lr_scheduler_type: cosine
|
| 37 |
+
warmup_ratio: 0.1
|
| 38 |
+
bf16: true
|
| 39 |
+
ddp_timeout: 180000000
|
| 40 |
+
resume_from_checkpoint: null
|
| 41 |
+
|
examples/EditScore-train/config/editscore_qwen3_vl_8B_instruct.yaml
ADDED
|
@@ -0,0 +1,41 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
model_name_or_path: Qwen/Qwen3-VL-8B-Instruct
|
| 2 |
+
image_max_pixels: 262144
|
| 3 |
+
video_max_pixels: 16384
|
| 4 |
+
trust_remote_code: true
|
| 5 |
+
|
| 6 |
+
### method
|
| 7 |
+
stage: sft
|
| 8 |
+
do_train: true
|
| 9 |
+
finetuning_type: lora
|
| 10 |
+
lora_rank: 32
|
| 11 |
+
lora_target: all
|
| 12 |
+
|
| 13 |
+
### dataset
|
| 14 |
+
dataset: EditScore-Reward-Data
|
| 15 |
+
template: qwen3_vl
|
| 16 |
+
cutoff_len: 8192
|
| 17 |
+
max_samples: 500000
|
| 18 |
+
overwrite_cache: true
|
| 19 |
+
preprocessing_num_workers: 16
|
| 20 |
+
dataloader_num_workers: 4
|
| 21 |
+
|
| 22 |
+
### output
|
| 23 |
+
output_dir: LLaMA-Factory/output/editscore_qwen3_vl_8B_instruct/
|
| 24 |
+
logging_steps: 1
|
| 25 |
+
save_steps: 250
|
| 26 |
+
plot_loss: true
|
| 27 |
+
overwrite_output_dir: true
|
| 28 |
+
save_only_model: false
|
| 29 |
+
report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow]
|
| 30 |
+
|
| 31 |
+
### train
|
| 32 |
+
per_device_train_batch_size: 1
|
| 33 |
+
gradient_accumulation_steps: 16
|
| 34 |
+
learning_rate: 1.0e-4
|
| 35 |
+
num_train_epochs: 3.0
|
| 36 |
+
lr_scheduler_type: cosine
|
| 37 |
+
warmup_ratio: 0.1
|
| 38 |
+
bf16: true
|
| 39 |
+
ddp_timeout: 180000000
|
| 40 |
+
resume_from_checkpoint: null
|
| 41 |
+
|
examples/EditScore-train/train.sh
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
SHELL_FOLDER=$(cd "$(dirname "$0")";pwd)
|
| 3 |
+
cd $SHELL_FOLDER
|
| 4 |
+
|
| 5 |
+
# Activate conda environment of LLaMA-Factory
|
| 6 |
+
conda activate llama-factory
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
RANK=0
|
| 10 |
+
WORLD_SIZE=1
|
| 11 |
+
MASTER_ADDR="localhost"
|
| 12 |
+
MASTER_PORT=29500
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
while [[ $# -gt 0 ]]; do
|
| 16 |
+
case "$1" in
|
| 17 |
+
--rank=*)
|
| 18 |
+
RANK="${1#*=}"
|
| 19 |
+
shift
|
| 20 |
+
;;
|
| 21 |
+
--world_size=*)
|
| 22 |
+
WORLD_SIZE="${1#*=}"
|
| 23 |
+
shift
|
| 24 |
+
;;
|
| 25 |
+
--master_addr=*)
|
| 26 |
+
MASTER_ADDR="${1#*=}"
|
| 27 |
+
shift
|
| 28 |
+
;;
|
| 29 |
+
--master_port=*)
|
| 30 |
+
MASTER_PORT="${1#*=}"
|
| 31 |
+
shift
|
| 32 |
+
;;
|
| 33 |
+
*)
|
| 34 |
+
echo "Unknown parameter: $1"
|
| 35 |
+
exit 1
|
| 36 |
+
;;
|
| 37 |
+
esac
|
| 38 |
+
done
|
| 39 |
+
|
| 40 |
+
name=experiment_name
|
| 41 |
+
log_dir="LLaMA-Factory/logs"
|
| 42 |
+
|
| 43 |
+
log_file="${log_dir}/${name}_${RANK}.log"
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
CONFIG_YAML="LLaMA-Factory/examples/train_editscore/${name}.yaml"
|
| 47 |
+
FORCE_TORCHRUN=1 \
|
| 48 |
+
NNODES=${WORLD_SIZE} \
|
| 49 |
+
NODE_RANK=${RANK} \
|
| 50 |
+
MASTER_ADDR=${MASTER_ADDR} \
|
| 51 |
+
MASTER_PORT=${MASTER_PORT} \
|
| 52 |
+
llamafactory-cli train ${CONFIG_YAML} 2>&1 | tee ${log_file}
|
examples/OmniGen2-RL/.gitignore
ADDED
|
@@ -0,0 +1,233 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Created by https://www.toptal.com/developers/gitignore/api/macos,python
|
| 2 |
+
# Edit at https://www.toptal.com/developers/gitignore?templates=macos,python
|
| 3 |
+
|
| 4 |
+
### macOS ###
|
| 5 |
+
# General
|
| 6 |
+
.DS_Store
|
| 7 |
+
.AppleDouble
|
| 8 |
+
.LSOverride
|
| 9 |
+
|
| 10 |
+
# Icon must end with two \r
|
| 11 |
+
Icon
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
# Thumbnails
|
| 15 |
+
._*
|
| 16 |
+
|
| 17 |
+
# Files that might appear in the root of a volume
|
| 18 |
+
.DocumentRevisions-V100
|
| 19 |
+
.fseventsd
|
| 20 |
+
.Spotlight-V100
|
| 21 |
+
.TemporaryItems
|
| 22 |
+
.Trashes
|
| 23 |
+
.VolumeIcon.icns
|
| 24 |
+
.com.apple.timemachine.donotpresent
|
| 25 |
+
|
| 26 |
+
# Directories potentially created on remote AFP share
|
| 27 |
+
.AppleDB
|
| 28 |
+
.AppleDesktop
|
| 29 |
+
Network Trash Folder
|
| 30 |
+
Temporary Items
|
| 31 |
+
.apdisk
|
| 32 |
+
|
| 33 |
+
### macOS Patch ###
|
| 34 |
+
# iCloud generated files
|
| 35 |
+
*.icloud
|
| 36 |
+
|
| 37 |
+
### Python ###
|
| 38 |
+
# Byte-compiled / optimized / DLL files
|
| 39 |
+
__pycache__/
|
| 40 |
+
*.py[cod]
|
| 41 |
+
*$py.class
|
| 42 |
+
|
| 43 |
+
# C extensions
|
| 44 |
+
*.so
|
| 45 |
+
|
| 46 |
+
# Distribution / packaging
|
| 47 |
+
.Python
|
| 48 |
+
build/
|
| 49 |
+
develop-eggs/
|
| 50 |
+
dist/
|
| 51 |
+
downloads/
|
| 52 |
+
eggs/
|
| 53 |
+
.eggs/
|
| 54 |
+
lib/
|
| 55 |
+
lib64/
|
| 56 |
+
parts/
|
| 57 |
+
sdist/
|
| 58 |
+
var/
|
| 59 |
+
wheels/
|
| 60 |
+
share/python-wheels/
|
| 61 |
+
*.egg-info/
|
| 62 |
+
.installed.cfg
|
| 63 |
+
*.egg
|
| 64 |
+
MANIFEST
|
| 65 |
+
|
| 66 |
+
# PyInstaller
|
| 67 |
+
# Usually these files are written by a python script from a template
|
| 68 |
+
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
| 69 |
+
*.manifest
|
| 70 |
+
*.spec
|
| 71 |
+
|
| 72 |
+
# Installer logs
|
| 73 |
+
pip-log.txt
|
| 74 |
+
pip-delete-this-directory.txt
|
| 75 |
+
|
| 76 |
+
# Unit test / coverage reports
|
| 77 |
+
htmlcov/
|
| 78 |
+
.tox/
|
| 79 |
+
.nox/
|
| 80 |
+
.coverage
|
| 81 |
+
.coverage.*
|
| 82 |
+
.cache
|
| 83 |
+
nosetests.xml
|
| 84 |
+
coverage.xml
|
| 85 |
+
*.cover
|
| 86 |
+
*.py,cover
|
| 87 |
+
.hypothesis/
|
| 88 |
+
.pytest_cache/
|
| 89 |
+
cover/
|
| 90 |
+
|
| 91 |
+
# Translations
|
| 92 |
+
*.mo
|
| 93 |
+
*.pot
|
| 94 |
+
|
| 95 |
+
# Django stuff:
|
| 96 |
+
*.log
|
| 97 |
+
local_settings.py
|
| 98 |
+
db.sqlite3
|
| 99 |
+
db.sqlite3-journal
|
| 100 |
+
|
| 101 |
+
# Flask stuff:
|
| 102 |
+
instance/
|
| 103 |
+
.webassets-cache
|
| 104 |
+
|
| 105 |
+
# Scrapy stuff:
|
| 106 |
+
.scrapy
|
| 107 |
+
|
| 108 |
+
# Sphinx documentation
|
| 109 |
+
docs/_build/
|
| 110 |
+
|
| 111 |
+
# PyBuilder
|
| 112 |
+
.pybuilder/
|
| 113 |
+
target/
|
| 114 |
+
|
| 115 |
+
# Jupyter Notebook
|
| 116 |
+
.ipynb_checkpoints
|
| 117 |
+
|
| 118 |
+
# IPython
|
| 119 |
+
profile_default/
|
| 120 |
+
ipython_config.py
|
| 121 |
+
|
| 122 |
+
# pyenv
|
| 123 |
+
# For a library or package, you might want to ignore these files since the code is
|
| 124 |
+
# intended to run in multiple environments; otherwise, check them in:
|
| 125 |
+
# .python-version
|
| 126 |
+
|
| 127 |
+
# pipenv
|
| 128 |
+
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
|
| 129 |
+
# However, in case of collaboration, if having platform-specific dependencies or dependencies
|
| 130 |
+
# having no cross-platform support, pipenv may install dependencies that don't work, or not
|
| 131 |
+
# install all needed dependencies.
|
| 132 |
+
#Pipfile.lock
|
| 133 |
+
|
| 134 |
+
# poetry
|
| 135 |
+
# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
|
| 136 |
+
# This is especially recommended for binary packages to ensure reproducibility, and is more
|
| 137 |
+
# commonly ignored for libraries.
|
| 138 |
+
# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
|
| 139 |
+
#poetry.lock
|
| 140 |
+
|
| 141 |
+
# pdm
|
| 142 |
+
# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control.
|
| 143 |
+
#pdm.lock
|
| 144 |
+
# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it
|
| 145 |
+
# in version control.
|
| 146 |
+
# https://pdm.fming.dev/#use-with-ide
|
| 147 |
+
.pdm.toml
|
| 148 |
+
|
| 149 |
+
# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm
|
| 150 |
+
__pypackages__/
|
| 151 |
+
|
| 152 |
+
# Celery stuff
|
| 153 |
+
celerybeat-schedule
|
| 154 |
+
celerybeat.pid
|
| 155 |
+
|
| 156 |
+
# SageMath parsed files
|
| 157 |
+
*.sage.py
|
| 158 |
+
|
| 159 |
+
# Environments
|
| 160 |
+
.env
|
| 161 |
+
.venv
|
| 162 |
+
env/
|
| 163 |
+
venv/
|
| 164 |
+
ENV/
|
| 165 |
+
env.bak/
|
| 166 |
+
venv.bak/
|
| 167 |
+
|
| 168 |
+
# Spyder project settings
|
| 169 |
+
.spyderproject
|
| 170 |
+
.spyproject
|
| 171 |
+
|
| 172 |
+
# Rope project settings
|
| 173 |
+
.ropeproject
|
| 174 |
+
|
| 175 |
+
# mkdocs documentation
|
| 176 |
+
/site
|
| 177 |
+
|
| 178 |
+
# mypy
|
| 179 |
+
.mypy_cache/
|
| 180 |
+
.dmypy.json
|
| 181 |
+
dmypy.json
|
| 182 |
+
|
| 183 |
+
# Pyre type checker
|
| 184 |
+
.pyre/
|
| 185 |
+
|
| 186 |
+
# pytype static type analyzer
|
| 187 |
+
.pytype/
|
| 188 |
+
|
| 189 |
+
# Cython debug symbols
|
| 190 |
+
cython_debug/
|
| 191 |
+
|
| 192 |
+
# PyCharm
|
| 193 |
+
# JetBrains specific template is maintained in a separate JetBrains.gitignore that can
|
| 194 |
+
# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore
|
| 195 |
+
# and can be added to the global gitignore or merged into this file. For a more nuclear
|
| 196 |
+
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
|
| 197 |
+
#.idea/
|
| 198 |
+
|
| 199 |
+
### Python Patch ###
|
| 200 |
+
# Poetry local configuration file - https://python-poetry.org/docs/configuration/#local-configuration
|
| 201 |
+
poetry.toml
|
| 202 |
+
|
| 203 |
+
# ruff
|
| 204 |
+
.ruff_cache/
|
| 205 |
+
|
| 206 |
+
# LSP config files
|
| 207 |
+
pyrightconfig.json
|
| 208 |
+
|
| 209 |
+
# End of https://www.toptal.com/developers/gitignore/api/macos,python
|
| 210 |
+
|
| 211 |
+
local_scripts/
|
| 212 |
+
|
| 213 |
+
omnigen2/utils/vpn_utils.py
|
| 214 |
+
|
| 215 |
+
test_tokenizer.py
|
| 216 |
+
save_pipeline.py
|
| 217 |
+
app.sh
|
| 218 |
+
logs/
|
| 219 |
+
results/
|
| 220 |
+
test_jsonl*
|
| 221 |
+
pbs_files/
|
| 222 |
+
convert_ckpt_to_pipeline.py
|
| 223 |
+
inference_test_efficiency.py
|
| 224 |
+
upload_pipeline*
|
| 225 |
+
example_images_resized/
|
| 226 |
+
example_t2i_test_efficiency*.sh
|
| 227 |
+
example_edit_test_efficiency*.sh
|
| 228 |
+
example_in_context_generation_test_efficiency*.sh
|
| 229 |
+
intro*
|
| 230 |
+
resize_example_images.py
|
| 231 |
+
save_pipeline.py
|
| 232 |
+
outputs_gradio/*
|
| 233 |
+
test.py
|
examples/OmniGen2-RL/LICENSE
ADDED
|
@@ -0,0 +1,201 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Apache License
|
| 2 |
+
Version 2.0, January 2004
|
| 3 |
+
http://www.apache.org/licenses/
|
| 4 |
+
|
| 5 |
+
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
| 6 |
+
|
| 7 |
+
1. Definitions.
|
| 8 |
+
|
| 9 |
+
"License" shall mean the terms and conditions for use, reproduction,
|
| 10 |
+
and distribution as defined by Sections 1 through 9 of this document.
|
| 11 |
+
|
| 12 |
+
"Licensor" shall mean the copyright owner or entity authorized by
|
| 13 |
+
the copyright owner that is granting the License.
|
| 14 |
+
|
| 15 |
+
"Legal Entity" shall mean the union of the acting entity and all
|
| 16 |
+
other entities that control, are controlled by, or are under common
|
| 17 |
+
control with that entity. For the purposes of this definition,
|
| 18 |
+
"control" means (i) the power, direct or indirect, to cause the
|
| 19 |
+
direction or management of such entity, whether by contract or
|
| 20 |
+
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
| 21 |
+
outstanding shares, or (iii) beneficial ownership of such entity.
|
| 22 |
+
|
| 23 |
+
"You" (or "Your") shall mean an individual or Legal Entity
|
| 24 |
+
exercising permissions granted by this License.
|
| 25 |
+
|
| 26 |
+
"Source" form shall mean the preferred form for making modifications,
|
| 27 |
+
including but not limited to software source code, documentation
|
| 28 |
+
source, and configuration files.
|
| 29 |
+
|
| 30 |
+
"Object" form shall mean any form resulting from mechanical
|
| 31 |
+
transformation or translation of a Source form, including but
|
| 32 |
+
not limited to compiled object code, generated documentation,
|
| 33 |
+
and conversions to other media types.
|
| 34 |
+
|
| 35 |
+
"Work" shall mean the work of authorship, whether in Source or
|
| 36 |
+
Object form, made available under the License, as indicated by a
|
| 37 |
+
copyright notice that is included in or attached to the work
|
| 38 |
+
(an example is provided in the Appendix below).
|
| 39 |
+
|
| 40 |
+
"Derivative Works" shall mean any work, whether in Source or Object
|
| 41 |
+
form, that is based on (or derived from) the Work and for which the
|
| 42 |
+
editorial revisions, annotations, elaborations, or other modifications
|
| 43 |
+
represent, as a whole, an original work of authorship. For the purposes
|
| 44 |
+
of this License, Derivative Works shall not include works that remain
|
| 45 |
+
separable from, or merely link (or bind by name) to the interfaces of,
|
| 46 |
+
the Work and Derivative Works thereof.
|
| 47 |
+
|
| 48 |
+
"Contribution" shall mean any work of authorship, including
|
| 49 |
+
the original version of the Work and any modifications or additions
|
| 50 |
+
to that Work or Derivative Works thereof, that is intentionally
|
| 51 |
+
submitted to Licensor for inclusion in the Work by the copyright owner
|
| 52 |
+
or by an individual or Legal Entity authorized to submit on behalf of
|
| 53 |
+
the copyright owner. For the purposes of this definition, "submitted"
|
| 54 |
+
means any form of electronic, verbal, or written communication sent
|
| 55 |
+
to the Licensor or its representatives, including but not limited to
|
| 56 |
+
communication on electronic mailing lists, source code control systems,
|
| 57 |
+
and issue tracking systems that are managed by, or on behalf of, the
|
| 58 |
+
Licensor for the purpose of discussing and improving the Work, but
|
| 59 |
+
excluding communication that is conspicuously marked or otherwise
|
| 60 |
+
designated in writing by the copyright owner as "Not a Contribution."
|
| 61 |
+
|
| 62 |
+
"Contributor" shall mean Licensor and any individual or Legal Entity
|
| 63 |
+
on behalf of whom a Contribution has been received by Licensor and
|
| 64 |
+
subsequently incorporated within the Work.
|
| 65 |
+
|
| 66 |
+
2. Grant of Copyright License. Subject to the terms and conditions of
|
| 67 |
+
this License, each Contributor hereby grants to You a perpetual,
|
| 68 |
+
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 69 |
+
copyright license to reproduce, prepare Derivative Works of,
|
| 70 |
+
publicly display, publicly perform, sublicense, and distribute the
|
| 71 |
+
Work and such Derivative Works in Source or Object form.
|
| 72 |
+
|
| 73 |
+
3. Grant of Patent License. Subject to the terms and conditions of
|
| 74 |
+
this License, each Contributor hereby grants to You a perpetual,
|
| 75 |
+
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
| 76 |
+
(except as stated in this section) patent license to make, have made,
|
| 77 |
+
use, offer to sell, sell, import, and otherwise transfer the Work,
|
| 78 |
+
where such license applies only to those patent claims licensable
|
| 79 |
+
by such Contributor that are necessarily infringed by their
|
| 80 |
+
Contribution(s) alone or by combination of their Contribution(s)
|
| 81 |
+
with the Work to which such Contribution(s) was submitted. If You
|
| 82 |
+
institute patent litigation against any entity (including a
|
| 83 |
+
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
| 84 |
+
or a Contribution incorporated within the Work constitutes direct
|
| 85 |
+
or contributory patent infringement, then any patent licenses
|
| 86 |
+
granted to You under this License for that Work shall terminate
|
| 87 |
+
as of the date such litigation is filed.
|
| 88 |
+
|
| 89 |
+
4. Redistribution. You may reproduce and distribute copies of the
|
| 90 |
+
Work or Derivative Works thereof in any medium, with or without
|
| 91 |
+
modifications, and in Source or Object form, provided that You
|
| 92 |
+
meet the following conditions:
|
| 93 |
+
|
| 94 |
+
(a) You must give any other recipients of the Work or
|
| 95 |
+
Derivative Works a copy of this License; and
|
| 96 |
+
|
| 97 |
+
(b) You must cause any modified files to carry prominent notices
|
| 98 |
+
stating that You changed the files; and
|
| 99 |
+
|
| 100 |
+
(c) You must retain, in the Source form of any Derivative Works
|
| 101 |
+
that You distribute, all copyright, patent, trademark, and
|
| 102 |
+
attribution notices from the Source form of the Work,
|
| 103 |
+
excluding those notices that do not pertain to any part of
|
| 104 |
+
the Derivative Works; and
|
| 105 |
+
|
| 106 |
+
(d) If the Work includes a "NOTICE" text file as part of its
|
| 107 |
+
distribution, then any Derivative Works that You distribute must
|
| 108 |
+
include a readable copy of the attribution notices contained
|
| 109 |
+
within such NOTICE file, excluding those notices that do not
|
| 110 |
+
pertain to any part of the Derivative Works, in at least one
|
| 111 |
+
of the following places: within a NOTICE text file distributed
|
| 112 |
+
as part of the Derivative Works; within the Source form or
|
| 113 |
+
documentation, if provided along with the Derivative Works; or,
|
| 114 |
+
within a display generated by the Derivative Works, if and
|
| 115 |
+
wherever such third-party notices normally appear. The contents
|
| 116 |
+
of the NOTICE file are for informational purposes only and
|
| 117 |
+
do not modify the License. You may add Your own attribution
|
| 118 |
+
notices within Derivative Works that You distribute, alongside
|
| 119 |
+
or as an addendum to the NOTICE text from the Work, provided
|
| 120 |
+
that such additional attribution notices cannot be construed
|
| 121 |
+
as modifying the License.
|
| 122 |
+
|
| 123 |
+
You may add Your own copyright statement to Your modifications and
|
| 124 |
+
may provide additional or different license terms and conditions
|
| 125 |
+
for use, reproduction, or distribution of Your modifications, or
|
| 126 |
+
for any such Derivative Works as a whole, provided Your use,
|
| 127 |
+
reproduction, and distribution of the Work otherwise complies with
|
| 128 |
+
the conditions stated in this License.
|
| 129 |
+
|
| 130 |
+
5. Submission of Contributions. Unless You explicitly state otherwise,
|
| 131 |
+
any Contribution intentionally submitted for inclusion in the Work
|
| 132 |
+
by You to the Licensor shall be under the terms and conditions of
|
| 133 |
+
this License, without any additional terms or conditions.
|
| 134 |
+
Notwithstanding the above, nothing herein shall supersede or modify
|
| 135 |
+
the terms of any separate license agreement you may have executed
|
| 136 |
+
with Licensor regarding such Contributions.
|
| 137 |
+
|
| 138 |
+
6. Trademarks. This License does not grant permission to use the trade
|
| 139 |
+
names, trademarks, service marks, or product names of the Licensor,
|
| 140 |
+
except as required for reasonable and customary use in describing the
|
| 141 |
+
origin of the Work and reproducing the content of the NOTICE file.
|
| 142 |
+
|
| 143 |
+
7. Disclaimer of Warranty. Unless required by applicable law or
|
| 144 |
+
agreed to in writing, Licensor provides the Work (and each
|
| 145 |
+
Contributor provides its Contributions) on an "AS IS" BASIS,
|
| 146 |
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
| 147 |
+
implied, including, without limitation, any warranties or conditions
|
| 148 |
+
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
| 149 |
+
PARTICULAR PURPOSE. You are solely responsible for determining the
|
| 150 |
+
appropriateness of using or redistributing the Work and assume any
|
| 151 |
+
risks associated with Your exercise of permissions under this License.
|
| 152 |
+
|
| 153 |
+
8. Limitation of Liability. In no event and under no legal theory,
|
| 154 |
+
whether in tort (including negligence), contract, or otherwise,
|
| 155 |
+
unless required by applicable law (such as deliberate and grossly
|
| 156 |
+
negligent acts) or agreed to in writing, shall any Contributor be
|
| 157 |
+
liable to You for damages, including any direct, indirect, special,
|
| 158 |
+
incidental, or consequential damages of any character arising as a
|
| 159 |
+
result of this License or out of the use or inability to use the
|
| 160 |
+
Work (including but not limited to damages for loss of goodwill,
|
| 161 |
+
work stoppage, computer failure or malfunction, or any and all
|
| 162 |
+
other commercial damages or losses), even if such Contributor
|
| 163 |
+
has been advised of the possibility of such damages.
|
| 164 |
+
|
| 165 |
+
9. Accepting Warranty or Additional Liability. While redistributing
|
| 166 |
+
the Work or Derivative Works thereof, You may choose to offer,
|
| 167 |
+
and charge a fee for, acceptance of support, warranty, indemnity,
|
| 168 |
+
or other liability obligations and/or rights consistent with this
|
| 169 |
+
License. However, in accepting such obligations, You may act only
|
| 170 |
+
on Your own behalf and on Your sole responsibility, not on behalf
|
| 171 |
+
of any other Contributor, and only if You agree to indemnify,
|
| 172 |
+
defend, and hold each Contributor harmless for any liability
|
| 173 |
+
incurred by, or claims asserted against, such Contributor by reason
|
| 174 |
+
of your accepting any such warranty or additional liability.
|
| 175 |
+
|
| 176 |
+
END OF TERMS AND CONDITIONS
|
| 177 |
+
|
| 178 |
+
APPENDIX: How to apply the Apache License to your work.
|
| 179 |
+
|
| 180 |
+
To apply the Apache License to your work, attach the following
|
| 181 |
+
boilerplate notice, with the fields enclosed by brackets "[]"
|
| 182 |
+
replaced with your own identifying information. (Don't include
|
| 183 |
+
the brackets!) The text should be enclosed in the appropriate
|
| 184 |
+
comment syntax for the file format. We also recommend that a
|
| 185 |
+
file or class name and description of purpose be included on the
|
| 186 |
+
same "printed page" as the copyright notice for easier
|
| 187 |
+
identification within third-party archives.
|
| 188 |
+
|
| 189 |
+
Copyright [yyyy] [name of copyright owner]
|
| 190 |
+
|
| 191 |
+
Licensed under the Apache License, Version 2.0 (the "License");
|
| 192 |
+
you may not use this file except in compliance with the License.
|
| 193 |
+
You may obtain a copy of the License at
|
| 194 |
+
|
| 195 |
+
http://www.apache.org/licenses/LICENSE-2.0
|
| 196 |
+
|
| 197 |
+
Unless required by applicable law or agreed to in writing, software
|
| 198 |
+
distributed under the License is distributed on an "AS IS" BASIS,
|
| 199 |
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 200 |
+
See the License for the specific language governing permissions and
|
| 201 |
+
limitations under the License.
|
examples/OmniGen2-RL/README.md
ADDED
|
@@ -0,0 +1,189 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# 🚀 Advanced Applications of EditScore
|
| 2 |
+
|
| 3 |
+
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*.
|
| 4 |
+
|
| 5 |
+
This guide covers two primary downstream applications:
|
| 6 |
+
1. **Best-of-N selection**: A simple, training-free method to instantly boost the output quality of any image editing model.
|
| 7 |
+
2. **Reinforcement Learning (RL) Fine-Tuning**: Using EditScore as a high-fidelity reward signal to train models for significantly better performance.
|
| 8 |
+
|
| 9 |
+
## 🛠️ Setup for Examples
|
| 10 |
+
The examples require libraries for RL, data handling, and potentially experiment tracking.
|
| 11 |
+
```bash
|
| 12 |
+
# Navigate to this directory if you are in the root
|
| 13 |
+
cd examples/OmniGen2-RL
|
| 14 |
+
|
| 15 |
+
# Install the required packages
|
| 16 |
+
pip install -r requirements.txt
|
| 17 |
+
```
|
| 18 |
+
|
| 19 |
+
## Application 1: Best-of-N for Superior Outputs
|
| 20 |
+
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.
|
| 21 |
+
|
| 22 |
+
This acts as a powerful "reranker" that filters out suboptimal results, significantly improving the perceived quality of the model without any extra training.
|
| 23 |
+
|
| 24 |
+
### How to Use
|
| 25 |
+
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.
|
| 26 |
+
|
| 27 |
+
**1. Generate Candidates**
|
| 28 |
+
```bash
|
| 29 |
+
bash evaluation/GEdit-Bench/omnigen2_16samples.sh # default using 8 GPUs
|
| 30 |
+
```
|
| 31 |
+
|
| 32 |
+
> **⚠️ Important Note on Resource Usage**
|
| 33 |
+
>
|
| 34 |
+
> 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**.
|
| 35 |
+
|
| 36 |
+
<details>
|
| 37 |
+
<summary><strong>👉 Click here for tips on the usage of the script</strong></summary>
|
| 38 |
+
|
| 39 |
+
- **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:
|
| 40 |
+
```bash
|
| 41 |
+
# On the first machine (rank 0)
|
| 42 |
+
bash evaluation/GEdit-Bench/omnigen2_16samples.sh --world_size 4 --rank 0
|
| 43 |
+
|
| 44 |
+
# On the second machine (rank 1)
|
| 45 |
+
bash evaluation/GEdit-Bench/omnigen2_16samples.sh --world_size 4 --rank 1
|
| 46 |
+
|
| 47 |
+
# ...and so on for ranks 2 and 3.
|
| 48 |
+
```
|
| 49 |
+
|
| 50 |
+
- **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.
|
| 51 |
+
</details>
|
| 52 |
+
|
| 53 |
+
**2. Score and Select**
|
| 54 |
+
Next, use EditScore to evaluate all N candidates and identify the one with the highest score.
|
| 55 |
+
|
| 56 |
+
```bash
|
| 57 |
+
bash evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass1.sh # EditScore-7B, single pass
|
| 58 |
+
bash evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass4.sh # EditScore-7B, Avg@4
|
| 59 |
+
```
|
| 60 |
+
|
| 61 |
+
**3. Evaluate the Final Selections**
|
| 62 |
+
Finally, evaluate the performance of the images selected by EditScore on GEdit-Bench to quantify the improvement.
|
| 63 |
+
```bash
|
| 64 |
+
bash evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass1_eval.sh
|
| 65 |
+
bash evaluation/GEdit-Bench/omnigen2_16samples_select_best_editscore_pass4_eval.sh
|
| 66 |
+
```
|
| 67 |
+
|
| 68 |
+
By comparing these results to the baseline performance of the original model, you will see the benefits of applying EditScore as a reranker.
|
| 69 |
+
|
| 70 |
+
## Application 2: Reinforcement Fine-Tuning
|
| 71 |
+
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.
|
| 72 |
+
|
| 73 |
+
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.
|
| 74 |
+
|
| 75 |
+
### 1. Prepare Training Data
|
| 76 |
+
First, set up the dataset for RL fine-tuning.
|
| 77 |
+
1. Download the Data
|
| 78 |
+
Downlaod the official RL training data from [EditScore-RL-Data](https://huggingface.co/datasets/EditScore/EditScore-RL-Data).
|
| 79 |
+
2. Create Meta File
|
| 80 |
+
The uploaded dataset uses relative image paths. Run the following script to convert them to absolute paths based on your local environment:
|
| 81 |
+
```bash
|
| 82 |
+
# Then
|
| 83 |
+
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
|
| 84 |
+
|
| 85 |
+
# Due to the limitation of base model (OmniGen2), we discard text change and portrait beautification, as these tasks harm RL training.
|
| 86 |
+
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
|
| 87 |
+
```
|
| 88 |
+
3. Configure the Data Path
|
| 89 |
+
Specify the path to your processed `.jsonl` file in the data configuration located at `data_configs/train/example/edit/all.yml`.
|
| 90 |
+
For example:
|
| 91 |
+
```yaml
|
| 92 |
+
ratio_type: inside_ratio
|
| 93 |
+
|
| 94 |
+
data:
|
| 95 |
+
-
|
| 96 |
+
path: '/path/to/EditScore-RL-Data/rl_abs_9tasks.jsonl' # <-- Ensure this path is correct
|
| 97 |
+
type: 'edit'
|
| 98 |
+
ratio: !!float 1
|
| 99 |
+
```
|
| 100 |
+
|
| 101 |
+
### 2. Prepare the Base Model (OmniGen2)
|
| 102 |
+
```bash
|
| 103 |
+
python scripts/misc/extract_bin_from_pipe.py
|
| 104 |
+
```
|
| 105 |
+
|
| 106 |
+
### 3. Launch the Reward Server
|
| 107 |
+
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.
|
| 108 |
+
|
| 109 |
+
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.
|
| 110 |
+
|
| 111 |
+
We provide a convenient script to launch the entire server stack across multiple machines, assuming you have `ssh` access to all reward server nodes.
|
| 112 |
+
|
| 113 |
+
```bash
|
| 114 |
+
# Launch EditScore-7B Reward Server
|
| 115 |
+
bash reward_server/start_multi_machines.sh --model_name=editscore_7B --config_path=reward_server/server_configs/editscore_7B.yml
|
| 116 |
+
|
| 117 |
+
# Launch EditScore-7B (Avg@4) Reward Server
|
| 118 |
+
bash reward_server/start_multi_machines.sh --model_name=editscore_7B_pass4 --config_path=reward_server/server_configs/editscore_7B_pass4.yml
|
| 119 |
+
|
| 120 |
+
# Launch EditScore-72B Reward Server
|
| 121 |
+
bash reward_server/start_multi_machines.sh --model_name=editscore_72B --config_path=reward_server/server_configs/editscore_72B.yml
|
| 122 |
+
```
|
| 123 |
+
|
| 124 |
+
> **⚠️ Important Notes**
|
| 125 |
+
>
|
| 126 |
+
> * Before running the script, you **must** specify the IP addresses of your reward server machines in the corresponding `.yml` configuration file.
|
| 127 |
+
> * 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.
|
| 128 |
+
> * You can monitor the status of the proxy and servers by checking the log files in the `reward_server/logs/` directory.
|
| 129 |
+
|
| 130 |
+
## 3.5 (Optional) Reward Server Sanity Check
|
| 131 |
+
To ensure the reward server is configured correctly and running as expected, we provide a sanity check script.
|
| 132 |
+
```bash
|
| 133 |
+
python reward_server/scripts/utils/reward_server_sanity_check.py --config_path=reward_server/server_configs/editscore_7B.yml
|
| 134 |
+
```
|
| 135 |
+
Once these steps are complete, your environment is ready to begin the reinforcement learning fine-tuning process.
|
| 136 |
+
|
| 137 |
+
### 4. Start RL Fine-Tuning
|
| 138 |
+
|
| 139 |
+
**Configure Training Parameters**
|
| 140 |
+
Before launching, you may need to adjust key parameters in the configuration file: `options/omnigen2_edit_rl_4machine_editscore7b_avg4.yml`.
|
| 141 |
+
|
| 142 |
+
Here are some important settings:
|
| 143 |
+
- `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`.
|
| 144 |
+
- `train.batch_size`: Batch size per GPU (`batch_size_per_forward * gradient_accumulation_steps * num_update_steps_per_sampling`)
|
| 145 |
+
- `train.rl.num_images_per_prompt`: The number of candidate images to generate for each unique prompt.
|
| 146 |
+
- `train.rl.num_unique_prompts_per_sampling`: The number of unique prompts in a global batch
|
| 147 |
+
- `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.
|
| 148 |
+
- `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.
|
| 149 |
+
|
| 150 |
+
**Launch Distributed Training**
|
| 151 |
+
We provide scripts for both single and multi-machine distributed training based on **FSDP**.
|
| 152 |
+
```bash
|
| 153 |
+
# Single-machine training (8 GPUs) using EditScore-7B as the reward model
|
| 154 |
+
bash scripts/train/omnigen2_edit_rl_single_machine_editscore7b.sh
|
| 155 |
+
|
| 156 |
+
# Multi-machine training (e.g., 4 machines with 8 GPUs each) using EditScore-7B (Avg@4)
|
| 157 |
+
bash scripts/train/omnigen2_edit_rl_4machine_editscore7b_avg4.sh
|
| 158 |
+
```
|
| 159 |
+
|
| 160 |
+
### 4. Training Outputs and Monitoring
|
| 161 |
+
All training artifacts, including logs and model checkpoints, are saved to the `experiments/` directory.
|
| 162 |
+
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).
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
### 5. Evaluate your RL Fine-Tuned Model
|
| 166 |
+
After training, you must convert the FSDP-saved checkpoint (`.bin`) into the standard Hugging Face format before you can use it for inference.
|
| 167 |
+
|
| 168 |
+
#### Step 1: Convert the Checkpoint
|
| 169 |
+
We provide a script to automatically handle the conversion from the distributed FSDP format to the standard Hugging Face format (`.bin`).
|
| 170 |
+
Run the following command, replacing the arguments with your experiment's details:
|
| 171 |
+
|
| 172 |
+
```shell
|
| 173 |
+
bash scripts/misc/convert_dist_ckpt_to_hf_format.sh [EXPERIMENT_NAME] [STEP_NUMBER]
|
| 174 |
+
```
|
| 175 |
+
- [EXPERIMENT_NAME]: The name of your training experiment (e.g., omnigen2_edit_rl_single_machine_editscore7b).
|
| 176 |
+
- [STEP_NUMBER]: The specific training step of the checkpoint you wish to evaluate (e.g., 500).
|
| 177 |
+
|
| 178 |
+
This will create a new directory containing the converted model weights in the standard format, ready for inference.
|
| 179 |
+
|
| 180 |
+
#### Step 2: Run Evaluation on GEdit-Bench
|
| 181 |
+
Once the checkpoint is converted, you can benchmark its performance. We provide evaluation scripts tailored for GEdit-Bench.
|
| 182 |
+
|
| 183 |
+
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.
|
| 184 |
+
```shell
|
| 185 |
+
# Run evaluation for the converted model from step 500
|
| 186 |
+
bash evaluation/GEdit-Bench/omnigen2.sh --experiment_name=omnigen2_edit_rl_4machine_editscore7b_avg4 --step=500
|
| 187 |
+
bash evaluation/GEdit-Bench/omnigen2_eval.sh --experiment_name=omnigen2_edit_rl_4machine_editscore7b_avg4 --step=500
|
| 188 |
+
```
|
| 189 |
+
By comparing the results to the baseline model's performance, you can quantify the improvements achieved through RL fine-tuning with EditScore.
|
examples/OmniGen2-RL/data_configs/train/example/edit/all.yml
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
ratio_type: inside_ratio
|
| 2 |
+
|
| 3 |
+
data:
|
| 4 |
+
-
|
| 5 |
+
path: '/path/to/EditScore-RL-Data/rl_abs_9tasks.jsonl'
|
| 6 |
+
type: 'edit'
|
| 7 |
+
ratio: !!float 1
|
examples/OmniGen2-RL/data_configs/train/example/train.yml
ADDED
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
data:
|
| 2 |
+
-
|
| 3 |
+
path: 'data_configs/train/example/edit/all.yml'
|
| 4 |
+
type: 'edit'
|
| 5 |
+
ratio: !!float 1
|
examples/OmniGen2-RL/docs/README.md
ADDED
|
@@ -0,0 +1,2 @@
|
|
|
|
|
|
|
|
|
|
| 1 |
+
## Apply EditScore to Image Editing
|
| 2 |
+
### Best-of-N selection**
|
examples/OmniGen2-RL/evaluation/GEdit-Bench/calculate_statistics.py
ADDED
|
@@ -0,0 +1,223 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import pandas as pd
|
| 3 |
+
from collections import defaultdict
|
| 4 |
+
import sys
|
| 5 |
+
import numpy as np
|
| 6 |
+
import math
|
| 7 |
+
|
| 8 |
+
GROUPS = [
|
| 9 |
+
"background_change",
|
| 10 |
+
"color_alter",
|
| 11 |
+
"style_change",
|
| 12 |
+
"subject-add",
|
| 13 |
+
"subject-remove",
|
| 14 |
+
"subject-replace",
|
| 15 |
+
"material_alter",
|
| 16 |
+
"motion_change",
|
| 17 |
+
"ps_human",
|
| 18 |
+
"text_change",
|
| 19 |
+
"tone_transfer",
|
| 20 |
+
]
|
| 21 |
+
|
| 22 |
+
GROUPS2 = [
|
| 23 |
+
"background_change",
|
| 24 |
+
"color_alter",
|
| 25 |
+
"material_alter",
|
| 26 |
+
"motion_change",
|
| 27 |
+
"ps_human",
|
| 28 |
+
"style_change",
|
| 29 |
+
"subject-add",
|
| 30 |
+
"subject-remove",
|
| 31 |
+
"subject-replace",
|
| 32 |
+
"text_change",
|
| 33 |
+
"tone_transfer",
|
| 34 |
+
]
|
| 35 |
+
|
| 36 |
+
def analyze_scores(result_dir, language, num_samples): # 这些 group_scores 字典用于存储每个 group 的最终平均分
|
| 37 |
+
group_scores_semantics = {}
|
| 38 |
+
group_scores_quality = {}
|
| 39 |
+
group_scores_overall = {}
|
| 40 |
+
group_scores_semantics_intersection = {}
|
| 41 |
+
group_scores_quality_intersection = {}
|
| 42 |
+
group_scores_overall_intersection = {}
|
| 43 |
+
|
| 44 |
+
# 外部循环,处理每一个 group
|
| 45 |
+
for group_name in GROUPS:
|
| 46 |
+
data_point_samples = defaultdict(list)
|
| 47 |
+
|
| 48 |
+
# 循环读取 num_samples 个评分文件
|
| 49 |
+
for turn in range(num_samples):
|
| 50 |
+
csv_path = os.path.join(result_dir, f"{group_name}_gpt_score{'_sample' + str(turn) if turn > 0 else ''}.csv")
|
| 51 |
+
if not os.path.exists(csv_path):
|
| 52 |
+
print(f"Warning: File not found, skipping: {csv_path}")
|
| 53 |
+
continue
|
| 54 |
+
|
| 55 |
+
with open(csv_path, 'r') as f:
|
| 56 |
+
df = pd.read_csv(f)
|
| 57 |
+
|
| 58 |
+
for _, row in df.iterrows():
|
| 59 |
+
# 过滤语言
|
| 60 |
+
if row['instruction_language'] != language:
|
| 61 |
+
continue
|
| 62 |
+
|
| 63 |
+
# 定义唯一标识符
|
| 64 |
+
unique_key = os.path.basename(row['source_image']).split('_SRCIMG')[0]
|
| 65 |
+
|
| 66 |
+
# 计算 overall_score
|
| 67 |
+
semantics_score = row['sementics_score']
|
| 68 |
+
quality_score = row['quality_score']
|
| 69 |
+
overall_score = math.sqrt(semantics_score * quality_score)
|
| 70 |
+
|
| 71 |
+
# 将当前样本的分数信息存入字典
|
| 72 |
+
sample_data = {
|
| 73 |
+
'semantics_score': semantics_score,
|
| 74 |
+
'quality_score': quality_score,
|
| 75 |
+
'overall_score': overall_score,
|
| 76 |
+
'intersection_exist': row['intersection_exist']
|
| 77 |
+
}
|
| 78 |
+
|
| 79 |
+
# 按唯一标识符聚合所有样本
|
| 80 |
+
data_point_samples[unique_key].append(sample_data)
|
| 81 |
+
|
| 82 |
+
# --- 核心改动部分:第二阶段 - 筛选与计算 ---
|
| 83 |
+
# 现在 data_point_samples 已经收集了所有测试项的所有样本数据。
|
| 84 |
+
# 我们需要遍历它,为每个测试项找到最佳样本,然后将最佳分数存入最终列表。
|
| 85 |
+
|
| 86 |
+
best_semantics_scores = []
|
| 87 |
+
best_quality_scores = []
|
| 88 |
+
best_overall_scores = []
|
| 89 |
+
|
| 90 |
+
for unique_key, samples in data_point_samples.items():
|
| 91 |
+
if not samples:
|
| 92 |
+
continue
|
| 93 |
+
|
| 94 |
+
# 从当前测试项的所有样本中,找到 overall_score 最高的那个
|
| 95 |
+
# max() 函数的 key 参数可以让我们指定按字典中的哪个值来比较
|
| 96 |
+
best_sample = max(samples, key=lambda s: s['overall_score'])
|
| 97 |
+
|
| 98 |
+
# 将这个最佳样本的分数添加到最终列表中
|
| 99 |
+
best_semantics_scores.append(best_sample['semantics_score'])
|
| 100 |
+
best_quality_scores.append(best_sample['quality_score'])
|
| 101 |
+
best_overall_scores.append(best_sample['overall_score'])
|
| 102 |
+
|
| 103 |
+
|
| 104 |
+
group_scores_semantics[group_name] = np.mean(best_semantics_scores)
|
| 105 |
+
group_scores_quality[group_name] = np.mean(best_quality_scores)
|
| 106 |
+
group_scores_overall[group_name] = np.mean(best_overall_scores)
|
| 107 |
+
|
| 108 |
+
print("\n--- Overall Model Averages ---")
|
| 109 |
+
|
| 110 |
+
print("\nSemantics:")
|
| 111 |
+
model_scores = [group_scores_semantics[group] for group in GROUPS]
|
| 112 |
+
model_avg = np.mean(model_scores)
|
| 113 |
+
group_scores_semantics["avg_semantics"] = model_avg
|
| 114 |
+
|
| 115 |
+
# print("\nSemantics Valid Num:")
|
| 116 |
+
# model_scores = [group_scores_semantics_valid_num[group] for group in GROUPS]
|
| 117 |
+
# model_avg = np.mean(model_scores)
|
| 118 |
+
# group_scores_semantics_valid_num["avg_semantics_valid_num"] = model_avg
|
| 119 |
+
|
| 120 |
+
# print("\nSemantics Intersection:")
|
| 121 |
+
# model_scores = [group_scores_semantics_intersection[group] for group in GROUPS]
|
| 122 |
+
# model_avg = np.mean(model_scores)
|
| 123 |
+
# group_scores_semantics_intersection["avg_semantics"] = model_avg
|
| 124 |
+
|
| 125 |
+
# print("\nSemantics Valid Num Intersection:")
|
| 126 |
+
# model_scores = [group_scores_semantics_valid_num_intersection[group] for group in GROUPS]
|
| 127 |
+
# model_avg = np.mean(model_scores)
|
| 128 |
+
# group_scores_semantics_valid_num_intersection["avg_semantics_valid_num"] = model_avg
|
| 129 |
+
|
| 130 |
+
print("\nQuality:")
|
| 131 |
+
model_scores = [group_scores_quality[group] for group in GROUPS]
|
| 132 |
+
model_avg = np.mean(model_scores)
|
| 133 |
+
group_scores_quality["avg_quality"] = model_avg
|
| 134 |
+
|
| 135 |
+
# print("\nQuality Valid Num:")
|
| 136 |
+
# model_scores = [group_scores_quality_valid_num[group] for group in GROUPS]
|
| 137 |
+
# model_avg = np.mean(model_scores)
|
| 138 |
+
# group_scores_quality_valid_num["avg_quality_valid_num"] = model_avg
|
| 139 |
+
|
| 140 |
+
# print("\nQuality Intersection:")
|
| 141 |
+
# model_scores = [group_scores_quality_intersection[group] for group in GROUPS]
|
| 142 |
+
# model_avg = np.mean(model_scores)
|
| 143 |
+
# group_scores_quality_intersection["avg_quality"] = model_avg
|
| 144 |
+
|
| 145 |
+
# print("\nQuality Valid Num Intersection:")
|
| 146 |
+
# model_scores = [group_scores_quality_valid_num_intersection[group] for group in GROUPS]
|
| 147 |
+
# model_avg = np.mean(model_scores)
|
| 148 |
+
# group_scores_quality_valid_num_intersection["avg_quality_valid_num"] = model_avg
|
| 149 |
+
|
| 150 |
+
print("\nOverall:")
|
| 151 |
+
model_scores = [group_scores_overall[group] for group in GROUPS]
|
| 152 |
+
model_avg = np.mean(model_scores)
|
| 153 |
+
group_scores_overall["avg_overall"] = model_avg
|
| 154 |
+
|
| 155 |
+
# print("\nOverall Valid Num:")
|
| 156 |
+
# model_scores = [group_scores_overall_valid_num[group] for group in GROUPS]
|
| 157 |
+
# model_avg = np.mean(model_scores)
|
| 158 |
+
# group_scores_overall_valid_num["avg_overall_valid_num"] = model_avg
|
| 159 |
+
|
| 160 |
+
|
| 161 |
+
return (
|
| 162 |
+
group_scores_semantics,
|
| 163 |
+
group_scores_quality,
|
| 164 |
+
group_scores_overall,
|
| 165 |
+
# group_scores_semantics_valid_num,
|
| 166 |
+
# group_scores_quality_valid_num,
|
| 167 |
+
# group_scores_overall_valid_num
|
| 168 |
+
)
|
| 169 |
+
|
| 170 |
+
if __name__ == "__main__":
|
| 171 |
+
import argparse
|
| 172 |
+
parser = argparse.ArgumentParser()
|
| 173 |
+
parser.add_argument("--result_dir", type=str, default="/results/")
|
| 174 |
+
parser.add_argument("--language", type=str, default="en", choices=["en", "cn"])
|
| 175 |
+
parser.add_argument("--num_samples", type=int, default=1)
|
| 176 |
+
parser.add_argument("--groups", type=str, default="GROUPS", choices=["GROUPS", "GROUPS2"])
|
| 177 |
+
args = parser.parse_args()
|
| 178 |
+
result_dir = args.result_dir
|
| 179 |
+
|
| 180 |
+
# result_dir = os.path.join(result_dir, "viescore")
|
| 181 |
+
|
| 182 |
+
print("\nOverall:")
|
| 183 |
+
|
| 184 |
+
(
|
| 185 |
+
group_scores_semantics,
|
| 186 |
+
group_scores_quality,
|
| 187 |
+
group_scores_overall,
|
| 188 |
+
# group_scores_semantics_valid_num,
|
| 189 |
+
# group_scores_quality_valid_num,
|
| 190 |
+
# group_scores_overall_valid_num
|
| 191 |
+
) = analyze_scores(result_dir, language=args.language, num_samples=args.num_samples)
|
| 192 |
+
|
| 193 |
+
if args.groups == "GROUPS":
|
| 194 |
+
groups = GROUPS
|
| 195 |
+
else:
|
| 196 |
+
groups = GROUPS2
|
| 197 |
+
|
| 198 |
+
for group_name in groups:
|
| 199 |
+
print(f"{group_name}: {group_scores_semantics[group_name]:.2f}, {group_scores_quality[group_name]:.2f}, {group_scores_overall[group_name]:.2f}")
|
| 200 |
+
|
| 201 |
+
print(f"Average: {group_scores_semantics['avg_semantics']:.2f}, {group_scores_quality['avg_quality']:.2f}, {group_scores_overall['avg_overall']:.2f}")
|
| 202 |
+
|
| 203 |
+
print("Semantics: " + " & ".join([f"{group_scores_semantics[group_name]:.2f}" for group_name in groups] + [f"{group_scores_semantics['avg_semantics']:.2f}"]))
|
| 204 |
+
print("Quality: " + " & ".join([f"{group_scores_quality[group_name]:.2f}" for group_name in groups] + [f"{group_scores_quality['avg_quality']:.2f}"]))
|
| 205 |
+
print("Overall: " + " & ".join([f"{group_scores_overall[group_name]:.2f}" for group_name in groups] + [f"{group_scores_overall['avg_overall']:.2f}"]))
|
| 206 |
+
|
| 207 |
+
# print("\nValid Num:")
|
| 208 |
+
# for group_name in GROUPS:
|
| 209 |
+
# 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}")
|
| 210 |
+
|
| 211 |
+
# 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}")
|
| 212 |
+
|
| 213 |
+
# print("\nIntersection:")
|
| 214 |
+
# for group_name in GROUPS:
|
| 215 |
+
# 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}")
|
| 216 |
+
|
| 217 |
+
# 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}")
|
| 218 |
+
|
| 219 |
+
# print("\nValid Num Intersection:")
|
| 220 |
+
# for group_name in GROUPS:
|
| 221 |
+
# 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}")
|
| 222 |
+
|
| 223 |
+
# 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}")
|
examples/OmniGen2-RL/evaluation/GEdit-Bench/flux_kontext_dev_16samples.sh
ADDED
|
@@ -0,0 +1,83 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# !/bin/bash
|
| 2 |
+
SHELL_FOLDER=$(cd "$(dirname "$0")";pwd)
|
| 3 |
+
cd $(dirname $SHELL_FOLDER)
|
| 4 |
+
cd ../
|
| 5 |
+
|
| 6 |
+
RANK=0
|
| 7 |
+
MASTER_ADDR=1
|
| 8 |
+
MASTER_PORT=29500
|
| 9 |
+
WORLD_SIZE=1
|
| 10 |
+
|
| 11 |
+
# 处理命名参数
|
| 12 |
+
while [[ $# -gt 0 ]]; do
|
| 13 |
+
case "$1" in
|
| 14 |
+
--rank=*)
|
| 15 |
+
RANK="${1#*=}"
|
| 16 |
+
shift
|
| 17 |
+
;;
|
| 18 |
+
--master_addr=*)
|
| 19 |
+
MASTER_ADDR="${1#*=}"
|
| 20 |
+
shift
|
| 21 |
+
;;
|
| 22 |
+
--master_port=*)
|
| 23 |
+
MASTER_PORT="${1#*=}"
|
| 24 |
+
shift
|
| 25 |
+
;;
|
| 26 |
+
--world_size=*)
|
| 27 |
+
WORLD_SIZE="${1#*=}"
|
| 28 |
+
shift
|
| 29 |
+
;;
|
| 30 |
+
*)
|
| 31 |
+
echo "未知参数: $1"
|
| 32 |
+
shift
|
| 33 |
+
;;
|
| 34 |
+
esac
|
| 35 |
+
done
|
| 36 |
+
|
| 37 |
+
# 输出配置
|
| 38 |
+
echo "RANK: $RANK"
|
| 39 |
+
echo "MASTER_ADDR: $MASTER_ADDR"
|
| 40 |
+
echo "MASTER_PORT: $MASTER_PORT"
|
| 41 |
+
echo "WORLD_SIZE: $WORLD_SIZE"
|
| 42 |
+
|
| 43 |
+
global_shift_index=0
|
| 44 |
+
total_num_images=606
|
| 45 |
+
|
| 46 |
+
num_gpus_per_machine=$(python -c "import torch; print(torch.cuda.device_count())")
|
| 47 |
+
# Calculate images per machine, rounding up to ensure all data is covered
|
| 48 |
+
num_images_per_machine=$(( (total_num_images + WORLD_SIZE - 1) / WORLD_SIZE ))
|
| 49 |
+
shift_index=$((RANK * num_images_per_machine))
|
| 50 |
+
|
| 51 |
+
if [ $((total_num_images - shift_index)) -lt $num_images_per_machine ]; then
|
| 52 |
+
num_images_per_machine=$((total_num_images - shift_index))
|
| 53 |
+
fi
|
| 54 |
+
|
| 55 |
+
# Calculate base number of images per GPU (for first 7 GPUs)
|
| 56 |
+
num_images_per_gpu=$(( (num_images_per_machine + num_gpus_per_machine - 1) / num_gpus_per_machine ))
|
| 57 |
+
|
| 58 |
+
guidance_scale=2.5
|
| 59 |
+
|
| 60 |
+
for ((i=0; i<num_gpus_per_machine; i++)); do
|
| 61 |
+
if [ $i -lt $((num_gpus_per_machine - 1)) ]; then
|
| 62 |
+
# First 7 GPUs process equal amounts
|
| 63 |
+
start_idx=$((global_shift_index + i * num_images_per_gpu + shift_index))
|
| 64 |
+
end_idx=$((start_idx + num_images_per_gpu))
|
| 65 |
+
else
|
| 66 |
+
# Last GPU processes remaining data
|
| 67 |
+
start_idx=$((global_shift_index + (num_gpus_per_machine - 1) * num_images_per_gpu + shift_index))
|
| 68 |
+
end_idx=$((global_shift_index + shift_index + num_images_per_machine))
|
| 69 |
+
fi
|
| 70 |
+
echo ${start_idx} ${end_idx}
|
| 71 |
+
|
| 72 |
+
CUDA_VISIBLE_DEVICES=${i} WORLD_SIZE=1 nohup accelerate launch --num_processes 1 --num_machines 1 \
|
| 73 |
+
evaluation/GEdit-Bench/inference_flux_kontext_dev.py \
|
| 74 |
+
--pipeline_path black-forest-labs/FLUX.1-Kontext-dev \
|
| 75 |
+
--num_inference_step 28 \
|
| 76 |
+
--height 1024 \
|
| 77 |
+
--width 1024 \
|
| 78 |
+
--guidance_scale ${guidance_scale} \
|
| 79 |
+
--result_dir evaluation/GEdit-Bench/results/FLUX-Kontext-dev/results_gs${guidance_scale}_16samples \
|
| 80 |
+
--start_index ${start_idx} --end_index ${end_idx} \
|
| 81 |
+
--num_samples 16 \
|
| 82 |
+
> logs/gedit_FLUX-Kontext-dev_gs${guidance_scale}_16samples_${start_idx}_${end_idx}.log 2>&1 &
|
| 83 |
+
done
|