diff --git a/.eslintignore b/.eslintignore new file mode 100644 index 000000000..c098feaa7 --- /dev/null +++ b/.eslintignore @@ -0,0 +1,4 @@ +extensions +extensions-disabled +repositories +venv diff --git a/.eslintrc.json b/.eslintrc.json new file mode 100644 index 000000000..2958fbc79 --- /dev/null +++ b/.eslintrc.json @@ -0,0 +1,89 @@ +{ + "env": { + "browser": true, + "commonjs": false, + "node": false, + "jquery": false, + "es2020": true + }, + "parserOptions": { "ecmaVersion": 2020 }, + "plugins": [], + "extends": ["eslint:recommended", "airbnb-base"], + "rules": { + "arrow-spacing": "error", + "block-spacing": "error", + "brace-style": "error", + "comma-dangle": ["error", "only-multiline"], + "comma-spacing": "error", + "comma-style": ["error", "last"], + "curly": ["error", "multi-line", "consistent"], + "eol-last": "error", + "func-call-spacing": "error", + "function-call-argument-newline": ["error", "consistent"], + "function-paren-newline": ["error", "consistent"], + "indent": ["error", 4], + "key-spacing": "error", + "keyword-spacing": "error", + "linebreak-style": ["error", "unix"], + "no-extra-semi": "error", + "no-mixed-spaces-and-tabs": "error", + "no-multi-spaces": "error", + "no-redeclare": ["error", { "builtinGlobals": false }], + "no-trailing-spaces": "error", + "no-unused-vars": "off", + "no-whitespace-before-property": "error", + "object-curly-newline": ["error", { "consistent": true, "multiline": true }], + "object-curly-spacing": ["error", "never"], + "operator-linebreak": ["error", "after"], + "quote-props": ["error", "consistent-as-needed"], + "semi": ["error", "always"], + "semi-spacing": "error", + "semi-style": ["error", "last"], + "space-before-blocks": "error", + "space-before-function-paren": ["error", "never"], + "space-in-parens": ["error", "never"], + "space-infix-ops": "error", + "space-unary-ops": "error", + "switch-colon-spacing": "error", + "template-curly-spacing": ["error", "never"], + "unicode-bom": "error" + }, + "globals": { + //script.js + "gradioApp": "readonly", + "onUiLoaded": "readonly", + "onUiUpdate": "readonly", + "onOptionsChanged": "readonly", + "uiCurrentTab": "writable", + "uiElementIsVisible": "readonly", + "uiElementInSight": "readonly", + "executeCallbacks": "readonly", + //ui.js + "opts": "writable", + "all_gallery_buttons": "readonly", + "selected_gallery_button": "readonly", + "selected_gallery_index": "readonly", + "switch_to_txt2img": "readonly", + "switch_to_img2img_tab": "readonly", + "switch_to_img2img": "readonly", + "switch_to_sketch": "readonly", + "switch_to_inpaint": "readonly", + "witch_to_inpaint_sketch": "readonly", + "switch_to_extras": "readonly", + "get_tab_index": "readonly", + "create_submit_args": "readonly", + "restart_reload": "readonly", + "updateInput": "readonly", + //extraNetworks.js + "requestGet": "readonly", + "popup": "readonly", + // from python + "localization": "readonly", + // progrssbar.js + "randomId": "readonly", + "requestProgress": "readonly", + // imageviewer.js + "modalPrevImage": "readonly", + "modalNextImage": "readonly" + } +} diff --git a/.github/ISSUE_TEMPLATE/bug_report.yml b/.github/ISSUE_TEMPLATE/bug_report.yml index f0b02904a..add464bca 100644 --- a/.github/ISSUE_TEMPLATE/bug_report.yml +++ b/.github/ISSUE_TEMPLATE/bug_report.yml @@ -1,7 +1,7 @@ name: Issue Report description: Something is broken title: "[Issue]: " -labels: ["issue"] +labels: [] body: - type: textarea @@ -23,22 +23,19 @@ body: Look at console log and copy the version string from there For example: `Version: f256fb8b Fri Apr 14 17:41:30 2023 -0400` - Additionally provide any relevant platorm information (OS, browser, versions) + Additionally provide any relevant platorm information: + - OS details, GPU information, browser used, app versions - type: markdown attributes: value: | - If issue is setup, installation or startup related, please check `setup.log` before reporting + If issue is setup, installation or startup related, please check `webui.log` before reporting And when posting console logs, please use code blocks ( \`\`\` ) to format them insead of uploading screenshots - - type: markdown - attributes: - value: | - Verify that issue persists after restarting server from command line - Restarting from UI or reloading browser is insufficient to create reproducible environment - type: markdown attributes: value: | If you have additional extensions installed, try to reproduce the issue with user extensions disabled And if the issue is with compatibility with specific extension, mark it as such when creating the issue + Try running with `--safe` command line flag with disables loading of user-installed extensions - type: markdown attributes: value: | diff --git a/.github/ISSUE_TEMPLATE/config.yml b/.github/ISSUE_TEMPLATE/config.yml index afdb9109a..4c439e8cf 100644 --- a/.github/ISSUE_TEMPLATE/config.yml +++ b/.github/ISSUE_TEMPLATE/config.yml @@ -1,5 +1,5 @@ blank_issues_enabled: false contact_links: - - name: WebUI Community Support + - name: SD.Next Community Support url: https://github.com/vladmandic/automatic/discussions about: Please ask and answer questions here. diff --git a/.github/ISSUE_TEMPLATE/extension_report.yml b/.github/ISSUE_TEMPLATE/extension_report.yml new file mode 100644 index 000000000..21e95688a --- /dev/null +++ b/.github/ISSUE_TEMPLATE/extension_report.yml @@ -0,0 +1,33 @@ +name: Extension Compatibility Report +description: Extension is not working as expected +title: "[Extension]: " +labels: ["extension"] + +body: + - type: textarea + id: description + attributes: + label: Issue Description + description: Tell us what happened in a very clear and simple way + value: Please fill this form with as much information as possible + - type: textarea + id: platform + attributes: + label: Version Platform Description + description: Describe your platform (program version, OS, browser) + value: + - type: textarea + id: url + attributes: + label: URL link of the extension + description: URL link of the extension + value: + - type: markdown + attributes: + value: | + Any issues without version information and actual link to extension will be closed + - type: markdown + attributes: + value: | + If issue is extension installation or startup related, please check `webui.log` before reporting + And when posting console logs, please use code blocks ( \`\`\` ) to format them insead of uploading screenshots diff --git a/.gitignore b/.gitignore index 06c8f11c5..ddcab30e3 100644 --- a/.gitignore +++ b/.gitignore @@ -1,18 +1,22 @@ # defaults __pycache__ -/params.txt /cache.json /config.json -/ui-config.json -/setup.log +/params.txt /styles.csv +/ui-config.json /user.css /webui-user.bat /webui-user.sh -/javascript/themes.json +/html/extensions.json +/html/themes.json +node_modules +pnpm-lock.yaml +package-lock.json venv # all models and temp files +*.log *.bak *.ckpt *.safetensors @@ -27,6 +31,8 @@ venv /*.bat /*.sh /*.txt +/*.mp3 +/*.lnk !webui.bat !webui.sh diff --git a/.gitmodules b/.gitmodules index ab36c3699..93009f907 100644 --- a/.gitmodules +++ b/.gitmodules @@ -4,24 +4,10 @@ [submodule "extensions-builtin/sd-extension-system-info"] path = extensions-builtin/sd-extension-system-info url = https://github.com/vladmandic/sd-extension-system-info -[submodule "extensions-builtin/sd-extension-aesthetic-scorer"] - path = extensions-builtin/sd-extension-aesthetic-scorer - url = https://github.com/vladmandic/sd-extension-aesthetic-scorer -[submodule "extensions-builtin/sd-extension-steps-animation"] - path = extensions-builtin/sd-extension-steps-animation - url = https://github.com/vladmandic/sd-extension-steps-animation [submodule "extensions-builtin/stable-diffusion-webui-images-browser"] path = extensions-builtin/stable-diffusion-webui-images-browser url = https://github.com/AlUlkesh/stable-diffusion-webui-images-browser ignore = dirty -[submodule "extensions-builtin/seed_travel"] - path = extensions-builtin/seed_travel - url = https://github.com/yownas/seed_travel - ignore = dirty -[submodule "extensions-builtin/sd-webui-model-converter"] - path = extensions-builtin/sd-webui-model-converter - url = https://github.com/Akegarasu/sd-webui-model-converter - ignore = dirty [submodule "extensions-builtin/sd-dynamic-thresholding"] path = extensions-builtin/sd-dynamic-thresholding url = https://github.com/mcmonkeyprojects/sd-dynamic-thresholding @@ -33,18 +19,27 @@ [submodule "extensions-builtin/clip-interrogator-ext"] path = extensions-builtin/clip-interrogator-ext url = https://github.com/pharmapsychotic/clip-interrogator-ext.git + ignore = dirty [submodule "extensions-builtin/sd-webui-controlnet"] path = extensions-builtin/sd-webui-controlnet url = https://github.com/Mikubill/sd-webui-controlnet + ignore = dirty [submodule "modules/lycoris"] path = modules/lycoris url = https://github.com/KohakuBlueleaf/LyCORIS + ignore = dirty [submodule "extensions-builtin/stable-diffusion-webui-rembg"] path = extensions-builtin/stable-diffusion-webui-rembg url = https://github.com/AUTOMATIC1111/stable-diffusion-webui-rembg + ignore = dirty [submodule "extensions-builtin/a1111-sd-webui-lycoris"] path = extensions-builtin/a1111-sd-webui-lycoris url = https://github.com/KohakuBlueleaf/a1111-sd-webui-lycoris + ignore = dirty [submodule "extensions-builtin/multidiffusion-upscaler-for-automatic1111"] path = extensions-builtin/multidiffusion-upscaler-for-automatic1111 url = https://github.com/pkuliyi2015/multidiffusion-upscaler-for-automatic1111 + ignore = dirty +[submodule "extensions-builtin/sd-webui-agent-scheduler"] + path = extensions-builtin/sd-webui-agent-scheduler + url = https://github.com/ArtVentureX/sd-webui-agent-scheduler diff --git a/.pylintrc b/.pylintrc index 515dce01d..1700d9ec8 100644 --- a/.pylintrc +++ b/.pylintrc @@ -75,16 +75,16 @@ valid-metaclass-classmethod-first-arg=mcs [DESIGN] exclude-too-few-public-methods= ignored-parents= -max-args=5 -max-attributes=7 -max-bool-expr=5 -max-branches=12 -max-locals=15 -max-parents=7 +max-args=99 +max-attributes=99 +max-bool-expr=99 +max-branches=99 +max-locals=99 +max-parents=99 max-public-methods=20 -max-returns=6 -max-statements=50 -min-public-methods=2 +max-returns=99 +max-statements=99 +min-public-methods=1 [EXCEPTIONS] overgeneral-exceptions=builtins.BaseException,builtins.Exception @@ -124,26 +124,29 @@ confidence=HIGH, # disable=C,R,W disable=raw-checker-failed, bad-inline-option, - locally-disabled, - file-ignored, - suppressed-message, - useless-suppression, - deprecated-pragma, - use-symbolic-message-instead, - line-too-long, - missing-function-docstring, - missing-module-docstring, - missing-class-docstring, - logging-fstring-interpolation, - import-outside-toplevel, + bare-except, + broad-exception-caught, consider-iterating-dictionary, - wrong-import-position, - unnecessary-lambda, consider-using-dict-items, dangerous-default-value, - unnecessary-dunder-call, + deprecated-pragma, + file-ignored, + import-outside-toplevel, invalid-name, - R0801, + line-too-long, + locally-disabled, + logging-fstring-interpolation, + missing-class-docstring, + missing-function-docstring, + missing-module-docstring, + duplicate-code, + suppressed-message, + unnecessary-dunder-call, + unnecessary-lambda, + use-symbolic-message-instead, + too-many-nested-blocks, + useless-suppression, + wrong-import-position, enable=c-extension-no-member [METHOD_ARGS] diff --git a/CHANGELOG.md b/CHANGELOG.md new file mode 100644 index 000000000..fbeb78785 --- /dev/null +++ b/CHANGELOG.md @@ -0,0 +1,395 @@ +# Change Log for SD.Next + +## Update for 06/03/2023 + +- new vae decode method to help with larger batch sizes, thanks @bigdog +- profiling of scripts/extensions callbacks +- additional exception handling so bad exception does not crash main app +- additional background removal models + +## Update for 06/02/2023 + +Some quality-of-life improvements while working on larger stuff in the background... + +- redesign action box to be uniform accross all themes +- add **pause** option next to stop/skip +- redesigned progress bar +- add new built-in extension: **agent-scheduler** + very elegant way to getting full queing capabilies, thank @artventurdev +- enable more image formats + note: not all are understood by browser so previews and images may appear as blank + unless you have some browser extensions that can handle them + but they are saved correctly. and cant beat raw quality of 32-bit `tiff` or `psd` :) +- change in behavior: `xformers` will be uninstalled on startup if they are not active + if you do have `xformers` selected as your desired cross-optimization method, then they will be used + reason is that a lot of libaries try to blindy import xformers even if they are not selected or not functional + +## Update for 05/30/2023 + +Another bigger one...And more to come in the next few days... + +- new live preview mode: taesd + i really like this one, so its enabled as default for new installs +- settings search feature +- new sampler: sde++ 2m sde +- fully common save/zip/delete (new) options in all tabs + which (again) meant rework of process image tab +- system info tab: live gpu utilization/memory graphs for nvidia gpus +- updated controlnet interface +- minor style changes +- updated lora, swinir, scunet and ldsr code from upstream +- start of merge from a1111 v1.3 + +## Update for 05/26/2023 + +Some quality-of-life improvements... + +- updated [README](https://github.com/vladmandic/automatic/blob/master/README.md) +- created [CHANGELOG](https://github.com/vladmandic/automatic/blob/master/CHANGELOG.md) + this will be the source for all info about new things moving forward + and cross-posted to [Discussions#99](https://github.com/vladmandic/automatic/discussions/99) as well as discord [announcements](https://discord.com/channels/1101998836328697867/1109953953396957286) +- optimize model loading on startup + this should reduce startup time significantly +- set default cross-optimization method for each platform backend + applicable for new installs only + - `cuda` => Scaled-Dot-Product + - `rocm` => Sub-quadratic + - `directml` => Sub-quadratic + - `ipex` => InvokeAI's + - `mps` => Doggettx's + - `cpu` => Doggettx's +- optimize logging +- optimize profiling + now includes startup profiling as well as `cuda` profiling during generate +- minor lightbox improvements +- bugfixes...i don't recall when was a release with at least several of those + +other than that - first stage of [Diffusers](https://github.com/huggingface/diffusers) integration is now in master branch +i don't recommend anyone to try it (and dont even think reporting issues for it) +but if anyone wants to contribute, take a look at [project page](https://github.com/users/vladmandic/projects/1/views/1) + +## Update for 05/23/2023 + +Major internal work with perhaps not that much user-facing to show for it ;) + +- update core repos: **stability-ai**, **taming-transformers**, **k-diffusion, blip**, **codeformer** + note: to avoid disruptions, this is applicable for new installs only +- tested with **torch 2.1**, **cuda 12.1**, **cudnn 8.9** + (production remains on torch2.0.1+cuda11.8+cudnn8.8) +- fully extend support of `--data-dir` + allows multiple installations to share pretty much everything, not just models + especially useful if you want to run in a stateless container or cloud instance +- redo api authentication + now api authentication will use same user/pwd (if specified) for ui and strictly enforce it using httpbasicauth + new authentication is also fully supported in combination with ssl for both sync and async calls + if you want to use api programatically, see examples in `cli/sdapi.py` +- add dark/light theme mode toggle +- redo some `clip-skip` functionality +- better matching for vae vs model +- update to `xyz grid` to allow creation of large number of images without creating grid itself +- update `gradio` (again) +- more prompt parser optimizations +- better error handling when importing image settings which are not compatible with current install + for example, when upscaler or sampler originally used is not available +- fixes...amazing how many issues were introduced by porting a1111 v1.20 code without adding almost no new functionality + next one is v1.30 (still in dev) which does bring a lot of new features + +## Update for 05/17/2023 + +This is a massive one due to huge number of changes, +but hopefully it will go ok... + +- new **prompt parsers** + select in UI -> Settings -> Stable Diffusion + - **Full**: my new implementation + - **A1111**: for backward compatibility + - **Compel**: as used in ComfyUI and InvokeAI (a.k.a *Temporal Weighting*) + - **Fixed**: for really old backward compatibility +- monitor **extensions** install/startup and + log if they modify any packages/requirements + this is a *deep-experimental* python hack, but i think its worth it as extensions modifying requirements + is one of most common causes of issues +- added `--safe` command line flag mode which skips loading user extensions + please try to use it before opening new issue +- reintroduce `--api-only` mode to start server without ui +- port *all* upstream changes from [A1111](https://github.com/AUTOMATIC1111/stable-diffusion-webui) + up to today - commit hash `89f9faa` + +## Update for 05/15/2023 + +- major work on **prompt parsing** + this can cause some differences in results compared to what you're used to, but its all about fixes & improvements + - prompt parser was adding commas and spaces as separate words and tokens and/or prefixes + - negative prompt weight using `[word:weight]` was ignored, it was always `0.909` + - bracket matching was anything but correct. complex nested attention brackets are now working. + - btw, if you run with `--debug` flag, you'll now actually see parsed prompt & schedule +- updated all scripts in `/cli` +- add option in settings to force different **latent sampler** instead of using primary only +- add **interrupt/skip** capabilities to process images + +## Update for 05/13/2023 + +This is mostly about optimizations... + +- improved `torch-directml` support + especially interesting for **amd** users on **windows** where **torch+rocm** is not yet available + dont forget to run using `--use-directml` or default is **cpu** +- improved compatibility with **nvidia** rtx 1xxx/2xxx series gpus +- fully working `torch.compile` with **torch 2.0.1** + using `inductor` compile takes a while on first run, but does result in 5-10% performance increase +- improved memory handling + for highest performance, you can also disable aggressive **gc** in settings +- improved performance + especially *after* generate as image handling has been moved to separate thread +- allow per-extension updates in extension manager +- option to reset configuration in settings + +## Update for 05/11/2023 + +- brand new **extension manager** + this is pretty much a complete rewrite, so new issues are possible +- support for `torch` 2.0.1 + note that if you are experiencing frequent hangs, this may be a worth a try +- updated `gradio` to 3.29.0 +- added `--reinstall` flag to force reinstall of all packages +- auto-recover & re-attempt when `--upgrade` is requested but fails +- check for duplicate extensions + +## Update for 05/08/2023 + +Back online with few updates: + +- bugfixes. yup, quite a lot of those +- auto-detect some cpu/gpu capabilities on startup + this should reduce need to tweak and tune settings like no-half, no-half-vae, fp16 vs fp32, etc +- configurable order of top level tabs +- configurable order of scripts in txt2img and img2img + for both, see sections in ui-> settings -> user interface + +## Update for 05/04/2023 + +Again, few days later... + +- reviewed/ported **all** commits from **A1111** upstream + some a few are not applicable as i already have alternative implementations + and very few i choose not to implement (save/restore last-known-good-config is a bad hack) + otherwise, we're fully up to date (its doesn't show on fork status as code merges were mostly manual due to conflicts) + but...due to sheer size of the updates, this may introduce some temporary issues +- redesigned server restart function + now available and working in ui + actually, since server restart is now a true restart and not ui restart, it can be used much more flexibly +- faster model load + plus support for slower devices via stream-load function (in ui settings) +- better logging + this includes new `--debug` flag for more verbose logging when troubleshooting + +## Update for 05/01/2023 + +Been a bit quieter for last few days as changes were quite significant, but finally here we are... + +- Updated core libraries: Gradio, Diffusers, Transformers +- Added support for **Intel ARC** GPUs via Intel OneAPI IPEX (auto-detected) +- Added support for **TorchML** (set by default when running on non-compatible GPU or on CPU) +- Enhanced support for AMD GPUs with **ROCm** +- Enhanced support for Apple **M1/M2** +- Redesigned command params: run `webui --help` for details +- Redesigned API and script processing +- Experimental support for multiple **Torch compile** options +- Improved sampler support +- Google Colab: + Maintained by +- Fixes, fixes, fixes... + +To take advantage of new out-of-the-box tunings, its recommended to delete your `config.json` so new defaults are applied. Its not necessary, but otherwise you may need to play with UI Settings to get the best of Intel ARC, TorchML, ROCm or Apple M1/M2. + +## Update for 04/27/2023 + +a bit shorter list as: + +- i've been busy with buxfixing + there are a lot of them, not going to list each here. + but seems like critical issues backlog is quieting down and soon i can focus on new features development. +- i've started collaboration with couple of major projects, + hopefully this will accelerate future development. + +what's new: + +- ability to view/add/edit model description shown in extra networks cards +- add option to specify fallback sampler if primary sampler is not compatible with desired operation +- make clip skip a local parameter +- remove obsolete items from UI settings +- set defaults for AMD ROCm + if you have issues, you may want to start with a fresh install so configuration can be created from scratch +- set defaults for Apple M1/M2 + if you have issues, you may want to start with a fresh install so configuration can be created from scratch + +## Update for 04/25/2023 + +- update process image -> info +- add VAE info to metadata +- update GPU utility search paths for better GPU type detection +- update git flags for wider compatibility +- update environment tuning +- update ti training defaults +- update VAE search paths +- add compatibility opts for some old extensions +- validate script args for always-on scripts + fixes: deforum with controlnet + +## Update for 04/24/2023 + +- identify race condition where generate locks up while fetching preview +- add pulldowns to x/y/z script +- add VAE rollback feature in case of NaNs +- use samples format for live preview +- add token merging +- use **Approx NN** for live preview +- create default `styles.csv` +- fix setup not installing `tensorflow` dependencies +- update default git flags to reduce number of warnings + +## Update for 04/23/2023 + +- fix VAE dtype + should fix most issues with NaN or black images +- add built-in Gradio themes +- reduce requirements +- more AMD specific work +- initial work on Apple platform support +- additional PR merges +- handle torch cuda crashing in setup +- fix setup race conditions +- fix ui lightbox +- mark tensorflow as optional +- add additional image name templates + +## Update for 04/22/2023 + +- autodetect which system libs should be installed + this is a first pass of autoconfig for **nVidia** vs **AMD** environments +- fix parse cmd line args from extensions +- only install `xformers` if actually selected as desired cross-attention method +- do not attempt to use `xformers` or `sdp` if running on cpu +- merge tomesd token merging +- merge 23 PRs pending from a1111 backlog (!!) + +*expect shorter updates for the next few days as i'll be partically ooo* + +## Update for 04/20/2023 + +- full CUDA tuning section in UI Settings +- improve exif/pnginfo metadata parsing + it can now handle 3rd party images or images edited in external software +- optimized setup performance and logging +- improve compatibility with some 3rd party extensions + for example handle extensions that install packages directly from github urls +- fix initial model download if no models found +- fix vae not found issues +- fix multiple git issues + +note: if you previously had command line optimizations such as --no-half, those are now ignored and moved to ui settings + +## Update for 04/19/2023 + +- fix live preview +- fix model merge +- fix handling of user-defined temp folders +- fix submit benchmark +- option to override `torch` and `xformers` installer +- separate benchmark data for system-info extension +- minor css fixes +- created initial merge backlog from pending prs on a1111 repo + see #258 for details + +## Update for 04/18/2023 + +- reconnect ui to active session on browser restart + this is one of most frequently asked for items, finally figured it out + works for text and image generation, but not for process as there is no progress bar reported there to start with +- force unload `xformers` when not used + improves compatibility with AMD/M1 platforms +- add `styles.csv` to UI settings to allow customizing path +- add `--skip-git` to cmd flags for power users that want + to skip all git checks and operations and perform manual updates +- add `--disable-queue` to cmd flags that disables Gradio queues (experimental) + this forces it to use HTTP instead of WebSockets and can help on unreliable network connections +- set scripts & extensions loading priority and allow custom priorities + fixes random extension issues: + `ScuNet` upscaler dissapearing, `Additional Networks` not showing up on XYZ axis, etc. +- improve html loading order +- remove some `asserts` causing runtime errors and replace with user-friendly messages +- update README.md +- update TODO.md + +## Update for 04/17/2023 + +- **themes** are now dynamic and discovered from list of available gradio themes on huggingface + its quite a list of 30+ supported themes so far +- added option to see **theme preview** without the need to apply it or restart server +- integrated **image info** functionality into **process image** tab and removed separate **image info** tab +- more installer improvements +- fix urls +- updated github integration +- make model download as optional if no models found + +## Update for 04/16/2023 + +- support for ui themes! to to *settings* -> *user interface* -> "ui theme* + includes 12 predefined themes +- ability to restart server from ui +- updated requirements +- removed `styles.csv` from repo, its now fully under user control +- removed model-keyword extension as overly aggresive +- rewrite of the fastapi middleware handlers +- install bugfixes, hopefully new installer is now ok \ + i really want to focus on features and not troubleshooting installer + +## Update for 04/15/2023 + +- update default values +- remove `ui-config.json` from repo, its not fully under user control +- updated extensions mangager +- updated locon/lycoris plugin +- enable quick launch by default +- add multidiffusion upscaler extensions +- add model keyword extension +- enable strong linting +- fix circular imports +- fix extensions updated +- fix git update issues +- update github templates + +## Update for 04/14/2023 + +- handle duplicate extensions +- redo exception handler +- fix generate forever +- enable cmdflags compatibility +- change default css font +- fix ti previews on initial start +- enhance tracebacks +- pin transformers version to last known good version +- fix extension loader + +## Update for 04/12/2023 + +This has been pending for a while, but finally uploaded some massive changes + +- New launcher + - `webui.bat` and `webui.sh`: + Platform specific wrapper scripts that starts `launch.py` in Python virtual environment + *Note*: Server can run without virtual environment, but it is recommended to use it + This is carry-over from original repo + **If you're unsure which launcher to use, this is the one you want** + - `launch.py`: + Main startup script + Can be used directly to start server in manually activated `venv` or to run it without `venv` + - `installer.py`: + Main installer, used by `launch.py` + - `webui.py`: + Main server script +- New logger +- New exception handler +- Built-in performance profiler +- New requirements handling +- Move of most of command line flags into UI Settings diff --git a/README.md b/README.md index 5b98cde34..6b06f20df 100644 --- a/README.md +++ b/README.md @@ -3,163 +3,101 @@ ![License](https://img.shields.io/github/license/vladmandic/human?style=flat-square&svg=true) ![GitHub Status Checks](https://img.shields.io/github/checks-status/vladmandic/human/main?style=flat-square&svg=true) -# Stable Diffusion - Automatic -*Heavily opinionated custom fork of* +# SD.Next -Fork is as close as up-to-date with origin as time allows -All code changes are merged upstream whenever possible +**Stable Diffusion implementation with modern UI and advanced features** -The idea behind the fork is to enable latest technologies and advances in text-to-image generation -*Sometimes this is not the same as "as simple as possible to use"* -If you are looking an amazing simple-to-use Stable Diffusion tool, I'd suggest [InvokeAI](https://invoke-ai.github.io/InvokeAI/) specifically due to its automated installer and ease of use +This project started as a form from [Automatic1111 WebUI](https://github.com/AUTOMATIC1111/stable-diffusion-webui/) and it grew siginificantly since then, but although it diverged significanly, any substantial features to original work is ported to this repository as well -
+Individual features are not listed here, instead check [Changelog](CHANGELOG.md) for full list of changes -### Follow [Development updates](https://github.com/vladmandic/automatic/discussions/99) for daily updates on new features/fixes +## Platform support -
- -![screenshot](javascript/black-orange.jpg) - -
- -## Notes - -### Fork does differ in few things - -- New installer -- Advanced CUDA tuning - Available in UI Settings -- Advanced environment tuning -- Optimized startup and models lazy-loading -- Built-in performance profiler -- Updated libraries to latest known compatible versions -- Includes opinionated **System** and **Options** configuration -- Does not rely on `Accelerate` as it only affects distributed systems - Gradio web server will be initialized much earlier which model load is done in the background - Faster model loading plus ability to fallback on corrupt models -- Uses simplified folder structure - e.g. `/train`, `/outputs/*`, `/models/*`, etc. -- Enhanced training templates -- Built-in `LoRA`, `LyCORIS`, `Custom Diffusion`, `Dreambooth` training -- Majority of settings configurable via UI without the need for command line flags - e.g, cross-optimization methods, system folders, etc. -- New logger -- New error and exception handlers - -### Optimizations - -- Optimized for `Torch` 2.0 -- Runs with `SDP` memory attention enabled by default if supported by system - *Note*: `xFormers` and other cross-optimization methods are still available -- Auto-adjust parameters when running on **CPU** or **CUDA** - *Note:* AMD and M1 platforms are supported, but without out-of-the-box optimizations - -### Integrated Extensions - -Hand-picked list of extensions that are deeply integrated into core workflows: - -- [System Info](https://github.com/vladmandic/sd-extension-system-info) -- [ControlNet](https://github.com/Mikubill/sd-webui-controlnet) -- [Image Browser](https://github.com/AlUlkesh/stable-diffusion-webui-images-browser) -- [LORA](https://github.com/kohya-ss/sd-scripts) *(both training and inference)* -- [LyCORIS](https://github.com/KohakuBlueleaf/LyCORIS) *(both training and inference)* -- [Model Converter](https://github.com/Akegarasu/sd-webui-model-converter) -- [CLiP Interrogator](https://github.com/pharmapsychotic/clip-interrogator-ext) -- [Dynamic Thresholding](https://github.com/mcmonkeyprojects/sd-dynamic-thresholding) -- [Steps Animation](https://github.com/vladmandic/sd-extension-steps-animation) -- [Seed Travel](https://github.com/yownas/seed_travel) -- [Multi-Diffusion Upscaler](https://github.com/pkuliyi2015/multidiffusion-upscaler-for-automatic1111) - -### User Interface - -- Includes updated **UI**: reskinned and reorganized - Black and orange dark theme with fixed width options panels and larger previews -- Includes support for **Gradio themes** - *Settings* -> *User interface* -> *UI theme* - Link to themes list & previews: - -### Removed - -- Drops compatibility with older versions of `python` and requires **3.9** or **3.10** -- Drops localizations - -### Integrated CLI/API tools - -Fork adds extra functionality: - -- New skin and UI layout -- Ships with set of **CLI** tools that rely on *SD API* for execution: - e.g. `generate`, `train`, `bench`, etc. - [Full list]() - -
+- **nVidia** GPUs using **CUDA** libraries on both *Windows and Linux* +- **AMD** GPUs using **ROCm** libraries on *Linux* + Support will be extended to *Windows* once AMD releases ROCm for Windows +- Any GPU compatibile with **DirectX** on *Windows* using **DirectML** libraries + This includes support for AMD GPUs that are not supported by native ROCm libraries +- **Intel Arc** GPUs using Intel OneAPI **Ipex/XPU** libraries +- **Apple M1/M2** on *OSX* using built-in support in Torch with **MPS** optimizations ## Install 1. Install first: **Python** & **Git** -2. If you have nVidia GPU, install nVidia CUDA toolkit: - -3. Clone repository +2. Clone repository `git clone https://github.com/vladmandic/automatic` +3. Run launcher + `webui.bat` or `webui.sh`: + - Platform specific wrapper scripts For Windows, Linux and OSX + - Starts `launch.py` in a Python virtual environment (`venv`) + - Uses `install.py` to handle all actual requirements and dependencies + - *Note*: Server can run without virtual environment, but it is recommended to use it to avoid library version conflicts with other applications -## Run +*Note*: **nVidia/CUDA** and **AMD/ROCm** are auto-detected is present and available, but for any other use case specify required parameter explicitly or wrong packages may be installed as installer will assume CPU-only environment -Run desired startup script to install dependencies and extensions and start server: +Full startup sequence is logged in `webui.log`, so if you encounter any issues, please check it first -- `webui.bat` and `webui.sh`: - Platform specific wrapper scripts For Windows, Linux and OSX - Starts `launch.py` in a Python virtual environment (venv) - *Note*: Server can run without virtual environment, but it is recommended to use it to avoid library version conflicts with other applications - **If you're unsure which launcher to use, this is the one you want** -- `launch.py`: - Main startup script - Can be used directly to start server in a manually activated `venv` or to run server without `venv` -- `setup.py`: - Main installer, used by `launch.py` - Can also be used directly to update repository or extensions - If running manually, make sure to activate `venv` first (if used) -- `webui.py`: - Main server script +Below is partial list of all available parameters, run `webui --help` for the full list: -Any of the above scripts can be used with `--help` to display detailed usage information and available parameters -For example: -> webui.bat --help + Setup options: + --use-ipex Use Intel OneAPI XPU backend, default: False + --use-directml Use DirectML if no compatible GPU is detected, default: False + --use-cuda Force use nVidia CUDA backend, default: False + --use-rocm Force use AMD ROCm backend, default: False + --skip-update Skip update of extensions and submodules, default: False + --skip-requirements Skips checking and installing requirements, default: False + --skip-extensions Skips running individual extension installers, default: False + --skip-git Skips running all GIT operations, default: False + --skip-torch Skips running Torch checks, default: False + --reinstall Force reinstallation of all requirements, default: False + --debug Run installer with debug logging, default: False + --reset Reset main repository to latest version, default: False + --upgrade Upgrade main repository to latest version, default: False + --safe Run in safe mode with no user extensions -Full startup sequence is logged in `setup.log`, so if you encounter any issues, please check it first +
![screenshot](html/black-orange.jpg)
-## Update +## Notes -The launcher can perform automatic update of main repository, requirements, extensions and submodules: +### **Collab** -- **Main repository**: - Update is *not* performed by default, enable with `--upgrade` flag -- **Requirements**: - Check is performed on each startup and missing requirements are auto-installed - Can be skipped with `--skip-requirements` flag -- **Extensions and submodules**: - Update is performed on each startup and installer for each extension is started - Can be skipped with `--skip-extensions` flag -- **Quick mode**: Automatically enabled if timestamp of last sucessful setup is newer than actual repository version or version of newest extension - -
- -## Other - -### Scripts - -This repository comes with a large collection of scripts that can be used to process inputs, train, generate, and benchmark models -As well as number of auxiliary scripts that do not rely on **WebUI**, but can be used for end-to-end solutions such as extract frames from videos, etc. -For full details see [Docs](cli/README.md) - -
- -### Docs - -- Scripts are in [Scripts](cli/README.md) -- Everything else is in [Wiki](https://github.com/vladmandic/automatic/wiki) -- Except my current [TODO](TODO.md) +- To avoid having this repo rely just on me, I'd love to have additional maintainers with full admin rights. If you're interested, ping me! +- In addition to general cross-platform code, desire is to have a lead for each of the main platforms +This should be fully cross-platform, but I would really love to have additional contibutors and/or maintainers to join and help lead the effords on different platforms + +### **Goals** + +The idea behind the fork is to enable latest technologies and advances in text-to-image generation +*Sometimes this is not the same as "as simple as possible to use"* +If you are looking an amazing simple-to-use Stable Diffusion tool, I'd suggest [InvokeAI](https://invoke-ai.github.io/InvokeAI/) specifically due to its automated installer and ease of use + +General goals: + +- Cross-platform + - Create uniform experience while automatically managing any platform specific differences +- Performance + - Enable best possible performance on all platforms +- Ease-of-Use + - Automatically handle all requirements, dependencies, flags regardless of platform + - Integrate all best options for uniform out-of-the-box experience without the need to tweak anything manually +- Look-and-Feel + - Create modern, intuitive and clean UI +- Up-to-Date + - Keep code up to date with latest advanced in text-to-image generation + +## Credits + +- Main credit goes to [Automatic1111 WebUI](https://github.com/AUTOMATIC1111/stable-diffusion-webui) +- Additional credits are listed in [Credits](https://github.com/AUTOMATIC1111/stable-diffusion-webui/#credits) +- Licenses for modules are listed in [Licenses](html/licenses.html) + +### **Docs** + +- [Radme](README.md) +- [ToDo](TODO.md) +- [Changelog](CHANGELOG.md) +- [CLI Tools](cli/README.md)
diff --git a/TODO.md b/TODO.md index 0a5b0e99d..e543fe3b6 100644 --- a/TODO.md +++ b/TODO.md @@ -4,35 +4,26 @@ Stuff to be fixed... -- Move Restart Server from WebUI to Launch and reload modules -- Mdularize `cli` scripts ## Features Stuff to be added... -- Update README -- Add Gradio theme maker -- Transformers version -- Create new GitHub hooks/actions for CI/CD -- Redo Extensions tab: see -- Stream-load models as option for slow storage -- Auto-test `torch.layer_norm` for FP16 +- Update `Wiki` +- Create new `GitHub` hooks/actions for CI/CD +- Import core repos +- Update `train.py` to use `interrogator` +- Update `train.py` to use `rembg` ## Investigate Stuff to be investigated... -- Torch Compile -- `Torch-DirectML` -- `TensorRT` -- [Temporal Weighing](https://github.com/comfyanonymous/ComfyUI/discussions/473) - ## Merge PRs Pick & merge PRs from main repo... -- Merge backlog: +- Compare commits: ## Models @@ -50,10 +41,18 @@ Tech that can be integrated as part of the core workflow... - [Null-text inversion](https://github.com/ouhenio/null-text-inversion-colab) - [Custom diffusion](https://github.com/guaneec/custom-diffusion-webui), [Custom diffusion](https://www.cs.cmu.edu/~custom-diffusion/) - [Dream artist](https://github.com/7eu7d7/DreamArtist-sd-webui-extension) +- [QuickEmbedding](https://github.com/ethansmith2000/QuickEmbedding) +- `TensorRT` ## Random - Bunch of stuff: - -### Pending Code Updates - +- +- +- +- shared.info +- hints +- localization +- docker +- port `p.all_hr_prompts` +- test `lyco_patch_lora` diff --git a/cli/.pylintrc b/cli/.pylintrc deleted file mode 100644 index 1231dff7a..000000000 --- a/cli/.pylintrc +++ /dev/null @@ -1,3 +0,0 @@ -# See https://pylint.pycqa.org/en/latest/user_guide/messages/message_control.html -[MESSAGES CONTROL] -enable=C,R,W,E,I diff --git a/cli/README.md b/cli/README.md index 1081e3215..70de255b1 100644 --- a/cli/README.md +++ b/cli/README.md @@ -1,9 +1,6 @@ # Stable-Diffusion Productivity Scripts -*Notes*: -- Offline scripts can be used with or without **Automatic WebUI** -- Online scripts rely on **Automatic WebUI** API which should be started with `--api` parameter -- All scripts have built-in `--help` parameter that can be used to get more information +Note: All scripts have built-in `--help` parameter that can be used to get more information
@@ -18,32 +15,29 @@ Supports upsampling, face restoration and grid creation By default uses parameters from `generate.json` Parameters that are not specified will be randomized: + - Prompt will be dynamically created from template of random samples: `random.json` - Sampler/Scheduler will be randomly picked from available ones - CFG Scale set to 5-10 ### Train -Textual inversion embedding training -> python train-ti.py +Combined pipeline for **embeddings**, **lora**, **lycoris**, **dreambooth** and **hypernetwork** +Optionally runs several image processing steps before training: -Combined pipeline: -1. Creates embedding -2. Extracts images if input is movie -3. Preprocesses images -4. Runs training +- keep original image +- detect and extract face +- detect and extract body +- detect blur +- detect dynamic range +- attempt to upscale low resolution images +- attempt to restore quality of low quality images +- automatically generate captions using interrogate +- resize image +- square image +- run image segmentation to remove background -LoRA training -> python train-lora.py - -Combined pipeline: -1. Preprocesses images -2. Runs training - -[Detailed documentation](https://github.com/vladmandic/automatic/wiki/Process.md) - -LoRA extract from model -> python moidules/lora-extract.py +> python train.py
@@ -51,107 +45,58 @@ LoRA extract from model ### Benchmark -Benchmark your **Automatic WebUI** -Note: Requires SD API +> python run-benchmark.py -> python modules/bench.py +### Create Previews -### Embedding Previews +Create previews for **embeddings**, **lora**, **lycoris**, **dreambooth** and **hypernetwork** -Create previews of embeddings using preview templates -Note: Requires SD API +> python create-previews.py -> python modules/preview-embeddings.py +## Image Grid -## Grid - -Create flexible image grids from any number of images -Note: Offline tool - -> python modiles/grid.py +> python image-grid.py ### Image Watermark Create invisible image watermark and remove existing EXIF tags -Note: Offline tool -> python modules/image-watermark.py +> python image-watermark.py -### Interrogate +### Image Interrogate Runs CLiP and Booru image interrogation -Note: Requires SD API -> python modules/interrogate.py - -### Interrogate-Offline - -Standalone implementation of GiT, CLiP and ViT image interrogation -Note: Offline tool - -> python modules/interrogate-offline.py - -### Models Previews - -Create previews of models using built-in templates -Note: Requires SD API - -> python modules/preview-models.py +> python image-interrogate.py ### Palette Extract Extract color palette from image(s) -Note: Offline tool -> python modules/palette-extract.py - -### Image Process - -Run image processing to extract face/body segments and run resolution/blur/dynamic-range checks -Note: Offline except for interrogate to generate caption files which requires SD API - -> python modules/process.py - -[Detailed documentation](https://github.com/vladmandic/automatic/wiki/Process.md) +> python image-palette.py ### Prompt Ideas Generate complex prompt ideas -Note: Offline tool -> python modules/prompt-ideas.py +> python prompt-ideas.py ### Prompt Promptist Attempts to beautify the provided prompt -Note: Offline tool -> python modules/promptist.py - -### Training Loss-Chart - -Create loss-chart from training log -Note: Offline tool, may require adjustment to train paths if used with other repos - -> python modules/train-losschart.py - -### Training Loss-Rate - -Create customizable loss rate to be used in training -Note: Offline tool - -> python modules/train-lossrate.py +> python prompt-promptist.py ### Video Extract Extract frames from video files -Note: Offline tool -> python modules/video-extract.py +> python video-extract.py
## Utility Scripts + ### SDAPI Utility module that handles async communication to Automatic API endpoints diff --git a/cli/modules/preview-models.py b/cli/create-previews.py similarity index 63% rename from cli/modules/preview-models.py rename to cli/create-previews.py index 6786d46aa..06c4eb097 100755 --- a/cli/modules/preview-models.py +++ b/cli/create-previews.py @@ -1,17 +1,15 @@ -#!/bin/env python +#!/usr/bin/env python import os -import sys import json import time +import importlib import asyncio import argparse from pathlib import Path from util import Map, log from sdapi import get, post, close -from grid import grid - -sys.path.append(os.path.join(os.path.dirname(__file__), '..')) -from generate import sd, generate +from generate import generate # pylint: disable=import-error +grid = importlib.import_module('image-grid').grid default = 'sd-v15-runwayml.ckpt [cc6cb27103]' @@ -63,7 +61,7 @@ options = Map({ "sd_vae": "vae-ft-mse-840000-ema-pruned.ckpt", }, 'lora': { - 'strength': 0.9, + 'strength': 1.0, }, 'hypernetwork': { 'keyword': 'beautiful sexy woman', @@ -72,13 +70,12 @@ options = Map({ }) -async def models(params): - global sd +async def preview_models(params): data = await get('/sdapi/v1/sd-models') - all = [m['title'] for m in data] + allmodels = [m['title'] for m in data] models = [] excluded = [] - for m in all: # loop through all registered models + for m in allmodels: # loop through all registered models ok = True for e in exclude: # check if model is excluded if e in m: @@ -100,30 +97,31 @@ async def models(params): models = filtered log.info({ 'models preview' }) log.info({ 'models': len(models), 'excluded': len(excluded) }) - cmdflags = await get('/sdapi/v1/cmd-flags') opt = await get('/sdapi/v1/options') if params.output != '': - dir = params.output + folder = params.output else: - dir = os.path.abspath(os.path.join(cmdflags['hypernetwork_dir'], '..', 'Stable-diffusion')) - log.info({ 'output directory': dir }) + folder = os.path.abspath(os.path.join(opt['hypernetwork_dir'], '..', 'Stable-diffusion')) + log.info({ 'output directory': folder }) log.info({ 'total jobs': len(models) * options.generate.batch_size, 'per-model': options.generate.batch_size }) log.info(json.dumps(options, indent=2)) for model in models: - fn = os.path.join(dir, os.path.basename(model) + options.format) + fn = os.path.join(folder, os.path.basename(model) + options.format) if os.path.exists(fn) and len(params.input) == 0: # if model preview exists and not manually included log.info({ 'model preview exists': model }) continue log.info({ 'model load': model }) opt['sd_model_checkpoint'] = model + del opt['sd_lora'] + del opt['sd_lyco'] await post('/sdapi/v1/options', opt) opt = await get('/sdapi/v1/options') images = [] labels = [] t0 = time.time() - for label, prompt in prompts: - options.generate.prompt = prompt + for label, p in prompts: + options.generate.prompt = p log.info({ 'model generating': model, 'label': label, 'prompt': options.generate.prompt }) data = await generate(options = options, quiet=True) if 'image' in data: @@ -139,26 +137,28 @@ async def models(params): t = t1 - t0 its = 1.0 * options.generate.steps * len(images) / t log.info({ 'model preview created': model, 'image': fn, 'images': len(images), 'grid': [image.width, image.height], 'time': round(t, 2), 'its': round(its, 2) }) - + opt = await get('/sdapi/v1/options') if opt['sd_model_checkpoint'] != default and not params.fixed: log.info({ 'model set default': default }) opt['sd_model_checkpoint'] = default + del opt['sd_lora'] + del opt['sd_lyco'] await post('/sdapi/v1/options', opt) async def lora(params): - cmdflags = await get('/sdapi/v1/cmd-flags') - dir = cmdflags['lora_dir'] - if not os.path.exists(dir): - log.error({ 'lora directory not found': dir }) + opt = await get('/sdapi/v1/options') + folder = opt['lora_dir'] + if not os.path.exists(folder): + log.error({ 'lora directory not found': folder }) return - models1 = [f for f in Path(dir).glob('*.safetensors')] - models2 = [f for f in Path(dir).glob('*.ckpt')] + models1 = [f for f in Path(folder).glob('*.safetensors')] + models2 = [f for f in Path(folder).glob('*.ckpt')] models = [f.stem for f in models1 + models2] log.info({ 'loras': len(models) }) for model in models: - fn = os.path.join(dir, model + options.format) + fn = os.path.join(folder, model + options.format) if os.path.exists(fn) and len(params.input) == 0: # if model preview exists and not manually included log.info({ 'lora preview exists': model }) continue @@ -166,7 +166,7 @@ async def lora(params): labels = [] t0 = time.time() import re - keywords = re.sub('\d', '', model) + keywords = re.sub(r'\d', '', model) keywords = keywords.replace('-v', ' ').replace('-', ' ').strip().split(' ') keyword = '\"' + '\" \"'.join(keywords) + '\"' options.generate.prompt = prompt.replace('', keyword) @@ -188,16 +188,57 @@ async def lora(params): log.info({ 'lora preview created': model, 'image': fn, 'images': len(images), 'grid': [image.width, image.height], 'time': round(t, 2), 'its': round(its, 2) }) -async def hypernetwork(params): - cmdflags = await get('/sdapi/v1/cmd-flags') - dir = cmdflags['hypernetwork_dir'] - if not os.path.exists(dir): - log.error({ 'hypernetwork directory not found': dir }) +async def lyco(params): + opt = await get('/sdapi/v1/options') + folder = opt['lyco_dir'] + if not os.path.exists(folder): + log.error({ 'lyco directory not found': folder }) return - models = [f.stem for f in Path(dir).glob('*.pt')] - log.info({ 'loras': len(models) }) + models1 = [f for f in Path(folder).glob('*.safetensors')] + models2 = [f for f in Path(folder).glob('*.ckpt')] + models = [f.stem for f in models1 + models2] + log.info({ 'lycos': len(models) }) for model in models: - fn = os.path.join(dir, model + options.format) + fn = os.path.join(folder, model + options.format) + if os.path.exists(fn) and len(params.input) == 0: # if model preview exists and not manually included + log.info({ 'lyco preview exists': model }) + continue + images = [] + labels = [] + t0 = time.time() + import re + keywords = re.sub(r'\d', '', model) + keywords = keywords.replace('-v', ' ').replace('-', ' ').strip().split(' ') + keyword = '\"' + '\" \"'.join(keywords) + '\"' + options.generate.prompt = prompt.replace('', keyword) + options.generate.prompt = options.generate.prompt.replace('', '') + options.generate.prompt += f' ' + log.info({ 'lyco generating': model, 'keyword': keyword, 'prompt': options.generate.prompt }) + data = await generate(options = options, quiet=True) + if 'image' in data: + for img in data['image']: + images.append(img) + labels.append(keyword) + else: + log.error({ 'lyco': model, 'keyword': keyword, 'error': data }) + t1 = time.time() + image = grid(images = images, labels = labels, border = 8) + image.save(fn) + t = t1 - t0 + its = 1.0 * options.generate.steps * len(images) / t + log.info({ 'lyco preview created': model, 'image': fn, 'images': len(images), 'grid': [image.width, image.height], 'time': round(t, 2), 'its': round(its, 2) }) + + +async def hypernetwork(params): + opt = await get('/sdapi/v1/options') + folder = opt['hypernetwork_dir'] + if not os.path.exists(folder): + log.error({ 'hypernetwork directory not found': folder }) + return + models = [f.stem for f in Path(folder).glob('*.pt')] + log.info({ 'hypernetworks': len(models) }) + for model in models: + fn = os.path.join(folder, model + options.format) if os.path.exists(fn) and len(params.input) == 0: # if model preview exists and not manually included log.info({ 'hypernetwork preview exists': model }) continue @@ -224,10 +265,49 @@ async def hypernetwork(params): log.info({ 'hypernetwork preview created': model, 'image': fn, 'images': len(images), 'grid': [image.width, image.height], 'time': round(t, 2), 'its': round(its, 2) }) +async def embedding(params): + opt = await get('/sdapi/v1/options') + folder = opt['embeddings_dir'] + if not os.path.exists(folder): + log.error({ 'embeddings directory not found': folder }) + return + models = [f.stem for f in Path(folder).glob('*.pt')] + log.info({ 'embeddings': len(models) }) + for model in models: + fn = os.path.join(folder, model + '.preview' + options.format) + if os.path.exists(fn) and len(params.input) == 0: # if model preview exists and not manually included + log.info({ 'embedding preview exists': model }) + continue + images = [] + labels = [] + t0 = time.time() + import re + keyword = '\"' + re.sub(r'\d', '', model) + '\"' + options.generate.batch_size = 4 + options.generate.prompt = prompt.replace('', keyword) + options.generate.prompt = options.generate.prompt.replace('', '') + log.info({ 'embedding generating': model, 'keyword': keyword, 'prompt': options.generate.prompt }) + data = await generate(options = options, quiet=True) + if 'image' in data: + for img in data['image']: + images.append(img) + labels.append(keyword) + else: + log.error({ 'lyco': model, 'keyword': keyword, 'error': data }) + t1 = time.time() + image = grid(images = images, labels = labels, border = 8) + image.save(fn) + t = t1 - t0 + its = 1.0 * options.generate.steps * len(images) / t + log.info({ 'embeding preview created': model, 'image': fn, 'images': len(images), 'grid': [image.width, image.height], 'time': round(t, 2), 'its': round(its, 2) }) + + async def create_previews(params): - await models(params) + await preview_models(params) await lora(params) + await lyco(params) await hypernetwork(params) + await embedding(params) await close() @@ -236,5 +316,5 @@ if __name__ == '__main__': parser.add_argument('--output', type = str, default = '', required = False, help = 'output directory') parser.add_argument('--fixed', default = False, action='store_true', help = "do not change model") parser.add_argument('input', type = str, nargs = '*') - params = parser.parse_args() - asyncio.run(create_previews(params)) + args = parser.parse_args() + asyncio.run(create_previews(args)) diff --git a/cli/generate.py b/cli/generate.py index 85860acb3..929a87634 100755 --- a/cli/generate.py +++ b/cli/generate.py @@ -1,4 +1,4 @@ -#!/bin/env python +#!/usr/bin/env python # pylint: disable=no-member """generate batches of images from prompts and upscale them @@ -32,9 +32,8 @@ from PIL import Image from PIL.ExifTags import TAGS from PIL.TiffImagePlugin import ImageFileDirectory_v2 -sys.path.append(os.path.join(os.path.dirname(__file__), 'modules')) -from modules.sdapi import close, get, interrupt, post, session -from modules.util import Map, log, safestring +from sdapi import close, get, interrupt, post, session +from util import Map, log, safestring sd = {} @@ -121,7 +120,7 @@ def sampler(params, options): # find sampler async def generate(prompt = None, options = None, quiet = False): # pylint: disable=redefined-outer-name - global sd + global sd # pylint: disable=global-statement if options: sd = Map(options) if prompt is not None: diff --git a/cli/hfsearch.py b/cli/hfsearch.py new file mode 100644 index 000000000..b0fc9040c --- /dev/null +++ b/cli/hfsearch.py @@ -0,0 +1,17 @@ +import sys +import huggingface_hub as hf +from rich import print # pylint: disable=redefined-builtin + +if __name__ == "__main__": + sys.argv.pop(0) + keyword = sys.argv[0] if len(sys.argv) > 0 else '' + hf_api = hf.HfApi() + model_filter = hf.ModelFilter( + model_name=keyword, + task='text-to-image', + tags='stable-diffusion', + library=['diffusers', 'stable-diffusion'], + ) + res = hf_api.list_models(filter=model_filter, full=True, limit=50, sort="downloads", direction=-1) + models = [{ 'name': m.modelId, 'downloads': m.downloads, 'mtime': m.lastModified, 'url': f'https://huggingface.co/{m.modelId}' } for m in res] + print('Online', models) diff --git a/cli/modules/grid.py b/cli/image-grid.py similarity index 99% rename from cli/modules/grid.py rename to cli/image-grid.py index 4922409b3..fdf3a51e0 100755 --- a/cli/modules/grid.py +++ b/cli/image-grid.py @@ -1,4 +1,4 @@ -#!/bin/env python +#!/usr/bin/env python """ Create image grid """ diff --git a/cli/modules/interrogate.py b/cli/image-interrogate.py similarity index 96% rename from cli/modules/interrogate.py rename to cli/image-interrogate.py index a96c8cf42..211980046 100755 --- a/cli/modules/interrogate.py +++ b/cli/image-interrogate.py @@ -1,4 +1,4 @@ -#!/bin/env python +#!/usr/bin/env python """ use clip to interrogate image(s) """ @@ -11,7 +11,7 @@ import asyncio import filetype from PIL import Image from util import log, Map -import sdapi as sdapi +import sdapi stats = { 'captions': {}, 'keywords': {} } @@ -96,7 +96,7 @@ async def main(): elif os.path.isdir(arg): for root, _dirs, files in os.walk(arg): for f in files: - caption, keywords, _style = await interrogate(os.path.join(root, f)) + _caption, _keywords, _style = await interrogate(os.path.join(root, f)) else: log.error({ 'interrogate unknown file type': arg }) else: diff --git a/cli/modules/palette-extract.py b/cli/image-palette.py similarity index 82% rename from cli/modules/palette-extract.py rename to cli/image-palette.py index cef0d4e0a..77eb5e130 100755 --- a/cli/modules/palette-extract.py +++ b/cli/image-palette.py @@ -1,27 +1,27 @@ -#!/bin/env python +#!/usr/bin/env python # based on import os import io import pathlib import argparse +import importlib import pandas as pd import numpy as np import extcolors import filetype import matplotlib.pyplot as plt import matplotlib.patches as patches -import matplotlib.image as mpimg from matplotlib.offsetbox import OffsetImage, AnnotationBbox from colormap import rgb2hex from PIL import Image from util import log -from grid import grid +grid = importlib.import_module('image-grid').grid -def color_to_df(input): - colors_pre_list = str(input).replace('([(','').split(', (')[0:-1] +def color_to_df(param): + colors_pre_list = str(param).replace('([(','').split(', (')[0:-1] df_rgb = [i.split('), ')[0] + ')' for i in colors_pre_list] - df_percent = [i.split('), ')[1].replace(')','') for i in colors_pre_list] + df_percent = [i.split('), ')[1].replace(')','') for i in colors_pre_list] #convert RGB to HEX code df_color_up = [rgb2hex(int(i.split(", ")[0].replace("(","")), int(i.split(", ")[1]), @@ -30,14 +30,14 @@ def color_to_df(input): return df -def palette(img, args, output): +def palette(img, params, output): size = 1024 img.thumbnail((size, size), Image.HAMMING) - + #crate dataframe - colors_x = extcolors.extract_from_image(img, tolerance = args.color, limit = 13) + colors_x = extcolors.extract_from_image(img, tolerance = params.color, limit = 13) df_color = color_to_df(colors_x) - + #annotate text list_color = list(df_color['c_code']) list_precent = [int(i) for i in list(df_color['occurence'])] @@ -54,7 +54,7 @@ def palette(img, args, output): imagebox = OffsetImage(data, zoom=2.5) ab = AnnotationBbox(imagebox, (0, 0)) ax1.add_artist(ab) - + #color palette x_posi, y_posi, y_posi2 = 160, -260, -260 for c in list_color: @@ -100,20 +100,20 @@ if __name__ == '__main__': args = parser.parse_args() log.info({ 'palette args': vars(args) }) if args.output != '': - pathlib.Path(args.output).mkdir(parents = True, exist_ok = True) + pathlib.Path(args.output).mkdir(parents = True, exist_ok = True) if not args.grid: for arg in args.input: if os.path.isfile(arg) and filetype.is_image(arg): - img = Image.open(arg) - output = os.path.join(args.output, pathlib.Path(arg).stem + '-' + args.suffix + '.jpg') - palette(img, args, output) + image = Image.open(arg) + fn = os.path.join(args.output, pathlib.Path(arg).stem + '-' + args.suffix + '.jpg') + palette(image, args, fn) elif os.path.isdir(arg): for root, _dirs, files in os.walk(arg): for f in files: if filetype.is_image(os.path.join(root, f)): - img = Image.open(os.path.join(root, f)) - output = os.path.join(args.output, pathlib.Path(f).stem + '-' + args.suffix + '.jpg') - palette(img, args, output) + image = Image.open(os.path.join(root, f)) + fn = os.path.join(args.output, pathlib.Path(f).stem + '-' + args.suffix + '.jpg') + palette(image, args, fn) else: images = [] for arg in args.input: @@ -124,6 +124,6 @@ if __name__ == '__main__': for f in files: if filetype.is_image(os.path.join(root, f)): images.append(Image.open(os.path.join(root, f))) - img = grid(images) - output = os.path.join(args.output, args.suffix + '.jpg') - palette(img, args, output) + image = grid(images) + fn = os.path.join(args.output, args.suffix + '.jpg') + palette(image, args, fn) diff --git a/cli/modules/image-watermark.py b/cli/image-watermark.py similarity index 68% rename from cli/modules/image-watermark.py rename to cli/image-watermark.py index 73e891238..16f3fb14a 100755 --- a/cli/modules/image-watermark.py +++ b/cli/image-watermark.py @@ -1,4 +1,4 @@ -#!/bin/env python +#!/usr/bin/env python import os import io import pathlib @@ -44,31 +44,31 @@ def set_exif(d: dict): ifd[_TAGS[k]] = v exif_stream = io.BytesIO() ifd.save(exif_stream) - bytes = b'Exif\x00\x00' + exif_stream.getvalue() - return bytes + encoded = b'Exif\x00\x00' + exif_stream.getvalue() + return encoded -def get_watermark(image, args): +def get_watermark(image, params): data = np.asarray(image) - decoder = WatermarkDecoder(options.type, args.length) - bytes = decoder.decode(data, options.method) + decoder = WatermarkDecoder(options.type, params.length) + decoded = decoder.decode(data, options.method) try: - watermark = str(bytes, 'UTF-8').replace('\x00', '') + s = str(decoded, 'UTF-8').replace('\x00', '') except: - watermark = '' - return watermark + s = '' + return s -def set_watermark(image, args): +def set_watermark(image, params): data = np.asarray(image) encoder = WatermarkEncoder() - encoder.set_watermark(options.type, args.wm.encode('utf-8')) + encoder.set_watermark(options.type, params.wm.encode('utf-8')) encoded = encoder.encode(data, options.method) image = Image.fromarray(encoded) return image -def watermark(args, file): +def watermark(params, file): if not os.path.exists(file): log.error({ 'watermark': 'file not found' }) return @@ -82,30 +82,30 @@ def watermark(args, file): exif = get_exif(image) - if args.command == 'read': - watermark = get_watermark(image, args) - log.info({ 'file': file, 'watermark': watermark, 'exif': exif, 'resolution': f'{image.width}x{image.height}' }) + if params.command == 'read': + wm = get_watermark(image, params) + log.info({ 'file': file, 'watermark': wm, 'exif': exif, 'resolution': f'{image.width}x{image.height}' }) - elif args.command == 'write': - metadata = b'' if args.strip else set_exif(exif) - if args.output != '': - pathlib.Path(args.output).mkdir(parents = True, exist_ok = True) - image=set_watermark(image, args) - fn = os.path.join(args.output, file) + elif params.command == 'write': + metadata = b'' if params.strip else set_exif(exif) + if params.output != '': + pathlib.Path(params.output).mkdir(parents = True, exist_ok = True) + image=set_watermark(image, params) + fn = os.path.join(params.output, file) image.save(fn, exif=metadata) - if args.verify: + if params.verify: data = np.asarray(image) - decoder = WatermarkDecoder(options.type, args.length) - bytes = decoder.decode(data, options.method) - if bytes.startswith(b'\xff'): - watermark = '' + decoder = WatermarkDecoder(options.type, params.length) + decoded = decoder.decode(data, options.method) + if decoded.startswith(b'\xff'): + wm = '' else: - watermark = str(bytes, 'UTF-8').replace('\x00', '') + wm = str(decoded, 'UTF-8').replace('\x00', '') else: - watermark = args.wm + wm = params.wm - log.info({ 'file': fn, 'watermark': watermark, 'exif': None if args.strip else exif, 'resolution': f'{image.width}x{image.height}' }) + log.info({ 'file': fn, 'watermark': wm, 'exif': None if params.strip else exif, 'resolution': f'{image.width}x{image.height}' }) if __name__ == '__main__': diff --git a/cli/train/latents.py b/cli/latents.py similarity index 94% rename from cli/train/latents.py rename to cli/latents.py index 94249b18a..53be33527 100755 --- a/cli/train/latents.py +++ b/cli/latents.py @@ -1,4 +1,4 @@ -#!/bin/env python +#!/usr/bin/env python import os import sys @@ -23,7 +23,7 @@ console = Console(log_time=True, log_time_format='%H:%M:%S-%f') pretty_install(console=console) traceback_install(console=console, extra_lines=1, width=console.width, word_wrap=False, indent_guides=False) -sys.path.append(os.path.join(os.path.dirname(__file__), '..', '..', 'modules', 'lora')) +sys.path.append(os.path.join(os.path.dirname(__file__), '..', 'modules', 'lora')) import library.model_util as model_util import library.train_util as train_util @@ -44,13 +44,13 @@ options = Map({ vae = None -def get_latents(vae, images, weight_dtype): +def get_latents(local_vae, images, weight_dtype): image_transforms = transforms.Compose([ transforms.ToTensor(), transforms.Normalize([0.5], [0.5]) ]) img_tensors = [image_transforms(image) for image in images] img_tensors = torch.stack(img_tensors) img_tensors = img_tensors.to(device, weight_dtype) with torch.no_grad(): - latents = vae.encode(img_tensors).latent_dist.sample().float().to('cpu').numpy() + latents = local_vae.encode(img_tensors).latent_dist.sample().float().to('cpu').numpy() return latents @@ -58,8 +58,8 @@ def get_npz_filename_wo_ext(data_dir, image_key): return os.path.join(data_dir, os.path.splitext(os.path.basename(image_key))[0]) -def create_vae_latents(params): - args = Map({**options, **params}) +def create_vae_latents(local_params): + args = Map({**options, **local_params}) console.log(f'create vae latents args: {args}') image_paths = train_util.glob_images(args.input) if os.path.exists(args.json): @@ -73,7 +73,7 @@ def create_vae_latents(params): weight_dtype = torch.bfloat16 else: weight_dtype = torch.float32 - global vae + global vae # pylint: disable=global-statement if vae is None: vae = model_util.load_vae(args.vae, weight_dtype) vae.eval() @@ -142,7 +142,7 @@ def create_vae_latents(params): def unload_vae(): - global vae + global vae # pylint: disable=global-statement vae = None diff --git a/cli/modules/interrogate-offline.py b/cli/modules/interrogate-offline.py deleted file mode 100755 index 6d9ae56fa..000000000 --- a/cli/modules/interrogate-offline.py +++ /dev/null @@ -1,166 +0,0 @@ -#!/bin/env python - -import os -import gc -import json -import time -import argparse -import torch -import filetype -from PIL import Image -import transformers -from transformers import AutoProcessor, AutoModelForCausalLM -from transformers import BlipProcessor, BlipForConditionalGeneration -from transformers import VisionEncoderDecoderModel, ViTFeatureExtractor, AutoTokenizer -from util import log, Map - - -model = None -processor = None -extractor = None -dtype = torch.float32 -device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') - -options = Map({ - 'input': '', - 'min': 8, - 'max': 256, - 'beams': 1, - 'json': '', - 'txt': False, - 'tag': '', - 'git': True, - 'blip': True, - 'precision': 'fp16', - 'model': 'git', -}) - - -def cleanup(s: str): - s = s.split('"')[0].split('.')[0].split(' that')[0] - s = s.split(' with a letter')[0].split(' with the number')[0].split(' with the word')[0] - s = s.replace('arafed image of ', '') - return s.replace('a ', '') - - -def load_model(args): - global model - global processor - global extractor - transformers.logging.set_verbosity_error() - if args.model == 'git': - model_name = "microsoft/git-large-textcaps" - if model is None: - model = AutoModelForCausalLM.from_pretrained(model_name, torch_dtype=dtype) - model.to(device) - processor = AutoProcessor.from_pretrained(model_name, torch_dtype=dtype) - log.info( { 'interrogate loaded model': model_name }) - elif args.model == 'blip': - model_name = "Salesforce/blip-image-captioning-large" - if model is None: - model = BlipForConditionalGeneration.from_pretrained(model_name, torch_dtype=dtype) - model.to(device) - processor = BlipProcessor.from_pretrained(model_name, torch_dtype=dtype) - log.info( { 'interrogate loaded model': model_name }) - elif args.model == 'vit': - model_name = "nlpconnect/vit-gpt2-image-captioning" - if model is None: - model = VisionEncoderDecoderModel.from_pretrained(model_name, torch_dtype=dtype) - model.to(device) - extractor = ViTFeatureExtractor.from_pretrained(model_name, torch_dtype=dtype) - processor = AutoTokenizer.from_pretrained(model_name, torch_dtype=dtype) - log.info( { 'interrogate loaded model': model_name }) - else: - log.info( { 'interrogate unknown model': args.model }) - - -def interrogate_files(params, files): - args = Map({**options, **params}) - data = [f for f in files if filetype.is_image(f)] - log.info({ 'interrogate files': len(files), 'images': len(data), 'args': args }) - load_model(args) - metadata = {} - for image_path in data: - image = Image.open(image_path).convert('RGB') - caption = '' - if args.model == 'git': - inputs = processor(images=[image], return_tensors="pt").to(device) - ids = model.generate(pixel_values=inputs.pixel_values, num_beams=args.beams, min_length=args.min, max_length=args.max) - caption = processor.batch_decode(ids, skip_special_tokens=True)[0] - elif args.model == 'blip': - inputs = processor(image, return_tensors="pt").to(device, dtype) - ids = model.generate(**inputs, num_beams=args.beams, min_length=args.min, max_length=args.max) - caption = processor.decode(ids[0], skip_special_tokens=True) - elif args.model == 'vit': - inputs = extractor(images=[image], return_tensors="pt").pixel_values.to(device) - ids = model.generate(inputs, num_beams=args.beams, min_length=args.min, max_length=args.max) - caption = processor.batch_decode(ids, skip_special_tokens=True)[0] - else: - log.error({ 'interrogate unknown model': args.model }) - - caption = cleanup(caption) - tags = '' - if args.tag != '': - tags += args.tag + ',' - tags += caption.split(' ')[0] - if args.txt: - with open(os.path.splitext(image_path)[0] + '.txt', "wt", encoding='utf-8') as f: - f.write(caption + "\n") - metadata[image_path] = { 'caption': caption, 'tags': tags } - log.info({ 'interrogate image': image_path, 'moodel': args.model, 'caption': caption, 'tags': tags }) - - if args.json != '': - with open(args.json, "wt", encoding='utf-8') as f: - f.write(json.dumps(metadata, indent=2) + "\n") - return metadata - - -def unload_model(): - global processor - global model - global extractor - if model is not None: - del model - model = None - if processor is not None: - del processor - processor = None - if extractor is not None: - del extractor - extractor = None - gc.collect() - if torch.cuda.is_available(): - with torch.no_grad(): - torch.cuda.empty_cache() - with torch.cuda.device('cuda'): - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - -if __name__ == '__main__': - parser = argparse.ArgumentParser(description = 'image interrogate') - parser.add_argument('input', type=str, nargs='*', help='input file or directory') - parser.add_argument('--model', default = 'git', choices = ['git', 'blip', 'vit'], help = "which model to use") - parser.add_argument("--min", type=int, default=8, help="min length of caption") - parser.add_argument("--max", type=int, default=256, help="max length of caption") - parser.add_argument("--beams", type=int, default=1, help="number of beams to use") - parser.add_argument("--json", type=str, default='', help="output json file") - parser.add_argument("--tag", type=str, default='', help="append tag") - parser.add_argument('--txt', default = False, action='store_true', help = "write captions to text files") - params = parser.parse_args() - log.info({ 'interrogate args': vars(params) }) - if len(params.input) == 0: - parser.print_help() - exit(1) - files = [] - for loc in params.input: - if os.path.isfile(loc): - files.append(loc) - elif os.path.isdir(loc): - for root, _sub_dirs, dir in os.walk(loc): - files = [os.path.join(root, f) for f in dir] - t0 = time.time() - metadata = interrogate_files(vars(params), files) - t1 = time.time() - log.info({ 'interrogate files': len(files), 'time': round(t1 - t0, 2) }) - unload_model() diff --git a/cli/modules/lora-extract.py b/cli/modules/lora-extract.py deleted file mode 100755 index 102728308..000000000 --- a/cli/modules/lora-extract.py +++ /dev/null @@ -1,144 +0,0 @@ -#!/bin/env python - -""" -Extract approximating LoRA by SVD from two SD models -Based on: -""" - -import os -import sys -import time -import argparse -import torch -import transformers -from tqdm import tqdm -from util import log - -sys.path.append(os.path.join(os.path.dirname(__file__), '..', '..', 'modules', 'lora')) -import library.model_util as model_util -import networks.lora as lora - - -def svd(args): # pylint: disable=redefined-outer-name - device = 'cuda' if torch.cuda.is_available() and args.device == 'cuda' else 'cpu' - transformers.logging.set_verbosity_error() - CLAMP_QUANTILE = 0.99 - MIN_DIFF = 1e-6 - if args.precision == 'fp32': - save_dtype = torch.float - elif args.precision == 'fp16': - save_dtype = torch.float16 - elif args.precision == 'bf16': - save_dtype = torch.bfloat16 - else: - save_dtype = None - t0 = time.time() - log.info({ 'loading model': args.original }) - text_encoder_o, _, unet_o = model_util.load_models_from_stable_diffusion_checkpoint(args.v2, args.original) - log.info({ 'loading model': args.tuned }) - text_encoder_t, _, unet_t = model_util.load_models_from_stable_diffusion_checkpoint(args.v2, args.tuned) - with torch.no_grad(): - torch.cuda.empty_cache() - # create LoRA network to extract weights: Use dim (rank) as alpha - lora_network_o = lora.create_network(1.0, args.dim, args.dim, None, text_encoder_o, unet_o) - lora_network_t = lora.create_network(1.0, args.dim, args.dim, None, text_encoder_t, unet_t) - assert len(lora_network_o.text_encoder_loras) == len(lora_network_t.text_encoder_loras), 'model version is different' - # get diffs - diffs = {} - text_encoder_different = False - for i, (lora_o, lora_t) in enumerate(zip(lora_network_o.text_encoder_loras, lora_network_t.text_encoder_loras)): - lora_name = lora_o.lora_name - module_o = lora_o.org_module - module_t = lora_t.org_module - diff = module_t.weight - module_o.weight - # Text Encoder might be same - if torch.max(torch.abs(diff)) > MIN_DIFF: - text_encoder_different = True - diff = diff.float() - diffs[lora_name] = diff - - if not text_encoder_different: - log.info({ 'lora': 'text encoder is same, extract U-Net only' }) - lora_network_o.text_encoder_loras = [] - diffs = {} - - for i, (lora_o, lora_t) in enumerate(zip(lora_network_o.unet_loras, lora_network_t.unet_loras)): - lora_name = lora_o.lora_name - module_o = lora_o.org_module - module_t = lora_t.org_module - diff = module_t.weight - module_o.weight - diff = diff.float() - diff = diff.to(device) - diffs[lora_name] = diff - t1 = time.time() - log.info({ 'lora models': 'ready', 'time': round(t1 - t0, 2) }) - - # make LoRA with svd - log.info({ 'lora': 'calculating by svd' }) - rank = args.dim - lora_weights = {} - with torch.no_grad(): - for lora_name, mat in tqdm(list(diffs.items())): - conv2d = len(mat.size()) == 4 - if conv2d: - mat = mat.squeeze() - U, S, Vh = torch.linalg.svd(mat) - U = U[:, :rank] - S = S[:rank] - U = U @ torch.diag(S) - Vh = Vh[:rank, :] - dist = torch.cat([U.flatten(), Vh.flatten()]) - hi_val = torch.quantile(dist, CLAMP_QUANTILE) - low_val = -hi_val - U = U.clamp(low_val, hi_val) - Vh = Vh.clamp(low_val, hi_val) - lora_weights[lora_name] = (U, Vh) - t2 = time.time() - - # make state dict for LoRA - lora_network_o.apply_to(text_encoder_o, unet_o, text_encoder_different, True) - lora_sd = lora_network_o.state_dict() - log.info({ 'lora extracted weights': len(lora_sd), 'time': round(t2 - t1, 2) }) - - for key in list(lora_sd.keys()): - if 'alpha' in key: - continue - lora_name = key.split('.')[0] - i = 0 if 'lora_up' in key else 1 - weights = lora_weights[lora_name][i] - # print(key, i, weights.size(), lora_sd[key].size()) - if len(lora_sd[key].size()) == 4: # pylint: disable=unsubscriptable-object - weights = weights.unsqueeze(2).unsqueeze(3) - assert weights.size() == lora_sd[key].size(), f'size unmatch: {key}' # pylint: disable=unsubscriptable-object - lora_sd[key] = weights # pylint: disable=unsupported-assignment-operation - - # load state dict to LoRA and save it - info = lora_network_o.load_state_dict(lora_sd) - log.info({ 'lora loading extracted weights': info }) - - dir_name = os.path.dirname(args.save) - if dir_name and not os.path.exists(dir_name): - os.makedirs(dir_name, exist_ok=True) - - # minimum metadata - metadata = {'ss_network_dim': str(args.dim), 'ss_network_alpha': str(args.dim)} - lora_network_o.save_weights(args.save, save_dtype, metadata) - t3 = time.time() - log.info({ 'lora saved weights': args.save, 'time': round(t3 - t2, 2) }) - - -if __name__ == '__main__': - parser = argparse.ArgumentParser(description = 'extract lora weights') - parser.add_argument('--v2', action='store_true', help='load Stable Diffusion v2.x model / Stable Diffusion') - parser.add_argument('--precision', type=str, default='fp16', choices=[None, 'fp32', 'fp16', 'bf16'], help='precision in saving, same to merging if omitted') - parser.add_argument('--device', type=str, default='cuda', choices=['cpu', 'cuda'], help='use cpu or cuda if available') - parser.add_argument('--original', type=str, default=None, required=True, help='Stable Diffusion original model: ckpt or safetensors file') - parser.add_argument('--tuned', type=str, default=None, required=True, help='Stable Diffusion tuned model, LoRA is difference of `original to tuned`: ckpt or safetensors file') - parser.add_argument('--save', type=str, default=None, required=True, help='destination file name: ckpt or safetensors file') - parser.add_argument('--dim', type=int, default=4, help='dimension (rank) of LoRA') - args = parser.parse_args() - log.info({ 'extract lora args': vars(args) }) - if not os.path.exists(args.original) or not os.path.exists(args.tuned): - log.error({ 'models not found': [args.original, args.tuned] }) - else: - svd(args) diff --git a/cli/modules/lora-latents.py b/cli/modules/lora-latents.py deleted file mode 100755 index d556d596b..000000000 --- a/cli/modules/lora-latents.py +++ /dev/null @@ -1,160 +0,0 @@ -#!/bin/env python - -import os -import sys -import json -import pathlib -import argparse -import warnings - -import cv2 -import numpy as np -import torch -from PIL import Image -from torchvision import transforms -from tqdm import tqdm -from util import log, Map - -sys.path.append(os.path.join(os.path.dirname(__file__), '..', '..', 'modules', 'lora')) -import library.model_util as model_util -import library.train_util as train_util - -warnings.filterwarnings('ignore') -device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') -options = Map({ - 'batch': 1, - 'input': '', - 'json': '', - 'max': 1024, - 'min': 256, - 'noupscale': False, - 'precision': 'fp32', - 'resolution': '512,512', - 'steps': 64, - 'vae': 'stabilityai/sd-vae-ft-mse' -}) -vae = None - - -def get_latents(vae, images, weight_dtype): - image_transforms = transforms.Compose([ transforms.ToTensor(), transforms.Normalize([0.5], [0.5]) ]) - img_tensors = [image_transforms(image) for image in images] - img_tensors = torch.stack(img_tensors) - img_tensors = img_tensors.to(device, weight_dtype) - with torch.no_grad(): - latents = vae.encode(img_tensors).latent_dist.sample().float().to('cpu').numpy() - return latents - - -def get_npz_filename_wo_ext(data_dir, image_key): - return os.path.join(data_dir, os.path.splitext(os.path.basename(image_key))[0]) - - -def create_vae_latents(params): - args = Map({**options, **params}) - log.info({ 'latents args': args }) - if args.steps % 8 > 0: - log.warning({ 'latents': 'resolution is not multiple of 8' }) - image_paths = train_util.glob_images(args.input) - if os.path.exists(args.json): - log.info({ 'latents metadata': args.json, 'images': len(image_paths) }) - with open(args.json, 'rt', encoding='utf-8') as f: - metadata = json.load(f) - else: - log.error({ 'latents metadata missing': args.json }) - return - if args.precision == 'fp16': - weight_dtype = torch.float16 - elif args.precision == 'bf16': - weight_dtype = torch.bfloat16 - else: - weight_dtype = torch.float32 - global vae - if vae is None: - vae = model_util.load_vae(args.vae, weight_dtype) - vae.eval() - vae.to(device, dtype=weight_dtype) - max_reso = tuple([int(t) for t in args.resolution.split(',')]) - assert len(max_reso) == 2, f'illegal resolution: {args.resolution}' - bucket_manager = train_util.BucketManager(args.noupscale, max_reso, args.min, args.max, args.steps) - if not args.noupscale: - bucket_manager.make_buckets() - else: - log.warning({ 'latents': 'min and max are ignored if noupscale is set' }) - img_ar_errors = [] - def process_batch(is_last): - for bucket in bucket_manager.buckets: - if (is_last and len(bucket) > 0) or len(bucket) >= args.batch: - latents = get_latents(vae, [img for _, img in bucket], weight_dtype) - assert latents.shape[2] == bucket[0][1].shape[0] // 8 and latents.shape[3] == bucket[0][1].shape[1] // 8, f'latent shape {latents.shape}, {bucket[0][1].shape}' - for (image_key, _), latent in zip(bucket, latents): - npz_file_name = get_npz_filename_wo_ext(args.input, image_key) - np.savez(npz_file_name, latent) - bucket.clear() - data = [[(None, ip)] for ip in image_paths] - bucket_counts = {} - for data_entry in tqdm(data, smoothing=0.0): - if data_entry[0] is None: - continue - img_tensor, image_path = data_entry[0] - if img_tensor is not None: - image = transforms.functional.to_pil_image(img_tensor) - else: - image = Image.open(image_path) - image_key = os.path.basename(image_path) - image_key = os.path.join(os.path.basename(pathlib.Path(image_path).parent), pathlib.Path(image_path).stem) - if image_key not in metadata: - metadata[image_key] = {} - reso, resized_size, ar_error = bucket_manager.select_bucket(image.width, image.height) - img_ar_errors.append(abs(ar_error)) - bucket_counts[reso] = bucket_counts.get(reso, 0) + 1 - metadata[image_key]['train_resolution'] = (reso[0] - reso[0] % 8, reso[1] - reso[1] % 8) - if not args.noupscale: - assert resized_size[0] == reso[0] or resized_size[1] == reso[1], f'internal error, resized size not match: {reso}, {resized_size}, {image.width}, {image.height}' - assert resized_size[0] >= reso[0] and resized_size[1] >= reso[1], f'internal error, resized size too small: {reso}, {resized_size}, {image.width}, {image.height}' - assert resized_size[0] >= reso[0] and resized_size[1] >= reso[1], f'internal error resized size is small: {resized_size}, {reso}' - image = np.array(image) - if resized_size[0] != image.shape[1] or resized_size[1] != image.shape[0]: - image = cv2.resize(image, resized_size, interpolation=cv2.INTER_AREA) - if resized_size[0] > reso[0]: - trim_size = resized_size[0] - reso[0] - image = image[:, trim_size//2:trim_size//2 + reso[0]] - if resized_size[1] > reso[1]: - trim_size = resized_size[1] - reso[1] - image = image[trim_size//2:trim_size//2 + reso[1]] - assert image.shape[0] == reso[1] and image.shape[1] == reso[0], f'internal error, illegal trimmed size: {image.shape}, {reso}' - bucket_manager.add_image(reso, (image_key, image)) - process_batch(False) - - process_batch(True) - vae.to('cpu') - - bucket_manager.sort() - img_ar_errors = np.array(img_ar_errors) - for i, reso in enumerate(bucket_manager.resos): - count = bucket_counts.get(reso, 0) - if count > 0: - log.info({ 'latents bucket': i, 'resolution': reso, 'count': count, 'mean ar error': np.mean(img_ar_errors) }) - with open(args.json, 'wt', encoding='utf-8') as f: - json.dump(metadata, f, indent=2) - - -def unload_vae(): - global vae - vae = None - - -if __name__ == '__main__': - parser = argparse.ArgumentParser() - parser.add_argument('input', type=str, help='directory for train images') - parser.add_argument('--json', type=str, required=True, help='metadata file to input') - parser.add_argument('--vae', type=str, required=True, help='model name or path to encode latents') - parser.add_argument('--batch', type=int, default=1, help='batch size in inference') - parser.add_argument('--resolution', type=str, default='512,512', help='max resolution in fine tuning (width,height)') - parser.add_argument('--min', type=int, default=256, help='minimum resolution for buckets') - parser.add_argument('--max', type=int, default=1024, help='maximum resolution for buckets') - parser.add_argument('--steps', type=int, default=64, help='steps of resolution for buckets, divisible by 8') - parser.add_argument('--noupscale', action='store_true', help='make bucket for each image without upscaling') - parser.add_argument('--precision', type=str, default='fp32', choices=['fp32', 'fp16', 'bf16'], help='use precision') - params = parser.parse_args() - create_vae_latents(vars(params)) diff --git a/cli/modules/models-diff.py b/cli/modules/models-diff.py deleted file mode 100755 index f491926e9..000000000 --- a/cli/modules/models-diff.py +++ /dev/null @@ -1,74 +0,0 @@ -#!/bin/env python -# based on - -import safetensors -import sys -import torch -from pathlib import Path -import torch.nn as nn -import torch.nn.functional as F -import warnings -from util import log - -warnings.filterwarnings("ignore", category=UserWarning) - -def cal_cross_attn(to_q, to_k, to_v, rand_input): - hidden_dim, embed_dim = to_q.shape - attn_to_q = nn.Linear(hidden_dim, embed_dim, bias=False) - attn_to_k = nn.Linear(hidden_dim, embed_dim, bias=False) - attn_to_v = nn.Linear(hidden_dim, embed_dim, bias=False) - attn_to_q.load_state_dict({"weight": to_q}) - attn_to_k.load_state_dict({"weight": to_k}) - attn_to_v.load_state_dict({"weight": to_v}) - - return torch.einsum( - "ik, jk -> ik", - F.softmax(torch.einsum("ij, kj -> ik", attn_to_q(rand_input), attn_to_k(rand_input)), dim=-1), - attn_to_v(rand_input) - ) - -def load_model(path): - if path.suffix == ".safetensors": - return safetensors.torch.load_file(path, device="cpu") - else: - ckpt = torch.load(path, map_location="cpu") - return ckpt["state_dict"] if "state_dict" in ckpt else ckpt - -def eval(model, n, input): - qk = f"model.diffusion_model.output_blocks.{n}.1.transformer_blocks.0.attn1.to_q.weight" - uk = f"model.diffusion_model.output_blocks.{n}.1.transformer_blocks.0.attn1.to_k.weight" - vk = f"model.diffusion_model.output_blocks.{n}.1.transformer_blocks.0.attn1.to_v.weight" - atoq, atok, atov = model[qk], model[uk], model[vk] - attn = cal_cross_attn(atoq, atok, atov, input) - return attn - -def main(): - file1 = Path(sys.argv[1]) - files = sys.argv[2:] - seed = 114514 - torch.manual_seed(seed) - model_a = load_model(file1) - log.info(f"base: {file1.name}") - - map_attn_a = {} - map_rand_input = {} - for n in range(3, 11): - hidden_dim, embed_dim = model_a[f"model.diffusion_model.output_blocks.{n}.1.transformer_blocks.0.attn1.to_q.weight"].shape - rand_input = torch.randn([embed_dim, hidden_dim]) - map_attn_a[n] = eval(model_a, n, rand_input) - map_rand_input[n] = rand_input - del model_a - - for file2 in files: - file2 = Path(file2) - model_b = load_model(file2) - sims = [] - for n in range(3, 11): - attn_a = map_attn_a[n] - attn_b = eval(model_b, n, map_rand_input[n]) - sim = torch.mean(torch.cosine_similarity(attn_a, attn_b)) - sims.append(sim) - log.info(f"{file2}: {torch.mean(torch.stack(sims)) * 1e2:.2f}%") - -if __name__ == "__main__": - main() diff --git a/cli/modules/preview-embeddings.py b/cli/modules/preview-embeddings.py deleted file mode 100755 index 8076d1a61..000000000 --- a/cli/modules/preview-embeddings.py +++ /dev/null @@ -1,101 +0,0 @@ -#!/bin/env python -""" -create preview images from embeddings -""" -import os -import io -import sys -import json -import base64 -import argparse -from pathlib import Path -from PIL import Image -from inspect import getsourcefile -from util import Map, log -from sdapi import getsync, postsync -from grid import grid - -template = 'photo of "{name}", {suffix}, high detailed, skin texture, looking forward, facing camera, 135mm, shot on dslr, 4k, modelshoot style' -img2img_options = Map({ - 'prompt': None, - 'negative_prompt': 'cartoon, drawing, cgi, sketch, comic, disfigured, deformed', - 'init_images': [], - 'sampler_name': 'DPM2 Karras', - 'batch_size': 4, - 'n_iter': 1, - 'steps': 30, - 'cfg_scale': 6, - 'width': 512, - 'height': 512, - 'restore_faces': False -}) - -def encode(f): - img = Image.open(f) - with io.BytesIO() as stream: - img.save(stream, 'JPEG') - values = stream.getvalue() - encoded = base64.b64encode(values).decode() - return encoded - -def create_preview(name: str, suffix: str): - options = getsync('/sdapi/v1/options') - cmdflags = getsync('/sdapi/v1/cmd-flags') - img2img_options['prompt'] = template.format(name = name, suffix = suffix) - log.debug({ 'preview options': img2img_options }) - if len(img2img_options['init_images']) == 0: - for i in range(img2img_options.batch_size): - mask = os.path.join(os.path.dirname(getsourcefile(lambda:0)), 'preview-template'+ str(i+1) +'.jpg') - if (not os.path.isfile(mask)): - log.error({ 'preview': 'missing preview mask' }) - return - img2img_options['init_images'].append(encode(mask)) - data = postsync('/sdapi/v1/img2img', img2img_options) - if 'error' in data: - log.error({ 'preview': data['error'], 'reason': data['reason'] }) - return - info = Map(json.loads(data['info'])) - if not 'images' in data: - log.error({ 'preview': 'no images' }) - return - fn = os.path.join(cmdflags.embeddings_dir, name + '.preview.png') - log.info({ 'preview': { 'name': fn, 'model': options.sd_model_checkpoint, 'seed': info.seed } }) - images = [] - for b64 in data['images']: - images.append(Image.open(io.BytesIO(base64.b64decode(b64.split(",",1)[0])))) - image = grid(images, None, square=True) - image.save(fn) - -if __name__ == "__main__": - log.info({ 'preview': 'start' }) - cmdflags = getsync('/sdapi/v1/cmd-flags') - - parser = argparse.ArgumentParser(description = 'generate embeddings previews') - parser.add_argument('--overwrite', default = False, action='store_true', help = 'overwrite existing previews') - parser.add_argument('input', type=str, nargs='*') - params = parser.parse_args() - - if len(params.input) == 0: - files = list(Path(cmdflags.embeddings_dir).glob('*.pt')) - else: - files = list(os.path.join(cmdflags.embeddings_dir, a + '.pt') for a in params.input if os.path.isfile(os.path.join(cmdflags.embeddings_dir, a + '.pt'))) - candidates = [str(f) for f in files] - candidates.sort(key=os.path.getctime, reverse=True) - - files = [] - for f in candidates: - fn = f.replace('.pt', '.preview.png') - if os.path.isfile(f.replace('.pt', '.preview.png')): - if params.overwrite: - log.info({ 'preview add': fn }) - files.append(f) - else: - log.info({ 'preview skip': fn }) - else: - log.info({ 'preview add': fn }) - files.append(f) - - log.info({ 'preview embeddings': len(files) }) - for f in files: - name = Path(f).stem - create_preview(name, 'person') diff --git a/cli/modules/preview-template1.jpg b/cli/modules/preview-template1.jpg deleted file mode 100644 index d6da3f5e9..000000000 Binary files a/cli/modules/preview-template1.jpg and /dev/null differ diff --git a/cli/modules/preview-template2.jpg b/cli/modules/preview-template2.jpg deleted file mode 100644 index adcdd3487..000000000 Binary files a/cli/modules/preview-template2.jpg and /dev/null differ diff --git a/cli/modules/preview-template3.jpg b/cli/modules/preview-template3.jpg deleted file mode 100644 index 4982d127a..000000000 Binary files a/cli/modules/preview-template3.jpg and /dev/null differ diff --git a/cli/modules/preview-template4.jpg b/cli/modules/preview-template4.jpg deleted file mode 100644 index 140dc9aaa..000000000 Binary files a/cli/modules/preview-template4.jpg and /dev/null differ diff --git a/cli/modules/process.py b/cli/modules/process.py deleted file mode 100755 index 67453ae4c..000000000 --- a/cli/modules/process.py +++ /dev/null @@ -1,500 +0,0 @@ -#!/bin/env python -""" -process people images -- check image resolution -- runs detection of face and body -- extracts crop and performs checks: - - visible: is face or body detected - - in frame: for face based on box, for body based on number of visible keypoints - - resolution: is cropped image still of sufficient resolution - - optionaly upsample and restore face quality - - blur: is image sharp enough - - dynamic range: is image bright enough - - similarity: compares image to all previously processed images to see if its unique enough -- images are resized and optionally squared -- face additionally runs through semantic segmentation to remove background -- if image passes checks - image padded and saved as extracted image -- body requires that face is detected and in-frame, - but does not have to pass all other checks as body performs its own checks -- runs clip interrogation on extracted images to generate filewords -""" - -import os -import sys -import io -import math -import base64 -import pathlib -import argparse -import logging -import filetype -import numpy as np -import mediapipe as mp -from PIL import Image, ImageOps -from skimage.metrics import structural_similarity as ssim -from scipy.stats import beta -sys.path.append(os.path.join(os.path.dirname(__file__))) - -from util import log, Map -from sdapi import postsync - - -params = Map({ - # general settings, do not modify - 'src': '', # source folder - 'dst': '', # destination folder - 'clear_dst': True, # remove all files from destination at the start - 'format': '.jpg', # image format - 'target_size': 512, # target resolution - 'square_images': True, # should output images be squared - 'segmentation_model': 0, # segmentation model 0/general 1/landscape - 'segmentation_background': (192, 192, 192), # segmentation background color - 'blur_samplesize': 60, # sample size to use for blur detection - 'similarity_size': 64, # base similarity detection on reduced images - # original image processing settings - 'keep_original': False, # keep original image - # face processing settings - 'extract_face': False, # extract face from image - 'face_score': 0.7, # min face detection score - 'face_pad': 0.1, # pad face image percentage - 'face_model': 1, # which face model to use 0/close-up 1/standard - 'face_blur': False, # check for body blur - 'face_blur_score': 1.5, # max score for face blur detection - 'face_range': False, # check for body blur - 'face_range_score': 0.15, # min score for face dynamic range detection - 'face_restore': False, # attempt to restore face quality - 'face_upscale': False, # attempt to scale small faces - 'face_segmentation': False, # segmentation enabled - # body processing settings - 'extract_body': False, # extract body from image - 'body_score': 0.9, # min body detection score - 'body_visibility': 0.5, # min visibility score for each detected body part - 'body_parts': 15, # min number of detected body parts with sufficient visibility - 'body_pad': 0.2, # pad body image percentage - 'body_model': 2, # body model to use 0/low 1/medium 2/high - 'body_blur': False, # check for body blur - 'body_blur_score': 1.8, # max score for body blur detection - 'body_range': False, # check for body blur - 'body_range_score': 0.15, # min score for body dynamic range detection - 'body_segmentation': False, # segmentation enabled - # similarity detection settings - 'similarity_score': 0.8, # maximum similarity score before image is discarded - # interrogate settings - 'interrogate_model': ['clip', 'deepdanbooru'], # interrogate models - 'interrogate_captions': True, # write captions to file - 'tag_limit': 5, # number of tags to extract -}) -face_model = None -body_model = None -segmentation_model = None - - -def detect_blur(image): - # based on - bw = ImageOps.grayscale(image) - cx, cy = image.size[0] // 2, image.size[1] // 2 - fft = np.fft.fft2(bw) - fftShift = np.fft.fftshift(fft) - fftShift[cy - params.blur_samplesize: cy + params.blur_samplesize, cx - params.blur_samplesize: cx + params.blur_samplesize] = 0 - fftShift = np.fft.ifftshift(fftShift) - recon = np.fft.ifft2(fftShift) - magnitude = np.log(np.abs(recon)) - mean = round(np.mean(magnitude), 2) - return mean - - -def detect_dynamicrange(image): - # based on - data = np.asarray(image) - image = np.float32(data) - RGB = [0.299, 0.587, 0.114] - height, width = image.shape[:2] - brightness_image = np.sqrt(image[..., 0] ** 2 * RGB[0] + image[..., 1] ** 2 * RGB[1] + image[..., 2] ** 2 * RGB[2]) - hist, _ = np.histogram(brightness_image, bins=256, range=(0, 255)) - img_brightness_pmf = hist / (height * width) - dist = beta(2, 2) - ys = dist.pdf(np.linspace(0, 1, 256)) - ref_pmf = ys / np.sum(ys) - dot_product = np.dot(ref_pmf, img_brightness_pmf) - squared_dist_a = np.sum(ref_pmf ** 2) - squared_dist_b = np.sum(img_brightness_pmf ** 2) - res = dot_product / math.sqrt(squared_dist_a * squared_dist_b) - return round(res, 2) - - -images = [] -def detect_simmilar(image): - img = image.resize((params.similarity_size, params.similarity_size)) - img = ImageOps.grayscale(img) - data = np.array(img) - similarity = 0 - for i in images: - val = ssim(data, i, data_range=255, channel_axis=None, gradient=False, full=False) - if val > similarity: - similarity = val - images.append(data) - return similarity - - -def segmentation(image): - global segmentation_model - if segmentation_model is None: - segmentation_model = mp.solutions.selfie_segmentation.SelfieSegmentation(model_selection=params.segmentation_model) - data = np.array(image) - results = segmentation_model.process(data) - condition = np.stack((results.segmentation_mask,) * 3, axis=-1) > 0.1 - background = np.zeros(data.shape, dtype=np.uint8) - background[:] = params.segmentation_background - data = np.where(condition, data, background) # consider using a joint bilateral filter instead of pure combine - segmented = Image.fromarray(data) - return segmented - - -def extract_face(img): - if not params.extract_face: - return None, True - if img.mode == 'RGBA': - img = img.convert('RGB') - scale = max(img.size[0], img.size[1]) / params.target_size - resized = img.copy() - resized.thumbnail((params.target_size, params.target_size), Image.HAMMING) - - global face_model - if face_model is None: - face_model = mp.solutions.face_detection.FaceDetection(min_detection_confidence=params.face_score, model_selection=params.face_model) - results = face_model.process(np.array(resized)) - if results.detections is None: - return None, False - box = results.detections[0].location_data.relative_bounding_box - if box.xmin < 0 or box.ymin < 0 or (box.width - box.xmin) > 1 or (box.height - box.ymin) > 1: - log.info({ 'process face skip': 'out of frame' }) - return None, False - x = (box.xmin - params.face_pad / 2) * resized.width - y = (box.ymin - params.face_pad / 2)* resized.height - w = (box.width + params.face_pad) * resized.width - h = (box.height + params.face_pad) * resized.height - cx = x + w / 2 - cy = y + h / 2 - l = max(w, h) / 2 - square = [scale * (cx - l), scale * (cy - l), scale * (cx + l), scale * (cy + l)] - square = [max(square[0], 0), max(square[1], 0), min(square[2], img.width), min(square[3], img.height)] - cropped = img.crop(tuple(square)) - - upscale = 1 - if params.face_restore or params.face_upscale: - if (cropped.size[0] < params.target_size or cropped.size[1] < params.target_size) and params.face_upscale: - upscale = 2 - kwargs = Map({ - 'image': encode(cropped), - 'upscaler_1': 'SwinIR_4x' if params.face_upscale else None, - 'codeformer_visibility': 1.0 if params.face_restore else 0.0, - 'codeformer_weight': 0.15 if params.face_restore else 0.0, - 'upscaling_resize': upscale, - }) - original = [cropped.size[0], cropped.size[1]] - res = postsync('/sdapi/v1/extra-single-image', kwargs) - if 'image' not in res: - log.error({ 'process face': 'upscale failed' }) - raise ValueError('upscale failed') - cropped = Image.open(io.BytesIO(base64.b64decode(res['image']))) - kwargs.image = [cropped.size[0], cropped.size[1]] - upscaled = [cropped.size[0], cropped.size[1]] - upscale = False if upscale == 1 else { 'original': original, 'upscaled': upscaled } - log.info({ 'process face restore': params.face_restore, 'upscale': upscale }) - - if cropped.size[0] < params.target_size and cropped.size[1] < params.target_size: - log.info({ 'process face skip': 'low resolution', 'size': [cropped.size[0], cropped.size[1]] }) - return None, True - cropped.thumbnail((params.target_size, params.target_size), Image.HAMMING) - - if params.square_images: - squared = Image.new('RGB', (params.target_size, params.target_size)) - squared.paste(cropped, ((params.target_size - cropped.width) // 2, (params.target_size - cropped.height) // 2)) - if params.face_segmentation: - squared = segmentation(squared) - else: - squared = cropped - - if params.face_blur: - blur = detect_blur(squared) - if blur > params.face_blur_score: - log.info({ 'process face skip': 'blur check fail', 'blur': blur }) - return None, True - else: - log.debug({ 'process face blur': blur }) - - if params.face_range: - range = detect_dynamicrange(squared) - if range < params.face_range_score: - log.info({ 'process face skip': 'dynamic range check fail', 'range': range }) - return None, True - else: - log.debug({ 'process face dynamic range': range }) - - similarity = detect_simmilar(squared) - if similarity > params.similarity_score: - log.info({ 'process face skip': 'similarity check fail', 'score': round(similarity, 2) }) - return None, True - - return squared, True - - -def extract_body(img): - if not params.extract_body: - return None, True - if img.mode == 'RGBA': - img = img.convert('RGB') - scale = max(img.size[0], img.size[1]) / params.target_size - resized = img.copy() - resized.thumbnail((params.target_size, params.target_size), Image.HAMMING) - - global body_model - if body_model is None: - body_model = mp.solutions.pose.Pose(static_image_mode=True, min_detection_confidence=params.body_score, model_complexity=params.body_model) - results = body_model.process(np.array(resized)) - if results.pose_landmarks is None: - return None, False - x = [resized.width * (i.x - params.body_pad / 2) for i in results.pose_landmarks.landmark if i.visibility > params.body_visibility] - y = [resized.height * (i.y - params.body_pad / 2) for i in results.pose_landmarks.landmark if i.visibility > params.body_visibility] - if len(x) < params.body_parts: - log.info({ 'process body skip': 'insufficient body parts', 'detected': len(x) }) - return None, True - w = max(x) - min(x) + resized.width * params.body_pad - h = max(y) - min(y) + resized.height * params.body_pad - cx = min(x) + w / 2 - cy = min(y) + h / 2 - l = max(w, h) / 2 - square = [scale * (cx - l), scale * (cy - l), scale * (cx + l), scale * (cy + l)] - square = [max(square[0], 0), max(square[1], 0), min(square[2], img.width), min(square[3], img.height)] - cropped = img.crop(tuple(square)) - if cropped.size[0] < params.target_size and cropped.size[1] < params.target_size: - log.info({ 'process body skip': 'low resolution', 'size': [cropped.size[0], cropped.size[1]] }) - return None, True - cropped.thumbnail((params.target_size, params.target_size), Image.HAMMING) - - if params.square_images: - squared = Image.new('RGB', (params.target_size, params.target_size)) - squared.paste(cropped, ((params.target_size - cropped.width) // 2, (params.target_size - cropped.height) // 2)) - if params.body_segmentation: - squared = segmentation(squared) - else: - squared = cropped - - if params.body_blur: - blur = detect_blur(squared) - if blur > params.body_blur_score: - log.info({ 'process body skip': 'blur check fail', 'blur': blur }) - return None, True - else: - log.debug({ 'process body blur': blur }) - - if params.body_range: - range = detect_dynamicrange(squared) - if range < params.body_range_score: - log.info({ 'process body skip': 'dynamic range check fail', 'range': range }) - return None, True - else: - log.debug({ 'process body dynamic range': range }) - - similarity = detect_simmilar(squared) - if similarity > params.similarity_score: - log.info({ 'process body skip': 'similarity check fail', 'score': round(similarity, 2) }) - return None, True - - return squared, True - - -def save_original(img): - if img.mode == 'RGBA': - img = img.convert('RGB') - resized = img.copy() - resized.thumbnail((params.target_size, params.target_size), Image.HAMMING) - if params.square_images: - squared = Image.new('RGB', (params.target_size, params.target_size)) - squared.paste(resized, ((params.target_size - resized.width) // 2, (params.target_size - resized.height) // 2)) - else: - squared = resized - return squared - - -def encode(img): - with io.BytesIO() as stream: - img.save(stream, 'JPEG') - values = stream.getvalue() - encoded = base64.b64encode(values).decode() - return encoded - - -def interrogate(img, fn, intag = None): - if len(params.interrogate_model) == 0: - return - caption = '' - tags = [] - for model in params.interrogate_model: - json = Map({ 'image': encode(img), 'model': model }) - res = postsync('/sdapi/v1/interrogate', json) - if model == 'clip': - caption = res.caption if 'caption' in res else '' - caption = caption.split(',')[0].replace('a ', '') - if intag is not None: - caption = intag + ', ' + caption - if model == 'deepdanbooru': - tag = res.caption if 'caption' in res else '' - tags = tag.split(',') - tags = [t.replace('(', '').replace(')', '').replace('\\', '').split(':')[0].strip() for t in tags] - if intag is not None: - for t in intag.split(',')[::-1]: - tags.insert(0, t.strip()) - if params.interrogate_captions: - file = fn.replace(params.format, '.txt') - f = open(file, 'w') - f.write(caption) - f.close() - pos = 0 if len(tags) == 0 else 1 - tags.insert(pos, caption.split(' ')[1]) - if len(tags) > params.tag_limit: - tags = tags[:params.tag_limit] - log.info({ 'interrogate': caption, 'tags': tags }) - return caption, tags - - -i = {} -metadata = Map({}) - -# entry point when used as module -def process_file(f: str, dst: str = None, preview: bool = False, offline: bool = False, txt = None, tag = None, opts = []): - def save(img, f, what): - i[what] = i.get(what, 0) + 1 - if dst is None: - dir = os.path.dirname(f) - else: - dir = dst - base = os.path.basename(f).split('.')[0] - parent = os.path.basename(pathlib.Path(dir)) - basename = str(i[what]).rjust(3, '0') + '-' + what + '-' + base - fn = basename + params.format - # log.debug({ 'save': fn }) - caption = '' - tags = '' - if not preview: - img.save(os.path.join(dir, fn)) - if not offline: - caption, tags = interrogate(img, os.path.join(dir, fn), tag) - metadata[os.path.join(parent, basename)] = { 'caption': caption, 'tags': ','.join(tags) } - return fn - - # overrides - if len(opts) > 0: - params.keep_original = True if 'original' in opts else False - params.extract_face = True if 'face' in opts else False - params.extract_body = True if 'body' in opts else False - params.face_blur = True if 'blur' in opts else False - params.body_blur = True if 'blur' in opts else False - params.face_range = True if 'range' in opts else False - params.body_range = True if 'range' in opts else False - params.face_upscale = True if 'upscale' in opts else False - params.face_restore = True if 'restore' in opts else False - - log.info({ 'processing': f }) - try: - image = Image.open(f) - except Exception as err: - log.error({ 'image': f, 'error': err }) - return 0, {} - - image = ImageOps.exif_transpose(image) # rotate image according to EXIF orientation - if txt is not None: - params.interrogate_captions = txt - - if image.width < 512 or image.height < 512: - log.info({ 'process skip': 'low resolution', 'resolution': [image.width, image.height] }) - return 0, {} - log.debug({ 'resolution': [image.width, image.height], 'mp': round((image.width * image.height) / 1024 / 1024, 1) }) - - face, ok = extract_face(image) - if face is not None: - fn = save(face, f, 'face') - log.info({ 'extract face': fn }) - else: - log.debug({ 'no face': f }) - - if not ok: - return 0, {} - - body, ok = extract_body(image) - if body is not None: - fn = save(body, f, 'body') - log.info({ 'extract body': fn }) - else: - log.debug({ 'no body': f }) - - if params.keep_original: - resized = save_original(image) - fn = save(resized, f, 'original') - log.info({ 'original': fn }) - - image.close() - return i, metadata - -def process_images(src: str, dst: str, args = None): - params.src = src - params.dst = dst - if args is not None: - params.update(args) - log.info({ 'processing': params }) - if not os.path.isdir(src): - log.error({ 'process': 'not a folder', 'src': src }) - else: - if os.path.isdir(dst) and params.clear_dst: - log.info({ 'clear dst': dst }) - i = [os.path.join(dst, f) for f in os.listdir(dst) if os.path.isfile(os.path.join(dst, f)) and filetype.is_image(os.path.join(dst, f))] - for f in i: - os.remove(f) - pathlib.Path(dst).mkdir(parents=True, exist_ok=True) - for root, _sub_dirs, files in os.walk(src): - for f in files: - i, _metadata = process_file(os.path.join(root, f), dst) - return i - - -def unload_models(): - global face_model - if face_model is not None: - face_model = None - global body_model - if body_model is not None: - body_model = None - global segmentation_model - if segmentation_model is not None: - segmentation_model = None - - -if __name__ == '__main__': - # log.setLevel(logging.DEBUG) - parser = argparse.ArgumentParser(description = 'dataset processor') - parser.add_argument('--output', type=str, required=True, help='folder to store images') - parser.add_argument('--preview', default=False, action='store_true', help = "run processing but do not store results") - parser.add_argument('--offline', default=False, action='store_true', help = "run only processing steps that do not require running server") - parser.add_argument('--debug', default=False, action='store_true', help = "enable debug logging") - parser.add_argument('input', type=str, nargs='*') - args = parser.parse_args() - params.dst = args.output - if args.debug: - log.setLevel(logging.DEBUG) - log.debug({ 'debug': True }) - log.info({ 'processing': params }) - if not os.path.exists(params.dst) and not args.preview: - pathlib.Path(params.dst).mkdir(parents=True, exist_ok=True) - files = [] - for loc in args.input: - if os.path.isfile(loc): - files.append(loc) - elif os.path.isdir(loc): - for root, _sub_dirs, dir in os.walk(loc): - for f in dir: - files.append(os.path.join(root, f)) - for f in files: - process_file(f, params.dst, args.preview, args.offline) - log.info({ 'processed': i, 'inputs': len(files) }) - # print(json.dumps(metadata, indent=2)) diff --git a/cli/modules/train-losschart.py b/cli/modules/train-losschart.py deleted file mode 100755 index 2d9c7a3f1..000000000 --- a/cli/modules/train-losschart.py +++ /dev/null @@ -1,191 +0,0 @@ -#!/bin/env python - -import io -import os -import sys -import json -import pathlib -import logging -import torch -import numpy as np -from PIL import Image, ImageFont, ImageDraw -from matplotlib import pyplot as plt -from util import log, Map - - -def settings(logdir: str, name: str): - filename = os.path.join(logdir, name, 'settings.json') - with open(filename, 'r', encoding='utf-8') as f: - data = json.load(f) - data = Map(data) - log.debug({ 'settings': data }) - return data - - -def plot(logdir: str, name: str): - f = os.path.join(logdir, name, 'train.csv') - if not os.path.isfile(f): - log.debug({ 'train log missing': f }) - return - name = pathlib.Path(f).parent.name - img = os.path.join(logdir, f"{name}.train.png") - - step, loss, rate = plt.np.loadtxt(f, delimiter = ',', skiprows = 1, usecols = [0, 3, 4], unpack = True) - d = settings(logdir, name) - # window = d.get('gradient_step', 1) * d.get('batch_size', 1) - window = d.get('save_embedding_every', 1) - try: - log.debug({ 'loss plot': name, 'output': img, 'data': f, 'records': len(step) }) - except: - return # no data - if len(step) < 5: - return - - plt.rcParams.update({'font.variant':'small-caps'}) - plt.rc('axes', edgecolor='gray') - plt.rc('font', size=10) - plt.rc('font', variant='small-caps') - plt.grid(color='gray', linewidth=1, axis='both', alpha=0.5) - plt.rcParams['figure.figsize'] = [14, 14] - plt.rcParams['figure.facecolor'] = 'black' - figure, axis = plt.subplots(2, 1) - - # create top graph - ax0 = axis[0] - ax0.set_facecolor(color = (0.1, 0.1, 0.1, 0.5)) - ax0.tick_params(axis='x', labelcolor='white') - ax0.set_xlabel('step'.upper(), color='white') - ax0.set_xlim(0, d.steps) - ax0.set_ylim(0, 0.5) - ax0.set_axisbelow(True) - ax0.xaxis.grid(color='gray', linestyle='dashed') - ax0.yaxis.grid(color='gray', linestyle='dashed') - - # loss values - ax0.plot(step, loss, color='gray', label='loss value') - ax0.set_ylabel('loss value'.upper(), color='#CE6400') - ax0.tick_params(axis='y', labelcolor='#CE6400') - - # trendline - z = np.polyfit(step, loss, 1) - p = np.poly1d(z) - ax0.plot(step, p(loss), color='#5020F0', linewidth=3, label='loss trendline') - - # moving average - if len(loss) > window: - maval = [] - minval = [] - maxval = [] - for ind in range(window - 1): - maval.insert(0, np.nan) - minval.insert(0, np.nan) - maxval.insert(0, np.nan) - for ind in range(len(loss) - window + 1): - maval.append(np.mean(loss[ind:ind+window])) - minval.append(np.min(loss[ind:ind+window])) - maxval.append(np.max(loss[ind:ind+window])) - ax0.plot(step, maval, color='#CE6400', linewidth=5, label='average loss value') - ax0.plot(step, minval, color='#500010', linewidth=5, label='min loss per epoch') - ax0.plot(step, maxval, color='#005010', linewidth=5, label='max loss per epoch') - - # learning rate - ax0_right = ax0.twinx() - ax0_right.set_ylabel('learn rate'.upper(), color='cyan') - ax0_right.plot(step, rate, color='cyan', linewidth=3, linestyle='dashed', label='learn rate') - ax0_right.tick_params(axis='y', labelcolor='cyan') - - # axis legend - handles0, labels0 = ax0.get_legend_handles_labels() # because ax2 is twin, both are included - ax0.legend(handles0, labels0, loc="best") - - # embeddings - ax1 = axis[1] - ax1.set_facecolor(color = (0.1, 0.1, 0.1, 0.5)) - ax1.set_ylabel('vector average'.upper(), color=(1, 0.2, 0.5)) - ax1.tick_params(axis='y', labelcolor=(1, 0.2, 0.5)) - ax1.set_xlim(0, d.steps) - ax1_right = ax1.twinx() - ax1_right.set_ylabel('vector norm'.upper(), color=(0.2, 1.0, 0.5)) - ax1_right.tick_params(axis='y', labelcolor=(0.2, 1.0, 0.5)) - ax1.xaxis.grid(color='gray', linestyle='dashed') - ax1.yaxis.grid(color='gray', linestyle='dashed') - x = [] - avg = [] - norm = [] - avg_v = [[] for y in range(d.num_vectors_per_token)] - norm_v = [[] for y in range(d.num_vectors_per_token)] - embedding_files = sorted(pathlib.Path(os.path.join(logdir, name, 'embeddings')).glob('*.pt'), key=os.path.getmtime) - for f in embedding_files: - embed = torch.load(f, map_location=torch.device("cpu")) # pylint: disable=no-member - x.append(embed["step"] + 1) - token = list(embed["string_to_token"].keys())[0] - tensors = embed["string_to_param"][token] - val = tensors.detach().numpy() - data = val.flatten() - avg.append(np.average(np.abs(data))) - norm.append(np.linalg.norm(data)) - for i in range(val.shape[0]): - avg_v[i].append(np.average(np.abs(val[i]))) - norm_v[i].append(np.linalg.norm(val[i])) - ax1.plot(x, avg, color=(1, 0.2, 0.5), linewidth=3, label='all vectors average value') - ax1_right.plot(x, norm, color=(0.2, 1, 0.5), linewidth=3, label='all vectors norm value', linestyle='dashed') - for i in range(d.num_vectors_per_token): - ax1.plot(x, avg_v[i], color= (i / (d.num_vectors_per_token + 1), 0.2, 0.5), linewidth=1, label=f'vector={i} average value') - ax1_right.plot(x, norm_v[i], color= (0.2, i / (d.num_vectors_per_token + 1), 0.5), linewidth=1, label=f'vector={i} norm value', linestyle='dashed') - - # axis legend - handles1, labels1 = ax1.get_legend_handles_labels() # because ax2 is twin, both are included - ax1.legend(handles1, labels1, loc="upper left") - ax1_right.legend(loc="upper right") - - # create chart and convert to pil - figure.tight_layout() - buf = io.BytesIO() - plt.savefig(buf, format='png') - pltimg = Image.open(buf) - size = (pltimg.size[0], pltimg.size[1] + 240) - image = Image.new('RGB', size = size, color = (206, 100, 0)) - font = ImageFont.truetype('DejaVuSansMono', 18) - image.paste(pltimg, box=(0, 240)) - buf.close() - plt.close() - - # text - textl = f"""NAME: {d.embedding_name.upper()} -IMAGES: {d.num_of_dataset_images} -VECTORS: {d.num_vectors_per_token} -STEPS: {d.steps} -BATCH-SIZE: {d.batch_size} -GRADIENT-STEP: {d.gradient_step} -SAMPLING-METHOD: {d.latent_sampling_method} -MODEL: {d.model_name.upper()} -LEARN-RATE: {d.learn_rate.replace(' ', '')} -""" - - minval = f"{round(np.min(loss), 4)} @ {round(step[np.argmin(loss)])}" - maxval = f"{round(np.max(loss), 4)} @ {round(step[np.argmax(loss)])}" - textr = f"""{d.datetime} -LOSS: {round(loss[-1], 4)} -MIN: {minval} -MAX: {maxval} -TREND: {z[0]:.5f} -""" - if len(avg) > 0: - textr += f"VECTOR AVG: {avg[-1]:.3f}\n" - if len(norm) > 0: - textr += f"VECTOR NORM: {norm[-1]:.3f}\n" - ctx = ImageDraw.Draw(image) - ctx.text((8, 8), textl, font = font, fill = (255, 255, 255), spacing = 8) - ctx.text((image.size[0] - 220, 8), textr, font = font, fill = (255, 255, 255)) - - image.save(img) - - -if __name__ == "__main__": - log.setLevel(logging.DEBUG) - if len(sys.argv) == 2: - arg = sys.argv[1] - log.debug({ 'args': arg }) - plot(os.path.dirname(arg), os.path.basename(arg)) - else: - log.error({ 'loss chart': 'specify embedding name'}) diff --git a/cli/modules/train-lossrate.py b/cli/modules/train-lossrate.py deleted file mode 100755 index 4a6165ff5..000000000 --- a/cli/modules/train-lossrate.py +++ /dev/null @@ -1,132 +0,0 @@ -#!/bin/env python -""" -auto-generate learn-rate -""" -import io -import math -import logging -import numpy as np -from PIL import Image, ImageFont, ImageDraw -from matplotlib import pyplot as plt -from util import log, Map - - -loss_types = ['linear', 'log', 'linalg', 'power'] - - -def gen_steps(steps, step): - return [x for x in range(1, steps + step) if x % step == 0] - - -def gen_loss_rate(steps: int, step: int, loss_start: float, loss_end: float, loss_type: loss_types, power: int = 3): - def norm(val): - return ((loss_start - loss_end) * val) / val.max() + loss_end - - steps_val = gen_steps(steps, step) - - if loss_type == 'linear': - loss_val = np.interp(steps_val, [steps_val[0], steps_val[-1]], [loss_start, loss_end]) - - elif loss_type == 'log': - loss_val = np.logspace(loss_start, 0, num=len(steps_val), base=math.e) - loss_val = norm(loss_val - loss_val.min()) - - elif loss_type == 'linalg': - loss_val = np.array(steps_val[::-1], dtype='float') - loss_val = norm(loss_val / np.linalg.norm(loss_val)) - - elif loss_type == 'power': - loss_val = np.array([math.pow(x, power) for x in range(len(steps_val))][::-1]) - loss_val = norm(loss_val) - - else: - return [] - - return loss_val - - -def gen_loss_rate_str(steps: int, step: int, loss_start: float, loss_end: float, loss_type: loss_types, power: int = 3): - steps_val = gen_steps(steps, step) - loss_val = gen_loss_rate(steps, step, loss_start, loss_end, loss_type, power) - loss_rate = [f"{loss_val[i]:.4f}:{steps_val[i]}" for i in range(len(steps_val))] - loss_rate = ', '.join(loss_rate) - log.debug({ 'loss_rate': loss_rate, 'function': loss_type, 'power': power }) - return loss_rate - - -def example_plot(steps: int, step: int, loss_start: float, loss_end: float): - plt.rcParams.update({'font.variant':'small-caps'}) - plt.rc('axes', edgecolor='gray') - plt.rc('font', size=10) - plt.rc('font', variant='small-caps') - plt.grid(color='gray', linewidth=1, axis='both', alpha=0.5) - plt.rcParams['figure.figsize'] = [14, 6] - plt.figure(facecolor='black') - - loss_rates = [] - - ax1 = plt.subplot(1, 2, 1) - ax1.set_facecolor('grey') - ax1.set_xlabel('step'.upper(), color='white') - ax1.set_ylabel('loss value'.upper(), color='white') - ax1.xaxis.grid(color='gray', linestyle='dashed') - ax1.tick_params(axis='x', labelcolor='white') - ax1.tick_params(axis='y', labelcolor='white') - ax1.legend(loc="best") - for loss_type in [x for x in loss_types if x != 'power']: - col = (np.random.random(), np.random.random(), np.random.random()) - x = gen_steps(steps, step) - y = gen_loss_rate(loss_type = loss_type, steps = steps, step = step, loss_start = loss_start, loss_end = loss_end) - ax1.plot(x, y, label=loss_type, color = col) - loss = gen_loss_rate_str(loss_type = loss_type, steps = steps, step = step, loss_start = loss_start, loss_end = loss_end) - loss_rates.append(f"LOSS {loss} TYPE {loss_type}") - handles, labels = ax1.get_legend_handles_labels() - ax1.legend(handles, labels) - - ax2 = plt.subplot(1, 2, 2) - ax2.set_facecolor('grey') - ax2.set_xlabel('step'.upper(), color='white') - ax2.set_ylabel('loss value'.upper(), color='white') - ax2.xaxis.grid(color='gray', linestyle='dashed') - ax2.tick_params(axis='x', labelcolor='white') - ax2.tick_params(axis='y', labelcolor='white') - ax2.legend(loc="best") - for power in range(1, 11): - col = (np.random.random(), np.random.random(), np.random.random()) - x = gen_steps(steps, step) - y = gen_loss_rate(loss_type = 'power', power = power, steps = steps, step = step, loss_start = loss_start, loss_end = loss_end) - ax2.plot(x, y, label=f"power={pow}", color = col) - loss = gen_loss_rate_str(loss_type = 'power', power = power, steps = steps, step = step, loss_start = loss_start, loss_end = loss_end) - loss_rates.append(f"LOSS {loss} TYPE power={power}") - handles, labels = ax2.get_legend_handles_labels() - ax2.legend(handles, labels) - - plt.tight_layout() - buf = io.BytesIO() - plt.savefig(buf, format='png') - pltimg = Image.open(buf) - size = (pltimg.size[0], pltimg.size[1] + 300) - image = Image.new('RGB', size = size, color = (206, 100, 0)) - font = ImageFont.truetype('DejaVuSansMono', 14) - image.paste(pltimg, box=(0, 300)) - buf.close() - - # text - rates = "\n".join(loss_rates) - text = f"STEPS {steps} STEP {step} LOSS-START {loss_start} LOSS-END {loss_end}\n" + rates - - ctx = ImageDraw.Draw(image) - ctx.text((8, 8), text, font = font, fill = (255, 255, 255), spacing = 8) - image.save('lossrate.jpg') - - -if __name__ == "__main__": - log.setLevel(logging.DEBUG) - arg = Map({ - "steps": 500, - "step": 50, - "loss_start": 0.01, - "loss_end": 0.001 - }) - log.debug({ 'options': arg }) - example_plot(**arg) diff --git a/cli/nvidia-smi.py b/cli/nvidia-smi.py new file mode 100755 index 000000000..f7b11311a --- /dev/null +++ b/cli/nvidia-smi.py @@ -0,0 +1,32 @@ +#!/usr/bin/env python +import os +import json +import shutil +import subprocess +import xmltodict +from rich import print # pylint: disable=redefined-builtin +from util import log, Map + +def get_nvidia_smi(output='dict'): + smi = shutil.which('nvidia-smi') + if smi is None: + log.error("nvidia-smi not found") + return None + result = subprocess.run(f'"{smi}" -q -x', shell=True, check=False, env=os.environ, stdout=subprocess.PIPE, stderr=subprocess.PIPE) + xml = result.stdout.decode(encoding="utf8", errors="ignore") + d = xmltodict.parse(xml) + if 'nvidia_smi_log' in d: + d = d['nvidia_smi_log'] + if 'gpu' in d and 'supported_clocks' in d['gpu']: + del d['gpu']['supported_clocks'] + if output == 'dict': + return d + elif output == 'class' or output == 'map': + d = Map(d) + return d + elif output == 'json': + return json.dumps(d, indent=4) + +if __name__ == "__main__": + res = get_nvidia_smi(output='dict') + print(type(res), res) diff --git a/cli/train/options.py b/cli/options.py similarity index 100% rename from cli/train/options.py rename to cli/options.py diff --git a/cli/train/process.py b/cli/process.py similarity index 96% rename from cli/train/process.py rename to cli/process.py index 8aedcd872..a2104bf0c 100644 --- a/cli/train/process.py +++ b/cli/process.py @@ -1,15 +1,13 @@ + # pylint: disable=global-statement import os -import sys import io import math import base64 -import pathlib import numpy as np import mediapipe as mp from PIL import Image, ImageOps from skimage.metrics import structural_similarity as ssim from scipy.stats import beta -sys.path.append(os.path.join(os.path.dirname(__file__))) import util import sdapi @@ -23,9 +21,9 @@ all_images_by_type = {} class Result(object): - def __init__(self, type: str, input: str, tag: str = None, requested: list = []): - self.type = type - self.input = input + def __init__(self, typ: str, fn: str, tag: str = None, requested: list = []): + self.type = typ + self.input = fn self.output = '' self.basename = '' self.message = '' @@ -56,8 +54,8 @@ def detect_dynamicrange(image: Image): data = np.asarray(image) image = np.float32(data) RGB = [0.299, 0.587, 0.114] - height, width = image.shape[:2] - brightness_image = np.sqrt(image[..., 0] ** 2 * RGB[0] + image[..., 1] ** 2 * RGB[1] + image[..., 2] ** 2 * RGB[2]) + height, width = image.shape[:2] # pylint: disable=unsubscriptable-object + brightness_image = np.sqrt(image[..., 0] ** 2 * RGB[0] + image[..., 1] ** 2 * RGB[1] + image[..., 2] ** 2 * RGB[2]) # pylint: disable=unsubscriptable-object hist, _ = np.histogram(brightness_image, bins=256, range=(0, 255)) img_brightness_pmf = hist / (height * width) dist = beta(2, 2) @@ -264,7 +262,7 @@ def save_image(res: Result, folder: str): def file(filename: str, folder: str, tag = None, requested = []): # initialize result dict - res = Result(input = filename, type='unknown', tag=tag, requested = requested) + res = Result(fn = filename, typ='unknown', tag=tag, requested = requested) # open image try: res.image = Image.open(filename) diff --git a/cli/modules/prompt-ideas.py b/cli/prompt-ideas.py similarity index 87% rename from cli/modules/prompt-ideas.py rename to cli/prompt-ideas.py index ff70123da..55b9a2d7b 100755 --- a/cli/modules/prompt-ideas.py +++ b/cli/prompt-ideas.py @@ -1,4 +1,4 @@ -#!/bin/env python +#!/usr/bin/env python """ generate prompt ideas model from: @@ -10,18 +10,13 @@ from transformers import GPT2Tokenizer, GPT2LMHeadModel from util import log -tokenizer = None -model = None - - def prompt(text: str, temp: float = 0.9, top: int = 8, penalty: float = 1.2, alpha: float = 0.6, num: int = 5, length: int = 80): - global tokenizer, model # pylint: disable=global-statement - if tokenizer is None: - tokenizer = GPT2Tokenizer.from_pretrained('distilgpt2') - tokenizer.add_special_tokens({'pad_token': '[PAD]'}) - if model is None: - model = GPT2LMHeadModel.from_pretrained('FredZhang7/distilgpt2-stable-diffusion-v2') + log.info({ 'loading': 'tokenizer' }) + tokenizer = GPT2Tokenizer.from_pretrained('distilgpt2') + tokenizer.add_special_tokens({'pad_token': '[PAD]'}) input_ids = tokenizer(text, return_tensors='pt').input_ids + log.info({ 'loading': 'model' }) + model = GPT2LMHeadModel.from_pretrained('FredZhang7/distilgpt2-stable-diffusion-v2') output = model.generate(input_ids, do_sample = True, temperature = temp, diff --git a/cli/modules/prompt-promptist.py b/cli/prompt-promptist.py similarity index 79% rename from cli/modules/prompt-promptist.py rename to cli/prompt-promptist.py index 60c5ee680..23e670e59 100755 --- a/cli/modules/prompt-promptist.py +++ b/cli/prompt-promptist.py @@ -1,28 +1,31 @@ -#!/bin/env python +#!/usr/bin/env python """ use microsoft promptist to beautify prompt - """ import sys -from transformers import AutoModelForCausalLM, AutoTokenizer from util import log - -def load_prompter(): +def load_model(): + log.info({ 'loading': 'model' }) + from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("microsoft/Promptist") # pylint: disable=redefined-outer-name + return model + +def load_tokenizer(): + log.info({ 'loading': 'tokenizer' }) + from transformers import AutoTokenizer tokenizer = AutoTokenizer.from_pretrained("gpt2") # pylint: disable=redefined-outer-name tokenizer.pad_token = tokenizer.eos_token tokenizer.padding_side = "left" - return model, tokenizer - - -model, tokenizer = load_prompter() - + return tokenizer def beautify(plain_text): + tokenizer = load_tokenizer() input_ids = tokenizer(plain_text.strip() + " Rephrase:", return_tensors = "pt").input_ids eos_id = tokenizer.eos_token_id + model = load_model() outputs = model.generate(input_ids, do_sample = False, max_new_tokens = 75, num_beams = 8, num_return_sequences = 8, eos_token_id = eos_id, pad_token_id = eos_id, length_penalty = -1.0) output_texts = tokenizer.batch_decode(outputs, skip_special_tokens = True) texts = [] diff --git a/cli/random/detectmodel.py b/cli/random/detectmodel.py deleted file mode 100755 index c104fbea5..000000000 --- a/cli/random/detectmodel.py +++ /dev/null @@ -1,43 +0,0 @@ -#!/bin/env python -""" -Detect model type - -Works for v1 and v2-base (EPS models), both standard inference and inpainting - -But looking at model dumps between EPS and V type models, its only about parametrization, there are no differences in actual model (its just weighted differently without any structural difference) -So i don't see easy way to auto-detect if model should be run in `EPS` or `V` mode -Only difference are some calculations in `ldm/models/diffusion/ddpm.py` and by then we already need to know which code-path to trigger -(maaaybe there could be a way by looking at some cherry-picked base tensors min/max range, but I dont see that as reliable) -""" - -import os -import sys - -import torch - -def signature(model): - if model is None: - return None - try: - size = model['state_dict']['model.diffusion_model.input_blocks.1.1.transformer_blocks.0.attn2.to_k.weight'].shape[1] - unet = model['state_dict']['model.diffusion_model.input_blocks.0.0.weight'].shape[1] - except: - return 'unknown' - guess = 'v1' if size == 768 else 'v2' # 768 for v1 and 1024 for v2 - guess += '-inference' if unet == 4 else '-inpainting' # inference models have shorter inputs, 4 for inference 9 for inpainting - return guess - -def load(file: str): - try: - model = torch.load(file, map_location='cpu') - return model - except Exception as err: - print(f"Error loading {f}: {err}") - -if __name__ == "__main__": - sys.argv.pop(0) - for f in sys.argv: - if os.path.isfile(f): - print(f"Model {f} is of type {signature(load(f))}") - else: - print(f"{f} is not a file") diff --git a/cli/random/versions.py b/cli/random/versions.py deleted file mode 100755 index 13e7b0f66..000000000 --- a/cli/random/versions.py +++ /dev/null @@ -1,33 +0,0 @@ -#!/bin/env python -""" -print module versions -""" - -import importlib -import pkg_resources - -modules = [ - 'diffusers', 'xformers', 'tokenizers', 'accelerate', 'safetensors' -] - -def get_torch(): - try: - torch = importlib.import_module('torch') - print('torch:', { 'version': torch.__version__ }) - print('cuda:', { 'available': torch.cuda.is_available(), 'version': torch.version.cuda, 'arch': torch.cuda.get_arch_list() }) - print('device:', { 'name': torch.cuda.get_device_name(torch.cuda.current_device()) }) - except Exception as err: - print('torch:', { 'error': err }) - - -def version(name: str): - try: - ver = pkg_resources.get_distribution(name).version - print(f"{name}: {ver}") - except Exception as err: - print(f"{name} error: {err}") - -if __name__ == "__main__": # create & train test embedding when used from cli - get_torch() - for module in modules: - version(module) diff --git a/cli/requirements.txt b/cli/requirements.txt index b8f807b41..50732dc8b 100644 --- a/cli/requirements.txt +++ b/cli/requirements.txt @@ -2,3 +2,4 @@ mediapipe colormap invisible-watermark filetype +albumentations diff --git a/cli/modules/bench.py b/cli/run-benchmark.py similarity index 97% rename from cli/modules/bench.py rename to cli/run-benchmark.py index 094b73f63..498f0cff3 100755 --- a/cli/modules/bench.py +++ b/cli/run-benchmark.py @@ -1,4 +1,4 @@ -#!/bin/env python +#!/usr/bin/env python """ sd api txt2img benchmark """ @@ -87,7 +87,6 @@ async def main(): 'vae': opts.sd_vae, 'hypernetwork': opts.sd_hypernetwork, 'sampler': options.sampler_name, - 'clip-stop': opts.CLIP_stop_at_last_layers, 'preview': opts.show_progress_every_n_steps } }) cpu, gpu = memstats() diff --git a/cli/modules/sdapi.py b/cli/sdapi.py similarity index 71% rename from cli/modules/sdapi.py rename to cli/sdapi.py index a8930c20b..62edb3038 100755 --- a/cli/modules/sdapi.py +++ b/cli/sdapi.py @@ -1,23 +1,59 @@ -#!/bin/env python +#!/usr/bin/env python +#pylint: disable=redefined-outer-name """ helper methods that creates HTTP session with managed connection pool provides async HTTP get/post methods and several helper methods """ +import os import sys -import json -import aiohttp +import ssl import asyncio import logging +import aiohttp import requests +import urllib3 from util import Map, log -sd_url = "http://127.0.0.1:7860" # automatic1111 api url root +sd_url = os.environ.get('SDAPI_URL', "http://127.0.0.1:7860") # automatic1111 api url root + use_session = True +urllib3.disable_warnings(urllib3.exceptions.InsecureRequestWarning) +ssl.create_default_context = ssl._create_unverified_context # pylint: disable=protected-access timeout = aiohttp.ClientTimeout(total = None, sock_connect = 10, sock_read = None) # default value is 5 minutes, we need longer for training sess = None quiet = False +BaseThreadPolicy = asyncio.WindowsSelectorEventLoopPolicy if sys.platform == "win32" and hasattr(asyncio, "WindowsSelectorEventLoopPolicy") else asyncio.DefaultEventLoopPolicy + + +class AnyThreadEventLoopPolicy(BaseThreadPolicy): + def get_event_loop(self) -> asyncio.AbstractEventLoop: + try: + return super().get_event_loop() + except (RuntimeError, AssertionError): + loop = self.new_event_loop() + self.set_event_loop(loop) + return loop + +asyncio.set_event_loop_policy(AnyThreadEventLoopPolicy()) + + +def authsync(): + sd_username = os.environ.get('SDAPI_USR', None) + sd_password = os.environ.get('SDAPI_PWD', None) + if sd_username is not None and sd_password is not None: + return requests.auth.HTTPBasicAuth(sd_username, sd_password) + return None + + +def auth(): + sd_username = os.environ.get('SDAPI_USR', None) + sd_password = os.environ.get('SDAPI_PWD', None) + if sd_username is not None and sd_password is not None: + return aiohttp.BasicAuth(sd_username, sd_password) + return None + async def result(req): @@ -60,7 +96,7 @@ async def get(endpoint: str, json: dict = None): global sess # pylint: disable=global-statement sess = sess if sess is not None else await session() try: - async with sess.get(url = endpoint, json = json) as req: + async with sess.get(url=endpoint, json=json, verify_ssl=False) as req: res = await result(req) return res except Exception as err: @@ -70,13 +106,12 @@ async def get(endpoint: str, json: dict = None): def getsync(endpoint: str, json: dict = None): try: - req = requests.get(f'{sd_url}{endpoint}', json = json) # pylint: disable=missing-timeout + req = requests.get(f'{sd_url}{endpoint}', json=json, verify=False, auth=authsync()) # pylint: disable=missing-timeout res = resultsync(req) return res except Exception as err: log.error({ 'session': err }) return {} - async def post(endpoint: str, json: dict = None): @@ -86,7 +121,7 @@ async def post(endpoint: str, json: dict = None): await sess.close() sess = await session() try: - async with sess.post(url = endpoint, json = json) as req: + async with sess.post(url=endpoint, json=json, verify_ssl=False) as req: res = await result(req) return res except Exception as err: @@ -95,7 +130,7 @@ async def post(endpoint: str, json: dict = None): def postsync(endpoint: str, json: dict = None): - req = requests.post(f'{sd_url}{endpoint}', json = json) # pylint: disable=missing-timeout + req = requests.post(f'{sd_url}{endpoint}', json=json, verify=False, auth=authsync()) # pylint: disable=missing-timeout res = resultsync(req) return res @@ -136,9 +171,9 @@ def progresssync(): def options(): - options = getsync('/sdapi/v1/options') + opts = getsync('/sdapi/v1/options') flags = getsync('/sdapi/v1/cmd-flags') - return { 'options': options, 'flags': flags } + return { 'options': opts, 'flags': flags } def shutdown(): @@ -151,7 +186,7 @@ def shutdown(): async def session(): global sess # pylint: disable=global-statement time = aiohttp.ClientTimeout(total = None, sock_connect = 10, sock_read = None) # default value is 5 minutes, we need longer for training - sess = aiohttp.ClientSession(timeout = time, base_url = sd_url) + sess = aiohttp.ClientSession(timeout = time, base_url = sd_url, auth=auth()) log.debug({ 'sdapi': 'session created', 'endpoint': sd_url }) """ sess = await aiohttp.ClientSession(timeout = timeout).__aenter__() @@ -171,6 +206,7 @@ async def session(): async def close(): if sess is not None: await asyncio.sleep(0) + await sess.close() await sess.__aexit__(None, None, None) log.debug({ 'sdapi': 'session closed', 'endpoint': sd_url }) @@ -181,12 +217,16 @@ if __name__ == "__main__": asyncio.run(interrupt()) if 'progress' in sys.argv: asyncio.run(progress()) + if 'progresssync' in sys.argv: + progresssync() if 'options' in sys.argv: opt = options() log.debug({ 'options' }) + import json print(json.dumps(opt['options'], indent = 2)) log.debug({ 'cmd-flags' }) print(json.dumps(opt['flags'], indent = 2)) if 'shutdown' in sys.argv: shutdown() - asyncio.run(close()) + asyncio.run(close(), debug=True) + asyncio.run(asyncio.sleep(0.5)) diff --git a/cli/simple-img2txt.py b/cli/simple-img2txt.py new file mode 100755 index 000000000..ceb89fd81 --- /dev/null +++ b/cli/simple-img2txt.py @@ -0,0 +1,62 @@ +#!/usr/bin/env python +import io +import sys +import base64 +import logging +import requests +from PIL import Image + +logging.basicConfig(level = logging.INFO, format = '%(asctime)s %(levelname)s: %(message)s') +log = logging.getLogger(__name__) +sd_url = "http://127.0.0.1:7860" +options = { + "init_images": [], + "prompt": "city at night", + "negative_prompt": "foggy, blurry", + "steps": 1, + "batch_size": 1, + "n_iter": 1, + "seed": -1, + "sampler_name": "Euler a", + "cfg_scale": 6, + "width": 512, + "height": 512, + "save_images": False, + "send_images": True, +} + +def post(endpoint: str, dct: dict = None): + req = requests.post(f'{sd_url}{endpoint}', json = dct, timeout=300) + if req.status_code != 200: + return { 'error': req.status_code, 'reason': req.reason, 'url': req.url } + else: + return req.json() + +def encode(f): + image = Image.open(f) + if image.mode == 'RGBA': + image = image.convert('RGB') + with io.BytesIO() as stream: + image.save(stream, 'JPEG') + values = stream.getvalue() + encoded = base64.b64encode(values).decode() + return encoded + +def generate(num: int = 0): + log.info(f'sending generate request: {num+1} {options}') + options['init_images'] = [encode('../html/logo.png')] + data = post('/sdapi/v1/img2img', options) + if 'images' in data: + for i in range(len(data['images'])): + b64 = data['images'][i].split(',',1)[0] + image = Image.open(io.BytesIO(base64.b64decode(b64))) + log.info(f'received image: {image.size}') + else: + log.warning(f'no images received: {data}') + +if __name__ == "__main__": + sys.argv.pop(0) + repeats = int(''.join(sys.argv) or '1') + log.info(f'repeats: {repeats}') + for n in range(repeats): + generate(n) diff --git a/cli/simple-txt2txt.py b/cli/simple-txt2txt.py new file mode 100755 index 000000000..34fadb4c5 --- /dev/null +++ b/cli/simple-txt2txt.py @@ -0,0 +1,50 @@ +#!/usr/bin/env python +import io +import sys +import base64 +import logging +import requests +from PIL import Image + +logging.basicConfig(level = logging.INFO, format = '%(asctime)s %(levelname)s: %(message)s') +log = logging.getLogger(__name__) +sd_url = "http://127.0.0.1:7860" +options = { + "prompt": "city at night", + "negative_prompt": "foggy, blurry", + "steps": 20, + "batch_size": 1, + "n_iter": 1, + "seed": -1, + "sampler_name": "UniPC", + "cfg_scale": 6, + "width": 512, + "height": 512, + "save_images": False, + "send_images": True, +} + +def post(endpoint: str, dct: dict = None): + req = requests.post(f'{sd_url}{endpoint}', json = dct, timeout=300) + if req.status_code != 200: + return { 'error': req.status_code, 'reason': req.reason, 'url': req.url } + else: + return req.json() + +def generate(num: int = 0): + log.info(f'sending generate request: {num+1} {options}') + data = post('/sdapi/v1/txt2img', options) + if 'images' in data: + for i in range(len(data['images'])): + b64 = data['images'][i].split(',',1)[0] + image = Image.open(io.BytesIO(base64.b64decode(b64))) + log.info(f'received image: {image.size}') + else: + log.warning(f'no images received: {data}') + +if __name__ == "__main__": + sys.argv.pop(0) + repeats = int(''.join(sys.argv) or '1') + log.info(f'repeats: {repeats}') + for n in range(repeats): + generate(n) diff --git a/cli/random/dynamotest.py b/cli/torch-compile.py similarity index 82% rename from cli/random/dynamotest.py rename to cli/torch-compile.py index 82b1143c6..f4f4d194b 100755 --- a/cli/random/dynamotest.py +++ b/cli/torch-compile.py @@ -1,4 +1,5 @@ -#!/bin/env python +#!/usr/bin/env python +# pylint: disable=cell-var-from-loop """ Test Torch Dynamo functionality and backends """ @@ -12,13 +13,14 @@ from torchvision.models import resnet18 print('torch:', torch.__version__) try: - import torch._dynamo as dynamo # must be imported explicitly or namespace is not found + # must be imported explicitly or namespace is not found + import torch._dynamo as dynamo # pylint: disable=ungrouped-imports except Exception as err: print('torch without dynamo support', err) N_ITERS = 20 -torch._dynamo.config.verbose=True +torch._dynamo.config.verbose=True # pylint: disable=protected-access warnings.filterwarnings('ignore', category=UserWarning) # disable those for now as many backends reports tons # torch.set_float32_matmul_precision('high') # enable to test in fp32 @@ -44,8 +46,8 @@ def init_model(): return resnet18().to(torch.float32).cuda() -def eval(mod, inp): - return mod(inp) +def evaluate(mod, val): + return mod(val) if __name__ == '__main__': @@ -56,18 +58,19 @@ if __name__ == '__main__': # repeat test results = {} times = [] - print('eager initial eval:', timed(lambda: eval(model, inp))[1]) + print('eager initial eval:', timed(lambda: evaluate(model, inp))[1]) for i in range(N_ITERS): inp = generate_data(16)[0] - _res, time = timed(lambda: eval(model, inp)) + _res, time = timed(lambda: evaluate(model, inp)) times.append(time) results['default'] = np.median(times) print('dynamo available backends:', dynamo.list_backends()) for backend in dynamo.list_backends(): try: - torch._dynamo.reset() # required before changing backends - eval_dyn = dynamo.optimize(backend)(eval) + # required before changing backends + torch._dynamo.reset() # pylint: disable=protected-access + eval_dyn = dynamo.optimize(backend)(evaluate) print('dynamo initial eval:', backend, timed(lambda: eval_dyn(model, inp))[1]) times = [] for i in range(N_ITERS): diff --git a/cli/train-lora.py b/cli/train-lora.py deleted file mode 100755 index 6f82063a6..000000000 --- a/cli/train-lora.py +++ /dev/null @@ -1,274 +0,0 @@ -#!/bin/env python - -""" -Extract approximating LoRA by SVD from two SD models -Based on: - -Train LoRA with custom preprocessing, tagging and bucketing - -Disabled/broken: -- `accelerate` with *dynamo* enabled -- `xformers` due to *faketensors* requirement -- `mem_eff_attn` due to *forwardfunc* mismatch -- 'use_8bit_adam` due to *bitsandbyttes* CUDA errors -""" - -import os -import re -import gc -import sys -import json -import time -import shutil -import argparse -import tempfile -import torch -import logging -import importlib -import transformers -from pathlib import Path -from modules.util import log, Map, get_memory -import modules.process -import modules.sdapi - -latents = importlib.import_module('modules.lora-latents') - -lora_path = os.path.abspath(os.path.join(os.path.dirname(__file__), os.pardir, 'modules', 'lora')) -sys.path.append(lora_path) -lycoris_path = os.path.abspath(os.path.join(os.path.dirname(__file__), os.pardir, 'modules', 'lycoris')) -sys.path.append(lycoris_path) -from train_network import train - -options = Map({ - "bucket_no_upscale": False, - "bucket_reso_steps": 64, - "cache_latents": True, - "caption_dropout_every_n_epochs": None, - "caption_dropout_rate": 0.0, - "caption_extension": ".txt", - "caption_extention": ".txt", - "caption_tag_dropout_rate": 0.0, - "clip_skip": None, - "color_aug": False, - "dataset_repeats": 1, - "debug_dataset": False, - "enable_bucket": False, - "face_crop_aug_range": None, - "flip_aug": False, - "full_fp16": False, - "gradient_accumulation_steps": 1, - "gradient_checkpointing": False, - "in_json": "", - "keep_tokens": None, - "learning_rate": 5e-05, - "log_prefix": None, - "logging_dir": None, - "lr_scheduler_num_cycles": 1, - "lr_scheduler_power": 1, - "lr_scheduler": "cosine", - "lr_warmup_steps": 0, - "max_bucket_reso": 1024, - "max_data_loader_n_workers": 8, - "max_grad_norm": 0.0, - "max_token_length": None, - "max_train_epochs": None, - "max_train_steps": 5000, - "mem_eff_attn": False, - "min_bucket_reso": 256, - "mixed_precision": "fp16", - "network_alpha": 1.0, - "network_args": None, - "network_dim": 16, - "network_module": "networks.lora", - "network_train_text_encoder_only": False, - "network_train_unet_only": False, - "network_weights": None, - "no_metadata": False, - "output_dir": "", - "output_name": "", - "persistent_data_loader_workers": False, - "pretrained_model_name_or_path": "", - "prior_loss_weight": 1.0, - "random_crop": False, - "reg_data_dir": None, - "resolution": "512,512", - "resume": None, - "save_every_n_epochs": None, - "save_last_n_epochs_state": None, - "save_last_n_epochs": None, - "save_model_as": "ckpt", - "save_n_epoch_ratio": None, - "save_precision": "fp16", - "save_state": False, - "seed": 42, - "shuffle_caption": False, - "text_encoder_lr": 5e-05, - "train_batch_size": 1, - "train_data_dir": "", - "training_comment": "mood-magic", - "unet_lr": 0.001, - "use_8bit_adam": False, - "v_parameterization": False, - "v2": False, - "vae": None, - "xformers": False, -}) - - -def mem_stats(): - gc.collect() - if torch.cuda.is_available(): - with torch.no_grad(): - torch.cuda.empty_cache() - with torch.cuda.device('cuda'): - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - mem = get_memory() - log.info({ 'memory': { 'ram': mem.ram, 'gpu': mem.gpu } }) - - -if __name__ == '__main__': - parser = argparse.ArgumentParser(description = 'train lora') - parser.add_argument('--model', type=str, default=None, required=False, help='original model to use a base for training, default: active model') - parser.add_argument('--input', '--dataset', type=str, default=None, required=True, help='input folder with training images') - parser.add_argument('--output', '--lora', type=str, default=None, required=True, help='lora name') - parser.add_argument('--tag', type=str, default=None, required=False, help='primary tag') - parser.add_argument('--dir', type=str, default=None, required=False, help='folder containing lora checkpoints') - parser.add_argument('--interim', type=int, default=0, help = 'save interim checkpoints after n epoch') - parser.add_argument('--process', type=str, default='original', required=False, help='list of processing steps: original,face,body,blur,range,upscale,restore') - parser.add_argument('--noprocess', default = False, action='store_true', help = 'skip processing and use existing input data') - parser.add_argument('--notrain', default = False, action='store_true', help = 'just run processing and skip training') - parser.add_argument('--nocaptions', default = False, action='store_true', help = 'skip creating captions and tags') - parser.add_argument('--nolatents', default = False, action='store_true', help = 'skip generating vae latents') - parser.add_argument('--offline', default = False, action='store_true', help = 'do not use webui server for processing') - parser.add_argument('--shutdown', default = False, action='store_true', help = 'shutdown webui server') - parser.add_argument('--gradient', type=int, default=1, required=False, help='gradient accumulation steps, default: %(default)s') - parser.add_argument('--steps', type=int, default=4000, required=False, help='training steps, default: %(default)s') - parser.add_argument('--dim', type=int, default=40, required=False, help='network dimension, default: %(default)s') - parser.add_argument('--repeats', type=int, default=10, required=False, help='number of repeats per image, default: %(default)s') - parser.add_argument('--alpha', type=float, default=0, required=False, help='alpha for weights scaling, default: half of dim') - parser.add_argument('--batch', type=int, default=1, required=False, help='batch size, default: %(default)s') - parser.add_argument('--lr', type=float, default=1e-04, required=False, help='model learning rate, default: %(default)s') - parser.add_argument('--unetlr', type=float, default=1e-04, required=False, help='unet learning rate, default: %(default)s') - parser.add_argument('--textlr', type=float, default=5e-05, required=False, help='text encoder learning rate, default: %(default)s') - parser.add_argument('--dreambooth', default=False, action='store_true', help = "use dreambooth style training") - parser.add_argument('--lycoris', default=False, action='store_true', help = "use lycoris style training") - parser.add_argument('--debug', default=False, action='store_true', help = "enable debug logging") - args = parser.parse_args() - defaults = Map({ 'options': {}, 'flags': {} }) if args.offline else Map(modules.sdapi.options()) - - if args.debug: - log.setLevel(logging.DEBUG) - log.debug({ 'debug': True }) - if args.model is None: - args.model = defaults.options.get('sd_model_checkpoint', None) - args.model = args.model.split(' [')[0] if args.model is not None else None - if args.dir is None: - args.dir = defaults.flags.get('lora_dir', None) - if not os.path.isabs(args.model) and args.dir is not None and not os.path.exists(args.model): - args.model = os.path.abspath(os.path.join(args.dir, os.pardir, 'Stable-diffusion', args.model)) - if args.dir is None: - args.dir = os.path.join(args.input, 'lora') - if not os.path.exists(args.model) or not os.path.isfile(args.model): - log.error({ 'lora cannot find model': args.model }) - exit(1) - if not os.path.exists(args.input) or not os.path.isdir(args.input): - log.error({ 'lora cannot find training dir': args.input }) - exit(1) - if not os.path.exists(args.dir) or not os.path.isdir(args.dir): - log.error({ 'lora cannot find training dir': args.dir }) - exit(1) - options.pretrained_model_name_or_path = args.model - options.output_dir = args.dir - options.output_name = args.output - options.max_train_steps = args.steps - options.network_dim = args.dim - options.network_alpha = args.dim // 2 if args.alpha == 0 else args.alpha - options.gradient_accumulation_steps = args.gradient - options.save_every_n_epochs = args.interim if args.interim > 0 else None - options.learning_rate = args.lr - options.unet_lr = args.unetlr - options.text_encoder_lr = args.textlr - options.train_batch_size = args.batch - log.info({ 'train lora args': vars(options) }) - transformers.logging.set_verbosity_error() - mem_stats() - - json_file = os.path.join(tempfile.gettempdir(), args.output, args.output + '.json') - base = os.path.join(tempfile.gettempdir(), args.output) - options.train_data_dir = base - res = None - - if args.dreambooth: - log.info({ 'using dreambooth style training': True }) - options.in_json = None - else: - options.in_json = json_file - - for root, _sub_dirs, folder in os.walk(args.input): - files = [os.path.join(root, f) for f in folder] - - if not args.noprocess: - # preprocess - processing_options = args.process.split(',') - processing_options = [opt.strip() for opt in re.split(',| ', args.process)] - log.info({ 'processing steps': processing_options }) - - if os.path.exists(json_file): - os.remove(json_file) - - steps = [step for step in processing_options if step in ['face', 'body', 'original']] - for step in steps: - # processing_options = [step for step in processing_options if step not in ['face', 'body', 'original']].append(step) - if step == 'face': - opts = [step for step in processing_options if step not in ['body', 'original']] - if step == 'body': - opts = [step for step in processing_options if step not in ['face', 'original', 'upscale', 'restore']] - if step == 'original': - opts = [step for step in processing_options if step not in ['face', 'body', 'upscale', 'restore', 'blur', 'range']] - log.info({ 'processing step': opts }) - concept = step - if concept == 'original' and args.tag is not None: - concept = args.tag.split(',')[0].strip() - dir = os.path.join(base, str(args.repeats) + '_' + concept) - if os.path.exists(dir): - shutil.rmtree(dir, ignore_errors=True) - Path(dir).mkdir(parents=True, exist_ok=True) - - for f in files: - try: - res, metadata = modules.process.process_file(f = f, dst = dir, preview = False, offline = args.offline, txt = args.dreambooth, tag = args.tag, opts = opts) - if not args.dreambooth: - with open(json_file, "w") as outfile: - outfile.write(json.dumps(metadata, indent=2)) - except ValueError as e: - exit(1) - log.info({ 'processed step': step, 'outputs': res, 'inputs': len(files), 'metadata': json_file, 'path': dir }) - - modules.process.unload_models() - mem_stats() - - - dirs = [os.path.join(base, dir) for dir in os.listdir(base) if os.path.isdir(os.path.join(base, dir))] - log.info({ 'input datasets': dirs, 'metadata': json_file }) - - if not args.nolatents and not args.dreambooth: - # create latents - for dir in dirs: - latents.create_vae_latents(Map({ 'input': dir, 'json': json_file })) - latents.unload_vae() - mem_stats() - else: - log.info({ 'skip processing': len(files), 'metadata': json_file, 'path': dir }) - - if args.shutdown: - log.info({ 'server shutdown required': True }) - modules.sdapi.shutdown() - time.sleep(1) - - if args.lycoris: - log.info({ 'using lycoris network': True }) - options.network_module = 'lycoris.kohya' - if not args.notrain: - train(options) - mem_stats() diff --git a/cli/train-ti.py b/cli/train-ti.py deleted file mode 100755 index 983b77139..000000000 --- a/cli/train-ti.py +++ /dev/null @@ -1,591 +0,0 @@ -#!/bin/env python -# pylint: disable=no-member -""" -simple implementation of training api: `/sdapi/v1/train` -- supports: create embedding, image preprocess, train embedding (with all known parameters) -- does not (yet) support: create hyper-network, train hyper-network -- compatible with progress api: `/sdapi/v1/progress` -- if interrupted, auto-continues from last known step -- create and preprocess executed as sync jobs -- train is executed as async job with progress monitoring -""" - -import argparse -import asyncio -import logging -import math -import os -import sys -import time -import importlib -from pathlib import Path, PurePath - -import filetype -from PIL import Image - -sys.path.append(os.path.join(os.path.dirname(__file__), 'modules')) -from modules.util import Map, log, set_logfile -from modules.sdapi import close, get, interrupt, post, progress, session -from modules.process import process_images -from modules.grid import grid -create_preview = importlib.import_module('modules.preview-embeddings').create_preview -plot = importlib.import_module('modules.train-losschart').plot -extract = importlib.import_module('modules.video-extract').extract -gen_loss_rate_str = importlib.import_module('modules.train-lossrate').gen_loss_rate_str - -images = [] -args = {} -options = None -cmdflags = None -args = Map({ - "training_model": "sd-v15-runwayml.ckpt", - "extract_video": { - "rate": 0, - "fps": 5, - "vstart": 0, - "vend": 0 - }, - "create_embedding": { - "name": "test", - "num_vectors_per_token": 1, - "overwrite_old": False, - "init_text": "*" - }, - "preprocess": { - "id_task": 0, - "process_src": "", - "process_dst": "", - "process_width": 512, - "process_height": 512, - "process_flip": False, - "process_split": False, - "process_caption": True, - "process_caption_deepbooru": False, - "preprocess_txt_action": "ignore", - "process_focal_crop": True, - "process_focal_crop_face_weight": 0.9, - "process_focal_crop_entropy_weight": 0.3, - "process_focal_crop_edges_weight": 0.5, - "process_focal_crop_debug": False, - "split_threshold": 0.5, - "overlap_ratio": 0.2, - "process_multicrop": None, - "process_multicrop_mindim": None, - "process_multicrop_maxdim": None, - "process_multicrop_minarea": None, - "process_multicrop_maxarea": None, - "process_multicrop_objective": None, - "process_multicrop_threshold": None, - }, - "train_embedding": { - "id_task": 0, - "embedding_name": "", - "learn_rate": -1, - "batch_size": 1, - "steps": 500, - "data_root": "", - "log_directory": "train/log", - "template_filename": "subject_filewords.txt", - "gradient_step": 20, - "training_width": 512, - "training_height": 512, - "shuffle_tags": False, - "tag_drop_out": 0, - "clip_grad_mode": "disabled", - "clip_grad_value": "0.1", - "latent_sampling_method": "once", - "create_image_every": -1, - "save_embedding_every": -1, - "save_image_with_stored_embedding": False, - "preview_from_txt2img": False, - "preview_prompt": "", - "preview_negative_prompt": "blurry, duplicate, ugly, deformed, low res, watermark, text", - "preview_steps": 20, - "preview_sampler_index": 0, - "preview_cfg_scale": 6, - "preview_seed": -1, - "preview_width": 512, - "preview_height": 512, - "varsize": False, - "use_weight": False, - }, -}) - - -async def plotloss(params): - logdir = os.path.abspath(os.path.join(cmdflags['embeddings_dir'], '../train/log')) - try: - plot(logdir, params.name) - except Exception as err: - log.warning({ 'loss chart error': err }) - - -async def captions(docs: list): - exclude = ['a', 'in', 'on', 'out', 'at', 'the', 'and', 'with', 'next', 'to', 'it', 'for', 'of', 'into', 'that'] - d = dict() - for f in docs: - text = open(f, 'r', encoding='utf-8') - for line in text: - line = line.strip() - line = line.lower() - words = line.split(" ") - for word in words: - if word in exclude: - continue - d[word] = d[word] + 1 if word in d else 1 - pairs = ((value, key) for (key,value) in d.items()) - sort = sorted(pairs, reverse = True) - if len(sort) > 10: - del sort[10:] - d = {k: v for v, k in sort} - log.info({ 'top captions': d }) - - -async def preprocess_cleanup(params): - log.info({ 'preprocess cleanup': params.dst }) - for f in Path(params.dst).glob('*.png'): - f.unlink() - for f in Path(params.dst).glob('*.jpg'): - f.unlink() - for f in Path(params.dst).glob('*.txt'): - f.unlink() - try: - if os.path.isdir(params.dst): - Path(params.dst).rmdir() - except Exception as err: - log.warning({ 'preprocess cleanup': params.dst, 'error': err }) - - -async def preprocess_builtin(params): - global images # pylint: disable=global-statement - log.debug({ 'preprocess start' }) - files = [os.path.join(params.src, f) for f in os.listdir(params.src) if os.path.isfile(os.path.join(params.src, f))] - candidates = [f for f in files if filetype.is_image(f)] - not_images = [f for f in files if (not filetype.is_image(f) and not f.endswith('.txt'))] - images = [] - low_res = [] - for f in candidates: - img = Image.open(f) - mp = (img.size[0] * img.size[1]) / 1024 / 1024 - if mp < 1 or img.size[0] < 512 or img.size[1] < 512: - low_res.append(f) - os.rename(f, f + '.skip') - else: - images.append(f) - log.debug({ 'preprocess skipping': not_images }) - log.debug({ 'preprocess low res': low_res }) - args.preprocess.process_src = params.src - args.preprocess.process_dst = params.dst - log.debug({ 'preprocess args': args.preprocess }) - _res = await post('/sdapi/v1/preprocess', json = args.preprocess) - processed = [os.path.join(params.dst, f) for f in os.listdir(params.dst) if os.path.isfile(os.path.join(params.dst, f))] - processed_imgs = [f for f in processed if f.endswith('.png')] - processed_docs = [f for f in processed if f.endswith('.txt')] - log.info({ 'preprocess': { - 'source': params.src, - 'destination': params.dst, - 'files': len(files), - 'images': len(images), - 'processed': len(processed_imgs), - 'captions': len(processed_docs), - 'skipped': len(not_images), - 'low-res': len(low_res) } - }) - if len(processed_docs) > 0: - await captions(processed_docs) - return len(processed_imgs) - - -async def preprocess(params): - global images # pylint: disable=global-statement - res = 0 - if os.path.isfile(params.src): - if not filetype.is_video(params.src): - kind = filetype.guess(params.src) - log.error({ 'preprocess error': { 'not a valid movie file': params.src, 'guess': kind } }) - else: - extract_dst = os.path.join(params.dst, 'extract') - log.debug({ 'preprocess args': args.extract_video }) - images = extract(params.src, extract_dst, rate = args.extract_video.rate, fps = args.extract_video.fps, start = args.extract_video.vstart, end = args.extract_video.vend) # extract keyframes from movie - if images > 0: - params.src = extract_dst - res = await preprocess(params) # call again but now with keyframes - else: - log.error({ 'preprocess video extract': 'no images' }) - elif os.path.isdir(params.src): - if params.overwrite: - await preprocess_cleanup(params) - elif os.path.isdir(params.dst): - log.error({ 'preprocess output folder already exists': params.dst }) - return 0 - - if params.preprocess == 'builtin': - res = await preprocess_builtin(params) - i = [os.path.join(params.dst, f) for f in os.listdir(params.dst) if os.path.isfile(os.path.join(params.dst, f)) and filetype.is_image(os.path.join(params.dst, f))] - images = [Image.open(img) for img in i] - res = len(images) - - elif params.preprocess == 'custom': - t0 = time.perf_counter() - args.preprocess.process_src = params.src - args.preprocess.process_dst = params.dst - process_images(src = params.src, dst = params.dst) - i = [os.path.join(params.dst, f) for f in os.listdir(params.dst) if os.path.isfile(os.path.join(params.dst, f)) and filetype.is_image(os.path.join(params.dst, f))] - images = [Image.open(img) for img in i] - t1 = time.perf_counter() - log.info({ 'preprocess': { 'source': params.src, 'destination': params.dst, 'images': len(images), 'time': round(t1 - t0, 2) } }) - res = len(images) - - else: - args.preprocess.process_dst = params.src - i = [os.path.join(params.src, f) for f in os.listdir(params.src) if os.path.isfile(os.path.join(params.src, f)) and filetype.is_image(os.path.join(params.src, f))] - images = [Image.open(img) for img in i] - res = len(images) - - else: - log.error({ 'preprocess error': { 'not a valid input': params.src } }) - if len(images) > 0: - img = grid(images, labels = None, width = 2048, height = 2048, border = 8, square = True) - logdir = os.path.abspath(os.path.join(cmdflags['embeddings_dir'], '../train/log')) - Path(logdir).mkdir(parents = True, exist_ok = True) - fn = os.path.join(logdir, params.name + '.inputs.jpg') - img.save(fn) - log.info({ 'preprocess input grid': fn }) - return res - - -async def check(params): - global options # pylint: disable=global-statement - options = await get('/sdapi/v1/options') - global cmdflags # pylint: disable=global-statement - cmdflags = await get('/sdapi/v1/cmd-flags') - - logdir = os.path.abspath(os.path.join(cmdflags['embeddings_dir'], '../train/log', params.name)) - logfile = os.path.abspath(os.path.join(cmdflags['embeddings_dir'], '../train/log', params.name + '.train.log')) - set_logfile(logfile) - - log.info({ 'checking server options' }) - - options['training_image_repeats_per_epoch'] = 1 - - if params.skipmodel: - log.info({ 'using model': options['sd_model_checkpoint'] }) - else: - log.debug({ 'check model': args.training_model }) - if len(args.training_model) > 0 and not options['sd_model_checkpoint'].startswith(args.training_model): - models = await get('/sdapi/v1/sd-models') - models = [obj["title"] for obj in models] - found = [i for i in models if i.startswith(args.training_model)] - if len(found) == 0: - log.error({ 'model not found': args.training_model, 'available': models }) - exit() - else: - log.warning({ 'switching model': found[0] }) - options['sd_model_checkpoint'] = found[0] - - log.debug({ 'check embedding': params.name }) - - lst = os.path.join(cmdflags['embeddings_dir']) - log.debug({ 'embeddings folder': lst }) - path = Path(cmdflags['embeddings_dir']).glob(f'{params.name}.pt*') - matches = [f for f in path] - for match in matches: - if params.overwrite: - log.info({ 'delete embedding': match.name }) - os.remove(os.path.join(cmdflags['embeddings_dir'], match.name)) - else: - log.error({ 'embedding exists': match.name }) - await close() - exit() - f = os.path.join(logdir, 'train.csv') - if os.path.isfile(f): - if params.overwrite: - log.info({ 'delete training log': f }) - os.remove(os.path.join(logdir, 'train.csv')) - else: - log.warning({ 'training log exists': f }) - f = os.path.join(logdir, '..', params.name, '.png') - if os.path.isfile(f): - if params.overwrite: - log.info({ 'delete training graph': f }) - os.remove(f) - - log.debug({ 'options': 'update' }) - await post('/sdapi/v1/options', options) - - return - - -async def create(params): - log.debug({ 'create start' }) - if not os.path.isdir(args.preprocess.process_dst): - log.error({ 'train source not found': args.preprocess.process_dst }) - exit() - if params.vectors == -1: # dynamically determine number of vectors depending on number of input images - if len(images) <= 20: - vectors = 2 - elif len(images) <= 100: - vectors = 4 - else: - vectors = 6 - else: - vectors = params.vectors - if os.path.exists(params.name) and os.path.isfile(params.name): - log.info({ 'deleting existing embedding': { 'name': params.name } }) - os.remove(params.name) - args.create_embedding.name = params.name - words = params.init.split(',') - if len(words) > vectors: - params.init = ','.join(words[:vectors]) - log.warning({ 'create embedding init words cut': params.init }) - args.create_embedding.init_text = params.init - args.create_embedding.num_vectors_per_token = vectors - log.debug({ 'create args': args.create_embedding }) - res = await post('/sdapi/v1/create/embedding', args.create_embedding) - if 'info' in res: - log.info({ 'create embedding': { 'name': params.name, 'init': params.init, 'vectors': vectors, 'message': res.info } }) - else: - log.error({ 'create failed:', res }) - return None - log.debug({ 'create end' }) - return params.name - - -async def train(params): - log.debug({ 'train start' }) - args.train_embedding.embedding_name = params.name - - imgs = [f for f in os.listdir(args.preprocess.process_dst) if os.path.isfile(os.path.join(args.preprocess.process_dst, f)) and filetype.is_image(os.path.join(args.preprocess.process_dst, f))] - args.train_embedding.data_root = args.preprocess.process_dst - if len(imgs) == 0: - log.error({ 'train no input images in folder': args.preprocess.process_dst }) - return - - if params.grad == -1: - args.train_embedding.gradient_step = len(imgs) // args.train_embedding.batch_size - divisor = args.train_embedding.gradient_step // 60 - args.train_embedding.gradient_step = args.train_embedding.gradient_step // (1 + divisor) - log.info({ 'dynamic gradient step': args.train_embedding.gradient_step }) - if params.steps == -1: - args.train_embedding.steps = params.maxsteps // args.train_embedding.gradient_step - log.info({ 'dynamic steps': args.train_embedding.steps, 'estimated total steps': args.train_embedding.steps * args.train_embedding.gradient_step * args.train_embedding.batch_size }) - - epoch_size = args.train_embedding.batch_size * args.train_embedding.gradient_step - if args.train_embedding.create_image_every == -1: - args.train_embedding.create_image_every = args.train_embedding.steps // 10 - if args.train_embedding.save_embedding_every == -1: - args.train_embedding.save_embedding_every = args.train_embedding.steps // 10 - if args.train_embedding.learn_rate == -1: - loss_args = { - "steps": args.train_embedding.steps, - "step": args.train_embedding.create_image_every, - "loss_start": params.rstart, - "loss_end": params.rend, - "loss_type": 'power', - "power": params.rdescend - } - args.train_embedding.learn_rate = gen_loss_rate_str(**loss_args) - log.info({ 'dynamic learn-rate': loss_args }) - log.debug({ 'learn rate': args.train_embedding.learn_rate, 'params': loss_args }) - - log.info({ 'train embedding': { - 'name': params.name, - 'source': args.preprocess.process_dst, - 'images': len(imgs), - 'steps': args.train_embedding.steps, - 'batch': args.train_embedding.batch_size, - 'gradient-step': args.train_embedding.gradient_step, - 'sampling': args.train_embedding.latent_sampling_method, - 'epoch-size': epoch_size } - }) - log.info({ 'learn rate': args.train_embedding.learn_rate }) - log.debug({ 'train args': args.train_embedding }) - t0 = time.time() - res = await post('/sdapi/v1/train/embedding', args.train_embedding) - log.info({ 'train result': res }) - t1 = time.time() - log.info({ 'train embedding finished': { 'name': params.name, 'time': round(t1 - t0) } }) - log.debug({ 'train end': res.info if 'info' in res else res }) - return - - -async def pipeline(params): - log.debug({ 'pipeline start' }) - - # interrupt - await interrupt() - - # preprocess - num = await preprocess(params) - if num == 0: - log.warning({ 'preprocess': 'no resulting images'}) - return - - # create embedding - name = await create(params) - if not params.name in name: - log.error({ 'create embedding failed': name }) - return - - # train embedding - await train(params) - - await plotloss(params) - - # create_preview(params.name, params.init) - - log.debug({ 'pipeline end' }) - return - - -async def monitor(params): - step = 0 - t0 = time.perf_counter() - t1 = time.perf_counter() - log.info({' starting monitor': t0 }) - finished = 0 - while True: - await asyncio.sleep(params.monitor) - res = await progress() - if not 'state' in res: - log.info({ 'monitor disconnected': res }) - break - if (res.state.job_count == params.steps and res.state.job_no >= res.state.job_count) or (res.eta_relative < 0) or (res.interrupted) or (res.state.job_count == 0): # need exit case if interrupted or failed - if res.interrupted: - log.info({ 'monitor interrupted': { 'embedding': params.name } }) - break # exit for monitor job - else: - finished += 1 - if finished >= 2: # do it more than once since preprocessing job can finish just in time for monitor to finish - log.info({ 'monitor finished': { 'embedding': params.name } }) - break - else: - if res.state.job_no == 0: - step = 0 - t0 = time.perf_counter() - t1 = time.perf_counter() - try: - if 'Loss:' in res.textinfo: - text = res.textinfo.split('
')[0].split() - loss = float(text[-1]) - else: - loss = -1 - except: - loss = -1 - if math.isnan(loss): - log.error({ 'monitor': { 'progress': round(100 * res.progress), 'embedding': params.name, 'eta': round(res.eta_relative), 'step': res.state.job_no, 'steps': res.state.job_count, 'loss': 'nan' } }) - await interrupt() - else: - elapsed = t1 - t0 - log.info({ 'monitor': { - 'job': res.state.job, - 'progress': round(100 * res.progress), - 'embedding': params.name, - 'epoch': (1 + res.state.job_no // len(images)) if len(images) > 0 else 'n/a', - 'step': res.state.job_no, - 'steps': res.state.job_count, - 'loss': loss if loss > -1 else 'n/a', - 'total': round(1.0 * elapsed * res.state.job_count / res.state.job_no) if res.state.job_no > 0 and t1 != t0 else 'n/a', - 'elapsed': round(elapsed), - 'remaining': round(res.eta_relative), - 'it/s': round((res.state.job_no - step) / (time.perf_counter() - t1), 2) } - }) - if step % 10 == 0: - await plotloss(params) - step = res.state.job_no - t1 = time.perf_counter() - return - - -async def main(): - parser = argparse.ArgumentParser(description="sd train ti pipeline") - parser.add_argument("--name", type = str, required = True, help = "embedding name, set to auto to use src folder name") - parser.add_argument("--src", type = str, required = True, help = "source image folder or movie file") - parser.add_argument("--init", type = str, default = "person", required = False, help = "initialization class, default: %(default)s") - parser.add_argument("--dst", type = str, default = "/tmp", required = False, help = "destination image folder for processed images, default: %(default)s") - parser.add_argument("--steps", type = int, default = -1, required = False, help = "training steps, default: %(default)s") - parser.add_argument("--maxsteps", type = int, default = 5000, required = False, help = "max training steps used when dynamic gradient is active, default: %(default)s") - parser.add_argument("--vectors", type = int, default = -1, required = False, help = "number of vectors per token, default: dynamic based on number of input images") - parser.add_argument("--batch", type = int, default = 1, required = False, help = "batch size, default: %(default)s") - parser.add_argument("--rate", type = str, default = "", required = False, help = "learn rate, default: dynamic") - parser.add_argument("--rstart", type = float, default = 0.02, required = False, help = "starting learn rate if using dynamic rate, default: %(default)s") - parser.add_argument("--rend", type = float, default = 0.0005, required = False, help = "ending learn rate if using dynamic rate, default: %(default)s") - parser.add_argument("--rdescend", type = float, default = 2, required = False, help = "learn rate descend power when using dynamic rate, default: %(default)s") - parser.add_argument("--grad", type = int, default = -1, required = False, help = "accumulate gradient over n images, default: : %(default)s") - parser.add_argument("--type", type = str, default = 'subject', required = False, help = "training type: subject/style/unknown, default: %(default)s") - parser.add_argument('--overwrite', default = False, action='store_true', help = "overwrite existing embedding, default: %(default)s") - parser.add_argument("--vstart", type = float, default = 0, required = False, help = "if processing video skip first n seconds, default: %(default)s") - parser.add_argument("--vend", type = float, default = 0, required = False, help = "if processing video skip last n seconds, default: %(default)s") - parser.add_argument('--skipcaption', default = False, action='store_true', help = "do not auto-generate captions, default: %(default)s") - parser.add_argument('--skipmodel', default = False, action='store_true', help = "skip model validation and switch, default: %(default)s") - parser.add_argument('--preprocess', type = str, choices=['builtin', 'custom', 'none'], default = 'custom', help = "preprocessing type, default: %(default)s") - parser.add_argument('--nocleanup', default = False, action='store_true', help = "skip cleanup after completion, default: %(default)s") - parser.add_argument("--monitor", type = int, default = 30, required = False, help = "progress monitor frequency, default: : %(default)s") - parser.add_argument('--debug', default = False, action='store_true', help = "print extra debug information, default: %(default)s") - params = parser.parse_args() - if params.debug: - log.setLevel(logging.DEBUG) - log.debug({ 'debug': True }) - log.debug({ 'args': params.__dict__ }) - home = Path(sys.argv[0]).parent - global args # pylint: disable=global-statement - if params.vstart > 0: - args.extract_video.vstart = params.vstart - if params.vend > 0: - args.extract_video.vend = params.vend - if params.steps > -1: - args.train_embedding.steps = params.steps - if params.batch > -1: - args.train_embedding.batch_size = params.batch - if params.rate != '': - args.train_embedding.learn_rate = params.rate - if params.grad > -1: - args.train_embedding.gradient_step = params.grad - if params.type == 'subject': - if params.skipcaption: - args.train_embedding.template_filename = 'subject.txt' - args.preprocess.process_caption = False - else: - args.train_embedding.template_filename = 'subject_filewords.txt' - elif params.type == 'style': - if params.skipcaption: - args.train_embedding.template_filename = 'style.txt' - args.preprocess.process_caption = False - else: - args.train_embedding.template_filename = 'style_filewords.txt' - else: - if params.skipcaption: - args.train_embedding.template_filename = 'unknown.txt' - args.preprocess.process_caption = False - else: - args.train_embedding.template_filename = 'unknown_filewords.txt' - if params.name == 'auto': - params.name = PurePath(params.src).name - log.info({ 'training name': params.name }) - if params.dst == "/tmp": - params.dst = os.path.join("/tmp/train", params.name) - log.debug({ 'args': params.__dict__ }) - params.src = os.path.abspath(params.src) - params.dst = os.path.abspath(params.dst) - - try: - await session() - await check(params) - a = asyncio.create_task(pipeline(params)) - b = asyncio.create_task(monitor(params)) - await asyncio.gather(a, b) # wait for both pipeline and monitor to finish - except Exception as e: - log.error({ 'exception': e }) - finally: - if not params.nocleanup: - await preprocess_cleanup(params) - await close() - return - -if __name__ == "__main__": - log.info({ 'train textual inversion' }) - try: - asyncio.run(main()) - except KeyboardInterrupt: - log.warning({ 'interrupted': 'keyboard request' }) - # asyncio.run(interrupt()) diff --git a/cli/train.py b/cli/train.py new file mode 100755 index 000000000..8ee5a675e --- /dev/null +++ b/cli/train.py @@ -0,0 +1,395 @@ +#!/usr/bin/env python + +# system imports +import os +import re +import gc +import sys +import json +import time +import shutil +import pathlib +import asyncio +import logging +import tempfile +import argparse +import warnings +warnings.filterwarnings(action="ignore", category=DeprecationWarning) +warnings.filterwarnings(action="ignore", category=UserWarning) +warnings.filterwarnings(action="ignore", category=FutureWarning) + +# 3rd party imports +import filetype +import torch +from tqdm.rich import tqdm + +# local imports +import util +import sdapi +import options +import process +import latents + + +# globals +args = None +log = logging.getLogger('train') +valid_steps = ['original', 'face', 'body', 'blur', 'range', 'upscale', 'restore', 'interrogate', 'resize', 'square', 'segment'] +log_file = os.path.join(os.path.dirname(__file__), 'train.log') + +# methods + +def setup_logging(clean=False): + try: + if clean and os.path.isfile(log_file): + os.remove(log_file) + time.sleep(0.1) # prevent race condition + except: + pass + from rich.theme import Theme + from rich.logging import RichHandler + from rich.console import Console + from rich.pretty import install as pretty_install + from rich.traceback import install as traceback_install + console = Console(log_time=True, log_time_format='%H:%M:%S-%f', theme=Theme({ + "traceback.border": "black", + "traceback.border.syntax_error": "black", + "inspect.value.border": "black", + })) + # logging.getLogger("urllib3").setLevel(logging.ERROR) + # logging.getLogger("httpx").setLevel(logging.ERROR) + level = logging.DEBUG if args.debug else logging.INFO + logging.basicConfig(level=logging.ERROR, format='%(asctime)s | %(name)s | %(levelname)s | %(module)s | %(message)s', filename=log_file, filemode='a', encoding='utf-8', force=True) + log.setLevel(logging.DEBUG) # log to file is always at level debug for facility `sd` + pretty_install(console=console) + traceback_install(console=console, extra_lines=1, width=console.width, word_wrap=False, indent_guides=False, suppress=[]) + rh = RichHandler(show_time=True, omit_repeated_times=False, show_level=True, show_path=False, markup=False, rich_tracebacks=True, log_time_format='%H:%M:%S-%f', level=level, console=console) + rh.set_name(level) + while log.hasHandlers() and len(log.handlers) > 0: + log.removeHandler(log.handlers[0]) + log.addHandler(rh) + + +def mem_stats(): + gc.collect() + if torch.cuda.is_available(): + with torch.no_grad(): + torch.cuda.empty_cache() + with torch.cuda.device('cuda'): + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + mem = util.get_memory() + peak = { 'active': mem['gpu-active']['peak'], 'allocated': mem['gpu-allocated']['peak'], 'reserved': mem['gpu-reserved']['peak'] } + log.debug(f"memory cpu: {mem.ram} gpu current: {mem.gpu} gpu peak: {peak}") + + +def parse_args(): + global args # pylint: disable=global-statement + parser = argparse.ArgumentParser(description = 'SD.Next Train') + + group_main = parser.add_argument_group('Main') + group_main.add_argument('--type', type=str, choices=['embedding', 'ti', 'lora', 'lyco', 'dreambooth', 'hypernetwork'], default=None, required=True, help='training type') + group_main.add_argument('--model', type=str, default='', required=False, help='base model to use for training, default: current loaded model') + group_main.add_argument('--name', type=str, default=None, required=True, help='output filename') + group_main.add_argument('--tag', type=str, default='person', required=False, help='primary tags, default: %(default)s') + + group_data = parser.add_argument_group('Dataset') + group_data.add_argument('--input', type=str, default=None, required=True, help='input folder with training images') + group_data.add_argument('--output', type=str, default='', required=False, help='where to store processed images, default is system temp/train') + group_data.add_argument('--process', type=str, default='original,interrogate,resize,square', required=False, help=f'list of possible processing steps: {valid_steps}, default: %(default)s') + + group_train = parser.add_argument_group('Train') + group_train.add_argument('--gradient', type=int, default=1, required=False, help='gradient accumulation steps, default: %(default)s') + group_train.add_argument('--steps', type=int, default=2500, required=False, help='training steps, default: %(default)s') + group_train.add_argument('--batch', type=int, default=1, required=False, help='batch size, default: %(default)s') + group_train.add_argument('--lr', type=float, default=1e-04, required=False, help='model learning rate, default: %(default)s') + group_train.add_argument('--dim', type=int, default=32, required=False, help='network dimension or number of vectors, default: %(default)s') + + # lora params + group_train.add_argument('--repeats', type=int, default=10, required=False, help='number of repeats per image, default: %(default)s') + group_train.add_argument('--alpha', type=float, default=0, required=False, help='lora/lyco alpha for weights scaling, default: dim/2') + group_train.add_argument('--algo', type=str, default=None, choices=['locon', 'loha', 'lokr', 'ia3'], required=False, help='alternative lyco algoritm, default: %(default)s') + group_train.add_argument('--args', type=str, default=None, required=False, help='lora/lyco additional network arguments, default: %(default)s') + + group_other = parser.add_argument_group('Other') + group_other.add_argument('--overwrite', default = False, action='store_true', help = "overwrite existing training, default: %(default)s") + group_other.add_argument('--debug', default = False, action='store_true', help = "enable debug level logging, default: %(default)s") + + args = parser.parse_args() + + +def prepare_server(): + try: + server_status = util.Map(sdapi.progresssync()) + server_state = server_status['state'] + except: + log.error(f'server error: {server_status}') + exit(1) + if server_state['job_count'] > 0: + log.error(f'server not idle: {server_state}') + exit(1) + + server_options = util.Map(sdapi.options()) + server_options.options.save_training_settings_to_txt = False + server_options.options.training_enable_tensorboard = False + server_options.options.training_tensorboard_save_images = False + server_options.options.pin_memory = True + server_options.options.save_optimizer_state = False + server_options.options.training_image_repeats_per_epoch = args.repeats + server_options.options.training_write_csv_every = 0 + sdapi.postsync('/sdapi/v1/options', server_options.options) + log.info('updated server options') + + +def verify_args(): + server_options = util.Map(sdapi.options()) + if args.model != '': + if not os.path.isfile(args.model): + log.error(f'cannot find loaded model: {args.model}') + exit(1) + server_options.options.sd_model_checkpoint = args.model + sdapi.postsync('/sdapi/v1/options', server_options.options) + else: + args.model = server_options.options.sd_model_checkpoint.split(' [')[0] + args.lora_dir = server_options.options.lora_dir + args.lyco_dir = server_options.options.lyco_dir + args.ckpt_dir = server_options.options.ckpt_dir + args.embeddings_dir = server_options.options.embeddings_dir + if not os.path.isfile(args.model): + attempt = os.path.abspath(os.path.join(args.ckpt_dir, args.model)) + args.model = attempt if os.path.isfile(attempt) else args.model + if not os.path.isfile(args.model): + attempt = os.path.abspath(os.path.join(args.ckpt_dir, '..', args.model)) + args.model = attempt if os.path.isfile(attempt) else args.model + if not os.path.isfile(args.model): + log.error(f'cannot find loaded model: {args.model}') + exit(1) + if not os.path.exists(args.ckpt_dir) or not os.path.isdir(args.ckpt_dir): + log.error(f'cannot find models folder: {args.ckpt_dir}') + exit(1) + if not os.path.exists(args.input) or not os.path.isdir(args.input): + log.error(f'cannot find training folder: {args.input}') + exit(1) + if not os.path.exists(args.lora_dir) or not os.path.isdir(args.lora_dir): + log.error(f'cannot find lora folder: {args.lora_dir}') + exit(1) + if not os.path.exists(args.lyco_dir) or not os.path.isdir(args.lyco_dir): + log.error(f'cannot find lyco folder: {args.lyco_dir}') + exit(1) + if args.output != '': + args.process_dir = args.output + else: + args.process_dir = os.path.join(tempfile.gettempdir(), 'train', args.name) + log.debug(f'args: {vars(args)}') + log.debug(f'server flags: {server_options.flags}') + log.debug(f'server options: {server_options.options}') + + +async def training_loop(): + async def async_train(): + res = await sdapi.post('/sdapi/v1/train/embedding', options.embedding) + log.info(f'train embedding result: {res}') + + async def async_monitor(): + await asyncio.sleep(3) + res = util.Map(sdapi.progress()) + with tqdm(desc='train embedding', total=res.state.job_count) as pbar: + while res.state.job_no < res.state.job_count and not res.state.interrupted and not res.state.skipped: + await asyncio.sleep(2) + prev_job = res.state.job_no + res = util.Map(sdapi.progress()) + loss = re.search(r"Loss: (.*?)(?=\<)", res.textinfo) + if loss: + pbar.set_postfix({ 'loss': loss.group(0) }) + pbar.update(res.state.job_no - prev_job) + + a = asyncio.create_task(async_train()) + b = asyncio.create_task(async_monitor()) + await asyncio.gather(a, b) # wait for both pipeline and monitor to finish + + +def train_embedding(): + log.info(f'{args.type} options: {options.embedding}') + create_options = util.Map({ + "name": args.name, + "num_vectors_per_token": args.dim, + "overwrite_old": False, + "init_text": args.tag, + }) + fn = os.path.join(args.embeddings_dir, args.name) + '.pt' + if os.path.exists(fn) and args.overwrite: + log.warning(f'delete existing embedding {fn}') + os.remove(fn) + else: + log.error(f'embedding exists {fn}') + return + log.info(f'create embedding {create_options}') + res = sdapi.postsync('/sdapi/v1/create/embedding', create_options) + if 'info' in res and 'error' in res['info']: # formatted error + log.error(res.info) + elif 'info' in res: # no error + asyncio.run(training_loop()) + else: # unknown error + log.error(f'create embedding error {res}') + + +def train_lora(): + fn = os.path.join(options.lora.output_dir, args.name) + for ext in ['.ckpt', '.pt', '.safetensors']: + if os.path.exists(fn + ext): + if args.overwrite: + log.warning(f'delete existing lora: {fn + ext}') + os.remove(fn + ext) + else: + log.error(f'lora exists: {fn + ext}') + return + log.info(f'{args.type} options: {options.lora}') + # lora imports + lora_path = os.path.abspath(os.path.join(os.path.dirname(__file__), os.pardir, 'modules', 'lora')) + lycoris_path = os.path.abspath(os.path.join(os.path.dirname(__file__), os.pardir, 'modules', 'lycoris')) + sys.path.append(lora_path) + if args.type == 'lyco': + sys.path.append(lycoris_path) + log.debug('importing lora lib') + import train_network + train_network.train(options.lora) + if args.type == 'lyco': + log.debug('importing lycoris lib') + import importlib + _network_module = importlib.import_module(options.lora.network_module) + + +def prepare_options(): + if args.type == 'embedding': + log.info('train embedding') + options.lora.in_json = None + if args.type == 'dreambooth': + log.info('train using dreambooth style training') + options.lora.in_json = None + if args.type == 'lora': + log.info('train using lora style training') + options.lora.output_dir = args.lora_dir + options.lora.in_json = os.path.join(args.process_dir, args.name + '.json') + if args.type == 'lyco': + log.info('train using lycoris network') + options.lora.output_dir = args.lyco_dir + options.lora.network_module = 'lycoris.kohya' + options.lora.in_json = os.path.join(args.process_dir, args.name + '.json') + # lora specific + options.lora.pretrained_model_name_or_path = args.model + options.lora.output_name = args.name + options.lora.max_train_steps = args.steps + options.lora.network_dim = args.dim + options.lora.network_alpha = args.dim // 2 if args.alpha == 0 else args.alpha + options.lora.netwoork_args = [] + if args.algo is not None: + options.lora.netwoork_args.append(f'algo={args.algo}') + if args.args is not None: + for net_arg in args.args: + options.lora.netwoork_args.append(net_arg) + options.lora.gradient_accumulation_steps = args.gradient + options.lora.learning_rate = args.lr + options.lora.train_batch_size = args.batch + options.lora.train_data_dir = args.process_dir + # embedding specific + options.embedding.embedding_name = args.name + options.embedding.learn_rate = str(args.lr) + options.embedding.batch_size = args.batch + options.embedding.steps = args.steps + options.embedding.data_root = args.process_dir + options.embedding.log_directory = os.path.join(args.process_dir, 'log') + options.embedding.gradient_step = args.gradient + + +def process_inputs(): + pathlib.Path(args.process_dir).mkdir(parents=True, exist_ok=True) + processing_options = args.process.split(',') if isinstance(args.process, str) else args.process + processing_options = [opt.strip() for opt in re.split(',| ', args.process)] + log.info(f'processing steps: {processing_options}') + for step in processing_options: + if step not in valid_steps: + log.error(f'invalid processing step: {[step]}') + exit(1) + for root, _sub_dirs, folder in os.walk(args.input): + files = [os.path.join(root, f) for f in folder if filetype.is_image(os.path.join(root, f))] + log.info(f'processing input images: {len(files)}') + if os.path.exists(args.process_dir): + if args.overwrite: + log.warning(f'removing existing processed folder: {args.process_dir}') + shutil.rmtree(args.process_dir, ignore_errors=True) + else: + log.info(f'processed folder exists: {args.process_dir}') + steps = [step for step in processing_options if step in ['face', 'body', 'original']] + process.reset() + metadata = {} + for step in steps: + if step == 'face': + opts = [step for step in processing_options if step not in ['body', 'original']] + if step == 'body': + opts = [step for step in processing_options if step not in ['face', 'original', 'upscale', 'restore']] # body does not perform upscale or restore + if step == 'original': + opts = [step for step in processing_options if step not in ['face', 'body', 'upscale', 'restore', 'blur', 'range', 'segment']] # original does not perform most steps + log.info(f'processing current step: {opts}') + tag = step + if tag == 'original' and args.tag is not None: + concept = args.tag.split(',')[0].strip() + else: + concept = step + if args.type in ['lora', 'lyco', 'dreambooth']: + folder = os.path.join(args.process_dir, str(args.repeats) + '_' + concept) # separate concepts per folder + if args.type in ['embedding']: + folder = os.path.join(args.process_dir) # everything into same folder + log.info(f'processing concept: {concept}') + log.info(f'processing output folder: {folder}') + pathlib.Path(folder).mkdir(parents=True, exist_ok=True) + results = {} + for f in files: + res = process.file(filename = f, folder = folder, tag = args.tag, requested = opts) + if res.image: # valid result + results[res.type] = results.get(res.type, 0) + 1 + results['total'] = results.get('total', 0) + 1 + rel_path = res.basename.replace(os.path.commonpath([res.basename, args.process_dir]), '') + if rel_path.startswith(os.path.sep): + rel_path = rel_path[1:] + metadata[rel_path] = { 'caption': res.caption, 'tags': ','.join(res.tags) } + if options.lora.in_json is None: + with open(res.output.replace(options.process.format, '.txt'), "w", encoding='utf-8') as outfile: + outfile.write(res.caption) + log.info(f"processing {'saved' if res.image is not None else 'skipped'}: {f} => {res.output} {res.ops} {res.message}") + folders = [os.path.join(args.process_dir, folder) for folder in os.listdir(args.process_dir) if os.path.isdir(os.path.join(args.process_dir, folder))] + log.info(f'input datasets {folders}') + if options.lora.in_json is not None: + with open(options.lora.in_json, "w", encoding='utf-8') as outfile: # write json at the end only + outfile.write(json.dumps(metadata, indent=2)) + for folder in folders: # create latents + latents.create_vae_latents(util.Map({ 'input': folder, 'json': options.lora.in_json })) + latents.unload_vae() + r = { 'inputs': len(files), 'outputs': results, 'metadata': options.lora.in_json } + log.info(f'processing steps result: {r}') + if args.gradient < 0: + log.info(f"setting gradient accumulation to number of images: {results['total']}") + options.lora.gradient_accumulation_steps = results['total'] + options.embedding.gradient_step = results['total'] + process.unload() + + +if __name__ == '__main__': + log.info('SD.Next train script') + parse_args() + setup_logging() + prepare_server() + verify_args() + prepare_options() + mem_stats() + process_inputs() + mem_stats() + try: + if args.type == 'embedding': + train_embedding() + if args.type == 'lora' or args.type == 'lyco' or args.type == 'dreambooth': + train_lora() + except KeyboardInterrupt as e: + log.error('interrupt requested') + sdapi.interrupt() + mem_stats() + log.info('done') diff --git a/cli/train/console.py b/cli/train/console.py deleted file mode 100644 index e69de29bb..000000000 diff --git a/cli/train/sdapi.py b/cli/train/sdapi.py deleted file mode 100644 index e85148e59..000000000 --- a/cli/train/sdapi.py +++ /dev/null @@ -1,111 +0,0 @@ -import sys -import json -import aiohttp -import asyncio -import requests -from util import Map - - -sd_url = "http://127.0.0.1:7860" # automatic1111 api url root -use_session = True -timeout = aiohttp.ClientTimeout(total = None, sock_connect = 10, sock_read = None) # default value is 5 minutes, we need longer for training -sess = None -quiet = False - - -async def result(req): - if req.status != 200: - if not use_session and sess is not None: - await sess.close() - return Map({ 'error': req.status, 'reason': req.reason, 'url': req.url }) - else: - json = await req.json() - if type(json) == list: - res = json - elif json is None: - res = {} - else: - res = Map(json) - return res - - -def resultsync(req: requests.Response): - if req.status_code != 200: - return Map({ 'error': req.status_code, 'reason': req.reason, 'url': req.url }) - else: - json = req.json() - if type(json) == list: - res = json - elif json is None: - res = {} - else: - res = Map(json) - return res - - -async def get(endpoint: str, json: dict = None): - global sess # pylint: disable=global-statement - sess = sess if sess is not None else await session() - async with sess.get(url = endpoint, json = json) as req: - res = await result(req) - return res - - -def getsync(endpoint: str, json: dict = None): - req = requests.get(f'{sd_url}{endpoint}', json = json) # pylint: disable=missing-timeout - res = resultsync(req) - return res - - -async def post(endpoint: str, json: dict = None): - global sess # pylint: disable=global-statement - # sess = sess if sess is not None else await session() - if sess and not sess.closed: - await sess.close() - sess = await session() - async with sess.post(url = endpoint, json = json) as req: - res = await result(req) - return res - - -def postsync(endpoint: str, json: dict = None): - req = requests.post(f'{sd_url}{endpoint}', json = json) # pylint: disable=missing-timeout - res = resultsync(req) - return res - - -def interrupt(): - res = getsync('/sdapi/v1/progress?skip_current_image=true') - if 'state' in res and res.state.job_count > 0: - res = postsync('/sdapi/v1/interrupt') - return res - else: - return { 'interrupt': 'idle' } - - -def progress(): - res = getsync('/sdapi/v1/progress?skip_current_image=true') - return res - - -def options(): - options = getsync('/sdapi/v1/options') - flags = getsync('/sdapi/v1/cmd-flags') - return { 'options': options, 'flags': flags } - - -def shutdown(): - postsync('/sdapi/v1/shutdown') - - -async def session(): - global sess # pylint: disable=global-statement - time = aiohttp.ClientTimeout(total = None, sock_connect = 10, sock_read = None) # default value is 5 minutes, we need longer for training - sess = aiohttp.ClientSession(timeout = time, base_url = sd_url) - return sess - - -async def close(): - if sess is not None: - await asyncio.sleep(0) - await sess.__aexit__(None, None, None) diff --git a/cli/train/train.py b/cli/train/train.py deleted file mode 100755 index b1e1d48fb..000000000 --- a/cli/train/train.py +++ /dev/null @@ -1,322 +0,0 @@ -#!/bin/env python - -# system imports -import os -import re -import gc -import sys -import json -import shutil -import pathlib -import asyncio -import tempfile -import argparse - -# 3rd party imports -import filetype -from tqdm.rich import tqdm - -# local imports -import util -import sdapi -import process -import latents -import options - -# console handler -from rich import print # pylint: disable=redefined-builtin -from rich.pretty import install as pretty_install -from rich.traceback import install as traceback_install -from rich.console import Console -console = Console(log_time=True, log_time_format='%H:%M:%S-%f') -pretty_install(console=console) -import torch, accelerate, diffusers, requests, urllib3, http -traceback_install(console=console, extra_lines=1, width=console.width, word_wrap=False, indent_guides=False, suppress=[torch,accelerate,diffusers,asyncio,http,urllib3,requests]) - -# lora imports -lora_path = os.path.abspath(os.path.join(os.path.dirname(__file__), os.pardir, 'modules', 'lora')) -sys.path.append(lora_path) -lycoris_path = os.path.abspath(os.path.join(os.path.dirname(__file__), os.pardir, 'modules', 'lycoris')) -sys.path.append(lycoris_path) -import train_network - - -# globals -args = None -valid_steps = ['original', 'face', 'body', 'blur', 'range', 'upscale', 'restore', 'interrogate', 'resize', 'square', 'segment'] - -# methods - -def mem_stats(): - gc.collect() - if torch.cuda.is_available(): - with torch.no_grad(): - torch.cuda.empty_cache() - with torch.cuda.device('cuda'): - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - mem = util.get_memory() - peak = { 'active': mem['gpu-active']['peak'], 'allocated': mem['gpu-allocated']['peak'], 'reserved': mem['gpu-reserved']['peak'] } - console.log(f"memory cpu: {mem.ram} gpu current: {mem.gpu} gpu peak: {peak}") - - -def parse_args(): - global args - parser = argparse.ArgumentParser(description = 'train lora') - # basic section - parser.add_argument('--output', '--name', type=str, default=None, required=True, help='output filename') - parser.add_argument('--type', type=str, choices=['embedding', 'lora', 'lycoris', 'dreambooth'], default=None, required=True, help='training type') - parser.add_argument('--tag', type=str, default='person', required=False, help='primary tag, default: %(default)s') - parser.add_argument('--process', type=str, default='original,interrogate,resize,square', required=False, help=f'list of possible processing steps: {valid_steps}, default: %(default)s') - parser.add_argument('--dir', type=str, default='', required=False, help='where to store processed images, default is system temp/train') - parser.add_argument('--input', '--dataset', type=str, default=None, required=True, help='input folder with training images') - parser.add_argument('--overwrite', default = False, action='store_true', help = "overwrite existing training, default: %(default)s") - - # global params - parser.add_argument('--gradient', type=int, default=1, required=False, help='gradient accumulation steps, default: %(default)s') - parser.add_argument('--steps', type=int, default=2500, required=False, help='training steps, default: %(default)s') - parser.add_argument('--batch', type=int, default=1, required=False, help='batch size, default: %(default)s') - parser.add_argument('--lr', type=float, default=1e-04, required=False, help='model learning rate, default: %(default)s') - parser.add_argument('--dim', '--vectors', type=int, default=40, required=False, help='network dimension, default: %(default)s') - - # lora params - parser.add_argument('--repeats', type=int, default=10, required=False, help='number of repeats per image, default: %(default)s') - parser.add_argument('--alpha', type=float, default=0, required=False, help='alpha for weights scaling, default: half of dim') - - args = parser.parse_args() - - -def prepare_server(): - try: - server_status = util.Map(sdapi.progress()) - server_state = server_status['state'] - except: - console.log('server error:', server_status) - exit(1) - if server_state['job_count'] > 0: - console.log('server not idle:', server_state) - exit(1) - - server_options = util.Map(sdapi.options()) - server_options.options.save_training_settings_to_txt = False - server_options.options.training_enable_tensorboard = False - server_options.options.training_tensorboard_save_images = False - server_options.options.pin_memory = True - server_options.options.save_optimizer_state = False - server_options.options.training_image_repeats_per_epoch = args.repeats - server_options.options.training_write_csv_every = 0 - sdapi.postsync('/sdapi/v1/options', server_options.options) - console.log(f'updated server options') - - -def verify_args(): - global args - server_options = util.Map(sdapi.options()) - args.model = server_options.options['sd_model_checkpoint'].split(' [')[0] - args.lora_dir = server_options.flags['lora_dir'] - if not os.path.isabs(args.model) and not os.path.exists(args.model): - args.model = os.path.abspath(os.path.join(args.lora_dir, os.pardir, 'Stable-diffusion', args.model)) - - if not os.path.exists(args.model) or not os.path.isfile(args.model): - console.log('cannot find model:', args.model) - exit(1) - if not os.path.exists(args.input) or not os.path.isdir(args.input): - console.log('cannot find training folder:', args.input) - exit(1) - if not os.path.exists(args.lora_dir) or not os.path.isdir(args.lora_dir): - console.log('cannot find lora folder:', args.dir) - exit(1) - if args.dir != '': - args.process_dir = args.dir - else: - args.process_dir = os.path.join(tempfile.gettempdir(), 'train', args.output) - console.log(f'args: {vars(args)}') - - -async def training_loop(): - async def async_train(): - res = await sdapi.post('/sdapi/v1/train/embedding', options.embedding) - console.log(f'train embedding result: {res}') - - async def async_monitor(): - await asyncio.sleep(3) - res = util.Map(sdapi.progress()) - with tqdm(desc='train embedding', total=res.state.job_count) as pbar: - while res.state.job_no < res.state.job_count and not res.state.interrupted and not res.state.skipped: - await asyncio.sleep(2) - prev_job = res.state.job_no - res = util.Map(sdapi.progress()) - loss = re.search(r"Loss: (.*?)(?=\<)", res.textinfo) - if loss: - pbar.set_postfix({ 'loss': loss.group(0) }) - pbar.update(res.state.job_no - prev_job) - - a = asyncio.create_task(async_train()) - b = asyncio.create_task(async_monitor()) - await asyncio.gather(a, b) # wait for both pipeline and monitor to finish - - -def train_embedding(): - console.log(f'{args.type} options: {options.embedding}') - create_options = util.Map({ - "name": args.output, - "num_vectors_per_token": args.dim, - "overwrite_old": False, - "init_text": args.tag, - }) - server_options = util.Map(sdapi.options()) - fn = os.path.join(server_options.flags.embeddings_dir, args.output) + '.pt' - if os.path.exists(fn) and args.overwrite: - console.log(f'delete existing embedding {fn}') - os.remove(fn) - else: - console.log(f'embedding exists {fn}') - return - console.log(f'create embedding {create_options}') - res = sdapi.postsync('/sdapi/v1/create/embedding', create_options) - if 'info' in res and 'error' in res['info']: # formatted error - console.log(res.info) - elif 'info' in res: # no error - asyncio.run(training_loop()) - else: # unknown error - console.log(f'create embedding error {res}') - - -def train_lora(): - fn = os.path.join(args.lora_dir, args.output) - for ext in ['.ckpt', '.pt', '.safetensors']: - if os.path.exists(fn + ext): - if args.overwrite: - console.log(f'delete existing lora: {fn + ext}') - os.remove(fn + ext) - else: - console.log(f'lora exists: {fn + ext}') - return - console.log(f'{args.type} options: {options.lora}') - train_network.train(options.lora) - - -def prepare_options(): - # lora specific - options.lora.pretrained_model_name_or_path = args.model - options.lora.output_dir = args.lora_dir - options.lora.output_name = args.output - options.lora.max_train_steps = args.steps - options.lora.network_dim = args.dim - options.lora.network_alpha = args.dim // 2 if args.alpha == 0 else args.alpha - options.lora.gradient_accumulation_steps = args.gradient - options.lora.learning_rate = args.lr - options.lora.train_batch_size = args.batch - options.lora.network_alpha = args.dim // 2 if args.alpha == 0 else args.alpha - options.lora.train_data_dir = args.process_dir - if args.type == 'lycoris': - console.log('train using lycoris network') - options.lora.network_module = 'lycoris.kohya' - options.lora.in_json = os.path.join(args.process_dir, args.output + '.json') - if args.type == 'dreambooth': - console.log('train using dreambooth style training') - options.lora.in_json = None - if args.type == 'lora': - console.log('train using lora style training') - options.lora.in_json = os.path.join(args.process_dir, args.output + '.json') - if args.type == 'embedding': - console.log('train embedding') - options.lora.in_json = None - pass - # embedding specific - options.embedding.embedding_name = args.output - options.embedding.learn_rate = str(args.lr) - options.embedding.batch_size = args.batch - options.embedding.steps = args.steps - options.embedding.data_root = args.process_dir - options.embedding.log_directory = os.path.join(args.process_dir, 'log') - options.embedding.gradient_step = args.gradient - - -def process_inputs(): - pathlib.Path(args.process_dir).mkdir(parents=True, exist_ok=True) - processing_options = args.process.split(',') if isinstance(args.process, str) else args.process - processing_options = [opt.strip() for opt in re.split(',| ', args.process)] - console.log(f'processing steps: {processing_options}') - for step in processing_options: - if step not in valid_steps: - console.log(f'invalid processing step: {[step]}') - exit(1) - for root, _sub_dirs, folder in os.walk(args.input): - files = [os.path.join(root, f) for f in folder if filetype.is_image(os.path.join(root, f))] - console.log(f'processing input images: {len(files)}') - if os.path.exists(args.process_dir): - console.log('removing existing processed folder:', args.process_dir) - shutil.rmtree(args.process_dir, ignore_errors=True) - steps = [step for step in processing_options if step in ['face', 'body', 'original']] - process.reset() - metadata = {} - for step in steps: - if step == 'face': - opts = [step for step in processing_options if step not in ['body', 'original']] - if step == 'body': - opts = [step for step in processing_options if step not in ['face', 'original', 'upscale', 'restore']] # body does not perform upscale or restore - if step == 'original': - opts = [step for step in processing_options if step not in ['face', 'body', 'upscale', 'restore', 'blur', 'range', 'segment']] # original does not perform most steps - console.log(f'processing current step: {opts}') - tag = step - if tag == 'original' and args.tag is not None: - concept = args.tag.split(',')[0].strip() - else: - concept = step - if args.type in ['lora', 'lycoris', 'dreambooth']: - dir = os.path.join(args.process_dir, str(args.repeats) + '_' + concept) # separate concepts per folder - if args.type in ['embedding']: - dir = os.path.join(args.process_dir) # everything into same folder - console.log('processing concept:', concept) - console.log('processing output folder:', dir) - pathlib.Path(dir).mkdir(parents=True, exist_ok=True) - results = {} - for f in files: - res = process.file(filename = f, folder = dir, tag = args.tag, requested = opts) - if res.image: # valid result - results[res.type] = results.get(res.type, 0) + 1 - results['total'] = results.get('total', 0) + 1 - rel_path = res.basename.replace(os.path.commonpath([res.basename, args.process_dir]), '') - if rel_path.startswith(os.path.sep): rel_path = rel_path[1:] - metadata[rel_path] = { 'caption': res.caption, 'tags': ','.join(res.tags) } - if options.lora.in_json is None: - with open(res.output.replace(options.process.format, '.txt'), "w") as outfile: - outfile.write(res.caption) - console.log(f"processing {'saved' if res.image is not None else 'skipped'}: {f} => {res.output} {res.ops} {res.message}") - dirs = [os.path.join(args.process_dir, dir) for dir in os.listdir(args.process_dir) if os.path.isdir(os.path.join(args.process_dir, dir))] - console.log(f'input datasets {dirs}') - if options.lora.in_json is not None: - with open(options.lora.in_json, "w") as outfile: # write json at the end only - outfile.write(json.dumps(metadata, indent=2)) - for dir in dirs: # create latents - latents.create_vae_latents(util.Map({ 'input': dir, 'json': options.lora.in_json })) - latents.unload_vae() - r = { 'inputs': len(files), 'outputs': results, 'metadata': options.lora.in_json } - console.log(f'processing steps result: {r}') - if args.gradient < 0: - console.log(f"setting gradient accumulation to number of images: {results['total']}") - options.lora.gradient_accumulation_steps = results['total'] - options.embedding.gradient_step = results['total'] - process.unload() - - -if __name__ == '__main__': - console.log('train script for stable diffusion') - parse_args() - prepare_server() - verify_args() - prepare_options() - mem_stats() - process_inputs() - mem_stats() - try: - if args.type == 'embedding': - train_embedding() - if args.type == 'lora' or args.type == 'lycoris' or args.type == 'dreambooth': - train_lora() - except KeyboardInterrupt as e: - console.log('interrupt requested') - sdapi.interrupt() - mem_stats() - console.log('done') diff --git a/cli/train/util.py b/cli/train/util.py deleted file mode 100644 index 8c9aeb6c9..000000000 --- a/cli/train/util.py +++ /dev/null @@ -1,85 +0,0 @@ -#!/bin/env python -import os - -import transformers -transformers.logging.set_verbosity_error() - - -def get_memory(): - def gb(val: float): - return round(val / 1024 / 1024 / 1024, 2) - mem = {} - try: - import psutil - process = psutil.Process(os.getpid()) - res = process.memory_info() - ram_total = 100 * res.rss / process.memory_percent() - ram = { 'free': gb(ram_total - res.rss), 'used': gb(res.rss), 'total': gb(ram_total) } - mem.update({ 'ram': ram }) - except Exception as e: - mem.update({ 'ram': e }) - try: - import torch - if torch.cuda.is_available(): - s = torch.cuda.mem_get_info() - gpu = { 'free': gb(s[0]), 'used': gb(s[1] - s[0]), 'total': gb(s[1]) } - s = dict(torch.cuda.memory_stats('cuda')) - allocated = { 'current': gb(s['allocated_bytes.all.current']), 'peak': gb(s['allocated_bytes.all.peak']) } - reserved = { 'current': gb(s['reserved_bytes.all.current']), 'peak': gb(s['reserved_bytes.all.peak']) } - active = { 'current': gb(s['active_bytes.all.current']), 'peak': gb(s['active_bytes.all.peak']) } - inactive = { 'current': gb(s['inactive_split_bytes.all.current']), 'peak': gb(s['inactive_split_bytes.all.peak']) } - warnings = { 'retries': s['num_alloc_retries'], 'oom': s['num_ooms'] } - mem.update({ - 'gpu': gpu, - 'gpu-active': active, - 'gpu-allocated': allocated, - 'gpu-reserved': reserved, - 'gpu-inactive': inactive, - 'events': warnings, - }) - except: - pass - return Map(mem) - - -class Map(dict): - __slots__ = ('__dict__') - def __init__(self, *args, **kwargs): - super(Map, self).__init__(*args, **kwargs) - for arg in args: - if isinstance(arg, dict): - for k, v in arg.items(): - if isinstance(v, dict): - v = Map(v) - if isinstance(v, list): - self.__convert(v) - self[k] = v - if kwargs: - for k, v in kwargs.items(): - if isinstance(v, dict): - v = Map(v) - elif isinstance(v, list): - self.__convert(v) - self[k] = v - def __convert(self, v): - for elem in range(0, len(v)): # pylint: disable=consider-using-enumerate - if isinstance(v[elem], dict): - v[elem] = Map(v[elem]) - elif isinstance(v[elem], list): - self.__convert(v[elem]) - def __getattr__(self, attr): - return self.get(attr) - def __setattr__(self, key, value): - self.__setitem__(key, value) - def __setitem__(self, key, value): - super(Map, self).__setitem__(key, value) - self.__dict__.update({key: value}) - def __delattr__(self, item): - self.__delitem__(item) - def __delitem__(self, key): - super(Map, self).__delitem__(key) - del self.__dict__[key] - - -if __name__ == "__main__": - pass diff --git a/cli/modules/util.py b/cli/util.py similarity index 88% rename from cli/modules/util.py rename to cli/util.py index 479b77233..0fafe664c 100755 --- a/cli/modules/util.py +++ b/cli/util.py @@ -1,4 +1,4 @@ -#!/bin/env python +#!/usr/bin/env python """ generic helper methods """ @@ -6,9 +6,13 @@ generic helper methods import os import string import logging +import warnings log_format = '%(asctime)s %(levelname)s: %(message)s' logging.basicConfig(level = logging.INFO, format = log_format) +warnings.filterwarnings(action="ignore", category=DeprecationWarning) +warnings.filterwarnings(action="ignore", category=FutureWarning) +warnings.filterwarnings(action="ignore", category=UserWarning) log = logging.getLogger("sd") @@ -52,22 +56,22 @@ def get_memory(): reserved = { 'current': gb(s['reserved_bytes.all.current']), 'peak': gb(s['reserved_bytes.all.peak']) } active = { 'current': gb(s['active_bytes.all.current']), 'peak': gb(s['active_bytes.all.peak']) } inactive = { 'current': gb(s['inactive_split_bytes.all.current']), 'peak': gb(s['inactive_split_bytes.all.peak']) } - warnings = { 'retries': s['num_alloc_retries'], 'oom': s['num_ooms'] } + events = { 'retries': s['num_alloc_retries'], 'oom': s['num_ooms'] } mem.update({ 'gpu': gpu, 'gpu-active': active, 'gpu-allocated': allocated, 'gpu-reserved': reserved, 'gpu-inactive': inactive, - 'events': warnings, + 'events': events, }) except: pass return Map(mem) -class Map(dict): - __slots__ = ('__dict__') +class Map(dict): # pylint: disable=C0205 + __slots__ = ('__dict__') # pylint: disable=C0325 def __init__(self, *args, **kwargs): super(Map, self).__init__(*args, **kwargs) for arg in args: diff --git a/cli/modules/video-extract.py b/cli/video-extract.py similarity index 96% rename from cli/modules/video-extract.py rename to cli/video-extract.py index 4a68c7440..9bc4544e6 100755 --- a/cli/modules/video-extract.py +++ b/cli/video-extract.py @@ -1,4 +1,4 @@ -#!/bin/env python +#!/usr/bin/env python """ use ffmpeg for animation processing """ @@ -16,8 +16,8 @@ def probe(src: str): result = subprocess.run(cmd, shell = True, capture_output = True, text = True, check = True) data = json.loads(result.stdout) stream = [x for x in data['streams'] if x["codec_type"] == "video"][0] - format = data['format'] if 'format' in data else {} - res = {**stream, **format} + fmt = data['format'] if 'format' in data else {} + res = {**stream, **fmt} video = Map({ 'codec': res.get('codec_name', 'unknown') + '/' + res.get('codec_tag_string', ''), 'resolution': [int(res.get('width', 0)), int(res.get('height', 0))], diff --git a/cli/xformers.sh b/cli/xformers.sh deleted file mode 100755 index c7073999f..000000000 --- a/cli/xformers.sh +++ /dev/null @@ -1,12 +0,0 @@ -#/bin/env bash -echo "Installing xformers" - -NVCC_FLAGS="--use_fast_math" -FORCE_CUDA="1" -TORCH_CUDA_ARCH_LIST="8.6" -pip install ninja -q -pip uninstall xformers -y 2>/dev/null -pip install -v -U git+https://github.com/facebookresearch/xformers.git@main#egg=xformers -pip show torch -pip show xformers -python -m xformers.info diff --git a/extensions-builtin/LDSR/ldsr_model_arch.py b/extensions-builtin/LDSR/ldsr_model_arch.py index c776fc720..7f450086f 100644 --- a/extensions-builtin/LDSR/ldsr_model_arch.py +++ b/extensions-builtin/LDSR/ldsr_model_arch.py @@ -40,7 +40,7 @@ class LDSR: model = model.to(shared.device) if half_attention: model = model.half() - if shared.opts.opt_channelslast: + if shared.cmd_opts.opt_channelslast: model = model.to(memory_format=torch.channels_last) sd_hijack.model_hijack.hijack(model) # apply optimization @@ -88,7 +88,7 @@ class LDSR: x_t = None logs = None - for n in range(n_runs): + for _ in range(n_runs): if custom_shape is not None: x_t = torch.randn(1, custom_shape[1], custom_shape[2], custom_shape[3]).to(model.device) x_t = repeat(x_t, '1 c h w -> b c h w', b=custom_shape[0]) @@ -110,7 +110,6 @@ class LDSR: diffusion_steps = int(steps) eta = 1.0 - down_sample_method = 'Lanczos' gc.collect() if torch.cuda.is_available: @@ -131,11 +130,11 @@ class LDSR: im_og = im_og.resize((width_downsampled_pre, height_downsampled_pre), Image.LANCZOS) else: print(f"Down sample rate is 1 from {target_scale} / 4 (Not downsampling)") - + # pad width and height to multiples of 64, pads with the edge values of image to avoid artifacts pad_w, pad_h = np.max(((2, 2), np.ceil(np.array(im_og.size) / 64).astype(int)), axis=0) * 64 - im_og.size im_padded = Image.fromarray(np.pad(np.array(im_og), ((0, pad_h), (0, pad_w), (0, 0)), mode='edge')) - + logs = self.run(model["model"], im_padded, diffusion_steps, eta) sample = logs["sample"] @@ -158,7 +157,7 @@ class LDSR: def get_cond(selected_path): - example = dict() + example = {} up_f = 4 c = selected_path.convert('RGB') c = torch.unsqueeze(torchvision.transforms.ToTensor()(c), 0) @@ -196,7 +195,7 @@ def convsample_ddim(model, cond, steps, shape, eta=1.0, callback=None, normals_s @torch.no_grad() def make_convolutional_sample(batch, model, custom_steps=None, eta=1.0, quantize_x0=False, custom_shape=None, temperature=1., noise_dropout=0., corrector=None, corrector_kwargs=None, x_T=None, ddim_use_x0_pred=False): - log = dict() + log = {} z, c, x, xrec, xc = model.get_input(batch, model.first_stage_key, return_first_stage_outputs=True, @@ -244,7 +243,7 @@ def make_convolutional_sample(batch, model, custom_steps=None, eta=1.0, quantize x_sample_noquant = model.decode_first_stage(sample, force_not_quantize=True) log["sample_noquant"] = x_sample_noquant log["sample_diff"] = torch.abs(x_sample_noquant - x_sample) - except: + except Exception: pass log["sample"] = x_sample diff --git a/extensions-builtin/LDSR/preload.py b/extensions-builtin/LDSR/preload.py new file mode 100644 index 000000000..cfd478d54 --- /dev/null +++ b/extensions-builtin/LDSR/preload.py @@ -0,0 +1,6 @@ +import os +from modules import paths + + +def preload(parser): + parser.add_argument("--ldsr-models-path", type=str, help="Path to directory with LDSR model file(s).", default=os.path.join(paths.models_path, 'LDSR')) diff --git a/extensions-builtin/LDSR/scripts/ldsr_model.py b/extensions-builtin/LDSR/scripts/ldsr_model.py index b8cff29b9..c4da79f31 100644 --- a/extensions-builtin/LDSR/scripts/ldsr_model.py +++ b/extensions-builtin/LDSR/scripts/ldsr_model.py @@ -7,7 +7,8 @@ from basicsr.utils.download_util import load_file_from_url from modules.upscaler import Upscaler, UpscalerData from ldsr_model_arch import LDSR from modules import shared, script_callbacks -import sd_hijack_autoencoder, sd_hijack_ddpm_v1 +import sd_hijack_autoencoder # noqa: F401 +import sd_hijack_ddpm_v1 # noqa: F401 class UpscalerLDSR(Upscaler): @@ -25,22 +26,28 @@ class UpscalerLDSR(Upscaler): yaml_path = os.path.join(self.model_path, "project.yaml") old_model_path = os.path.join(self.model_path, "model.pth") new_model_path = os.path.join(self.model_path, "model.ckpt") - safetensors_model_path = os.path.join(self.model_path, "model.safetensors") + + local_model_paths = self.find_models(ext_filter=[".ckpt", ".safetensors"]) + local_ckpt_path = next(iter([local_model for local_model in local_model_paths if local_model.endswith("model.ckpt")]), None) + local_safetensors_path = next(iter([local_model for local_model in local_model_paths if local_model.endswith("model.safetensors")]), None) + local_yaml_path = next(iter([local_model for local_model in local_model_paths if local_model.endswith("project.yaml")]), None) + if os.path.exists(yaml_path): statinfo = os.stat(yaml_path) if statinfo.st_size >= 10485760: print("Removing invalid LDSR YAML file.") os.remove(yaml_path) + if os.path.exists(old_model_path): print("Renaming model from model.pth to model.ckpt") os.rename(old_model_path, new_model_path) - if os.path.exists(safetensors_model_path): - model = safetensors_model_path + + if local_safetensors_path is not None and os.path.exists(local_safetensors_path): + model = local_safetensors_path else: - model = load_file_from_url(url=self.model_url, model_dir=self.model_path, - file_name="model.ckpt", progress=True) - yaml = load_file_from_url(url=self.yaml_url, model_dir=self.model_path, - file_name="project.yaml", progress=True) + model = local_ckpt_path if local_ckpt_path is not None else load_file_from_url(url=self.model_url, model_dir=self.model_download_path, file_name="model.ckpt", progress=True) + + yaml = local_yaml_path if local_yaml_path is not None else load_file_from_url(url=self.yaml_url, model_dir=self.model_download_path, file_name="project.yaml", progress=True) try: return LDSR(model, yaml) diff --git a/extensions-builtin/LDSR/sd_hijack_autoencoder.py b/extensions-builtin/LDSR/sd_hijack_autoencoder.py index 8e03c7f89..81c5101b7 100644 --- a/extensions-builtin/LDSR/sd_hijack_autoencoder.py +++ b/extensions-builtin/LDSR/sd_hijack_autoencoder.py @@ -1,16 +1,21 @@ # The content of this file comes from the ldm/models/autoencoder.py file of the compvis/stable-diffusion repo # The VQModel & VQModelInterface were subsequently removed from ldm/models/autoencoder.py when we moved to the stability-ai/stablediffusion repo # As the LDSR upscaler relies on VQModel & VQModelInterface, the hijack aims to put them back into the ldm.models.autoencoder - +import numpy as np import torch import pytorch_lightning as pl import torch.nn.functional as F from contextlib import contextmanager + +from torch.optim.lr_scheduler import LambdaLR + +from ldm.modules.ema import LitEma from taming.modules.vqvae.quantize import VectorQuantizer2 as VectorQuantizer from ldm.modules.diffusionmodules.model import Encoder, Decoder from ldm.util import instantiate_from_config import ldm.models.autoencoder +from packaging import version class VQModel(pl.LightningModule): def __init__(self, @@ -19,7 +24,7 @@ class VQModel(pl.LightningModule): n_embed, embed_dim, ckpt_path=None, - ignore_keys=[], + ignore_keys=None, image_key="image", colorize_nlabels=None, monitor=None, @@ -57,7 +62,7 @@ class VQModel(pl.LightningModule): print(f"Keeping EMAs of {len(list(self.model_ema.buffers()))}.") if ckpt_path is not None: - self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys) + self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys or []) self.scheduler_config = scheduler_config self.lr_g_factor = lr_g_factor @@ -76,11 +81,11 @@ class VQModel(pl.LightningModule): if context is not None: print(f"{context}: Restored training weights") - def init_from_ckpt(self, path, ignore_keys=list()): + def init_from_ckpt(self, path, ignore_keys=None): sd = torch.load(path, map_location="cpu")["state_dict"] keys = list(sd.keys()) for k in keys: - for ik in ignore_keys: + for ik in ignore_keys or []: if k.startswith(ik): print("Deleting key {} from state_dict.".format(k)) del sd[k] @@ -165,7 +170,7 @@ class VQModel(pl.LightningModule): def validation_step(self, batch, batch_idx): log_dict = self._validation_step(batch, batch_idx) with self.ema_scope(): - log_dict_ema = self._validation_step(batch, batch_idx, suffix="_ema") + self._validation_step(batch, batch_idx, suffix="_ema") return log_dict def _validation_step(self, batch, batch_idx, suffix=""): @@ -232,7 +237,7 @@ class VQModel(pl.LightningModule): return self.decoder.conv_out.weight def log_images(self, batch, only_inputs=False, plot_ema=False, **kwargs): - log = dict() + log = {} x = self.get_input(batch, self.image_key) x = x.to(self.device) if only_inputs: @@ -249,7 +254,8 @@ class VQModel(pl.LightningModule): if plot_ema: with self.ema_scope(): xrec_ema, _ = self(x) - if x.shape[1] > 3: xrec_ema = self.to_rgb(xrec_ema) + if x.shape[1] > 3: + xrec_ema = self.to_rgb(xrec_ema) log["reconstructions_ema"] = xrec_ema return log @@ -264,7 +270,7 @@ class VQModel(pl.LightningModule): class VQModelInterface(VQModel): def __init__(self, embed_dim, *args, **kwargs): - super().__init__(embed_dim=embed_dim, *args, **kwargs) + super().__init__(*args, embed_dim=embed_dim, **kwargs) self.embed_dim = embed_dim def encode(self, x): @@ -282,5 +288,5 @@ class VQModelInterface(VQModel): dec = self.decoder(quant) return dec -setattr(ldm.models.autoencoder, "VQModel", VQModel) -setattr(ldm.models.autoencoder, "VQModelInterface", VQModelInterface) +ldm.models.autoencoder.VQModel = VQModel +ldm.models.autoencoder.VQModelInterface = VQModelInterface diff --git a/extensions-builtin/LDSR/sd_hijack_ddpm_v1.py b/extensions-builtin/LDSR/sd_hijack_ddpm_v1.py index 5c0488e5f..631a08ef0 100644 --- a/extensions-builtin/LDSR/sd_hijack_ddpm_v1.py +++ b/extensions-builtin/LDSR/sd_hijack_ddpm_v1.py @@ -48,7 +48,7 @@ class DDPMV1(pl.LightningModule): beta_schedule="linear", loss_type="l2", ckpt_path=None, - ignore_keys=[], + ignore_keys=None, load_only_unet=False, monitor="val/loss", use_ema=True, @@ -100,7 +100,7 @@ class DDPMV1(pl.LightningModule): if monitor is not None: self.monitor = monitor if ckpt_path is not None: - self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys, only_model=load_only_unet) + self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys or [], only_model=load_only_unet) self.register_schedule(given_betas=given_betas, beta_schedule=beta_schedule, timesteps=timesteps, linear_start=linear_start, linear_end=linear_end, cosine_s=cosine_s) @@ -182,13 +182,13 @@ class DDPMV1(pl.LightningModule): if context is not None: print(f"{context}: Restored training weights") - def init_from_ckpt(self, path, ignore_keys=list(), only_model=False): + def init_from_ckpt(self, path, ignore_keys=None, only_model=False): sd = torch.load(path, map_location="cpu") if "state_dict" in list(sd.keys()): sd = sd["state_dict"] keys = list(sd.keys()) for k in keys: - for ik in ignore_keys: + for ik in ignore_keys or []: if k.startswith(ik): print("Deleting key {} from state_dict.".format(k)) del sd[k] @@ -375,7 +375,7 @@ class DDPMV1(pl.LightningModule): @torch.no_grad() def log_images(self, batch, N=8, n_row=2, sample=True, return_keys=None, **kwargs): - log = dict() + log = {} x = self.get_input(batch, self.first_stage_key) N = min(x.shape[0], N) n_row = min(x.shape[0], n_row) @@ -383,7 +383,7 @@ class DDPMV1(pl.LightningModule): log["inputs"] = x # get diffusion row - diffusion_row = list() + diffusion_row = [] x_start = x[:n_row] for t in range(self.num_timesteps): @@ -444,13 +444,13 @@ class LatentDiffusionV1(DDPMV1): conditioning_key = None ckpt_path = kwargs.pop("ckpt_path", None) ignore_keys = kwargs.pop("ignore_keys", []) - super().__init__(conditioning_key=conditioning_key, *args, **kwargs) + super().__init__(*args, conditioning_key=conditioning_key, **kwargs) self.concat_mode = concat_mode self.cond_stage_trainable = cond_stage_trainable self.cond_stage_key = cond_stage_key try: self.num_downs = len(first_stage_config.params.ddconfig.ch_mult) - 1 - except: + except Exception: self.num_downs = 0 if not scale_by_std: self.scale_factor = scale_factor @@ -460,7 +460,7 @@ class LatentDiffusionV1(DDPMV1): self.instantiate_cond_stage(cond_stage_config) self.cond_stage_forward = cond_stage_forward self.clip_denoised = False - self.bbox_tokenizer = None + self.bbox_tokenizer = None self.restarted_from_ckpt = False if ckpt_path is not None: @@ -792,7 +792,7 @@ class LatentDiffusionV1(DDPMV1): z = z.view((z.shape[0], -1, ks[0], ks[1], z.shape[-1])) # (bn, nc, ks[0], ks[1], L ) # 2. apply model loop over last dim - if isinstance(self.first_stage_model, VQModelInterface): + if isinstance(self.first_stage_model, VQModelInterface): output_list = [self.first_stage_model.decode(z[:, :, :, :, i], force_not_quantize=predict_cids or force_not_quantize) for i in range(z.shape[-1])] @@ -877,16 +877,6 @@ class LatentDiffusionV1(DDPMV1): c = self.q_sample(x_start=c, t=tc, noise=torch.randn_like(c.float())) return self.p_losses(x, c, t, *args, **kwargs) - def _rescale_annotations(self, bboxes, crop_coordinates): # TODO: move to dataset - def rescale_bbox(bbox): - x0 = clamp((bbox[0] - crop_coordinates[0]) / crop_coordinates[2]) - y0 = clamp((bbox[1] - crop_coordinates[1]) / crop_coordinates[3]) - w = min(bbox[2] / crop_coordinates[2], 1 - x0) - h = min(bbox[3] / crop_coordinates[3], 1 - y0) - return x0, y0, w, h - - return [rescale_bbox(b) for b in bboxes] - def apply_model(self, x_noisy, t, cond, return_ids=False): if isinstance(cond, dict): @@ -900,7 +890,7 @@ class LatentDiffusionV1(DDPMV1): if hasattr(self, "split_input_params"): assert len(cond) == 1 # todo can only deal with one conditioning atm - assert not return_ids + assert not return_ids ks = self.split_input_params["ks"] # eg. (128, 128) stride = self.split_input_params["stride"] # eg. (64, 64) @@ -1126,7 +1116,7 @@ class LatentDiffusionV1(DDPMV1): if cond is not None: if isinstance(cond, dict): cond = {key: cond[key][:batch_size] if not isinstance(cond[key], list) else - list(map(lambda x: x[:batch_size], cond[key])) for key in cond} + [x[:batch_size] for x in cond[key]] for key in cond} else: cond = [c[:batch_size] for c in cond] if isinstance(cond, list) else cond[:batch_size] @@ -1157,8 +1147,10 @@ class LatentDiffusionV1(DDPMV1): if i % log_every_t == 0 or i == timesteps - 1: intermediates.append(x0_partial) - if callback: callback(i) - if img_callback: img_callback(img, i) + if callback: + callback(i) + if img_callback: + img_callback(img, i) return img, intermediates @torch.no_grad() @@ -1205,8 +1197,10 @@ class LatentDiffusionV1(DDPMV1): if i % log_every_t == 0 or i == timesteps - 1: intermediates.append(img) - if callback: callback(i) - if img_callback: img_callback(img, i) + if callback: + callback(i) + if img_callback: + img_callback(img, i) if return_intermediates: return img, intermediates @@ -1221,7 +1215,7 @@ class LatentDiffusionV1(DDPMV1): if cond is not None: if isinstance(cond, dict): cond = {key: cond[key][:batch_size] if not isinstance(cond[key], list) else - list(map(lambda x: x[:batch_size], cond[key])) for key in cond} + [x[:batch_size] for x in cond[key]] for key in cond} else: cond = [c[:batch_size] for c in cond] if isinstance(cond, list) else cond[:batch_size] return self.p_sample_loop(cond, @@ -1253,7 +1247,7 @@ class LatentDiffusionV1(DDPMV1): use_ddim = ddim_steps is not None - log = dict() + log = {} z, c, x, xrec, xc = self.get_input(batch, self.first_stage_key, return_first_stage_outputs=True, force_c_encode=True, @@ -1280,7 +1274,7 @@ class LatentDiffusionV1(DDPMV1): if plot_diffusion_rows: # get diffusion row - diffusion_row = list() + diffusion_row = [] z_start = z[:n_row] for t in range(self.num_timesteps): if t % self.log_every_t == 0 or t == self.num_timesteps - 1: @@ -1322,7 +1316,7 @@ class LatentDiffusionV1(DDPMV1): if inpaint: # make a simple center square - b, h, w = z.shape[0], z.shape[2], z.shape[3] + h, w = z.shape[2], z.shape[3] mask = torch.ones(N, h, w).to(self.device) # zeros will be filled in mask[:, h // 4:3 * h // 4, w // 4:3 * w // 4] = 0. @@ -1424,10 +1418,10 @@ class Layout2ImgDiffusionV1(LatentDiffusionV1): # TODO: move all layout-specific hacks to this class def __init__(self, cond_stage_key, *args, **kwargs): assert cond_stage_key == 'coordinates_bbox', 'Layout2ImgDiffusion only for cond_stage_key="coordinates_bbox"' - super().__init__(cond_stage_key=cond_stage_key, *args, **kwargs) + super().__init__(*args, cond_stage_key=cond_stage_key, **kwargs) def log_images(self, batch, N=8, *args, **kwargs): - logs = super().log_images(batch=batch, N=N, *args, **kwargs) + logs = super().log_images(*args, batch=batch, N=N, **kwargs) key = 'train' if self.training else 'validation' dset = self.trainer.datamodule.datasets[key] @@ -1443,7 +1437,7 @@ class Layout2ImgDiffusionV1(LatentDiffusionV1): logs['bbox_image'] = cond_img return logs -setattr(ldm.models.diffusion.ddpm, "DDPMV1", DDPMV1) -setattr(ldm.models.diffusion.ddpm, "LatentDiffusionV1", LatentDiffusionV1) -setattr(ldm.models.diffusion.ddpm, "DiffusionWrapperV1", DiffusionWrapperV1) -setattr(ldm.models.diffusion.ddpm, "Layout2ImgDiffusionV1", Layout2ImgDiffusionV1) +ldm.models.diffusion.ddpm.DDPMV1 = DDPMV1 +ldm.models.diffusion.ddpm.LatentDiffusionV1 = LatentDiffusionV1 +ldm.models.diffusion.ddpm.DiffusionWrapperV1 = DiffusionWrapperV1 +ldm.models.diffusion.ddpm.Layout2ImgDiffusionV1 = Layout2ImgDiffusionV1 diff --git a/extensions-builtin/Lora/extra_networks_lora.py b/extensions-builtin/Lora/extra_networks_lora.py index 45f899fc4..b5fea4d2e 100644 --- a/extensions-builtin/Lora/extra_networks_lora.py +++ b/extensions-builtin/Lora/extra_networks_lora.py @@ -1,6 +1,7 @@ from modules import extra_networks, shared import lora + class ExtraNetworkLora(extra_networks.ExtraNetwork): def __init__(self): super().__init__('lora') @@ -22,5 +23,23 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): lora.load_loras(names, multipliers) + if shared.opts.lora_add_hashes_to_infotext: + lora_hashes = [] + for item in lora.loaded_loras: + shorthash = item.lora_on_disk.shorthash + if not shorthash: + continue + + alias = item.mentioned_name + if not alias: + continue + + alias = alias.replace(":", "").replace(",", "") + + lora_hashes.append(f"{alias}: {shorthash}") + + if lora_hashes: + p.extra_generation_params["Lora hashes"] = ", ".join(lora_hashes) + def deactivate(self, p): pass diff --git a/extensions-builtin/Lora/lora.py b/extensions-builtin/Lora/lora.py index 3cbd91646..eec147122 100644 --- a/extensions-builtin/Lora/lora.py +++ b/extensions-builtin/Lora/lora.py @@ -1,10 +1,9 @@ -import glob import os import re -from typing import Union import torch +from typing import Union -from modules import shared, devices, sd_models, errors +from modules import shared, devices, sd_models, errors, scripts, sd_hijack, hashes metadata_tags_order = {"ss_sd_model_name": 1, "ss_resolution": 2, "ss_clip_skip": 3, "ss_num_train_images": 10, "ss_tag_frequency": 20} @@ -77,9 +76,9 @@ class LoraOnDisk: self.name = name self.filename = filename self.metadata = {} + self.is_safetensors = os.path.splitext(filename)[1].lower() == ".safetensors" - _, ext = os.path.splitext(filename) - if ext.lower() == ".safetensors": + if self.is_safetensors: try: self.metadata = sd_models.read_metadata_from_safetensors(filename) except Exception as e: @@ -93,15 +92,45 @@ class LoraOnDisk: self.metadata = m self.ssmd_cover_images = self.metadata.pop('ssmd_cover_images', None) # those are cover images and they are too big to display in UI as text + self.alias = self.metadata.get('ss_output_name', self.name) + + self.hash = None + self.shorthash = None + self.set_hash( + self.metadata.get('sshs_model_hash') or + hashes.sha256_from_cache(self.filename, "lora/" + self.name, use_addnet_hash=self.is_safetensors) or + '' + ) + + def set_hash(self, v): + self.hash = v + self.shorthash = self.hash[0:12] + + if self.shorthash: + available_lora_hash_lookup[self.shorthash] = self + + def read_hash(self): + if not self.hash: + self.set_hash(hashes.sha256(self.filename, "lora/" + self.name, use_addnet_hash=self.is_safetensors) or '') + + def get_alias(self): + if shared.opts.lora_preferred_name == "Filename" or self.alias.lower() in forbidden_lora_aliases: + return self.name + else: + return self.alias class LoraModule: - def __init__(self, name): + def __init__(self, name, lora_on_disk: LoraOnDisk): self.name = name + self.lora_on_disk = lora_on_disk self.multiplier = 1.0 self.modules = {} self.mtime = None + self.mentioned_name = None + """the text that was used to add lora to prompt - can be either name or an alias""" + class LoraUpDownModule: def __init__(self): @@ -126,20 +155,23 @@ def assign_lora_names_to_compvis_modules(sd_model): sd_model.lora_layer_mapping = lora_layer_mapping -def load_lora(name, filename): - lora = LoraModule(name) - lora.mtime = os.path.getmtime(filename) +def load_lora(name, lora_on_disk): + lora = LoraModule(name, lora_on_disk) + lora.mtime = os.path.getmtime(lora_on_disk.filename) - sd = sd_models.read_state_dict(filename) + sd = sd_models.read_state_dict(lora_on_disk.filename) + + # this should not be needed but is here as an emergency fix for an unknown error people are experiencing in 1.2.0 + if not hasattr(shared.sd_model, 'lora_layer_mapping'): + assign_lora_names_to_compvis_modules(shared.sd_model) keys_failed_to_match = {} is_sd2 = 'model_transformer_resblocks' in shared.sd_model.lora_layer_mapping for key_diffusers, weight in sd.items(): - lora_key_parts = key_diffusers.split(".", 1) - key_diffusers_without_lora_parts = lora_key_parts[0] - lora_key = lora_key_parts[1] if len(lora_key_parts) > 1 else "" + key_diffusers_without_lora_parts, lora_key = key_diffusers.split(".", 1) key = convert_diffusers_name_to_compvis(key_diffusers_without_lora_parts, is_sd2) + sd_module = shared.sd_model.lora_layer_mapping.get(key, None) if sd_module is None: @@ -166,11 +198,14 @@ def load_lora(name, filename): module = torch.nn.Linear(weight.shape[1], weight.shape[0], bias=False) elif type(sd_module) == torch.nn.MultiheadAttention: module = torch.nn.Linear(weight.shape[1], weight.shape[0], bias=False) - elif type(sd_module) == torch.nn.Conv2d: - module = torch.nn.Conv2d(weight.shape[1], weight.shape[0], (weight.shape[2], weight.shape[3]), bias=False) + elif type(sd_module) == torch.nn.Conv2d and weight.shape[2:] == (1, 1): + module = torch.nn.Conv2d(weight.shape[1], weight.shape[0], (1, 1), bias=False) + elif type(sd_module) == torch.nn.Conv2d and weight.shape[2:] == (3, 3): + module = torch.nn.Conv2d(weight.shape[1], weight.shape[0], (3, 3), bias=False) else: print(f'Lora layer {key_diffusers} matched a layer with unsupported type: {type(sd_module).__name__}') continue + raise AssertionError(f"Lora layer {key_diffusers} matched a layer with unsupported type: {type(sd_module).__name__}") with torch.no_grad(): module.weight.copy_(weight) @@ -182,10 +217,10 @@ def load_lora(name, filename): elif lora_key == "lora_down.weight": lora_module.down = module else: - assert False, f'Bad Lora layer name: {key_diffusers} - must end in lora_up.weight, lora_down.weight or alpha' + raise AssertionError(f"Bad Lora layer name: {key_diffusers} - must end in lora_up.weight, lora_down.weight or alpha") if len(keys_failed_to_match) > 0: - print(f"Failed to match keys when loading Lora {filename}: {keys_failed_to_match}") + print(f"Failed to match keys when loading Lora {lora_on_disk.filename}: {keys_failed_to_match}") return lora @@ -199,27 +234,42 @@ def load_loras(names, multipliers=None): loaded_loras.clear() - loras_on_disk = [available_loras.get(name, None) for name in names] - if any([x is None for x in loras_on_disk]): + loras_on_disk = [available_lora_aliases.get(name, None) for name in names] + if any(x is None for x in loras_on_disk): list_available_loras() - loras_on_disk = [available_loras.get(name, None) for name in names] + loras_on_disk = [available_lora_aliases.get(name, None) for name in names] + + failed_to_load_loras = [] for i, name in enumerate(names): lora = already_loaded.get(name, None) lora_on_disk = loras_on_disk[i] + if lora_on_disk is not None: if lora is None or os.path.getmtime(lora_on_disk.filename) > lora.mtime: - lora = load_lora(name, lora_on_disk.filename) + try: + lora = load_lora(name, lora_on_disk) + except Exception as e: + errors.display(e, f"loading Lora {lora_on_disk.filename}") + continue + + lora.mentioned_name = name + + lora_on_disk.read_hash() if lora is None: + failed_to_load_loras.append(name) print(f"Couldn't find Lora with name {name}") continue lora.multiplier = multipliers[i] if multipliers else 1.0 loaded_loras.append(lora) + if len(failed_to_load_loras) > 0: + sd_hijack.model_hijack.comments.append("Failed to find Loras: " + ", ".join(failed_to_load_loras)) + def lora_calc_updown(lora, module, target): with torch.no_grad(): @@ -228,31 +278,29 @@ def lora_calc_updown(lora, module, target): if up.shape[2:] == (1, 1) and down.shape[2:] == (1, 1): updown = (up.squeeze(2).squeeze(2) @ down.squeeze(2).squeeze(2)).unsqueeze(2).unsqueeze(3) + elif up.shape[2:] == (3, 3) or down.shape[2:] == (3, 3): + updown = torch.nn.functional.conv2d(down.permute(1, 0, 2, 3), up).permute(1, 0, 2, 3) else: - permute, h, w = False, 1, 1 - if len(up.shape) == 4 and len(down.shape) == 4: - if up.shape[2:] == (1, 1): - up = up.squeeze(2).squeeze(2) - else: - n, c, h, w = up.shape - up = up.view(n, c, -1).permute(2, 0, 1) - permute = True - if down.shape[2:] == (1, 1): - down = down.squeeze(2).squeeze(2) - else: - n, c, h, w = down.shape - down = down.view(n, c, -1).permute(2, 0, 1) - permute = True updown = up @ down - if permute: - nh, nw = updown.shape[1:] - updown = updown.permute(1, 2, 0).view(nh, nw, h, w) updown = updown * lora.multiplier * (module.alpha / module.up.weight.shape[1] if module.alpha else 1.0) return updown +def lora_restore_weights_from_backup(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.MultiheadAttention]): + weights_backup = getattr(self, "lora_weights_backup", None) + + if weights_backup is None: + return + + if isinstance(self, torch.nn.MultiheadAttention): + self.in_proj_weight.copy_(weights_backup[0]) + self.out_proj.weight.copy_(weights_backup[1]) + else: + self.weight.copy_(weights_backup) + + def lora_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.MultiheadAttention]): """ Applies the currently selected set of Loras to the weights of torch layer self. @@ -277,12 +325,7 @@ def lora_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.Mu self.lora_weights_backup = weights_backup if current_names != wanted_names: - if weights_backup is not None: - if isinstance(self, torch.nn.MultiheadAttention): - self.in_proj_weight.copy_(weights_backup[0]) - self.out_proj.weight.copy_(weights_backup[1]) - else: - self.weight.copy_(weights_backup) + lora_restore_weights_from_backup(self) for lora in loaded_loras: module = lora.modules.get(lora_layer_name, None) @@ -310,18 +353,51 @@ def lora_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.Mu print(f'failed to calculate lora weights for layer {lora_layer_name}') - setattr(self, "lora_current_names", wanted_names) + self.lora_current_names = wanted_names + + +def lora_forward(module, input, original_forward): + """ + Old way of applying Lora by executing operations during layer's forward. + Stacking many loras this way results in big performance degradation. + """ + + if len(loaded_loras) == 0: + return original_forward(module, input) + + input = devices.cond_cast_unet(input) + + lora_restore_weights_from_backup(module) + lora_reset_cached_weight(module) + + res = original_forward(module, input) + + lora_layer_name = getattr(module, 'lora_layer_name', None) + for lora in loaded_loras: + module = lora.modules.get(lora_layer_name, None) + if module is None: + continue + + module.up.to(device=devices.device) + module.down.to(device=devices.device) + + res = res + module.up(module.down(input)) * lora.multiplier * (module.alpha / module.up.weight.shape[1] if module.alpha else 1.0) + + return res def lora_reset_cached_weight(self: Union[torch.nn.Conv2d, torch.nn.Linear]): - setattr(self, "lora_current_names", ()) - setattr(self, "lora_weights_backup", None) + self.lora_current_names = () + self.lora_weights_backup = None -def lora_Linear_forward(self, lora_input): +def lora_Linear_forward(self, input): + if shared.opts.lora_functional: + return lora_forward(self, input, torch.nn.Linear_forward_before_lora) + lora_apply_weights(self) - return torch.nn.Linear_forward_before_lora(self, lora_input) + return torch.nn.Linear_forward_before_lora(self, input) def lora_Linear_load_state_dict(self, *args, **kwargs): @@ -330,10 +406,13 @@ def lora_Linear_load_state_dict(self, *args, **kwargs): return torch.nn.Linear_load_state_dict_before_lora(self, *args, **kwargs) -def lora_Conv2d_forward(self, lora_input): +def lora_Conv2d_forward(self, input): + if shared.opts.lora_functional: + return lora_forward(self, input, torch.nn.Conv2d_forward_before_lora) + lora_apply_weights(self) - return torch.nn.Conv2d_forward_before_lora(self, lora_input) + return torch.nn.Conv2d_forward_before_lora(self, input) def lora_Conv2d_load_state_dict(self, *args, **kwargs): @@ -356,23 +435,68 @@ def lora_MultiheadAttention_load_state_dict(self, *args, **kwargs): def list_available_loras(): available_loras.clear() + available_lora_aliases.clear() + forbidden_lora_aliases.clear() + available_lora_hash_lookup.clear() + forbidden_lora_aliases.update({"none": 1, "Addams": 1}) + os.makedirs(shared.cmd_opts.lora_dir, exist_ok=True) - candidates = \ - glob.glob(os.path.join(shared.cmd_opts.lora_dir, '**/*.pt'), recursive=True) + \ - glob.glob(os.path.join(shared.cmd_opts.lora_dir, '**/*.safetensors'), recursive=True) + \ - glob.glob(os.path.join(shared.cmd_opts.lora_dir, '**/*.ckpt'), recursive=True) - + candidates = list(shared.walk_files(shared.cmd_opts.lora_dir, allowed_extensions=[".pt", ".ckpt", ".safetensors"])) for filename in sorted(candidates, key=str.lower): if os.path.isdir(filename): continue name = os.path.splitext(os.path.basename(filename))[0] + entry = LoraOnDisk(name, filename) - available_loras[name] = LoraOnDisk(name, filename) + available_loras[name] = entry + + if entry.alias in available_lora_aliases: + forbidden_lora_aliases[entry.alias.lower()] = 1 + + available_lora_aliases[name] = entry + available_lora_aliases[entry.alias] = entry + + +re_lora_name = re.compile(r"(.*)\s*\([0-9a-fA-F]+\)") + + +def infotext_pasted(infotext, params): + if "AddNet Module 1" in [x[1] for x in scripts.scripts_txt2img.infotext_fields]: + return # if the other extension is active, it will handle those fields, no need to do anything + + added = [] + + for k in params: + if not k.startswith("AddNet Model "): + continue + + num = k[13:] + + if params.get("AddNet Module " + num) != "LoRA": + continue + + name = params.get("AddNet Model " + num) + if name is None: + continue + + m = re_lora_name.match(name) + if m: + name = m.group(1) + + multiplier = params.get("AddNet Weight A " + num, "1.0") + + added.append(f"") + + if added: + params["Prompt"] += "\n" + "".join(added) available_loras = {} +available_lora_aliases = {} +available_lora_hash_lookup = {} +forbidden_lora_aliases = {} loaded_loras = [] list_available_loras() diff --git a/extensions-builtin/Lora/preload.py b/extensions-builtin/Lora/preload.py new file mode 100644 index 000000000..863dc5c0b --- /dev/null +++ b/extensions-builtin/Lora/preload.py @@ -0,0 +1,6 @@ +import os +from modules import paths + + +def preload(parser): + parser.add_argument("--lora-dir", type=str, help="Path to directory with Lora networks.", default=os.path.join(paths.models_path, 'Lora')) diff --git a/extensions-builtin/Lora/scripts/lora_script.py b/extensions-builtin/Lora/scripts/lora_script.py index 3fc38ab9d..cd61c5175 100644 --- a/extensions-builtin/Lora/scripts/lora_script.py +++ b/extensions-builtin/Lora/scripts/lora_script.py @@ -1,12 +1,14 @@ +import re + import torch import gradio as gr +from fastapi import FastAPI import lora import extra_networks_lora import ui_extra_networks_lora from modules import script_callbacks, ui_extra_networks, extra_networks, shared - def unload(): torch.nn.Linear.forward = torch.nn.Linear_forward_before_lora torch.nn.Linear._load_from_state_dict = torch.nn.Linear_load_state_dict_before_lora @@ -49,8 +51,61 @@ torch.nn.MultiheadAttention._load_from_state_dict = lora.lora_MultiheadAttention script_callbacks.on_model_loaded(lora.assign_lora_names_to_compvis_modules) script_callbacks.on_script_unloaded(unload) script_callbacks.on_before_ui(before_ui) +script_callbacks.on_infotext_pasted(lora.infotext_pasted) shared.options_templates.update(shared.options_section(('extra_networks', "Extra Networks"), { - "sd_lora": shared.OptionInfo("None", "Add Lora to prompt", gr.Dropdown, lambda: {"choices": ["None"] + [x for x in lora.available_loras]}, refresh=lora.list_available_loras), + "sd_lora": shared.OptionInfo("None", "Add Lora to prompt", gr.Dropdown, lambda: {"choices": ["None", *lora.available_loras]}, refresh=lora.list_available_loras), + "lora_preferred_name": shared.OptionInfo("Alias from file", "When adding to prompt, refer to Lora by", gr.Radio, {"choices": ["Alias from file", "Filename"]}), + "lora_add_hashes_to_infotext": shared.OptionInfo(True, "Add Lora hashes to infotext"), })) + + +def create_lora_json(obj: lora.LoraOnDisk): + return { + "name": obj.name, + "alias": obj.alias, + "path": obj.filename, + "metadata": obj.metadata, + } + + +def api_loras(_: gr.Blocks, app: FastAPI): + @app.get("/sdapi/v1/loras") + async def get_loras(): + return [create_lora_json(obj) for obj in lora.available_loras.values()] + + @app.post("/sdapi/v1/refresh-loras") + async def refresh_loras(): + return lora.list_available_loras() + + +script_callbacks.on_app_started(api_loras) + +re_lora = re.compile(""), + "prompt": json.dumps(f""), "local_preview": f"{path}.{shared.opts.samples_format}", "metadata": json.dumps(lora_on_disk.metadata, indent=4) if lora_on_disk.metadata else None, } def allowed_directories_for_previews(self): - return [shared.opts.lora_dir] + return [shared.cmd_opts.lora_dir] + diff --git a/extensions-builtin/ScuNET/preload.py b/extensions-builtin/ScuNET/preload.py new file mode 100644 index 000000000..4ce82b1d4 --- /dev/null +++ b/extensions-builtin/ScuNET/preload.py @@ -0,0 +1,6 @@ +import os +from modules import paths + + +def preload(parser): + parser.add_argument("--scunet-models-path", type=str, help="Path to directory with ScuNET model file(s).", default=os.path.join(paths.models_path, 'ScuNET')) diff --git a/extensions-builtin/ScuNET/scripts/scunet_model.py b/extensions-builtin/ScuNET/scripts/scunet_model.py index e0fbf3a33..45d9297b6 100644 --- a/extensions-builtin/ScuNET/scripts/scunet_model.py +++ b/extensions-builtin/ScuNET/scripts/scunet_model.py @@ -5,11 +5,14 @@ import traceback import PIL.Image import numpy as np import torch +from tqdm import tqdm + from basicsr.utils.download_util import load_file_from_url import modules.upscaler -from modules import devices, modelloader +from modules import devices, modelloader, script_callbacks from scunet_model_arch import SCUNet as net +from modules.shared import opts class UpscalerScuNET(modules.upscaler.Upscaler): @@ -42,34 +45,83 @@ class UpscalerScuNET(modules.upscaler.Upscaler): scalers.append(scaler_data2) self.scalers = scalers - def do_upscale(self, img: PIL.Image, selected_file): + @staticmethod + @torch.no_grad() + def tiled_inference(img, model): + # test the image tile by tile + h, w = img.shape[2:] + tile = opts.SCUNET_tile + tile_overlap = opts.SCUNET_tile_overlap + if tile == 0: + return model(img) + + device = devices.get_device_for('scunet') + assert tile % 8 == 0, "tile size should be a multiple of window_size" + sf = 1 + + stride = tile - tile_overlap + h_idx_list = list(range(0, h - tile, stride)) + [h - tile] + w_idx_list = list(range(0, w - tile, stride)) + [w - tile] + E = torch.zeros(1, 3, h * sf, w * sf, dtype=img.dtype, device=device) + W = torch.zeros_like(E, dtype=devices.dtype, device=device) + + with tqdm(total=len(h_idx_list) * len(w_idx_list), desc="ScuNET tiles") as pbar: + for h_idx in h_idx_list: + + for w_idx in w_idx_list: + + in_patch = img[..., h_idx: h_idx + tile, w_idx: w_idx + tile] + + out_patch = model(in_patch) + out_patch_mask = torch.ones_like(out_patch) + + E[ + ..., h_idx * sf: (h_idx + tile) * sf, w_idx * sf: (w_idx + tile) * sf + ].add_(out_patch) + W[ + ..., h_idx * sf: (h_idx + tile) * sf, w_idx * sf: (w_idx + tile) * sf + ].add_(out_patch_mask) + pbar.update(1) + output = E.div_(W) + + return output + + def do_upscale(self, img: PIL.Image.Image, selected_file): + torch.cuda.empty_cache() model = self.load_model(selected_file) if model is None: + print(f"ScuNET: Unable to load model from {selected_file}", file=sys.stderr) return img device = devices.get_device_for('scunet') - img = np.array(img) - img = img[:, :, ::-1] - img = np.moveaxis(img, 2, 0) / 255 - img = torch.from_numpy(img).float() - img = img.unsqueeze(0).to(device) + tile = opts.SCUNET_tile + h, w = img.height, img.width + np_img = np.array(img) + np_img = np_img[:, :, ::-1] # RGB to BGR + np_img = np_img.transpose((2, 0, 1)) / 255 # HWC to CHW + torch_img = torch.from_numpy(np_img).float().unsqueeze(0).to(device) # type: ignore - with torch.no_grad(): - output = model(img) - output = output.squeeze().float().cpu().clamp_(0, 1).numpy() - output = 255. * np.moveaxis(output, 0, 2) - output = output.astype(np.uint8) - output = output[:, :, ::-1] + if tile > h or tile > w: + _img = torch.zeros(1, 3, max(h, tile), max(w, tile), dtype=torch_img.dtype, device=torch_img.device) + _img[:, :, :h, :w] = torch_img # pad image + torch_img = _img + + torch_output = self.tiled_inference(torch_img, model).squeeze(0) + torch_output = torch_output[:, :h * 1, :w * 1] # remove padding, if any + np_output: np.ndarray = torch_output.float().cpu().clamp_(0, 1).numpy() + del torch_img, torch_output torch.cuda.empty_cache() - return PIL.Image.fromarray(output, 'RGB') + + output = np_output.transpose((1, 2, 0)) # CHW to HWC + output = output[:, :, ::-1] # BGR to RGB + return PIL.Image.fromarray((output * 255).astype(np.uint8)) def load_model(self, path: str): device = devices.get_device_for('scunet') if "http" in path: - filename = load_file_from_url(url=self.model_url, model_dir=self.model_path, file_name="%s.pth" % self.name, - progress=True) + filename = load_file_from_url(url=self.model_url, model_dir=self.model_download_path, file_name="%s.pth" % self.name, progress=True) else: filename = path if not os.path.exists(os.path.join(self.model_path, filename)) or filename is None: @@ -79,9 +131,19 @@ class UpscalerScuNET(modules.upscaler.Upscaler): model = net(in_nc=3, config=[4, 4, 4, 4, 4, 4, 4], dim=64) model.load_state_dict(torch.load(filename), strict=True) model.eval() - for k, v in model.named_parameters(): + for _, v in model.named_parameters(): v.requires_grad = False model = model.to(device) return model + +def on_ui_settings(): + import gradio as gr + from modules import shared + + shared.opts.add_option("SCUNET_tile", shared.OptionInfo(256, "Tile size for SCUNET upscalers.", gr.Slider, {"minimum": 0, "maximum": 512, "step": 16}, section=('upscaling', "Upscaling")).info("0 = no tiling")) + shared.opts.add_option("SCUNET_tile_overlap", shared.OptionInfo(8, "Tile overlap for SCUNET upscalers.", gr.Slider, {"minimum": 0, "maximum": 64, "step": 1}, section=('upscaling', "Upscaling")).info("Low values = visible seam")) + + +script_callbacks.on_ui_settings(on_ui_settings) diff --git a/extensions-builtin/ScuNET/scunet_model_arch.py b/extensions-builtin/ScuNET/scunet_model_arch.py index 43ca8d36f..b51a88062 100644 --- a/extensions-builtin/ScuNET/scunet_model_arch.py +++ b/extensions-builtin/ScuNET/scunet_model_arch.py @@ -61,7 +61,9 @@ class WMSA(nn.Module): Returns: output: tensor shape [b h w c] """ - if self.type != 'W': x = torch.roll(x, shifts=(-(self.window_size // 2), -(self.window_size // 2)), dims=(1, 2)) + if self.type != 'W': + x = torch.roll(x, shifts=(-(self.window_size // 2), -(self.window_size // 2)), dims=(1, 2)) + x = rearrange(x, 'b (w1 p1) (w2 p2) c -> b w1 w2 p1 p2 c', p1=self.window_size, p2=self.window_size) h_windows = x.size(1) w_windows = x.size(2) @@ -85,8 +87,9 @@ class WMSA(nn.Module): output = self.linear(output) output = rearrange(output, 'b (w1 w2) (p1 p2) c -> b (w1 p1) (w2 p2) c', w1=h_windows, p1=self.window_size) - if self.type != 'W': output = torch.roll(output, shifts=(self.window_size // 2, self.window_size // 2), - dims=(1, 2)) + if self.type != 'W': + output = torch.roll(output, shifts=(self.window_size // 2, self.window_size // 2), dims=(1, 2)) + return output def relative_embedding(self): @@ -262,4 +265,4 @@ class SCUNet(nn.Module): nn.init.constant_(m.bias, 0) elif isinstance(m, nn.LayerNorm): nn.init.constant_(m.bias, 0) - nn.init.constant_(m.weight, 1.0) \ No newline at end of file + nn.init.constant_(m.weight, 1.0) diff --git a/extensions-builtin/SwinIR/preload.py b/extensions-builtin/SwinIR/preload.py new file mode 100644 index 000000000..e912c6402 --- /dev/null +++ b/extensions-builtin/SwinIR/preload.py @@ -0,0 +1,6 @@ +import os +from modules import paths + + +def preload(parser): + parser.add_argument("--swinir-models-path", type=str, help="Path to directory with SwinIR model file(s).", default=os.path.join(paths.models_path, 'SwinIR')) diff --git a/extensions-builtin/SwinIR/scripts/swinir_model.py b/extensions-builtin/SwinIR/scripts/swinir_model.py index 86672cd9a..1c7bf325e 100644 --- a/extensions-builtin/SwinIR/scripts/swinir_model.py +++ b/extensions-builtin/SwinIR/scripts/swinir_model.py @@ -1,13 +1,13 @@ import os + import numpy as np import torch from PIL import Image from basicsr.utils.download_util import load_file_from_url from tqdm import tqdm -from rich import print, progress # pylint: disable=redefined-builtin from modules import modelloader, devices, script_callbacks, shared -from modules.shared import cmd_opts, opts, state +from modules.shared import opts, state from swinir_model_arch import SwinIR as net from swinir_model_arch_v2 import Swin2SR as net2 from modules.upscaler import Upscaler, UpscalerData @@ -44,31 +44,31 @@ class UpscalerSwinIR(Upscaler): img = upscale(img, model) try: torch.cuda.empty_cache() - except: + except Exception: pass return img def load_model(self, path, scale=4): if "http" in path: dl_name = "%s%s" % (self.model_name.replace(" ", "_"), ".pth") - filename = load_file_from_url(url=path, model_dir=self.model_path, file_name=dl_name, progress=True) + filename = load_file_from_url(url=path, model_dir=self.model_download_path, file_name=dl_name, progress=True) else: filename = path if filename is None or not os.path.exists(filename): return None if filename.endswith(".v2.pth"): model = net2( - upscale=scale, - in_chans=3, - img_size=64, - window_size=8, - img_range=1.0, - depths=[6, 6, 6, 6, 6, 6], - embed_dim=180, - num_heads=[6, 6, 6, 6, 6, 6], - mlp_ratio=2, - upsampler="nearest+conv", - resi_connection="1conv", + upscale=scale, + in_chans=3, + img_size=64, + window_size=8, + img_range=1.0, + depths=[6, 6, 6, 6, 6, 6], + embed_dim=180, + num_heads=[6, 6, 6, 6, 6, 6], + mlp_ratio=2, + upsampler="nearest+conv", + resi_connection="1conv", ) params = None else: @@ -87,9 +87,8 @@ class UpscalerSwinIR(Upscaler): ) params = "params_ema" - with progress.open(filename, 'rb', description=f'Loading weights: [cyan]{filename}', auto_refresh=True) as f: - pretrained_model = torch.load(filename) - if params is not None and params in pretrained_model: + pretrained_model = torch.load(filename) + if params is not None: model.load_state_dict(pretrained_model[params], strict=True) else: model.load_state_dict(pretrained_model, strict=True) @@ -151,7 +150,7 @@ def inference(img, model, tile, tile_overlap, window_size, scale): for w_idx in w_idx_list: if state.interrupted or state.skipped: break - + in_patch = img[..., h_idx: h_idx + tile, w_idx: w_idx + tile] out_patch = model(in_patch) out_patch_mask = torch.ones_like(out_patch) diff --git a/extensions-builtin/SwinIR/swinir_model_arch.py b/extensions-builtin/SwinIR/swinir_model_arch.py index 863f42db6..93b932747 100644 --- a/extensions-builtin/SwinIR/swinir_model_arch.py +++ b/extensions-builtin/SwinIR/swinir_model_arch.py @@ -644,7 +644,7 @@ class SwinIR(nn.Module): """ def __init__(self, img_size=64, patch_size=1, in_chans=3, - embed_dim=96, depths=[6, 6, 6, 6], num_heads=[6, 6, 6, 6], + embed_dim=96, depths=(6, 6, 6, 6), num_heads=(6, 6, 6, 6), window_size=7, mlp_ratio=4., qkv_bias=True, qk_scale=None, drop_rate=0., attn_drop_rate=0., drop_path_rate=0.1, norm_layer=nn.LayerNorm, ape=False, patch_norm=True, @@ -805,7 +805,7 @@ class SwinIR(nn.Module): def forward(self, x): H, W = x.shape[2:] x = self.check_image_size(x) - + self.mean = self.mean.type_as(x) x = (x - self.mean) * self.img_range @@ -844,7 +844,7 @@ class SwinIR(nn.Module): H, W = self.patches_resolution flops += H * W * 3 * self.embed_dim * 9 flops += self.patch_embed.flops() - for i, layer in enumerate(self.layers): + for layer in self.layers: flops += layer.flops() flops += H * W * 3 * self.embed_dim * self.embed_dim flops += self.upsample.flops() diff --git a/extensions-builtin/SwinIR/swinir_model_arch_v2.py b/extensions-builtin/SwinIR/swinir_model_arch_v2.py index 0e28ae6ee..dad22cca2 100644 --- a/extensions-builtin/SwinIR/swinir_model_arch_v2.py +++ b/extensions-builtin/SwinIR/swinir_model_arch_v2.py @@ -74,7 +74,7 @@ class WindowAttention(nn.Module): """ def __init__(self, dim, window_size, num_heads, qkv_bias=True, attn_drop=0., proj_drop=0., - pretrained_window_size=[0, 0]): + pretrained_window_size=(0, 0)): super().__init__() self.dim = dim @@ -241,7 +241,7 @@ class SwinTransformerBlock(nn.Module): attn_mask = None self.register_buffer("attn_mask", attn_mask) - + def calculate_mask(self, x_size): # calculate attention mask for SW-MSA H, W = x_size @@ -263,7 +263,7 @@ class SwinTransformerBlock(nn.Module): attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2) attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0)) - return attn_mask + return attn_mask def forward(self, x, x_size): H, W = x_size @@ -288,7 +288,7 @@ class SwinTransformerBlock(nn.Module): attn_windows = self.attn(x_windows, mask=self.attn_mask) # nW*B, window_size*window_size, C else: attn_windows = self.attn(x_windows, mask=self.calculate_mask(x_size).to(x.device)) - + # merge windows attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C) shifted_x = window_reverse(attn_windows, self.window_size, H, W) # B H' W' C @@ -369,7 +369,7 @@ class PatchMerging(nn.Module): H, W = self.input_resolution flops = (H // 2) * (W // 2) * 4 * self.dim * 2 * self.dim flops += H * W * self.dim // 2 - return flops + return flops class BasicLayer(nn.Module): """ A basic Swin Transformer layer for one stage. @@ -447,7 +447,7 @@ class BasicLayer(nn.Module): nn.init.constant_(blk.norm1.weight, 0) nn.init.constant_(blk.norm2.bias, 0) nn.init.constant_(blk.norm2.weight, 0) - + class PatchEmbed(nn.Module): r""" Image to Patch Embedding Args: @@ -492,7 +492,7 @@ class PatchEmbed(nn.Module): flops = Ho * Wo * self.embed_dim * self.in_chans * (self.patch_size[0] * self.patch_size[1]) if self.norm is not None: flops += Ho * Wo * self.embed_dim - return flops + return flops class RSTB(nn.Module): """Residual Swin Transformer Block (RSTB). @@ -531,7 +531,7 @@ class RSTB(nn.Module): num_heads=num_heads, window_size=window_size, mlp_ratio=mlp_ratio, - qkv_bias=qkv_bias, + qkv_bias=qkv_bias, drop=drop, attn_drop=attn_drop, drop_path=drop_path, norm_layer=norm_layer, @@ -622,7 +622,7 @@ class Upsample(nn.Sequential): else: raise ValueError(f'scale {scale} is not supported. ' 'Supported scales: 2^n and 3.') super(Upsample, self).__init__(*m) - + class Upsample_hf(nn.Sequential): """Upsample module. @@ -642,7 +642,7 @@ class Upsample_hf(nn.Sequential): m.append(nn.PixelShuffle(3)) else: raise ValueError(f'scale {scale} is not supported. ' 'Supported scales: 2^n and 3.') - super(Upsample_hf, self).__init__(*m) + super(Upsample_hf, self).__init__(*m) class UpsampleOneStep(nn.Sequential): @@ -667,8 +667,8 @@ class UpsampleOneStep(nn.Sequential): H, W = self.input_resolution flops = H * W * self.num_feat * 3 * 9 return flops - - + + class Swin2SR(nn.Module): r""" Swin2SR @@ -698,8 +698,8 @@ class Swin2SR(nn.Module): """ def __init__(self, img_size=64, patch_size=1, in_chans=3, - embed_dim=96, depths=[6, 6, 6, 6], num_heads=[6, 6, 6, 6], - window_size=7, mlp_ratio=4., qkv_bias=True, + embed_dim=96, depths=(6, 6, 6, 6), num_heads=(6, 6, 6, 6), + window_size=7, mlp_ratio=4., qkv_bias=True, drop_rate=0., attn_drop_rate=0., drop_path_rate=0.1, norm_layer=nn.LayerNorm, ape=False, patch_norm=True, use_checkpoint=False, upscale=2, img_range=1., upsampler='', resi_connection='1conv', @@ -764,7 +764,7 @@ class Swin2SR(nn.Module): num_heads=num_heads[i_layer], window_size=window_size, mlp_ratio=self.mlp_ratio, - qkv_bias=qkv_bias, + qkv_bias=qkv_bias, drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[sum(depths[:i_layer]):sum(depths[:i_layer + 1])], # no impact on SR results norm_layer=norm_layer, @@ -776,7 +776,7 @@ class Swin2SR(nn.Module): ) self.layers.append(layer) - + if self.upsampler == 'pixelshuffle_hf': self.layers_hf = nn.ModuleList() for i_layer in range(self.num_layers): @@ -787,7 +787,7 @@ class Swin2SR(nn.Module): num_heads=num_heads[i_layer], window_size=window_size, mlp_ratio=self.mlp_ratio, - qkv_bias=qkv_bias, + qkv_bias=qkv_bias, drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[sum(depths[:i_layer]):sum(depths[:i_layer + 1])], # no impact on SR results norm_layer=norm_layer, @@ -799,7 +799,7 @@ class Swin2SR(nn.Module): ) self.layers_hf.append(layer) - + self.norm = norm_layer(self.num_features) # build the last conv layer in deep feature extraction @@ -829,10 +829,10 @@ class Swin2SR(nn.Module): self.conv_aux = nn.Conv2d(num_feat, num_out_ch, 3, 1, 1) self.conv_after_aux = nn.Sequential( nn.Conv2d(3, num_feat, 3, 1, 1), - nn.LeakyReLU(inplace=True)) + nn.LeakyReLU(inplace=True)) self.upsample = Upsample(upscale, num_feat) self.conv_last = nn.Conv2d(num_feat, num_out_ch, 3, 1, 1) - + elif self.upsampler == 'pixelshuffle_hf': self.conv_before_upsample = nn.Sequential(nn.Conv2d(embed_dim, num_feat, 3, 1, 1), nn.LeakyReLU(inplace=True)) @@ -846,7 +846,7 @@ class Swin2SR(nn.Module): nn.Conv2d(embed_dim, num_feat, 3, 1, 1), nn.LeakyReLU(inplace=True)) self.conv_last_hf = nn.Conv2d(num_feat, num_out_ch, 3, 1, 1) - + elif self.upsampler == 'pixelshuffledirect': # for lightweight SR (to save parameters) self.upsample = UpsampleOneStep(upscale, embed_dim, num_out_ch, @@ -905,7 +905,7 @@ class Swin2SR(nn.Module): x = self.patch_unembed(x, x_size) return x - + def forward_features_hf(self, x): x_size = (x.shape[2], x.shape[3]) x = self.patch_embed(x) @@ -919,7 +919,7 @@ class Swin2SR(nn.Module): x = self.norm(x) # B L C x = self.patch_unembed(x, x_size) - return x + return x def forward(self, x): H, W = x.shape[2:] @@ -951,7 +951,7 @@ class Swin2SR(nn.Module): x = self.conv_after_body(self.forward_features(x)) + x x_before = self.conv_before_upsample(x) x_out = self.conv_last(self.upsample(x_before)) - + x_hf = self.conv_first_hf(x_before) x_hf = self.conv_after_body_hf(self.forward_features_hf(x_hf)) + x_hf x_hf = self.conv_before_upsample_hf(x_hf) @@ -977,15 +977,15 @@ class Swin2SR(nn.Module): x_first = self.conv_first(x) res = self.conv_after_body(self.forward_features(x_first)) + x_first x = x + self.conv_last(res) - + x = x / self.img_range + self.mean if self.upsampler == "pixelshuffle_aux": return x[:, :, :H*self.upscale, :W*self.upscale], aux - + elif self.upsampler == "pixelshuffle_hf": x_out = x_out / self.img_range + self.mean return x_out[:, :, :H*self.upscale, :W*self.upscale], x[:, :, :H*self.upscale, :W*self.upscale], x_hf[:, :, :H*self.upscale, :W*self.upscale] - + else: return x[:, :, :H*self.upscale, :W*self.upscale] @@ -994,7 +994,7 @@ class Swin2SR(nn.Module): H, W = self.patches_resolution flops += H * W * 3 * self.embed_dim * 9 flops += self.patch_embed.flops() - for i, layer in enumerate(self.layers): + for layer in self.layers: flops += layer.flops() flops += H * W * 3 * self.embed_dim * self.embed_dim flops += self.upsample.flops() @@ -1014,4 +1014,4 @@ if __name__ == '__main__': x = torch.randn((1, 3, height, width)) x = model(x) - print(x.shape) \ No newline at end of file + print(x.shape) diff --git a/extensions-builtin/a1111-sd-webui-lycoris b/extensions-builtin/a1111-sd-webui-lycoris index ce584a0ff..21e9ea0f5 160000 --- a/extensions-builtin/a1111-sd-webui-lycoris +++ b/extensions-builtin/a1111-sd-webui-lycoris @@ -1 +1 @@ -Subproject commit ce584a0ff863de98233ee135dcc17f2fb44703c3 +Subproject commit 21e9ea0f58dfdbee12b08389fd2cd1aa545e058d diff --git a/extensions-builtin/multidiffusion-upscaler-for-automatic1111 b/extensions-builtin/multidiffusion-upscaler-for-automatic1111 index 6931b89cb..70b3c5ea3 160000 --- a/extensions-builtin/multidiffusion-upscaler-for-automatic1111 +++ b/extensions-builtin/multidiffusion-upscaler-for-automatic1111 @@ -1 +1 @@ -Subproject commit 6931b89cb4507c7dc8fa81ac36c2c19d0691c44e +Subproject commit 70b3c5ea3c9f684d04e7ff59167565974415735c diff --git a/extensions-builtin/prompt-bracket-checker/javascript/prompt-bracket-checker.js b/extensions-builtin/prompt-bracket-checker/javascript/prompt-bracket-checker.js deleted file mode 100644 index f0918e260..000000000 --- a/extensions-builtin/prompt-bracket-checker/javascript/prompt-bracket-checker.js +++ /dev/null @@ -1,103 +0,0 @@ -// Stable Diffusion WebUI - Bracket checker -// Version 1.0 -// By Hingashi no Florin/Bwin4L -// Counts open and closed brackets (round, square, curly) in the prompt and negative prompt text boxes in the txt2img and img2img tabs. -// If there's a mismatch, the keyword counter turns red and if you hover on it, a tooltip tells you what's wrong. - -function checkBrackets(evt, textArea, counterElt) { - errorStringParen = '(...) - Different number of opening and closing parentheses detected.\n'; - errorStringSquare = '[...] - Different number of opening and closing square brackets detected.\n'; - errorStringCurly = '{...} - Different number of opening and closing curly brackets detected.\n'; - - openBracketRegExp = /\(/g; - closeBracketRegExp = /\)/g; - - openSquareBracketRegExp = /\[/g; - closeSquareBracketRegExp = /\]/g; - - openCurlyBracketRegExp = /\{/g; - closeCurlyBracketRegExp = /\}/g; - - totalOpenBracketMatches = 0; - totalCloseBracketMatches = 0; - totalOpenSquareBracketMatches = 0; - totalCloseSquareBracketMatches = 0; - totalOpenCurlyBracketMatches = 0; - totalCloseCurlyBracketMatches = 0; - - openBracketMatches = textArea.value.match(openBracketRegExp); - if(openBracketMatches) { - totalOpenBracketMatches = openBracketMatches.length; - } - - closeBracketMatches = textArea.value.match(closeBracketRegExp); - if(closeBracketMatches) { - totalCloseBracketMatches = closeBracketMatches.length; - } - - openSquareBracketMatches = textArea.value.match(openSquareBracketRegExp); - if(openSquareBracketMatches) { - totalOpenSquareBracketMatches = openSquareBracketMatches.length; - } - - closeSquareBracketMatches = textArea.value.match(closeSquareBracketRegExp); - if(closeSquareBracketMatches) { - totalCloseSquareBracketMatches = closeSquareBracketMatches.length; - } - - openCurlyBracketMatches = textArea.value.match(openCurlyBracketRegExp); - if(openCurlyBracketMatches) { - totalOpenCurlyBracketMatches = openCurlyBracketMatches.length; - } - - closeCurlyBracketMatches = textArea.value.match(closeCurlyBracketRegExp); - if(closeCurlyBracketMatches) { - totalCloseCurlyBracketMatches = closeCurlyBracketMatches.length; - } - - if(totalOpenBracketMatches != totalCloseBracketMatches) { - if(!counterElt.title.includes(errorStringParen)) { - counterElt.title += errorStringParen; - } - } else { - counterElt.title = counterElt.title.replace(errorStringParen, ''); - } - - if(totalOpenSquareBracketMatches != totalCloseSquareBracketMatches) { - if(!counterElt.title.includes(errorStringSquare)) { - counterElt.title += errorStringSquare; - } - } else { - counterElt.title = counterElt.title.replace(errorStringSquare, ''); - } - - if(totalOpenCurlyBracketMatches != totalCloseCurlyBracketMatches) { - if(!counterElt.title.includes(errorStringCurly)) { - counterElt.title += errorStringCurly; - } - } else { - counterElt.title = counterElt.title.replace(errorStringCurly, ''); - } - - if(counterElt.title != '') { - counterElt.classList.add('error'); - } else { - counterElt.classList.remove('error'); - } -} - -function setupBracketChecking(id_prompt, id_counter){ - var textarea = gradioApp().querySelector("#" + id_prompt + " > label > textarea"); - var counter = gradioApp().getElementById(id_counter) - - textarea.addEventListener("input", function(evt){ - checkBrackets(evt, textarea, counter) - }); -} - -onUiLoaded(function(){ - setupBracketChecking('txt2img_prompt', 'txt2img_token_counter') - setupBracketChecking('txt2img_neg_prompt', 'txt2img_negative_token_counter') - setupBracketChecking('img2img_prompt', 'img2img_token_counter') - setupBracketChecking('img2img_neg_prompt', 'img2img_negative_token_counter') -}) \ No newline at end of file diff --git a/extensions-builtin/sd-dynamic-thresholding b/extensions-builtin/sd-dynamic-thresholding index 5d3465c4b..f02cacfc9 160000 --- a/extensions-builtin/sd-dynamic-thresholding +++ b/extensions-builtin/sd-dynamic-thresholding @@ -1 +1 @@ -Subproject commit 5d3465c4b2d5dacfdd6caa38b71b23b4f88401c0 +Subproject commit f02cacfc923e8bbf73f25327d722d50c458d66bb diff --git a/extensions-builtin/sd-extension-aesthetic-scorer b/extensions-builtin/sd-extension-aesthetic-scorer deleted file mode 160000 index 896b6aa03..000000000 --- a/extensions-builtin/sd-extension-aesthetic-scorer +++ /dev/null @@ -1 +0,0 @@ -Subproject commit 896b6aa034127e8064ae5fdd6370f351e0755475 diff --git a/extensions-builtin/sd-extension-steps-animation b/extensions-builtin/sd-extension-steps-animation deleted file mode 160000 index 90663eb74..000000000 --- a/extensions-builtin/sd-extension-steps-animation +++ /dev/null @@ -1 +0,0 @@ -Subproject commit 90663eb7450c3487b693cf20e76ec4d7edd78cd5 diff --git a/extensions-builtin/sd-extension-system-info b/extensions-builtin/sd-extension-system-info index 70ab5cf31..064c856ac 160000 --- a/extensions-builtin/sd-extension-system-info +++ b/extensions-builtin/sd-extension-system-info @@ -1 +1 @@ -Subproject commit 70ab5cf312be0fa913c5ba6ab85fbb64430507e2 +Subproject commit 064c856acaed8c73c3a232e39df8c37d0bb9a15c diff --git a/extensions-builtin/sd-webui-agent-scheduler b/extensions-builtin/sd-webui-agent-scheduler new file mode 160000 index 000000000..b07031958 --- /dev/null +++ b/extensions-builtin/sd-webui-agent-scheduler @@ -0,0 +1 @@ +Subproject commit b0703195833b8ea055e07b80902cbda1a01363f6 diff --git a/extensions-builtin/sd-webui-controlnet b/extensions-builtin/sd-webui-controlnet index 4d4b1f8c0..d8551e447 160000 --- a/extensions-builtin/sd-webui-controlnet +++ b/extensions-builtin/sd-webui-controlnet @@ -1 +1 @@ -Subproject commit 4d4b1f8c00a0355d1517465ac3c0e801d5a2d194 +Subproject commit d8551e447d8718e15b8ff5de04036d3fd1b3c5ce diff --git a/extensions-builtin/sd-webui-model-converter b/extensions-builtin/sd-webui-model-converter deleted file mode 160000 index d19e28168..000000000 --- a/extensions-builtin/sd-webui-model-converter +++ /dev/null @@ -1 +0,0 @@ -Subproject commit d19e28168268b0f2f50c8a5b7a4fa0a0d9b42b8c diff --git a/extensions-builtin/seed_travel b/extensions-builtin/seed_travel deleted file mode 160000 index 4bc8b2f10..000000000 --- a/extensions-builtin/seed_travel +++ /dev/null @@ -1 +0,0 @@ -Subproject commit 4bc8b2f10d5c12958f48b67ad23fb445aff074df diff --git a/extensions-builtin/stable-diffusion-webui-images-browser b/extensions-builtin/stable-diffusion-webui-images-browser index 2c988c08c..5795886be 160000 --- a/extensions-builtin/stable-diffusion-webui-images-browser +++ b/extensions-builtin/stable-diffusion-webui-images-browser @@ -1 +1 @@ -Subproject commit 2c988c08c7fc2f1c0f572bc4209f0baa1fac4fee +Subproject commit 5795886bee895c2e69e5c64e67aa643da423511c diff --git a/extensions-builtin/stable-diffusion-webui-rembg b/extensions-builtin/stable-diffusion-webui-rembg index 64821f047..3d9eedbbf 160000 --- a/extensions-builtin/stable-diffusion-webui-rembg +++ b/extensions-builtin/stable-diffusion-webui-rembg @@ -1 +1 @@ -Subproject commit 64821f04767b04b92ac157f3c02f28d269dba5d8 +Subproject commit 3d9eedbbf0d585207f97d5b21e42f32c0042df70 diff --git a/javascript/black-orange.jpg b/html/black-orange.jpg similarity index 100% rename from javascript/black-orange.jpg rename to html/black-orange.jpg diff --git a/html/card-no-preview.png b/html/card-no-preview.png index e2beb2692..952d30580 100644 Binary files a/html/card-no-preview.png and b/html/card-no-preview.png differ diff --git a/javascript/gradio-base.jpg b/html/gradio-base.jpg similarity index 100% rename from javascript/gradio-base.jpg rename to html/gradio-base.jpg diff --git a/javascript/gradio-default.jpg b/html/gradio-default.jpg similarity index 100% rename from javascript/gradio-default.jpg rename to html/gradio-default.jpg diff --git a/javascript/gradio-glass.jpg b/html/gradio-glass.jpg similarity index 100% rename from javascript/gradio-glass.jpg rename to html/gradio-glass.jpg diff --git a/javascript/gradio-monochrome.jpg b/html/gradio-monochrome.jpg similarity index 100% rename from javascript/gradio-monochrome.jpg rename to html/gradio-monochrome.jpg diff --git a/javascript/gradio-soft.jpg b/html/gradio-soft.jpg similarity index 100% rename from javascript/gradio-soft.jpg rename to html/gradio-soft.jpg diff --git a/html/licenses.html b/html/licenses.html index bc995aa07..ef6f2c0a4 100644 --- a/html/licenses.html +++ b/html/licenses.html @@ -661,4 +661,30 @@ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. + + +

TAESD

+Tiny AutoEncoder for Stable Diffusion option for live previews +
+MIT License
+
+Copyright (c) 2023 Ollin Boer Bohan
+
+Permission is hereby granted, free of charge, to any person obtaining a copy
+of this software and associated documentation files (the "Software"), to deal
+in the Software without restriction, including without limitation the rights
+to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+copies of the Software, and to permit persons to whom the Software is
+furnished to do so, subject to the following conditions:
+
+The above copyright notice and this permission notice shall be included in all
+copies or substantial portions of the Software.
+
+THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+SOFTWARE.
 
\ No newline at end of file diff --git a/automatic.ico b/html/logo.ico similarity index 100% rename from automatic.ico rename to html/logo.ico diff --git a/automatic.png b/html/logo.png similarity index 100% rename from automatic.png rename to html/logo.png diff --git a/automatic.svg b/html/logo.svg similarity index 100% rename from automatic.svg rename to html/logo.svg diff --git a/html/notification.mp3 b/html/notification.mp3 new file mode 100644 index 000000000..dba530bf7 Binary files /dev/null and b/html/notification.mp3 differ diff --git a/javascript/roboto.ttf b/html/roboto.ttf similarity index 100% rename from javascript/roboto.ttf rename to html/roboto.ttf diff --git a/installer.py b/installer.py new file mode 100644 index 000000000..6eb6d7640 --- /dev/null +++ b/installer.py @@ -0,0 +1,782 @@ +import os +import sys +import json +import time +import shutil +import logging +import platform +import subprocess +import io +import pstats +import cProfile +import pkg_resources + +try: + from modules.cmd_args import parser +except: + import argparse + parser = argparse.ArgumentParser(description="SD.Next", conflict_handler='resolve', formatter_class=lambda prog: argparse.HelpFormatter(prog, max_help_position=55, indent_increment=2, width=200)) + +class Dot(dict): # dot notation access to dictionary attributes + __getattr__ = dict.get + __setattr__ = dict.__setitem__ + __delattr__ = dict.__delitem__ + + +log = logging.getLogger("sd") +log_file = os.path.join(os.path.dirname(__file__), 'webui.log') +quick_allowed = True +errors = 0 +opts = {} +args = Dot({ + 'debug': False, + 'reset': False, + 'upgrade': False, + 'skip_update': False, + 'skip_extensions': False, + 'skip_requirements': False, + 'skip_git': False, + 'skip_torch': False, + 'use_directml': False, + 'use_ipex': False, + 'use_cuda': False, + 'use_rocm': False, + 'experimental': False, + 'test': False, + 'tls_selfsign': False, + 'reinstall': False, + 'version': False, + 'ignore': False, +}) +git_commit = "unknown" + + +# setup console and file logging +def setup_logging(clean=False): + try: + if clean and os.path.isfile(log_file): + os.remove(log_file) + time.sleep(0.1) # prevent race condition + except: + pass + from rich.theme import Theme + from rich.logging import RichHandler + from rich.console import Console + from rich.pretty import install as pretty_install + from rich.traceback import install as traceback_install + console = Console(log_time=True, log_time_format='%H:%M:%S-%f', theme=Theme({ + "traceback.border": "black", + "traceback.border.syntax_error": "black", + "inspect.value.border": "black", + })) + # logging.getLogger("urllib3").setLevel(logging.ERROR) + # logging.getLogger("httpx").setLevel(logging.ERROR) + level = logging.DEBUG if args.debug else logging.INFO + logging.basicConfig(level=logging.ERROR, format='%(asctime)s | %(name)s | %(levelname)s | %(module)s | %(message)s', filename=log_file, filemode='a', encoding='utf-8', force=True) + log.setLevel(logging.DEBUG) # log to file is always at level debug for facility `sd` + pretty_install(console=console) + traceback_install(console=console, extra_lines=1, width=console.width, word_wrap=False, indent_guides=False, suppress=[]) + rh = RichHandler(show_time=True, omit_repeated_times=False, show_level=True, show_path=False, markup=False, rich_tracebacks=True, log_time_format='%H:%M:%S-%f', level=level, console=console) + rh.set_name(level) + while log.hasHandlers() and len(log.handlers) > 0: + log.removeHandler(log.handlers[0]) + log.addHandler(rh) + + +def print_profile(profile: cProfile.Profile, msg: str): + try: + from rich import print # pylint: disable=redefined-builtin + except: + pass + profile.disable() + stream = io.StringIO() + ps = pstats.Stats(profile, stream=stream) + ps.sort_stats(pstats.SortKey.CUMULATIVE).print_stats(15) + profile = None + lines = stream.getvalue().split('\n') + lines = [l for l in lines if '=' in pkg: + p = pkg.split('>=') + else: + p = pkg.split('==') + spec = pkg_resources.working_set.by_key.get(p[0], None) # more reliable than importlib + if spec is None: + spec = pkg_resources.working_set.by_key.get(p[0].lower(), None) # check name variations + if spec is None: + spec = pkg_resources.working_set.by_key.get(p[0].replace('_', '-'), None) # check name variations + ok = ok and spec is not None + if ok: + version = pkg_resources.get_distribution(p[0]).version + log.debug(f"Package version found: {p[0]} {version}") + if len(p) > 1: + ok = ok and version == p[1] + if not ok: + log.warning(f"Package wrong version: {p[0]} {version} required {p[1]}") + else: + log.debug(f"Package version not found: {p[0]}") + return ok + except ModuleNotFoundError: + log.debug(f"Package not installed: {pkgs}") + return False + + +def pip(arg: str, ignore: bool = False, quiet: bool = False): + arg = arg.replace('>=', '==') + if not quiet: + log.info(f'Installing package: {arg.replace("install", "").replace("--upgrade", "").replace("--no-deps", "").replace("--force", "").replace(" ", " ").strip()}') + log.debug(f"Running pip: {arg}") + result = subprocess.run(f'"{sys.executable}" -m pip {arg}', shell=True, check=False, env=os.environ, stdout=subprocess.PIPE, stderr=subprocess.PIPE) + txt = result.stdout.decode(encoding="utf8", errors="ignore") + if len(result.stderr) > 0: + txt += ('\n' if len(txt) > 0 else '') + result.stderr.decode(encoding="utf8", errors="ignore") + txt = txt.strip() + if result.returncode != 0 and not ignore: + global errors # pylint: disable=global-statement + errors += 1 + log.error(f'Error running pip: {arg}') + log.debug(f'Pip output: {txt}') + return txt + + +# install package using pip if not already installed +def install(package, friendly: str = None, ignore: bool = False): + if args.reinstall: + global quick_allowed # pylint: disable=global-statement + quick_allowed = False + if args.use_ipex and package == "pytorch_lightning==1.9.4": + package = "pytorch_lightning==1.8.6" + if args.reinstall or not installed(package, friendly): + pip(f"install --upgrade {package}", ignore=ignore) + + +# execute git command +def git(arg: str, folder: str = None, ignore: bool = False): + if args.skip_git: + return '' + git_cmd = os.environ.get('GIT', "git") + result = subprocess.run(f'"{git_cmd}" {arg}', check=False, shell=True, env=os.environ, stdout=subprocess.PIPE, stderr=subprocess.PIPE, cwd=folder or '.') + txt = result.stdout.decode(encoding="utf8", errors="ignore") + if len(result.stderr) > 0: + txt += ('\n' if len(txt) > 0 else '') + result.stderr.decode(encoding="utf8", errors="ignore") + txt = txt.strip() + if result.returncode != 0 and not ignore: + global errors # pylint: disable=global-statement + errors += 1 + log.error(f'Error running git: {folder} / {arg}') + if 'or stash them' in txt: + log.error(f'Local changes detected: check log for details: {log_file}') + log.debug(f'Git output: {txt}') + return txt + + +# update switch to main branch as head can get detached and update repository +def update(folder): + if not os.path.exists(os.path.join(folder, '.git')): + return + branch = git('branch', folder) + if 'main' in branch: + branch = 'main' + elif 'master' in branch: + branch = 'master' + else: + branch = branch.split('\n')[0].replace('*', '').strip() + # log.debug(f'Setting branch: {folder} / {branch}') + git(f'checkout {branch}', folder) + if branch is None: + git('pull --autostash --rebase --force', folder) + else: + git(f'pull origin {branch} --autostash --rebase --force', folder) + # branch = git('branch', folder) + + +# clone git repository +def clone(url, folder, commithash=None): + if os.path.exists(folder): + if args.skip_update: + return + if commithash is None: + update(folder) + else: + current_hash = git('rev-parse HEAD', folder).strip() + if current_hash != commithash: + git('fetch', folder) + git(f'checkout {commithash}', folder) + return + else: + log.info(f'Cloning repository: {url}') + git(f'clone "{url}" "{folder}"') + if commithash is not None: + git(f'-C "{folder}" checkout {commithash}') + + +# check python version +def check_python(): + supported_minors = [9, 10] + if args.experimental: + supported_minors.append(11) + log.info(f'Python {platform.python_version()} on {platform.system()}') + if not (int(sys.version_info.major) == 3 and int(sys.version_info.minor) in supported_minors): + log.error(f"Incompatible Python version: {sys.version_info.major}.{sys.version_info.minor}.{sys.version_info.micro} required 3.{supported_minors}") + if not args.ignore: + sys.exit(1) + if not args.skip_git: + git_cmd = os.environ.get('GIT', "git") + if shutil.which(git_cmd) is None: + log.error('Git not found') + if not args.ignore: + sys.exit(1) + else: + git_version = git('--version', folder=None, ignore=False) + log.debug(f'Git {git_version.replace("git version", "").strip()}') + + +# check torch version +def check_torch(): + if args.profile: + pr = cProfile.Profile() + pr.enable() + allow_cuda = not (args.use_rocm or args.use_directml or args.use_ipex) + allow_rocm = not (args.use_cuda or args.use_directml or args.use_ipex) + allow_ipex = not (args.use_cuda or args.use_rocm or args.use_directml) + allow_directml = not (args.use_cuda or args.use_rocm or args.use_ipex) + log.debug(f'Torch overrides: cuda={args.use_cuda} rocm={args.use_rocm} ipex={args.use_ipex} diml={args.use_directml}') + log.debug(f'Torch allowed: cuda={allow_cuda} rocm={allow_rocm} ipex={allow_ipex} diml={allow_directml}') + torch_command = os.environ.get('TORCH_COMMAND', '') + if torch_command != '': + pass + elif allow_cuda and (shutil.which('nvidia-smi') is not None or os.path.exists(os.path.join(os.environ.get('SystemRoot') or r'C:\Windows', 'System32', 'nvidia-smi.exe'))): + log.info('nVidia CUDA toolkit detected') + torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision --index-url https://download.pytorch.org/whl/cu118') + xformers_package = os.environ.get('XFORMERS_PACKAGE', 'xformers==0.0.17' if opts.get('cross_attention_optimization', '') == 'xFormers' else 'none') + elif allow_rocm and (shutil.which('rocminfo') is not None or os.path.exists('/opt/rocm/bin/rocminfo') or os.path.exists('/dev/kfd')): + log.info('AMD ROCm toolkit detected') + os.environ.setdefault('HSA_OVERRIDE_GFX_VERSION', '10.3.0') + os.environ.setdefault('PYTORCH_HIP_ALLOC_CONF', 'garbage_collection_threshold:0.8,max_split_size_mb:512') + torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.0.1 torchvision==0.15.2 --index-url https://download.pytorch.org/whl/rocm5.4.2') + xformers_package = os.environ.get('XFORMERS_PACKAGE', 'none') + elif allow_ipex and args.use_ipex and shutil.which('sycl-ls') is not None: + log.info('Intel OneAPI Toolkit detected') + torch_command = os.environ.get('TORCH_COMMAND', 'torch==1.13.0a0 torchvision==0.14.1a0 intel_extension_for_pytorch==1.13.120+xpu -f https://developer.intel.com/ipex-whl-stable-xpu') + xformers_package = os.environ.get('XFORMERS_PACKAGE', 'none') + else: + machine = platform.machine() + if sys.platform == 'darwin': + torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.0.1 torchvision==0.15.2') + elif allow_directml and args.use_directml and ('arm' not in machine and 'aarch' not in machine): + log.info('Using DirectML Backend') + torch_command = os.environ.get('TORCH_COMMAND', 'torch-directml') + xformers_package = os.environ.get('XFORMERS_PACKAGE', 'none') + if 'torch' in torch_command and not args.version: + install(torch_command, 'torch torchvision') + else: + log.info('Using CPU-only Torch') + torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision') + xformers_package = os.environ.get('XFORMERS_PACKAGE', 'none') + if 'torch' in torch_command and not args.version: + install(torch_command, 'torch torchvision') + if args.skip_torch: + log.info('Skipping Torch tests') + else: + try: + import torch + log.info(f'Torch {torch.__version__}') + if args.use_ipex and allow_ipex: + import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import + log.info(f'Torch backend: Intel IPEX {ipex.__version__}') + log.info(f'{os.popen("icpx --version").read().rstrip()}') + for device in [torch.xpu.device(i) for i in range(torch.xpu.device_count())]: + log.info(f'Torch detected GPU: {torch.xpu.get_device_name(device)} VRAM {round(torch.xpu.get_device_properties(device).total_memory / 1024 / 1024)}') + elif torch.cuda.is_available() and (allow_cuda or allow_rocm): + # log.debug(f'Torch allocator: {torch.cuda.get_allocator_backend()}') + if torch.version.cuda and allow_cuda: + log.info(f'Torch backend: nVidia CUDA {torch.version.cuda} cuDNN {torch.backends.cudnn.version() if torch.backends.cudnn.is_available() else "N/A"}') + elif torch.version.hip and allow_rocm: + log.info(f'Torch backend: AMD ROCm HIP {torch.version.hip}') + else: + log.warning('Unknown Torch backend') + for device in [torch.cuda.device(i) for i in range(torch.cuda.device_count())]: + log.info(f'Torch detected GPU: {torch.cuda.get_device_name(device)} VRAM {round(torch.cuda.get_device_properties(device).total_memory / 1024 / 1024)} Arch {torch.cuda.get_device_capability(device)} Cores {torch.cuda.get_device_properties(device).multi_processor_count}') + else: + try: + if args.use_directml and allow_directml: + import torch_directml # pylint: disable=import-error + version = pkg_resources.get_distribution("torch-directml") + log.info(f'Torch backend: DirectML ({version})') + for i in range(0, torch_directml.device_count()): + log.info(f'Torch detected GPU: {torch_directml.device_name(i)}') + except: + log.warning("Torch reports CUDA not available") + except Exception as e: + log.error(f'Could not load torch: {e}') + if not args.ignore: + sys.exit(1) + if args.version: + return + try: + if 'xformers' in xformers_package: + install(f'--no-deps {xformers_package}', ignore=True) + else: + x = pkg_resources.working_set.by_key.get('xformers', None) + if x is not None: + log.warning(f'Not used, uninstalling: {x}') + pip('uninstall xformers --yes --quiet', ignore=True, quiet=True) + except Exception as e: + log.debug(f'Cannot install xformers package: {e}') + try: + tensorflow_package = os.environ.get('TENSORFLOW_PACKAGE', 'tensorflow==2.12.0') + install(tensorflow_package, 'tensorflow', ignore=True) + except Exception as e: + log.debug(f'Cannot install tensorflow package: {e}') + if opts.get('cuda_compile_mode', '') == 'hidet': + install('hidet', 'hidet') + if args.profile: + print_profile(pr, 'Torch') + + +# install required packages +def install_packages(): + if args.profile: + pr = cProfile.Profile() + pr.enable() + log.info('Installing packages') + # gfpgan_package = os.environ.get('GFPGAN_PACKAGE', "git+https://github.com/TencentARC/GFPGAN.git@8d2447a2d918f8eba5a4a01463fd48e45126a379") + # openclip_package = os.environ.get('OPENCLIP_PACKAGE', "git+https://github.com/mlfoundations/open_clip.git@bb6e834e9c70d9c27d0dc3ecedeebeaeb1ffad6b") + # install(gfpgan_package, 'gfpgan') + # install(openclip_package, 'open-clip-torch') + clip_package = os.environ.get('CLIP_PACKAGE', "git+https://github.com/openai/CLIP.git") + install(clip_package, 'clip') + install('onnxruntime==1.14.0', 'onnxruntime', ignore=True) + if args.profile: + print_profile(pr, 'Packages') + + +# clone required repositories +def install_repositories(): + if args.profile: + pr = cProfile.Profile() + pr.enable() + def d(name): + return os.path.join(os.path.dirname(__file__), 'repositories', name) + log.info('Installing repositories') + os.makedirs(os.path.join(os.path.dirname(__file__), 'repositories'), exist_ok=True) + stable_diffusion_repo = os.environ.get('STABLE_DIFFUSION_REPO', "https://github.com/Stability-AI/stablediffusion.git") + # stable_diffusion_commit = os.environ.get('STABLE_DIFFUSION_COMMIT_HASH', "cf1d67a6fd5ea1aa600c4df58e5b47da45f6bdbf") + stable_diffusion_commit = os.environ.get('STABLE_DIFFUSION_COMMIT_HASH', None) + clone(stable_diffusion_repo, d('stable-diffusion-stability-ai'), stable_diffusion_commit) + taming_transformers_repo = os.environ.get('TAMING_TRANSFORMERS_REPO', "https://github.com/CompVis/taming-transformers.git") + # taming_transformers_commit = os.environ.get('TAMING_TRANSFORMERS_COMMIT_HASH', "3ba01b241669f5ade541ce990f7650a3b8f65318") + taming_transformers_commit = os.environ.get('TAMING_TRANSFORMERS_COMMIT_HASH', None) + clone(taming_transformers_repo, d('taming-transformers'), taming_transformers_commit) + k_diffusion_repo = os.environ.get('K_DIFFUSION_REPO', 'https://github.com/crowsonkb/k-diffusion.git') + # k_diffusion_commit = os.environ.get('K_DIFFUSION_COMMIT_HASH', "b43db16749d51055f813255eea2fdf1def801919") + k_diffusion_commit = os.environ.get('K_DIFFUSION_COMMIT_HASH', None) + clone(k_diffusion_repo, d('k-diffusion'), k_diffusion_commit) + codeformer_repo = os.environ.get('CODEFORMER_REPO', 'https://github.com/sczhou/CodeFormer.git') + # codeformer_commit = os.environ.get('CODEFORMER_COMMIT_HASH', "c5b4593074ba6214284d6acd5f1719b6c5d739af") + codeformer_commit = os.environ.get('CODEFORMER_COMMIT_HASH', "7a584fd") + clone(codeformer_repo, d('CodeFormer'), codeformer_commit) + blip_repo = os.environ.get('BLIP_REPO', 'https://github.com/salesforce/BLIP.git') + # blip_commit = os.environ.get('BLIP_COMMIT_HASH', "48211a1594f1321b00f14c9f7a5b4813144b2fb9") + blip_commit = os.environ.get('BLIP_COMMIT_HASH', None) + clone(blip_repo, d('BLIP'), blip_commit) + if args.profile: + print_profile(pr, 'Repositories') + + +# run extension installer +def run_extension_installer(folder): + path_installer = os.path.join(folder, "install.py") + if not os.path.isfile(path_installer): + return + try: + log.debug(f"Running extension installer: {path_installer}") + env = os.environ.copy() + env['PYTHONPATH'] = os.path.abspath(".") + result = subprocess.run(f'"{sys.executable}" "{path_installer}"', shell=True, env=env, check=False, stdout=subprocess.PIPE, stderr=subprocess.PIPE, cwd=folder) + if result.returncode != 0: + global errors # pylint: disable=global-statement + errors += 1 + txt = result.stdout.decode(encoding="utf8", errors="ignore") + if len(result.stderr) > 0: + txt = txt + '\n' + result.stderr.decode(encoding="utf8", errors="ignore") + log.error(f'Error running extension installer: {path_installer}') + log.debug(txt) + except Exception as e: + log.error(f'Exception running extension installer: {e}') + +# get list of all enabled extensions +def list_extensions(folder): + disabled_extensions = opts.get('disable_all_extensions', 'none') + if disabled_extensions != 'none': + log.debug(f'Disabled extensions: {disabled_extensions}') + return [] + disabled_extensions = set(opts.get('disabled_extensions', [])) + if len(disabled_extensions) > 0: + log.debug(f'Disabled extensions: {disabled_extensions}') + return [x for x in os.listdir(folder) if x not in disabled_extensions and not x.startswith('.')] + + +# run installer for each installed and enabled extension and optionally update them +def install_extensions(): + if args.profile: + pr = cProfile.Profile() + pr.enable() + pkg_resources._initialize_master_working_set() # pylint: disable=protected-access + pkgs = [f'{p.project_name}=={p._version}' for p in pkg_resources.working_set] # pylint: disable=protected-access,not-an-iterable + log.debug(f'Installed packages: {len(pkgs)}') + from modules.paths_internal import extensions_builtin_dir, extensions_dir + extensions_duplicates = [] + extensions_enabled = [] + extension_folders = [extensions_builtin_dir] if args.safe else [extensions_builtin_dir, extensions_dir] + for folder in extension_folders: + if not os.path.isdir(folder): + continue + extensions = list_extensions(folder) + log.debug(f'Extensions all: {extensions}') + for ext in extensions: + if ext in extensions_enabled: + extensions_duplicates.append(ext) + continue + extensions_enabled.append(ext) + if not args.skip_update: + try: + update(os.path.join(folder, ext)) + except: + log.error(f'Error updating extension: {os.path.join(folder, ext)}') + if not args.skip_extensions: + run_extension_installer(os.path.join(folder, ext)) + pkg_resources._initialize_master_working_set() # pylint: disable=protected-access + updated = [f'{p.project_name}=={p._version}' for p in pkg_resources.working_set] # pylint: disable=protected-access,not-an-iterable + diff = [x for x in updated if x not in pkgs] + pkgs = updated + if len(diff) > 0: + log.info(f'Extension installed packages: {ext} {diff}') + log.info(f'Extensions enabled: {extensions_enabled}') + if len(extensions_duplicates) > 0: + log.warning(f'Extensions duplicates: {extensions_duplicates}') + if args.profile: + print_profile(pr, 'Extensions') + + +# initialize and optionally update submodules +def install_submodules(): + if args.profile: + pr = cProfile.Profile() + pr.enable() + log.info('Installing submodules') + txt = git('submodule') + log.debug(f'Submodules list: {txt}') + if 'no submodule mapping found' in txt: + log.warning('Attempting repository recover') + git('add .') + git('stash') + git('merge --abort', folder=None, ignore=True) + git('fetch --all') + git('reset --hard origin/master') + git('checkout master') + txt = git('submodule') + log.info('Continuing setup') + git('submodule --quiet update --init --recursive') + if not args.skip_update: + log.info('Updating submodules') + submodules = txt.splitlines() + for submodule in submodules: + try: + name = submodule.split()[1].strip() + update(name) + except: + log.error(f'Error updating submodule: {submodule}') + if args.profile: + print_profile(pr, 'Submodule') + + +def ensure_base_requirements(): + try: + import rich # pylint: disable=unused-import + except ImportError: + install('rich', 'rich') + + +def install_requirements(): + if args.profile: + pr = cProfile.Profile() + pr.enable() + if args.skip_requirements: + return + log.info('Verifying requirements') + with open('requirements.txt', 'r', encoding='utf8') as f: + lines = [line.strip() for line in f.readlines() if line.strip() != '' and not line.startswith('#') and line is not None] + for line in lines: + install(line) + if args.profile: + print_profile(pr, 'Requirements') + + +# set environment variables controling the behavior of various libraries +def set_environment(): + log.info('Setting environment tuning') + os.environ.setdefault('USE_TORCH', '1') + os.environ.setdefault('TF_CPP_MIN_LOG_LEVEL', '2') + os.environ.setdefault('ACCELERATE', 'True') + os.environ.setdefault('FORCE_CUDA', '1') + os.environ.setdefault('ATTN_PRECISION', 'fp16') + os.environ.setdefault('PYTORCH_CUDA_ALLOC_CONF', 'garbage_collection_threshold:0.8,max_split_size_mb:512') + os.environ.setdefault('CUDA_LAUNCH_BLOCKING', '0') + os.environ.setdefault('CUDA_CACHE_DISABLE', '0') + os.environ.setdefault('CUDA_AUTO_BOOST', '1') + os.environ.setdefault('CUDA_MODULE_LOADING', 'LAZY') + os.environ.setdefault('CUDA_DEVICE_DEFAULT_PERSISTING_L2_CACHE_PERCENTAGE_LIMIT', '0') + os.environ.setdefault('GRADIO_ANALYTICS_ENABLED', 'False') + os.environ.setdefault('SAFETENSORS_FAST_GPU', '1') + os.environ.setdefault('NUMEXPR_MAX_THREADS', '16') + os.environ.setdefault('PYTHONHTTPSVERIFY', '0') + os.environ.setdefault('HF_HUB_DISABLE_TELEMETRY', '1') + if sys.platform == 'darwin': + os.environ.setdefault('PYTORCH_ENABLE_MPS_FALLBACK', '1') + + +def check_extensions(): + newest_all = os.path.getmtime('requirements.txt') + from modules.paths_internal import extensions_builtin_dir, extensions_dir + extension_folders = [extensions_builtin_dir] if args.safe else [extensions_builtin_dir, extensions_dir] + for folder in extension_folders: + if not os.path.isdir(folder): + continue + extensions = list_extensions(folder) + for ext in extensions: + newest = 0 + extension_dir = os.path.join(folder, ext) + if not os.path.isdir(extension_dir): + log.debug(f'Extension listed as installed but folder missing: {extension_dir}') + continue + for f in os.listdir(extension_dir): + if '.json' in f or '.csv' in f or '__pycache__' in f: + continue + ts = os.path.getmtime(os.path.join(extension_dir, f)) + newest = max(newest, ts) + newest_all = max(newest_all, newest) + log.debug(f'Extension version: {time.ctime(newest)} {folder}{os.pathsep}{ext}') + return round(newest_all) + + +# check version of the main repo and optionally upgrade it +def check_version(offline=False, reset=True): # pylint: disable=unused-argument + if not os.path.exists('.git'): + log.error('Not a git repository') + if not args.ignore: + sys.exit(1) + # status = git('status') + # if 'branch' not in status: + # log.error('Cannot get git repository status') + # sys.exit(1) + ver = git('log -1 --pretty=format:"%h %ad"') + log.info(f'Version: {ver}') + if args.version: + return + commit = git('rev-parse HEAD') + global git_commit # pylint: disable=global-statement + git_commit = commit[:7] + try: + import requests + except ImportError: + return + commits = None + try: + commits = requests.get('https://api.github.com/repos/vladmandic/automatic/branches/master', timeout=10).json() + if commits['commit']['sha'] != commit: + if args.upgrade: + global quick_allowed # pylint: disable=global-statement + quick_allowed = False + log.info('Updating main repository') + try: + git('add .') + git('stash') + update('.') + # git('git stash pop') + ver = git('log -1 --pretty=format:"%h %ad"') + log.info(f'Upgraded to version: {ver}') + except Exception: + if not reset: + log.error('Error during repository upgrade') + else: + log.warning('Retrying repository upgrade...') + git_reset() + check_version(offline=offline, reset=False) + else: + log.info(f'Latest published version: {commits["commit"]["sha"]} {commits["commit"]["commit"]["author"]["date"]}') + except Exception as e: + log.error(f'Failed to check version: {e} {commits}') + + +def update_wiki(): + if not args.skip_update: + log.info('Updating Wiki') + try: + update(os.path.join(os.path.dirname(__file__), "wiki")) + update(os.path.join(os.path.dirname(__file__), "wiki", "origin-wiki")) + except: + log.error('Error updating wiki') + + +# check if we can run setup in quick mode +def check_timestamp(): + if not quick_allowed or not os.path.isfile(log_file): + return False + if args.skip_git: + return True + ok = True + setup_time = -1 + with open(log_file, 'r', encoding='utf8') as f: + lines = f.readlines() + for line in lines: + if 'Setup complete without errors' in line: + setup_time = int(line.split(' ')[-1]) + try: + version_time = int(git('log -1 --pretty=format:"%at"')) + except Exception as e: + log.error(f'Error getting local repository version: {e}') + if not args.ignore: + sys.exit(1) + log.debug(f'Repository update time: {time.ctime(int(version_time))}') + if setup_time == -1: + return False + log.debug(f'Previous setup time: {time.ctime(setup_time)}') + if setup_time < version_time: + ok = False + extension_time = check_extensions() + log.debug(f'Latest extensions time: {time.ctime(extension_time)}') + if setup_time < extension_time: + ok = False + log.debug(f'Timestamps: version:{version_time} setup:{setup_time} extension:{extension_time}') + if args.reinstall: + ok = False + return ok + + +def add_args(): + group = parser.add_argument_group('Setup options') + group.add_argument('--debug', default = False, action='store_true', help = "Run installer with debug logging, default: %(default)s") + group.add_argument('--reset', default = False, action='store_true', help = "Reset main repository to latest version, default: %(default)s") + group.add_argument('--upgrade', default = False, action='store_true', help = "Upgrade main repository to latest version, default: %(default)s") + group.add_argument("--use-ipex", default = False, action='store_true', help="Use Intel OneAPI XPU backend, default: %(default)s") + group.add_argument('--use-directml', default = False, action='store_true', help = "Use DirectML if no compatible GPU is detected, default: %(default)s") + group.add_argument("--use-cuda", default=False, action='store_true', help="Force use nVidia CUDA backend, default: %(default)s") + group.add_argument("--use-rocm", default=False, action='store_true', help="Force use AMD ROCm backend, default: %(default)s") + group.add_argument('--skip-update', default = False, action='store_true', help = "Skip update of extensions and submodules, default: %(default)s") + group.add_argument('--skip-requirements', default = False, action='store_true', help = "Skips checking and installing requirements, default: %(default)s") + group.add_argument('--skip-extensions', default = False, action='store_true', help = "Skips running individual extension installers, default: %(default)s") + group.add_argument('--skip-git', default = False, action='store_true', help = "Skips running all GIT operations, default: %(default)s") + group.add_argument('--skip-torch', default = False, action='store_true', help = "Skips running Torch checks, default: %(default)s") + group.add_argument('--experimental', default = False, action='store_true', help = "Allow unsupported versions of libraries, default: %(default)s") + group.add_argument('--reinstall', default = False, action='store_true', help = "Force reinstallation of all requirements, default: %(default)s") + group.add_argument('--test', default = False, action='store_true', help = "Run test only and exit") + group.add_argument('--version', default = False, action='store_true', help = "Print version information") + group.add_argument('--ignore', default = False, action='store_true', help = "Ignore any errors and attempt to continue") + group.add_argument('--safe', default = False, action='store_true', help = "Run in safe mode with no user extensions") + + +def parse_args(): + # command line args + global args # pylint: disable=global-statement + args = parser.parse_args() + + +def extensions_preload(force = False): + if args.profile: + pr = cProfile.Profile() + pr.enable() + setup_time = 0 + if not force: + if os.path.isfile(log_file): + with open(log_file, 'r', encoding='utf8') as f: + lines = f.readlines() + for line in lines: + if 'Setup complete without errors' in line: + setup_time = int(line.split(' ')[-1]) + if setup_time > 0 or force: + log.info('Running extension preloading') + if args.safe: + log.info('Running in safe mode without user extensions') + try: + from modules.script_loading import preload_extensions + from modules.paths_internal import extensions_builtin_dir, extensions_dir + extension_folders = [extensions_builtin_dir] if args.safe else [extensions_builtin_dir, extensions_dir] + for ext_dir in extension_folders: + t0 = time.time() + preload_extensions(ext_dir, parser) + t1 = time.time() + log.debug(f'Extension preload: {round(t1 - t0, 1)}s {ext_dir}') + except: + log.error('Error running extension preloading') + if args.profile: + print_profile(pr, 'Preload') + +def git_reset(): + log.warning('Running GIT reset') + global quick_allowed # pylint: disable=global-statement + quick_allowed = False + git('merge --abort') + git('fetch --all') + git('reset --hard origin/master') + git('checkout master') + log.info('GIT reset complete') + + +def read_options(): + global opts # pylint: disable=global-statement + if os.path.isfile(args.config): + with open(args.config, "r", encoding="utf8") as file: + opts = json.load(file) + + +# entry method when used as module +def run_setup(): + # setup_logging(args.upgrade) + log.info('Starting SD.Next') + read_options() + check_python() + if args.reset: + git_reset() + if args.skip_git: + log.info('Skipping GIT operations') + check_version() + set_environment() + if args.reinstall: + log.info('Forcing reinstall of all packages') + check_torch() + install_requirements() + install_packages() + if check_timestamp(): + log.info('No changes detected: Quick launch active') + return + log.info("Running setup") + log.debug(f"Args: {vars(args)}") + install_repositories() + install_submodules() + install_extensions() + update_wiki() + if errors == 0: + log.debug(f'Setup complete without errors: {round(time.time())}') + else: + log.warning(f'Setup complete with errors: {errors}') + log.warning(f'See log file for more details: {log_file}') + + +if __name__ == "__main__": + add_args() + ensure_base_requirements() + parse_args() + run_setup() diff --git a/javascript/aspectRatioOverlay.js b/javascript/aspectRatioOverlay.js index a8278cca2..1159cb697 100644 --- a/javascript/aspectRatioOverlay.js +++ b/javascript/aspectRatioOverlay.js @@ -1,116 +1,75 @@ - -let currentWidth = null; -let currentHeight = null; -let arFrameTimeout = setTimeout(function(){},0); - -function dimensionChange(e, is_width, is_height){ - - if(is_width){ - currentWidth = e.target.value*1.0 - } - if(is_height){ - currentHeight = e.target.value*1.0 - } - - var inImg2img = gradioApp().querySelector("#tab_img2img").style.display == "block"; - - if(!inImg2img){ - return; - } - - var targetElement = null; - - var tabIndex = get_tab_index('mode_img2img') - if(tabIndex == 0){ // img2img - targetElement = gradioApp().querySelector('#img2img_image div[data-testid=image] img'); - } else if(tabIndex == 1){ //Sketch - targetElement = gradioApp().querySelector('#img2img_sketch div[data-testid=image] img'); - } else if(tabIndex == 2){ // Inpaint - targetElement = gradioApp().querySelector('#img2maskimg div[data-testid=image] img'); - } else if(tabIndex == 3){ // Inpaint sketch - targetElement = gradioApp().querySelector('#inpaint_sketch div[data-testid=image] img'); - } - - - if(targetElement){ - - var arPreviewRect = gradioApp().querySelector('#imageARPreview'); - if(!arPreviewRect){ - arPreviewRect = document.createElement('div') - arPreviewRect.id = "imageARPreview"; - gradioApp().appendChild(arPreviewRect) - } - - - - var viewportOffset = targetElement.getBoundingClientRect(); - - viewportscale = Math.min( targetElement.clientWidth/targetElement.naturalWidth, targetElement.clientHeight/targetElement.naturalHeight ) - - scaledx = targetElement.naturalWidth*viewportscale - scaledy = targetElement.naturalHeight*viewportscale - - cleintRectTop = (viewportOffset.top+window.scrollY) - cleintRectLeft = (viewportOffset.left+window.scrollX) - cleintRectCentreY = cleintRectTop + (targetElement.clientHeight/2) - cleintRectCentreX = cleintRectLeft + (targetElement.clientWidth/2) - - viewRectTop = cleintRectCentreY-(scaledy/2) - viewRectLeft = cleintRectCentreX-(scaledx/2) - arRectWidth = scaledx - arRectHeight = scaledy - - arscale = Math.min( arRectWidth/currentWidth, arRectHeight/currentHeight ) - arscaledx = currentWidth*arscale - arscaledy = currentHeight*arscale - - arRectTop = cleintRectCentreY-(arscaledy/2) - arRectLeft = cleintRectCentreX-(arscaledx/2) - arRectWidth = arscaledx - arRectHeight = arscaledy - - arPreviewRect.style.top = arRectTop+'px'; - arPreviewRect.style.left = arRectLeft+'px'; - arPreviewRect.style.width = arRectWidth+'px'; - arPreviewRect.style.height = arRectHeight+'px'; - - clearTimeout(arFrameTimeout); - arFrameTimeout = setTimeout(function(){ - arPreviewRect.style.display = 'none'; - },2000); - - arPreviewRect.style.display = 'block'; - - } - -} - - -onUiUpdate(function(){ - var arPreviewRect = gradioApp().querySelector('#imageARPreview'); - if(arPreviewRect){ - arPreviewRect.style.display = 'none'; - } - var tabImg2img = gradioApp().querySelector("#tab_img2img"); - if (tabImg2img) { - var inImg2img = tabImg2img.style.display == "block"; - if(inImg2img){ - let inputs = gradioApp().querySelectorAll('input'); - inputs.forEach(function(e){ - var is_width = e.parentElement.id == "img2img_width" - var is_height = e.parentElement.id == "img2img_height" - - if((is_width || is_height) && !e.classList.contains('scrollwatch')){ - e.addEventListener('input', function(e){dimensionChange(e, is_width, is_height)} ) - e.classList.add('scrollwatch') - } - if(is_width){ - currentWidth = e.value*1.0 - } - if(is_height){ - currentHeight = e.value*1.0 - } - }) - } - } -}); +/* global gradioApp, onUiUpdate, get_tab_index */ + +let currentWidth = null; +let currentHeight = null; +let arFrameTimeout = setTimeout(() => {}, 0); + +function dimensionChange(e, is_width, is_height) { + if (is_width) currentWidth = e.target.value * 1.0; + if (is_height) currentHeight = e.target.value * 1.0; + const inImg2img = gradioApp().querySelector('#tab_img2img').style.display === 'block'; + if (!inImg2img) return; + let targetElement = null; + const tabIndex = get_tab_index('mode_img2img'); + if (tabIndex === 0) targetElement = gradioApp().querySelector('#img2img_image div[data-testid=image] img'); // img2img + else if (tabIndex === 1) targetElement = gradioApp().querySelector('#img2img_sketch div[data-testid=image] img'); // Sketch + else if (tabIndex === 2) targetElement = gradioApp().querySelector('#img2maskimg div[data-testid=image] img'); // Inpaint + else if (tabIndex === 3) targetElement = gradioApp().querySelector('#inpaint_sketch div[data-testid=image] img'); // Inpaint sketch + + if (targetElement) { + let arPreviewRect = gradioApp().querySelector('#imageARPreview'); + if (!arPreviewRect) { + arPreviewRect = document.createElement('div'); + arPreviewRect.id = 'imageARPreview'; + gradioApp().appendChild(arPreviewRect); + } + + const viewportOffset = targetElement.getBoundingClientRect(); + const viewportscale = Math.min(targetElement.clientWidth / targetElement.naturalWidth, targetElement.clientHeight / targetElement.naturalHeight); + const scaledx = targetElement.naturalWidth * viewportscale; + const scaledy = targetElement.naturalHeight * viewportscale; + const cleintRectTop = (viewportOffset.top + window.scrollY); + const cleintRectLeft = (viewportOffset.left + window.scrollX); + const cleintRectCentreY = cleintRectTop + (targetElement.clientHeight / 2); + const cleintRectCentreX = cleintRectLeft + (targetElement.clientWidth / 2); + const arscale = Math.min(scaledx / currentWidth, scaledy / currentHeight); + const arscaledx = currentWidth * arscale; + const arscaledy = currentHeight * arscale; + const arRectTop = cleintRectCentreY - (arscaledy / 2); + const arRectLeft = cleintRectCentreX - (arscaledx / 2); + const arRectWidth = arscaledx; + const arRectHeight = arscaledy; + arPreviewRect.style.top = `${arRectTop}px`; + arPreviewRect.style.left = `${arRectLeft}px`; + arPreviewRect.style.width = `${arRectWidth}px`; + arPreviewRect.style.height = `${arRectHeight}px`; + + clearTimeout(arFrameTimeout); + arFrameTimeout = setTimeout(() => { + arPreviewRect.style.display = 'none'; + }, 2000); + arPreviewRect.style.display = 'block'; + } +} + +onUiUpdate(() => { + const arPreviewRect = gradioApp().querySelector('#imageARPreview'); + if (arPreviewRect) arPreviewRect.style.display = 'none'; + const tabImg2img = gradioApp().querySelector('#tab_img2img'); + if (tabImg2img) { + const inImg2img = tabImg2img.style.display === 'block'; + if (inImg2img) { + const inputs = gradioApp().querySelectorAll('input'); + inputs.forEach((e) => { + const is_width = e.parentElement.id === 'img2img_width'; + const is_height = e.parentElement.id === 'img2img_height'; + if ((is_width || is_height) && !e.classList.contains('scrollwatch')) { + e.addEventListener('input', (evt) => { dimensionChange(evt, is_width, is_height); }); + e.classList.add('scrollwatch'); + } + if (is_width) currentWidth = e.value * 1.0; + if (is_height) currentHeight = e.value * 1.0; + }); + } + } +}); diff --git a/javascript/black-orange.css b/javascript/black-orange.css index a055e2e51..de7429e82 100644 --- a/javascript/black-orange.css +++ b/javascript/black-orange.css @@ -7,8 +7,8 @@ img { background-color: black; } input[type=range] { height: 18px; appearance: none; margin-top: 0; min-width: 160px; background-color: black; width: 100%; background: transparent; } input[type=range]::-webkit-slider-runnable-track { width: 100%; height: 18px; cursor: pointer; box-shadow: 2px 2px 3px #111111; background: #50555C; border-radius: 2px; border: 0px solid #222222; } input[type=range]::-moz-range-track { width: 100%; height: 18px; cursor: pointer; box-shadow: 2px 2px 3px #111111; background: #50555C; border-radius: 2px; border: 0px solid #222222; } -input[type=range]::-webkit-slider-thumb { box-shadow: 2px 2px 3px #111111; border: 0px solid #000000; height: 18px; width: 40px; border-radius: 2px; background: var(--highlight-color); cursor: pointer; -webkit-appearance: none; margin-top: 0px; } -input[type=range]::-moz-range-thumb { box-shadow: 2px 2px 3px #111111; border: 0px solid #000000; height: 18px; width: 40px; border-radius: 2px; background: var(--highlight-color); cursor: pointer; -webkit-appearance: none; margin-top: 0px; } +input[type=range]::-webkit-slider-thumb { box-shadow: 2px 2px 3px #111111; border: 0px solid #000000; height: 18px; width: 40px; border-radius: 2px; background: var(--highlight-color); cursor: pointer; appearance: none; margin-top: 0px; } +input[type=range]::-moz-range-thumb { box-shadow: 2px 2px 3px #111111; border: 0px solid #000000; height: 18px; width: 40px; border-radius: 2px; background: var(--highlight-color); cursor: pointer; appearance: none; margin-top: 0px; } ::-webkit-scrollbar { width: 12px; } ::-webkit-scrollbar-track { background: #333333; } ::-webkit-scrollbar-thumb { background-color: var(--highlight-color); border-radius: 2px; border-width: 0; box-shadow: 2px 2px 3px #111111; } @@ -42,7 +42,7 @@ div.gradio-container.dark > div.w-full.flex.flex-col.min-h-screen > div { backgr .extra-networks .tab-nav .search { max-height: 2em; } .extra-network-cards .card .actions .name { font-weight: 400; font-size: 1.2rem; } .extra-network-cards .card { width: 18em; height: 12em; box-shadow: none; } -.extra-network-cards .card:hover { transform: scale(1.5); transition: all 0.3s ease; z-index: 99; box-shadow: none; } +.extra-network-cards .card:hover { transform: scale(1.15); transition: all 0.2s ease; z-index: 99; box-shadow: none; } svg.feather.feather-image, .feather .feather-image { display: none } .gap-2 { padding-top: 8px; } .gr-box > div > div > input.gr-text-input { right: 0; width: 4em; padding: 0; top: -12px; border: none; max-height: 20px; } @@ -54,7 +54,7 @@ svg.feather.feather-image, .feather .feather-image { display: none } .py-6 { padding-bottom: 0; } .rounded-lg { border-radius: 0; } .tabs { background-color: black; } -.gradio-button.tool { border-radius: 0; height: 2em; } +.gradio-button.tool { border-radius: 0; } .block.token-counter span { background-color: #222 !important; box-shadow: 2px 2px 2px #111; border: none !important; border-radius: 0; font-size: 0.8rem; } .tab-nav { zoom: 130%; margin-bottom: 16px; border-bottom: 2px solid #CE6400 !important; padding-bottom: 2px; } .label-wrap { margin: 16px 0px 8px 0px; } @@ -64,11 +64,10 @@ svg.feather.feather-image, .feather .feather-image { display: none } #tab_extensions table thead { background-color: var(--neutral-700); } /* automatic style classes */ -.progressDiv { border-radius: 0 !important; position: fixed; top: 318px; right: 26px; max-width: 262px; height: 48px; z-index: 99; } +.progressDiv { border-radius: 0 !important; position: fixed; top: 44px; right: 26px; max-width: 262px; height: 48px; z-index: 99; box-shadow: var(--button-shadow); } .progressDiv .progress { border-radius: 0 !important; background: var(--highlight-color); line-height: 3rem; height: 48px; } .gallery-item { box-shadow: none !important; } .performance { color: #888; } -.modalControls { background-color: #4E1400; } /* gradio elements overrides */ #div.gradio-container.dark { overflow-x: hidden; } @@ -76,34 +75,32 @@ svg.feather.feather-image, .feather .feather-image { display: none } #txt2img_prompt, #txt2img_neg_prompt, #img2img_prompt, #img2img_neg_prompt { background-color: black; box-shadow: 4px 4px 4px 0px #333333 !important; } #txt2img_prompt > label > textarea, #txt2img_neg_prompt > label > textarea, #img2img_prompt > label > textarea, #img2img_neg_prompt > label > textarea { font-size: 1.2rem; } #img2img_settings { min-width: calc(2 * var(--left-column)); max-width: calc(2 * var(--left-column)); background-color: #111111; padding-top: 16px; } -#interrogate, #deepbooru { margin: 16px 0px 16px 0px; max-width: 100px; max-height: 74px; font-weight: normal; font-size: 14px; } +#interrogate, #deepbooru { margin: 0 0px 10px 0px; max-width: 80px; max-height: 80px; font-weight: normal; font-size: 0.95em; } #lightboxModal { background-color: rgba(20, 20, 20, 0.8) } #quicksettings .gr-button-tool { font-size: 1.6rem; box-shadow: none; margin-left: -20px; margin-top: -2px; height: 2.4em; } -#quicksettings > div, #quicksettings > fieldset { min-width: 26em; max-width: 26em; line-height: 2em; } -#refresh_sd_model_checkpoint { height: 40px; margin-left: -14px; background: #333333; box-shadow: none; } -#refresh_txt2img_styles, #refresh_img2img_styles, #open_folder_txt2img, #open_folder_img2img, #open_folder_extras, #footer, #style_pos_col, #style_neg_col, #roll_col, #extras_upscaler_2, #extras_upscaler_2_visibility, #txt2img_res_switch_btn, #img2img_res_switch_btn, #txt2img_seed_resize_from_w, #txt2img_seed_resize_from_h, #txt2img_tiling { display: none; } +#quicksettings > div, #quicksettings > fieldset { min-width: 24em; max-width: 26em; line-height: 2em; } +#refresh_sd_model_checkpoint { height: 48px; margin-left: -14px; background: #333333; box-shadow: none; } +#open_folder_extras, #footer, #style_pos_col, #style_neg_col, #roll_col, #extras_upscaler_2, #extras_upscaler_2_visibility, #txt2img_seed_resize_from_w, #txt2img_seed_resize_from_h { display: none; } #save-animation { border-radius: 0 !important; margin-bottom: 16px; background-color: #111111; } #script_list { padding: 4px; margin-top: 20px; margin-bottom: 20px; } #settings > div.flex-wrap { width: 15em; } #tab_extensions table { background-color: #222222; } -#txt2img_actions_column, #img2img_actions_column { min-width: 260px !important; max-width: 260px !important; } #txt2img_cfg_scale { min-width: 200px; } #txt2img_checkboxes, #img2img_checkboxes { background-color: transparent; } #txt2img_checkboxes, #img2img_checkboxes { margin-bottom: 0.2em; } #txt2img_gallery, #img2img_gallery, #extras_gallery { background: black !important; padding: 0; margin: 0; object-fit: contain; box-shadow: none; min-height: 0; } -#txt2img_generate, #img2img_generate { height: 36px; border: none; border-radius: 0; min-height: 36px; padding: 0; } -#txt2img_interrupt, #img2img_interrupt, #txt2img_skip, #img2img_skip { height: 36px; min-width: 116px; max-width: 116px; border: none; border-radius: 0; background-color: var(--inactive-color); margin-top: 46px; display: block !important; padding: 0; } +#extras_upscale { margin-top: 10px } #txt2img_progress_row > div { min-width: var(--left-column); max-width: var(--left-column); } #txt2img_results, #img2img_results, #extras_results { background-color: black; padding: 0; } #txt2img_seed_row { padding: 0; margin-top: 8px; } #txt2img_settings { min-width: var(--left-column); max-width: var(--left-column); background-color: #111111; padding-top: 16px; } #txt2img_subseed_row { padding: 0; margin-top: 16px; } -#txt2img_subseed_show { min-width: 74px; padding: 8px 0 0 0 } +#txt2img_subseed_show, #img2img_subseed_show { display: None } #txt2img_subseed_strength { margin-top: 0; } -#txt2img_tools, #img2img_tools { margin-top: 54px; scale: 120%; margin-left: 26px; } #txtimg_hr_finalres { max-width: 200px; } #pnginfo_html2_info { margin-top: -18px; background-color: var(--input-background-fill); padding: var(--input-padding) } -#txt2img_extra_refresh, #txt2img_extra_close { height: 1.7em } +#txt2img_extra_refresh, #txt2img_extra_close { height: 1.7em; } +#txt2img_tools > div > button, #img2img_tools > div > button { filter: hue-rotate(180deg) saturate(0.5); } /* custom elements overrides */ #steps-animation, #controlnet { border-width: 0; } diff --git a/javascript/contextMenus.js b/javascript/contextMenus.js index 517bacac8..880ff3172 100644 --- a/javascript/contextMenus.js +++ b/javascript/contextMenus.js @@ -1,178 +1,140 @@ - -contextMenuInit = function(){ - let eventListenerApplied=false; - let menuSpecs = new Map(); - - const uid = function(){ - return Date.now().toString(36) + Math.random().toString(36).substr(2); - } - - function showContextMenu(event,element,menuEntries){ - let posx = event.clientX + document.body.scrollLeft + document.documentElement.scrollLeft; - let posy = event.clientY + document.body.scrollTop + document.documentElement.scrollTop; - - let oldMenu = gradioApp().querySelector('#context-menu') - if(oldMenu){ - oldMenu.remove() - } - - let tabButton = uiCurrentTab - let baseStyle = window.getComputedStyle(tabButton) - - const contextMenu = document.createElement('nav') - contextMenu.id = "context-menu" - contextMenu.style.background = baseStyle.background - contextMenu.style.color = baseStyle.color - contextMenu.style.fontFamily = baseStyle.fontFamily - contextMenu.style.top = posy+'px' - contextMenu.style.left = posx+'px' - - - - const contextMenuList = document.createElement('ul') - contextMenuList.className = 'context-menu-items'; - contextMenu.append(contextMenuList); - - menuEntries.forEach(function(entry){ - let contextMenuEntry = document.createElement('a') - contextMenuEntry.innerHTML = entry['name'] - contextMenuEntry.addEventListener("click", function(e) { - entry['func'](); - }) - contextMenuList.append(contextMenuEntry); - - }) - - gradioApp().appendChild(contextMenu) - - let menuWidth = contextMenu.offsetWidth + 4; - let menuHeight = contextMenu.offsetHeight + 4; - - let windowWidth = window.innerWidth; - let windowHeight = window.innerHeight; - - if ( (windowWidth - posx) < menuWidth ) { - contextMenu.style.left = windowWidth - menuWidth + "px"; - } - - if ( (windowHeight - posy) < menuHeight ) { - contextMenu.style.top = windowHeight - menuHeight + "px"; - } - - } - - function appendContextMenuOption(targetElementSelector,entryName,entryFunction){ - - currentItems = menuSpecs.get(targetElementSelector) - - if(!currentItems){ - currentItems = [] - menuSpecs.set(targetElementSelector,currentItems); - } - let newItem = {'id':targetElementSelector+'_'+uid(), - 'name':entryName, - 'func':entryFunction, - 'isNew':true} - - currentItems.push(newItem) - return newItem['id'] - } - - function removeContextMenuOption(uid){ - menuSpecs.forEach(function(v,k) { - let index = -1 - v.forEach(function(e,ei){if(e['id']==uid){index=ei}}) - if(index>=0){ - v.splice(index, 1); - } - }) - } - - function addContextMenuEventListener(){ - if(eventListenerApplied){ - return; - } - gradioApp().addEventListener("click", function(e) { - let source = e.composedPath()[0] - if(source.id && source.id.indexOf('check_progress')>-1){ - return - } - - let oldMenu = gradioApp().querySelector('#context-menu') - if(oldMenu){ - oldMenu.remove() - } - }); - gradioApp().addEventListener("contextmenu", function(e) { - let oldMenu = gradioApp().querySelector('#context-menu') - if(oldMenu){ - oldMenu.remove() - } - menuSpecs.forEach(function(v,k) { - if(e.composedPath()[0].matches(k)){ - showContextMenu(e,e.composedPath()[0],v) - e.preventDefault() - return - } - }) - }); - eventListenerApplied=true - - } - - return [appendContextMenuOption, removeContextMenuOption, addContextMenuEventListener] -} - -initResponse = contextMenuInit(); -appendContextMenuOption = initResponse[0]; -removeContextMenuOption = initResponse[1]; -addContextMenuEventListener = initResponse[2]; - -(function(){ - //Start example Context Menu Items - let generateOnRepeat = function(genbuttonid,interruptbuttonid){ - let genbutton = gradioApp().querySelector(genbuttonid); - const busy = document.getElementById('progressbar')?.style.display == "block" - if(!busy){ - genbutton.click(); - } - clearInterval(window.generateOnRepeatInterval) - window.generateOnRepeatInterval = setInterval(function(){ - const busy = document.getElementById('progressbar')?.style.display == "block" - if(!busy){ - genbutton.click(); - } - }, - 500) - } - - appendContextMenuOption('#txt2img_generate','Generate forever',function(){ - generateOnRepeat('#txt2img_generate','#txt2img_interrupt'); - }) - appendContextMenuOption('#img2img_generate','Generate forever',function(){ - generateOnRepeat('#img2img_generate','#img2img_interrupt'); - }) - - let cancelGenerateForever = function(){ - clearInterval(window.generateOnRepeatInterval) - } - - appendContextMenuOption('#txt2img_interrupt','Cancel generate forever',cancelGenerateForever) - appendContextMenuOption('#txt2img_generate', 'Cancel generate forever',cancelGenerateForever) - appendContextMenuOption('#img2img_interrupt','Cancel generate forever',cancelGenerateForever) - appendContextMenuOption('#img2img_generate', 'Cancel generate forever',cancelGenerateForever) - - appendContextMenuOption('#roll','Roll three', - function(){ - let rollbutton = get_uiCurrentTabContent().querySelector('#roll'); - setTimeout(function(){rollbutton.click()},100) - setTimeout(function(){rollbutton.click()},200) - setTimeout(function(){rollbutton.click()},300) - } - ) -})(); -//End example Context Menu Items - -onUiUpdate(function(){ - addContextMenuEventListener() -}); +/* global gradioApp, uiCurrentTab, onUiUpdate, get_uiCurrentTabContent */ + +const contextMenuInit = () => { + let eventListenerApplied = false; + const menuSpecs = new Map(); + + const uid = () => Date.now().toString(36) + Math.random().toString(36).substring(2); + + function showContextMenu(event, element, menuEntries) { + const posx = event.clientX + document.body.scrollLeft + document.documentElement.scrollLeft; + const posy = event.clientY + document.body.scrollTop + document.documentElement.scrollTop; + const oldMenu = gradioApp().querySelector('#context-menu'); + if (oldMenu) oldMenu.remove(); + const tabButton = uiCurrentTab; + const baseStyle = window.getComputedStyle(tabButton); + const contextMenu = document.createElement('nav'); + contextMenu.id = 'context-menu'; + contextMenu.style.background = baseStyle.background; + contextMenu.style.color = baseStyle.color; + contextMenu.style.fontFamily = baseStyle.fontFamily; + contextMenu.style.top = `${posy}px`; + contextMenu.style.left = `${posx}px`; + const contextMenuList = document.createElement('ul'); + contextMenuList.className = 'context-menu-items'; + contextMenu.append(contextMenuList); + menuEntries.forEach((entry) => { + const contextMenuEntry = document.createElement('a'); + contextMenuEntry.innerHTML = entry.name; + contextMenuEntry.addEventListener('click', (e) => entry.func()); + contextMenuList.append(contextMenuEntry); + }); + gradioApp().appendChild(contextMenu); + const menuWidth = contextMenu.offsetWidth + 4; + const menuHeight = contextMenu.offsetHeight + 4; + const windowWidth = window.innerWidth; + const windowHeight = window.innerHeight; + if ((windowWidth - posx) < menuWidth) contextMenu.style.left = `${windowWidth - menuWidth}px`; + if ((windowHeight - posy) < menuHeight) contextMenu.style.top = `${windowHeight - menuHeight}px`; + } + + function appendContextMenuOption(targetElementSelector, entryName, entryFunction) { + let currentItems = menuSpecs.get(targetElementSelector); + if (!currentItems) { + currentItems = []; + menuSpecs.set(targetElementSelector, currentItems); + } + const newItem = { + id: `${targetElementSelector}_${uid()}`, + name: entryName, + func: entryFunction, + isNew: true, + }; + currentItems.push(newItem); + return newItem.id; + } + + function removeContextMenuOption(uid) { + menuSpecs.forEach((v, k) => { + let index = -1; + v.forEach((e, ei) => { if (e.id === uid) { index = ei; } }); + if (index >= 0) { + v.splice(index, 1); + } + }); + } + + function addContextMenuEventListener() { + if (eventListenerApplied) return; + gradioApp().addEventListener('click', (e) => { + if (!e.isTrusted) return; + const oldMenu = gradioApp().querySelector('#context-menu'); + if (oldMenu) oldMenu.remove(); + }); + gradioApp().addEventListener('contextmenu', (e) => { + const oldMenu = gradioApp().querySelector('#context-menu'); + if (oldMenu) oldMenu.remove(); + menuSpecs.forEach((v, k) => { + if (e.composedPath()[0].matches(k)) { + showContextMenu(e, e.composedPath()[0], v); + e.preventDefault(); + } + }); + }); + eventListenerApplied = true; + } + return [appendContextMenuOption, removeContextMenuOption, addContextMenuEventListener]; +}; + +const initResponse = contextMenuInit(); +const appendContextMenuOption = initResponse[0]; +const removeContextMenuOption = initResponse[1]; +const addContextMenuEventListener = initResponse[2]; + +(function () { + // Start example Context Menu Items + const generateOnRepeat = function (genbuttonid, interruptbuttonid) { + const genbutton = gradioApp().querySelector(genbuttonid); + const busy = document.getElementById('progressbar')?.style.display === 'block'; + if (!busy) { + genbutton.click(); + } + clearInterval(window.generateOnRepeatInterval); + window.generateOnRepeatInterval = setInterval( + () => { + const busy = document.getElementById('progressbar')?.style.display === 'block'; + if (!busy) genbutton.click(); + }, + 500, + ); + }; + + appendContextMenuOption('#txt2img_generate', 'Generate forever', () => { + generateOnRepeat('#txt2img_generate', '#txt2img_interrupt'); + }); + appendContextMenuOption('#img2img_generate', 'Generate forever', () => { + generateOnRepeat('#img2img_generate', '#img2img_interrupt'); + }); + + const cancelGenerateForever = function () { + clearInterval(window.generateOnRepeatInterval); + }; + + appendContextMenuOption('#txt2img_interrupt', 'Cancel generate forever', cancelGenerateForever); + appendContextMenuOption('#txt2img_generate', 'Cancel generate forever', cancelGenerateForever); + appendContextMenuOption('#img2img_interrupt', 'Cancel generate forever', cancelGenerateForever); + appendContextMenuOption('#img2img_generate', 'Cancel generate forever', cancelGenerateForever); + appendContextMenuOption( + '#roll', + 'Roll three', + () => { + const rollbutton = get_uiCurrentTabContent().querySelector('#roll'); + setTimeout(() => { rollbutton.click(); }, 100); + setTimeout(() => { rollbutton.click(); }, 200); + setTimeout(() => { rollbutton.click(); }, 300); + }, + ); +}()); +// End example Context Menu Items + +onUiUpdate(() => addContextMenuEventListener()); diff --git a/javascript/dragdrop.js b/javascript/dragdrop.js index 9015e4bd5..55e229e89 100644 --- a/javascript/dragdrop.js +++ b/javascript/dragdrop.js @@ -1,75 +1,75 @@ // allows drag-dropping files into gradio image elements, and also pasting images from clipboard -function isValidImageList( files ) { - return files && files?.length === 1 && ['image/png', 'image/gif', 'image/jpeg'].includes(files[0].type); +function isValidImageList(files) { + return files && files?.length === 1 && ['image/png', 'image/gif', 'image/jpeg'].includes(files[0].type); } -function dropReplaceImage( imgWrap, files ) { - if (!isValidImageList(files)) return; - const tmpFile = files[0]; - imgWrap.querySelector('.modify-upload button + button, .touch-none + div button + button')?.click(); - const callback = () => { - const fileInput = imgWrap.querySelector('input[type="file"]'); - if (fileInput) { - if (files.length === 0) { - files = new DataTransfer(); - files.items.add(tmpFile); - fileInput.files = files.files; - } else { - fileInput.files = files; - } - fileInput.dispatchEvent(new Event('change')); - } - }; - - if (imgWrap.closest('#pnginfo_image')) { - // special treatment for PNG Info tab, wait for fetch request to finish - const oldFetch = window.fetch; - window.fetch = async (input, options) => { - const response = await oldFetch(input, options); - if ( 'api/predict/' === input ) { - const content = await response.text(); - window.fetch = oldFetch; - window.requestAnimationFrame( () => callback() ); - return new Response(content, { - status: response.status, - statusText: response.statusText, - headers: response.headers - }) - } - return response; - }; - } else { - window.requestAnimationFrame(() => callback()); +function dropReplaceImage(imgWrap, files) { + if (!isValidImageList(files)) return; + const tmpFile = files[0]; + imgWrap.querySelector('.modify-upload button + button, .touch-none + div button + button')?.click(); + const callback = () => { + const fileInput = imgWrap.querySelector('input[type="file"]'); + if (fileInput) { + if (files.length === 0) { + files = new DataTransfer(); + files.items.add(tmpFile); + fileInput.files = files.files; + } else { + fileInput.files = files; + } + fileInput.dispatchEvent(new Event('change')); } + }; + + if (imgWrap.closest('#pnginfo_image')) { + // special treatment for PNG Info tab, wait for fetch request to finish + const oldFetch = window.fetch; + window.fetch = async (input, options) => { + const response = await oldFetch(input, options); + if (input === 'api/predict/') { + const content = await response.text(); + window.fetch = oldFetch; + window.requestAnimationFrame(() => callback()); + return new Response(content, { + status: response.status, + statusText: response.statusText, + headers: response.headers, + }); + } + return response; + }; + } else { + window.requestAnimationFrame(() => callback()); + } } -window.document.addEventListener('dragover', e => { - const target = e.composedPath()[0]; - const imgWrap = target.closest('[data-testid="image"]'); - if ( !imgWrap && target.placeholder && target.placeholder.indexOf("Prompt") == -1) return; - e.stopPropagation(); - e.preventDefault(); - e.dataTransfer.dropEffect = 'copy'; +window.document.addEventListener('dragover', (e) => { + const target = e.composedPath()[0]; + const imgWrap = target.closest('[data-testid="image"]'); + if (!imgWrap && target.placeholder && target.placeholder.indexOf('Prompt') === -1) return; + e.stopPropagation(); + e.preventDefault(); + e.dataTransfer.dropEffect = 'copy'; }); -window.document.addEventListener('drop', e => { - const target = e.composedPath()[0]; - if (!target.placeholder) return; - if (target.placeholder.indexOf("Prompt") == -1) return; - const imgWrap = target.closest('[data-testid="image"]'); - if (!imgWrap) return; - e.stopPropagation(); - e.preventDefault(); - const files = e.dataTransfer.files; - dropReplaceImage(imgWrap, files); +window.document.addEventListener('drop', (e) => { + const target = e.composedPath()[0]; + if (!target.placeholder) return; + if (target.placeholder.indexOf('Prompt') === -1) return; + const imgWrap = target.closest('[data-testid="image"]'); + if (!imgWrap) return; + e.stopPropagation(); + e.preventDefault(); + const { files } = e.dataTransfer; + dropReplaceImage(imgWrap, files); }); -window.addEventListener('paste', e => { - const files = e.clipboardData.files; - if ( ! isValidImageList( files ) ) return; - const visibleImageFields = [...gradioApp().querySelectorAll('[data-testid="image"]')].filter(el => uiElementIsVisible(el)); - if ( ! visibleImageFields.length ) return; - const firstFreeImageField = visibleImageFields.filter(el => el.querySelector('input[type=file]'))?.[0]; - dropReplaceImage(firstFreeImageField ? firstFreeImageField : visibleImageFields[visibleImageFields.length - 1], files); +window.addEventListener('paste', (e) => { + const { files } = e.clipboardData; + if (!isValidImageList(files)) return; + const visibleImageFields = [...gradioApp().querySelectorAll('[data-testid="image"]')].filter((el) => uiElementIsVisible(el)); + if (!visibleImageFields.length) return; + const firstFreeImageField = visibleImageFields.filter((el) => el.querySelector('input[type=file]'))?.[0]; + dropReplaceImage(firstFreeImageField || visibleImageFields[visibleImageFields.length - 1], files); }); diff --git a/javascript/edit-attention.js b/javascript/edit-attention.js index 20a5aadfb..467a99842 100644 --- a/javascript/edit-attention.js +++ b/javascript/edit-attention.js @@ -1,96 +1,98 @@ -function keyupEditAttention(event){ - let target = event.originalTarget || event.composedPath()[0]; - if (! target.matches("[id*='_toprow'] [id*='_prompt'] textarea")) return; - if (! (event.metaKey || event.ctrlKey)) return; - - let isPlus = event.key == "ArrowUp" - let isMinus = event.key == "ArrowDown" - if (!isPlus && !isMinus) return; - - let selectionStart = target.selectionStart; - let selectionEnd = target.selectionEnd; - let text = target.value; - - function selectCurrentParenthesisBlock(OPEN, CLOSE){ - if (selectionStart !== selectionEnd) return false; - - // Find opening parenthesis around current cursor - const before = text.substring(0, selectionStart); - let beforeParen = before.lastIndexOf(OPEN); - if (beforeParen == -1) return false; - let beforeParenClose = before.lastIndexOf(CLOSE); - while (beforeParenClose !== -1 && beforeParenClose > beforeParen) { - beforeParen = before.lastIndexOf(OPEN, beforeParen - 1); - beforeParenClose = before.lastIndexOf(CLOSE, beforeParenClose - 1); - } - - // Find closing parenthesis around current cursor - const after = text.substring(selectionStart); - let afterParen = after.indexOf(CLOSE); - if (afterParen == -1) return false; - let afterParenOpen = after.indexOf(OPEN); - while (afterParenOpen !== -1 && afterParen > afterParenOpen) { - afterParen = after.indexOf(CLOSE, afterParen + 1); - afterParenOpen = after.indexOf(OPEN, afterParenOpen + 1); - } - if (beforeParen === -1 || afterParen === -1) return false; - - // Set the selection to the text between the parenthesis - const parenContent = text.substring(beforeParen + 1, selectionStart + afterParen); - const lastColon = parenContent.lastIndexOf(":"); - selectionStart = beforeParen + 1; - selectionEnd = selectionStart + lastColon; - target.setSelectionRange(selectionStart, selectionEnd); - return true; - } - - // If the user hasn't selected anything, let's select their current parenthesis block - if(! selectCurrentParenthesisBlock('<', '>')){ - selectCurrentParenthesisBlock('(', ')') - } - - event.preventDefault(); - - closeCharacter = ')' - delta = opts.keyedit_precision_attention - - if (selectionStart > 0 && text[selectionStart - 1] == '<'){ - closeCharacter = '>' - delta = opts.keyedit_precision_extra - } else if (selectionStart == 0 || text[selectionStart - 1] != "(") { - - // do not include spaces at the end - while(selectionEnd > selectionStart && text[selectionEnd-1] == ' '){ - selectionEnd -= 1; - } - if(selectionStart == selectionEnd){ - return - } - - text = text.slice(0, selectionStart) + "(" + text.slice(selectionStart, selectionEnd) + ":1.0)" + text.slice(selectionEnd); - - selectionStart += 1; - selectionEnd += 1; - } - - end = text.slice(selectionEnd + 1).indexOf(closeCharacter) + 1; - weight = parseFloat(text.slice(selectionEnd + 1, selectionEnd + 1 + end)); - if (isNaN(weight)) return; - - weight += isPlus ? delta : -delta; - weight = parseFloat(weight.toPrecision(12)); - if(String(weight).length == 1) weight += ".0" - - text = text.slice(0, selectionEnd + 1) + weight + text.slice(selectionEnd + 1 + end - 1); - - target.focus(); - target.value = text; - target.selectionStart = selectionStart; - target.selectionEnd = selectionEnd; - - updateInput(target) -} - -addEventListener('keydown', (event) => { - keyupEditAttention(event); -}); \ No newline at end of file +function keyupEditAttention(event) { + const target = event.originalTarget || event.composedPath()[0]; + if (!target.matches("[id*='_toprow'] [id*='_prompt'] textarea")) return; + if (!(event.metaKey || event.ctrlKey)) return; + const isPlus = event.key === 'ArrowUp'; + const isMinus = event.key === 'ArrowDown'; + if (!isPlus && !isMinus) return; + let { selectionStart } = target; + let { selectionEnd } = target; + let text = target.value; + + function selectCurrentParenthesisBlock(OPEN, CLOSE) { + if (selectionStart !== selectionEnd) return false; + + // Find opening parenthesis around current cursor + const before = text.substring(0, selectionStart); + let beforeParen = before.lastIndexOf(OPEN); + if (beforeParen === -1) return false; + let beforeParenClose = before.lastIndexOf(CLOSE); + while (beforeParenClose !== -1 && beforeParenClose > beforeParen) { + beforeParen = before.lastIndexOf(OPEN, beforeParen - 1); + beforeParenClose = before.lastIndexOf(CLOSE, beforeParenClose - 1); + } + + // Find closing parenthesis around current cursor + const after = text.substring(selectionStart); + let afterParen = after.indexOf(CLOSE); + if (afterParen === -1) return false; + let afterParenOpen = after.indexOf(OPEN); + while (afterParenOpen !== -1 && afterParen > afterParenOpen) { + afterParen = after.indexOf(CLOSE, afterParen + 1); + afterParenOpen = after.indexOf(OPEN, afterParenOpen + 1); + } + if (beforeParen === -1 || afterParen === -1) return false; + + // Set the selection to the text between the parenthesis + const parenContent = text.substring(beforeParen + 1, selectionStart + afterParen); + const lastColon = parenContent.lastIndexOf(':'); + selectionStart = beforeParen + 1; + selectionEnd = selectionStart + lastColon; + target.setSelectionRange(selectionStart, selectionEnd); + return true; + } + + function selectCurrentWord() { + if (selectionStart !== selectionEnd) return false; + const delimiters = `${opts.keyedit_delimiters} \r\n\t`; + // seek backward until to find beggining + while (!delimiters.includes(text[selectionStart - 1]) && selectionStart > 0) selectionStart--; + // seek forward to find end + while (!delimiters.includes(text[selectionEnd]) && selectionEnd < text.length) selectionEnd++; + target.setSelectionRange(selectionStart, selectionEnd); + return true; + } + + // If the user hasn't selected anything, let's select their current parenthesis block or word + if (!selectCurrentParenthesisBlock('<', '>') && !selectCurrentParenthesisBlock('(', ')')) selectCurrentWord(); + event.preventDefault(); + + let closeCharacter = ')'; + let delta = opts.keyedit_precision_attention; + + if (selectionStart > 0 && text[selectionStart - 1] === '<') { + closeCharacter = '>'; + delta = opts.keyedit_precision_extra; + } else if (selectionStart === 0 || text[selectionStart - 1] !== '(') { + while (selectionEnd > selectionStart && text[selectionEnd - 1] === ' ') selectionEnd -= 1; + if (selectionStart === selectionEnd) return; + text = `${text.slice(0, selectionStart)}(${text.slice(selectionStart, selectionEnd)}:1.0)${text.slice(selectionEnd)}`; + selectionStart += 1; + selectionEnd += 1; + } + + const end = text.slice(selectionEnd + 1).indexOf(closeCharacter) + 1; + let weight = parseFloat(text.slice(selectionEnd + 1, selectionEnd + 1 + end)); + if (Number.isNaN(weight)) return; + + weight += isPlus ? delta : -delta; + weight = parseFloat(weight.toPrecision(12)); + if (String(weight).length === 1) weight += '.0'; + + if (closeCharacter == ')' && weight == 1) { + text = text.slice(0, selectionStart - 1) + text.slice(selectionStart, selectionEnd) + text.slice(selectionEnd + 5); + selectionStart--; + selectionEnd--; + } else { + text = text.slice(0, selectionEnd + 1) + weight + text.slice(selectionEnd + 1 + end - 1); + } + + target.focus(); + target.value = text; + target.selectionStart = selectionStart; + target.selectionEnd = selectionEnd; + + updateInput(target); +} + +addEventListener('keydown', (event) => keyupEditAttention(event)); diff --git a/javascript/extensions.js b/javascript/extensions.js index e64a0c795..b827299f8 100644 --- a/javascript/extensions.js +++ b/javascript/extensions.js @@ -1,33 +1,53 @@ - -function extensions_apply(_, _, disable_all){ - var disable = [] - var update = [] - gradioApp().querySelectorAll('#extensions input[type="checkbox"]').forEach(function(x){ - if(x.name.startsWith("enable_") && ! x.checked) disable.push(x.name.substr(7)) - if(x.name.startsWith("update_") && x.checked) update.push(x.name.substr(7)) - }) - restart_reload() - return [JSON.stringify(disable), JSON.stringify(update), disable_all] -} - -function extensions_check(_, _){ - var disable = [] - gradioApp().querySelectorAll('#extensions input[type="checkbox"]').forEach(function(x){ - if(x.name.startsWith("enable_") && ! x.checked) disable.push(x.name.substr(7)) - }) - gradioApp().querySelectorAll('#extensions .extension_status').forEach(function(x){ - x.innerHTML = "Loading..." - }) - var id = randomId() - requestProgress(id, gradioApp().getElementById('extensions_installed_top'), null, null, null, false) - return [id, JSON.stringify(disable)] -} - -function install_extension_from_index(button, url){ - button.disabled = "disabled" - button.value = "Installing..." - textarea = gradioApp().querySelector('#extension_to_install textarea') - textarea.value = url - updateInput(textarea) - gradioApp().querySelector('#install_extension_button').click() -} +function extensions_apply(extensions_disabled_list, extensions_update_list, disable_all) { + console.log('Extensions apply:', extensions_disabled_list, extensions_update_list, disable_all); + const disable = []; + const update = []; + gradioApp().querySelectorAll('#extensions input[type="checkbox"]').forEach((x) => { + if (x.name.startsWith('enable_') && !x.checked) disable.push(x.name.substring(7)); + if (x.name.startsWith('update_') && x.checked) update.push(x.name.substring(7)); + }); + restart_reload(); + return [JSON.stringify(disable), JSON.stringify(update), disable_all]; +} + +function extensions_check(info, extensions_disabled_list, search_text, sort_column) { + console.log('Extensions check:', info, extensions_disabled_list); + const disable = []; + gradioApp().querySelectorAll('#extensions input[type="checkbox"]').forEach((x) => { + if (x.name.startsWith('enable_') && !x.checked) disable.push(x.name.substring(7)); + }); + const id = randomId(); + return [id, JSON.stringify(disable), search_text, sort_column]; +} + +function install_extension(button, url) { + console.log('Extension install:', url); + button.disabled = 'disabled'; + button.value = 'Installing...'; + button.innerHTML = 'installing'; + const textarea = gradioApp().querySelector('#extension_to_install textarea'); + textarea.value = url; + updateInput(textarea); + gradioApp().querySelector('#install_extension_button').click(); +} + +function uninstall_extension(button, url) { + console.log('Extension uninstall:', url, JSON.stringify(url), decodeURIComponent(url), encodeURI(url)); + button.disabled = 'disabled'; + button.value = 'Uninstalling...'; + button.innerHTML = 'uninstalling'; + const textarea = gradioApp().querySelector('#extension_to_install textarea'); + textarea.value = url; + updateInput(textarea); + gradioApp().querySelector('#uninstall_extension_button').click(); +} + +function update_extension(button, url) { + console.log('Extension update:', url); + button.value = 'Updating...'; + button.innerHTML = 'updating'; + const textarea = gradioApp().querySelector('#extension_to_install textarea'); + textarea.value = url; + updateInput(textarea); + gradioApp().querySelector('#update_extension_button').click(); +} diff --git a/javascript/extraNetworks.js b/javascript/extraNetworks.js index bff379beb..a8c54bf7b 100644 --- a/javascript/extraNetworks.js +++ b/javascript/extraNetworks.js @@ -1,170 +1,174 @@ -function setupExtraNetworksForTab(tabname){ - gradioApp().querySelector('#'+tabname+'_extra_tabs').classList.add('extra-networks') - var tabs = gradioApp().querySelector('#'+tabname+'_extra_tabs > div') - var search = gradioApp().querySelector('#'+tabname+'_extra_search textarea') - var refresh = gradioApp().getElementById(tabname+'_extra_refresh') - var descriptInput = gradioApp().getElementById(tabname+ '_description_input') - var close = gradioApp().getElementById(tabname+'_extra_close') - search.classList.add('search') - tabs.appendChild(search) - tabs.appendChild(refresh) - tabs.appendChild(close) - tabs.appendChild(descriptInput) - search.addEventListener("input", function(evt){ - searchTerm = search.value.toLowerCase() - gradioApp().querySelectorAll('#'+tabname+'_extra_tabs div.card').forEach(function(elem){ - text = elem.querySelector('.name').textContent.toLowerCase() + " " + elem.querySelector('.search_term').textContent.toLowerCase() - elem.style.display = text.indexOf(searchTerm) == -1 ? "none" : "" - }) - }); -} - -var activePromptTextarea = {}; - -function setupExtraNetworks(){ - setupExtraNetworksForTab('txt2img') - setupExtraNetworksForTab('img2img') - function registerPrompt(tabname, id){ - var textarea = gradioApp().querySelector("#" + id + " > label > textarea"); - if ( !activePromptTextarea[tabname]) activePromptTextarea[tabname] = textarea - textarea.addEventListener("focus", function(){ - activePromptTextarea[tabname] = textarea; - }); - } - registerPrompt('txt2img', 'txt2img_prompt') - registerPrompt('txt2img', 'txt2img_neg_prompt') - registerPrompt('img2img', 'img2img_prompt') - registerPrompt('img2img', 'img2img_neg_prompt') -} - -onUiLoaded(setupExtraNetworks) -var re_extranet = /<([^:]+:[^:]+):[\d\.]+>/; -var re_extranet_g = /\s+<([^:]+:[^:]+):[\d\.]+>/g; - -function tryToRemoveExtraNetworkFromPrompt(textarea, text){ - var m = text.match(re_extranet) - if(! m) return false - var partToSearch = m[1] - var replaced = false - var newTextareaText = textarea.value.replaceAll(re_extranet_g, function(found, index){ - m = found.match(re_extranet); - if(m[1] == partToSearch){ - replaced = true; - return "" - } - return found; - }) - if(replaced){ - textarea.value = newTextareaText - return true; - } - return false -} - -function cardClicked(tabname, textToAdd, allowNegativePrompt){ - var textarea = allowNegativePrompt ? activePromptTextarea[tabname] : gradioApp().querySelector("#" + tabname + "_prompt > label > textarea") - if (!tryToRemoveExtraNetworkFromPrompt(textarea, textToAdd)) textarea.value = textarea.value + opts.extra_networks_add_text_separator + textToAdd - updateInput(textarea) -} - -function saveCardPreview(event, tabname, filename){ - var textarea = gradioApp().querySelector("#" + tabname + '_preview_filename > label > textarea') - var button = gradioApp().getElementById(tabname + '_save_preview') - textarea.value = filename - updateInput(textarea) - button.click() - event.stopPropagation() - event.preventDefault() -} - -function saveCardDescription(event, tabname, filename, descript){ - var textarea = gradioApp().querySelector("#" + tabname + '_description_filename > label > textarea') - var button = gradioApp().getElementById(tabname + '_save_description') - var description = gradioApp().getElementById(tabname+ '_description_input') - textarea.value = filename - description.value=descript - updateInput(textarea) - button.click() - event.stopPropagation() - event.preventDefault() -} - -function readCardDescription(event, tabname, filename, descript, extraPage, cardName){ - var textarea = gradioApp().querySelector("#" + tabname + '_description_filename > label > textarea') - var description_textarea = gradioApp().querySelector("#" + tabname+ '_description_input > label > textarea') - var button = gradioApp().getElementById(tabname + '_read_description') - textarea.value = filename - description_textarea.value = descript - updateInput(textarea) - updateInput(description_textarea) - button.click() - event.stopPropagation() - event.preventDefault() -} - -function extraNetworksSearchButton(tabs_id, event){ - searchTextarea = gradioApp().querySelector("#" + tabs_id + ' > div > textarea') - button = event.target - text = button.classList.contains("search-all") ? "" : button.textContent.trim() - searchTextarea.value = text - updateInput(searchTextarea) -} - -var globalPopup = null; -var globalPopupInner = null; -function popup(contents){ - if(! globalPopup){ - globalPopup = document.createElement('div') - globalPopup.onclick = function(){ globalPopup.style.display = "none"; }; - globalPopup.classList.add('global-popup'); - var close = document.createElement('div') - close.classList.add('global-popup-close'); - close.onclick = function(){ globalPopup.style.display = "none"; }; - close.title = "Close"; - globalPopup.appendChild(close) - globalPopupInner = document.createElement('div') - globalPopupInner.onclick = function(event){ event.stopPropagation(); return false; }; - globalPopupInner.classList.add('global-popup-inner'); - globalPopup.appendChild(globalPopupInner) - gradioApp().appendChild(globalPopup); - } - globalPopupInner.innerHTML = ''; - globalPopupInner.appendChild(contents); - globalPopup.style.display = "flex"; -} - -function readCardMetadata(event, extraPage, cardName){ - requestGet("./sd_extra_networks/metadata", {"page": extraPage, "item": cardName}, function(data){ - if (data && data.metadata){ - elem = document.createElement('pre') - elem.classList.add('popup-metadata'); - elem.textContent = data.metadata; - popup(elem); - } - }, () => {}) - event.stopPropagation() - event.preventDefault() -} - -function requestGet(url, data, handler, errorHandler){ - var xhr = new XMLHttpRequest(); - var args = Object.keys(data).map(function(k){ return encodeURIComponent(k) + '=' + encodeURIComponent(data[k]) }).join('&') - xhr.open("GET", url + "?" + args, true); - xhr.onreadystatechange = function () { - if (xhr.readyState === 4) { - if (xhr.status === 200) { - try { - var js = JSON.parse(xhr.responseText); - handler(js) - } catch (error) { - console.error(error); - errorHandler() - } - } else{ - errorHandler() - } - } - }; - var js = JSON.stringify(data); - xhr.send(js); -} +function setupExtraNetworksForTab(tabname) { + gradioApp().querySelector(`#${tabname}_extra_tabs`).classList.add('extra-networks'); + const tabs = gradioApp().querySelector(`#${tabname}_extra_tabs > div`); + const search = gradioApp().querySelector(`#${tabname}_extra_search textarea`); + const refresh = gradioApp().getElementById(`${tabname}_extra_refresh`); + const descriptInput = gradioApp().getElementById(`${tabname}_description_input`); + const close = gradioApp().getElementById(`${tabname}_extra_close`); + search.classList.add('search'); + tabs.appendChild(search); + tabs.appendChild(refresh); + tabs.appendChild(close); + tabs.appendChild(descriptInput); + search.addEventListener('input', (evt) => { + searchTerm = search.value.toLowerCase(); + gradioApp().querySelectorAll(`#${tabname}_extra_tabs div.card`).forEach((elem) => { + text = `${elem.querySelector('.name').textContent.toLowerCase()} ${elem.querySelector('.search_term').textContent.toLowerCase()}`; + elem.style.display = text.indexOf(searchTerm) == -1 ? 'none' : ''; + }); + }); +} + +const activePromptTextarea = {}; + +function setupExtraNetworks() { + setupExtraNetworksForTab('txt2img'); + setupExtraNetworksForTab('img2img'); + function registerPrompt(tabname, id) { + const textarea = gradioApp().querySelector(`#${id} > label > textarea`); + if (!activePromptTextarea[tabname]) activePromptTextarea[tabname] = textarea; + textarea.addEventListener('focus', () => { + activePromptTextarea[tabname] = textarea; + }); + } + registerPrompt('txt2img', 'txt2img_prompt'); + registerPrompt('txt2img', 'txt2img_neg_prompt'); + registerPrompt('img2img', 'img2img_prompt'); + registerPrompt('img2img', 'img2img_neg_prompt'); +} + +onUiLoaded(setupExtraNetworks); +const re_extranet = /<([^:]+:[^:]+):[\d\.]+>/; +const re_extranet_g = /\s+<([^:]+:[^:]+):[\d\.]+>/g; + +function tryToRemoveExtraNetworkFromPrompt(textarea, text) { + let m = text.match(re_extranet); + if (!m) return false; + const partToSearch = m[1]; + let replaced = false; + const newTextareaText = textarea.value.replaceAll(re_extranet_g, (found, index) => { + m = found.match(re_extranet); + if (m[1] === partToSearch) { + replaced = true; + return ''; + } + return found; + }); + if (replaced) { + textarea.value = newTextareaText; + return true; + } + return false; +} + +function refreshExtraNetworks(tabname) { + gradioApp().querySelector(`#${tabname}_extra_networks textarea`)?.dispatchEvent(new Event('input')); +} + +function cardClicked(tabname, textToAdd, allowNegativePrompt) { + const textarea = allowNegativePrompt ? activePromptTextarea[tabname] : gradioApp().querySelector(`#${tabname}_prompt > label > textarea`); + if (!tryToRemoveExtraNetworkFromPrompt(textarea, textToAdd)) textarea.value = textarea.value + opts.extra_networks_add_text_separator + textToAdd; + updateInput(textarea); +} + +function saveCardPreview(event, tabname, filename) { + const textarea = gradioApp().querySelector(`#${tabname}_preview_filename > label > textarea`); + const button = gradioApp().getElementById(`${tabname}_save_preview`); + textarea.value = filename; + updateInput(textarea); + button.click(); + event.stopPropagation(); + event.preventDefault(); +} + +function saveCardDescription(event, tabname, filename, descript) { + const textarea = gradioApp().querySelector(`#${tabname}_description_filename > label > textarea`); + const button = gradioApp().getElementById(`${tabname}_save_description`); + const description = gradioApp().getElementById(`${tabname}_description_input`); + textarea.value = filename; + description.value = descript; + updateInput(textarea); + button.click(); + event.stopPropagation(); + event.preventDefault(); +} + +function readCardDescription(event, tabname, filename, descript, extraPage, cardName) { + const textarea = gradioApp().querySelector(`#${tabname}_description_filename > label > textarea`); + const description_textarea = gradioApp().querySelector(`#${tabname}_description_input > label > textarea`); + const button = gradioApp().getElementById(`${tabname}_read_description`); + textarea.value = filename; + description_textarea.value = descript; + updateInput(textarea); + updateInput(description_textarea); + button.click(); + event.stopPropagation(); + event.preventDefault(); +} + +function extraNetworksSearchButton(tabs_id, event) { + searchTextarea = gradioApp().querySelector(`#${tabs_id} > div > textarea`); + button = event.target; + text = button.classList.contains('search-all') ? '' : button.textContent.trim(); + searchTextarea.value = text; + updateInput(searchTextarea); +} + +let globalPopup = null; +let globalPopupInner = null; +function popup(contents) { + if (!globalPopup) { + globalPopup = document.createElement('div'); + globalPopup.onclick = function () { globalPopup.style.display = 'none'; }; + globalPopup.classList.add('global-popup'); + const close = document.createElement('div'); + close.classList.add('global-popup-close'); + close.onclick = function () { globalPopup.style.display = 'none'; }; + close.title = 'Close'; + globalPopup.appendChild(close); + globalPopupInner = document.createElement('div'); + globalPopupInner.onclick = function (event) { event.stopPropagation(); return false; }; + globalPopupInner.classList.add('global-popup-inner'); + globalPopup.appendChild(globalPopupInner); + gradioApp().appendChild(globalPopup); + } + globalPopupInner.innerHTML = ''; + globalPopupInner.appendChild(contents); + globalPopup.style.display = 'flex'; +} + +function readCardMetadata(event, extraPage, cardName) { + requestGet('./sd_extra_networks/metadata', { page: extraPage, item: cardName }, (data) => { + if (data && data.metadata) { + elem = document.createElement('pre'); + elem.classList.add('popup-metadata'); + elem.textContent = data.metadata; + popup(elem); + } + }, () => {}); + event.stopPropagation(); + event.preventDefault(); +} + +function requestGet(url, data, handler, errorHandler) { + const xhr = new XMLHttpRequest(); + const args = Object.keys(data).map((k) => `${encodeURIComponent(k)}=${encodeURIComponent(data[k])}`).join('&'); + xhr.open('GET', `${url}?${args}`, true); + xhr.onreadystatechange = function () { + if (xhr.readyState === 4) { + if (xhr.status === 200) { + try { + const js = JSON.parse(xhr.responseText); + handler(js); + } catch (error) { + console.error(error); + errorHandler(); + } + } else { + errorHandler(); + } + } + }; + const js = JSON.stringify(data); + xhr.send(js); +} diff --git a/javascript/generationParams.js b/javascript/generationParams.js index 5d4f996bf..4d9b06f75 100644 --- a/javascript/generationParams.js +++ b/javascript/generationParams.js @@ -1,30 +1,41 @@ +/* global gradioApp, onUiUpdate */ // attaches listeners to the txt2img and img2img galleries to update displayed generation param text when the image changes -let txt2img_gallery, img2img_gallery, modal = undefined; -onUiUpdate(function(){ - if (!txt2img_gallery) txt2img_gallery = attachGalleryListeners("txt2img") - if (!img2img_gallery) img2img_gallery = attachGalleryListeners("img2img") - if (!modal) { - modal = gradioApp().getElementById('lightboxModal') - modalObserver.observe(modal, { attributes : true, attributeFilter : ['style'] }); - } -}); - -let modalObserver = new MutationObserver(function(mutations) { - mutations.forEach((mutationRecord) => { - let selectedTab = gradioApp().querySelector('#tabs div button.selected')?.innerText - if (!selectedTab) selectedTab = gradioApp().querySelector('#tabs div button')?.innerText - if (mutationRecord.target.style.display === 'none' && (selectedTab === 'txt2img' || selectedTab === 'img2img')) - gradioApp().getElementById(selectedTab+"_generation_info_button")?.click() - }); -}); - function attachGalleryListeners(tab_name) { - gallery = gradioApp().querySelector('#'+tab_name+'_gallery') - gallery?.addEventListener('click', () => gradioApp().getElementById(tab_name+"_generation_info_button").click()); - gallery?.addEventListener('keydown', (e) => { - if (e.keyCode == 37 || e.keyCode == 39) // left or right arrow - gradioApp().getElementById(tab_name+"_generation_info_button").click() - }); - return gallery; + const gallery = gradioApp().querySelector(`#${tab_name}_gallery`); + gallery?.addEventListener('click', () => setTimeout(() => { + gradioApp() + .getElementById(`${tab_name}_generation_info_button`) + ?.click(); + }, 500)); + gallery?.addEventListener('keydown', (e) => { + if (e.keyCode == 37 || e.keyCode == 39) { // left or right arrow + gradioApp() + .getElementById(`${tab_name}_generation_info_button`) + .click(); + } + }); + return gallery; } + +let txt2img_gallery; +let img2img_gallery; +let modal; + +onUiUpdate(() => { + if (!txt2img_gallery) txt2img_gallery = attachGalleryListeners('txt2img'); + if (!img2img_gallery) img2img_gallery = attachGalleryListeners('img2img'); + if (!modal) { + modal = gradioApp().getElementById('lightboxModal'); + modalObserver.observe(modal, { attributes: true, attributeFilter: ['style'] }); + } +}); + +let modalObserver = new MutationObserver((mutations) => { + mutations.forEach((mutationRecord) => { + let selectedTab = gradioApp().querySelector('#tabs div button.selected')?.innerText; + if (!selectedTab) selectedTab = gradioApp().querySelector('#tabs div button')?.innerText; + if (mutationRecord.target.style.display === 'none' && (selectedTab === 'txt2img' || selectedTab === 'img2img')) { gradioApp().getElementById(`${selectedTab}_generation_info_button`)?.click(); } + }); +}); + diff --git a/javascript/hints.js b/javascript/hints.js index f48a0eb69..026836b67 100644 --- a/javascript/hints.js +++ b/javascript/hints.js @@ -1,147 +1,97 @@ -// mouseover tooltips for various UI elements +// HTML tooltips for various UI elements titles = { - "Sampling steps": "How many times to improve the generated image iteratively; higher values take longer; very low values can produce bad results", - "Sampling method": "Which algorithm to use to produce the image", - "GFPGAN": "Restore low quality faces using GFPGAN neural network", - "Euler a": "Euler Ancestral - very creative, each can get a completely different picture depending on step count, setting steps higher than 30-40 does not help", - "DDIM": "Denoising Diffusion Implicit Models - best at inpainting", - "UniPC": "Unified Predictor-Corrector Framework for Fast Sampling of Diffusion Models", - "DPM adaptive": "Ignores step count - uses a number of steps determined by the CFG and resolution", - - "Batch count": "How many batches of images to create (has no impact on generation performance or VRAM usage)", - "Batch size": "How many image to create in a single batch (increases generation performance at cost of higher VRAM usage)", - "CFG Scale": "Classifier Free Guidance Scale - how strongly the image should conform to prompt - lower values produce more creative results", - "Seed": "A value that determines the output of random number generator - if you create an image with same parameters and seed as another image, you'll get the same result", - "\u{1f3b2}\ufe0f": "Set seed to -1, which will cause a new random number to be used every time", - "\u267b\ufe0f": "Reuse seed from last generation, mostly useful if it was randomed", - "\u2199\ufe0f": "Read generation parameters from prompt or last generation if prompt is empty into user interface.", - "\u{1f4c2}": "Open images output directory", - "\u{1f4be}": "Save style", - "\u{1f5d1}\ufe0f": "Clear prompt", - "\u{1f4cb}": "Apply selected styles to current prompt", - "\u{1f4d2}": "Paste available values into the field", - "\u{1f3b4}": "Show/hide extra networks", - - "Inpaint a part of image": "Draw a mask over an image, and the script will regenerate the masked area with content according to prompt", - "SD upscale": "Upscale image normally, split result into tiles, improve each tile using img2img, merge whole image back", - - "Just resize": "Resize image to target resolution. Unless height and width match, you will get incorrect aspect ratio.", - "Crop and resize": "Resize the image so that entirety of target resolution is filled with the image. Crop parts that stick out.", - "Resize and fill": "Resize the image so that entirety of image is inside target resolution. Fill empty space with image's colors.", - - "Mask blur": "How much to blur the mask before processing, in pixels.", - "Masked content": "What to put inside the masked area before processing it with Stable Diffusion.", - "fill": "fill it with colors of the image", - "original": "keep whatever was there originally", - "latent noise": "fill it with latent space noise", - "latent nothing": "fill it with latent space zeroes", - "Inpaint at full resolution": "Upscale masked region to target resolution, do inpainting, downscale back and paste into original image", - - "Denoising strength": "Determines how little respect the algorithm should have for image's content. At 0, nothing will change, and at 1 you'll get an unrelated image. With values below 1.0, processing will take less steps than the Sampling Steps slider specifies.", - - "Skip": "Stop processing current image and continue processing.", - "Interrupt": "Stop processing images and return any results accumulated so far.", - "Save": "Write image to a directory (default - log/images) and generation parameters into csv file.", - - "X values": "Separate values for X axis using commas.", - "Y values": "Separate values for Y axis using commas.", - - "None": "Do not do anything special", - "Prompt matrix": "Separate prompts into parts using vertical pipe character (|) and the script will create a picture for every combination of them (except for the first part, which will be present in all combinations)", - "X/Y/Z plot": "Create grid(s) where images will have different parameters. Use inputs below to specify which parameters will be shared by columns and rows", - "Custom code": "Run Python code. Advanced user only. Must run program with --allow-code for this to work", - - "Prompt S/R": "Separate a list of words with commas, and the first word will be used as a keyword: script will search for this word in the prompt, and replace it with others", - "Prompt order": "Separate a list of words with commas, and the script will make a variation of prompt with those words for their every possible order", - - "Tiling": "Produce an image that can be tiled.", - "Tile overlap": "For SD upscale, how much overlap in pixels should there be between tiles. Tiles overlap so that when they are merged back into one picture, there is no clearly visible seam.", - - "Variation seed": "Seed of a different picture to be mixed into the generation.", - "Variation strength": "How strong of a variation to produce. At 0, there will be no effect. At 1, you will get the complete picture with variation seed (except for ancestral samplers, where you will just get something).", - "Resize seed from height": "Make an attempt to produce a picture similar to what would have been produced with same seed at specified resolution", - "Resize seed from width": "Make an attempt to produce a picture similar to what would have been produced with same seed at specified resolution", - - "Interrogate": "Reconstruct prompt from existing image and put it into the prompt field.", - - "Images filename pattern": "Use following tags to define how filenames for images are chosen: [steps], [cfg], [prompt_hash], [prompt], [prompt_no_styles], [prompt_spaces], [width], [height], [styles], [sampler], [seed], [model_hash], [model_name], [prompt_words], [date], [datetime], [datetime], [datetime