diff --git a/.eslintrc.json b/.eslintrc.json deleted file mode 100644 index 77b6fb242..000000000 --- a/.eslintrc.json +++ /dev/null @@ -1,138 +0,0 @@ -{ - "root": true, - "parserOptions": { - "ecmaVersion": "latest", - "sourceType": "module" - }, - "plugins": [ - "html", - "json" - ], - "extends": [ - "eslint:recommended", - "airbnb-base", - "plugin:css/recommended", - "plugin:json/recommended", - "plugin:node/recommended" - ], - "env": { - "browser": true, - "commonjs": true, - "node": true, - "jquery": true, - "es2024": true - }, - "rules": { - "max-len": [1, 275, 3], - "camelcase":"off", - "default-case":"off", - "no-await-in-loop":"off", - "no-bitwise":"off", - "no-continue":"off", - "no-confusing-arrow":"off", - "no-console":"off", - "no-empty":"off", - "no-loop-func":"off", - "no-mixed-operators":"off", - "no-param-reassign":"off", - "no-process-exit":"off", - "no-plusplus":"off", - "no-restricted-globals":"off", - "no-restricted-syntax":"off", - "no-return-assign":"off", - "no-unused-vars":"off", - "no-useless-escape":"off", - "object-curly-newline":"off", - "prefer-rest-params":"off", - "prefer-destructuring":"off", - "radix":"off", - "node/shebang": "off" - }, - "globals": { - "panzoom": "readonly", - "authFetch": "readonly", - "log": "readonly", - "debug": "readonly", - "error": "readonly", - "xhrGet": "readonly", - "xhrPost": "readonly", - "gradioApp": "readonly", - "executeCallbacks": "readonly", - "onAfterUiUpdate": "readonly", - "onOptionsChanged": "readonly", - "optionsChangedCallbacks": "readonly", - "onUiLoaded": "readonly", - "onUiUpdate": "readonly", - "onUiTabChange": "readonly", - "onUiReady": "readonly", - "uiCurrentTab": "writable", - "uiElementIsVisible": "readonly", - "uiElementInSight": "readonly", - "getUICurrentTabContent": "readonly", - "waitForFlag": "readonly", - "logFn": "readonly", - "generateForever": "readonly", - "showContributors": "readonly", - "opts": "writable", - "sortUIElements": "readonly", - "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", - "restartReload": "readonly", - "markSelectedCards": "readonly", - "updateInput": "readonly", - "toggleCompact": "readonly", - "setFontSize": "readonly", - "setTheme": "readonly", - "registerDragDrop": "readonly", - "getToken": "readonly", - "getENActiveTab": "readonly", - "quickApplyStyle": "readonly", - "quickSaveStyle": "readonly", - "setupExtraNetworks": "readonly", - "showNetworks": "readonly", - "localization": "readonly", - "randomId": "readonly", - "requestProgress": "readonly", - "setRefreshInterval": "readonly", - "modalPrevImage": "readonly", - "modalNextImage": "readonly", - "galleryClickEventHandler": "readonly", - "getExif": "readonly", - "jobStatusEl": "readonly", - "removeSplash": "readonly", - "initGPU": "readonly", - "startGPU": "readonly", - "disableNVML": "readonly", - "idbGet": "readonly", - "idbPut": "readonly", - "idbDel": "readonly", - "idbAdd": "readonly", - "idbCount": "readonly", - "idbFolderCleanup": "readonly", - "initChangelog": "readonly", - "sendNotification": "readonly", - "monitorConnection": "readonly" - }, - "ignorePatterns": [ - "node_modules", - "extensions", - "repositories", - "venv", - "panzoom.js", - "split.js", - "exifr.js", - "jquery.js", - "sparkline.js", - "iframeResizer.min.js" - ] -} diff --git a/.vscode/settings.json b/.vscode/settings.json index d62ba10b8..924535e5a 100644 --- a/.vscode/settings.json +++ b/.vscode/settings.json @@ -2,5 +2,14 @@ "python.analysis.extraPaths": [".", "./modules", "./scripts", "./pipelines"], "python.analysis.typeCheckingMode": "off", "editor.formatOnSave": false, - "python.REPL.enableREPLSmartSend": false + "python.REPL.enableREPLSmartSend": false, + "eslint.enable": true, + "eslint.validate": [ + "javascript", + "typescript", + "html", + "css", + "json", + "markdown" + ] } diff --git a/CHANGELOG.md b/CHANGELOG.md index ddcb0bb39..7f20d52bc 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,115 @@ # Change Log for SD.Next +## Update for 2025-01-20 + +### Highlights for 2025-01-20 + +First release of 2026 brings quite a few new models: **Flux.2-Klein, Qwen-Image-2512, LTX-2-Dev, GLM-Image** +There are also improvements to *SDNQ* quantization engine, updated *Prompt Enhance*, *Image Preview* and many others. +Plus some significant under-the-hood changes to improve code coverage and quality which resulted in more than usual levels of bug-fixes and some ~330 commits! +For full list of changes, see full changelog. + +[ReadMe](https://github.com/vladmandic/automatic/blob/master/README.md) | [ChangeLog](https://github.com/vladmandic/automatic/blob/master/CHANGELOG.md) | [Docs](https://vladmandic.github.io/sdnext-docs/) | [WiKi](https://github.com/vladmandic/automatic/wiki) | [Discord](https://discord.com/invite/sd-next-federal-batch-inspectors-1101998836328697867) | [Sponsor](https://github.com/sponsors/vladmandic) + +### Details for 2025-01-20 + +- **Models** + - [Flux.2 Klein](https://bfl.ai/blog/flux2-klein-towards-interactive-visual-intelligence) + Flux.2-Klein is a new family of compact models from BFL in *4B and 9B sizes* and avaialable as *destilled and base* variants + also includes are *sdnq prequantized variants* + *note*: 9B variant is [gated](https://vladmandic.github.io/sdnext-docs/Gated/) + - [Qwen-Image-2512](https://qwen.ai/blog?id=qwen-image-2512) + Qwen-Image successor, significantly reduces the AI-generated look and adds finer natural detailils and improved text rendering + available in both *original*, *sdnq-svd prequantized* and *sdnq-dynamic prequantized* variants + thanks @CalamitousFelicitousness + - [LTX-2 19B Dev](https://ltx.io/model/ltx-2) + LTX-2 is a new very large 19B parameter video generation model from Lightricks using Gemma-3 text encoder + available for T2I/I2I workflows in original and sdnq prequantized variants + *note*: model is very sensitive to input params and will result in errors otherwise + - [GLM-Image](https://z.ai/blog/glm-image) + GLM-image is a new image generation model that adopts a hybrid autoregressive with diffusion decoder architecture + available in both *original* and *sdnq-dynamic prequantized* variants + thanks @CalamitousFelicitousness + *note*: model requires pre-release versions of `transformers` package: + > pip install --upgrade git+https://github.com/huggingface/transformers.git + > ./webui.sh --experimental + - [Nunchaku Z-Image Turbo](https://huggingface.co/nunchaku-tech/nunchaku-z-image-turbo) + nunchaku optimized z-image turbo +- **Feaures** + - **SDNQ**: add *dynamic* quantization method + sdnq can dynamically determine best quantization method for each module layer + slower to quantize on-the-fly, but results in better quality with minimal resource usage + - **SDNQ** now has *19 int* based and *69 float* based quantization types + *note*: not all are exposed via ui purely for simplicity, but all are available via api and scripts + - **wildcards**: allow weights, thanks @Tillerz + - **sampler**: add laplace beta schedule + results in better prompt adherence and smoother infills + - **prompt enhance**: improve handling and refresh ui, thanks @CalamitousFelicitousness + new models such moondream-3 and xiaomo-mimo + add support for *thinking* mode where model can reason about the prompt + add support for *vision* processing where prompt enhance can also optionally analyze input image + add support for *pre-fill* mode where prompt enhance can continue from existing caption + - **chroma**: add inpaint pipeline support + - **taesd preview**: support for more models, thanks @alerikaisattera + - **image ouput paths**: better handling of relative/absolute paths, thanks @CalamitousFelicitousness +- **UI** + - kanvas add send-to functionality + - kanvas improve support for standardui + - improve extensions tab layout and behavior, thanks @awsr + - indicate collapsed/hidden sections + - persistent panel minimize/maximize state + - gallery improve sorting behavior + - gallery implement prev/next navigation in full screen viewer, thanks @ryanmeador +- **Internal** + - **lora** native support by default will now skip text-encoder + can be enabled in *settings -> networks* + - update core js linting to `eslint9`, thanks @awsr + - update modernui js linting to `eslint9`, thanks @awsr + - update kanvas js linting to `eslint9`, thanks @awsr + - update strong typing checks, thanks @awsr + - update reference models previews, thanks @liutyi + - update models specs page, thanks @alerikaisattera + - sdnq improvements + - startup sequence optimizations + - rocm/hip/hipblast detection and initialization improvements + - zluda detection and initialization improvements + - new env variable `SD_VAE_DEFAULT` to force default vae processing + - update `nunchaku==1.1.0` + - lora switch logic from force-diffusers to allow-native + - split `reference.json` + - print system env on startup + - disable fallback on models with custom loaders + - refactor triggering of prompt parser and set secondary prompts when needed + - refactor handling of seeds + - allow unsafe ssl context for downloads +- **Fixes** + - controlnet: controlnet with non-english ui locales + - core: add skip_keys to offloading logic, fixes wan frames mismatch, thanks @ryanmeador + - core: force model move on offload=none + - core: hidiffusion tracing + - core: hip device name detection + - core: reduce triton test verbosity + - core: switch processing class not restoring params + - extension tab: update checker, date handling, formatting etc., thanks @awsr + - lora force unapply on change + - lora handle null description, thanks @CalamitousFelicitousness + - lora loading when using torch without distributed support + - lora skip with strength zero + - lora: generate slowdown when consequtive lora-diffusers enabled + - model: google-genai auth, thanks @CalamitousFelicitousness + - model: improve qwen i2i handling + - model: kandinsky-5 image and video on non-cuda platforms + - model: meituan-longca-image-edit missing image param + - model: wan 2.2 i2v + - model: z-image single-file loader + - other: update civitai base models, thanks @trojaner + - ui: gallery save/delete + - ui: mobile auto-collapse when using side panel, thanks @awsr + - ui: networks filter by model type + - ui: networks icon/list view type switch, thanks @awsr + - vae: force align width/height to vae scale factor + - wildards with folder specification + ## Update for 2025-12-26 ### Highlights for 2025-12-26 diff --git a/TODO.md b/TODO.md index 626fc7a6e..4c3f94f73 100644 --- a/TODO.md +++ b/TODO.md @@ -1,52 +1,54 @@ # TODO -## Known issues - -- z-image-turbo controlnet device mismatch: -- z-image-turbo safetensors loader: -- kandinsky-image-5 hardcoded cuda: -- peft lora with torch-rocm-windows: - ## Project Board - ## Internal -- Reimplement `llama` remover for kanvas -- Deploy: Create executable for SD.Next -- Feature: Integrate natural language image search - [ImageDB](https://github.com/vladmandic/imagedb) -- Feature: Transformers unified cache handler -- Feature: Remote Text-Encoder support -- Refactor: [Modular pipelines and guiders](https://github.com/huggingface/diffusers/issues/11915) -- Refactor: move sampler options to settings to config -- Refactor: [GGUF](https://huggingface.co/docs/diffusers/main/en/quantization/gguf) -- Feature: LoRA add OMI format support for SD35/FLUX.1 +- Feature: Move `nunchaku` models to refernce instead of internal decision +- Update: `transformers==5.0.0` +- Feature: Unify *huggingface* and *diffusers* model folders +- Reimplement `llama` remover for Kanvas +- Deploy: Create executable for SD.Next +- Feature: Integrate natural language image search + [ImageDB](https://github.com/vladmandic/imagedb) +- Feature: Remote Text-Encoder support +- Refactor: move sampler options to settings to config +- Refactor: [GGUF](https://huggingface.co/docs/diffusers/main/en/quantization/gguf) +- Feature: LoRA add OMI format support for SD35/FLUX.1 - Refactor: remove `CodeFormer` -- Refactor: remove `GFPGAN` -- UI: Lite vs Expert mode -- Video tab: add full API support -- Control tab: add overrides handling +- Refactor: remove `GFPGAN` +- UI: Lite vs Expert mode +- Video tab: add full API support +- Control tab: add overrides handling - Engine: `TensorRT` acceleration - Engine: [mmgp](https://github.com/deepbeepmeep/mmgp) - Engine: [sharpfin](https://github.com/drhead/sharpfin) instead of `torchvision` +## Modular + +- Switch to modular pipelines +- Feature: Transformers unified cache handler +- Refactor: [Modular pipelines and guiders](https://github.com/huggingface/diffusers/issues/11915) +- [MagCache](https://github.com/lllyasviel/FramePack/pull/673/files) +- [SmoothCache](https://github.com/huggingface/diffusers/issues/11135) + ## Features - [Flux.2 TinyVAE](https://huggingface.co/fal/FLUX.2-Tiny-AutoEncoder) -- [IPAdapter composition](https://huggingface.co/ostris/ip-composition-adapter) -- [IPAdapter negative guidance](https://github.com/huggingface/diffusers/discussions/7167) -- [MagCache](https://github.com/lllyasviel/FramePack/pull/673/files) -- [SmoothCache](https://github.com/huggingface/diffusers/issues/11135) -- [STG](https://github.com/huggingface/diffusers/blob/main/examples/community/README.md#spatiotemporal-skip-guidance) +- [IPAdapter composition](https://huggingface.co/ostris/ip-composition-adapter) +- [IPAdapter negative guidance](https://github.com/huggingface/diffusers/discussions/7167) +- [STG](https://github.com/huggingface/diffusers/blob/main/examples/community/README.md#spatiotemporal-skip-guidance) - [Video Inpaint Pipeline](https://github.com/huggingface/diffusers/pull/12506) - [Sonic Inpaint](https://github.com/ubc-vision/sonic) ### New models / Pipelines -TODO: *Prioritize*! +TODO: Investigate which models are diffusers-compatible and prioritize! +- [Bria FiboEdit](https://github.com/huggingface/diffusers/commit/d7a1c31f4f85bae5a9e01cdce49bd7346bd8ccd6) +- [LTXVideo 0.98 LongMulti](https://github.com/huggingface/diffusers/pull/12614) - [Cosmos-Predict-2.5](https://huggingface.co/nvidia/Cosmos-Predict2.5-2B) - [NewBie Image Exp0.1](https://github.com/huggingface/diffusers/pull/12803) - [Sana-I2V](https://github.com/huggingface/diffusers/pull/12634#issuecomment-3540534268) @@ -54,7 +56,8 @@ TODO: *Prioritize*! - [Bytedance Lynx](https://github.com/bytedance/lynx) - [ByteDance OneReward](https://github.com/bytedance/OneReward) - [ByteDance USO](https://github.com/bytedance/USO) -- [Chroma1 Radiance](https://huggingface.co/lodestones/Chroma1-Radiance) +- [Chroma Radiance](https://huggingface.co/lodestones/Chroma1-Radiance) +- [Chroma Zeta](https://huggingface.co/lodestones/Zeta-Chroma) - [DiffSynth Studio](https://github.com/modelscope/DiffSynth-Studio) - [DiffusionForcing](https://github.com/kwsong0113/diffusion-forcing-transformer) - [Dream0 guidance](https://huggingface.co/ByteDance/DreamO) diff --git a/cli/api-pulid.js b/cli/api-pulid.js index 033824e9b..23e4eb094 100755 --- a/cli/api-pulid.js +++ b/cli/api-pulid.js @@ -2,9 +2,9 @@ // simple nodejs script to test sdnext api -const fs = require('fs'); -const path = require('path'); -const process = require('process'); +const fs = require('node:fs'); +const path = require('node:path'); + const argparse = require('argparse'); const sd_url = process.env.SDAPI_URL || 'http://127.0.0.1:7860'; diff --git a/cli/api-txt2img.js b/cli/api-txt2img.js index 7b0f6994a..093943c53 100755 --- a/cli/api-txt2img.js +++ b/cli/api-txt2img.js @@ -2,8 +2,7 @@ // simple nodejs script to test sdnext api -const fs = require('fs'); // eslint-disable-line no-undef -const process = require('process'); // eslint-disable-line no-undef +const fs = require('node:fs'); const sd_url = process.env.SDAPI_URL || 'http://127.0.0.1:7860'; const sd_username = process.env.SDAPI_USR; diff --git a/cli/download.py b/cli/download.py index c7a34d4b3..3abef8c29 100755 --- a/cli/download.py +++ b/cli/download.py @@ -40,7 +40,6 @@ def download_urllib(args): fn = '' req = urllib.request.Request(args.url, headers=headers) res = urllib.request.urlopen(req) - res.getheader('content-length') content_length = int(res.getheader('content-length') or 0) fn = get_filename(args, res) print(f'downloading: url={args.url} file={fn} size={content_length if content_length > 0 else "unknown"} lib=urllib block={args.block}') diff --git a/cli/hf-search.py b/cli/hf-search.py index 9ee696602..638af0aa1 100755 --- a/cli/hf-search.py +++ b/cli/hf-search.py @@ -9,10 +9,19 @@ if __name__ == "__main__": keyword = sys.argv[0] if len(sys.argv) > 0 else '' hf.logging.set_verbosity_info() hf_api = hf.HfApi() - res = hf_api.list_models(model_name=keyword, full=True, limit=100, sort="downloads", direction=-1) + res = hf_api.list_models( + model_name=keyword, + full=True, + limit=100, + sort="downloads", + direction=-1, + ) res = sorted(res, key=lambda x: x.id) + exact = [m for m in res if keyword.lower() in m.id.lower()] + if len(exact) > 0: + res = exact for m in res: meta = hf_api.model_info(m.id, files_metadata=True) m.files = [f.rfilename for f in meta.siblings if f.rfilename.endswith('.bin') or f.rfilename.endswith('.safetensors')] - m.size = sum([f.size for f in meta.siblings]) / 1024 / 1024 / 1024 # in GB - print({ 'name': m.id, 'files': len(m.files), 'size': m.size, 'downloads': m.downloads, 'mtime': m.lastModified, 'url': f'https://huggingface.co/{m.id}', 'pipeline': m.pipeline_tag }) + m.size = round(sum([f.size for f in meta.siblings]) / 1024 / 1024 / 1024, 2) + print({ 'name': m.id, 'files': len(m.files), 'size': m.size, 'downloads': m.downloads, 'ctime': m.created_at.isoformat(), 'url': f'https://huggingface.co/{m.id}', 'pipeline': m.pipeline_tag }) diff --git a/cli/localize.js b/cli/localize.js index 12d289867..48c869185 100755 --- a/cli/localize.js +++ b/cli/localize.js @@ -1,8 +1,8 @@ #!/usr/bin/env node // script used to localize sdnext ui and hints to multiple languages using google gemini ai -const fs = require('fs'); -const process = require('process'); +const fs = require('node:fs'); + const { GoogleGenerativeAI } = require('@google/generative-ai'); const api_key = process.env.GOOGLE_AI_API_KEY; diff --git a/eslint.config.mjs b/eslint.config.mjs new file mode 100644 index 000000000..fddda6ca1 --- /dev/null +++ b/eslint.config.mjs @@ -0,0 +1,339 @@ +import path from 'node:path'; + +import { includeIgnoreFile } from '@eslint/compat'; +import css from '@eslint/css'; +import js from '@eslint/js'; +import json from '@eslint/json'; +import markdown from '@eslint/markdown'; +import html from '@html-eslint/eslint-plugin'; +import { configs, helpers, plugins, rules } from 'eslint-config-airbnb-extended'; +import pluginPromise from 'eslint-plugin-promise'; +import { defineConfig, globalIgnores } from 'eslint/config'; +import globals from 'globals'; + +const gitignorePath = path.resolve('.', '.gitignore'); + +const jsConfig = defineConfig([ + // ESLint recommended config + { + name: 'js/config', + files: helpers.extensions.allFiles, + ...js.configs.recommended, + languageOptions: { + ecmaVersion: 'latest', + parserOptions: { + ecmaVersion: 'latest', + }, + globals: { // Set per project + ...globals.builtin, + ...globals.browser, + ...globals.jquery, + panzoom: 'readonly', + authFetch: 'readonly', + log: 'readonly', + debug: 'readonly', + error: 'readonly', + xhrGet: 'readonly', + xhrPost: 'readonly', + gradioApp: 'readonly', + executeCallbacks: 'readonly', + onAfterUiUpdate: 'readonly', + onOptionsChanged: 'readonly', + optionsChangedCallbacks: 'readonly', + onUiLoaded: 'readonly', + onUiUpdate: 'readonly', + onUiTabChange: 'readonly', + onUiReady: 'readonly', + uiCurrentTab: 'writable', + uiElementIsVisible: 'readonly', + uiElementInSight: 'readonly', + getUICurrentTabContent: 'readonly', + waitForFlag: 'readonly', + logFn: 'readonly', + generateForever: 'readonly', + showContributors: 'readonly', + opts: 'writable', + sortUIElements: 'readonly', + 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', + restartReload: 'readonly', + markSelectedCards: 'readonly', + updateInput: 'readonly', + toggleCompact: 'readonly', + setFontSize: 'readonly', + setTheme: 'readonly', + registerDragDrop: 'readonly', + getToken: 'readonly', + getENActiveTab: 'readonly', + quickApplyStyle: 'readonly', + quickSaveStyle: 'readonly', + setupExtraNetworks: 'readonly', + showNetworks: 'readonly', + localization: 'readonly', + randomId: 'readonly', + requestProgress: 'readonly', + setRefreshInterval: 'readonly', + modalPrevImage: 'readonly', + modalNextImage: 'readonly', + galleryClickEventHandler: 'readonly', + getExif: 'readonly', + jobStatusEl: 'readonly', + removeSplash: 'readonly', + initGPU: 'readonly', + startGPU: 'readonly', + disableNVML: 'readonly', + idbGet: 'readonly', + idbPut: 'readonly', + idbDel: 'readonly', + idbAdd: 'readonly', + idbCount: 'readonly', + idbFolderCleanup: 'readonly', + initChangelog: 'readonly', + sendNotification: 'readonly', + monitorConnection: 'readonly', + }, + }, + }, + pluginPromise.configs['flat/recommended'], + // Stylistic plugin + plugins.stylistic, + // Import X plugin + plugins.importX, + // Airbnb base recommended config + ...configs.base.recommended, + { + name: 'sdnext/js', + files: helpers.extensions.allFiles, + languageOptions: { + ecmaVersion: 'latest', + parserOptions: { + ecmaVersion: 'latest', + }, + }, + rules: { + camelcase: 'off', + 'default-case': 'off', + 'max-classes-per-file': 'warn', + 'no-await-in-loop': 'off', + 'no-bitwise': 'off', + 'no-continue': 'off', + 'no-console': 'off', + 'no-loop-func': 'off', + 'no-param-reassign': 'off', + 'no-plusplus': 'off', + 'no-redeclare': 'off', + 'no-restricted-globals': 'off', + 'no-restricted-syntax': 'off', + 'no-unused-vars': 'off', + 'no-use-before-define': 'warn', + 'no-useless-escape': 'warn', + 'prefer-destructuring': 'off', + 'prefer-rest-params': 'off', + 'prefer-template': 'warn', + 'promise/no-nesting': 'off', + radix: 'off', + '@stylistic/brace-style': [ + 'error', + '1tbs', + { + allowSingleLine: true, + }, + ], + '@stylistic/indent': ['error', 2], + '@stylistic/lines-between-class-members': [ + 'error', + 'always', + { + exceptAfterSingleLine: true, + }, + ], + '@stylistic/max-len': [ + 'warn', + { + code: 275, + tabWidth: 2, + }, + ], + '@stylistic/max-statements-per-line': 'off', + '@stylistic/no-mixed-operators': 'off', + '@stylistic/object-curly-newline': [ + 'error', + { + multiline: true, + consistent: true, + }, + ], + '@stylistic/quotes': [ + 'error', + 'single', + { + avoidEscape: true, + }, + ], + '@stylistic/semi': [ + 'error', + 'always', + { + omitLastInOneLineBlock: false, + }, + ], + 'promise/always-return': 'off', + 'promise/catch-or-return': 'off', + }, + }, +]); + +// const typescriptConfig = defineConfig([ +// // TypeScript ESLint plugin +// plugins.typescriptEslint, +// // Airbnb base TypeScript config +// ...configs.base.typescript, +// { +// name: 'sdnext/typescript', +// files: helpers.extensions.tsFiles, +// rules: { +// '@typescript-eslint/ban-ts-comment': 'off', +// '@typescript-eslint/explicit-module-boundary-types': 'off', +// '@typescript-eslint/no-shadow': 'error', +// '@typescript-eslint/no-var-requires': 'off', +// }, +// }, +// ]); + +const nodeConfig = defineConfig([ + // Node plugin + plugins.node, + { + name: 'sdnext/node', + files: ['**/cli/*.js'], + languageOptions: { + globals: { + ...globals.node, + }, + }, + rules: { + // Import as rule sets to override the `files` setting from default config + ...rules.node.base.rules, + ...rules.node.globals.rules, + ...rules.node.noUnsupportedFeatures.rules, + ...rules.node.promises.rules, + 'n/no-sync': 'off', + 'n/no-process-exit': 'off', + 'n/hashbang': 'off', + }, + }, +]); + +const jsonConfig = defineConfig([ + { + files: ['**/*.json'], + ignores: ['package-lock.json'], + plugins: { json }, + language: 'json/json', + extends: ['json/recommended'], + }, +]); + +const markdownConfig = defineConfig([ + { + files: ['**/*.md'], + plugins: { markdown }, + language: 'markdown/gfm', + processor: 'markdown/markdown', + extends: ['markdown/recommended'], + }, +]); + +const cssConfig = defineConfig([ + { + files: ['**/*.css'], + language: 'css/css', + plugins: { css }, + extends: ['css/recommended'], + // languageOptions: { + // tolerant: true, + // }, + rules: { + 'css/font-family-fallbacks': 'off', + 'css/no-invalid-properties': [ + 'error', + { + allowUnknownVariables: true, + }, + ], + 'css/no-important': 'off', + 'css/use-baseline': 'off', + }, + }, +]); + +const htmlConfig = defineConfig([ + { + files: ['**/*.html'], + plugins: { + html, + }, + extends: ['html/recommended'], + language: 'html/html', + rules: { + 'html/attrs-newline': 'off', + 'html/element-newline': 'off', + 'html/indent': [ + 'warn', + 2, + ], + 'html/no-duplicate-class': 'error', + 'html/no-extra-spacing-attrs': [ + 'error', + { + enforceBeforeSelfClose: true, + disallowMissing: true, + disallowTabs: true, + disallowInAssignment: true, + }, + ], + 'html/require-closing-tags': [ + 'error', + { + selfClosing: 'always', + }, + ], + 'html/use-baseline': 'off', + }, + }, +]); + +export default defineConfig([ + // Ignore files and folders listed in .gitignore + includeIgnoreFile(gitignorePath), + globalIgnores([ + '**/node_modules', + '**/extensions', + '**/extensions-builtin', + '**/repositories', + '**/venv', + '**/panZoom.js', + '**/split.js', + '**/exifr.js', + '**/jquery.js', + '**/sparkline.js', + '**/iframeResizer.min.js', + ]), + ...jsConfig, + // ...typescriptConfig, + ...nodeConfig, + ...jsonConfig, + ...markdownConfig, + ...cssConfig, + ...htmlConfig, +]); diff --git a/extensions-builtin/sdnext-kanvas b/extensions-builtin/sdnext-kanvas index 989a54a5b..79cae1944 160000 --- a/extensions-builtin/sdnext-kanvas +++ b/extensions-builtin/sdnext-kanvas @@ -1 +1 @@ -Subproject commit 989a54a5b2ae4ba12fefbf48c9ed61c3663c4c0c +Subproject commit 79cae1944646e57cfbfb126a971a04e44e45d776 diff --git a/extensions-builtin/sdnext-modernui b/extensions-builtin/sdnext-modernui index ded112e94..f8cb233f3 160000 --- a/extensions-builtin/sdnext-modernui +++ b/extensions-builtin/sdnext-modernui @@ -1 +1 @@ -Subproject commit ded112e94a94bf64daefa027376e0335fb43e0b7 +Subproject commit f8cb233f39e406befe70f2130f626bfa413e641a diff --git a/html/locale_en.json b/html/locale_en.json index 57b5a7c80..0894db703 100644 --- a/html/locale_en.json +++ b/html/locale_en.json @@ -97,8 +97,11 @@ {"id":"","label":"Build info on first access","localized":"","reload":"","hint":"Prevents server from building EN page on server startup and instead build it when requested"}, {"id":"","label":"Show reference styles","localized":"","reload":"","hint":"Show or hide build-it styles"}, {"id":"","label":"LoRA load using Diffusers method","localized":"","reload":"","hint":"Alternative method uses diffusers built-in LoRA capabilities instead of native SD.Next implementation (may reduce LoRA compatibility)"}, - {"id":"","label":"LoRA fuse directly to model","localized":"","reload":"","hint":"When loading LoRAs, immediately merge weights with underlying model instead of applying them on-the-fly"}, + {"id":"","label":"LoRA native fuse with model","localized":"","reload":"","hint":"Merge LoRA into the model for lower memory usage.

Warning: After removing or switching a LoRA, you may still see its style in generated images. To get a clean model, reload it from the model selector."}, {"id":"","label":"LoRA memory cache","localized":"","reload":"","hint":"How many LoRAs to keep in network for future use before requiring reloading from storage"}, + {"id":"","label":"LoRA force reload always","localized":"","reload":"","hint":"Forces LoRA networks to reload from storage on every generation, even if already cached.
Useful for debugging or when LoRA files are being modified externally.
Disable for normal use to benefit from caching."}, + {"id":"","label":"LoRA diffusers fuse with model","localized":"","reload":"","hint":"Merge LoRA into the model for lower memory usage and torch.compile compatibility.

Warning: After removing or switching a LoRA, you may still see its style in generated images. To get a clean model, reload it from the model selector."}, + {"id":"","label":"LoRA precision when quantized","localized":"","reload":"","hint":"When using a BnB 4-bit model, LoRA is applied by decompressing the weights, adding the LoRA, then recompressing. This controls the format used for recompression.

Only affects BnB 4-bit models. SDNQ models keep their original format."}, {"id":"","label":"Local","localized":"","reload":"","hint":"Models that are downlaoded and ready to use"}, {"id":"","label":"Gallery","localized":"","reload":"","hint":"Image gallery"}, {"id":"","label":"Reference","localized":"","reload":"","hint":"List of reference models that can be automatically downloaded on first use"}, @@ -156,7 +159,7 @@ {"id":"","label":"HDR Maximize","localized":"","reload":"","hint":"Calculates a 'normalization factor' by dividing the maximum tensor value by the specified range multiplied by 4. This factor is then used to shift the channels within the given boundary, ensuring maximum dynamic range for subsequent processing. The objective is to optimize dynamic range for external applications like Photoshop, particularly for adjusting levels, contrast, and brightness"}, {"id":"","label":"Enable refine pass","localized":"","reload":"","hint":"Use a similar process as image to image to upscale and/or add detail to the final image. Optionally uses refiner model to enhance image details."}, {"id":"","label":"Enable detailer pass","localized":"","reload":"","hint":"Detect target objects such as face and reprocess it at higher resolution"}, - {"id":"","label":"Include detection results","localized":"","reload":"","hint":"Include original image with detected areas marked"}, + {"id":"","label":"Include detections","localized":"","reload":"","hint":"Include original image with detected areas marked"}, {"id":"","label":"Sort detections","localized":"","reload":"","hint":"Sort detected areas by from left to right instead of detection score"}, {"id":"","label":"Denoising strength","localized":"","reload":"","hint":"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"}, {"id":"","label":"Denoise start","localized":"","reload":"","hint":"Override denoise strength by stating how early base model should finish and when refiner should start. Only applicable to refiner usage. If set to 0 or 1, denoising strength will be used"}, @@ -241,6 +244,21 @@ {"id":"","label":"Sort by","localized":"","reload":"","hint":"Sort by"}, {"id":"","label":"Nudenet","localized":"","reload":"","hint":"Flexible extension that can detect and obfustate nudity in images"}, {"id":"","label":"Prompt enhance","localized":"","reload":"","hint":"Extension that can use different LLMs to rewrite prompt for improved results"}, + {"id":"","label":"Enhance now","localized":"","reload":"","hint":"Run prompt enhancement using the selected LLM model"}, + {"id":"","label":"Apply to prompt","localized":"","reload":"","hint":"Automatically copy enhanced result to the prompt input box"}, + {"id":"","label":"Auto enhance","localized":"","reload":"","hint":"Automatically enhance prompt before every image generation"}, + {"id":"","label":"Use vision","localized":"","reload":"","hint":"Include input image when enhancing prompt.

Only available for vision-capable models, marked with \uf06e icon."}, + {"id":"","label":"LLM model","localized":"","reload":"","hint":"Select the language model to use for prompt enhancement.

Models supporting vision are marked with \uf06e icon.
Models supporting thinking mode are marked with \uf0eb icon."}, + {"id":"","label":"Model repo","localized":"","reload":"","hint":"HuggingFace repository ID for the model"}, + {"id":"","label":"Model gguf","localized":"","reload":"","hint":"Optional GGUF quantized model repository on HuggingFace"}, + {"id":"","label":"Model type","localized":"","reload":"","hint":"Optional GGUF model quantization type"}, + {"id":"","label":"Model file","localized":"","reload":"","hint":"Optional specific GGUF model file inside the repository"}, + {"id":"","label":"Load custom model","localized":"","reload":"","hint":"Load a custom model with the specified configuration"}, + {"id":"","label":"NSFW allowed","localized":"","reload":"","hint":"Allow the model to generate adult content in enhanced prompts"}, + {"id":"","label":"Prompt prefix","localized":"","reload":"","hint":"Text prepended at the beginning of the enhanced prompt result.

Useful for adding prompt elements which need to be copied to the image prompt unchanged, like quality tags 'masterpiece, best quality' or artist names, which would otherwise be rewritten by the LLM."}, + {"id":"","label":"Prompt suffix","localized":"","reload":"","hint":"Text appended to the end of the enhanced prompt result.

Useful for adding prompt elements which need to be copied to the image prompt unchanged, which would otherwise be rewritten by the LLM."}, + {"id":"","label":"Enhanced prompt","localized":"","reload":"","hint":"The enhanced prompt output from the LLM"}, + {"id":"","label":"Set prompt","localized":"","reload":"","hint":"Copy the enhanced prompt to the main prompt input"}, {"id":"","label":"Manage extensions","localized":"","reload":"","hint":"Manage extensions"}, {"id":"","label":"Manual install","localized":"","reload":"","hint":"Manually install extension"}, {"id":"","label":"Extension GIT repository URL","localized":"","reload":"","hint":"Specify extension repository URL on GitHub"}, @@ -876,8 +894,8 @@ {"id":"","label":"loaded lora","localized":"","reload":"","hint":"loaded lora"}, {"id":"","label":"logsnr","localized":"","reload":"","hint":"logsnr"}, {"id":"","label":"loop","localized":"","reload":"","hint":"loop"}, - {"id":"","label":"lora add hash info to metadata","localized":"","reload":"","hint":"lora add hash info to metadata"}, - {"id":"","label":"lora auto-apply tags","localized":"","reload":"","hint":"lora auto-apply tags"}, + {"id":"","label":"LoRA add hash info to metadata","localized":"","reload":"","hint":"Include LoRA file hashes in generated image metadata.
Useful for reproducibility and tracking which exact LoRA versions were used."}, + {"id":"","label":"LoRA auto-apply tags","localized":"","reload":"","hint":"Automatically add trigger words/tags from LoRA metadata to your prompt.
Set to the number of tags to auto-apply, e.g., 3 = add top 3 trigger tags.
Set to 0 to disable, -1 to add all available tags."}, {"id":"","label":"lora load using diffusers method for selected models","localized":"","reload":"","hint":"lora load using diffusers method for selected models"}, {"id":"","label":"lora load using legacy method","localized":"","reload":"","hint":"lora load using legacy method"}, {"id":"","label":"lora target filename","localized":"","reload":"","hint":"lora target filename"}, @@ -1209,7 +1227,7 @@ {"id":"","label":"tdd","localized":"","reload":"","hint":"tdd"}, {"id":"","label":"te","localized":"","reload":"","hint":"te"}, {"id":"","label":"temperature","localized":"","reload":"","hint":"Controls randomness in token selection by reshaping the probability distribution.
Like adjusting a dial between cautious predictability (low values ~0.4) and creative exploration (higher values ~1). Higher temperatures increase willingness to choose less obvious options, but makes outputs more unpredictable.

Set to 0 to disable, resulting in silent switch to greedy decoding, disabling sampling."}, - {"id":"","label":"Thinking mode","localized":"","reload":"","hint":"Enables thinking/reasoning, allowing the model to take more time to generate responses.
This can lead to more thoughtful and detailed answers, but will increase response time.
This setting affects both hybrid and thinking-only models, and in some may result in lower overall quality than expected. For thinking-only models like Qwen3-VL this setting might have to be combined with prefill to guarantee preventing thinking.

Models supporting this feature are marked with an \uf0eb icon."}, + {"id":"","label":"Thinking Mode","localized":"","reload":"","hint":"Enables thinking/reasoning, allowing the model to take more time to generate responses.
This can lead to more thoughtful and detailed answers, but will increase response time.
This setting affects both hybrid and thinking-only models, and in some may result in lower overall quality than expected. For thinking-only models like Qwen3-VL this setting might have to be combined with prefill to guarantee preventing thinking.

Models supporting this feature are marked with an \uf0eb icon."}, {"id":"","label":"Repetition penalty","localized":"","reload":"","hint":"Discourages reusing tokens that already appear in the prompt or output by penalizing their probabilities.
Like adding friction to revisiting previous choices. Helps break repetitive loops but may reduce coherence at aggressive values.

Set to 1 to disable."}, {"id":"","label":"text guidance scale","localized":"","reload":"","hint":"text guidance scale"}, {"id":"","label":"template","localized":"","reload":"","hint":"template"}, diff --git a/html/reference-cloud.json b/html/reference-cloud.json new file mode 100644 index 000000000..6c7ce0981 --- /dev/null +++ b/html/reference-cloud.json @@ -0,0 +1,16 @@ +{ + "Google Gemini 2.5 Flash Nano Banana": { + "path": "gemini-2.5-flash-image", + "desc": "Gemini can generate and process images conversationally. You can prompt Gemini with text, images, or a combination of both allowing you to create, edit, and iterate on visuals with unprecedented control.", + "preview": "gemini-2.5-flash-image.jpg", + "tags": "cloud", + "skip": true + }, + "Google Gemini 3.0 Pro Nano Banana": { + "path": "gemini-3-pro-image-preview", + "desc": "Built on Gemini 3. Create and edit images with studio-quality levels of precision and control", + "preview": "gemini-3-pro-image-preview.jpg", + "tags": "cloud", + "skip": true + } +} diff --git a/html/reference-community.json b/html/reference-community.json new file mode 100644 index 000000000..b76aab420 --- /dev/null +++ b/html/reference-community.json @@ -0,0 +1,132 @@ +{ + "Tempest-by-Vlad XL": { + "path": "tempestByVlad_baseV01.safetensors@https://civitai.com/api/download/models/1301775", + "preview": "tempestByVlad_baseV01.jpg", + "desc": "Flexible SDXL model with custom encoder and finetuned for larger landscape resolutions with high details and high contrast.", + "tags": "community", + "size": 6.94, + "date": "2025 January", + "extras": "" + }, + "Tempest-by-Vlad XL Hyper": { + "path": "tempestByVlad_hyperV01.safetensors@https://civitai.com/api/download/models/1343512", + "preview": "tempestByVlad_hyperV01.jpg", + "desc": "Custom distilled variant with goal to get as-normal-as-possible model that works with low steps and guidance-free", + "tags": "community", + "size": 6.94, + "date": "2025 January", + "extras": "" + }, + "Juggernaut XL XI": { + "path": "juggernautXL_juggXIByRundiffusion.safetensors@https://civitai.com/api/download/models/782002", + "preview": "juggernautXL_juggXIByRundiffusion.jpg", + "desc": "Showcase finetuned model based on Stable diffusion XL", + "date": "2024 August", + "size": 6.94, + "tags": "community", + "extras": "sampler: DEIS, steps: 20, cfg_scale: 6.0" + }, + "Juggernaut XL XI Lightning": { + "path": "juggernautXL_juggXILightningByRD.safetensors@https://civitai.com/api/download/models/920957", + "preview": "juggernautXL_juggXILightningByRD.jpg", + "desc": "Showcase finetuned model based on Stable diffusion XL", + "date": "2024 August", + "size": 6.94, + "tags": "community", + "extras": "sampler: DPM SDE, steps: 6, cfg_scale: 2.0" + }, + "Juggernaut SD Reborn": { + "original": true, + "path": "juggernaut_reborn.safetensors@https://civitai.com/api/download/models/274039", + "preview": "juggernaut_reborn.jpg", + "desc": "Showcase finetuned model based on Stable diffusion 1.5", + "date": "2023 December", + "size": 2.28, + "tags": "community", + "extras": "width: 512, height: 512, sampler: DEIS, steps: 20, cfg_scale: 6.0" + }, + "WAI Illustrious XL v15": { + "path": "waiIllustriousSDXL_v150.safetensors@https://civitai.com/api/download/models/2167369", + "preview": "waiIllustriousSDXL_v150.jpg", + "desc": "", + "tags": "community", + "size": 6.94, + "date": "2025 August", + "extras": "" + }, + "Pony Realism XL v2.3": { + "path": "ponyRealism_V23.safetensors@https://civitai.com/api/download/models/1763661", + "preview": "ponyRealism_V23.jpg", + "desc": "", + "tags": "community", + "size": 6.94, + "date": "2025 May", + "extras": "" + }, + "NoobAI XL 1.0 V-Pred": { + "path": "noobaiXLNAIXL_vPred10Version.safetensors@https://huggingface.co/Laxhar/noobai-XL-Vpred-1.0/resolve/main/NoobAI-XL-Vpred-v1.0.safetensors", + "preview": "noobaiXLNAIXL_vPred10Version.jpg", + "desc": "", + "tags": "community", + "size": 6.94, + "date": "2024 December", + "extras": "" + }, + "NoobAI XL 1.1 Epsilon": { + "path": "noobaiXLNAIXL_epsilonPred11Version.safetensors@https://huggingface.co/Laxhar/noobai-XL-1.1/resolve/main/NoobAI-XL-v1.1.safetensors", + "preview": "noobaiXLNAIXL_epsilonPred11Version.jpg", + "desc": "", + "tags": "community", + "size": 6.94, + "date": "2024 November", + "extras": "" + }, + "WAI-Ani-Pony XL v14": { + "path": "waiANIPONYXL_v140.safetensors.safetensors@https://civitai.com/api/download/models/1767402", + "preview": "waiANIPONYXL_v140.jpg", + "desc": "", + "tags": "community", + "size": 6.94, + "date": "2025 May", + "extras": "" + }, + "Tiwaz CenKreChro": { + "path": "Tiwaz/CenKreChro", + "preview": "Tiwaz--CenKreChro.jpg", + "skip": true, + "desc": "Based Centerfold Flux 5, trying to merge in Chroma and Krea.", + "extras": "", + "tags": "community", + "date": "2025 September" + }, + "purplesmartai Pony 7": { + "path": "purplesmartai/pony-v7-base", + "preview": "purplesmartai--pony-v7-base.jpg", + "skip": true, + "desc": "Pony V7 is a versatile character generation model based on AuraFlow architecture. It supports a wide range of styles and species types (humanoid, anthro, feral, and more) and handles character interactions through natural language prompts.", + "extras": "", + "tags": "community", + "date": "October September" + }, + "ShuttleAI Shuttle 3.0 Diffusion": { + "path": "shuttleai/shuttle-3-diffusion", + "desc": "Shuttle uses Flux.1 Schnell as its base. It can produce images similar to Flux Dev or Pro in just 4 steps, and it is licensed under Apache 2. The model was partially de-distilled during training. When used beyond 10 steps, it enters refiner mode enhancing image details without altering the composition", + "preview": "shuttleai--shuttle-3-diffusion.jpg", + "tags": "community", + "skip": true + }, + "ShuttleAI Shuttle 3.1 Aesthetic": { + "path": "shuttleai/shuttle-3.1-aesthetic", + "desc": "Shuttle uses Flux.1 Schnell as its base. It can produce images similar to Flux Dev or Pro in just 4 steps, and it is licensed under Apache 2. The model was partially de-distilled during training. When used beyond 10 steps, it enters refiner mode enhancing image details without altering the composition", + "preview": "shuttleai--shuttle-3_1-aestetic.jpg", + "tags": "community", + "skip": true + }, + "ShuttleAI Shuttle Jaguar": { + "path": "shuttleai/shuttle-jaguar", + "desc": "Shuttle uses Flux.1 Schnell as its base. It can produce images similar to Flux Dev or Pro in just 4 steps, and it is licensed under Apache 2. The model was partially de-distilled during training. When used beyond 10 steps, it enters refiner mode enhancing image details without altering the composition", + "preview": "shuttleai--shuttle-jaguar.jpg", + "tags": "community", + "skip": true + } +} diff --git a/html/reference-distilled.json b/html/reference-distilled.json new file mode 100644 index 000000000..dc26e97bd --- /dev/null +++ b/html/reference-distilled.json @@ -0,0 +1,185 @@ +{ + "StabilityAI StableDiffusion XL Turbo": { + "path": "stabilityai/sdxl-turbo", + "preview": "stabilityai--sdxl-turbo.jpg", + "desc": "SDXL-Turbo is a fast generative text-to-image model that can synthesize photorealistic images from a text prompt in a 1-4 steps.", + "skip": true, + "variant": "fp16", + "tags": "distilled", + "extras": "steps: 4, cfg_scale: 0.0" + }, + "StabilityAI Stable Cascade Lite": { + "path": "huggingface/stabilityai/stable-cascade-lite", + "skip": true, + "variant": "bf16", + "desc": "Stable Cascade is a diffusion model built upon the Würstchen architecture and its main difference to other models like Stable Diffusion is that it is working at a much smaller latent space. Why is this important? The smaller the latent space, the faster you can run inference and the cheaper the training becomes. How small is the latent space? Stable Diffusion uses a compression factor of 8, resulting in a 1024x1024 image being encoded to 128x128. Stable Cascade achieves a compression factor of 42, meaning that it is possible to encode a 1024x1024 image to 24x24, while maintaining crisp reconstructions. The text-conditional model is then trained in the highly compressed latent space. Previous versions of this architecture, achieved a 16x cost reduction over Stable Diffusion 1.5", + "preview": "stabilityai--stable-cascade-lite.jpg", + "extras": "sampler: Default, cfg_scale: 4.0, image_cfg_scale: 1.0", + "size": 4.97, + "tags": "distilled", + "date": "2024 February" + }, + "StabilityAI Stable Diffusion 3.5 Turbo": { + "path": "stabilityai/stable-diffusion-3.5-large-turbo", + "skip": true, + "variant": "fp16", + "desc": "Stable Diffusion 3.5 Large Turbo is a Multimodal Diffusion Transformer (MMDiT) text-to-image model with Adversarial Diffusion Distillation (ADD) that features improved performance in image quality, typography, complex prompt understanding, and resource-efficiency, with a focus on fewer inference steps.", + "preview": "stabilityai--stable-diffusion-3_5-large-turbo.jpg", + "tags": "distilled", + "extras": "sampler: Default, cfg_scale: 7.0" + }, + "Tencent FLUX.1 Dev SRPO": { + "path": "vladmandic/flux.1-dev-SRPO", + "preview": "vladmandic--flux.1-dev-SRPO.jpg", + "desc": "FLUX.1 Dev SRPO is Tencent trained with specific technique: Directly Aligning the Full Diffusion Trajectory with Fine-Grained Human Preference", + "tags": "distilled", + "skip": true, + "extras": "sampler: Default, cfg_scale: 4.5" + }, + "Qwen-Image-Lightning": { + "path": "vladmandic/Qwen-Lightning", + "preview": "vladmandic--Qwen-Lightning.jpg", + "desc": "Qwen-Lightning is step-distilled from Qwen-Image to allow for generation in 8 steps.", + "skip": true, + "extras": "steps: 8", + "size": 56.1, + "tags": "distilled", + "date": "2025 August" + }, + "Qwen-Image-Distill": { + "path": "SahilCarterr/Qwen-Image-Distill-Full", + "preview": "SahilCarterr--Qwen-Image-Distill-Full.jpg", + "desc": "Qwen-Image-Distill is a distilled and accelerated version of Qwen-Image by DiffSynth-Studio.", + "skip": true, + "extras": "steps: 15", + "size": 56.1, + "tags": "distilled", + "date": "2025 August" + }, + "Qwen-Image-Lightning-Edit": { + "path": "vladmandic/Qwen-Lightning-Edit", + "preview": "vladmandic--Qwen-Lightning-Edit.jpg", + "desc": "Qwen-Lightning-Edit is step-distilled from Qwen-Image-Edit to allow for generation in 8 steps.", + "skip": true, + "extras": "steps: 8", + "size": 56.1, + "tags": "distilled", + "date": "2025 August" + }, + "Qwen-Image Pruning-12B": { + "path": "OPPOer/Qwen-Image-Pruning", + "subfolder": "Qwen-Image-12B-8steps", + "preview": "OPPOer--Qwen-Image-Pruning.jpg", + "desc": "This open-source project is based on Qwen-Image and has attempted model pruning, removing 20 layers while retaining the weights of 40 layers, resulting in a model size of 12B parameters.", + "skip": true, + "tags": "distilled", + "date": "2025 Ocotober" + }, + "Qwen-Image-Edit Pruning-13B": { + "path": "OPPOer/Qwen-Image-Edit-Pruning", + "subfolder": "Qwen-Image-Edit-13B-4steps", + "preview": "OPPOer--Qwen-Image-Edit-Pruning.jpg", + "desc": "This open-source project is based on Qwen-Image-Edit and has attempted model pruning, removing 20 layers while retaining the weights of 40 layers, resulting in a model size of 13.6B parameters.", + "skip": true, + "tags": "distilled", + "date": "2025 Ocotober" + }, + "Qwen-Image-Edit-2509 Pruning-13B": { + "path": "OPPOer/Qwen-Image-Edit-2509-Pruning", + "subfolder": "Qwen-Image-Edit-2509-13B-4steps", + "preview": "OPPOer--Qwen-Image-Edit-2509-Pruning.jpg", + "desc": "This open-source project is based on Qwen-Image-Edit and has attempted model pruning, removing 20 layers while retaining the weights of 40 layers, resulting in a model size of 13.6B parameters.", + "skip": true, + "tags": "distilled", + "date": "2025 Ocotober" + }, + "lodestones Chroma1 Flash": { + "path": "lodestones/Chroma1-Flash", + "preview": "lodestones--Chroma1-Flash.jpg", + "desc": "Chroma is a 8.9B parameter model based on FLUX.1-schnell. It’s fully Apache 2.0 licensed, ensuring that anyone can use, modify, and build on top of it—no corporate gatekeeping. A fine-tuned version of the Chroma1-Base made to find the best way to make these flow matching models faster.", + "skip": true, + "extras": "", + "size": 26.84, + "tags": "distilled", + "date": "2025 July" + }, + "SDXL Flash Mini": { + "path": "SDXL-Flash_Mini.safetensors@https://huggingface.co/sd-community/sdxl-flash-mini/resolve/main/SDXL-Flash_Mini.safetensors?download=true", + "preview": "SDXL-Flash_Mini.jpg", + "desc": "Introducing the new fast model SDXL Flash (Mini), we learned that all fast XL models work fast, but the quality decreases, and we also made a fast model, but it is not as fast as LCM, Turbo, Lightning and Hyper, but the quality is higher.", + "extras": "width: 2048, height: 1024, sampler: DEIS, steps: 40, cfg_scale: 6.0", + "tags": "distilled", + "experimental": true + }, + "NVLabs Sana 1.5 1.6B 1k Sprint": { + "path": "Efficient-Large-Model/Sana_Sprint_1.6B_1024px_diffusers", + "desc": "SANA-Sprint is an ultra-efficient diffusion model for text-to-image (T2I) generation, reducing inference steps from 20 to 1-4 while achieving state-of-the-art performance.", + "preview": "Efficient-Large-Model--Sana15_Sprint_1600M_1024px_diffusers.jpg", + "tags": "distilled", + "skip": true + }, + "Segmind SSD-1B": { + "path": "huggingface/segmind/SSD-1B", + "preview": "segmind--SSD-1B.jpg", + "desc": "The Segmind Stable Diffusion Model (SSD-1B) offers a compact, efficient, and distilled version of the SDXL model. At 50% smaller and 60% faster than Stable Diffusion XL (SDXL), it provides quick and seamless performance without sacrificing image quality.", + "variant": "fp16", + "skip": true, + "extras": "sampler: Default, cfg_scale: 9.0", + "size": 8.72, + "tags": "distilled", + "date": "2023 October" + }, + "Segmind Tiny": { + "path": "segmind/tiny-sd", + "preview": "segmind--tiny-sd.jpg", + "desc": "Segmind's Tiny-SD offers a compact, efficient, and distilled version of Realistic Vision 4.0 and is up to 80% faster than SD1.5", + "extras": "width: 512, height: 512, sampler: Default, cfg_scale: 9.0", + "size": 1.03, + "tags": "distilled", + "date": "2023 July" + }, + "Tencent HunyuanImage 2.1 Distilled": { + "path": "hunyuanvideo-community/HunyuanImage-2.1-Distilled-Diffusers", + "desc": "HunyuanImage-2.1, a highly efficient text-to-image model that is capable of generating 2K (2048 × 2048) resolution images.", + "preview": "hunyuanvideo-community--HunyuanImage-2.1-Distilled-Diffusers.jpg", + "extras": "", + "tags": "distilled", + "skip": true, + "size": 0, + "date": "2025 August" + }, + "Tencent HunyuanDiT 1.2 Distilled": { + "path": "Tencent-Hunyuan/HunyuanDiT-v1.2-Diffusers-Distilled", + "desc": "Hunyuan-DiT : A Powerful Multi-Resolution Diffusion Transformer with Fine-Grained Chinese Understanding.", + "preview": "Tencent-Hunyuan--HunyuanDiT-v1.2-Diffusers-Distilled.jpg", + "tags": "distilled", + "extras": "sampler: Default, cfg_scale: 2.0" + }, + "Tencent HunyuanDiT 1.1 Distilled": { + "path": "Tencent-Hunyuan/HunyuanDiT-v1.1-Diffusers-Distilled", + "desc": "Hunyuan-DiT : A Powerful Multi-Resolution Diffusion Transformer with Fine-Grained Chinese Understanding.", + "preview": "Tencent-Hunyuan--HunyuanDiT-v1.1-Diffusers-Distilled.jpg", + "tags": "distilled", + "extras": "sampler: Default, cfg_scale: 2.0" + }, + "Black Forest Labs FLUX.2 Klein 4B": { + "path": "black-forest-labs/FLUX.2-klein-4B", + "preview": "black-forest-labs--FLUX.2-klein-4B.jpg", + "desc": "FLUX.2-klein-4B is a 4 billion parameter size-distilled version of FLUX.2-dev optimized for consumer GPUs. Achieves sub-second inference with 4 steps. Supports both text-to-image generation and multi-reference image editing. Apache 2.0 licensed.", + "skip": true, + "tags": "distilled", + "extras": "sampler: Default, cfg_scale: 1.0, steps: 4", + "size": 8.5, + "date": "2025 January" + }, + "Black Forest Labs FLUX.2 Klein 9B": { + "path": "black-forest-labs/FLUX.2-klein-9B", + "preview": "black-forest-labs--FLUX.2-klein-9B.jpg", + "desc": "FLUX.2-klein-9B is a 9 billion parameter size-distilled version of FLUX.2-dev. Higher quality than 4B variant with sub-second inference using 4 steps. Supports text-to-image and multi-reference editing. Non-commercial license.", + "skip": true, + "tags": "distilled", + "extras": "sampler: Default, cfg_scale: 1.0, steps: 4", + "size": 18.5, + "date": "2025 January" + } +} diff --git a/html/reference-quant.json b/html/reference-quant.json new file mode 100644 index 000000000..1ce214013 --- /dev/null +++ b/html/reference-quant.json @@ -0,0 +1,239 @@ +{ + "FLUX.1-Dev sdnq-svd-uint4": { + "path": "Disty0/FLUX.1-dev-SDNQ-uint4-svd-r32", + "preview": "Disty0--FLUX.1-dev-SDNQ-uint4-svd-r32.jpg", + "desc": "Quantization of black-forest-labs/FLUX.1-dev using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "skip": true, + "tags": "quantized", + "size": 12.60, + "date": "2025 October", + "extras": "" + }, + "FLUX.1-Schnell sdnq-svd-uint4": { + "path": "Disty0/FLUX.1-schnell-SDNQ-uint4-svd-r32", + "preview": "Disty0--FLUX.1-schnell-SDNQ-uint4-svd-r32.jpg", + "desc": "Quantization of black-forest-labs/FLUX.1-schnell using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "skip": true, + "tags": "quantized", + "size": 12.60, + "date": "2025 October", + "extras": "" + }, + "FLUX.1-Dev Krea sdnq-svd-uint4": { + "path": "Disty0/FLUX.1-Krea-dev-SDNQ-uint4-svd-r32", + "preview": "Disty0--FLUX.1-Krea-dev-SDNQ-uint4-svd-r32.jpg", + "desc": "Quantization of black-forest-labs/FLUX.1-Krea-dev using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "skip": true, + "tags": "quantized", + "size": 12.60, + "date": "2025 October", + "extras": "" + }, + "FLUX.1-Dev Kontext sdnq-svd-uint4": { + "path": "Disty0/FLUX.1-Kontext-dev-SDNQ-uint4-svd-r32", + "preview": "Disty0--FLUX.1-Kontext-dev-SDNQ-uint4-svd-r32.jpg", + "desc": "Quantization of black-forest-labs/FLUX.1-Kontext-dev using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "skip": true, + "tags": "quantized", + "size": 12.60, + "date": "2025 October", + "extras": "" + }, + "FLUX.2 Dev sdnq-svd-uint4": { + "path": "Disty0/FLUX.2-dev-SDNQ-uint4-svd-r32", + "preview": "Disty0--FLUX.2-dev-SDNQ-uint4-svd-r32.jpg", + "desc": "Quantization of black-forest-labs/FLUX.2-dev using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "skip": true, + "tags": "quantized", + "extras": "", + "size": 31.58, + "date": "2025 November" + }, + "Black Forest Labs FLUX.2 Klein 4B sdnq-uint4-dynamic": { + "path": "Disty0/FLUX.2-klein-4B-SDNQ-4bit-dynamic", + "preview": "Disty0--FLUX.2-klein-4B-SDNQ-4bit-dynamic.jpg", + "desc": "Dynamic 4-bit quantization of black-forest-labs/FLUX.2-klein-4B using SDNQ.", + "skip": true, + "extras": "sampler: Default, cfg_scale: 1.0, steps: 4", + "tags": "quantized", + "size": 5.1, + "date": "2026 January" + }, + "Black Forest Labs FLUX.2 Klein 9B sdnq-uint4-dynamic-svd": { + "path": "Disty0/FLUX.2-klein-9B-SDNQ-4bit-dynamic-svd-r32", + "preview": "Disty0--FLUX.2-klein-9B-SDNQ-4bit-dynamic-svd-r32.jpg", + "desc": "Dynamic 4-bit quantization of black-forest-labs/FLUX.2-klein-9B using SDNQ with SVD rank 32.", + "skip": true, + "extras": "sampler: Default, cfg_scale: 1.0, steps: 4", + "tags": "quantized", + "size": 11.7, + "date": "2026 January" + }, + "Chroma1-HD sdnq-svd-uint4": { + "path": "Disty0/Chroma1-HD-SDNQ-uint4-svd-r32", + "preview": "Disty0--Chroma1-HD-SDNQ-uint4-svd-r32.jpg", + "desc": "Quantization of lodestones/Chroma1-HD using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "skip": true, + "tags": "quantized", + "size": 11.89, + "date": "2025 October", + "extras": "" + }, + "Wan-AI Wan2.2 A14B T2I sdnq-svd-uint4": { + "path": "Disty0/Wan2.2-T2V-A14B-SDNQ-uint4-svd-r32", + "preview": "Wan-AI--Wan2.2-T2V-A14B-Diffusers.jpg", + "desc": "Quantization of black-forest-labs/FLUX.1-dev using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "skip": true, + "tags": "quantized", + "date": "2025 October", + "size": 23.54, + "extras": "" + }, + "Wan-AI Wan2.2 A14B I2I sdnq-svd-uint4": { + "path": "Disty0/Wan2.2-I2V-A14B-SDNQ-uint4-svd-r32", + "preview": "Wan-AI--Wan2.2-T2V-A14B-Diffusers.jpg", + "desc": "Quantization of Laxhar/noobai-XL-1.1 using SDNQ: sdnq-svd 4-bit uint with svd rank 128", + "skip": true, + "tags": "quantized", + "date": "2025 October", + "size": 23.55, + "extras": "" + }, + "Z-Image-Turbo sdnq-svd-uint4": { + "path": "Disty0/Z-Image-Turbo-SDNQ-uint4-svd-r32", + "preview": "Disty0--Z-Image-Turbo-SDNQ-uint4-svd-r32.jpg", + "desc": "Quantization of Tongyi-MAI/Z-Image-Turbo using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "skip": true, + "tags": "quantized", + "extras": "sampler: Default, cfg_scale: 1.0, steps: 9", + "size": 6.5, + "date": "2025 November" + }, + "Qwen-Image sdnq-svd-uint4": { + "path": "Disty0/Qwen-Image-SDNQ-uint4-svd-r32", + "preview": "Qwen--Qwen-Image.jpg", + "desc": "Quantization of Qwen/Qwen-Image using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "skip": true, + "tags": "quantized", + "date": "2025 October", + "size": 16.09, + "extras": "" + }, + "Qwen-Image-2512 sdnq-svd-uint4": { + "path": "Disty0/Qwen-Image-2512-SDNQ-uint4-svd-r32", + "preview": "Disty0--Qwen-Image-2512-SDNQ-uint4-svd-r32.jpg", + "desc": "Quantization of Qwen/Qwen-Image-2512 using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "skip": true, + "tags": "quantized", + "extras": "", + "size": 16.10, + "date": "2025 December" + }, + "Qwen-Image-2512 sdnq-dynamic-uint4": { + "path": "Disty0/Qwen-Image-2512-SDNQ-4bit-dynamic", + "preview": "Disty0--Qwen-Image-2512-SDNQ-uint4-svd-r32.jpg", + "desc": "Quantization of Qwen/Qwen-Image-2512 using SDNQ: sdnq-dynamic 4-bit uint", + "skip": true, + "tags": "quantized", + "extras": "", + "size": 17.2, + "date": "2026 January" + }, + "Qwen-Image-Edit sdnq-svd-uint4": { + "path": "Disty0/Qwen-Image-Edit-SDNQ-uint4-svd-r32", + "preview": "Qwen--Qwen-Image-Edit.jpg", + "desc": "Quantization of Qwen/Qwen-Image-Edit using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "skip": true, + "tags": "quantized", + "date": "2025 October", + "size": 16.10, + "extras": "" + }, + "Qwen-Image-Edit-2509 sdnq-svd-uint4": { + "path": "Disty0/Qwen-Image-Edit-2509-SDNQ-uint4-svd-r32", + "preview": "Qwen--Qwen-Image-Edit-2509.jpg", + "desc": "Quantization of Qwen/Qwen-Image-Edit-2509 using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "skip": true, + "tags": "quantized", + "date": "2025 October", + "size": 16.10, + "extras": "" + }, + "Qwen-Image-Edit-2511 sdnq-svd-uint4": { + "path": "Disty0/Qwen-Image-Edit-2511-SDNQ-uint4-svd-r32", + "preview": "Disty0--Qwen-Image-Edit-2511-SDNQ-uint4-svd-r32.jpg", + "desc": "Quantization of Qwen/Qwen-Image-Edit-2511 using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "skip": true, + "tags": "quantized", + "date": "2025 December", + "size": 16.10, + "extras": "" + }, + "Qwen-Image-Layered sdnq-svd-uint4": { + "path": "Disty0/Qwen-Image-Layered-SDNQ-uint4-svd-r32", + "preview": "Disty0--Qwen-Image-Layered-SDNQ-uint4-svd-r32.jpg", + "desc": "Quantization of Qwen/Qwen-Image-Layered using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "skip": true, + "tags": "quantized", + "date": "2025 December", + "size": 16.10, + "extras": "" + }, + "nVidia ChronoEdit sdnq-svd-uint4": { + "path": "Disty0/ChronoEdit-14B-SDNQ-uint4-svd-r32", + "preview": "Disty0--ChronoEdit-14B-SDNQ-uint4-svd-r32.jpg", + "desc": "Quantization of nvidia/ChronoEdit-14B-Diffusers using SDNQ: sdnq-svd 4-bit uint with svd rank 32.", + "skip": true, + "tags": "quantized", + "date": "2025 October", + "size": 18.10, + "extras": "" + }, + "Tencent HunyuanImage 3.0 sdnq-svd-uint4": { + "path": "Disty0/HunyuanImage3-SDNQ-uint4-svd-r32", + "desc": "Quantization of tencent/HunyuanImage-3.0 using SDNQ: sdnq-svd 4-bit uint with svd rank 32.", + "preview": "Disty0--HunyuanImage3-SDNQ-uint4-svd-r32.jpg", + "extras": "", + "skip": true, + "tags": "quantized", + "size": 57.06, + "date": "2025 September" + }, + "Tempest-by-Vlad XL sdnq-svd-uint4": { + "path": "vladmandic/tempestByVlad_baseV01-SDNQ-uint4-svd", + "preview": "vladmandic--tempestByVlad_baseV01-SDNQ-uint4-svd.jpg", + "desc": "Quantization of vladmandic/tempestByVlad_baseV01 using SDNQ: sdnq-svd 4-bit uint with svd rank 128", + "tags": "quantized", + "size": 3.37, + "date": "2025 October", + "extras": "" + }, + "NoobAI-XL v1.1 epsilon sdnq-svd-uint4": { + "path": "Disty0/NoobAI-XL-v1.1-SDNQ-uint4-svd-r128", + "preview": "Disty0--NoobAI-XL-v1.1-SDNQ-uint4-svd-r128.jpg", + "desc": "Quantization of Laxhar/noobai-XL-1.1 using SDNQ: sdnq-svd 4-bit uint with svd rank 128", + "tags": "quantized", + "size": 3.37, + "date": "2025 October", + "extras": "" + }, + "NoobAI-XL v1.0 v-pred sdnq-svd-uint4": { + "path": "Disty0/NoobAI-XL-Vpred-v1.0-SDNQ-uint4-svd-r128", + "preview": "Disty0--NoobAI-XL-Vpred-v1.0-SDNQ-uint4-svd-r128.jpg", + "desc": "Quantization of Laxhar/noobai-XL-Vpred-1.0 using SDNQ: sdnq-svd 4-bit uint with svd rank 128", + "tags": "quantized", + "size": 3.37, + "date": "2025 October", + "extras": "" + }, + "ZAI GLM-Image sdnq-dynamic-uint4": { + "path": "Disty0/GLM-Image-SDNQ-4bit-dynamic", + "preview": "zai-org--GLM-Image.jpg", + "desc": "Quantization of ZAI GLM-Image using SDNQ: sdnq-dynamic 4-bit uint", + "skip": true, + "extras": "sampler: Default, cfg_scale: 1.5, steps: 50", + "size": 11.6, + "tags": "quantized", + "date": "2026 January" + } +} diff --git a/html/reference.json b/html/reference.json index 492584ad8..2f1f6562b 100644 --- a/html/reference.json +++ b/html/reference.json @@ -38,15 +38,6 @@ "size": 6.94, "date": "2023 July" }, - "StabilityAI StableDiffusion XL Turbo": { - "path": "stabilityai/sdxl-turbo", - "preview": "stabilityai--sdxl-turbo.jpg", - "desc": "SDXL-Turbo is a fast generative text-to-image model that can synthesize photorealistic images from a text prompt in a 1-4 steps.", - "skip": true, - "variant": "fp16", - "tags": "distilled", - "extras": "steps: 4, cfg_scale: 0.0" - }, "StabilityAI Stable Cascade": { "path": "huggingface/stabilityai/stable-cascade", "skip": true, @@ -57,17 +48,6 @@ "size": 11.82, "date": "2024 February" }, - "StabilityAI Stable Cascade Lite": { - "path": "huggingface/stabilityai/stable-cascade-lite", - "skip": true, - "variant": "bf16", - "desc": "Stable Cascade is a diffusion model built upon the Würstchen architecture and its main difference to other models like Stable Diffusion is that it is working at a much smaller latent space. Why is this important? The smaller the latent space, the faster you can run inference and the cheaper the training becomes. How small is the latent space? Stable Diffusion uses a compression factor of 8, resulting in a 1024x1024 image being encoded to 128x128. Stable Cascade achieves a compression factor of 42, meaning that it is possible to encode a 1024x1024 image to 24x24, while maintaining crisp reconstructions. The text-conditional model is then trained in the highly compressed latent space. Previous versions of this architecture, achieved a 16x cost reduction over Stable Diffusion 1.5", - "preview": "stabilityai--stable-cascade-lite.jpg", - "extras": "sampler: Default, cfg_scale: 4.0, image_cfg_scale: 1.0", - "size": 4.97, - "tags": "distilled", - "date": "2024 February" - }, "StabilityAI Stable Diffusion 3.0 Medium": { "path": "stabilityai/stable-diffusion-3-medium-diffusers", "skip": true, @@ -98,15 +78,6 @@ "size": 26.98, "date": "2024 October" }, - "StabilityAI Stable Diffusion 3.5 Turbo": { - "path": "stabilityai/stable-diffusion-3.5-large-turbo", - "skip": true, - "variant": "fp16", - "desc": "Stable Diffusion 3.5 Large Turbo is a Multimodal Diffusion Transformer (MMDiT) text-to-image model with Adversarial Diffusion Distillation (ADD) that features improved performance in image quality, typography, complex prompt understanding, and resource-efficiency, with a focus on fewer inference steps.", - "preview": "stabilityai--stable-diffusion-3_5-large-turbo.jpg", - "tags": "distilled", - "extras": "sampler: Default, cfg_scale: 7.0" - }, "Black Forest Labs FLUX.1 Dev": { "path": "black-forest-labs/FLUX.1-dev", @@ -153,20 +124,29 @@ "size": 104.74, "date": "2025 November" }, - "Tencent FLUX.1 Dev SRPO": { - "path": "vladmandic/flux.1-dev-SRPO", - "preview": "vladmandic--flux.1-dev-SRPO.jpg", - "desc": "FLUX.1 Dev SRPO is Tencent trained with specific technique: Directly Aligning the Full Diffusion Trajectory with Fine-Grained Human Preference", - "tags": "distilled", + "Black Forest Labs FLUX.2 Klein Base 4B": { + "path": "black-forest-labs/FLUX.2-klein-base-4B", + "preview": "black-forest-labs--FLUX.2-klein-base-4B.jpg", + "desc": "FLUX.2-klein-base-4B is the undistilled 4 billion parameter base model of FLUX.2-klein. Requires 50 inference steps for full quality but offers flexibility for fine-tuning. Supports text-to-image and multi-reference editing. Apache 2.0 licensed.", "skip": true, - "extras": "sampler: Default, cfg_scale: 4.5" + "extras": "sampler: Default, cfg_scale: 4.0, steps: 50", + "size": 8.5, + "date": "2025 January" + }, + "Black Forest Labs FLUX.2 Klein Base 9B": { + "path": "black-forest-labs/FLUX.2-klein-base-9B", + "preview": "black-forest-labs--FLUX.2-klein-base-9B.jpg", + "desc": "FLUX.2-klein-base-9B is the undistilled 9 billion parameter base model of FLUX.2-klein. Requires 50 inference steps for full quality but offers flexibility for fine-tuning. Supports text-to-image and multi-reference editing. Non-commercial license.", + "skip": true, + "extras": "sampler: Default, cfg_scale: 4.0, steps: 50", + "size": 18.5, + "date": "2025 January" }, "Z-Image-Turbo": { "path": "Tongyi-MAI/Z-Image-Turbo", "preview": "Tongyi-MAI--Z-Image-Turbo.jpg", - "desc": "Z-Image-Turbo, a distilled version of Z-Image that matches or exceeds leading competitors with only 8 NFEs (Number of Function Evaluations). It offers sub-second inference latency on enterprise-grade H800 GPUs and fits comfortably within 16G VRAM consumer devices. It excels in photorealistic image generation, bilingual text rendering (English & Chinese), and robust instruction adherence.", - "tags": "distilled", + "desc": "Z-Image-Turbo, a distilled version of Z-Image that matches or exceeds leading competitors with only 8 NFEs (Number of Function Evaluations). It excels in photorealistic image generation, bilingual text rendering (English & Chinese), and robust instruction adherence.", "skip": true, "extras": "sampler: Default, cfg_scale: 1.0, steps: 9", "size": 20.3, @@ -182,6 +162,15 @@ "size": 56.1, "date": "2025 August" }, + "Qwen-Image-2512": { + "path": "Qwen/Qwen-Image-2512", + "preview": "Qwen--Qwen-Image-2512.jpg", + "desc": "Qwen-Image-2512 is an Qwen Image successor, that significantly reduces the AI-generated look, got finer natural detailils and improved text rendering.", + "skip": true, + "extras": "", + "size": 53.7, + "date": "2025 December" + }, "Qwen-Image-Edit": { "path": "Qwen/Qwen-Image-Edit", "preview": "Qwen--Qwen-Image-Edit.jpg", @@ -202,7 +191,7 @@ }, "Qwen-Image-Edit-2511": { "path": "Qwen/Qwen-Image-Edit-2511", - "preview": "Qwen--Qwen-Image-Edit-2509.jpg", + "preview": "Qwen--Qwen-Image-Edit-2511.jpg", "desc": "Key enhancements: mitigate image drift, improved character consistency, enhanced industrial design generation, and strengthened geometric reasoning ability.", "skip": true, "extras": "", @@ -211,74 +200,17 @@ }, "Qwen-Image-Layered": { "path": "Qwen/Qwen-Image-Layered", - "preview": "Qwen--Qwen-Image-Edit-2509.jpg", + "preview": "Qwen--Qwen-Image-Layered.jpg", "desc": "Qwen-Image-Layered, a model capable of decomposing an image into multiple RGBA layers", "skip": true, "extras": "", "size": 53.7, "date": "2025 December" }, - "Qwen-Image-Lightning": { - "path": "vladmandic/Qwen-Lightning", - "preview": "vladmandic--Qwen-Lightning.jpg", - "desc": "Qwen-Lightning is step-distilled from Qwen-Image to allow for generation in 8 steps.", - "skip": true, - "extras": "steps: 8", - "size": 56.1, - "tags": "distilled", - "date": "2025 August" - }, - "Qwen-Image-Distill": { - "path": "SahilCarterr/Qwen-Image-Distill-Full", - "preview": "SahilCarterr--Qwen-Image-Distill-Full.jpg", - "desc": "Qwen-Image-Distill is a distilled and accelerated version of Qwen-Image by DiffSynth-Studio.", - "skip": true, - "extras": "steps: 15", - "size": 56.1, - "tags": "distilled", - "date": "2025 August" - }, - "Qwen-Image-Lightning-Edit": { - "path": "vladmandic/Qwen-Lightning-Edit", - "preview": "vladmandic--Qwen-Lightning-Edit.jpg", - "desc": "Qwen-Lightning-Edit is step-distilled from Qwen-Image-Edit to allow for generation in 8 steps.", - "skip": true, - "extras": "steps: 8", - "size": 56.1, - "tags": "distilled", - "date": "2025 August" - }, - "Qwen-Image Pruning-12B": { - "path": "OPPOer/Qwen-Image-Pruning", - "subfolder": "Qwen-Image-12B-8steps", - "preview": "OPPOer--Qwen-Image-Pruning.jpg", - "desc": "This open-source project is based on Qwen-Image and has attempted model pruning, removing 20 layers while retaining the weights of 40 layers, resulting in a model size of 12B parameters.", - "skip": true, - "tags": "distilled", - "date": "2025 Ocotober" - }, - "Qwen-Image-Edit Pruning-13B": { - "path": "OPPOer/Qwen-Image-Edit-Pruning", - "subfolder": "Qwen-Image-Edit-13B-4steps", - "preview": "OPPOer--Qwen-Image-Edit-Pruning.jpg", - "desc": "This open-source project is based on Qwen-Image-Edit and has attempted model pruning, removing 20 layers while retaining the weights of 40 layers, resulting in a model size of 13.6B parameters.", - "skip": true, - "tags": "distilled", - "date": "2025 Ocotober" - }, - "Qwen-Image-Edit-2509 Pruning-13B": { - "path": "OPPOer/Qwen-Image-Edit-2509-Pruning", - "subfolder": "Qwen-Image-Edit-2509-13B-4steps", - "preview": "OPPOer--Qwen-Image-Edit-2509-Pruning.jpg", - "desc": "This open-source project is based on Qwen-Image-Edit and has attempted model pruning, removing 20 layers while retaining the weights of 40 layers, resulting in a model size of 13.6B parameters.", - "skip": true, - "tags": "distilled", - "date": "2025 Ocotober" - }, "lodestones Chroma1 HD": { "path": "lodestones/Chroma1-HD", - "preview": "lodestones--Chroma-HD.jpg", + "preview": "lodestones--Chroma1-HD.jpg", "desc": "Chroma is a 8.9B parameter model based on FLUX.1-schnell. It’s fully Apache 2.0 licensed, ensuring that anyone can use, modify, and build on top of it—no corporate gatekeeping. This is the high-res fine-tune of the Chroma1-Base at a 1024x1024 resolution.", "skip": true, "extras": "", @@ -287,50 +219,22 @@ }, "lodestones Chroma1 Base": { "path": "lodestones/Chroma1-Base", - "preview": "lodestones--Chroma-Base.jpg", + "preview": "lodestones--Chroma1-Base.jpg", "desc": "Chroma is a 8.9B parameter model based on FLUX.1-schnell. It’s fully Apache 2.0 licensed, ensuring that anyone can use, modify, and build on top of it—no corporate gatekeeping. This is the core 512x512 model. It's a solid, all-around foundation for pretty much any creative project.", "skip": true, "extras": "", "size": 26.84, "date": "2025 July" }, - "lodestones Chroma1 Flash": { - "path": "lodestones/Chroma1-Flash", - "preview": "lodestones--Chroma-flash.jpg", - "desc": "Chroma is a 8.9B parameter model based on FLUX.1-schnell. It’s fully Apache 2.0 licensed, ensuring that anyone can use, modify, and build on top of it—no corporate gatekeeping. A fine-tuned version of the Chroma1-Base made to find the best way to make these flow matching models faster.", - "skip": true, - "extras": "", - "size": 26.84, - "tags": "distilled", - "date": "2025 July" - }, "lodestones Chroma1 v50 Preview Annealed": { "path": "vladmandic/chroma-unlocked-v50-annealed", - "preview": "lodestones--Chroma-annealed.jpg", + "preview": "vladmandic--chroma-unlocked-v50-annealed.jpg", "desc": "Chroma is a 8.9B parameter model based on FLUX.1-schnell. It’s fully Apache 2.0 licensed, ensuring that anyone can use, modify, and build on top of it—no corporate gatekeeping. Re-tweaked variant with extra noise added.", "skip": true, "extras": "", "size": 26.84, "date": "2025 July" }, - "lodestones Chroma1 v48 Preview": { - "path": "vladmandic/chroma-unlocked-v48", - "preview": "lodestones--Chroma.jpg", - "desc": "Chroma is a 8.9B parameter model based on FLUX.1-schnell. It’s fully Apache 2.0 licensed, ensuring that anyone can use, modify, and build on top of it—no corporate gatekeeping. Last raw version of Chroma before final finetuning.", - "skip": true, - "extras": "", - "size": 26.84, - "date": "2025 July" - }, - "lodestones Chroma1 v48 Preview Calibrated": { - "path": "vladmandic/chroma-unlocked-v48-detail-calibrated", - "preview": "lodestones--Chroma-detail.jpg", - "desc": "Chroma is a 8.9B parameter model based on FLUX.1-schnell. It’s fully Apache 2.0 licensed, ensuring that anyone can use, modify, and build on top of it—no corporate gatekeeping. Last raw version of Chroma before final finetuning but with some detail calibration.", - "skip": true, - "extras": "", - "size": 26.84, - "date": "2025 July" - }, "Meituan LongCat Image": { "path": "meituan-longcat/LongCat-Image", @@ -451,14 +355,6 @@ "desc": "SDXS: Real-Time One-Step Latent Diffusion Models with Image Conditions", "extras": "width: 512, height: 512, sampler: CMSI, steps: 1, cfg_scale: 0.0" }, - "SDXL Flash Mini": { - "path": "SDXL-Flash_Mini.safetensors@https://huggingface.co/sd-community/sdxl-flash-mini/resolve/main/SDXL-Flash_Mini.safetensors?download=true", - "preview": "SDXL-Flash_Mini.jpg", - "desc": "Introducing the new fast model SDXL Flash (Mini), we learned that all fast XL models work fast, but the quality decreases, and we also made a fast model, but it is not as fast as LCM, Turbo, Lightning and Hyper, but the quality is higher.", - "extras": "width: 2048, height: 1024, sampler: DEIS, steps: 40, cfg_scale: 6.0", - "tags": "distilled", - "experimental": true - }, "NVLabs Sana 1.5 1.6B 1k": { "path": "Efficient-Large-Model/SANA1.5_1.6B_1024px_diffusers", @@ -476,13 +372,6 @@ "size": 15.58, "date": "2025 March" }, - "NVLabs Sana 1.5 1.6B 1k Sprint": { - "path": "Efficient-Large-Model/Sana_Sprint_1.6B_1024px_diffusers", - "desc": "SANA-Sprint is an ultra-efficient diffusion model for text-to-image (T2I) generation, reducing inference steps from 20 to 1-4 while achieving state-of-the-art performance.", - "preview": "Efficient-Large-Model--Sana15_Sprint_1600M_1024px_diffusers.jpg", - "tags": "distilled", - "skip": true - }, "NVLabs Sana 1.0 1.6B 4k": { "path": "Efficient-Large-Model/Sana_1600M_4Kpx_BF16_diffusers", "desc": "Sana is a text-to-image framework that can efficiently generate images up to 4096 × 4096 resolution. Sana can synthesize high-resolution, high-quality images with strong text-image alignment at a remarkably fast speed, deployable on laptop GPU.", @@ -593,26 +482,6 @@ "size": 6.43, "date": "2023 November" }, - "Segmind SSD-1B": { - "path": "huggingface/segmind/SSD-1B", - "preview": "segmind--SSD-1B.jpg", - "desc": "The Segmind Stable Diffusion Model (SSD-1B) offers a compact, efficient, and distilled version of the SDXL model. At 50% smaller and 60% faster than Stable Diffusion XL (SDXL), it provides quick and seamless performance without sacrificing image quality.", - "variant": "fp16", - "skip": true, - "extras": "sampler: Default, cfg_scale: 9.0", - "size": 8.72, - "tags": "distilled", - "date": "2023 October" - }, - "Segmind Tiny": { - "path": "segmind/tiny-sd", - "preview": "segmind--tiny-sd.jpg", - "desc": "Segmind's Tiny-SD offers a compact, efficient, and distilled version of Realistic Vision 4.0 and is up to 80% faster than SD1.5", - "extras": "width: 512, height: 512, sampler: Default, cfg_scale: 9.0", - "size": 1.03, - "tags": "distilled", - "date": "2023 July" - }, "Segmind SegMoE SD 4x2": { "path": "segmind/SegMoE-SD-4x2-v0", "preview": "segmind--SegMoE-SD-4x2-v0.jpg", @@ -672,16 +541,6 @@ "size": 0, "date": "2025 August" }, - "Tencent HunyuanImage 2.1 Distilled": { - "path": "hunyuanvideo-community/HunyuanImage-2.1-Distilled-Diffusers", - "desc": "HunyuanImage-2.1, a highly efficient text-to-image model that is capable of generating 2K (2048 × 2048) resolution images.", - "preview": "hunyuanvideo-community--HunyuanImage-2.1-Distilled-Diffusers.jpg", - "extras": "", - "tags": "distilled", - "skip": true, - "size": 0, - "date": "2025 August" - }, "Tencent HunyuanImage 2.1 Refiner": { "path": "hunyuanvideo-community/HunyuanImage-2.1-Refiner-Diffusers", "desc": "HunyuanImage-2.1, a highly efficient text-to-image model that is capable of generating 2K (2048 × 2048) resolution images.", @@ -699,26 +558,12 @@ "size": 14.09, "date": "2024 May" }, - "Tencent HunyuanDiT 1.2 Distilled": { - "path": "Tencent-Hunyuan/HunyuanDiT-v1.2-Diffusers-Distilled", - "desc": "Hunyuan-DiT : A Powerful Multi-Resolution Diffusion Transformer with Fine-Grained Chinese Understanding.", - "preview": "Tencent-Hunyuan--HunyuanDiT-v1.2-Diffusers-Distilled.jpg", - "tags": "distilled", - "extras": "sampler: Default, cfg_scale: 2.0" - }, "Tencent HunyuanDiT 1.1": { "path": "Tencent-Hunyuan/HunyuanDiT-v1.1-Diffusers", "desc": "Hunyuan-DiT : A Powerful Multi-Resolution Diffusion Transformer with Fine-Grained Chinese Understanding.", "preview": "Tencent-Hunyuan--HunyuanDiT-v1.1-Diffusers.jpg", "extras": "sampler: Default, cfg_scale: 2.0" }, - "Tencent HunyuanDiT 1.1 Distilled": { - "path": "Tencent-Hunyuan/HunyuanDiT-v1.1-Diffusers-Distilled", - "desc": "Hunyuan-DiT : A Powerful Multi-Resolution Diffusion Transformer with Fine-Grained Chinese Understanding.", - "preview": "Tencent-Hunyuan--HunyuanDiT-v1.1-Diffusers-Distilled.jpg", - "tags": "distilled", - "extras": "sampler: Default, cfg_scale: 2.0" - }, "AlphaVLLM Lumina Next SFT": { "path": "Alpha-VLLM/Lumina-Next-SFT-diffusers", @@ -1002,338 +847,14 @@ "skip": true }, - "FLUX.1-Dev sdnq-svd-uint4": { - "path": "Disty0/FLUX.1-dev-SDNQ-uint4-svd-r32", - "preview": "Disty0--FLUX.1-dev-SDNQ-uint4-svd-r32.jpg", - "desc": "Quantization of black-forest-labs/FLUX.1-dev using SDNQ: sdnq-svd 4-bit uint with svd rank 32", + "ZAI GLM-Image": { + "path": "zai-org/GLM-Image", + "preview": "zai-org--GLM-Image.jpg", + "desc": "GLM-Image is a two-stage image generation model combining autoregressive token generation (9B vision-language encoder) with diffusion refinement (7B DiT transformer). Features strong text rendering and compositional capabilities.", "skip": true, - "tags": "quantized", - "size": 12.60, - "date": "2025 October", - "extras": "" - }, - "FLUX.1-Schnell sdnq-svd-uint4": { - "path": "Disty0/FLUX.1-schnell-SDNQ-uint4-svd-r32", - "preview": "Disty0--FLUX.1-schnell-SDNQ-uint4-svd-r32.jpg", - "desc": "Quantization of black-forest-labs/FLUX.1-schnell using SDNQ: sdnq-svd 4-bit uint with svd rank 32", - "skip": true, - "tags": "quantized", - "size": 12.60, - "date": "2025 October", - "extras": "" - }, - "FLUX.1-Dev Krea sdnq-svd-uint4": { - "path": "Disty0/FLUX.1-Krea-dev-SDNQ-uint4-svd-r32", - "preview": "Disty0--FLUX.1-Krea-dev-SDNQ-uint4-svd-r32.jpg", - "desc": "Quantization of black-forest-labs/FLUX.1-Krea-dev using SDNQ: sdnq-svd 4-bit uint with svd rank 32", - "skip": true, - "tags": "quantized", - "size": 12.60, - "date": "2025 October", - "extras": "" - }, - "FLUX.1-Dev Kontext sdnq-svd-uint4": { - "path": "Disty0/FLUX.1-Kontext-dev-SDNQ-uint4-svd-r32", - "preview": "Disty0--FLUX.1-Kontext-dev-SDNQ-uint4-svd-r32.jpg", - "desc": "Quantization of black-forest-labs/FLUX.1-Kontext-dev using SDNQ: sdnq-svd 4-bit uint with svd rank 32", - "skip": true, - "tags": "quantized", - "size": 12.60, - "date": "2025 October", - "extras": "" - }, - "FLUX.2 Dev sdnq-svd-uint4": { - "path": "Disty0/FLUX.2-dev-SDNQ-uint4-svd-r32", - "preview": "Disty0--FLUX.2-dev-SDNQ-uint4-svd-r32.jpg", - "desc": "Quantization of black-forest-labs/FLUX.2-dev using SDNQ: sdnq-svd 4-bit uint with svd rank 32", - "skip": true, - "tags": "quantized", - "extras": "", - "size": 31.58, - "date": "2025 November" - }, - "Chroma1-HD sdnq-svd-uint4": { - "path": "Disty0/Chroma1-HD-SDNQ-uint4-svd-r32", - "preview": "Disty0--Chroma1-HD-SDNQ-uint4-svd-r32.jpg", - "desc": "Quantization of lodestones/Chroma1-HD using SDNQ: sdnq-svd 4-bit uint with svd rank 32", - "skip": true, - "tags": "quantized", - "size": 11.89, - "date": "2025 October", - "extras": "" - }, - "Wan-AI Wan2.2 A14B T2I sdnq-svd-uint4": { - "path": "Disty0/Wan2.2-T2V-A14B-SDNQ-uint4-svd-r32", - "preview": "Wan-AI--Wan2.2-T2V-A14B-Diffusers.jpg", - "desc": "Quantization of black-forest-labs/FLUX.1-dev using SDNQ: sdnq-svd 4-bit uint with svd rank 32", - "skip": true, - "tags": "quantized", - "date": "2025 October", - "size": 23.54, - "extras": "" - }, - "Wan-AI Wan2.2 A14B I2I sdnq-svd-uint4": { - "path": "Disty0/Wan2.2-I2V-A14B-SDNQ-uint4-svd-r32", - "preview": "Wan-AI--Wan2.2-T2V-A14B-Diffusers.jpg", - "desc": "Quantization of Laxhar/noobai-XL-1.1 using SDNQ: sdnq-svd 4-bit uint with svd rank 128", - "skip": true, - "tags": "quantized", - "date": "2025 October", - "size": 23.55, - "extras": "" - }, - "Z-Image-Turbo sdnq-svd-uint4": { - "path": "Disty0/Z-Image-Turbo-SDNQ-uint4-svd-r32", - "preview": "Disty0--Z-Image-Turbo-SDNQ-uint4-svd-r32.jpg", - "desc": "Quantization of Tongyi-MAI/Z-Image-Turbo using SDNQ: sdnq-svd 4-bit uint with svd rank 32", - "skip": true, - "tags": "quantized", - "extras": "sampler: Default, cfg_scale: 1.0, steps: 9", - "size": 6.5, - "date": "2025 November" - }, - "Qwen-Image sdnq-svd-uint4": { - "path": "Disty0/Qwen-Image-SDNQ-uint4-svd-r32", - "preview": "Qwen--Qwen-Image.jpg", - "desc": "Quantization of Qwen/Qwen-Image using SDNQ: sdnq-svd 4-bit uint with svd rank 32", - "skip": true, - "tags": "quantized", - "date": "2025 October", - "size": 16.09, - "extras": "" - }, - "Qwen-Image-Edit sdnq-svd-uint4": { - "path": "Disty0/Qwen-Image-Edit-SDNQ-uint4-svd-r32", - "preview": "Qwen--Qwen-Image-Edit.jpg", - "desc": "Quantization of Qwen/Qwen-Image-Edit using SDNQ: sdnq-svd 4-bit uint with svd rank 32", - "skip": true, - "tags": "quantized", - "date": "2025 October", - "size": 16.10, - "extras": "" - }, - "Qwen-Image-Edit-2509 sdnq-svd-uint4": { - "path": "Disty0/Qwen-Image-Edit-2509-SDNQ-uint4-svd-r32", - "preview": "Qwen--Qwen-Image-Edit-2509.jpg", - "desc": "Quantization of Qwen/Qwen-Image-Edit-2509 using SDNQ: sdnq-svd 4-bit uint with svd rank 32", - "skip": true, - "tags": "quantized", - "date": "2025 October", - "size": 16.10, - "extras": "" - }, - "Qwen-Image-Edit-2511 sdnq-svd-uint4": { - "path": "Disty0/Qwen-Image-Edit-2511-SDNQ-uint4-svd-r32", - "preview": "Qwen--Qwen-Image-Edit-2509.jpg", - "desc": "Quantization of Qwen/Qwen-Image-Edit-2511 using SDNQ: sdnq-svd 4-bit uint with svd rank 32", - "skip": true, - "tags": "quantized", - "date": "2025 December", - "size": 16.10, - "extras": "" - }, - "Qwen-Image-Layered sdnq-svd-uint4": { - "path": "Disty0/Qwen-Image-Layered-SDNQ-uint4-svd-r32", - "preview": "Qwen--Qwen-Image-Edit-2509.jpg", - "desc": "Quantization of Qwen/Qwen-Image-Layered using SDNQ: sdnq-svd 4-bit uint with svd rank 32", - "skip": true, - "tags": "quantized", - "date": "2025 December", - "size": 16.10, - "extras": "" - }, - "nVidia ChronoEdit sdnq-svd-uint4": { - "path": "Disty0/ChronoEdit-14B-SDNQ-uint4-svd-r32", - "preview": "Disty0--ChronoEdit-14B-SDNQ-uint4-svd-r32.jpg", - "desc": "Quantization of nvidia/ChronoEdit-14B-Diffusers using SDNQ: sdnq-svd 4-bit uint with svd rank 32.", - "skip": true, - "tags": "quantized", - "date": "2025 October", - "size": 18.10, - "extras": "" - }, - "Tencent HunyuanImage 3.0 sdnq-svd-uint4": { - "path": "Disty0/HunyuanImage3-SDNQ-uint4-svd-r32", - "desc": "Quantization of tencent/HunyuanImage-3.0 using SDNQ: sdnq-svd 4-bit uint with svd rank 32.", - "preview": "Disty0--HunyuanImage3-SDNQ-uint4-svd-r32.jpg", - "extras": "", - "skip": true, - "tags": "quantized", - "size": 57.06, - "date": "2025 September" - }, - "Tempest-by-Vlad XL sdnq-svd-uint4": { - "path": "vladmandic/tempestByVlad_baseV01-SDNQ-uint4-svd", - "preview": "vladmandic--tempestByVlad_baseV01-SDNQ-uint4-svd.jpg", - "desc": "Quantization of vladmandic/tempestByVlad_baseV01 using SDNQ: sdnq-svd 4-bit uint with svd rank 128", - "tags": "quantized", - "size": 3.37, - "date": "2025 October", - "extras": "" - }, - "NoobAI-XL v1.1 epsilon sdnq-svd-uint4": { - "path": "Disty0/NoobAI-XL-v1.1-SDNQ-uint4-svd-r128", - "preview": "Disty0--NoobAI-XL-v1.1-SDNQ-uint4-svd-r128.jpg", - "desc": "Quantization of Laxhar/noobai-XL-1.1 using SDNQ: sdnq-svd 4-bit uint with svd rank 128", - "tags": "quantized", - "size": 3.37, - "date": "2025 October", - "extras": "" - }, - "NoobAI-XL v1.0 v-pred sdnq-svd-uint4": { - "path": "Disty0/NoobAI-XL-Vpred-v1.0-SDNQ-uint4-svd-r128", - "preview": "Disty0--NoobAI-XL-Vpred-v1.0-SDNQ-uint4-svd-r128.jpg", - "desc": "Quantization of Laxhar/noobai-XL-Vpred-1.0 using SDNQ: sdnq-svd 4-bit uint with svd rank 128", - "tags": "quantized", - "size": 3.37, - "date": "2025 October", - "extras": "" - }, - - "Tempest-by-Vlad XL": { - "path": "tempestByVlad_baseV01.safetensors@https://civitai.com/api/download/models/1301775", - "preview": "tempestByVlad_baseV01.jpg", - "desc": "Flexible SDXL model with custom encoder and finetuned for larger landscape resolutions with high details and high contrast.", - "tags": "community", - "size": 6.94, - "date": "2025 January", - "extras": "" - }, - "Tempest-by-Vlad XL Hyper": { - "path": "tempestByVlad_hyperV01.safetensors@https://civitai.com/api/download/models/1343512", - "preview": "tempestByVlad_hyperV01.jpg", - "desc": "Custom distilled variant with goal to get as-normal-as-possible model that works with low steps and guidance-free", - "tags": "community", - "size": 6.94, - "date": "2025 January", - "extras": "" - }, - "Juggernaut XL XI": { - "path": "juggernautXL_juggXIByRundiffusion.safetensors@https://civitai.com/api/download/models/782002", - "preview": "juggernautXL_juggXIByRundiffusion.jpg", - "desc": "Showcase finetuned model based on Stable diffusion XL", - "date": "2024 August", - "size": 6.94, - "tags": "community", - "extras": "sampler: DEIS, steps: 20, cfg_scale: 6.0" - }, - "Juggernaut XL XI Lightning": { - "path": "juggernautXL_juggXILightningByRD.safetensors@https://civitai.com/api/download/models/920957", - "preview": "juggernautXL_juggXILightningByRD.jpg", - "desc": "Showcase finetuned model based on Stable diffusion XL", - "date": "2024 August", - "size": 6.94, - "tags": "community", - "extras": "sampler: DPM SDE, steps: 6, cfg_scale: 2.0" - }, - "Juggernaut SD Reborn": { - "original": true, - "path": "juggernaut_reborn.safetensors@https://civitai.com/api/download/models/274039", - "preview": "juggernaut_reborn.jpg", - "desc": "Showcase finetuned model based on Stable diffusion 1.5", - "date": "2023 December", - "size": 2.28, - "tags": "community", - "extras": "width: 512, height: 512, sampler: DEIS, steps: 20, cfg_scale: 6.0" - }, - "WAI Illustrious XL v15": { - "path": "waiIllustriousSDXL_v150.safetensors@https://civitai.com/api/download/models/2167369", - "preview": "waiIllustriousSDXL_v150.jpg", - "desc": "", - "tags": "community", - "size": 6.94, - "date": "2025 August", - "extras": "" - }, - "Pony Realism XL v2.3": { - "path": "ponyRealism_V23.safetensors@https://civitai.com/api/download/models/1763661", - "preview": "ponyRealism_V23.jpg", - "desc": "", - "tags": "community", - "size": 6.94, - "date": "2025 May", - "extras": "" - }, - "NoobAI XL 1.0 V-Pred": { - "path": "noobaiXLNAIXL_vPred10Version.safetensors@https://huggingface.co/Laxhar/noobai-XL-Vpred-1.0/resolve/main/NoobAI-XL-Vpred-v1.0.safetensors", - "preview": "noobaiXLNAIXL_vPred10Version.jpg", - "desc": "", - "tags": "community", - "size": 6.94, - "date": "2024 December", - "extras": "" - }, - "NoobAI XL 1.1 Epsilon": { - "path": "noobaiXLNAIXL_epsilonPred11Version.safetensors@https://huggingface.co/Laxhar/noobai-XL-1.1/resolve/main/NoobAI-XL-v1.1.safetensors", - "preview": "noobaiXLNAIXL_epsilonPred11Version.jpg", - "desc": "", - "tags": "community", - "size": 6.94, - "date": "2024 November", - "extras": "" - }, - "WAI-Ani-Pony XL v14": { - "path": "waiANIPONYXL_v140.safetensors.safetensors@https://civitai.com/api/download/models/1767402", - "preview": "waiANIPONYXL_v140.jpg", - "desc": "", - "tags": "community", - "size": 6.94, - "date": "2025 May", - "extras": "" - }, - "Tiwaz CenKreChro": { - "path": "Tiwaz/CenKreChro", - "preview": "Tiwaz--CenKreChro.jpg", - "skip": true, - "desc": "Based Centerfold Flux 5, trying to merge in Chroma and Krea.", - "extras": "", - "tags": "community", - "date": "2025 September" - }, - "purplesmartai Pony 7": { - "path": "purplesmartai/pony-v7-base", - "preview": "purplesmartai--pony-v7-base.jpg", - "skip": true, - "desc": "Pony V7 is a versatile character generation model based on AuraFlow architecture. It supports a wide range of styles and species types (humanoid, anthro, feral, and more) and handles character interactions through natural language prompts.", - "extras": "", - "tags": "community", - "date": "October September" - }, - "ShuttleAI Shuttle 3.0 Diffusion": { - "path": "shuttleai/shuttle-3-diffusion", - "desc": "Shuttle uses Flux.1 Schnell as its base. It can produce images similar to Flux Dev or Pro in just 4 steps, and it is licensed under Apache 2. The model was partially de-distilled during training. When used beyond 10 steps, it enters refiner mode enhancing image details without altering the composition", - "preview": "shuttleai--shuttle-3-diffusion.jpg", - "tags": "community", - "skip": true - }, - "ShuttleAI Shuttle 3.1 Aesthetic": { - "path": "shuttleai/shuttle-3.1-aesthetic", - "desc": "Shuttle uses Flux.1 Schnell as its base. It can produce images similar to Flux Dev or Pro in just 4 steps, and it is licensed under Apache 2. The model was partially de-distilled during training. When used beyond 10 steps, it enters refiner mode enhancing image details without altering the composition", - "preview": "shuttleai--shuttle-3_1-aestetic.jpg", - "tags": "community", - "skip": true - }, - "ShuttleAI Shuttle Jaguar": { - "path": "shuttleai/shuttle-jaguar", - "desc": "Shuttle uses Flux.1 Schnell as its base. It can produce images similar to Flux Dev or Pro in just 4 steps, and it is licensed under Apache 2. The model was partially de-distilled during training. When used beyond 10 steps, it enters refiner mode enhancing image details without altering the composition", - "preview": "shuttleai--shuttle-jaguar.jpg", - "tags": "community", - "skip": true - }, - - "Google Gemini 2.5 Flash Nano Banana": { - "path": "gemini-2.5-flash-image", - "desc": "Gemini can generate and process images conversationally. You can prompt Gemini with text, images, or a combination of both allowing you to create, edit, and iterate on visuals with unprecedented control.", - "preview": "gemini-2.5-flash-image.jpg", - "tags": "cloud", - "skip": true - }, - "Google Gemini 3.0 Pro Nano Banana": { - "path": "gemini-3-pro-image-preview", - "desc": "Built on Gemini 3. Create and edit images with studio-quality levels of precision and control", - "preview": "gemini-3-pro-image-preview.jpg", - "tags": "cloud", - "skip": true + "extras": "sampler: Default, cfg_scale: 1.5, steps: 50", + "size": 15.3, + "date": "2025 January" } } diff --git a/installer.py b/installer.py index 94346883e..8a5f15b82 100644 --- a/installer.py +++ b/installer.py @@ -648,7 +648,7 @@ def check_diffusers(): t_start = time.time() if args.skip_all: return - sha = 'f6b6a7181eb44f0120b29cd897c129275f366c2a' # diffusers commit hash + sha = 'd7a1c31f4f85bae5a9e01cdce49bd7346bd8ccd6' # diffusers commit hash # if args.use_rocm or args.use_zluda or args.use_directml: # sha = '043ab2520f6a19fce78e6e060a68dbc947edb9f9' # lock diffusers versions for now pkg = pkg_resources.working_set.by_key.get('diffusers', None) @@ -678,9 +678,12 @@ def check_transformers(): if args.use_directml: target_transformers = '4.52.4' target_tokenizers = '0.21.4' + elif args.new: + target_transformers = '5.0.0rc2' + target_tokenizers = '0.22.2' else: - target_transformers = '4.57.3' - target_tokenizers = '0.22.1' + target_transformers = '4.57.5' + target_tokenizers = '0.22.2' if (pkg_transformers is None) or ((pkg_transformers.version != target_transformers) or (pkg_tokenizers is None) or ((pkg_tokenizers.version != target_tokenizers) and (not args.experimental))): if pkg_transformers is None: log.info(f'Transformers install: version={target_transformers}') @@ -757,7 +760,7 @@ def install_rocm_zluda(): msg = f'ROCm: version={rocm.version}' if device is not None: - msg += f', using agent {device.name}' + msg += f', using agent {device}' log.info(msg) if sys.platform == "win32": @@ -778,20 +781,20 @@ def install_rocm_zluda(): zluda_installer.install() zluda_installer.set_default_agent(device) except Exception as e: - log.warning(f'Failed to install ZLUDA: {e}') + log.error(f'Install ZLUDA: {e}') try: zluda_installer.load() except Exception as e: - log.warning(f'Failed to load ZLUDA: {e}') + log.error(f'Load ZLUDA: {e}') else: # TODO rocm: switch to pytorch source when it becomes available if device is None: - log.warning('No ROCm agent was found. Please make sure that graphics driver is installed and up to date.') + log.error('ROCm: no agent found - make sure that graphics driver is installed and up to date') if isinstance(rocm.environment, rocm.PythonPackageEnvironment): - check_python(supported_minors=[11, 12, 13], reason='ROCm backend requires a Python version between 3.11 and 3.13') + check_python(supported_minors=[11, 12, 13], reason='ROCm: python==3.11/3.12/3.13 required') torch_command = os.environ.get('TORCH_COMMAND', f'torch torchvision --index-url https://rocm.nightlies.amd.com/{device.therock}') else: - check_python(supported_minors=[12], reason='ROCm Windows preview requires Python version 3.12') + check_python(supported_minors=[12], reason='ROCm: Windows preview python==3.12 required') torch_command = os.environ.get('TORCH_COMMAND', '--no-cache-dir https://repo.radeon.com/rocm/windows/rocm-rel-6.4.4/torch-2.8.0a0%2Bgitfc14c65-cp312-cp312-win_amd64.whl https://repo.radeon.com/rocm/windows/rocm-rel-6.4.4/torchvision-0.24.0a0%2Bc85f008-cp312-cp312-win_amd64.whl') else: #check_python(supported_minors=[10, 11, 12, 13, 14], reason='ROCm backend requires a Python version between 3.10 and 3.13') @@ -818,12 +821,12 @@ def install_rocm_zluda(): log.warning("ROCm: minimum supported version=6.0") if device is None or os.environ.get("HSA_OVERRIDE_GFX_VERSION", None) is not None: - log.info(f'ROCm: HSA_OVERRIDE_GFX_VERSION auto config skipped: device={device.name if device is not None else None} version={os.environ.get("HSA_OVERRIDE_GFX_VERSION", None)}') + log.info(f'ROCm: HSA_OVERRIDE_GFX_VERSION auto config skipped: device={device} version={os.environ.get("HSA_OVERRIDE_GFX_VERSION", None)}') else: gfx_ver = device.get_gfx_version() if gfx_ver is not None and device.name.removeprefix("gfx") != gfx_ver.replace(".", ""): os.environ.setdefault('HSA_OVERRIDE_GFX_VERSION', gfx_ver) - log.info(f'ROCm: HSA_OVERRIDE_GFX_VERSION config overridden: device={device.name} version={os.environ.get("HSA_OVERRIDE_GFX_VERSION", None)}') + log.info(f'ROCm: HSA_OVERRIDE_GFX_VERSION config overridden: device={device} version={os.environ.get("HSA_OVERRIDE_GFX_VERSION", None)}') ts('amd', t_start) return torch_command @@ -946,7 +949,7 @@ def check_torch(): if not is_cuda_available and not is_ipex_available and allow_rocm: from modules import rocm - is_rocm_available = allow_rocm and (args.use_rocm or args.use_zluda or (len(rocm.agents) != 0 if sys.platform == "win32" else rocm.is_installed)) # late eval to avoid unnecessary import + is_rocm_available = allow_rocm and (args.use_rocm or args.use_zluda or rocm.is_installed) # late eval to avoid unnecessary import if is_cuda_available and args.use_cuda: # prioritize cuda torch_command = install_cuda() @@ -1397,6 +1400,7 @@ def set_environment(): log.debug('Setting environment tuning') os.environ.setdefault('ACCELERATE', 'True') os.environ.setdefault('ATTN_PRECISION', 'fp16') + os.environ.setdefault('ClDeviceGlobalMemSizeAvailablePercent', '100') os.environ.setdefault('CUDA_AUTO_BOOST', '1') os.environ.setdefault('CUDA_CACHE_DISABLE', '0') os.environ.setdefault('CUDA_DEVICE_DEFAULT_PERSISTING_L2_CACHE_PERCENTAGE_LIMIT', '0') @@ -1407,20 +1411,27 @@ def set_environment(): os.environ.setdefault('GRADIO_ANALYTICS_ENABLED', 'False') os.environ.setdefault('K_DIFFUSION_USE_COMPILE', '0') os.environ.setdefault('KINETO_LOG_LEVEL', '3') + os.environ.setdefault('NEOReadDebugKeys', '1') os.environ.setdefault('NUMEXPR_MAX_THREADS', '16') os.environ.setdefault('PYTHONHTTPSVERIFY', '0') + os.environ.setdefault('PYTORCH_ENABLE_MPS_FALLBACK', '1') + os.environ.setdefault('PYTORCH_ENABLE_XPU_FALLBACK', '1') + os.environ.setdefault('RUNAI_STREAMER_CHUNK_BYTESIZE', '2097152') + os.environ.setdefault('RUNAI_STREAMER_LOG_LEVEL', 'DEBUG' if os.environ.get('SD_LOAD_DEBUG') else 'WARNING') + os.environ.setdefault('RUNAI_STREAMER_MEMORY_LIMIT', '-1') os.environ.setdefault('SAFETENSORS_FAST_GPU', '1') + os.environ.setdefault('SYCL_CACHE_PERSISTENT', '1') os.environ.setdefault('TF_CPP_MIN_LOG_LEVEL', '2') os.environ.setdefault('TF_ENABLE_ONEDNN_OPTS', '0') + os.environ.setdefault('TOKENIZERS_PARALLELISM', '0') os.environ.setdefault('TORCH_CUDNN_V8_API_ENABLED', '1') os.environ.setdefault('TORCH_FORCE_NO_WEIGHTS_ONLY_LOAD', '1') + os.environ.setdefault('TORCH_ROCM_AOTRITON_ENABLE_EXPERIMENTAL', '1') + os.environ.setdefault('UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS', '1') os.environ.setdefault('USE_TORCH', '1') os.environ.setdefault('UV_INDEX_STRATEGY', 'unsafe-any-match') os.environ.setdefault('UV_NO_BUILD_ISOLATION', '1') os.environ.setdefault('UVICORN_TIMEOUT_KEEP_ALIVE', '60') - os.environ.setdefault('RUNAI_STREAMER_CHUNK_BYTESIZE', '2097152') - os.environ.setdefault('RUNAI_STREAMER_MEMORY_LIMIT', '-1') - os.environ.setdefault('RUNAI_STREAMER_LOG_LEVEL', 'DEBUG' if os.environ.get('SD_LOAD_DEBUG') else 'WARNING') allocator = f'garbage_collection_threshold:{opts.get("torch_gc_threshold", 80)/100:0.2f},max_split_size_mb:512' if opts.get("torch_malloc", "native") == 'cudaMallocAsync': allocator += ',backend:cudaMallocAsync' @@ -1430,14 +1441,6 @@ def set_environment(): os.environ.setdefault('PYTORCH_CUDA_ALLOC_CONF', allocator) os.environ.setdefault('PYTORCH_HIP_ALLOC_CONF', allocator) log.debug(f'Torch allocator: "{allocator}"') - os.environ.setdefault('TORCH_ROCM_AOTRITON_ENABLE_EXPERIMENTAL', '1') - os.environ.setdefault('NEOReadDebugKeys', '1') - os.environ.setdefault('ClDeviceGlobalMemSizeAvailablePercent', '100') - os.environ.setdefault('SYCL_CACHE_PERSISTENT', '1') - os.environ.setdefault('UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS', '1') - os.environ.setdefault('PYTORCH_ENABLE_XPU_FALLBACK', '1') - os.environ.setdefault('PYTORCH_ENABLE_MPS_FALLBACK', '1') - os.environ.setdefault('TOKENIZERS_PARALLELISM', '0') def check_extensions(): @@ -1644,6 +1647,7 @@ def check_version(reset=True): # pylint: disable=unused-argument # git('git stash pop') ver = git('log -1 --pretty=format:"%h %ad"') log.info(f'Repository upgraded: {ver}') + log.warning('Server restart is recommended to apply changes') if ver == latest: # double check restart() except Exception: diff --git a/javascript/black-gray.css b/javascript/black-gray.css index 784626d60..299875b1b 100644 --- a/javascript/black-gray.css +++ b/javascript/black-gray.css @@ -46,8 +46,8 @@ img { background-color: var(--background-color); } input[type=range] { height: var(--line-xs) !important; appearance: none !important; margin-top: 0 !important; min-width: max(4em, 100%) !important; width: 100% !important; background: transparent !important; } input[type=range]::-webkit-slider-runnable-track { width: 100% !important; height: var(--line-xs) !important; cursor: pointer !important; background: var(--input-background-fill) !important; border: 0px solid var(--primary-900) !important; } input[type=range]::-moz-range-track { width: 100% !important; height: var(--line-xs) !important; cursor: pointer !important; background: var(--input-background-fill) !important; border: 0px solid var(--primary-900) !important; } -input[type=range]::-webkit-slider-thumb { border: 0px solid var(--primary-950)0 !important; height: var(--line-sm) !important; width: var(--line-sm) !important; background: var(--highlight-color) !important; cursor: pointer !important; appearance: none !important; margin-top: -4px; border-radius: 4px; } -input[type=range]::-moz-range-thumb { border: 0px solid var(--primary-950)0 !important; height: var(--line-sm) !important; width: var(--line-sm) !important; background: var(--highlight-color) !important; cursor: pointer !important; appearance: none !important; margin-top: -4px; border-radius: 4px; } +input[type=range]::-webkit-slider-thumb { border: 0px solid var(--primary-950) !important; height: var(--line-sm) !important; width: var(--line-sm) !important; background: var(--highlight-color) !important; cursor: pointer !important; appearance: none !important; margin-top: -4px; border-radius: 4px; } +input[type=range]::-moz-range-thumb { border: 0px solid var(--primary-950) !important; height: var(--line-sm) !important; width: var(--line-sm) !important; background: var(--highlight-color) !important; cursor: pointer !important; appearance: none !important; margin-top: -4px; border-radius: 4px; } ::-webkit-scrollbar { width: 12px; height: 12px; } ::-webkit-scrollbar-track { background: var(--primary-800); } diff --git a/javascript/black-teal-reimagined.css b/javascript/black-teal-reimagined.css index 6f6b7e24e..df5b2ccc5 100644 --- a/javascript/black-teal-reimagined.css +++ b/javascript/black-teal-reimagined.css @@ -553,7 +553,7 @@ svg.feather.feather-image, } .output-html { - line-height: 1.2 rem; + line-height: 1.2rem; overflow-x: hidden; } @@ -1016,7 +1016,7 @@ svg.feather.feather-image, .loading { color: white; - position: border-box; + position: absolute; top: 85%; font-size: 1.5em; } @@ -1028,7 +1028,7 @@ svg.feather.feather-image, border-radius: 50%; border-top: var(--spacing-md) solid var(--primary-600); animation: spin 2s linear infinite, pulse 1.5s ease-in-out infinite; - position: border-box; + position: absolute; } .loader::before, diff --git a/javascript/black-teal.css b/javascript/black-teal.css index 13f2b5143..bcfe7f74d 100644 --- a/javascript/black-teal.css +++ b/javascript/black-teal.css @@ -43,12 +43,12 @@ --button-secondary-background-fill-hover: var(--neutral-600); --block-title-text-color: var(--neutral-300); --radius-xxs: 0; - --radius-xs: 1; + --radius-xs: 1px; --radius-sm: 2px; - --radius-md: 3; + --radius-md: 3px; --radius-lg: 4px; - --radius-xl: 5; - --radius-xxl: 6; + --radius-xl: 5px; + --radius-xxl: 6px; --line-xs: 1.0em; --line-sm: 1.2em; --line-md: 1.4em; @@ -77,8 +77,10 @@ button { max-width: 400px; white-space: nowrap; } img { background-color: var(--background-color); } input[type='range'] { display: block; margin: 0; padding: 0; height: 0.8em; background-color: transparent; overflow: hidden; cursor: pointer; box-shadow: 0 0 0 0 transparent; -webkit-appearance: none; appearance: none; } +/* eslint-disable-next-line css/no-invalid-properties */ input[type='range']::-webkit-slider-thumb { height: .9em; width: .9em; background-color: hsl(180, 54%, 61%); box-shadow: var(--range-shadow); border-radius: var(--radius-xs); } input[type='range']::-webkit-slider-runnable-track, input[type='range']::-webkit-slider-thumb { -webkit-appearance: none; } +/* eslint-disable-next-line css/no-invalid-properties */ input[type='range']::-moz-range-thumb { height: .9em; width: .9em; background-color: hsl(180, 54%, 61%); box-shadow: var(--range-shadow); } input[type='range']::-moz-range-track, input[type='range']::-webkit-slider-runnable-track { border: none; background: none; width: 100%; height: 100%; } diff --git a/javascript/control.js b/javascript/control.js index 3706a71b6..4d94fc643 100644 --- a/javascript/control.js +++ b/javascript/control.js @@ -3,7 +3,6 @@ function controlInputMode(inputMode, ...args) { if (updateEl) updateEl.click(); const tab = gradioApp().querySelector('#control-tab-input button.selected'); if (!tab) return ['Image', ...args]; - // let inputTab = tab.innerText; const tabs = Array.from(gradioApp().querySelectorAll('#control-tab-input button')); const tabIdx = tabs.findIndex((btn) => btn.classList.contains('selected')); const tabNames = ['Image', 'Video', 'Batch', 'Folder']; @@ -21,7 +20,7 @@ async function setupControlUI() { const tabs = ['input', 'output', 'preview']; for (const tab of tabs) { const btn = gradioApp().getElementById(`control-${tab}-button`); - if (!btn) continue; // eslint-disable-line no-continue + if (!btn) continue; btn.style.cursor = 'pointer'; btn.onclick = () => { const t = gradioApp().getElementById(`control-tab-${tab}`); diff --git a/javascript/extraNetworks.js b/javascript/extraNetworks.js index 6f3f4cdc9..acb305da8 100644 --- a/javascript/extraNetworks.js +++ b/javascript/extraNetworks.js @@ -140,39 +140,55 @@ async function filterExtraNetworksForTab(searchTerm) { const cards = Array.from(pg.querySelectorAll('.card') || []); items += cards.length; if (searchTerm === '' || searchTerm === 'all/') { - cards.forEach((elem) => elem.style.display = ''); + cards.forEach((elem) => { elem.style.display = ''; }); } else if (searchTerm === 'reference/') { - cards.forEach((elem) => elem.style.display = elem.dataset.name - .toLowerCase() - .includes('reference/') && elem.dataset.tags === '' ? '' : 'none'); + cards.forEach((elem) => { + elem.style.display = elem.dataset.name + .toLowerCase() + .includes('reference/') && elem.dataset.tags === '' ? '' : 'none'; + }); } else if (searchTerm === 'distilled/') { - cards.forEach((elem) => elem.style.display = elem.dataset.tags - .toLowerCase() - .includes('distilled') ? '' : 'none'); + cards.forEach((elem) => { + elem.style.display = elem.dataset.tags + .toLowerCase() + .includes('distilled') ? '' : 'none'; + }); } else if (searchTerm === 'community/') { - cards.forEach((elem) => elem.style.display = elem.dataset.tags - .toLowerCase() - .includes('community') ? '' : 'none'); + cards.forEach((elem) => { + elem.style.display = elem.dataset.tags + .toLowerCase() + .includes('community') ? '' : 'none'; + }); } else if (searchTerm === 'cloud/') { - cards.forEach((elem) => elem.style.display = elem.dataset.tags - .toLowerCase() - .includes('cloud') ? '' : 'none'); + cards.forEach((elem) => { + elem.style.display = elem.dataset.tags + .toLowerCase() + .includes('cloud') ? '' : 'none'; + }); } else if (searchTerm === 'quantized/') { - cards.forEach((elem) => elem.style.display = elem.dataset.tags - .toLowerCase() - .includes('quantized') ? '' : 'none'); + cards.forEach((elem) => { + elem.style.display = elem.dataset.tags + .toLowerCase() + .includes('quantized') ? '' : 'none'; + }); } else if (searchTerm === 'local/') { - cards.forEach((elem) => elem.style.display = elem.dataset.name - .toLowerCase() - .includes('reference/') ? 'none' : ''); + cards.forEach((elem) => { + elem.style.display = elem.dataset.name + .toLowerCase() + .includes('reference/') ? 'none' : ''; + }); } else if (searchTerm === 'diffusers/') { - cards.forEach((elem) => elem.style.display = elem.dataset.name - .toLowerCase().replace('models--', 'diffusers').replaceAll('\\', '/') - .includes('diffusers/') ? '' : 'none'); + cards.forEach((elem) => { + elem.style.display = elem.dataset.name + .toLowerCase().replace('models--', 'diffusers').replaceAll('\\', '/') + .includes('diffusers/') ? '' : 'none'; + }); } else if (searchTerm.startsWith('r#')) { searchTerm = searchTerm.substring(2); const re = new RegExp(searchTerm, 'i'); - cards.forEach((elem) => elem.style.display = re.test(`filename: ${elem.dataset.filename}|name: ${elem.dataset.name}|tags: ${elem.dataset.tags}`) ? '' : 'none'); + cards.forEach((elem) => { + elem.style.display = re.test(`filename: ${elem.dataset.filename}|name: ${elem.dataset.name}|tags: ${elem.dataset.tags}`) ? '' : 'none'; + }); } else { const searchList = searchTerm.split('|').filter((s) => s !== '' && !s.startsWith('-')).map((s) => s.trim()); const excludeList = searchTerm.split('|').filter((s) => s !== '' && s.trim().startsWith('-')).map((s) => s.trim().substring(1).trim()); @@ -198,8 +214,8 @@ async function filterExtraNetworksForTab(searchTerm) { } function tryToRemoveExtraNetworkFromPrompt(textarea, text) { - const re_extranet = /<([^:]+:[^:]+):[\d\.]+>/; - const re_extranet_g = /\s+<([^:]+:[^:]+):[\d\.]+>/g; + const re_extranet = /<([^:]+:[^:]+):[\d.]+>/; + const re_extranet_g = /\s+<([^:]+:[^:]+):[\d.]+>/g; let m = text.match(re_extranet); let replaced = false; let newTextareaText; @@ -243,7 +259,7 @@ function sortExtraNetworks(fixed = 'no') { const cards = Array.from(pg.querySelectorAll('.card') || []); if (cards.length === 0) return 'sort: no cards'; num += cards.length; - cards.sort((a, b) => { // eslint-disable-line no-loop-func + cards.sort((a, b) => { switch (sortVal) { case 0: return 0; case 1: return a.dataset.name ? a.dataset.name.localeCompare(b.dataset.name) : 0; @@ -310,21 +326,21 @@ function extraNetworksSearchButton(event) { function extraNetworksFilterVersion(event) { const version = event.target.textContent.trim(); - const activeTab = getENActiveTab(); const activePage = getENActivePage().toLowerCase(); - let cardContainer = gradioApp().querySelector(`#${activeTab}_${activePage}_cards`); - if (!cardContainer) cardContainer = gradioApp().querySelector(`#txt2img_extra_networks_${activePage}_cards`); - log('extraNetworksFilterVersion', { version, activeTab, activePage, cardContainer }); - if (!cardContainer) return; - if (cardContainer.dataset.activeVersion === version) { - cardContainer.dataset.activeVersion = ''; - cardContainer.querySelectorAll('.card').forEach((card) => card.style.display = ''); - } else { - cardContainer.dataset.activeVersion = version; - cardContainer.querySelectorAll('.card').forEach((card) => { - if (card.dataset.version === version) card.style.display = ''; - else card.style.display = 'none'; - }); + const cardContainers = gradioApp().querySelectorAll('.extra-network-cards'); + log('extraNetworksFilterVersion', { activePage, version }); + for (const cardContainer of cardContainers) { + if (!cardContainer.id.includes(activePage)) continue; + if (cardContainer.dataset.activeVersion === version) { + cardContainer.dataset.activeVersion = ''; + cardContainer.querySelectorAll('.card').forEach((card) => { card.style.display = ''; }); + } else { + cardContainer.dataset.activeVersion = version; + cardContainer.querySelectorAll('.card').forEach((card) => { + if (card.dataset.version === version) card.style.display = ''; + else card.style.display = 'none'; + }); + } } } diff --git a/javascript/gallery.js b/javascript/gallery.js index fc36f5152..fb81bb58a 100644 --- a/javascript/gallery.js +++ b/javascript/gallery.js @@ -1,12 +1,12 @@ /* eslint-disable max-classes-per-file */ -/* eslint lines-between-class-members: ["error", "always", { "exceptAfterSingleLine": true }] */ let ws; let url; -let currentImage; +let currentImage = null; let pruneImagesTimer; let outstanding = 0; let lastSort = 0; let lastSortName = 'None'; +let gallerySelection = { files: [], index: -1 }; const galleryHashes = new Set(); let maintenanceController = new AbortController(); const folderStylesheet = new CSSStyleSheet(); @@ -24,6 +24,64 @@ const el = { const SUPPORTED_EXTENSIONS = ['jpg', 'jpeg', 'png', 'webp', 'tiff', 'jp2', 'jxl', 'gif', 'mp4', 'mkv', 'avi', 'mjpeg', 'mpg', 'avr']; +function getVisibleGalleryFiles() { + if (!el.files) return []; + return Array.from(el.files.children).filter((node) => node.name && node.offsetParent); +} + +function updateGallerySelectionClasses(files = gallerySelection.files, index = gallerySelection.index) { + files.forEach((file, i) => { + file.classList.toggle('gallery-file-selected', i === index); + }); +} + +function refreshGallerySelection() { + updateGallerySelectionClasses(gallerySelection.files, -1); + const files = getVisibleGalleryFiles(); + const index = files.findIndex((file) => file.src === currentImage); + gallerySelection = { files, index }; + updateGallerySelectionClasses(files, index); +} + +function resetGallerySelection() { + updateGallerySelectionClasses(gallerySelection.files, -1); + gallerySelection = { files: [], index: -1 }; + currentImage = null; +} + +function applyGallerySelection(index, { send = true } = {}) { + if (!gallerySelection.files.length) refreshGallerySelection(); + const files = gallerySelection.files; + if (!files.length) return; + if (!Number.isInteger(index) || index < 0 || index >= files.length) { + log('gallery selection index out of range', index, files.length); + resetGallerySelection(); + return; + } + gallerySelection.index = index; + currentImage = files[index].src; + updateGallerySelectionClasses(files, index); + if (send && el.btnSend) el.btnSend.click(); +} + +function setGallerySelectionByElement(element, options) { + if (!gallerySelection.files.length) refreshGallerySelection(); + let index = gallerySelection.files.findIndex((file) => file === element); + if (index < 0) { + refreshGallerySelection(); + index = gallerySelection.files.findIndex((file) => file === element); + } + if (index >= 0) applyGallerySelection(index, options); +} + +function buildGalleryFileUrl(path) { + return new URL(`/file=${encodeURI(path)}`, window.location.origin).toString(); +} + +window.getGallerySelection = () => ({ index: gallerySelection.index, files: gallerySelection.files }); +window.setGallerySelection = (index, options) => applyGallerySelection(index, options); +window.getGallerySelectedUrl = () => (currentImage ? buildGalleryFileUrl(currentImage) : null); + /** * Wait for the `outstanding` variable to be below the specified value * @param {number} num - Threshold for `outstanding` @@ -103,11 +161,91 @@ function updateGalleryStyles() { .gallery-file:hover { filter: grayscale(100%); } + :host(.gallery-file-selected) .gallery-file { + box-shadow: 0 0 0 2px var(--sd-button-selected-color); + } `); } // Classes +class SimpleProgressBar { + #container = document.createElement('div'); + #progress = document.createElement('div'); + #textDiv = document.createElement('div'); + #text = document.createElement('span'); + #visible = false; + #hideTimeout = null; + #interval = null; + #max = 0; + /** @type {Set} */ + #monitoredSet; + + constructor(monitoredSet) { + this.#monitoredSet = monitoredSet; // This is required because incrementing a variable with a class method turned out to not be an atomic operation + this.#container.style.cssText = 'position:relative;overflow:hidden;border-radius:var(--sd-border-radius);width:100%;background-color:hsla(0,0%,36%,0.3);height:1.2rem;margin:0;padding:0;display:none;'; + this.#progress.style.cssText = 'position:absolute;left:0;height:100%;width:0;transition:width 200ms;'; + this.#progress.style.backgroundColor = 'hsla(110, 32%, 35%, 0.80)'; // alt: '#27911d' + this.#textDiv.style.cssText = 'position:relative;margin:auto;width:max-content;height:100%;'; + this.#text.style.cssText = 'user-select:none;color:white;'; + + this.#textDiv.append(this.#text); + this.#container.append(this.#progress, this.#textDiv); + } + + start(total) { + this.clear(); + this.#max = total; + this.#interval = setInterval(() => { + this.#update(this.#monitoredSet.size, this.#max); + }, 250); + } + + attachTo(element) { + if (element.hasChildNodes) { + element.innerHTML = ''; + } + element.appendChild(this.#container); + } + + clear() { + this.#stop(); + clearTimeout(this.#hideTimeout); + this.#hideTimeout = null; + this.#container.style.display = 'none'; + this.#visible = false; + this.#progress.style.width = '0'; + this.#text.textContent = ''; + } + + #update(loaded, max) { + if (this.#hideTimeout) { + this.#hideTimeout = null; + } + + this.#progress.style.width = `${Math.floor((loaded / max) * 100)}%`; + this.#text.textContent = `${loaded}/${max}`; + + if (!this.#visible) { + this.#container.style.display = 'block'; + this.#visible = true; + } + if (loaded >= max) { + this.#stop(); + this.#hideTimeout = setTimeout(() => { + this.clear(); + }, 1000); + } + } + + #stop() { + clearInterval(this.#interval); + this.#interval = null; + } +} + +const galleryProgressBar = new SimpleProgressBar(galleryHashes); + /* This isn't as robust as the Web Locks API, but it will at least work if accessing a remote machine without HTTPS */ class SimpleFunctionQueue { #id; @@ -163,9 +301,16 @@ class SimpleFunctionQueue { // HTML Elements class GalleryFolder extends HTMLElement { - constructor(name) { + constructor(folder) { super(); - this.name = decodeURI(name); + // Support both old format (string) and new format (object with path and label) + if (typeof folder === 'object' && folder !== null) { + this.name = decodeURI(folder.path || ''); + this.label = decodeURI(folder.label || folder.path || ''); + } else { + this.name = decodeURI(folder); + this.label = this.name; + } this.style.overflowX = 'hidden'; this.shadow = this.attachShadow({ mode: 'open' }); this.shadow.adoptedStyleSheets = [folderStylesheet]; @@ -174,8 +319,8 @@ class GalleryFolder extends HTMLElement { connectedCallback() { const div = document.createElement('div'); div.className = 'gallery-folder'; - div.innerHTML = `\uf03e ${this.name}`; - div.title = this.name; + div.innerHTML = `\uf03e ${this.label}`; + div.title = this.name; // Show full path on hover div.addEventListener('click', () => { for (const folder of el.folders.children) { if (folder.name === this.name) folder.shadow.firstElementChild.classList.add('gallery-folder-selected'); @@ -220,7 +365,7 @@ async function handleSeparator(separator) { if (!f.name) continue; // Skip separators // Check if file belongs to this exact directory - const fileDir = f.name.match(/(.*)[\/\\]/); + const fileDir = f.name.match(/(.*)[/\\]/); const fileDirPath = fileDir ? fileDir[1] : ''; if (separator.title.length > 0 && fileDirPath === separator.title) { @@ -232,14 +377,18 @@ async function handleSeparator(separator) { } async function addSeparators() { - document.querySelectorAll('.gallery-separator').forEach((node) => el.files.removeChild(node)); + document.querySelectorAll('.gallery-separator').forEach((node) => { el.files.removeChild(node); }); const all = Array.from(el.files.children); let lastDir; - let isFirstSeparator = true; // Flag to open the first separator by default + + // Count root files (files without a directory path) + const hasRootFiles = all.some((f) => f.name && !f.name.match(/[/\\]/)); + // Only auto-open first separator if there are no root files to display + let isFirstSeparator = !hasRootFiles; // First pass: create separators for (const f of all) { - let dir = f.name?.match(/(.*)[\/\\]/); + let dir = f.name?.match(/(.*)[/\\]/); if (!dir) dir = ''; else dir = dir[1]; if (dir !== lastDir) { @@ -249,7 +398,7 @@ async function addSeparators() { let fileCount = 0; for (const file of all) { if (!file.name) continue; - const fileDir = file.name.match(/(.*)[\/\\]/); + const fileDir = file.name.match(/(.*)[/\\]/); const fileDirPath = fileDir ? fileDir[1] : ''; if (fileDirPath === dir) fileCount++; } @@ -299,7 +448,7 @@ async function addSeparators() { for (const f of all) { if (!f.name) continue; // Skip separators - const dir = f.name.match(/(.*)[\/\\]/); + const dir = f.name.match(/(.*)[/\\]/); if (dir && dir[1]) { const dirPath = dir[1]; const isOpen = separatorStates.get(dirPath); @@ -357,7 +506,7 @@ class GalleryFile extends HTMLElement { } // Check separator state early to hide the element immediately - const dir = this.name.match(/(.*)[\/\\]/); + const dir = this.name.match(/(.*)[/\\]/); if (dir && dir[1]) { const dirPath = dir[1]; const isOpen = separatorStates.get(dirPath); @@ -366,7 +515,9 @@ class GalleryFile extends HTMLElement { } } - this.hash = await getHash(`${this.folder}/${this.name}/${this.size}/${this.mtime}`); // eslint-disable-line no-use-before-define + // Normalize path to ensure consistent hash regardless of which folder view is used + const normalizedPath = this.src.replace(/\/+/g, '/').replace(/\/$/, ''); + this.hash = await getHash(`${normalizedPath}/${this.size}/${this.mtime}`); // eslint-disable-line no-use-before-define const cachedData = (this.hash && opts.browser_cache) ? await idbGet(this.hash).catch(() => undefined) : undefined; const img = document.createElement('img'); img.className = 'gallery-file'; @@ -402,9 +553,11 @@ class GalleryFile extends HTMLElement { this.size = json.size; this.mtime = new Date(json.mtime); if (opts.browser_cache) { + // Store file's actual parent directory (not browsed folder) for consistent cleanup + const fileDir = this.src.replace(/\/+/g, '/').replace(/\/[^/]+$/, ''); await idbAdd({ hash: this.hash, - folder: this.folder, + folder: fileDir, file: this.name, size: this.size, mtime: this.mtime, @@ -430,8 +583,7 @@ class GalleryFile extends HTMLElement { return; } // ... to here unless modifications are also being made to maintenance functionality and the usage of AbortController/AbortSignal img.onclick = () => { - currentImage = this.src; - el.btnSend.click(); + setGallerySelectionByElement(this, { send: true }); }; img.title = `Folder: ${this.folder}\nFile: ${this.name}\nSize: ${this.size.toLocaleString()} bytes\nModified: ${this.mtime.toLocaleString()}`; if (this.shadow.children.length > 0) { @@ -538,7 +690,7 @@ async function wsConnect(socket, timeout = 5000) { let loop = 0; while (socket.readyState === WebSocket.CONNECTING && loop < ttl) { - await new Promise((resolve) => setTimeout(resolve, intrasleep)); // eslint-disable-line no-promise-executor-return + await new Promise((resolve) => { setTimeout(resolve, intrasleep); }); loop++; } return isOpened(); @@ -569,7 +721,7 @@ async function gallerySearch() { }); allFiles.forEach((f) => { - const dir = f.name.match(/(.*)[\/\\]/); + const dir = f.name.match(/(.*)[/\\]/); const dirPath = (dir && dir[1]) ? dir[1] : ''; const isOpen = separatorStates.get(dirPath); f.style.display = (!dirPath || isOpen) ? 'unset' : 'none'; @@ -603,7 +755,7 @@ async function gallerySearch() { if (isMatch) { fileMatches.add(f); totalFound++; - const dir = f.name.match(/(.*)[\/\\]/); + const dir = f.name.match(/(.*)[/\\]/); const dirPath = (dir && dir[1]) ? dir[1] : ''; directoryMatches.set(dirPath, (directoryMatches.get(dirPath) || 0) + 1); } @@ -634,6 +786,7 @@ async function gallerySearch() { const t1 = performance.now(); updateStatusWithSort('Filter', ['Images', `${totalFound.toLocaleString()} / ${allFiles.length.toLocaleString()}`], `${iconStopwatch} ${Math.floor(t1 - t0).toLocaleString()}ms`); + refreshGallerySelection(); }, 250); } @@ -653,59 +806,84 @@ async function gallerySort(btn) { if (arr.length === 0) return; // no files to sort if (btn) lastSort = btn.charCodeAt(0); const fragment = document.createDocumentFragment(); + + // Helper to get directory path from a file node + const getDirPath = (node) => { + const match = node.name.match(/(.*)[/\\]/); + return match ? match[1] : ''; + }; + + // Partition into root files and subfolder files - root files always stay at top + const rootFiles = arr.filter((node) => !getDirPath(node)); + const subfolderFiles = arr.filter((node) => getDirPath(node)); + + // Group subfolder files by directory + const folderGroups = new Map(); + for (const file of subfolderFiles) { + const dir = getDirPath(file); + if (!folderGroups.has(dir)) { + folderGroups.set(dir, []); + } + folderGroups.get(dir).push(file); + } + + // Sort function based on current sort mode + let sortFn; switch (lastSort) { case 61789: // name asc lastSortName = 'Name Ascending'; - arr - .sort((a, b) => a.name.localeCompare(b.name)) - .forEach((node) => fragment.appendChild(node)); + sortFn = (a, b) => a.name.localeCompare(b.name); break; case 61790: // name dsc lastSortName = 'Name Descending'; - arr - .sort((b, a) => a.name.localeCompare(b.name)) - .forEach((node) => fragment.appendChild(node)); + sortFn = (a, b) => b.name.localeCompare(a.name); break; case 61792: // size asc lastSortName = 'Size Ascending'; - arr - .sort((a, b) => a.size - b.size) - .forEach((node) => fragment.appendChild(node)); + sortFn = (a, b) => a.size - b.size; break; case 61793: // size dsc lastSortName = 'Size Descending'; - arr - .sort((b, a) => a.size - b.size) - .forEach((node) => fragment.appendChild(node)); + sortFn = (a, b) => b.size - a.size; break; case 61794: // resolution asc lastSortName = 'Resolution Ascending'; - arr - .sort((a, b) => a.width * a.height - b.width * b.height) - .forEach((node) => fragment.appendChild(node)); + sortFn = (a, b) => a.width * a.height - b.width * b.height; break; case 61795: // resolution dsc lastSortName = 'Resolution Descending'; - arr - .sort((b, a) => a.width * a.height - b.width * b.height) - .forEach((node) => fragment.appendChild(node)); + sortFn = (a, b) => b.width * b.height - a.width * a.height; break; case 61662: lastSortName = 'Modified Ascending'; - arr - .sort((a, b) => a.mtime - b.mtime) - .forEach((node) => fragment.appendChild(node)); + sortFn = (a, b) => a.mtime - b.mtime; break; case 61661: lastSortName = 'Modified Descending'; - arr - .sort((b, a) => a.mtime - b.mtime) - .forEach((node) => fragment.appendChild(node)); + sortFn = (a, b) => b.mtime - a.mtime; break; default: lastSortName = 'None'; + sortFn = null; break; } + + // Sort root files + if (sortFn) { + rootFiles.sort(sortFn); + } + rootFiles.forEach((node) => fragment.appendChild(node)); + + // Sort folder names alphabetically, then sort files within each folder + const sortedFolderNames = Array.from(folderGroups.keys()).sort((a, b) => a.localeCompare(b)); + for (const folderName of sortedFolderNames) { + const files = folderGroups.get(folderName); + if (sortFn) { + files.sort(sortFn); + } + files.forEach((node) => fragment.appendChild(node)); + } + if (fragment.children.length === 0) return; el.files.innerHTML = ''; el.files.appendChild(fragment); @@ -716,7 +894,7 @@ async function gallerySort(btn) { for (const f of all) { if (!f.name) continue; // Skip separators - const dir = f.name.match(/(.*)[\/\\]/); + const dir = f.name.match(/(.*)[/\\]/); if (dir && dir[1]) { const dirPath = dir[1]; const isOpen = separatorStates.get(dirPath); @@ -729,6 +907,7 @@ async function gallerySort(btn) { const t1 = performance.now(); log(`gallerySort: char=${lastSort} len=${arr.length} time=${Math.floor(t1 - t0)} sort=${lastSortName}`); updateStatusWithSort(['Images', arr.length.toLocaleString()], `${iconStopwatch} ${Math.floor(t1 - t0).toLocaleString()}ms`); + refreshGallerySelection(); } /** @@ -856,12 +1035,15 @@ async function fetchFilesHT(evt, controller) { } } + if (controller.signal.aborted) return; el.files.appendChild(fragment); const t1 = performance.now(); log(`gallery: folder=${evt.target.name} num=${numFiles} time=${Math.floor(t1 - t0)}ms`); updateStatusWithSort(['Folder', evt.target.name], ['Images', numFiles.toLocaleString()], `${iconStopwatch} ${Math.floor(t1 - t0).toLocaleString()}ms`); + galleryProgressBar.start(numFiles); addSeparators(); + refreshGallerySelection(); thumbCacheCleanup(evt.target.name, numFiles, controller); } @@ -871,6 +1053,8 @@ async function fetchFilesWS(evt) { // fetch file-by-file list over websockets maintenanceController.abort('Gallery update'); // Abort previous controller maintenanceController = controller; // Point to new controller for next time galleryHashes.clear(); // Must happen AFTER the AbortController steps + galleryProgressBar.clear(); + resetGallerySelection(); el.files.innerHTML = ''; updateGalleryStyles(); @@ -915,11 +1099,14 @@ async function fetchFilesWS(evt) { // fetch file-by-file list over websockets } }; ws.onclose = (event) => { + if (controller.signal.aborted) return; el.files.appendChild(fragment); // gallerySort(); log(`gallery: folder=${evt.target.name} num=${numFiles} time=${Math.floor(t1 - t0)}ms`); updateStatusWithSort(['Folder', evt.target.name], ['Images', numFiles.toLocaleString()], `${iconStopwatch} ${Math.floor(t1 - t0).toLocaleString()}ms`); + galleryProgressBar.start(numFiles); addSeparators(); + refreshGallerySelection(); thumbCacheCleanup(evt.target.name, numFiles, controller); }; ws.onerror = (event) => { @@ -973,7 +1160,8 @@ async function monitorGalleries() { async function setOverlayAnimation() { const busyAnimation = document.createElement('style'); - busyAnimation.textContent = '.idbBusyAnim{width:16px;height:16px;border-radius:50%;display:block;margin:40px;position:relative;background:#ff3d00;color:#fff;box-shadow:-24px 0,24px 0;box-sizing:border-box;animation:2s ease-in-out infinite overlayRotation}@keyframes overlayRotation{0%{transform:rotate(0)}100%{transform:rotate(360deg)}}'; // eslint-disable-line max-len + // eslint-disable-next-line @stylistic/max-len + busyAnimation.textContent = '.idbBusyAnim{width:16px;height:16px;border-radius:50%;display:block;margin:40px;position:relative;background:#ff3d00;color:#fff;box-shadow:-24px 0,24px 0;box-sizing:border-box;animation:2s ease-in-out infinite overlayRotation}@keyframes overlayRotation{0%{transform:rotate(0)}100%{transform:rotate(360deg)}}'; document.head.append(busyAnimation); } @@ -990,6 +1178,12 @@ async function initGallery() { // triggered on gradio change to monitor when ui updateGalleryStyles(); injectGalleryStatusCSS(); setOverlayAnimation(); + const progress = gradioApp().getElementById('tab-gallery-progress'); + if (progress) { + galleryProgressBar.attachTo(progress); + } else { + log('initGallery', 'Failed to attach loading progress bar'); + } el.search.addEventListener('input', gallerySearch); el.btnSend = gradioApp().getElementById('tab-gallery-send-image'); document.getElementById('tab-gallery-files').style.height = opts.logmonitor_show ? '75vh' : '85vh'; diff --git a/javascript/gpu.js b/javascript/gpu.js index 6e73c7fc9..9faaf9a80 100644 --- a/javascript/gpu.js +++ b/javascript/gpu.js @@ -1,9 +1,9 @@ -let gpuInterval = null; // eslint-disable-line prefer-const +let gpuInterval = null; const chartData = { mem: [], load: [] }; async function updateGPUChart(mem, load) { const maxLen = 120; - const colorRangeMap = $.range_map({ // eslint-disable-line no-undef + const colorRangeMap = $.range_map({ '0:5': '#fffafa', '6:10': '#fff7ed', '11:20': '#fed7aa', @@ -22,8 +22,8 @@ async function updateGPUChart(mem, load) { chartData.load.push(load); if (chartData.mem.length > maxLen) chartData.mem.shift(); chartData.mem.push(mem); - $('#gpuChart').sparkline(chartData.load, sparklineConfigLOAD); // eslint-disable-line no-undef - $('#gpuChart').sparkline(chartData.mem, sparklineConfigMEM); // eslint-disable-line no-undef + $('#gpuChart').sparkline(chartData.load, sparklineConfigLOAD); + $('#gpuChart').sparkline(chartData.mem, sparklineConfigMEM); } async function updateGPU() { diff --git a/javascript/imageViewer.js b/javascript/imageViewer.js index 8788ca689..53caa3c6c 100644 --- a/javascript/imageViewer.js +++ b/javascript/imageViewer.js @@ -32,6 +32,7 @@ function closeModal(evt, force = false) { } function modalImageSwitch(offset) { + const negmod = (n, m) => ((n % m) + m) % m; const galleryButtons = all_gallery_buttons(); if (galleryButtons.length > 1) { const currentButton = selected_gallery_button(); @@ -39,7 +40,6 @@ function modalImageSwitch(offset) { galleryButtons.forEach((v, i) => { if (v === currentButton) result = i; }); - const negmod = (n, m) => ((n % m) + m) % m; if (result !== -1) { const nextButton = galleryButtons[negmod((result + offset), galleryButtons.length)]; nextButton.click(); @@ -47,8 +47,24 @@ function modalImageSwitch(offset) { const modal = gradioApp().getElementById('lightboxModal'); modalImage.src = nextButton.children[0].src; if (modalImage.style.display === 'none') modal.style.setProperty('background-image', `url(${modalImage.src})`); + return; } } + + const galleryFilesContainer = gradioApp().getElementById('tab-gallery-files'); + if (!galleryFilesContainer || !galleryFilesContainer.offsetParent) return; + const gallerySelection = window.getGallerySelection(); + if (!gallerySelection.files.length || gallerySelection.files.length <= 1) return; + const baseIndex = gallerySelection.index >= 0 ? gallerySelection.index : 0; + const nextIndex = negmod((baseIndex + offset), gallerySelection.files.length); + window.setGallerySelection(nextIndex, { send: true }); + const modalImage = gradioApp().getElementById('modalImage'); + const modal = gradioApp().getElementById('lightboxModal'); + const directSrc = window.getGallerySelectedUrl(); + if (modalImage && modal && directSrc) { + modalImage.src = directSrc; + if (modalImage.style.display === 'none') modal.style.setProperty('background-image', `url(${directSrc})`); + } } function modalSaveImage(event) { diff --git a/javascript/indexdb.js b/javascript/indexdb.js index 21fbdfa61..ace3d5843 100644 --- a/javascript/indexdb.js +++ b/javascript/indexdb.js @@ -150,7 +150,10 @@ async function idbFolderCleanup(keepSet, folder, signal) { throw new Error('IndexedDB cleaning function must be told the current active folder'); } - let removals = new Set(await idbGetAllKeys('folder', folder)); + // Use range query to match folder and all its subdirectories + const folderNormalized = folder.replace(/\/+/g, '/').replace(/\/$/, ''); + const range = IDBKeyRange.bound(folderNormalized, `${folderNormalized}\uffff`, false, true); + let removals = new Set(await idbGetAllKeys('folder', range)); removals = removals.difference(keepSet); // Don't need to keep full set in memory const totalRemovals = removals.size; if (signal.aborted) { diff --git a/javascript/logMonitor.js b/javascript/logMonitor.js index c3bcf5673..ee3bf9369 100644 --- a/javascript/logMonitor.js +++ b/javascript/logMonitor.js @@ -56,7 +56,7 @@ async function logMonitor() { if (imgGallery) imgGallery.style.height = opts.logmonitor_show ? '50vh' : '55vh'; if (!opts.logmonitor_show) { - Array.from(document.getElementsByClassName('log-monitor')).forEach((el) => el.style.display = 'none'); + Array.from(document.getElementsByClassName('log-monitor')).forEach((el) => { el.style.display = 'none'; }); return; } diff --git a/javascript/logger.js b/javascript/logger.js index 1d89294dc..739afbb73 100644 --- a/javascript/logger.js +++ b/javascript/logger.js @@ -10,7 +10,7 @@ const log = async (...msg) => { window.logger.innerHTML += window.logPrettyPrint(...msg); scrollBottom(window.logger); } - console.log(ts, ...msg); // eslint-disable-line no-console + console.log(ts, ...msg); }; const debug = async (...msg) => { @@ -20,7 +20,7 @@ const debug = async (...msg) => { window.logger.innerHTML += window.logPrettyPrint(...msg); scrollBottom(window.logger); } - console.debug(ts, ...msg); // eslint-disable-line no-console + console.debug(ts, ...msg); }; const error = async (...msg) => { @@ -30,7 +30,7 @@ const error = async (...msg) => { window.logger.innerHTML += window.logPrettyPrint(...msg); scrollBottom(window.logger); } - console.error(ts, ...msg); // eslint-disable-line no-console + console.error(ts, ...msg); // const txt = msg.join(' '); // if (!txt.includes('asctime') && !txt.includes('xhr.')) xhrPost('/sdapi/v1/log', { error: txt }); // eslint-disable-line no-use-before-define }; diff --git a/javascript/progressBar.js b/javascript/progressBar.js index 5062517e0..c12d71346 100644 --- a/javascript/progressBar.js +++ b/javascript/progressBar.js @@ -7,7 +7,7 @@ function setRefreshInterval() { document.addEventListener('visibilitychange', () => { if (document.hidden) refreshInterval = Math.max(2500, opts.live_preview_refresh_period || 1000); else refreshInterval = opts.live_preview_refresh_period || 1000; - log('refreshInterval', document.visibilityState, refreshInterval); + // log('refreshInterval', document.visibilityState, refreshInterval); }); } diff --git a/javascript/sdnext.css b/javascript/sdnext.css index d41f22a18..fea7396ac 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -2031,42 +2031,43 @@ div:has(>#tab-gallery-folders) { .loading { color: white; - position: border-box; - top: 85%; - font-size: 1.5em; + left: 50%; + position: absolute; + top: 20%; + transform: translateX(-50%); } .loader { - width: 100px; - height: 100px; + animation: spin 4s linear infinite, hue 5s infinite alternate; border: var(--spacing-md) solid transparent; border-radius: 50%; - border-top: var(--spacing-md) solid var(--primary-600); - animation: spin 2s linear infinite, hue 5s infinite alternate; - position: border-box; + border-top: var(--spacing-md) solid var(--sd-main-accent-color); + height: 300px; + position: relative; + width: 300px; } -.loader::before, -.loader::after { - content: ""; - position: absolute; - top: 6px; - bottom: 6px; - left: 6px; - right: 6px; - border-radius: 50%; +.loader::before, .loader::after { border: var(--spacing-md) solid transparent; - animation: hue 5s infinite alternate; + border-radius: 50%; + bottom: 6px; + content: ""; + left: 6px; + position: absolute; + right: 6px; + top: 6px; } .loader::before { - border-top-color: var(--primary-900); - animation: spin 3s linear infinite; + animation: 3s spin linear infinite, hue 5s infinite alternate; + border-top-color: var(--sd-main-accent-color); + filter: brightness(50%); } .loader::after { - border-top-color: var(--primary-300); - animation: spin 1.5s linear infinite; + animation: spin 1.5s linear infinite, hue 5s infinite alternate; + border-top-color: var(--sd-main-accent-color); + filter: brightness(150%); } .docs-search textarea { @@ -2126,11 +2127,11 @@ div:has(>#tab-gallery-folders) { .docs-card-h1 { font-weight: bold; - font-size: 1.0; + font-size: 1.0em; } .docs-card-h2 { - font-size: 1.0; + font-size: 1.0em; max-height: 4em; overflow: hidden; } diff --git a/javascript/settings.js b/javascript/settings.js index 1ffe449e9..537bfd7f0 100644 --- a/javascript/settings.js +++ b/javascript/settings.js @@ -173,7 +173,7 @@ async function initModels() { `; el.innerHTML = res.length > 0 ? ready() : warn(); el.style.display = 'block'; - setTimeout(() => el.style.display = 'none', res.length === 0 ? 30000 : 1500); + setTimeout(() => { el.style.display = 'none'; }, res.length === 0 ? 30000 : 1500); if (res.length === 0) { if (en.classList.contains('hide')) gradioApp().getElementById('txt2img_extra_networks_btn').click(); const repeat = setInterval(() => { diff --git a/javascript/startup.js b/javascript/startup.js index 5a1af8f5e..8f512572b 100644 --- a/javascript/startup.js +++ b/javascript/startup.js @@ -6,7 +6,7 @@ async function waitForOpts() { // make sure all of the ui is ready and options are loaded const t0 = performance.now(); let t1 = performance.now(); - while (true) { // eslint-disable-line no-constant-condition + while (true) { if (t1 - t0 > 120000) { log('waitForOpts timeout'); break; diff --git a/javascript/trainMonitor.js b/javascript/trainMonitor.js index dbe9a172d..b0ee61c50 100644 --- a/javascript/trainMonitor.js +++ b/javascript/trainMonitor.js @@ -1,7 +1,7 @@ function startTrainMonitor() { gradioApp().querySelector('#train_error').innerHTML = ''; const id = randomId(); - const onProgress = (progress) => gradioApp().getElementById('train_progress').innerHTML = progress.textinfo; + const onProgress = (progress) => { gradioApp().getElementById('train_progress').innerHTML = progress.textinfo; }; requestProgress(id, gradioApp().getElementById('train_gallery'), null, onProgress, false); const res = Array.from(arguments); res[0] = id; diff --git a/javascript/ui.js b/javascript/ui.js index a89c3a249..a62fba54f 100644 --- a/javascript/ui.js +++ b/javascript/ui.js @@ -33,9 +33,10 @@ function clip_gallery_urls(gallery) { } function isVisible(el) { + if (!el) return false; const rect = el.getBoundingClientRect(); if (rect.width === 0 && rect.height === 0) return false; - return rect.top >= 0 && rect.left >= 0 && rect.bottom <= (window.innerHeight || document.documentElement.clientHeight) && rect.right <= (window.innerWidth || document.documentElement.clientWidth); + return (rect.top >= 0) && (rect.left >= 0) && (rect.bottom <= (window.innerHeight || document.documentElement.clientHeight)) && (rect.right <= (window.innerWidth || document.documentElement.clientWidth)); } function all_gallery_buttons() { @@ -63,18 +64,31 @@ function selected_gallery_index() { const button = selected_gallery_button(); let result = -1; buttons.forEach((v, i) => { if (v === button) { result = i; } }); + if (result === -1 && gradioApp().getElementById('tab-gallery-search')?.checkVisibility()) { + const gallerySelection = window.getGallerySelection(); + if (Number.isInteger(gallerySelection.index)) result = gallerySelection.index; + } return result; } -function selected_gallery_files() { +function selected_gallery_files(tabname) { let allImages = []; + let allThumbnails; + if (tabname && tabname !== 'gallery') allThumbnails = gradioApp().querySelectorAll('div[id$=_gallery].gradio-gallery .thumbnail-item.thumbnail-small'); + else allThumbnails = gradioApp().querySelectorAll('.gradio-gallery .thumbnails > .thumbnail-item.thumbnail-small'); try { - let allCurrentButtons = gradioApp().querySelectorAll('[style="display: block;"].tabitem div[id$=_gallery].gradio-gallery .thumbnail-item.thumbnail-small'); - if (allCurrentButtons.length === 0) allCurrentButtons = gradioApp().querySelectorAll('.gradio-gallery .thumbnails > .thumbnail-item.thumbnail-small'); - allImages = Array.from(allCurrentButtons).map((v) => v.querySelector('img')?.src); - allImages = allImages.filter((el) => isVisible(el)); - } catch { /**/ } - const selectedIndex = selected_gallery_index(); + allImages = Array.from(allThumbnails).map((v) => v.querySelector('img')); + if (tabname && tabname !== 'gallery') allImages = allImages.filter((img) => isVisible(img)); + allImages = allImages.map((img) => { + let fn = img.src; + if (fn.includes('file=')) fn = fn.split('file=')[1]; + return decodeURI(fn); + }); + } catch (err) { + error(`selected_gallery_files: ${err}`); + } + let selectedIndex = -1; + if (tabname && tabname !== 'gallery') selectedIndex = selected_gallery_index(); return [allImages, selectedIndex]; } @@ -86,6 +100,16 @@ function extract_image_from_gallery(gallery) { return [gallery[index]]; } +function send_to_kanvas(gallery) { + const [image] = extract_image_from_gallery(gallery); + log('sendToKanvas', image); + if (window.loadFromURL && image.data) window.loadFromURL(image.data); + // const inputPanelEl = gradioApp().getElementById('control-template-column-input'); + // if (inputPanelEl) inputPanelEl.classList.remove('hidden'); + const inputPanelCb = gradioApp().getElementById('control_dynamic_input'); + if (inputPanelCb && !inputPanelCb.checked) inputPanelCb.click(); +} + async function setTheme(val, old) { if (!old || val === old) return; old = old.replace('modern/', ''); @@ -99,7 +123,7 @@ async function setTheme(val, old) { const href = link.href.replace(old, val); const res = await fetch(href); if (res.ok) { - log('setTheme:', old, val); + log('setTheme', old, val); link.href = link.href.replace(old, val); } else { log('setTheme: CSS not found', val); @@ -254,7 +278,12 @@ function submit_control(...args) { const res = create_submit_args(args); res[0] = id; res[1] = window.submit_state; - res[2] = gradioApp().querySelector('#control-tabs > .tab-nav > .selected')?.innerText.toLowerCase() || ''; // selected tab name + + const tabs = Array.from(gradioApp().querySelectorAll('#control-tabs > .tab-nav > button')); + const tabIdx = tabs.findIndex((btn) => btn.classList.contains('selected')); + const tabNames = ['ControlNet', 'T2I Adapter', 'XS', 'Lite', 'Reference']; + const selectedTab = tabNames[tabIdx] || 'ControlNet'; + res[2] = selectedTab.toLowerCase(); window.submit_state = ''; return res; } diff --git a/javascript/uiConfig.js b/javascript/uiConfig.js index 79356d4ac..113e94d06 100644 --- a/javascript/uiConfig.js +++ b/javascript/uiConfig.js @@ -14,7 +14,7 @@ async function getUIDefaults() { const btn = gradioApp().getElementById('ui_defaults_view'); if (!btn) return; const intersectionObserver = new IntersectionObserver((entries) => { - if (entries[0].intersectionRatio <= 0) { } + if (entries[0].intersectionRatio <= 0) { /* Pass */ } if (entries[0].intersectionRatio > 0) btn.click(); }); intersectionObserver.observe(btn); // monitor visibility of tab diff --git a/launch.py b/launch.py index 57a1930af..e0fbc61ed 100755 --- a/launch.py +++ b/launch.py @@ -67,11 +67,8 @@ def get_custom_args(): del env['PS1'] installer.log.trace(f'Environment: {installer.print_dict(env)}') env = [f'{k}={v}' for k, v in os.environ.items() if k.startswith('SD_')] - installer.log.debug(f'Env flags: {env}') - ldpreload = os.environ.get('LD_PRELOAD', None) - ldpath = os.environ.get('LD_LIBRARY_PATH', None) - if ldpreload is not None or ldpath is not None: - installer.log.debug(f'Linker flags: preload="{ldpreload}" path="{ldpath}"') + ld = [f'{k}={v}' for k, v in os.environ.items() if k.startswith('LD_')] + installer.log.debug(f'Flags: sd={env} ld={ld}') rec('args') diff --git a/models/Reference/AIDC-AI--Ovis-Image-7B.jpg b/models/Reference/AIDC-AI--Ovis-Image-7B.jpg index 30323c196..bdc806453 100644 Binary files a/models/Reference/AIDC-AI--Ovis-Image-7B.jpg and b/models/Reference/AIDC-AI--Ovis-Image-7B.jpg differ diff --git a/models/Reference/Disty0--Chroma1-HD-SDNQ-uint4-svd-r32.jpg b/models/Reference/Disty0--Chroma1-HD-SDNQ-uint4-svd-r32.jpg index 4df709ada..bb93553f2 100644 Binary files a/models/Reference/Disty0--Chroma1-HD-SDNQ-uint4-svd-r32.jpg and b/models/Reference/Disty0--Chroma1-HD-SDNQ-uint4-svd-r32.jpg differ diff --git a/models/Reference/Disty0--FLUX.2-klein-4B-SDNQ-4bit-dynamic.jpg b/models/Reference/Disty0--FLUX.2-klein-4B-SDNQ-4bit-dynamic.jpg new file mode 100644 index 000000000..ac11cfc01 Binary files /dev/null and b/models/Reference/Disty0--FLUX.2-klein-4B-SDNQ-4bit-dynamic.jpg differ diff --git a/models/Reference/Disty0--FLUX.2-klein-9B-SDNQ-4bit-dynamic-svd-r32.jpg b/models/Reference/Disty0--FLUX.2-klein-9B-SDNQ-4bit-dynamic-svd-r32.jpg new file mode 100644 index 000000000..f24eb6ae9 Binary files /dev/null and b/models/Reference/Disty0--FLUX.2-klein-9B-SDNQ-4bit-dynamic-svd-r32.jpg differ diff --git a/models/Reference/Disty0--Qwen-Image-2512-SDNQ-uint4-svd-r32.jpg b/models/Reference/Disty0--Qwen-Image-2512-SDNQ-uint4-svd-r32.jpg new file mode 100644 index 000000000..28f0aba76 Binary files /dev/null and b/models/Reference/Disty0--Qwen-Image-2512-SDNQ-uint4-svd-r32.jpg differ diff --git a/models/Reference/Disty0--Qwen-Image-Edit-2511-SDNQ-uint4-svd-r32.jpg b/models/Reference/Disty0--Qwen-Image-Edit-2511-SDNQ-uint4-svd-r32.jpg new file mode 100644 index 000000000..8f12470fd Binary files /dev/null and b/models/Reference/Disty0--Qwen-Image-Edit-2511-SDNQ-uint4-svd-r32.jpg differ diff --git a/models/Reference/Disty0--Qwen-Image-Layered-SDNQ-uint4-svd-r32.jpg b/models/Reference/Disty0--Qwen-Image-Layered-SDNQ-uint4-svd-r32.jpg new file mode 100644 index 000000000..e951ce2d2 Binary files /dev/null and b/models/Reference/Disty0--Qwen-Image-Layered-SDNQ-uint4-svd-r32.jpg differ diff --git a/models/Reference/Qwen--Qwen-Image-2512.jpg b/models/Reference/Qwen--Qwen-Image-2512.jpg new file mode 100644 index 000000000..dc1839828 Binary files /dev/null and b/models/Reference/Qwen--Qwen-Image-2512.jpg differ diff --git a/models/Reference/Qwen--Qwen-Image-Edit-2511.jpg b/models/Reference/Qwen--Qwen-Image-Edit-2511.jpg new file mode 100644 index 000000000..8f786035d Binary files /dev/null and b/models/Reference/Qwen--Qwen-Image-Edit-2511.jpg differ diff --git a/models/Reference/Qwen--Qwen-Image-Layered.jpg b/models/Reference/Qwen--Qwen-Image-Layered.jpg new file mode 100644 index 000000000..041f4e1b6 Binary files /dev/null and b/models/Reference/Qwen--Qwen-Image-Layered.jpg differ diff --git a/models/Reference/black-forest-labs--FLUX.2-klein-4B.jpg b/models/Reference/black-forest-labs--FLUX.2-klein-4B.jpg new file mode 100644 index 000000000..a9a24ad99 Binary files /dev/null and b/models/Reference/black-forest-labs--FLUX.2-klein-4B.jpg differ diff --git a/models/Reference/black-forest-labs--FLUX.2-klein-9B.jpg b/models/Reference/black-forest-labs--FLUX.2-klein-9B.jpg new file mode 100644 index 000000000..73476b20d Binary files /dev/null and b/models/Reference/black-forest-labs--FLUX.2-klein-9B.jpg differ diff --git a/models/Reference/black-forest-labs--FLUX.2-klein-base-4B.jpg b/models/Reference/black-forest-labs--FLUX.2-klein-base-4B.jpg new file mode 100644 index 000000000..43e6c2ac7 Binary files /dev/null and b/models/Reference/black-forest-labs--FLUX.2-klein-base-4B.jpg differ diff --git a/models/Reference/black-forest-labs--FLUX.2-klein-base-9B.jpg b/models/Reference/black-forest-labs--FLUX.2-klein-base-9B.jpg new file mode 100644 index 000000000..a794b244c Binary files /dev/null and b/models/Reference/black-forest-labs--FLUX.2-klein-base-9B.jpg differ diff --git a/models/Reference/lodestones--Chroma-Base.jpg b/models/Reference/lodestones--Chroma-Base.jpg deleted file mode 100644 index 14683e907..000000000 Binary files a/models/Reference/lodestones--Chroma-Base.jpg and /dev/null differ diff --git a/models/Reference/lodestones--Chroma-HD.jpg b/models/Reference/lodestones--Chroma-HD.jpg deleted file mode 100644 index 72509efe1..000000000 Binary files a/models/Reference/lodestones--Chroma-HD.jpg and /dev/null differ diff --git a/models/Reference/lodestones--Chroma-annealed.jpg b/models/Reference/lodestones--Chroma-annealed.jpg deleted file mode 100644 index ceee2286a..000000000 Binary files a/models/Reference/lodestones--Chroma-annealed.jpg and /dev/null differ diff --git a/models/Reference/lodestones--Chroma-detail.jpg b/models/Reference/lodestones--Chroma-detail.jpg deleted file mode 100644 index 78e6e33f7..000000000 Binary files a/models/Reference/lodestones--Chroma-detail.jpg and /dev/null differ diff --git a/models/Reference/lodestones--Chroma-flash.jpg b/models/Reference/lodestones--Chroma-flash.jpg deleted file mode 100644 index 0f9cbf974..000000000 Binary files a/models/Reference/lodestones--Chroma-flash.jpg and /dev/null differ diff --git a/models/Reference/lodestones--Chroma.jpg b/models/Reference/lodestones--Chroma.jpg deleted file mode 100644 index 78e6e33f7..000000000 Binary files a/models/Reference/lodestones--Chroma.jpg and /dev/null differ diff --git a/models/Reference/lodestones--Chroma1-Base.jpg b/models/Reference/lodestones--Chroma1-Base.jpg new file mode 100644 index 000000000..1c32cfffd Binary files /dev/null and b/models/Reference/lodestones--Chroma1-Base.jpg differ diff --git a/models/Reference/lodestones--Chroma1-Flash.jpg b/models/Reference/lodestones--Chroma1-Flash.jpg new file mode 100644 index 000000000..4021fc0be Binary files /dev/null and b/models/Reference/lodestones--Chroma1-Flash.jpg differ diff --git a/models/Reference/lodestones--Chroma1-HD.jpg b/models/Reference/lodestones--Chroma1-HD.jpg new file mode 100644 index 000000000..b8c810fc1 Binary files /dev/null and b/models/Reference/lodestones--Chroma1-HD.jpg differ diff --git a/models/Reference/vladmandic--chroma-unlocked-v50-annealed.jpg b/models/Reference/vladmandic--chroma-unlocked-v50-annealed.jpg new file mode 100644 index 000000000..872fb78b9 Binary files /dev/null and b/models/Reference/vladmandic--chroma-unlocked-v50-annealed.jpg differ diff --git a/models/Reference/zai-org--GLM-Image.jpg b/models/Reference/zai-org--GLM-Image.jpg new file mode 100644 index 000000000..526e50338 Binary files /dev/null and b/models/Reference/zai-org--GLM-Image.jpg differ diff --git a/modules/apg/pipeline_stable_diffision_xl_apg.py b/modules/apg/pipeline_stable_diffision_xl_apg.py index eda0e3cf0..3371877fd 100644 --- a/modules/apg/pipeline_stable_diffision_xl_apg.py +++ b/modules/apg/pipeline_stable_diffision_xl_apg.py @@ -1081,7 +1081,7 @@ class StableDiffusionXLPipelineAPG( latents, ) - # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 6. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 7. Prepare added time ids & embeddings diff --git a/modules/apg/pipeline_stable_diffusion_apg.py b/modules/apg/pipeline_stable_diffusion_apg.py index 4b615ffb8..6eb6bae90 100644 --- a/modules/apg/pipeline_stable_diffusion_apg.py +++ b/modules/apg/pipeline_stable_diffusion_apg.py @@ -965,7 +965,7 @@ class StableDiffusionPipelineAPG( latents, ) - # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 6. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 6.1 Add image embeds for IP-Adapter diff --git a/modules/api/gallery.py b/modules/api/gallery.py index d3fdf85ef..e4dc8ba0b 100644 --- a/modules/api/gallery.py +++ b/modules/api/gallery.py @@ -10,6 +10,7 @@ from starlette.websockets import WebSocket, WebSocketState from pydantic import BaseModel, Field # pylint: disable=no-name-in-module from PIL import Image from modules import shared, images, files_cache, modelstats +from modules.paths import resolve_output_path debug = shared.log.debug if os.environ.get('SD_BROWSER_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -128,20 +129,51 @@ def register_api(app: FastAPI): # register api # @app.get('/sdapi/v1/browser/folders', response_model=List[str]) def get_folders(): + def make_folder(path, label=None): + """Create folder entry with path and display label.""" + if label is None: + label = os.path.basename(path) or path + return {"path": path, "label": label} + reference_dir = os.path.join('models', 'Reference') - folders = [shared.opts.data.get(f, '') for f in OPTS_FOLDERS] - folders += list(shared.opts.browser_folders.split(',')) - folders += [reference_dir] - folders = [f.strip() for f in folders if f != ''] - folders = list(dict.fromkeys(folders)) # filter duplicates - folders = [f for f in folders if os.path.isdir(f)] - if shared.demo is not None: - for f in folders: - if f not in shared.demo.allowed_paths: - debug(f'Browser folders allow: {f}') - shared.demo.allowed_paths.append(quote(f)) - debug(f'Browser folders: {folders}') - return JSONResponse(content=folders) + base_samples = shared.opts.outdir_samples + base_grids = shared.opts.outdir_grids + # Build list of resolved output paths with labels + folders = [] + if base_samples: + folders.append(make_folder(base_samples, os.path.basename(base_samples.rstrip('/\\')))) + if base_grids and base_grids != base_samples: + folders.append(make_folder(base_grids, os.path.basename(base_grids.rstrip('/\\')))) + # Use the specific folder setting values as labels (e.g., "outputs/text" -> "outputs/text") + folders.append(make_folder(resolve_output_path(base_samples, shared.opts.outdir_txt2img_samples), shared.opts.outdir_txt2img_samples)) + folders.append(make_folder(resolve_output_path(base_samples, shared.opts.outdir_img2img_samples), shared.opts.outdir_img2img_samples)) + folders.append(make_folder(resolve_output_path(base_samples, shared.opts.outdir_control_samples), shared.opts.outdir_control_samples)) + folders.append(make_folder(resolve_output_path(base_samples, shared.opts.outdir_extras_samples), shared.opts.outdir_extras_samples)) + folders.append(make_folder(resolve_output_path(base_samples, shared.opts.outdir_save), shared.opts.outdir_save)) + folders.append(make_folder(resolve_output_path(base_samples, shared.opts.outdir_video), shared.opts.outdir_video)) + folders.append(make_folder(resolve_output_path(base_samples, shared.opts.outdir_init_images), shared.opts.outdir_init_images)) + folders.append(make_folder(resolve_output_path(base_grids, shared.opts.outdir_txt2img_grids), shared.opts.outdir_txt2img_grids)) + folders.append(make_folder(resolve_output_path(base_grids, shared.opts.outdir_img2img_grids), shared.opts.outdir_img2img_grids)) + folders.append(make_folder(resolve_output_path(base_grids, shared.opts.outdir_control_grids), shared.opts.outdir_control_grids)) + # Custom browser folders and reference dir + for f in shared.opts.browser_folders.split(','): + f = f.strip() + if f: + folders.append(make_folder(f)) + folders.append(make_folder(reference_dir, 'Reference')) + # Filter empty and duplicates (by path) + seen_paths = set() + unique_folders = [] + for f in folders: + path = f["path"].strip() + if path and path not in seen_paths and os.path.isdir(path): + seen_paths.add(path) + unique_folders.append(f) + if shared.demo is not None and path not in shared.demo.allowed_paths: + debug(f'Browser folders allow: {path}') + shared.demo.allowed_paths.append(quote(path)) + debug(f'Browser folders: {unique_folders}') + return JSONResponse(content=unique_folders) # @app.get("/sdapi/v1/browser/thumb", response_model=dict) async def get_thumb(file: str): diff --git a/modules/api/generate.py b/modules/api/generate.py index f03a3a360..102b15f2c 100644 --- a/modules/api/generate.py +++ b/modules/api/generate.py @@ -3,6 +3,7 @@ from fastapi.responses import JSONResponse from modules import errors, shared, scripts_manager, ui from modules.api import models, script, helpers from modules.processing import StableDiffusionProcessingTxt2Img, StableDiffusionProcessingImg2Img, process_images +from modules.paths import resolve_output_path errors.install() @@ -105,8 +106,8 @@ class APIGenerate(): p = StableDiffusionProcessingTxt2Img(sd_model=shared.sd_model, **args) self.prepare_ip_adapter(txt2imgreq, p) p.scripts = script_runner - p.outpath_grids = shared.opts.outdir_grids or shared.opts.outdir_txt2img_grids - p.outpath_samples = shared.opts.outdir_samples or shared.opts.outdir_txt2img_samples + p.outpath_grids = resolve_output_path(shared.opts.outdir_grids, shared.opts.outdir_txt2img_grids) + p.outpath_samples = resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_txt2img_samples) for key, value in getattr(txt2imgreq, "extra", {}).items(): setattr(p, key, value) jobid = shared.state.begin('API-TXT', api=True) @@ -157,8 +158,8 @@ class APIGenerate(): self.prepare_ip_adapter(img2imgreq, p) p.init_images = [helpers.decode_base64_to_image(x) for x in init_images] p.scripts = script_runner - p.outpath_grids = shared.opts.outdir_img2img_grids - p.outpath_samples = shared.opts.outdir_img2img_samples + p.outpath_grids = resolve_output_path(shared.opts.outdir_grids, shared.opts.outdir_img2img_grids) + p.outpath_samples = resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_img2img_samples) for key, value in getattr(img2imgreq, "extra", {}).items(): setattr(p, key, value) jobid = shared.state.begin('API-IMG', api=True) diff --git a/modules/api/server.py b/modules/api/server.py index a01fb744d..69de6d18d 100644 --- a/modules/api/server.py +++ b/modules/api/server.py @@ -1,6 +1,6 @@ import os import time -from typing import Any, Dict +from typing import Any from fastapi import Request, Depends from fastapi.exceptions import HTTPException from fastapi.responses import FileResponse @@ -95,11 +95,11 @@ def get_config(): del options['sd_lora'] return options -def set_config(req: Dict[str, Any]): +def set_config(req: dict[str, Any]): updated = [] for k, v in req.items(): updated.append({ k: shared.opts.set(k, v) }) - shared.opts.save(shared.config_filename) + shared.opts.save() return { "updated": updated } def get_cmd_flags(): diff --git a/modules/attention.py b/modules/attention.py index e92eb95dc..a0a29bfb1 100644 --- a/modules/attention.py +++ b/modules/attention.py @@ -89,7 +89,7 @@ def set_ck_flash_attention(backend: str, device: torch.device): if backend == "rocm": if not installed('flash-attn'): log.info('Torch attention: type="Flash attention" building...') - agent = rocm.Agent(getattr(torch.cuda.get_device_properties(device), "gcnArchName", "gfx0000")) + agent = rocm.Agent(device) install(rocm.get_flash_attention_command(agent), reinstall=True) else: install('flash-attn') diff --git a/modules/call_queue.py b/modules/call_queue.py index 38b372ca2..7ed03e5b9 100644 --- a/modules/call_queue.py +++ b/modules/call_queue.py @@ -1,15 +1,26 @@ +import os +import sys import html import threading import time import cProfile from modules import shared, progress, errors, timer + queue_lock = threading.Lock() +debug = os.environ.get('SD_QUEUE_DEBUG', None) is not None + + +def get_lock(): + if debug: + fn = f'{sys._getframe(3).f_code.co_name}:{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + errors.log.debug(f'Queue: fn={fn} lock={queue_lock.locked()}') + return queue_lock def wrap_queued_call(func): def f(*args, **kwargs): - with queue_lock: + with get_lock(): res = func(*args, **kwargs) return res return f @@ -24,7 +35,7 @@ def wrap_gradio_gpu_call(func, extra_outputs=None, name=None): progress.add_task_to_queue(id_task) else: id_task = None - with queue_lock: + with get_lock(): progress.start_task(id_task) res = [None, '', '', ''] try: diff --git a/modules/civitai/search_civitai.py b/modules/civitai/search_civitai.py index 488a05cdd..3b534662d 100644 --- a/modules/civitai/search_civitai.py +++ b/modules/civitai/search_civitai.py @@ -7,8 +7,7 @@ from installer import install, log full_dct = False full_html = False -base_models = ['', 'ODOR', 'SD 1.4', 'SD 1.5', 'SD 1.5 LCM', 'SD 1.5 Hyper', 'SD 2.0', 'SD 2.0 768', 'SD 2.1', 'SD 2.1 768', 'SD 2.1 Unclip', 'SDXL 0.9', 'SDXL 1.0', 'SD 3', 'SD 3.5', 'SD 3.5 Medium', 'SD 3.5 Large', 'SD 3.5 Large Turbo', 'Pony', 'Flux.1 S', 'Flux.1 D', 'Flux.1 Kontext', 'AuraFlow', 'SDXL 1.0 LCM', 'SDXL Distilled', 'SDXL Turbo', 'SDXL Lightning', 'SDXL Hyper', 'Stable Cascade', 'SVD', 'SVD XT', 'Playground v2', 'PixArt a', 'PixArt E', 'Hunyuan 1', 'Hunyuan Video', 'Lumina', 'Kolors', 'Illustrious', 'Mochi', 'LTXV', 'CogVideoX', 'NoobAI', 'Wan Video', 'Wan Video 1.3B t2v', 'Wan Video 14B t2v', 'Wan Video 14B i2v 480p', 'Wan Video 14B i2v 720p', 'HiDream', 'OpenAI', 'Imagen4', 'Other'] - +base_models = ['', 'AuraFlow', 'Chroma', 'CogVideoX', 'Flux.1 S', 'Flux.1 D', 'Flux.1 Krea', 'Flux.1 Kontext', 'Flux.2 D', 'HiDream', 'Hunyuan 1', 'Hunyuan Video', 'Illustrious', 'Kolors', 'LTXV', 'Lumina', 'Mochi', 'NoobAI', 'PixArt a', 'PixArt E', 'Pony', 'Pony V7', 'Qwen', 'SD 1.4', 'SD 1.5', 'SD 1.5 LCM', 'SD 1.5 Hyper', 'SD 2.0', 'SD 2.1', 'SDXL 1.0', 'SDXL Lightning', 'SDXL Hyper', 'Wan Video 1.3B t2v', 'Wan Video 14B t2v', 'Wan Video 14B i2v 480p', 'Wan Video 14B i2v 720p', 'Wan Video 2.2 TI2V-5B', 'Wan Video 2.2 I2V-A14B', 'Wan Video 2.2 T2V-A14B', 'Wan Video 2.5 T2V', 'Wan Video 2.5 I2V', 'ZImageTurbo', 'Other'] @dataclass class ModelImage(): diff --git a/modules/control/proc/dwpose/config/yolox_l_8xb8-300e_coco.py b/modules/control/proc/dwpose/config/yolox_l_8xb8-300e_coco.py index 7b4cb5a4b..090394015 100644 --- a/modules/control/proc/dwpose/config/yolox_l_8xb8-300e_coco.py +++ b/modules/control/proc/dwpose/config/yolox_l_8xb8-300e_coco.py @@ -194,7 +194,6 @@ param_scheduler = [ dict( # use quadratic formula to warm up 5 epochs # and lr is updated by iteration - # TODO: fix default scope in get function type='mmdet.QuadraticWarmupLR', by_epoch=True, begin=0, diff --git a/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/activations/activations.py b/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/activations/activations.py index eca58933b..207db52c5 100644 --- a/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/activations/activations.py +++ b/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/activations/activations.py @@ -12,8 +12,6 @@ from torch.nn import functional as F def swish(x, inplace: bool = False): """Swish - Described originally as SiLU (https://arxiv.org/abs/1702.03118v3) and also as Swish (https://arxiv.org/abs/1710.05941). - - TODO Rename to SiLU with addition to PyTorch """ return x.mul_(x.sigmoid()) if inplace else x.mul(x.sigmoid()) diff --git a/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/activations/activations_jit.py b/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/activations/activations_jit.py index 7176b05e7..8ea154e47 100644 --- a/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/activations/activations_jit.py +++ b/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/activations/activations_jit.py @@ -22,8 +22,6 @@ __all__ = ['swish_jit', 'SwishJit', 'mish_jit', 'MishJit', def swish_jit(x, inplace: bool = False): """Swish - Described originally as SiLU (https://arxiv.org/abs/1702.03118v3) and also as Swish (https://arxiv.org/abs/1710.05941). - - TODO Rename to SiLU with addition to PyTorch """ return x.mul(x.sigmoid()) diff --git a/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/activations/activations_me.py b/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/activations/activations_me.py index e91df5a50..e16cc6258 100644 --- a/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/activations/activations_me.py +++ b/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/activations/activations_me.py @@ -36,8 +36,6 @@ class SwishJitAutoFn(torch.autograd.Function): Swish - Described originally as SiLU (https://arxiv.org/abs/1702.03118v3) and also as Swish (https://arxiv.org/abs/1710.05941). - - TODO Rename to SiLU with addition to PyTorch """ @staticmethod diff --git a/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/efficientnet_builder.py b/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/efficientnet_builder.py index 0343e3f44..637090945 100644 --- a/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/efficientnet_builder.py +++ b/modules/control/proc/normalbae/nets/submodules/efficientnet_repo/geffnet/efficientnet_builder.py @@ -483,7 +483,7 @@ def _decode_block_str(block_str): Returns: A list of block args (dicts) Raises: - ValueError: if the string def not properly specified (TODO) + ValueError: if the string def not properly specified """ assert isinstance(block_str, str) ops = block_str.split('_') diff --git a/modules/control/proc/zoe/zoedepth/models/layers/localbins_layers.py b/modules/control/proc/zoe/zoedepth/models/layers/localbins_layers.py index 91d08de0f..b70ae562e 100644 --- a/modules/control/proc/zoe/zoedepth/models/layers/localbins_layers.py +++ b/modules/control/proc/zoe/zoedepth/models/layers/localbins_layers.py @@ -155,7 +155,7 @@ class LinearSplitter(nn.Module): b_prev = b_prev / b_prev.sum(dim=1, keepdim=True) # renormalize for gurantees # print(b_prev.shape, S_normed.shape) - # if is_for_query:(1).expand(-1, b_prev.size(0)//n, -1, -1, -1, -1).flatten(0,1) # TODO ? can replace all this with a single torch.repeat? + # if is_for_query:(1).expand(-1, b_prev.size(0)//n, -1, -1, -1, -1).flatten(0,1) b = b_prev.unsqueeze(2) * S_normed b = b.flatten(1,2) # .shape n, prev_nbins * split_factor, h, w diff --git a/modules/control/proc/zoe/zoedepth/utils/config.py b/modules/control/proc/zoe/zoedepth/utils/config.py index 24525d947..dde747eff 100644 --- a/modules/control/proc/zoe/zoedepth/utils/config.py +++ b/modules/control/proc/zoe/zoedepth/utils/config.py @@ -395,7 +395,7 @@ def get_config(model_name, mode='train', dataset=None, **overwrite_kwargs): overwrite_kwargs = split_combined_args(overwrite_kwargs) config = {**config, **overwrite_kwargs} - # Casting to bool # TODO: Not necessary. Remove and test + # Casting to bool for key in KEYS_TYPE_BOOL: if key in config: config[key] = bool(config[key]) diff --git a/modules/control/run.py b/modules/control/run.py index 726d3aa4f..243840ff1 100644 --- a/modules/control/run.py +++ b/modules/control/run.py @@ -18,6 +18,7 @@ from modules.processing_class import StableDiffusionProcessingControl from modules.ui_common import infotext_to_html from modules.api import script from modules.generation_parameters_copypaste import create_override_settings_dict +from modules.paths import resolve_output_path debug = os.environ.get('SD_CONTROL_DEBUG', None) is not None @@ -402,8 +403,8 @@ def control_run(state: str = '', # pylint: disable=keyword-arg-before-vararg hdr_mode=hdr_mode, hdr_brightness=hdr_brightness, hdr_color=hdr_color, hdr_sharpen=hdr_sharpen, hdr_clamp=hdr_clamp, hdr_boundary=hdr_boundary, hdr_threshold=hdr_threshold, hdr_maximize=hdr_maximize, hdr_max_center=hdr_max_center, hdr_max_boundary=hdr_max_boundary, hdr_color_picker=hdr_color_picker, hdr_tint_ratio=hdr_tint_ratio, # path - outpath_samples=shared.opts.outdir_samples or shared.opts.outdir_control_samples, - outpath_grids=shared.opts.outdir_grids or shared.opts.outdir_control_grids, + outpath_samples=resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_control_samples), + outpath_grids=resolve_output_path(shared.opts.outdir_grids, shared.opts.outdir_control_grids), # overrides override_settings=extra ) @@ -572,13 +573,13 @@ def control_run(state: str = '', # pylint: disable=keyword-arg-before-vararg # what are we doing? if 'control' in p.ops: - p.outpath_samples = shared.opts.outdir_samples or shared.opts.outdir_control_samples + p.outpath_samples = resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_control_samples) elif 'img2img' in p.ops: - p.outpath_samples = shared.opts.outdir_samples or shared.opts.outdir_img2img_samples + p.outpath_samples = resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_img2img_samples) elif 'txt2img' in p.ops: - p.outpath_samples = shared.opts.outdir_samples or shared.opts.outdir_txt2img_samples + p.outpath_samples = resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_txt2img_samples) else: # fallback to txt2img - p.outpath_samples = shared.opts.outdir_samples or shared.opts.outdir_txt2img_samples + p.outpath_samples = resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_txt2img_samples) # pipeline output = None diff --git a/modules/control/units/controlnet.py b/modules/control/units/controlnet.py index 9e026923f..fcd610396 100644 --- a/modules/control/units/controlnet.py +++ b/modules/control/units/controlnet.py @@ -181,7 +181,7 @@ def api_list_models(model_type: str = None): model_list += list(predefined_qwen) if model_type == 'hunyuandit' or model_type == 'all': model_list += list(predefined_hunyuandit) - if model_type == 'z_image': + if model_type == 'zimage': model_list += list(predefined_zimage) model_list += sorted(find_models()) return model_list @@ -207,7 +207,7 @@ def list_models(refresh=False): models = ['None'] + list(predefined_qwen) + sorted(find_models()) elif modules.shared.sd_model_type == 'hunyuandit': models = ['None'] + list(predefined_hunyuandit) + sorted(find_models()) - elif modules.shared.sd_model_type == 'z_image': + elif modules.shared.sd_model_type == 'zimage': models = ['None'] + list(predefined_zimage) + sorted(find_models()) else: log.warning(f'Control {what} model list failed: unknown model type') @@ -273,7 +273,7 @@ class ControlNet(): elif shared.sd_model_type == 'hunyuandit': from diffusers import HunyuanDiT2DControlNetModel as cls config = 'Tencent-Hunyuan/HunyuanDiT-v1.2-ControlNet-Diffusers-Canny' - elif shared.sd_model_type == 'z_image': + elif shared.sd_model_type == 'zimage': from diffusers import ZImageControlNetModel as cls if '2.0' in model_id: config = 'hlky/Z-Image-Turbo-Fun-Controlnet-Union-2.0' diff --git a/modules/devices.py b/modules/devices.py index 15b65368c..9fa446ebf 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -67,6 +67,10 @@ def has_triton(early:bool=False) -> bool: return test_triton(early=early) +def get_hip_agent() -> rocm.Agent: + return rocm.Agent(device) + + def get_backend(shared_cmd_opts): global args # pylint: disable=global-statement args = shared_cmd_opts @@ -91,7 +95,6 @@ def get_backend(shared_cmd_opts): def get_gpu_info(): def get_driver(): - import subprocess if torch.xpu.is_available(): try: return torch.xpu.get_device_properties(torch.xpu.current_device()).driver_version @@ -99,6 +102,7 @@ def get_gpu_info(): return '' elif torch.cuda.is_available() and torch.version.cuda: try: + import subprocess result = subprocess.run('nvidia-smi --query-gpu=driver_version --format=csv,noheader', shell=True, check=False, env=os.environ, stdout=subprocess.PIPE, stderr=subprocess.PIPE) version = result.stdout.decode(encoding="utf8", errors="ignore").strip() return version @@ -110,7 +114,7 @@ def get_gpu_info(): def get_package_version(pkg: str): import pkg_resources spec = pkg_resources.working_set.by_key.get(pkg, None) # more reliable than importlib - version = pkg_resources.get_distribution(pkg).version if spec is not None else '' + version = pkg_resources.get_distribution(pkg).version if spec is not None else None return version if not torch.cuda.is_available(): @@ -324,9 +328,9 @@ def test_fp16(): elif backend == 'rocm': # gfx1102 (RX 7600, 7500, 7650 and 7700S) causes segfaults with fp16 # agent can be overriden to gfx1100 to get gfx1102 working with ROCm so check the gpu name as well - agent = getattr(torch.cuda.get_device_properties(device), "gcnArchName", "gfx0000") + agent = get_hip_agent() agent_name = getattr(torch.cuda.get_device_properties(device), "name", "AMD Radeon RX 0000") - if agent == "gfx1102" or (agent == "gfx1100" and any(i in agent_name for i in ("7600", "7500", "7650", "7700S"))): + if agent.gfx_version == 0x1102 or (agent.gfx_version == 0x1100 and any(i in agent_name for i in ("7600", "7500", "7650", "7700S"))): fp16_ok = False return fp16_ok try: @@ -355,7 +359,7 @@ def test_bf16(): elif backend == 'rocm' or backend == 'zluda': agent = None if backend == 'rocm': - agent = rocm.Agent(getattr(torch.cuda.get_device_properties(device), "gcnArchName", "gfx0000")) + agent = get_hip_agent() else: from modules.zluda_installer import default_agent agent = default_agent @@ -396,7 +400,8 @@ def test_triton(early: bool = False): triton_ok = False except Exception as e: triton_ok = False - log.warning(f"Triton test fail: {e}") + line = str(e).splitlines()[0] + log.warning(f"Triton test fail: {line}") if debug: from modules import errors errors.display(e, 'Triton') diff --git a/modules/extensions.py b/modules/extensions.py index 1063e5eb7..b29a6cbb0 100644 --- a/modules/extensions.py +++ b/modules/extensions.py @@ -1,6 +1,6 @@ from __future__ import annotations import os -from datetime import datetime +from datetime import datetime, timezone import git from modules import shared, errors from modules.paths import extensions_dir, extensions_builtin_dir @@ -11,6 +11,34 @@ if not os.path.exists(extensions_dir): os.makedirs(extensions_dir) +def parse_isotime(time_string: str) -> datetime: + # If Python minimum version is 3.11+, this function can be replaced with datetime.fromisoformat() + trimmed = time_string.rstrip("Z") + if "." in trimmed: + trimmed = trimmed.split(".")[0] + match len(trimmed): + case 16: + return datetime.strptime(trimmed, "%Y-%m-%dT%H:%M").replace(tzinfo=timezone.utc) + case 19: + return datetime.strptime(trimmed, "%Y-%m-%dT%H:%M:%S").replace(tzinfo=timezone.utc) + case _: + raise ValueError(f"Unexpected time string format: '{time_string}'") + + +def format_dt(d: datetime, seconds = False) -> str: + if d.tzinfo is None: + return d.strftime('%Y-%m-%d %H:%M') + if seconds: + return d.astimezone(timezone.utc).strftime('%Y-%m-%d %H:%M:%S') + return d.astimezone(timezone.utc).strftime('%Y-%m-%d %H:%M') + + +def ts2utc(timestamp: int) -> datetime: + try: + return datetime.fromtimestamp(timestamp, timezone.utc) + except Exception: + return "unknown" + def active(): if shared.opts.disable_all_extensions == "all": return [] @@ -106,8 +134,8 @@ class Extension: self.branch = None self.remote = None self.have_info_from_repo = False - self.mtime = 0 - self.ctime = 0 + self.mtime = "2000-01-01T00:00Z" + self.ctime = "2000-01-01T00:00Z" def read_info(self, force=False): if self.have_info_from_repo and not force: @@ -142,7 +170,7 @@ class Extension: except Exception: self.branch = 'unknown' self.commit_hash = head.hexsha - self.version = f"

{self.commit_hash[:8]}

{datetime.fromtimestamp(self.commit_date).strftime('%a %b%d %Y %H:%M')}

" + self.version = f"

{self.commit_hash[:8]}

{format_dt(ts2utc(self.commit_date))}

" except Exception as ex: shared.log.error(f"Extension: failed reading data from git repo={self.name}: {ex}") self.remote = None diff --git a/modules/extra_networks.py b/modules/extra_networks.py index 8e4019ac4..054bc5c2b 100644 --- a/modules/extra_networks.py +++ b/modules/extra_networks.py @@ -108,7 +108,6 @@ def activate(p, extra_network_data=None, step=0, include=[], exclude=[]): if args is not None: continue try: - # extra_network.activate(p, []) signature = list(inspect.signature(extra_network.activate).parameters) if 'include' in signature and 'exclude' in signature: extra_network.activate(p, [], include=include, exclude=exclude) diff --git a/modules/face/instantid_model.py b/modules/face/instantid_model.py index 8df9075d9..8af9a2907 100644 --- a/modules/face/instantid_model.py +++ b/modules/face/instantid_model.py @@ -882,7 +882,7 @@ class StableDiffusionXLInstantIDPipeline(StableDiffusionXLControlNetPipeline): guidance_scale_tensor, embedding_dim=self.unet.config.time_cond_proj_dim ).to(device=device, dtype=latents.dtype) - # 7. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 7. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 7.1 Create tensor stating which controlnets to keep diff --git a/modules/face/photomaker_pipeline.py b/modules/face/photomaker_pipeline.py index 4b4e36214..45006a7e1 100644 --- a/modules/face/photomaker_pipeline.py +++ b/modules/face/photomaker_pipeline.py @@ -679,7 +679,7 @@ class PhotoMakerStableDiffusionXLPipeline(StableDiffusionXLPipeline): latents, ) - # 9. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 9. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 10. Prepare added time ids & embeddings diff --git a/modules/facelib/detection/retinaface/retinaface_utils.py b/modules/facelib/detection/retinaface/retinaface_utils.py index 8c3577577..da38d6d13 100644 --- a/modules/facelib/detection/retinaface/retinaface_utils.py +++ b/modules/facelib/detection/retinaface/retinaface_utils.py @@ -181,7 +181,6 @@ def match(threshold, truths, priors, variances, labels, landms, loc_t, conf_t, l best_prior_idx_filter.squeeze_(1) best_prior_overlap.squeeze_(1) best_truth_overlap.index_fill_(0, best_prior_idx_filter, 2) # ensure best prior - # TODO refactor: index best_prior_idx with long tensor # ensure every gt matches with its prior of max overlap for j in range(best_prior_idx.size(0)): # 判别此anchor是预测哪一个boxes best_truth_idx[best_prior_idx[j]] = j diff --git a/modules/facelib/utils/face_utils.py b/modules/facelib/utils/face_utils.py index 25ff853e3..3470a4c7e 100644 --- a/modules/facelib/utils/face_utils.py +++ b/modules/facelib/utils/face_utils.py @@ -99,7 +99,7 @@ def align_crop_face_landmarks(img, # - np.flipud(eye_to_mouth) * [-1, 1]: rotate 90 clockwise # norm with the hypotenuse: get the direction x /= np.hypot(*x) # get the hypotenuse of a right triangle - rect_scale = 1 # TODO: you can edit it to get larger rect + rect_scale = 1 x *= max(np.hypot(*eye_to_eye) * 2.0 * rect_scale, np.hypot(*eye_to_mouth) * 1.8 * rect_scale) # y: half height of the oriented crop rectangle y = np.flipud(x) * [-1, 1] @@ -116,7 +116,6 @@ def align_crop_face_landmarks(img, quad_ori = np.copy(quad) # Shrink, for large face - # TODO: do we really need shrink shrink = int(np.floor(qsize / output_size * 0.5)) if shrink > 1: h, w = img.shape[0:2] diff --git a/modules/flash_attn_triton_amd/fwd_prefill.py b/modules/flash_attn_triton_amd/fwd_prefill.py index 38589f016..074df4132 100644 --- a/modules/flash_attn_triton_amd/fwd_prefill.py +++ b/modules/flash_attn_triton_amd/fwd_prefill.py @@ -52,7 +52,6 @@ def _attn_fwd_inner(acc, l_i, m_i, q, k_ptrs, v_ptrs, bias_ptrs, stride_kn, stri qk = tl.zeros([BLOCK_M, BLOCK_N], dtype=ACCUMULATOR_TYPE) # We start from end of seqlen_k so only the first iteration would need # to be checked for padding if it is not a multiple of block_n - # TODO: This can be optimized to only be true for the padded block. if MASK_STEPS: # If this is the last block / iteration, we want to # mask if the sequence length is not a multiple of block size @@ -105,7 +104,7 @@ def _attn_fwd_inner(acc, l_i, m_i, q, k_ptrs, v_ptrs, bias_ptrs, stride_kn, stri # CAVEAT: Must update l_ij before applying dropout l_ij = tl.sum(p, 1) if ENABLE_DROPOUT: - rng_output = tl.rand(philox_seed, philox_ptrs) # TODO: use tl.randint for better performance + rng_output = tl.rand(philox_seed, philox_ptrs) dropout_mask = rng_output > dropout_p # return scores with negative values for dropped vals @@ -304,7 +303,6 @@ def attn_fwd(Q, K, V, bias, Cache_seqlens, Cache_batch_idx, # pylint: disable=un # softmax_lse = tl.where(lse_mask, 0.0, softmax_lse) l_ptrs_mask = offs_m < MAX_SEQLENS_Q tl.store(l_ptrs, l, mask=l_ptrs_mask) - # TODO: Should dropout and return encoded softmax be handled here too? return # If MQA / GQA, set the K and V head offsets appropriately. diff --git a/modules/flash_attn_triton_amd/utils.py b/modules/flash_attn_triton_amd/utils.py index 09f6a6d78..8b1640ff2 100644 --- a/modules/flash_attn_triton_amd/utils.py +++ b/modules/flash_attn_triton_amd/utils.py @@ -116,7 +116,6 @@ class MetaData(): assert self.cu_seqlens_q is not None assert self.cu_seqlens_k is not None assert len(self.cu_seqlens_q) == len(self.cu_seqlens_k) - # TODO: Remove once bias is supported with varlen assert self.bias is None # assert not self.return_scores else: @@ -125,7 +124,6 @@ class MetaData(): assert self.cu_seqlens_q is None and self.cu_seqlens_k is None assert k.shape == v.shape assert q.shape[-1] == k.shape[-1] and q.shape[-1] == v.shape[-1] - # TODO: Change assert if we support qkl f8 and v f16 assert q.dtype == k.dtype and q.dtype == v.dtype assert o.shape == q.shape assert (nheads_q % nheads_k) == 0 @@ -243,7 +241,6 @@ def input_helper( equal_seqlens=False # gen tensors - # TODO: the gen functions should maybe have different gen modes like random, ones, increasing seqlen q, cu_seqlens_q, max_seqlen_q = generate_varlen_tensor(TOTAL_SEQLENS_Q, HQ, D_HEAD, batch_size=BATCH, dtype=dtype, device=device, equal_seqlens=equal_seqlens, DEBUG_INPUT=DEBUG_INPUT) k, cu_seqlens_k, max_seqlen_k = generate_varlen_tensor(TOTAL_SEQLENS_K, HK, D_HEAD, batch_size=BATCH, dtype=dtype, device=device, equal_seqlens=equal_seqlens, DEBUG_INPUT=DEBUG_INPUT) v, _, _ = generate_varlen_tensor(TOTAL_SEQLENS_K, HK, D_HEAD, batch_size=BATCH, dtype=dtype, device=device, equal_seqlens=equal_seqlens, DEBUG_INPUT=DEBUG_INPUT) diff --git a/modules/framepack/framepack_wrappers.py b/modules/framepack/framepack_wrappers.py index ad4745846..7259db5ea 100644 --- a/modules/framepack/framepack_wrappers.py +++ b/modules/framepack/framepack_wrappers.py @@ -1,12 +1,11 @@ import os import re import random -import threading import numpy as np import torch import gradio as gr from PIL import Image -from modules import shared, processing, timer, paths, extra_networks, progress, ui_video_vlm +from modules import shared, processing, timer, paths, extra_networks, progress, ui_video_vlm, call_queue from modules.video_models.video_utils import check_av from modules.framepack import framepack_install # pylint: disable=wrong-import-order from modules.framepack import framepack_load # pylint: disable=wrong-import-order @@ -18,7 +17,6 @@ tmp_dir = os.path.join(paths.data_path, 'tmp', 'framepack') git_dir = os.path.join(os.path.dirname(__file__), 'framepack') git_repo = 'https://github.com/lllyasviel/framepack' git_commit = 'c5d375661a2557383f0b8da9d11d14c23b0c4eaf' -queue_lock = threading.Lock() loaded_variant = None @@ -131,7 +129,7 @@ def run_framepack(task_id, _ui_state, init_image, end_image, start_weight, end_w return progress.add_task_to_queue(task_id) - with queue_lock: + with call_queue.get_lock(): progress.start_task(task_id) yield from load_model(variant, attention) diff --git a/modules/generation_parameters_copypaste.py b/modules/generation_parameters_copypaste.py index 9ccf44ee7..9662024a3 100644 --- a/modules/generation_parameters_copypaste.py +++ b/modules/generation_parameters_copypaste.py @@ -1,3 +1,4 @@ +from __future__ import annotations import base64 import io import os @@ -8,9 +9,9 @@ from modules.infotext import parse, mapping, quote, unquote # pylint: disable=un type_of_gr_update = type(gr.update()) -paste_fields = {} +paste_fields: dict[str, dict] = {} field_names = {} -registered_param_bindings = [] +registered_param_bindings: list[ParamBinding] = [] debug = shared.log.trace if os.environ.get('SD_PASTE_DEBUG', None) is not None else lambda *args, **kwargs: None debug('Trace: PASTE') parse_generation_parameters = parse # compatibility @@ -18,7 +19,7 @@ infotext_to_setting_name_mapping = mapping # compatibility class ParamBinding: - def __init__(self, paste_button, tabname, source_text_component=None, source_image_component=None, source_tabname=None, override_settings_component=None, paste_field_names=None): + def __init__(self, paste_button, tabname: str, source_text_component=None, source_image_component=None, source_tabname=None, override_settings_component=None, paste_field_names=None): self.paste_button = paste_button self.tabname = tabname self.source_text_component = source_text_component @@ -60,7 +61,7 @@ def image_from_url_text(filedata): if len(filedata) == 0: return None filedata = filedata[0] - if type(filedata) == dict: + if not isinstance(filedata, str): shared.log.warning('Incorrect filedata received') return None if filedata.startswith("data:image/png;base64,"): @@ -71,13 +72,13 @@ def image_from_url_text(filedata): filedata = filedata[len("data:image/jpeg;base64,"):] if filedata.startswith("data:image/jxl;base64,"): filedata = filedata[len("data:image/jxl;base64,"):] - filedata = base64.decodebytes(filedata.encode('utf-8')) - image = Image.open(io.BytesIO(filedata)) + filebytes = base64.decodebytes(filedata.encode('utf-8')) + image = Image.open(io.BytesIO(filebytes)) images.read_info_from_image(image) return image -def add_paste_fields(tabname, init_img, fields, override_settings_component=None): +def add_paste_fields(tabname: str, init_img: gr.Image | gr.HTML | None, fields: list[tuple[gr.components.Component, str]] | None, override_settings_component=None): paste_fields[tabname] = {"init_img": init_img, "fields": fields, "override_settings_component": override_settings_component} try: field_names[tabname] = [f[1] for f in fields if f[1] is not None and not callable(f[1])] if fields is not None else [] # tuple (component, label) @@ -108,7 +109,7 @@ def get_all_fields(): return all_fields -def create_buttons(tabs_list): +def create_buttons(tabs_list: list[str]) -> dict[str, gr.Button]: buttons = {} for tab in tabs_list: name = tab @@ -128,7 +129,7 @@ def create_buttons(tabs_list): return buttons -def should_skip(param): +def should_skip(param: str): skip_params = [p.strip().lower() for p in shared.opts.disable_apply_params.split(",")] if not shared.opts.clip_skip_enabled: skip_params += ['clip skip'] @@ -149,18 +150,26 @@ def connect_paste_params_buttons(): if binding.tabname not in paste_fields: debug(f"Not not registered: tab={binding.tabname}") continue - fields = paste_fields[binding.tabname]["fields"] + fields: list[tuple[gr.components.Component, str]] = paste_fields[binding.tabname]["fields"] destination_image_component = paste_fields[binding.tabname]["init_img"] - if binding.source_image_component and destination_image_component: - binding.paste_button.click( - _js="extract_image_from_gallery" if isinstance(binding.source_image_component, gr.Gallery) else None, - fn=send_image, - inputs=[binding.source_image_component], - outputs=[destination_image_component], - show_progress='hidden', - ) - + if binding.source_image_component: + if isinstance(destination_image_component, gr.Image): + binding.paste_button.click( + _js="extract_image_from_gallery" if isinstance(binding.source_image_component, gr.Gallery) else None, + fn=send_image, + inputs=[binding.source_image_component], + outputs=[destination_image_component], + show_progress='hidden', + ) + elif isinstance(destination_image_component, gr.HTML): # kanvas + binding.paste_button.click( + _js="send_to_kanvas", + fn=None, + inputs=[binding.source_image_component], + outputs=[], + show_progress='hidden', + ) override_settings_component = binding.override_settings_component or paste_fields[binding.tabname]["override_settings_component"] if binding.source_text_component is not None and fields is not None: connect_paste(binding.paste_button, fields, binding.source_text_component, override_settings_component, binding.tabname) diff --git a/modules/gr_tempdir.py b/modules/gr_tempdir.py index f19fd53a3..eacf782f5 100644 --- a/modules/gr_tempdir.py +++ b/modules/gr_tempdir.py @@ -19,19 +19,32 @@ def check_tmp_file(gradio, filename): ok = False if hasattr(gradio, 'temp_file_sets'): ok = ok or any(filename in fileset for fileset in gradio.temp_file_sets) - if shared.opts.outdir_samples != '': - ok = ok or Path(shared.opts.outdir_samples).resolve() in Path(filename).resolve().parents - else: - ok = ok or Path(shared.opts.outdir_txt2img_samples).resolve() in Path(filename).resolve().parents - ok = ok or Path(shared.opts.outdir_img2img_samples).resolve() in Path(filename).resolve().parents - ok = ok or Path(shared.opts.outdir_extras_samples).resolve() in Path(filename).resolve().parents - if shared.opts.outdir_grids != '': - ok = ok or Path(shared.opts.outdir_grids).resolve() in Path(filename).resolve().parents - else: - ok = ok or Path(shared.opts.outdir_txt2img_grids).resolve() in Path(filename).resolve().parents - ok = ok or Path(shared.opts.outdir_img2img_grids).resolve() in Path(filename).resolve().parents - ok = ok or Path(shared.opts.outdir_save).resolve() in Path(filename).resolve().parents - ok = ok or Path(shared.opts.outdir_init_images).resolve() in Path(filename).resolve().parents + # Check resolved output paths (base + specific) + base_samples = shared.opts.outdir_samples + base_grids = shared.opts.outdir_grids + resolved_paths = [ + paths.resolve_output_path(base_samples, shared.opts.outdir_txt2img_samples), + paths.resolve_output_path(base_samples, shared.opts.outdir_img2img_samples), + paths.resolve_output_path(base_samples, shared.opts.outdir_extras_samples), + paths.resolve_output_path(base_samples, shared.opts.outdir_control_samples), + paths.resolve_output_path(base_samples, shared.opts.outdir_save), + paths.resolve_output_path(base_samples, shared.opts.outdir_video), + paths.resolve_output_path(base_samples, shared.opts.outdir_init_images), + paths.resolve_output_path(base_grids, shared.opts.outdir_txt2img_grids), + paths.resolve_output_path(base_grids, shared.opts.outdir_img2img_grids), + paths.resolve_output_path(base_grids, shared.opts.outdir_control_grids), + ] + # Also check base folders directly if set + if base_samples: + resolved_paths.append(base_samples) + if base_grids: + resolved_paths.append(base_grids) + for path in resolved_paths: + if path: + try: + ok = ok or Path(path).resolve() in Path(filename).resolve().parents + except Exception: + pass return ok diff --git a/modules/hashes.py b/modules/hashes.py index 423fa51b9..ecfb9c914 100644 --- a/modules/hashes.py +++ b/modules/hashes.py @@ -1,9 +1,11 @@ import hashlib import os.path from rich import progress, errors -from modules import shared +from installer import log, console +from modules.json_helpers import readfile, writefile from modules.paths import data_path + cache_filename = os.path.join(data_path, "cache.json") cache_data = None progress_ok = True @@ -12,17 +14,17 @@ progress_ok = True def init_cache(): global cache_data # pylint: disable=global-statement if cache_data is None: - cache_data = {} if not os.path.isfile(cache_filename) else shared.readfile(cache_filename, lock=True, as_type="dict") + cache_data = {} if not os.path.isfile(cache_filename) else readfile(cache_filename, lock=True, as_type="dict") def dump_cache(): - shared.writefile(cache_data, cache_filename) + writefile(cache_data, cache_filename) def cache(subsection): global cache_data # pylint: disable=global-statement if cache_data is None: - cache_data = {} if not os.path.isfile(cache_filename) else shared.readfile(cache_filename, lock=True, as_type="dict") + cache_data = {} if not os.path.isfile(cache_filename) else readfile(cache_filename, lock=True, as_type="dict") s = cache_data.get(subsection, {}) cache_data[subsection] = s return s @@ -35,11 +37,11 @@ def calculate_sha256(filename, quiet=False): if not quiet: if progress_ok: try: - with progress.open(filename, 'rb', description=f'[cyan]Calculating hash: [yellow]{filename}', auto_refresh=True, console=shared.console) as f: + with progress.open(filename, 'rb', description=f'[cyan]Calculating hash: [yellow]{filename}', auto_refresh=True, console=console) as f: for chunk in iter(lambda: f.read(blksize), b""): hash_sha256.update(chunk) except errors.LiveError: - shared.log.warning('Hash: attempting to use function in a thread') + log.warning('Hash: attempting to use function in a thread') progress_ok = False if not progress_ok: with open(filename, 'rb') as f: @@ -65,6 +67,7 @@ def sha256_from_cache(filename, title, use_addnet_hash=False): def sha256(filename, title, use_addnet_hash=False): + from modules import shared global progress_ok # pylint: disable=global-statement hashes = cache("hashes-addnet") if use_addnet_hash else cache("hashes") sha256_value = sha256_from_cache(filename, title, use_addnet_hash) @@ -81,7 +84,7 @@ def sha256(filename, title, use_addnet_hash=False): with progress.open(filename, 'rb', description=f'[cyan]Calculating hash: [yellow]{filename}', auto_refresh=True, console=shared.console) as f: sha256_value = addnet_hash_safetensors(f) except errors.LiveError: - shared.log.warning('Hash: attempting to use function in a thread') + log.warning('Hash: attempting to use function in a thread') progress_ok = False if not progress_ok: with open(filename, 'rb') as f: diff --git a/modules/hidiffusion/__init__.py b/modules/hidiffusion/__init__.py index 004d54b33..858aacf88 100644 --- a/modules/hidiffusion/__init__.py +++ b/modules/hidiffusion/__init__.py @@ -10,6 +10,7 @@ def apply(p, model_type): shared.log.warning(f'HiDiffusion: class={shared.sd_model.__class__.__name__} not supported') return unapply() + pipe = shared.sd_model.pipe if hasattr(shared.sd_model, 'pipe') else shared.sd_model if getattr(p, 'hidiffusion', False) is True: t0 = time.time() hidiffusion.is_aggressive_raunet = shared.opts.hidiffusion_steps > 0 @@ -30,11 +31,12 @@ def apply(p, model_type): hidiffusion.switching_threshold_ratio_dict['sdxl_4096']['T2_ratio'] = t2 hidiffusion.switching_threshold_ratio_dict['sdxl_turbo_1024']['T2_ratio'] = t2 p.extra_generation_params['HiDiffusion Ratios'] = f'{shared.opts.hidiffusion_t1}/{shared.opts.hidiffusion_t2}' - pipe = shared.sd_model.pipe if hasattr(shared.sd_model, 'pipe') else shared.sd_model hidiffusion.apply_hidiffusion(pipe, apply_raunet=shared.opts.hidiffusion_raunet, apply_window_attn=shared.opts.hidiffusion_attn, model_type=model_type, steps=p.steps) p.extra_generation_params['HiDiffusion'] = f'{shared.opts.hidiffusion_raunet}/{shared.opts.hidiffusion_attn}/{shared.opts.hidiffusion_steps > 0}:{shared.opts.hidiffusion_steps}' t1 = time.time() shared.log.debug(f'Applying HiDiffusion: raunet={shared.opts.hidiffusion_raunet} attn={shared.opts.hidiffusion_attn} aggressive={shared.opts.hidiffusion_steps > 0}:{shared.opts.hidiffusion_steps} t1={shared.opts.hidiffusion_t1} t2={shared.opts.hidiffusion_t2} time={t1-t0:.2f} type={shared.sd_model_type} width={p.width} height={p.height}') + elif hasattr(pipe, 'unet') and getattr(pipe.unet, 'hidiffusion', False): + shared.log.warning('HiDiffusion: model reload recomended') def unapply(): diff --git a/modules/hidiffusion/hidiffusion.py b/modules/hidiffusion/hidiffusion.py index e0be29aee..b00d19132 100644 --- a/modules/hidiffusion/hidiffusion.py +++ b/modules/hidiffusion/hidiffusion.py @@ -2,7 +2,6 @@ from typing import Type, Dict, Any, Tuple, Optional import math import torch import torch.nn.functional as F -from diffusers.utils.torch_utils import is_torch_version from diffusers.pipelines import auto_pipeline @@ -85,7 +84,6 @@ def make_diffusers_transformer_block(block_class: Type[torch.nn.Module]) -> Type class transformer_block(block_class): # Save for unpatching later _parent = block_class - _forward = block_class.forward def forward( self, @@ -98,7 +96,6 @@ def make_diffusers_transformer_block(block_class: Type[torch.nn.Module]) -> Type class_labels: Optional[torch.LongTensor] = None, added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None, ) -> torch.FloatTensor: - # reference: https://github.com/microsoft/Swin-Transformer def window_partition(x, window_size, shift_size, H, W): B, _N, C = x.shape @@ -158,7 +155,11 @@ def make_diffusers_transformer_block(block_class: Type[torch.nn.Module]) -> Type # MSW-MSA rand_num = torch.rand(1) _B, N, _C = hidden_states.shape - ori_H, ori_W = self.info['size'] + try: + ori_H, ori_W = self.info['size'] + except Exception as e: + raise RuntimeError(f'HiDiffusion: cls={self.__class__.__name__} info={hasattr(self, "info")} parent={hasattr(self, "_parent")} orphaned call') from e + downsample_ratio = round(((ori_H*ori_W) / N)**0.5) H, W = (math.ceil(ori_H/downsample_ratio), math.ceil(ori_W/downsample_ratio)) widow_size = (math.ceil(H/2), math.ceil(W/2)) @@ -249,7 +250,6 @@ def make_diffusers_transformer_block(block_class: Type[torch.nn.Module]) -> Type hidden_states = hidden_states.squeeze(1) return hidden_states - _patched_forward = forward return transformer_block @@ -257,7 +257,6 @@ def make_diffusers_cross_attn_down_block(block_class: Type[torch.nn.Module]) -> # replace conventional downsampler with resolution-aware downsampler class cross_attn_down_block(block_class): _parent = block_class # Save for unpatching later - _forward = block_class.forward timestep = 0 aggressive_raunet = False T1_ratio = 0 @@ -280,7 +279,10 @@ def make_diffusers_cross_attn_down_block(block_class: Type[torch.nn.Module]) -> self.info['pipeline']._num_timesteps = self.max_timestep # pylint: disable=protected-access self.max_timestep = self.info['pipeline']._num_timesteps # pylint: disable=protected-access # self.max_timestep = len(self.info['scheduler'].timesteps) - ori_H, ori_W = self.info['size'] + try: + ori_H, ori_W = self.info['size'] + except Exception as e: + raise RuntimeError(f'HiDiffusion: cls={self.__class__.__name__} info={hasattr(self, "info")} parent={hasattr(self, "_parent")} orphaned call') from e if self.model == 'sd15': if ori_H < 256 or ori_W < 256: self.T1_ratio = switching_threshold_ratio_dict['sd15_1024'][self.switching_threshold_ratio] @@ -294,8 +296,6 @@ def make_diffusers_cross_attn_down_block(block_class: Type[torch.nn.Module]) -> self.T1_ratio = switching_threshold_ratio_dict['sdxl_2048'][self.switching_threshold_ratio] if self.info['is_inpainting_task']: self.aggressive_raunet = inpainting_is_aggressive_raunet - elif self.info['is_playground']: - self.aggressive_raunet = playground_is_aggressive_raunet else: self.aggressive_raunet = is_aggressive_raunet else: @@ -329,7 +329,7 @@ def make_diffusers_cross_attn_down_block(block_class: Type[torch.nn.Module]) -> return custom_forward - ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} hidden_states = torch.utils.checkpoint.checkpoint( create_custom_forward(resnet), hidden_states, @@ -382,7 +382,6 @@ def make_diffusers_cross_attn_down_block(block_class: Type[torch.nn.Module]) -> return hidden_states, output_states - _patched_forward = forward return cross_attn_down_block @@ -391,7 +390,6 @@ def make_diffusers_cross_attn_up_block(block_class: Type[torch.nn.Module]) -> Ty class cross_attn_up_block(block_class): # Save for unpatching later _parent = block_class - _forward = block_class.forward timestep = 0 aggressive_raunet = False T1_ratio = 0 @@ -411,7 +409,6 @@ def make_diffusers_cross_attn_up_block(block_class: Type[torch.nn.Module]) -> Ty attention_mask: Optional[torch.FloatTensor] = None, encoder_attention_mask: Optional[torch.FloatTensor] = None, ) -> torch.FloatTensor: - def fix_scale(first, second): if (first.shape[-1] != second.shape[-1] or first.shape[-2] != second.shape[-2]): rescale = min(second.shape[-2] / first.shape[-2], second.shape[-1] / first.shape[-1]) @@ -420,7 +417,10 @@ def make_diffusers_cross_attn_up_block(block_class: Type[torch.nn.Module]) -> Ty return first self.max_timestep = self.info['pipeline']._num_timesteps # pylint: disable=protected-access - ori_H, ori_W = self.info['size'] + try: + ori_H, ori_W = self.info['size'] + except Exception as e: + raise RuntimeError(f'HiDiffusion: cls={self.__class__.__name__} info={hasattr(self, "info")} parent={hasattr(self, "_parent")} orphaned call') from e if self.model == 'sd15': if ori_H < 256 or ori_W < 256: self.T1_ratio = switching_threshold_ratio_dict['sd15_1024'][self.switching_threshold_ratio] @@ -435,8 +435,6 @@ def make_diffusers_cross_attn_up_block(block_class: Type[torch.nn.Module]) -> Ty if self.info['is_inpainting_task']: self.aggressive_raunet = inpainting_is_aggressive_raunet - elif self.info['is_playground']: - self.aggressive_raunet = playground_is_aggressive_raunet else: self.aggressive_raunet = is_aggressive_raunet @@ -483,7 +481,6 @@ def make_diffusers_cross_attn_up_block(block_class: Type[torch.nn.Module]) -> Ty self.timestep = 0 return hidden_states - _patched_forward = forward return cross_attn_up_block @@ -492,7 +489,6 @@ def make_diffusers_downsampler_block(block_class: Type[torch.nn.Module]) -> Type class downsampler_block(block_class): # Save for unpatching later _parent = block_class - _forward = block_class.forward T1_ratio = 0 T1 = 0 timestep = 0 @@ -502,7 +498,10 @@ def make_diffusers_downsampler_block(block_class: Type[torch.nn.Module]) -> Type def forward(self, hidden_states: torch.Tensor, scale = 1.0) -> torch.Tensor: # pylint: disable=unused-argument self.max_timestep = self.info['pipeline']._num_timesteps # pylint: disable=protected-access # self.max_timestep = len(self.info['scheduler'].timesteps) - ori_H, ori_W = self.info['size'] + try: + ori_H, ori_W = self.info['size'] + except Exception as e: + raise RuntimeError(f'HiDiffusion: cls={self.__class__.__name__} info={hasattr(self, "info")} parent={hasattr(self, "_parent")} orphaned call') from e if self.model == 'sd15': if ori_H < 256 or ori_W < 256: self.T1_ratio = switching_threshold_ratio_dict['sd15_1024'][self.switching_threshold_ratio] @@ -516,8 +515,6 @@ def make_diffusers_downsampler_block(block_class: Type[torch.nn.Module]) -> Type self.T1_ratio = switching_threshold_ratio_dict['sdxl_2048'][self.switching_threshold_ratio] if self.info['is_inpainting_task']: self.aggressive_raunet = inpainting_is_aggressive_raunet - elif self.info['is_playground']: - self.aggressive_raunet = playground_is_aggressive_raunet else: self.aggressive_raunet = is_aggressive_raunet else: @@ -551,7 +548,6 @@ def make_diffusers_downsampler_block(block_class: Type[torch.nn.Module]) -> Type self.timestep = 0 return hidden_states - _patched_forward = forward return downsampler_block @@ -560,7 +556,6 @@ def make_diffusers_upsampler_block(block_class: Type[torch.nn.Module]) -> Type[t class upsampler_block(block_class): # Save for unpatching later _parent = block_class - _forward = block_class.forward T1_ratio = 0 T1 = 0 timestep = 0 @@ -570,7 +565,10 @@ def make_diffusers_upsampler_block(block_class: Type[torch.nn.Module]) -> Type[t def forward(self, hidden_states: torch.Tensor, scale = 1.0) -> torch.Tensor: # pylint: disable=unused-argument self.max_timestep = self.info['pipeline']._num_timesteps # pylint: disable=protected-access # self.max_timestep = len(self.info['scheduler'].timesteps) - ori_H, ori_W = self.info['size'] + try: + ori_H, ori_W = self.info['size'] + except Exception as e: + raise RuntimeError(f'HiDiffusion: cls={self.__class__.__name__} info={hasattr(self, "info")} parent={hasattr(self, "_parent")} orphaned call') from e if self.model == 'sd15': if ori_H < 256 or ori_W < 256: self.T1_ratio = switching_threshold_ratio_dict['sd15_1024'][self.switching_threshold_ratio] @@ -585,8 +583,6 @@ def make_diffusers_upsampler_block(block_class: Type[torch.nn.Module]) -> Type[t if self.info['is_inpainting_task']: self.aggressive_raunet = inpainting_is_aggressive_raunet - elif self.info['is_playground']: - self.aggressive_raunet = playground_is_aggressive_raunet else: self.aggressive_raunet = is_aggressive_raunet else: @@ -606,7 +602,6 @@ def make_diffusers_upsampler_block(block_class: Type[torch.nn.Module]) -> Type[t return F.conv2d(hidden_states, self.weight, self.bias, self.stride, self.padding, self.dilation, self.groups) - _patched_forward = forward return upsampler_block @@ -632,12 +627,13 @@ def apply_hidiffusion( """ global current_steps # pylint: disable=global-statement current_steps = steps - if hasattr(model, 'controlnet'): + if hasattr(model, 'controlnet') and (model_type == 'sd' or model_type == 'sdxl'): from .hidiffusion_controlnet import make_diffusers_sdxl_contrtolnet_ppl, make_diffusers_unet_2d_condition make_ppl_fn = make_diffusers_sdxl_contrtolnet_ppl model.__class__ = make_ppl_fn(model.__class__) make_block_fn = make_diffusers_unet_2d_condition model.unet.__class__ = make_block_fn(model.unet.__class__) + diffusion_model = model.unet if hasattr(model, "unet") else model diffusion_model.num_upsamplers += 12 diffusion_model.info = { @@ -646,14 +642,13 @@ def apply_hidiffusion( 'hooks': [], 'text_to_img_controlnet': hasattr(model, 'controlnet'), 'is_inpainting_task': model.__class__ in auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING.values(), - 'is_playground': False, 'pipeline': model} - model.info = diffusion_model.info - hook_diffusion_model(diffusion_model) if model_type == 'sd': modified_key = sd15_hidiffusion_key() for key, module in diffusion_model.named_modules(): + if hasattr(module, "_parent"): + raise RuntimeError(f'HiDiffusion: key={key} module={module.__class__} already patched') if apply_raunet and key in modified_key['down_module_key']: module.__class__ = make_diffusers_downsampler_block(module.__class__) module.switching_threshold_ratio = 'T1_ratio' @@ -668,15 +663,15 @@ def apply_hidiffusion( module.switching_threshold_ratio = 'T2_ratio' if apply_window_attn and key in modified_key['windown_attn_module_key']: module.__class__ = make_diffusers_transformer_block(module.__class__) - if hasattr(module, "_patched_forward"): - module.forward = module._patched_forward # pylint: disable=protected-access - module.model = 'sd15' - module.info = diffusion_model.info - + if hasattr(module, "_parent"): + module.model = 'sd15' + module.info = diffusion_model.info elif model_type == 'sdxl': modified_key = sdxl_hidiffusion_key() for key, module in diffusion_model.named_modules(): + if hasattr(module, "_parent"): + raise RuntimeError(f'HiDiffusion: key={key} module={module.__class__} already patched') if apply_raunet and key in modified_key['down_module_key']: module.__class__ = make_diffusers_cross_attn_down_block(module.__class__) module.switching_threshold_ratio = 'T1_ratio' @@ -691,25 +686,26 @@ def apply_hidiffusion( module.switching_threshold_ratio = 'T2_ratio' if apply_window_attn and key in modified_key['windown_attn_module_key']: module.__class__ = make_diffusers_transformer_block(module.__class__) - if hasattr(module, "_patched_forward"): - module.forward = module._patched_forward # pylint: disable=protected-access - module.model = 'sdxl' - module.info = diffusion_model.info + if hasattr(module, "_parent"): + module.model = 'sdxl' + module.info = diffusion_model.info else: raise RuntimeError('HiDiffusion: unsupported model type') - return model + + model.info = diffusion_model.info + model.hidiffusion = True + hook_diffusion_model(diffusion_model) def remove_hidiffusion(model: torch.nn.Module): """ Removes hidiffusion from a Diffusion module if it was already patched. """ - for _, module in model.unet.named_modules(): + model = model.unet if hasattr(model, "unet") else model + for _, module in model.named_modules(): + while hasattr(module, "_parent"): + model.hidiffusion = True + module.__class__ = module._parent # pylint: disable=protected-access if hasattr(module, "info"): - for hook in module.info["hooks"]: + for hook in module.info.get("hooks", []): hook.remove() module.info["hooks"].clear() del module.info - if hasattr(module, "_forward"): - module.forward = module._forward # pylint: disable=protected-access - if hasattr(module, "_parent"): - module.__class__ = module._parent # pylint: disable=protected-access - return model diff --git a/modules/hidiffusion/hidiffusion_controlnet.py b/modules/hidiffusion/hidiffusion_controlnet.py index dd3dcd115..7a81ab066 100644 --- a/modules/hidiffusion/hidiffusion_controlnet.py +++ b/modules/hidiffusion/hidiffusion_controlnet.py @@ -15,6 +15,7 @@ def make_diffusers_unet_2d_condition(block_class): class unet_2d_condition(block_class): # Save for unpatching later _parent = block_class + def forward( self, sample: torch.FloatTensor, @@ -549,7 +550,7 @@ def make_diffusers_sdxl_contrtolnet_ppl(block_class): # # scale the initial noise by the standard deviation required by the scheduler # latents = latents * self.scheduler.init_noise_sigma - # 7. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 7. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 7.1 Create tensor stating which controlnets to keep diff --git a/modules/history.py b/modules/history.py index d9f2b5d0b..9c0c9c0c6 100644 --- a/modules/history.py +++ b/modules/history.py @@ -1,8 +1,5 @@ """ -TODO: -- apply metadata -- preview -- load/save +TODO: apply metadata, preview, load/save """ import sys diff --git a/modules/images.py b/modules/images.py index f16ae9380..c54f982eb 100644 --- a/modules/images.py +++ b/modules/images.py @@ -164,7 +164,7 @@ def save_image(image, if not check_grid_size([image]): return None, None, None if path is None or path == '': # set default path to avoid errors when functions are triggered manually or via api and param is not set - path = shared.opts.outdir_save + path = paths.resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_save) namegen = FilenameGenerator(p, seed, prompt, image, grid=grid) suffix = suffix if suffix is not None else '' basename = '' if basename is None else basename @@ -311,7 +311,7 @@ def parse_novelai_metadata(data: dict): return geninfo -def read_info_from_image(image: Image, watermark: bool = False): +def read_info_from_image(image: Image.Image, watermark: bool = False): if image is None: return '', {} if isinstance(image, str): @@ -419,7 +419,7 @@ def draw_overlay(im, text: str = '', y_offset: int = 0): return im -def set_watermark(image, wm_text: str = None, wm_image: Image.Image = None): +def set_watermark(image, wm_text: str | None = None, wm_image: Image.Image | None = None): if shared.opts.image_watermark_position != 'none' and wm_image is not None: # visible watermark if isinstance(wm_image, str): try: diff --git a/modules/img2img.py b/modules/img2img.py index 4809d9fa0..bc080f732 100644 --- a/modules/img2img.py +++ b/modules/img2img.py @@ -7,6 +7,7 @@ from modules import scripts_manager, shared, processing, images, errors from modules.generation_parameters_copypaste import create_override_settings_dict from modules.ui_common import plaintext_to_html from modules.memstats import memory_stats +from modules.paths import resolve_output_path debug = shared.log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -239,8 +240,8 @@ def img2img(id_task: str, state: str, mode: int, p = processing.StableDiffusionProcessingImg2Img( sd_model=shared.sd_model, - outpath_samples=shared.opts.outdir_samples or shared.opts.outdir_img2img_samples, - outpath_grids=shared.opts.outdir_grids or shared.opts.outdir_img2img_grids, + outpath_samples=resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_img2img_samples), + outpath_grids=resolve_output_path(shared.opts.outdir_grids, shared.opts.outdir_img2img_grids), prompt=prompt, negative_prompt=negative_prompt, styles=prompt_styles, diff --git a/modules/intel/ipex/hijacks.py b/modules/intel/ipex/hijacks.py index 56999da21..327d9f18a 100644 --- a/modules/intel/ipex/hijacks.py +++ b/modules/intel/ipex/hijacks.py @@ -143,12 +143,11 @@ def torch_tensor(data, *args, dtype=None, device=None, **kwargs): global device_supports_fp64 if check_cuda(device): device = return_xpu(device) - if not device_supports_fp64: - if check_device_type(device, "xpu"): - if dtype == torch.float64: - dtype = torch.float32 - elif dtype is None and (hasattr(data, "dtype") and (data.dtype == torch.float64 or data.dtype == float)): - dtype = torch.float32 + if not device_supports_fp64 and check_device_type(device, "xpu"): + if dtype == torch.float64: + dtype = torch.float32 + elif dtype is None and (hasattr(data, "dtype") and (data.dtype == torch.float64 or data.dtype == float)): + dtype = torch.float32 return original_torch_tensor(data, *args, dtype=dtype, device=device, **kwargs) @@ -223,8 +222,6 @@ def torch_empty(*args, device=None, **kwargs): original_torch_randn = torch.randn @wraps(torch.randn) def torch_randn(*args, device=None, dtype=None, **kwargs): - if dtype is bytes: - dtype = None if check_cuda(device): return original_torch_randn(*args, device=return_xpu(device), dtype=dtype, **kwargs) else: @@ -258,13 +255,32 @@ def torch_full(*args, device=None, **kwargs): return original_torch_full(*args, device=device, **kwargs) +original_torch_arange = torch.arange +@wraps(torch.arange) +def torch_arange(*args, device=None, dtype=None, **kwargs): + global device_supports_fp64 + if check_cuda(device): + if not device_supports_fp64 and dtype == torch.float64: + dtype = torch.float32 + return original_torch_arange(*args, device=return_xpu(device), dtype=dtype, **kwargs) + else: + if not device_supports_fp64 and check_device_type(device, "xpu") and dtype == torch.float64: + dtype = torch.float32 + return original_torch_arange(*args, device=device, dtype=dtype, **kwargs) + + original_torch_linspace = torch.linspace @wraps(torch.linspace) -def torch_linspace(*args, device=None, **kwargs): +def torch_linspace(*args, device=None, dtype=None, **kwargs): + global device_supports_fp64 if check_cuda(device): - return original_torch_linspace(*args, device=return_xpu(device), **kwargs) + if not device_supports_fp64 and dtype == torch.float64: + dtype = torch.float32 + return original_torch_linspace(*args, device=return_xpu(device), dtype=dtype, **kwargs) else: - return original_torch_linspace(*args, device=device, **kwargs) + if not device_supports_fp64 and check_device_type(device, "xpu") and dtype == torch.float64: + dtype = torch.float32 + return original_torch_linspace(*args, device=device, dtype=dtype, **kwargs) original_torch_eye = torch.eye @@ -358,6 +374,7 @@ def ipex_hijacks(): torch.ones = torch_ones torch.zeros = torch_zeros torch.full = torch_full + torch.arange = torch_arange torch.linspace = torch_linspace torch.eye = torch_eye torch.load = torch_load diff --git a/modules/interrogate/vqa.py b/modules/interrogate/vqa.py index ded9ad9c4..036fd7ced 100644 --- a/modules/interrogate/vqa.py +++ b/modules/interrogate/vqa.py @@ -365,12 +365,14 @@ class VQA: def unload(self): """Release VLM model from GPU/memory.""" if self.model is not None: - shared.log.debug(f'VQA unload: model="{self.loaded}"') + model_name = self.loaded + shared.log.debug(f'VQA unload: unloading model="{model_name}"') sd_models.move_model(self.model, devices.cpu, force=True) self.model = None self.processor = None self.loaded = None devices.torch_gc(force=True, reason='vqa unload') + shared.log.debug(f'VQA unload: model="{model_name}" unloaded') else: shared.log.debug('VQA unload: no model loaded') @@ -521,9 +523,6 @@ class VQA: cls_name = self.model.__class__.__name__ debug(f'VQA interrogate: handler=qwen model_name="{model_name}" model_class="{cls_name}" repo="{repo}" question="{question}" system_prompt="{system_prompt}" image_size={image.size if image else None}') - # Warn if using Florence-2 task tokens with non-Florence-2 models - if is_florence_task(question): - shared.log.warning(f'Interrogate: Florence-2 task token "{question}" is designed for Florence-2 models. Using it anyway, but results may vary.') question = question.replace('<', '').replace('>', '').replace('_', ' ') system_prompt = system_prompt or shared.opts.interrogate_vlm_system conversation = [ @@ -658,9 +657,6 @@ class VQA: cls_name = self.model.__class__.__name__ debug(f'VQA interrogate: handler=gemma model_name="{model_name}" model_class="{cls_name}" repo="{repo}" question="{question}" system_prompt="{system_prompt}" image_size={image.size if image else None}') - # Warn if using Florence-2 task tokens with non-Florence-2 models - if is_florence_task(question): - shared.log.warning(f'Interrogate: Florence-2 task token "{question}" is designed for Florence-2 models. Using it anyway, but results may vary.') question = question.replace('<', '').replace('>', '').replace('_', ' ') system_prompt = system_prompt or shared.opts.interrogate_vlm_system @@ -860,9 +856,6 @@ class VQA: cls_name = self.model.__class__.__name__ debug(f'VQA interrogate: handler=smol model_name="{model_name}" model_class="{cls_name}" repo="{repo}" question="{question}" system_prompt="{system_prompt}" image_size={image.size if image else None}') - # Warn if using Florence-2 task tokens with non-Florence-2 models - if is_florence_task(question): - shared.log.warning(f'Interrogate: Florence-2 task token "{question}" is designed for Florence-2 models. Using it anyway, but results may vary.') question = question.replace('<', '').replace('>', '').replace('_', ' ') system_prompt = system_prompt or shared.opts.interrogate_vlm_system conversation = [ diff --git a/modules/json_helpers.py b/modules/json_helpers.py index e06c835f4..7d28b3e01 100644 --- a/modules/json_helpers.py +++ b/modules/json_helpers.py @@ -19,13 +19,13 @@ def readfile(filename: str, silent: bool = False, lock: bool = False, *, as_type def readfile(filename: str, silent: bool = False, lock: bool = False) -> dict | list: ... def readfile(filename: str, silent: bool = False, lock: bool = False, *, as_type="") -> dict | list: global locking_available # pylint: disable=global-statement - data = {} + data = {} if as_type == "dict" else [] lock_file = None locked = False if lock and locking_available: try: lock_file = fasteners.InterProcessReaderWriterLock(f"{filename}.lock") - lock_file.logger.disabled = True + lock_file.logger.disabled = True # type: ignore - False positive. Bad typing in Fasteners. locked = lock_file.acquire_read_lock(blocking=True, timeout=3) except Exception as err: lock_file = None @@ -59,11 +59,17 @@ def readfile(filename: str, silent: bool = False, lock: bool = False, *, as_type except Exception: locking_available = False if isinstance(data, list) and as_type == "dict": + if not data: + return {} + log.warning(f"Read: Expected dictionary from '{filename}' but got list") data0 = data[0] if isinstance(data0, dict): return data0 return {} if isinstance(data, dict) and as_type == "list": + if not data: + return [] + log.warning(f"Read: Expected list from '{filename}' but got dictionary") return [data] return data @@ -99,7 +105,7 @@ def writefile(data, filename, mode='w', silent=False, atomic=False): try: if locking_available: lock_file = fasteners.InterProcessReaderWriterLock(f"{filename}.lock") if locking_available else None - lock_file.logger.disabled = True + lock_file.logger.disabled = True # type: ignore - False positive. Bad typing in Fasteners. locked = lock_file.acquire_write_lock(blocking=True, timeout=3) if lock_file is not None else False except Exception as err: locking_available = False @@ -118,7 +124,8 @@ def writefile(data, filename, mode='w', silent=False, atomic=False): file.write(output) t1 = time.time() if not silent: - log.debug(f'Save: file="{filename}" json={len(data)} bytes={len(output)} time={t1-t0:.3f}') + datalength = len(data) if isinstance(data, (dict, list)) else (len(data.__dict__)) + log.debug(f'Save: file="{filename}" json={datalength} bytes={len(output)} time={t1-t0:.3f}') except Exception as err: log.error(f'Save failed: file="{filename}" {err}') try: diff --git a/modules/loader.py b/modules/loader.py index 5814baabb..c6e25a1d3 100644 --- a/modules/loader.py +++ b/modules/loader.py @@ -13,7 +13,7 @@ initialized = False errors.install() logging.getLogger("DeepSpeed").disabled = True timer.startup.record("loader") - +errors.log.debug('Initializing: libraries') np = None try: @@ -99,14 +99,23 @@ except Exception: _bnb = False timer.startup.record("bnb") +import huggingface_hub # pylint: disable=W0611,C0411 +logging.getLogger("huggingface_hub.file_download").setLevel(logging.ERROR) +if huggingface_hub.__version__.startswith('0.'): + huggingface_hub.is_offline_mode = lambda: False +timer.startup.record("hfhub") + +import accelerate # pylint: disable=W0611,C0411 +timer.startup.record("accelerate") + +import pydantic # pylint: disable=W0611,C0411 +timer.startup.record("pydantic") + import transformers # pylint: disable=W0611,C0411 from transformers import logging as transformers_logging # pylint: disable=W0611,C0411 transformers_logging.set_verbosity_error() timer.startup.record("transformers") -import accelerate # pylint: disable=W0611,C0411 -timer.startup.record("accelerate") - try: import onnxruntime # pylint: disable=W0611,C0411 onnxruntime.set_default_logger_severity(4) @@ -121,9 +130,6 @@ import gradio # pylint: disable=W0611,C0411 timer.startup.record("gradio") errors.install([gradio]) -import pydantic # pylint: disable=W0611,C0411 -timer.startup.record("pydantic") - # patch different progress bars import tqdm as tqdm_lib # pylint: disable=C0411 from tqdm.rich import tqdm # pylint: disable=W0611,C0411 @@ -145,10 +151,6 @@ except Exception as e: errors.log.error('Please restart re-run the installer') sys.exit(1) -import huggingface_hub # pylint: disable=W0611,C0411 -logging.getLogger("huggingface_hub.file_download").setLevel(logging.ERROR) -timer.startup.record("hfhub") - try: import pillow_jxl # pylint: disable=W0611,C0411 except Exception: @@ -185,6 +187,7 @@ def get_packages(): "gradio": gradio.__version__, "transformers": transformers.__version__, "accelerate": accelerate.__version__, + "hub": huggingface_hub.__version__, } try: diff --git a/modules/lora/extra_networks_lora.py b/modules/lora/extra_networks_lora.py index 1c73712cd..882c0d91b 100644 --- a/modules/lora/extra_networks_lora.py +++ b/modules/lora/extra_networks_lora.py @@ -97,7 +97,7 @@ def parse(p, params_list, step=0): dyn_dims = [] lora_modules = [] for params in params_list: - names.append(params.positional[0]) + name = params.positional[0] default_multiplier = params.positional[1] if len(params.positional) > 1 else shared.opts.extra_networks_default_multiplier default_multiplier = to_float(default_multiplier) @@ -121,6 +121,11 @@ def parse(p, params_list, step=0): unet_multiplier[i] = to_float(unet_multiplier[i]) dyn_dim = int(params.named["dyn"]) if "dyn" in params.named else None + + if (te_multiplier == 0) and all(u == 0 for u in unet_multiplier): # skip lora with strength zero + continue + + names.append(name) te_multipliers.append(te_multiplier) unet_multipliers.append(unet_multiplier) dyn_dims.append(dyn_dim) @@ -180,7 +185,7 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): exclude = ['none'] key = f'include={",".join(include)}:exclude={",".join(exclude)}' loaded = sd_model.loaded_loras.get(key, []) - debug_log(f'Network load: type=LoRA key="{key}" requested={requested} loaded={loaded}') + debug_log(f'Network check: type=LoRA key="{key}" requested={requested} loaded={loaded}') if len(requested) != len(loaded): sd_model.loaded_loras[key] = requested return True @@ -223,18 +228,19 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): if has_changed: jobid = shared.state.begin('LoRA') if len(l.previously_loaded_networks) > 0: - shared.log.info(f'Network unload: type=LoRA apply={[n.name for n in l.previously_loaded_networks]} mode={"fuse" if shared.opts.lora_fuse_diffusers else "backup"}') + shared.log.info(f'Network unload: type=LoRA networks={[n.name for n in l.previously_loaded_networks]} mode={"fuse" if shared.opts.lora_fuse_native else "backup"}') networks.network_deactivate(include, exclude) networks.network_activate(include, exclude) - l.previously_loaded_networks = l.loaded_networks.copy() - debug_log(f'Network load: type=LoRA previous={[n.name for n in l.previously_loaded_networks]} current={[n.name for n in l.loaded_networks]} changed') + debug_log(f'Network change: type=LoRA previous={[n.name for n in l.previously_loaded_networks]} current={[n.name for n in l.loaded_networks]}') + if len(include) == 0: + l.previously_loaded_networks = l.loaded_networks.copy() shared.state.end(jobid) if len(l.loaded_networks) > 0 and (len(networks.applied_layers) > 0 or load_method=='diffusers' or load_method=='nunchaku') and step == 0: infotext(p) prompt(p) if has_changed and len(include) == 0: # print only once - shared.log.info(f'Network load: type=LoRA apply={[n.name for n in l.loaded_networks]} method={load_method} mode={"fuse" if shared.opts.lora_fuse_native else "backup"} te={te_multipliers} unet={unet_multipliers} time={l.timer.summary}') + shared.log.info(f'Network load: type=LoRA networks={[n.name for n in l.loaded_networks]} method={load_method} mode={"fuse" if shared.opts.lora_fuse_native else "backup"} te={te_multipliers} unet={unet_multipliers} time={l.timer.summary}') def deactivate(self, p, force=False): if len(lora_diffusers.diffuser_loaded) > 0: diff --git a/modules/lora/lora_apply.py b/modules/lora/lora_apply.py index 5d9a3829b..06b896349 100644 --- a/modules/lora/lora_apply.py +++ b/modules/lora/lora_apply.py @@ -213,7 +213,7 @@ def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.G quantization_device=devices.device, return_device=device, param_name=getattr(self, 'network_layer_name', None), - ).to(device) + )[0].to(device) weight = None del dequant_weight except Exception as e: diff --git a/modules/lora/lora_load.py b/modules/lora/lora_load.py index 2a54707f7..14ee012ad 100644 --- a/modules/lora/lora_load.py +++ b/modules/lora/lora_load.py @@ -304,6 +304,7 @@ def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=Non shared.log.error(f'Network load: type=LoRA action=fuse {str(e)}') if l.debug: errors.display(e, 'LoRA') + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, force=True) # some layers may end up on cpu without hook if len(l.loaded_networks) > 0 and l.debug: shared.log.debug(f'Network load: type=LoRA loaded={[n.name for n in l.loaded_networks]} cache={list(lora_cache)} fuse={shared.opts.lora_fuse_native}:{shared.opts.lora_fuse_diffusers}') diff --git a/modules/lora/lora_overrides.py b/modules/lora/lora_overrides.py index 640181977..d877889ca 100644 --- a/modules/lora/lora_overrides.py +++ b/modules/lora/lora_overrides.py @@ -1,12 +1,14 @@ from modules import shared -maybe_diffusers = [ # forced if lora_maybe_diffusers is enabled +force_hashes_diffusers = [ # forced always + # '816d0eed49fd', # flash-sdxl + # 'c2ec22757b46', # flash-sd15 + # '22c8339e7666', # spo-sdxl-10ep # 'aaebf6360f7d', # sd15-lcm # '3d18b05e4f56', # sdxl-lcm # 'b71dcb732467', # sdxl-tcd # '813ea5fb1c67', # sdxl-turbo - # not really needed, but just in case # '5a48ac366664', # hyper-sd15-1step # 'ee0ff23dcc42', # hyper-sd15-2step # 'e476eb1da5df', # hyper-sd15-4step @@ -19,41 +21,14 @@ maybe_diffusers = [ # forced if lora_maybe_diffusers is enabled # '8cca3706050b', # hyper-sdxl-1step ] -force_diffusers = [ # forced always - '816d0eed49fd', # flash-sdxl - 'c2ec22757b46', # flash-sd15 - '22c8339e7666', # spo-sdxl-10ep +allow_native = [ + 'sd', + 'sdxl', + 'sd3', + 'f1', + 'chroma', ] -force_models_diffusers = [ # forced always - # 'sd3', - 'sc', - 'h1', - 'kandinsky5', - 'kandinsky3', - 'kandinsky', - 'hunyuandit', - 'hunyuanimage', - 'auraflow', - 'lumina2', - 'qwen', - 'bria', - 'flite', - 'cosmos', - 'chrono', - 'z_image', - 'f2', - 'longcat', - # video models - 'hunyuanvideo', - 'hunyuanvideo15' - 'cogvideo', - 'wanai', - 'chrono', - 'ltxvideo', - 'mochivideo', - 'allegrovideo', -] force_classes_diffusers = [ # forced always 'FluxKontextPipeline', 'FluxKontextInpaintPipeline', @@ -65,11 +40,9 @@ fuse_ignore = [ def get_method(shorthash=''): - use_diffusers = shared.opts.lora_force_diffusers or (shared.sd_model_type in force_models_diffusers) or (shared.sd_model.__class__.__name__ in force_classes_diffusers) - if shared.opts.lora_maybe_diffusers and len(shorthash) > 4: - use_diffusers = use_diffusers or any(x.startswith(shorthash) for x in maybe_diffusers) - if shared.opts.lora_force_diffusers and len(shorthash) > 4: - use_diffusers = use_diffusers or any(x.startswith(shorthash) for x in force_diffusers) + use_diffusers = shared.opts.lora_force_diffusers or (shared.sd_model.__class__.__name__ in force_classes_diffusers) or (shared.sd_model_type not in allow_native) + if len(shorthash) > 4: + use_diffusers = use_diffusers or any(x.startswith(shorthash) for x in force_hashes_diffusers) nunchaku_dit = hasattr(shared.sd_model, 'transformer') and 'Nunchaku' in shared.sd_model.transformer.__class__.__name__ nunchaku_unet = hasattr(shared.sd_model, 'unet') and 'Nunchaku' in shared.sd_model.unet.__class__.__name__ use_nunchaku = nunchaku_dit or nunchaku_unet diff --git a/modules/lora/network.py b/modules/lora/network.py index f272242f1..b8a09913b 100644 --- a/modules/lora/network.py +++ b/modules/lora/network.py @@ -65,6 +65,10 @@ class NetworkOnDisk: return 'hv' if base.startswith("chroma"): return 'chroma' + if base.startswith('zimage'): + return 'zimage' + if base.startswith('qwen'): + return 'qwen' if arch.startswith("stable-diffusion-v1"): return 'sd1' diff --git a/modules/lora/network_lora.py b/modules/lora/network_lora.py index 0e980e1bf..fa93a4aaa 100644 --- a/modules/lora/network_lora.py +++ b/modules/lora/network_lora.py @@ -27,8 +27,8 @@ class NetworkModuleLora(network.NetworkModule): return None linear_modules = [torch.nn.Linear, torch.nn.modules.linear.NonDynamicallyQuantizableLinear, torch.nn.MultiheadAttention, diffusers_lora.LoRACompatibleLinear] typ = type(self.sd_module) - is_linear = typ in linear_modules or self.sd_module.__class__.__name__ in ["NNCFLinear", "QLinear", "Linear4bit"] - is_conv = (typ in [torch.nn.Conv2d, diffusers_lora.LoRACompatibleConv]) or (self.sd_module.__class__.__name__ in ["NNCFConv2d", "QConv2d"]) or (typ.__name__ in ['downsampler_block', 'upsampler_block']) + is_linear = typ in linear_modules or self.sd_module.__class__.__name__ in ["SDNQLinear", "QLinear", "Linear4bit"] + is_conv = (typ in [torch.nn.Conv2d, diffusers_lora.LoRACompatibleConv]) or (self.sd_module.__class__.__name__ in ["SDNQConv2d", "QConv2d"]) or (typ.__name__ in ['downsampler_block', 'upsampler_block']) if is_linear: weight = weight.reshape(weight.shape[0], -1) module = torch.nn.Linear(weight.shape[1], weight.shape[0], bias=False) diff --git a/modules/lora/networks.py b/modules/lora/networks.py index 4294615c9..6d37fd656 100644 --- a/modules/lora/networks.py +++ b/modules/lora/networks.py @@ -7,6 +7,7 @@ from modules import shared, devices, sd_models applied_layers: list[str] = [] +default_components = ['text_encoder', 'text_encoder_2', 'text_encoder_3', 'text_encoder_4', 'unet', 'transformer', 'transformer_2'] def network_activate(include=[], exclude=[]): @@ -17,7 +18,7 @@ def network_activate(include=[], exclude=[]): sd_models.move_model(sd_model, device=devices.cpu) device = None modules = {} - components = include if len(include) > 0 else ['text_encoder', 'text_encoder_2', 'text_encoder_3', 'unet', 'transformer'] + components = include if len(include) > 0 else default_components components = [x for x in components if x not in exclude] active_components = [] for name in components: @@ -43,7 +44,7 @@ def network_activate(include=[], exclude=[]): for _, module in modules[component]: network_layer_name = getattr(module, 'network_layer_name', None) current_names = getattr(module, "network_current_names", ()) - if getattr(module, 'weight', None) is None or shared.state.interrupted or network_layer_name is None or current_names == wanted_names: + if getattr(module, 'weight', None) is None or shared.state.interrupted or (network_layer_name is None) or (current_names == wanted_names): if task is not None: pbar.update(task, advance=1) continue diff --git a/modules/ltx/ltx_process.py b/modules/ltx/ltx_process.py index 8b429a593..d2c077607 100644 --- a/modules/ltx/ltx_process.py +++ b/modules/ltx/ltx_process.py @@ -1,11 +1,10 @@ -""" -- modernui -- teacache and others -""" import os import time -import threading -from modules import shared, errors, timer, memstats, progress, processing, sd_models, sd_samplers, extra_networks +import torch +from PIL import Image + +from modules import shared, errors, timer, memstats, progress, processing, sd_models, sd_samplers, extra_networks, call_queue +from modules.video_models.video_vae import set_vae_params from modules.video_models.video_save import save_video from modules.video_models.video_utils import check_av from modules.processing_callbacks import diffusers_callback @@ -16,7 +15,6 @@ debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None e # engine, model = 'LTX Video', 'LTXVideo 0.9.7 13B' upsample_repo_id = "a-r-r-o-w/LTX-Video-0.9.7-Latent-Spatial-Upsampler-diffusers" upsample_pipe = None -queue_lock = threading.Lock() def run_ltx(task_id, @@ -52,6 +50,7 @@ def run_ltx(task_id, mp4_video:bool, mp4_frames:bool, mp4_sf:bool, + audio_enable:bool, _overrides, ): @@ -73,7 +72,7 @@ def run_ltx(task_id, # from diffusers import LTXConditionPipeline # pylint: disable=unused-import check_av() progress.add_task_to_queue(task_id) - with queue_lock: + with call_queue.get_lock(): progress.start_task(task_id) memstats.reset_stats() timer.process.reset() @@ -123,11 +122,18 @@ def run_ltx(task_id, sampler_name = processing.get_sampler_name(sampler_index) sd_samplers.create_sampler(sampler_name, shared.sd_model) shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} op=init styles={styles} networks={networks} sampler={shared.sd_model.scheduler.__class__.__name__}') + extra_networks.activate(p, networks) + framewise = 'LTX2' not in shared.sd_model.__class__.__name__ + set_vae_params(p, framewise=framewise) t0 = time.time() shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) t1 = time.time() + if 'LTX2' in shared.sd_model.__class__.__name__: + output_type = 'np' + else: + output_type = 'latent' base_args = { "prompt": prompt, "negative_prompt": negative, @@ -135,24 +141,41 @@ def run_ltx(task_id, "height": get_bucket(height), "num_frames": get_frames(frames), "num_inference_steps": steps, - "image_cond_noise_scale": image_cond_noise_scale, "generator": get_generator(seed), "callback_on_step_end": diffusers_callback, - "output_type": "latent", + "output_type": output_type, } + if 'LTX2' in shared.sd_model.__class__.__name__: + base_args["frame_rate"] = float(mp4_fps) + if 'Condition' in shared.sd_model.__class__.__name__: + base_args["image_cond_noise_scale"] = image_cond_noise_scale shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} op=base {base_args}') if len(conditions) > 0: base_args["conditions"] = conditions + + if debug: + shared.log.trace(f'LTX args: {base_args}') yield None, 'LTX: Generate in progress...' samplejob = shared.state.begin('Sample') try: - latents = shared.sd_model(**base_args).frames[0] + result = shared.sd_model(**base_args) + latents = result.frames[0] except AssertionError as e: yield from abort(e, ok=True, p=p) return except Exception as e: yield from abort(e, ok=False, p=p) return + if audio_enable and hasattr(result, 'audio') and result.audio is not None: + audio = result.audio[0].float().cpu() + else: + audio = None + try: + if debug: + shared.log.trace(f'LTX result frames={latents.shape if latents is not None else None} audio={audio.shape if audio is not None else None}') + except Exception: + pass + t2 = time.time() shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) t3 = time.time() @@ -171,7 +194,7 @@ def run_ltx(task_id, "width": get_bucket(upsample_ratio * width), "height": get_bucket(upsample_ratio * height), "generator": get_generator(seed), - "output_type": "latent", + "output_type": output_type, } if latents.ndim == 4: latents = latents.unsqueeze(0) # add batch dimension @@ -208,10 +231,11 @@ def run_ltx(task_id, "image_cond_noise_scale": image_cond_noise_scale, "generator": get_generator(seed), "callback_on_step_end": diffusers_callback, - "output_type": "latent", + "output_type": output_type, } if latents.ndim == 4: latents = latents.unsqueeze(0) # add batch dimension + shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} op=refine latents={latents.shape} {refine_args}') if len(conditions) > 0: refine_args["conditions"] = conditions @@ -236,7 +260,12 @@ def run_ltx(task_id, yield None, 'LTX: VAE decode in progress...' try: - frames = vae_decode(latents, decode_timestep, seed) + if torch.is_tensor(latents): + frames = vae_decode(latents, decode_timestep, seed) + else: + frames = latents + except TypeError: + frames = latents # likely because the latents are already decoded except AssertionError as e: yield from abort(e, ok=True, p=p) return @@ -248,9 +277,15 @@ def run_ltx(task_id, t11 = time.time() timer.process.add('offload', t11 - t10) + try: + aac_sample_rate = shared.sd_model.vocoder.config.output_sampling_rate + except Exception: + aac_sample_rate = 24000 + num_frames, video_file = save_video( p=p, pixels=frames, + audio=audio, mp4_fps=mp4_fps, mp4_codec=mp4_codec, mp4_opt=mp4_opt, @@ -259,11 +294,19 @@ def run_ltx(task_id, mp4_video=mp4_video, mp4_frames=mp4_frames, mp4_interpolate=mp4_interpolate, + aac_sample_rate=aac_sample_rate, metadata={}, ) t_end = time.time() - _n, _c, _t, h, w = frames.shape + if isinstance(frames, list) and isinstance(frames[0], Image.Image): + w, h = frames[0].size + elif frames.ndim == 5: + _n, _c, _t, h, w = frames.shape + elif frames.ndim == 4: + _n, h, w, _c = frames.shape + else: + h, w = frames.shape[-2], frames.shape[-1] resolution = f'{w}x{h}' if num_frames > 0 else None summary = timer.process.summary(min_time=0.25, total=False).replace('=', ' ') memory = shared.mem_mon.summary() diff --git a/modules/ltx/ltx_ui.py b/modules/ltx/ltx_ui.py index 7d99bef65..d175e4be6 100644 --- a/modules/ltx/ltx_ui.py +++ b/modules/ltx/ltx_ui.py @@ -38,6 +38,9 @@ def create_ui(prompt, negative, styles, overrides, init_image, init_strength, la with gr.Row(): decode_timestep = gr.Slider(label='LTX decode timestep', minimum=0.01, maximum=1.0, step=0.01, value=0.05, elem_id="ltx_decode_timestep") image_cond_noise_scale = gr.Slider(label='Noise scale', minimum=0.01, maximum=1.0, step=0.01, value=0.025, elem_id="ltx_image_cond_noise_scale") + with gr.Accordion(open=False, label="Audio", elem_id='ltx_audio_accordion'): + with gr.Row(): + audio_enable = gr.Checkbox(label='LTX enable audio', value=False, elem_id="ltx_audio_enable") with gr.Column(elem_id='ltx-output-column', scale=2) as _column_output: with gr.Row(): @@ -60,6 +63,7 @@ def create_ui(prompt, negative, styles, overrides, init_image, init_strength, la init_strength, init_image, last_image, condition_files, condition_video, condition_video_frames, condition_video_skip, decode_timestep, image_cond_noise_scale, mp4_fps, mp4_interpolate, mp4_codec, mp4_ext, mp4_opt, mp4_video, mp4_frames, mp4_sf, + audio_enable, overrides, ] video_outputs = [ diff --git a/modules/ltx/ltx_util.py b/modules/ltx/ltx_util.py index a329373fd..5264b11dc 100644 --- a/modules/ltx/ltx_util.py +++ b/modules/ltx/ltx_util.py @@ -9,7 +9,7 @@ loaded_model: str = None def get_bucket(size: int): if not hasattr(shared.sd_model, 'vae_temporal_compression_ratio'): - return int(size) - (int(size) % 16) + return int(size) - (int(size) % 32) return int(size) - (int(size) % shared.sd_model.vae_temporal_compression_ratio) diff --git a/modules/masking.py b/modules/masking.py index a23fee32f..ea1844c19 100644 --- a/modules/masking.py +++ b/modules/masking.py @@ -118,7 +118,7 @@ def fill(image, mask): """ [docs](https://huggingface.co/docs/transformers/v4.36.1/en/model_doc/sam#overview) -TODO: +TODO: additional masking algorithms - PerSAM - REMBG - https://huggingface.co/docs/transformers/tasks/semantic_segmentation diff --git a/modules/mit_nunchaku.py b/modules/mit_nunchaku.py index 268e29f56..b5e82c1da 100644 --- a/modules/mit_nunchaku.py +++ b/modules/mit_nunchaku.py @@ -4,7 +4,7 @@ from installer import log, pip from modules import devices -nunchaku_ver = '1.0.2' +nunchaku_ver = '1.1.0' ok = False diff --git a/modules/model_quant.py b/modules/model_quant.py index 940594890..1a501be0a 100644 --- a/modules/model_quant.py +++ b/modules/model_quant.py @@ -221,10 +221,12 @@ def create_sdnq_config(kwargs = None, allow: bool = True, module: str = 'Model', group_size=shared.opts.sdnq_quantize_weights_group_size, svd_rank=shared.opts.sdnq_svd_rank, svd_steps=shared.opts.sdnq_svd_steps, + dynamic_loss_threshold=shared.opts.sdnq_dynamic_loss_threshold, use_svd=shared.opts.sdnq_use_svd, quant_conv=shared.opts.sdnq_quantize_conv_layers, use_quantized_matmul=shared.opts.sdnq_use_quantized_matmul, use_quantized_matmul_conv=shared.opts.sdnq_use_quantized_matmul_conv, + use_dynamic_quantization=shared.opts.sdnq_use_dynamic_quantization, dequantize_fp32=shared.opts.sdnq_dequantize_fp32, non_blocking=shared.opts.diffusers_offload_nonblocking, quantization_device=quantization_device, @@ -234,7 +236,8 @@ def create_sdnq_config(kwargs = None, allow: bool = True, module: str = 'Model', ) if quantized_matmul_dtype is None: quantized_matmul_dtype = "auto" # set for logging - log.debug(f'Quantization: module="{module}" type=sdnq mode=pre dtype={weights_dtype} matmul_dtype={quantized_matmul_dtype} matmul={shared.opts.sdnq_use_quantized_matmul} group_size={shared.opts.sdnq_quantize_weights_group_size} svd_rank={shared.opts.sdnq_svd_rank} svd_steps={shared.opts.sdnq_svd_steps} use_svd={shared.opts.sdnq_use_svd} quant_conv={shared.opts.sdnq_quantize_conv_layers} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} dequantize_fp32={shared.opts.sdnq_dequantize_fp32} quantize_with_gpu={shared.opts.sdnq_quantize_with_gpu} quantization_device={quantization_device} return_device={return_device} device_map={shared.opts.device_map} offload_mode={shared.opts.diffusers_offload_mode} non_blocking={shared.opts.diffusers_offload_nonblocking} modules_to_not_convert={modules_to_not_convert} modules_dtype_dict={modules_dtype_dict}') + svd = f'{shared.opts.sdnq_use_svd} rank={shared.opts.sdnq_svd_rank} steps={shared.opts.sdnq_svd_steps}' if shared.opts.sdnq_use_svd else f'{shared.opts.sdnq_use_svd}' + log.debug(f'Quantization: module="{module}" type=sdnq mode=pre dtype={weights_dtype} svd={svd} dynamic={shared.opts.sdnq_use_dynamic_quantization} group={shared.opts.sdnq_quantize_weights_group_size} loss={shared.opts.sdnq_dynamic_loss_threshold} matmul_dtype={quantized_matmul_dtype} matmul_quant={shared.opts.sdnq_use_quantized_matmul} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} quant_conv={shared.opts.sdnq_quantize_conv_layers} fp32={shared.opts.sdnq_dequantize_fp32} device={quantization_device} return={return_device} use_gpu={shared.opts.sdnq_quantize_with_gpu} map={shared.opts.device_map} offload={shared.opts.diffusers_offload_mode} non_blocking={shared.opts.diffusers_offload_nonblocking} skip_modules={modules_to_not_convert} dict={modules_dtype_dict}') if kwargs is None: return sdnq_config else: @@ -556,10 +559,12 @@ def sdnq_quantize_model(model, op=None, sd_model=None, do_gc: bool = True, weigh group_size=shared.opts.sdnq_quantize_weights_group_size, svd_rank=shared.opts.sdnq_svd_rank, svd_steps=shared.opts.sdnq_svd_steps, + dynamic_loss_threshold=shared.opts.sdnq_dynamic_loss_threshold, use_svd=shared.opts.sdnq_use_svd, quant_conv=shared.opts.sdnq_quantize_conv_layers, use_quantized_matmul=shared.opts.sdnq_use_quantized_matmul, use_quantized_matmul_conv=shared.opts.sdnq_use_quantized_matmul_conv, + use_dynamic_quantization=shared.opts.sdnq_use_dynamic_quantization, dequantize_fp32=shared.opts.sdnq_dequantize_fp32, non_blocking=shared.opts.diffusers_offload_nonblocking, quantization_device=quantization_device, @@ -594,7 +599,7 @@ def sdnq_quantize_model(model, op=None, sd_model=None, do_gc: bool = True, weigh if quantized_matmul_dtype is None: quantized_matmul_dtype = "auto" # set for logging - log.debug(f'Quantization: module="{op if op is not None else model.__class__}" type=sdnq mode=post dtype={weights_dtype} matmul_dtype={quantized_matmul_dtype} matmul={shared.opts.sdnq_use_quantized_matmul} svd={shared.opts.sdnq_use_svd}:group={shared.opts.sdnq_quantize_weights_group_size}:rank={shared.opts.sdnq_svd_rank}:steps={shared.opts.sdnq_svd_steps} quant_conv={shared.opts.sdnq_quantize_conv_layers} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} fp32={shared.opts.sdnq_dequantize_fp32} gpu={shared.opts.sdnq_quantize_with_gpu} device={quantization_device} return={return_device} map={shared.opts.device_map} non_blocking={shared.opts.diffusers_offload_nonblocking} modules_skip={modules_to_not_convert} modules_dtype={modules_dtype_dict}') + log.debug(f'Quantization: module="{op if op is not None else model.__class__}" type=sdnq mode=post dtype={weights_dtype} matmul_dtype={quantized_matmul_dtype} matmul={shared.opts.sdnq_use_quantized_matmul} svd={shared.opts.sdnq_use_svd}dynamic={shared.opts.sdnq_use_dynamic_quantization}:group={shared.opts.sdnq_quantize_weights_group_size}:rank={shared.opts.sdnq_svd_rank}:steps={shared.opts.sdnq_svd_steps}:loss={shared.opts.sdnq_dynamic_loss_threshold} quant_conv={shared.opts.sdnq_quantize_conv_layers} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} fp32={shared.opts.sdnq_dequantize_fp32} gpu={shared.opts.sdnq_quantize_with_gpu} device={quantization_device} return={return_device} map={shared.opts.device_map} non_blocking={shared.opts.diffusers_offload_nonblocking} modules_skip={modules_to_not_convert} modules_dtype={modules_dtype_dict}') return model diff --git a/modules/modeldata.py b/modules/modeldata.py index 29ab6b659..2fcfac26a 100644 --- a/modules/modeldata.py +++ b/modules/modeldata.py @@ -22,7 +22,7 @@ def get_model_type(pipe): model_type = 'sd' # instaflow is compatible with sd elif "AnimateDiffPipeline" in name: model_type = 'sd' # animatediff is compatible with sd - elif "Kandinsky5" in name: + elif "Kandinsky5" in name and '2I' in name: model_type = 'kandinsky5' elif "Kandinsky3" in name: model_type = 'kandinsky3' @@ -41,7 +41,7 @@ def get_model_type(pipe): elif "Flux" in name or "Flex1" in name or "Flex2" in name: model_type = 'f1' elif "ZImage" in name or "Z-Image" in name: - model_type = 'z_image' + model_type = 'zimage' elif "Lumina2" in name: model_type = 'lumina2' elif "Lumina" in name: @@ -78,9 +78,23 @@ def get_model_type(pipe): model_type = 'prx' elif 'LongCat' in name: model_type = 'longcat' + elif 'GlmImage' in name: + model_type = 'glmimage' elif 'Ovis-Image' in name: model_type = 'ovis' + elif 'Wan' in name: + model_type = 'wanai' + elif 'ChronoEdit' in name: + model_type = 'chrono' + elif 'HDM-xut' in name: + model_type = 'hdm' + elif 'HunyuanImage3' in name: + model_type = 'hunyuanimage3' + elif 'HunyuanImage' in name: + model_type = 'hunyuanimage' # video models + elif "Kandinsky5" in name and '2V' in name: + model_type = 'kandinsky5video' elif "CogVideo" in name: model_type = 'cogvideo' elif 'HunyuanVideo15' in name: @@ -93,17 +107,6 @@ def get_model_type(pipe): model_type = 'mochivideo' elif "Allegro" in name: model_type = 'allegrovideo' - # hybrid models - elif 'Wan' in name: - model_type = 'wanai' - elif 'ChronoEdit' in name: - model_type = 'chrono' - elif 'HDM-xut' in name: - model_type = 'hdm' - elif 'HunyuanImage3' in name: - model_type = 'hunyuanimage3' - elif 'HunyuanImage' in name: - model_type = 'hunyuanimage' # cloud models elif 'GoogleVeo' in name: model_type = 'veo3' diff --git a/modules/modelloader.py b/modules/modelloader.py index 9edf51a9d..26cb228c7 100644 --- a/modules/modelloader.py +++ b/modules/modelloader.py @@ -45,7 +45,7 @@ def hf_login(token=None): except Exception: pass with contextlib.redirect_stdout(stdout): - hf.login(token=token, add_to_git_credential=False, write_permission=False) + hf.login(token=token, add_to_git_credential=False) os.environ['HF_TOKEN'] = token text = stdout.getvalue() or '' obfuscated_token = 'hf_...' + token[-4:] @@ -274,13 +274,16 @@ def load_civitai(model: str, url: str): def download_url_to_file(url: str, dst: str): # based on torch.hub.download_url_to_file + import ssl import uuid import tempfile from urllib.request import urlopen, Request from rich.progress import Progress, TextColumn, BarColumn, TaskProgressColumn, TimeRemainingColumn, TimeElapsedColumn file_size = None req = Request(url, headers={"User-Agent": "sdnext"}) - u = urlopen(req) # pylint: disable=R1732 + + context = ssl._create_unverified_context() # pylint: disable=protected-access + u = urlopen(req, context=context) # pylint: disable=R1732 meta = u.info() if hasattr(meta, 'getheaders'): content_length = meta.getheaders("Content-Length") diff --git a/modules/onnx_impl/pipelines/onnx_stable_diffusion_upscale_pipeline.py b/modules/onnx_impl/pipelines/onnx_stable_diffusion_upscale_pipeline.py index 45365deb1..5bdc09794 100644 --- a/modules/onnx_impl/pipelines/onnx_stable_diffusion_upscale_pipeline.py +++ b/modules/onnx_impl/pipelines/onnx_stable_diffusion_upscale_pipeline.py @@ -128,7 +128,7 @@ class OnnxStableDiffusionUpscalePipeline(diffusers.OnnxStableDiffusionUpscalePip " `pipeline.unet` or your `image` input." ) - # 8. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 8. Prepare extra step kwargs. accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys()) extra_step_kwargs = {} if accepts_eta: diff --git a/modules/options.py b/modules/options.py index 6e31c8660..6b551385b 100644 --- a/modules/options.py +++ b/modules/options.py @@ -21,7 +21,7 @@ def options_section(section_identifier: tuple[str, str], options_dict: dict[str, class OptionInfo: def __init__( self, - default: Any | None = None, + default: Any = None, label="", component: type[Component] | type[DropdownEditable] | None = None, component_args: dict | Callable[..., dict] | None = None, diff --git a/modules/options_handler.py b/modules/options_handler.py index 6d7ae67c6..b087c0529 100644 --- a/modules/options_handler.py +++ b/modules/options_handler.py @@ -12,55 +12,54 @@ from installer import log if TYPE_CHECKING: from collections.abc import Callable from modules.options import OptionInfo + from typing import Any cmd_opts = cmd_args.parse_args() compatibility_opts = ['clip_skip', 'uni_pc_lower_order_final', 'uni_pc_order'] class Options(): - data = None - data_labels = None - filename = None + data_labels: dict[str, OptionInfo | LegacyOption] + data: dict[str, Any] typemap = {int: float} debug = os.environ.get('SD_CONFIG_DEBUG', None) is not None - def __init__(self, options_templates: dict[str, OptionInfo | LegacyOption] = {}, restricted_opts: set[str] | None = None): + def __init__(self, options_templates: dict[str, OptionInfo | LegacyOption] = {}, restricted_opts: set[str] | None = None, *, filename = ''): if restricted_opts is None: restricted_opts = set() - self.data_labels = options_templates + super().__setattr__('data_labels', options_templates) + super().__setattr__('data', {k: v.default for k, v in options_templates.items()}) + self.filename: str = filename or cmd_opts.config self.restricted_opts = restricted_opts - self.data = {k: v.default for k, v in self.data_labels.items()} - self.legacy = [k for k, v in self.data_labels.items() if isinstance(v, LegacyOption)] + self.legacy = [k for k, v in options_templates.items() if isinstance(v, LegacyOption)] + self.load() def __setattr__(self, key, value): # pylint: disable=inconsistent-return-statements - if self.data is not None: - if key in self.data or key in self.data_labels: - if cmd_opts.freeze: - log.warning(f'Settings are frozen: {key}') - return - if cmd_opts.hide_ui_dir_config and key in self.restricted_opts: - log.warning(f'Settings key is restricted: {key}') - return - if self.debug: - log.trace(f'Settings set: {key}={value}') - if key in self.legacy: - log.warning(f'Settings set: {key}={value} legacy') - self.data[key] = value + if key in self.data or key in self.data_labels: + if cmd_opts.freeze: + log.warning(f'Settings are frozen: {key}') return + if cmd_opts.hide_ui_dir_config and key in self.restricted_opts: + log.warning(f'Settings key is restricted: {key}') + return + if self.debug: + log.trace(f'Settings set: {key}={value}') + if key in self.legacy: + log.warning(f'Settings set: {key}={value} legacy') + self.data[key] = value + return return super(Options, self).__setattr__(key, value) # pylint: disable=super-with-arguments def get(self, item): - if self.data is not None: - if item in self.data: - return self.data[item] + if item in self.data: + return self.data[item] if item in self.data_labels: return self.data_labels[item].default return super(Options, self).__getattribute__(item) # pylint: disable=super-with-arguments def __getattr__(self, item): - if self.data is not None: - if item in self.data: - return self.data[item] + if item in self.data: + return self.data[item] if item in self.data_labels: return self.data_labels[item].default return super(Options, self).__getattribute__(item) # pylint: disable=super-with-arguments @@ -80,9 +79,10 @@ class Options(): setattr(self, key, value) except RuntimeError: return False - if self.data_labels[key].onchange is not None: + func = self.data_labels[key].onchange + if func is not None: try: - self.data_labels[key].onchange() + func() except Exception as err: log.error(f'Error in onchange callback: {key} {value} {err}') errors.display(err, 'Error in onchange callback') @@ -103,8 +103,6 @@ class Options(): def save_atomic(self, filename=None, silent=False): if self.debug: log.debug(f'Settings: save settings="{self.filename}" override="{filename}" cmd="{cmd_opts.config}" cwd="{os.getcwd()}"') - if self.filename is None: - self.filename = cmd_opts.config if filename is None: filename = self.filename filename = os.path.abspath(filename) @@ -170,10 +168,10 @@ class Options(): return self.data = readfile(filename, lock=True, as_type="dict") if self.data.get('quicksettings') is not None and self.data.get('quicksettings_list') is None: - self.data['quicksettings_list'] = [i.strip() for i in self.data.get('quicksettings').split(',')] + self.data['quicksettings_list'] = [i.strip() for i in self.data.get('quicksettings', '').split(',')] unknown_settings = [] for k, v in self.data.items(): - info: OptionInfo | None = self.data_labels.get(k, None) + info = self.data_labels.get(k, None) if info is not None: if not info.validate(k, v): self.data[k] = info.default diff --git a/modules/pag/pipe_sd.py b/modules/pag/pipe_sd.py index 11f4fb0cf..4393a24d3 100644 --- a/modules/pag/pipe_sd.py +++ b/modules/pag/pipe_sd.py @@ -1268,7 +1268,7 @@ class StableDiffusionPAGPipeline( latents, ) - # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 6. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 6.1 Add image embeds for IP-Adapter diff --git a/modules/pag/pipe_sdxl.py b/modules/pag/pipe_sdxl.py index 3a47af3e5..653fecf30 100644 --- a/modules/pag/pipe_sdxl.py +++ b/modules/pag/pipe_sdxl.py @@ -1366,7 +1366,7 @@ class StableDiffusionXLPAGPipeline( latents, ) - # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 6. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 7. Prepare added time ids & embeddings diff --git a/modules/paths.py b/modules/paths.py index 503f6eeb0..e3141814d 100644 --- a/modules/paths.py +++ b/modules/paths.py @@ -57,6 +57,23 @@ def create_path(folder): log.error(f'Create failed: folder="{folder}" {e}') +def resolve_output_path(base_path: str, specific_path: str) -> str: + """ + Resolve output path by combining base and specific paths. + + - If specific_path is absolute, return it directly (base is ignored) + - If base_path is set and specific_path is relative, join them + - If base_path is empty/None, return specific_path as-is + """ + if not specific_path: + return base_path or '' + if os.path.isabs(specific_path): + return specific_path + if base_path: + return os.path.normpath(os.path.join(base_path, specific_path)) + return specific_path + + def create_paths(opts): def fix_path(folder): tgt = None @@ -111,6 +128,22 @@ def create_paths(opts): create_path(fix_path('yolo_dir')) create_path(fix_path('wildcards_dir')) + # Create resolved output paths (base + specific) + base_samples = opts.data.get('outdir_samples', '') + base_grids = opts.data.get('outdir_grids', '') + if base_samples: + create_path(resolve_output_path(base_samples, opts.data.get('outdir_txt2img_samples', ''))) + create_path(resolve_output_path(base_samples, opts.data.get('outdir_img2img_samples', ''))) + create_path(resolve_output_path(base_samples, opts.data.get('outdir_control_samples', ''))) + create_path(resolve_output_path(base_samples, opts.data.get('outdir_extras_samples', ''))) + create_path(resolve_output_path(base_samples, opts.data.get('outdir_save', ''))) + create_path(resolve_output_path(base_samples, opts.data.get('outdir_video', ''))) + create_path(resolve_output_path(base_samples, opts.data.get('outdir_init_images', ''))) + if base_grids: + create_path(resolve_output_path(base_grids, opts.data.get('outdir_txt2img_grids', ''))) + create_path(resolve_output_path(base_grids, opts.data.get('outdir_img2img_grids', ''))) + create_path(resolve_output_path(base_grids, opts.data.get('outdir_control_grids', ''))) + class Prioritize: def __init__(self, name): diff --git a/modules/postprocess/yolo.py b/modules/postprocess/yolo.py index 9f171b19f..1d20dac35 100644 --- a/modules/postprocess/yolo.py +++ b/modules/postprocess/yolo.py @@ -421,6 +421,7 @@ class YoloRestorer(Detailer): pc.init_images = [image] pc.image_mask = [item.mask] pc.overlay_images = [] + pc.enable_hr = False # explictly disable hires for detailer pass pc.recursion = True # process @@ -490,7 +491,7 @@ class YoloRestorer(Detailer): shared.opts.detailer_sort = sort shared.opts.detailer_seg = seg # shared.opts.detailer_resolution = resolution - shared.opts.save(shared.config_filename, silent=True) + shared.opts.save(silent=True) shared.log.debug(f'Detailer settings: models={detailers} classes={classes} strength={strength} conf={min_confidence} max={max_detected} iou={iou} size={min_size}-{max_size} padding={padding} steps={steps} resolution={resolution} save={save} sort={sort} seg={seg}') if not self.ui_mode: shared.log.debug(f'Detailer expert: {text}') @@ -500,7 +501,7 @@ class YoloRestorer(Detailer): enabled = gr.Checkbox(label="Enable detailer pass", elem_id=f"{tab}_detailer_enabled", value=False) with gr.Row(): seg = gr.Checkbox(label="Use segmentation", elem_id=f"{tab}_detailer_seg", value=shared.opts.detailer_seg, visible=True) - save = gr.Checkbox(label="Include detection results", elem_id=f"{tab}_detailer_save", value=shared.opts.detailer_save, visible=True) + save = gr.Checkbox(label="Include detections", elem_id=f"{tab}_detailer_save", value=shared.opts.detailer_save, visible=True) with gr.Row(): merge = gr.Checkbox(label="Merge detailers", elem_id=f"{tab}_detailer_merge", value=shared.opts.detailer_merge, visible=True) sort = gr.Checkbox(label="Sort detections", elem_id=f"{tab}_detailer_sort", value=shared.opts.detailer_sort, visible=True) diff --git a/modules/postprocessing.py b/modules/postprocessing.py index 296db6eda..1f04905c5 100644 --- a/modules/postprocessing.py +++ b/modules/postprocessing.py @@ -6,6 +6,7 @@ from PIL import Image from modules import shared, images, devices, scripts_manager, scripts_postprocessing, infotext from modules.shared import opts +from modules.paths import resolve_output_path def run_postprocessing(extras_mode, image, image_folder: List[tempfile.NamedTemporaryFile], input_dir, output_dir, show_extras_results, *args, save_output: bool = True): @@ -58,7 +59,7 @@ def run_postprocessing(extras_mode, image, image_folder: List[tempfile.NamedTemp if extras_mode == 2 and output_dir != '': outpath = output_dir else: - outpath = opts.outdir_samples or opts.outdir_extras_samples + outpath = resolve_output_path(opts.outdir_samples, opts.outdir_extras_samples) processed_images = [] for image, name, ext in zip(image_data, image_names, image_ext): # pylint: disable=redefined-argument-from-local shared.log.debug(f'Process: image={image} {args}') diff --git a/modules/processing.py b/modules/processing.py index 79e62d669..523915942 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -31,7 +31,7 @@ processed = None # last known processed results class Processed: - def __init__(self, p: StableDiffusionProcessing, images_list, seed=-1, info=None, subseed=None, all_prompts=None, all_negative_prompts=None, all_seeds=None, all_subseeds=None, index_of_first_image=0, infotexts=None, comments="", binary=None): + def __init__(self, p: StableDiffusionProcessing, images_list, seed=-1, info=None, subseed=None, all_prompts=None, all_negative_prompts=None, all_seeds=None, all_subseeds=None, index_of_first_image=0, infotexts=None, comments="", binary=None, audio=None): self.sd_model_hash = getattr(shared.sd_model, 'sd_model_hash', '') if model_data.sd_model is not None else '' self.prompt = p.prompt or '' @@ -53,6 +53,8 @@ class Processed: self.batch_size = max(1, p.batch_size) self.denoising_strength = p.denoising_strength + self.audio = audio + self.restore_faces = p.restore_faces or False self.face_restoration_model = shared.opts.face_restoration_model if p.restore_faces else None self.detailer = p.detailer_enabled or False @@ -323,9 +325,6 @@ def process_samples(p: StableDiffusionProcessing, samples): if p.color_corrections is not None and i < len(p.color_corrections): p.ops.append('color') if not p.do_not_save_samples and shared.opts.save_images_before_color_correction: - orig = p.color_corrections - p.color_corrections = None - p.color_corrections = orig image_without_cc = apply_overlay(image, p.paste_to, i, p.overlay_images) info = create_infotext(p, p.prompts, p.seeds, p.subseeds, index=i) images.save_image(image_without_cc, path=p.outpath_samples, basename="", seed=p.seeds[i], prompt=p.prompts[i], extension=shared.opts.samples_format, info=info, p=p, suffix="-before-color-correct") @@ -398,6 +397,7 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed: infotexts = [] output_images = [] output_binary = None + audio = None process_init(p) if p.scripts is not None and isinstance(p.scripts, scripts_manager.ScriptRunner): @@ -484,6 +484,8 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed: output_images.append(batch_image) infotexts.append(batch_infotext) + audio = getattr(samples, 'audio', None) + if shared.cmd_opts.lowvram: devices.torch_gc(force=True, reason='lowvram') timer.process.record('post') @@ -522,6 +524,7 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed: subseed=p.all_subseeds[0], index_of_first_image=index_of_first_image, infotexts=infotexts, + audio=audio, ) if p.scripts is not None and isinstance(p.scripts, scripts_manager.ScriptRunner) and not (shared.state.interrupted or shared.state.skipped): p.scripts.postprocess(p, results) diff --git a/modules/processing_args.py b/modules/processing_args.py index 869b1309e..ebd75d751 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -7,9 +7,10 @@ import inspect import torch import numpy as np from PIL import Image -from modules import shared, errors, sd_models, processing, processing_vae, processing_helpers, sd_hijack_hypertile, prompt_parser_diffusers, timer, extra_networks, sd_vae +from modules import shared, sd_models, processing, processing_vae, processing_helpers, sd_hijack_hypertile, extra_networks, sd_vae from modules.processing_callbacks import diffusers_callback_legacy, diffusers_callback, set_callbacks_p -from modules.processing_helpers import resize_hires, fix_prompts, calculate_base_steps, calculate_hires_steps, calculate_refiner_steps, get_generator, set_latents, apply_circular # pylint: disable=unused-import +from modules.processing_helpers import resize_hires, calculate_base_steps, calculate_hires_steps, calculate_refiner_steps, get_generator, set_latents, apply_circular # pylint: disable=unused-import +from modules.processing_prompt import set_prompt from modules.api import helpers @@ -49,12 +50,13 @@ def task_specific_kwargs(p, model): p.init_images = [helpers.decode_base64_to_image(i, quiet=True) for i in p.init_images] if isinstance(p.init_images[0], Image.Image): p.init_images = [i.convert('RGB') if i.mode != 'RGB' else i for i in p.init_images if i is not None] + width, height = processing_helpers.resize_init_images(p) if (task_type == sd_models.DiffusersTaskType.TEXT_2_IMAGE or len(getattr(p, 'init_images', [])) == 0) and not is_img2img_model and 'video' not in p.ops: p.ops.append('txt2img') if hasattr(p, 'width') and hasattr(p, 'height'): task_args = { - 'width': vae_scale_factor * math.ceil(p.width / vae_scale_factor), - 'height': vae_scale_factor * math.ceil(p.height / vae_scale_factor), + 'width': width, + 'height': height, } elif (task_type == sd_models.DiffusersTaskType.IMAGE_2_IMAGE or is_img2img_model) and len(getattr(p, 'init_images', [])) > 0: if shared.sd_model_type == 'sdxl' and hasattr(model, 'register_to_config'): @@ -73,18 +75,23 @@ def task_specific_kwargs(p, model): } if model_cls == 'FluxImg2ImgPipeline' or model_cls == 'FluxKontextPipeline': # needs explicit width/height if torch.is_tensor(p.init_images[0]): - p.width, p.height = p.init_images[0].shape[-1] * vae_scale_factor, p.init_images[0].shape[-2] * vae_scale_factor + p.width = p.init_images[0].shape[-1] * vae_scale_factor + p.height = p.init_images[0].shape[-2] * vae_scale_factor else: - p.width, p.height = 8 * math.ceil(p.init_images[0].width / vae_scale_factor), 8 * math.ceil(p.init_images[0].height / vae_scale_factor) + p.width = width + p.height = height if model_cls == 'FluxKontextPipeline': aspect_ratio = p.width / p.height max_area = max(p.width, p.height)**2 - p.width, p.height = round((max_area * aspect_ratio) ** 0.5), round((max_area / aspect_ratio) ** 0.5) - p.width, p.height = p.width // vae_scale_factor * vae_scale_factor, p.height // vae_scale_factor * vae_scale_factor + p.width = round((max_area * aspect_ratio) ** 0.5) + p.height = round((max_area / aspect_ratio) ** 0.5) + p.width = p.width // vae_scale_factor * vae_scale_factor + p.height = p.height // vae_scale_factor * vae_scale_factor task_args['max_area'] = max_area task_args['width'], task_args['height'] = p.width, p.height elif model_cls == 'OmniGenPipeline' or model_cls == 'OmniGen2Pipeline': - p.width, p.height = vae_scale_factor * math.ceil(p.init_images[0].width / vae_scale_factor), vae_scale_factor * math.ceil(p.init_images[0].height / vae_scale_factor) + p.width = width + p.height = height task_args = { 'width': p.width, 'height': p.height, @@ -93,8 +100,8 @@ def task_specific_kwargs(p, model): elif task_type == sd_models.DiffusersTaskType.INSTRUCT and len(getattr(p, 'init_images', [])) > 0: p.ops.append('instruct') task_args = { - 'width': vae_scale_factor * math.ceil(p.width / vae_scale_factor) if hasattr(p, 'width') else None, - 'height': vae_scale_factor * math.ceil(p.height / vae_scale_factor) if hasattr(p, 'height') else None, + 'width': width if hasattr(p, 'width') else None, + 'height': height if hasattr(p, 'height') else None, 'image': p.init_images, 'strength': p.denoising_strength, } @@ -108,7 +115,6 @@ def task_specific_kwargs(p, model): p.ops.append('detailer') else: p.ops.append('inpaint') - width, height = processing_helpers.resize_init_images(p) mask_image = p.task_args.get('image_mask', None) or getattr(p, 'image_mask', None) or getattr(p, 'mask', None) if p.vae_type == 'Remote': from modules.sd_vae_remote import remote_encode @@ -151,6 +157,8 @@ def task_specific_kwargs(p, model): task_args['reference_images'] = p.init_images if ('GoogleNanoBananaPipeline' in model_cls) and (p.init_images is not None) and (len(p.init_images) > 0): task_args['image'] = p.init_images[0] + if ('GlmImagePipeline' in model_cls) and (p.init_images is not None) and (len(p.init_images) > 0): + task_args['image'] = p.init_images if 'BlipDiffusionPipeline' in model_cls: if len(p.init_images) == 0: shared.log.error('BLiP diffusion requires init image') @@ -184,6 +192,7 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t apply_circular(p.tiling, model) args = {} has_vae = hasattr(model, 'vae') or (hasattr(model, 'pipe') and hasattr(model.pipe, 'vae')) + cls = model.__class__.__name__ if hasattr(model, 'pipe') and not hasattr(model, 'no_recurse'): # recurse model = model.pipe has_vae = has_vae or hasattr(model, 'vae') @@ -197,87 +206,19 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t if debug_enabled: debug_log(f'Process pipeline possible: {possible}') - prompts, negative_prompts, prompts_2, negative_prompts_2 = fix_prompts(p, prompts, negative_prompts, prompts_2, negative_prompts_2) steps = kwargs.get("num_inference_steps", None) or len(getattr(p, 'timesteps', ['1'])) clip_skip = kwargs.pop("clip_skip", 1) - extra_networks.activate(p, include=['text_encoder', 'text_encoder_2', 'text_encoder_3']) + prompt_attention, args = set_prompt(p, args, possible, cls, prompt_attention, steps, clip_skip, prompts, negative_prompts, prompts_2, negative_prompts_2) - parser = 'fixed' - prompt_attention = prompt_attention or shared.opts.prompt_attention - if (prompt_attention != 'fixed') and ('Onnx' not in model.__class__.__name__) and ('prompt' not in p.task_args) and ( - 'StableDiffusion' in model.__class__.__name__ or - 'StableCascade' in model.__class__.__name__ or - ('Flux' in model.__class__.__name__ and 'Flux2' not in model.__class__.__name__) or - 'Chroma' in model.__class__.__name__ or - 'HiDreamImagePipeline' in model.__class__.__name__ - ): - jobid = shared.state.begin('TE Encode') - try: - prompt_parser_diffusers.embedder = prompt_parser_diffusers.PromptEmbedder(prompts, negative_prompts, steps, clip_skip, p) - parser = shared.opts.prompt_attention - except Exception as e: - shared.log.error(f'Prompt parser encode: {e}') - if os.environ.get('SD_PROMPT_DEBUG', None) is not None: - errors.display(e, 'Prompt parser encode') - timer.process.record('prompt', reset=False) - shared.state.end(jobid) - else: - prompt_parser_diffusers.embedder = None + if 'clip_skip' in possible: + if clip_skip == 1: + pass # clip_skip = None + else: + args['clip_skip'] = clip_skip - 1 - if 'prompt' in possible: - if 'OmniGen' in model.__class__.__name__: - prompts = [p.replace('|image|', '<|image_1|>') for p in prompts] - if ('HiDreamImage' in model.__class__.__name__) and (prompt_parser_diffusers.embedder is not None): - args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds') - prompt_embeds = prompt_parser_diffusers.embedder('prompt_embeds') - args['prompt_embeds_t5'] = prompt_embeds[0] - args['prompt_embeds_llama3'] = prompt_embeds[1] - elif hasattr(model, 'text_encoder') and hasattr(model, 'tokenizer') and ('prompt_embeds' in possible) and (prompt_parser_diffusers.embedder is not None): - embeds = prompt_parser_diffusers.embedder('prompt_embeds') - if embeds is None: - shared.log.warning('Prompt parser encode: empty prompt embeds') - prompt_parser_diffusers.embedder = None - args['prompt'] = prompts - elif embeds.device == torch.device('meta'): - shared.log.warning('Prompt parser encode: embeds on meta device') - prompt_parser_diffusers.embedder = None - args['prompt'] = prompts - else: - args['prompt_embeds'] = embeds - if 'StableCascade' in model.__class__.__name__: - args['prompt_embeds_pooled'] = prompt_parser_diffusers.embedder('positive_pooleds').unsqueeze(0) - elif 'XL' in model.__class__.__name__: - args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds') - elif 'StableDiffusion3' in model.__class__.__name__: - args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds') - elif 'Flux' in model.__class__.__name__: - args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds') - elif 'Chroma' in model.__class__.__name__: - args['prompt_attention_mask'] = prompt_parser_diffusers.embedder('prompt_attention_masks') - else: - args['prompt'] = prompts - if 'negative_prompt' in possible: - if 'HiDreamImage' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None: - args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds') - negative_prompt_embeds = prompt_parser_diffusers.embedder('negative_prompt_embeds') - args['negative_prompt_embeds_t5'] = negative_prompt_embeds[0] - args['negative_prompt_embeds_llama3'] = negative_prompt_embeds[1] - elif hasattr(model, 'text_encoder') and hasattr(model, 'tokenizer') and 'negative_prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None: - args['negative_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_prompt_embeds') - if 'StableCascade' in model.__class__.__name__: - args['negative_prompt_embeds_pooled'] = prompt_parser_diffusers.embedder('negative_pooleds').unsqueeze(0) - elif 'XL' in model.__class__.__name__: - args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds') - elif 'StableDiffusion3' in model.__class__.__name__: - args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds') - elif 'Chroma' in model.__class__.__name__: - args['negative_prompt_attention_mask'] = prompt_parser_diffusers.embedder('negative_prompt_attention_masks') - else: - if 'PixArtSigmaPipeline' in model.__class__.__name__: # pixart-sigma pipeline throws list-of-list for negative prompt - args['negative_prompt'] = negative_prompts[0] - else: - args['negative_prompt'] = negative_prompts + if shared.opts.lora_apply_te: + extra_networks.activate(p, include=['text_encoder', 'text_encoder_2', 'text_encoder_3']) if 'complex_human_instruction' in possible: chi = shared.opts.te_complex_human_instruction @@ -288,14 +229,6 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t args['use_resolution_binning'] = False if 'use_mask_in_transformer' in possible: args['use_mask_in_transformer'] = shared.opts.te_use_mask - if prompt_parser_diffusers.embedder is not None and not prompt_parser_diffusers.embedder.scheduled_prompt: # not scheduled so we dont need it anymore - prompt_parser_diffusers.embedder = None - - if 'clip_skip' in possible and parser == 'fixed': - if clip_skip == 1: - pass # clip_skip = None - else: - args['clip_skip'] = clip_skip - 1 timesteps = re.split(',| ', shared.opts.schedulers_timesteps) if len(timesteps) > 2: @@ -482,7 +415,7 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t clean['negative_prompt'] = len(clean['negative_prompt']) if generator is not None: clean['generator'] = f'{generator[0].device}:{[g.initial_seed() for g in generator]}' - clean['parser'] = parser + clean['parser'] = prompt_attention for k, v in clean.copy().items(): if v is None: clean[k] = None diff --git a/modules/processing_callbacks.py b/modules/processing_callbacks.py index d62d464dc..eed7f985f 100644 --- a/modules/processing_callbacks.py +++ b/modules/processing_callbacks.py @@ -116,18 +116,48 @@ def diffusers_callback(pipe, step: int = 0, timestep: int = 0, kwargs: dict = {} if current_noise_pred is None: current_noise_pred = kwargs.get("predicted_image_embedding", None) - if hasattr(pipe, "_unpack_latents") and hasattr(pipe, "vae_scale_factor"): # FLUX + if hasattr(pipe, "_unpack_latents") and hasattr(pipe, "vae_scale_factor"): # FLUX.1 if p.hr_resize_mode > 0 and (p.hr_upscaler != 'None' or p.hr_resize_mode == 5) and p.is_hr_pass: width = max(getattr(p, 'width', 0), getattr(p, 'hr_upscale_to_x', 0)) height = max(getattr(p, 'height', 0), getattr(p, 'hr_upscale_to_y', 0)) else: - width = getattr(p, 'width', 0) - height = getattr(p, 'height', 0) + width = getattr(p, 'width', 1024) + height = getattr(p, 'height', 1024) shared.state.current_latent = pipe._unpack_latents(kwargs['latents'], height, width, pipe.vae_scale_factor) # pylint: disable=protected-access if current_noise_pred is not None: shared.state.current_noise_pred = pipe._unpack_latents(current_noise_pred, height, width, pipe.vae_scale_factor) # pylint: disable=protected-access else: shared.state.current_noise_pred = current_noise_pred + elif hasattr(pipe, "_unpatchify_latents"): # FLUX.2 - unpack [B, seq, patch_ch] to [B, ch, H, W] + vae_scale = getattr(pipe, 'vae_scale_factor', 8) + if p.hr_resize_mode > 0 and (p.hr_upscaler != 'None' or p.hr_resize_mode == 5) and p.is_hr_pass: + width = max(getattr(p, 'width', 0), getattr(p, 'hr_upscale_to_x', 0)) + height = max(getattr(p, 'height', 0), getattr(p, 'hr_upscale_to_y', 0)) + else: + width = getattr(p, 'width', 1024) + height = getattr(p, 'height', 1024) + latents = kwargs['latents'] + if len(latents.shape) == 3: # packed format [B, seq_len, patch_channels] + b, seq_len, patch_ch = latents.shape + channels = patch_ch // 4 # 4 = 2x2 patch + h_patches = height // vae_scale // 2 + w_patches = width // vae_scale // 2 + if h_patches * w_patches != seq_len: # fallback to square assumption + h_patches = w_patches = int(seq_len ** 0.5) + # [B, h*w, C*4] -> [B, h, w, C, 2, 2] -> [B, C, h, 2, w, 2] -> [B, C, H, W] + latents = latents.view(b, h_patches, w_patches, channels, 2, 2) + latents = latents.permute(0, 3, 1, 4, 2, 5).reshape(b, channels, h_patches * 2, w_patches * 2) + shared.state.current_latent = latents + if current_noise_pred is not None and len(current_noise_pred.shape) == 3: + b, seq_len, patch_ch = current_noise_pred.shape + channels = patch_ch // 4 + h_patches = height // vae_scale // 2 + w_patches = width // vae_scale // 2 + if h_patches * w_patches != seq_len: + h_patches = w_patches = int(seq_len ** 0.5) + current_noise_pred = current_noise_pred.view(b, h_patches, w_patches, channels, 2, 2) + current_noise_pred = current_noise_pred.permute(0, 3, 1, 4, 2, 5).reshape(b, channels, h_patches * 2, w_patches * 2) + shared.state.current_noise_pred = current_noise_pred else: shared.state.current_latent = kwargs['latents'] shared.state.current_noise_pred = current_noise_pred diff --git a/modules/processing_class.py b/modules/processing_class.py index 836af33b8..d4305d52f 100644 --- a/modules/processing_class.py +++ b/modules/processing_class.py @@ -7,6 +7,7 @@ from dataclasses import dataclass, field import numpy as np from PIL import Image, ImageOps from modules import shared, images, scripts_manager, masking, sd_models, sd_vae, processing_helpers +from modules.paths import resolve_output_path debug = shared.log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -312,6 +313,8 @@ class StableDiffusionProcessing: self.negative_prompts = None self.all_prompts = None self.all_negative_prompts = None + self.seeds = [] + self.subseeds = [] self.all_seeds = None self.all_subseeds = None @@ -537,7 +540,7 @@ class StableDiffusionProcessingImg2Img(StableDiffusionProcessing): self.init_img_width = getattr(self, 'init_img_width', img.width) # pylint: disable=attribute-defined-outside-init self.init_img_height = getattr(self, 'init_img_height', img.height) # pylint: disable=attribute-defined-outside-init if shared.opts.save_init_img: - images.save_image(img, path=shared.opts.outdir_init_images, basename=None, forced_filename=self.init_img_hash, suffix="-init-image") + images.save_image(img, path=resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_init_images), basename=None, forced_filename=self.init_img_hash, suffix="-init-image") image = images.flatten(img, shared.opts.img2img_background_color) if crop_region is None and self.resize_mode > 0: image = images.resize_image(self.resize_mode, image, self.width, self.height, upscaler_name=self.resize_name, context=self.resize_context) @@ -592,13 +595,18 @@ class StableDiffusionProcessingControl(StableDiffusionProcessingImg2Img): def switch_class(p: StableDiffusionProcessing, new_class: type, dct: dict = None): - signature = inspect.signature(type(new_class).__init__, follow_wrapped=True) - possible = list(signature.parameters) kwargs = {} + signature = inspect.signature(StableDiffusionProcessing.__init__, follow_wrapped=True) # base class + possible = list(signature.parameters) for k, v in p.__dict__.copy().items(): if k in possible: kwargs[k] = v - if dct is not None: + signature = inspect.signature(type(new_class).__init__, follow_wrapped=True) # target class + possible = list(signature.parameters) + for k, v in p.__dict__.copy().items(): + if k in possible: + kwargs[k] = v + if dct is not None: # overrides for k, v in dct.items(): if k in possible: kwargs[k] = v diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index d727a3363..269351120 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -13,6 +13,7 @@ from modules.lora import lora_common debug = os.environ.get('SD_DIFFUSERS_DEBUG', None) is not None +output_type = 'np' if os.environ.get('SD_VAE_DEFAULT', None) is not None else 'latent' last_p = None orig_pipeline = shared.sd_model @@ -157,7 +158,7 @@ def process_base(p: processing.StableDiffusionProcessing): denoising_start=0 if use_refiner_start else p.refiner_start if use_denoise_start else None, denoising_end=p.refiner_start if use_refiner_start else 1 if use_denoise_start else None, num_frames=getattr(p, 'frames', 1), - output_type='latent', + output_type=output_type, clip_skip=p.clip_skip, desc=desc, ) @@ -307,7 +308,7 @@ def process_hires(p: processing.StableDiffusionProcessing, output): eta=shared.opts.scheduler_eta, guidance_scale=p.image_cfg_scale if p.image_cfg_scale is not None else p.cfg_scale, guidance_rescale=p.diffusers_guidance_rescale, - output_type='latent', + output_type=output_type, clip_skip=p.clip_skip, image=output.images, strength=strength, @@ -377,11 +378,11 @@ def process_refine(p: processing.StableDiffusionProcessing, output): for i in range(len(output.images)): image = output.images[i] noise_level = round(350 * p.denoising_strength) - output_type = 'latent' + refiner_output_type = output_type if 'Upscale' in shared.sd_refiner.__class__.__name__ or 'Flux' in shared.sd_refiner.__class__.__name__ or 'Kandinsky' in shared.sd_refiner.__class__.__name__: image = processing_vae.vae_decode(latents=image, model=shared.sd_model, vae_type=p.vae_type, output_type='pil', width=p.width, height=p.height) p.extra_generation_params['Noise level'] = noise_level - output_type = 'np' + refiner_output_type = 'np' update_sampler(p, shared.sd_refiner, second_pass=True) shared.opts.prompt_attention = 'fixed' refiner_args = set_pipeline_args( @@ -398,7 +399,7 @@ def process_refine(p: processing.StableDiffusionProcessing, output): denoising_start=p.refiner_start if p.refiner_start > 0 and p.refiner_start < 1 else None, denoising_end=1 if p.refiner_start > 0 and p.refiner_start < 1 else None, image=image, - output_type=output_type, + output_type=refiner_output_type, clip_skip=p.clip_skip, prompt_attention='fixed', desc='Refiner', diff --git a/modules/processing_helpers.py b/modules/processing_helpers.py index 19f8352c6..222409738 100644 --- a/modules/processing_helpers.py +++ b/modules/processing_helpers.py @@ -236,6 +236,10 @@ def get_fixed_seed(seed): def fix_seed(p): p.seed = get_fixed_seed(p.seed) p.subseed = get_fixed_seed(p.subseed) + for i in range(len(p.all_seeds)): + p.all_seeds[i] = get_fixed_seed(p.all_seeds[i]) + for i in range(len(p.all_subseeds)): + p.all_subseeds[i] = get_fixed_seed(p.all_subseeds[i]) def old_hires_fix_first_pass_dimensions(width, height): @@ -299,37 +303,43 @@ def decode_images(image): return helpers.decode_base64_to_image(image, quiet=True) except Exception as e: shared.log.error(f'Decode image: {e}') - elif isinstance(image, Image.Image): - return image + # elif isinstance(image, Image.Image): + # return image + # elif torch.is_tensor(image): + # return image else: - shared.log.error(f'Decode image: {type(image)} unknown type') + return image + # shared.log.error(f'Decode image: {type(image)} unknown type') return None def resize_init_images(p): - if getattr(p, 'image', None) is not None and getattr(p, 'init_images', None) is None: - p.init_images = [p.image] - - if getattr(p, 'init_images', None) is not None and len(p.init_images) > 0: - p.init_images = decode_images(p.init_images) - vae_scale_factor = sd_vae.get_vae_scale_factor() - tgt_width, tgt_height = vae_scale_factor * math.ceil(p.init_images[0].width / vae_scale_factor), vae_scale_factor * math.ceil(p.init_images[0].height / vae_scale_factor) - if p.init_images[0].size != (tgt_width, tgt_height): - shared.log.debug(f'Resizing init images: original={p.init_images[0].width}x{p.init_images[0].height} target={tgt_width}x{tgt_height}') - p.init_images = [images.resize_image(1, image, tgt_width, tgt_height, upscaler_name=None) for image in p.init_images] - p.height = tgt_height - p.width = tgt_width - sd_hijack_hypertile.hypertile_set(p) - if getattr(p, 'mask', None) is not None and p.mask is not None and p.mask.size != (tgt_width, tgt_height): - p.mask = decode_images(p.mask) - p.mask = images.resize_image(1, p.mask, tgt_width, tgt_height, upscaler_name=None) - if getattr(p, 'init_mask', None) is not None and p.init_mask is not None and p.init_mask.size != (tgt_width, tgt_height): - p.init_mask = decode_images(p.init_mask) - p.init_mask = images.resize_image(1, p.init_mask, tgt_width, tgt_height, upscaler_name=None) - if getattr(p, 'mask_for_overlay', None) is not None and p.mask_for_overlay is not None and p.mask_for_overlay.size != (tgt_width, tgt_height): - p.mask_for_overlay = decode_images(p.mask_for_overlay) - p.mask_for_overlay = images.resize_image(1, p.mask_for_overlay, tgt_width, tgt_height, upscaler_name=None) - return tgt_width, tgt_height + try: + if getattr(p, 'image', None) is not None and getattr(p, 'init_images', None) is None: + p.init_images = [p.image] + if getattr(p, 'init_images', None) is not None and len(p.init_images) > 0: + p.init_images = decode_images(p.init_images) + vae_scale_factor = sd_vae.get_vae_scale_factor() + tgt_width = vae_scale_factor * math.ceil(p.init_images[0].width / vae_scale_factor) + tgt_height = vae_scale_factor * math.ceil(p.init_images[0].height / vae_scale_factor) + if p.init_images[0].size != (tgt_width, tgt_height): + shared.log.debug(f'Resizing init images: original={p.init_images[0].width}x{p.init_images[0].height} target={tgt_width}x{tgt_height}') + p.init_images = [images.resize_image(1, image, tgt_width, tgt_height, upscaler_name=None) for image in p.init_images] + p.height = tgt_height + p.width = tgt_width + sd_hijack_hypertile.hypertile_set(p) + if getattr(p, 'mask', None) is not None and p.mask is not None and p.mask.size != (tgt_width, tgt_height): + p.mask = decode_images(p.mask) + p.mask = images.resize_image(1, p.mask, tgt_width, tgt_height, upscaler_name=None) + if getattr(p, 'init_mask', None) is not None and p.init_mask is not None and p.init_mask.size != (tgt_width, tgt_height): + p.init_mask = decode_images(p.init_mask) + p.init_mask = images.resize_image(1, p.init_mask, tgt_width, tgt_height, upscaler_name=None) + if getattr(p, 'mask_for_overlay', None) is not None and p.mask_for_overlay is not None and p.mask_for_overlay.size != (tgt_width, tgt_height): + p.mask_for_overlay = decode_images(p.mask_for_overlay) + p.mask_for_overlay = images.resize_image(1, p.mask_for_overlay, tgt_width, tgt_height, upscaler_name=None) + return tgt_width, tgt_height + except Exception: + pass return p.width, p.height @@ -366,44 +376,6 @@ def resize_hires(p, latents): # input=latents output=pil if not latent_upscaler return resized -def fix_prompts(p, prompts, negative_prompts, prompts_2, negative_prompts_2): - if hasattr(p, 'keep_prompts'): - return prompts, negative_prompts, prompts_2, negative_prompts_2 - - if type(prompts) is str: - prompts = [prompts] - if type(negative_prompts) is str: - negative_prompts = [negative_prompts] - - if hasattr(p, '[init_images]') and p.init_images is not None and len(p.init_images) > 1: - while len(prompts) < len(p.init_images): - prompts.append(prompts[-1]) - while len(negative_prompts) < len(p.init_images): - negative_prompts.append(negative_prompts[-1]) - - while len(prompts) < p.batch_size: - prompts.append(prompts[-1]) - while len(negative_prompts) < p.batch_size: - negative_prompts.append(negative_prompts[-1]) - - while len(negative_prompts) < len(prompts): - negative_prompts.append(negative_prompts[-1]) - while len(prompts) < len(negative_prompts): - prompts.append(prompts[-1]) - - if type(prompts_2) is str: - prompts_2 = [prompts_2] - if type(prompts_2) is list: - while len(prompts_2) < len(prompts): - prompts_2.append(prompts_2[-1]) - if type(negative_prompts_2) is str: - negative_prompts_2 = [negative_prompts_2] - if type(negative_prompts_2) is list: - while len(negative_prompts_2) < len(prompts_2): - negative_prompts_2.append(negative_prompts_2[-1]) - return prompts, negative_prompts, prompts_2, negative_prompts_2 - - def calculate_base_steps(p, use_denoise_start, use_refiner_start): if len(getattr(p, 'timesteps', [])) > 0: return None diff --git a/modules/processing_info.py b/modules/processing_info.py index c2a6566ec..3c4e06e40 100644 --- a/modules/processing_info.py +++ b/modules/processing_info.py @@ -153,7 +153,7 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No args["Color correction"] = True if shared.opts.token_merging_method == 'ToMe': # tome/todo args['ToMe'] = shared.opts.tome_ratio if shared.opts.tome_ratio != 0 else None - else: + elif shared.opts.token_merging_method == 'ToDo': args['ToDo'] = shared.opts.todo_ratio if shared.opts.todo_ratio != 0 else None if hasattr(shared.sd_model, 'embedding_db') and len(shared.sd_model.embedding_db.embeddings_used) > 0: # register used embeddings args['Embeddings'] = ', '.join(shared.sd_model.embedding_db.embeddings_used) diff --git a/modules/processing_prompt.py b/modules/processing_prompt.py new file mode 100644 index 000000000..81b713934 --- /dev/null +++ b/modules/processing_prompt.py @@ -0,0 +1,182 @@ +import os +import torch +from modules import shared, errors, timer, prompt_parser_diffusers + + +debug_enabled = os.environ.get('SD_PROMPT_DEBUG', None) is not None +debug_log = shared.log.trace if debug_enabled else lambda *args, **kwargs: None + + +def fix_prompt_batch(p, prompts, negative_prompts, prompts_2, negative_prompts_2): + if hasattr(p, 'keep_prompts'): + return prompts, negative_prompts, prompts_2, negative_prompts_2 + + if type(prompts) is str: + prompts = [prompts] + if type(negative_prompts) is str: + negative_prompts = [negative_prompts] + + if hasattr(p, '[init_images]') and p.init_images is not None and len(p.init_images) > 1: + while len(prompts) < len(p.init_images): + prompts.append(prompts[-1]) + while len(negative_prompts) < len(p.init_images): + negative_prompts.append(negative_prompts[-1]) + + while len(prompts) < p.batch_size: + prompts.append(prompts[-1]) + while len(negative_prompts) < p.batch_size: + negative_prompts.append(negative_prompts[-1]) + + while len(negative_prompts) < len(prompts): + negative_prompts.append(negative_prompts[-1]) + while len(prompts) < len(negative_prompts): + prompts.append(prompts[-1]) + + if type(prompts_2) is str: + prompts_2 = [prompts_2] + if type(prompts_2) is list: + while len(prompts_2) < len(prompts): + prompts_2.append(prompts_2[-1]) + if type(negative_prompts_2) is str: + negative_prompts_2 = [negative_prompts_2] + if type(negative_prompts_2) is list: + while len(negative_prompts_2) < len(prompts_2): + negative_prompts_2.append(negative_prompts_2[-1]) + return prompts, negative_prompts, prompts_2, negative_prompts_2 + + +def fix_prompt_model(cls, prompts, negative_prompts, prompts_2, negative_prompts_2): + if 'OmniGen' in cls: + prompts = [p.replace('|image|', '<|image_1|>') for p in prompts] + if 'PixArtSigmaPipeline' in cls: # pixart-sigma pipeline throws list-of-list for negative prompt + negative_prompts = negative_prompts[0] + return prompts, negative_prompts, prompts_2, negative_prompts_2 + + +def set_fallback_prompt(args: dict, possible: list[str], prompts, negative_prompts, prompts_2, negative_prompts_2) -> dict: + if ('prompt' in possible) and ('prompt' not in args) and (prompts is not None) and len(prompts) > 0: + debug_log(f'Prompt fallback: prompt={prompts}') + args['prompt'] = prompts + if ('negative_prompt' in possible) and ('negative_prompt' not in args) and (negative_prompts is not None) and len(negative_prompts) > 0: + debug_log(f'Prompt fallback: negative_prompt={negative_prompts}') + args['negative_prompt'] = negative_prompts + if ('prompt_2' in possible) and ('prompt_2' not in args) and (prompts_2 is not None) and len(prompts_2) > 0: + debug_log(f'Prompt fallback: prompt_2={prompts_2}') + args['prompt_2'] = prompts_2 + if ('negative_prompt_2' in possible) and ('negative_prompt_2' not in args) and (negative_prompts_2 is not None) and len(negative_prompts_2) > 0: + debug_log(f'Prompt fallback: negative_prompt_2={negative_prompts_2}') + args['negative_prompt_2'] = negative_prompts_2 + return args + + +def set_prompt(p, + args: dict, + possible: list[str], + cls: str, + prompt_attention: str, + steps: int, + clip_skip: int, + prompts: list[str], + negative_prompts: list[str], + prompts_2: list[str], + negative_prompts_2: list[str], + ) -> dict: + prompt_attention = prompt_attention or shared.opts.prompt_attention + if (prompt_attention != 'fixed') and ('Onnx' not in cls) and ('prompt' not in p.task_args) and ( + ('StableDiffusion' in cls) or + ('StableCascade' in cls) or + ('Flux' in cls and 'Flux2' not in cls) or + ('Chroma' in cls) or + ('HiDreamImagePipeline' in cls) + ): + jobid = shared.state.begin('TE Encode') + try: + prompt_parser_diffusers.embedder = prompt_parser_diffusers.PromptEmbedder(prompts, negative_prompts, steps, clip_skip, p) + except Exception as e: + prompt_parser_diffusers.embedder = None + shared.log.error(f'Prompt parser encode: {e}') + if debug_enabled: + errors.display(e, 'Prompt parser encode') + timer.process.record('prompt', reset=False) + shared.state.end(jobid) + else: + prompt_parser_diffusers.embedder = None + prompt_attention = 'fixed' + + prompts, negative_prompts, prompts_2, negative_prompts_2 = fix_prompt_batch(p, prompts, negative_prompts, prompts_2, negative_prompts_2) + prompts, negative_prompts, prompts_2, negative_prompts_2 = fix_prompt_model(cls, prompts, negative_prompts, prompts_2, negative_prompts_2) + + if prompt_parser_diffusers.embedder is not None: + if 'prompt' in possible: + debug_log(f'Prompt set embeds: positive={prompts}') + prompt_embeds = prompt_parser_diffusers.embedder('prompt_embeds') + prompt_pooled_embeds = prompt_parser_diffusers.embedder('positive_pooleds') + prompt_attention_masks = prompt_parser_diffusers.embedder('prompt_attention_masks') + + if prompt_embeds is None: + shared.log.warning('Prompt parser encode: empty prompt embeds') + prompt_parser_diffusers.embedder = None + args = set_fallback_prompt(args, possible, prompts=prompts, negative_prompts=None, prompts_2=None, negative_prompts_2=None) + prompt_attention = 'fixed' + elif prompt_embeds.device == torch.device('meta'): + shared.log.warning('Prompt parser encode: embeds on meta device') + prompt_parser_diffusers.embedder = None + args = set_fallback_prompt(args, possible, prompts=prompts, negative_prompts=None, prompts_2=None, negative_prompts_2=None) + prompt_attention = 'fixed' + else: + if 'prompt_embeds' in possible: + args['prompt_embeds'] = prompt_embeds + else: + args = set_fallback_prompt(args, possible, prompts=prompts, negative_prompts=None, prompts_2=None, negative_prompts_2=None) + if 'pooled_prompt_embeds' in possible: + args['pooled_prompt_embeds'] = prompt_pooled_embeds + if 'StableCascade' in cls: + args['prompt_embeds_pooled'] = prompt_pooled_embeds.unsqueeze(0) + if 'HiDreamImage' in cls: + args['prompt_embeds_t5'] = prompt_embeds[0] + args['prompt_embeds_llama3'] = prompt_embeds[1] + if 'prompt_attention_mask' in possible: + args['prompt_attention_mask'] = prompt_attention_masks + + if 'negative_prompt' in possible: + debug_log(f'Prompt set embeds: negative={negative_prompts}') + negative_embeds = prompt_parser_diffusers.embedder('negative_prompt_embeds') + negative_pooled_embeds = prompt_parser_diffusers.embedder('negative_pooleds') + negative_attention_masks = prompt_parser_diffusers.embedder('negative_prompt_attention_masks') + + if negative_embeds is None: + shared.log.warning('Prompt parser encode: empty negative prompt embeds') + prompt_parser_diffusers.embedder = None + args = set_fallback_prompt(args, possible, prompts=None, negative_prompts=negative_prompts, prompts_2=None, negative_prompts_2=None) + prompt_attention = 'fixed' + elif negative_embeds.device == torch.device('meta'): + shared.log.warning('Prompt parser encode: negative embeds on meta device') + prompt_parser_diffusers.embedder = None + args = set_fallback_prompt(args, possible, prompts=None, negative_prompts=negative_prompts, prompts_2=None, negative_prompts_2=None) + prompt_attention = 'fixed' + else: + if 'negative_prompt_embeds' in possible: + args['negative_prompt_embeds'] = negative_embeds + else: + args = set_fallback_prompt(args, possible, prompts=None, negative_prompts=negative_prompts, prompts_2=None, negative_prompts_2=None) + if 'negative_pooled_prompt_embeds' in possible: + args['negative_pooled_prompt_embeds'] = negative_pooled_embeds + if 'StableCascade' in cls: + args['negative_prompt_embeds_pooled'] = negative_pooled_embeds.unsqueeze(0) + if 'HiDreamImage' in cls: + args['negative_prompt_embeds_t5'] = negative_embeds[0] + args['negative_prompt_embeds_llama3'] = negative_embeds[1] + if 'negative_prompt_attention_mask' in possible: + args['negative_prompt_attention_mask'] = negative_attention_masks + else: + debug_log('Prompt fallback: no embedder') + args = set_fallback_prompt(args, possible, prompts=prompts, negative_prompts=negative_prompts, prompts_2=None, negative_prompts_2=None) + prompt_attention = 'fixed' + + if 'prompt_embeds' not in args and 'negative_prompt_embeds' not in args: # pass secondary prompts as-in + args = set_fallback_prompt(args, possible, prompts=None, negative_prompts=None, prompts_2=prompts_2, negative_prompts_2=negative_prompts_2) + + if (prompt_parser_diffusers.embedder is not None) and (not prompt_parser_diffusers.embedder.scheduled_prompt): + prompt_parser_diffusers.embedder = None # not scheduled so we dont need it anymore + + return prompt_attention, args diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index a3408c61c..25349e583 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -48,7 +48,13 @@ def prepare_model(pipe = None): class PromptEmbedder: - def __init__(self, prompts, negative_prompts, steps, clip_skip, p): + def __init__(self, + prompts, + negative_prompts, + steps, + clip_skip, + p, + ): t0 = time.time() self.prompts = prompts self.negative_prompts = negative_prompts diff --git a/modules/ras/ras_attention.py b/modules/ras/ras_attention.py index ca3a083e2..4989cc931 100644 --- a/modules/ras/ras_attention.py +++ b/modules/ras/ras_attention.py @@ -126,7 +126,7 @@ class RASLuminaAttnProcessor2_0: else: softmax_scale = attn.scale - # perform Grouped-qurey Attention (GQA) # TODO replace with GQA + # perform Grouped-qurey Attention (GQA) n_rep = attn.heads // kv_heads if n_rep >= 1: key = key.unsqueeze(3).repeat(1, 1, 1, n_rep, 1).flatten(2, 3) diff --git a/modules/rocm.py b/modules/rocm.py index ca74810d5..c4c205bcb 100644 --- a/modules/rocm.py +++ b/modules/rocm.py @@ -1,11 +1,18 @@ import os import sys +import glob import ctypes import shutil import subprocess -from typing import Union, List +from types import ModuleType +from typing import Union, overload, TYPE_CHECKING from enum import Enum from functools import wraps +if TYPE_CHECKING: + import torch + + +rocm_sdk: Union[ModuleType, None] = None def resolve_link(path_: str) -> str: @@ -20,8 +27,8 @@ def dirname(path_: str, r: int = 1) -> str: return path_ -def spawn(command: Union[str, List[str]], cwd: os.PathLike = '.') -> str: - process = subprocess.run(command, cwd=cwd, shell=True, check=False, stdout=subprocess.PIPE, stderr=subprocess.PIPE) +def spawn(command: Union[str, list[str]], cwd: os.PathLike = '.') -> str: + process = subprocess.run(command, cwd=cwd, shell=True, check=False, stdout=subprocess.PIPE, stderr=subprocess.DEVNULL) return process.stdout.decode(encoding="utf8", errors="ignore") @@ -45,14 +52,14 @@ class ROCmEnvironment(Environment): class PythonPackageEnvironment(Environment): hip: ctypes.CDLL - def __init__(self): - import _rocm_sdk_core - if sys.platform == "win32": - path = os.path.join(_rocm_sdk_core.__path__[0], "bin", "amdhip64_7.dll") - else: - raise NotImplementedError - # This library will be loaded/used by PyTorch. So it won't make conflicts. - self.hip = ctypes.CDLL(path) + def __init__(self, rocm_sdk_module: ModuleType): + spec = rocm_sdk_module._dist_info.ALL_PACKAGES['core'].get_py_package() # pylint: disable=protected-access + lib = rocm_sdk_module._dist_info.ALL_LIBRARIES['amdhip64'] # pylint: disable=protected-access + pattern = os.path.join(os.path.dirname(spec.origin), lib.windows_relpath if sys.platform == "win32" else lib.posix_relpath, lib.dll_pattern if sys.platform == "win32" else lib.so_pattern) + candidates = glob.glob(pattern) + if len(candidates) == 0: + raise FileNotFoundError("Could not find amdhip64 in rocm-sdk package") + self.hip = ctypes.CDLL(candidates[0]) class MicroArchitecture(Enum): @@ -83,9 +90,19 @@ class Agent: break return result - def __init__(self, name: str): - self.name = name - self.gfx_version = Agent.parse_gfx_version(name) + @overload + def __init__(self, name: str): ... + @overload + def __init__(self, device: 'torch.types.Device'): ... + + def __init__(self, arg): + if isinstance(arg, str): + name = arg + else: # assume arg is device-like object + import torch + name = getattr(torch.cuda.get_device_properties(arg), "gcnArchName", "gfx0000") + self.name = name.split(':')[0] + self.gfx_version = Agent.parse_gfx_version(self.name) if self.gfx_version > 0x1000: self.arch = MicroArchitecture.RDNA elif self.gfx_version in (0x908, 0x90a, 0x942,): @@ -93,14 +110,17 @@ class Agent: else: self.arch = MicroArchitecture.GCN self.is_apu = (self.gfx_version & 0xFFF0 == 0x1150) or self.gfx_version in (0x801, 0x902, 0x90c, 0x1013, 0x1033, 0x1035, 0x1036, 0x1103,) - self.blaslt_supported = os.path.exists(os.path.join(blaslt_tensile_libpath, f"Kernels.so-000-{name}.hsaco" if sys.platform == "win32" else f"extop_{name}.co")) + self.blaslt_supported = False if blaslt_tensile_libpath is None else os.path.exists(os.path.join(blaslt_tensile_libpath, f"Kernels.so-000-{self.name}.hsaco" if sys.platform == "win32" else f"extop_{self.name}.co")) + + def __str__(self) -> str: + return self.name @property def therock(self) -> Union[str, None]: if (self.gfx_version & 0xFFF0) == 0x1200: return "v2/gfx120X-all" if (self.gfx_version & 0xFFF0) == 0x1100: - return "v2/gfx110X-all" + return "v2/gfx110X-" + ("all" if self.is_apu else "dgpu") if self.gfx_version == 0x1150: return "v2-staging/gfx1150" if self.gfx_version == 0x1151: @@ -127,14 +147,7 @@ class Agent: return None -def find() -> Union[Environment, None]: - try: # TheRock - import _rocm_sdk_core # pylint: disable=unused-import - return PythonPackageEnvironment() - except ImportError: - pass - - # system-wide installation +def find() -> Union[ROCmEnvironment, None]: hip_path = shutil.which("hipconfig") if hip_path is not None: return ROCmEnvironment(dirname(resolve_link(hip_path), 2)) @@ -218,42 +231,44 @@ def get_flash_attention_command(agent: Agent) -> str: def refresh(): - global environment, blaslt_tensile_libpath, is_installed, version # pylint: disable=global-statement - if sys.platform == "win32": - global agents # pylint: disable=global-statement + global rocm_sdk, environment, blaslt_tensile_libpath, is_installed, version # pylint: disable=global-statement + try: + import rocm_sdk + environment = PythonPackageEnvironment(rocm_sdk) try: - agents = driver_get_agents() + target_family = rocm_sdk._dist_info.determine_target_family() # pylint: disable=protected-access + spec = rocm_sdk._dist_info.ALL_PACKAGES['libraries'].get_py_package(target_family) # pylint: disable=protected-access + blaslt_tensile_libpath = os.path.join(os.path.dirname(spec.origin), "bin", "hipblaslt", "library") except Exception: - agents = [] - environment = find() + blaslt_tensile_libpath = None + spawn(["rocm-sdk", "init"]) + except ImportError: + rocm_sdk = None + environment = find() + if environment is not None: + blaslt_tensile_libpath = os.path.join(environment.path, "bin" if sys.platform == "win32" else "lib", "hipblaslt", "library") + if environment is not None: - if isinstance(environment, ROCmEnvironment): - blaslt_tensile_libpath = os.environ.get("HIPBLASLT_TENSILE_LIBPATH", os.path.join(environment.path, "bin" if sys.platform == "win32" else "lib", "hipblaslt", "library")) - elif isinstance(environment, PythonPackageEnvironment): - spawn(["rocm-sdk", "init"]) + blaslt_tensile_libpath = os.environ.get("HIPBLASLT_TENSILE_LIBPATH", blaslt_tensile_libpath) is_installed = True version = get_version() if sys.platform == "win32": - def get_agents() -> List[Agent]: - return agents - #if isinstance(environment, ROCmEnvironment): - # out = spawn("amdgpu-arch", cwd=os.path.join(environment.path, 'bin')) - #else: - # # Assume that amdgpu-arch is in PATH (venv/Scripts/amdgpu-arch.exe) - # out = spawn("amdgpu-arch") - #out = out.strip() - #if out == "": - # return [] - #return [Agent(x.split(' ')[-1].strip()) for x in out.split("\n")] + import tempfile - def driver_get_agents() -> List[Agent]: - # unsafe and experimental feature - from modules import windows_hip_ffi - archs = windows_hip_ffi.get_archs() - # filter out None (is there any better way?) - return [Agent(x) for x in archs if x is not None] + def get_agents() -> list[Agent]: + name = None + with tempfile.NamedTemporaryFile(mode='w', encoding='utf-8', delete=False) as f: + name = f.name + f.write(CODE_AMDGPU_ARCH) + f.flush() + out = spawn([sys.executable, name]) + os.unlink(name) + out = out.strip() + if out == "": + return [] + return [Agent(x.split(' ')[-1].strip()) for x in out.split("\n")] def postinstall(): import torch @@ -268,24 +283,31 @@ if sys.platform == "win32": os.environ["PATH"] = ";".join(paths_no_rocm) return - build_targets = torch.cuda.get_arch_list() - for available in agents: - if available.name in build_targets: - return - - # use cpu instead of crashing - torch.cuda.is_available = lambda: False - def rocm_init(): try: import torch import numpy as np - from modules.devices import get_optimal_device + from installer import log + from modules.devices import get_hip_agent + from modules.rocm_triton_windows import apply_triton_patches - gfx_version = Agent.parse_gfx_version(getattr(torch.cuda.get_device_properties(get_optimal_device()), "gcnArchName", "gfx0000")) - if (gfx_version & 0xFFF0) == 0x1200: + build_targets = torch.cuda.get_arch_list() + agents = get_agents() + if all(available.name not in build_targets for available in agents): + log.warning('ROCm: torch-rocm is installed, but none of build targets is available') + # use cpu instead of crashing + torch.cuda.is_available = lambda: False + + agent = get_hip_agent() + if not agent.blaslt_supported: + log.warning(f'ROCm: hipBLASLt unavailable agent={agent}') + if (agent.gfx_version & 0xFFF0) == 0x1200: # disable MIOpen for gfx120x torch.backends.cudnn.enabled = False + log.debug('ROCm: disabled MIOpen') + + if sys.platform == "win32": + apply_triton_patches() original_cholesky_ex = torch.linalg.cholesky_ex @wraps(original_cholesky_ex) @@ -314,9 +336,8 @@ if sys.platform == "win32": return True, None is_wsl: bool = False - agents: List[Agent] = [] # temp else: # sys.platform != "win32" - def get_agents() -> List[Agent]: + def get_agents() -> list[Agent]: try: _agents = spawn("rocm_agent_enumerator").split("\n") _agents = [x for x in _agents if x and x != 'gfx000'] @@ -330,6 +351,7 @@ else: # sys.platform != "win32" try: if shutil.which("conda") is not None: # Preload stdc++ library. This will bypass Anaconda stdc++ library. + # (hsa-runtime64 depends on stdc++) load_library_global("/lib/x86_64-linux-gnu/libstdc++.so.6") # Preload rocr4wsl. The user don't have to replace the library file. load_library_global("/opt/rocm/lib/libhsa-runtime64.so") @@ -337,12 +359,107 @@ else: # sys.platform != "win32" pass def rocm_init(): + try: + import torch + from installer import log + from modules.devices import get_hip_agent + + agent = get_hip_agent() + if not agent.blaslt_supported: + log.debug(f'ROCm: hipBLASLt unavailable agent={agent}') + except Exception as e: + return False, e return True, None is_wsl: bool = os.environ.get('WSL_DISTRO_NAME', 'unknown' if spawn('wslpath -w /') else None) is not None -environment = None -blaslt_tensile_libpath = "" -is_installed = False -version = None +environment: Union[Environment, None] = None +blaslt_tensile_libpath: Union[str, None] = None +is_installed: bool = False +version: Union[str, None] = None refresh() + +# amdgpu-arch.exe written in Python +CODE_AMDGPU_ARCH = """ +import os +import sys +import ctypes +import ctypes.wintypes +import contextlib +hipDeviceProp = ctypes.c_byte * 1472 +@contextlib.contextmanager +def mute(fd): + s = os.dup(fd) + try: + with open(os.devnull, 'w') as devnull: + os.dup2(devnull.fileno(), fd) + yield + finally: + os.dup2(s, fd) + os.close(s) +class HIP: + def __init__(self): + ctypes.windll.kernel32.LoadLibraryA.restype = ctypes.wintypes.HMODULE + ctypes.windll.kernel32.LoadLibraryA.argtypes = [ctypes.c_char_p] + self.handle = None + path = os.environ.get("windir", "C:\\\\Windows") + "\\\\System32\\\\amdhip64_7.dll" + if not os.path.isfile(path): + path = os.environ.get("windir", "C:\\\\Windows") + "\\\\System32\\\\amdhip64_6.dll" + if not os.path.isfile(path): + path = os.environ.get("windir", "C:\\\\Windows") + "\\\\System32\\\\amdhip64.dll" + assert os.path.isfile(path) + self.handle = ctypes.windll.kernel32.LoadLibraryA(path.encode('utf-8')) + ctypes.windll.kernel32.GetLastError.restype = ctypes.wintypes.DWORD + ctypes.windll.kernel32.GetLastError.argtypes = [] + assert ctypes.windll.kernel32.GetLastError() == 0 + ctypes.windll.kernel32.GetProcAddress.restype = ctypes.c_void_p + ctypes.windll.kernel32.GetProcAddress.argtypes = [ctypes.wintypes.HMODULE, ctypes.c_char_p] + hipInit = ctypes.CFUNCTYPE(ctypes.c_int, ctypes.c_uint)( + ctypes.windll.kernel32.GetProcAddress(self.handle, b"hipInit")) + with mute(sys.stdout.fileno()): + hipInit(0) + self.hipGetDeviceCount = ctypes.CFUNCTYPE(ctypes.c_int, ctypes.POINTER(ctypes.c_int))( + ctypes.windll.kernel32.GetProcAddress(self.handle, b"hipGetDeviceCount")) + self.hipGetDeviceProperties = ctypes.CFUNCTYPE(ctypes.c_int, ctypes.POINTER(hipDeviceProp), ctypes.c_int)( + ctypes.windll.kernel32.GetProcAddress(self.handle, b"hipGetDeviceProperties")) + def get_device_count(self) -> int: + count = ctypes.c_int() + assert self.hipGetDeviceCount(ctypes.byref(count)) == 0 + return count.value + def get_device_properties(self, device_id) -> bytes: + prop = hipDeviceProp() + assert self.hipGetDeviceProperties(ctypes.byref(prop), device_id) == 0 + return bytes(prop) +if __name__ == "__main__": + hip = HIP() + count = hip.get_device_count() + archs: list[str | None] = [None] * count + for i in range(count): + prop = hip.get_device_properties(i) + name = "" + idx = 0 + while idx < len(prop): + try: + idx = prop.index(0x67, idx) + 1 + except ValueError: + break + if prop[idx] != 0x66: + continue + if prop[idx + 1] != 0x78: + continue + idx = idx + 2 + while prop[idx] != 0x00: + c = prop[idx] + idx += 1 + if (c < 0x30 or c > 0x39) and (c < 0x61 or c > 0x66): + name = "" + continue + name += chr(c) + break + if name: + archs[i] = "gfx" + name + del hip + for arch in archs: + if arch is not None: + print(arch) +""" diff --git a/modules/rocm_triton_windows.py b/modules/rocm_triton_windows.py index 939713da9..88c509b6d 100644 --- a/modules/rocm_triton_windows.py +++ b/modules/rocm_triton_windows.py @@ -1,6 +1,8 @@ import sys +from typing import Union import torch -from modules import shared +from modules import shared, devices +from modules.rocm import Agent if sys.platform == "win32": @@ -35,7 +37,7 @@ if sys.platform == "win32": class DeviceProperties: PROPERTIES_OVERRIDE = { # sometimes gcnArchName contains device name ("AMD Radeon RX ..."), not architecture name ("gfx...") - "gcnArchName": "UNKNOWN ARCHITECTURE", + "gcnArchName": "gfx0000", } internal: torch._C._CudaDeviceProperties @@ -56,20 +58,17 @@ if sys.platform == "win32": from modules import zluda return zluda.core.to_hip_stream(_cuda_getCurrentRawStream(device)) - def get_default_agent_name(): + def get_default_agent() -> Union[Agent, None]: if shared.devices.backend == "rocm": - device = shared.devices.get_optimal_device() - return getattr(torch.cuda.get_device_properties(device), "gcnArchName", None) + return devices.get_hip_agent() else: from modules import zluda - if zluda.default_agent is None: - return None - return zluda.default_agent.name + return zluda.default_agent def apply_triton_patches(): - arch_name = get_default_agent_name() - if arch_name is not None: - DeviceProperties.PROPERTIES_OVERRIDE["gcnArchName"] = arch_name + agent = get_default_agent() + if agent is not None: + DeviceProperties.PROPERTIES_OVERRIDE["gcnArchName"] = agent.name torch.cuda._get_device_properties = torch_cuda__get_device_properties # pylint: disable=protected-access if shared.devices.backend == "zluda": torch._C._cuda_getCurrentRawStream = torch__C__cuda_getCurrentRawStream # pylint: disable=protected-access diff --git a/modules/schedulers/scheduler_flashflow.py b/modules/schedulers/scheduler_flashflow.py index 122f8ed74..edc63016f 100644 --- a/modules/schedulers/scheduler_flashflow.py +++ b/modules/schedulers/scheduler_flashflow.py @@ -349,7 +349,6 @@ class FlashFlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin): """Constructs the noise schedule of Karras et al. (2022).""" # Hack to make sure that other schedulers which copy this function don't break - # TODO: Add this logic to the other schedulers if hasattr(self.config, "sigma_min"): sigma_min = self.config.sigma_min else: @@ -375,7 +374,6 @@ class FlashFlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin): """Constructs an exponential noise schedule.""" # Hack to make sure that other schedulers which copy this function don't break - # TODO: Add this logic to the other schedulers if hasattr(self.config, "sigma_min"): sigma_min = self.config.sigma_min else: @@ -399,7 +397,6 @@ class FlashFlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin): """From "Beta Sampling is All You Need" [arXiv:2407.12173] (Lee et. al, 2024)""" # Hack to make sure that other schedulers which copy this function don't break - # TODO: Add this logic to the other schedulers if hasattr(self.config, "sigma_min"): sigma_min = self.config.sigma_min else: diff --git a/modules/sd_checkpoint.py b/modules/sd_checkpoint.py index ff192b668..375c1c58f 100644 --- a/modules/sd_checkpoint.py +++ b/modules/sd_checkpoint.py @@ -116,11 +116,15 @@ def setup_model(): # sd_hijack_accelerate.hijack_torch_conv() -def checkpoint_titles(): +def checkpoint_titles(use_short=False): def convert(name): return int(name) if name.isdigit() else name.lower() + def alphanumeric_key(key): - return [convert(c) for c in re.split('([0-9]+)', key)] + return [convert(c) for c in re.split("([0-9]+)", key)] + + if use_short: + return sorted([x.title.rsplit("\\", 1)[-1].rsplit("/", 1)[-1] for x in checkpoints_list.values()], key=alphanumeric_key) return sorted([x.title for x in checkpoints_list.values()], key=alphanumeric_key) diff --git a/modules/sd_detect.py b/modules/sd_detect.py index 2143e34b3..a1cb6e913 100644 --- a/modules/sd_detect.py +++ b/modules/sd_detect.py @@ -92,6 +92,8 @@ def guess_by_name(fn, current_guess): new_guess = 'HiDream' elif 'chroma' in fn.lower() and 'xl' not in fn.lower(): new_guess = 'Chroma' + elif 'flux.2' in fn.lower() and 'klein' in fn.lower(): + new_guess = 'FLUX2 Klein' elif 'flux.2' in fn.lower(): new_guess = 'FLUX2' elif 'flux' in fn.lower() or 'flex.1' in fn.lower(): @@ -143,6 +145,8 @@ def guess_by_name(fn, current_guess): new_guess = 'LongCat' elif 'ovis-image' in fn.lower(): new_guess = 'Ovis-Image' + elif 'glm-image' in fn.lower(): + new_guess = 'GLM-Image' if debug_load: shared.log.trace(f'Autodetect: method=name file="{fn}" previous="{current_guess}" current="{new_guess}"') return new_guess or current_guess diff --git a/modules/sd_hijack_vae.py b/modules/sd_hijack_vae.py index 915e64846..8034a6428 100644 --- a/modules/sd_hijack_vae.py +++ b/modules/sd_hijack_vae.py @@ -29,7 +29,10 @@ def hijack_vae_decode(*args, **kwargs): else: res = shared.sd_model.vae.orig_decode(latents, *args[1:], **kwargs) t1 = time.time() - shared.log.debug(f'Decode: vae={shared.sd_model.vae.__class__.__name__} slicing={getattr(shared.sd_model.vae, "use_slicing", None)} tiling={getattr(shared.sd_model.vae, "use_tiling", None)} latents={list(latents.shape)}:{latents.device} dtype={latents.dtype} time={t1-t0:.3f}') + try: + shared.log.debug(f'Decode: vae={shared.sd_model.vae.__class__.__name__} dtype={latents.dtype} latents={list(latents.shape)}:{latents.device} decoded={list(res[0].shape)} slicing={getattr(shared.sd_model.vae, "use_slicing", None)} tiling={getattr(shared.sd_model.vae, "use_tiling", None)} time={t1-t0:.3f}') + except Exception: + pass else: res = shared.sd_model.vae.orig_decode(*args, **kwargs) except Exception as e: diff --git a/modules/sd_models.py b/modules/sd_models.py index da0888431..89a437463 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -359,6 +359,10 @@ def load_diffuser_force(detected_model_type, checkpoint_info, diffusers_load_con from pipelines.model_flux2 import load_flux2 sd_model = load_flux2(checkpoint_info, diffusers_load_config) allow_post_quant = False + elif model_type in ['FLUX2 Klein']: + from pipelines.model_flux2_klein import load_flux2_klein + sd_model = load_flux2_klein(checkpoint_info, diffusers_load_config) + allow_post_quant = False elif model_type in ['FLEX']: from pipelines.model_flex import load_flex sd_model = load_flex(checkpoint_info, diffusers_load_config) @@ -439,7 +443,7 @@ def load_diffuser_force(detected_model_type, checkpoint_info, diffusers_load_con from pipelines.model_kandinsky import load_kandinsky3 sd_model = load_kandinsky3(checkpoint_info, diffusers_load_config) allow_post_quant = False - elif model_type in ['Kandinsky 5.0']: + elif model_type in ['Kandinsky 5.0'] and '2I' in model_type: from pipelines.model_kandinsky import load_kandinsky5 sd_model = load_kandinsky5(checkpoint_info, diffusers_load_config) allow_post_quant = False @@ -483,12 +487,19 @@ def load_diffuser_force(detected_model_type, checkpoint_info, diffusers_load_con from pipelines.model_ovis import load_ovis sd_model = load_ovis(checkpoint_info, diffusers_load_config) allow_post_quant = False + elif model_type in ['GLM-Image']: + from pipelines.model_glm import load_glm_image + sd_model = load_glm_image(checkpoint_info, diffusers_load_config) + allow_post_quant = False except Exception as e: shared.log.error(f'Load {op}: path="{checkpoint_info.path}" {e}') if debug_load: errors.display(e, 'Load') - return None - return sd_model + return None, True + if sd_model is not None: + return sd_model, True + else: + return sd_model, False def load_diffuser_folder(model_type, pipeline, checkpoint_info, diffusers_load_config, op='model'): @@ -789,6 +800,7 @@ def load_diffuser(checkpoint_info=None, op='model', revision=None): # pylint: di return sd_model = None + handled = False try: # initial load only if sd_model is None: @@ -829,25 +841,25 @@ def load_diffuser(checkpoint_info=None, op='model', revision=None): # pylint: di timer.load.record("vae") # load with custom loader - if sd_model is None: - sd_model = load_diffuser_force(model_type, checkpoint_info, diffusers_load_config, op) + if sd_model is None and not handled: + sd_model, handled = load_diffuser_force(model_type, checkpoint_info, diffusers_load_config, op) if sd_model is not None and not sd_model: shared.log.error(f'Load {op}: type="{model_type}" pipeline="{pipeline}" not loaded') return # load sdnq-prequantized model - if sd_model is None: + if sd_model is None and not handled: if model_type.endswith('SDNQ'): sd_model = load_sdnq_model(checkpoint_info, pipeline, diffusers_load_config, op) model_type = model_type.replace(' SDNQ', '') # load from single-file - if sd_model is None: + if sd_model is None and not handled: if os.path.isfile(checkpoint_info.path) and checkpoint_info.path.lower().endswith('.safetensors'): sd_model = load_diffuser_file(model_type, pipeline, checkpoint_info, diffusers_load_config, op) # load from hf folder-style - if sd_model is None: + if sd_model is None and not handled: if os.path.isdir(checkpoint_info.path) or (checkpoint_info.type == 'huggingface') or (checkpoint_info.type == 'transformer') or (checkpoint_info.type == 'reference'): sd_model = load_diffuser_folder(model_type, pipeline, checkpoint_info, diffusers_load_config, op) diff --git a/modules/sd_offload.py b/modules/sd_offload.py index 846a38773..c7b2b0e66 100644 --- a/modules/sd_offload.py +++ b/modules/sd_offload.py @@ -7,18 +7,22 @@ import torch import accelerate.hooks import accelerate.utils.modeling from installer import log -from modules import shared, devices, errors, model_quant +from modules import shared, devices, errors, model_quant, sd_models from modules.timer import process as process_timer debug = os.environ.get('SD_MOVE_DEBUG', None) is not None verbose = os.environ.get('SD_MOVE_VERBOSE', None) is not None debug_move = log.trace if debug else lambda *args, **kwargs: None -offload_warn = ['sc', 'sd3', 'f1', 'f2', 'h1', 'hunyuandit', 'auraflow', 'omnigen', 'omnigen2', 'cogview4', 'cosmos', 'chroma', 'x-omni', 'hunyuanimage', 'hunyuanimage3', 'longcat'] +offload_allow_none = ['sd', 'sdxl'] offload_post = ['h1'] offload_hook_instance = None balanced_offload_exclude = ['CogView4Pipeline', 'MeissonicPipeline'] -no_split_module_classes = ["Linear", "Conv1d", "Conv2d", "Conv3d", "ConvTranspose1d", "ConvTranspose2d", "ConvTranspose3d", "WanTransformerBlock"] +no_split_module_classes = [ + "Linear", "Conv1d", "Conv2d", "Conv3d", "ConvTranspose1d", "ConvTranspose2d", "ConvTranspose3d", + "SDNQLinear", "SDNQConv1d", "SDNQConv2d", "SDNQConv3d", "SDNQConvTranspose1d", "SDNQConvTranspose2d", "SDNQConvTranspose3d", + "WanTransformerBlock", +] accelerate_dtype_byte_size = None move_stream = None @@ -92,6 +96,60 @@ def apply_group_offload(sd_model, op:str='model'): return sd_model +def apply_model_offload(sd_model, op:str='model', quiet:bool=False): + try: + shared.log.quiet(quiet, f'Setting {op}: offload={shared.opts.diffusers_offload_mode} limit={shared.opts.cuda_mem_fraction}') + if shared.opts.diffusers_move_base or shared.opts.diffusers_move_unet or shared.opts.diffusers_move_refiner: + shared.opts.diffusers_move_base = False + shared.opts.diffusers_move_unet = False + shared.opts.diffusers_move_refiner = False + shared.log.warning(f'Disabling {op} "Move model to CPU" since "Model CPU offload" is enabled') + if not hasattr(sd_model, "_all_hooks") or len(sd_model._all_hooks) == 0: # pylint: disable=protected-access + sd_model.enable_model_cpu_offload(device=devices.device) + else: + sd_model.maybe_free_model_hooks() + set_accelerate(sd_model) + except Exception as e: + shared.log.error(f'Setting {op}: offload={shared.opts.diffusers_offload_mode} {e}') + + +def apply_sequential_offload(sd_model, op:str='model', quiet:bool=False): + try: + shared.log.quiet(quiet, f'Setting {op}: offload={shared.opts.diffusers_offload_mode} limit={shared.opts.cuda_mem_fraction}') + if shared.opts.diffusers_move_base or shared.opts.diffusers_move_unet or shared.opts.diffusers_move_refiner: + shared.opts.diffusers_move_base = False + shared.opts.diffusers_move_unet = False + shared.opts.diffusers_move_refiner = False + shared.log.warning(f'Disabling {op} "Move model to CPU" since "Sequential CPU offload" is enabled') + if sd_model.has_accelerate: + if op == "vae": # reapply sequential offload to vae + from accelerate import cpu_offload + sd_model.vae.to(devices.cpu) + cpu_offload(sd_model.vae, devices.device, offload_buffers=len(sd_model.vae._parameters) > 0) # pylint: disable=protected-access + else: + pass # do nothing if offload is already applied + else: + sd_model.enable_sequential_cpu_offload(device=devices.device) + set_accelerate(sd_model) + except Exception as e: + shared.log.error(f'Setting {op}: offload={shared.opts.diffusers_offload_mode} {e}') + + +def apply_none_offload(sd_model, op:str='model', quiet:bool=False): + if shared.sd_model_type not in offload_allow_none: + shared.log.warning(f'Setting {op}: offload={shared.opts.diffusers_offload_mode} type={shared.sd_model.__class__.__name__} large model') + else: + shared.log.quiet(quiet, f'Setting {op}: offload={shared.opts.diffusers_offload_mode} limit={shared.opts.cuda_mem_fraction}') + try: + sd_model.has_accelerate = False + if hasattr(sd_model, 'maybe_free_model_hooks'): + sd_model.maybe_free_model_hooks() + sd_model = accelerate.hooks.remove_hook_from_module(sd_model, recurse=True) + except Exception: + pass + sd_models.move_model(sd_model, devices.device) + + def set_diffuser_offload(sd_model, op:str='model', quiet:bool=False, force:bool=False): global accelerate_dtype_byte_size # pylint: disable=global-statement t0 = time.time() @@ -105,50 +163,13 @@ def set_diffuser_offload(sd_model, op:str='model', quiet:bool=False, force:bool= accelerate.utils.modeling.dtype_byte_size = dtype_byte_size if shared.opts.diffusers_offload_mode == "none": - if shared.sd_model_type in offload_warn or 'video' in shared.sd_model_type: - shared.log.warning(f'Setting {op}: offload={shared.opts.diffusers_offload_mode} type={shared.sd_model.__class__.__name__} large model') - else: - shared.log.quiet(quiet, f'Setting {op}: offload={shared.opts.diffusers_offload_mode} limit={shared.opts.cuda_mem_fraction}') - if hasattr(sd_model, 'maybe_free_model_hooks'): - sd_model.maybe_free_model_hooks() - sd_model.has_accelerate = False + apply_none_offload(sd_model, op=op, quiet=quiet) if shared.opts.diffusers_offload_mode == "model" and hasattr(sd_model, "enable_model_cpu_offload"): - try: - shared.log.quiet(quiet, f'Setting {op}: offload={shared.opts.diffusers_offload_mode} limit={shared.opts.cuda_mem_fraction}') - if shared.opts.diffusers_move_base or shared.opts.diffusers_move_unet or shared.opts.diffusers_move_refiner: - shared.opts.diffusers_move_base = False - shared.opts.diffusers_move_unet = False - shared.opts.diffusers_move_refiner = False - shared.log.warning(f'Disabling {op} "Move model to CPU" since "Model CPU offload" is enabled') - if not hasattr(sd_model, "_all_hooks") or len(sd_model._all_hooks) == 0: # pylint: disable=protected-access - sd_model.enable_model_cpu_offload(device=devices.device) - else: - sd_model.maybe_free_model_hooks() - set_accelerate(sd_model) - except Exception as e: - shared.log.error(f'Setting {op}: offload={shared.opts.diffusers_offload_mode} {e}') + apply_model_offload(sd_model, op=op, quiet=quiet) if shared.opts.diffusers_offload_mode == "sequential" and hasattr(sd_model, "enable_sequential_cpu_offload"): - try: - shared.log.debug(f'Setting {op}: offload={shared.opts.diffusers_offload_mode} limit={shared.opts.cuda_mem_fraction}') - if shared.opts.diffusers_move_base or shared.opts.diffusers_move_unet or shared.opts.diffusers_move_refiner: - shared.opts.diffusers_move_base = False - shared.opts.diffusers_move_unet = False - shared.opts.diffusers_move_refiner = False - shared.log.warning(f'Disabling {op} "Move model to CPU" since "Sequential CPU offload" is enabled') - if sd_model.has_accelerate: - if op == "vae": # reapply sequential offload to vae - from accelerate import cpu_offload - sd_model.vae.to(devices.cpu) - cpu_offload(sd_model.vae, devices.device, offload_buffers=len(sd_model.vae._parameters) > 0) # pylint: disable=protected-access - else: - pass # do nothing if offload is already applied - else: - sd_model.enable_sequential_cpu_offload(device=devices.device) - set_accelerate(sd_model) - except Exception as e: - shared.log.error(f'Setting {op}: offload={shared.opts.diffusers_offload_mode} {e}') + apply_sequential_offload(sd_model, op=op, quiet=quiet) if shared.opts.diffusers_offload_mode == "group": sd_model = apply_group_offload(sd_model, op=op) @@ -254,10 +275,12 @@ class OffloadHook(accelerate.hooks.ModelHook): if debug: shared.log.trace(f'Offload: type=balanced op=dispatch map={device_map}') if device_map is not None: + skip_keys = getattr(module, "_skip_keys", None) module = accelerate.dispatch_model(module, main_device=torch.device(devices.device), device_map=device_map, offload_dir=offload_dir, + skip_keys=skip_keys, force_hooks=True, ) module._hf_hook.execution_device = torch.device(devices.device) # pylint: disable=protected-access diff --git a/modules/sd_samplers_common.py b/modules/sd_samplers_common.py index 8089ad8e1..3cdc91943 100644 --- a/modules/sd_samplers_common.py +++ b/modules/sd_samplers_common.py @@ -9,7 +9,7 @@ from modules import shared, devices, processing, images, sd_vae_approx, sd_vae_t SamplerData = namedtuple('SamplerData', ['name', 'constructor', 'aliases', 'options']) approximation_indexes = { "Simple": 0, "Approximate": 1, "TAESD": 2, "Full VAE": 3 } -flow_models = ['f1', 'f2', 'sd3', 'lumina', 'auraflow', 'sana', 'z_image', 'lumina2', 'cogview4', 'h1', 'cosmos', 'chroma', 'omnigen', 'omnigen2', 'longcat'] +flow_models = ['f1', 'f2', 'sd3', 'lumina', 'auraflow', 'sana', 'zimage', 'lumina2', 'cogview4', 'h1', 'cosmos', 'chroma', 'omnigen', 'omnigen2', 'longcat'] warned = False queue_lock = threading.Lock() diff --git a/modules/sd_vae.py b/modules/sd_vae.py index 246f52bb3..9c7255493 100644 --- a/modules/sd_vae.py +++ b/modules/sd_vae.py @@ -10,7 +10,7 @@ base_vae = None loaded_vae_file = None checkpoint_info = None vae_path = os.path.abspath(os.path.join(paths.models_path, 'VAE')) -debug = os.environ.get('SD_LOAD_DEBUG', None) is not None +debug = os.environ.get('SD_VAE_DEBUG', None) is not None unspecified = object() vae_scale_override = { 'WanPipeline': 16, diff --git a/modules/sd_vae_remote.py b/modules/sd_vae_remote.py index e7a179ec4..7a03bf383 100644 --- a/modules/sd_vae_remote.py +++ b/modules/sd_vae_remote.py @@ -1,4 +1,3 @@ -from typing import List import io import time import json @@ -21,7 +20,7 @@ hf_decode_endpoints['auraflow'] = hf_decode_endpoints['sdxl'] hf_decode_endpoints['omnigen'] = hf_decode_endpoints['sdxl'] hf_decode_endpoints['h1'] = hf_decode_endpoints['f1'] hf_decode_endpoints['chroma'] = hf_decode_endpoints['f1'] -hf_decode_endpoints['z_image'] = hf_decode_endpoints['f1'] +hf_decode_endpoints['zimage'] = hf_decode_endpoints['f1'] hf_decode_endpoints['lumina2'] = hf_decode_endpoints['f1'] hf_encode_endpoints = { @@ -36,7 +35,7 @@ hf_encode_endpoints['hunyuandit'] = hf_encode_endpoints['sdxl'] hf_encode_endpoints['auraflow'] = hf_encode_endpoints['sdxl'] hf_encode_endpoints['omnigen'] = hf_encode_endpoints['sdxl'] hf_encode_endpoints['h1'] = hf_encode_endpoints['f1'] -hf_encode_endpoints['z_image'] = hf_encode_endpoints['f1'] +hf_encode_endpoints['zimage'] = hf_encode_endpoints['f1'] hf_encode_endpoints['lumina2'] = hf_encode_endpoints['f1'] dtypes = { @@ -47,7 +46,7 @@ dtypes = { } -def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_type: str = None) -> Image.Image: +def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_type: str | None = None): from modules import devices, shared, errors, modelloader tensors = [] content = 0 @@ -93,7 +92,7 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_ params["output_type"] = "pt" params["output_tensor_type"] = "binary" headers["Accept"] = "tensor/binary" - if model_type in {'f1', 'h1', 'z_image', 'lumina2', 'chroma'} and (width > 0) and (height > 0): + if model_type in {'f1', 'h1', 'zimage', 'lumina2', 'chroma'} and (width > 0) and (height > 0): params['width'] = width params['height'] = height if shared.sd_model.vae is not None and shared.sd_model.vae.config is not None: @@ -128,7 +127,7 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_ return tensors -def remote_encode(images: List[Image.Image], model_type: str = None) -> torch.Tensor: +def remote_encode(images: list[Image.Image], model_type: str | None = None): from diffusers.utils import remote_utils from modules import devices, shared, errors, modelloader if not shared.opts.remote_vae_encode: diff --git a/modules/sd_vae_taesd.py b/modules/sd_vae_taesd.py index 00294d560..1c5dd0dd0 100644 --- a/modules/sd_vae_taesd.py +++ b/modules/sd_vae_taesd.py @@ -12,11 +12,17 @@ import torch from modules import devices, paths, shared +debug = os.environ.get('SD_PREVIEW_DEBUG', None) is not None + + TAESD_MODELS = { 'TAESD 1.3 Mocha Croissant': { 'fn': 'taesd_13_', 'uri': 'https://github.com/madebyollin/taesd/raw/7f572ca629c9b0d3c9f71140e5f501e09f9ea280', 'model': None }, 'TAESD 1.2 Chocolate-Dipped Shortbread': { 'fn': 'taesd_12_', 'uri': 'https://github.com/madebyollin/taesd/raw/8909b44e3befaa0efa79c5791e4fe1c4d4f7884e', 'model': None }, 'TAESD 1.1 Fruit Loops': { 'fn': 'taesd_11_', 'uri': 'https://github.com/madebyollin/taesd/raw/3e8a8a2ab4ad4079db60c1c7dc1379b4cc0c6b31', 'model': None }, 'TAESD 1.0': { 'fn': 'taesd_10_', 'uri': 'https://github.com/madebyollin/taesd/raw/88012e67cf0454e6d90f98911fe9d4aef62add86', 'model': None }, + 'TAE FLUX.1': { 'fn': 'taef1.pth', 'uri': 'https://github.com/madebyollin/taesd/raw/main/taef1_decoder.pth', 'model': None }, + 'TAE FLUX.2': { 'fn': 'taef2.pth', 'uri': 'https://github.com/madebyollin/taesd/raw/main/taef2_decoder.pth', 'model': None }, + 'TAE SD3': { 'fn': 'taesd3.pth', 'uri': 'https://github.com/madebyollin/taesd/raw/main/taesd3_decoder.pth', 'model': None }, 'TAE HunyuanVideo': { 'fn': 'taehv.pth', 'uri': 'https://github.com/madebyollin/taehv/raw/refs/heads/main/taehv.pth', 'model': None }, 'TAE WanVideo': { 'fn': 'taew1.pth', 'uri': 'https://github.com/madebyollin/taehv/raw/refs/heads/main/taew2_1.pth', 'model': None }, 'TAE MochiVideo': { 'fn': 'taem1.pth', 'uri': 'https://github.com/madebyollin/taem1/raw/refs/heads/main/taem1.pth', 'model': None }, @@ -38,7 +44,7 @@ prev_cls = '' prev_type = '' prev_model = '' lock = threading.Lock() -supported = ['sd', 'sdxl', 'sd3', 'f1', 'h1', 'z_image', 'lumina2', 'hunyuanvideo', 'wanai', 'chrono', 'mochivideo', 'pixartsigma', 'pixartalpha', 'hunyuandit', 'omnigen', 'qwen', 'longcat'] +supported = ['sd', 'sdxl', 'sd3', 'f1', 'f2', 'h1', 'zimage', 'lumina2', 'hunyuanvideo', 'wanai', 'chrono', 'cosmos', 'mochivideo', 'pixartsigma', 'pixartalpha', 'hunyuandit', 'omnigen', 'qwen', 'longcat', 'omnigen2', 'flite', 'ovis', 'kandinsky5', 'glmimage', 'cogview3', 'cogview4'] def warn_once(msg, variant=None): @@ -59,9 +65,15 @@ def get_model(model_type = 'decoder', variant = None): model_cls = 'sd' elif model_cls in {'pixartsigma', 'hunyuandit', 'omnigen', 'auraflow'}: model_cls = 'sdxl' - elif model_cls in {'h1', 'z_image', 'lumina2', 'chroma', 'longcat'}: + elif model_cls in {'f1', 'h1', 'zimage', 'lumina2', 'chroma', 'longcat', 'omnigen2', 'flite', 'ovis', 'kandinsky5', 'glmimage', 'cogview3', 'cogview4'}: model_cls = 'f1' - elif model_cls in {'wanai', 'qwen', 'chrono'}: + variant = 'TAE FLUX.1' + elif model_cls == 'f2': + model_cls = 'f2' + variant = 'TAE FLUX.2' + elif model_cls == 'sd3': + variant = 'TAE SD3' + elif model_cls in {'wanai', 'qwen', 'chrono', 'cosmos'}: variant = variant or 'TAE WanVideo' elif model_cls not in supported: warn_once(f'cls={shared.sd_model.__class__.__name__} type={model_cls} unsuppported', variant=variant) @@ -149,7 +161,13 @@ def decode(latents): dtype = devices.dtype_vae if devices.dtype_vae != torch.bfloat16 else torch.float16 # taesd does not support bf16 tensor = latents.unsqueeze(0) if len(latents.shape) == 3 else latents tensor = tensor.detach().clone().to(devices.device, dtype=dtype) - if variant.startswith('TAESD'): + if debug: + shared.log.debug(f'Decode: type="taesd" variant="{variant}" input={latents.shape} tensor={tensor.shape}') + # Fallback: reshape packed 128-channel latents to 32 channels if not already unpacked + if variant == 'TAE FLUX.2' and len(tensor.shape) == 4 and tensor.shape[1] == 128: + b, _c, h, w = tensor.shape + tensor = tensor.reshape(b, 32, h * 2, w * 2) + if variant.startswith('TAESD') or variant in {'TAE FLUX.1', 'TAE FLUX.2', 'TAE SD3'}: image = vae.decoder(tensor).clamp(0, 1).detach() image = image[0] else: diff --git a/modules/sdnq/common.py b/modules/sdnq/common.py index 4349bc307..3afb60a67 100644 --- a/modules/sdnq/common.py +++ b/modules/sdnq/common.py @@ -5,21 +5,25 @@ import torch from modules import shared, devices -sdnq_version = "0.1.3" +sdnq_version = "0.1.4" dtype_dict = { + ### Integers "int32": {"min": -2147483648, "max": 2147483647, "num_bits": 32, "sign": 1, "exponent": 0, "mantissa": 31, "target_dtype": torch.int32, "torch_dtype": torch.int32, "storage_dtype": torch.int32, "is_unsigned": False, "is_integer": True, "is_packed": False}, "int16": {"min": -32768, "max": 32767, "num_bits": 16, "sign": 1, "exponent": 0, "mantissa": 15, "target_dtype": torch.int16, "torch_dtype": torch.int16, "storage_dtype": torch.int16, "is_unsigned": False, "is_integer": True, "is_packed": False}, "int8": {"min": -128, "max": 127, "num_bits": 8, "sign": 1, "exponent": 0, "mantissa": 7, "target_dtype": torch.int8, "torch_dtype": torch.int8, "storage_dtype": torch.int8, "is_unsigned": False, "is_integer": True, "is_packed": False}, + ### Custom Integers "int7": {"min": -64, "max": 63, "num_bits": 7, "sign": 1, "exponent": 0, "mantissa": 6, "target_dtype": "int7", "torch_dtype": torch.int8, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": True, "is_packed": True}, "int6": {"min": -32, "max": 31, "num_bits": 6, "sign": 1, "exponent": 0, "mantissa": 5, "target_dtype": "int6", "torch_dtype": torch.int8, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": True, "is_packed": True}, "int5": {"min": -16, "max": 15, "num_bits": 5, "sign": 1, "exponent": 0, "mantissa": 4, "target_dtype": "int5", "torch_dtype": torch.int8, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": True, "is_packed": True}, "int4": {"min": -8, "max": 7, "num_bits": 4, "sign": 1, "exponent": 0, "mantissa": 3, "target_dtype": "int4", "torch_dtype": torch.int8, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": True, "is_packed": True}, "int3": {"min": -4, "max": 3, "num_bits": 3, "sign": 1, "exponent": 0, "mantissa": 2, "target_dtype": "int3", "torch_dtype": torch.int8, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": True, "is_packed": True}, "int2": {"min": -2, "max": 1, "num_bits": 2, "sign": 1, "exponent": 0, "mantissa": 1, "target_dtype": "int2", "torch_dtype": torch.int8, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": True, "is_packed": True}, + ### Unsigned Integers "uint32": {"min": 0, "max": 4294967295, "num_bits": 32, "sign": 0, "exponent": 0, "mantissa": 32, "target_dtype": torch.uint32, "torch_dtype": torch.uint32, "storage_dtype": torch.uint32, "is_unsigned": True, "is_integer": True, "is_packed": False}, "uint16": {"min": 0, "max": 65535, "num_bits": 16, "sign": 0, "exponent": 0, "mantissa": 16, "target_dtype": torch.uint16, "torch_dtype": torch.uint16, "storage_dtype": torch.uint16, "is_unsigned": True, "is_integer": True, "is_packed": False}, "uint8": {"min": 0, "max": 255, "num_bits": 8, "sign": 0, "exponent": 0, "mantissa": 8, "target_dtype": torch.uint8, "torch_dtype": torch.uint8, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": True, "is_packed": False}, + ### Custom Unsigned Integers "uint7": {"min": 0, "max": 127, "num_bits": 7, "sign": 0, "exponent": 0, "mantissa": 7, "target_dtype": "uint7", "torch_dtype": torch.uint8, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": True, "is_packed": True}, "uint6": {"min": 0, "max": 63, "num_bits": 6, "sign": 0, "exponent": 0, "mantissa": 6, "target_dtype": "uint6", "torch_dtype": torch.uint8, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": True, "is_packed": True}, "uint5": {"min": 0, "max": 31, "num_bits": 5, "sign": 0, "exponent": 0, "mantissa": 5, "target_dtype": "uint5", "torch_dtype": torch.uint8, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": True, "is_packed": True}, @@ -27,18 +31,107 @@ dtype_dict = { "uint3": {"min": 0, "max": 7, "num_bits": 3, "sign": 0, "exponent": 0, "mantissa": 3, "target_dtype": "uint3", "torch_dtype": torch.uint8, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": True, "is_packed": True}, "uint2": {"min": 0, "max": 3, "num_bits": 2, "sign": 0, "exponent": 0, "mantissa": 2, "target_dtype": "uint2", "torch_dtype": torch.uint8, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": True, "is_packed": True}, "uint1": {"min": 0, "max": 1, "num_bits": 1, "sign": 0, "exponent": 0, "mantissa": 1, "target_dtype": torch.bool, "torch_dtype": torch.bool, "storage_dtype": torch.bool, "is_unsigned": True, "is_integer": True, "is_packed": True}, + ### Floats "float32": {"min": -3.40282e+38, "max": 3.40282e+38, "num_bits": 32, "sign": 1, "exponent": 8, "mantissa": 23, "target_dtype": torch.float32, "torch_dtype": torch.float32, "storage_dtype": torch.float32, "is_unsigned": False, "is_integer": False, "is_packed": False}, "bfloat16": {"min": -3.38953e+38, "max": 3.38953e+38, "num_bits": 16, "sign": 1, "exponent": 8, "mantissa": 7, "target_dtype": torch.bfloat16, "torch_dtype": torch.bfloat16, "storage_dtype": torch.bfloat16, "is_unsigned": False, "is_integer": False, "is_packed": False}, - "float16": {"min": -65504, "max": 65504, "num_bits": 16, "sign": 1, "exponent": 5, "mantissa": 10, "target_dtype": torch.float16, "torch_dtype": torch.float16, "storage_dtype": torch.float16, "is_unsigned": False, "is_integer": False, "is_packed": False}, - "float8_e4m3fn": {"min": -448, "max": 448, "num_bits": 8, "sign": 1, "exponent": 4, "mantissa": 3, "target_dtype": torch.float8_e4m3fn, "torch_dtype": torch.float8_e4m3fn, "storage_dtype": torch.float8_e4m3fn, "is_unsigned": False, "is_integer": False, "is_packed": False}, - "float8_e5m2": {"min": -57344, "max": 57344, "num_bits": 8, "sign": 1, "exponent": 5, "mantissa": 2, "target_dtype": torch.float8_e5m2, "torch_dtype": torch.float8_e5m2, "storage_dtype": torch.float8_e5m2, "is_unsigned": False, "is_integer": False, "is_packed": False}, + "float16": {"min": -65504.0, "max": 65504.0, "num_bits": 16, "sign": 1, "exponent": 5, "mantissa": 10, "target_dtype": torch.float16, "torch_dtype": torch.float16, "storage_dtype": torch.float16, "is_unsigned": False, "is_integer": False, "is_packed": False}, + "float8_e4m3fn": {"min": -448.0, "max": 448.0, "num_bits": 8, "sign": 1, "exponent": 4, "mantissa": 3, "target_dtype": torch.float8_e4m3fn, "torch_dtype": torch.float8_e4m3fn, "storage_dtype": torch.float8_e4m3fn, "is_unsigned": False, "is_integer": False, "is_packed": False}, + "float8_e5m2": {"min": -57344.0, "max": 57344.0, "num_bits": 8, "sign": 1, "exponent": 5, "mantissa": 2, "target_dtype": torch.float8_e5m2, "torch_dtype": torch.float8_e5m2, "storage_dtype": torch.float8_e5m2, "is_unsigned": False, "is_integer": False, "is_packed": False}, + ### Custom Floats + "float16_e1m14fn": {"min": -3.9998779296875, "max": 3.9998779296875, "num_bits": 16, "sign": 1, "exponent": 1, "mantissa": 14, "min_normal": 1.00006103515625, "target_dtype": torch.float16, "torch_dtype": torch.float32, "storage_dtype": torch.uint16, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float16_e2m13fn": {"min": -7.99951171875, "max": 7.99951171875, "num_bits": 16, "sign": 1, "exponent": 2, "mantissa": 13, "min_normal": 0.50006103515625, "target_dtype": torch.float16, "torch_dtype": torch.float32, "storage_dtype": torch.uint16, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float16_e3m12fn": {"min": -31.99609375, "max": 31.99609375, "num_bits": 16, "sign": 1, "exponent": 3, "mantissa": 12, "min_normal": 0.125030517578125, "target_dtype": torch.float16, "torch_dtype": torch.float32, "storage_dtype": torch.uint16, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float16_e4m11fn": {"min": -511.875, "max": 511.875, "num_bits": 16, "sign": 1, "exponent": 4, "mantissa": 11, "min_normal": 0.007816314697265625, "target_dtype": torch.float16, "torch_dtype": torch.float32, "storage_dtype": torch.uint16, "is_unsigned": False, "is_integer": False, "is_packed": True}, + # float16_e5m10 is native in PyTorch + "float8_e1m6fn": {"min": -3.96875, "max": 3.96875, "num_bits": 8, "sign": 1, "exponent": 1, "mantissa": 6, "min_normal": 1.015625, "target_dtype": "fp8", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float8_e2m5fn": {"min": -7.875, "max": 7.875, "num_bits": 8, "sign": 1, "exponent": 2, "mantissa": 5, "min_normal": 0.515625, "target_dtype": "fp8", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float8_e3m4fn": {"min": -31.0, "max": 31.0, "num_bits": 8, "sign": 1, "exponent": 3, "mantissa": 4, "min_normal": 0.1328125, "target_dtype": "fp8", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + # float8_e4m3fn is native in PyTorch + # float8_e5m2fn is native in PyTorch + "float7_e1m5fn": {"min": -3.9375, "max": 3.9375, "num_bits": 7, "sign": 1, "exponent": 1, "mantissa": 5, "min_normal": 1.03125, "target_dtype": "fp7", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float7_e2m4fn": {"min": -7.75, "max": 7.75, "num_bits": 7, "sign": 1, "exponent": 2, "mantissa": 4, "min_normal": 0.53125, "target_dtype": "fp7", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float7_e3m3fn": {"min": -30.0, "max": 30.0, "num_bits": 7, "sign": 1, "exponent": 3, "mantissa": 3, "min_normal": 0.140625, "target_dtype": "fp7", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float7_e4m2fn": {"min": -448.0, "max": 448.0, "num_bits": 7, "sign": 1, "exponent": 4, "mantissa": 2, "min_normal": 0.009765625, "target_dtype": "fp7", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float7_e5m1fn": {"min": -98304.0, "max": 98304.0, "num_bits": 7, "sign": 1, "exponent": 5, "mantissa": 1, "min_normal": 4.57763671875e-05, "target_dtype": "fp7", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + # + "float6_e1m4fn": {"min": -3.875, "max": 3.875, "num_bits": 6, "sign": 1, "exponent": 1, "mantissa": 4, "min_normal": 1.0625, "target_dtype": "fp6", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float6_e2m3fn": {"min": -7.5, "max": 7.5, "num_bits": 6, "sign": 1, "exponent": 2, "mantissa": 3, "min_normal": 0.5625, "target_dtype": "fp6", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float6_e3m2fn": {"min": -28.0, "max": 28.0, "num_bits": 6, "sign": 1, "exponent": 3, "mantissa": 2, "min_normal": 0.15625, "target_dtype": "fp6", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float6_e4m1fn": {"min": -384.0, "max": 384.0, "num_bits": 6, "sign": 1, "exponent": 4, "mantissa": 1, "min_normal": 0.01171875, "target_dtype": "fp6", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float6_e5m0fn": {"min": -65536.0, "max": 65536.0, "num_bits": 6, "sign": 1, "exponent": 5, "mantissa": 0, "min_normal": 6.103515625e-05, "target_dtype": "fp6", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + # + "float5_e1m3fn": {"min": -3.75, "max": 3.75, "num_bits": 5, "sign": 1, "exponent": 1, "mantissa": 3, "min_normal": 1.125, "target_dtype": "fp5", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float5_e2m2fn": {"min": -7.0, "max": 7.0, "num_bits": 5, "sign": 1, "exponent": 2, "mantissa": 2, "min_normal": 0.625, "target_dtype": "fp5", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float5_e3m1fn": {"min": -24.0, "max": 24.0, "num_bits": 5, "sign": 1, "exponent": 3, "mantissa": 1, "min_normal": 0.1875, "target_dtype": "fp5", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float5_e4m0fn": {"min": -256.0, "max": 256.0, "num_bits": 5, "sign": 1, "exponent": 4, "mantissa": 0, "min_normal": 0.015625, "target_dtype": "fp5", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + # + "float4_e1m2fn": {"min": -3.5, "max": 3.5, "num_bits": 4, "sign": 1, "exponent": 1, "mantissa": 2, "min_normal": 1.25, "target_dtype": "fp4", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float4_e2m1fn": {"min": -6.0, "max": 6.0, "num_bits": 4, "sign": 1, "exponent": 2, "mantissa": 1, "min_normal": 0.75, "target_dtype": "fp4", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float4_e3m0fn": {"min": -16.0, "max": 16.0, "num_bits": 4, "sign": 1, "exponent": 3, "mantissa": 0, "min_normal": 0.25, "target_dtype": "fp4", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + # + "float3_e1m1fn": {"min": -3.0, "max": 3.0, "num_bits": 3, "sign": 1, "exponent": 1, "mantissa": 1, "min_normal": 1.5, "target_dtype": "fp3", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + "float3_e2m0fn": {"min": -4.0, "max": 4.0, "num_bits": 3, "sign": 1, "exponent": 2, "mantissa": 0, "min_normal": 1.0, "target_dtype": "fp3", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + # + "float2_e1m0fn": {"min": -2.0, "max": 2.0, "num_bits": 2, "sign": 1, "exponent": 1, "mantissa": 0, "min_normal": 2.0, "target_dtype": "fp2", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": False, "is_integer": False, "is_packed": True}, + ### Custom Usigned Floats + "float16_e1m15fnu": {"min": 0, "max": 3.99993896484375, "num_bits": 16, "sign": 0, "exponent": 1, "mantissa": 15, "min_normal": 1.000030517578125, "target_dtype": torch.float16, "torch_dtype": torch.float32, "storage_dtype": torch.uint16, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float16_e2m14fnu": {"min": 0, "max": 7.999755859375, "num_bits": 16, "sign": 0, "exponent": 2, "mantissa": 14, "min_normal": 0.500030517578125, "target_dtype": torch.float16, "torch_dtype": torch.float32, "storage_dtype": torch.uint16, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float16_e3m13fnu": {"min": 0, "max": 31.998046875, "num_bits": 16, "sign": 0, "exponent": 3, "mantissa": 13, "min_normal": 0.1250152587890625, "target_dtype": torch.float16, "torch_dtype": torch.float32, "storage_dtype": torch.uint16, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float16_e4m12fnu": {"min": 0, "max": 511.9375, "num_bits": 16, "sign": 0, "exponent": 4, "mantissa": 12, "min_normal": 0.007814407348632812, "target_dtype": torch.float16, "torch_dtype": torch.float32, "storage_dtype": torch.uint16, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float16_e5m11fnu": {"min": 0, "max": 131040.0, "num_bits": 16, "sign": 0, "exponent": 5, "mantissa": 11, "min_normal": 3.053247928619385e-05, "target_dtype": torch.float16, "torch_dtype": torch.float32, "storage_dtype": torch.uint16, "is_unsigned": True, "is_integer": False, "is_packed": True}, + # + "float8_e1m7fnu": {"min": 0, "max": 3.984375, "num_bits": 8, "sign": 0, "exponent": 1, "mantissa": 7, "min_normal": 1.0078125, "target_dtype": "fp8", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float8_e2m6fnu": {"min": 0, "max": 7.9375, "num_bits": 8, "sign": 0, "exponent": 2, "mantissa": 6, "min_normal": 0.5078125, "target_dtype": "fp8", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float8_e3m5fnu": {"min": 0, "max": 31.5, "num_bits": 8, "sign": 0, "exponent": 3, "mantissa": 5, "min_normal": 0.12890625, "target_dtype": "fp8", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float8_e4m4fnu": {"min": 0, "max": 496.0, "num_bits": 8, "sign": 0, "exponent": 4, "mantissa": 4, "min_normal": 0.00830078125, "target_dtype": "fp8", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float8_e5m3fnu": {"min": 0, "max": 122880.0, "num_bits": 8, "sign": 0, "exponent": 5, "mantissa": 3, "min_normal": 3.4332275390625e-05, "target_dtype": "fp8", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + # + "float7_e1m6fnu": {"min": 0, "max": 3.96875, "num_bits": 7, "sign": 0, "exponent": 1, "mantissa": 6, "min_normal": 1.015625, "target_dtype": "fp7", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float7_e2m5fnu": {"min": 0, "max": 7.875, "num_bits": 7, "sign": 0, "exponent": 2, "mantissa": 5, "min_normal": 0.515625, "target_dtype": "fp7", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float7_e3m4fnu": {"min": 0, "max": 31.0, "num_bits": 7, "sign": 0, "exponent": 3, "mantissa": 4, "min_normal": 0.1328125, "target_dtype": "fp7", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float7_e4m3fnu": {"min": 0, "max": 480.0, "num_bits": 7, "sign": 0, "exponent": 4, "mantissa": 3, "min_normal": 0.0087890625, "target_dtype": "fp7", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float7_e5m2fnu": {"min": 0, "max": 114688.0, "num_bits": 7, "sign": 0, "exponent": 5, "mantissa": 2, "min_normal": 3.814697265625e-05, "target_dtype": "fp7", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + # + "float6_e1m5fnu": {"min": 0, "max": 3.9375, "num_bits": 6, "sign": 0, "exponent": 1, "mantissa": 5, "min_normal": 1.03125, "target_dtype": "fp6", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float6_e2m4fnu": {"min": 0, "max": 7.75, "num_bits": 6, "sign": 0, "exponent": 2, "mantissa": 4, "min_normal": 0.53125, "target_dtype": "fp6", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float6_e3m3fnu": {"min": 0, "max": 30.0, "num_bits": 6, "sign": 0, "exponent": 3, "mantissa": 3, "min_normal": 0.140625, "target_dtype": "fp6", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float6_e4m2fnu": {"min": 0, "max": 448.0, "num_bits": 6, "sign": 0, "exponent": 4, "mantissa": 2, "min_normal": 0.009765625, "target_dtype": "fp6", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float6_e5m1fnu": {"min": 0, "max": 98304.0, "num_bits": 6, "sign": 0, "exponent": 5, "mantissa": 1, "min_normal": 4.57763671875e-05, "target_dtype": "fp6", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + # + "float5_e1m4fnu": {"min": 0, "max": 3.875, "num_bits": 5, "sign": 0, "exponent": 1, "mantissa": 4, "min_normal": 1.0625, "target_dtype": "fp5", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float5_e2m3fnu": {"min": 0, "max": 7.5, "num_bits": 5, "sign": 0, "exponent": 2, "mantissa": 3, "min_normal": 0.5625, "target_dtype": "fp5", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float5_e3m2fnu": {"min": 0, "max": 28.0, "num_bits": 5, "sign": 0, "exponent": 3, "mantissa": 2, "min_normal": 0.15625, "target_dtype": "fp5", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float5_e4m1fnu": {"min": 0, "max": 384.0, "num_bits": 5, "sign": 0, "exponent": 4, "mantissa": 1, "min_normal": 0.01171875, "target_dtype": "fp5", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float5_e5m0fnu": {"min": 0, "max": 65536.0, "num_bits": 5, "sign": 0, "exponent": 5, "mantissa": 0, "min_normal": 6.103515625e-05, "target_dtype": "fp5", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + # + "float4_e1m3fnu": {"min": 0, "max": 3.75, "num_bits": 4, "sign": 0, "exponent": 1, "mantissa": 3, "min_normal": 1.125, "target_dtype": "fp4", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float4_e2m2fnu": {"min": 0, "max": 7.0, "num_bits": 4, "sign": 0, "exponent": 2, "mantissa": 2, "min_normal": 0.625, "target_dtype": "fp4", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float4_e3m1fnu": {"min": 0, "max": 24.0, "num_bits": 4, "sign": 0, "exponent": 3, "mantissa": 1, "min_normal": 0.1875, "target_dtype": "fp4", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float4_e4m0fnu": {"min": 0, "max": 256.0, "num_bits": 4, "sign": 0, "exponent": 4, "mantissa": 0, "min_normal": 0.015625, "target_dtype": "fp4", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + # + "float3_e1m2fnu": {"min": 0, "max": 3.5, "num_bits": 3, "sign": 0, "exponent": 1, "mantissa": 2, "min_normal": 1.25, "target_dtype": "fp3", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float3_e2m1fnu": {"min": 0, "max": 6.0, "num_bits": 3, "sign": 0, "exponent": 2, "mantissa": 1, "min_normal": 0.75, "target_dtype": "fp3", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float3_e3m0fnu": {"min": 0, "max": 16.0, "num_bits": 3, "sign": 0, "exponent": 3, "mantissa": 0, "min_normal": 0.25, "target_dtype": "fp3", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + # + "float2_e1m1fnu": {"min": 0, "max": 3.0, "num_bits": 2, "sign": 0, "exponent": 1, "mantissa": 1, "min_normal": 1.5, "target_dtype": "fp2", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + "float2_e2m0fnu": {"min": 0, "max": 4.0, "num_bits": 2, "sign": 0, "exponent": 2, "mantissa": 0, "min_normal": 1.0, "target_dtype": "fp2", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, + # + "float1_e1m0fnu": {"min": 0, "max": 2.0, "num_bits": 1, "sign": 0, "exponent": 1, "mantissa": 0, "min_normal": 2.0, "target_dtype": "fp1", "torch_dtype": torch.float32, "storage_dtype": torch.uint8, "is_unsigned": True, "is_integer": False, "is_packed": True}, } dtype_dict["fp32"] = dtype_dict["float32"] dtype_dict["bf16"] = dtype_dict["bfloat16"] dtype_dict["fp16"] = dtype_dict["float16"] dtype_dict["fp8"] = dtype_dict["float8_e4m3fn"] +dtype_dict["fp7"] = dtype_dict["float7_e3m3fn"] +dtype_dict["fp6"] = dtype_dict["float6_e3m2fn"] +dtype_dict["fp5"] = dtype_dict["float5_e2m2fn"] +dtype_dict["fp4"] = dtype_dict["float4_e2m1fn"] +dtype_dict["fp3"] = dtype_dict["float3_e1m1fn"] +dtype_dict["fp2"] = dtype_dict["float2_e1m0fn"] +dtype_dict["fp1"] = dtype_dict["float1_e1m0fnu"] dtype_dict["bool"] = dtype_dict["uint1"] +dtype_dict["int1"] = dtype_dict["uint1"] torch_dtype_dict = { torch.int32: "int32", @@ -55,20 +148,43 @@ torch_dtype_dict = { } if hasattr(torch, "float8_e4m3fnuz"): - dtype_dict["float8_e4m3fnuz"] = {"min": -240, "max": 240, "num_bits": 8, "sign": 1, "exponent": 4, "mantissa": 3, "target_dtype": "fp8", "torch_dtype": torch.float8_e4m3fnuz, "storage_dtype": torch.float8_e4m3fnuz, "is_unsigned": False, "is_integer": False, "is_packed": False} + dtype_dict["float8_e4m3fnuz"] = {"min": -240.0, "max": 240.0, "num_bits": 8, "sign": 1, "exponent": 4, "mantissa": 3, "target_dtype": "fp8", "torch_dtype": torch.float8_e4m3fnuz, "storage_dtype": torch.float8_e4m3fnuz, "is_unsigned": False, "is_integer": False, "is_packed": False} torch_dtype_dict[torch.float8_e4m3fnuz] = "float8_e4m3fnuz" if hasattr(torch, "float8_e5m2fnuz"): - dtype_dict["float8_e5m2fnuz"] = {"min": -57344, "max": 57344, "num_bits": 8, "sign": 1, "exponent": 5, "mantissa": 2, "target_dtype": "fp8", "torch_dtype": torch.float8_e5m2fnuz, "storage_dtype": torch.float8_e5m2fnuz, "is_unsigned": False, "is_integer": False, "is_packed": False} + dtype_dict["float8_e5m2fnuz"] = {"min": -57344.0, "max": 57344.0, "num_bits": 8, "sign": 1, "exponent": 5, "mantissa": 2, "target_dtype": "fp8", "torch_dtype": torch.float8_e5m2fnuz, "storage_dtype": torch.float8_e5m2fnuz, "is_unsigned": False, "is_integer": False, "is_packed": False} torch_dtype_dict[torch.float8_e5m2fnuz] = "float8_e5m2fnuz" -linear_types = {"Linear"} -conv_types = {"Conv1d", "Conv2d", "Conv3d"} -conv_transpose_types = {"ConvTranspose1d", "ConvTranspose2d", "ConvTranspose3d"} +linear_types = {"Linear", "SDNQLinear"} +conv_types = {"Conv1d", "Conv2d", "Conv3d", "SDNQConv1d", "SDNQConv2d", "SDNQConv3d"} +conv_transpose_types = {"ConvTranspose1d", "ConvTranspose2d", "ConvTranspose3d", "SDNQConvTranspose1d", "SDNQConvTranspose2d", "SDNQConvTranspose3d"} allowed_types = set.union(linear_types, conv_types, conv_transpose_types) -accepted_weight_dtypes = set(dtype_dict.keys()) -accepted_matmul_dtypes = {"int8", "fp8", "fp16", "float8_e4m3fnuz", "float16"} -is_rdna2 = bool(devices.backend == "rocm" and int(getattr(torch.cuda.get_device_properties(devices.device), "gcnArchName", "gfx0000")[3:]) < 1100) +accepted_weight_dtypes = set(dtype_dict.keys()) +accepted_matmul_dtypes = {"int8", "fp8", "fp16", "float8_e4m3fn", "float16"} + +weights_dtype_order = [ + "uint1", "float1_e1m0fnu", + "int2", "float2_e1m0fn", + "uint2", "float2_e1m1fnu", "float2_e2m0fnu", + "int3", "float3_e1m1fn", "float3_e2m0fn", + "uint3", "float3_e1m2fnu", "float3_e2m1fnu", "float3_e3m0fnu", + "int4", "float4_e1m2fn", "float4_e2m1fn", "float4_e3m0fn", + "uint4", "float4_e1m3fnu", "float4_e2m2fnu", "float4_e3m1fnu", "float4_e4m0fnu", + "int5", "float5_e1m3fn", "float5_e2m2fn", "float5_e3m1fn", "float5_e4m0fn", + "uint5", "float5_e1m4fnu", "float5_e2m3fnu", "float5_e3m2fnu", "float5_e4m1fnu", "float5_e5m0fnu", + "int6", "float6_e1m4fn", "float6_e2m3fn", "float6_e3m2fn", "float6_e4m1fn", "float6_e5m0fn", + "uint6", "float6_e1m5fnu", "float6_e2m4fnu", "float6_e3m3fnu", "float6_e4m2fnu", "float6_e5m1fnu", + "int7", "float7_e1m5fn", "float7_e2m4fn", "float7_e3m3fn", "float7_e4m2fn", "float7_e5m1fn", + "uint7", "float7_e1m6fnu", "float7_e2m5fnu", "float7_e3m4fnu", "float7_e4m3fnu", "float7_e5m2fnu", + "int8", "float8_e4m3fn", "float8_e5m2", "float8_e1m6fn", "float8_e2m5fn", "float8_e3m4fn", + "uint8", "float8_e1m7fnu", "float8_e2m6fnu", "float8_e3m5fnu", "float8_e4m4fnu", "float8_e5m3fnu", +] +weights_dtype_order_fp32 = weights_dtype_order + [ + "int16", "float16", "float16_e1m14fn", "float16_e2m13fn", "float16_e3m12fn", "float16_e4m11fn", + "uint16", "float16_e1m15fnu", "float16_e2m14fnu", "float16_e3m13fnu", "float16_e4m12fnu", "float16_e5m11fnu", +] + +is_rdna2 = bool(devices.backend == "rocm" and devices.get_hip_agent().gfx_version < 0x1100) use_torch_compile = shared.opts.sdnq_dequantize_compile # this setting requires a full restart of the webui to apply def check_torch_compile(): # dynamo can be disabled after startup @@ -118,8 +234,8 @@ if fp_mm_func is None: if use_torch_compile: - torch._dynamo.config.cache_size_limit = max(8192, torch._dynamo.config.cache_size_limit) - torch._dynamo.config.accumulated_recompile_limit = max(8192, torch._dynamo.config.accumulated_recompile_limit) + torch._dynamo.config.cache_size_limit = max(8192, getattr(torch._dynamo.config, "cache_size_limit", 0)) + torch._dynamo.config.accumulated_recompile_limit = max(8192, getattr(torch._dynamo.config, "accumulated_recompile_limit", 0)) def compile_func(fn, **kwargs): if kwargs.get("fullgraph", None) is None: kwargs["fullgraph"] = True @@ -183,6 +299,13 @@ module_skip_keys_dict = { ["blocks.0.adaLN_modulation.1.weight", "x_embedder", "t_embedder", "y_embedder", "final_layer"], {} ], + "LTX2VideoTransformer3DModel": [ + [ + "audio_time_embed", "time_embed", "audio_caption_projection", "caption_projection", "proj_in", "audio_proj_in", "proj_out", "audio_proj_out", + "av_cross_attn_audio_scale_shift", "av_cross_attn_audio_v2a_gate", "av_cross_attn_video_a2v_gate", "av_cross_attn_video_scale_shift", + ], + {} + ], "Lumina2Transformer2DModel": [ ["layers.0.norm1.linear.weight", "time_caption_embed", "x_embedder", "norm_out"], {} @@ -191,6 +314,14 @@ module_skip_keys_dict = { ["layers.0.adaLN_modulation.0.weight", "t_embedder", "cap_embedder", "siglip_embedder", "all_x_embedder", "all_final_layer"], {} ], + "GlmImageTransformer2DModel": [ + ["transformer_blocks.0.norm1.linear.weight", "image_projector", "glyph_projector", "prior_projector", "time_condition_embed", "norm_out", "proj_out"], + {} + ], + "GlmImageForConditionalGeneration": [ + ["lm_head", "patch_embed", "embeddings", "embed_tokens", "vqmodel"], + {} + ], "HunyuanImage3ForCausalMM": [ ["lm_head", "patch_embed", "time_embed", "time_embed_2", "final_layer", "wte", "ln_f", "timestep_emb", "vae", "vision_aligner", "head", "post_layernorm", "embeddings"], {} diff --git a/modules/sdnq/dequantizer.py b/modules/sdnq/dequantizer.py index a6ddcc7c5..b298ac1a2 100644 --- a/modules/sdnq/dequantizer.py +++ b/modules/sdnq/dequantizer.py @@ -8,6 +8,7 @@ import torch from modules import devices from .common import dtype_dict, compile_func, use_contiguous_mm, use_tensorwise_fp8_matmul from .packed_int import unpack_int_symetric, unpack_int_asymetric +from .packed_float import unpack_float @devices.inference_context() @@ -25,7 +26,7 @@ def dequantize_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, ze if result.ndim > 2 and weight.ndim > 2: # convs result = result.add_(torch.mm(svd_up, svd_down).unflatten(-1, (*result.shape[1:],))) else: - result = result.addmm_(svd_up, svd_down) + result = result.to(dtype=svd_up.dtype).addmm_(svd_up, svd_down) if dtype is not None: result = result.to(dtype=dtype) return result @@ -48,7 +49,7 @@ def dequantize_symmetric(weight: torch.CharTensor, scale: torch.FloatTensor, svd if result.ndim > 2 and weight.ndim > 2: # convs result = result.add_(torch.mm(svd_up, svd_down).unflatten(-1, (*result.shape[1:],))) else: - result = result.addmm_(svd_up, svd_down) + result = result.to(dtype=svd_up.dtype).addmm_(svd_up, svd_down) if dtype is not None: result = result.to(dtype=dtype) return result @@ -70,10 +71,20 @@ def dequantize_packed_int_asymmetric(weight: torch.ByteTensor, scale: torch.Floa @devices.inference_context() -def dequantize_packed_int_symmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, dtype: Optional[torch.dtype] = None, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None, result_shape: Optional[torch.Size] = None, skip_quantized_matmul: bool = False, re_quantize_for_matmul: bool = False) -> torch.FloatTensor: +def dequantize_packed_int_symmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None, dtype: Optional[torch.dtype] = None, result_shape: Optional[torch.Size] = None, skip_quantized_matmul: bool = False, re_quantize_for_matmul: bool = False) -> torch.FloatTensor: return dequantize_symmetric(unpack_int_symetric(weight, shape, weights_dtype, dtype=scale.dtype), scale, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul) +@devices.inference_context() +def dequantize_packed_float_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, weights_dtype: str, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None, dtype: Optional[torch.dtype] = None, result_shape: Optional[torch.Size] = None, skip_quantized_matmul: bool = False) -> torch.FloatTensor: + return dequantize_asymmetric(unpack_float(weight, shape, weights_dtype), scale, zero_point, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=result_shape, skip_quantized_matmul=skip_quantized_matmul) + + +@devices.inference_context() +def dequantize_packed_float_symmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None, dtype: Optional[torch.dtype] = None, result_shape: Optional[torch.Size] = None, skip_quantized_matmul: bool = False, re_quantize_for_matmul: bool = False) -> torch.FloatTensor: + return dequantize_symmetric(unpack_float(weight, shape, weights_dtype), scale, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul) + + @devices.inference_context() def quantize_int_mm(input: torch.FloatTensor, dim: int = -1, matmul_dtype: str = "int8") -> Tuple[torch.Tensor, torch.FloatTensor]: scale = torch.amax(input.abs(), dim=dim, keepdims=True).div_(dtype_dict[matmul_dtype]["max"]) @@ -84,7 +95,7 @@ def quantize_int_mm(input: torch.FloatTensor, dim: int = -1, matmul_dtype: str = @devices.inference_context() def quantize_int_mm_sr(input: torch.FloatTensor, dim: int = -1, matmul_dtype: str = "int8") -> Tuple[torch.Tensor, torch.FloatTensor]: scale = torch.amax(input.abs(), dim=dim, keepdims=True).div_(dtype_dict[matmul_dtype]["max"]) - input = torch.div(input, scale).add_(torch.randn_like(input), alpha=0.1).round_().clamp_(dtype_dict[matmul_dtype]["min"], dtype_dict[matmul_dtype]["max"]).to(dtype=dtype_dict[matmul_dtype]["torch_dtype"]) + input = torch.div(input, scale).add_(torch.rand_like(input), alpha=0.1).round_().clamp_(dtype_dict[matmul_dtype]["min"], dtype_dict[matmul_dtype]["max"]).to(dtype=dtype_dict[matmul_dtype]["torch_dtype"]) return input, scale @@ -156,6 +167,16 @@ def re_quantize_matmul_packed_int_symmetric(weight: torch.ByteTensor, scale: tor return re_quantize_matmul_symmetric(unpack_int_symetric(weight, shape, weights_dtype, dtype=scale.dtype), scale, matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=result_shape) +@devices.inference_context() +def re_quantize_matmul_packed_float_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, weights_dtype: str, matmul_dtype: str, result_shape: torch.Size, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None) -> Tuple[torch.Tensor, torch.FloatTensor]: + return re_quantize_matmul_asymmetric(unpack_float(weight, shape, weights_dtype), scale, zero_point, matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=result_shape) + + +@devices.inference_context() +def re_quantize_matmul_packed_float_symmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, matmul_dtype: str, result_shape: Optional[torch.Size] = None, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None) -> Tuple[torch.Tensor, torch.FloatTensor]: + return re_quantize_matmul_symmetric(unpack_float(weight, shape, weights_dtype), scale, matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=result_shape) + + @devices.inference_context() def dequantize_layer_weight(self: torch.nn.Module, inplace: bool = False): weight = torch.nn.Parameter(self.sdnq_dequantizer(self.weight, self.scale, self.zero_point, self.svd_up, self.svd_down, skip_quantized_matmul=self.sdnq_dequantizer.use_quantized_matmul), requires_grad=True) @@ -209,6 +230,25 @@ def dequantize_sdnq_model(model: torch.nn.Module): # SDNQDequantizer has to be a dataclass for torch.compile @dataclass class SDNQDequantizer: + result_dtype: torch.dtype + result_shape: torch.Size + original_shape: torch.Size + original_stride: List[int] + quantized_weight_shape: torch.Size + weights_dtype: str + quantized_matmul_dtype: str + group_size: int + svd_rank: int + svd_steps: int + use_quantized_matmul: bool + re_quantize_for_matmul: bool + use_stochastic_rounding: bool + layer_class_name: str + is_packed: bool + is_unsigned: bool + is_integer: bool + is_integer_matmul: bool + def __init__( self, result_dtype: torch.dtype, @@ -226,10 +266,6 @@ class SDNQDequantizer: use_stochastic_rounding: bool, layer_class_name: str, ): - self.is_packed = dtype_dict[weights_dtype]["is_packed"] - self.is_unsigned = dtype_dict[weights_dtype]["is_unsigned"] - self.is_integer = dtype_dict[weights_dtype]["is_integer"] - self.is_integer_matmul = dtype_dict[quantized_matmul_dtype]["is_integer"] self.result_dtype = result_dtype self.result_shape = result_shape self.original_shape = original_shape @@ -244,14 +280,24 @@ class SDNQDequantizer: self.re_quantize_for_matmul = re_quantize_for_matmul self.use_stochastic_rounding = use_stochastic_rounding self.layer_class_name = layer_class_name + self.is_packed = dtype_dict[weights_dtype]["is_packed"] + self.is_unsigned = dtype_dict[weights_dtype]["is_unsigned"] + self.is_integer = dtype_dict[weights_dtype]["is_integer"] + self.is_integer_matmul = dtype_dict[quantized_matmul_dtype]["is_integer"] @devices.inference_context() def re_quantize_matmul(self, weight, scale, zero_point, svd_up, svd_down): # pylint: disable=unused-argument if self.is_packed: - if self.is_unsigned: - return re_quantize_matmul_packed_int_asymmetric_compiled(weight, scale, zero_point, self.quantized_weight_shape, self.weights_dtype, self.quantized_matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape) + if self.is_integer: + if self.is_unsigned: + return re_quantize_matmul_packed_int_asymmetric_compiled(weight, scale, zero_point, self.quantized_weight_shape, self.weights_dtype, self.quantized_matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape) + else: + return re_quantize_matmul_packed_int_symmetric_compiled(weight, scale, self.quantized_weight_shape, self.weights_dtype, self.quantized_matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape) else: - return re_quantize_matmul_packed_int_symmetric_compiled(weight, scale, self.quantized_weight_shape, self.weights_dtype, self.quantized_matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape) + if self.is_unsigned: + return re_quantize_matmul_packed_float_asymmetric_compiled(weight, scale, zero_point, self.quantized_weight_shape, self.weights_dtype, self.quantized_matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape) + else: + return re_quantize_matmul_packed_float_symmetric_compiled(weight, scale, self.quantized_weight_shape, self.weights_dtype, self.quantized_matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape) else: if self.is_unsigned: return re_quantize_matmul_asymmetric_compiled(weight, scale, zero_point, self.quantized_matmul_dtype, svd_up=svd_up, svd_down=svd_down, result_shape=self.result_shape) @@ -262,17 +308,30 @@ class SDNQDequantizer: def __call__(self, weight, scale, zero_point, svd_up, svd_down, skip_quantized_matmul: bool = False, skip_compile: bool = False, dtype: torch.dtype = None): # pylint: disable=unused-argument if dtype is None: dtype = self.result_dtype + re_quantize_for_matmul = self.re_quantize_for_matmul or self.is_packed if self.is_packed: - if self.is_unsigned: - if skip_compile: # compiled training needs to be traced with the original function - return dequantize_packed_int_asymmetric(weight, scale, zero_point, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul) + if self.is_integer: + if self.is_unsigned: + if skip_compile: # compiled training needs to be traced with the original function + return dequantize_packed_int_asymmetric(weight, scale, zero_point, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul) + else: + return dequantize_packed_int_asymmetric_compiled(weight, scale, zero_point, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul) else: - return dequantize_packed_int_asymmetric_compiled(weight, scale, zero_point, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul) + if skip_compile: + return dequantize_packed_int_symmetric(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul) + else: + return dequantize_packed_int_symmetric_compiled(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul) else: - if skip_compile: - return dequantize_packed_int_symmetric(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=self.re_quantize_for_matmul) + if self.is_unsigned: + if skip_compile: # compiled training needs to be traced with the original function + return dequantize_packed_float_asymmetric(weight, scale, zero_point, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul) + else: + return dequantize_packed_float_asymmetric_compiled(weight, scale, zero_point, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul) else: - return dequantize_packed_int_symmetric_compiled(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=self.re_quantize_for_matmul) + if skip_compile: + return dequantize_packed_float_symmetric(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul) + else: + return dequantize_packed_float_symmetric_compiled(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul) else: if self.is_unsigned: if skip_compile: @@ -281,16 +340,22 @@ class SDNQDequantizer: return dequantize_asymmetric_compiled(weight, scale, zero_point, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul) else: if skip_compile: - return dequantize_symmetric(weight, scale, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=self.re_quantize_for_matmul) + return dequantize_symmetric(weight, scale, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul) else: - return dequantize_symmetric_compiled(weight, scale, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=self.re_quantize_for_matmul) + return dequantize_symmetric_compiled(weight, scale, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul) dequantize_asymmetric_compiled = compile_func(dequantize_asymmetric) dequantize_symmetric_compiled = compile_func(dequantize_symmetric) dequantize_packed_int_asymmetric_compiled = compile_func(dequantize_packed_int_asymmetric) dequantize_packed_int_symmetric_compiled = compile_func(dequantize_packed_int_symmetric) +dequantize_packed_float_asymmetric_compiled = compile_func(dequantize_packed_float_asymmetric) +dequantize_packed_float_symmetric_compiled = compile_func(dequantize_packed_float_symmetric) re_quantize_matmul_asymmetric_compiled = compile_func(re_quantize_matmul_asymmetric) re_quantize_matmul_symmetric_compiled = compile_func(re_quantize_matmul_symmetric) re_quantize_matmul_packed_int_asymmetric_compiled = compile_func(re_quantize_matmul_packed_int_asymmetric) re_quantize_matmul_packed_int_symmetric_compiled = compile_func(re_quantize_matmul_packed_int_symmetric) +re_quantize_matmul_packed_float_asymmetric_compiled = compile_func(re_quantize_matmul_packed_float_asymmetric) +re_quantize_matmul_packed_float_symmetric_compiled = compile_func(re_quantize_matmul_packed_float_symmetric) + +torch.serialization.add_safe_globals([SDNQDequantizer]) diff --git a/modules/sdnq/layers/__init__.py b/modules/sdnq/layers/__init__.py new file mode 100644 index 000000000..47f4420af --- /dev/null +++ b/modules/sdnq/layers/__init__.py @@ -0,0 +1,69 @@ +import torch + + +class SDNQLayer(torch.nn.Module): + def __init__(self, original_layer, forward_func): + torch.nn.Module.__init__(self) + for key, value in original_layer.__dict__.items(): + if key not in {"forward", "forward_func", "original_class"}: + setattr(self, key, value) + self.original_class = original_layer.__class__ + self.forward_func = forward_func + + def forward(self, *args, **kwargs) -> torch.Tensor: + return self.forward_func(self, *args, **kwargs) + + def __repr__(self): + return f"{self.__class__.__name__}(original_class={self.original_class.__name__} forward_func={self.forward_func} sdnq_dequantizer={repr(getattr(self, 'sdnq_dequantizer', None))})" + + +class SDNQLinear(SDNQLayer, torch.nn.Linear): + original_class: torch.nn.Linear + +class SDNQConv1d(SDNQLayer, torch.nn.Conv1d): + original_class: torch.nn.Conv1d + +class SDNQConv2d(SDNQLayer, torch.nn.Conv2d): + original_class: torch.nn.Conv2d + +class SDNQConv3d(SDNQLayer, torch.nn.Conv3d): + original_class: torch.nn.Conv3d + +class SDNQConvTranspose1d(SDNQLayer, torch.nn.ConvTranspose1d): + original_class: torch.nn.ConvTranspose1d + +class SDNQConvTranspose2d(SDNQLayer, torch.nn.ConvTranspose2d): + original_class: torch.nn.ConvTranspose2d + +class SDNQConvTranspose3d(SDNQLayer, torch.nn.ConvTranspose3d): + original_class: torch.nn.ConvTranspose3d + + +torch.serialization.add_safe_globals([SDNQLayer]) +torch.serialization.add_safe_globals([SDNQLinear]) +torch.serialization.add_safe_globals([SDNQConv1d]) +torch.serialization.add_safe_globals([SDNQConv2d]) +torch.serialization.add_safe_globals([SDNQConv3d]) +torch.serialization.add_safe_globals([SDNQConvTranspose1d]) +torch.serialization.add_safe_globals([SDNQConvTranspose2d]) +torch.serialization.add_safe_globals([SDNQConvTranspose3d]) + + +def get_sdnq_wrapper_class(original_layer, forward_func): + match original_layer.__class__.__name__: + case "Linear": + return SDNQLinear(original_layer, forward_func) + case "Conv1d": + return SDNQConv1d(original_layer, forward_func) + case "Conv2d": + return SDNQConv2d(original_layer, forward_func) + case "Conv3d": + return SDNQConv3d(original_layer, forward_func) + case "ConvTranspose1d": + return SDNQConvTranspose1d(original_layer, forward_func) + case "ConvTranspose2d": + return SDNQConvTranspose2d(original_layer, forward_func) + case "ConvTranspose3d": + return SDNQConvTranspose3d(original_layer, forward_func) + case _: + return SDNQLayer(original_layer, forward_func) diff --git a/modules/sdnq/layers/conv/conv_fp16.py b/modules/sdnq/layers/conv/conv_fp16.py index 71ddf3162..8b60767cc 100644 --- a/modules/sdnq/layers/conv/conv_fp16.py +++ b/modules/sdnq/layers/conv/conv_fp16.py @@ -5,6 +5,7 @@ from typing import List import torch from ...common import compile_func, fp_mm_func # noqa: TID252 +from ...packed_float import unpack_float # noqa: TID252 from ...dequantizer import dequantize_symmetric, dequantize_symmetric_with_bias # noqa: TID252 from .forward import get_conv_args, process_conv_input @@ -24,6 +25,8 @@ def conv_fp16_matmul( bias: torch.FloatTensor = None, svd_up: torch.FloatTensor = None, svd_down: torch.FloatTensor = None, + quantized_weight_shape: torch.Size = None, + weights_dtype: str = None, ) -> torch.FloatTensor: return_dtype = input.dtype input, mm_output_shape = process_conv_input(conv_type, input, reversed_padding_repeated_twice, padding_mode, result_shape, stride, padding, dilation) @@ -34,8 +37,12 @@ def conv_fp16_matmul( else: bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up) + if quantized_weight_shape is not None: + weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float16).t_() + scale = scale.t() + elif weight.dtype != torch.float16: + weight = weight.to(dtype=torch.float16) # fp8 weights input, scale = quantize_fp_mm_input_tensorwise(input, scale, matmul_dtype="float16") - weight = weight.to(dtype=torch.float16) # fp8 weights input, weight = check_mats(input, weight) if groups == 1: @@ -64,8 +71,10 @@ def conv_fp16_matmul( def quantized_conv_forward_fp16_matmul(self, input) -> torch.FloatTensor: if self.sdnq_dequantizer.re_quantize_for_matmul: weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None) + quantized_weight_shape = None else: weight, scale = self.weight, self.scale + quantized_weight_shape = self.sdnq_dequantizer.quantized_weight_shape if self.sdnq_dequantizer.is_packed else None conv_type, stride, padding, dilation = get_conv_args(input.ndim, self.stride, self.padding, self.dilation) return conv_fp16_matmul( input, weight, scale, @@ -76,6 +85,8 @@ def quantized_conv_forward_fp16_matmul(self, input) -> torch.FloatTensor: bias=self.bias, svd_up=self.svd_up, svd_down=self.svd_down, + quantized_weight_shape=quantized_weight_shape, + weights_dtype=self.sdnq_dequantizer.weights_dtype, ) diff --git a/modules/sdnq/layers/conv/conv_fp8.py b/modules/sdnq/layers/conv/conv_fp8.py index 4fcad6509..994850fb1 100644 --- a/modules/sdnq/layers/conv/conv_fp8.py +++ b/modules/sdnq/layers/conv/conv_fp8.py @@ -5,6 +5,7 @@ from typing import List import torch from ...common import compile_func # noqa: TID252 +from ...packed_float import unpack_float # noqa: TID252 from .forward import get_conv_args, process_conv_input from ..linear.linear_fp8 import quantize_fp_mm_input # noqa: TID252 @@ -23,6 +24,8 @@ def conv_fp8_matmul( bias: torch.FloatTensor = None, svd_up: torch.FloatTensor = None, svd_down: torch.FloatTensor = None, + quantized_weight_shape: torch.Size = None, + weights_dtype: str = None, ) -> torch.FloatTensor: return_dtype = input.dtype input, mm_output_shape = process_conv_input(conv_type, input, reversed_padding_repeated_twice, padding_mode, result_shape, stride, padding, dilation) @@ -30,6 +33,9 @@ def conv_fp8_matmul( input = input.flatten(0,-2) svd_bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up) + if quantized_weight_shape is not None: + weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn).t_() + scale = scale.t() input, input_scale = quantize_fp_mm_input(input) input, weight = check_mats(input, weight) @@ -71,8 +77,10 @@ def quantized_conv_forward_fp8_matmul(self, input) -> torch.FloatTensor: return self._conv_forward(input, self.sdnq_dequantizer(self.weight, self.scale, self.zero_point, self.svd_up, self.svd_down, skip_quantized_matmul=True), self.bias) if self.sdnq_dequantizer.re_quantize_for_matmul: weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None) + quantized_weight_shape = None else: weight, scale = self.weight, self.scale + quantized_weight_shape = self.sdnq_dequantizer.quantized_weight_shape if self.sdnq_dequantizer.is_packed else None conv_type, stride, padding, dilation = get_conv_args(input.ndim, self.stride, self.padding, self.dilation) return conv_fp8_matmul( input, weight, scale, @@ -83,6 +91,8 @@ def quantized_conv_forward_fp8_matmul(self, input) -> torch.FloatTensor: bias=self.bias, svd_up=self.svd_up, svd_down=self.svd_down, + quantized_weight_shape=quantized_weight_shape, + weights_dtype=self.sdnq_dequantizer.weights_dtype, ) diff --git a/modules/sdnq/layers/conv/conv_fp8_tensorwise.py b/modules/sdnq/layers/conv/conv_fp8_tensorwise.py index 2079bea33..9be958923 100644 --- a/modules/sdnq/layers/conv/conv_fp8_tensorwise.py +++ b/modules/sdnq/layers/conv/conv_fp8_tensorwise.py @@ -5,6 +5,7 @@ from typing import List import torch from ...common import compile_func # noqa: TID252 +from ...packed_float import unpack_float # noqa: TID252 from ...dequantizer import dequantize_symmetric, dequantize_symmetric_with_bias # noqa: TID252 from .forward import get_conv_args, process_conv_input @@ -24,6 +25,8 @@ def conv_fp8_matmul_tensorwise( bias: torch.FloatTensor = None, svd_up: torch.FloatTensor = None, svd_down: torch.FloatTensor = None, + quantized_weight_shape: torch.Size = None, + weights_dtype: str = None, ) -> torch.FloatTensor: return_dtype = input.dtype input, mm_output_shape = process_conv_input(conv_type, input, reversed_padding_repeated_twice, padding_mode, result_shape, stride, padding, dilation) @@ -34,6 +37,9 @@ def conv_fp8_matmul_tensorwise( else: bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up) + if quantized_weight_shape is not None: + weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn).t_() + scale = scale.t() input, scale = quantize_fp_mm_input_tensorwise(input, scale) input, weight = check_mats(input, weight) dummy_input_scale = torch.ones(1, device=input.device, dtype=torch.float32) @@ -66,8 +72,10 @@ def quantized_conv_forward_fp8_matmul_tensorwise(self, input) -> torch.FloatTens return self._conv_forward(input, self.sdnq_dequantizer(self.weight, self.scale, self.zero_point, self.svd_up, self.svd_down, skip_quantized_matmul=True), self.bias) if self.sdnq_dequantizer.re_quantize_for_matmul: weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None) + quantized_weight_shape = None else: weight, scale = self.weight, self.scale + quantized_weight_shape = self.sdnq_dequantizer.quantized_weight_shape if self.sdnq_dequantizer.is_packed else None conv_type, stride, padding, dilation = get_conv_args(input.ndim, self.stride, self.padding, self.dilation) return conv_fp8_matmul_tensorwise( input, weight, scale, @@ -78,6 +86,8 @@ def quantized_conv_forward_fp8_matmul_tensorwise(self, input) -> torch.FloatTens bias=self.bias, svd_up=self.svd_up, svd_down=self.svd_down, + quantized_weight_shape=quantized_weight_shape, + weights_dtype=self.sdnq_dequantizer.weights_dtype, ) diff --git a/modules/sdnq/layers/conv/conv_int8.py b/modules/sdnq/layers/conv/conv_int8.py index 9eaee44c7..9777b3d9b 100644 --- a/modules/sdnq/layers/conv/conv_int8.py +++ b/modules/sdnq/layers/conv/conv_int8.py @@ -15,18 +15,18 @@ from ..linear.forward import check_mats # noqa: TID252 def conv_int8_matmul( input: torch.FloatTensor, - weight: torch.CharTensor, - bias: torch.FloatTensor, + weight: torch.Tensor, scale: torch.FloatTensor, - svd_up: torch.FloatTensor, - svd_down: torch.FloatTensor, - quantized_weight_shape: torch.Size, result_shape: torch.Size, - weights_dtype: str, reversed_padding_repeated_twice: List[int], padding_mode: str, conv_type: int, groups: int, stride: List[int], padding: List[int], dilation: List[int], + bias: torch.FloatTensor = None, + svd_up: torch.FloatTensor = None, + svd_down: torch.FloatTensor = None, + quantized_weight_shape: torch.Size = None, + weights_dtype: str = None, ) -> torch.FloatTensor: return_dtype = input.dtype input, mm_output_shape = process_conv_input(conv_type, input, reversed_padding_repeated_twice, padding_mode, result_shape, stride, padding, dilation) @@ -37,9 +37,10 @@ def conv_int8_matmul( else: bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up) - input, scale = quantize_int_mm_input(input, scale) if quantized_weight_shape is not None: - weight = unpack_int_symetric(weight, quantized_weight_shape, weights_dtype, dtype=torch.int8) + weight = unpack_int_symetric(weight, quantized_weight_shape, weights_dtype, dtype=torch.int8).t_() + scale = scale.t() + input, scale = quantize_int_mm_input(input, scale) input, weight = check_mats(input, weight) if groups == 1: @@ -73,18 +74,19 @@ def quantized_conv_forward_int8_matmul(self, input) -> torch.FloatTensor: weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None) quantized_weight_shape = None else: - weight = self.weight - scale = self.scale + weight, scale = self.weight, self.scale quantized_weight_shape = self.sdnq_dequantizer.quantized_weight_shape if self.sdnq_dequantizer.is_packed else None return conv_int8_matmul( - input, weight, self.bias, - scale, self.svd_up, self.svd_down, - quantized_weight_shape, + input, weight, scale, self.sdnq_dequantizer.result_shape, - self.sdnq_dequantizer.weights_dtype, self._reversed_padding_repeated_twice, self.padding_mode, conv_type, self.groups, stride, padding, dilation, + bias=self.bias, + svd_up=self.svd_up, + svd_down=self.svd_down, + quantized_weight_shape=quantized_weight_shape, + weights_dtype=self.sdnq_dequantizer.weights_dtype, ) diff --git a/modules/sdnq/layers/linear/linear_fp16.py b/modules/sdnq/layers/linear/linear_fp16.py index db8a5cab0..3e04d3be6 100644 --- a/modules/sdnq/layers/linear/linear_fp16.py +++ b/modules/sdnq/layers/linear/linear_fp16.py @@ -3,6 +3,7 @@ import torch from ...common import compile_func, fp_mm_func # noqa: TID252 +from ...packed_float import unpack_float # noqa: TID252 from ...dequantizer import dequantize_symmetric, dequantize_symmetric_with_bias # noqa: TID252 from .forward import check_mats @@ -16,7 +17,14 @@ def fp16_matmul( bias: torch.FloatTensor = None, svd_up: torch.FloatTensor = None, svd_down: torch.FloatTensor = None, + quantized_weight_shape: torch.Size = None, + weights_dtype: str = None, ) -> torch.FloatTensor: + if quantized_weight_shape is not None: + weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float16).t_() + scale = scale.t() + elif weight.dtype != torch.float16: + weight = weight.to(dtype=torch.float16) # fp8 weights return_dtype = input.dtype output_shape = (*input.shape[:-1], weight.shape[-1]) if svd_up is not None: @@ -26,7 +34,6 @@ def fp16_matmul( else: bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up) input, scale = quantize_fp_mm_input_tensorwise(input, scale, matmul_dtype="float16") - weight = weight.to(dtype=torch.float16) # fp8 weights input, weight = check_mats(input, weight) if bias is not None: return dequantize_symmetric_with_bias(fp_mm_func(input, weight), scale, bias, dtype=return_dtype, result_shape=output_shape) @@ -37,9 +44,18 @@ def fp16_matmul( def quantized_linear_forward_fp16_matmul(self, input: torch.FloatTensor) -> torch.FloatTensor: if self.sdnq_dequantizer.re_quantize_for_matmul: weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None) + quantized_weight_shape = None else: weight, scale = self.weight, self.scale - return fp16_matmul(input, weight, scale, bias=self.bias, svd_up=self.svd_up, svd_down=self.svd_down) + quantized_weight_shape = self.sdnq_dequantizer.quantized_weight_shape if self.sdnq_dequantizer.is_packed else None + return fp16_matmul( + input, weight, scale, + bias=self.bias, + svd_up=self.svd_up, + svd_down=self.svd_down, + quantized_weight_shape=quantized_weight_shape, + weights_dtype=self.sdnq_dequantizer.weights_dtype, + ) fp16_matmul = compile_func(fp16_matmul) diff --git a/modules/sdnq/layers/linear/linear_fp8.py b/modules/sdnq/layers/linear/linear_fp8.py index d8f65ad6f..c037ff1c0 100644 --- a/modules/sdnq/layers/linear/linear_fp8.py +++ b/modules/sdnq/layers/linear/linear_fp8.py @@ -5,6 +5,7 @@ from typing import Tuple import torch from ...common import compile_func # noqa: TID252 +from ...packed_float import unpack_float # noqa: TID252 from ...dequantizer import quantize_fp_mm # noqa: TID252 from .forward import check_mats @@ -23,7 +24,12 @@ def fp8_matmul( bias: torch.FloatTensor = None, svd_up: torch.FloatTensor = None, svd_down: torch.FloatTensor = None, + quantized_weight_shape: torch.Size = None, + weights_dtype: str = None, ) -> torch.FloatTensor: + if quantized_weight_shape is not None: + weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn).t_() + scale = scale.t() return_dtype = input.dtype output_shape = (*input.shape[:-1], weight.shape[-1]) if svd_up is not None: @@ -45,9 +51,18 @@ def quantized_linear_forward_fp8_matmul(self, input: torch.FloatTensor) -> torch return torch.nn.functional.linear(input, self.sdnq_dequantizer(self.weight, self.scale, self.zero_point, self.svd_up, self.svd_down, skip_quantized_matmul=True), self.bias) if self.sdnq_dequantizer.re_quantize_for_matmul: weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None) + quantized_weight_shape = None else: weight, scale = self.weight, self.scale - return fp8_matmul(input, weight, scale, bias=self.bias, svd_up=self.svd_up, svd_down=self.svd_down) + quantized_weight_shape = self.sdnq_dequantizer.quantized_weight_shape if self.sdnq_dequantizer.is_packed else None + return fp8_matmul( + input, weight, scale, + bias=self.bias, + svd_up=self.svd_up, + svd_down=self.svd_down, + quantized_weight_shape=quantized_weight_shape, + weights_dtype=self.sdnq_dequantizer.weights_dtype, + ) fp8_matmul = compile_func(fp8_matmul) diff --git a/modules/sdnq/layers/linear/linear_fp8_tensorwise.py b/modules/sdnq/layers/linear/linear_fp8_tensorwise.py index 14560db48..9c65a3cd5 100644 --- a/modules/sdnq/layers/linear/linear_fp8_tensorwise.py +++ b/modules/sdnq/layers/linear/linear_fp8_tensorwise.py @@ -5,6 +5,7 @@ from typing import Tuple import torch from ...common import compile_func # noqa: TID252 +from ...packed_float import unpack_float # noqa: TID252 from ...dequantizer import quantize_fp_mm, dequantize_symmetric, dequantize_symmetric_with_bias # noqa: TID252 from .forward import check_mats @@ -26,7 +27,12 @@ def fp8_matmul_tensorwise( bias: torch.FloatTensor = None, svd_up: torch.FloatTensor = None, svd_down: torch.FloatTensor = None, + quantized_weight_shape: torch.Size = None, + weights_dtype: str = None, ) -> torch.FloatTensor: + if quantized_weight_shape is not None: + weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn).t_() + scale = scale.t() return_dtype = input.dtype output_shape = (*input.shape[:-1], weight.shape[-1]) if svd_up is not None: @@ -49,9 +55,18 @@ def quantized_linear_forward_fp8_matmul_tensorwise(self, input: torch.FloatTenso return torch.nn.functional.linear(input, self.sdnq_dequantizer(self.weight, self.scale, self.zero_point, self.svd_up, self.svd_down, skip_quantized_matmul=True), self.bias) if self.sdnq_dequantizer.re_quantize_for_matmul: weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None) + quantized_weight_shape = None else: weight, scale = self.weight, self.scale - return fp8_matmul_tensorwise(input, weight, scale, bias=self.bias, svd_up=self.svd_up, svd_down=self.svd_down) + quantized_weight_shape = self.sdnq_dequantizer.quantized_weight_shape if self.sdnq_dequantizer.is_packed else None + return fp8_matmul_tensorwise( + input, weight, scale, + bias=self.bias, + svd_up=self.svd_up, + svd_down=self.svd_down, + quantized_weight_shape=quantized_weight_shape, + weights_dtype=self.sdnq_dequantizer.weights_dtype, + ) fp8_matmul_tensorwise = compile_func(fp8_matmul_tensorwise) diff --git a/modules/sdnq/layers/linear/linear_int8.py b/modules/sdnq/layers/linear/linear_int8.py index 14efcea34..2d26a6086 100644 --- a/modules/sdnq/layers/linear/linear_int8.py +++ b/modules/sdnq/layers/linear/linear_int8.py @@ -31,7 +31,8 @@ def int8_matmul( weights_dtype: str = None, ) -> torch.FloatTensor: if quantized_weight_shape is not None: - weight = unpack_int_symetric(weight, quantized_weight_shape, weights_dtype, dtype=torch.int8) + weight = unpack_int_symetric(weight, quantized_weight_shape, weights_dtype, dtype=torch.int8).t_() + scale = scale.t() return_dtype = input.dtype output_shape = (*input.shape[:-1], weight.shape[-1]) if svd_up is not None: @@ -55,8 +56,7 @@ def quantized_linear_forward_int8_matmul(self, input: torch.FloatTensor) -> torc weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None) quantized_weight_shape = None else: - weight = self.weight - scale = self.scale + weight, scale = self.weight, self.scale quantized_weight_shape = self.sdnq_dequantizer.quantized_weight_shape if self.sdnq_dequantizer.is_packed else None return int8_matmul( input, weight, scale, @@ -64,7 +64,7 @@ def quantized_linear_forward_int8_matmul(self, input: torch.FloatTensor) -> torc svd_up=self.svd_up, svd_down=self.svd_down, quantized_weight_shape=quantized_weight_shape, - weights_dtype=self.sdnq_dequantizer.weights_dtype + weights_dtype=self.sdnq_dequantizer.weights_dtype, ) diff --git a/modules/sdnq/loader.py b/modules/sdnq/loader.py index e264a9b1f..91be08394 100644 --- a/modules/sdnq/loader.py +++ b/modules/sdnq/loader.py @@ -3,7 +3,7 @@ import json import torch from diffusers.models.modeling_utils import ModelMixin -from .common import dtype_dict, use_tensorwise_fp8_matmul, check_torch_compile +from .common import dtype_dict, use_tensorwise_fp8_matmul, check_torch_compile, conv_types, linear_types from .quantizer import SDNQConfig, sdnq_post_load_quant, prepare_weight_for_matmul, prepare_svd_for_matmul, get_quant_args_from_config from .forward import get_forward_func from .file_loader import load_files @@ -25,7 +25,7 @@ def unset_config_on_save(quantization_config: SDNQConfig) -> SDNQConfig: return quantization_config -def save_sdnq_model(model: ModelMixin, model_path: str, max_shard_size: str = "10GB", is_pipeline: bool = False, sdnq_config: SDNQConfig = None) -> None: +def save_sdnq_model(model: ModelMixin, model_path: str, max_shard_size: str = "5GB", is_pipeline: bool = False, sdnq_config: SDNQConfig = None) -> None: if is_pipeline: for module_name in get_module_names(model): module = getattr(model, module_name, None) @@ -106,7 +106,7 @@ def load_sdnq_model(model_path: str, model_cls: ModelMixin = None, file_name: st else: model = model_cls(**model_config) - model = sdnq_post_load_quant(model, torch_dtype=dtype, add_skip_keys=False, **get_quant_args_from_config(quantization_config)) + model = sdnq_post_load_quant(model, torch_dtype=dtype, add_skip_keys=False, use_dynamic_quantization=False, **get_quant_args_from_config(quantization_config)) key_mapping = getattr(model, "_checkpoint_conversion_mapping", None) files = [] @@ -170,6 +170,18 @@ def apply_sdnq_options_to_module(model, dtype: torch.dtype = None, dequantize_fp return model for module_name, module in model.named_children(): if hasattr(module, "sdnq_dequantizer"): + layer_class_name = module.original_class.__name__ + current_use_quantized_matmul = use_quantized_matmul + if current_use_quantized_matmul: + if layer_class_name in conv_types: + output_channel_size, channel_size = module.sdnq_dequantizer.original_shape[:2] + elif layer_class_name in linear_types: + output_channel_size, channel_size = module.sdnq_dequantizer.original_shape + else: + current_use_quantized_matmul = False + current_use_quantized_matmul = current_use_quantized_matmul and channel_size >= 32 and output_channel_size >= 32 # pylint: disable=possibly-used-before-assignment + current_use_quantized_matmul = current_use_quantized_matmul and output_channel_size % 16 == 0 and channel_size % 16 == 0 # pylint: disable=possibly-used-before-assignment + if dtype is not None and module.sdnq_dequantizer.result_dtype != torch.float32: module.sdnq_dequantizer.result_dtype = dtype @@ -177,7 +189,7 @@ def apply_sdnq_options_to_module(model, dtype: torch.dtype = None, dequantize_fp dequantize_fp32 or dtype_dict[module.sdnq_dequantizer.weights_dtype]["num_bits"] > 8 or ( - (use_quantized_matmul or (use_quantized_matmul is None and module.sdnq_dequantizer.use_quantized_matmul)) + (current_use_quantized_matmul or (current_use_quantized_matmul is None and module.sdnq_dequantizer.use_quantized_matmul)) and not dtype_dict[module.sdnq_dequantizer.quantized_matmul_dtype]["is_integer"] and (not use_tensorwise_fp8_matmul or dtype_dict[module.sdnq_dequantizer.quantized_matmul_dtype]["num_bits"] == 16) ) @@ -191,20 +203,19 @@ def apply_sdnq_options_to_module(model, dtype: torch.dtype = None, dequantize_fp module.svd_up.data = module.svd_up.to(dtype=scale_dtype) module.svd_down.data = module.svd_down.to(dtype=scale_dtype) - if use_quantized_matmul is not None and use_quantized_matmul != module.sdnq_dequantizer.use_quantized_matmul: - if not module.sdnq_dequantizer.re_quantize_for_matmul: + if current_use_quantized_matmul is not None and current_use_quantized_matmul != module.sdnq_dequantizer.use_quantized_matmul: + if not module.sdnq_dequantizer.re_quantize_for_matmul and not dtype_dict[module.sdnq_dequantizer.weights_dtype]["is_packed"]: module.scale.t_() module.weight.t_() - if use_quantized_matmul: + if current_use_quantized_matmul: module.weight.data = prepare_weight_for_matmul(module.weight) else: module.scale.data = module.scale.contiguous() module.weight.data = module.weight.contiguous() if module.svd_up is not None: - module.svd_up.data, module.svd_down.data = prepare_svd_for_matmul(module.svd_up.t_(), module.svd_down.t_(), use_quantized_matmul) - module.sdnq_dequantizer.use_quantized_matmul = use_quantized_matmul - module.forward = get_forward_func(module.__class__.__name__, module.sdnq_dequantizer.quantized_matmul_dtype, use_quantized_matmul) - module.forward = module.forward.__get__(module, module.__class__) + module.svd_up.data, module.svd_down.data = prepare_svd_for_matmul(module.svd_up.t_(), module.svd_down.t_(), current_use_quantized_matmul) + module.sdnq_dequantizer.use_quantized_matmul = current_use_quantized_matmul + module.forward_func = get_forward_func(module.original_class.__name__, module.sdnq_dequantizer.quantized_matmul_dtype, current_use_quantized_matmul) setattr(model, module_name, module) else: setattr(model, module_name, apply_sdnq_options_to_module(module, dtype=dtype, dequantize_fp32=dequantize_fp32, use_quantized_matmul=use_quantized_matmul)) diff --git a/modules/sdnq/packed_float.py b/modules/sdnq/packed_float.py new file mode 100644 index 000000000..8cd7940d7 --- /dev/null +++ b/modules/sdnq/packed_float.py @@ -0,0 +1,102 @@ +import torch + +from .common import dtype_dict +from .packed_int import pack_int_asymetric, unpack_int_asymetric + + +float_bits_to_uint_dict = { + 1: "uint1", + 2: "uint2", + 3: "uint3", + 4: "uint4", + 5: "uint5", + 6: "uint6", + 7: "uint7", +} + + +def pack_float(x: torch.FloatTensor, weights_dtype: str) -> torch.Tensor: + exponent_bits = dtype_dict[weights_dtype]["exponent"] + mantissa_bits = dtype_dict[weights_dtype]["mantissa"] + total_bits = dtype_dict[weights_dtype]["num_bits"] + + if dtype_dict[weights_dtype]["is_unsigned"]: + sign_mask = (1 << (total_bits-1)) # pylint: disable=superfluous-parens + else: + sign_mask = (1 << (total_bits-1)) + (1 << (total_bits-2)) + + mantissa_difference = 23 - mantissa_bits + exponent_difference = 8 - exponent_bits + mantissa_mask = (1 << mantissa_difference) # pylint: disable=superfluous-parens + + x = x.to(dtype=torch.float32).view(torch.int32) + + x = torch.where( + torch.gt( + torch.bitwise_and(x, -(1 << (mantissa_difference-4)) & ~(-mantissa_mask)), + (1 << (mantissa_difference-1)), + ), + torch.add(x, mantissa_mask), + x, + ) + + x = torch.where(torch.lt(x.view(torch.float32).abs(), dtype_dict[weights_dtype]["min_normal"]), 0, x) + + x = torch.bitwise_right_shift(x, mantissa_difference) + x = torch.bitwise_and( + torch.bitwise_or( + torch.bitwise_and(torch.bitwise_right_shift(x, exponent_difference), sign_mask), + torch.bitwise_and(x, ~sign_mask), + ), + ~(-(1 << total_bits)), + ).view(torch.uint32) + + if total_bits < 8: + x = pack_int_asymetric(x, float_bits_to_uint_dict[total_bits]) + else: + x = x.to(dtype=dtype_dict[weights_dtype]["storage_dtype"]) + + return x + + +def unpack_float(x: torch.Tensor, shape: torch.Size, weights_dtype: str) -> torch.FloatTensor: + exponent_bits = dtype_dict[weights_dtype]["exponent"] + mantissa_bits = dtype_dict[weights_dtype]["mantissa"] + total_bits = dtype_dict[weights_dtype]["num_bits"] + + if dtype_dict[weights_dtype]["is_unsigned"]: + sign_mask = (1 << (total_bits-1)) # pylint: disable=superfluous-parens + else: + sign_mask = (1 << (total_bits-1)) + (1 << (total_bits-2)) + + mantissa_difference = 23 - mantissa_bits + exponent_difference = 8 - exponent_bits + + if total_bits < 8: + x = unpack_int_asymetric(x, shape, float_bits_to_uint_dict[total_bits]) + + x = x.to(dtype=torch.uint32).view(torch.int32) + x = torch.bitwise_left_shift( + torch.bitwise_or( + torch.bitwise_left_shift(torch.bitwise_and(x, sign_mask), exponent_difference), + torch.bitwise_and(x, ~sign_mask), + ), + mantissa_difference, + ) + + x = torch.bitwise_or( + x, + torch.bitwise_and( + torch.bitwise_right_shift( + -torch.bitwise_and(torch.bitwise_not(x), 1073741824), + exponent_difference, + ), + 1065353216, + ), + ) + + overflow_mask = (~(-(1 << (22 + exponent_bits))) | -1073741824) + x = torch.where(torch.bitwise_and(x, overflow_mask).to(dtype=torch.bool), x, 0) + x = x.view(torch.float32) + + return x diff --git a/modules/sdnq/quantizer.py b/modules/sdnq/quantizer.py index 02184ed8d..9414ff6f6 100644 --- a/modules/sdnq/quantizer.py +++ b/modules/sdnq/quantizer.py @@ -15,10 +15,12 @@ from diffusers.utils import get_module_from_name from accelerate import init_empty_weights from modules import devices, shared -from .common import sdnq_version, dtype_dict, common_skip_keys, module_skip_keys_dict, accepted_weight_dtypes, accepted_matmul_dtypes, allowed_types, linear_types, conv_types, conv_transpose_types, compile_func, use_tensorwise_fp8_matmul, use_contiguous_mm, check_torch_compile +from .common import sdnq_version, dtype_dict, common_skip_keys, module_skip_keys_dict, accepted_weight_dtypes, accepted_matmul_dtypes, weights_dtype_order, weights_dtype_order_fp32, allowed_types, linear_types, conv_types, conv_transpose_types, compile_func, use_tensorwise_fp8_matmul, use_contiguous_mm, check_torch_compile from .dequantizer import SDNQDequantizer, dequantize_sdnq_model from .packed_int import pack_int_symetric, pack_int_asymetric +from .packed_float import pack_float from .forward import get_forward_func +from .layers import get_sdnq_wrapper_class class QuantizationMethod(str, Enum): @@ -54,7 +56,7 @@ def quantize_weight(weight: torch.FloatTensor, reduction_axes: Union[int, List[i if dtype_dict[weights_dtype]["is_integer"]: if use_stochastic_rounding: - quantized_weight.add_(torch.randn_like(quantized_weight), alpha=0.1) + quantized_weight.add_(torch.rand_like(quantized_weight), alpha=0.1) quantized_weight.round_() else: if use_stochastic_rounding: @@ -131,6 +133,7 @@ def get_quant_args_from_config(quantization_config: Union["SDNQConfig", dict]) - quantization_config_dict.pop("return_device", None) quantization_config_dict.pop("non_blocking", None) quantization_config_dict.pop("add_skip_keys", None) + quantization_config_dict.pop("use_dynamic_quantization", None) quantization_config_dict.pop("use_static_quantization", None) quantization_config_dict.pop("use_stochastic_rounding", None) quantization_config_dict.pop("use_grad_ckpt", None) @@ -202,7 +205,7 @@ def add_module_skip_keys(model, modules_to_not_convert: List[str] = None, module @devices.inference_context() -def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int8", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, use_svd=False, use_quantized_matmul=False, use_stochastic_rounding=False, dequantize_fp32=False, param_name=None): # pylint: disable=unused-argument +def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int8", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, use_svd=False, use_quantized_matmul=False, use_stochastic_rounding=False, dequantize_fp32=False, using_pre_calculated_svd=False, param_name=None): # pylint: disable=unused-argument num_of_groups = 1 is_conv_type = False is_conv_transpose_type = False @@ -226,6 +229,15 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int dtype_dict[weights_dtype]["is_unsigned"] or dtype_dict[weights_dtype]["is_integer"] != dtype_dict[quantized_matmul_dtype]["is_integer"] or dtype_dict[weights_dtype]["num_bits"] > dtype_dict[quantized_matmul_dtype]["num_bits"] + or ( + dtype_dict[weights_dtype]["is_packed"] + and not dtype_dict[weights_dtype]["is_integer"] + and not dtype_dict[quantized_matmul_dtype]["is_integer"] + and ( + dtype_dict[weights_dtype]["num_bits"] >= dtype_dict[quantized_matmul_dtype]["num_bits"] + or dtype_dict[weights_dtype]["max"] > dtype_dict[quantized_matmul_dtype]["max"] + ) + ) ) if layer_class_name in conv_types: @@ -278,9 +290,9 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int if use_quantized_matmul and not re_quantize_for_matmul and dtype_dict[weights_dtype]["num_bits"] >= 6: group_size = -1 elif is_linear_type: - group_size = 2 ** ((2 if svd_up is None else 3) + dtype_dict[weights_dtype]["num_bits"]) + group_size = 2 ** ((3 if (svd_up is not None or using_pre_calculated_svd) else 2) + dtype_dict[weights_dtype]["num_bits"]) else: - group_size = 2 ** ((1 if svd_up is None else 2) + dtype_dict[weights_dtype]["num_bits"]) + group_size = 2 ** ((2 if (svd_up is not None or using_pre_calculated_svd) else 1) + dtype_dict[weights_dtype]["num_bits"]) if group_size > 0: if group_size >= channel_size: @@ -341,11 +353,11 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int svd_down = svd_down.to(dtype=torch_dtype) re_quantize_for_matmul = re_quantize_for_matmul or num_of_groups > 1 - if use_quantized_matmul and not re_quantize_for_matmul: + if use_quantized_matmul and not re_quantize_for_matmul and not dtype_dict[weights_dtype]["is_packed"]: scale.t_() weight.t_() weight = prepare_weight_for_matmul(weight) - if not use_tensorwise_fp8_matmul and not dtype_dict[weights_dtype]["is_integer"]: + if not use_tensorwise_fp8_matmul and not dtype_dict[quantized_matmul_dtype]["is_integer"]: scale = scale.to(dtype=torch.float32) sdnq_dequantizer = SDNQDequantizer( @@ -366,10 +378,13 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int ) if dtype_dict[weights_dtype]["is_packed"]: - if dtype_dict[weights_dtype]["is_unsigned"]: - weight = pack_int_asymetric(weight, weights_dtype) + if dtype_dict[weights_dtype]["is_integer"]: + if dtype_dict[weights_dtype]["is_unsigned"]: + weight = pack_int_asymetric(weight, weights_dtype) + else: + weight = pack_int_symetric(weight, weights_dtype) else: - weight = pack_int_symetric(weight, weights_dtype) + weight = pack_float(weight, weights_dtype) else: weight = weight.to(dtype=dtype_dict[weights_dtype]["torch_dtype"]) @@ -377,11 +392,63 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int @devices.inference_context() -def sdnq_quantize_layer(layer, weights_dtype="int8", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, use_svd=False, quant_conv=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, use_stochastic_rounding=False, dequantize_fp32=False, non_blocking=False, quantization_device=None, return_device=None, param_name=None): # pylint: disable=unused-argument +def sdnq_quantize_layer_weight_dynamic(weight, layer_class_name=None, weights_dtype="int2", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, dynamic_loss_threshold=1e-2, use_svd=False, use_quantized_matmul=False, use_dynamic_quantization=False, use_stochastic_rounding=False, dequantize_fp32=False, param_name=None): # pylint: disable=unused-argument + if torch_dtype is None: + torch_dtype = weight.dtype + weights_dtype_order_to_use = weights_dtype_order_fp32 if torch_dtype in {torch.float32, torch.float64} else weights_dtype_order + weight = weight.to(dtype=torch.float32) + weight_std = weight.std().square() + + if use_svd: + try: + svd_weight, svd_up, svd_down = apply_svdquant(weight, rank=svd_rank, niter=svd_steps) + svd_up, svd_down = prepare_svd_for_matmul(svd_up, svd_down, use_quantized_matmul) + svd_up = svd_up.to(dtype=torch_dtype) + svd_down = svd_down.to(dtype=torch_dtype) + except Exception: + svd_up, svd_down = None, None + svd_weight = weight + else: + svd_up, svd_down = None, None + svd_weight = weight + + quantization_loss = None + svd_is_transposed = False + for i in range(weights_dtype_order_to_use.index(weights_dtype), len(weights_dtype_order_to_use)): + quantized_weight, scale, zero_point, _, _, sdnq_dequantizer = sdnq_quantize_layer_weight( + svd_weight, + layer_class_name=layer_class_name, + weights_dtype=weights_dtype_order_to_use[i], + quantized_matmul_dtype=quantized_matmul_dtype, + torch_dtype=torch_dtype, + group_size=group_size, + svd_rank=svd_rank, + svd_steps=svd_steps, + use_svd=False, + using_pre_calculated_svd=use_svd, + use_quantized_matmul=use_quantized_matmul, + use_stochastic_rounding=use_stochastic_rounding, + dequantize_fp32=dequantize_fp32, + param_name=param_name, + ) + + if use_svd and not svd_is_transposed and sdnq_dequantizer.use_quantized_matmul: + svd_up = svd_up.t_() + svd_down = svd_down.t_() + svd_is_transposed = True + + quantization_loss = torch.nn.functional.mse_loss(weight, sdnq_dequantizer(quantized_weight, scale, zero_point, svd_up, svd_down, skip_quantized_matmul=sdnq_dequantizer.use_quantized_matmul, dtype=torch.float32, skip_compile=True)).div_(weight_std) + if quantization_loss <= dynamic_loss_threshold: + return (quantized_weight, scale, zero_point, svd_up, svd_down, sdnq_dequantizer) + return None + + +@devices.inference_context() +def sdnq_quantize_layer(layer, weights_dtype="int8", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, dynamic_loss_threshold=1e-2, use_svd=False, quant_conv=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, use_dynamic_quantization=False, use_stochastic_rounding=False, dequantize_fp32=False, non_blocking=False, modules_to_not_convert=None, modules_dtype_dict=None, quantization_device=None, return_device=None, param_name=None): # pylint: disable=unused-argument layer_class_name = layer.__class__.__name__ if layer_class_name in conv_transpose_types or layer_class_name in conv_types: if not quant_conv: - return layer + return layer, modules_to_not_convert, modules_dtype_dict use_quantized_matmul = use_quantized_matmul_conv layer.weight.requires_grad_(False) @@ -390,46 +457,81 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", quantized_matmul_dtype=None if quantization_device is not None: layer.weight.data = layer.weight.to(quantization_device, non_blocking=non_blocking) - ( - layer.weight.data, - layer.scale, layer.zero_point, - layer.svd_up, layer.svd_down, - layer.sdnq_dequantizer, - ) = sdnq_quantize_layer_weight( - layer.weight, - layer_class_name=layer_class_name, - weights_dtype=weights_dtype, - quantized_matmul_dtype=quantized_matmul_dtype, - torch_dtype=torch_dtype, - group_size=group_size, - svd_rank=svd_rank, - svd_steps=svd_steps, - use_svd=use_svd, - use_quantized_matmul=use_quantized_matmul, - use_stochastic_rounding=use_stochastic_rounding, - dequantize_fp32=dequantize_fp32, - param_name=param_name, - ) + if use_dynamic_quantization: + weight_data = sdnq_quantize_layer_weight_dynamic( + layer.weight, + layer_class_name=layer_class_name, + weights_dtype=weights_dtype, + quantized_matmul_dtype=quantized_matmul_dtype, + torch_dtype=torch_dtype, + group_size=group_size, + svd_rank=svd_rank, + svd_steps=svd_steps, + dynamic_loss_threshold=dynamic_loss_threshold, + use_svd=use_svd, + use_quantized_matmul=use_quantized_matmul, + use_stochastic_rounding=use_stochastic_rounding, + dequantize_fp32=dequantize_fp32, + param_name=param_name, + ) + else: + weight_data = sdnq_quantize_layer_weight( + layer.weight, + layer_class_name=layer_class_name, + weights_dtype=weights_dtype, + quantized_matmul_dtype=quantized_matmul_dtype, + torch_dtype=torch_dtype, + group_size=group_size, + svd_rank=svd_rank, + svd_steps=svd_steps, + use_svd=use_svd, + use_quantized_matmul=use_quantized_matmul, + use_stochastic_rounding=use_stochastic_rounding, + dequantize_fp32=dequantize_fp32, + param_name=param_name, + ) - layer.weight = torch.nn.Parameter(layer.weight.to(return_device, non_blocking=non_blocking), requires_grad=False) - layer.scale = torch.nn.Parameter(layer.scale.to(return_device, non_blocking=non_blocking), requires_grad=False) - if layer.zero_point is not None: - layer.zero_point = torch.nn.Parameter(layer.zero_point.to(return_device, non_blocking=non_blocking), requires_grad=False) - if layer.svd_up is not None: - layer.svd_up = torch.nn.Parameter(layer.svd_up.to(return_device, non_blocking=non_blocking), requires_grad=False) - layer.svd_down = torch.nn.Parameter(layer.svd_down.to(return_device, non_blocking=non_blocking), requires_grad=False) + if weight_data is not None: + ( + layer.weight.data, + layer.scale, layer.zero_point, + layer.svd_up, layer.svd_down, + layer.sdnq_dequantizer, + ) = weight_data + del weight_data - layer = layer.to(return_device, non_blocking=non_blocking) - layer.forward = get_forward_func(layer_class_name, layer.sdnq_dequantizer.quantized_matmul_dtype, layer.sdnq_dequantizer.use_quantized_matmul) - layer.forward = layer.forward.__get__(layer, layer.__class__) - return layer + layer = get_sdnq_wrapper_class(layer, get_forward_func(layer_class_name, layer.sdnq_dequantizer.quantized_matmul_dtype, layer.sdnq_dequantizer.use_quantized_matmul)) + layer.weight = torch.nn.Parameter(layer.weight.to(return_device, non_blocking=non_blocking), requires_grad=False) + layer.scale = torch.nn.Parameter(layer.scale.to(return_device, non_blocking=non_blocking), requires_grad=False) + if layer.zero_point is not None: + layer.zero_point = torch.nn.Parameter(layer.zero_point.to(return_device, non_blocking=non_blocking), requires_grad=False) + if layer.svd_up is not None: + layer.svd_up = torch.nn.Parameter(layer.svd_up.to(return_device, non_blocking=non_blocking), requires_grad=False) + layer.svd_down = torch.nn.Parameter(layer.svd_down.to(return_device, non_blocking=non_blocking), requires_grad=False) + layer = layer.to(return_device, non_blocking=non_blocking) + + if use_dynamic_quantization: + if modules_dtype_dict is None: + modules_dtype_dict = {} + if layer.sdnq_dequantizer.weights_dtype not in modules_dtype_dict.keys(): + modules_dtype_dict[layer.sdnq_dequantizer.weights_dtype] = [param_name] + else: + modules_dtype_dict[layer.sdnq_dequantizer.weights_dtype].append(param_name) + else: + layer = layer.to(return_device, dtype=torch_dtype, non_blocking=non_blocking) + if use_dynamic_quantization: + if modules_to_not_convert is None: + modules_to_not_convert = [] + modules_to_not_convert.append(param_name) + + return layer, modules_to_not_convert, modules_dtype_dict @devices.inference_context() -def apply_sdnq_to_module(model, weights_dtype="int8", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, use_svd=False, quant_conv=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, use_stochastic_rounding=False, dequantize_fp32=False, non_blocking=False, quantization_device=None, return_device=None, modules_to_not_convert: List[str] = None, modules_dtype_dict: Dict[str, List[str]] = None, full_param_name=""): # pylint: disable=unused-argument +def apply_sdnq_to_module(model, weights_dtype="int8", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, dynamic_loss_threshold=1e-2, use_svd=False, quant_conv=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, use_dynamic_quantization=False, use_stochastic_rounding=False, dequantize_fp32=False, non_blocking=False, modules_to_not_convert: List[str] = None, modules_dtype_dict: Dict[str, List[str]] = None, quantization_device=None, return_device=None, full_param_name=""): # pylint: disable=unused-argument has_children = list(model.children()) if not has_children: - return model + return model, modules_to_not_convert, modules_dtype_dict if modules_to_not_convert is None: modules_to_not_convert = [] if modules_dtype_dict is None: @@ -447,7 +549,7 @@ def apply_sdnq_to_module(model, weights_dtype="int8", quantized_matmul_dtype=Non if layer_class_name in allowed_types and module.weight.dtype in {torch.float32, torch.float16, torch.bfloat16}: if (layer_class_name in conv_types or layer_class_name in conv_transpose_types) and not quant_conv: continue - setattr(model, module_name, sdnq_quantize_layer( + module, modules_to_not_convert, modules_dtype_dict = sdnq_quantize_layer( module, weights_dtype=get_minimum_dtype(weights_dtype, param_name, modules_dtype_dict), quantized_matmul_dtype=quantized_matmul_dtype, @@ -455,39 +557,48 @@ def apply_sdnq_to_module(model, weights_dtype="int8", quantized_matmul_dtype=Non group_size=group_size, svd_rank=svd_rank, svd_steps=svd_steps, + dynamic_loss_threshold=dynamic_loss_threshold, use_svd=use_svd, quant_conv=quant_conv, use_quantized_matmul=use_quantized_matmul, use_quantized_matmul_conv=use_quantized_matmul_conv, + use_dynamic_quantization=use_dynamic_quantization, use_stochastic_rounding=use_stochastic_rounding, dequantize_fp32=dequantize_fp32, non_blocking=non_blocking, quantization_device=quantization_device, return_device=return_device, + modules_to_not_convert=modules_to_not_convert, + modules_dtype_dict=modules_dtype_dict, param_name=param_name, - )) - setattr(model, module_name, apply_sdnq_to_module( - module, - weights_dtype=weights_dtype, - quantized_matmul_dtype=quantized_matmul_dtype, - torch_dtype=torch_dtype, - group_size=group_size, - svd_rank=svd_rank, - svd_steps=svd_steps, - use_svd=use_svd, - quant_conv=quant_conv, - use_quantized_matmul=use_quantized_matmul, - use_quantized_matmul_conv=use_quantized_matmul_conv, - use_stochastic_rounding=use_stochastic_rounding, - dequantize_fp32=dequantize_fp32, - non_blocking=non_blocking, - quantization_device=quantization_device, - return_device=return_device, - modules_to_not_convert=modules_to_not_convert, - modules_dtype_dict=modules_dtype_dict, - full_param_name=param_name, - )) - return model + ) + setattr(model, module_name, module) + + module, modules_to_not_convert, modules_dtype_dict = apply_sdnq_to_module( + module, + dynamic_loss_threshold=dynamic_loss_threshold, + weights_dtype=weights_dtype, + quantized_matmul_dtype=quantized_matmul_dtype, + torch_dtype=torch_dtype, + group_size=group_size, + svd_rank=svd_rank, + svd_steps=svd_steps, + use_svd=use_svd, + quant_conv=quant_conv, + use_quantized_matmul=use_quantized_matmul, + use_quantized_matmul_conv=use_quantized_matmul_conv, + use_dynamic_quantization=use_dynamic_quantization, + use_stochastic_rounding=use_stochastic_rounding, + dequantize_fp32=dequantize_fp32, + non_blocking=non_blocking, + quantization_device=quantization_device, + return_device=return_device, + modules_to_not_convert=modules_to_not_convert, + modules_dtype_dict=modules_dtype_dict, + full_param_name=param_name, + ) + setattr(model, module_name, module) + return model, modules_to_not_convert, modules_dtype_dict @devices.inference_context() @@ -499,18 +610,20 @@ def sdnq_post_load_quant( group_size: int = 0, svd_rank: int = 32, svd_steps: int = 8, + dynamic_loss_threshold: float = 1e-2, use_svd: bool = False, quant_conv: bool = False, use_quantized_matmul: bool = False, use_quantized_matmul_conv: bool = False, + use_dynamic_quantization: bool = False, use_stochastic_rounding: bool = False, dequantize_fp32: bool = False, non_blocking: bool = False, add_skip_keys:bool = True, - quantization_device: Optional[torch.device] = None, - return_device: Optional[torch.device] = None, modules_to_not_convert: List[str] = None, modules_dtype_dict: Dict[str, List[str]] = None, + quantization_device: Optional[torch.device] = None, + return_device: Optional[torch.device] = None, ): if modules_to_not_convert is None: modules_to_not_convert = [] @@ -527,22 +640,24 @@ def sdnq_post_load_quant( group_size=group_size, svd_rank=svd_rank, svd_steps=svd_steps, + dynamic_loss_threshold=dynamic_loss_threshold, use_svd=use_svd, quant_conv=quant_conv, use_quantized_matmul=use_quantized_matmul, use_quantized_matmul_conv=use_quantized_matmul_conv, + use_dynamic_quantization=use_dynamic_quantization, use_stochastic_rounding=use_stochastic_rounding, dequantize_fp32=dequantize_fp32, non_blocking=non_blocking, add_skip_keys=add_skip_keys, - quantization_device=quantization_device, - return_device=return_device, modules_to_not_convert=modules_to_not_convert, modules_dtype_dict=modules_dtype_dict, + quantization_device=quantization_device, + return_device=return_device, ) model.eval() - model = apply_sdnq_to_module( + model, modules_to_not_convert, modules_dtype_dict = apply_sdnq_to_module( model, weights_dtype=weights_dtype, quantized_matmul_dtype=quantized_matmul_dtype, @@ -550,19 +665,24 @@ def sdnq_post_load_quant( group_size=group_size, svd_rank=svd_rank, svd_steps=svd_steps, + dynamic_loss_threshold=dynamic_loss_threshold, use_svd=use_svd, quant_conv=quant_conv, use_quantized_matmul=use_quantized_matmul, use_quantized_matmul_conv=use_quantized_matmul_conv, + use_dynamic_quantization=use_dynamic_quantization, use_stochastic_rounding=use_stochastic_rounding, dequantize_fp32=dequantize_fp32, non_blocking=non_blocking, - quantization_device=quantization_device, - return_device=return_device, modules_to_not_convert=modules_to_not_convert, modules_dtype_dict=modules_dtype_dict, + quantization_device=quantization_device, + return_device=return_device, ) + quantization_config.modules_to_not_convert = modules_to_not_convert + quantization_config.modules_dtype_dict = modules_dtype_dict + model.quantization_config = quantization_config if hasattr(model, "config"): try: @@ -693,9 +813,9 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): else: param_value = param_value.to(target_device, non_blocking=self.quantization_config.non_blocking).to(dtype=torch.float32) - layer, _ = get_module_from_name(model, param_name) + layer, tensor_name = get_module_from_name(model, param_name) layer.weight = torch.nn.Parameter(param_value, requires_grad=False) - layer = sdnq_quantize_layer( + layer, self.quantization_config.modules_to_not_convert, self.quantization_config.modules_dtype_dict = sdnq_quantize_layer( layer, weights_dtype=weights_dtype, quantized_matmul_dtype=self.quantization_config.quantized_matmul_dtype, @@ -703,25 +823,32 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): group_size=self.quantization_config.group_size, svd_rank=self.quantization_config.svd_rank, svd_steps=self.quantization_config.svd_steps, + dynamic_loss_threshold=self.quantization_config.dynamic_loss_threshold, use_svd=self.quantization_config.use_svd, quant_conv=self.quantization_config.quant_conv, use_quantized_matmul=self.quantization_config.use_quantized_matmul, use_quantized_matmul_conv=self.quantization_config.use_quantized_matmul_conv, + use_dynamic_quantization=self.quantization_config.use_dynamic_quantization, use_stochastic_rounding=self.quantization_config.use_stochastic_rounding, dequantize_fp32=self.quantization_config.dequantize_fp32, non_blocking=self.quantization_config.non_blocking, + modules_to_not_convert=self.quantization_config.modules_to_not_convert, + modules_dtype_dict=self.quantization_config.modules_dtype_dict, quantization_device=None, return_device=return_device, param_name=param_name, ) layer.weight._is_hf_initialized = True # pylint: disable=protected-access - layer.scale._is_hf_initialized = True # pylint: disable=protected-access - if layer.zero_point is not None: - layer.zero_point._is_hf_initialized = True # pylint: disable=protected-access - if layer.svd_up is not None: - layer.svd_up._is_hf_initialized = True # pylint: disable=protected-access - layer.svd_down._is_hf_initialized = True # pylint: disable=protected-access + if hasattr(layer, "scale"): + layer.scale._is_hf_initialized = True # pylint: disable=protected-access + if layer.zero_point is not None: + layer.zero_point._is_hf_initialized = True # pylint: disable=protected-access + if layer.svd_up is not None: + layer.svd_up._is_hf_initialized = True # pylint: disable=protected-access + layer.svd_down._is_hf_initialized = True # pylint: disable=protected-access + parent_module, tensor_name = get_module_from_name(model, param_name.removesuffix(tensor_name).removesuffix(".")) + setattr(parent_module, tensor_name, layer) def get_quantize_ops(self): return SDNQQuantize(self) @@ -757,7 +884,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): self.quantization_config.add_skip_keys = False with init_empty_weights(): - model = sdnq_post_load_quant(model, torch_dtype=self.torch_dtype, add_skip_keys=False, **get_quant_args_from_config(self.quantization_config)) + model = sdnq_post_load_quant(model, torch_dtype=self.torch_dtype, add_skip_keys=False, use_dynamic_quantization=False, **get_quant_args_from_config(self.quantization_config)) if self.quantization_config.add_skip_keys: if keep_in_fp32_modules is not None: @@ -838,8 +965,9 @@ class SDNQConfig(QuantizationConfigMixin): Args: weights_dtype (`str`, *optional*, defaults to `"int8"`): - The target dtype for the weights after quantization. Supported values are: - ("int16", "int8", "int7", "int6", "int5", "int4", "int3", "int2", "uint16", "uint8", "uint7", "uint6", "uint5", "uint4", "uint3", "uint2", "uint1", "bool", "float16", "float8_e4m3fn", "float8_e4m3fnuz", "float8_e5m2", "float8_e5m2fnuz") + The target dtype for the weights after quantization. + Check out `sdnq.common.accepted_weight_dtypes` for all the supported values. + These are some of the recommended values to use: ("int8", "int7", "int6", "uint5", "uint4", "uint3", "uint2", "float8_e4m3fn", "float7_e3m3fn", "float6_e3m2fn", "float5_e2m2fn", "float4_e2m1fn", "float3_e1m1fn", "float2_e1m0fn") quantized_matmul_dtype (`str`, *optional*, defaults to `None`): The target dtype for quantized matmul. `None` will use "int8" with integer weight dtypes and "float8_e4m3fn" or "float16" with float weight dtypes. @@ -849,6 +977,8 @@ class SDNQConfig(QuantizationConfigMixin): group_size = 0 will automatically select a group size based on weights_dtype. svd_rank (`int`, *optional*, defaults to `32`): The rank size used for the SVDQuant algorithm. + dynamic_loss_threshold (`float`, *optional*, defaults to `1e-2`): + The target quantization mse loss threshold to use for dynamic quantization. svd_steps (`int`, *optional*, defaults to `8`): The number of iterations to use in svd lowrank estimation. use_svd (`bool`, *optional*, defaults to `False`): @@ -861,6 +991,9 @@ class SDNQConfig(QuantizationConfigMixin): Same as use_quantized_matmul_conv but for the convolutional layers with UNets like SDXL. use_stochastic_rounding (`bool`, *optional*, defaults to `False`): Enabling this option will use stochastic rounding on the quantization step. + use_dynamic_quantization (`bool`, *optional*, defaults to `False`): + Enabling this option will dynamically select a per layer quantization type based on the dynamic_loss_threshold. + weights_dtype will be used as the minimum allowed quantization type when this option is enabled. dequantize_fp32 (`bool`, *optional*, defaults to `False`): Enabling this option will use FP32 on the dequantization step. non_blocking (`bool`, *optional*, defaults to `False`): @@ -885,12 +1018,14 @@ class SDNQConfig(QuantizationConfigMixin): group_size: int = 0, svd_rank: int = 32, svd_steps: int = 8, + dynamic_loss_threshold: float = 1e-2, use_svd: bool = False, use_grad_ckpt: bool = True, quant_conv: bool = False, use_quantized_matmul: bool = False, use_quantized_matmul_conv: bool = False, use_static_quantization: bool = True, + use_dynamic_quantization: bool = False, use_stochastic_rounding: bool = False, dequantize_fp32: bool = False, non_blocking: bool = False, @@ -911,6 +1046,7 @@ class SDNQConfig(QuantizationConfigMixin): self.quant_method = QuantizationMethod.SDNQ self.group_size = group_size self.svd_rank = svd_rank + self.dynamic_loss_threshold = dynamic_loss_threshold self.svd_steps = svd_steps self.use_svd = use_svd self.use_grad_ckpt = use_grad_ckpt @@ -918,6 +1054,7 @@ class SDNQConfig(QuantizationConfigMixin): self.use_quantized_matmul = use_quantized_matmul self.use_quantized_matmul_conv = use_quantized_matmul_conv self.use_static_quantization = use_static_quantization + self.use_dynamic_quantization = use_dynamic_quantization self.use_stochastic_rounding = use_stochastic_rounding self.dequantize_fp32 = dequantize_fp32 self.non_blocking = non_blocking @@ -970,18 +1107,24 @@ class SDNQConfig(QuantizationConfigMixin): self.modules_dtype_dict = self.modules_dtype_dict.copy() def to_dict(self): - dct = self.__dict__.copy() # make serializable - dct["quantization_device"] = str(dct["quantization_device"]) if dct["quantization_device"] is not None else None - dct["return_device"] = str(dct["return_device"]) if dct["return_device"] is not None else None - return dct + quantization_config_dict = self.__dict__.copy() # make serializable + quantization_config_dict["quantization_device"] = str(quantization_config_dict["quantization_device"]) if quantization_config_dict["quantization_device"] is not None else None + quantization_config_dict["return_device"] = str(quantization_config_dict["return_device"]) if quantization_config_dict["return_device"] is not None else None + return quantization_config_dict import diffusers.quantizers.auto # noqa: E402,RUF100 # pylint: disable=wrong-import-order diffusers.quantizers.auto.AUTO_QUANTIZER_MAPPING["sdnq"] = SDNQQuantizer diffusers.quantizers.auto.AUTO_QUANTIZATION_CONFIG_MAPPING["sdnq"] = SDNQConfig +diffusers.quantizers.auto.AUTO_QUANTIZER_MAPPING["sdnq_training"] = SDNQQuantizer +diffusers.quantizers.auto.AUTO_QUANTIZATION_CONFIG_MAPPING["sdnq_training"] = SDNQConfig + import transformers.quantizers.auto # noqa: E402,RUF100 # pylint: disable=wrong-import-order transformers.quantizers.auto.AUTO_QUANTIZER_MAPPING["sdnq"] = SDNQQuantizer transformers.quantizers.auto.AUTO_QUANTIZATION_CONFIG_MAPPING["sdnq"] = SDNQConfig +transformers.quantizers.auto.AUTO_QUANTIZER_MAPPING["sdnq_training"] = SDNQQuantizer +transformers.quantizers.auto.AUTO_QUANTIZATION_CONFIG_MAPPING["sdnq_training"] = SDNQConfig + sdnq_quantize_layer_weight_compiled = compile_func(sdnq_quantize_layer_weight) diff --git a/modules/shared.py b/modules/shared.py index 924fccea5..5deb96825 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -4,25 +4,27 @@ import os import sys import time import contextlib + from enum import Enum from typing import TYPE_CHECKING import gradio as gr +from installer import log, print_dict, console, get_version # pylint: disable=unused-import +log.debug('Initializing: shared module') + +import modules.memmon +import modules.paths as paths from modules.json_helpers import readfile, writefile # pylint: disable=W0611 from modules.shared_helpers import listdir, walk_files, html_path, html, req, total_tqdm # pylint: disable=W0611 +from modules import errors, devices, shared_state, cmd_args, theme, history, files_cache from modules.shared_defaults import get_default_modes -from modules import errors, devices, shared_items, shared_state, cmd_args, theme, history, files_cache from modules.paths import models_path, script_path, data_path, sd_configs_path, sd_default_config, sd_model_file, default_sd_model_file, extensions_dir, extensions_builtin_dir # pylint: disable=W0611 -from modules.dml import memory_providers, default_memory_provider, directml_do_hijack -from modules.onnx_impl import execution_providers from modules.memstats import memory_stats, ram_stats # pylint: disable=unused-import + +log.debug('Initializing: pipelines') +from modules import shared_items from modules.interrogate.openclip import caption_models, caption_types, get_clip_models, refresh_clip_models from modules.interrogate.vqa import vlm_models, vlm_prompts, vlm_system, vlm_default -from modules.ui_components import DropdownEditable -from modules.options import OptionInfo, options_section -import modules.memmon -import modules.styles -import modules.paths as paths -from installer import log, print_dict, console, get_version # pylint: disable=unused-import + if TYPE_CHECKING: # Behavior modified by __future__.annotations @@ -52,7 +54,6 @@ face_restorers = [] yolo = None tab_names = [] extra_networks: list[ExtraNetworksPage] = [] -options_templates: dict[str, OptionInfo | LegacyOption] = {} hypernetworks = {} settings_components = {} restricted_opts = { @@ -70,7 +71,7 @@ restricted_opts = { } resize_modes = ["None", "Fixed", "Crop", "Fill", "Outpaint", "Context aware"] max_workers = 12 -sdnq_quant_modes = ["int8", "float8_e4m3fn", "int7", "int6", "int5", "uint4", "uint3", "uint2", "float8_e5m2", "float8_e4m3fnuz", "float8_e5m2fnuz", "float16", "int16", "uint16", "uint8", "uint7", "uint6", "uint5", "int4", "int3", "int2", "uint1"] +sdnq_quant_modes = ["int8", "int7", "int6", "uint5", "uint4", "uint3", "uint2", "float8_e4m3fn", "float7_e3m3fn", "float6_e3m2fn", "float5_e2m2fn", "float4_e2m1fn", "float3_e1m1fn", "float2_e1m0fn"] sdnq_matmul_modes = ["auto", "int8", "float8_e4m3fn", "float16"] default_hfcache_dir = os.environ.get("SD_HFCACHEDIR", None) or os.path.join(paths.models_path, 'huggingface') state = shared_state.State() @@ -143,8 +144,15 @@ def list_samplers(): modules.sd_samplers.set_samplers() return modules.sd_samplers.all_samplers - +log.debug('Initializing: default modes') startup_offload_mode, startup_offload_min_gpu, startup_offload_max_gpu, startup_cross_attention, startup_sdp_options, startup_sdp_choices, startup_sdp_override_options, startup_sdp_override_choices, startup_offload_always, startup_offload_never = get_default_modes(cmd_opts=cmd_opts, mem_stat=mem_stat) +from modules.dml import memory_providers, default_memory_provider, directml_do_hijack +from modules.onnx_impl import execution_providers + +log.debug('Initializing: settings') +from modules.ui_components import DropdownEditable +from modules.options import OptionInfo, options_section +options_templates: dict[str, OptionInfo | LegacyOption] = {} options_templates.update(options_section(('sd', "Model Loading"), { "sd_backend": OptionInfo('diffusers', "Execution backend", gr.Radio, {"choices": ['diffusers', 'original'], "visible": False }), @@ -224,7 +232,9 @@ options_templates.update(options_section(("quantization", "Model Quantization"), "sdnq_quantize_weights_group_size": OptionInfo(0, "Group size", gr.Slider, {"minimum": -1, "maximum": 4096, "step": 1}), "sdnq_svd_rank": OptionInfo(32, "SVD rank size", gr.Slider, {"minimum": 1, "maximum": 512, "step": 1}), "sdnq_svd_steps": OptionInfo(8, "SVD steps", gr.Slider, {"minimum": 1, "maximum": 128, "step": 1}), + "sdnq_dynamic_loss_threshold": OptionInfo(1e-2, "Dynamic loss threshold", gr.Slider, {"minimum": 1e-4, "maximum": 1e-1, "step": 1e-4}), "sdnq_use_svd": OptionInfo(False, "Use SVD quantization", gr.Checkbox), + "sdnq_use_dynamic_quantization": OptionInfo(False, "Use Dynamic quantization", gr.Checkbox), "sdnq_quantize_conv_layers": OptionInfo(False, "Quantize convolutional layers", gr.Checkbox), "sdnq_dequantize_compile": OptionInfo(devices.has_triton(early=True), "Dequantize using torch.compile", gr.Checkbox), "sdnq_use_quantized_matmul": OptionInfo(False, "Use quantized MatMul", gr.Checkbox), @@ -548,7 +558,7 @@ options_templates.update(options_section(('saving-paths', "Image Paths"), { "directories_max_prompt_words": OptionInfo(8, "Max words", gr.Slider, {"minimum": 1, "maximum": 99, "step": 1, **hide_dirs}), "outdir_sep_dirs": OptionInfo("

Folders

", "", gr.HTML), - "outdir_samples": OptionInfo("", "Images folder", component_args=hide_dirs, folder=True), + "outdir_samples": OptionInfo("", "Base images folder", component_args=hide_dirs, folder=True), "outdir_txt2img_samples": OptionInfo("outputs/text", 'Folder for text generate', component_args=hide_dirs, folder=True), "outdir_img2img_samples": OptionInfo("outputs/image", 'Folder for image generate', component_args=hide_dirs, folder=True), "outdir_control_samples": OptionInfo("outputs/control", 'Folder for control generate', component_args=hide_dirs, folder=True), @@ -558,7 +568,7 @@ options_templates.update(options_section(('saving-paths', "Image Paths"), { "outdir_init_images": OptionInfo("outputs/inputs", "Folder for init images", component_args=hide_dirs, folder=True), "outdir_sep_grids": OptionInfo("

Grids

", "", gr.HTML), - "outdir_grids": OptionInfo("", "Grids folder", component_args=hide_dirs, folder=True), + "outdir_grids": OptionInfo("", "Base grids folde", component_args=hide_dirs, folder=True), "outdir_txt2img_grids": OptionInfo("outputs/grids", 'Folder for txt2img grids', component_args=hide_dirs, folder=True), "outdir_img2img_grids": OptionInfo("outputs/grids", 'Folder for img2img grids', component_args=hide_dirs, folder=True), "outdir_control_grids": OptionInfo("outputs/grids", 'Folder for control grids', component_args=hide_dirs, folder=True), @@ -679,8 +689,8 @@ options_templates.update(options_section(('interrogate', "Interrogate"), { "interrogate_vlm_system": OptionInfo(vlm_system, "VLM: default prompt"), "interrogate_vlm_num_beams": OptionInfo(1, "VLM: num beams", gr.Slider, {"minimum": 1, "maximum": 16, "step": 1, "visible": False}), "interrogate_vlm_max_length": OptionInfo(512, "VLM: max length", gr.Slider, {"minimum": 1, "maximum": 4096, "step": 1, "visible": False}), - "interrogate_vlm_do_sample": OptionInfo(False, "VLM: use sample method"), - "interrogate_vlm_temperature": OptionInfo(0, "VLM: temperature", gr.Slider, {"minimum": 0, "maximum": 1.0, "step": 0.01, "visible": False}), + "interrogate_vlm_do_sample": OptionInfo(True, "VLM: use sample method"), + "interrogate_vlm_temperature": OptionInfo(0.8, "VLM: temperature", gr.Slider, {"minimum": 0, "maximum": 1.0, "step": 0.01, "visible": False}), "interrogate_vlm_top_k": OptionInfo(0, "VLM: top-k", gr.Slider, {"minimum": 0, "maximum": 99, "step": 1, "visible": False}), "interrogate_vlm_top_p": OptionInfo(0, "VLM: top-p", gr.Slider, {"minimum": 0, "maximum": 1.0, "step": 0.01, "visible": False}), "interrogate_vlm_keep_prefill": OptionInfo(False, "VLM: keep prefill text in output", gr.Checkbox, {"visible": False}), @@ -733,13 +743,13 @@ options_templates.update(options_section(('extra_networks', "Networks"), { "extra_networks_default_multiplier": OptionInfo(1.0, "Default strength", gr.Slider, {"minimum": 0.0, "maximum": 2.0, "step": 0.01}), "lora_force_reload": OptionInfo(False, "LoRA force reload always"), "lora_force_diffusers": OptionInfo(False if not cmd_opts.use_openvino else True, "LoRA load using Diffusers method"), + + "lora_apply_te": OptionInfo(False, "LoRA native apply to text encoder"), "lora_fuse_native": OptionInfo(True, "LoRA native fuse with model"), "lora_fuse_diffusers": OptionInfo(False, "LoRA diffusers fuse with model"), "lora_apply_tags": OptionInfo(0, "LoRA auto-apply tags", gr.Slider, {"minimum": -1, "maximum": 32, "step": 1}), "lora_in_memory_limit": OptionInfo(1, "LoRA memory cache", gr.Slider, {"minimum": 0, "maximum": 32, "step": 1}), "lora_add_hashes_to_infotext": OptionInfo(False, "LoRA add hash info to metadata"), - "lora_quant": OptionInfo("NF4","LoRA precision when quantized", gr.Radio, {"choices": ["NF4", "FP4"]}), - "lora_maybe_diffusers": OptionInfo(False, "LoRA load using Diffusers method for selected models", gr.Checkbox, {"visible": False}), "extra_networks_styles_sep": OptionInfo("

Styles

", "", gr.HTML), "extra_networks_styles": OptionInfo(True, "Show reference styles"), @@ -823,9 +833,8 @@ options_templates.update(options_section(('hidden_options', "Hidden options"), { from modules.shared_legacy import get_legacy_options options_templates.update(get_legacy_options()) from modules.options_handler import Options -opts = Options(options_templates, restricted_opts) config_filename = cmd_opts.config -opts.load(config_filename) +opts = Options(options_templates, restricted_opts, filename=config_filename) cmd_opts = cmd_args.settings_args(opts, cmd_opts) if cmd_opts.locale is not None: opts.data['ui_locale'] = cmd_opts.locale @@ -836,9 +845,12 @@ opts.data['uni_pc_order'] = max(2, opts.schedulers_solver_order) # compatibility log.info(f'Engine: backend={backend} compute={devices.backend} device={devices.get_optimal_device_name()} attention="{opts.cross_attention_optimization}" mode={devices.inference_context.__name__}') profiler = None +import modules.styles prompt_styles = modules.styles.StyleDatabase(opts) reference_models = readfile(os.path.join('html', 'reference.json'), as_type="dict") if opts.extra_network_reference_enable else {} cmd_opts.disable_extension_access = (cmd_opts.share or cmd_opts.listen or (cmd_opts.server_name or False)) and not cmd_opts.insecure + +log.debug('Initializing: devices') devices.args = cmd_opts devices.opts = opts devices.onnx = [opts.onnx_execution_provider] @@ -851,13 +863,6 @@ mem_mon = modules.memmon.MemUsageMonitor("MemMon", devices.device) history = history.History() if devices.backend == "directml": directml_do_hijack() -elif sys.platform == "win32" and (devices.backend == "zluda" or devices.backend == "rocm"): - from modules.rocm_triton_windows import apply_triton_patches - apply_triton_patches() - - if devices.backend == "zluda": - from modules.zluda import initialize_zluda - initialize_zluda() from modules import sdnq # pylint: disable=unused-import # register to diffusers and transformers log.debug('Quantization: registered=SDNQ') diff --git a/modules/shared_defaults.py b/modules/shared_defaults.py index 59058b566..5c7179a0e 100644 --- a/modules/shared_defaults.py +++ b/modules/shared_defaults.py @@ -52,13 +52,12 @@ def get_default_modes(cmd_opts, mem_stat): default_sdp_override_choices.append('Triton Flash attention') elif devices.backend == "rocm": default_sdp_override_choices.append('Triton Flash attention') - import torch - if int(getattr(torch.cuda.get_device_properties(devices.device), "gcnArchName", "gfx0000")[3:]) < 1100: + agent = devices.get_hip_agent() + if agent.gfx_version < 0x1100: default_sdp_override_options = ['Dynamic attention'] # only RDNA2 and older GPUs needs this elif devices.backend in {"directml", "cpu", "mps"}: default_sdp_override_options = ['Dynamic attention'] - return ( default_offload_mode, default_diffusers_offload_min_gpu_memory, diff --git a/modules/shared_items.py b/modules/shared_items.py index 40e3b768f..b973e4886 100644 --- a/modules/shared_items.py +++ b/modules/shared_items.py @@ -48,7 +48,10 @@ pipelines = { 'Qwen': getattr(diffusers, 'QwenImagePipeline', None), 'HunyuanImage': getattr(diffusers, 'HunyuanImagePipeline', None), 'Z-Image': getattr(diffusers, 'ZImagePipeline', None), + 'FLUX2': getattr(diffusers, 'Flux2Pipeline', None), + 'FLUX2 Klein': getattr(diffusers, 'Flux2KleinPipeline', None), 'LongCat': getattr(diffusers, 'LongCatImagePipeline', None), + 'GLM-Image': getattr(diffusers, 'GlmImagePipeline', None), # dynamically imported and redefined later 'Meissonic': getattr(diffusers, 'DiffusionPipeline', None), 'Monetico': getattr(diffusers, 'DiffusionPipeline', None), @@ -134,7 +137,7 @@ def get_pipelines(): for k, v in pipelines.items(): if k != 'Autodetect' and v is None: from installer import log # pylint: disable=redefined-outer-name - log.error(f'Pipeline={k} diffusers={diffusers.__version__} path={diffusers.__file__} not available') + log.error(f'Model="{k}" diffusers={diffusers.__version__} path={diffusers.__file__} pipeline not available') return pipelines diff --git a/modules/styles.py b/modules/styles.py index 02187f581..5164f6870 100644 --- a/modules/styles.py +++ b/modules/styles.py @@ -1,4 +1,3 @@ -# We need this so Python doesn't complain about the unknown StableDiffusionProcessing-typehint at runtime from __future__ import annotations import re import os @@ -6,6 +5,7 @@ import csv import json import time import random +from typing import Dict from modules import files_cache, shared, infotext, sd_models, sd_vae @@ -45,8 +45,71 @@ def apply_styles_to_prompt(prompt, styles): return prompt +def select_from_weighted_list(inner: str) -> str: + if not inner: + return '' + + parts = [p.strip() for p in inner.split('|') if p.strip()] + weighted: Dict[str, float] = {} + unweighted = [] + + for p in parts: + if ':' in p and not p.startswith('(') and not p.endswith(')'): + name, wstr = p.split(':', 1) + name = name.strip() + try: + w = float(wstr.strip()) + except Exception: + w = 0.0 + w = max(0.0, min(1.0, w)) + weighted[name] = weighted.get(name, 0.0) + w + else: + unweighted.append(p) + + W = sum(weighted.values()) + U = len(unweighted) + + if U == 0: + # Only weighted options + keys = list(weighted.keys()) + if not keys: + return '' + if W == 0.0: + return random.choice(keys) + if abs(W - 1.0) > 1e-12: + for k in weighted: + weighted[k] = weighted[k] / W + else: + # Mix of weighted and unweighted + if W >= 1.0: + # Weighted probabilities consume whole mass -> normalize them, unweighted get 0 + for k in weighted: + weighted[k] = weighted[k] / W + else: + remaining = 1.0 - W + per = remaining / U + for name in unweighted: + weighted[name] = weighted.get(name, 0.0) + per + + items = list(weighted.items()) + if not items: + return '' + total = sum(v for _, v in items) + if total <= 0.0: + return items[0][0] + + r = random.random() * total + cum = 0.0 + for name, prob in items: + cum += prob + if r <= cum: + return name + return items[-1][0] + + def apply_curly_braces_to_prompt(prompt, seed=-1): - # woman with {white|green|{purple|yellow}} highlights and {red|blue} dress + # unweighted: woman with {white|green|{purple|yellow}} highlights and {red|blue} dress + # weighted: woman with {white:0.6|green:0.2|{purple|yellow}} highlights and {red:.6|blue:.4} dress if not isinstance(prompt, str) or len(prompt) == 0: return prompt old_state = None @@ -60,8 +123,7 @@ def apply_curly_braces_to_prompt(prompt, seed=-1): if not m: break inner = m.group(1) - options = [opt.strip() for opt in inner.split('|')] - choice = random.choice([o for o in options if o != '']) if options else '' + choice = select_from_weighted_list(inner) prompt = prompt[:m.start()] + choice + prompt[m.end():] # replace this specific span (slice-based) to avoid accidental other replacements finally: if old_state is not None: @@ -71,12 +133,13 @@ def apply_curly_braces_to_prompt(prompt, seed=-1): def apply_file_wildcards(prompt, replaced = [], not_found = [], recursion=0, seed=-1): def check_wildcard_files(prompt, wildcard, files, file_only=True): - trimmed = wildcard.replace('\\', os.path.sep).strip().lower() + trimmed = wildcard.replace('\\', os.path.sep).replace('/', os.path.sep).strip().lower() for file in files: if file_only: paths = [os.path.splitext(file)[0].lower(), os.path.splitext(os.path.basename(file).lower())[0]] # fullname and basename else: paths = [os.path.splitext(p.lower())[0] for p in os.path.normpath(file).split(os.path.sep)] # every path component + paths.insert(0, os.path.splitext(file)[0].lower()) if (trimmed in paths) or (os.path.sep in trimmed and trimmed in paths[0]): try: with open(file, 'r', encoding='utf-8') as f: diff --git a/modules/taesd/taehv.py b/modules/taesd/taehv.py index a8ba17469..4e974deae 100644 --- a/modules/taesd/taehv.py +++ b/modules/taesd/taehv.py @@ -81,8 +81,6 @@ def apply_model_with_memblocks(model, x, parallel, show_progress_bar): T = NT // N x = x.view(N, T, C, H, W) else: - # TODO(oboerbohan): at least on macos this still gradually uses more memory during decode... - # need to fix :( out = [] # iterate over input timesteps and also iterate over blocks. # because of the cursed TPool/TGrow blocks, this is not a nested loop, diff --git a/modules/taesd/taem1.py b/modules/taesd/taem1.py index c92c86fb5..682d4de26 100644 --- a/modules/taesd/taem1.py +++ b/modules/taesd/taem1.py @@ -77,8 +77,6 @@ def apply_model_with_memblocks(model, x, parallel, show_progress_bar): T = NT // N x = x.view(N, T, C, H, W) else: - # TODO(oboerbohan): at least on macos this still gradually uses more memory during decode... - # need to fix :( out = [] # iterate over input timesteps and also iterate over blocks. # because of the cursed TPool/TGrow blocks, this is not a nested loop, diff --git a/modules/taesd/taesd.py b/modules/taesd/taesd.py index 8e391a8fb..f066f4cfd 100644 --- a/modules/taesd/taesd.py +++ b/modules/taesd/taesd.py @@ -77,7 +77,11 @@ class TAESD(nn.Module): # pylint: disable=abstract-method self.decoder = self.decoder.to(devices.device, dtype=self.dtype) def guess_latent_channels(self, decoder_path, encoder_path): - return 16 if ("f1" in encoder_path or "f1" in decoder_path) or ("sd3" in encoder_path or "sd3" in decoder_path) else 4 + if "f2" in encoder_path or "f2" in decoder_path: + return 32 # FLUX.2 uses 32 latent channels + if ("f1" in encoder_path or "f1" in decoder_path) or ("sd3" in encoder_path or "sd3" in decoder_path): + return 16 + return 4 @staticmethod def scale_latents(x): diff --git a/modules/txt2img.py b/modules/txt2img.py index 08d4f1bf2..021b586e5 100644 --- a/modules/txt2img.py +++ b/modules/txt2img.py @@ -2,6 +2,7 @@ import os from modules import shared, processing, scripts_manager from modules.generation_parameters_copypaste import create_override_settings_dict from modules.ui_common import plaintext_to_html +from modules.paths import resolve_output_path debug = shared.log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -41,8 +42,8 @@ def txt2img(id_task, state, p = processing.StableDiffusionProcessingTxt2Img( sd_model=shared.sd_model, - outpath_samples=shared.opts.outdir_samples or shared.opts.outdir_txt2img_samples, - outpath_grids=shared.opts.outdir_grids or shared.opts.outdir_txt2img_grids, + outpath_samples=resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_txt2img_samples), + outpath_grids=resolve_output_path(shared.opts.outdir_grids, shared.opts.outdir_txt2img_grids), prompt=prompt, styles=prompt_styles, negative_prompt=negative_prompt, diff --git a/modules/ui_caption.py b/modules/ui_caption.py index aad6deddb..d27b76ce6 100644 --- a/modules/ui_caption.py +++ b/modules/ui_caption.py @@ -40,7 +40,7 @@ def update_vlm_params(*args): shared.opts.interrogate_vlm_keep_prefill = bool(vlm_keep_prefill) shared.opts.interrogate_vlm_keep_thinking = bool(vlm_keep_thinking) shared.opts.interrogate_vlm_thinking_mode = bool(vlm_thinking_mode) - shared.opts.save(shared.config_filename) + shared.opts.save() def update_clip_params(*args): @@ -52,7 +52,7 @@ def update_clip_params(*args): shared.opts.interrogate_clip_num_beams = int(clip_num_beams) shared.opts.interrogate_clip_flavor_count = int(clip_flavor_count) shared.opts.interrogate_clip_chunk_size = int(clip_chunk_size) - shared.opts.save(shared.config_filename) + shared.opts.save() openclip.update_interrogate_params() diff --git a/modules/ui_common.py b/modules/ui_common.py index 98cda43d8..4c0c6979e 100644 --- a/modules/ui_common.py +++ b/modules/ui_common.py @@ -6,6 +6,7 @@ import platform import subprocess import gradio as gr from modules import call_queue, shared, errors, ui_sections, ui_symbols, ui_components, generation_parameters_copypaste, images, scripts_manager, script_callbacks, infotext, processing +from modules.paths import resolve_output_path folder_symbol = ui_symbols.folder @@ -67,9 +68,14 @@ def delete_files(js_data, files, all_files, index): except Exception: data = { 'index_of_first_image': 0 } start_index = 0 - if index > -1 and shared.opts.save_selected_only and (index >= data['index_of_first_image']): - files = [files[index]] - start_index = index + first_index = data['index_of_first_image'] + if (index > -1) and shared.opts.save_selected_only and (index >= first_index): # ensures we are looking at a specific non-grid picture, and we have save_selected_only # pylint: disable=no-member + if index < len(files): + files = [files[index]] + start_index = index + else: + shared.log.error(f'Delete: index={index} first={first_index} files={len(files)} out of range') + files = [] deleted = [] all_files = [f.split('/file=')[1] if 'file=' in f else f for f in all_files] if isinstance(all_files, list) else [] all_files = [os.path.normpath(f) for f in all_files] @@ -82,7 +88,7 @@ def delete_files(js_data, files, all_files, index): continue if os.path.exists(fn) and os.path.isfile(fn): deleted.append(fn) - os.remove(fn) + # os.remove(fn) if fn in all_files: all_files.remove(fn) shared.log.info(f'Delete: image="{fn}"') @@ -100,7 +106,7 @@ def delete_files(js_data, files, all_files, index): def save_files(js_data, files, html_info, index): - os.makedirs(shared.opts.outdir_save, exist_ok=True) + os.makedirs(resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_save), exist_ok=True) class PObject: # pylint: disable=too-few-public-methods def __init__(self, d=None): @@ -122,7 +128,7 @@ def save_files(js_data, files, html_info, index): self.styles = getattr(self, 'styles', None) or getattr(self, 'Styles', None) or [] self.styles = [s.strip() for s in self.styles.split(',')] if isinstance(self.styles, str) else self.styles - self.outpath_grids = shared.opts.outdir_grids or shared.opts.outdir_txt2img_grids + self.outpath_grids = resolve_output_path(shared.opts.outdir_grids, shared.opts.outdir_txt2img_grids) self.infotexts = getattr(self, 'infotexts', [html_info]) self.infotext = self.infotexts[0] if len(self.infotexts) > 0 else html_info self.all_negative_prompt = getattr(self, 'all_negative_prompts', [self.negative_prompt]) @@ -135,9 +141,13 @@ def save_files(js_data, files, html_info, index): data = {} p = PObject(data) start_index = 0 - if index > -1 and shared.opts.save_selected_only and (index >= p.index_of_first_image): # ensures we are looking at a specific non-grid picture, and we have save_selected_only # pylint: disable=no-member - files = [files[index]] - start_index = index + if (index > -1) and shared.opts.save_selected_only and (index >= p.index_of_first_image): # ensures we are looking at a specific non-grid picture, and we have save_selected_only # pylint: disable=no-member + if index < len(files): + files = [files[index]] + start_index = index + else: + shared.log.error(f'Save: index={index} first={p.index_of_first_image} files={len(files)} out of range') + files = [] filenames = [] fullfns = [] for image_index, filedata in enumerate(files, start_index): @@ -152,14 +162,14 @@ def save_files(js_data, files, html_info, index): if 'name' in filedata and ('tmp' not in filedata['name']) and os.path.isfile(filedata['name']): fullfn = filedata['name'] fullfns.append(fullfn) - destination = shared.opts.outdir_save + destination = resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_save) namegen = images.FilenameGenerator(p, seed=p.all_seeds[i], prompt=p.all_prompts[i], image=None) # pylint: disable=no-member dirname = namegen.apply(shared.opts.directories_filename_pattern or "[prompt_words]").lstrip(' ').rstrip('\\ /') destination = os.path.join(destination, dirname) destination = namegen.sanitize(destination) os.makedirs(destination, exist_ok = True) tgt_filename = os.path.join(destination, os.path.basename(fullfn)) - relfn = os.path.relpath(tgt_filename, shared.opts.outdir_save) + relfn = os.path.relpath(tgt_filename, resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_save)) filenames.append(relfn) if not os.path.exists(tgt_filename): try: @@ -192,14 +202,14 @@ def save_files(js_data, files, html_info, index): try: seed = p.all_seeds[i] if i < len(p.all_seeds) else p.seed prompt = p.all_prompts[i] if i < len(p.all_prompts) else p.prompt - fullfn, txt_fullfn, _exif = images.save_image(image, shared.opts.outdir_save, "", seed=seed, prompt=prompt, info=info, extension=shared.opts.samples_format, grid=is_grid, p=p) + fullfn, txt_fullfn, _exif = images.save_image(image, resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_save), "", seed=seed, prompt=prompt, info=info, extension=shared.opts.samples_format, grid=is_grid, p=p) except Exception as e: fullfn, txt_fullfn = None, None shared.log.error(f'Save: image={image} i={i} seeds={p.all_seeds} prompts={p.all_prompts}') errors.display(e, 'save') if fullfn is None: continue - filename = os.path.relpath(fullfn, shared.opts.outdir_save) + filename = os.path.relpath(fullfn, resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_save)) filenames.append(filename) fullfns.append(fullfn) if txt_fullfn: @@ -207,7 +217,7 @@ def save_files(js_data, files, html_info, index): # fullfns.append(txt_fullfn) script_callbacks.image_save_btn_callback(filename) if shared.opts.samples_save_zip and len(fullfns) > 1: - zip_filepath = os.path.join(shared.opts.outdir_save, "images.zip") + zip_filepath = os.path.join(resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_save), "images.zip") from zipfile import ZipFile with ZipFile(zip_filepath, "w") as zip_file: for i in range(len(fullfns)): @@ -277,7 +287,7 @@ def create_output_panel(tabname, preview=True, prompt=None, height=None, transfe save = gr.Button('Save', elem_id=f'save_{tabname}') delete = gr.Button('Delete', elem_id=f'delete_{tabname}') if transfer: - buttons = generation_parameters_copypaste.create_buttons(["txt2img", "img2img", "control", "extras", "caption"]) + buttons = generation_parameters_copypaste.create_buttons(["control", "txt2img", "img2img", "extras", "caption"]) else: buttons = None @@ -302,7 +312,7 @@ def create_output_panel(tabname, preview=True, prompt=None, height=None, transfe outputs=[download_files, html_log], ) delete.click(fn=call_queue.wrap_gradio_call(delete_files), show_progress='hidden', - _js="(x, y, i, j) => [x, y, ...selected_gallery_files()]", + _js=f'(x, y, i, j) => [x, y, ...selected_gallery_files("{tabname}")]', inputs=[generation_info, result_gallery, html_info, html_info], outputs=[result_gallery, html_log], ) diff --git a/modules/ui_control.py b/modules/ui_control.py index 5480c6dce..8e967621b 100644 --- a/modules/ui_control.py +++ b/modules/ui_control.py @@ -11,7 +11,7 @@ import installer gr_height = 512 max_units = shared.opts.control_max_units units: list[unit.Unit] = [] # main state variable -controls: list[gr.component] = [] # list of gr controls +controls: list[gr.components.Component] = [] # list of gr controls debug = shared.log.trace if os.environ.get('SD_CONTROL_DEBUG', None) is not None else lambda *args, **kwargs: None debug('Trace: CONTROL') @@ -89,11 +89,12 @@ def get_units(*values): def generate_click(job_id: str, state: str, active_tab: str, *args): while helpers.busy: - time.sleep(0.01) + debug(f'Control: tab="{active_tab}" job={job_id} busy') + time.sleep(0.1) from modules.control.run import control_run debug(f'Control: tab="{active_tab}" job={job_id} args={args}') progress.add_task_to_queue(job_id) - with call_queue.queue_lock: + with call_queue.get_lock(): yield [None, None, None, None, 'Control: starting', ''] shared.mem_mon.reset() jobid = shared.state.begin('Control') @@ -103,12 +104,15 @@ def generate_click(job_id: str, state: str, active_tab: str, *args): for results in control_run(state, units, helpers.input_source, helpers.input_init, helpers.input_mask, active_tab, True, *args): progress.record_results(job_id, results) yield return_controls(results, t) + except GeneratorExit: + shared.log.error("Control: generator exit") except Exception as e: shared.log.error(f"Control exception: {e}") errors.display(e, 'Control') yield [None, None, None, None, f'Control: Exception: {e}', ''] - progress.finish_task(job_id) - shared.state.end(jobid) + finally: + progress.finish_task(job_id) + shared.state.end(jobid) def create_ui(_blocks: gr.Blocks=None): @@ -248,8 +252,18 @@ def create_ui(_blocks: gr.Blocks=None): show_input.change(fn=lambda x: gr.update(visible=x), inputs=[show_input], outputs=[column_input]) show_preview.change(fn=lambda x: gr.update(visible=x), inputs=[show_preview], outputs=[column_preview]) input_type.change(fn=lambda x: gr.update(visible=x == 2), inputs=[input_type], outputs=[column_init]) - btn_prompt_counter.click(fn=call_queue.wrap_queued_call(ui_common.update_token_counter), inputs=[prompt], outputs=[prompt_counter], show_progress = 'hidden') - btn_negative_counter.click(fn=call_queue.wrap_queued_call(ui_common.update_token_counter), inputs=[negative], outputs=[negative_counter], show_progress = 'hidden') + btn_prompt_counter.click( + fn=call_queue.wrap_queued_call(ui_common.update_token_counter), + inputs=[prompt], + outputs=[prompt_counter], + show_progress = 'hidden', + ) + btn_negative_counter.click( + fn=call_queue.wrap_queued_call(ui_common.update_token_counter), + inputs=[negative], + outputs=[negative_counter], + show_progress = 'hidden', + ) select_dict = dict( fn=helpers.select_input, @@ -305,7 +319,8 @@ def create_ui(_blocks: gr.Blocks=None): _js="submit_control", inputs=[tabs_state, state, tabs_state] + input_fields + input_script_args, outputs=output_fields, - show_progress='full', + show_progress='hidden', + # queue=not shared.cmd_opts.listen, ) prompt.submit(**control_dict) negative.submit(**control_dict) diff --git a/modules/ui_control_helpers.py b/modules/ui_control_helpers.py index 6940411ef..f1545dd07 100644 --- a/modules/ui_control_helpers.py +++ b/modules/ui_control_helpers.py @@ -131,7 +131,7 @@ def process_kanvas(x): # only used when kanvas overrides gr.Image object def select_input(input_mode, input_image, init_image, init_type, input_video, input_batch, input_folder): global busy, input_source, input_init, input_mask # pylint: disable=global-statement t0 = time.time() - busy = True + busy = False selected_input = input_image # default: Image or Kanvas if input_mode == 'Video': selected_input = input_video @@ -142,8 +142,9 @@ def select_input(input_mode, input_image, init_image, init_type, input_video, in size = [gr.update(), gr.update()] if selected_input is None: input_source = None - busy = False return [gr.Tabs.update(), None, ''] + size + + busy = True input_type = type(selected_input) input_mask = None status = 'Control input | Unknown' diff --git a/modules/ui_extensions.py b/modules/ui_extensions.py index a44211d7b..112591121 100644 --- a/modules/ui_extensions.py +++ b/modules/ui_extensions.py @@ -3,7 +3,8 @@ import json import shutil import errno import html -from datetime import datetime, timedelta +import re +from datetime import datetime, timezone, timedelta import gradio as gr from modules import extensions, shared, paths, errors, ui_symbols, call_queue @@ -29,6 +30,10 @@ sort_ordering = { } +re_snake_case = re.compile(r'_(?=[a-zA-z0-9])') +re_camelCase = re.compile(r'(?<=[a-z])([A-Z])') + + def get_installed(ext): installed = [e for e in extensions.extensions if (e.remote or '').startswith(ext['url'].replace('.git', ''))] return installed[0] if len(installed) > 0 else None @@ -37,12 +42,9 @@ def get_installed(ext): def list_extensions(): global extensions_list # pylint: disable=global-statement fn = os.path.join(paths.script_path, "html", "extensions.json") - extensions_list = shared.readfile(fn, silent=True) or [] - if type(extensions_list) != list: - shared.log.warning(f'Invalid extensions list: file="{fn}"') - extensions_list = [] + extensions_list = shared.readfile(fn, silent=True, as_type="list") if len(extensions_list) == 0: - shared.log.info('Extension list is empty: refresh required') + shared.log.info("Extension List: No information found. Refresh required.") found = [] for ext in extensions.extensions: ext.read_info() @@ -90,7 +92,7 @@ def apply_changes(disable_list, update_list, disable_all): errors.display(e, f'extensions apply update: {ext.name}') shared.opts.disabled_extensions = disabled shared.opts.disable_all_extensions = disable_all - shared.opts.save(shared.config_filename) + shared.opts.save() shared.restart_server(restart=True) @@ -111,10 +113,10 @@ def check_updates(_id_task, disable_list, search_text, sort_column): ext.git_fetch() ext.read_info() commit_date = ext.commit_date or 1577836800 - shared.log.info(f'Extensions updated: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}') + shared.log.info(f'Extensions updated: {ext.name} {ext.commit_hash[:8]} {extensions.format_dt(extensions.ts2utc(commit_date), seconds=True)}') else: commit_date = ext.commit_date or 1577836800 - shared.log.debug(f'Extensions no update available: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}') + shared.log.debug(f'Extensions no update available: {ext.name} {ext.commit_hash[:8]} {extensions.format_dt(extensions.ts2utc(commit_date), seconds=True)}') except FileNotFoundError as e: if 'FETCH_HEAD' not in str(e): raise @@ -124,26 +126,27 @@ def check_updates(_id_task, disable_list, search_text, sort_column): return create_html(search_text, sort_column), "Extension update complete | Restart required" -def normalize_git_url(url) -> str: - return '' if url is None else url.removesuffix('.git') +def normalize_git_url(url: str | None) -> str: + return '' if url is None else url.strip().removesuffix('.git') def install_extension_from_url(dirname, url, branch_name, search_text, sort_column): if shared.cmd_opts.disable_extension_access: shared.log.error('Extension: apply changes disallowed because public access is enabled and insecure is not specified') return ['', ''] - if url is None or len(url) == 0: + url = normalize_git_url(url) + if not url: shared.log.error('Extension: url is not specified') return ['', ''] - if dirname is None or dirname == "": - dirname = normalize_git_url(url.split('/')[-1]) + if not dirname: + dirname = url.split('/')[-1] target_dir = os.path.join(extensions.extensions_dir, dirname) shared.log.info(f'Installing extension: {url} into {target_dir}') if os.path.exists(target_dir): shared.log.error(f'Extension: path="{target_dir}" directory already exists') return ['', ''] - url = normalize_git_url(url) - assert len([x for x in extensions.extensions if normalize_git_url(x.remote) == url]) == 0, 'Extension with this URL is already installed' + if any(normalize_git_url(x.remote) == url for x in extensions.extensions): + return ['', "Extension with this URL is already installed"] tmpdir = os.path.join(paths.data_path, "tmp", dirname) try: import git @@ -232,10 +235,10 @@ def update_extension(extension_path, search_text, sort_column): ext.git_fetch() ext.read_info() commit_date = ext.commit_date or 1577836800 - shared.log.info(f'Extensions updated: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}') + shared.log.info(f'Extensions updated: {ext.name} {ext.commit_hash[:8]} {extensions.format_dt(extensions.ts2utc(commit_date), seconds=True)}') else: commit_date = ext.commit_date or 1577836800 - shared.log.info(f'Extensions no update available: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}') + shared.log.info(f'Extensions no update available: {ext.name} {ext.commit_hash[:8]} {extensions.format_dt(extensions.ts2utc(commit_date), seconds=True)}') except FileNotFoundError as e: if 'FETCH_HEAD' not in str(e): raise @@ -249,10 +252,12 @@ def update_extension(extension_path, search_text, sort_column): def refresh_extensions_list(search_text, sort_column): global extensions_list # pylint: disable=global-statement + import ssl import urllib.request try: shared.log.debug(f'Updating extensions list: url={extensions_index}') - with urllib.request.urlopen(extensions_index, timeout=3.0) as response: + context = ssl._create_unverified_context() # pylint: disable=protected-access + with urllib.request.urlopen(extensions_index, timeout=3.0, context=context) as response: text = response.read() extensions_list = json.loads(text) with open(os.path.join(paths.script_path, "html", "extensions.json"), "w", encoding="utf-8") as outfile: @@ -271,6 +276,12 @@ def search_extensions(search_text, sort_column): return code, f'Search | {search_text} | {sort_column}' +def make_wrappable_html(text: str) -> str: + text = html.escape(text) + text = re_snake_case.sub("_", text) + return re_camelCase.sub(r"\1", text) + + def create_html(search_text, sort_column): # shared.log.debug(f'Extensions manager: refresh list search="{search_text}" sort="{sort_column}"') code = """ @@ -287,8 +298,8 @@ def create_html(search_text, sort_column): - Status - Enabled + + Extension Description Type @@ -308,13 +319,13 @@ def create_html(search_text, sort_column): ext['enabled'] = installed.enabled if installed is not None else '' ext['remote'] = installed.remote if installed is not None else None ext['path'] = installed.path if installed is not None else '' - ext['sort_default'] = f"{'1' if ext['is_builtin'] else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00')}" + ext['sort_default'] = f"{'1' if ext['is_builtin'] else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00Z')}" sort_reverse, sort_function = sort_ordering[sort_column] def dt(x: str): val = ext.get(x, None) try: - return datetime.fromisoformat(val[:-1]).strftime('%a %b%d %Y %H:%M') if val is not None else "N/A" + return extensions.format_dt(extensions.parse_isotime(val)) if val is not None else "N/A" except Exception: return 'N/A' @@ -322,22 +333,23 @@ def create_html(search_text, sort_column): for ext in sorted(extensions_list, key=sort_function, reverse=sort_reverse): installed = get_installed(ext) author = '' - updated = datetime.timestamp(datetime.now()) + updated = datetime.now(timezone.utc) # TZ-aware try: if 'github' in ext['url']: author = 'Author: ' + ext['url'].split('/')[-2].split(':')[-1] if '/' in ext['url'] else ext['url'].split(':')[1].split('/')[0] - updated = datetime.timestamp(datetime.fromisoformat(ext.get('updated', '2000-01-01T00:00:00.000Z').rstrip('Z'))) + updated = extensions.parse_isotime(ext.get('updated', '2000-01-01T00:00:00Z')) # TZ-aware else: debug(f'Extension not from github: name={ext["name"]} url={ext["url"]}') except Exception as e: debug(f'Extension get updated error: name={ext["name"]} url={ext["url"]} {e}') - update_available = (installed is not None) and (ext['remote'] is not None) and (ext['commit_date'] > updated) + local_ver_date = extensions.ts2utc(ext['commit_date']) # TZ-aware + update_available = (installed is not None) and (not ext['is_builtin']) and (ext['remote'] is not None) and (updated > local_ver_date) # TZ-aware if update_available: - debug(f'Extension update available: name={ext["name"]} updated={updated}/{datetime.utcfromtimestamp(updated)} commit={ext["commit_date"]}/{datetime.utcfromtimestamp(ext["commit_date"])}') + debug(f'Extension update available: name={ext["name"]} updated={extensions.format_dt(updated, seconds=True)} commit={extensions.format_dt(local_ver_date, seconds=True)}') # TZ-aware ext['sort_user'] = f"{'0' if ext['is_builtin'] else '1'}{'1' if ext['installed'] else '0'}{ext.get('name', '')}" - ext['sort_enabled'] = f"{'0' if ext['enabled'] else '1'}{'1' if ext['is_builtin'] else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00')}" - ext['sort_update'] = f"{'1' if update_available else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00')}" - delta = datetime.now() - datetime.fromisoformat(ext.get('created', '2000-01-01T00:00Z')[:-1]) + ext['sort_enabled'] = f"{'0' if ext['enabled'] else '1'}{'1' if ext['is_builtin'] else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00Z')}" + ext['sort_update'] = f"{'1' if update_available else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00Z')}" + delta = datetime.now(timezone.utc) - extensions.parse_isotime(ext.get('created', '2000-01-01T00:00Z')) # TZ-aware to prep for 3.11+ datetime.fromisoformat() behavior ext['sort_trending'] = round(ext.get('stars', 0) / max(delta.days, 5), 1) tags = ext.get("tags", []) if not isinstance(tags, list): @@ -370,7 +382,7 @@ def create_html(search_text, sort_column): stats['enabled'] += 1 type_code = f"""
{"SYSTEM" if ext['is_builtin'] else 'USER'}
""" version_code = f"""
{ext['version']}
""" - enabled_code = f"""""" + enabled_code = f"""""" masked_path = html.escape(ext.get("path", "").replace('\\', '/')) if not ext['is_builtin']: install_code = f"""""" @@ -382,36 +394,36 @@ def create_html(search_text, sort_column): if ext.get('status', None) is None or type(ext['status']) == str: # old format ext['status'] = 0 if ext['url'] is None or ext['url'] == '': - status = f"
{ui_symbols.svg_bullet.color('#00C0FD')}
" + status = f"
{ui_symbols.svg_bullet.style('#00C0FD')}
" elif ext['status'] > 0: if ext['status'] == 1: - status = f"
{ui_symbols.svg_bullet.color('#00FD9C')}
" + status = f"
{ui_symbols.svg_bullet.style('#00FD9C')}
" elif ext['status'] == 2: - status = f"
{ui_symbols.svg_bullet.color('#FFC300')}
" + status = f"
{ui_symbols.svg_bullet.style('#FFC300')}
" elif ext['status'] == 3: - status = f"
{ui_symbols.svg_bullet.color('#FFC300')}
" + status = f"
{ui_symbols.svg_bullet.style('#FFC300')}
" elif ext['status'] == 4: - status = f"
{ui_symbols.svg_bullet.color('#4E22FF')}
" + status = f"
{ui_symbols.svg_bullet.style('#4E22FF')}
" elif ext['status'] == 5: - status = f"
{ui_symbols.svg_bullet.color('#CE0000')}
" + status = f"
{ui_symbols.svg_bullet.style('#CE0000')}
" elif ext['status'] == 6: - status = f"
{ui_symbols.svg_bullet.color('#AEAEAE')}
" + status = f"
{ui_symbols.svg_bullet.style('#AEAEAE')}
" else: - status = f"
{ui_symbols.svg_bullet.color('#008EBC')}
" + status = f"
{ui_symbols.svg_bullet.style('#008EBC')}
" else: - if updated < datetime.timestamp(datetime.now() - timedelta(6*30)): - status = f"
{ui_symbols.svg_bullet.color('#C000CF')}
" + if updated < datetime.now(timezone.utc) - timedelta(6*30): # TZ-aware + status = f"
{ui_symbols.svg_bullet.style('#C000CF')}
" else: - status = f"
{ui_symbols.svg_bullet.color('#7C7C7C')}
" + status = f"
{ui_symbols.svg_bullet.style('#7C7C7C')}
" code += f""" - {status} {enabled_code} - {html.escape(ext.get("name", "unknown"))}
{tags_text} + {status} + {make_wrappable_html(ext.get("name", "unknown"))}
{tags_text} {html.escape(ext.get("description", ""))} -

Created {html.escape(dt('created'))} | Added {html.escape(dt('added'))} | Pushed {html.escape(dt('pushed'))} | Updated {html.escape(dt('updated'))}

-

{author} | Stars {html.escape(str(ext.get('stars', 0)))} | Size {html.escape(str(ext.get('size', 0)))} | Commits {html.escape(str(ext.get('commits', 0)))} | Issues {html.escape(str(ext.get('issues', 0)))} | Trending {html.escape(str(ext['sort_trending']))}

+

Created: {html.escape(dt('created'))} | Added: {html.escape(dt('added'))} | Pushed: {html.escape(dt('pushed'))} | Updated: {html.escape(dt('updated'))}

+

{author} | Stars: {html.escape(str(ext.get('stars', 0)))} | Size: {html.escape(str(ext.get('size', 0)))} | Commits: {html.escape(str(ext.get('commits', 0)))} | Issues: {html.escape(str(ext.get('issues', 0)))} | Trending: {html.escape(str(ext['sort_trending']))}

{type_code} {version_code} diff --git a/modules/ui_extra_networks.py b/modules/ui_extra_networks.py index 7d2eccae2..155f182c2 100644 --- a/modules/ui_extra_networks.py +++ b/modules/ui_extra_networks.py @@ -70,49 +70,49 @@ def init_api(): return FileResponse(filename, headers={"Accept-Ranges": "bytes"}) def get_metadata(page: str = "", item: str = ""): - page = next(iter([x for x in shared.extra_networks if x.name.lower() == page.lower()]), None) - if page is None: + page_dict = next(iter([x for x in shared.extra_networks if x.name.lower() == page.lower()]), None) + if page_dict is None: return JSONResponse({ 'metadata': 'none' }) - metadata = page.metadata.get(item, 'none') + metadata = page_dict.metadata.get(item, 'none') if metadata is None: metadata = '' # shared.log.debug(f"Networks metadata: page='{page}' item={item} len={len(metadata)}") return JSONResponse({"metadata": metadata}) def get_info(page: str = "", item: str = ""): - page = next(iter([x for x in get_pages() if x.name.lower() == page.lower()]), None) - if page is None: + page_dict = next(iter([x for x in get_pages() if x.name.lower() == page.lower()]), None) + if page_dict is None: return JSONResponse({ 'info': 'none' }) - item = next(iter([x for x in page.items if x['name'].lower() == item.lower()]), None) - if item is None: + item_dict = next(iter([x for x in page_dict.items if x['name'].lower() == item.lower()]), None) + if item_dict is None: return JSONResponse({ 'info': 'none' }) - info = page.find_info(item.get('filename', None) or item.get('name', None)) + info = page_dict.find_info(item_dict.get('filename', None) or item_dict.get('name', None)) if info is None: info = {} # shared.log.debug(f"Networks info: page='{page.name}' item={item['name']} len={len(info)}") return JSONResponse({"info": info}) def get_desc(page: str = "", item: str = ""): - page = next(iter([x for x in get_pages() if x.name.lower() == page.lower()]), None) - if page is None: + page_dict = next(iter([x for x in get_pages() if x.name.lower() == page.lower()]), None) + if page_dict is None: return JSONResponse({ 'description': 'none' }) - item = next(iter([x for x in page.items if x['name'].lower() == item.lower()]), None) - if item is None: + item_dict = next(iter([x for x in page_dict.items if x['name'].lower() == item.lower()]), None) + if item_dict is None: return JSONResponse({ 'description': 'none' }) - desc = page.find_description(item.get('filename', None) or item.get('name', None)) + desc = page_dict.find_description(item_dict.get('filename', None) or item_dict.get('name', None)) if desc is None: desc = '' # shared.log.debug(f"Networks desc: page='{page.name}' item={item['name']} len={len(desc)}") return JSONResponse({"description": desc}) def get_network(page: str = "", item: str = ""): - page = next(iter([x for x in get_pages() if x.name.lower() == page.lower()]), None) - if page is None: + page_dict = next(iter([x for x in get_pages() if x.name.lower() == page.lower()]), None) + if page_dict is None: return JSONResponse({ 'page': 'none' }) - item = next(iter([x for x in page.items if (x['alias'].lower() == item.lower() or x['name'].lower() == item.lower())]), None) - if item is None: + item_dict = next(iter([x for x in page_dict.items if (x['alias'].lower() == item.lower() or x['name'].lower() == item.lower())]), None) + if item_dict is None: return JSONResponse({ 'item': 'none' }) - obj = json.dumps(item, cls=DateTimeEncoder) + obj = json.dumps(item_dict, cls=DateTimeEncoder) return JSONResponse(obj) shared.api.add_api_route("/sdapi/v1/network", get_network, methods=["GET"]) @@ -153,11 +153,21 @@ class ExtraNetworksPage: def __str__(self): return f'Page(title="{self.title}" name="{self.name}" items={len(self.items)})' + def switch_view(self, tabname: str): + new_view = 'gallery' if self.view == 'list' else 'list' + self.view = new_view + self.card = card_full if new_view == 'gallery' else card_list + self.html = '' + self.create_page(tabname) + if shared.opts.extra_networks_view != new_view: + shared.opts.extra_networks_view = new_view + shared.opts.save() + def refresh(self): pass def patch(self, text: str, tabname: str): - return text.replace('~tabname', tabname) + return text.replace('~tabname', tabname).replace('txt2img', tabname) def create_xyz_grid(self): pass @@ -193,6 +203,7 @@ class ExtraNetworksPage: def get_exif(self, image: Image.Image): import piexif + import piexif.helper try: exifinfo = image.getexif() if exifinfo is not None and len(exifinfo) > 0: @@ -326,6 +337,7 @@ class ExtraNetworksPage: else: style = 'network-folder' subdirs_html += f'
' + self.html = '' self.create_items(tabname) versions = sorted({item.get("version", "") for item in self.items if item.get("version")}) @@ -506,6 +518,9 @@ class ExtraNetworksPage: pass if info is None: info = self.find_info(path) + if not isinstance(info, dict): + self.desc_time += time.time() - t0 + return '' desc = info.get('description', '') or '' f = HTMLFilter() f.feed(desc) @@ -573,7 +588,7 @@ def register_pages(): def get_pages(title=None): visible = shared.opts.extra_networks - pages = [] + pages: list[ExtraNetworksPage] = [] if 'All' in visible or visible == []: # default en sort order visible = ['Model', 'Lora', 'Style', 'Wildcards', 'Embedding', 'VAE', 'History', 'Hypernetwork'] @@ -1007,12 +1022,7 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): def ui_view_cards(title): pages = [] for page in get_pages(): - shared.opts.extra_networks_view = page.view - # shared.opts.save(shared.config_filename) - page.view = 'gallery' if page.view == 'list' else 'list' - page.card = card_full if page.view == 'gallery' else card_list - page.html = '' - page.create_page(ui.tabname) + page.switch_view(ui.tabname) shared.log.debug(f'Networks: refresh page="{page.title}" items={len(page.items)} tab={ui.tabname} view={page.view}') pages.append(page.html) ui.search.update(title) @@ -1071,7 +1081,7 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): def ui_sort_cards(sort_order): if shared.opts.extra_networks_sort != sort_order: shared.opts.extra_networks_sort = sort_order - shared.opts.save(shared.config_filename) + shared.opts.save() return f'Networks: sort={sort_order}' dummy = gr.State(value=False) # pylint: disable=abstract-class-instantiated diff --git a/modules/ui_extra_networks_checkpoints.py b/modules/ui_extra_networks_checkpoints.py index 7ade3743e..df6681bca 100644 --- a/modules/ui_extra_networks_checkpoints.py +++ b/modules/ui_extra_networks_checkpoints.py @@ -19,6 +19,8 @@ version_map = { "StableDiffusionXL": "SD XL", "WanToVideo": "Wan", "WanVACE": "Wan", + "Z": "Z-Image", + "Glm": "GLM-Image", } class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage): @@ -40,7 +42,19 @@ class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage): shared.log.debug(f'Networks: type="reference" autodownload={shared.opts.sd_checkpoint_autodownload} enable={shared.opts.extra_network_reference_enable}') return [] count = { 'total': 0, 'ready': 0, 'hidden': 0, 'experimental': 0, 'base': 0 } - shared.reference_models = readfile(os.path.join('html', 'reference.json'), as_type="dict") + + reference_base = readfile(os.path.join('html', 'reference.json'), as_type="dict") + reference_quant = readfile(os.path.join('html', 'reference-quant.json'), as_type="dict") + reference_distilled = readfile(os.path.join('html', 'reference-distilled.json'), as_type="dict") + reference_community = readfile(os.path.join('html', 'reference-community.json'), as_type="dict") + reference_cloud = readfile(os.path.join('html', 'reference-cloud.json'), as_type="dict") + shared.reference_models = {} + shared.reference_models.update(reference_base) + shared.reference_models.update(reference_quant) + shared.reference_models.update(reference_community) + shared.reference_models.update(reference_distilled) + shared.reference_models.update(reference_cloud) + for k, v in shared.reference_models.items(): count['total'] += 1 url = v['path'] @@ -79,7 +93,7 @@ class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage): ready = reference_downloaded(url) version = "ready" if ready else "download" if tag == 'cloud': - version = 'cloud' + version = 'Cloud' if not ready and shared.opts.offline_mode: count['hidden'] += 1 continue @@ -103,7 +117,7 @@ class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage): "version": version, "tags": tag, } - shared.log.debug(f'Networks: type="reference" items={count}') + shared.log.debug(f'Networks: type="reference" {count}') def create_item(self, name): record = None diff --git a/modules/ui_extra_networks_lora.py b/modules/ui_extra_networks_lora.py index 6b9705674..5a11f85ab 100644 --- a/modules/ui_extra_networks_lora.py +++ b/modules/ui_extra_networks_lora.py @@ -64,6 +64,11 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): clean_tags.pop('dataset', None) return clean_tags + def cleanup_version(self, dct, lora): + ver = dct.get("baseModel", lora.sd_version) + ver = ver.replace(' 0.9', '').replace(' 1.0', '').replace(' ', '') + return ver + def create_item(self, name): l = lora_load.available_networks.get(name) if l is None: @@ -74,7 +79,7 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): name = os.path.splitext(os.path.relpath(l.filename, shared.cmd_opts.lora_dir))[0] size, mtime = modelstats.stat(l.filename) info = self.find_info(l.filename) - version = self.find_version(l, info) + ver_dct = self.find_version(l, info) item = { "type": 'Lora', "name": name, @@ -85,10 +90,10 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): "metadata": json.dumps(l.metadata, indent=4) if l.metadata else None, "mtime": mtime, "size": size, - "version": version.get("baseModel", l.sd_version), + "version": self.cleanup_version(ver_dct, l), "info": info, "description": self.find_description(l.filename, info), - "tags": self.get_tags(l, info, version), + "tags": self.get_tags(l, info, ver_dct), } return item except Exception as e: diff --git a/modules/ui_gallery.py b/modules/ui_gallery.py index 283b720a8..f22bc1a1a 100644 --- a/modules/ui_gallery.py +++ b/modules/ui_gallery.py @@ -68,6 +68,7 @@ def create_ui(): sort_buttons.append(ToolButton(value=ui_symbols.sort_time_dsc, elem_classes=['gallery-sort'])) gr.Textbox(show_label=False, placeholder='Search', elem_id='tab-gallery-search') gr.HTML('', elem_id='tab-gallery-status') + gr.HTML('', elem_id='tab-gallery-progress') for btn in sort_buttons: btn.click(fn=None, _js='gallerySort', inputs=[btn], outputs=[]) with gr.Row(): diff --git a/modules/ui_javascript.py b/modules/ui_javascript.py index dd482e2a9..dcbf14731 100644 --- a/modules/ui_javascript.py +++ b/modules/ui_javascript.py @@ -31,7 +31,7 @@ def html_head(): for script in modules.scripts_manager.list_scripts("javascript", ".js"): if script.filename in main or script.filename in skip: continue - if '.esm' in js or '.mjs' in js: + if '.esm' in script.filename or '.mjs' in script.filename: head += f'\n' else: head += f'\n' diff --git a/modules/ui_loadsave.py b/modules/ui_loadsave.py index fbb1cbeba..ff054f686 100644 --- a/modules/ui_loadsave.py +++ b/modules/ui_loadsave.py @@ -1,3 +1,4 @@ +from typing import TYPE_CHECKING, cast import os import gradio as gr from modules import errors @@ -90,10 +91,11 @@ class UiLoadsave: apply_field(x, 'value', check_dropdown, getattr(x, 'init_field', None)) def check_tab_id(tab_id): - tab_items = list(filter(lambda e: isinstance(e, gr.TabItem), x.children)) + if TYPE_CHECKING: + assert isinstance(x, gr.Tabs) + tab_items = cast('list[gr.TabItem]', list(filter(lambda e: isinstance(e, gr.TabItem), x.children))) # Force static type checker to get correct type if type(tab_id) == str: - tab_ids = [t.id for t in tab_items] - return tab_id in tab_ids + return tab_id in [t.id for t in tab_items] elif type(tab_id) == int: return 0 <= tab_id < len(tab_items) else: @@ -290,11 +292,12 @@ class UiLoadsave: self.ui_defaults_review = gr.HTML("", elem_id="ui_defaults_review") def setup_ui(self): + review = [self.ui_defaults_review] if self.ui_defaults_review is not None else None if self.ui_defaults_view: - self.ui_defaults_view.click(fn=self.ui_view, inputs=list(self.component_mapping.values()), outputs=[self.ui_defaults_review]) + self.ui_defaults_view.click(fn=self.ui_view, inputs=list(self.component_mapping.values()), outputs=review) if self.ui_defaults_apply: - self.ui_defaults_apply.click(fn=self.ui_apply, inputs=list(self.component_mapping.values()), outputs=[self.ui_defaults_review]) + self.ui_defaults_apply.click(fn=self.ui_apply, inputs=list(self.component_mapping.values()), outputs=review) if self.ui_defaults_restore: - self.ui_defaults_restore.click(fn=self.ui_restore, inputs=[], outputs=[self.ui_defaults_review]) + self.ui_defaults_restore.click(fn=self.ui_restore, inputs=[], outputs=review) if self.ui_defaults_submenu: - self.ui_defaults_submenu.click(fn=self.ui_submenu_apply, _js='uiOpenSubmenus', inputs=[self.ui_defaults_review], outputs=[self.ui_defaults_review]) + self.ui_defaults_submenu.click(fn=self.ui_submenu_apply, _js='uiOpenSubmenus', inputs=review, outputs=review) diff --git a/modules/ui_models.py b/modules/ui_models.py index 909613c14..def5ec3c8 100644 --- a/modules/ui_models.py +++ b/modules/ui_models.py @@ -1,5 +1,6 @@ import os import inspect +from typing import cast import gradio as gr from modules import errors, sd_models, sd_vae, extras, sd_samplers, ui_symbols, modelstats from modules.ui_components import ToolButton @@ -143,7 +144,7 @@ def create_ui(): model_table = gr.HTML(value='', elem_id="model_list_table") model_checkhash_btn.click(fn=sd_models.update_model_hashes, inputs=[], outputs=[model_table]) - model_list_btn.click(fn=lambda: create_models_table(sd_models.checkpoints_list.values()), inputs=[], outputs=[model_table]) + model_list_btn.click(fn=lambda: create_models_table(list(sd_models.checkpoints_list.values())), inputs=[], outputs=[model_table]) with gr.Tab(label="Metadata", elem_id="models_metadata_tab"): from modules.civitai.metadata_civitai import civit_search_metadata, civit_update_metadata @@ -178,7 +179,7 @@ def create_ui(): custom_name = gr.Textbox(label="New model name") with gr.Row(): merge_mode = gr.Dropdown(choices=merge_methods.__all__, value="weighted_sum", label="Interpolation Method") - merge_mode_docs = gr.HTML(value=getattr(merge_methods, "weighted_sum", "").__doc__.replace("\n", "
")) + merge_mode_docs = gr.HTML(value=merge_methods.weighted_sum.__doc__.strip().replace("\n", "
")) # pylint: disable=no-member # pyright: ignore[reportOptionalMemberAccess] with gr.Row(): primary_model_name = gr.Dropdown(sd_model_choices(), label="Primary model", value="None") create_refresh_button(primary_model_name, sd_models.list_models, lambda: {"choices": sd_model_choices()}, "checkpoint_A_refresh") @@ -318,7 +319,11 @@ def create_ui(): return gr.Slider.update(value=None, visible=False) def show_help(mode): - doc = getattr(merge_methods, mode).__doc__.replace("\n", "
") + try: + doc = getattr(merge_methods, mode).__doc__.strip().replace("\n", "
") + except AttributeError: + log.warning(f'Merge mode "{mode}" is missing documentation') + doc = "Error: Documentation missing" return gr.update(value=doc, visible=True) def show_unload(device): @@ -358,8 +363,8 @@ def create_ui(): merge_mode.input(fn=tertiary, inputs=merge_mode, outputs=[tertiary_model_name, tertiary_refresh]) merge_mode.input(fn=beta_visibility, inputs=merge_mode, outputs=[beta, alpha_label, beta_label, beta_apply_preset, beta_preset, beta_base, beta_in_blocks, beta_mid_block, beta_out_blocks]) re_basin.change(fn=show_iters, inputs=re_basin, outputs=re_basin_iterations) - apply_preset.click(fn=load_presets, inputs=[alpha_preset, alpha_preset_lambda], outputs=[alpha_base, alpha_in_blocks, alpha_mid_block, alpha_out_blocks, tabs]) - beta_apply_preset.click(fn=load_presets, inputs=[beta_preset, beta_preset_lambda], outputs=[beta_base, beta_in_blocks, beta_mid_block, beta_out_blocks, tabs]) + apply_preset.click(fn=load_presets, inputs=[alpha_preset, alpha_preset_lambda], outputs=[alpha_base, alpha_in_blocks, alpha_mid_block, alpha_out_blocks, cast("gr.components.Component", tabs)]) # Casting because Tabs has an update method. + beta_apply_preset.click(fn=load_presets, inputs=[beta_preset, beta_preset_lambda], outputs=[beta_base, beta_in_blocks, beta_mid_block, beta_out_blocks, cast("gr.components.Component", tabs)]) # Casting because Tabs has an update method. modelmerger_merge.click( fn=wrap_gradio_gpu_call(modelmerger, extra_outputs=lambda: [gr.update() for _ in range(4)], name='Models'), diff --git a/modules/ui_sections.py b/modules/ui_sections.py index 0d57eccb2..2c4e11b78 100644 --- a/modules/ui_sections.py +++ b/modules/ui_sections.py @@ -200,51 +200,51 @@ def create_sampler_options(tabname): shared.opts.data['schedulers_use_loworder'] = 'low order' in sampler_options shared.opts.data['schedulers_rescale_betas'] = 'rescale' in sampler_options shared.log.debug(f'Sampler set options: {sampler_options}') - shared.opts.save(shared.config_filename, silent=True) + shared.opts.save(silent=True) def set_sampler_timesteps(timesteps): shared.log.debug(f'Sampler set options: timesteps={timesteps}') shared.opts.schedulers_timesteps = timesteps - shared.opts.save(shared.config_filename, silent=True) + shared.opts.save(silent=True) def set_sampler_spacing(spacing): shared.log.debug(f'Sampler set options: spacing={spacing}') shared.opts.schedulers_timestep_spacing = spacing - shared.opts.save(shared.config_filename, silent=True) + shared.opts.save(silent=True) def set_sampler_sigma(sampler_sigma): shared.log.debug(f'Sampler set options: sigma={sampler_sigma}') shared.opts.schedulers_sigma = sampler_sigma - shared.opts.save(shared.config_filename, silent=True) + shared.opts.save(silent=True) def set_sampler_order(sampler_order): shared.log.debug(f'Sampler set options: order={sampler_order}') shared.opts.schedulers_solver_order = sampler_order - shared.opts.save(shared.config_filename, silent=True) + shared.opts.save(silent=True) def set_sampler_prediction(sampler_prediction): shared.log.debug(f'Sampler set options: prediction={sampler_prediction}') shared.opts.schedulers_prediction_type = sampler_prediction - shared.opts.save(shared.config_filename, silent=True) + shared.opts.save(silent=True) def set_sampler_beta(sampler_beta): shared.log.debug(f'Sampler set options: beta={sampler_beta}') shared.opts.schedulers_beta_schedule = sampler_beta - shared.opts.save(shared.config_filename, silent=True) + shared.opts.save(silent=True) def set_sampler_shift(sampler_shift, sampler_base_shift, sampler_max_shift): shared.log.debug(f'Sampler set options: shift={sampler_shift} base={sampler_base_shift} max={sampler_max_shift}') shared.opts.schedulers_shift = sampler_shift shared.opts.schedulers_base_shift = sampler_base_shift shared.opts.schedulers_max_shift = sampler_max_shift - shared.opts.save(shared.config_filename, silent=True) + shared.opts.save(silent=True) def set_sigma_adjust(val, start, end): shared.log.debug(f'Sampler set options: sigma={val} min={start} max={end}') shared.opts.schedulers_sigma_adjust = val shared.opts.schedulers_sigma_adjust_min = start shared.opts.schedulers_sigma_adjust_max = end - shared.opts.save(shared.config_filename, silent=True) + shared.opts.save(silent=True) # 'linear', 'scaled_linear', 'squaredcos_cap_v2' def set_sampler_preset(preset): @@ -258,7 +258,7 @@ def create_sampler_options(tabname): sampler_sigma = gr.Dropdown(label='Sigma method', elem_id=f"{tabname}_sampler_sigma", choices=['default', 'karras', 'betas', 'exponential', 'lambdas', 'flowmatch'], value=shared.opts.schedulers_sigma, type='value') sampler_spacing = gr.Dropdown(label='Timestep spacing', elem_id=f"{tabname}_sampler_spacing", choices=['default', 'linspace', 'leading', 'trailing'], value=shared.opts.schedulers_timestep_spacing, type='value') with gr.Row(elem_classes=['flex-break']): - sampler_beta = gr.Dropdown(label='Beta schedule', elem_id=f"{tabname}_sampler_beta", choices=['default', 'linear', 'scaled', 'cosine', 'sigmoid'], value=shared.opts.schedulers_beta_schedule, type='value') + sampler_beta = gr.Dropdown(label='Beta schedule', elem_id=f"{tabname}_sampler_beta", choices=['default', 'linear', 'scaled', 'cosine', 'sigmoid', 'laplace'], value=shared.opts.schedulers_beta_schedule, type='value') sampler_prediction = gr.Dropdown(label='Prediction method', elem_id=f"{tabname}_sampler_prediction", choices=['default', 'epsilon', 'sample', 'v_prediction', 'flow_prediction'], value=shared.opts.schedulers_prediction_type, type='value') with gr.Row(elem_classes=['flex-break']): sampler_presets = gr.Dropdown(label='Timesteps presets', elem_id=f"{tabname}_sampler_presets", choices=['None', 'AYS SD15', 'AYS SDXL'], value='None', type='value') diff --git a/modules/ui_settings.py b/modules/ui_settings.py index fc6bfa1e2..345077ec9 100644 --- a/modules/ui_settings.py +++ b/modules/ui_settings.py @@ -37,7 +37,7 @@ def apply_setting(key, value): shared.opts.data[key] = valtype(value) if valtype != type(None) else value if oldval != value and shared.opts.data_labels[key].onchange is not None: shared.opts.data_labels[key].onchange() - shared.opts.save(shared.config_filename) + shared.opts.save() return getattr(shared.opts, key) @@ -149,7 +149,7 @@ def run_settings(*args): shared.opts.sd_backend = "diffusers" try: if len(changed) > 0: - shared.opts.save(shared.config_filename) + shared.opts.save() shared.log.info(f'Settings: changed={len(changed)} {changed}') except RuntimeError: shared.log.error(f'Settings failed: change={len(changed)} {changed}') @@ -167,7 +167,7 @@ def run_settings_single(value, key, progress=False): if shared.cmd_opts.use_directml: from modules.dml import directml_override_opts directml_override_opts() - shared.opts.save(shared.config_filename) + shared.opts.save() if key not in ['sd_model_checkpoint', 'sd_model_refiner', 'sd_vae', 'sd_te', 'sd_unet']: shared.log.debug(f'Setting changed: {key}={value} progress={progress}') return get_value_for_setting(key), shared.opts.dumpjson() diff --git a/modules/ui_symbols.py b/modules/ui_symbols.py index ddc16c273..c880a4a8f 100644 --- a/modules/ui_symbols.py +++ b/modules/ui_symbols.py @@ -1,3 +1,10 @@ +import re +from functools import lru_cache +from typing import final + + +# Basic symbols + refresh = '⟲' close = '✕' load = '⇧' @@ -41,23 +48,45 @@ sort_time_dsc = '\uf0dd' style_apply = '↶' style_save = '↷' +# Configurable symbols + +@final class SVGSymbol: + __created = [] + __re_display = re.compile(r"(?<=display:)\s*([\w\-]+)(?=;)") + + @classmethod + @lru_cache # Class method due to B019, but also mostly so the `style` method shows params in IDE + def __stylize(cls, svg: str, color: str | None = None, display: str | None = None): + if color: + svg = re.sub("currentColor", color, svg) + if display: + svg = cls.__re_display.sub(display, svg, count=1) + return svg + def __init__(self, svg: str): + svg = re.sub(r"\s{2,}", " ", svg.replace("\n", "")).replace("> <", "><").strip() + if svg in self.__created: + raise RuntimeError("SVGSymbol class was created with an existing value. There should only be one instance per symbol.", svg) + else: + self.__created.append(svg) self.svg = svg - self.before = "" - self.after = "" self.supports_color = False + self.supports_display = False if "currentColor" in self.svg: self.supports_color = True - self.before, self.after = self.svg.split("currentColor", maxsplit=1) + if self.__re_display.search(self.svg): + self.supports_display = True - def color(self, color: str): - if self.supports_color: - return self.before + color + self.after - else: - return self.svg + def style(self, color: str | None = None, display: str | None = None) -> str: + style_args = { + "color": color if color and self.supports_color else None, + "display": display if display and self.supports_display else None + } + return self.__stylize(self.svg, **style_args) def __str__(self): return self.svg -svg_bullet = SVGSymbol("") + +svg_bullet = SVGSymbol("") diff --git a/modules/video_models/google_veo.py b/modules/video_models/google_veo.py index 34ba99d55..aebc3f22f 100644 --- a/modules/video_models/google_veo.py +++ b/modules/video_models/google_veo.py @@ -71,24 +71,42 @@ class GoogleVeoVideoPipeline(): def get_args(self): from modules.shared import opts - api_key = os.getenv("GOOGLE_API_KEY") or opts.google_api_key - vertex_credentials = os.getenv("GOOGLE_APPLICATION_CREDENTIALS") - if (api_key is None or len(api_key) == 0) and (vertex_credentials is None or len(vertex_credentials) == 0): - log.error(f'Cloud: model="{self.model}" API key not provided') - return None - use_vertexai = (os.getenv("GOOGLE_GENAI_USE_VERTEXAI") is not None) or opts.google_use_vertexai - project_id = os.getenv("GOOGLE_CLOUD_PROJECT") or opts.google_project_id - location_id = os.getenv("GOOGLE_CLOUD_LOCATION") or opts.google_location_id - args = { - 'api_key': api_key, - 'vertexai': use_vertexai, - 'project': project_id if len(project_id) > 0 else None, - 'location': location_id if len(location_id) > 0 else None, - } - args_copy = args.copy() - args_copy['api_key'] = '...' + args_copy['api_key'][-4:] # last 4 chars - args_copy['credentials'] = vertex_credentials - log.debug(f'Cloud: model="{self.model}" args={args_copy}') + # Use UI settings only - env vars are intentionally ignored + api_key = opts.google_api_key + project_id = opts.google_project_id + location_id = opts.google_location_id + use_vertexai = opts.google_use_vertexai + + has_api_key = api_key and len(api_key) > 0 + has_project = project_id and len(project_id) > 0 + has_location = location_id and len(location_id) > 0 + + if use_vertexai: + if has_api_key and (has_project or has_location): + # Invalid: can't have both api_key AND project/location + log.error(f'Cloud: model="{self.model}" API key and project/location are mutually exclusive') + return None + elif has_api_key: + # Vertex AI Express Mode: api_key + vertexai, no project/location + args = {'api_key': api_key, 'vertexai': True} + elif has_project and has_location: + # Standard Vertex AI: project/location, no api_key + args = {'vertexai': True, 'project': project_id, 'location': location_id} + else: + log.error(f'Cloud: model="{self.model}" Vertex AI requires either API key (Express Mode) or project ID + location ID') + return None + else: + # Gemini Developer API: api_key only + if not has_api_key: + log.error(f'Cloud: model="{self.model}" API key not provided') + return None + args = {'api_key': api_key} + + # Debug logging + args_log = args.copy() + if args_log.get('api_key'): + args_log['api_key'] = '...' + args_log['api_key'][-4:] + log.debug(f'Cloud: model="{self.model}" args={args_log}') return args def __call__(self, prompt: list[str], width: int, height: int, image: Image.Image = None, num_frames: int = 4*24): diff --git a/modules/video_models/models_def.py b/modules/video_models/models_def.py index c89995da8..06b64fdd1 100644 --- a/modules/video_models/models_def.py +++ b/modules/video_models/models_def.py @@ -128,13 +128,37 @@ try: ], 'LTX Video': [ Model(name='None'), - Model(name='LTXVideo 0.9.8 13B', + Model(name='LTXVideo 2 19B T2V Dev', + url='https://huggingface.co/Lightricks/LTX-2', + repo='Lightricks/LTX-2', + repo_cls=getattr(diffusers, 'LTX2Pipeline', None), + te_cls=getattr(transformers, 'Gemma3ForConditionalGeneration', None), + dit_cls=getattr(diffusers, 'LTX2VideoTransformer3DModel', None)), + Model(name='LTXVideo 2 19B I2V Dev', + url='https://huggingface.co/Lightricks/LTX-2', + repo='Lightricks/LTX-2', + repo_cls=getattr(diffusers, 'LTX2ImageToVideoPipeline', None), + te_cls=getattr(transformers, 'Gemma3ForConditionalGeneration', None), + dit_cls=getattr(diffusers, 'LTX2VideoTransformer3DModel', None)), + Model(name='LTXVideo 2 19B T2V Dev SDNQ', + url='https://huggingface.co/Disty0/LTX-2-SDNQ-4bit-dynamic', + repo='Disty0/LTX-2-SDNQ-4bit-dynamic', + repo_cls=getattr(diffusers, 'LTX2Pipeline', None), + te_cls=getattr(transformers, 'Gemma3ForConditionalGeneration', None), + dit_cls=getattr(diffusers, 'LTX2VideoTransformer3DModel', None)), + Model(name='LTXVideo 2 19B I2V Dev SDNQ', + url='https://huggingface.co/Disty0/LTX-2-SDNQ-4bit-dynamic', + repo='Disty0/LTX-2-SDNQ-4bit-dynamic', + repo_cls=getattr(diffusers, 'LTX2ImageToVideoPipeline', None), + te_cls=getattr(transformers, 'Gemma3ForConditionalGeneration', None), + dit_cls=getattr(diffusers, 'LTX2VideoTransformer3DModel', None)), + Model(name='LTXVideo 0.9.8 13B Distilled', url='https://huggingface.co/Lightricks/LTX-Video-0.9.8-13B-distilled', repo='Lightricks/LTX-Video-0.9.8-13B-distilled', repo_cls=getattr(diffusers, 'LTXConditionPipeline', None), te_cls=getattr(transformers, 'T5EncoderModel', None), dit_cls=getattr(diffusers, 'LTXVideoTransformer3DModel', None)), - Model(name='LTXVideo 0.9.7 13B', + Model(name='LTXVideo 0.9.7 13B Dev', url='https://huggingface.co/Lightricks/LTX-Video-0.9.7-dev', repo='a-r-r-o-w/LTX-Video-0.9.7-diffusers', repo_cls=getattr(diffusers, 'LTXConditionPipeline', None), diff --git a/modules/video_models/video_run.py b/modules/video_models/video_run.py index 596497b6f..8f478fa15 100644 --- a/modules/video_models/video_run.py +++ b/modules/video_models/video_run.py @@ -3,6 +3,7 @@ import copy import time from modules import shared, errors, sd_models, processing, devices, images, ui_common from modules.video_models import models_def, video_utils, video_load, video_vae, video_overrides, video_save, video_prompt +from modules.paths import resolve_output_path debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -56,7 +57,7 @@ def generate(*args, **kwargs): p.state = ui_state p.do_not_save_grid = True p.do_not_save_samples = not mp4_frames - p.outpath_samples = shared.opts.outdir_samples or shared.opts.outdir_video + p.outpath_samples = resolve_output_path(shared.opts.outdir_samples, shared.opts.outdir_video) if 'T2V' in model: if init_image is not None: shared.log.warning('Video: op=T2V init image not supported') @@ -154,9 +155,15 @@ def generate(*args, **kwargs): pixels = video_save.images_to_tensor(processed.images) else: pixels = None + if hasattr(processed, 'audio') and processed.audio is not None: + audio = processed.audio[0].float().cpu() + else: + audio = None + _num_frames, video_file = video_save.save_video( p=p, pixels=pixels, + audio=audio, binary=processed.bytes, mp4_fps=mp4_fps, mp4_codec=mp4_codec, diff --git a/modules/video_models/video_save.py b/modules/video_models/video_save.py index 38bc961ef..ac4b9b725 100644 --- a/modules/video_models/video_save.py +++ b/modules/video_models/video_save.py @@ -1,9 +1,11 @@ +from fractions import Fraction import os import time import cv2 import numpy as np import torch import einops +from PIL import Image from modules import shared, errors ,timer, rife, processing from modules.video_models.video_utils import check_av @@ -62,7 +64,81 @@ def images_to_tensor(images): return tensor -def atomic_save_video(filename, tensor:torch.Tensor, fps:float=24, codec:str='libx264', pix_fmt:str='yuv420p', options:str='', metadata:dict={}, pbar=None): +def numpy_to_tensor(images): + if images is None or len(images) == 0: + return None + images = (2.0 * images) - 1.0 # from [0,1] to [-1,1] + array = [torch.from_numpy(images[i]) for i in range(images.shape[0])] + tensor = torch.stack(array, dim=0) # n h w c + tensor = tensor.unsqueeze(0) # 1, n, h, w, c + tensor = tensor.permute(0, 4, 1, 2, 3).contiguous() # 1, c, n, h, w + # tensor = (tensor.float() / 127.5) - 1.0 # from [0,255] to [-1,1] + # shared.log.debug(f'Video output: images={len(images)} tensor={tensor.shape}') + return tensor + + +def write_audio( + container, + samples: torch.Tensor, + audio_sample_rate: int, +) -> None: + av = check_av() + # create stream + audio_options = { 'time_base': f'1/{audio_sample_rate}' } + audio_stream = container.add_stream("aac", rate=audio_sample_rate, options=audio_options) + audio_stream.codec_context.sample_rate = audio_sample_rate + audio_stream.codec_context.layout = "stereo" + audio_stream.codec_context.format = "fltp" + audio_stream.codec_context.time_base = Fraction(1, audio_sample_rate) + # audio_stream.time_base = audio_stream.codec_context.time_base # TODO audio set time-base + shared.log.debug(f'Audio: codec={audio_stream.codec_context.name} rate={audio_stream.codec_context.sample_rate} layout={audio_stream.codec_context.layout} format={audio_stream.codec_context.format} base={audio_stream.codec_context.time_base}') + # init input samples + if samples.ndim == 1: + samples = samples[:, None] + if samples.shape[1] != 2 and samples.shape[0] == 2: + samples = samples.T + if samples.shape[1] != 2: + raise ValueError(f"Expected samples with 2 channels; got shape {samples.shape}.") + if samples.dtype != torch.int16: + samples = torch.clip(samples, -1.0, 1.0) + samples = (samples * 32767.0).to(torch.int16) + audio_frames = av.AudioFrame.from_ndarray( + samples.contiguous().reshape(1, -1).cpu().numpy(), + format="s16", + layout="stereo", + ) + audio_frames.sample_rate = audio_sample_rate + # init resampler + audio_resampler = av.audio.resampler.AudioResampler( + format=audio_stream.codec_context.format, + layout=audio_stream.codec_context.layout, + rate=audio_stream.codec_context.sample_rate, + ) + # resample + pts = 0 + for resampled in audio_resampler.resample(audio_frames): + resampled.pts = resampled.pts or 0 + resampled.sample_rate = audio_frames.sample_rate + packets = audio_stream.encode(resampled) + for packet in packets: + container.mux(packet) + pts += resampled.samples + # flush audio encoder + for packet in audio_stream.encode(): + container.mux(packet) + + +def atomic_save_video(filename: str, + tensor:torch.Tensor, + audio:torch.Tensor=None, + fps:float=24, + codec:str='libx264', + pix_fmt:str='yuv420p', + options:str='', + aac:int=24000, + metadata:dict={}, + pbar=None, + ): av = check_av() if av is None or av is False: shared.log.error('Video: ffmpeg/av not available') @@ -81,11 +157,13 @@ def atomic_save_video(filename, tensor:torch.Tensor, fps:float=24, codec:str='li else: continue options[key.strip()] = value.strip() - shared.log.info(f'Video: file="{filename}" codec={codec} frames={frames} width={width} height={height} fps={rate} options={options}') + shared.log.info(f'Video: file="{filename}" codec={codec} frames={frames} width={width} height={height} fps={rate} audio={audio is not None} aac={aac} options={options}') video_array = torch.as_tensor(tensor, dtype=torch.uint8).numpy(force=True) + task = pbar.add_task('encoding', total=frames) if pbar is not None else None if task is not None: pbar.update(task, description='video encoding') + with av.open(filename, mode="w") as container: for k, v in metadata.items(): container.metadata[k] = v @@ -101,6 +179,13 @@ def atomic_save_video(filename, tensor:torch.Tensor, fps:float=24, codec:str='li pbar.update(task, advance=1) for packet in stream.encode(): # flush container.mux(packet) + if audio is not None: + try: + write_audio(container, audio, aac) + except Exception as e: + shared.log.error(f'Video audio encoding: {e}') + errors.display(e, 'Audio') + shared.state.outputs(filename) shared.state.end(savejob) @@ -108,6 +193,7 @@ def atomic_save_video(filename, tensor:torch.Tensor, fps:float=24, codec:str='li def save_video( p:processing.StableDiffusionProcessingVideo, pixels:torch.Tensor=None, + audio:torch.Tensor=None, binary:bytes=None, mp4_fps:int=24, mp4_codec:str='libx264', @@ -117,6 +203,7 @@ def save_video( mp4_video:bool=True, # save video mp4_frames:bool=False, # save frames mp4_interpolate:int=0, # rife interpolation + aac_sample_rate:int=24000, # audio sample rate stream=None, # async progress reporting stream metadata:dict={}, # metadata for video pbar=None, # progress bar for video @@ -141,6 +228,10 @@ def save_video( if pixels is None: return 0, output_video + if isinstance(pixels, np.ndarray): + pixels = numpy_to_tensor(pixels) + if isinstance(pixels, list) and isinstance(pixels[0], Image.Image): + pixels = images_to_tensor(pixels) if not torch.is_tensor(pixels): shared.log.error(f'Video: type={type(pixels)} not a tensor') return 0, output_video @@ -148,7 +239,7 @@ def save_video( n, _c, t, h, w = pixels.shape size = pixels.element_size() * pixels.numel() shared.log.debug(f'Video: video={mp4_video} export={mp4_frames} safetensors={mp4_sf} interpolate={mp4_interpolate}') - shared.log.debug(f'Video: encode={t} raw={size} latent={pixels.shape} fps={mp4_fps} codec={mp4_codec} ext={mp4_ext} options="{mp4_opt}"') + shared.log.debug(f'Video: encode={t} raw={size} latent={pixels.shape} audio={audio.shape if audio is not None else None} fps={mp4_fps} codec={mp4_codec} ext={mp4_ext} options="{mp4_opt}"') try: preparejob = shared.state.begin('Prepare video') if stream is not None: @@ -189,7 +280,7 @@ def save_video( if mp4_video and (mp4_codec != 'none'): output_video = f'{output_filename}.{mp4_ext}' - atomic_save_video(output_video, tensor=x, fps=mp4_fps, codec=mp4_codec, options=mp4_opt, metadata=metadata, pbar=pbar) + atomic_save_video(output_video, tensor=x, audio=audio, fps=mp4_fps, codec=mp4_codec, options=mp4_opt, aac=aac_sample_rate, metadata=metadata, pbar=pbar) if stream is not None: stream.output_queue.push(('progress', (None, f'Video {os.path.basename(output_video)} | Codec {mp4_codec} | Size {w}x{h}x{t} | FPS {mp4_fps}'))) stream.output_queue.push(('file', output_video)) diff --git a/modules/video_models/video_vae.py b/modules/video_models/video_vae.py index fefb116d9..e31108088 100644 --- a/modules/video_models/video_vae.py +++ b/modules/video_models/video_vae.py @@ -6,25 +6,27 @@ debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None e vae_type = None -def set_vae_params(p): +def set_vae_params(p, slicing:bool=True, tiling:bool=True, framewise:bool=True) -> None: global vae_type # pylint: disable=global-statement vae_type = p.vae_type if not hasattr(shared.sd_model, 'vae'): return - if hasattr(shared.sd_model.vae, 'enable_slicing'): + if slicing and hasattr(shared.sd_model.vae, 'enable_slicing'): shared.sd_model.vae.enable_slicing() - if p.frames > p.vae_tile_frames: + if (p.frames > p.vae_tile_frames) and (p.vae_tile_frames > 0): if hasattr(shared.sd_model.vae, 'tile_sample_min_num_frames'): shared.sd_model.vae.tile_sample_min_num_frames = p.vae_tile_frames - if hasattr(shared.sd_model.vae, 'use_framewise_decoding'): + if framewise and hasattr(shared.sd_model.vae, 'use_framewise_decoding'): shared.sd_model.vae.use_framewise_decoding = True - if hasattr(shared.sd_model.vae, 'enable_tiling'): + if tiling and hasattr(shared.sd_model.vae, 'enable_tiling'): shared.sd_model.vae.enable_tiling() + debug(f'VAE params: type={vae_type} tiling=True frames={p.frames} tile_frames={p.vae_tile_frames} framewise={getattr(shared.sd_model.vae, "use_framewise_decoding", None)}') else: if hasattr(shared.sd_model.vae, 'use_framewise_decoding'): shared.sd_model.vae.use_framewise_decoding = False if hasattr(shared.sd_model.vae, 'disable_tiling'): shared.sd_model.vae.disable_tiling() + debug(f'VAE params: type={vae_type} tiling=False frames={p.frames} tile_frames={p.vae_tile_frames} framewise={getattr(shared.sd_model.vae, "use_framewise_decoding", None)}') def vae_decode_tiny(latents): diff --git a/modules/windows_hip_ffi.py b/modules/windows_hip_ffi.py deleted file mode 100644 index 13cbf3f84..000000000 --- a/modules/windows_hip_ffi.py +++ /dev/null @@ -1,87 +0,0 @@ -import sys - -if sys.platform == "win32": - import os - import ctypes - import ctypes.wintypes - - class hipDeviceProp(ctypes.Structure): - _fields_ = [ - ('bytes', ctypes.c_byte * 1472) # 1472 in amdhip64_6.dll, shorter in amdhip64_7.dll? - ] - - class HIP: - def __init__(self): - ctypes.windll.kernel32.LoadLibraryA.restype = ctypes.wintypes.HMODULE - ctypes.windll.kernel32.LoadLibraryA.argtypes = [ctypes.c_char_p] - self.handle = None - path = os.environ.get("windir", "C:\\Windows") + "\\System32\\amdhip64_7.dll" - if not os.path.isfile(path): - path = os.environ.get("windir", "C:\\Windows") + "\\System32\\amdhip64_6.dll" - if not os.path.isfile(path): - path = os.environ.get("windir", "C:\\Windows") + "\\System32\\amdhip64.dll" - assert os.path.isfile(path) - self.handle = ctypes.windll.kernel32.LoadLibraryA(path.encode('utf-8')) - ctypes.windll.kernel32.GetLastError.restype = ctypes.wintypes.DWORD - ctypes.windll.kernel32.GetLastError.argtypes = [] - assert ctypes.windll.kernel32.GetLastError() == 0 - ctypes.windll.kernel32.GetProcAddress.restype = ctypes.c_void_p - ctypes.windll.kernel32.GetProcAddress.argtypes = [ctypes.wintypes.HMODULE, ctypes.c_char_p] - self.hipGetDeviceCount = ctypes.CFUNCTYPE(ctypes.c_int, ctypes.POINTER(ctypes.c_int))( - ctypes.windll.kernel32.GetProcAddress(self.handle, b"hipGetDeviceCount")) - self.hipGetDeviceProperties = ctypes.CFUNCTYPE(ctypes.c_int, ctypes.POINTER(hipDeviceProp), ctypes.c_int)( - ctypes.windll.kernel32.GetProcAddress(self.handle, b"hipGetDeviceProperties")) - - def __del__(self): - if self.handle is None: - return - # Hopefully this will prevent conflicts with amdhip64_7.dll from ROCm Python packages or HIP SDK - ctypes.windll.kernel32.FreeLibrary.argtypes = [ctypes.wintypes.HMODULE] - ctypes.windll.kernel32.FreeLibrary(self.handle) - - def get_device_count(self): - count = ctypes.c_int() - assert self.hipGetDeviceCount(ctypes.byref(count)) == 0 - return count.value - - def get_device_properties(self, device_id): - prop = hipDeviceProp() - assert self.hipGetDeviceProperties(ctypes.byref(prop), device_id) == 0 - return prop.bytes - - def get_archs(): - hip = HIP() - - count = hip.get_device_count() - archs = [None] * count - for i in range(count): - prop = hip.get_device_properties(i)[:] - - name = "" - idx = 0 - while idx < len(prop): - try: - idx = prop.index(0x67, idx) + 1 # 'g' - except ValueError: - break - if prop[idx] != 0x66: # 'f' - continue - if prop[idx + 1] != 0x78: # 'x' - continue - - idx = idx + 2 - while prop[idx] != 0x00: - c = prop[idx] - idx += 1 - if (c < 0x30 or c > 0x39) and (c < 0x61 or c > 0x66): # hexadecimal - name = "" - continue - name += chr(c) - break - - # if name == "", hipDeviceProp does not contain arch name - if name: - archs[i] = "gfx" + name - - del hip - return archs diff --git a/modules/zluda.py b/modules/zluda.py index 4bb263f22..7b85eec62 100644 --- a/modules/zluda.py +++ b/modules/zluda.py @@ -1,17 +1,14 @@ import sys from typing import Union -import torch -from torch._prims_common import DeviceLikeType -from modules import shared, devices, zluda_installer from modules.zluda_installer import core, default_agent # pylint: disable=unused-import -from modules.onnx_impl.execution_providers import available_execution_providers, ExecutionProvider PLATFORM = sys.platform do_nothing = lambda _: None # pylint: disable=unnecessary-lambda-assignment -def test(device: DeviceLikeType) -> Union[Exception, None]: +def test(device) -> Union[Exception, None]: + import torch device = torch.device(device) try: ten1 = torch.randn((2, 4,), device=device) @@ -23,40 +20,42 @@ def test(device: DeviceLikeType) -> Union[Exception, None]: return e -def initialize_zluda(): - shared.cmd_opts.device_id = None - if not devices.cuda_ok or not devices.has_zluda(): - return - - torch.backends.cudnn.enabled = zluda_installer.MIOpen_enabled if shared.opts.cudnn_enabled == 'default' else shared.opts.cudnn_enabled == 'true' - if hasattr(torch.backends.cuda, "enable_cudnn_sdp"): - if not zluda_installer.MIOpen_enabled: - torch.backends.cuda.enable_cudnn_sdp(False) - torch.backends.cuda.enable_cudnn_sdp = do_nothing - else: - torch.backends.cuda.enable_cudnn_sdp = do_nothing - torch.backends.cuda.enable_flash_sdp(False) - torch.backends.cuda.enable_flash_sdp = torch.backends.cuda.enable_cudnn_sdp - torch.backends.cuda.enable_mem_efficient_sdp(False) - torch.backends.cuda.enable_mem_efficient_sdp = do_nothing - - # ONNX Runtime is not supported +def zluda_init(): try: - import onnxruntime as ort - ort.capi._pybind_state.get_available_providers = lambda: [v for v in available_execution_providers if v != ExecutionProvider.CUDA] # pylint: disable=protected-access - ort.get_available_providers = ort.capi._pybind_state.get_available_providers # pylint: disable=protected-access - if shared.opts.onnx_execution_provider == ExecutionProvider.CUDA: - shared.opts.onnx_execution_provider = ExecutionProvider.CPU - except Exception as e: - shared.log.warning(f'ZLUDA ONNX runtime: {e}') - shared.opts.onnx_execution_provider = ExecutionProvider.CPU + import torch + from installer import log + from modules import devices, zluda_installer + from modules.shared import cmd_opts + from modules.rocm_triton_windows import apply_triton_patches - device = devices.get_optimal_device() - result = test(device) - if result is not None: - shared.log.warning(f'ZLUDA device failed to pass basic operation test: index={device.index}, device_name={torch.cuda.get_device_name(device)}') - shared.log.error(result) - torch.cuda.is_available = lambda: False - devices.cuda_ok = False - devices.backend = 'cpu' - devices.device = devices.cpu + cmd_opts.device_id = None + + device = devices.get_optimal_device() + result = test(device) + if result is not None: + log.warning(f'ZLUDA device failed to pass basic operation test: index={device.index}, device_name={torch.cuda.get_device_name(device)}') + torch.cuda.is_available = lambda: False + devices.cuda_ok = False + devices.backend = 'cpu' + devices.device = devices.cpu + return False, result + + if not zluda_installer.default_agent.blaslt_supported: + log.debug(f'ROCm: hipBLASLt unavailable agent={zluda_installer.default_agent}') + + apply_triton_patches() + + torch.backends.cudnn.enabled = zluda_installer.MIOpen_enabled + if hasattr(torch.backends.cuda, "enable_cudnn_sdp"): + if not zluda_installer.MIOpen_enabled: + torch.backends.cuda.enable_cudnn_sdp(False) + torch.backends.cuda.enable_cudnn_sdp = do_nothing + else: + torch.backends.cuda.enable_cudnn_sdp = do_nothing + torch.backends.cuda.enable_flash_sdp(False) + torch.backends.cuda.enable_flash_sdp = torch.backends.cuda.enable_cudnn_sdp + torch.backends.cuda.enable_mem_efficient_sdp(False) + torch.backends.cuda.enable_mem_efficient_sdp = do_nothing + except Exception as e: + return False, e + return True, None diff --git a/modules/zluda_installer.py b/modules/zluda_installer.py index c6fb55c23..b5055b049 100644 --- a/modules/zluda_installer.py +++ b/modules/zluda_installer.py @@ -1,5 +1,6 @@ import os import sys +import ssl import site import ctypes import shutil @@ -84,14 +85,23 @@ def install(): args.use_nightly = True if args.use_nightly: platform = "nightly-" + platform - urllib.request.urlretrieve(f'https://github.com/lshqqytiger/ZLUDA/releases/download/rel.{commit}/ZLUDA-{platform}-rocm{rocm.version[0]}-amd64.zip', '_zluda') - with zipfile.ZipFile('_zluda', 'r') as archive: - infos = archive.infolist() - for info in infos: - if not info.is_dir(): - info.filename = os.path.basename(info.filename) - archive.extract(info, path) - os.remove('_zluda') + log.debug(f'Install ZLUDA: rocm={rocm.version} platform={platform} commit={commit}') + ssl._create_default_https_context = ssl._create_unverified_context # pylint: disable=protected-access + try: + urllib.request.urlretrieve(f'https://github.com/lshqqytiger/ZLUDA/releases/download/rel.{commit}/ZLUDA-{platform}-rocm{rocm.version[0]}-amd64.zip', '_zluda') + if not os.path.exists('_zluda'): + raise RuntimeError('ZLUDA download failed') + with zipfile.ZipFile('_zluda', 'r') as archive: + infos = archive.infolist() + for info in infos: + if not info.is_dir(): + info.filename = os.path.basename(info.filename) + archive.extract(info, path) + except Exception as e: + raise RuntimeError(f'Install ZLUDA: {e}') from e + finally: + if os.path.exists('_zluda'): + os.remove('_zluda') def uninstall(): @@ -124,14 +134,9 @@ def load(): core = Core(ctypes.windll.LoadLibrary(os.path.join(path, 'nvcuda.dll'))) ml = ZLUDALibrary(ctypes.windll.LoadLibrary(os.path.join(path, 'nvml.dll'))) is_nightly = core.get_nightly_flag() == 1 - hipBLASLt_enabled = is_nightly and os.path.exists(rocm.blaslt_tensile_libpath) and os.path.exists(os.path.join(rocm.environment.path, "bin", "hipblaslt.dll")) and default_agent is not None + hipBLASLt_enabled = is_nightly and os.path.exists(rocm.blaslt_tensile_libpath) and os.path.exists(os.path.join(rocm.environment.path, "bin", "hipblaslt.dll")) and default_agent is not None and default_agent.blaslt_supported MIOpen_enabled = is_nightly and os.path.exists(os.path.join(rocm.environment.path, "bin", "MIOpen.dll")) - if hipBLASLt_enabled: - if not default_agent.blaslt_supported: - hipBLASLt_enabled = False - log.debug(f'ROCm hipBLASLt: arch={default_agent.name} available={hipBLASLt_enabled}') - for k, v in DLL_MAPPING.items(): if not os.path.exists(os.path.join(path, v)): link_or_copy(os.path.join(path, k), os.path.join(path, v)) @@ -166,6 +171,7 @@ def load(): def postinstall(): import torch torch.version.hip = rocm.version + platform = sys.platform sys.platform = "" from torch.utils import cpp_extension @@ -177,3 +183,6 @@ def load(): return os.path.join(cpp_extension.ROCM_HOME, *paths) cpp_extension._join_rocm_home = _join_rocm_home # pylint: disable=protected-access rocm.postinstall = postinstall + + from modules.zluda import zluda_init + rocm.rocm_init = zluda_init diff --git a/package.json b/package.json index cfbf92e6a..4a7d2d0d8 100644 --- a/package.json +++ b/package.json @@ -20,33 +20,43 @@ "start": ". venv/bin/activate; python launch.py --debug", "localize": "node cli/localize.js", "packages": ". venv/bin/activate && pip install --upgrade transformers accelerate huggingface_hub safetensors tokenizers peft pytorch_lightning pylint ruff", - "eslint": "eslint . javascript/", - "ruff": ". venv/bin/activate && ruff check", - "pylint": ". venv/bin/activate && pylint *.py modules/ pipelines/ scripts/ extensions-builtin/ | grep -v '^*'", "format": ". venv/bin/activate && pre-commit run --all-files", - "lint": "npm run eslint && npm run format && npm run ruff && npm run pylint | grep -v TODO", - "todo": "npm run pylint | grep W0511 | awk -F'TODO ' '{print \"- \"$NF}' | sed 's/ (fixme)//g' | sort", - "eslint-win": "eslint . javascript/ --rule \"linebreak-style: off\"", - "ruff-win": "venv\\scripts\\activate && ruff check", - "pylint-win": "venv\\scripts\\activate && pylint *.py modules/ pipelines/ scripts/ extensions-builtin/", "format-win": "venv\\scripts\\activate && pre-commit run --all-files", - "lint-win": "npm run eslint-win && npm run format-win && npm run ruff-win && npm run pylint-win", - "test": ". venv/bin/activate; python launch.py --debug --test" + "eslint": "eslint . javascript/", + "eslint-win": "eslint . javascript/ --rule \"@stylistic/linebreak-style: off\"", + "eslint-ui": "cd extensions-builtin/sdnext-modernui && eslint . javascript/", + "eslint-ui-win": "cd extensions-builtin/sdnext-modernui && eslint . javascript/ --rule \"@stylistic/linebreak-style: off\"", + "ruff": ". venv/bin/activate && ruff check", + "ruff-win": "venv\\scripts\\activate && ruff check", + "pylint": ". venv/bin/activate && pylint --disable=W0511 *.py modules/ pipelines/ scripts/ extensions-builtin/ | grep -v '^*'", + "pylint-win": "venv\\scripts\\activate && pylint --disable=W0511 *.py modules/ pipelines/ scripts/ extensions-builtin/", + "lint": "npm run format && npm run eslint && npm run eslint-ui && npm run ruff && npm run pylint | grep -v TODO", + "lint-win": "npm run format-win && npm run eslint-win && npm run eslint-ui-win && npm run ruff-win && npm run pylint-win", + "test": ". venv/bin/activate; python launch.py --debug --test", + "todo": "grep -oIPR 'TODO.*' *.py modules/ pipelines/ | sort -u", + "debug": "grep -ohIPR 'SD_.*?_DEBUG' *.py modules/ pipelines/ | sort -u" }, "devDependencies": { - "esbuild": "^0.18.20" + "@eslint/compat": "^2.0.0", + "@eslint/css": "^0.14.1", + "@eslint/js": "^9.39.2", + "@eslint/json": "^0.14.0", + "@eslint/markdown": "^7.5.1", + "@html-eslint/eslint-plugin": "^0.52.1", + "esbuild": "^0.27.2", + "eslint": "^9.39.2", + "eslint-config-airbnb-extended": "^3.0.0", + "eslint-plugin-promise": "^7.2.1", + "globals": "^17.0.0" }, "dependencies": { - "@google/generative-ai": "^0.21.0", - "@typescript-eslint/eslint-plugin": "^8.47.0", - "argparse": "^2.0.1", - "eslint": "^8.57.1", - "eslint-config-airbnb-base": "^15.0.0", - "eslint-plugin-css": "^0.9.2", - "eslint-plugin-html": "^8.1.3", - "eslint-plugin-import": "^2.32.0", - "eslint-plugin-json": "^3.1.0", - "eslint-plugin-markdown": "^4.0.1", - "eslint-plugin-node": "^11.1.0" + "@google/generative-ai": "^0.24.1", + "argparse": "^2.0.1" + }, + "//": { + "disabled": { + "typescript": "^5.9.3", + "@types/node": "^25.0.3" + } } } diff --git a/pipelines/bria/bria_pipeline.py b/pipelines/bria/bria_pipeline.py index c07f6e4de..beec4e2dc 100644 --- a/pipelines/bria/bria_pipeline.py +++ b/pipelines/bria/bria_pipeline.py @@ -94,7 +94,6 @@ class BriaPipeline(FluxPipeline): scheduler=scheduler, ) - # TODO - why different than offical flux (-1) self.vae_scale_factor = ( 2 ** (len(self.vae.config.block_out_channels)) if hasattr(self, "vae") and self.vae is not None else 16 ) diff --git a/pipelines/generic.py b/pipelines/generic.py index d6dad46c8..6ad9a3fd5 100644 --- a/pipelines/generic.py +++ b/pipelines/generic.py @@ -92,7 +92,7 @@ def load_transformer(repo_id, cls_name, load_config=None, subfolder="transformer transformer.quantization_config = quant_args.get('quantization_config', None) except Exception as e: shared.log.error(f'Load model: transformer="{repo_id}" cls={cls_name.__name__} {e}') - errors.display(e, 'Load:') + errors.display(e, 'Load') raise devices.torch_gc() shared.state.end(jobid) @@ -200,6 +200,24 @@ def load_text_encoder(repo_id, cls_name, load_config=None, subfolder="text_encod **load_args, **quant_args, ) + # Qwen3ForCausalLM - shared text encoders by hidden_size: + # - Z-Image, Klein-4B: Qwen3-4B (hidden_size=2560) + # - Klein-9B: Qwen3-8B (hidden_size=4096) + # SDNQ repos for Klein and Z-Image contain text encoders pre-quantized with different quantization methods, skip shared loading + elif cls_name == transformers.Qwen3ForCausalLM and allow_shared and shared.opts.te_shared_t5 and 'sdnq' not in repo_id.lower(): + if '-9b' in repo_id.lower(): + shared_repo = 'black-forest-labs/FLUX.2-klein-9B' # 9B variants use Qwen3-8B + else: + shared_repo = 'Tongyi-MAI/Z-Image-Turbo' # 4B variants and Z-Image use Qwen3-4B + subfolder = 'text_encoder' + shared.log.debug(f'Load model: text_encoder="{shared_repo}" cls={cls_name.__name__} quant="{quant_type}" loader={_loader("transformers")} shared={shared.opts.te_shared_t5}') + text_encoder = cls_name.from_pretrained( + shared_repo, + cache_dir=shared.opts.hfcache_dir, + subfolder=subfolder, + **load_args, + **quant_args, + ) # load from repo if text_encoder is None: @@ -226,7 +244,7 @@ def load_text_encoder(repo_id, cls_name, load_config=None, subfolder="text_encod text_encoder.quantization_config = quant_args.get('quantization_config', None) except Exception as e: shared.log.error(f'Load model: text_encoder="{repo_id}" cls={cls_name.__name__} {e}') - errors.display(e, 'Load:') + errors.display(e, 'Load') raise devices.torch_gc() shared.state.end(jobid) diff --git a/pipelines/hdm/hdm/data/base.py b/pipelines/hdm/hdm/data/base.py index 91d025117..f72315495 100644 --- a/pipelines/hdm/hdm/data/base.py +++ b/pipelines/hdm/hdm/data/base.py @@ -12,7 +12,6 @@ class BaseDataset(Data.Dataset): samples = torch.stack([x["sample"] for x in batch]) caption = [x["caption"] for x in batch] tokenizer_outs = [x["tokenizer_out"] for x in batch] - # TODO: change to stack and reduce dim? add_time_ids = [x["add_time_ids"] for x in batch] tokenizer_outputs = [] for tokenizer_out in zip(*tokenizer_outs): diff --git a/pipelines/hdm/hdm/loader.py b/pipelines/hdm/hdm/loader.py index d285556a3..ef60708d3 100644 --- a/pipelines/hdm/hdm/loader.py +++ b/pipelines/hdm/hdm/loader.py @@ -119,6 +119,5 @@ def load_all(conf: dict): trainer = load_trainer( conf.pop("trainer"), unet=unet, te=te, vae=vae, scheduler=scheduler ) - # TODO: there might be a better way to handle this dataset.tokenizers = tokenizers return dataset, trainer, (unet, te, tokenizers, vae, scheduler) diff --git a/pipelines/hdm/hdm/modules/unet_patch.py b/pipelines/hdm/hdm/modules/unet_patch.py index a42edc0e4..9fdc0dd25 100644 --- a/pipelines/hdm/hdm/modules/unet_patch.py +++ b/pipelines/hdm/hdm/modules/unet_patch.py @@ -141,7 +141,6 @@ class RoPEAttnProcessor2_0(AttnProcessor2_0): key = key.transpose(1, 2) # the output of sdp = (batch, num_heads, seq_len, head_dim) - # TODO: add support for attn.scale when we move to Torch 2.1 hidden_states = F.scaled_dot_product_attention( query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False ) diff --git a/pipelines/model_chroma.py b/pipelines/model_chroma.py index a2a21c34a..bcb2cdcd5 100644 --- a/pipelines/model_chroma.py +++ b/pipelines/model_chroma.py @@ -26,6 +26,7 @@ def load_chroma(checkpoint_info, diffusers_load_config=None): diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["chroma"] = diffusers.ChromaPipeline diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["chroma"] = diffusers.ChromaImg2ImgPipeline + diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["chroma"] = diffusers.ChromaInpaintPipeline del text_encoder del transformer sd_hijack_te.init_hijack(pipe) diff --git a/pipelines/model_flux2_klein.py b/pipelines/model_flux2_klein.py new file mode 100644 index 000000000..d810821d9 --- /dev/null +++ b/pipelines/model_flux2_klein.py @@ -0,0 +1,42 @@ +import transformers +import diffusers +from modules import shared, devices, sd_models, model_quant, sd_hijack_te, sd_hijack_vae +from pipelines import generic + + +def load_flux2_klein(checkpoint_info, diffusers_load_config=None): + if diffusers_load_config is None: + diffusers_load_config = {} + repo_id = sd_models.path_to_repo(checkpoint_info) + sd_models.hf_auth_check(checkpoint_info) + + load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False) + shared.log.debug(f'Load model: type=Flux2Klein repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}') + + # Load transformer - Klein uses Flux2Transformer2DModel (same class as Flux2, different size) + transformer = generic.load_transformer(repo_id, cls_name=diffusers.Flux2Transformer2DModel, load_config=diffusers_load_config) + + # Load text encoder - Klein uses Qwen3 (4B for Klein-4B, 8B for Klein-9B) + text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.Qwen3ForCausalLM, load_config=diffusers_load_config) + + pipe = diffusers.Flux2KleinPipeline.from_pretrained( + repo_id, + transformer=transformer, + text_encoder=text_encoder, + cache_dir=shared.opts.diffusers_dir, + **load_args, + ) + pipe.task_args = { + 'output_type': 'np', + } + diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["flux2klein"] = diffusers.Flux2KleinPipeline + diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["flux2klein"] = diffusers.Flux2KleinPipeline + diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["flux2klein"] = diffusers.Flux2KleinPipeline + + del text_encoder + del transformer + sd_hijack_te.init_hijack(pipe) + sd_hijack_vae.init_hijack(pipe) + + devices.torch_gc(force=True, reason='load') + return pipe diff --git a/pipelines/model_glm.py b/pipelines/model_glm.py new file mode 100644 index 000000000..d849ebb19 --- /dev/null +++ b/pipelines/model_glm.py @@ -0,0 +1,141 @@ +import time +import rich.progress as rp +import transformers +import diffusers +from modules import shared, devices, sd_models, model_quant, sd_hijack_te +from pipelines import generic + + +class GLMTokenProgressProcessor(transformers.LogitsProcessor): + """LogitsProcessor that tracks autoregressive token generation progress for GLM-Image.""" + + def __init__(self): + self.total_tokens = 0 + self.current_step = 0 + self.task_id = None + self.pbar = None + self.pbar_task = None + self.start_time = 0 + + def set_total(self, total_tokens: int): + self.total_tokens = total_tokens + self.current_step = 0 + + def __call__(self, input_ids, scores): + if self.current_step == 0: + self.task_id = shared.state.begin('AR Generation') + self.start_time = time.time() + self.pbar = rp.Progress( + rp.TextColumn('[cyan]AR Generation'), + rp.TextColumn('{task.fields[speed]}'), + rp.BarColumn(bar_width=40, complete_style='#327fba', finished_style='#327fba'), + rp.TaskProgressColumn(), + rp.MofNCompleteColumn(), + rp.TimeElapsedColumn(), + rp.TimeRemainingColumn(), + console=shared.console, + ) + self.pbar.start() + self.pbar_task = self.pbar.add_task(description='', total=self.total_tokens, speed='') + self.current_step += 1 + shared.state.sampling_step = self.current_step + shared.state.sampling_steps = self.total_tokens + if self.pbar is not None and self.pbar_task is not None: + elapsed = time.time() - self.start_time + speed = f'{self.current_step / elapsed:.2f}tok/s' if elapsed > 0 else '' + self.pbar.update(self.pbar_task, completed=self.current_step, speed=speed) + if self.current_step >= self.total_tokens: + if self.pbar is not None: + self.pbar.stop() + self.pbar = None + if self.task_id is not None: + shared.state.end(self.task_id) + self.task_id = None + return scores + + +def hijack_vision_language_generate(pipe): + """Wrap vision_language_encoder.generate to add progress tracking.""" + if not hasattr(pipe, 'vision_language_encoder') or pipe.vision_language_encoder is None: + return + + original_generate = pipe.vision_language_encoder.generate + progress_processor = GLMTokenProgressProcessor() + + def wrapped_generate(*args, **kwargs): + # Get max_new_tokens to determine total tokens + max_new_tokens = kwargs.get('max_new_tokens', 0) + progress_processor.set_total(max_new_tokens) + + # Add progress processor to logits_processor list + existing_processors = kwargs.get('logits_processor', None) + if existing_processors is None: + existing_processors = [] + elif not isinstance(existing_processors, list): + existing_processors = list(existing_processors) + kwargs['logits_processor'] = existing_processors + [progress_processor] + + return original_generate(*args, **kwargs) + + pipe.vision_language_encoder.generate = wrapped_generate + + +def load_glm_image(checkpoint_info, diffusers_load_config=None): + if diffusers_load_config is None: + diffusers_load_config = {} + repo_id = sd_models.path_to_repo(checkpoint_info) + sd_models.hf_auth_check(checkpoint_info) + + if not hasattr(transformers, 'GlmImageForConditionalGeneration'): + shared.log.error(f'Load model: type=GLM-Image repo="{repo_id}" transformers={transformers.__version__} not supported') + return None + + load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False) + shared.log.debug(f'Load model: type=GLM-Image repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}') + + # Load transformer (DiT decoder - 7B) with quantization support + transformer = generic.load_transformer( + repo_id, + cls_name=diffusers.GlmImageTransformer2DModel, + load_config=diffusers_load_config + ) + + # Load text encoder (ByT5 for glyph) - cannot use shared T5 as GLM-Image requires specific ByT5 encoder (1472 hidden size) + text_encoder = generic.load_text_encoder( + repo_id, + cls_name=transformers.T5EncoderModel, + load_config=diffusers_load_config, + allow_shared=False + ) + + # Load vision-language encoder (AR model - 9B) + # Note: This is a conditional generation model, different from typical text encoders + vision_language_encoder = generic.load_text_encoder( + repo_id, + cls_name=transformers.GlmImageForConditionalGeneration, # pylint: disable=no-member + subfolder="vision_language_encoder", + load_config=diffusers_load_config, + allow_shared=False + ) + + pipe = diffusers.GlmImagePipeline.from_pretrained( + repo_id, + cache_dir=shared.opts.diffusers_dir, + transformer=transformer, + text_encoder=text_encoder, + vision_language_encoder=vision_language_encoder, + **load_args, + ) + + pipe.task_args = { + 'output_type': 'np', + 'generate_kwargs': { + 'eos_token_id': None, # Disable EOS early stopping to ensure all required tokens are generated + }, + } + + del transformer, text_encoder, vision_language_encoder + sd_hijack_te.init_hijack(pipe) + hijack_vision_language_generate(pipe) # Add progress tracking for AR token generation + devices.torch_gc(force=True, reason='load') + return pipe diff --git a/pipelines/model_google.py b/pipelines/model_google.py index b44e3d38e..627003fd6 100644 --- a/pipelines/model_google.py +++ b/pipelines/model_google.py @@ -70,24 +70,42 @@ class GoogleNanoBananaPipeline(): def get_args(self): from modules.shared import opts - api_key = os.getenv("GOOGLE_API_KEY") or opts.google_api_key - vertex_credentials = os.getenv("GOOGLE_APPLICATION_CREDENTIALS") - if (api_key is None or len(api_key) == 0) and (vertex_credentials is None or len(vertex_credentials) == 0): - log.error(f'Cloud: model="{self.model}" API key not provided') - return None - use_vertexai = (os.getenv("GOOGLE_GENAI_USE_VERTEXAI") is not None) or opts.google_use_vertexai - project_id = os.getenv("GOOGLE_CLOUD_PROJECT") or opts.google_project_id - location_id = os.getenv("GOOGLE_CLOUD_LOCATION") or opts.google_location_id - args = { - 'api_key': api_key, - 'vertexai': use_vertexai, - 'project': project_id if len(project_id) > 0 else None, - 'location': location_id if len(location_id) > 0 else None, - } - args_copy = args.copy() - args_copy['api_key'] = '...' + args_copy['api_key'][-4:] # last 4 chars - args_copy['credentials'] = vertex_credentials - log.debug(f'Cloud: model="{self.model}" args={args_copy}') + # Use UI settings only - env vars are intentionally ignored + api_key = opts.google_api_key + project_id = opts.google_project_id + location_id = opts.google_location_id + use_vertexai = opts.google_use_vertexai + + has_api_key = api_key and len(api_key) > 0 + has_project = project_id and len(project_id) > 0 + has_location = location_id and len(location_id) > 0 + + if use_vertexai: + if has_api_key and (has_project or has_location): + # Invalid: can't have both api_key AND project/location + log.error(f'Cloud: model="{self.model}" API key and project/location are mutually exclusive') + return None + elif has_api_key: + # Vertex AI Express Mode: api_key + vertexai, no project/location + args = {'api_key': api_key, 'vertexai': True} + elif has_project and has_location: + # Standard Vertex AI: project/location, no api_key + args = {'vertexai': True, 'project': project_id, 'location': location_id} + else: + log.error(f'Cloud: model="{self.model}" Vertex AI requires either API key (Express Mode) or project ID + location ID') + return None + else: + # Gemini Developer API: api_key only + if not has_api_key: + log.error(f'Cloud: model="{self.model}" API key not provided') + return None + args = {'api_key': api_key} + + # Debug logging + args_log = args.copy() + if args_log.get('api_key'): + args_log['api_key'] = '...' + args_log['api_key'][-4:] + log.debug(f'Cloud: model="{self.model}" args={args_log}') return args def __call__(self, prompt: list[str], width: int, height: int, image: Image.Image = None): diff --git a/pipelines/model_longcat.py b/pipelines/model_longcat.py index de103c054..d8dd7ece4 100644 --- a/pipelines/model_longcat.py +++ b/pipelines/model_longcat.py @@ -30,6 +30,9 @@ def load_longcat(checkpoint_info, diffusers_load_config=None): text_processor=text_processor, **load_args, ) + diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["longcat"] = cls + diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["longcat"] = cls + diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["longcat"] = cls del transformer del text_encoder diff --git a/pipelines/model_qwen.py b/pipelines/model_qwen.py index 31eb76f93..3bea5c121 100644 --- a/pipelines/model_qwen.py +++ b/pipelines/model_qwen.py @@ -30,7 +30,7 @@ def load_qwen(checkpoint_info, diffusers_load_config=None): diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["qwen-layered"] = diffusers.QwenImageLayeredPipeline diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["qwen-layered"] = diffusers.QwenImageLayeredPipeline diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["qwen-layered"] = diffusers.QwenImageLayeredPipeline - else: + else: # qwen-image, qwen-image-2512 cls_name = diffusers.QwenImagePipeline diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["qwen-image"] = diffusers.QwenImagePipeline diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["qwen-image"] = diffusers.QwenImageImg2ImgPipeline diff --git a/pipelines/model_z_image.py b/pipelines/model_z_image.py index a839715e8..9f8dd51e1 100644 --- a/pipelines/model_z_image.py +++ b/pipelines/model_z_image.py @@ -4,6 +4,20 @@ from modules import shared, devices, sd_models, model_quant, sd_hijack_te from pipelines import generic +def load_nunchaku(): + import nunchaku + nunchaku_precision = nunchaku.utils.get_precision() + nunchaku_rank = 128 + nunchaku_repo = f"nunchaku-tech/nunchaku-z-image-turbo/svdq-{nunchaku_precision}_r{nunchaku_rank}-z-image-turbo.safetensors" + shared.log.debug(f'Load module: quant=Nunchaku module=transformer repo="{nunchaku_repo}" attention={shared.opts.nunchaku_attention}') + transformer = nunchaku.NunchakuZImageTransformer2DModel.from_pretrained( + nunchaku_repo, + torch_dtype=devices.dtype, + cache_dir=shared.opts.hfcache_dir, + ) + return transformer + + def load_z_image(checkpoint_info, diffusers_load_config=None): if diffusers_load_config is None: diffusers_load_config = {} @@ -13,7 +27,11 @@ def load_z_image(checkpoint_info, diffusers_load_config=None): load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, allow_quant=False) shared.log.debug(f'Load model: type=ZImage repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={diffusers_load_config}') - transformer = generic.load_transformer(repo_id, cls_name=diffusers.ZImageTransformer2DModel, load_config=diffusers_load_config) + if model_quant.check_nunchaku('Model'): # only available model + transformer = load_nunchaku() + else: + transformer = generic.load_transformer(repo_id, cls_name=diffusers.ZImageTransformer2DModel, load_config=diffusers_load_config) + text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.Qwen3ForCausalLM, load_config=diffusers_load_config) pipe = diffusers.ZImagePipeline.from_pretrained( diff --git a/requirements.txt b/requirements.txt index 0acc4e891..38018fcf3 100644 --- a/requirements.txt +++ b/requirements.txt @@ -36,7 +36,7 @@ fastapi==0.124.4 rich==14.1.0 safetensors==0.7.0 tensordict==0.8.3 -peft==0.18.0 +peft==0.18.1 httpx==0.28.1 compel==2.2.1 torchsde==0.2.6 diff --git a/scripts/consistory/consistory_pipeline.py b/scripts/consistory/consistory_pipeline.py index cf020c254..9ad065db8 100644 --- a/scripts/consistory/consistory_pipeline.py +++ b/scripts/consistory/consistory_pipeline.py @@ -322,7 +322,7 @@ class ConsistoryExtendAttnSDXLPipeline( latents, ) - # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 6. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) if share_queries: diff --git a/scripts/differential_diffusion.py b/scripts/differential_diffusion.py index db67e8e00..f351571a6 100644 --- a/scripts/differential_diffusion.py +++ b/scripts/differential_diffusion.py @@ -15,7 +15,7 @@ import PIL.Image import numpy as np import torch import torchvision -from transformers import CLIPFeatureExtractor, CLIPTextModel, CLIPTextModelWithProjection, CLIPTokenizer +from transformers import CLIPImageProcessor, CLIPTextModel, CLIPTextModelWithProjection, CLIPTokenizer from diffusers.image_processor import VaeImageProcessor from diffusers.loaders import FromSingleFileMixin, LoraLoaderMixin, TextualInversionLoaderMixin from diffusers.models import AutoencoderKL, UNet2DConditionModel @@ -1059,7 +1059,7 @@ class StableDiffusionDiffImg2ImgPipeline(DiffusionPipeline): unet: UNet2DConditionModel, scheduler: KarrasDiffusionSchedulers, safety_checker: StableDiffusionSafetyChecker, - feature_extractor: CLIPFeatureExtractor, + feature_extractor: CLIPImageProcessor, requires_safety_checker: bool = False, ): super().__init__() @@ -1353,17 +1353,6 @@ class StableDiffusionDiffImg2ImgPipeline(DiffusionPipeline): return prompt_embeds - # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.run_safety_checker - def run_safety_checker(self, image, device, dtype): - if self.safety_checker is not None: - safety_checker_input = self.feature_extractor(self.numpy_to_pil(image), return_tensors="pt").to(device) - image, has_nsfw_concept = self.safety_checker( - images=image, clip_input=safety_checker_input.pixel_values.to(dtype) - ) - else: - has_nsfw_concept = None - return image, has_nsfw_concept - # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.decode_latents def decode_latents(self, latents): latents = 1 / self.vae.config.scaling_factor * latents @@ -1767,7 +1756,7 @@ class StableDiffusionDiffImg2ImgPipeline(DiffusionPipeline): timesteps, num_inference_steps = self.get_timesteps(num_inference_steps, strength, device) - # 7. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 7. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) map = torchvision.transforms.Resize(tuple(s // self.vae_scale_factor for s in image.shape[2:]),antialias=None)(map) diff --git a/scripts/freescale/freescale_pipeline.py b/scripts/freescale/freescale_pipeline.py index 9b7a68b68..df91c014b 100644 --- a/scripts/freescale/freescale_pipeline.py +++ b/scripts/freescale/freescale_pipeline.py @@ -873,7 +873,7 @@ class StableDiffusionXLFreeScale(DiffusionPipeline, FromSingleFileMixin, LoraLoa latents, ) - # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 6. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 7. Prepare added time ids & embeddings diff --git a/scripts/freescale/freescale_pipeline_img2img.py b/scripts/freescale/freescale_pipeline_img2img.py index df4c3f0c1..7f2964cc2 100644 --- a/scripts/freescale/freescale_pipeline_img2img.py +++ b/scripts/freescale/freescale_pipeline_img2img.py @@ -902,7 +902,7 @@ class StableDiffusionXLFreeScaleImg2Img(DiffusionPipeline, FromSingleFileMixin, latents, ) - # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 6. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 7. Prepare added time ids & embeddings diff --git a/scripts/instantir/sdxl_instantir.py b/scripts/instantir/sdxl_instantir.py index 5e5df3ece..bfb4f84e9 100644 --- a/scripts/instantir/sdxl_instantir.py +++ b/scripts/instantir/sdxl_instantir.py @@ -1405,7 +1405,7 @@ class InstantIRPipeline( guidance_scale_tensor, embedding_dim=self.unet.config.time_cond_proj_dim ).to(device=device, dtype=latents.dtype) - # 7. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 7. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 7.1 Create tensor stating which controlnets to keep diff --git a/scripts/mod/__init__.py b/scripts/mod/__init__.py index ea64c075e..db9445b8c 100644 --- a/scripts/mod/__init__.py +++ b/scripts/mod/__init__.py @@ -1039,7 +1039,7 @@ class StableDiffusionXLTilingPipeline( if isinstance(self.scheduler, LMSDiscreteScheduler): latents = latents * self.scheduler.sigmas[0] - # 5. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 5. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 6. Prepare added time ids & embeddings diff --git a/scripts/pixelsmith/pixelsmith_pipeline.py b/scripts/pixelsmith/pixelsmith_pipeline.py index 702ee67f6..474fcd4e8 100644 --- a/scripts/pixelsmith/pixelsmith_pipeline.py +++ b/scripts/pixelsmith/pixelsmith_pipeline.py @@ -1384,7 +1384,7 @@ class PixelSmithXLPipeline( latents, ) - # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 6. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 7. Prepare added time ids & embeddings diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index 1925a05d8..245773e78 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -9,7 +9,8 @@ import torch import transformers import gradio as gr from PIL import Image -from modules import scripts_manager, shared, devices, errors, processing, sd_models, sd_modules, timer +from modules import scripts_manager, shared, devices, errors, processing, sd_models, sd_modules, timer, ui_symbols +from modules import ui_control_helpers debug_enabled = os.environ.get('SD_LLM_DEBUG', None) is not None @@ -28,6 +29,66 @@ def b64(image): return encoded +def is_vision_model(model_name: str) -> bool: + """Check if model supports vision/image input.""" + if not model_name: + return False + return model_name in Options.img2img + + +def is_thinking_model(model_name: str) -> bool: + """Check if model supports thinking/reasoning mode.""" + if not model_name: + return False + model_lower = model_name.lower() + # Match VQA's detection patterns for consistency + thinking_indicators = [ + 'thinking', # Qwen3-VL-*-Thinking models + 'moondream3', # Moondream 3 supports thinking + 'moondream 3', + 'moondream2', # Moondream 2 supports reasoning mode + 'moondream 2', + 'mimo', # XiaomiMiMo models + ] + return any(indicator in model_lower for indicator in thinking_indicators) + + +def get_model_display_name(model_repo: str) -> str: + """Generate display name with vision/reasoning symbols.""" + symbols = [] + if model_repo in Options.img2img: + symbols.append(ui_symbols.vision) + if is_thinking_model(model_repo): + symbols.append(ui_symbols.reasoning) + return f"{model_repo} {' '.join(symbols)}" if symbols else model_repo + + +def get_model_repo_from_display(display_name: str) -> str: + """Strip symbols from display name to get repo.""" + if not display_name: + return display_name + result = display_name + for symbol in [ui_symbols.vision, ui_symbols.reasoning]: + result = result.replace(symbol, '') + return result.strip() + + +def keep_think_block_open(text_prompt: str) -> str: + """Remove closing so model can continue reasoning with prefill.""" + think_open = "" + think_close = "" + last_open = text_prompt.rfind(think_open) + if last_open == -1: + return text_prompt + close_index = text_prompt.find(think_close, last_open) + if close_index == -1: + return text_prompt + end_close = close_index + len(think_close) + while end_close < len(text_prompt) and text_prompt[end_close] in ' \t\r\n': + end_close += 1 + return text_prompt[:close_index] + text_prompt[end_close:] + + @dataclass class Options: img2img = [ @@ -98,12 +159,24 @@ class Options: censored = ["i cannot", "i can't", "i am sorry", "against my programming", "i am not able", "i am unable", 'i am not allowed'] max_delim_index: int = 60 - max_tokens: int = 50 + max_tokens: int = 512 do_sample: bool = True - temperature: float = 0.15 + temperature: float = 0.8 repetition_penalty: float = 1.2 + top_k: int = 0 + top_p: float = 0.0 thinking_mode: bool = False + @staticmethod + def get_model_choices(): + """Return list of display names for dropdown.""" + return [get_model_display_name(repo) for repo in Options.models.keys()] + + @staticmethod + def get_default_display(): + """Return display name for default model.""" + return get_model_display_name(Options.default) + class Script(scripts_manager.Script): prompt: gr.Textbox = None @@ -127,7 +200,8 @@ class Script(scripts_manager.Script): self.llm = compile_torch(self.llm) def load(self, name:str=None, model_repo:str=None, model_gguf:str=None, model_type:str=None, model_file:str=None): - name = name or self.options.default + # Strip symbols from display name if present + name = get_model_repo_from_display(name) if name else self.options.default if self.busy: shared.log.debug('Prompt enhance: busy') return @@ -204,7 +278,7 @@ class Script(scripts_manager.Script): if debug_enabled: modules = sd_modules.get_model_stats(self.llm) + sd_modules.get_model_stats(self.tokenizer) for m in modules: - shared.log.trace(f'Prompt enhance: {m}') + debug_log(f'Prompt enhance: {m}') self.model = name t1 = time.time() shared.log.info(f'Prompt enhance: cls={self.llm.__class__.__name__} name="{name}" repo="{model_repo}" fn="{model_file}" time={t1-t0:.2f} loaded') @@ -221,40 +295,75 @@ class Script(scripts_manager.Script): def unload(self): if self.llm is not None: - sd_models.move_model(self.llm, devices.cpu) - self.model = None - self.llm = None - self.tokenizer = None - devices.torch_gc() - shared.log.debug('Prompt enhance: model unloaded') + model_name = self.model + shared.log.debug(f'Prompt enhance: unloading model="{model_name}"') + sd_models.move_model(self.llm, devices.cpu, force=True) + self.model = None + self.llm = None + self.tokenizer = None + devices.torch_gc(force=True, reason='prompt enhance unload') + shared.log.debug(f'Prompt enhance: model="{model_name}" unloaded') + else: + shared.log.debug('Prompt enhance: no model loaded') + + def clean(self, response, keep_thinking=False, prefill_text='', keep_prefill=False): + # Handle thinking tags FIRST (before generic tag removal) + if '' in response or '' in response: + if keep_thinking: + # Format: handle partial tags ( without means thinking was in prompt) + if '' in response and '' not in response: + response = 'Reasoning:\n' + response.replace('', '\n\nAnswer:\n') + else: + response = response.replace('', 'Reasoning:\n').replace('', '\n\nAnswer:\n') + else: + # Strip all thinking content + response = re.sub(r'.*?', '', response, flags=re.DOTALL) + response = response.replace('', '') # Handle orphaned closing tags - def clean(self, response): # remove special characters - response = response.replace('"', '').replace("'", "").replace('“', '').replace('”', '').replace('**', '') + response = response.replace('"', '').replace("'", "").replace('"', '').replace('"', '').replace('**', '') # remove repeating characters response = response.replace('\n\n', '\n').replace(' ', ' ').replace('...', '.') - # remove comments between brackets + # remove comments between brackets (but not Reasoning:/Answer: which we may have added) response = re.sub(r'<.*?>', '', response) - response = re.sub(r'\[.*?\]', '', response) # Fixed regex for brackets - response = re.sub(r'\/.*?\/', '', response) # Fixed regex for slashes + response = re.sub(r'\[.*?\]', '', response) + response = re.sub(r'\/.*?\/', '', response) # remove llm commentary removed = '' if response.startswith('Prompt'): removed, response = response.split('Prompt', maxsplit=1) if 0 <= response.find(':') < self.options.max_delim_index: - removed, response = response.split(':', maxsplit=1) + # Don't split on "Reasoning:" or "Answer:" if we're keeping thinking + colon_pos = response.find(':') + prefix_text = response[:colon_pos].strip() + if not keep_thinking or (prefix_text not in ['Reasoning', 'Answer']): + removed, response = response.split(':', maxsplit=1) if 0 <= response.find('---') < self.options.max_delim_index: response, removed = response.split('---', maxsplit=1) if len(removed) > 0: debug_log(f'Prompt enhance: max={self.options.max_delim_index} removed="{removed}"') # remove bullets and lists - lines = [re.sub(r'^(\s*[-*]|\s*\d+)\s+', '', line).strip() for line in response.splitlines()] # Fixed regex + lines = [re.sub(r'^(\s*[-*]|\s*\d+)\s+', '', line).strip() for line in response.splitlines()] response = '\n'.join(lines) response = response.strip() + + # Handle prefill retention/removal + prefill_text = (prefill_text or '').strip() + if prefill_text: + if keep_prefill: + # Add prefill if it's missing from the cleaned response + if not response.startswith(prefill_text): + sep = '' if (not response or response[0] in '.,!?;:') else ' ' + response = f'{prefill_text}{sep}{response}' + else: + # Remove prefill if it's present in the cleaned response + if response.startswith(prefill_text): + response = response[len(prefill_text):].strip() + return response def post(self, response, prefix, suffix, networks): @@ -275,18 +384,26 @@ class Script(scripts_manager.Script): filtered = re.sub(pattern, '', prompt) return filtered, matches - def enhance(self, model: str=None, prompt:str=None, system:str=None, prefix:str=None, suffix:str=None, sample:bool=None, tokens:int=None, temperature:float=None, penalty:float=None, thinking:bool=False, seed:int=-1, image=None, nsfw:bool=None): - model = model or self.options.default + def enhance(self, model: str=None, prompt:str=None, system:str=None, prefix:str=None, suffix:str=None, sample:bool=None, tokens:int=None, temperature:float=None, penalty:float=None, top_k:int=None, top_p:float=None, thinking:bool=False, seed:int=-1, image=None, nsfw:bool=None, use_vision:bool=True, prefill:str='', keep_prefill:bool=False, keep_thinking:bool=False): + # Strip symbols from model name if present + model = get_model_repo_from_display(model) if model else self.options.default prompt = prompt or (self.prompt.value if self.prompt else "") # Check if self.prompt is None - image = image or self.image + # Handle vision toggle - if disabled or non-VL model, don't use image + if use_vision and is_vision_model(model): + image = image or self.image + else: + image = None prefix = prefix or '' suffix = suffix or '' tokens = tokens or self.options.max_tokens penalty = penalty or self.options.repetition_penalty temperature = temperature or self.options.temperature + top_k = top_k if top_k is not None else self.options.top_k + top_p = top_p if top_p is not None else self.options.top_p thinking = thinking or self.options.thinking_mode sample = sample if sample is not None else self.options.do_sample nsfw = nsfw if nsfw is not None else True # Default nsfw to True if not provided + debug_log(f'Prompt enhance: model="{model}" model_class="{self.llm.__class__.__name__ if self.llm else "not loaded"}" nsfw={nsfw} thinking={thinking} prefill="{prefill[:30] if prefill else ""}" use_vision={use_vision} image={image is not None}') while self.busy: time.sleep(0.1) @@ -302,15 +419,29 @@ class Script(scripts_manager.Script): debug_log(f'Prompt enhance: networks={networks}') current_image = None - try: - if image is not None and isinstance(image, gr.Image): - current_image = image.value - elif image is not None and isinstance(image, Image.Image): # if image is already a PIL image - current_image = image - if current_image is not None and (current_image.width <= 64 or current_image.height <= 64): + # Only process images if vision is enabled and model supports it + if use_vision and is_vision_model(model): + try: + if image is not None and isinstance(image, gr.Image): + current_image = image.value + elif image is not None and isinstance(image, Image.Image): # if image is already a PIL image + current_image = image + if current_image is not None and (current_image.width <= 64 or current_image.height <= 64): + current_image = None + # Fallback to Kanvas/Control input if no image from Gradio component (e.g., when Kanvas is active) + if current_image is None and ui_control_helpers.input_source is not None: + if isinstance(ui_control_helpers.input_source, list) and len(ui_control_helpers.input_source) > 0: + current_image = ui_control_helpers.input_source[0] + elif isinstance(ui_control_helpers.input_source, Image.Image): + current_image = ui_control_helpers.input_source + except Exception: current_image = None - except Exception: - current_image = None + debug_log(f'Prompt enhance: current_image={current_image is not None} size={f"{current_image.width}x{current_image.height}" if current_image else "N/A"}') + + # Check if vision was requested but no image is available + if use_vision and is_vision_model(model) and current_image is None: + shared.log.error(f'Prompt enhance: model="{model}" error="No input image provided"') + return 'Error: No input image provided. Please upload or select an image.' # Resize large images to match VQA performance (Qwen3-VL performance is sensitive to resolution) # Create a copy to avoid modifying the original image used by img2img @@ -332,7 +463,6 @@ class Script(scripts_manager.Script): debug_log('Prompt enhance: Converted image to RGB mode') has_system = system is not None and len(system) > 4 - mode = 'custom' if has_system else '' if current_image is not None and isinstance(current_image, Image.Image): if (self.tokenizer is None) or (not self.tokenizer.is_processor): @@ -340,7 +470,6 @@ class Script(scripts_manager.Script): return prompt_text # Return original text part if image cannot be processed if prompt_text is not None and len(prompt_text) > 0: if not has_system: - mode = 'i2i-prompt' system = self.options.i2i_prompt system += self.options.nsfw_ok if nsfw else self.options.nsfw_no system += self.options.details_prompt @@ -355,7 +484,6 @@ class Script(scripts_manager.Script): ] else: if not has_system: - mode = 'i2i-noprompt' system = self.options.i2i_noprompt system += self.options.nsfw_ok if nsfw else self.options.nsfw_no system += self.options.details_prompt @@ -373,13 +501,11 @@ class Script(scripts_manager.Script): system += self.options.nsfw_ok if nsfw else self.options.nsfw_no system += self.options.details_prompt if not self.tokenizer.is_processor: - mode = 't2i+tokenizer' chat_template = [ { "role": "system", "content": system }, { "role": "user", "content": prompt_text }, ] else: - mode = 't2i+processor' chat_template = [ { "role": "system", "content": [ {"type": "text", "text": system } @@ -389,18 +515,63 @@ class Script(scripts_manager.Script): ] }, ] + # Prepare prefill (VQA approach: string concatenation, not assistant message) + prefill_text = (prefill or '').strip() + use_prefill = len(prefill_text) > 0 + is_thinking = is_thinking_model(model) + + debug_log(f'Prompt enhance: chat_template roles={[msg["role"] for msg in chat_template]} is_thinking={is_thinking} thinking={thinking} use_prefill={use_prefill}') t0 = time.time() self.busy = True try: - inputs = self.tokenizer.apply_chat_template( - chat_template, - add_generation_prompt=True, - enable_thinking=thinking, - tokenize=True, - return_dict=True, - return_tensors="pt", - ).to(devices.device).to(devices.dtype) + # Generate text prompt using template (WITHOUT enable_thinking parameter) + # Let template naturally generate for thinking models + try: + text_prompt = self.tokenizer.apply_chat_template( + chat_template, + add_generation_prompt=True, + tokenize=False, + ) + except TypeError: + text_prompt = self.tokenizer.apply_chat_template( + chat_template, + tokenize=False, + ) + + # Manually handle thinking tags and prefill (VQA Qwen approach) + if is_thinking: + if not thinking: + # User wants to SKIP thinking + # Template opened the block with , close it immediately + text_prompt += "\n" + if use_prefill: + text_prompt += prefill_text + debug_log('Prompt enhance: forced thinking off, appended ') + else: + # User wants thinking - prefill becomes part of thought process + if use_prefill: + text_prompt += prefill_text + debug_log('Prompt enhance: thinking enabled, prefill inside think block') + else: + # Standard model (no block) + if use_prefill: + text_prompt += prefill_text + + debug_log(f'Prompt enhance: final text_prompt (last 200 chars)="{text_prompt[-200:]}"') + + # Tokenize the final prompt + # For VL models with images, pass the image to the processor (like VQA does) + if self.tokenizer.is_processor and current_image is not None: + inputs = self.tokenizer(text=[text_prompt], images=[current_image], padding=True, return_tensors="pt") + elif self.tokenizer.is_processor: + # VL processor without image - must use explicit text= parameter + inputs = self.tokenizer(text=[text_prompt], images=None, padding=True, return_tensors="pt") + else: + inputs = self.tokenizer(text_prompt, return_tensors="pt") + inputs = inputs.to(devices.device).to(devices.dtype) + input_len = inputs['input_ids'].shape[1] + debug_log(f'Prompt enhance: input_len={input_len} input_ids_shape={inputs["input_ids"].shape} sample={sample} temp={temperature} penalty={penalty} max_tokens={tokens}') except Exception as e: shared.log.error(f'Prompt enhance tokenize: {e}') errors.display(e, 'Prompt enhance') @@ -409,25 +580,29 @@ class Script(scripts_manager.Script): try: with devices.inference_context(): sd_models.move_model(self.llm, devices.device) - outputs = self.llm.generate( - **inputs, - do_sample=sample, - temperature=float(temperature), - max_new_tokens=int(input_len + tokens), - repetition_penalty=float(penalty), - ) + gen_kwargs = { + 'do_sample': sample, + 'temperature': float(temperature), + 'max_new_tokens': int(tokens), + 'repetition_penalty': float(penalty), + } + if top_k > 0: + gen_kwargs['top_k'] = int(top_k) + if top_p > 0: + gen_kwargs['top_p'] = float(top_p) + outputs = self.llm.generate(**inputs, **gen_kwargs) if shared.opts.diffusers_offload_mode != 'none': - sd_models.move_model(self.llm, devices.cpu) - devices.torch_gc() - if debug_enabled: - raw_response = self.tokenizer.batch_decode(outputs, skip_special_tokens=True, clean_up_tokenization_spaces=True) - shared.log.trace(f'Prompt enhance: raw="{raw_response}"') + sd_models.move_model(self.llm, devices.cpu, force=True) + devices.torch_gc(force=True, reason='prompt enhance offload') outputs_cropped = outputs[:, input_len:] response = self.tokenizer.batch_decode( outputs_cropped, skip_special_tokens=True, clean_up_tokenization_spaces=True, ) + if debug_enabled: + response_before_clean = response[0] if isinstance(response, list) else response + debug_log(f'Prompt enhance: response_before_clean="{response_before_clean}"') except Exception as e: outputs = None shared.log.error(f'Prompt enhance generate: {e}') @@ -440,19 +615,18 @@ class Script(scripts_manager.Script): response = response[0] is_censored = self.censored(response) if not is_censored: - response = self.clean(response) + response = self.clean(response, keep_thinking=keep_thinking, prefill_text=prefill_text, keep_prefill=keep_prefill) response = self.post(response, prefix, suffix, networks) - shared.log.info(f'Prompt enhance: model="{model}" mode="{mode}" nsfw={nsfw} time={t1-t0:.2f} seed={seed} sample={sample} temperature={temperature} penalty={penalty} thinking={thinking} tokens={tokens} inputs={input_len} outputs={outputs.shape[-1] if isinstance(outputs, torch.Tensor) else 0} prompt={len(prompt_text)} response={len(response)}') # Added check for outputs - if debug_enabled: - shared.log.trace(f'Prompt enhance: prompt="{prompt_text}"') - shared.log.trace(f'Prompt enhance: response="{response}"') + shared.log.info(f'Prompt enhance: model="{model}" nsfw={nsfw} time={t1-t0:.2f} seed={seed} sample={sample} temperature={temperature} penalty={penalty} thinking={thinking} keep_thinking={keep_thinking} prefill="{prefill_text[:20] if prefill_text else ""}" keep_prefill={keep_prefill} tokens={tokens} inputs={input_len} outputs={outputs.shape[-1] if isinstance(outputs, torch.Tensor) else 0} prompt={len(prompt_text)} response={len(response)}') + debug_log(f'Prompt enhance: prompt="{prompt_text}"') + debug_log(f'Prompt enhance: response_after_clean="{response}"') self.busy = False if is_censored: shared.log.warning(f'Prompt enhance: censored response="{response}"') return prompt # Return original full prompt on censorship return response - def apply(self, prompt, image, apply_prompt, llm_model, prompt_system, prompt_prefix, prompt_suffix, max_tokens, do_sample, temperature, repetition_penalty, thinking_mode, nsfw_mode): # Added nsfw_mode + def apply(self, prompt, image, apply_prompt, llm_model, prompt_system, prompt_prefix, prompt_suffix, max_tokens, do_sample, temperature, repetition_penalty, top_k, top_p, thinking_mode, nsfw_mode, use_vision, prefill_text, keep_prefill, keep_thinking): response = self.enhance( prompt=prompt, image=image, @@ -464,31 +638,51 @@ class Script(scripts_manager.Script): tokens=max_tokens, temperature=temperature, penalty=repetition_penalty, + top_k=top_k, + top_p=top_p, thinking=thinking_mode, - nsfw=nsfw_mode # Pass nsfw_mode here + nsfw=nsfw_mode, + use_vision=use_vision, + prefill=prefill_text, + keep_prefill=keep_prefill, + keep_thinking=keep_thinking, ) if apply_prompt: return [response, response] return [response, gr.update()] def get_custom(self, name): - model_repo = self.options.models.get(name, {}).get('repo', None) or name - model_gguf = self.options.models.get(name, {}).get('gguf', None) - model_type = self.options.models.get(name, {}).get('type', None) - model_file = self.options.models.get(name, {}).get('file', None) + # Strip symbols from display name to get repo + repo_name = get_model_repo_from_display(name) + model_repo = self.options.models.get(repo_name, {}).get('repo', None) or repo_name + model_gguf = self.options.models.get(repo_name, {}).get('gguf', None) + model_type = self.options.models.get(repo_name, {}).get('type', None) + model_file = self.options.models.get(repo_name, {}).get('file', None) return [model_repo, model_gguf, model_type, model_file] + def update_vision_toggle(self, model_name): + """Update vision toggle interactivity and value based on model selection.""" + repo_name = get_model_repo_from_display(model_name) + is_vl = is_vision_model(repo_name) + # When non-VL model: disable and uncheck. When VL model: enable and check. + return gr.update(interactive=is_vl, value=is_vl) + def ui(self, _is_img2img): with gr.Accordion('Prompt enhance', open=False, elem_id='prompt_enhance'): + gr.HTML('') with gr.Row(): apply_btn = gr.Button(value='Enhance now', elem_id='prompt_enhance_apply', variant='primary') with gr.Row(): apply_prompt = gr.Checkbox(label='Apply to prompt', value=False) apply_auto = gr.Checkbox(label='Auto enhance', value=False) + with gr.Row(): + # Set initial state based on whether default model supports vision + default_is_vl = is_vision_model(Options.default) + use_vision = gr.Checkbox(label='Use vision', value=default_is_vl, interactive=default_is_vl, elem_id='prompt_enhance_use_vision') gr.HTML('
') with gr.Group(): with gr.Row(): - llm_model = gr.Dropdown(label='LLM model', choices=list(self.options.models), value=self.options.default, interactive=True, allow_custom_value=True, elem_id='prompt_enhance_model') + llm_model = gr.Dropdown(label='LLM model', choices=Options.get_model_choices(), value=Options.get_default_display(), interactive=True, allow_custom_value=True, elem_id='prompt_enhance_model') with gr.Row(): load_btn = gr.Button(value='Load model', elem_id='prompt_enhance_load', variant='secondary') load_btn.click(fn=self.load, inputs=[llm_model], outputs=[]) @@ -511,24 +705,32 @@ class Script(scripts_manager.Script): with gr.Accordion('Options', open=False, elem_id='prompt_enhance_options'): with gr.Row(): max_tokens = gr.Slider(label='Max tokens', value=self.options.max_tokens, minimum=10, maximum=1024, step=1, interactive=True) - do_sample = gr.Checkbox(label='Do sample', value=self.options.do_sample, interactive=True) + do_sample = gr.Checkbox(label='Use samplers', value=self.options.do_sample, interactive=True) with gr.Row(): temperature = gr.Slider(label='Temperature', value=self.options.temperature, minimum=0.0, maximum=1.0, step=0.01, interactive=True) repetition_penalty = gr.Slider(label='Repetition penalty', value=self.options.repetition_penalty, minimum=0.0, maximum=2.0, step=0.01, interactive=True) + with gr.Row(): + top_k = gr.Slider(label='Top-K', value=self.options.top_k, minimum=0, maximum=100, step=1, interactive=True) + top_p = gr.Slider(label='Top-P', value=self.options.top_p, minimum=0.0, maximum=1.0, step=0.01, interactive=True) with gr.Row(): nsfw_mode = gr.Checkbox(label='NSFW allowed', value=True, interactive=True) thinking_mode = gr.Checkbox(label='Thinking mode', value=False, interactive=True) + with gr.Row(): + keep_thinking = gr.Checkbox(label='Keep Thinking Trace', value=False, interactive=True) + keep_prefill = gr.Checkbox(label='Keep Prefill', value=False, interactive=True) + with gr.Row(): + prefill_text = gr.Textbox(label='Prefill text', value='', placeholder='Optional: pre-fill start of model response', interactive=True, lines=1) gr.HTML('
') with gr.Accordion('Input', open=False, elem_id='prompt_enhance_system_prompt'): # Corrected elem_id reference with gr.Row(): - prompt_prefix = gr.Textbox(label='Prompt prefix', value='', placeholder='Optional prompt prefix', interactive=True, lines=2, elem_id='prompt_enhance_prefix') + prompt_prefix = gr.Textbox(label='Prompt prefix', value='', placeholder='Text prepended to the enhanced result', interactive=True, lines=2, elem_id='prompt_enhance_prefix') with gr.Row(): - prompt_suffix = gr.Textbox(label='Prompt suffix', value='', placeholder='Optional prompt suffix', interactive=True, lines=2, elem_id='prompt_enhance_suffix') + prompt_suffix = gr.Textbox(label='Prompt suffix', value='', placeholder='Text appended to the enhanced result', interactive=True, lines=2, elem_id='prompt_enhance_suffix') with gr.Row(): - prompt_system = gr.Textbox(label='System prompt', value='', interactive=True, lines=4, elem_id='prompt_enhance_system') # Default to empty as per diff + prompt_system = gr.Textbox(label='System prompt', value='', placeholder='Leave empty to use built-in enhancement instructions', interactive=True, lines=4, elem_id='prompt_enhance_system') with gr.Accordion('Output', open=True, elem_id='prompt_enhance_output'): # Corrected elem_id reference with gr.Row(): - prompt_output = gr.Textbox(label='Enhanced prompt', value='', interactive=True, lines=4) + prompt_output = gr.Textbox(label='Enhanced prompt', value='', placeholder='Enhanced prompt will appear here', interactive=True, lines=4, max_lines=12, elem_id='prompt_enhance_result') with gr.Row(): clear_btn = gr.Button(value='Clear', elem_id='prompt_enhance_clear', variant='secondary') clear_btn.click(fn=lambda: '', inputs=[], outputs=[prompt_output]) @@ -536,8 +738,10 @@ class Script(scripts_manager.Script): copy_btn.click(fn=lambda x: x, inputs=[prompt_output], outputs=[self.prompt]) if self.image is None: self.image = gr.Image(type='pil', interactive=False, visible=False, width=64, height=64) # dummy image - apply_btn.click(fn=self.apply, inputs=[self.prompt, self.image, apply_prompt, llm_model, prompt_system, prompt_prefix, prompt_suffix, max_tokens, do_sample, temperature, repetition_penalty, thinking_mode, nsfw_mode], outputs=[prompt_output, self.prompt]) - return [self.prompt, self.image, apply_auto, llm_model, prompt_system, prompt_prefix, prompt_suffix, max_tokens, do_sample, temperature, repetition_penalty, thinking_mode, nsfw_mode] + # Update vision toggle interactivity when model changes + llm_model.change(fn=self.update_vision_toggle, inputs=[llm_model], outputs=[use_vision], show_progress=False) + apply_btn.click(fn=self.apply, inputs=[self.prompt, self.image, apply_prompt, llm_model, prompt_system, prompt_prefix, prompt_suffix, max_tokens, do_sample, temperature, repetition_penalty, top_k, top_p, thinking_mode, nsfw_mode, use_vision, prefill_text, keep_prefill, keep_thinking], outputs=[prompt_output, self.prompt]) + return [self.prompt, self.image, apply_auto, llm_model, prompt_system, prompt_prefix, prompt_suffix, max_tokens, do_sample, temperature, repetition_penalty, top_k, top_p, thinking_mode, nsfw_mode, use_vision, prefill_text, keep_prefill, keep_thinking] def after_component(self, component, **kwargs): # searching for actual ui prompt components if getattr(component, 'elem_id', '') in ['txt2img_prompt', 'img2img_prompt', 'control_prompt', 'video_prompt']: @@ -548,7 +752,7 @@ class Script(scripts_manager.Script): self.image.use_original = True def before_process(self, p: processing.StableDiffusionProcessing, *args, **kwargs): # pylint: disable=unused-argument - _self_prompt, self_image, apply_auto, llm_model, prompt_system, prompt_prefix, prompt_suffix, max_tokens, do_sample, temperature, repetition_penalty, thinking_mode, nsfw_mode = args + _self_prompt, self_image, apply_auto, llm_model, prompt_system, prompt_prefix, prompt_suffix, max_tokens, do_sample, temperature, repetition_penalty, top_k, top_p, thinking_mode, nsfw_mode, use_vision, prefill_text, keep_prefill, keep_thinking = args if not apply_auto and not p.enhance_prompt: return if shared.state.skipped or shared.state.interrupted: @@ -570,8 +774,14 @@ class Script(scripts_manager.Script): tokens=max_tokens, temperature=temperature, penalty=repetition_penalty, + top_k=top_k, + top_p=top_p, thinking=thinking_mode, nsfw=nsfw_mode, + use_vision=use_vision, + prefill=prefill_text, + keep_prefill=keep_prefill, + keep_thinking=keep_thinking, ) timer.process.record('prompt') p.extra_generation_params['LLM'] = llm_model diff --git a/scripts/xadapter/pipeline_sd_xl_adapter.py b/scripts/xadapter/pipeline_sd_xl_adapter.py index 757681972..788235517 100644 --- a/scripts/xadapter/pipeline_sd_xl_adapter.py +++ b/scripts/xadapter/pipeline_sd_xl_adapter.py @@ -833,7 +833,7 @@ class StableDiffusionXLAdapterPipeline(DiffusionPipeline, FromSingleFileMixin, L latents_sd1_5, ) - # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 6. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 7. Prepare added time ids & embeddings diff --git a/scripts/xadapter/pipeline_sd_xl_adapter_controlnet.py b/scripts/xadapter/pipeline_sd_xl_adapter_controlnet.py index a07982925..853b49e26 100644 --- a/scripts/xadapter/pipeline_sd_xl_adapter_controlnet.py +++ b/scripts/xadapter/pipeline_sd_xl_adapter_controlnet.py @@ -934,7 +934,7 @@ class StableDiffusionXLAdapterControlnetPipeline(DiffusionPipeline, FromSingleFi latents_sd1_5, ) - # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 6. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 7. Prepare added time ids & embeddings diff --git a/scripts/xadapter/pipeline_sd_xl_adapter_controlnet_img2img.py b/scripts/xadapter/pipeline_sd_xl_adapter_controlnet_img2img.py index b2eee115f..932ef5f1a 100644 --- a/scripts/xadapter/pipeline_sd_xl_adapter_controlnet_img2img.py +++ b/scripts/xadapter/pipeline_sd_xl_adapter_controlnet_img2img.py @@ -941,7 +941,7 @@ class StableDiffusionXLAdapterControlnetI2IPipeline(DiffusionPipeline, FromSingl generator, ) - # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + # 6. Prepare extra step kwargs. extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) # 7. Prepare added time ids & embeddings diff --git a/scripts/xyz/xyz_grid_classes.py b/scripts/xyz/xyz_grid_classes.py index d2331a731..37767cdbe 100644 --- a/scripts/xyz/xyz_grid_classes.py +++ b/scripts/xyz/xyz_grid_classes.py @@ -221,7 +221,7 @@ axis_options = [ AxisOption("[Sampler] Timestep spacing", str, apply_setting("schedulers_timestep_spacing"), choices=lambda: ['default', 'linspace', 'leading', 'trailing']), AxisOption("[Sampler] Timestep range", int, apply_setting("schedulers_timesteps_range")), AxisOption("[Sampler] Solver order", int, apply_setting("schedulers_solver_order")), - AxisOption("[Sampler] Beta schedule", str, apply_setting("schedulers_beta_schedule"), choices=lambda: ['default', 'linear', 'scaled', 'cosine']), + AxisOption("[Sampler] Beta schedule", str, apply_setting("schedulers_beta_schedule"), choices=lambda: ['default', 'linear', 'scaled', 'cosine', 'sigmoid', 'laplace']), AxisOption("[Sampler] Beta start", float, apply_setting("schedulers_beta_start")), AxisOption("[Sampler] Beta end", float, apply_setting("schedulers_beta_end")), AxisOption("[Sampler] Flow shift", float, apply_setting("schedulers_shift")), diff --git a/webui.py b/webui.py index db1e2bb22..7606f187a 100644 --- a/webui.py +++ b/webui.py @@ -9,13 +9,18 @@ import logging import importlib import contextlib from threading import Thread +from installer import log, git_commit, custom_excepthook, version +from modules import timer import modules.loader import modules.hashes - -from installer import log, git_commit, custom_excepthook, version -from modules import timer, paths, shared, extensions, gr_tempdir, modelloader, modeldata -from modules.call_queue import queue_lock, wrap_queued_call, wrap_gradio_gpu_call # pylint: disable=unused-import +import modules.paths import modules.devices +from modules import shared +from modules.call_queue import queue_lock, wrap_queued_call, wrap_gradio_gpu_call # pylint: disable=unused-import +import modules.gr_tempdir +import modules.modeldata +import modules.extensions +import modules.modelloader import modules.sd_checkpoint import modules.sd_samplers import modules.scripts_manager @@ -63,7 +68,7 @@ fastapi_args = { def initialize(): - log.debug('Initializing') + log.debug('Initializing: modules') modules.sd_checkpoint.init_metadata() modules.hashes.init_cache() @@ -80,7 +85,7 @@ def initialize(): modules.model_te.refresh_te_list() timer.startup.record("te") - modelloader.cleanup_models() + modules.modelloader.cleanup_models() modules.sd_models.setup_model() timer.startup.record("models") @@ -100,7 +105,7 @@ def initialize(): yolo.initialize() timer.startup.record("detailer") - extensions.list_extensions() + modules.extensions.list_extensions() timer.startup.record("extensions") log.info('Load extensions') @@ -110,7 +115,7 @@ def initialize(): timer.startup.records["extensions"] = t_total # scripts can reset the time log.debug(f'Extensions init time: {t_timer.summary()}') - modelloader.load_upscalers() + modules.modelloader.load_upscalers() timer.startup.record("upscalers") modules.ui_extra_networks.initialize() @@ -151,7 +156,7 @@ def initialize(): def load_model(): - modeldata.model_data.locked = False + modules.modeldata.model_data.locked = False autoload = shared.opts.sd_checkpoint_autoload or shared.cmd_opts.ckpt is not None log.info(f'Model: autoload={autoload} selected="{shared.opts.sd_model_checkpoint}"') if autoload: @@ -169,7 +174,7 @@ def load_model(): shared.opts.onchange("sd_vae", wrap_queued_call(lambda: modules.sd_vae.reload_vae_weights()), call=False) shared.opts.onchange("sd_unet", wrap_queued_call(lambda: modules.sd_unet.load_unet(shared.sd_model)), call=False) shared.opts.onchange("sd_text_encoder", wrap_queued_call(lambda: modules.sd_models.reload_text_encoder()), call=False) - shared.opts.onchange("temp_dir", gr_tempdir.on_tmpdir_changed) + shared.opts.onchange("temp_dir", modules.gr_tempdir.on_tmpdir_changed) timer.startup.record("onchange") @@ -232,7 +237,7 @@ def start_common(): log.info(f'Base path: data="{shared.cmd_opts.data_dir}"') if shared.cmd_opts.models_dir is not None and len(shared.cmd_opts.models_dir) > 0 and shared.cmd_opts.models_dir != 'models': log.info(f'Base path: models="{shared.cmd_opts.models_dir}"') - paths.create_paths(shared.opts) + modules.paths.create_paths(shared.opts) async_policy() initialize() if shared.cmd_opts.backend == 'original': @@ -245,7 +250,7 @@ def start_common(): except Exception: pass if shared.opts.clean_temp_dir_at_start: - gr_tempdir.cleanup_tmpdr() + modules.gr_tempdir.cleanup_tmpdr() timer.startup.record("cleanup") @@ -317,7 +322,7 @@ def start_ui(): _frontend=True and shared.cmd_opts.share, ) if shared.cmd_opts.data_dir is not None: - gr_tempdir.register_tmp_file(shared.demo, os.path.join(shared.cmd_opts.data_dir, 'x')) + modules.gr_tempdir.register_tmp_file(shared.demo, os.path.join(shared.cmd_opts.data_dir, 'x')) shared.log.info(f'Local URL: {local_url}') if shared.cmd_opts.listen: if not gradio_auth_creds: @@ -373,7 +378,7 @@ def webui(restart=False): load_model() mount_subpath(app) - shared.opts.save(shared.config_filename) + shared.opts.save() if shared.cmd_opts.profile: for k, v in modules.script_callbacks.callback_map.items(): diff --git a/wiki b/wiki index 01a5b7af7..e18041f2b 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 01a5b7af78897212a8d1b32def6ff4bd3d03a352 +Subproject commit e18041f2bab7709706fe9205a8f27695e0a5af8f