diff --git a/.eslintrc.json b/.eslintrc.json index 81c196021..90efb00d3 100644 --- a/.eslintrc.json +++ b/.eslintrc.json @@ -17,10 +17,10 @@ ], "env": { "browser": true, - "commonjs": false, - "node": false, - "jquery": false, - "es2021": true + "commonjs": true, + "node": true, + "jquery": true, + "es2024": true }, "rules": { "max-len": [1, 275, 3], @@ -50,6 +50,7 @@ }, "globals": { "panzoom": "readonly", + "authFetch": "readonly", "log": "readonly", "debug": "readonly", "error": "readonly", @@ -122,12 +123,13 @@ "ignorePatterns": [ "node_modules", "extensions", - "extensions-builtin", "repositories", "venv", "panzoom.js", "split.js", "exifr.js", + "jquery.js", + "sparkline.js", "iframeResizer.min.js" ] } diff --git a/.gitmodules b/.gitmodules index 41ed87300..a68f25540 100644 --- a/.gitmodules +++ b/.gitmodules @@ -17,3 +17,6 @@ [submodule "extensions-builtin/sdnext-modernui"] path = extensions-builtin/sdnext-modernui url = https://github.com/BinaryQuantumSoul/sdnext-modernui +[submodule "extensions-builtin/sdnext-kanvas"] + path = extensions-builtin/sdnext-kanvas + url = https://github.com/vladmandic/sdnext-kanvas diff --git a/.pylintrc b/.pylintrc index 27d3e7a52..5f22da840 100644 --- a/.pylintrc +++ b/.pylintrc @@ -61,10 +61,10 @@ ignore-paths=/usr/lib/.*$, scripts/pulid, scripts/xadapter, repositories, - extensions-builtin/Lora, extensions-builtin/sd-extension-chainner/nodes, extensions-builtin/sd-webui-agent-scheduler, extensions-builtin/sdnext-modernui/node_modules, + extensions-builtin/sdnext-kanvas/node_modules, ignore-patterns=.*test*.py$, .*_model.py$, .*_arch.py$, diff --git a/CHANGELOG.md b/CHANGELOG.md index d02eab090..51b1d588f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,62 @@ # Change Log for SD.Next +## Update for 2025-11-13 + +### Highlights for 2025-11-13 + +New native [kanvas](https://vladmandic.github.io/sdnext-docs/Kanvas/) module for image manipulation that fully replaces img2img, inpaint and outpaint controls +And a first cloud model with **Google's Nano Banana** + +![Screenshot](https://github.com/user-attachments/assets/54b25586-b611-4d70-a28f-ee3360944034) + +[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-11-13 + +- **Models** + - [Google Gemini 2.5 Flash Nano Banana](https://blog.google/products/gemini/gemini-nano-banana-examples/) + first cloud-based model directly supported in SD.Next UI + *note*: need to set `GOOGLE_API_KEY` environment variable with your key to use this model + - [Photoroom PRX 1024 Beta](https://huggingface.co/Photoroom/prx-1024-t2i-beta) + PRX (Photoroom Experimental) is a small 1.3-billion-parameter text-to-image model trained entirely from scratch, it uses T5-Gemma text-encoder +- **Features** + - **kanvas**: new module for native canvas-based image manipulation + kanvas is a full replacement for *img2img, inpaint and outpaint* controls + see [docs](https://vladmandic.github.io/sdnext-docs/Kanvas/) for details + *experimental*: report any feedback in master [issue](https://github.com/vladmandic/sdnext/issues/4358) + - **wildcards**: allow recursive inline wildcards using curly braces syntax + - **sdnq**: simplify pre-quantization saved config + - **attention**: additional torch attention settings + - **lora**: separate fuse setting for native-vs-diffuser implementations + - **auth**: strong-enforce auth check on all api endpoints + - **amdgpu**: prefer rocm-on-windows over zluda +- **Internal** + - refactor attention handling + - remove obsolete video scripts + - update global lint rules + - chrono: switch to official pipeline + - pipeline: add optional preprocess and postprocess hooks + - auth: wrap all internal api calls with auth check and use token when possible + - installer: reduce requirements +- **Fixes** + - hires: strength save/load in metadata, thanks @awsr + - imgi2img: fix initial scale tab, thanks @awsr + - img2img: fix restoring refine sampler from metadata, thanks @awsr + - log: client log formatting, thanks @awsr + - rocm: check if installed before forcing install + - pony-v7: fix text-encoder + - detailer: with face-restorers + - detailer: using lora in detailer prompt + - detailer: fail on unsupported models instead of corrputing results + - ui: fix collapsible panels + - svd: fix stable-video-diffusion dtype mismatch + - animatediff: disable sdnq if used + - lora: restore pipeline type if reload/recompile needed + - process: improve send-to functionality + - control: safe load non-sparse controlnet + - control: fix marigold preprocessor with bfloat16 + - auth: fix password being shown in clear text during login + ## Update for 2025-11-06 ### Highlights for 2025-11-06 diff --git a/TODO.md b/TODO.md index 9405eb231..5ca4347b0 100644 --- a/TODO.md +++ b/TODO.md @@ -4,12 +4,17 @@ - +## Kanvas + +- server-side mask handling vs ui mask handling +- implement different auto-masking options + ## Internal - UI: New inpaint/outpaint interface [Kanvas](https://github.com/vladmandic/kanvas) - Deploy: Create executable for SD.Next -- Feature: Integrate natural language imagesearch +- Feature: Integrate natural language image search [ImageDB](https://github.com/vladmandic/imagedb) - Feature: Transformers unified cache handler - Feature: Remote Text-Encoder support diff --git a/cli/api-txt2img.js b/cli/api-txt2img.js index 8d0e9f5d1..7b0f6994a 100755 --- a/cli/api-txt2img.js +++ b/cli/api-txt2img.js @@ -30,10 +30,15 @@ async function main() { const headers = new Headers(); const body = JSON.stringify(sd_options); headers.set('Content-Type', 'application/json'); - if (sd_username && sd_password) headers.set({ Authorization: `Basic ${btoa('sd_username:sd_password')}` }); + if (sd_username && sd_password) { + // const credentials = btoa(`${sd_username}:${sd_password}`); + const credentials = Buffer.from(`${sd_username}:${sd_password}`).toString('base64'); + headers.set('Authorization', `Basic ${credentials}`); + } const res = await fetch(`${sd_url}/sdapi/v1/txt2img`, { method, headers, body }); if (res.status !== 200) { - console.log('Error', res.status); + const err = await res.text(); + console.log('Error', res.status, res.statusText, err); } else { const json = await res.json(); console.log('result:', json.info); diff --git a/extensions-builtin/sd-extension-system-info b/extensions-builtin/sd-extension-system-info index 19a1b1b72..90abd719f 160000 --- a/extensions-builtin/sd-extension-system-info +++ b/extensions-builtin/sd-extension-system-info @@ -1 +1 @@ -Subproject commit 19a1b1b722b768b821f67831a6d3ba83847bcb61 +Subproject commit 90abd719f54578664317325c83781226217d68fd diff --git a/extensions-builtin/sdnext-kanvas b/extensions-builtin/sdnext-kanvas new file mode 160000 index 000000000..ec77ebd4a --- /dev/null +++ b/extensions-builtin/sdnext-kanvas @@ -0,0 +1 @@ +Subproject commit ec77ebd4a50ad5bdc977a7f163c623c9f7bc22e1 diff --git a/extensions-builtin/sdnext-modernui b/extensions-builtin/sdnext-modernui index 5bf81ed7f..404a1e53a 160000 --- a/extensions-builtin/sdnext-modernui +++ b/extensions-builtin/sdnext-modernui @@ -1 +1 @@ -Subproject commit 5bf81ed7fad61833272b6b9523640942c5227351 +Subproject commit 404a1e53a85d64a64c499690c71360735b710a8d diff --git a/html/reference.json b/html/reference.json index c11d627fc..a2baf795d 100644 --- a/html/reference.json +++ b/html/reference.json @@ -912,6 +912,12 @@ "size": 15.48, "date": "2023 April" }, + "Photoroom PRX 1024": { + "path": "Photoroom/prx-1024-t2i-beta", + "desc": "PRX (Photoroom Experimental) is a 1.3-billion-parameter text-to-image model trained entirely from scratch and released under an Apache 2.0 license.", + "preview": "gemini-2.5-flash-image.jpg", + "skip": true + }, "FLUX.1-Dev sdnq-svd-uint4": { "path": "Disty0/FLUX.1-dev-SDNQ-uint4-svd-r32", @@ -1190,6 +1196,14 @@ "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 } } diff --git a/installer.py b/installer.py index 6e7aad790..57ae7a2e2 100644 --- a/installer.py +++ b/installer.py @@ -20,7 +20,7 @@ class Dot(dict): # dot notation access to dictionary attributes pkg_resources, setuptools, distutils = None, None, None # defined via ensure_base_requirements -version = None +version = { 'app': 'sd.next', 'updated': 'unknown', 'commit': 'unknown', 'branch': 'unknown', 'url': 'unknown', 'kanvas': 'unknown' } current_branch = None log = logging.getLogger("sd") console = None @@ -502,6 +502,12 @@ def branch(folder=None): return b +# restart process +def restart(): + log.critical('Restarting process...') + os.execv(sys.executable, ['python'] + sys.argv) + + # update git repository def update(folder, keep_branch = False, rebase = True): t_start = time.time() @@ -513,19 +519,19 @@ def update(folder, keep_branch = False, rebase = True): if keep_branch: res = git(f'pull {arg}', folder) debug(f'Install update: folder={folder} args={arg} {res}') - return res - b = branch(folder) - if branch is None: - res = git(f'pull {arg}', folder) - debug(f'Install update: folder={folder} branch={b} args={arg} {res}') else: - res = git(f'pull origin {b} {arg}', folder) - debug(f'Install update: folder={folder} branch={b} args={arg} {res}') - if not args.experimental: - commit = extensions_commit.get(os.path.basename(folder), None) - if commit is not None: - res = git(f'checkout {commit}', folder) - debug(f'Install update: folder={folder} branch={b} args={arg} commit={commit} {res}') + b = branch(folder) + if branch is None: + res = git(f'pull {arg}', folder) + debug(f'Install update: folder={folder} branch={b} args={arg} {res}') + else: + res = git(f'pull origin {b} {arg}', folder) + debug(f'Install update: folder={folder} branch={b} args={arg} {res}') + if not args.experimental: + commit = extensions_commit.get(os.path.basename(folder), None) + if commit is not None: + res = git(f'checkout {commit}', folder) + debug(f'Install update: folder={folder} branch={b} args={arg} commit={commit} {res}') ts('update', t_start) return res @@ -613,7 +619,7 @@ def check_diffusers(): t_start = time.time() if args.skip_all: return - sha = 'b3e9dfced7c9e8d00f646c710766b532383f04c6' # diffusers commit hash + sha = 'cd3bbe2910666880307b84729176203f5785ff7e' # 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) @@ -690,7 +696,6 @@ def install_rocm_zluda(): amd_gpus = [] try: amd_gpus = rocm.get_agents() - log.info('ROCm: AMD toolkit detected') except Exception as e: log.warning(f'ROCm agent enumerator failed: {e}') @@ -716,7 +721,7 @@ def install_rocm_zluda(): if device_id < len(amd_gpus): device = amd_gpus[device_id] - if sys.platform == "win32" and args.use_rocm and device is not None and device.therock is not None: + if sys.platform == "win32" and not args.use_zluda and device is not None and device.therock is not None and not installed("rocm"): check_python(supported_minors=[11, 12, 13], reason='ROCm backend requires a Python version between 3.11 and 3.13') install(f"rocm rocm-sdk-core --index-url https://rocm.nightlies.amd.com/v2-staging/{device.therock}") rocm.refresh() @@ -727,16 +732,7 @@ def install_rocm_zluda(): log.info(msg) if sys.platform == "win32": - if args.use_rocm: # TODO install: 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.') - 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') - torch_command = os.environ.get('TORCH_COMMAND', f'torch torchvision --index-url https://rocm.nightlies.amd.com/v2-staging/{device.therock}') - else: - check_python(supported_minors=[12], reason='ROCm Windows preview requires Python version 3.12') - 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: + if args.use_zluda: #check_python(supported_minors=[10, 11, 12, 13], reason='ZLUDA backend requires a Python version between 3.10 and 3.13') torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.7.1+cu118 torchvision==0.22.1+cu118 --index-url https://download.pytorch.org/whl/cu118') @@ -759,6 +755,15 @@ def install_rocm_zluda(): zluda_installer.load() except Exception as e: log.warning(f'Failed to load ZLUDA: {e}') + else: # TODO install: 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.') + 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') + torch_command = os.environ.get('TORCH_COMMAND', f'torch torchvision --index-url https://rocm.nightlies.amd.com/v2-staging/{device.therock}') + else: + check_python(supported_minors=[12], reason='ROCm Windows preview requires Python version 3.12') + 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') if args.use_nightly: @@ -912,7 +917,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 rocm.is_installed) # late eval to avoid unnecessary import + 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 if is_cuda_available and args.use_cuda: # prioritize cuda torch_command = install_cuda() @@ -1209,6 +1214,7 @@ def ensure_base_requirements(): setuptools_version = '69.5.1' def update_setuptools(): + local_log = logging.getLogger('sdnext.installer') global pkg_resources, setuptools, distutils # pylint: disable=global-statement # python may ship with incompatible setuptools subprocess.run(f'"{sys.executable}" -m pip install setuptools=={setuptools_version}', shell=True, check=False, env=os.environ, stdout=subprocess.PIPE, stderr=subprocess.PIPE) @@ -1216,12 +1222,27 @@ def ensure_base_requirements(): modules = [m for m in sys.modules if m.startswith('setuptools') or m.startswith('pkg_resources') or m.startswith('distutils')] for m in modules: del sys.modules[m] - setuptools = importlib.import_module('setuptools') - sys.modules['setuptools'] = setuptools - distutils = importlib.import_module('distutils') - sys.modules['distutils'] = distutils - pkg_resources = importlib.import_module('pkg_resources') - sys.modules['pkg_resources'] = pkg_resources + try: + setuptools = importlib.import_module('setuptools') + sys.modules['setuptools'] = setuptools + except ImportError as e: + local_log.info(f'Python: version={platform.python_version()} platform={platform.system()} bin="{sys.executable}" venv="{sys.prefix}"') + local_log.critical(f'Import: setuptools {e}') + os._exit(1) + try: + distutils = importlib.import_module('distutils') + sys.modules['distutils'] = distutils + except ImportError as e: + local_log.info(f'Python: version={platform.python_version()} platform={platform.system()} bin="{sys.executable}" venv="{sys.prefix}"') + local_log.critical(f'Import: distutils {e}') + os._exit(1) + try: + pkg_resources = importlib.import_module('pkg_resources') + sys.modules['pkg_resources'] = pkg_resources + except ImportError as e: + local_log.info(f'Python: version={platform.python_version()} platform={platform.system()} bin="{sys.executable}" venv="{sys.prefix}"') + local_log.critical(f'Import: pkg_resources {e}') + os._exit(1) try: global pkg_resources, setuptools # pylint: disable=global-statement @@ -1420,8 +1441,7 @@ def check_extensions(): def get_version(force=False): t_start = time.time() - global version # pylint: disable=global-statement - if version is None or force: + if (version is None) or (version.get('branch', 'unknown') == 'unknown') or force: try: subprocess.run('git config log.showsignature false', stdout = subprocess.PIPE, stderr = subprocess.PIPE, shell=True, check=True) except Exception: @@ -1429,30 +1449,52 @@ def get_version(force=False): try: res = subprocess.run('git log --pretty=format:"%h %ad" -1 --date=short', stdout = subprocess.PIPE, stderr = subprocess.PIPE, shell=True, check=True) ver = res.stdout.decode(encoding = 'utf8', errors='ignore') if len(res.stdout) > 0 else ' ' - githash, updated = ver.split(' ') + commit, updated = ver.split(' ') + version['commit'], version['updated'] = commit, updated + except Exception as e: + log.warning(f'Version: where=commit {e}') + try: res = subprocess.run('git remote get-url origin', stdout = subprocess.PIPE, stderr = subprocess.PIPE, shell=True, check=True) origin = res.stdout.decode(encoding = 'utf8', errors='ignore') if len(res.stdout) > 0 else '' res = subprocess.run('git rev-parse --abbrev-ref HEAD', stdout = subprocess.PIPE, stderr = subprocess.PIPE, shell=True, check=True) branch_name = res.stdout.decode(encoding = 'utf8', errors='ignore') if len(res.stdout) > 0 else '' - version = { - 'app': 'sd.next', - 'updated': updated, - 'hash': githash, - 'branch': branch_name.replace('\n', ''), - 'url': origin.replace('\n', '').removesuffix('.git') + '/tree/' + branch_name.replace('\n', '') - } - except Exception: - version = { 'app': 'sd.next', 'version': 'unknown', 'branch': 'unknown' } + version['url'] = origin.replace('\n', '').removesuffix('.git') + '/tree/' + branch_name.replace('\n', '') + version['branch'] = branch_name.replace('\n', '') + if version['branch'] == 'HEAD': + log.warning('Version: detached state detected') + except Exception as e: + log.warning(f'Version: where=branch {e}') cwd = os.getcwd() try: - os.chdir('extensions-builtin/sdnext-modernui') - res = subprocess.run('git rev-parse --abbrev-ref HEAD', stdout = subprocess.PIPE, stderr = subprocess.PIPE, shell=True, check=True) - branch_ui = res.stdout.decode(encoding = 'utf8', errors='ignore') if len(res.stdout) > 0 else '' - branch_ui = 'dev' if 'dev' in branch_ui else 'main' - version['ui'] = branch_ui - except Exception: + if os.path.exists('extensions-builtin/sdnext-modernui'): + os.chdir('extensions-builtin/sdnext-modernui') + res = subprocess.run('git rev-parse --abbrev-ref HEAD', stdout = subprocess.PIPE, stderr = subprocess.PIPE, shell=True, check=True) + branch_ui = res.stdout.decode(encoding = 'utf8', errors='ignore') if len(res.stdout) > 0 else '' + branch_ui = 'dev' if 'dev' in branch_ui else 'main' + version['ui'] = branch_ui + else: + version['ui'] = 'unavailable' + except Exception as e: + log.warning(f'Version: where=modernui {e}') version['ui'] = 'unknown' - os.chdir(cwd) + finally: + os.chdir(cwd) + try: + if os.environ.get('SD_KANVAS_DISABLE', None) is not None: + version['kanvas'] = 'disabled' + elif os.path.exists('extensions-builtin/sdnext-kanvas'): + os.chdir('extensions-builtin/sdnext-kanvas') + res = subprocess.run('git rev-parse --abbrev-ref HEAD', stdout = subprocess.PIPE, stderr = subprocess.PIPE, shell=True, check=True) + branch_kanvas = res.stdout.decode(encoding = 'utf8', errors='ignore') if len(res.stdout) > 0 else '' + branch_kanvas = 'dev' if 'dev' in branch_kanvas else 'main' + version['kanvas'] = branch_kanvas + else: + version['kanvas'] = 'unavailable' + except Exception as e: + log.warning(f'Version: where=kanvas {e}') + version['kanvas'] = 'unknown' + finally: + os.chdir(cwd) ts('version', t_start) return version @@ -1461,7 +1503,7 @@ def check_ui(ver): def same(ver): core = ver['branch'] if ver is not None and 'branch' in ver else 'unknown' ui = ver['ui'] if ver is not None and 'ui' in ver else 'unknown' - return (core == ui) or (core == 'master' and ui == 'main') or (core == 'dev' and ui == 'dev') + return (core == ui) or (core == 'master' and ui == 'main') or (core == 'dev' and ui == 'dev') or (core == 'HEAD') t_start = time.time() if not same(ver): @@ -1528,7 +1570,9 @@ def check_version(reset=True): # pylint: disable=unused-argument args.skip_git = True # pylint: disable=attribute-defined-outside-init ver = get_version() log.info(f'Version: {print_dict(ver)}') - branch_name = ver['branch'] if ver is not None and 'branch' in ver else 'master' + branch_name = ver.get('branch', None) if ver is not None else 'master' + if branch_name is None or branch_name == 'unknown': + branch_name = 'master' if args.version or args.skip_git: return check_ui(ver) @@ -1542,9 +1586,24 @@ def check_version(reset=True): # pylint: disable=unused-argument except ImportError: return commits = None + branch_names = [] + try: + branches = requests.get('https://api.github.com/repos/vladmandic/sdnext/branches', timeout=10).json() + branch_names = [b['name'] for b in branches if 'name' in b] + log.trace(f'Repository branches: active={branch_name} available={branch_names}') + except Exception as e: + log.error(f'Repository: failed to get branches: {e}') + return + if branch_name not in branch_names: + log.warning(f'Repository: branch={branch_name} skipping update') + ts('latest', t_start) + return try: commits = requests.get(f'https://api.github.com/repos/vladmandic/sdnext/branches/{branch_name}', timeout=10).json() - if commits['commit']['sha'] != commit and args.upgrade: + latest = commits['commit']['sha'] + if len(latest) != 40: + log.error(f'Repository error: commit={latest} invalid') + elif latest != commit and args.upgrade: global quick_allowed # pylint: disable=global-statement quick_allowed = False log.info('Updating main repository') @@ -1555,6 +1614,8 @@ 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}') + if (ver == latest): # double check + restart() except Exception: if not reset: log.error('Repository error upgrading') diff --git a/javascript/authWrap.js b/javascript/authWrap.js new file mode 100644 index 000000000..d5f78f7b6 --- /dev/null +++ b/javascript/authWrap.js @@ -0,0 +1,20 @@ +let user = null; +let token = null; + +async function authFetch(url, options = {}) { + if (!token) { + const res = await fetch(`${window.subpath}/token`); + if (res.ok) { + const data = await res.json(); + user = data.user; + token = data.token; + } + } + if (user && token) { + if (!options.headers) options.headers = {}; + const encoded = btoa(`${user}:${token}`); + options.headers.Authorization = `Basic ${encoded}`; + } + const res = await fetch(url, options); + return res; +} diff --git a/javascript/civitai.js b/javascript/civitai.js index 71e957380..cf86c069c 100644 --- a/javascript/civitai.js +++ b/javascript/civitai.js @@ -73,7 +73,7 @@ async function modelCardClick(id) { log('modelCardClick id', id); const el = gradioApp().getElementById('model-details') || gradioApp().getElementById('civitai_models_output') || gradioApp().getElementById('models_outcome'); if (!el) return; - const res = await fetch(`${window.api}/civitai?model_id=${encodeURI(id)}`); + const res = await authFetch(`${window.api}/civitai?model_id=${encodeURI(id)}`); if (!res || res.status !== 200) { error(`modelCardClick: id=${id} status=${res ? res.status : 'unknown'}`); return; diff --git a/javascript/contextMenus.js b/javascript/contextMenus.js index 96c57a231..26d10d522 100644 --- a/javascript/contextMenus.js +++ b/javascript/contextMenus.js @@ -132,7 +132,7 @@ const getStatus = async () => { log('progressInternal:', data); if (el) el.innerText += '\nProgress internal:\n' + JSON.stringify(data, null, 2); // eslint-disable-line prefer-template } - res = await fetch('./sdapi/v1/progress?skip_current_image=true', { method: 'GET', headers }); + res = await authFetch('./sdapi/v1/progress?skip_current_image=true', { method: 'GET', headers }); if (res?.ok) { data = await res.json(); log('progressAPI:', data); diff --git a/javascript/control.js b/javascript/control.js index e0bb550ba..ce1ba04c7 100644 --- a/javascript/control.js +++ b/javascript/control.js @@ -3,8 +3,14 @@ function controlInputMode(inputMode, ...args) { if (updateEl) updateEl.click(); const tab = gradioApp().querySelector('#control-tab-input button.selected'); if (!tab) return ['Image', ...args]; - inputMode = tab.innerText; - return [inputMode, ...args]; + let inputTab = tab.innerText; + log('controlInputMode', { mode: inputMode, tab: inputTab, kanvas: typeof Kanvas }); + if ((inputTab === 'Image') && (typeof 'Kanvas' !== 'undefined')) { + inputTab = 'Kanvas'; + const imageData = window.kanvas.getImage(); + args[0] = imageData; + } + return [inputTab, ...args]; } async function setupControlUI() { diff --git a/javascript/extraNetworks.js b/javascript/extraNetworks.js index 204422d78..0143210ea 100644 --- a/javascript/extraNetworks.js +++ b/javascript/extraNetworks.js @@ -153,6 +153,10 @@ async function filterExtraNetworksForTab(searchTerm) { 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'); } else if (searchTerm === 'quantized/') { cards.forEach((elem) => elem.style.display = elem.dataset.tags .toLowerCase() @@ -368,7 +372,7 @@ function selectHistory(id) { const headers = new Headers(); headers.set('Content-Type', 'application/json'); const init = { method: 'POST', body: { name: id }, headers }; - fetch(`${window.api}/history`, { method: 'POST', body: JSON.stringify({ name: id }), headers }); + authFetch(`${window.api}/history`, { method: 'POST', body: JSON.stringify({ name: id }), headers }); } let enDirty = false; diff --git a/javascript/gallery.js b/javascript/gallery.js index 38797551b..da4da5fd3 100644 --- a/javascript/gallery.js +++ b/javascript/gallery.js @@ -185,7 +185,7 @@ async function delayFetchThumb(fn) { while (outstanding > 16) await new Promise((resolve) => setTimeout(resolve, 50)); // eslint-disable-line no-promise-executor-return outstanding++; const ts = Date.now().toString(); - const res = await fetch(`${window.api}/browser/thumb?file=${encodeURI(fn)}&ts=${ts}`, { priority: 'low' }); + const res = await authFetch(`${window.api}/browser/thumb?file=${encodeURI(fn)}&ts=${ts}`, { priority: 'low' }); if (!res.ok) { error(`fetchThumb: ${res.statusText}`); outstanding--; @@ -552,7 +552,7 @@ async function fetchFilesHT(evt) { updateStatusWithSort(`Folder: ${evt.target.name} | in-progress`); let numFiles = 0; - const res = await fetch(`${window.api}/browser/files?folder=${encodeURI(evt.target.name)}`); + const res = await authFetch(`${window.api}/browser/files?folder=${encodeURI(evt.target.name)}`); if (!res || res.status !== 200) { updateStatusWithSort(`Folder: ${evt.target.name} | failed: ${res?.statusText}`); return; @@ -639,7 +639,7 @@ async function pruneImages() { async function galleryVisible() { // if (el.folders.children.length > 0) return; - const res = await fetch(`${window.api}/browser/folders`); + const res = await authFetch(`${window.api}/browser/folders`); if (!res || res.status !== 200) return; el.folders.innerHTML = ''; url = res.url.split('/sdapi')[0].replace('http', 'ws'); // update global url as ws need fqdn diff --git a/javascript/gpu.js b/javascript/gpu.js index 724299375..6e73c7fc9 100644 --- a/javascript/gpu.js +++ b/javascript/gpu.js @@ -30,7 +30,7 @@ async function updateGPU() { const gpuEl = document.getElementById('gpu'); const gpuTable = document.getElementById('gpu-table'); try { - const res = await fetch(`${window.api}/gpu`); + const res = await authFetch(`${window.api}/gpu`); if (!res.ok) { clearInterval(gpuInterval); gpuEl.style.display = 'none'; diff --git a/javascript/history.js b/javascript/history.js index 2a40b0859..c7429ff2b 100644 --- a/javascript/history.js +++ b/javascript/history.js @@ -3,7 +3,7 @@ const ioTypes = ['load', 'save']; function refreshHistory() { log('refreshHistory'); - fetch(`${window.api}/history`, { priority: 'low' }).then((res) => { + authFetch(`${window.api}/history`, { priority: 'low' }).then((res) => { const timeline = document.getElementById('history_timeline'); const table = document.getElementById('history_table'); timeline.innerHTML = ''; diff --git a/javascript/loader.js b/javascript/loader.js index 6d0ddb160..597854fd1 100644 --- a/javascript/loader.js +++ b/javascript/loader.js @@ -52,7 +52,7 @@ async function createSplash() { } const imgEl = `
`; document.getElementById('splash').insertAdjacentHTML('afterbegin', imgEl); - fetch(`${window.api}/motd`) + authFetch(`${window.api}/motd`) .then((res) => res.text()) .then((text) => { const motdEl = document.getElementById('motd'); diff --git a/javascript/logMonitor.js b/javascript/logMonitor.js index ab70400af..c3bcf5673 100644 --- a/javascript/logMonitor.js +++ b/javascript/logMonitor.js @@ -17,20 +17,24 @@ function dateToStr(ts) { return s; } +function htmlEscape(text) { + return text.replaceAll('&', '&').replaceAll('<', '<').replaceAll('>', '>'); +} + async function logMonitor() { const addLogLine = (line) => { try { - const l = JSON.parse(line.replaceAll('\n', ' ')); + const l = JSON.parse(line.replaceAll('\n', ' ').replaceAll('\\', '\\\\')); const row = document.createElement('tr'); // row.style = 'padding: 10px; margin: 0;'; const level = `${l.level}`; if (l.level === 'WARNING') logWarnings++; if (l.level === 'ERROR') logErrors++; const module = `${l.module}`; - row.innerHTML = `${dateToStr(l.created)}${level}${l.facility}${module}${l.msg}`; + row.innerHTML = `${dateToStr(l.created)}${level}${l.facility}${module}${htmlEscape(l.msg)}`; logMonitorEl.appendChild(row); } catch (e) { - error(`logMonitor: ${line}`); + error(`logMonitor: ${e}\n${line}`); } }; @@ -70,7 +74,7 @@ async function logMonitor() { if (!logMonitorEl) return; const atBottom = logMonitorEl.scrollHeight <= (logMonitorEl.scrollTop + logMonitorEl.clientHeight); try { - const res = await fetch(`${window.api}/log?clear=True`); + const res = await authFetch(`${window.api}/log?clear=True`); if (res?.ok) { logMonitorStatus = true; const lines = await res.json(); @@ -119,7 +123,7 @@ async function initLogMonitor() { `; el.style.display = 'none'; - fetch(`${window.api}/start?agent=${encodeURI(navigator.userAgent)}`); + authFetch(`${window.api}/start?agent=${encodeURI(navigator.userAgent)}`); logMonitor(); log('initLogMonitor'); } diff --git a/javascript/logger.js b/javascript/logger.js index d917178bb..c4ac1d78d 100644 --- a/javascript/logger.js +++ b/javascript/logger.js @@ -54,11 +54,11 @@ const xhrInternal = (xhrObj, data, handler = undefined, errorHandler = undefined try { const json = JSON.parse(xhrObj.responseText); if (handler) handler(json); - } catch (e) { - error(`xhr.onreadystatechange: ${e}`); + } catch { + // error(`xhr.onreadystatechange: ${e}`); } } else { - err(`xhr.onreadystatechange: state=${xhrObj.readyState} status=${xhrObj.status} response=${xhrObj.responseText}`); + // err(`xhr.onreadystatechange: state=${xhrObj.readyState} status=${xhrObj.status} response=${xhrObj.responseText}`); } } }; diff --git a/javascript/login.js b/javascript/login.js index 3f2ab1f64..3f684f0fc 100644 --- a/javascript/login.js +++ b/javascript/login.js @@ -4,21 +4,21 @@ const loginCSS = ` left: 0; width: 100%; height: 100%; - background: var(--background-fill-primary); - color: var(--body-text-color-subdued); + background: #222; + color: #ddd; font-family: monospace; z-index: 100; `; const loginHTML = ` -
+

Login

- + - +
- +
`; diff --git a/javascript/monitor.js b/javascript/monitor.js index 276e4bc0e..6cd815945 100644 --- a/javascript/monitor.js +++ b/javascript/monitor.js @@ -23,7 +23,7 @@ async function updateIndicator(online, data, msg) { async function monitorConnection() { try { - const res = await fetch(`${window.api}/version`); + const res = await authFetch(`${window.api}/version`); const data = await res.json(); const url = res.url.split('/sdapi')[0].replace('http', 'ws'); // update global url as ws need fqdn const ws = new WebSocket(`${url}/queue/join`); diff --git a/javascript/settings.js b/javascript/settings.js index 540e65af0..1ffe449e9 100644 --- a/javascript/settings.js +++ b/javascript/settings.js @@ -164,7 +164,7 @@ async function initModels() { const el = gradioApp().getElementById('main_info'); const en = gradioApp().getElementById('txt2img_extra_networks'); if (!el || !en) return; - const req = await fetch(`${window.api}/sd-models`); + const req = await authFetch(`${window.api}/sd-models`); const res = req.ok ? await req.json() : []; log('initModels', res.length); const ready = () => ` @@ -178,7 +178,7 @@ async function initModels() { if (en.classList.contains('hide')) gradioApp().getElementById('txt2img_extra_networks_btn').click(); const repeat = setInterval(() => { const buttons = Array.from(gradioApp().querySelectorAll('#txt2img_model_subdirs > button')) || []; - const reference = buttons.find((b) => (b.innerText === 'Reference') || (b.innerText === 'Distilled') || (b.innerText === 'Community') || (b.innerText === 'Quantized')); + const reference = buttons.find((b) => (b.innerText === 'Reference') || (b.innerText === 'Distilled') || (b.innerText === 'Community') || (b.innerText === 'Quantized') || (b.innerText === 'Cloud')); if (reference) { clearInterval(repeat); reference.click(); diff --git a/javascript/ui.js b/javascript/ui.js index 8c5cd6630..a89c3a249 100644 --- a/javascript/ui.js +++ b/javascript/ui.js @@ -463,7 +463,7 @@ function monitorServerStatus() { function restartReload() { document.body.style = 'background: #222222; font-size: 1rem; font-family:monospace; margin-top:20%; color:lightgray; text-align:center'; document.body.innerHTML = '

Server shutdown in progress...

'; - fetch(`${window.api}/progress?skip_current_image=true`) + authFetch(`${window.api}/progress?skip_current_image=true`) .then((res) => setTimeout(restartReload, 1000)) .catch((e) => setTimeout(monitorServerStatus, 500)); return []; diff --git a/models/Reference/Disty0--NoobAI-XL-Vpred-v1.0-SDNQ-uint4-svd-r128.jpg b/models/Reference/Disty0--NoobAI-XL-Vpred-v1.0-SDNQ-uint4-svd-r128.jpg index 46d04b1c4..d12e7a058 100644 Binary files a/models/Reference/Disty0--NoobAI-XL-Vpred-v1.0-SDNQ-uint4-svd-r128.jpg and b/models/Reference/Disty0--NoobAI-XL-Vpred-v1.0-SDNQ-uint4-svd-r128.jpg differ diff --git a/models/Reference/Disty0--NoobAI-XL-v1.1-SDNQ-uint4-svd-r128.jpg b/models/Reference/Disty0--NoobAI-XL-v1.1-SDNQ-uint4-svd-r128.jpg index 2c96c3cda..ab6c97f3e 100644 Binary files a/models/Reference/Disty0--NoobAI-XL-v1.1-SDNQ-uint4-svd-r128.jpg and b/models/Reference/Disty0--NoobAI-XL-v1.1-SDNQ-uint4-svd-r128.jpg differ diff --git a/models/Reference/Photoroom--prx-1024-t2i-beta.jpg b/models/Reference/Photoroom--prx-1024-t2i-beta.jpg new file mode 100644 index 000000000..a2fc05758 Binary files /dev/null and b/models/Reference/Photoroom--prx-1024-t2i-beta.jpg differ diff --git a/models/Reference/gemini-2.5-flash-image.jpg b/models/Reference/gemini-2.5-flash-image.jpg new file mode 100644 index 000000000..9ed6ec95c Binary files /dev/null and b/models/Reference/gemini-2.5-flash-image.jpg differ diff --git a/models/Reference/noobaiXLNAIXL_epsilonPred11Version.jpg b/models/Reference/noobaiXLNAIXL_epsilonPred11Version.jpg index cbb7cc1fa..acadff10a 100644 Binary files a/models/Reference/noobaiXLNAIXL_epsilonPred11Version.jpg and b/models/Reference/noobaiXLNAIXL_epsilonPred11Version.jpg differ diff --git a/models/Reference/noobaiXLNAIXL_vPred10Version.jpg b/models/Reference/noobaiXLNAIXL_vPred10Version.jpg index 17dd2c8d4..6b6f8fbaa 100644 Binary files a/models/Reference/noobaiXLNAIXL_vPred10Version.jpg and b/models/Reference/noobaiXLNAIXL_vPred10Version.jpg differ diff --git a/models/Reference/ponyRealism_V23.jpg b/models/Reference/ponyRealism_V23.jpg index c92446fd6..42af4f3cd 100644 Binary files a/models/Reference/ponyRealism_V23.jpg and b/models/Reference/ponyRealism_V23.jpg differ diff --git a/models/Reference/waiANIPONYXL_v140.jpg b/models/Reference/waiANIPONYXL_v140.jpg index 7888d1eeb..699c0f285 100644 Binary files a/models/Reference/waiANIPONYXL_v140.jpg and b/models/Reference/waiANIPONYXL_v140.jpg differ diff --git a/models/Reference/waiIllustriousSDXL_v150.jpg b/models/Reference/waiIllustriousSDXL_v150.jpg index 4d697cf70..d6e032388 100644 Binary files a/models/Reference/waiIllustriousSDXL_v150.jpg and b/models/Reference/waiIllustriousSDXL_v150.jpg differ diff --git a/modules/api/api.py b/modules/api/api.py index a260c22ae..ed43190ff 100644 --- a/modules/api/api.py +++ b/modules/api/api.py @@ -4,7 +4,7 @@ from secrets import compare_digest from fastapi import FastAPI, APIRouter, Depends, Request from fastapi.security import HTTPBasic, HTTPBasicCredentials from fastapi.exceptions import HTTPException -from modules import errors, shared, postprocessing +from modules import errors, shared from modules.api import models, endpoints, script, helpers, server, generate, process, control, docs, gpu @@ -60,8 +60,8 @@ class Api: self.add_api_route("/sdapi/v1/txt2img", self.generate.post_text2img, methods=["POST"], response_model=models.ResTxt2Img) self.add_api_route("/sdapi/v1/img2img", self.generate.post_img2img, methods=["POST"], response_model=models.ResImg2Img) self.add_api_route("/sdapi/v1/control", self.control.post_control, methods=["POST"], response_model=control.ResControl) - self.add_api_route("/sdapi/v1/extra-single-image", self.extras_single_image_api, methods=["POST"], response_model=models.ResProcessImage) - self.add_api_route("/sdapi/v1/extra-batch-images", self.extras_batch_images_api, methods=["POST"], response_model=models.ResProcessBatch) + self.add_api_route("/sdapi/v1/extra-single-image", self.process.extras_single_image_api, methods=["POST"], response_model=models.ResProcessImage) + self.add_api_route("/sdapi/v1/extra-batch-images", self.process.extras_batch_images_api, methods=["POST"], response_model=models.ResProcessBatch) self.add_api_route("/sdapi/v1/preprocess", self.process.post_preprocess, methods=["POST"]) self.add_api_route("/sdapi/v1/mask", self.process.post_mask, methods=["POST"]) self.add_api_route("/sdapi/v1/detect", self.process.post_detect, methods=["POST"]) @@ -117,17 +117,25 @@ class Api: from modules.civitai import api_civitai api_civitai.register_api() - - def add_api_route(self, path: str, endpoint, **kwargs): + def add_api_route(self, path: str, fn, auth: bool = True, **kwargs): + if auth and self.credentials: + deps = list(kwargs.get('dependencies', [])) + deps.append(Depends(self.auth)) + kwargs['dependencies'] = deps if shared.opts.subpath is not None and len(shared.opts.subpath) > 0: - self.app.add_api_route(f'{shared.opts.subpath}{path}', endpoint, **kwargs) - self.app.add_api_route(path, endpoint, **kwargs) + self.app.add_api_route(f'{shared.opts.subpath}{path}', endpoint=fn, **kwargs) + self.app.add_api_route(path, endpoint=fn, **kwargs) def auth(self, credentials: HTTPBasicCredentials = Depends(HTTPBasic())): - # this is only needed for api-only since otherwise auth is handled in gradio/routes.py + if not self.credentials: + return True if credentials.username in self.credentials: if compare_digest(credentials.password, self.credentials[credentials.username]): return True + if hasattr(self.app, 'tokens') and (self.app.tokens is not None): + if credentials.password in self.app.tokens.keys(): + return True + shared.log.error(f'API authentication: user="{credentials.username}" password="{credentials.password}"') raise HTTPException(status_code=401, detail="Unauthorized", headers={"WWW-Authenticate": "Basic"}) def get_session_start(self, req: Request, agent: Optional[str] = None): @@ -136,27 +144,6 @@ class Api: shared.log.info(f'Browser session: user={user} client={req.client.host} agent={agent}') return {} - def set_upscalers(self, req: dict): - reqDict = vars(req) - reqDict['extras_upscaler_1'] = reqDict.pop('upscaler_1', None) - reqDict['extras_upscaler_2'] = reqDict.pop('upscaler_2', None) - return reqDict - - def extras_single_image_api(self, req: models.ReqProcessImage): - reqDict = self.set_upscalers(req) - reqDict['image'] = helpers.decode_base64_to_image(reqDict['image']) - with self.queue_lock: - result = postprocessing.run_extras(extras_mode=0, image_folder="", input_dir="", output_dir="", save_output=False, **reqDict) - return models.ResProcessImage(image=helpers.encode_pil_to_base64(result[0][0]), html_info=result[1]) - - def extras_batch_images_api(self, req: models.ReqProcessBatch): - reqDict = self.set_upscalers(req) - image_list = reqDict.pop('imageList', []) - image_folder = [helpers.decode_base64_to_image(x.data) for x in image_list] - with self.queue_lock: - result = postprocessing.run_extras(extras_mode=1, image_folder=image_folder, image="", input_dir="", output_dir="", save_output=False, **reqDict) - return models.ResProcessBatch(images=list(map(helpers.encode_pil_to_base64, result[0])), html_info=result[1]) - def launch(self): config = { "listen": shared.cmd_opts.listen, diff --git a/modules/api/middleware.py b/modules/api/middleware.py index 6270a9ff2..8dc10e31e 100644 --- a/modules/api/middleware.py +++ b/modules/api/middleware.py @@ -81,7 +81,8 @@ def setup_middleware(app: FastAPI, cmd_opts): if err['code'] == 404 and 'file=html/' in req.url.path: # dont spam with locales return JSONResponse(status_code=err['code'], content=jsonable_encoder(err)) - log.error(f"API error: {req.method}: {req.url} {err}") + if not any([req.url.path.endswith(x) for x in ignore_endpoints]): # noqa C419 # pylint: disable=use-a-generator + log.error(f"API error: {req.method}: {req.url} {err}") if not isinstance(e, HTTPException) and err['error'] != 'TypeError': # do not print backtrace on known httpexceptions errors.display(e, 'HTTP API', [anyio, fastapi, uvicorn, starlette]) diff --git a/modules/api/process.py b/modules/api/process.py index 3151907a5..c106a18e2 100644 --- a/modules/api/process.py +++ b/modules/api/process.py @@ -4,8 +4,8 @@ from pydantic import BaseModel, Field # pylint: disable=no-name-in-module from fastapi.responses import JSONResponse from fastapi.exceptions import HTTPException from modules.api.helpers import decode_base64_to_image, encode_pil_to_base64 -from modules import errors, shared -from modules.api import models +from modules import errors, shared, postprocessing +from modules.api import models, helpers processor = None # cached instance of processor @@ -175,3 +175,24 @@ class APIProcess(): raise HTTPException(status_code=400, detail="prompt enhancement: invalid type") res = models.ResPromptEnhance(prompt=prompt, seed=seed) return res + + def set_upscalers(self, req: dict): + reqDict = vars(req) + reqDict['extras_upscaler_1'] = reqDict.pop('upscaler_1', None) + reqDict['extras_upscaler_2'] = reqDict.pop('upscaler_2', None) + return reqDict + + def extras_single_image_api(self, req: models.ReqProcessImage): + reqDict = self.set_upscalers(req) + reqDict['image'] = helpers.decode_base64_to_image(reqDict['image']) + with self.queue_lock: + result = postprocessing.run_extras(extras_mode=0, image_folder="", input_dir="", output_dir="", save_output=False, **reqDict) + return models.ResProcessImage(image=helpers.encode_pil_to_base64(result[0][0]), html_info=result[1]) + + def extras_batch_images_api(self, req: models.ReqProcessBatch): + reqDict = self.set_upscalers(req) + image_list = reqDict.pop('imageList', []) + image_folder = [helpers.decode_base64_to_image(x.data) for x in image_list] + with self.queue_lock: + result = postprocessing.run_extras(extras_mode=1, image_folder=image_folder, image="", input_dir="", output_dir="", save_output=False, **reqDict) + return models.ResProcessBatch(images=list(map(helpers.encode_pil_to_base64, result[0])), html_info=result[1]) diff --git a/modules/api/server.py b/modules/api/server.py index b459ed605..7fa9eca3b 100644 --- a/modules/api/server.py +++ b/modules/api/server.py @@ -15,7 +15,7 @@ def get_motd(): motd = '' ver = shared.get_version() if ver.get('updated', None) is not None: - motd = f"version {ver['hash']} {ver['updated']} {ver['url'].split('/')[-1]}
" + motd = f"version {ver['commit']} {ver['updated']} {ver['url'].split('/')[-1]}
" # pylint: disable=use-maxsplit-arg if shared.opts.motd: try: res = requests.get('https://vladmandic.github.io/sdnext/motd', timeout=3) diff --git a/modules/attention.py b/modules/attention.py new file mode 100644 index 000000000..e92eb95dc --- /dev/null +++ b/modules/attention.py @@ -0,0 +1,234 @@ +from typing import Optional +from functools import wraps +import torch +from modules import rocm +from modules.errors import log +from installer import install, installed + + +def set_dynamic_attention(): + try: + sdpa_pre_dyanmic_atten = torch.nn.functional.scaled_dot_product_attention + from modules.sd_hijack_dynamic_atten import dynamic_scaled_dot_product_attention + torch.nn.functional.scaled_dot_product_attention = dynamic_scaled_dot_product_attention + return sdpa_pre_dyanmic_atten + except Exception as err: + log.error(f'Torch attention: type="dynamic attention" {err}') + return None + + +def set_triton_flash_attention(backend: str): + try: + if backend in {"rocm", "zluda"}: # flash_attn_triton_amd only works with AMD + from modules.flash_attn_triton_amd import interface_fa + sdpa_pre_triton_flash_atten = torch.nn.functional.scaled_dot_product_attention + @wraps(sdpa_pre_triton_flash_atten) + def sdpa_triton_flash_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.Tensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor: + if query.shape[-1] <= 128 and attn_mask is None and query.dtype != torch.float32: + if scale is None: + scale = query.shape[-1] ** (-0.5) + head_size_og = query.size(3) + if head_size_og % 8 != 0: + query = torch.nn.functional.pad(query, [0, 8 - head_size_og % 8]) + key = torch.nn.functional.pad(key, [0, 8 - head_size_og % 8]) + value = torch.nn.functional.pad(value, [0, 8 - head_size_og % 8]) + query = query.transpose(1, 2) + key = key.transpose(1, 2) + value = value.transpose(1, 2) + out_padded = torch.zeros_like(query) + interface_fa.fwd(query, key, value, out_padded, dropout_p, scale, is_causal) + return out_padded[..., :head_size_og].transpose(1, 2) + else: + if enable_gqa: + kwargs["enable_gqa"] = enable_gqa + return sdpa_pre_triton_flash_atten(query=query, key=key, value=value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, scale=scale, **kwargs) + torch.nn.functional.scaled_dot_product_attention = sdpa_triton_flash_atten + log.debug('Torch attention: type="Triton Flash attention"') + except Exception as err: + log.error(f'Torch attention: type="Triton Flash attention" {err}') + + +def set_flex_attention(): + try: + from torch.nn.attention.flex_attention import flex_attention, create_block_mask + def flex_attention_causal_mask(b, h, q_idx, kv_idx): # pylint: disable=unused-argument + return q_idx >= kv_idx + + sdpa_pre_flex_atten = torch.nn.functional.scaled_dot_product_attention + @wraps(sdpa_pre_flex_atten) + def sdpa_flex_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.Tensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor: # pylint: disable=unused-argument + score_mod = None + block_mask = None + if attn_mask is not None: + batch_size, num_heads = query.shape[:2] + seq_len_q = query.shape[-2] + seq_len_kv = key.shape[-2] + if attn_mask.ndim == 2: + attn_mask = attn_mask.view(attn_mask.shape[0], 1, attn_mask.size[1], 1) + attn_mask = attn_mask.expand(batch_size, num_heads, seq_len_q, seq_len_kv) + if attn_mask.dtype == torch.bool: + def mask_mod(batch_idx, head_idx, q_idx, kv_idx): + return attn_mask[batch_idx, head_idx, q_idx, kv_idx] + block_mask = create_block_mask(mask_mod, batch_size, None, seq_len_q, seq_len_kv, device=query.device) + else: + def score_mod_fn(score, batch_idx, head_idx, q_idx, kv_idx): + return score + attn_mask[batch_idx, head_idx, q_idx, kv_idx] + score_mod = score_mod_fn + elif is_causal: + block_mask = create_block_mask(flex_attention_causal_mask, query.shape[0], query.shape[1], query.shape[-2], key.shape[-2], device=query.device) + return flex_attention(query, key, value, score_mod=score_mod, block_mask=block_mask, scale=scale, enable_gqa=enable_gqa) + + torch.nn.functional.scaled_dot_product_attention = sdpa_flex_atten + log.debug('Torch attention: type="Flex attention"') + except Exception as err: + log.error(f'Torch attention: type="Flex attention" {err}') + + +def set_ck_flash_attention(backend: str, device: torch.device): + try: + 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")) + install(rocm.get_flash_attention_command(agent), reinstall=True) + else: + install('flash-attn') + from flash_attn import flash_attn_func + sdpa_pre_flash_atten = torch.nn.functional.scaled_dot_product_attention + @wraps(sdpa_pre_flash_atten) + def sdpa_flash_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.Tensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor: + if query.shape[-1] <= 128 and attn_mask is None and query.dtype != torch.float32: + is_unsqueezed = False + if query.dim() == 3: + query = query.unsqueeze(0) + is_unsqueezed = True + if key.dim() == 3: + key = key.unsqueeze(0) + if value.dim() == 3: + value = value.unsqueeze(0) + if enable_gqa: + key = key.repeat_interleave(query.size(-3)//key.size(-3), -3) + value = value.repeat_interleave(query.size(-3)//value.size(-3), -3) + query = query.transpose(1, 2) + key = key.transpose(1, 2) + value = value.transpose(1, 2) + attn_output = flash_attn_func(q=query, k=key, v=value, dropout_p=dropout_p, causal=is_causal, softmax_scale=scale).transpose(1, 2) + if is_unsqueezed: + attn_output = attn_output.squeeze(0) + return attn_output + else: + if enable_gqa: + kwargs["enable_gqa"] = enable_gqa + return sdpa_pre_flash_atten(query=query, key=key, value=value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, scale=scale, **kwargs) + torch.nn.functional.scaled_dot_product_attention = sdpa_flash_atten + log.debug('Torch attention: type="Flash attention"') + except Exception as err: + log.error(f'Torch attention: type="Flash attention" {err}') + + +def set_sage_attention(backend: str, device: torch.device): + try: + install('sageattention') + + use_cuda_backend = False + if (backend == "cuda") and (torch.cuda.get_device_capability(device) == (8, 6)): + use_cuda_backend = True # Detect GPU architecture - sm86 confirmed to need CUDA backend workaround as Sage Attention + Triton causes NaNs + try: + from sageattention import sageattn_qk_int8_pv_fp16_cuda + except Exception: + use_cuda_backend = False + + if use_cuda_backend: + from sageattention import sageattn_qk_int8_pv_fp16_cuda + def sage_attn_impl(query, key, value, is_causal, scale): + return sageattn_qk_int8_pv_fp16_cuda( + q=query, k=key, v=value, + tensor_layout="HND", + is_causal=is_causal, + sm_scale=scale, + return_lse=False, + pv_accum_dtype="fp32", + ) + else: + from sageattention import sageattn + def sage_attn_impl(query, key, value, is_causal, scale): + return sageattn( + q=query, k=key, v=value, + attn_mask=None, + dropout_p=0.0, + is_causal=is_causal, + scale=scale, + ) + + sdpa_pre_sage_atten = torch.nn.functional.scaled_dot_product_attention + @wraps(sdpa_pre_sage_atten) + def sdpa_sage_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.Tensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor: + if (query.shape[-1] in {128, 96, 64}) and (attn_mask is None) and (query.dtype != torch.float32): + if enable_gqa: + key = key.repeat_interleave(query.size(-3)//key.size(-3), -3) + value = value.repeat_interleave(query.size(-3)//value.size(-3), -3) + + # Call pre-selected sage attention implementation + return sage_attn_impl(query, key, value, is_causal, scale) + else: + if enable_gqa: + kwargs["enable_gqa"] = enable_gqa + return sdpa_pre_sage_atten(query=query, key=key, value=value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, scale=scale, **kwargs) + torch.nn.functional.scaled_dot_product_attention = sdpa_sage_atten + log.debug(f'Torch attention: type="Sage attention" backend={"cuda" if use_cuda_backend else "auto"}') + except Exception as err: + log.error(f'Torch attention: type="Sage attention" {err}') + + +def set_diffusers_attention(pipe, quiet:bool=False): + from modules import shared + import diffusers.models.attention_processor as p + + def set_attn(pipe, attention, name:str=None): + if attention is None: + return + # other models uses their own attention processor + if getattr(pipe, "unet", None) is not None and hasattr(pipe.unet, "set_attn_processor"): + try: + pipe.unet.set_attn_processor(attention) + except Exception as e: + if 'Nunchaku' in pipe.unet.__class__.__name__: + pass + else: + shared.log.error(f'Torch attention: type="{name}" cls={attention.__class__.__name__} pipe={pipe.__class__.__name__} {e}') + """ # each transformer typically has its own attention processor + if getattr(pipe, "transformer", None) is not None and hasattr(pipe.transformer, "set_attn_processor"): + try: + pipe.transformer.set_attn_processor(attention) + except Exception as e: + if 'Nunchaku' in pipe.transformer.__class__.__name__: + pass + else: + shared.log.error(f'Torch attention: type="{name}" cls={attention.__class__.__name__} pipe={pipe.__class__.__name__} {e}') + """ + + shared.log.quiet(quiet, f'Setting model: attention="{shared.opts.cross_attention_optimization}"') + if shared.opts.cross_attention_optimization == "Disabled": + pass # do nothing + elif shared.opts.cross_attention_optimization == "Scaled-Dot-Product": # The default set by Diffusers + # set_attn(pipe, p.AttnProcessor2_0(), name="Scaled-Dot-Product") + pass + elif shared.opts.cross_attention_optimization == "xFormers": + if hasattr(pipe, 'enable_xformers_memory_efficient_attention'): + pipe.enable_xformers_memory_efficient_attention() + else: + shared.log.warning(f"Attention: xFormers is not compatible with {pipe.__class__.__name__}") + elif shared.opts.cross_attention_optimization == "Batch matrix-matrix": + set_attn(pipe, p.AttnProcessor(), name="Batch matrix-matrix") + elif shared.opts.cross_attention_optimization == "Dynamic Attention BMM": + from modules.sd_hijack_dynamic_atten import DynamicAttnProcessorBMM + set_attn(pipe, DynamicAttnProcessorBMM(), name="Dynamic Attention BMM") + + if shared.opts.attention_slicing != "Default" and hasattr(pipe, "enable_attention_slicing") and hasattr(pipe, "disable_attention_slicing"): + if shared.opts.attention_slicing: + pipe.enable_attention_slicing() + else: + pipe.disable_attention_slicing() + shared.log.debug(f"Torch attention: slicing={shared.opts.attention_slicing}") + + pipe.current_attn_name = shared.opts.cross_attention_optimization diff --git a/modules/control/proc/marigold/__init__.py b/modules/control/proc/marigold/__init__.py index 6ea346430..af29be777 100644 --- a/modules/control/proc/marigold/__init__.py +++ b/modules/control/proc/marigold/__init__.py @@ -1,3 +1,4 @@ +import torch from PIL import Image from modules.control.util import HWC3, resize_image from modules import devices @@ -28,7 +29,7 @@ class MarigoldDetector: color_map: str = "Spectral", output_type=None, ): - self.model.to(device=devices.device, dtype=devices.dtype) + self.model.to(device=devices.device, dtype=torch.float16) res = self.model( input_image, denoising_steps=denoising_steps, diff --git a/modules/control/proc/marigold/marigold_pipeline.py b/modules/control/proc/marigold/marigold_pipeline.py index a8cec453e..ac60f833e 100644 --- a/modules/control/proc/marigold/marigold_pipeline.py +++ b/modules/control/proc/marigold/marigold_pipeline.py @@ -228,7 +228,7 @@ class MarigoldPipeline(DiffusionPipeline): depth_pred = (depth_pred - min_d) / (max_d - min_d) # Convert to numpy - depth_pred = depth_pred.cpu().numpy().astype(np.float32) + depth_pred = depth_pred.to(torch.float32).cpu().numpy() # Resize back to original resolution if match_input_res: diff --git a/modules/control/proc/marigold/util/ensemble.py b/modules/control/proc/marigold/util/ensemble.py index d7a9023ca..710db1cc2 100644 --- a/modules/control/proc/marigold/util/ensemble.py +++ b/modules/control/proc/marigold/util/ensemble.py @@ -64,8 +64,9 @@ def ensemble_depths( input_images = downscaler(torch.from_numpy(input_images)).numpy() # init guess - _min = np.min(input_images.reshape((n_img, -1)).cpu().numpy(), axis=1) - _max = np.max(input_images.reshape((n_img, -1)).cpu().numpy(), axis=1) + np_img = input_images.reshape((n_img, -1)).to(torch.float32).cpu().numpy() + _min = np.min(np_img, axis=1) + _max = np.max(np_img, axis=1) s_init = 1.0 / (_max - _min).reshape((-1, 1, 1)) t_init = (-1 * s_init.flatten() * _min.flatten()).reshape((-1, 1, 1)) x = np.concatenate([s_init, t_init]).reshape(-1).astype(np_dtype) @@ -95,7 +96,7 @@ def ensemble_depths( far_err = torch.sqrt((1 - torch.max(pred)) ** 2) err = sqrt_dist + (near_err + far_err) * regularizer_strength - err = err.detach().cpu().numpy().astype(np_dtype) + err = err.to(torch.float32).detach().cpu().numpy().astype(np_dtype) return err res = minimize( diff --git a/modules/control/run.py b/modules/control/run.py index b8078ecaa..b33a2774d 100644 --- a/modules/control/run.py +++ b/modules/control/run.py @@ -474,7 +474,7 @@ def control_run(state: str = '', # pylint: disable=keyword-arg-before-vararg try: with devices.inference_context(): - if isinstance(inputs, str): # only video, the rest is a list + if isinstance(inputs, str) and os.path.exists(inputs): # only video, the rest is a list if input_type == 2: # separate init image if isinstance(inits, str) and inits != inputs: shared.log.warning('Control: separate init video not support for video input') @@ -521,7 +521,7 @@ def control_run(state: str = '', # pylint: disable=keyword-arg-before-vararg yield terminate('Interrupted') return [], '', '', 'Interrupted' # get input - if isinstance(input_image, str): + if isinstance(input_image, str) and os.path.exists(input_image): try: input_image = Image.open(input_image) except Exception as e: diff --git a/modules/control/units/controlnet.py b/modules/control/units/controlnet.py index 0108cbfa4..571491987 100644 --- a/modules/control/units/controlnet.py +++ b/modules/control/units/controlnet.py @@ -353,7 +353,12 @@ class ControlNet(): if self.dtype is not None: self.model.to(self.dtype) if self.device is not None: - self.model.to_empty(device=self.device) # model could be sparse + if (opts.diffusers_offload_mode != 'balanced') and hasattr(self.model, 'to'): + try: + self.model.to(self.device) + except Exception as e: + if 'Cannot copy out of meta tensor' in str(e): + self.model.to_empty(device=self.device) if "Control" in opts.sdnq_quantize_weights: try: log.debug(f'Control {what} model SDNQ quantize: id="{model_id}"') diff --git a/modules/devices.py b/modules/devices.py index 8945c0a58..76d9a8652 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -1,14 +1,10 @@ -from typing import Optional - import os import sys import time import contextlib -from functools import wraps import torch -from modules import rocm +from modules import rocm, attention from modules.errors import log, display, install as install_traceback -from installer import install, installed debug = os.environ.get('SD_DEVICE_DEBUG', None) is not None @@ -462,148 +458,33 @@ def set_sdpa_params(): log.warning(f'Torch attention: type="sdpa" {err}') try: - torch.backends.cuda.enable_flash_sdp('Flash attention' in opts.sdp_options) - torch.backends.cuda.enable_mem_efficient_sdp('Memory attention' in opts.sdp_options) - torch.backends.cuda.enable_math_sdp('Math attention' in opts.sdp_options) + torch.backends.cuda.enable_flash_sdp('Flash' in opts.sdp_options or 'Flash attention' in opts.sdp_options) + torch.backends.cuda.enable_mem_efficient_sdp('Memory' in opts.sdp_options or 'Memory attention' in opts.sdp_options) + torch.backends.cuda.enable_math_sdp('Math' in opts.sdp_options or 'Math attention' in opts.sdp_options) if hasattr(torch.backends.cuda, "allow_fp16_bf16_reduction_math_sdp"): # only valid for torch >= 2.5 torch.backends.cuda.allow_fp16_bf16_reduction_math_sdp(True) - log.debug(f'Torch attention: type="sdpa" opts={opts.sdp_options}') + log.debug(f'Torch attention: type="sdpa" kernels={opts.sdp_options} overrides={opts.sdp_overrides}') except Exception as err: log.warning(f'Torch attention: type="sdpa" {err}') # Stack hijcaks in reverse order. This gives priority to the last added hijack. # If the last hijack is not compatible, it will use the one before it and so on. - if 'Dynamic attention' in opts.sdp_options: - try: - global sdpa_pre_dyanmic_atten # pylint: disable=global-statement - sdpa_pre_dyanmic_atten = torch.nn.functional.scaled_dot_product_attention - from modules.sd_hijack_dynamic_atten import dynamic_scaled_dot_product_attention - torch.nn.functional.scaled_dot_product_attention = dynamic_scaled_dot_product_attention - except Exception as err: - log.error(f'Torch attention: type="dynamic attention" {err}') + if 'Dynamic attention' in opts.sdp_overrides: + global sdpa_pre_dyanmic_atten # pylint: disable=global-statement + sdpa_pre_dyanmic_atten = attention.set_dynamic_attention() - if 'Triton Flash attention' in opts.sdp_options: - try: - if backend in {"zluda", "rocm"}: - from modules.flash_attn_triton_amd import interface_fa - sdpa_pre_triton_flash_atten = torch.nn.functional.scaled_dot_product_attention - @wraps(sdpa_pre_triton_flash_atten) - def sdpa_triton_flash_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.FloatTensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor: - if query.shape[-1] <= 128 and attn_mask is None and query.dtype != torch.float32: - if scale is None: - scale = query.shape[-1] ** (-0.5) - head_size_og = query.size(3) - if head_size_og % 8 != 0: - query = torch.nn.functional.pad(query, [0, 8 - head_size_og % 8]) - key = torch.nn.functional.pad(key, [0, 8 - head_size_og % 8]) - value = torch.nn.functional.pad(value, [0, 8 - head_size_og % 8]) - query = query.transpose(1, 2) - key = key.transpose(1, 2) - value = value.transpose(1, 2) - out_padded = torch.zeros_like(query) - interface_fa.fwd(query, key, value, out_padded, dropout_p, scale, is_causal) - return out_padded[..., :head_size_og].transpose(1, 2) - else: - if enable_gqa: - kwargs["enable_gqa"] = enable_gqa - return sdpa_pre_triton_flash_atten(query=query, key=key, value=value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, scale=scale, **kwargs) - torch.nn.functional.scaled_dot_product_attention = sdpa_triton_flash_atten - log.debug('Torch attention: type="triton flash attention"') - except Exception as err: - log.error(f'Torch attention: type="triton flash attention" {err}') + if 'Flex attention' in opts.sdp_overrides: + attention.set_flex_attention() - if 'CK Flash attention' in opts.sdp_options: - try: - if backend == "rocm": - if not installed('flash-attn'): - log.info('Building CK Flash attention...') - agent = rocm.Agent(getattr(torch.cuda.get_device_properties(device), "gcnArchName", "gfx0000")) - install(rocm.get_flash_attention_command(agent), reinstall=True) - else: - install('flash-attn') - from flash_attn import flash_attn_func - sdpa_pre_flash_atten = torch.nn.functional.scaled_dot_product_attention - @wraps(sdpa_pre_flash_atten) - def sdpa_flash_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.FloatTensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor: - if query.shape[-1] <= 128 and attn_mask is None and query.dtype != torch.float32: - is_unsqueezed = False - if query.dim() == 3: - query = query.unsqueeze(0) - is_unsqueezed = True - if key.dim() == 3: - key = key.unsqueeze(0) - if value.dim() == 3: - value = value.unsqueeze(0) - if enable_gqa: - key = key.repeat_interleave(query.size(-3)//key.size(-3), -3) - value = value.repeat_interleave(query.size(-3)//value.size(-3), -3) - query = query.transpose(1, 2) - key = key.transpose(1, 2) - value = value.transpose(1, 2) - attn_output = flash_attn_func(q=query, k=key, v=value, dropout_p=dropout_p, causal=is_causal, softmax_scale=scale).transpose(1, 2) - if is_unsqueezed: - attn_output = attn_output.squeeze(0) - return attn_output - else: - if enable_gqa: - kwargs["enable_gqa"] = enable_gqa - return sdpa_pre_flash_atten(query=query, key=key, value=value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, scale=scale, **kwargs) - torch.nn.functional.scaled_dot_product_attention = sdpa_flash_atten - log.debug('Torch attention: type="ck flash attention"') - except Exception as err: - log.error(f'Torch attention: type="ck flash attention" {err}') + if 'Triton Flash attention' in opts.sdp_overrides: + attention.set_triton_flash_attention(backend) - if 'Sage attention' in opts.sdp_options: - try: - install('sageattention') - from sageattention import sageattn, sageattn_qk_int8_pv_fp16_cuda - - use_cuda_backend = False - if (backend == "cuda") and (torch.cuda.get_device_capability(device) == (8, 6)): - use_cuda_backend = True # Detect GPU architecture - sm86 confirmed to need CUDA backend workaround as Sage Attention + Triton causes NaNs - - if use_cuda_backend: - log.debug('Torch attention: type=SageAttention backend=cuda') - def sage_attn_impl(query, key, value, is_causal, scale): - return sageattn_qk_int8_pv_fp16_cuda( - q=query, k=key, v=value, - tensor_layout="HND", - is_causal=is_causal, - sm_scale=scale, - return_lse=False, - pv_accum_dtype="fp32", - ) - else: - log.debug('Torch attention: type=SageAttention backend=auto') - def sage_attn_impl(query, key, value, is_causal, scale): - return sageattn( - q=query, k=key, v=value, - attn_mask=None, - dropout_p=0.0, - is_causal=is_causal, - scale=scale, - ) - - sdpa_pre_sage_atten = torch.nn.functional.scaled_dot_product_attention - @wraps(sdpa_pre_sage_atten) - def sdpa_sage_atten(query: torch.FloatTensor, key: torch.FloatTensor, value: torch.FloatTensor, attn_mask: Optional[torch.FloatTensor] = None, dropout_p: float = 0.0, is_causal: bool = False, scale: Optional[float] = None, enable_gqa: bool = False, **kwargs) -> torch.FloatTensor: - if (query.shape[-1] in {128, 96, 64}) and (attn_mask is None) and (query.dtype != torch.float32): - if enable_gqa: - key = key.repeat_interleave(query.size(-3)//key.size(-3), -3) - value = value.repeat_interleave(query.size(-3)//value.size(-3), -3) - - # Call pre-selected sage attention implementation - return sage_attn_impl(query, key, value, is_causal, scale) - else: - if enable_gqa: - kwargs["enable_gqa"] = enable_gqa - return sdpa_pre_sage_atten(query=query, key=key, value=value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, scale=scale, **kwargs) - torch.nn.functional.scaled_dot_product_attention = sdpa_sage_atten - log.debug('Torch attention: type="sage attention"') - except Exception as err: - log.error(f'Torch attention: type="sage attention" {err}') + if 'Flash attention' in opts.sdp_overrides: + attention.set_ck_flash_attention(backend, device) + if 'Sage attention' in opts.sdp_overrides: + attention.set_sage_attention(backend, device) from importlib.metadata import version try: diff --git a/modules/extra_networks.py b/modules/extra_networks.py index bb00c7847..8e4019ac4 100644 --- a/modules/extra_networks.py +++ b/modules/extra_networks.py @@ -58,7 +58,7 @@ class ExtraNetwork: """ raise NotImplementedError - def deactivate(self, p): + def deactivate(self, p, force=False): """ Called at the end of processing for housekeeping. No need to do anything here. """ @@ -123,7 +123,7 @@ def activate(p, extra_network_data=None, step=0, include=[], exclude=[]): shared.opts.data['lora_functional'] = functional -def deactivate(p, extra_network_data=None): +def deactivate(p, extra_network_data=None, force=shared.opts.lora_force_reload): """call deactivate for extra networks in extra_network_data in specified order, then call deactivate for all remaining registered networks""" if p.disable_extra_networks: return @@ -135,7 +135,7 @@ def deactivate(p, extra_network_data=None): if extra_network is None: continue try: - extra_network.deactivate(p) + extra_network.deactivate(p, force=force) except Exception as e: errors.display(e, f"deactivating extra network {extra_network_name}") @@ -144,7 +144,7 @@ def deactivate(p, extra_network_data=None): if args is not None: continue try: - extra_network.deactivate(p) + extra_network.deactivate(p, force=force) except Exception as e: errors.display(e, f"deactivating unmentioned extra network {extra_network_name}") @@ -154,6 +154,8 @@ re_extra_net = re.compile(r"<(\w+):([^>]+)>") def parse_prompt(prompt): res = defaultdict(list) + if prompt is None: + return prompt, res def found(m): name = m.group(1) @@ -170,6 +172,8 @@ def parse_prompt(prompt): def parse_prompts(prompts): res = [] extra_data = None + if prompts is None: + return prompts, extra_data for prompt in prompts: updated_prompt, parsed_extra_data = parse_prompt(prompt) diff --git a/modules/gr_hijack.py b/modules/gr_hijack.py index b4607cc11..7fa1cecdd 100644 --- a/modules/gr_hijack.py +++ b/modules/gr_hijack.py @@ -1,3 +1,4 @@ +import time from PIL import Image import gradio as gr import gradio.processing_utils @@ -11,13 +12,49 @@ original_BlockContext_init = None original_Blocks_get_config_file = None +def process_kanvas(self, x): # only used when kanvas overrides gr.Image object + import numpy as np + from modules import errors + t0 = time.time() + image_data = list(x.get('image', {}).values()) + image = None + mask = None + if image_data: + width = x['imageWidth'] + height = x['imageHeight'] + array = np.array(image_data, dtype=np.uint8).reshape((height, width, 4)) + image = Image.fromarray(array, 'RGBA') + image = image.convert('RGB') + mask_data = list(x.get('mask', {}).values()) + if mask_data: + width = x['maskWidth'] + height = x['maskHeight'] + array = np.array(mask_data, dtype=np.uint8).reshape((height, width, 4)) + mask = Image.fromarray(array, 'RGBA') + # alpha = mask.getchannel("A").convert("L") + # mask = Image.merge("RGB", [alpha, alpha, alpha]) + mask = mask.convert('L') + t1 = time.time() + errors.log.debug(f'Kanvas: image={image} mask={mask} time={t1-t0:.2f}') + if image is None: + return None + if mask is None: + return self._format_image(image) # pylint: disable=protected-access + return { "image": self._format_image(image), "mask": self._format_image(mask) } # pylint: disable=protected-access + + def gr_image_preprocess(self, x): if x is None: return x mask = None - if isinstance(x, dict): + if isinstance(x, dict) and "kanvas" in x: + return process_kanvas(self, x) + if isinstance(x, dict) and "image" in x: x, mask = x["image"], x["mask"] - im = gradio.processing_utils.decode_base64_to_image(x) + if isinstance(x, str): + im = gradio.processing_utils.decode_base64_to_image(x) + else: + im = x im = im.convert(self.image_mode) if self.shape is not None: im = gradio.processing_utils.resize_and_crop(im, self.shape) diff --git a/modules/images_resize.py b/modules/images_resize.py index 03646484e..bbe25ba68 100644 --- a/modules/images_resize.py +++ b/modules/images_resize.py @@ -95,6 +95,8 @@ def resize_image(resize_mode: int, im: Union[Image.Image, torch.Tensor], width: return res def context_aware(im: Image.Image, width, height, context): + from installer import install + install('seam-carving') width, height = int(width), int(height) import seam_carving # https://github.com/li-plus/seam-carving if 'forward' in context.lower(): @@ -116,13 +118,14 @@ def resize_image(resize_mode: int, im: Union[Image.Image, torch.Tensor], width: src_image = resize(im, src_w, src_h) else: return im - res = Image.fromarray(seam_carving.resize( + np_image = seam_carving.resize( src_image, # source image (rgb or gray) size=(width, height), # target size energy_mode=energy_mode, # choose from {backward, forward} order="width-first", # choose from {width-first, height-first} keep_mask=None, # object mask to protect from removal - )) + ) + res = Image.fromarray(np_image) return res t0 = time.time() @@ -154,5 +157,6 @@ def resize_image(resize_mode: int, im: Union[Image.Image, torch.Tensor], width: shared.log.error(f'Invalid resize mode: {resize_mode}') t1 = time.time() fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access - shared.log.debug(f'Image resize: source={im.width}:{im.height} target={width}:{height} mode="{shared.resize_modes[resize_mode]}" upscaler="{upscaler_name}" type={output_type} time={t1-t0:.2f} fn={fn}') # pylint: disable=protected-access + if im.width != width or im.height != height: + shared.log.debug(f'Image resize: source={im.width}:{im.height} target={width}:{height} mode="{shared.resize_modes[resize_mode]}" upscaler="{upscaler_name}" type={output_type} time={t1-t0:.2f} fn={fn}') # pylint: disable=protected-access return np.array(res) if output_type == 'np' else res diff --git a/modules/intel/ipex/hijacks.py b/modules/intel/ipex/hijacks.py index 53a0d661d..56999da21 100644 --- a/modules/intel/ipex/hijacks.py +++ b/modules/intel/ipex/hijacks.py @@ -3,8 +3,8 @@ from functools import wraps from contextlib import nullcontext import torch import numpy as np -from modules import devices +from modules import devices from .device_prop import cache_size_dict torch_version = torch.__version__[:4] @@ -42,9 +42,7 @@ def return_xpu(device): # keep the device instance type, aka return string if th original_autocast_init = torch.amp.autocast_mode.autocast.__init__ @wraps(torch.amp.autocast_mode.autocast.__init__) def autocast_init(self, device_type=None, dtype=None, enabled=True, cache_enabled=None): - if device_type is None or check_cuda(device_type) or check_device_type(device_type, "xpu"): - if dtype is None: - dtype = devices.dtype + if device_type is None or check_cuda(device_type): return original_autocast_init(self, device_type="xpu", dtype=dtype, enabled=enabled, cache_enabled=cache_enabled) else: return original_autocast_init(self, device_type=device_type, dtype=dtype, enabled=enabled, cache_enabled=cache_enabled) @@ -72,7 +70,7 @@ original_get_autocast_dtype = torch.get_autocast_dtype @wraps(torch.get_autocast_dtype) def torch_get_autocast_dtype(device_type=None): if device_type is None or check_cuda(device_type) or check_device_type(device_type, "xpu"): - return devices.dtype + return devices.dtype or torch.bfloat16 else: return original_get_autocast_dtype(device_type) diff --git a/modules/intel/openvino/__init__.py b/modules/intel/openvino/__init__.py index 8d787d0e0..e14b26ec0 100644 --- a/modules/intel/openvino/__init__.py +++ b/modules/intel/openvino/__init__.py @@ -73,6 +73,13 @@ if hasattr(torch, "float8_e8m0fnu"): dtype_mapping[torch.float8_e8m0fnu] = Type.f8e8m0 +warned = False +def warn_once(msg): + global warned + if not warned: + shared.log.warning(msg) + warned = True + class OpenVINOGraphModule(torch.nn.Module): def __init__(self, gm, partition_id, use_python_fusion_cache, model_hash_str: str = None, file_name="", int_inputs=[]): super().__init__() @@ -128,7 +135,7 @@ def get_device(): device = "GPU.0" else: device = core.available_devices[-1] - shared.log.warning(f"OpenVINO: No compatible GPU detected! Using {device}") + warn_once(f"OpenVINO: device={device} no compatible GPU detected") return device diff --git a/modules/interrogate/vqa.py b/modules/interrogate/vqa.py index 877740a76..7c93c56f3 100644 --- a/modules/interrogate/vqa.py +++ b/modules/interrogate/vqa.py @@ -178,7 +178,7 @@ def fastvlm(question: str, image: Image.Image, repo: str = None): def qwen(question: str, image: Image.Image, repo: str = None, system_prompt: str = None): global processor, model, loaded # pylint: disable=global-statement - if model is None or loaded != repo: + if (model is None) or (loaded != repo): shared.log.debug(f'Interrogate load: vlm="{repo}"') model = None if 'Qwen3-VL' in repo or 'Qwen3VL' in repo: @@ -633,8 +633,7 @@ def interrogate(question:str='', system_prompt:str=None, prompt:str=None, image: global quant_args # pylint: disable=global-statement jobid = shared.state.begin('Interrogate LLM') t0 = time.time() - if quant_args is None: - quant_args = model_quant.create_config(module='LLM') + quant_args = model_quant.create_config(module='LLM') model_name = model_name or shared.opts.interrogate_vlm_model if isinstance(image, list): image = image[0] if len(image) > 0 else None diff --git a/modules/linfusion/__init__.py b/modules/linfusion/__init__.py index 34495936b..01b894af9 100644 --- a/modules/linfusion/__init__.py +++ b/modules/linfusion/__init__.py @@ -1,4 +1,4 @@ -from modules import shared, sd_models, devices +from modules import shared, sd_models, devices, attention from .linfusion import LinFusion from .attention import GeneralizedLinearAttention @@ -41,6 +41,6 @@ def unapply(pipeline): if applied is None: return # shared.log.debug('LinFusion: unapply') - sd_models.set_diffusers_attention(pipeline) + attention.set_diffusers_attention(pipeline) devices.torch_gc() applied = None diff --git a/modules/lora/extra_networks_lora.py b/modules/lora/extra_networks_lora.py index 99ab20d15..1c73712cd 100644 --- a/modules/lora/extra_networks_lora.py +++ b/modules/lora/extra_networks_lora.py @@ -226,8 +226,7 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): 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"}') networks.network_deactivate(include, exclude) networks.network_activate(include, exclude) - if len(exclude) > 0: # only update on last activation - l.previously_loaded_networks = l.loaded_networks.copy() + 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') shared.state.end(jobid) @@ -235,12 +234,14 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): 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_diffusers else "backup"} te={te_multipliers} unet={unet_multipliers} time={l.timer.summary}') + 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}') - def deactivate(self, p): + def deactivate(self, p, force=False): if len(lora_diffusers.diffuser_loaded) > 0: if not (shared.compiled_model_state is not None and shared.compiled_model_state.is_compiled is True): unload_diffusers() + if force: + networks.network_deactivate() if self.active and l.debug: shared.log.debug(f"Network end: type=LoRA time={l.timer.summary}") if self.errors: diff --git a/modules/lora/lora_apply.py b/modules/lora/lora_apply.py index f18922040..ca78006e0 100644 --- a/modules/lora/lora_apply.py +++ b/modules/lora/lora_apply.py @@ -20,7 +20,7 @@ def network_backup_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.n weights_backup = getattr(self, "network_weights_backup", None) bias_backup = getattr(self, "network_bias_backup", None) if weights_backup is not None or bias_backup is not None: - if (shared.opts.lora_fuse_diffusers and not isinstance(weights_backup, bool)) or (not shared.opts.lora_fuse_diffusers and isinstance(weights_backup, bool)): # invalidate so we can change direct/backup on-the-fly + if (shared.opts.lora_fuse_native and not isinstance(weights_backup, bool)) or (not shared.opts.lora_fuse_native and isinstance(weights_backup, bool)): # invalidate so we can change direct/backup on-the-fly weights_backup = None bias_backup = None self.network_weights_backup = weights_backup @@ -33,15 +33,15 @@ def network_backup_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.n if bnb is None: bnb = model_quant.load_bnb('Network load: type=LoRA', silent=True) if bnb is not None: - if shared.opts.lora_fuse_diffusers: + if shared.opts.lora_fuse_native: self.network_weights_backup = True else: self.network_weights_backup = bnb.functional.dequantize_4bit(weight, quant_state=weight.quant_state, quant_type=weight.quant_type, blocksize=weight.blocksize,) self.quant_state, self.quant_type, self.blocksize = weight.quant_state, weight.quant_type, weight.blocksize else: - self.network_weights_backup = weight.clone().to(devices.cpu) if not shared.opts.lora_fuse_diffusers else True + self.network_weights_backup = weight.clone().to(devices.cpu) if not shared.opts.lora_fuse_native else True else: - if shared.opts.lora_fuse_diffusers: + if shared.opts.lora_fuse_native: self.network_weights_backup = True else: self.network_weights_backup = weight.clone().to(devices.cpu) @@ -61,7 +61,7 @@ def network_backup_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.n if bias_backup is None: if getattr(self, 'bias', None) is not None: - if shared.opts.lora_fuse_diffusers: + if shared.opts.lora_fuse_native: self.network_bias_backup = True else: bias_backup = self.bias.clone() @@ -167,23 +167,27 @@ def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.G try: from modules.sdnq import sdnq_quantize_layer if hasattr(self, "sdnq_dequantizer_backup"): - weights_dtype = self.sdnq_dequantizer_backup.weights_dtype + use_svd = bool(self.sdnq_svd_up_backup is not None) + dequantize_fp32 = bool(self.sdnq_scale_backup.dtype == torch.float32) + sdnq_dequantizer = self.sdnq_dequantizer_backup dequant_weight = self.sdnq_dequantizer_backup.to(devices.device)( model_weights.to(devices.device), self.sdnq_scale_backup.to(devices.device), self.sdnq_zero_point_backup.to(devices.device) if self.sdnq_zero_point_backup is not None else None, - self.sdnq_svd_up_backup.to(devices.device) if self.sdnq_svd_up_backup is not None else None, - self.sdnq_svd_down_backup.to(devices.device) if self.sdnq_svd_down_backup is not None else None, + self.sdnq_svd_up_backup.to(devices.device) if use_svd else None, + self.sdnq_svd_down_backup.to(devices.device) if use_svd else None, skip_quantized_matmul=self.sdnq_dequantizer_backup.use_quantized_matmul ) else: - weights_dtype = self.sdnq_dequantizer.weights_dtype + use_svd = bool(self.svd_up is not None) + dequantize_fp32 = bool(self.scale.dtype == torch.float32) + sdnq_dequantizer = self.sdnq_dequantizer dequant_weight = self.sdnq_dequantizer.to(devices.device)( model_weights.to(devices.device), self.scale.to(devices.device), self.zero_point.to(devices.device) if self.zero_point is not None else None, - self.svd_up.to(devices.device) if self.svd_up is not None else None, - self.svd_down.to(devices.device) if self.svd_down is not None else None, + self.svd_up.to(devices.device) if use_svd else None, + self.svd_down.to(devices.device) if use_svd else None, skip_quantized_matmul=self.sdnq_dequantizer.use_quantized_matmul ) @@ -192,16 +196,16 @@ def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.G del self.sdnq_dequantizer, self.scale, self.zero_point, self.svd_up, self.svd_down self = sdnq_quantize_layer( self, - weights_dtype=weights_dtype, - torch_dtype=devices.dtype, - group_size=shared.opts.sdnq_quantize_weights_group_size, - svd_rank=shared.opts.sdnq_svd_rank, + weights_dtype=sdnq_dequantizer.weights_dtype, + torch_dtype=sdnq_dequantizer.result_dtype, + group_size=sdnq_dequantizer.group_size, + svd_rank=sdnq_dequantizer.svd_rank, + use_quantized_matmul=sdnq_dequantizer.use_quantized_matmul, + use_quantized_matmul_conv=sdnq_dequantizer.use_quantized_matmul, + use_svd=use_svd, + dequantize_fp32=dequantize_fp32, svd_steps=shared.opts.sdnq_svd_steps, - 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, - dequantize_fp32=shared.opts.sdnq_dequantize_fp32, + quant_conv=True, # quant_conv is True if conv layers ends up here non_blocking=False, quantization_device=devices.device, return_device=device, diff --git a/modules/lora/lora_diffusers.py b/modules/lora/lora_diffusers.py index a97272370..eb1515ca0 100644 --- a/modules/lora/lora_diffusers.py +++ b/modules/lora/lora_diffusers.py @@ -54,7 +54,7 @@ def load_diffusers(name: str, network_on_disk: network.NetworkOnDisk, lora_scale t0 = time.time() name = name.replace(".", "_") sd_model: diffusers.DiffusionPipeline = getattr(shared.sd_model, "pipe", shared.sd_model) - shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" detected={network_on_disk.sd_version} method=diffusers scale={lora_scale} fuse={shared.opts.lora_fuse_diffusers}') + shared.log.debug(f'Network load: type=LoRA name="{name}" file="{network_on_disk.filename}" detected={network_on_disk.sd_version} method=diffusers scale={lora_scale} fuse={shared.opts.lora_fuse_native}:{shared.opts.lora_fuse_diffusers}') if not hasattr(sd_model, 'load_lora_weights'): shared.log.error(f'Network load: type=LoRA class={sd_model.__class__} does not implement load lora') return None diff --git a/modules/lora/lora_load.py b/modules/lora/lora_load.py index 85c66208d..2a54707f7 100644 --- a/modules/lora/lora_load.py +++ b/modules/lora/lora_load.py @@ -128,7 +128,7 @@ def load_safetensors(name, network_on_disk: network.NetworkOnDisk) -> Union[netw if l.debug: shared.log.debug(f'Network load: type=LoRA name="{name}" unmatched={keys_failed_to_match}') else: - shared.log.debug(f'Network load: type=LoRA name="{name}" type={set(network_types)} keys={len(matched_networks)} dtypes={dtypes} fuse={shared.opts.lora_fuse_diffusers}') + shared.log.debug(f'Network load: type=LoRA name="{name}" type={set(network_types)} keys={len(matched_networks)} dtypes={dtypes} fuse={shared.opts.lora_fuse_native}:{shared.opts.lora_fuse_diffusers}') if len(matched_networks) == 0: return None lora_cache[name] = net @@ -157,11 +157,14 @@ def maybe_recompile_model(names, te_multipliers): recompile_model = True shared.compiled_model_state.lora_model = [] if recompile_model: + current_task = sd_models.get_diffusers_task(shared.sd_model) + shared.log.debug(f'Compile: task={current_task} force model reload') backup_cuda_compile = shared.opts.cuda_compile backup_scheduler = getattr(sd_model, "scheduler", None) sd_models.unload_model_weights(op='model') shared.opts.cuda_compile = [] sd_models.reload_model_weights(op='model') + shared.sd_model = sd_models.set_diffuser_pipe(shared.sd_model, current_task) shared.opts.cuda_compile = backup_cuda_compile if backup_scheduler is not None: sd_model.scheduler = backup_scheduler @@ -247,7 +250,7 @@ def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=Non try: lora_scale = te_multipliers[i] if te_multipliers else shared.opts.extra_networks_default_multiplier lora_module = lora_modules[i] if lora_modules and len(lora_modules) > i else None - if recompile_model: + if recompile_model and shared.compiled_model_state is not None: shared.compiled_model_state.lora_model.append(f"{name}:{lora_scale}") lora_method = lora_overrides.get_method(shorthash) if lora_method == 'diffusers': @@ -303,13 +306,17 @@ def network_load(names, te_multipliers=None, unet_multipliers=None, dyn_dims=Non errors.display(e, 'LoRA') 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)}') + 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}') if recompile_model: shared.log.info("Network load: type=LoRA recompiling model") - backup_lora_model = shared.compiled_model_state.lora_model + if shared.compiled_model_state is not None: + backup_lora_model = shared.compiled_model_state.lora_model + else: + backup_lora_model = [] if 'Model' in shared.opts.cuda_compile: sd_model = sd_models_compile.compile_diffusers(sd_model) - shared.compiled_model_state.lora_model = backup_lora_model + if shared.compiled_model_state is not None: + shared.compiled_model_state.lora_model = backup_lora_model l.timer.load = time.time() - t0 diff --git a/modules/lora/networks.py b/modules/lora/networks.py index f03063e2d..4294615c9 100644 --- a/modules/lora/networks.py +++ b/modules/lora/networks.py @@ -49,7 +49,7 @@ def network_activate(include=[], exclude=[]): continue backup_size += network_backup_weights(module, network_layer_name, wanted_names) batch_updown, batch_ex_bias = network_calc_weights(module, network_layer_name) - if shared.opts.lora_fuse_diffusers: + if shared.opts.lora_fuse_native: network_apply_direct(module, batch_updown, batch_ex_bias, device=device) else: network_apply_weights(module, batch_updown, batch_ex_bias, device=device) @@ -68,14 +68,14 @@ def network_activate(include=[], exclude=[]): pbar.remove_task(task) # hide progress bar for no action l.timer.activate += time.time() - t0 if l.debug and len(l.loaded_networks) > 0: - shared.log.debug(f'Network load: type=LoRA networks={[n.name for n in l.loaded_networks]} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={round(backup_size/1024/1024/1024, 2)} fuse={shared.opts.lora_fuse_diffusers} device={device} time={l.timer.summary}') + shared.log.debug(f'Network load: type=LoRA networks={[n.name for n in l.loaded_networks]} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={round(backup_size/1024/1024/1024, 2)} fuse={shared.opts.lora_fuse_native}:{shared.opts.lora_fuse_diffusers} device={device} time={l.timer.summary}') modules.clear() if len(applied_layers) > 0 or shared.opts.diffusers_offload_mode == "sequential": sd_models.set_diffuser_offload(sd_model, op="model") def network_deactivate(include=[], exclude=[]): - if not shared.opts.lora_fuse_diffusers or shared.opts.lora_force_diffusers: + if not shared.opts.lora_fuse_native or shared.opts.lora_force_diffusers: return if len(l.previously_loaded_networks) == 0: return @@ -112,7 +112,7 @@ def network_deactivate(include=[], exclude=[]): pbar.update(task, advance=1) continue batch_updown, batch_ex_bias = network_calc_weights(module, network_layer_name, use_previous=True) - if shared.opts.lora_fuse_diffusers: + if shared.opts.lora_fuse_native: network_apply_direct(module, batch_updown, batch_ex_bias, device=device, deactivate=True) else: network_apply_weights(module, batch_updown, batch_ex_bias, device=device, deactivate=True) @@ -125,7 +125,7 @@ def network_deactivate(include=[], exclude=[]): l.timer.deactivate = time.time() - t0 if l.debug and len(l.previously_loaded_networks) > 0: - shared.log.debug(f'Network deactivate: type=LoRA networks={[n.name for n in l.previously_loaded_networks]} modules={active_components} layers={total} apply={len(applied_layers)} fuse={shared.opts.lora_fuse_diffusers} time={l.timer.summary}') + shared.log.debug(f'Network deactivate: type=LoRA networks={[n.name for n in l.previously_loaded_networks]} modules={active_components} layers={total} apply={len(applied_layers)} fuse={shared.opts.lora_fuse_native}:{shared.opts.lora_fuse_diffusers} time={l.timer.summary}') modules.clear() if len(applied_layers) > 0 or shared.opts.diffusers_offload_mode == "sequential": sd_models.set_diffuser_offload(sd_model, op="model") diff --git a/modules/model_quant.py b/modules/model_quant.py index 93f3accc4..e844a0ccf 100644 --- a/modules/model_quant.py +++ b/modules/model_quant.py @@ -246,10 +246,10 @@ def check_nunchaku(module: str = ''): def create_config(kwargs = None, allow: bool = True, module: str = 'Model', modules_to_not_convert: list = None, modules_dtype_dict: dict = None): - if dont_quant(): - return kwargs if kwargs is None: kwargs = {} + if module == 'Model' and dont_quant(): + return kwargs kwargs = create_sdnq_config(kwargs, allow=allow, module=module, modules_to_not_convert=modules_to_not_convert, modules_dtype_dict=modules_dtype_dict) if kwargs is not None and 'quantization_config' in kwargs: if debug: diff --git a/modules/modeldata.py b/modules/modeldata.py index 5ec225816..6b0fb89da 100644 --- a/modules/modeldata.py +++ b/modules/modeldata.py @@ -14,6 +14,8 @@ def get_model_type(pipe): model_type = 'sdxl' elif "StableDiffusion" in name: model_type = 'sd' + elif "StableVideoDiffusion" in name: + model_type = 'svd' elif "LatentConsistencyModel" in name: model_type = 'sd' # lcm is compatible with sd elif "InstaFlowPipeline" in name: @@ -64,6 +66,8 @@ def get_model_type(pipe): model_type = 'nextstep' elif 'X-Omni' in name: model_type = 'x-omni' + elif 'Photoroom' in name: + model_type = 'prx' # video models elif "CogVideo" in name: model_type = 'cogvideo' @@ -86,6 +90,9 @@ def get_model_type(pipe): model_type = 'hunyuanimage3' elif 'HunyuanImage' in name: model_type = 'hunyuanimage' + # cloud models + elif 'NanoBanana' in name: + model_type = 'nanobanana' else: model_type = name return model_type diff --git a/modules/postprocess/yolo.py b/modules/postprocess/yolo.py index e6ceb1794..8098b1a7e 100644 --- a/modules/postprocess/yolo.py +++ b/modules/postprocess/yolo.py @@ -6,7 +6,7 @@ from copy import copy import numpy as np import gradio as gr from PIL import Image, ImageDraw -from modules import shared, processing, devices, processing_class, ui_common, ui_components, ui_symbols, images +from modules import shared, processing, devices, processing_class, ui_common, ui_components, ui_symbols, images, extra_networks, sd_models from modules.detailer import Detailer @@ -239,6 +239,12 @@ class YoloRestorer(Detailer): p.detailer_active = 0 if np_image is None or p.detailer_active >= p.batch_size * p.n_iter: return np_image + + shared.sd_model = sd_models.set_diffuser_pipe(shared.sd_model, sd_models.DiffusersTaskType.INPAINTING) + if (sd_models.get_diffusers_task(shared.sd_model) != sd_models.DiffusersTaskType.INPAINTING) and (shared.sd_model.__class__.__name__ not in sd_models.pipe_switch_task_exclude): + shared.log.error(f'Detailer: model="{shared.sd_model.__class__.__name__}" not compatible') + return np_image + models = [] if len(shared.opts.detailer_args) > 0: models = [m.strip() for m in re.split(r'[\n,;]+', shared.opts.detailer_args)] @@ -257,6 +263,7 @@ class YoloRestorer(Detailer): models_used = [] np_images = [] annotated = Image.fromarray(np_image) + image = None for i, model_val in enumerate(models): if ':' in model_val: @@ -271,9 +278,10 @@ class YoloRestorer(Detailer): shared.log.warning(f'Detailer: model="{name}" not loaded') continue - if name.endswith('.fp16'): + if name.endswith('.fp16'): # run gfpgan or codeformer directly and skip detailer processing from modules.postprocess import restorer np_image = restorer.restore(np_image, name, model, p.detailer_strength) + image = Image.fromarray(np_image) continue image = Image.fromarray(np_image) @@ -302,10 +310,8 @@ class YoloRestorer(Detailer): else: negative = negative.replace('[PROMPT]', orig_negative) negative = negative.replace('[prompt]', orig_negative) - prompt_lines = prompt.split('\n') - negative_lines = negative.split('\n') - prompt = prompt_lines[i % len(prompt_lines)] - negative = negative_lines[i % len(negative_lines)] + prompt_lines = 99 * [p.strip() for p in prompt.split('\n')] + negative_lines = 99 * [n.strip() for n in negative.split('\n')] args = { 'detailer': True, @@ -367,16 +373,25 @@ class YoloRestorer(Detailer): if item.mask is None: continue pc.keep_prompts = True - pc.prompts = [prompt_lines[(i*len(items)+j) % len(prompt_lines)]] - pc.negative_prompts = [negative_lines[(i*len(items)+j) % len(negative_lines)]] + shared.sd_model.fail_on_switch_error = True + pc.prompt = prompt_lines[i*len(items)+j] + pc.negative_prompt = negative_lines[i*len(items)+j] + pc.prompts = [pc.prompt] + pc.negative_prompts = [pc.negative_prompt] + pc.prompts, pc.network_data = extra_networks.parse_prompts(pc.prompts) + extra_networks.activate(pc, pc.network_data) shared.log.debug(f'Detail: model="{i+1}:{name}" item={j+1}/{len(items)} box={item.box} label="{item.label} score={item.score:.2f} prompt="{pc.prompt}"') pc.init_images = [image] pc.image_mask = [item.mask] pc.overlay_images = [] pc.recursion = True + jobid = shared.state.begin('Detailer') pp = processing.process_images_inner(pc) + extra_networks.deactivate(pc, force=True) + shared.sd_model.fail_on_switch_error = False shared.state.end(jobid) + del pc.recursion if pp is not None and pp.images is not None and len(pp.images) > 0: image = pp.images[0] # update image to be reused for next item @@ -403,7 +418,8 @@ class YoloRestorer(Detailer): p.image_mask = blend([np.array(m) for m in mask_all]) p.image_mask = Image.fromarray(p.image_mask) - np_images.append(np.array(image)) + if image is not None: + np_images.append(np.array(image)) if shared.opts.detailer_save and annotated is not None: np_images.append(annotated) # save debug image with boxes return np_images diff --git a/modules/postprocessing.py b/modules/postprocessing.py index 4f919ea00..296db6eda 100644 --- a/modules/postprocessing.py +++ b/modules/postprocessing.py @@ -77,6 +77,10 @@ def run_postprocessing(extras_mode, image, image_folder: List[tempfile.NamedTemp pp.image.info[k] = v if 'parameters' in items: info = items['parameters'] + ', ' + if (params.get('size-1', 0) != pp.image.width) or (params.get('size-2', 0) != pp.image.height): + params['size-1'] = pp.image.width + params['size-2'] = pp.image.height + info += f"Size: {pp.image.width}x{pp.image.height}, " info = info + ", ".join([k if k == v else f'{k}: {infotext.quote(v)}' for k, v in pp.info.items() if v is not None]) pp.image.info["postprocessing"] = info processed_images.append(pp.image) diff --git a/modules/processing.py b/modules/processing.py index 22b852fd0..c363a58fa 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -311,7 +311,7 @@ def process_samples(p: StableDiffusionProcessing, samples): if len(sample) > 0: image = Image.fromarray(sample[0]) if len(sample) > 1: - annotated = Image.fromarray(sample[1]) + annotated = sample[1] if isinstance(sample[1], Image.Image) else Image.fromarray(sample[1]) out_images.append(annotated) out_infotexts.append("Detailer annotations") elif sample is not None: diff --git a/modules/processing_args.py b/modules/processing_args.py index 08f85fc94..74bfce1fd 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -145,6 +145,8 @@ def task_specific_kwargs(p, model): task_args['image'] = Image.new('RGB', (p.width, p.height), (0, 0, 0)) # monkey-patch so wan-i2i pipeline does not error-out on t2i if ('WanVACEPipeline' in model_cls) and (p.init_images is not None) and (len(p.init_images) > 0): 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 'BlipDiffusionPipeline' in model_cls: if len(p.init_images) == 0: shared.log.error('BLiP diffusion requires init image') @@ -391,6 +393,11 @@ def set_pipeline_args(p, model, prompts:list, negative_prompts:list, prompts_2:t continue args[arg] = kwargs[arg] + # optional preprocess + if hasattr(model, 'preprocess') and callable(model.preprocess): + model.preprocess(p, args) + + # handle task specific args if sd_models.get_diffusers_task(model) == sd_models.DiffusersTaskType.MODULAR: task_kwargs = task_modular_kwargs(p, model) diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index c2d852f25..1da33d732 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -5,7 +5,7 @@ import numpy as np import torch import torchvision.transforms.functional as TF from PIL import Image -from modules import shared, devices, processing, sd_models, errors, sd_hijack_hypertile, processing_vae, sd_models_compile, timer, modelstats, extra_networks +from modules import shared, devices, processing, sd_models, errors, sd_hijack_hypertile, processing_vae, sd_models_compile, timer, modelstats, extra_networks, attention from modules.processing_helpers import resize_hires, calculate_base_steps, calculate_hires_steps, calculate_refiner_steps, save_intermediate, update_sampler, is_txt2img, is_refiner_enabled, get_job_name from modules.processing_args import set_pipeline_args from modules.onnx_impl import preprocess_pipeline as preprocess_onnx_pipeline, check_parameters_changed as olive_check_parameters_changed @@ -497,7 +497,7 @@ def update_pipeline(sd_model, p: processing.StableDiffusionProcessing): orig_pipeline = sd_model # processed ONNX pipeline should not be replaced with original pipeline. if getattr(sd_model, "current_attn_name", None) != shared.opts.cross_attention_optimization: shared.log.info(f"Setting attention optimization: {shared.opts.cross_attention_optimization}") - sd_models.set_diffusers_attention(sd_model) + attention.set_diffusers_attention(sd_model) return sd_model @@ -543,6 +543,7 @@ def process_diffusers(p: processing.StableDiffusionProcessing): if len(getattr(p, 'init_images', [])) > 0: while len(p.init_images) < len(p.prompts): p.init_images.append(p.init_images[-1]) + # pipeline type is set earlier in processing, but check for sanity is_control = getattr(p, 'is_control', False) is True has_images = len(getattr(p, 'init_images', [])) > 0 diff --git a/modules/processing_info.py b/modules/processing_info.py index 129595408..c2a6566ec 100644 --- a/modules/processing_info.py +++ b/modules/processing_info.py @@ -105,8 +105,8 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No if p.hr_force or ('Latent' in p.hr_upscaler): args["Hires force"] = p.hr_force args["Hires steps"] = p.hr_second_pass_steps - args["Hires strength"] = p.denoising_strength - args["Hires sampler"] = p.hr_sampler_name if p.hr_sampler_name != p.sampler_name else None + args["Hires strength"] = p.hr_denoising_strength + args["Hires sampler"] = p.hr_sampler_name args["Hires CFG scale"] = p.image_cfg_scale if 'refine' in p.ops: args["Refine"] = p.enable_hr diff --git a/modules/sd_detect.py b/modules/sd_detect.py index 2e169ae11..9043edc5c 100644 --- a/modules/sd_detect.py +++ b/modules/sd_detect.py @@ -127,6 +127,12 @@ def guess_by_name(fn, current_guess): new_guess = 'X-Omni' elif 'sdxl-turbo' in fn.lower() or 'stable-diffusion-xl' in fn.lower(): new_guess = 'Stable Diffusion XL' + elif 'stable-video-diffusion' in fn.lower(): + new_guess = 'StableVideoDiffusion' + elif 'prx-' in fn.lower(): + new_guess = 'PRX' + elif 'gemini-2.5-flash-image' in fn.lower(): + new_guess = 'NanoBanana' if debug_load: shared.log.trace(f'Autodetect: method=name file="{fn}" previous="{current_guess}" current="{new_guess}"') return new_guess or current_guess @@ -156,11 +162,22 @@ def guess_by_diffusers(fn, current_guess): if folder.endswith('quantization_config.json'): is_quant = True break + if folder.endswith('config.json'): + quantization_config = shared.readfile(folder, silent=True).get("quantization_config", None) + if quantization_config is not None: + is_quant = True + break if os.path.isdir(folder): for f in os.listdir(folder): + f = os.path.join(folder, f) if f.endswith('quantization_config.json'): is_quant = True break + if f.endswith('config.json'): + quantization_config = shared.readfile(f, silent=True).get("quantization_config", None) + if quantization_config is not None: + is_quant = True + break pipelines = shared_items.get_pipelines() for k, v in pipelines.items(): if v is not None and v.__name__ == pipeline.__name__: diff --git a/modules/sd_models.py b/modules/sd_models.py index bd52deada..129cbb211 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -10,7 +10,7 @@ import diffusers.loaders.single_file_utils import torch import huggingface_hub as hf from installer import log -from modules import timer, paths, shared, shared_items, modelloader, devices, script_callbacks, sd_vae, sd_unet, errors, sd_models_compile, sd_detect, model_quant, sd_hijack_te, sd_hijack_accelerate, sd_hijack_safetensors +from modules import timer, paths, shared, shared_items, modelloader, devices, script_callbacks, sd_vae, sd_unet, errors, sd_models_compile, sd_detect, model_quant, sd_hijack_te, sd_hijack_accelerate, sd_hijack_safetensors, attention from modules.memstats import memory_stats from modules.modeldata import model_data from modules.sd_checkpoint import CheckpointInfo, select_checkpoint, list_models, checkpoints_list, checkpoint_titles, get_closest_checkpoint_match, model_hash, update_model_hashes, setup_model, write_metadata, read_metadata_from_safetensors # pylint: disable=unused-import @@ -49,6 +49,7 @@ pipe_switch_task_exclude = [ 'HunyuanImagePipeline', 'AuraFlowPipeline', 'ChronoEditPipeline', + 'GoogleNanoBananaPipeline', ] i2i_pipes = [ 'LEditsPPPipelineStableDiffusion', 'LEditsPPPipelineStableDiffusionXL', @@ -84,33 +85,39 @@ def set_vae_options(sd_model, vae=None, op:str='model', quiet:bool=False): ops['no-half'] = True if hasattr(sd_model, 'vae') and hasattr(sd_model.vae, 'enable_slicing') and hasattr(sd_model.vae, 'disable_slicing'): ops['slicing'] = shared.opts.diffusers_vae_slicing - if shared.opts.diffusers_vae_slicing: - sd_model.vae.enable_slicing() - else: - sd_model.vae.disable_slicing() + try: + if shared.opts.diffusers_vae_slicing: + sd_model.vae.enable_slicing() + else: + sd_model.vae.disable_slicing() + except Exception: + pass if hasattr(sd_model, 'vae') and hasattr(sd_model.vae, 'enable_tiling') and hasattr(sd_model.vae, 'disable_tiling'): ops['tiling'] = shared.opts.diffusers_vae_tiling - if shared.opts.diffusers_vae_tiling: - if hasattr(sd_model, 'vae') and hasattr(sd_model.vae, 'config') and hasattr(sd_model.vae.config, 'sample_size') and isinstance(sd_model.vae.config.sample_size, int): - if getattr(sd_model.vae, "tile_sample_min_size_backup", None) is None: - sd_model.vae.tile_sample_min_size_backup = sd_model.vae.tile_sample_min_size - sd_model.vae.tile_latent_min_size_backup = sd_model.vae.tile_latent_min_size - sd_model.vae.tile_overlap_factor_backup = sd_model.vae.tile_overlap_factor - if shared.opts.diffusers_vae_tile_size > 0: - sd_model.vae.tile_sample_min_size = int(shared.opts.diffusers_vae_tile_size) - sd_model.vae.tile_latent_min_size = int(shared.opts.diffusers_vae_tile_size / (2 ** (len(sd_model.vae.config.block_out_channels) - 1))) - else: - sd_model.vae.tile_sample_min_size = getattr(sd_model.vae, "tile_sample_min_size_backup", sd_model.vae.tile_sample_min_size) - sd_model.vae.tile_latent_min_size = getattr(sd_model.vae, "tile_latent_min_size_backup", sd_model.vae.tile_latent_min_size) - if shared.opts.diffusers_vae_tile_overlap != 0.25: - sd_model.vae.tile_overlap_factor = float(shared.opts.diffusers_vae_tile_overlap) - else: - sd_model.vae.tile_overlap_factor = getattr(sd_model.vae, "tile_overlap_factor_backup", sd_model.vae.tile_overlap_factor) - ops['tile'] = sd_model.vae.tile_sample_min_size - ops['overlap'] = sd_model.vae.tile_overlap_factor - sd_model.vae.enable_tiling() - else: - sd_model.vae.disable_tiling() + try: + if shared.opts.diffusers_vae_tiling: + if hasattr(sd_model, 'vae') and hasattr(sd_model.vae, 'config') and hasattr(sd_model.vae.config, 'sample_size') and isinstance(sd_model.vae.config.sample_size, int): + if getattr(sd_model.vae, "tile_sample_min_size_backup", None) is None: + sd_model.vae.tile_sample_min_size_backup = sd_model.vae.tile_sample_min_size + sd_model.vae.tile_latent_min_size_backup = sd_model.vae.tile_latent_min_size + sd_model.vae.tile_overlap_factor_backup = sd_model.vae.tile_overlap_factor + if shared.opts.diffusers_vae_tile_size > 0: + sd_model.vae.tile_sample_min_size = int(shared.opts.diffusers_vae_tile_size) + sd_model.vae.tile_latent_min_size = int(shared.opts.diffusers_vae_tile_size / (2 ** (len(sd_model.vae.config.block_out_channels) - 1))) + else: + sd_model.vae.tile_sample_min_size = getattr(sd_model.vae, "tile_sample_min_size_backup", sd_model.vae.tile_sample_min_size) + sd_model.vae.tile_latent_min_size = getattr(sd_model.vae, "tile_latent_min_size_backup", sd_model.vae.tile_latent_min_size) + if shared.opts.diffusers_vae_tile_overlap != 0.25: + sd_model.vae.tile_overlap_factor = float(shared.opts.diffusers_vae_tile_overlap) + else: + sd_model.vae.tile_overlap_factor = getattr(sd_model.vae, "tile_overlap_factor_backup", sd_model.vae.tile_overlap_factor) + ops['tile'] = sd_model.vae.tile_sample_min_size + ops['overlap'] = sd_model.vae.tile_overlap_factor + sd_model.vae.enable_tiling() + else: + sd_model.vae.disable_tiling() + except Exception: + pass if hasattr(sd_model, "vqvae"): ops['upcast'] = True sd_model.vqvae.to(torch.float32) # vqvae is producing nans in fp16 @@ -130,7 +137,7 @@ def set_diffuser_options(sd_model, vae=None, op:str='model', offload:bool=True, clear_caches() set_vae_options(sd_model, vae, op, quiet) - set_diffusers_attention(sd_model, quiet) + attention.set_diffusers_attention(sd_model, quiet) if shared.opts.diffusers_fuse_projections and hasattr(sd_model, 'fuse_qkv_projections'): try: @@ -415,6 +422,14 @@ def load_diffuser_force(model_type, checkpoint_info, diffusers_load_config, op=' from pipelines.model_xomni import load_xomni sd_model = load_xomni(checkpoint_info, diffusers_load_config) # pylint: disable=assignment-from-none allow_post_quant = False + elif model_type in ['NanoBanana']: + from pipelines.model_google import load_nanobanana + sd_model = load_nanobanana(checkpoint_info, diffusers_load_config) + allow_post_quant = False + elif model_type in ['PRX']: + from pipelines.model_prx import load_prx + sd_model = load_prx(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: @@ -559,11 +574,16 @@ def load_diffuser_file(model_type, pipeline, checkpoint_info, diffusers_load_con def load_sdnq_module(fn: str, module_name: str, load_method: str): from modules import sdnq t0 = time.time() + quantization_config = None quantization_config_path = os.path.join(fn, module_name, 'quantization_config.json') - if not os.path.exists(quantization_config_path): + model_config_path = os.path.join(fn, module_name, 'config.json') + if os.path.exists(quantization_config_path): + quantization_config = shared.readfile(quantization_config_path, silent=True) + elif os.path.exists(model_config_path): + quantization_config = shared.readfile(model_config_path, silent=True).get("quantization_config", None) + if quantization_config is None: return None, module_name, 0 model_name = os.path.join(fn, module_name) - quantization_config = shared.readfile(quantization_config_path, silent=True) try: module = sdnq.load_sdnq_model( model_path=model_name, @@ -852,6 +872,7 @@ def load_diffuser(checkpoint_info=None, op='model', revision=None): # pylint: di modelstats.analyze() shared.log.info(f"Load {op}: family={shared.sd_model_type} time={timer.load.dct()} native={get_native(sd_model)} memory={memory_stats()}") + shared.opts.save(silent=True) class DiffusersTaskType(Enum): @@ -1083,51 +1104,50 @@ def set_diffuser_pipe(pipe, new_pipe_type): if 'Onnx' in cls: return pipe - new_pipe = None # in some cases we want to reset the pipeline to parent as they dont have their own variants - if new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE or new_pipe_type == DiffusersTaskType.INPAINTING: + if (new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE) or (new_pipe_type == DiffusersTaskType.INPAINTING): if cls == 'StableDiffusionPAGPipeline': pipe = switch_pipe(diffusers.StableDiffusionPipeline, pipe) if cls == 'StableDiffusionXLPAGPipeline': pipe = switch_pipe(diffusers.StableDiffusionXLPipeline, pipe) + new_pipe = None components_backup = backup_pipe_components(pipe) - if new_pipe is None: - if hasattr(pipe, 'config'): # real pipeline which can be auto-switched - try: - if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE: - new_pipe = diffusers.AutoPipelineForText2Image.from_pipe(pipe) - elif new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE: - new_pipe = diffusers.AutoPipelineForImage2Image.from_pipe(pipe) - elif new_pipe_type == DiffusersTaskType.INPAINTING: - new_pipe = diffusers.AutoPipelineForInpainting.from_pipe(pipe) - else: - shared.log.warning(f'Pipeline class change failed: type={new_pipe_type} pipeline={cls}') - return pipe - except Exception as e: # pylint: disable=unused-variable - fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access - shared.log.trace(f"Pipeline class change requested: target={new_pipe_type} fn={fn}") # pylint: disable=protected-access - shared.log.warning(f'Pipeline class change failed: type={new_pipe_type} pipeline={cls} {e}') - has_errors = True - if not hasattr(pipe, 'config') or has_errors: - try: # maybe a wrapper pipeline so just change the class - if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE: - pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING, cls) # pylint: disable=protected-access - new_pipe = pipe - elif new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE: - pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING, cls) # pylint: disable=protected-access - new_pipe = pipe - elif new_pipe_type == DiffusersTaskType.INPAINTING: - pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING, cls) # pylint: disable=protected-access - new_pipe = pipe - else: - shared.log.error(f'Pipeline class set failed: type={new_pipe_type} pipeline={cls}') - return pipe - except Exception as e: # pylint: disable=unused-variable - shared.log.warning(f'Pipeline class set failed: type={new_pipe_type} pipeline={cls} {e}') - has_errors = True + if hasattr(pipe, 'config'): # real pipeline which can be auto-switched + try: + if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE: + new_pipe = diffusers.AutoPipelineForText2Image.from_pipe(pipe) + elif new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE: + new_pipe = diffusers.AutoPipelineForImage2Image.from_pipe(pipe) + elif new_pipe_type == DiffusersTaskType.INPAINTING: + new_pipe = diffusers.AutoPipelineForInpainting.from_pipe(pipe) + else: + shared.log.warning(f'Pipeline class change failed: type={new_pipe_type} pipeline={cls}') return pipe + except Exception as e: # pylint: disable=unused-variable + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + shared.log.trace(f"Pipeline class change requested: target={new_pipe_type} fn={fn}") # pylint: disable=protected-access + shared.log.warning(f'Pipeline class change failed: type={new_pipe_type} pipeline={cls} {e}') + has_errors = True + if not hasattr(pipe, 'config') or has_errors: + try: # maybe a wrapper pipeline so just change the class + if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE: + pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING, cls) # pylint: disable=protected-access + new_pipe = pipe + elif new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE: + pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING, cls) # pylint: disable=protected-access + new_pipe = pipe + elif new_pipe_type == DiffusersTaskType.INPAINTING: + pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING, cls) # pylint: disable=protected-access + new_pipe = pipe + else: + shared.log.error(f'Pipeline class set failed: type={new_pipe_type} pipeline={cls}') + return pipe + except Exception as e: # pylint: disable=unused-variable + shared.log.warning(f'Pipeline class set failed: type={new_pipe_type} pipeline={cls} {e}') + has_errors = True + return pipe if new_pipe is None: return pipe @@ -1152,60 +1172,6 @@ def set_diffuser_pipe(pipe, new_pipe_type): return pipe -def set_diffusers_attention(pipe, quiet:bool=False): - import diffusers.models.attention_processor as p - - def set_attn(pipe, attention, name:str=None, quiet:bool=False): - if attention is None: - return - # other models uses their own attention processor - if pipe.__class__.__name__.startswith("StableDiffusion") and getattr(pipe, "unet", None) is not None and hasattr(pipe.unet, "set_attn_processor"): - try: - pipe.unet.set_attn_processor(attention) - except Exception as e: - if 'Nunchaku' in pipe.unet.__class__.__name__: - pass - else: - shared.log.error(f"Attention: {name if name is not None else attention.__class__.__name__} pipe={pipe.__class__.__name__} {e}") - elif not quiet: - shared.log.warning(f"Attention: {name if name is not None else attention.__class__.__name__} is not compatible with {pipe.__class__.__name__}") - - # if hasattr(pipe, 'pipe'): - # set_diffusers_attention(pipe.pipe) - - if 'Control' in pipe.__class__.__name__ or 'Adapter' in pipe.__class__.__name__ or not (pipe.__class__.__name__.startswith("StableDiffusion") and hasattr(pipe, "unet")): - if shared.opts.cross_attention_optimization not in {"Scaled-Dot-Product", "Disabled"}: - shared.log.warning(f"Attention: {shared.opts.cross_attention_optimization} is not compatible with {pipe.__class__.__name__}") - else: - pipe.current_attn_name = shared.opts.cross_attention_optimization - return - - shared.log.quiet(quiet, f'Setting model: attention="{shared.opts.cross_attention_optimization}"') - if shared.opts.cross_attention_optimization == "Disabled": - pass # do nothing - elif shared.opts.cross_attention_optimization == "Scaled-Dot-Product": # The default set by Diffusers - set_attn(pipe, p.AttnProcessor2_0(), name="Scaled-Dot-Product", quiet=True) - elif shared.opts.cross_attention_optimization == "xFormers": - if hasattr(pipe, 'enable_xformers_memory_efficient_attention'): - pipe.enable_xformers_memory_efficient_attention() - else: - shared.log.warning(f"Attention: xFormers is not compatible with {pipe.__class__.__name__}") - elif shared.opts.cross_attention_optimization == "Batch matrix-matrix": - set_attn(pipe, p.AttnProcessor(), name="Batch matrix-matrix") - elif shared.opts.cross_attention_optimization == "Dynamic Attention BMM": - from modules.sd_hijack_dynamic_atten import DynamicAttnProcessorBMM - set_attn(pipe, DynamicAttnProcessorBMM(), name="Dynamic Attention BMM") - - if shared.opts.attention_slicing != "Default" and hasattr(pipe, "enable_attention_slicing") and hasattr(pipe, "disable_attention_slicing"): - if shared.opts.attention_slicing: - pipe.enable_attention_slicing() - else: - pipe.disable_attention_slicing() - shared.log.debug(f"Attention: slicing={shared.opts.attention_slicing}") - - pipe.current_attn_name = shared.opts.cross_attention_optimization - - def add_noise_pred_to_diffusers_callback(pipe): if not hasattr(pipe, "_callback_tensor_inputs"): return pipe @@ -1307,6 +1273,7 @@ def clear_caches(full:bool=False): def unload_model_weights(op='model'): + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access clear_caches(full=True) if shared.compiled_model_state is not None: shared.compiled_model_state.compiled_cache.clear() @@ -1319,14 +1286,14 @@ def unload_model_weights(op='model'): move_model(model_data.sd_model, 'meta') model_data.sd_model = None devices.torch_gc(force=True, reason='unload') - shared.log.debug(f'Unload {op}: {memory_stats()} after') + shared.log.debug(f'Unload {op}: {memory_stats()} fn={fn}') elif (op == 'refiner') and model_data.sd_refiner: shared.log.debug(f'Current {op}: {memory_stats()}') disable_offload(model_data.sd_refiner) move_model(model_data.sd_refiner, 'meta') model_data.sd_refiner = None devices.torch_gc(force=True, reason='unload') - shared.log.debug(f'Unload {op}: {memory_stats()}') + shared.log.debug(f'Unload {op}: {memory_stats()} fn={fn}') def hf_auth_check(checkpoint_info, force:bool=False): diff --git a/modules/sdnq/dequantizer.py b/modules/sdnq/dequantizer.py index 73afe67d2..164b3f24a 100644 --- a/modules/sdnq/dequantizer.py +++ b/modules/sdnq/dequantizer.py @@ -137,17 +137,25 @@ class AsymmetricWeightsDequantizer(torch.nn.Module): result_dtype: torch.dtype, result_shape: torch.Size, original_shape: torch.Size, + quantized_weight_shape: torch.Size, weights_dtype: str, - use_quantized_matmul: bool = False, - **kwargs, # pylint: disable=unused-argument + group_size: int, + svd_rank: int, + use_quantized_matmul: bool, + re_quantize_for_matmul: bool, ): super().__init__() - self.weights_dtype = weights_dtype - self.original_shape = original_shape - self.use_quantized_matmul = use_quantized_matmul - self.re_quantize_for_matmul = True + self.is_packed = False + self.is_asym = True self.result_dtype = result_dtype self.result_shape = result_shape + self.original_shape = original_shape + self.quantized_weight_shape = quantized_weight_shape + self.weights_dtype = weights_dtype + self.group_size = group_size + self.svd_rank = svd_rank + self.use_quantized_matmul = use_quantized_matmul + self.re_quantize_for_matmul = re_quantize_for_matmul def pack_weight(self, weight: torch.Tensor) -> torch.Tensor: return weight.to(dtype=dtype_dict[self.weights_dtype]["torch_dtype"]) @@ -165,18 +173,25 @@ class SymmetricWeightsDequantizer(torch.nn.Module): result_dtype: torch.dtype, result_shape: torch.Size, original_shape: torch.Size, + quantized_weight_shape: torch.Size, weights_dtype: str, - use_quantized_matmul: bool = False, - re_quantize_for_matmul: bool = False, - **kwargs, # pylint: disable=unused-argument + group_size: int, + svd_rank: int, + use_quantized_matmul: bool, + re_quantize_for_matmul: bool, ): super().__init__() - self.weights_dtype = weights_dtype - self.original_shape = original_shape - self.use_quantized_matmul = use_quantized_matmul - self.re_quantize_for_matmul = re_quantize_for_matmul + self.is_packed = False + self.is_asym = False self.result_dtype = result_dtype self.result_shape = result_shape + self.original_shape = original_shape + self.quantized_weight_shape = quantized_weight_shape + self.weights_dtype = weights_dtype + self.group_size = group_size + self.svd_rank = svd_rank + self.use_quantized_matmul = use_quantized_matmul + self.re_quantize_for_matmul = re_quantize_for_matmul def pack_weight(self, weight: torch.Tensor) -> torch.Tensor: return weight.to(dtype=dtype_dict[self.weights_dtype]["torch_dtype"]) @@ -192,22 +207,28 @@ class SymmetricWeightsDequantizer(torch.nn.Module): class PackedINTAsymmetricWeightsDequantizer(torch.nn.Module): def __init__( self, - quantized_weight_shape: torch.Size, result_dtype: torch.dtype, result_shape: torch.Size, original_shape: torch.Size, + quantized_weight_shape: torch.Size, weights_dtype: str, - use_quantized_matmul: bool = False, - **kwargs, # pylint: disable=unused-argument + group_size: int, + svd_rank: int, + use_quantized_matmul: bool, + re_quantize_for_matmul: bool, ): super().__init__() - self.weights_dtype = weights_dtype - self.use_quantized_matmul = use_quantized_matmul - self.re_quantize_for_matmul = True - self.original_shape = original_shape - self.quantized_weight_shape = quantized_weight_shape + self.is_packed = True + self.is_asym = True self.result_dtype = result_dtype self.result_shape = result_shape + self.original_shape = original_shape + self.quantized_weight_shape = quantized_weight_shape + self.weights_dtype = weights_dtype + self.group_size = group_size + self.svd_rank = svd_rank + self.use_quantized_matmul = use_quantized_matmul + self.re_quantize_for_matmul = re_quantize_for_matmul def pack_weight(self, weight: torch.Tensor) -> torch.Tensor: return pack_int_asymetric(weight, self.weights_dtype) @@ -222,23 +243,28 @@ class PackedINTAsymmetricWeightsDequantizer(torch.nn.Module): class PackedINTSymmetricWeightsDequantizer(torch.nn.Module): def __init__( self, - quantized_weight_shape: torch.Size, result_dtype: torch.dtype, result_shape: torch.Size, original_shape: torch.Size, + quantized_weight_shape: torch.Size, weights_dtype: str, - use_quantized_matmul: bool = False, - re_quantize_for_matmul: bool = False, - **kwargs, # pylint: disable=unused-argument + group_size: int, + svd_rank: int, + use_quantized_matmul: bool, + re_quantize_for_matmul: bool, ): super().__init__() - self.weights_dtype = weights_dtype - self.original_shape = original_shape - self.use_quantized_matmul = use_quantized_matmul - self.re_quantize_for_matmul = re_quantize_for_matmul - self.quantized_weight_shape = quantized_weight_shape + self.is_packed = True + self.is_asym = False self.result_dtype = result_dtype self.result_shape = result_shape + self.original_shape = original_shape + self.quantized_weight_shape = quantized_weight_shape + self.weights_dtype = weights_dtype + self.group_size = group_size + self.svd_rank = svd_rank + self.use_quantized_matmul = use_quantized_matmul + self.re_quantize_for_matmul = re_quantize_for_matmul def pack_weight(self, weight: torch.Tensor) -> torch.Tensor: return pack_int_symetric(weight, self.weights_dtype) diff --git a/modules/sdnq/layers/conv/conv_int8.py b/modules/sdnq/layers/conv/conv_int8.py index a1e297cc1..d566e74ff 100644 --- a/modules/sdnq/layers/conv/conv_int8.py +++ b/modules/sdnq/layers/conv/conv_int8.py @@ -74,7 +74,7 @@ def quantized_conv_forward_int8_matmul(self, input) -> torch.FloatTensor: else: weight = self.weight scale = self.scale - quantized_weight_shape = getattr(self.sdnq_dequantizer, "quantized_weight_shape", None) + 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, diff --git a/modules/sdnq/layers/linear/linear_int8.py b/modules/sdnq/layers/linear/linear_int8.py index 7fa60a87f..6ab859815 100644 --- a/modules/sdnq/layers/linear/linear_int8.py +++ b/modules/sdnq/layers/linear/linear_int8.py @@ -56,7 +56,7 @@ def quantized_linear_forward_int8_matmul(self, input: torch.FloatTensor) -> torc else: weight = self.weight scale = self.scale - quantized_weight_shape = getattr(self.sdnq_dequantizer, "quantized_weight_shape", None) + quantized_weight_shape = self.sdnq_dequantizer.quantized_weight_shape if self.sdnq_dequantizer.is_packed else None return int8_matmul(input, weight, self.bias, scale, self.svd_up, self.svd_down, quantized_weight_shape, self.sdnq_dequantizer.weights_dtype) diff --git a/modules/sdnq/loader.py b/modules/sdnq/loader.py index 749af9fb4..555f9d26f 100644 --- a/modules/sdnq/loader.py +++ b/modules/sdnq/loader.py @@ -2,8 +2,9 @@ import os import json import torch from diffusers.models.modeling_utils import ModelMixin -from .common import dtype_dict, use_tensorwise_fp8_matmul, use_contiguous_mm -from .quantizer import SDNQConfig, sdnq_post_load_quant + +from .common import dtype_dict, use_tensorwise_fp8_matmul +from .quantizer import SDNQConfig, sdnq_post_load_quant, prepare_weight_for_matmul, prepare_svd_for_matmul from .dequantizer import dequantize_symmetric, re_quantize_int8, re_quantize_fp8 from .forward import get_forward_func from .file_loader import load_files @@ -67,20 +68,25 @@ def load_sdnq_model(model_path: str, model_cls: ModelMixin = None, file_name: st from accelerate import init_empty_weights with init_empty_weights(): - if quantization_config is None: - try: - with open(os.path.join(model_path, "quantization_config.json"), "r", encoding="utf-8") as f: - quantization_config = json.load(f) - except Exception: - quantization_config = {} + model_config_path = os.path.join(model_path, "config.json") + quantization_config_path = os.path.join(model_path, "quantization_config.json") if model_config is None: - try: - with open(os.path.join(model_path, "config.json"), "r", encoding="utf-8") as f: + if os.path.exists(model_config_path): + with open(model_config_path, "r", encoding="utf-8") as f: model_config = json.load(f) - except Exception: + else: model_config = {} + if quantization_config is None: + if os.path.exists(quantization_config_path): + with open(quantization_config_path, "r", encoding="utf-8") as f: + quantization_config = json.load(f) + else: + quantization_config = model_config.get("quantization_config", None) + if quantization_config is None: + raise ValueError(f"Cannot determine quantization_config for {model_path}, please provide quantization_config argument") + if model_cls is None: import transformers import diffusers @@ -99,14 +105,14 @@ def load_sdnq_model(model_path: str, model_cls: ModelMixin = None, file_name: st quantization_config.pop("non_blocking", None) quantization_config.pop("add_skip_keys", None) - if hasattr(model_cls, "load_config"): + if hasattr(model_cls, "load_config") and hasattr(model_cls, "from_config"): config = model_cls.load_config(model_path) model = model_cls.from_config(config) elif hasattr(model_cls, "_from_config"): config = transformers.AutoConfig.from_pretrained(model_path) model = model_cls(config) else: - raise ValueError(f"Dont know how to load model for {model_cls}") + model = model_cls(**model_config) model = sdnq_post_load_quant(model, add_skip_keys=False, **quantization_config) @@ -127,11 +133,26 @@ def load_sdnq_model(model_path: str, model_cls: ModelMixin = None, file_name: st model.load_state_dict(state_dict, assign=True) del state_dict + model = post_process_model(model) if (dtype is not None) or (dequantize_fp32 is not None) or (use_quantized_matmul is not None): model = apply_options_to_model(model, dtype=dtype, dequantize_fp32=dequantize_fp32, use_quantized_matmul=use_quantized_matmul) return model +def post_process_model(model): + has_children = list(model.children()) + if not has_children: + return model + for module in model.children(): + if hasattr(module, "sdnq_dequantizer"): + if module.sdnq_dequantizer.use_quantized_matmul and not module.sdnq_dequantizer.re_quantize_for_matmul: + module.weight.data = prepare_weight_for_matmul(module.weight) + if module.svd_up is not None: + module.svd_up.data, module.svd_down.data = prepare_svd_for_matmul(module.svd_up, module.svd_down, module.sdnq_dequantizer.use_quantized_matmul) + module = post_process_model(module) + return model + + def apply_options_to_model(model, dtype: torch.dtype = None, dequantize_fp32: bool = None, use_quantized_matmul: bool = None): has_children = list(model.children()) if not has_children: @@ -168,30 +189,12 @@ def apply_options_to_model(model, dtype: torch.dtype = None, dequantize_fp32: bo if use_tensorwise_fp8_matmul: module.scale.data = module.scale.to(dtype=scale_dtype) elif not module.sdnq_dequantizer.re_quantize_for_matmul: - module.weight.data, module.scale.data = module.weight.t_(), module.scale.t_() + module.scale.t_() + module.weight.t_() if use_quantized_matmul: - if use_contiguous_mm: - module.weight.data = module.weight.contiguous() - elif module.weight.is_contiguous(): - module.weight.data = module.weight.t_().contiguous().t_() + module.weight.data = prepare_weight_for_matmul(module.weight) if module.svd_up is not None: - module.svd_up.data = module.svd_up.t_() - module.svd_down.data = module.svd_down.t_() - if use_quantized_matmul: - if use_contiguous_mm: - module.svd_up.data = module.svd_up.contiguous() - module.svd_down.data = module.svd_down.contiguous() - else: - if module.svd_up.is_contiguous(): - module.svd_up.data = module.svd_up.t_().contiguous().t_() - if module.svd_up.is_contiguous(): - module.svd_down.data = module.svd_down.t_().contiguous().t_() - else: - module.svd_up.data = module.svd_up.contiguous() - if use_contiguous_mm: - module.svd_down.data = module.svd_down.contiguous() - elif module.svd_down.is_contiguous(): - module.svd_down.data = module.svd_down.t_().contiguous().t_() + 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__, use_quantized_matmul, dtype_dict[module.sdnq_dequantizer.weights_dtype]["is_integer"], use_tensorwise_fp8_matmul) module.forward = module.forward.__get__(module, module.__class__) diff --git a/modules/sdnq/quantizer.py b/modules/sdnq/quantizer.py index 14ecfcc92..a4a96c24c 100644 --- a/modules/sdnq/quantizer.py +++ b/modules/sdnq/quantizer.py @@ -68,6 +68,25 @@ def apply_svdquant(weight: torch.FloatTensor, rank: int = 32, niter: int = 8) -> return weight, svd_up, svd_down +def prepare_weight_for_matmul(weight: torch.Tensor) -> torch.Tensor: + if use_contiguous_mm: + weight = weight.contiguous() + elif weight.is_contiguous(): + weight = weight.t_().contiguous().t_() + return weight + + +def prepare_svd_for_matmul(svd_up: torch.FloatTensor, svd_down: torch.FloatTensor, use_quantized_matmul: bool) -> Tuple[torch.FloatTensor, torch.FloatTensor]: + if svd_up is not None: + if use_quantized_matmul: + svd_up = prepare_weight_for_matmul(svd_up) + else: + svd_up = svd_up.contiguous() + if svd_down is not None: + svd_down = prepare_weight_for_matmul(svd_down) + return svd_up, svd_down + + def check_param_name_in(param_name: str, param_list: List[str]) -> bool: split_param_name = param_name.split(".") for param in param_list: @@ -212,20 +231,7 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_siz if use_quantized_matmul: svd_up = svd_up.t_() svd_down = svd_down.t_() - if use_contiguous_mm: - svd_up = svd_up.contiguous() - svd_down = svd_down.contiguous() - else: - if svd_up.is_contiguous(): - svd_up = svd_up.t_().contiguous().t_() - if svd_down.is_contiguous(): - svd_down = svd_down.t_().contiguous().t_() - else: - svd_up = svd_up.contiguous() - if use_contiguous_mm: - svd_down = svd_down.contiguous() - elif svd_down.is_contiguous(): - svd_down = svd_down.t_().contiguous().t_() + svd_up, svd_down = prepare_svd_for_matmul(svd_up, svd_down, use_quantized_matmul) except Exception: svd_up, svd_down = None, None else: @@ -295,10 +301,7 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_siz if use_quantized_matmul and not re_quantize_for_matmul: scale.t_() layer.weight.t_() - if use_contiguous_mm: - layer.weight.data = layer.weight.contiguous() - elif layer.weight.is_contiguous(): - layer.weight.data = layer.weight.t_().contiguous().t_() + layer.weight.data = prepare_weight_for_matmul(layer.weight) if not use_tensorwise_fp8_matmul and not dtype_dict[weights_dtype]["is_integer"]: scale = scale.to(dtype=torch.float32) @@ -318,11 +321,13 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_siz layer.svd_up, layer.svd_down = None, None layer.sdnq_dequantizer = dequantizer_dict[weights_dtype]( - quantized_weight_shape=layer.weight.shape, result_dtype=torch_dtype, result_shape=result_shape, original_shape=original_shape, + quantized_weight_shape=layer.weight.shape, weights_dtype=weights_dtype, + group_size=group_size, + svd_rank=svd_rank, use_quantized_matmul=use_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul, ) @@ -354,27 +359,23 @@ def apply_sdnq_to_module(model, weights_dtype="int8", torch_dtype=None, group_si if layer_class_name in allowed_types: if (layer_class_name in conv_types or layer_class_name in conv_transpose_types) and not quant_conv: continue - else: - continue - - weights_dtype = get_minimum_dtype(weights_dtype, param_name, modules_dtype_dict) - module = sdnq_quantize_layer( - module, - weights_dtype=weights_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, - dequantize_fp32=dequantize_fp32, - non_blocking=non_blocking, - quantization_device=quantization_device, - return_device=return_device, - param_name=param_name, - ) + module = sdnq_quantize_layer( + module, + weights_dtype=get_minimum_dtype(weights_dtype, param_name, modules_dtype_dict), + 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, + dequantize_fp32=dequantize_fp32, + non_blocking=non_blocking, + quantization_device=quantization_device, + return_device=return_device, + param_name=param_name, + ) module = apply_sdnq_to_module( module, weights_dtype=weights_dtype, @@ -418,6 +419,13 @@ def sdnq_post_load_quant( modules_dtype_dict: Dict[str, List[str]] = None, op=None, ): + if modules_to_not_convert is None: + modules_to_not_convert = [] + if modules_dtype_dict is None: + modules_dtype_dict = {} + + modules_to_not_convert = modules_to_not_convert.copy() + modules_dtype_dict = modules_dtype_dict.copy() if add_skip_keys: model, modules_to_not_convert, modules_dtype_dict = add_module_skip_keys(model, modules_to_not_convert, modules_dtype_dict) @@ -438,7 +446,7 @@ def sdnq_post_load_quant( quantization_device=quantization_device, return_device=return_device, modules_to_not_convert=modules_to_not_convert, - modules_dtype_dict=modules_dtype_dict.copy(), + modules_dtype_dict=modules_dtype_dict, op=op, ) model.quantization_config = SDNQConfig( @@ -455,12 +463,15 @@ def sdnq_post_load_quant( quantization_device=quantization_device, return_device=return_device, modules_to_not_convert=modules_to_not_convert, - modules_dtype_dict=modules_dtype_dict.copy(), + modules_dtype_dict=modules_dtype_dict, ) if hasattr(model, "config"): try: model.config.quantization_config = model.quantization_config + except Exception: + pass + try: model.config["quantization_config"] = model.quantization_config.to_dict() except Exception: pass @@ -543,6 +554,14 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): param_value = param_value.clone() else: param_value = param_value.to(target_device, dtype=return_dtype) + + if tensor_name == "weight" and layer.sdnq_dequantizer.use_quantized_matmul and not layer.sdnq_dequantizer.re_quantize_for_matmul: + param_value = prepare_weight_for_matmul(param_value) + elif tensor_name == "svd_up": + param_value, _ = prepare_svd_for_matmul(param_value, None, layer.sdnq_dequantizer.use_quantized_matmul) + elif tensor_name == "svd_down": + _, param_value = prepare_svd_for_matmul(None, param_value, layer.sdnq_dequantizer.use_quantized_matmul) + param_value = torch.nn.Parameter(param_value, requires_grad=False) setattr(layer, tensor_name, param_value) return @@ -626,6 +645,9 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): if hasattr(model, "config"): try: model.config.quantization_config = self.quantization_config + except Exception: + pass + try: model.config["quantization_config"] = self.quantization_config.to_dict() except Exception: pass @@ -655,8 +677,17 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): del model.quantization_method if hasattr(model, "quantization_config"): del model.quantization_config - if hasattr(model, "config") and hasattr(model.config, "quantization_config"): - del model.config.quantization_config + if hasattr(model, "config"): + try: + if hasattr(model.config, "quantization_config"): + del model.config.quantization_config + except Exception: + pass + try: + if hasattr(model.config, "pop"): + model.config.pop("quantization_config", None) + except Exception: + pass return model def is_serializable(self, *args, **kwargs) -> bool: # pylint: disable=unused-argument, invalid-overridden-method @@ -772,6 +803,7 @@ class SDNQConfig(QuantizationConfigMixin): elif not isinstance(self.modules_dtype_dict, dict): raise ValueError(f"modules_dtype_dict must be a dict but got {type(self.modules_dtype_dict)}") elif len(self.modules_dtype_dict.keys()) > 0: + self.modules_dtype_dict = self.modules_dtype_dict.copy() for key, value in self.modules_dtype_dict.items(): if isinstance(value, str): value = [value] @@ -782,6 +814,9 @@ class SDNQConfig(QuantizationConfigMixin): if not isinstance(key, str) or not isinstance(value, list): raise ValueError(f"modules_dtype_dict must be a dictionary of strings and lists but got {type(key)} and {type(value)}") + self.modules_to_not_convert = self.modules_to_not_convert.copy() + 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 diff --git a/modules/shared.py b/modules/shared.py index 24aad5271..8af5135e2 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -136,7 +136,7 @@ def list_samplers(): return modules.sd_samplers.all_samplers -startup_offload_mode, startup_offload_min_gpu, startup_offload_max_gpu, startup_cross_attention, startup_sdp_options, startup_sdp_choices, startup_offload_always, startup_offload_never = get_default_modes(cmd_opts=cmd_opts, mem_stat=mem_stat) +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) options_templates.update(options_section(('sd', "Model Loading"), { "sd_backend": OptionInfo('diffusers', "Execution backend", gr.Radio, {"choices": ['diffusers', 'original'], "visible": False }), @@ -296,13 +296,13 @@ options_templates.update(options_section(('cuda', "Compute Settings"), { "diffusers_generator_device": OptionInfo("GPU", "Generator device", gr.Radio, {"choices": ["GPU", "CPU", "Unset"]}), "cross_attention_sep": OptionInfo("

Cross Attention

", "", gr.HTML), - "cross_attention_optimization": OptionInfo(startup_cross_attention, "Attention optimization method", gr.Radio, lambda: {"choices": shared_items.list_crossattention()}), - "attention_": OptionInfo("

Cross Attention

", "", gr.HTML), + "cross_attention_optimization": OptionInfo(startup_cross_attention, "Attention method", gr.Radio, lambda: {"choices": shared_items.list_crossattention()}), + "sdp_options": OptionInfo(startup_sdp_options, "SDP kernels", gr.CheckboxGroup, {"choices": startup_sdp_choices}), + "sdp_overrides": OptionInfo(startup_sdp_override_options, "SDP overrides", gr.CheckboxGroup, {"choices": startup_sdp_override_choices}), "attention_slicing": OptionInfo('Default', "Attention slicing", gr.Radio, {"choices": ['Default', 'Enabled', 'Disabled']}), - "sdp_options": OptionInfo(startup_sdp_options, "SDP options", gr.CheckboxGroup, {"choices": startup_sdp_choices}), "xformers_options": OptionInfo(['Flash attention'], "xFormers options", gr.CheckboxGroup, {"choices": ['Flash attention'] }), - "dynamic_attention_slice_rate": OptionInfo(0.5, "Dynamic Attention slicing rate in GB", gr.Slider, {"minimum": 0.01, "maximum": max(gpu_memory,4), "step": 0.01}), - "dynamic_attention_trigger_rate": OptionInfo(1, "Dynamic Attention trigger rate in GB", gr.Slider, {"minimum": 0.01, "maximum": max(gpu_memory,4)*2, "step": 0.01}), + "dynamic_attention_slice_rate": OptionInfo(0.5, "Dynamic Attention slicing rate", gr.Slider, {"minimum": 0.01, "maximum": max(gpu_memory,4), "step": 0.01}), + "dynamic_attention_trigger_rate": OptionInfo(1, "Dynamic Attention trigger rate", gr.Slider, {"minimum": 0.01, "maximum": max(gpu_memory,4)*2, "step": 0.01}), })) options_templates.update(options_section(('backends', "Backend Settings"), { @@ -711,14 +711,15 @@ options_templates.update(options_section(('extra_networks', "Networks"), { "extra_networks_lora_sep": OptionInfo("

LoRA

", "", gr.HTML), "extra_networks_default_multiplier": OptionInfo(1.0, "Default strength", gr.Slider, {"minimum": 0.0, "maximum": 2.0, "step": 0.01}), - "lora_fuse_diffusers": OptionInfo(True, "LoRA fuse directly to model"), "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_maybe_diffusers": OptionInfo(False, "LoRA load using Diffusers method for selected models", gr.Checkbox, {"visible": False}), + "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"), diff --git a/modules/shared_defaults.py b/modules/shared_defaults.py index d3993d2e9..9eac6dc5e 100644 --- a/modules/shared_defaults.py +++ b/modules/shared_defaults.py @@ -40,16 +40,24 @@ def get_default_modes(cmd_opts, mem_stat): default_cross_attention = "Scaled-Dot-Product" - if devices.backend == "zluda": - default_sdp_options = ['Math attention', 'Dynamic attention'] - elif devices.backend in {"rocm", "directml", "cpu", "mps"}: - default_sdp_options = ['Flash attention', 'Memory attention', 'Math attention', 'Dynamic attention'] - else: - default_sdp_options = ['Flash attention', 'Memory attention', 'Math attention'] + default_sdp_choices = ['Flash', 'Memory', 'Math'] + default_sdp_options = ['Flash', 'Memory', 'Math'] + + default_sdp_override_choices = ['Dynamic attention', 'Flex attention', 'Flash attention', 'Sage attention'] + default_sdp_override_options = [] + + if devices.backend == "zluda": + default_sdp_options = ['Math'] + default_sdp_override_options = ['Dynamic attention'] + 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: + 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'] - default_sdp_choices = ['Flash attention', 'Memory attention', 'Math attention', 'Dynamic attention', 'CK Flash attention', 'Sage attention'] - if devices.backend in {"rocm", "zluda"}: - default_sdp_choices.insert(4, 'Triton Flash attention') # insert after Dynamic attention return ( default_offload_mode, @@ -58,6 +66,8 @@ def get_default_modes(cmd_opts, mem_stat): default_cross_attention, default_sdp_options, default_sdp_choices, + default_sdp_override_options, + default_sdp_override_choices, default_diffusers_offload_always, - default_diffusers_offload_never + default_diffusers_offload_never, ) diff --git a/modules/styles.py b/modules/styles.py index 4ff29b6e1..02187f581 100644 --- a/modules/styles.py +++ b/modules/styles.py @@ -46,17 +46,24 @@ def apply_styles_to_prompt(prompt, styles): def apply_curly_braces_to_prompt(prompt, seed=-1): - # woman with {blonde|brunette|red-head|purple highlights} hair - curly_braces_matches = re.findall(r'\{(.*?)\}', prompt) - for match in curly_braces_matches: - old_state = None - if seed > 0: - old_state = random.getstate() - random.seed(seed) - options = match.split('|') - if options: - choice = random.choice(options).strip() - prompt = prompt.replace(f'{{{match}}}', choice, 1) + # woman with {white|green|{purple|yellow}} highlights and {red|blue} dress + if not isinstance(prompt, str) or len(prompt) == 0: + return prompt + old_state = None + if seed > 0: + old_state = random.getstate() + random.seed(seed) + try: + pattern = re.compile(r'\{([^{}]*)\}', re.DOTALL) # innermost braces + while True: + m = pattern.search(prompt) + 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 '' + 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: random.setstate(old_state) return prompt diff --git a/modules/ui_control.py b/modules/ui_control.py index af5de9882..9b9f0f287 100644 --- a/modules/ui_control.py +++ b/modules/ui_control.py @@ -5,6 +5,7 @@ from modules.control import unit from modules import errors, shared, progress, generation_parameters_copypaste, call_queue, scripts_manager, masking, images, processing_vae, timer # pylint: disable=ungrouped-imports from modules import ui_common, ui_sections, ui_guidance from modules import ui_control_helpers as helpers +import installer gr_height = 512 @@ -185,12 +186,13 @@ def create_ui(_blocks: gr.Blocks=None): with gr.Tabs(elem_classes=['control-tabs'], elem_id='control-tab-input'): input_mode = gr.Label(value='select', visible=False) with gr.Tab('Image', id='in-image') as tab_image: - input_image = gr.Image(label="Input", show_label=False, type="pil", interactive=True, tool="editor", height=gr_height, image_mode='RGB', elem_id='control_input_select', elem_classes=['control-image']) + if (installer.version['kanvas'] == 'disabled') or (installer.version['kanvas'] == 'unavailable'): + shared.log.warning(f'Kanvas: status={installer.version["kanvas"]}') + input_image = gr.Image(label="Input", show_label=False, type="pil", interactive=True, tool="editor", height=gr_height, image_mode='RGB', elem_id='control_input_select', elem_classes=['control-image']) + else: + input_image = gr.HTML(value='

Kanvas not initialized

', elem_id='kanvas-container') + input_changed = gr.Button('Kanvas change', elem_id='kanvas-change-button', visible=False) btn_interrogate = ui_sections.create_interrogate_button('control', what='input') - with gr.Tab('Inpaint', id='in-inpaint') as _tab_inpaint: - input_inpaint = gr.Image(label="Input", show_label=False, type="pil", interactive=True, tool="sketch", height=gr_height, image_mode='RGB', elem_id='control_input_inpaint', brush_radius=32, mask_opacity=0.6, elem_classes=['control-image']) - with gr.Tab('Outpaint', id='in-outpaint') as _tab_outpaint: - input_resize = gr.Image(label="Input", show_label=False, type="pil", interactive=True, tool="select", height=gr_height, image_mode='RGB', elem_id='control_input_resize', elem_classes=['control-image']) with gr.Tab('Video', id='in-video') as tab_video: input_video = gr.Video(label="Input", show_label=False, interactive=True, height=gr_height, elem_classes=['control-image']) with gr.Tab('Batch', id='in-batch') as tab_batch: @@ -246,28 +248,29 @@ def create_ui(_blocks: gr.Blocks=None): 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 = False) btn_negative_counter.click(fn=call_queue.wrap_queued_call(ui_common.update_token_counter), inputs=[negative], outputs=[negative_counter], show_progress = False) - btn_interrogate.click(fn=helpers.interrogate, inputs=[], outputs=[prompt]) select_dict = dict( fn=helpers.select_input, _js="controlInputMode", - inputs=[input_mode, input_image, init_image, input_type, input_resize, input_inpaint, input_video, input_batch, input_folder], + inputs=[input_mode, input_image, init_image, input_type, input_video, input_batch, input_folder], outputs=[output_tabs, preview_process, result_txt, width_before, height_before], show_progress=False, queue=False, ) + input_changed.click(**select_dict) + btn_interrogate.click(**select_dict) # need to fetch input first + btn_interrogate.click(fn=helpers.interrogate, inputs=[], outputs=[prompt]) + + prompt.submit(**select_dict) negative.submit(**select_dict) btn_generate.click(**select_dict) - for ctrl in [input_image, input_resize, input_video, input_batch, input_folder, init_image, init_video, init_batch, init_folder, tab_image, tab_video, tab_batch, tab_folder, tab_image_init, tab_video_init, tab_batch_init, tab_folder_init]: + for ctrl in [input_image, input_video, input_batch, input_folder, init_image, init_video, init_batch, init_folder, tab_image, tab_video, tab_batch, tab_folder, tab_image_init, tab_video_init, tab_batch_init, tab_folder_init]: if hasattr(ctrl, 'change'): ctrl.change(**select_dict) if hasattr(ctrl, 'clear'): ctrl.clear(**select_dict) - for ctrl in [input_inpaint]: # gradio image mode inpaint triggeres endless loop on change event - if hasattr(ctrl, 'upload'): - ctrl.upload(**select_dict) tabs_state = gr.Textbox(value='none', visible=False) input_fields = [ @@ -382,8 +385,7 @@ def create_ui(_blocks: gr.Blocks=None): # second pass (enable_hr, "Second pass"), (enable_hr, "Refine"), - (denoising_strength, "Denoising strength"), - (denoising_strength, "Hires strength"), + (hr_denoising_strength, "Hires strength"), (hr_sampler_index, "Hires sampler"), (hr_resize_mode, "Hires mode"), (hr_resize_context, "Hires context"), @@ -410,7 +412,7 @@ def create_ui(_blocks: gr.Blocks=None): generation_parameters_copypaste.add_paste_fields("control", input_image, paste_fields, override_settings) bindings = generation_parameters_copypaste.ParamBinding(paste_button=btn_paste, tabname="control", source_text_component=prompt, source_image_component=output_gallery) generation_parameters_copypaste.register_paste_params_button(bindings) - masking.bind_controls([input_image, input_inpaint, input_resize], preview_process, output_image) + # masking.bind_controls([input_image], preview_process, output_image) if os.environ.get('SD_CONTROL_DEBUG', None) is not None: # debug only from modules.control.test import test_processors, test_controlnets, test_adapters, test_xs, test_lite diff --git a/modules/ui_control_helpers.py b/modules/ui_control_helpers.py index 05e31bf1c..6940411ef 100644 --- a/modules/ui_control_helpers.py +++ b/modules/ui_control_helpers.py @@ -1,4 +1,5 @@ import os +import time import gradio as gr from PIL import Image from modules import shared, scripts_manager, masking, video # pylint: disable=ungrouped-imports @@ -49,11 +50,14 @@ def initialize(): def interrogate(): prompt = None + if input_source is None or len(input_source) == 0: + shared.log.warning('Interrogate: no input source') + return prompt try: from modules.interrogate.interrogate import interrogate as interrogate_fn prompt = interrogate_fn(input_source[0]) - except Exception: - pass + except Exception as e: + shared.log.error(f'Interrogate: {e}') return prompt @@ -63,6 +67,8 @@ def display_units(num_units): def get_video(filepath: str): + if not os.path.exists(filepath): + return '' try: frames, fps, duration, w, h, codec, _cap = video.get_video_params(filepath) shared.log.debug(f'Control: input video: path={filepath} frames={frames} fps={fps} size={w}x{h} codec={codec}') @@ -74,28 +80,69 @@ def get_video(filepath: str): return msg -def select_input(input_mode, input_image, init_image, init_type, input_resize, input_inpaint, input_video, input_batch, input_folder): +def process_kanvas(x): # only used when kanvas overrides gr.Image object + image = None + mask = None + try: # try base64 decode + t0 = time.time() + image_data = x.get('image', '') + image_bytes = len(image_data) + if image_bytes > 0: + from modules.api import helpers + image = helpers.decode_base64_to_image(image_data) + image = image.convert('RGB') + mask_data = x.get('mask', '') + mask_bytes = len(mask_data) + if mask_bytes > 0: + from modules.api import helpers + mask = helpers.decode_base64_to_image(mask_data) + mask = mask.convert('L') + t1 = time.time() + shared.log.debug(f'Kanvas: image={image}:{image_bytes} mask={mask}:{mask_bytes} time={t1-t0:.2f}') + return image, mask + except Exception: + pass + try: # try raw pixel data + import numpy as np + t0 = time.time() + image_data = list(x.get('image', {}).values()) + if image_data: + width = x['imageWidth'] + height = x['imageHeight'] + array = np.array(image_data, dtype=np.uint8).reshape((height, width, 4)) + image = Image.fromarray(array, 'RGBA') + image = image.convert('RGB') + mask_data = list(x.get('mask', {}).values()) + if mask_data: + width = x['maskWidth'] + height = x['maskHeight'] + array = np.array(mask_data, dtype=np.uint8).reshape((height, width, 4)) + mask = Image.fromarray(array, 'RGBA') + # alpha = mask.getchannel("A").convert("L") + # mask = Image.merge("RGB", [alpha, alpha, alpha]) + mask = mask.convert('L') + t1 = time.time() + shared.log.debug(f'Kanvas: image={image} mask={mask} time={t1-t0:.2f}') + except Exception: + pass + return image, mask + + +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 - if input_mode == 'Image': - selected_input = input_image - elif input_mode == 'Outpaint': - selected_input = input_resize - elif input_mode == 'Inpaint': - selected_input = input_inpaint - elif input_mode == 'Video': + selected_input = input_image # default: Image or Kanvas + if input_mode == 'Video': selected_input = input_video elif input_mode == 'Batch': selected_input = input_batch elif input_mode == 'Folder': selected_input = input_folder - else: - selected_input = None size = [gr.update(), gr.update()] if selected_input is None: input_source = None busy = False - # debug('Control input: none') return [gr.Tabs.update(), None, ''] + size input_type = type(selected_input) input_mask = None @@ -108,21 +155,29 @@ def select_input(input_mode, input_image, init_image, init_type, input_resize, i selected_input, input_mask = masking.outpaint(input_image=selected_input) input_source = [selected_input] input_type = 'PIL.Image' - status = f'Control input | Image | Size {selected_input.width}x{selected_input.height} | Mode {selected_input.mode}' + status = f'Control input | Image | Size {selected_input.width if selected_input else 0}x{selected_input.height if selected_input else 0} | Mode {selected_input.mode if selected_input else "Unknown"}' size = [gr.update(value=selected_input.width), gr.update(value=selected_input.height)] res = [gr.Tabs.update(selected='out-gallery'), input_mask, status] - elif isinstance(selected_input, dict): # inpaint -> dict image+mask + elif isinstance(selected_input, dict) and 'kanvas' in selected_input: # kanvas via js -> kanvas dict + selected_input, input_mask = process_kanvas(selected_input) + input_source = [selected_input] + input_type = 'Kanvas' + status = f'Control input | Kanvas | Size {selected_input.width if selected_input else 0}x{selected_input.height if selected_input else 0} | Mode {selected_input.mode if selected_input else "Unknown"}' + if selected_input: + size = [gr.update(value=selected_input.width), gr.update(value=selected_input.height)] + res = [gr.Tabs.update(selected='out-gallery'), input_mask, status] + elif isinstance(selected_input, dict) and 'mask' in selected_input: # inpaint -> dict image+mask input_mask = selected_input['mask'] selected_input = selected_input['image'] input_source = [selected_input] input_type = 'PIL.Image' - status = f'Control input | Image | Size {selected_input.width}x{selected_input.height} | Mode {selected_input.mode}' + status = f'Control input | Image | Size {selected_input.width if selected_input else 0}x{selected_input.height if selected_input else 0} | Mode {selected_input.mode if selected_input else "Unknown"}' res = [gr.Tabs.update(selected='out-gallery'), input_mask, status] elif isinstance(selected_input, gr.components.image.Image): # not likely input_source = [selected_input.value] input_type = 'gr.Image' res = [gr.Tabs.update(selected='out-gallery'), input_mask, status] - elif isinstance(selected_input, str): # video via upload > tmp filepath to video + elif isinstance(selected_input, str) and os.path.exists(selected_input): # video via upload > tmp filepath to video input_source = selected_input input_type = 'gr.Video' status = get_video(input_source) @@ -138,14 +193,14 @@ def select_input(input_mode, input_image, init_image, init_type, input_resize, i res = [gr.Tabs.update(selected='out-gallery'), input_mask, status] else: # unknown input_source = None - # init inputs: optional if init_type == 0: # Control only input_init = None elif init_type == 1: # Init image same as control assigned during runtime input_init = None elif init_type == 2: # Separate init image input_init = [init_image] - debug_log(f'Control select input: type={input_type} source={input_source} init={input_init} mask={input_mask} mode={input_mode}') + t1 = time.time() + shared.log.debug(f'Select input: type={input_type} source={input_source} init={input_init} mask={input_mask} mode={input_mode} time={t1-t0:.2f}') busy = False return res + size diff --git a/modules/ui_extensions.py b/modules/ui_extensions.py index 43158211a..640c0d622 100644 --- a/modules/ui_extensions.py +++ b/modules/ui_extensions.py @@ -11,7 +11,7 @@ from modules import extensions, shared, paths, errors, ui_symbols, call_queue debug = shared.log.debug if os.environ.get('SD_EXT_DEBUG', None) is not None else lambda *args, **kwargs: None extensions_index = "https://vladmandic.github.io/sd-data/pages/extensions.json" hide_tags = ["localization"] -exclude_extensions = ['sdnext-modernui'] +exclude_extensions = ['sdnext-modernui', 'sdnext-kanvas'] extensions_list = [] sort_ordering = { "default": (True, lambda x: x.get('sort_default', '')), @@ -124,11 +124,8 @@ 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): - if url is None: - return "" - url = url.replace(".git", "") - return url +def normalize_git_url(url) -> str: + return '' if url is None else url.removesuffix('.git') def install_extension_from_url(dirname, url, branch_name, search_text, sort_column): @@ -139,19 +136,15 @@ def install_extension_from_url(dirname, url, branch_name, search_text, sort_colu shared.log.error('Extension: url is not specified') return ['', ''] if dirname is None or dirname == "": - *parts, last_part = url.split('/') # pylint: disable=unused-variable - last_part = normalize_git_url(last_part) - dirname = last_part + dirname = normalize_git_url(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 ['', ''] - normalized_url = normalize_git_url(url) - assert len([x for x in extensions.extensions if normalize_git_url(x.remote) == normalized_url]) == 0, 'Extension with this URL is already installed' + 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' tmpdir = os.path.join(paths.data_path, "tmp", dirname) - if url.endswith('.git'): - url = url.replace('.git', '') try: import git shutil.rmtree(tmpdir, True) diff --git a/modules/ui_extra_networks.py b/modules/ui_extra_networks.py index d5d5804cb..7ba05c907 100644 --- a/modules/ui_extra_networks.py +++ b/modules/ui_extra_networks.py @@ -115,10 +115,10 @@ def init_api(): return JSONResponse(obj) shared.api.add_api_route("/sdapi/v1/network", get_network, methods=["GET"]) - shared.api.add_api_route("/sdapi/v1/network/thumb", fetch_file, methods=["GET"]) - shared.api.add_api_route("/sdapi/v1/network/metadata", get_metadata, methods=["GET"]) - shared.api.add_api_route("/sdapi/v1/network/info", get_info, methods=["GET"]) - shared.api.add_api_route("/sdapi/v1/network/desc", get_desc, methods=["GET"]) + shared.api.add_api_route("/sdapi/v1/network/thumb", fetch_file, methods=["GET"], auth=False) + shared.api.add_api_route("/sdapi/v1/network/metadata", get_metadata, methods=["GET"], auth=False) + shared.api.add_api_route("/sdapi/v1/network/info", get_info, methods=["GET"], auth=False) + shared.api.add_api_route("/sdapi/v1/network/desc", get_desc, methods=["GET"], auth=False) class DateTimeEncoder(json.JSONEncoder): @@ -294,6 +294,7 @@ class ExtraNetworksPage: subdirs['Distilled'] = 1 subdirs['Quantized'] = 1 subdirs['Community'] = 1 + subdirs['Cloud'] = 1 subdirs[diffusers_base] = 1 if self.name == 'style' and shared.opts.extra_networks_styles: subdirs['Local'] = 1 @@ -313,11 +314,13 @@ class ExtraNetworksPage: subdirs.move_to_end('Quantized', last=True) if 'Community' in subdirs: subdirs.move_to_end('Community', last=True) + if 'Cloud' in subdirs: + subdirs.move_to_end('Cloud', last=True) subdirs_html = '' for subdir in subdirs: if len(subdir) == 0: continue - if subdir in ['All', 'Local', 'Diffusers', 'Reference', 'Distilled', 'Quantized', 'Community']: + if subdir in ['All', 'Local', 'Diffusers', 'Reference', 'Distilled', 'Quantized', 'Community', 'Cloud']: style = 'network-reference' else: style = 'network-folder' diff --git a/modules/ui_extra_networks_checkpoints.py b/modules/ui_extra_networks_checkpoints.py index 2fe4db376..26d81121a 100644 --- a/modules/ui_extra_networks_checkpoints.py +++ b/modules/ui_extra_networks_checkpoints.py @@ -68,13 +68,6 @@ class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage): else: path = f'{v.get("path", "")}' - ready = reference_downloaded(url) - if not ready and shared.opts.offline_mode: - count['hidden'] += 1 - continue - if ready: - count['ready'] += 1 - tag = v.get('tags', '') if tag in count: count[tag] += 1 @@ -83,6 +76,16 @@ class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage): else: count['base'] += 1 + ready = reference_downloaded(url) + version = "ready" if ready else "download" + if tag == 'cloud': + version = 'cloud' + if not ready and shared.opts.offline_mode: + count['hidden'] += 1 + continue + if ready: + count['ready'] += 1 + yield { "type": 'Model', "name": name, @@ -97,7 +100,7 @@ class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage): "info": {}, "metadata": {}, "description": v.get('desc', ''), - "version": "ready" if ready else "download", + "version": version, "tags": tag, } shared.log.debug(f'Networks: type="reference" items={count}') diff --git a/modules/ui_extra_networks_styles.py b/modules/ui_extra_networks_styles.py index 84d9f7b95..7c9bc40dd 100644 --- a/modules/ui_extra_networks_styles.py +++ b/modules/ui_extra_networks_styles.py @@ -135,5 +135,5 @@ class ExtraNetworkStyles(extra_networks.ExtraNetwork): styles.apply_styles_to_extra(p, style) - def deactivate(self, p): + def deactivate(self, p, force=False): pass diff --git a/modules/ui_img2img.py b/modules/ui_img2img.py index a22e66901..a89046df0 100644 --- a/modules/ui_img2img.py +++ b/modules/ui_img2img.py @@ -281,7 +281,7 @@ def create_ui(): (enable_hr, "Second pass"), (enable_hr, "Refine"), (denoising_strength, "Denoising strength"), - (denoising_strength, "Hires strength"), + (hr_denoising_strength, "Hires strength"), (hr_sampler_index, "Hires sampler"), (hr_resize_mode, "Hires mode"), (hr_resize_context, "Hires context"), diff --git a/modules/ui_javascript.py b/modules/ui_javascript.py index c17e1ad14..49ff5d157 100644 --- a/modules/ui_javascript.py +++ b/modules/ui_javascript.py @@ -20,12 +20,18 @@ def html_head(): skip = ['login.js'] for js in main: script_js = os.path.join(script_path, "javascript", js) - head += f'\n' + if '.esm' in js or '.mjs' in js: + head += f'\n' + else: + head += f'\n' added = [] for script in modules.scripts_manager.list_scripts("javascript", ".js"): if script.filename in main or script.filename in skip: continue - head += f'\n' + if '.esm' in js or '.mjs' in js: + head += f'\n' + else: + head += f'\n' added.append(script.path) for script in modules.scripts_manager.list_scripts("javascript", ".mjs"): head += f'\n' @@ -96,7 +102,6 @@ def reload_javascript(): css_base = theme.reload_gradio_theme() css_timesheet = "timesheet.css" - css = html_css([css_base, css_timesheet]) body = html_body() @@ -111,8 +116,6 @@ def reload_javascript(): for line in lines: if 'meta name="twitter:' in line: res.body = res.body.replace(line.encode("utf8"), b'') - # if 'href="https://fonts.googleapis.com"' in line or 'href="https://fonts.gstatic.com"' in line: - # res.body = res.body.replace(line.encode("utf8"), b'') if 'iframeResizer.contentWindow.min.js' in line: res.body = res.body.replace(line.encode("utf8"), b'src="file=javascript/iframeResizer.min.js"') res.init_headers() diff --git a/modules/ui_postprocessing.py b/modules/ui_postprocessing.py index 5ed90446a..92c80697e 100644 --- a/modules/ui_postprocessing.py +++ b/modules/ui_postprocessing.py @@ -1,4 +1,3 @@ -import json import gradio as gr from modules import scripts_manager, shared, ui_common, postprocessing, call_queue, generation_parameters_copypaste @@ -7,12 +6,15 @@ def submit_info(image): from modules.extras import run_pnginfo from modules.ui_common import infotext_to_html _, geninfo, info = run_pnginfo(image) + if hasattr(scripts_manager, 'scripts_postproc'): + scripts_manager.scripts_postproc.image_changed() return infotext_to_html(geninfo), info, geninfo def submit_process(tab_index, extras_image, image_batch, extras_batch_input_dir, extras_batch_output_dir, show_extras_results, save_output, *script_inputs): - result_images, geninfo, js_info = postprocessing.run_postprocessing(tab_index, extras_image, image_batch, extras_batch_input_dir, extras_batch_output_dir, show_extras_results, *script_inputs, save_output=save_output) - return result_images, geninfo, json.dumps(js_info), '' + from modules.ui_common import infotext_to_html + result_images, geninfo, _js_info = postprocessing.run_postprocessing(tab_index, extras_image, image_batch, extras_batch_input_dir, extras_batch_output_dir, show_extras_results, *script_inputs, save_output=save_output) + return result_images, geninfo, infotext_to_html(geninfo) def create_ui(): @@ -44,22 +46,20 @@ def create_ui(): skip.click(fn=shared.state.skip, inputs=[], outputs=[]) pause = gr.Button('Pause', elem_id=f"{id_part}_pause") pause.click(fn=shared.state.pause, _js='checkPaused', inputs=[], outputs=[]) - result_images, generation_info, html_info, html_info_formatted, html_log = ui_common.create_output_panel("extras") + result_images, generation_info, _html_info, html_info_formatted, _html_log = ui_common.create_output_panel("extras") gr.HTML('File metadata') exif_info = gr.HTML(elem_id="pnginfo_html_info") - gen_info = gr.Textbox(elem_id="pnginfo_gen_info", visible=False) with gr.Row(elem_id='copy_buttons_process'): copy_process_buttons = generation_parameters_copypaste.create_buttons(["txt2img", "img2img", "control", "caption"]) for tabname, button in copy_process_buttons.items(): - generation_parameters_copypaste.register_paste_params_button(generation_parameters_copypaste.ParamBinding(paste_button=button, tabname=tabname, source_text_component=gen_info, source_image_component=extras_image)) + generation_parameters_copypaste.register_paste_params_button(generation_parameters_copypaste.ParamBinding(paste_button=button, tabname=tabname, source_text_component=generation_info, source_image_component=extras_image)) generation_parameters_copypaste.add_paste_fields("extras", extras_image, None) tab_single.select(fn=lambda: 0, inputs=[], outputs=[tab_index]) tab_batch.select(fn=lambda: 1, inputs=[], outputs=[tab_index]) tab_batch_dir.select(fn=lambda: 2, inputs=[], outputs=[tab_index]) - extras_image.change(fn=submit_info, inputs=[extras_image], outputs=[html_info_formatted, exif_info, gen_info]) - extras_image.change(fn=scripts_manager.scripts_postproc.image_changed, inputs=[], outputs=[]) + extras_image.change(fn=submit_info, inputs=[extras_image], outputs=[html_info_formatted, exif_info, generation_info]) submit.click( _js="submit_postprocessing", fn=call_queue.wrap_gradio_gpu_call(submit_process, extra_outputs=[None, ''], name='Postprocess'), @@ -75,8 +75,7 @@ def create_ui(): ], outputs=[ result_images, - html_info, generation_info, - html_log, + html_info_formatted, ] ) diff --git a/modules/ui_sections.py b/modules/ui_sections.py index 7b95ce282..ffa0fbcfb 100644 --- a/modules/ui_sections.py +++ b/modules/ui_sections.py @@ -339,7 +339,7 @@ def create_resize_inputs(tab, images, accordion=True, latent=False, non_zero=Tru with gr.Row(visible=True) as _resize_group: with gr.Column(elem_id=f"{tab}_column_size"): - selected_scale_tab = gr.State(value=0) # pylint: disable=abstract-class-instantiated + selected_scale_tab = gr.State(value=0 if tab != 'img2img' else 1) # pylint: disable=abstract-class-instantiated with gr.Tabs(elem_id=f"{tab}_scale_tabs", selected=0 if non_zero else 1): with gr.Tab(label="Fixed", id=0, elem_id=f"{tab}_scale_tab_fixed") as tab_scale_to: with gr.Row(elem_id=f"{tab}_resize_row_fixed"): diff --git a/modules/ui_txt2img.py b/modules/ui_txt2img.py index f3e108666..8497491c9 100644 --- a/modules/ui_txt2img.py +++ b/modules/ui_txt2img.py @@ -35,7 +35,7 @@ def create_ui(): guidance_name, guidance_scale, guidance_rescale, guidance_start, guidance_stop, cfg_scale, image_cfg_scale, diffusers_guidance_rescale, pag_scale, pag_adaptive, cfg_end = ui_guidance.create_guidance_inputs('txt2img') vae_type, tiling, hidiffusion, clip_skip = ui_sections.create_advanced_inputs('txt2img') hdr_mode, hdr_brightness, hdr_color, hdr_sharpen, hdr_clamp, hdr_boundary, hdr_threshold, hdr_maximize, hdr_max_center, hdr_max_boundary, hdr_color_picker, hdr_tint_ratio = ui_sections.create_correction_inputs('txt2img') - enable_hr, hr_sampler_index, denoising_strength, hr_resize_mode, hr_resize_context, hr_upscaler, hr_force, hr_second_pass_steps, hr_scale, hr_resize_x, hr_resize_y, refiner_steps, refiner_start, refiner_prompt, refiner_negative = ui_sections.create_hires_inputs('txt2img') + enable_hr, hr_sampler_index, hr_denoising_strength, hr_resize_mode, hr_resize_context, hr_upscaler, hr_force, hr_second_pass_steps, hr_scale, hr_resize_x, hr_resize_y, refiner_steps, refiner_start, refiner_prompt, refiner_negative = ui_sections.create_hires_inputs('txt2img') detailer_enabled, detailer_prompt, detailer_negative, detailer_steps, detailer_strength, detailer_resolution = shared.yolo.ui('txt2img') override_settings = ui_common.create_override_inputs('txt2img') state = gr.Textbox(value='', visible=False) @@ -61,7 +61,7 @@ def create_ui(): clip_skip, seed, subseed, subseed_strength, seed_resize_from_h, seed_resize_from_w, height, width, - enable_hr, denoising_strength, + enable_hr, hr_denoising_strength, hr_scale, hr_resize_mode, hr_resize_context, hr_upscaler, hr_force, hr_second_pass_steps, hr_resize_x, hr_resize_y, refiner_steps, refiner_start, refiner_prompt, refiner_negative, hdr_mode, hdr_brightness, hdr_color, hdr_sharpen, hdr_clamp, hdr_boundary, hdr_threshold, hdr_maximize, hdr_max_center, hdr_max_boundary, hdr_color_picker, hdr_tint_ratio, @@ -132,8 +132,7 @@ def create_ui(): # second pass (enable_hr, "Second pass"), (enable_hr, "Refine"), - (denoising_strength, "Denoising strength"), - (denoising_strength, "Hires strength"), + (hr_denoising_strength, "Hires strength"), (hr_sampler_index, "Hires sampler"), (hr_resize_mode, "Hires mode"), (hr_resize_context, "Hires context"), diff --git a/modules/update.py b/modules/update.py index 220258add..d958f3e7d 100644 --- a/modules/update.py +++ b/modules/update.py @@ -19,7 +19,7 @@ def get_version(): origin = origin.splitlines()[0] version.branch = i.git('rev-parse --abbrev-ref HEAD') version.branch = version.branch.splitlines()[0] - version.url = origin + '/tree/' + version.branch + version.url = origin.removesuffix('.git') + '/tree/' + version.branch ver = i.git('log --pretty=format:"%h %ad" -1 --date=short') ver = ver.splitlines()[0] diff --git a/modules/video_models/models_def.py b/modules/video_models/models_def.py index 4eb46da27..bcb62bcf2 100644 --- a/modules/video_models/models_def.py +++ b/modules/video_models/models_def.py @@ -242,6 +242,12 @@ try: repo_cls=getattr(diffusers, 'WanVACEPipeline', None), te_cls=getattr(transformers, 'UMT5EncoderModel', None), dit_cls=getattr(diffusers, 'WanVACETransformer3DModel', None)), + Model(name='WAN 2.2 Animate 14B', + url='https://huggingface.co/Wan-AI/Wan2.2-Animate-14B-Diffusers', + repo='Wan-AI/Wan2.2-Animate-14B-Diffusers', + repo_cls=getattr(diffusers, 'WanAnimatePipeline', None), + te_cls=getattr(transformers, 'UMT5EncoderModel', None), + dit_cls=getattr(diffusers, 'WanAnimateTransformer3DModel', None)), ], 'SkyReels V2': [ Model(name='None'), diff --git a/modules/video_models/video_run.py b/modules/video_models/video_run.py index b0ce3060a..cdb1771d1 100644 --- a/modules/video_models/video_run.py +++ b/modules/video_models/video_run.py @@ -77,6 +77,14 @@ def generate(*args, **kwargs): if init_image is not None: p.task_args['reference_images'] = [images.resize_image(resize_mode=2, im=init_image, width=p.width, height=p.height, upscaler_name=None, output_type='pil')] shared.log.debug(f'Video: op=VACE reference={init_image} resized={p.task_args["reference_images"]}') + elif 'Animate' in model: + if init_image is None: + return video_utils.queue_err('init image not set') + p.task_args['image'] = images.resize_image(resize_mode=2, im=init_image, width=p.width, height=p.height, upscaler_name=None, output_type='pil') + p.task_args['mode'] = 'animate' + p.task_args['pose_video'] = [] # input pose video to condition the generation on. must be a list of PIL images. + p.task_args['face_video'] = [] # input face video to condition the generation on. must be a list of PIL images. + shared.log.debug(f'Video: op=Animate init={p.task_args["image"]} pose={p.task_args["pose_video"]} face={p.task_args["face_video"]}') else: shared.log.warning(f'Video: unknown model type "{model}"') diff --git a/package.json b/package.json index e42d5a2fd..12a76a8c4 100644 --- a/package.json +++ b/package.json @@ -9,7 +9,7 @@ "homepage": "https://github.com/vladmandic/sdnext", "license": "Apache-2.0", "engines": { - "node": ">=14.0.0" + "node": ">=22.0.0" }, "repository": { "type": "git", @@ -20,7 +20,7 @@ "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/ extensions-builtin/sdnext-modernui/javascript/", + "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", diff --git a/pipelines/chrono/__init__.py b/pipelines/chrono/__init__.py deleted file mode 100644 index 902450972..000000000 --- a/pipelines/chrono/__init__.py +++ /dev/null @@ -1,763 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import html -from typing import Any, Callable, Dict, List, Optional, Tuple, Union - -import PIL -import regex as re -import torch -from transformers import AutoTokenizer, CLIPImageProcessor, CLIPVisionModel, UMT5EncoderModel - -from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback -from diffusers.image_processor import PipelineImageInput -from diffusers.loaders import WanLoraLoaderMixin -from diffusers.models import AutoencoderKLWan, WanTransformer3DModel -from diffusers.schedulers import FlowMatchEulerDiscreteScheduler -from diffusers.utils import is_ftfy_available, is_torch_xla_available, logging, replace_example_docstring -from diffusers.utils.torch_utils import randn_tensor -from diffusers.video_processor import VideoProcessor -from diffusers.pipelines.pipeline_utils import DiffusionPipeline -from diffusers.pipelines.wan.pipeline_output import WanPipelineOutput - - -if is_torch_xla_available(): - import torch_xla.core.xla_model as xm - - XLA_AVAILABLE = True -else: - XLA_AVAILABLE = False - -logger = logging.get_logger(__name__) # pylint: disable=invalid-name - -if is_ftfy_available(): - import ftfy - -EXAMPLE_DOC_STRING = """ - Examples: - ```python - >>> import torch - >>> import numpy as np - >>> from diffusers import AutoencoderKLWan, WanImageToVideoPipeline - >>> from diffusers.utils import export_to_video, load_image - >>> from transformers import CLIPVisionModel - - >>> # Available models: Wan-AI/Wan2.1-I2V-14B-480P-Diffusers, Wan-AI/Wan2.1-I2V-14B-720P-Diffusers - >>> model_id = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers" - >>> image_encoder = CLIPVisionModel.from_pretrained( - ... model_id, subfolder="image_encoder", torch_dtype=torch.float32 - ... ) - >>> vae = AutoencoderKLWan.from_pretrained(model_id, subfolder="vae", torch_dtype=torch.float32) - >>> pipe = WanImageToVideoPipeline.from_pretrained( - ... model_id, vae=vae, image_encoder=image_encoder, torch_dtype=torch.bfloat16 - ... ) - >>> pipe.to("cuda") - - >>> image = load_image( - ... "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg" - ... ) - >>> max_area = 480 * 832 - >>> aspect_ratio = image.height / image.width - >>> mod_value = pipe.vae_scale_factor_spatial * pipe.transformer.config.patch_size[1] - >>> height = round(np.sqrt(max_area * aspect_ratio)) // mod_value * mod_value - >>> width = round(np.sqrt(max_area / aspect_ratio)) // mod_value * mod_value - >>> image = image.resize((width, height)) - >>> prompt = ( - ... "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in " - ... "the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot." - ... ) - >>> negative_prompt = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" - - >>> output = pipe( - ... image=image, - ... prompt=prompt, - ... negative_prompt=negative_prompt, - ... height=height, - ... width=width, - ... num_frames=81, - ... guidance_scale=5.0, - ... ).frames[0] - >>> export_to_video(output, "output.mp4", fps=16) - ``` -""" - - -def basic_clean(text): - text = ftfy.fix_text(text) - text = html.unescape(html.unescape(text)) - return text.strip() - - -def whitespace_clean(text): - text = re.sub(r"\s+", " ", text) - text = text.strip() - return text - - -def prompt_clean(text): - text = whitespace_clean(basic_clean(text)) - return text - - -# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_img2img.retrieve_latents -def retrieve_latents( - encoder_output: torch.Tensor, generator: Optional[torch.Generator] = None, sample_mode: str = "sample" -): - if hasattr(encoder_output, "latent_dist") and sample_mode == "sample": - return encoder_output.latent_dist.sample(generator) - elif hasattr(encoder_output, "latent_dist") and sample_mode == "argmax": - return encoder_output.latent_dist.mode() - elif hasattr(encoder_output, "latents"): - return encoder_output.latents - else: - raise AttributeError("Could not access latents of provided encoder_output") - - -class ChronoEditPipeline(DiffusionPipeline, WanLoraLoaderMixin): - r""" - Pipeline for image-to-video generation using Wan. - - This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods - implemented for all pipelines (downloading, saving, running on a particular device, etc.). - - Args: - tokenizer ([`T5Tokenizer`]): - Tokenizer from [T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5Tokenizer), - specifically the [google/umt5-xxl](https://huggingface.co/google/umt5-xxl) variant. - text_encoder ([`T5EncoderModel`]): - [T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5EncoderModel), specifically - the [google/umt5-xxl](https://huggingface.co/google/umt5-xxl) variant. - image_encoder ([`CLIPVisionModel`]): - [CLIP](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPVisionModel), specifically - the - [clip-vit-huge-patch14](https://github.com/mlfoundations/open_clip/blob/main/docs/PRETRAINED.md#vit-h14-xlm-roberta-large) - variant. - transformer ([`WanTransformer3DModel`]): - Conditional Transformer to denoise the input latents. - scheduler ([`UniPCMultistepScheduler`]): - A scheduler to be used in combination with `transformer` to denoise the encoded image latents. - vae ([`AutoencoderKLWan`]): - Variational Auto-Encoder (VAE) Model to encode and decode videos to and from latent representations. - """ - - model_cpu_offload_seq = "text_encoder->image_encoder->transformer->vae" - _callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"] - - def __init__( - self, - tokenizer: AutoTokenizer, - text_encoder: UMT5EncoderModel, - image_encoder: CLIPVisionModel, - image_processor: CLIPImageProcessor, - transformer: WanTransformer3DModel, - vae: AutoencoderKLWan, - scheduler: FlowMatchEulerDiscreteScheduler, - ): - super().__init__() - - self.register_modules( - vae=vae, - text_encoder=text_encoder, - tokenizer=tokenizer, - image_encoder=image_encoder, - transformer=transformer, - scheduler=scheduler, - image_processor=image_processor, - ) - - self.vae_scale_factor_temporal = 2 ** sum(self.vae.temperal_downsample) if getattr(self, "vae", None) else 4 - self.vae_scale_factor_spatial = 2 ** len(self.vae.temperal_downsample) if getattr(self, "vae", None) else 8 - self.video_processor = VideoProcessor(vae_scale_factor=self.vae_scale_factor_spatial) - self.image_processor = image_processor - - def _get_t5_prompt_embeds( - self, - prompt: Union[str, List[str]] = None, - num_videos_per_prompt: int = 1, - max_sequence_length: int = 512, - device: Optional[torch.device] = None, - dtype: Optional[torch.dtype] = None, - ): - device = device or self._execution_device - dtype = dtype or self.text_encoder.dtype - - prompt = [prompt] if isinstance(prompt, str) else prompt - prompt = [prompt_clean(u) for u in prompt] - batch_size = len(prompt) - - text_inputs = self.tokenizer( - prompt, - padding="max_length", - max_length=max_sequence_length, - truncation=True, - add_special_tokens=True, - return_attention_mask=True, - return_tensors="pt", - ) - text_input_ids, mask = text_inputs.input_ids, text_inputs.attention_mask - seq_lens = mask.gt(0).sum(dim=1).long() - - prompt_embeds = self.text_encoder(text_input_ids.to(device), mask.to(device)).last_hidden_state - prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) - prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)] - prompt_embeds = torch.stack( - [torch.cat([u, u.new_zeros(max_sequence_length - u.size(0), u.size(1))]) for u in prompt_embeds], dim=0 - ) - - # duplicate text embeddings for each generation per prompt, using mps friendly method - _, seq_len, _ = prompt_embeds.shape - prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1) - prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1) - - return prompt_embeds - - def encode_image( - self, - image: PipelineImageInput, - device: Optional[torch.device] = None, - ): - device = device or self._execution_device - image = self.image_processor(images=image, return_tensors="pt").to(device) - image_embeds = self.image_encoder(**image, output_hidden_states=True) - return image_embeds.hidden_states[-2] - - # Copied from diffusers.pipelines.wan.pipeline_wan.WanPipeline.encode_prompt - def encode_prompt( - self, - prompt: Union[str, List[str]], - negative_prompt: Optional[Union[str, List[str]]] = None, - do_classifier_free_guidance: bool = True, - num_videos_per_prompt: int = 1, - prompt_embeds: Optional[torch.Tensor] = None, - negative_prompt_embeds: Optional[torch.Tensor] = None, - max_sequence_length: int = 226, - device: Optional[torch.device] = None, - dtype: Optional[torch.dtype] = None, - ): - r""" - Encodes the prompt into text encoder hidden states. - - Args: - prompt (`str` or `List[str]`, *optional*): - prompt to be encoded - negative_prompt (`str` or `List[str]`, *optional*): - The prompt or prompts not to guide the image generation. If not defined, one has to pass - `negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is - less than `1`). - do_classifier_free_guidance (`bool`, *optional*, defaults to `True`): - Whether to use classifier free guidance or not. - num_videos_per_prompt (`int`, *optional*, defaults to 1): - Number of videos that should be generated per prompt. torch device to place the resulting embeddings on - prompt_embeds (`torch.Tensor`, *optional*): - Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not - provided, text embeddings will be generated from `prompt` input argument. - negative_prompt_embeds (`torch.Tensor`, *optional*): - Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt - weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input - argument. - device: (`torch.device`, *optional*): - torch device - dtype: (`torch.dtype`, *optional*): - torch dtype - """ - device = device or self._execution_device - - prompt = [prompt] if isinstance(prompt, str) else prompt - if prompt is not None: - batch_size = len(prompt) - else: - batch_size = prompt_embeds.shape[0] - - if prompt_embeds is None: - prompt_embeds = self._get_t5_prompt_embeds( - prompt=prompt, - num_videos_per_prompt=num_videos_per_prompt, - max_sequence_length=max_sequence_length, - device=device, - dtype=dtype, - ) - - if do_classifier_free_guidance and negative_prompt_embeds is None: - negative_prompt = negative_prompt or "" - negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt - - if prompt is not None and type(prompt) is not type(negative_prompt): - raise TypeError( - f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !=" - f" {type(prompt)}." - ) - elif batch_size != len(negative_prompt): - raise ValueError( - f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:" - f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches" - " the batch size of `prompt`." - ) - - negative_prompt_embeds = self._get_t5_prompt_embeds( - prompt=negative_prompt, - num_videos_per_prompt=num_videos_per_prompt, - max_sequence_length=max_sequence_length, - device=device, - dtype=dtype, - ) - - return prompt_embeds, negative_prompt_embeds - - def check_inputs( - self, - prompt, - negative_prompt, - image, - height, - width, - prompt_embeds=None, - negative_prompt_embeds=None, - image_embeds=None, - callback_on_step_end_tensor_inputs=None, - ): - if image is not None and image_embeds is not None: - raise ValueError( - f"Cannot forward both `image`: {image} and `image_embeds`: {image_embeds}. Please make sure to" - " only forward one of the two." - ) - if image is None and image_embeds is None: - raise ValueError( - "Provide either `image` or `prompt_embeds`. Cannot leave both `image` and `image_embeds` undefined." - ) - if image is not None and not isinstance(image, torch.Tensor) and not isinstance(image, PIL.Image.Image): - raise ValueError(f"`image` has to be of type `torch.Tensor` or `PIL.Image.Image` but is {type(image)}") - if height % 16 != 0 or width % 16 != 0: - raise ValueError(f"`height` and `width` have to be divisible by 16 but are {height} and {width}.") - - if callback_on_step_end_tensor_inputs is not None and not all( - k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs - ): - raise ValueError( - f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}" - ) - - if prompt is not None and prompt_embeds is not None: - raise ValueError( - f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to" - " only forward one of the two." - ) - elif negative_prompt is not None and negative_prompt_embeds is not None: - raise ValueError( - f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`: {negative_prompt_embeds}. Please make sure to" - " only forward one of the two." - ) - elif prompt is None and prompt_embeds is None: - raise ValueError( - "Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined." - ) - elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)): - raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}") - elif negative_prompt is not None and ( - not isinstance(negative_prompt, str) and not isinstance(negative_prompt, list) - ): - raise ValueError(f"`negative_prompt` has to be of type `str` or `list` but is {type(negative_prompt)}") - - def prepare_latents( - self, - image: PipelineImageInput, - batch_size: int, - num_channels_latents: int = 16, - height: int = 480, - width: int = 832, - num_frames: int = 81, - dtype: Optional[torch.dtype] = None, - device: Optional[torch.device] = None, - generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, - latents: Optional[torch.Tensor] = None, - ) -> Tuple[torch.Tensor, torch.Tensor]: - num_latent_frames = (num_frames - 1) // self.vae_scale_factor_temporal + 1 - latent_height = height // self.vae_scale_factor_spatial - latent_width = width // self.vae_scale_factor_spatial - - shape = (batch_size, num_channels_latents, num_latent_frames, latent_height, latent_width) - if isinstance(generator, list) and len(generator) != batch_size: - raise ValueError( - f"You have passed a list of generators of length {len(generator)}, but requested an effective batch" - f" size of {batch_size}. Make sure the batch size matches the length of the generators." - ) - - if latents is None: - latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) - else: - latents = latents.to(device=device, dtype=dtype) - - image = image.unsqueeze(2) - video_condition = torch.cat( - [image, image.new_zeros(image.shape[0], image.shape[1], num_frames - 1, height, width)], dim=2 - ) - video_condition = video_condition.to(device=device, dtype=dtype) - - latents_mean = ( - torch.tensor(self.vae.config.latents_mean) - .view(1, self.vae.config.z_dim, 1, 1, 1) - .to(latents.device, latents.dtype) - ) - latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(1, self.vae.config.z_dim, 1, 1, 1).to( - latents.device, latents.dtype - ) - - if isinstance(generator, list): - latent_condition = [ - retrieve_latents(self.vae.encode(video_condition), sample_mode="argmax") for _ in generator - ] - latent_condition = torch.cat(latent_condition) - else: - latent_condition = retrieve_latents(self.vae.encode(video_condition), sample_mode="argmax") - latent_condition = latent_condition.repeat(batch_size, 1, 1, 1, 1) - - latent_condition = (latent_condition - latents_mean) * latents_std - - mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height, latent_width) - mask_lat_size[:, :, list(range(1, num_frames))] = 0 - first_frame_mask = mask_lat_size[:, :, 0:1] - first_frame_mask = torch.repeat_interleave(first_frame_mask, dim=2, repeats=self.vae_scale_factor_temporal) - mask_lat_size = torch.concat([first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2) - mask_lat_size = mask_lat_size.view(batch_size, -1, self.vae_scale_factor_temporal, latent_height, latent_width) - mask_lat_size = mask_lat_size.transpose(1, 2) - mask_lat_size = mask_lat_size.to(latent_condition.device) - - return latents, torch.concat([mask_lat_size, latent_condition], dim=1) - - @property - def guidance_scale(self): - return self._guidance_scale - - @property - def do_classifier_free_guidance(self): - return self._guidance_scale > 1 - - @property - def num_timesteps(self): - return self._num_timesteps - - @property - def current_timestep(self): - return self._current_timestep - - @property - def interrupt(self): - return self._interrupt - - @property - def attention_kwargs(self): - return self._attention_kwargs - - @torch.no_grad() - @replace_example_docstring(EXAMPLE_DOC_STRING) - def __call__( - self, - image: PipelineImageInput, - prompt: Union[str, List[str]] = None, - negative_prompt: Union[str, List[str]] = None, - height: int = 480, - width: int = 832, - num_frames: int = 81, - num_inference_steps: int = 50, - guidance_scale: float = 5.0, - num_videos_per_prompt: Optional[int] = 1, - generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, - latents: Optional[torch.Tensor] = None, - prompt_embeds: Optional[torch.Tensor] = None, - negative_prompt_embeds: Optional[torch.Tensor] = None, - image_embeds: Optional[torch.Tensor] = None, - output_type: Optional[str] = "np", - return_dict: bool = True, - attention_kwargs: Optional[Dict[str, Any]] = None, - callback_on_step_end: Optional[ - Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks] - ] = None, - callback_on_step_end_tensor_inputs: List[str] = ["latents"], - max_sequence_length: int = 512, - enable_temporal_reasoning: bool = False, - num_temporal_reasoning_steps: int = 0, - offload_model: bool=False - ): - r""" - The call function to the pipeline for generation. - - Args: - image (`PipelineImageInput`): - The input image to condition the generation on. Must be an image, a list of images or a `torch.Tensor`. - prompt (`str` or `List[str]`, *optional*): - The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`. - instead. - negative_prompt (`str` or `List[str]`, *optional*): - The prompt or prompts not to guide the image generation. If not defined, one has to pass - `negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is - less than `1`). - height (`int`, defaults to `480`): - The height of the generated video. - width (`int`, defaults to `832`): - The width of the generated video. - num_frames (`int`, defaults to `81`): - The number of frames in the generated video. - num_inference_steps (`int`, defaults to `50`): - The number of denoising steps. More denoising steps usually lead to a higher quality image at the - expense of slower inference. - guidance_scale (`float`, defaults to `5.0`): - Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598). - `guidance_scale` is defined as `w` of equation 2. of [Imagen - Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale > - 1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`, - usually at the expense of lower image quality. - num_videos_per_prompt (`int`, *optional*, defaults to 1): - The number of images to generate per prompt. - generator (`torch.Generator` or `List[torch.Generator]`, *optional*): - A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make - generation deterministic. - latents (`torch.Tensor`, *optional*): - Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image - generation. Can be used to tweak the same generation with different prompts. If not provided, a latents - tensor is generated by sampling using the supplied random `generator`. - prompt_embeds (`torch.Tensor`, *optional*): - Pre-generated text embeddings. Can be used to easily tweak text inputs (prompt weighting). If not - provided, text embeddings are generated from the `prompt` input argument. - negative_prompt_embeds (`torch.Tensor`, *optional*): - Pre-generated text embeddings. Can be used to easily tweak text inputs (prompt weighting). If not - provided, text embeddings are generated from the `negative_prompt` input argument. - image_embeds (`torch.Tensor`, *optional*): - Pre-generated image embeddings. Can be used to easily tweak image inputs (weighting). If not provided, - image embeddings are generated from the `image` input argument. - output_type (`str`, *optional*, defaults to `"pil"`): - The output format of the generated image. Choose between `PIL.Image` or `np.array`. - return_dict (`bool`, *optional*, defaults to `True`): - Whether or not to return a [`WanPipelineOutput`] instead of a plain tuple. - attention_kwargs (`dict`, *optional*): - A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under - `self.processor` in - [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). - callback_on_step_end (`Callable`, `PipelineCallback`, `MultiPipelineCallbacks`, *optional*): - A function or a subclass of `PipelineCallback` or `MultiPipelineCallbacks` that is called at the end of - each denoising step during the inference. with the following arguments: `callback_on_step_end(self: - DiffusionPipeline, step: int, timestep: int, callback_kwargs: Dict)`. `callback_kwargs` will include a - list of all tensors as specified by `callback_on_step_end_tensor_inputs`. - callback_on_step_end_tensor_inputs (`List`, *optional*): - The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list - will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the - `._callback_tensor_inputs` attribute of your pipeline class. - max_sequence_length (`int`, *optional*, defaults to `512`): - The maximum sequence length of the prompt. - shift (`float`, *optional*, defaults to `5.0`): - The shift of the flow. - autocast_dtype (`torch.dtype`, *optional*, defaults to `torch.bfloat16`): - The dtype to use for the torch.amp.autocast. - Examples: - - Returns: - [`~WanPipelineOutput`] or `tuple`: - If `return_dict` is `True`, [`WanPipelineOutput`] is returned, otherwise a `tuple` is returned where - the first element is a list with the generated images and the second element is a list of `bool`s - indicating whether the corresponding generated image contains "not-safe-for-work" (nsfw) content. - """ - - if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)): - callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs - - # 1. Check inputs. Raise error if not correct - self.check_inputs( - prompt, - negative_prompt, - image, - height, - width, - prompt_embeds, - negative_prompt_embeds, - image_embeds, - callback_on_step_end_tensor_inputs, - ) - - if num_frames % self.vae_scale_factor_temporal != 1: - logger.warning( - f"`num_frames - 1` has to be divisible by {self.vae_scale_factor_temporal}. Rounding to the nearest number." - ) - num_frames = num_frames // self.vae_scale_factor_temporal * self.vae_scale_factor_temporal + 1 - num_frames = max(num_frames, 1) - - self._guidance_scale = guidance_scale - self._attention_kwargs = attention_kwargs - self._current_timestep = None - self._interrupt = False - - device = self._execution_device - - # 2. Define call parameters - if prompt is not None and isinstance(prompt, str): - batch_size = 1 - elif prompt is not None and isinstance(prompt, list): - batch_size = len(prompt) - else: - batch_size = prompt_embeds.shape[0] - - # 3. Encode input prompt - prompt_embeds, negative_prompt_embeds = self.encode_prompt( - prompt=prompt, - negative_prompt=negative_prompt, - do_classifier_free_guidance=self.do_classifier_free_guidance, - num_videos_per_prompt=num_videos_per_prompt, - prompt_embeds=prompt_embeds, - negative_prompt_embeds=negative_prompt_embeds, - max_sequence_length=max_sequence_length, - device=device, - ) - if offload_model: - self.text_encoder.cpu() - # Encode image embedding - transformer_dtype = self.transformer.dtype - prompt_embeds = prompt_embeds.to(transformer_dtype) - if negative_prompt_embeds is not None: - negative_prompt_embeds = negative_prompt_embeds.to(transformer_dtype) - - if image_embeds is None: - image_embeds = self.encode_image(image, device) - image_embeds = image_embeds.repeat(batch_size, 1, 1) - image_embeds = image_embeds.to(transformer_dtype) - - if offload_model: - self.image_encoder.cpu() - - - # 4. Prepare timesteps - self.scheduler.set_timesteps(num_inference_steps, device=device) - timesteps = self.scheduler.timesteps - - # 5. Prepare latent variables - num_channels_latents = self.vae.config.z_dim - image = self.video_processor.preprocess(image, height=height, width=width).to(device, dtype=torch.bfloat16) - latents, condition = self.prepare_latents( - image, - batch_size * num_videos_per_prompt, - num_channels_latents, - height, - width, - num_frames, - torch.bfloat16, - device, - generator, - latents, - ) - - # 6. Denoising loop - num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order - self._num_timesteps = len(timesteps) - - if offload_model: - torch.cuda.empty_cache() - - with self.progress_bar(total=num_inference_steps) as progress_bar: - for i, t in enumerate(timesteps): - - if self.interrupt: - continue - - if enable_temporal_reasoning and i == num_temporal_reasoning_steps: - latents = latents[:, :, [0, -1]] - condition = condition[:, :, [0, -1]] - - for j in range(len(self.scheduler.model_outputs)): - if self.scheduler.model_outputs[j] is not None: - if latents.shape[-3] != self.scheduler.model_outputs[j].shape[-3]: - self.scheduler.model_outputs[j] = self.scheduler.model_outputs[j][:,:,[0, -1]] - if self.scheduler.last_sample is not None: - self.scheduler.last_sample = self.scheduler.last_sample[:, :, [0, -1]] - - self._current_timestep = t - latent_model_input = torch.cat([latents, condition], dim=1).to(transformer_dtype) - timestep = t.expand(latents.shape[0]) - - noise_pred = self.transformer( - hidden_states=latent_model_input, - timestep=timestep, - encoder_hidden_states=prompt_embeds, - encoder_hidden_states_image=image_embeds, - attention_kwargs=attention_kwargs, - return_dict=False, - )[0] - - if offload_model: - torch.cuda.empty_cache() - - if self.do_classifier_free_guidance: - noise_uncond = self.transformer( - hidden_states=latent_model_input, - timestep=timestep, - encoder_hidden_states=negative_prompt_embeds, - encoder_hidden_states_image=image_embeds, - attention_kwargs=attention_kwargs, - return_dict=False, - )[0] - noise_pred = noise_uncond + guidance_scale * (noise_pred - noise_uncond) - - # compute the previous noisy sample x_t -> x_t-1 - latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0] - - if callback_on_step_end is not None: - callback_kwargs = {} - for k in callback_on_step_end_tensor_inputs: - callback_kwargs[k] = locals()[k] - callback_outputs = callback_on_step_end(self, i, t, callback_kwargs) - - latents = callback_outputs.pop("latents", latents) - prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds) - negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds) - - # call the callback, if provided - if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): - progress_bar.update() - - if XLA_AVAILABLE: - xm.mark_step() - - if offload_model: - self.transformer.cpu() - torch.cuda.empty_cache() - - self._current_timestep = None - - if output_type != "latent": - latents = latents.to(self.vae.dtype) - latents_mean = ( - torch.tensor(self.vae.config.latents_mean) - .view(1, self.vae.config.z_dim, 1, 1, 1) - .to(latents.device, latents.dtype) - ) - latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(1, self.vae.config.z_dim, 1, 1, 1).to( - latents.device, latents.dtype - ) - latents = latents / latents_std + latents_mean - - if enable_temporal_reasoning and num_temporal_reasoning_steps > 0: - video_edit = self.vae.decode(latents[:, :, [0, -1]], return_dict=False)[0] - video_reason = self.vae.decode(latents[:, :, :-1], return_dict=False)[0] - video = torch.cat([video_reason, video_edit[:, :, 1:]], dim=2) - else: - video = self.vae.decode(latents, return_dict=False)[0] - - # video = self.vae.decode(latents, return_dict=False)[0] - video = self.video_processor.postprocess_video(video, output_type=output_type) - else: - video = latents - - # Offload all models - self.maybe_free_model_hooks() - - if not return_dict: - return (video,) - - return WanPipelineOutput(frames=video) diff --git a/pipelines/generic.py b/pipelines/generic.py index 896d3c5c2..3ae0a154d 100644 --- a/pipelines/generic.py +++ b/pipelines/generic.py @@ -76,8 +76,11 @@ def load_transformer(repo_id, cls_name, load_config=None, subfolder="transformer if shared.opts.diffusers_offload_mode != 'none' and transformer is not None: sd_models.move_model(transformer, devices.cpu) - if (transformer is not None) and (quant_type is not None) and (quant_args.get('quantization_config', None) is not None): # attach quantization_config - transformer.quantization_config = quant_args.get('quantization_config', None) + if transformer is not None and not hasattr(transformer, 'quantization_config'): # attach quantization_config + if hasattr(transformer, 'config') and hasattr(transformer.config, 'quantization_config'): + transformer.quantization_config = transformer.config.quantization_config + elif (quant_type is not None) and (quant_args.get('quantization_config', None) is not None): + 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:') @@ -209,8 +212,11 @@ def load_text_encoder(repo_id, cls_name, load_config=None, subfolder="text_encod if shared.opts.diffusers_offload_mode != 'none' and text_encoder is not None: sd_models.move_model(text_encoder, devices.cpu) - if (text_encoder is not None) and (quant_type is not None) and (quant_args.get('quantization_config', None) is not None): # attach quantization_config - text_encoder.quantization_config = quant_args.get('quantization_config', None) + if text_encoder is not None and not hasattr(text_encoder, 'quantization_config'): # attach quantization_config + if hasattr(text_encoder, 'config') and hasattr(text_encoder.config, 'quantization_config'): + text_encoder.quantization_config = text_encoder.config.quantization_config + elif (quant_type is not None) and (quant_args.get('quantization_config', None) is not None): + 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:') diff --git a/pipelines/model_auraflow.py b/pipelines/model_auraflow.py index 957875dfa..ee9e1b633 100644 --- a/pipelines/model_auraflow.py +++ b/pipelines/model_auraflow.py @@ -14,7 +14,7 @@ def load_auraflow(checkpoint_info, diffusers_load_config=None): shared.log.debug(f'Load model: type=AuraFlow repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}') transformer = generic.load_transformer(repo_id, cls_name=diffusers.AuraFlowTransformer2DModel, load_config=diffusers_load_config) - text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.UMT5EncoderModel, load_config=diffusers_load_config) # auraflow uses EleutherAI/pile-t5-xl + text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.UMT5EncoderModel, load_config=diffusers_load_config, allow_shared=False) # auraflow uses EleutherAI/pile-t5-xl pipe = diffusers.AuraFlowPipeline.from_pretrained( repo_id, diff --git a/pipelines/model_chrono.py b/pipelines/model_chrono.py index 4522eeee2..2c66df916 100644 --- a/pipelines/model_chrono.py +++ b/pipelines/model_chrono.py @@ -20,12 +20,11 @@ def load_chrono(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=ChronoEdit repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}') - transformer = generic.load_transformer(repo_id, cls_name=diffusers.WanTransformer3DModel, load_config=diffusers_load_config, subfolder="transformer") + transformer = generic.load_transformer(repo_id, cls_name=diffusers.ChronoEditTransformer3DModel, load_config=diffusers_load_config, subfolder="transformer") text_encoder = generic.load_text_encoder(repo_id, cls_name=transformers.UMT5EncoderModel, load_config=diffusers_load_config, subfolder="text_encoder") try: - from pipelines.chrono import ChronoEditPipeline - pipe = ChronoEditPipeline.from_pretrained( + pipe = diffusers.ChronoEditPipeline.from_pretrained( repo_id, transformer=transformer, text_encoder=text_encoder, diff --git a/pipelines/model_google.py b/pipelines/model_google.py new file mode 100644 index 000000000..b08a28222 --- /dev/null +++ b/pipelines/model_google.py @@ -0,0 +1,105 @@ +import io +import os +from PIL import Image +from installer import install, reload, log + + +image_size_buckets = { + '1M': 1024*1024, + '2M': 2048*1024, + '4M': 4096*1024, +} +aspect_ratios_buckets = { + '1:1': 1/1, + '2:3': 2/3, + '3:2': 3/2, + '4:3': 4/3, + '3:4': 3/4, + '4:5': 4/5, + '5:4': 5/4, + '16:9': 16/9, + '9:16': 9/16, + '21:9': 21/9, + '9:21': 9/21, +} + + +def get_size_buckets(width: int, height: int) -> str: + aspect_ratio = width / height + closest_aspect_ratio = min(aspect_ratios_buckets.items(), key=lambda x: abs(x[1] - aspect_ratio))[0] + pixel_count = width * height + closest_size = min(image_size_buckets.items(), key=lambda x: abs(x[1] - pixel_count))[0] + closest_aspect_ratio = min(aspect_ratios_buckets.items(), key=lambda x: abs(x[1] - aspect_ratio))[0] + return closest_size, closest_aspect_ratio + + +class GoogleNanoBananaPipeline(): + def __init__(self, model_name: str): + self.model = model_name + self.client = None + self.config = None + install('google-genai') + install('pydantic==2.11.7', ignore=True, quiet=True) + reload('pydantic', '2.11.7') + log.debug(f'Load model: type=NanoBanana model="{model_name}"') + + def txt2img(self, prompt): + return self.client.models.generate_content( + model=self.model, + config=self.config, + contents=prompt, + ) + + def img2img(self, prompt, image): + from google import genai + image_bytes = io.BytesIO() + image.save(image_bytes, format='JPEG') + return self.client.models.generate_content( + model=self.model, + config=self.config, + contents=[ + genai.types.Part.from_bytes(data=image_bytes.getvalue(), mime_type='image/jpeg'), + prompt, + ], + ) + + def __call__(self, prompt: list[str], width: int, height: int, image: Image.Image = None): + from google import genai + if self.client is None: + api_key = os.getenv("GOOGLE_API_KEY", None) + if api_key is None: + log.error(f'Cloud: model="{self.model}" GOOGLE_API_KEY environment variable not set') + return None + self.client = genai.Client(api_key=api_key, vertexai=False) + + image_size, aspect_ratio = get_size_buckets(width, height) + log.debug(f'Cloud: prompt={prompt} size={image_size} ar={aspect_ratio} image={image} model="{self.model}"') + self.config=genai.types.GenerateContentConfig( + response_modalities=["IMAGE"], + image_config=genai.types.ImageConfig(aspect_ratio=aspect_ratio, image_size=image_size) + ) + + try: + if image is not None: + response = self.img2img(prompt, image) + else: + response = self.txt2img(prompt) + except Exception as e: + log.error(f'Cloud: model="{self.model}" {e}') + return None + + image = None + if getattr(response, 'prompt_feedback', None) is not None: + log.error(f'Cloud: model="{self.model}" {response.prompt_feedback}') + if not hasattr(response, 'candidates') or (response.candidates is None) or (len(response.candidates) == 0): + log.error(f'Cloud: model="{self.model}" no images received') + return None + for part in response.candidates[0].content.parts: + if part.inline_data is not None: + image = Image.open(io.BytesIO(part.inline_data.data)) + return image + + +def load_nanobanana(checkpoint_info, diffusers_load_config): # pylint: disable=unused-argument + pipe = GoogleNanoBananaPipeline(model_name = checkpoint_info.filename) + return pipe diff --git a/pipelines/model_prx.py b/pipelines/model_prx.py new file mode 100644 index 000000000..4bcc251cb --- /dev/null +++ b/pipelines/model_prx.py @@ -0,0 +1,32 @@ +import diffusers +from modules import shared, devices, sd_models, model_quant, sd_hijack_te +from pipelines import generic + + +def load_prx(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=PRX repo="{repo_id}" config={diffusers_load_config} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args}') + + from transformers.models.t5gemma.modeling_t5gemma import T5GemmaEncoder + transformer = generic.load_transformer(repo_id, cls_name=diffusers.PRXTransformer2DModel, load_config=diffusers_load_config) + text_encoder = generic.load_text_encoder(repo_id, cls_name=T5GemmaEncoder, load_config=diffusers_load_config) + + pipe = diffusers.PRXPipeline.from_pretrained( + repo_id, + transformer=transformer, + text_encoder=text_encoder, + cache_dir=shared.opts.diffusers_dir, + **load_args, + ) + + del text_encoder + del transformer + sd_hijack_te.init_hijack(pipe) + + devices.torch_gc() + return pipe diff --git a/requirements.txt b/requirements.txt index e29873aef..7f2aae720 100644 --- a/requirements.txt +++ b/requirements.txt @@ -61,9 +61,7 @@ sentencepiece==0.2.1 # additional blendmodes scipy==1.14.1 -torchdiffeq scikit-image -seam-carving # lint ruff diff --git a/scripts/allegrovideo.py b/scripts/allegrovideo.py deleted file mode 100644 index 4ff136762..000000000 --- a/scripts/allegrovideo.py +++ /dev/null @@ -1,121 +0,0 @@ -import time -import gradio as gr -import transformers -import diffusers -from modules import scripts_manager, processing, shared, images, devices, sd_models, sd_checkpoint, model_quant, timer, sd_hijack_te - - -repo_id = 'rhymes-ai/Allegro' - - -def hijack_decode(*args, **kwargs): - t0 = time.time() - vae: diffusers.AutoencoderKLAllegro = shared.sd_model.vae - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) - res = shared.sd_model.vae.orig_decode(*args, **kwargs) - t1 = time.time() - timer.process.add('vae', t1-t0) - shared.log.debug(f'Video: vae={vae.__class__.__name__} time={t1-t0:.2f}') - return res - - -class Script(scripts_manager.Script): - def title(self): - return 'Video: Allegro (Legacy)' - - def show(self, is_img2img): - return not is_img2img - - # return signature is array of gradio components - def ui(self, is_img2img): - with gr.Row(): - gr.HTML('  Allegro Video
') - with gr.Row(): - num_frames = gr.Slider(label='Frames', minimum=4, maximum=88, step=1, value=22) - with gr.Row(): - override_scheduler = gr.Checkbox(label='Override scheduler', value=True) - with gr.Row(): - from modules.ui_sections import create_video_inputs - video_type, duration, gif_loop, mp4_pad, mp4_interpolate = create_video_inputs(tab='img2img' if is_img2img else 'txt2img') - return [num_frames, override_scheduler, video_type, duration, gif_loop, mp4_pad, mp4_interpolate] - - def run(self, p: processing.StableDiffusionProcessing, num_frames, override_scheduler, video_type, duration, gif_loop, mp4_pad, mp4_interpolate): # pylint: disable=arguments-differ, unused-argument - # set params - num_frames = int(num_frames) - p.width = 8 * int(p.width // 8) - p.height = 8 * int(p.height // 8) - p.do_not_save_grid = True - p.ops.append('video') - - # load model - if shared.sd_model.__class__ != diffusers.AllegroPipeline: - sd_models.unload_model_weights() - t0 = time.time() - quant_args = model_quant.create_config() - transformer = diffusers.AllegroTransformer3DModel.from_pretrained( - repo_id, - subfolder="transformer", - torch_dtype=devices.dtype, - cache_dir=shared.opts.hfcache_dir, - **quant_args - ) - shared.log.debug(f'Video: module={transformer.__class__.__name__}') - text_encoder = transformers.T5EncoderModel.from_pretrained( - repo_id, - subfolder="text_encoder", - cache_dir=shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - **quant_args - ) - shared.log.debug(f'Video: module={text_encoder.__class__.__name__}') - shared.sd_model = diffusers.AllegroPipeline.from_pretrained( - repo_id, - # transformer=transformer, - # text_encoder=text_encoder, - cache_dir=shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - **quant_args - ) - t1 = time.time() - shared.log.debug(f'Video: load cls={shared.sd_model.__class__.__name__} repo="{repo_id}" dtype={devices.dtype} time={t1-t0:.2f}') - sd_models.set_diffuser_options(shared.sd_model) - shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(repo_id) - shared.sd_model.sd_model_hash = None - shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode - shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt - shared.sd_model.vae.decode = hijack_decode - shared.sd_model.vae.enable_tiling() - # shared.sd_model.vae.enable_slicing() - sd_hijack_te.init_hijack(shared.sd_model) - - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - devices.torch_gc(force=True) - - processing.fix_seed(p) - if override_scheduler: - p.sampler_name = 'Default' - p.steps = 100 - p.task_args['num_frames'] = num_frames - p.task_args['output_type'] = 'pil' - p.task_args['clean_caption'] = False - - p.all_prompts, p.all_negative_prompts = shared.prompt_styles.apply_styles_to_prompts([p.prompt], [p.negative_prompt], p.styles, [p.seed]) - p.task_args['prompt'] = p.all_prompts[0] - p.task_args['negative_prompt'] = p.all_negative_prompts[0] - - # w = shared.sd_model.transformer.config.sample_width * shared.sd_model.vae_scale_factor_spatial - # h = shared.sd_model.transformer.config.sample_height * shared.sd_model.vae_scale_factor_spatial - # n = shared.sd_model.transformer.config.sample_frames * shared.sd_model.vae_scale_factor_temporal - - # run processing - t0 = time.time() - shared.state.disable_preview = True - shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={num_frames}') - processed = processing.process_images(p) - shared.state.disable_preview = False - t1 = time.time() - if processed is not None and len(processed.images) > 0: - shared.log.info(f'Video: frames={len(processed.images)} time={t1-t0:.2f}') - if video_type != 'None': - images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=duration, loop=gif_loop, pad=mp4_pad, interpolate=mp4_interpolate) - return processed diff --git a/scripts/animatediff.py b/scripts/animatediff.py index 2d5f1bb9c..933d54049 100644 --- a/scripts/animatediff.py +++ b/scripts/animatediff.py @@ -82,6 +82,14 @@ def set_adapter(adapter_name: str = 'None'): sd_models.set_diffuser_options(motion_adapter, vae=None, op='adapter') loaded_adapter = adapter_name new_pipe = None + + if 'Model' in shared.opts.sdnq_quantize_weights: + shared.log.debug(f'AnimateDiff: sdnq={shared.opts.sdnq_quantize_weights} reloading model weights') + prev_opts = shared.opts.sdnq_quantize_weights + shared.opts.sdnq_quantize_weights = [] + sd_models.reload_model_weights(force=True) + shared.opts.sdnq_quantize_weights = prev_opts + if shared.sd_model_type == 'sd': new_pipe = diffusers.AnimateDiffPipeline( vae=shared.sd_model.vae, @@ -106,6 +114,7 @@ def set_adapter(adapter_name: str = 'None'): image_encoder=getattr(shared.sd_model, 'image_encoder', None), motion_adapter=motion_adapter, ) + if new_pipe is None: motion_adapter = None loaded_adapter = None @@ -122,6 +131,8 @@ def set_adapter(adapter_name: str = 'None'): motion_adapter = None loaded_adapter = None shared.log.error(f'AnimateDiff load error: adapter="{adapter_name}" {e}') + from modules import errors + errors.display('e', 'AnimateDiff') def set_scheduler(p, model, override: bool = False): diff --git a/scripts/cogvideo.py b/scripts/cogvideo.py deleted file mode 100644 index 1e7fbd076..000000000 --- a/scripts/cogvideo.py +++ /dev/null @@ -1,215 +0,0 @@ -""" -models: https://huggingface.co/THUDM/CogVideoX-2b https://huggingface.co/THUDM/CogVideoX-5b -source: https://github.com/THUDM/CogVideo -quanto: https://gist.github.com/a-r-r-o-w/31be62828b00a9292821b85c1017effa -torchao: https://gist.github.com/a-r-r-o-w/4d9732d17412888c885480c6521a9897 -venhancer: https://github.com/THUDM/CogVideo/blob/dcb82ae30b454ab898aeced0633172d75dbd55b8/tools/venhancer/README.md -""" -import os -import time -import cv2 -import gradio as gr -import torch -from torchvision import transforms -import diffusers -import numpy as np -from modules import scripts_manager, shared, devices, errors, sd_models, processing -from modules.processing_callbacks import diffusers_callback, set_callbacks_p - - -debug = (os.environ.get('SD_LOAD_DEBUG', None) is not None) or (os.environ.get('SD_PROCESS_DEBUG', None) is not None) - - -class Script(scripts_manager.Script): - def title(self): - return 'Video: CogVideoX (Legacy)' - - def show(self, is_img2img): - return True - - - def ui(self, is_img2img): - with gr.Row(): - gr.HTML("  CogVideoX
") - with gr.Row(): - model = gr.Dropdown(label='Model', choices=['None', 'THUDM/CogVideoX-2b', 'THUDM/CogVideoX-5b', 'THUDM/CogVideoX-5b-I2V'], value='THUDM/CogVideoX-2b') - sampler = gr.Dropdown(label='Sampler', choices=['DDIM', 'DPM'], value='DDIM') - with gr.Row(): - frames = gr.Slider(label='Frames', minimum=1, maximum=100, step=1, value=49) - guidance = gr.Slider(label='Guidance', minimum=0.0, maximum=14.0, step=0.5, value=6.0) - with gr.Row(): - offload = gr.Dropdown(label='Offload', choices=['none', 'balanced', 'model', 'sequential'], value='balanced') - override = gr.Checkbox(label='Override resolution', value=True) - with gr.Accordion('Optional init image or video', open=False): - with gr.Row(): - image = gr.Image(value=None, label='Image', type='pil', width=256, height=256) - video = gr.Video(value=None, label='Video', width=256, height=256) - with gr.Row(): - from modules.ui_sections import create_video_inputs - video_type, duration, loop, pad, interpolate = create_video_inputs(tab='img2img' if is_img2img else 'txt2img') - return [model, sampler, frames, guidance, offload, override, video_type, duration, loop, pad, interpolate, image, video] - - def load(self, model): - if (shared.sd_model_type != 'cogvideo' or shared.sd_model.sd_model_checkpoint != model) and model != 'None': - sd_models.unload_model_weights('model') - shared.log.info(f'CogVideoX load: model="{model}"') - try: - shared.sd_model = None - cls = diffusers.CogVideoXImageToVideoPipeline if 'I2V' in model else diffusers.CogVideoXPipeline - shared.sd_model = cls.from_pretrained(model, torch_dtype=devices.dtype, cache_dir=shared.opts.diffusers_dir) - shared.sd_model.sd_checkpoint_info = sd_models.CheckpointInfo(model) - shared.sd_model.sd_model_hash = '' - shared.sd_model.sd_model_checkpoint = model - except Exception as e: - shared.log.error(f'Load CogVideoX: {e}') - if debug: - errors.display(e, 'CogVideoX') - if shared.sd_model_type == 'cogvideo' and model != 'None': - shared.sd_model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m', ncols=80, colour='#327fba') - shared.log.debug(f'CogVideoX load: class="{shared.sd_model.__class__.__name__}"') - if shared.sd_model is not None and model == 'None': - shared.log.info(f'CogVideoX unload: model={model}') - shared.sd_model = None - devices.torch_gc(force=True) - devices.torch_gc() - - def offload(self, offload): - if shared.sd_model_type != 'cogvideo': - return - if offload == 'none': - sd_models.move_model(shared.sd_model, devices.device) - shared.log.debug(f'CogVideoX: offload={offload}') - if offload == 'balanced': - sd_models.apply_balanced_offload(shared.sd_model) - if offload == 'model': - shared.sd_model.enable_model_cpu_offload() - if offload == 'sequential': - shared.sd_model.enable_model_cpu_offload() - shared.sd_model.enable_sequential_cpu_offload() - shared.sd_model.vae.enable_slicing() - shared.sd_model.vae.enable_tiling() - - def video(self, p, fn): - frames = [] - try: - from modules.control.util import decode_fourcc - video = cv2.VideoCapture(fn) - if not video.isOpened(): - shared.log.error(f'Video: file="{fn}" open failed') - return frames - frame_count = int(video.get(cv2.CAP_PROP_FRAME_COUNT)) - fps = int(video.get(cv2.CAP_PROP_FPS)) - w, h = int(video.get(cv2.CAP_PROP_FRAME_WIDTH)), int(video.get(cv2.CAP_PROP_FRAME_HEIGHT)) - codec = decode_fourcc(video.get(cv2.CAP_PROP_FOURCC)) - shared.log.debug(f'CogVideoX input: video="{fn}" fps={fps} width={w} height={h} codec={codec} frames={frame_count} target={len(frames)}') - frames = [] - while True: - ok, frame = video.read() - if not ok: - break - frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) - frame = cv2.resize(frame, (p.width, p.height)) - frames.append(frame) - video.release() - if len(frames) > p.frames: - frames = np.asarray(frames) - indices = np.linspace(0, len(frames) - 1, p.frames).astype(int) # reduce array from n_frames to p_frames - frames = frames[indices] - shared.log.debug(f'CogVideoX input reduce: source={len(frames)} target={p.frames}') - frames = [transforms.ToTensor()(frame) for frame in frames] - except Exception as e: - shared.log.error(f'Video: file="{fn}" {e}') - if debug: - errors.display(e, 'CogVideoX') - return frames - - def image(self, p, img): - img = img.resize((p.width, p.height)) - shared.log.debug(f'CogVideoX input: image={img}') - # frames = [np.array(img)] - # frames = [transforms.ToTensor()(frame) for frame in frames] - return img - - def generate(self, p: processing.StableDiffusionProcessing, model: str): - if shared.sd_model_type != 'cogvideo': - return [] - shared.log.info(f'CogVideoX: sampler={p.sampler} steps={p.steps} frames={p.frames} width={p.width} height={p.height} seed={p.seed} guidance={p.guidance}') - if p.sampler == 'DDIM': - shared.sd_model.scheduler = diffusers.CogVideoXDDIMScheduler.from_config(shared.sd_model.scheduler.config, timestep_spacing="trailing") - if p.sampler == 'DPM': - shared.sd_model.scheduler = diffusers.CogVideoXDPMScheduler.from_config(shared.sd_model.scheduler.config, timestep_spacing="trailing") - t0 = time.time() - frames = [] - set_callbacks_p(p) - shared.state.job_count = 1 - shared.state.sampling_steps = p.steps - 1 - try: - args = dict( - prompt=p.prompt, - negative_prompt=p.negative_prompt, - height=p.height, - width=p.width, - num_videos_per_prompt=1, - num_inference_steps=p.steps, - guidance_scale=p.guidance, - generator=torch.Generator(device=devices.device).manual_seed(p.seed), - callback_on_step_end=diffusers_callback, - callback_on_step_end_tensor_inputs=['latents'], - ) - if 'I2V' in model: - if hasattr(p, 'video') and p.video is not None: - args['video'] = self.video(p, p.video) - shared.sd_model = sd_models.switch_pipe(diffusers.CogVideoXVideoToVideoPipeline, shared.sd_model) - elif (hasattr(p, 'image') and p.image is not None) or (hasattr(p, 'init_images') and len(p.init_images) > 0): - p.init_images = [p.image] if hasattr(p, 'image') and p.image is not None else p.init_images - args['image'] = self.image(p, p.init_images[0]) - shared.sd_model = sd_models.switch_pipe(diffusers.CogVideoXImageToVideoPipeline, shared.sd_model) - else: - shared.sd_model = sd_models.switch_pipe(diffusers.CogVideoXPipeline, shared.sd_model) - args['num_frames'] = p.frames # only txt2vid has num_frames - shared.log.info(f"CogVideoX: class={shared.sd_model.__class__.__name__} frames={p.frames} input={args.get('video', None) or args.get('image', None)}") - if debug: - shared.log.debug(f'CogVideoX args: {args}') - frames = shared.sd_model(**args).frames[0] - except AssertionError as e: - shared.log.info(f'CogVideoX: {e}') - except Exception as e: - shared.log.error(f'CogVideoX: {e}') - if debug: - errors.display(e, 'CogVideoX') - t1 = time.time() - its = (len(frames) * p.steps) / (t1 - t0) - shared.log.info(f'CogVideoX: frame={frames[0] if len(frames) > 0 else None} frames={len(frames)} its={its:.2f} time={t1 - t0:.2f}') - return frames - - # auto-executed by the script-callback - def run(self, p: processing.StableDiffusionProcessing, model, sampler, frames, guidance, offload, override, video_type, duration, loop, pad, interpolate, image, video): # pylint: disable=arguments-differ, unused-argument - processing.fix_seed(p) - p.extra_generation_params['CogVideoX'] = model - p.do_not_save_grid = True - if 'animatediff' not in p.ops: - p.ops.append('video') - if override: - p.width = 720 - p.height = 480 - p.sampler = sampler - p.guidance = guidance - p.frames = frames - p.use_dynamic_cfg = sampler == 'DPM' - p.prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) - p.negative_prompt = shared.prompt_styles.apply_negative_styles_to_prompt(p.negative_prompt, p.styles) - p.image = image - p.video = video - self.load(model) - self.offload(offload) - frames = self.generate(p, model) - devices.torch_gc() - processed = processing.get_processed(p, images_list=frames) - return processed - - # auto-executed by the script-callback - def after(self, p: processing.StableDiffusionProcessing, processed: processing.Processed, model, sampler, frames, guidance, offload, override, video_type, duration, loop, pad, interpolate, image, video): # pylint: disable=arguments-differ, unused-argument - if video_type != 'None' and processed is not None and len(processed.images) > 0: - from modules.images import save_video - shared.log.info(f'CogVideoX video: type={video_type} duration={duration} loop={loop} pad={pad} interpolate={interpolate}') - save_video(p, filename=None, images=processed.images, video_type=video_type, duration=duration, loop=loop, pad=pad, interpolate=interpolate) diff --git a/scripts/hunyuanvideo.py b/scripts/hunyuanvideo.py deleted file mode 100644 index b43c90b79..000000000 --- a/scripts/hunyuanvideo.py +++ /dev/null @@ -1,176 +0,0 @@ -import time -import torch -import gradio as gr -import transformers -import diffusers -from modules import scripts_manager, processing, shared, images, devices, sd_models, sd_checkpoint, sd_samplers, model_quant, timer, sd_hijack_te - - -default_template = """Describe the video by detailing the following aspects: -1. The main content and theme of the video. -2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects. -3. Actions, events, behaviors temporal relationships, physical movement changes of the objects. -4. Background environment, light, style and atmosphere. -5. Camera angles, movements, and transitions used in the video. -6. Thematic and aesthetic concepts associated with the scene, i.e. realistic, futuristic, fairy tale, etc. -""" - -models = { - 'HunyuanVideo': { 'repo': 'tencent/HunyuanVideo', 'revision': 'refs/pr/18' }, - 'FastHunyuan': { 'repo': 'FastVideo/FastHunyuan', 'revision': None }, -} -loaded_model = None - - -def get_template(template: str = None): - # diffusers.pipelines.hunyuan_video.pipeline_hunyuan_video.DEFAULT_PROMPT_TEMPLATE - base_template_pre = "<|start_header_id|>system<|end_header_id|>\n\n" - base_template_post = "<|eot_id|>\n" - base_template_end = "<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>" - if template is None or len(template) == 0: - template = default_template - template_lines = '\n'.join([line for line in template.split('\n') if len(line) > 0]) - prompt_template = { - "crop_start": 95, - "template": base_template_pre + template_lines + base_template_post + base_template_end - } - return prompt_template - - -def hijack_decode(*args, **kwargs): - t0 = time.time() - vae: diffusers.AutoencoderKLHunyuanVideo = shared.sd_model.vae - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) - res = shared.sd_model.vae.orig_decode(*args, **kwargs) - t1 = time.time() - timer.process.add('vae', t1-t0) - shared.log.debug(f'Video: vae={vae.__class__.__name__} tile={vae.tile_sample_min_width}:{vae.tile_sample_min_height}:{vae.tile_sample_min_num_frames} stride={vae.tile_sample_stride_width}:{vae.tile_sample_stride_height}:{vae.tile_sample_stride_num_frames} time={t1-t0:.2f}') - return res - - -class Script(scripts_manager.Script): - def title(self): - return 'Video: Hunyuan Video (Legacy)' - - def show(self, is_img2img): - return not is_img2img - - # return signature is array of gradio components - def ui(self, is_img2img): - with gr.Row(): - gr.HTML('  Hunyuan Video
') - with gr.Row(): - model = gr.Dropdown(label='Model', choices=list(models.keys()), value=list(models.keys())[0]) - with gr.Row(): - num_frames = gr.Slider(label='Frames', minimum=9, maximum=257, step=1, value=45) - tile_frames = gr.Slider(label='Tile frames', minimum=1, maximum=64, step=1, value=16) - with gr.Row(): - with gr.Column(): - override_scheduler = gr.Checkbox(label='HV override sampler', value=True) - with gr.Column(): - scheduler_shift = gr.Slider(label='HV sampler shift', minimum=0.0, maximum=20.0, step=0.1, value=7.0) - with gr.Row(): - template = gr.TextArea(label='HV prompt processor', lines=3, value=default_template, visible=False) - with gr.Row(): - from modules.ui_sections import create_video_inputs - video_type, duration, gif_loop, mp4_pad, mp4_interpolate = create_video_inputs(tab='img2img' if is_img2img else 'txt2img') - return [model, num_frames, tile_frames, override_scheduler, scheduler_shift, template, video_type, duration, gif_loop, mp4_pad, mp4_interpolate] - - def load(self, model:str): - global loaded_model # pylint: disable=global-statement - if shared.sd_model.__class__ != diffusers.HunyuanVideoPipeline or model != loaded_model: - sd_models.unload_model_weights() - t0 = time.time() - quant_args = model_quant.create_config() - transformer = diffusers.HunyuanVideoTransformer3DModel.from_pretrained( - pretrained_model_name_or_path='tencent/HunyuanVideo', - subfolder="transformer", - torch_dtype=devices.dtype, - revision='refs/pr/18', - cache_dir=shared.opts.hfcache_dir, - **quant_args - ) - shared.log.debug(f'Video: module={transformer.__class__.__name__}') - text_encoder = transformers.LlamaModel.from_pretrained( - pretrained_model_name_or_path=models.get(model)['repo'], - subfolder="text_encoder", - revision=models.get(model)['revision'], - cache_dir = shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - **quant_args - ) - text_encoder_2 = transformers.CLIPTextModel.from_pretrained( - pretrained_model_name_or_path=models.get(model)['repo'], - subfolder="text_encoder_2", - revision=models.get(model)['revision'], - cache_dir = shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - ) - shared.log.debug(f'Video: module={text_encoder.__class__.__name__}') - shared.sd_model = diffusers.HunyuanVideoPipeline.from_pretrained( - pretrained_model_name_or_path='tencent/HunyuanVideo', - transformer=transformer, - text_encoder=text_encoder, - text_encoder_2=text_encoder_2, - revision='refs/pr/18', - cache_dir = shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - **quant_args - ) - t1 = time.time() - shared.log.debug(f'Video: load cls={shared.sd_model.__class__.__name__} model="{model}" repo={models.get(model)["repo"]} dtype={devices.dtype} time={t1-t0:.2f}') - sd_models.set_diffuser_options(shared.sd_model) - shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(models.get(model)['repo']) - shared.sd_model.sd_model_hash = None - shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode - shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt - shared.sd_model.vae.decode = hijack_decode - shared.sd_model.vae.enable_slicing() - shared.sd_model.vae.enable_tiling() - shared.sd_model.vae.use_framewise_decoding = True - sd_hijack_te.init_hijack(shared.sd_model) - loaded_model = model - - def run(self, p: processing.StableDiffusionProcessing, model, num_frames, tile_frames, override_scheduler, scheduler_shift, template, video_type, duration, gif_loop, mp4_pad, mp4_interpolate): # pylint: disable=arguments-differ, unused-argument - # set params - num_frames = int(num_frames) - p.width = 16 * int(p.width // 16) - p.height = 16 * int(p.height // 16) - p.do_not_save_grid = True - p.ops.append('video') - - # load model - self.load(model) - - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - devices.torch_gc(force=True) - - if override_scheduler: - p.sampler_name = 'Default' - else: - shared.sd_model.scheduler = sd_samplers.create_sampler(p.sampler_name, shared.sd_model) - p.sampler_name = 'Default' # avoid double creation - if hasattr(shared.sd_model.scheduler, '_shift'): - shared.sd_model.scheduler._shift = scheduler_shift # pylint: disable=protected-access - - # encode prompt - processing.fix_seed(p) - p.task_args['num_frames'] = num_frames - p.task_args['output_type'] = 'pil' - p.task_args['generator'] = torch.manual_seed(p.seed) - # p.task_args['prompt'] = None - # p.task_args['prompt_embeds'], p.task_args['pooled_prompt_embeds'], p.task_args['prompt_attention_mask'] = shared.sd_model.encode_prompt(prompt=p.prompt, prompt_template=get_template(template), device=devices.device) - - # run processing - t0 = time.time() - shared.sd_model.vae.tile_sample_min_num_frames = tile_frames - shared.state.disable_preview = True - shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={num_frames}') - processed = processing.process_images(p) - shared.state.disable_preview = False - t1 = time.time() - if processed is not None and len(processed.images) > 0: - shared.log.info(f'Video: frames={len(processed.images)} time={t1-t0:.2f}') - if video_type != 'None': - images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=duration, loop=gif_loop, pad=mp4_pad, interpolate=mp4_interpolate) - return processed diff --git a/scripts/legacy_allegrovideo.py b/scripts/legacy_allegrovideo.py deleted file mode 100644 index 84c357589..000000000 --- a/scripts/legacy_allegrovideo.py +++ /dev/null @@ -1,121 +0,0 @@ -import time -import gradio as gr -import transformers -import diffusers -from modules import scripts_manager, processing, shared, images, devices, sd_models, sd_checkpoint, model_quant, timer, sd_hijack_te - - -repo_id = 'rhymes-ai/Allegro' - - -def hijack_decode(*args, **kwargs): - t0 = time.time() - vae: diffusers.AutoencoderKLAllegro = shared.sd_model.vae - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) - res = shared.sd_model.vae.orig_decode(*args, **kwargs) - t1 = time.time() - timer.process.add('vae', t1-t0) - shared.log.debug(f'Video: vae={vae.__class__.__name__} time={t1-t0:.2f}') - return res - - -class Script(scripts_manager.Script): - def title(self): - return 'Video: Allegro (Legacy)' - - def show(self, is_img2img): - return not is_img2img - - # return signature is array of gradio components - def ui(self, is_img2img): - with gr.Row(): - gr.HTML('  Allegro Video
') - with gr.Row(): - num_frames = gr.Slider(label='Frames', minimum=4, maximum=88, step=1, value=22) - with gr.Row(): - override_scheduler = gr.Checkbox(label='Override scheduler', value=True) - with gr.Row(): - from modules.ui_sections import create_video_inputs - video_type, duration, gif_loop, mp4_pad, mp4_interpolate = create_video_inputs(tab='img2img' if is_img2img else 'txt2img') - return [num_frames, override_scheduler, video_type, duration, gif_loop, mp4_pad, mp4_interpolate] - - def run(self, p: processing.StableDiffusionProcessing, num_frames, override_scheduler, video_type, duration, gif_loop, mp4_pad, mp4_interpolate): # pylint: disable=arguments-differ, unused-argument - # set params - num_frames = int(num_frames) - p.width = 8 * int(p.width // 8) - p.height = 8 * int(p.height // 8) - p.do_not_save_grid = True - p.ops.append('video') - - # load model - if shared.sd_model.__class__ != diffusers.AllegroPipeline: - sd_models.unload_model_weights() - t0 = time.time() - quant_args = model_quant.create_config() - transformer = diffusers.AllegroTransformer3DModel.from_pretrained( - repo_id, - subfolder="transformer", - torch_dtype=devices.dtype, - cache_dir=shared.opts.hfcache_dir, - **quant_args - ) - shared.log.debug(f'Video: module={transformer.__class__.__name__}') - text_encoder = transformers.T5EncoderModel.from_pretrained( - repo_id, - subfolder="text_encoder", - cache_dir=shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - **quant_args - ) - shared.log.debug(f'Video: module={text_encoder.__class__.__name__}') - shared.sd_model = diffusers.AllegroPipeline.from_pretrained( - repo_id, - # transformer=transformer, - # text_encoder=text_encoder, - cache_dir=shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - **quant_args - ) - t1 = time.time() - shared.log.debug(f'Video: load cls={shared.sd_model.__class__.__name__} repo="{repo_id}" dtype={devices.dtype} time={t1-t0:.2f}') - sd_models.set_diffuser_options(shared.sd_model) - shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(repo_id) - shared.sd_model.sd_model_hash = None - shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode - shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt - shared.sd_model.vae.decode = hijack_decode - shared.sd_model.vae.enable_tiling() - sd_hijack_te.init_hijack(shared.sd_model) - # shared.sd_model.vae.enable_slicing() - - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - devices.torch_gc(force=True) - - processing.fix_seed(p) - if override_scheduler: - p.sampler_name = 'Default' - p.steps = 100 - p.task_args['num_frames'] = num_frames - p.task_args['output_type'] = 'pil' - p.task_args['clean_caption'] = False - - p.all_prompts, p.all_negative_prompts = shared.prompt_styles.apply_styles_to_prompts([p.prompt], [p.negative_prompt], p.styles, [p.seed]) - p.task_args['prompt'] = p.all_prompts[0] - p.task_args['negative_prompt'] = p.all_negative_prompts[0] - - # w = shared.sd_model.transformer.config.sample_width * shared.sd_model.vae_scale_factor_spatial - # h = shared.sd_model.transformer.config.sample_height * shared.sd_model.vae_scale_factor_spatial - # n = shared.sd_model.transformer.config.sample_frames * shared.sd_model.vae_scale_factor_temporal - - # run processing - t0 = time.time() - shared.state.disable_preview = True - shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} width={p.width} height={p.height} frames={num_frames}') - processed = processing.process_images(p) - shared.state.disable_preview = False - t1 = time.time() - if processed is not None and len(processed.images) > 0: - shared.log.info(f'Video: frames={len(processed.images)} time={t1-t0:.2f}') - if video_type != 'None': - images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=duration, loop=gif_loop, pad=mp4_pad, interpolate=mp4_interpolate) - return processed diff --git a/scripts/ltxvideo.py b/scripts/ltxvideo.py deleted file mode 100644 index baaa150aa..000000000 --- a/scripts/ltxvideo.py +++ /dev/null @@ -1,152 +0,0 @@ -import os -import time -import torch -import gradio as gr -import diffusers -import transformers -from modules import scripts_manager, processing, shared, images, devices, sd_models, sd_checkpoint, model_quant, timer, sd_hijack_te - - -repos = { - '0.9.0': 'a-r-r-o-w/LTX-Video-diffusers', - '0.9.1': 'a-r-r-o-w/LTX-Video-0.9.1-diffusers', - '0.9.5': 'Lightricks/LTX-Video-0.9.5', - 'custom': None, -} - - -def load_quants(kwargs, repo_id): - quant_args = model_quant.create_config() - if not quant_args: - return kwargs - model_quant.load_bnb(f'Load model: type=LTX quant={quant_args}') - if 'transformer' not in kwargs and ('Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization): - kwargs['transformer'] = diffusers.LTXVideoTransformer3DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=shared.opts.hfcache_dir, torch_dtype=devices.dtype, **quant_args) - shared.log.debug(f'Quantization: module=transformer type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') - if 'text_encoder' not in kwargs and ('TE' in shared.opts.bnb_quantization or 'TE' in shared.opts.torchao_quantization): - kwargs['text_encoder'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder", cache_dir=shared.opts.hfcache_dir, torch_dtype=devices.dtype, **quant_args) - shared.log.debug(f'Quantization: module=t5 type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') - return kwargs - - -def hijack_decode(*args, **kwargs): - t0 = time.time() - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae']) - res = shared.sd_model.vae.orig_decode(*args, **kwargs) - t1 = time.time() - timer.process.add('vae', t1-t0) - shared.log.debug(f'Video: vae={shared.sd_model.vae.__class__.__name__} time={t1-t0:.2f}') - return res - - -class Script(scripts_manager.Script): - def title(self): - return 'Video: LTX Video (Legacy)' - - def show(self, is_img2img): - return True - - # return signature is array of gradio components - def ui(self, is_img2img): - def model_change(model): - return gr.update(visible=model == 'custom') - - with gr.Row(): - gr.HTML('  LTX Video
') - with gr.Row(): - model = gr.Dropdown(label='LTX Model', choices=list(repos), value='0.9.1') - decode = gr.Dropdown(label='Decode', choices=['diffusers', 'native'], value='diffusers', visible=False) - with gr.Row(): - num_frames = gr.Slider(label='Frames', minimum=9, maximum=257, step=1, value=41) - sampler = gr.Checkbox(label='Override sampler', value=True) - with gr.Row(): - teacache_enable = gr.Checkbox(label='Enable TeaCache', value=False) - teacache_threshold = gr.Slider(label='Threshold', minimum=0.01, maximum=0.1, step=0.01, value=0.03) - with gr.Row(): - model_custom = gr.Textbox(value='', label='Path to model file', visible=False) - with gr.Row(): - from modules.ui_sections import create_video_inputs - video_type, duration, gif_loop, mp4_pad, mp4_interpolate = create_video_inputs(tab='img2img' if is_img2img else 'txt2img') - model.change(fn=model_change, inputs=[model], outputs=[model_custom]) - return [model, model_custom, decode, sampler, num_frames, video_type, duration, gif_loop, mp4_pad, mp4_interpolate, teacache_enable, teacache_threshold] - - def run(self, p: processing.StableDiffusionProcessing, model, model_custom, decode, sampler, num_frames, video_type, duration, gif_loop, mp4_pad, mp4_interpolate, teacache_enable, teacache_threshold): # pylint: disable=arguments-differ, unused-argument - # set params - image = getattr(p, 'init_images', None) - image = None if image is None or len(image) == 0 else image[0] - if (p.width == 0 or p.height == 0) and image is not None: - p.width = image.width - p.height = image.height - num_frames = 8 * int(num_frames // 8) + 1 - p.width = 32 * int(p.width // 32) - p.height = 32 * int(p.height // 32) - processing.fix_seed(p) - if image: - image = images.resize_image(resize_mode=2, im=image, width=p.width, height=p.height, upscaler_name=None, output_type='pil') - p.task_args['image'] = image - p.task_args['output_type'] = 'latent' if decode == 'native' else 'pil' - p.task_args['generator'] = torch.Generator(devices.device).manual_seed(p.seed) - p.task_args['num_frames'] = num_frames - p.do_not_save_grid = True - if sampler: - p.sampler_name = 'Default' - p.ops.append('video') - - # load model - cls = diffusers.LTXPipeline if image is None else diffusers.LTXImageToVideoPipeline - diffusers.LTXTransformer3DModel = diffusers.LTXVideoTransformer3DModel - diffusers.AutoencoderKLLTX = diffusers.AutoencoderKLLTXVideo - repo_id = repos[model] - if repo_id is None: - repo_id = model_custom - if shared.sd_model.__class__ != cls: - sd_models.unload_model_weights() - kwargs = model_quant.create_config() - if os.path.isfile(repo_id): - shared.sd_model = cls.from_single_file( - repo_id, - cache_dir = shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - **kwargs - ) - else: - kwargs = load_quants(kwargs, repo_id) - shared.sd_model = cls.from_pretrained( - repo_id, - cache_dir = shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - **kwargs - ) - sd_models.set_diffuser_options(shared.sd_model) - shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode - shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt - shared.sd_model.vae.decode = hijack_decode - shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(repo_id) - shared.sd_model.sd_model_hash = None - sd_hijack_te.init_hijack(shared.sd_model) - - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - shared.sd_model.vae.enable_slicing() - shared.sd_model.vae.enable_tiling() - shared.sd_model.vae.use_framewise_decoding = True - devices.torch_gc(force=True) - - shared.sd_model.transformer.cnt = 0 - shared.sd_model.transformer.accumulated_rel_l1_distance = 0 - shared.sd_model.transformer.previous_modulated_input = None - shared.sd_model.transformer.previous_residual = None - shared.sd_model.transformer.enable_teacache = teacache_enable - shared.sd_model.transformer.rel_l1_thresh = teacache_threshold - shared.sd_model.transformer.num_steps = p.steps - - shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} args={p.task_args} steps={p.steps} teacache={teacache_enable} threshold={teacache_threshold}') - - # run processing - t0 = time.time() - processed = processing.process_images(p) - t1 = time.time() - if processed is not None and len(processed.images) > 0: - shared.log.info(f'Video: frames={len(processed.images)} time={t1-t0:.2f}') - if video_type != 'None': - images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=duration, loop=gif_loop, pad=mp4_pad, interpolate=mp4_interpolate) - return processed diff --git a/scripts/mochivideo.py b/scripts/mochivideo.py deleted file mode 100644 index 9a3d1f2fa..000000000 --- a/scripts/mochivideo.py +++ /dev/null @@ -1,70 +0,0 @@ -import time -import torch -import gradio as gr -import diffusers -from modules import scripts_manager, processing, shared, images, devices, sd_models, sd_checkpoint, model_quant - - -repo_id = 'genmo/mochi-1-preview' - - -class Script(scripts_manager.Script): - def title(self): - return 'Video: Mochi.1 Video (Legacy)' - - def show(self, is_img2img): - return not is_img2img - - # return signature is array of gradio components - def ui(self, is_img2img): - with gr.Row(): - gr.HTML('  Mochi.1 Video
') - with gr.Row(): - num_frames = gr.Slider(label='Frames', minimum=9, maximum=257, step=1, value=45) - with gr.Row(): - from modules.ui_sections import create_video_inputs - video_type, duration, gif_loop, mp4_pad, mp4_interpolate = create_video_inputs(tab='img2img' if is_img2img else 'txt2img') - return [num_frames, video_type, duration, gif_loop, mp4_pad, mp4_interpolate] - - def run(self, p: processing.StableDiffusionProcessing, num_frames, video_type, duration, gif_loop, mp4_pad, mp4_interpolate): # pylint: disable=arguments-differ, unused-argument - # set params - num_frames = int(num_frames) - p.width = 32 * int(p.width // 32) - p.height = 32 * int(p.height // 32) - p.task_args['output_type'] = 'pil' - p.task_args['generator'] = torch.manual_seed(p.seed) - p.task_args['num_frames'] = num_frames - p.sampler_name = 'Default' - p.do_not_save_grid = True - p.ops.append('video') - - # load model - cls = diffusers.MochiPipeline - if shared.sd_model.__class__ != cls: - sd_models.unload_model_weights() - kwargs = model_quant.create_config() - shared.sd_model = cls.from_pretrained( - repo_id, - cache_dir = shared.opts.hfcache_dir, - torch_dtype=devices.dtype, - **kwargs - ) - shared.sd_model.scheduler._shift = 7.0 # pylint: disable=protected-access - sd_models.set_diffuser_options(shared.sd_model) - shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(repo_id) - shared.sd_model.sd_model_hash = None - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - shared.sd_model.vae.enable_slicing() - shared.sd_model.vae.enable_tiling() - devices.torch_gc(force=True) - shared.log.debug(f'Video: cls={shared.sd_model.__class__.__name__} args={p.task_args}') - - # run processing - t0 = time.time() - processed = processing.process_images(p) - t1 = time.time() - if processed is not None and len(processed.images) > 0: - shared.log.info(f'Video: frames={len(processed.images)} time={t1-t0:.2f}') - if video_type != 'None': - images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=duration, loop=gif_loop, pad=mp4_pad, interpolate=mp4_interpolate) - return processed diff --git a/scripts/pulid_ext.py b/scripts/pulid_ext.py index ede691902..9fdb85e3d 100644 --- a/scripts/pulid_ext.py +++ b/scripts/pulid_ext.py @@ -28,9 +28,11 @@ class Script(scripts_manager.Script): return True def dependencies(self): - from installer import installed, install_insightface + from installer import installed, install, install_insightface if not installed('insightface', reload=False, quiet=True): install_insightface() + if not installed('torchdiffeq'): + install('torchdiffeq') def register(self): # register xyz grid elements global registered # pylint: disable=global-statement diff --git a/scripts/stablevideodiffusion.py b/scripts/stablevideodiffusion.py index a0b959824..0e2d25ab7 100644 --- a/scripts/stablevideodiffusion.py +++ b/scripts/stablevideodiffusion.py @@ -32,7 +32,7 @@ class Script(scripts_manager.Script): min_guidance_scale = gr.Slider(label='Min guidance', minimum=0.0, maximum=10.0, step=0.1, value=1.0) max_guidance_scale = gr.Slider(label='Max guidance', minimum=0.0, maximum=10.0, step=0.1, value=3.0) with gr.Row(): - decode_chunk_size = gr.Slider(label='Decode chunks', minimum=1, maximum=25, step=1, value=6) + decode_chunk_size = gr.Slider(label='Decode chunks', minimum=1, maximum=25, step=1, value=1) motion_bucket_id = gr.Slider(label='Motion level', minimum=0, maximum=1, step=0.05, value=0.5) noise_aug_strength = gr.Slider(label='Noise strength', minimum=0.0, maximum=1.0, step=0.01, value=0.1) with gr.Row(): @@ -42,6 +42,31 @@ class Script(scripts_manager.Script): video_type, duration, gif_loop, mp4_pad, mp4_interpolate = create_video_inputs(tab='img2img' if is_img2img else 'txt2img') return [model, num_frames, override_resolution, min_guidance_scale, max_guidance_scale, decode_chunk_size, motion_bucket_id, noise_aug_strength, video_type, duration, gif_loop, mp4_pad, mp4_interpolate] + def _encode_image(self, image: torch.Tensor, device, num_videos_per_prompt, do_classifier_free_guidance): + image = image.to(device=device, dtype=shared.sd_model.vae.dtype) + shared.log.debug(f'Video encode: type=svd input={image.shape} dtype={image.dtype} device={image.device}') + image_latents = shared.sd_model.vae.encode(image).latent_dist.mode() + image_latents = image_latents.repeat(num_videos_per_prompt, 1, 1, 1) + if do_classifier_free_guidance: + negative_image_latents = torch.zeros_like(image_latents) + image_latents = torch.cat([negative_image_latents, image_latents]) + return image_latents + + def _decode_latents(self, latents: torch.Tensor, num_frames: int, decode_chunk_size: int = 14): + shared.log.debug(f'Video decode: type=svd input={latents.shape} dtype={latents.dtype} device={latents.device} chunk={decode_chunk_size} frames={num_frames}') + latents = latents.flatten(0, 1) + latents = 1 / shared.sd_model.vae.config.scaling_factor * latents + frames = [] + for i in range(0, latents.shape[0], decode_chunk_size): + num_frames_in = latents[i : i + decode_chunk_size].shape[0] + decode_kwargs = { "num_frames": num_frames_in } + frame = shared.sd_model.vae.decode(latents[i : i + decode_chunk_size], **decode_kwargs).sample + frames.append(frame) + frames = torch.cat(frames, dim=0) + frames = frames.reshape(-1, num_frames, *frames.shape[1:]).permute(0, 2, 1, 3, 4) + frames = frames.float() + return frames + def run(self, p: processing.StableDiffusionProcessing, model, num_frames, override_resolution, min_guidance_scale, max_guidance_scale, decode_chunk_size, motion_bucket_id, noise_aug_strength, video_type, duration, gif_loop, mp4_pad, mp4_interpolate): # pylint: disable=arguments-differ, unused-argument image = getattr(p, 'init_images', None) if image is None or len(image) == 0: @@ -60,9 +85,11 @@ class Script(scripts_manager.Script): c = shared.sd_model.__class__.__name__ model_loaded = shared.sd_model.sd_checkpoint_info.model_name if shared.sd_loaded else None if model_name != model_loaded or c != 'StableVideoDiffusionPipeline': + from diffusers import StableVideoDiffusionPipeline # pylint: disable=unused-import shared.opts.sd_model_checkpoint = model_path sd_models.reload_model_weights() - shared.sd_model = shared.sd_model.to(torch.float32) # must run in fp32 due to dtype mismatch + shared.sd_model._encode_vae_image = self._encode_image # pylint: disable=protected-access + shared.sd_model.decode_latents = self._decode_latents # pylint: disable=protected-access # set params if override_resolution: diff --git a/webui.py b/webui.py index 8dc29d5bc..0b6e2b109 100644 --- a/webui.py +++ b/webui.py @@ -12,7 +12,7 @@ from threading import Thread import modules.loader import modules.hashes -from installer import log, git_commit, custom_excepthook +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.devices @@ -263,6 +263,7 @@ def mount_subpath(app): def start_ui(): log.debug('UI start sequence') + log.debug(f'UI image support: kanvas={version["kanvas"]}') modules.script_callbacks.before_ui_callback() timer.startup.record("before-ui") shared.demo = modules.ui.create_ui(timer.startup) diff --git a/wiki b/wiki index 6fc069677..21b5e921d 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 6fc0696770923e78e1ae07cb849942568ec11dde +Subproject commit 21b5e921d81c5dc31447f0869a785e367865b6ed