Author SHA1 Message Date
Tung Nguyen 210dc072b1 chore: bump version 1.1.7 2026-04-03 21:42:11 +07:00
Tung Nguyen 9bd02bd62b chore: add new models & support nano banana via openrouter 2026-04-03 21:41:48 +07:00
Tung Nguyen a445c82b2b chore: bump version 1.1.6 2026-03-11 15:06:16 +07:00
Tung Nguyen 91eceb7d57 fix: cannot import name 'apply_chunking_to_forward' 2026-03-11 15:05:43 +07:00
Tung Nguyen 090418eb8d chore: bump version 1.1.5 2026-02-01 23:04:24 +07:00
Tung Nguyen f1830ba85b feat: support nano banana pro 2026-02-01 23:03:33 +07:00
Tung Nguyen 51ed4a0bc4 chore: bump version 1.1.4 2025-12-31 11:52:13 +07:00
Tung Nguyen c62b79b6d9 fix: ImageScaleToMegapixels not available 2025-12-31 11:50:48 +07:00
Tung Nguyen 1138a1f4d9 rename node category 2025-11-04 15:43:03 +07:00
Tung Nguyen f37981bffb bump version to 1.1.3 2025-11-04 15:32:25 +07:00
Tung Nguyen e40e9244b3 Merge branch 'main' of https://github.com/sipherxyz/comfyui-art-venture 2025-11-04 15:31:55 +07:00
Tung Nguyen c01aafb508 fix: broken video style 2025-11-04 15:31:48 +07:00
Tung Nguyen (Blockchain) 61171a2a87 Merge pull request #116 from wzgrx/patch-1
Update requirements.txt
2025-11-04 14:38:13 +07:00
wzgrx 91b4761689 Update requirements.txt 2025-10-25 21:34:22 +08:00
Tung Nguyen 75b47d8eb4 update README.md 2025-10-24 14:57:04 +07:00
Tung Nguyen c6b58caca6 bump version 1.1.2 2025-10-24 14:29:17 +07:00
Tung Nguyen 8965cfb3c3 add gemini/nanobanana + openrouter support 2025-10-24 14:28:55 +07:00
Tung Nguyen 6df1bc2298 feat: add support for new GPT-5 models in chat module 2025-08-11 21:28:13 +07:00
Tung Nguyen (Blockchain) 6107619468 Merge pull request #112 from AlexK98/main
add more dir names to better find corresponding packages
2025-08-11 20:38:26 +07:00
AlexK98 7dbb8fb35c add more dir names to better find corresponding packages 2025-08-05 14:07:39 +03:00
Tung Nguyen 746c1cefbd chore: update log 2025-07-14 15:23:19 +07:00
Tung Nguyen e5e027e25c deprecate AV_ControlNetPreprocessor 2025-07-14 14:56:58 +07:00
Tung Nguyen dd6673b4d5 add license file 2025-07-14 14:50:32 +07:00
Tung Nguyen 42c116dbb0 fix and deprecate AVCheckpointMerge 2025-07-14 14:46:38 +07:00
Tung Nguyen 12ee9ebe16 chore: replace | with Union 2025-07-14 13:57:34 +07:00
Tung Nguyen 12aa820fdc bump version 1.1.1 2025-07-14 13:49:19 +07:00
Tung Nguyen beae36b7ed feat: support anime-manga lama model 2025-07-14 13:48:36 +07:00
Tung Nguyen b886928282 bump version to 1.1.0 2025-07-09 21:31:21 +07:00
Tung Nguyen 1721ff7a70 refactor(llm): update LLMMessage to support multiple base64 encoded images 2025-07-09 21:22:57 +07:00
Tung Nguyen 8e83109a4e Merge branch 'test' 2025-07-09 14:06:36 +00:00
Tung Nguyen e4510faffb refactor(llm): reorganize model lists 2025-07-09 14:04:00 +00:00
Tung Nguyen 27c0905ba4 refactor: enhance ensure_package function to support version constraints and improve error handling 2025-07-09 14:03:30 +00:00
Tung Nguyen 896a59a294 fix: add "optional" key not available in AV_FaceDetailer 2025-07-09 20:01:40 +07:00
Tung Nguyen b3c5a98603 refactor: update import paths for consistency in text-switch-case and upload modules 2025-07-09 15:45:35 +07:00
Tung Nguyen (Blockchain) 52d1b5c874 Merge pull request #81 from khengyun/feature-support-reason-model
Feature: Support Reasonning Model (openai)
2025-07-09 15:33:09 +07:00
Tung Nguyen (Blockchain) e3b2493b1f Merge pull request #92 from ComfyNodePRs/update-publish-yaml
Update Github Action for Publishing to Comfy Registry
2025-07-09 15:30:40 +07:00
Tung Nguyen (Blockchain) 20952fd8c1 Merge pull request #95 from O-oshir/fixing-image-not-used-in-imageurl
Fixed image not sent to LLM when using LLM Message node due to sending the text instead of the base64 image
2025-07-09 15:29:12 +07:00
Tung Nguyen (Blockchain) 43461905b2 Merge pull request #106 from m0rtus59/feature/add-controlnet-support-for-inpaint
Feature/add controlnet support for inpaint
2025-07-09 15:28:19 +07:00
m0rtus59andgemini-code-assist[bot] a2aaa32fbc Update modules/inpaint/nodes.py
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2025-07-07 00:02:00 +05:00
m0rtus59 707517937a Controlnet support for 'Prepare for inpaint' node
Optional input for controlnet pre-processed images to crop and resize alongside the inpaint_image
2025-07-06 23:41:07 +05:00
m0rtus59 71722e4c7f Merge branch 'sipherxyz:main' into main 2025-07-06 20:10:11 +05:00
Tung Nguyen (Blockchain) 0d7bcc5e23 Merge pull request #69 from Visionatrix/fix/broken-extra-paths-yaml 2025-07-04 03:54:17 +07:00
Tung Nguyen 2503bc2cf3 bump version 1.0.8 2025-07-01 05:51:57 +00:00
Tung Nguyen 3a6a0c2b52 fix(LoadImageFromUrl): has_image always return True 2025-07-01 05:47:24 +00:00
m0rtus59 efaa2dcd6c Merge pull request #1 from m0rtus59/fix/mask-fallback
Fix mask fallback for inpaint_masked=false in PrepareImageAndMaskForInpaint
2025-06-20 20:07:48 +05:00
m0rtus59 c77a9b386e Fix mask fallback for inpaint_masked=false in PrepareImageAndMaskForInpaint
Vibe-coded a solution for when the mask gets de-blurred if the 'inpaint_masked' is 'false'
2025-06-20 19:37:32 +05:00
Tung Nguyen c3bacdc0c4 fix(upload_from_url): double format 2025-06-04 14:13:23 +00:00
Tung Nguyen 4e97ff8c4a chore(upload): improve js code 2025-06-04 13:40:41 +00:00
Tung Nguyen 64fa05980d fix(LoadImageFromUrl): preview not auto load after refresh 2025-06-03 13:51:30 +00:00
Tung Nguyen d78b709e31 bump version 1.0. 2025-06-03 04:58:27 +00:00
Tung Nguyen 4d6caa301b improve image from url code 2025-06-03 04:49:55 +00:00
Yossi Starz ad57177bcd Fixed image not sent to LLM when using LLM Message node due to sending the text instead of the base64 image 2025-04-20 20:43:32 +03:00
Tung Nguyen fc00f4a094 fix(web): error when redefine value property 2025-04-15 08:23:05 +00:00
khaangnguyeen 633352bf5d Update chat.py 2025-02-07 09:53:26 +07:00
snomiao 3e97c544f2 chore(publish): update GitHub Actions workflow for node publishing
- Add permissions for writing issues
- Update action version to v1 for publish-node-action
- Add condition to run job only for 'sipherxyz' repository owner
2025-01-25 07:53:43 +00:00
bigcat88 3bf0cfa0cc do not overwrite "sams" in "folder_paths" if it is present 2024-12-17 14:04:48 +02:00
Tung Nguyen (Blockchain) 50abaace75 Merge pull request #58 from sipherxyz/develop
Release v1.0.6
2024-11-04 21:05:04 +07:00
Tung Nguyen d0cf2abdaf Bump version 1.0.6 2024-11-04 21:01:54 +07:00
Tung Nguyen 36ed264742 add new LLM models 2024-11-04 21:01:02 +07:00
Tung Nguyen fb8bc917ba add TextSwitchCase node 2024-11-04 21:01:02 +07:00
Tung Nguyen (Blockchain) 8d538c9678 Merge pull request #56 from sipherxyz/develop
Fix invalid path aux.py in windows
2024-10-31 20:34:42 +07:00
Tung Nguyen 8034af6478 Bump version to 1.0.5 2024-10-31 13:33:19 +00:00
Tung Nguyen 713f6de761 rename aux.py to preprocessor.py to fix windows issue 2024-10-31 13:32:42 +00:00
Tung Nguyen (Blockchain) 83c3732201 Update pyproject.toml
Bump version to 1.0.4
2024-10-31 15:00:55 +07:00
Tung Nguyen (Blockchain) de986044b3 Merge pull request #54 from sipherxyz/develop
fix controlnet preprocessor node
2024-10-31 15:00:05 +07:00
Tung Nguyen bb6994f677 fix controlnet preprocessor node 2024-10-31 07:57:50 +00:00
Tung Nguyen (Blockchain) 4405434b9c Merge pull request #53 from sipherxyz/develop
Update model download with sha256 validate
2024-10-30 17:25:21 +07:00
Tung Nguyen e7e9e58c66 bump version to 1.0.3 2024-10-30 17:22:14 +07:00
Tung Nguyen 51dd8fcb7c update model download with sha validate 2024-10-30 17:06:33 +07:00
Tung Nguyen (Blockchain) 4be544aa9e Merge pull request #52 from sipherxyz/develop
Release 1.0.2
2024-10-30 12:30:07 +07:00
Tung Nguyen ba06d209d4 bump version to 1.0.2 2024-10-30 12:27:10 +07:00
Tung Nguyen 5d22cae422 add README.md 2024-10-30 12:26:42 +07:00
Tung Nguyen b9dc7e59cb use original image as preview for UtilLoadImageFromUrl & improve JSON extract 2024-10-30 12:26:32 +07:00
Tung Nguyen 135e58a6e9 update lama model url 2024-10-30 12:25:00 +07:00
Tung Nguyen 539864865b fix load image as mask when channel is alpha and image has no alpha 2024-10-30 10:03:39 +07:00
Tung Nguyen c467bbe54f remove AV_StyleApple node 2024-10-30 10:03:24 +07:00
Tung Nguyen (Blockchain) 08fa873f2c Merge pull request #49 from sipherxyz/develop
use torch.jit to load Lama model
2024-10-25 14:35:20 +07:00
Tung Nguyen e1139c55c3 use torch.jit to load Lama model 2024-10-25 07:29:09 +00:00
Tung Nguyen (Blockchain) a8ceae60ea Update pyproject.toml 2024-10-22 11:05:56 +07:00
Tung Nguyen (Blockchain) ee47479097 Update pyproject.toml
Update project description and logo
2024-10-21 20:57:08 +07:00
Tung Nguyen (Blockchain) ee55653500 Update pyproject.toml
Update ComfyUI Registry info
2024-10-21 20:52:52 +07:00
Tung Nguyen (Blockchain) af630185c3 Merge pull request #48 from sipherxyz/add_textrandom_node
Add TextRandomMultiline node
2024-10-21 11:46:21 +07:00
Tung Nguyen (Blockchain) 6bf9ad1f3d Merge pull request #33 from haohaocreates/publish
Add Github Action for Publishing to Comfy Registry
2024-10-21 11:43:51 +07:00
haohaocreates be0a668549 chore(publish): Add Github Action for Publishing to Comfy Registry 2024-05-22 14:15:28 -04:00
66 changed files with 2037 additions and 5997 deletions
+25
View File
@@ -0,0 +1,25 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
paths:
- "pyproject.toml"
permissions:
issues: write
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'sipherxyz' }}
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@v1
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2025 VIXION
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+155
View File
@@ -0,0 +1,155 @@
# ArtVenture Custom Nodes
A comprehensive set of custom nodes for ComfyUI, focusing on utilities for image processing, JSON manipulation, model operations and working with object via URLs
### Image Nodes
#### LoadImageFromUrl
Loads images from URLs.
**Inputs:**
- `image`: List of URLs or base64 image data, separated by new lines
- `keep_alpha_channel`: Preserve alpha channel
- `output_mode`: List or batch output. Use `List` if you have different resolutions.
![load image from url](https://github.com/user-attachments/assets/9da4840c-925e-4e0c-984a-5412282aee79)
### JSON Nodes
#### LoadJsonFromUrl
Loads JSON data from URLs.
**Inputs:**
- `url`: JSON URL
- `print_to_console`: Print JSON to console
#### LoadJsonFromText
Loads JSON data from text.
**Inputs:**
- `data`: JSON text
- `print_to_console`: Print JSON to console
#### Get<\*>FromJson
Includes `GetObjectFromJson`, `GetTextFromJson`, `GetFloatFromJson`, `GetIntFromJson`, `GetBoolFromJson`.
Use key format `key.[index].subkey.[sub_index]` to access nested objects.
![get data from json](https://github.com/user-attachments/assets/a71793d6-9661-441c-a15c-66b2dcaa7972)
### Utility Nodes
#### StringToNumber
Converts strings to numbers.
**Inputs:**
- `string`: Input string
- `rounding`: Rounding method
#### TextRandomMultiline
Randomizes the order of lines in a multiline string.
**Inputs:**
- `text`: Input text
- `amount`: Number of lines to randomize
- `seed`: Random seed
![text random multiline](https://github.com/user-attachments/assets/86f811e3-579e-4ccc-81a3-e216cd851d3c)
#### TextSwitchCase
Switch between multiple cases based on a condition.
**Inputs:**
- `switch_cases`: Switch cases, separated by new lines
- `condition`: Condition to switch on
- `default_value`: Default value when no condition matches
- `delimiter`: Delimiter between case and value, default is `:`
The `switch_cases` format is `case<delimiter>value`, where `case` is the condition to match and `value` is the value to return when the condition matches. You can have new lines in the value to return multiple lines.
![text switch case](https://github.com/user-attachments/assets/4c5450a8-6a3a-4d3c-8c2a-c6e3a33cb95f)
### Inpainting Nodes
#### PrepareImageAndMaskForInpaint
Prepares images and masks for inpainting operations. It's to mimic the behavior of the inpainting in A1111.
**Inputs:**
- `image`: Input image tensor
- `mask`: Input mask tensor
- `mask_blur`: Blur amount for mask (0-64)
- `inpaint_masked`: Whether to inpaint only the masked regions, otherwise it will inpaint the whole image.
- `mask_padding`: Padding around mask (0-256)
- `width`: Manually set inpaint area width. Leave 0 default to the masked area plus padding. (0-2048)
- `height`: Manually set inpaint area height. (0-2048)
**Outputs:**
- `inpaint_image`: Processed image for inpainting
- `inpaint_mask`: Processed mask
- `overlay_image`: Preview overlay
- `crop_region`: Crop coordinates (input of OverlayInpaintedImage)
![inpaiting prepare](https://github.com/user-attachments/assets/38e87c04-7a64-4a62-a462-054396b3de14)
#### OverlayInpaintedImage
Overlays inpainted images with original images.
**Inputs:**
- `inpainted`: Inpainted image
- `overlay_image`: Original image
- `crop_region`: Crop region coordinates
**Outputs:**
- `IMAGE`: Final composited image
#### LaMaInpaint
Remove objects from images using LaMa model.
![lama remove object](https://github.com/user-attachments/assets/c28bbd8b-d55f-4fa5-bbc9-ace267382bd0)
### LLM Nodes
- **LLM API Config**: Generic model settings (model, tokens, temperature).
- **NanoBanana API Config**: Preset for Gemini 2.5 Flash Image (modalities/aspect ratio).
- **OpenAI API**: Connect to OpenAI chat/completions.
- **OpenRouter API**: Route to many providers via OpenRouter; supports text and images where available.
- **Gemini API**: Google Gemini (text and image generation).
- **Claude API**: Anthropic Claude messages.
- **AWS Bedrock Claude API**: Claude via AWS Bedrock.
- **AWS Bedrock Mistral API**: Mistral completions via AWS Bedrock.
- **LLM Message**: Build message lists (system/user/assistant) with optional images.
- **LLM Chat**: Run multi-turn chats; returns text and optional images.
- **LLM Completion**: Single-prompt completion.
![LLM chat workflow](https://github.com/user-attachments/assets/45b8d4fd-57cd-4bd9-8274-d3e6ac4ef938)
![NanoBanana workflow](https://github.com/user-attachments/assets/9d699b47-6239-419f-b778-348618c99c4a)
# Known Issues
## AV_controlnetPreprocessor is missing
`AV_controlnetPreprocessor` is a wrapper for [comfyui_controlnet_aux](https://github.com/Fannovel16/comfyui_controlnet_aux) that I created to quickly switch between multiple preprocessors, instead of having a separate node for each one. It requires `comfyui_controlnet_aux` to be installed, otherwise it will not be available.
Since `AIO_Preprocessor` is already implemented in `comfyui_controlnet_aux`, this node will be deprecated. You are recommended to switch to using `AIO_Preprocessor` node directly.
+15 -49
View File
@@ -3,13 +3,10 @@ from typing import List
import folder_paths
from nodes import ControlNetLoader, ControlNetApply, ControlNetApplyAdvanced
from .preprocessors import control_net_preprocessors, DummyPreprocessor
from .preprocessor import preprocessors, apply_preprocessor
from .advanced import comfy_load_controlnet
control_net_preprocessors["tile"] = (DummyPreprocessor, [])
def load_controlnet(control_net_name, control_net_override="None", timestep_keyframe=None):
if control_net_override != "None":
if control_net_override not in folder_paths.get_filename_list("controlnet"):
@@ -23,33 +20,6 @@ def load_controlnet(control_net_name, control_net_override="None", timestep_keyf
return comfy_load_controlnet(control_net_name, timestep_keyframe=timestep_keyframe)
def apply_preprocessor(image, preprocessor, resolution=512):
if preprocessor == "None":
return image
if preprocessor not in control_net_preprocessors:
raise Exception(f"Preprocessor {preprocessor} is not implemented")
preprocessor_class, default_args = control_net_preprocessors[preprocessor]
default_args: List = default_args.copy()
required_args = preprocessor_class.INPUT_TYPES()["required"].keys()
optional_args = preprocessor_class.INPUT_TYPES().get("optional", {}).keys()
resolution_idx = list(optional_args).index("resolution")
default_args.insert(resolution_idx, resolution)
default_args.insert(0, image)
preprocessor_args = {key: default_args[i] for i, key in enumerate(required_args)}
preprocessor_args.update({key: default_args[i + len(required_args)] for i, key in enumerate(optional_args)})
function_name = preprocessor_class.FUNCTION
res = getattr(preprocessor_class(), function_name)(**preprocessor_args)
if isinstance(res, dict):
res = res["result"]
return res[0]
def detect_controlnet(preprocessor: str, sd_version: str):
controlnets = folder_paths.get_filename_list("controlnet")
controlnets = filter(lambda x: sd_version in x, controlnets)
@@ -96,22 +66,20 @@ class AVControlNetLoader(ControlNetLoader):
RETURN_TYPES = ("CONTROL_NET",)
FUNCTION = "load_controlnet"
CATEGORY = "Art Venture/Loaders"
CATEGORY = "ArtVenture/Loaders"
def load_controlnet(self, control_net_name, control_net_override="None", timestep_keyframe=None):
return load_controlnet(control_net_name, control_net_override, timestep_keyframe=timestep_keyframe)
class AV_ControlNetPreprocessor:
preprocessors = list(control_net_preprocessors.keys())
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"preprocessor": (["None", "tile"] + s.preprocessors,),
"sd_version": (["sd15", "sdxl", "sdxl_t2i"],),
"preprocessor": (["None"] + preprocessors,),
"sd_version": (["sd15", "sdxl"],),
},
"optional": {
"resolution": ("INT", {"default": 512, "min": 64, "max": 2048, "step": 64}),
@@ -122,11 +90,12 @@ class AV_ControlNetPreprocessor:
RETURN_TYPES = ("IMAGE", "STRING")
RETURN_NAMES = ("IMAGE", "CNET_NAME")
FUNCTION = "detect_controlnet"
CATEGORY = "Art Venture/Loaders"
CATEGORY = "ArtVenture/Loaders"
DESCRIPTION = "DEPRECATED: Use comfyui_controlnet_aux's AIO Preprocessor instead"
def detect_controlnet(self, image, preprocessor, sd_version, resolution=512, preprocessor_override="None"):
if preprocessor_override != "None":
if preprocessor_override not in control_net_preprocessors:
if preprocessor_override not in preprocessors:
print(
f"Warning: Not found ControlNet preprocessor {preprocessor_override}. Use {preprocessor} instead."
)
@@ -141,7 +110,6 @@ class AV_ControlNetPreprocessor:
class AVControlNetEfficientStacker:
controlnets = folder_paths.get_filename_list("controlnet")
preprocessors = list(control_net_preprocessors.keys())
@classmethod
def INPUT_TYPES(s):
@@ -155,7 +123,7 @@ class AVControlNetEfficientStacker:
),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"preprocessor": (["None"] + s.preprocessors,),
"preprocessor": (["None"] + preprocessors,),
},
"optional": {
"cnet_stack": ("CONTROL_NET_STACK",),
@@ -169,7 +137,7 @@ class AVControlNetEfficientStacker:
RETURN_TYPES = ("CONTROL_NET_STACK",)
RETURN_NAMES = ("CNET_STACK",)
FUNCTION = "control_net_stacker"
CATEGORY = "Art Venture/Loaders"
CATEGORY = "ArtVenture/Loaders"
def control_net_stacker(
self,
@@ -219,7 +187,7 @@ class AVControlNetEfficientStackerSimple(AVControlNetEfficientStacker):
"FLOAT",
{"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01},
),
"preprocessor": (["None"] + s.preprocessors,),
"preprocessor": (["None"] + preprocessors,),
},
"optional": {
"cnet_stack": ("CONTROL_NET_STACK",),
@@ -242,7 +210,6 @@ class AVControlNetEfficientStackerSimple(AVControlNetEfficientStacker):
class AVControlNetEfficientLoader(ControlNetApply):
controlnets = folder_paths.get_filename_list("controlnet")
preprocessors = list(control_net_preprocessors.keys())
@classmethod
def INPUT_TYPES(s):
@@ -255,7 +222,7 @@ class AVControlNetEfficientLoader(ControlNetApply):
"FLOAT",
{"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01},
),
"preprocessor": (["None"] + s.preprocessors,),
"preprocessor": (["None"] + preprocessors,),
},
"optional": {
"control_net_override": ("STRING", {"default": "None"}),
@@ -267,7 +234,7 @@ class AVControlNetEfficientLoader(ControlNetApply):
RETURN_TYPES = ("CONDITIONING",)
FUNCTION = "load_controlnet"
CATEGORY = "Art Venture/Loaders"
CATEGORY = "ArtVenture/Loaders"
def load_controlnet(
self,
@@ -295,7 +262,6 @@ class AVControlNetEfficientLoader(ControlNetApply):
class AVControlNetEfficientLoaderAdvanced(ControlNetApplyAdvanced):
controlnets = folder_paths.get_filename_list("controlnet")
preprocessors = list(control_net_preprocessors.keys())
@classmethod
def INPUT_TYPES(s):
@@ -311,7 +277,7 @@ class AVControlNetEfficientLoaderAdvanced(ControlNetApplyAdvanced):
),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"preprocessor": (["None"] + s.preprocessors,),
"preprocessor": (["None"] + preprocessors,),
},
"optional": {
"control_net_override": ("STRING", {"default": "None"}),
@@ -324,7 +290,7 @@ class AVControlNetEfficientLoaderAdvanced(ControlNetApplyAdvanced):
RETURN_TYPES = ("CONDITIONING", "CONDITIONING")
RETURN_NAMES = ("positive", "negative")
FUNCTION = "load_controlnet"
CATEGORY = "Art Venture/Loaders"
CATEGORY = "ArtVenture/Loaders"
def load_controlnet(
self,
@@ -368,5 +334,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"AV_ControlNetEfficientLoaderAdvanced": "ControlNet Loader Adv.",
"AV_ControlNetEfficientStacker": "ControlNet Stacker Adv.",
"AV_ControlNetEfficientStackerSimple": "ControlNet Stacker",
"AV_ControlNetPreprocessor": "ControlNet Preprocessor",
"AV_ControlNetPreprocessor": "[Deprecated] ControlNet Preprocessor",
}
+1 -1
View File
@@ -8,7 +8,7 @@ import comfy.controlnet
from ..utils import load_module
custom_nodes = folder_paths.get_folder_paths("custom_nodes")
advanced_cnet_dir_names = ["AdvancedControlNet", "ComfyUI-Advanced-ControlNet"]
advanced_cnet_dir_names = ["AdvancedControlNet", "ComfyUI-Advanced-ControlNet", "comfyui-advanced-controlnet"]
def comfy_load_controlnet(control_net_name: str, **_):
+106
View File
@@ -0,0 +1,106 @@
import os
from typing import Dict
import folder_paths
from ..utils import load_module
custom_nodes = folder_paths.get_folder_paths("custom_nodes")
preprocessors_dir_names = ["ControlNetPreprocessors", "comfyui_controlnet_aux"]
preprocessors: list[str] = []
_preprocessors_map = {
"canny": "CannyEdgePreprocessor",
"canny_pyra": "PyraCannyPreprocessor",
"lineart": "LineArtPreprocessor",
"lineart_anime": "AnimeLineArtPreprocessor",
"lineart_manga": "Manga2Anime_LineArt_Preprocessor",
"lineart_any": "AnyLineArtPreprocessor_aux",
"scribble": "ScribblePreprocessor",
"scribble_xdog": "Scribble_XDoG_Preprocessor",
"scribble_pidi": "Scribble_PiDiNet_Preprocessor",
"scribble_hed": "FakeScribblePreprocessor",
"hed": "HEDPreprocessor",
"pidi": "PiDiNetPreprocessor",
"mlsd": "M-LSDPreprocessor",
"pose": "DWPreprocessor",
"openpose": "OpenposePreprocessor",
"dwpose": "DWPreprocessor",
"pose_dense": "DensePosePreprocessor",
"pose_animal": "AnimalPosePreprocessor",
"normalmap_bae": "BAE-NormalMapPreprocessor",
"normalmap_dsine": "DSINE-NormalMapPreprocessor",
"normalmap_midas": "MiDaS-NormalMapPreprocessor",
"depth": "DepthAnythingV2Preprocessor",
"depth_anything": "DepthAnythingPreprocessor",
"depth_anything_v2": "DepthAnythingV2Preprocessor",
"depth_anything_zoe": "Zoe_DepthAnythingPreprocessor",
"depth_zoe": "Zoe-DepthMapPreprocessor",
"depth_midas": "MiDaS-DepthMapPreprocessor",
"depth_leres": "LeReS-DepthMapPreprocessor",
"depth_metric3d": "Metric3D-DepthMapPreprocessor",
"depth_meshgraphormer": "MeshGraphormer-DepthMapPreprocessor",
"seg_ofcoco": "OneFormer-COCO-SemSegPreprocessor",
"seg_ofade20k": "OneFormer-ADE20K-SemSegPreprocessor",
"seg_ufade20k": "UniFormer-SemSegPreprocessor",
"seg_animeface": "AnimeFace_SemSegPreprocessor",
"shuffle": "ShufflePreprocessor",
"teed": "TEEDPreprocessor",
"color": "ColorPreprocessor",
"sam": "SAMPreprocessor",
"tile": "TilePreprocessor",
}
def apply_preprocessor(image, preprocessor, resolution=512):
raise NotImplementedError("apply_preprocessor is not implemented")
try:
module_path = None
for custom_node in custom_nodes:
custom_node = custom_node if not os.path.islink(custom_node) else os.readlink(custom_node)
for module_dir in preprocessors_dir_names:
if module_dir in os.listdir(custom_node):
module_path = os.path.abspath(os.path.join(custom_node, module_dir))
break
if module_path is None:
raise Exception("Could not find comfyui_controlnet_aux nodes, AV_ControlNetPreprocessor will not work. Please install comfyui_controlnet_aux first")
module = load_module(module_path)
print("Loaded comfyui_controlnet_aux nodes from", module_path)
nodes: Dict = getattr(module, "NODE_CLASS_MAPPINGS")
available_preprocessors: list[str] = getattr(module, "PREPROCESSOR_OPTIONS")
AIO_Preprocessor = nodes.get("AIO_Preprocessor", None)
if AIO_Preprocessor is None:
raise Exception("Could not find AIO_Preprocessor node")
for name, preprocessor in _preprocessors_map.items():
if preprocessor in available_preprocessors:
preprocessors.append(name)
aio_preprocessor = AIO_Preprocessor()
def apply_preprocessor(image, preprocessor, resolution=512):
if preprocessor == "None":
return image
if preprocessor not in preprocessors:
raise Exception(f"Preprocessor {preprocessor} is not implemented")
preprocessor_cls = _preprocessors_map[preprocessor]
args = {"preprocessor": preprocessor_cls, "image": image, "resolution": resolution}
function_name = AIO_Preprocessor.FUNCTION
res = getattr(aio_preprocessor, function_name)(**args)
if isinstance(res, dict):
res = res["result"]
return res[0]
except Exception as e:
print(e)
-142
View File
@@ -1,142 +0,0 @@
import os
import math
from typing import Dict
import folder_paths
from ..utils import load_module
custom_nodes = folder_paths.get_folder_paths("custom_nodes")
preprocessors_dir_names = ["ControlNetPreprocessors", "comfyui_controlnet_aux"]
control_net_preprocessors = {}
try:
module_path = None
for custom_node in custom_nodes:
custom_node = (
custom_node if not os.path.islink(custom_node) else os.readlink(custom_node)
)
for module_dir in preprocessors_dir_names:
if module_dir in os.listdir(custom_node):
module_path = os.path.abspath(os.path.join(custom_node, module_dir))
break
if module_path is None:
raise Exception("Could not find ControlNetPreprocessors nodes")
module = load_module(module_path)
print("Loaded ControlNetPreprocessors nodes from", module_path)
nodes: Dict = getattr(module, "NODE_CLASS_MAPPINGS")
if "CannyEdgePreprocessor" in nodes:
control_net_preprocessors["canny"] = (
nodes["CannyEdgePreprocessor"],
[100, 200],
)
if "LineArtPreprocessor" in nodes:
control_net_preprocessors["lineart"] = (
nodes["LineArtPreprocessor"],
["disable"],
)
control_net_preprocessors["lineart_coarse"] = (
nodes["LineArtPreprocessor"],
["enable"],
)
if "AnimeLineArtPreprocessor" in nodes:
control_net_preprocessors["lineart_anime"] = (
nodes["AnimeLineArtPreprocessor"],
[],
)
if "Manga2Anime_LineArt_Preprocessor" in nodes:
control_net_preprocessors["lineart_manga"] = (
nodes["Manga2Anime_LineArt_Preprocessor"],
[],
)
if "ScribblePreprocessor" in nodes:
control_net_preprocessors["scribble"] = (nodes["ScribblePreprocessor"], [])
if "FakeScribblePreprocessor" in nodes:
control_net_preprocessors["scribble_hed"] = (
nodes["FakeScribblePreprocessor"],
["enable"],
)
if "HEDPreprocessor" in nodes:
control_net_preprocessors["hed"] = (nodes["HEDPreprocessor"], ["disable"])
control_net_preprocessors["hed_safe"] = (nodes["HEDPreprocessor"], ["enable"])
if "PiDiNetPreprocessor" in nodes:
control_net_preprocessors["pidi"] = (
nodes["PiDiNetPreprocessor"],
["disable"],
)
control_net_preprocessors["pidi_safe"] = (
nodes["PiDiNetPreprocessor"],
["enable"],
)
if "M-LSDPreprocessor" in nodes:
control_net_preprocessors["mlsd"] = (nodes["M-LSDPreprocessor"], [0.1, 0.1])
if "OpenposePreprocessor" in nodes:
control_net_preprocessors["openpose"] = (
nodes["OpenposePreprocessor"],
["enable", "enable", "enable"],
)
control_net_preprocessors["pose"] = control_net_preprocessors["openpose"]
if "DWPreprocessor" in nodes:
control_net_preprocessors["dwpose"] = (
nodes["DWPreprocessor"],
["enable", "enable", "enable", "yolox_l.onnx", "dw-ll_ucoco_384.onnx"],
)
# use DWPreprocessor for pose by default if available
control_net_preprocessors["pose"] = control_net_preprocessors["dwpose"]
if "BAE-NormalMapPreprocessor" in nodes:
control_net_preprocessors["normalmap_bae"] = (
nodes["BAE-NormalMapPreprocessor"],
[],
)
if "MiDaS-NormalMapPreprocessor" in nodes:
control_net_preprocessors["normalmap_midas"] = (
nodes["MiDaS-NormalMapPreprocessor"],
[math.pi * 2.0, 0.1],
)
if "MiDaS-DepthMapPreprocessor" in nodes:
control_net_preprocessors["depth_midas"] = (
nodes["MiDaS-DepthMapPreprocessor"],
[math.pi * 2.0, 0.4],
)
if "Zoe-DepthMapPreprocessor" in nodes:
control_net_preprocessors["depth"] = (nodes["Zoe-DepthMapPreprocessor"], [])
control_net_preprocessors["depth_zoe"] = (nodes["Zoe-DepthMapPreprocessor"], [])
if "OneFormer-COCO-SemSegPreprocessor" in nodes:
control_net_preprocessors["seg_ofcoco"] = (
nodes["OneFormer-COCO-SemSegPreprocessor"],
[],
)
if "OneFormer-ADE20K-SemSegPreprocessor" in nodes:
control_net_preprocessors["seg_ofade20k"] = (
nodes["OneFormer-ADE20K-SemSegPreprocessor"],
[],
)
if "UniFormer-SemSegPreprocessor" in nodes:
control_net_preprocessors["seg_ufade20k"] = (
nodes["UniFormer-SemSegPreprocessor"],
[],
)
except Exception as e:
print(e)
class DummyPreprocessor:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",)
}
}
FUNCTION = "process"
def process(self, image):
return (image,)
+2 -2
View File
@@ -20,7 +20,7 @@ class KSamplerWithSharpness(KSampler):
return inputs
CATEGORY = "Art Venture/Sampling"
CATEGORY = "ArtVenture/Sampling"
def sample(self, *args, sharpness=2.0, **kwargs):
patch.sharpness = sharpness
@@ -43,7 +43,7 @@ class KSamplerAdvancedWithSharpness(KSamplerAdvanced):
return inputs
CATEGORY = "Art Venture/Sampling"
CATEGORY = "ArtVenture/Sampling"
def sample(self, *args, sharpness=2.0, **kwargs):
patch.sharpness = sharpness
+34 -50
View File
@@ -1,4 +1,5 @@
import torch
from typing import Union, Optional
Tensor = torch.Tensor
@@ -7,12 +8,12 @@ Dtype = torch.Type
pad = torch.nn.functional.pad
def _compute_zero_padding(kernel_size: tuple[int, int] | int) -> tuple[int, int]:
def _compute_zero_padding(kernel_size: Union[tuple[int, int], int]) -> tuple[int, int]:
ky, kx = _unpack_2d_ks(kernel_size)
return (ky - 1) // 2, (kx - 1) // 2
def _unpack_2d_ks(kernel_size: tuple[int, int] | int) -> tuple[int, int]:
def _unpack_2d_ks(kernel_size: Union[tuple[int, int], int]) -> tuple[int, int]:
if isinstance(kernel_size, int):
ky = kx = kernel_size
else:
@@ -26,17 +27,14 @@ def _unpack_2d_ks(kernel_size: tuple[int, int] | int) -> tuple[int, int]:
def gaussian(
window_size: int,
sigma: Tensor | float,
sigma: Union[Tensor, float],
*,
device: Device | None = None,
dtype: Dtype | None = None,
device: Optional[Device] = None,
dtype: Optional[Dtype] = None,
) -> Tensor:
batch_size = sigma.shape[0]
x = (
torch.arange(window_size, device=sigma.device, dtype=sigma.dtype)
- window_size // 2
).expand(batch_size, -1)
x = (torch.arange(window_size, device=sigma.device, dtype=sigma.dtype) - window_size // 2).expand(batch_size, -1)
if window_size % 2 == 0:
x = x + 0.5
@@ -48,68 +46,58 @@ def gaussian(
def get_gaussian_kernel1d(
kernel_size: int,
sigma: float | Tensor,
sigma: Union[float, Tensor],
force_even: bool = False,
*,
device: Device | None = None,
dtype: Dtype | None = None,
device: Optional[Device] = None,
dtype: Optional[Dtype] = None,
) -> Tensor:
return gaussian(kernel_size, sigma, device=device, dtype=dtype)
def get_gaussian_kernel2d(
kernel_size: tuple[int, int] | int,
sigma: tuple[float, float] | Tensor,
kernel_size: Union[tuple[int, int], int],
sigma: Union[tuple[float, float], Tensor],
force_even: bool = False,
*,
device: Device | None = None,
dtype: Dtype | None = None,
device: Optional[Device] = None,
dtype: Optional[Dtype] = None,
) -> Tensor:
sigma = torch.Tensor([[sigma, sigma]]).to(device=device, dtype=dtype)
ksize_y, ksize_x = _unpack_2d_ks(kernel_size)
sigma_y, sigma_x = sigma[:, 0, None], sigma[:, 1, None]
kernel_y = get_gaussian_kernel1d(
ksize_y, sigma_y, force_even, device=device, dtype=dtype
)[..., None]
kernel_x = get_gaussian_kernel1d(
ksize_x, sigma_x, force_even, device=device, dtype=dtype
)[..., None]
kernel_y = get_gaussian_kernel1d(ksize_y, sigma_y, force_even, device=device, dtype=dtype)[..., None]
kernel_x = get_gaussian_kernel1d(ksize_x, sigma_x, force_even, device=device, dtype=dtype)[..., None]
return kernel_y * kernel_x.view(-1, 1, ksize_x)
def _bilateral_blur(
input: Tensor,
guidance: Tensor | None,
kernel_size: tuple[int, int] | int,
sigma_color: float | Tensor,
sigma_space: tuple[float, float] | Tensor,
guidance: Union[Tensor, None],
kernel_size: Union[tuple[int, int], int],
sigma_color: Union[float, Tensor],
sigma_space: Union[tuple[float, float], Tensor],
border_type: str = "reflect",
color_distance_type: str = "l1",
) -> Tensor:
if isinstance(sigma_color, Tensor):
sigma_color = sigma_color.to(device=input.device, dtype=input.dtype).view(
-1, 1, 1, 1, 1
)
sigma_color = sigma_color.to(device=input.device, dtype=input.dtype).view(-1, 1, 1, 1, 1)
ky, kx = _unpack_2d_ks(kernel_size)
pad_y, pad_x = _compute_zero_padding(kernel_size)
padded_input = pad(input, (pad_x, pad_x, pad_y, pad_y), mode=border_type)
unfolded_input = (
padded_input.unfold(2, ky, 1).unfold(3, kx, 1).flatten(-2)
) # (B, C, H, W, Ky x Kx)
unfolded_input = padded_input.unfold(2, ky, 1).unfold(3, kx, 1).flatten(-2) # (B, C, H, W, Ky x Kx)
if guidance is None:
guidance = input
unfolded_guidance = unfolded_input
else:
padded_guidance = pad(guidance, (pad_x, pad_x, pad_y, pad_y), mode=border_type)
unfolded_guidance = (
padded_guidance.unfold(2, ky, 1).unfold(3, kx, 1).flatten(-2)
) # (B, C, H, W, Ky x Kx)
unfolded_guidance = padded_guidance.unfold(2, ky, 1).unfold(3, kx, 1).flatten(-2) # (B, C, H, W, Ky x Kx)
diff = unfolded_guidance - guidance.unsqueeze(-1)
if color_distance_type == "l1":
@@ -118,13 +106,9 @@ def _bilateral_blur(
color_distance_sq = diff.square().sum(1, keepdim=True)
else:
raise ValueError("color_distance_type only acceps l1 or l2")
color_kernel = (
-0.5 / sigma_color**2 * color_distance_sq
).exp() # (B, 1, H, W, Ky x Kx)
color_kernel = (-0.5 / sigma_color**2 * color_distance_sq).exp() # (B, 1, H, W, Ky x Kx)
space_kernel = get_gaussian_kernel2d(
kernel_size, sigma_space, device=input.device, dtype=input.dtype
)
space_kernel = get_gaussian_kernel2d(kernel_size, sigma_space, device=input.device, dtype=input.dtype)
space_kernel = space_kernel.view(-1, 1, 1, 1, kx * ky)
kernel = space_kernel * color_kernel
@@ -134,9 +118,9 @@ def _bilateral_blur(
def bilateral_blur(
input: Tensor,
kernel_size: tuple[int, int] | int = (13, 13),
sigma_color: float | Tensor = 3.0,
sigma_space: tuple[float, float] | Tensor = 3.0,
kernel_size: Union[tuple[int, int], int] = (13, 13),
sigma_color: Union[float, Tensor] = 3.0,
sigma_space: Union[tuple[float, float], Tensor] = 3.0,
border_type: str = "reflect",
color_distance_type: str = "l1",
) -> Tensor:
@@ -154,9 +138,9 @@ def bilateral_blur(
def joint_bilateral_blur(
input: Tensor,
guidance: Tensor,
kernel_size: tuple[int, int] | int,
sigma_color: float | Tensor,
sigma_space: tuple[float, float] | Tensor,
kernel_size: Union[tuple[int, int], int],
sigma_color: Union[float, Tensor],
sigma_space: Union[tuple[float, float], Tensor],
border_type: str = "reflect",
color_distance_type: str = "l1",
) -> Tensor:
@@ -174,9 +158,9 @@ def joint_bilateral_blur(
class _BilateralBlur(torch.nn.Module):
def __init__(
self,
kernel_size: tuple[int, int] | int,
sigma_color: float | Tensor,
sigma_space: tuple[float, float] | Tensor,
kernel_size: Union[tuple[int, int], int],
sigma_color: Union[float, Tensor],
sigma_space: Union[tuple[float, float], Tensor],
border_type: str = "reflect",
color_distance_type: str = "l1",
) -> None:
+6 -12
View File
@@ -16,14 +16,10 @@ try:
module_path = None
for custom_node in custom_nodes:
custom_node = (
custom_node if not os.path.islink(custom_node) else os.readlink(custom_node)
)
custom_node = custom_node if not os.path.islink(custom_node) else os.readlink(custom_node)
for module_dir in efficieny_dir_names:
if module_dir in os.listdir(custom_node):
module_path = os.path.abspath(
os.path.join(custom_node, module_dir)
)
module_path = os.path.abspath(os.path.join(custom_node, module_dir))
break
if module_path is None:
@@ -49,7 +45,7 @@ try:
return inputs
CATEGORY = "Art Venture/Sampling"
CATEGORY = "ArtVenture/Sampling"
def sample(self, *args, sharpness=2.0, **kwargs):
patch.sharpness = sharpness
@@ -69,7 +65,7 @@ try:
return inputs
CATEGORY = "Art Venture/Sampling"
CATEGORY = "ArtVenture/Sampling"
def sampleadv(self, *args, sharpness=2.0, **kwargs):
patch.sharpness = sharpness
@@ -87,7 +83,7 @@ try:
inputs["optional"]["lora_override"] = ("STRING", {"default": "None"})
return inputs
CATEGORY = "Art Venture/Loaders"
CATEGORY = "ArtVenture/Loaders"
def efficientloader(
self,
@@ -108,9 +104,7 @@ try:
if lora_override != "None":
lora_name = lora_override
return super().efficientloader(
ckpt_name, vae_name, clip_skip, lora_name, *args, **kwargs
)
return super().efficientloader(ckpt_name, vae_name, clip_skip, lora_name, *args, **kwargs)
NODE_CLASS_MAPPINGS.update(
{
+7 -2
View File
@@ -1,5 +1,4 @@
import os
import inspect
from typing import Dict
import folder_paths
@@ -7,7 +6,7 @@ import folder_paths
from ..utils import load_module
custom_nodes = folder_paths.get_folder_paths("custom_nodes")
efficieny_dir_names = ["ImpactPack", "ComfyUI-Impact-Pack"]
efficieny_dir_names = ["ImpactPack", "ComfyUI-Impact-Pack", "comfyui-impact-pack"]
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
@@ -36,6 +35,9 @@ try:
@classmethod
def INPUT_TYPES(s):
inputs = FaceDetailer.INPUT_TYPES()
if not "optional" in inputs:
inputs["optional"] = {}
inputs["optional"]["enabled"] = (
"BOOLEAN",
{"default": True, "label_on": "enabled", "label_off": "disabled"},
@@ -75,6 +77,9 @@ try:
@classmethod
def INPUT_TYPES(s):
inputs = FaceDetailerPipe.INPUT_TYPES()
if not "optional" in inputs:
inputs["optional"] = {}
inputs["optional"]["enabled"] = (
"BOOLEAN",
{"default": True, "label_on": "enabled", "label_off": "disabled"},
+59 -43
View File
@@ -1,22 +1,33 @@
# https://github.com/advimman/lama
import os
import yaml
import torch
import torch.nn.functional as F
import logging
import folder_paths
import comfy.model_management as model_management
from ...model_utils import download_model
from ...utils import ensure_package
from ...model_utils import download_file
lama = None
model_dir = os.path.join(folder_paths.models_dir, "lama")
if "lama" not in folder_paths.folder_names_and_paths:
folder_paths.folder_names_and_paths["lama"] = ([model_dir], folder_paths.supported_pt_extensions)
gpu = model_management.get_torch_device()
cpu = torch.device("cpu")
model_dir = os.path.join(folder_paths.models_dir, "lama")
model_url = "https://d111kwgh87c0gj.cloudfront.net/stable-diffusion/lama/big-lama.pt"
config_path = os.path.join(os.path.dirname(os.path.realpath(__file__)), "config.yaml")
_models = {
"big-lama.pt": {
"url": "https://github.com/Sanster/models/releases/download/add_big_lama/big-lama.pt",
"sha": "344c77bbcb158f17dd143070d1e789f38a66c04202311ae3a258ef66667a9ea9",
},
"anime-manga-big-lama.pt": {
"url": "https://github.com/Sanster/models/releases/download/AnimeMangaInpainting/anime-manga-big-lama.pt",
"sha": "479d3afdcb7ed2fd944ed4ebcc39ca45b33491f0f2e43eb1000bd623cfb41823",
},
}
def ceil_modulo(x, mod):
@@ -32,33 +43,43 @@ def pad_tensor_to_modulo(img, mod):
return F.pad(img, pad=(0, out_width - width, 0, out_height - height), mode="reflect")
def load_model():
global lama
if lama is None:
ensure_package("omegaconf")
class LoadLaMaModel:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model_name": (
list(set(["big-lama.pt", "anime-manga-big-lama.pt"] + folder_paths.get_filename_list("lama"))),
),
}
}
from omegaconf import OmegaConf
from .saicinpainting.training.trainers import load_checkpoint
RETURN_TYPES = ("LAMA",)
CATEGORY = "ArtVenture/Inpainting"
FUNCTION = "load_model"
files = download_model(
model_path=model_dir,
model_url=model_url,
ext_filter=[".pt"],
download_name="big-lama.pt",
)
def load_model(self, model_name: str):
model_path = folder_paths.get_full_path("lama", model_name)
if model_path is None:
if model_name in _models:
model_url = _models[model_name]["url"]
model_sha = _models[model_name]["sha"]
logging.info(f"Downloading {model_name} into {model_dir}")
model_path = os.path.join(model_dir, model_name)
download_file(model_url, model_path, model_sha)
else:
raise Exception(f"Not found model {model_name}")
cfg = yaml.safe_load(open(config_path, "rt"))
cfg = OmegaConf.create(cfg)
cfg.training_model.predict_only = True
cfg.visualizer.kind = "noop"
lama = torch.jit.load(model_path, map_location="cpu")
lama.eval()
lama = load_checkpoint(cfg, files[0], strict=False, map_location="cpu")
lama.freeze()
return lama
return (lama,)
class LaMaInpaint:
class LaMaInpaint(LoadLaMaModel):
def __init__(self):
self.model_name = "big-lama.pt"
@classmethod
def INPUT_TYPES(s):
return {
@@ -66,25 +87,22 @@ class LaMaInpaint:
"image": ("IMAGE",),
"mask": ("MASK",),
},
"optional": {"device_mode": (["AUTO", "Prefer GPU", "CPU"],)},
"optional": {
"device_mode": (["AUTO", "Prefer GPU", "CPU"],),
"lama_model": ("LAMA",),
},
}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "Art Venture/Inpainting"
FUNCTION = "lama_inpaint"
def lama_inpaint(
self,
image: torch.Tensor,
mask: torch.Tensor,
device_mode="AUTO",
):
def lama_inpaint(self, image: torch.Tensor, mask: torch.Tensor, device_mode="AUTO", lama_model=None):
if image.shape[0] != mask.shape[0]:
raise Exception("Image and mask must have the same batch size")
device = gpu if device_mode != "CPU" else cpu
model = load_model()
model = lama_model or self.load_model(self.model_name)
model.to(device)
try:
@@ -98,13 +116,11 @@ class LaMaInpaint:
msk = (msk > 0) * 1.0
msk = msk.unsqueeze(0).unsqueeze(0)
batch = {}
batch["image"] = pad_tensor_to_modulo(img, 8).to(device)
batch["mask"] = pad_tensor_to_modulo(msk, 8).to(device)
src_image = pad_tensor_to_modulo(img, 8).to(device)
src_mask = pad_tensor_to_modulo(msk, 8).to(device)
res = model(batch)
res = batch["inpainted"][0].permute(1, 2, 0)
res = res.detach().cpu()
res = model(src_image, src_mask)
res = res[0].permute(1, 2, 0).detach().cpu()
res = res[:orig_h, :orig_w]
inpainted.append(res)
-157
View File
@@ -1,157 +0,0 @@
run_title: b18_ffc075_batch8x15
training_model:
kind: default
visualize_each_iters: 1000
concat_mask: true
store_discr_outputs_for_vis: true
losses:
l1:
weight_missing: 0
weight_known: 10
perceptual:
weight: 0
adversarial:
kind: r1
weight: 10
gp_coef: 0.001
mask_as_fake_target: true
allow_scale_mask: true
feature_matching:
weight: 100
resnet_pl:
weight: 30
weights_path: ${env:TORCH_HOME}
optimizers:
generator:
kind: adam
lr: 0.001
discriminator:
kind: adam
lr: 0.0001
visualizer:
key_order:
- image
- predicted_image
- discr_output_fake
- discr_output_real
- inpainted
rescale_keys:
- discr_output_fake
- discr_output_real
kind: directory
outdir: /group-volume/User-Driven-Content-Generation/r.suvorov/inpainting/experiments/r.suvorov_2021-04-30_14-41-12_train_simple_pix2pix2_gap_sdpl_novgg_large_b18_ffc075_batch8x15/samples
location:
data_root_dir: /group-volume/User-Driven-Content-Generation/datasets/inpainting_data_root_large
out_root_dir: /group-volume/User-Driven-Content-Generation/${env:USER}/inpainting/experiments
tb_dir: /group-volume/User-Driven-Content-Generation/${env:USER}/inpainting/tb_logs
data:
batch_size: 15
val_batch_size: 2
num_workers: 3
train:
indir: ${location.data_root_dir}/train
out_size: 256
mask_gen_kwargs:
irregular_proba: 1
irregular_kwargs:
max_angle: 4
max_len: 200
max_width: 100
max_times: 5
min_times: 1
box_proba: 1
box_kwargs:
margin: 10
bbox_min_size: 30
bbox_max_size: 150
max_times: 3
min_times: 1
segm_proba: 0
segm_kwargs:
confidence_threshold: 0.5
max_object_area: 0.5
min_mask_area: 0.07
downsample_levels: 6
num_variants_per_mask: 1
rigidness_mode: 1
max_foreground_coverage: 0.3
max_foreground_intersection: 0.7
max_mask_intersection: 0.1
max_hidden_area: 0.1
max_scale_change: 0.25
horizontal_flip: true
max_vertical_shift: 0.2
position_shuffle: true
transform_variant: distortions
dataloader_kwargs:
batch_size: ${data.batch_size}
shuffle: true
num_workers: ${data.num_workers}
val:
indir: ${location.data_root_dir}/val
img_suffix: .png
dataloader_kwargs:
batch_size: ${data.val_batch_size}
shuffle: false
num_workers: ${data.num_workers}
visual_test:
indir: ${location.data_root_dir}/korean_test
img_suffix: _input.png
pad_out_to_modulo: 32
dataloader_kwargs:
batch_size: 1
shuffle: false
num_workers: ${data.num_workers}
generator:
kind: ffc_resnet
input_nc: 4
output_nc: 3
ngf: 64
n_downsampling: 3
n_blocks: 18
add_out_act: sigmoid
init_conv_kwargs:
ratio_gin: 0
ratio_gout: 0
enable_lfu: false
downsample_conv_kwargs:
ratio_gin: ${generator.init_conv_kwargs.ratio_gout}
ratio_gout: ${generator.downsample_conv_kwargs.ratio_gin}
enable_lfu: false
resnet_conv_kwargs:
ratio_gin: 0.75
ratio_gout: ${generator.resnet_conv_kwargs.ratio_gin}
enable_lfu: false
discriminator:
kind: pix2pixhd_nlayer
input_nc: 3
ndf: 64
n_layers: 4
evaluator:
kind: default
inpainted_key: inpainted
integral_kind: ssim_fid100_f1
trainer:
kwargs:
gpus: -1
accelerator: ddp
max_epochs: 200
gradient_clip_val: 1
log_gpu_memory: None
limit_train_batches: 25000
val_check_interval: ${trainer.kwargs.limit_train_batches}
log_every_n_steps: 1000
precision: 32
terminate_on_nan: false
check_val_every_n_epoch: 1
num_sanity_val_steps: 8
limit_val_batches: 1000
replace_sampler_ddp: false
checkpoint_kwargs:
verbose: true
save_top_k: 5
save_last: true
period: 1
monitor: val_ssim_fid100_f1_total_mean
mode: max
@@ -1,367 +0,0 @@
import math
import random
import hashlib
import logging
from enum import Enum
import cv2
import numpy as np
# from ..evaluation.masks.mask import SegmentationMask
from ...utils import LinearRamp
LOGGER = logging.getLogger(__name__)
class DrawMethod(Enum):
LINE = "line"
CIRCLE = "circle"
SQUARE = "square"
def make_random_irregular_mask(
shape, max_angle=4, max_len=60, max_width=20, min_times=0, max_times=10, draw_method=DrawMethod.LINE
):
draw_method = DrawMethod(draw_method)
height, width = shape
mask = np.zeros((height, width), np.float32)
times = np.random.randint(min_times, max_times + 1)
for i in range(times):
start_x = np.random.randint(width)
start_y = np.random.randint(height)
for j in range(1 + np.random.randint(5)):
angle = 0.01 + np.random.randint(max_angle)
if i % 2 == 0:
angle = 2 * 3.1415926 - angle
length = 10 + np.random.randint(max_len)
brush_w = 5 + np.random.randint(max_width)
end_x = np.clip((start_x + length * np.sin(angle)).astype(np.int32), 0, width)
end_y = np.clip((start_y + length * np.cos(angle)).astype(np.int32), 0, height)
if draw_method == DrawMethod.LINE:
cv2.line(mask, (start_x, start_y), (end_x, end_y), 1.0, brush_w)
elif draw_method == DrawMethod.CIRCLE:
cv2.circle(mask, (start_x, start_y), radius=brush_w, color=1.0, thickness=-1)
elif draw_method == DrawMethod.SQUARE:
radius = brush_w // 2
mask[start_y - radius : start_y + radius, start_x - radius : start_x + radius] = 1
start_x, start_y = end_x, end_y
return mask[None, ...]
class RandomIrregularMaskGenerator:
def __init__(
self,
max_angle=4,
max_len=60,
max_width=20,
min_times=0,
max_times=10,
ramp_kwargs=None,
draw_method=DrawMethod.LINE,
):
self.max_angle = max_angle
self.max_len = max_len
self.max_width = max_width
self.min_times = min_times
self.max_times = max_times
self.draw_method = draw_method
self.ramp = LinearRamp(**ramp_kwargs) if ramp_kwargs is not None else None
def __call__(self, img, iter_i=None, raw_image=None):
coef = self.ramp(iter_i) if (self.ramp is not None) and (iter_i is not None) else 1
cur_max_len = int(max(1, self.max_len * coef))
cur_max_width = int(max(1, self.max_width * coef))
cur_max_times = int(self.min_times + 1 + (self.max_times - self.min_times) * coef)
return make_random_irregular_mask(
img.shape[1:],
max_angle=self.max_angle,
max_len=cur_max_len,
max_width=cur_max_width,
min_times=self.min_times,
max_times=cur_max_times,
draw_method=self.draw_method,
)
def make_random_rectangle_mask(shape, margin=10, bbox_min_size=30, bbox_max_size=100, min_times=0, max_times=3):
height, width = shape
mask = np.zeros((height, width), np.float32)
bbox_max_size = min(bbox_max_size, height - margin * 2, width - margin * 2)
times = np.random.randint(min_times, max_times + 1)
for i in range(times):
box_width = np.random.randint(bbox_min_size, bbox_max_size)
box_height = np.random.randint(bbox_min_size, bbox_max_size)
start_x = np.random.randint(margin, width - margin - box_width + 1)
start_y = np.random.randint(margin, height - margin - box_height + 1)
mask[start_y : start_y + box_height, start_x : start_x + box_width] = 1
return mask[None, ...]
class RandomRectangleMaskGenerator:
def __init__(self, margin=10, bbox_min_size=30, bbox_max_size=100, min_times=0, max_times=3, ramp_kwargs=None):
self.margin = margin
self.bbox_min_size = bbox_min_size
self.bbox_max_size = bbox_max_size
self.min_times = min_times
self.max_times = max_times
self.ramp = LinearRamp(**ramp_kwargs) if ramp_kwargs is not None else None
def __call__(self, img, iter_i=None, raw_image=None):
coef = self.ramp(iter_i) if (self.ramp is not None) and (iter_i is not None) else 1
cur_bbox_max_size = int(self.bbox_min_size + 1 + (self.bbox_max_size - self.bbox_min_size) * coef)
cur_max_times = int(self.min_times + (self.max_times - self.min_times) * coef)
return make_random_rectangle_mask(
img.shape[1:],
margin=self.margin,
bbox_min_size=self.bbox_min_size,
bbox_max_size=cur_bbox_max_size,
min_times=self.min_times,
max_times=cur_max_times,
)
class RandomSegmentationMaskGenerator:
def __init__(self, **kwargs):
self.impl = None # will be instantiated in first call (effectively in subprocess)
self.kwargs = kwargs
def __call__(self, img, iter_i=None, raw_image=None):
if self.impl is None:
self.impl = SegmentationMask(**self.kwargs)
masks = self.impl.get_masks(np.transpose(img, (1, 2, 0)))
masks = [m for m in masks if len(np.unique(m)) > 1]
return np.random.choice(masks)
def make_random_superres_mask(shape, min_step=2, max_step=4, min_width=1, max_width=3):
height, width = shape
mask = np.zeros((height, width), np.float32)
step_x = np.random.randint(min_step, max_step + 1)
width_x = np.random.randint(min_width, min(step_x, max_width + 1))
offset_x = np.random.randint(0, step_x)
step_y = np.random.randint(min_step, max_step + 1)
width_y = np.random.randint(min_width, min(step_y, max_width + 1))
offset_y = np.random.randint(0, step_y)
for dy in range(width_y):
mask[offset_y + dy :: step_y] = 1
for dx in range(width_x):
mask[:, offset_x + dx :: step_x] = 1
return mask[None, ...]
class RandomSuperresMaskGenerator:
def __init__(self, **kwargs):
self.kwargs = kwargs
def __call__(self, img, iter_i=None):
return make_random_superres_mask(img.shape[1:], **self.kwargs)
class DumbAreaMaskGenerator:
min_ratio = 0.1
max_ratio = 0.35
default_ratio = 0.225
def __init__(self, is_training):
# Parameters:
# is_training(bool): If true - random rectangular mask, if false - central square mask
self.is_training = is_training
def _random_vector(self, dimension):
if self.is_training:
lower_limit = math.sqrt(self.min_ratio)
upper_limit = math.sqrt(self.max_ratio)
mask_side = round((random.random() * (upper_limit - lower_limit) + lower_limit) * dimension)
u = random.randint(0, dimension - mask_side - 1)
v = u + mask_side
else:
margin = (math.sqrt(self.default_ratio) / 2) * dimension
u = round(dimension / 2 - margin)
v = round(dimension / 2 + margin)
return u, v
def __call__(self, img, iter_i=None, raw_image=None):
c, height, width = img.shape
mask = np.zeros((height, width), np.float32)
x1, x2 = self._random_vector(width)
y1, y2 = self._random_vector(height)
mask[x1:x2, y1:y2] = 1
return mask[None, ...]
class OutpaintingMaskGenerator:
def __init__(
self,
min_padding_percent: float = 0.04,
max_padding_percent: int = 0.25,
left_padding_prob: float = 0.5,
top_padding_prob: float = 0.5,
right_padding_prob: float = 0.5,
bottom_padding_prob: float = 0.5,
is_fixed_randomness: bool = False,
):
"""
is_fixed_randomness - get identical paddings for the same image if args are the same
"""
self.min_padding_percent = min_padding_percent
self.max_padding_percent = max_padding_percent
self.probs = [left_padding_prob, top_padding_prob, right_padding_prob, bottom_padding_prob]
self.is_fixed_randomness = is_fixed_randomness
assert self.min_padding_percent <= self.max_padding_percent
assert self.max_padding_percent > 0
assert (
len([x for x in [self.min_padding_percent, self.max_padding_percent] if (x >= 0 and x <= 1)]) == 2
), f"Padding percentage should be in [0,1]"
assert sum(self.probs) > 0, f"At least one of the padding probs should be greater than 0 - {self.probs}"
assert (
len([x for x in self.probs if (x >= 0) and (x <= 1)]) == 4
), f"At least one of padding probs is not in [0,1] - {self.probs}"
if len([x for x in self.probs if x > 0]) == 1:
LOGGER.warning(
f"Only one padding prob is greater than zero - {self.probs}. That means that the outpainting masks will be always on the same side"
)
def apply_padding(self, mask, coord):
mask[
int(coord[0][0] * self.img_h) : int(coord[1][0] * self.img_h),
int(coord[0][1] * self.img_w) : int(coord[1][1] * self.img_w),
] = 1
return mask
def get_padding(self, size):
n1 = int(self.min_padding_percent * size)
n2 = int(self.max_padding_percent * size)
return self.rnd.randint(n1, n2) / size
@staticmethod
def _img2rs(img):
arr = np.ascontiguousarray(img.astype(np.uint8))
str_hash = hashlib.sha1(arr).hexdigest()
res = hash(str_hash) % (2**32)
return res
def __call__(self, img, iter_i=None, raw_image=None):
c, self.img_h, self.img_w = img.shape
mask = np.zeros((self.img_h, self.img_w), np.float32)
at_least_one_mask_applied = False
if self.is_fixed_randomness:
assert raw_image is not None, f"Cant calculate hash on raw_image=None"
rs = self._img2rs(raw_image)
self.rnd = np.random.RandomState(rs)
else:
self.rnd = np.random
coords = [
[(0, 0), (1, self.get_padding(size=self.img_h))],
[(0, 0), (self.get_padding(size=self.img_w), 1)],
[(0, 1 - self.get_padding(size=self.img_h)), (1, 1)],
[(1 - self.get_padding(size=self.img_w), 0), (1, 1)],
]
for pp, coord in zip(self.probs, coords):
if self.rnd.random() < pp:
at_least_one_mask_applied = True
mask = self.apply_padding(mask=mask, coord=coord)
if not at_least_one_mask_applied:
idx = self.rnd.choice(range(len(coords)), p=np.array(self.probs) / sum(self.probs))
mask = self.apply_padding(mask=mask, coord=coords[idx])
return mask[None, ...]
class MixedMaskGenerator:
def __init__(
self,
irregular_proba=1 / 3,
irregular_kwargs=None,
box_proba=1 / 3,
box_kwargs=None,
segm_proba=1 / 3,
segm_kwargs=None,
squares_proba=0,
squares_kwargs=None,
superres_proba=0,
superres_kwargs=None,
outpainting_proba=0,
outpainting_kwargs=None,
invert_proba=0,
):
self.probas = []
self.gens = []
if irregular_proba > 0:
self.probas.append(irregular_proba)
if irregular_kwargs is None:
irregular_kwargs = {}
else:
irregular_kwargs = dict(irregular_kwargs)
irregular_kwargs["draw_method"] = DrawMethod.LINE
self.gens.append(RandomIrregularMaskGenerator(**irregular_kwargs))
if box_proba > 0:
self.probas.append(box_proba)
if box_kwargs is None:
box_kwargs = {}
self.gens.append(RandomRectangleMaskGenerator(**box_kwargs))
if segm_proba > 0:
self.probas.append(segm_proba)
if segm_kwargs is None:
segm_kwargs = {}
self.gens.append(RandomSegmentationMaskGenerator(**segm_kwargs))
if squares_proba > 0:
self.probas.append(squares_proba)
if squares_kwargs is None:
squares_kwargs = {}
else:
squares_kwargs = dict(squares_kwargs)
squares_kwargs["draw_method"] = DrawMethod.SQUARE
self.gens.append(RandomIrregularMaskGenerator(**squares_kwargs))
if superres_proba > 0:
self.probas.append(superres_proba)
if superres_kwargs is None:
superres_kwargs = {}
self.gens.append(RandomSuperresMaskGenerator(**superres_kwargs))
if outpainting_proba > 0:
self.probas.append(outpainting_proba)
if outpainting_kwargs is None:
outpainting_kwargs = {}
self.gens.append(OutpaintingMaskGenerator(**outpainting_kwargs))
self.probas = np.array(self.probas, dtype="float32")
self.probas /= self.probas.sum()
self.invert_proba = invert_proba
def __call__(self, img, iter_i=None, raw_image=None):
kind = np.random.choice(len(self.probas), p=self.probas)
gen = self.gens[kind]
result = gen(img, iter_i=iter_i, raw_image=raw_image)
if self.invert_proba > 0 and random.random() < self.invert_proba:
result = 1 - result
return result
def get_mask_generator(kind, kwargs):
if kind is None:
kind = "mixed"
if kwargs is None:
kwargs = {}
if kind == "mixed":
cl = MixedMaskGenerator
elif kind == "outpainting":
cl = OutpaintingMaskGenerator
elif kind == "dumb":
cl = DumbAreaMaskGenerator
else:
raise NotImplementedError(f"No such generator kind = {kind}")
return cl(**kwargs)
@@ -1,204 +0,0 @@
from typing import Tuple, Dict, Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
class BaseAdversarialLoss:
def pre_generator_step(
self, real_batch: torch.Tensor, fake_batch: torch.Tensor, generator: nn.Module, discriminator: nn.Module
):
"""
Prepare for generator step
:param real_batch: Tensor, a batch of real samples
:param fake_batch: Tensor, a batch of samples produced by generator
:param generator:
:param discriminator:
:return: None
"""
def pre_discriminator_step(
self, real_batch: torch.Tensor, fake_batch: torch.Tensor, generator: nn.Module, discriminator: nn.Module
):
"""
Prepare for discriminator step
:param real_batch: Tensor, a batch of real samples
:param fake_batch: Tensor, a batch of samples produced by generator
:param generator:
:param discriminator:
:return: None
"""
def generator_loss(
self,
real_batch: torch.Tensor,
fake_batch: torch.Tensor,
discr_real_pred: torch.Tensor,
discr_fake_pred: torch.Tensor,
mask: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
"""
Calculate generator loss
:param real_batch: Tensor, a batch of real samples
:param fake_batch: Tensor, a batch of samples produced by generator
:param discr_real_pred: Tensor, discriminator output for real_batch
:param discr_fake_pred: Tensor, discriminator output for fake_batch
:param mask: Tensor, actual mask, which was at input of generator when making fake_batch
:return: total generator loss along with some values that might be interesting to log
"""
raise NotImplemented()
def discriminator_loss(
self,
real_batch: torch.Tensor,
fake_batch: torch.Tensor,
discr_real_pred: torch.Tensor,
discr_fake_pred: torch.Tensor,
mask: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
"""
Calculate discriminator loss and call .backward() on it
:param real_batch: Tensor, a batch of real samples
:param fake_batch: Tensor, a batch of samples produced by generator
:param discr_real_pred: Tensor, discriminator output for real_batch
:param discr_fake_pred: Tensor, discriminator output for fake_batch
:param mask: Tensor, actual mask, which was at input of generator when making fake_batch
:return: total discriminator loss along with some values that might be interesting to log
"""
raise NotImplemented()
def interpolate_mask(self, mask, shape):
assert mask is not None
assert self.allow_scale_mask or shape == mask.shape[-2:]
if shape != mask.shape[-2:] and self.allow_scale_mask:
if self.mask_scale_mode == "maxpool":
mask = F.adaptive_max_pool2d(mask, shape)
else:
mask = F.interpolate(mask, size=shape, mode=self.mask_scale_mode)
return mask
def make_r1_gp(discr_real_pred, real_batch):
if torch.is_grad_enabled():
grad_real = torch.autograd.grad(outputs=discr_real_pred.sum(), inputs=real_batch, create_graph=True)[0]
grad_penalty = (grad_real.view(grad_real.shape[0], -1).norm(2, dim=1) ** 2).mean()
else:
grad_penalty = 0
real_batch.requires_grad = False
return grad_penalty
class NonSaturatingWithR1(BaseAdversarialLoss):
def __init__(
self,
gp_coef=5,
weight=1,
mask_as_fake_target=False,
allow_scale_mask=False,
mask_scale_mode="nearest",
extra_mask_weight_for_gen=0,
use_unmasked_for_gen=True,
use_unmasked_for_discr=True,
):
self.gp_coef = gp_coef
self.weight = weight
# use for discr => use for gen;
# otherwise we teach only the discr to pay attention to very small difference
assert use_unmasked_for_gen or (not use_unmasked_for_discr)
# mask as target => use unmasked for discr:
# if we don't care about unmasked regions at all
# then it doesn't matter if the value of mask_as_fake_target is true or false
assert use_unmasked_for_discr or (not mask_as_fake_target)
self.use_unmasked_for_gen = use_unmasked_for_gen
self.use_unmasked_for_discr = use_unmasked_for_discr
self.mask_as_fake_target = mask_as_fake_target
self.allow_scale_mask = allow_scale_mask
self.mask_scale_mode = mask_scale_mode
self.extra_mask_weight_for_gen = extra_mask_weight_for_gen
def generator_loss(
self,
real_batch: torch.Tensor,
fake_batch: torch.Tensor,
discr_real_pred: torch.Tensor,
discr_fake_pred: torch.Tensor,
mask=None,
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
fake_loss = F.softplus(-discr_fake_pred)
if (
self.mask_as_fake_target and self.extra_mask_weight_for_gen > 0
) or not self.use_unmasked_for_gen: # == if masked region should be treated differently
mask = self.interpolate_mask(mask, discr_fake_pred.shape[-2:])
if not self.use_unmasked_for_gen:
fake_loss = fake_loss * mask
else:
pixel_weights = 1 + mask * self.extra_mask_weight_for_gen
fake_loss = fake_loss * pixel_weights
return fake_loss.mean() * self.weight, dict()
def pre_discriminator_step(
self, real_batch: torch.Tensor, fake_batch: torch.Tensor, generator: nn.Module, discriminator: nn.Module
):
real_batch.requires_grad = True
def discriminator_loss(
self,
real_batch: torch.Tensor,
fake_batch: torch.Tensor,
discr_real_pred: torch.Tensor,
discr_fake_pred: torch.Tensor,
mask=None,
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
real_loss = F.softplus(-discr_real_pred)
grad_penalty = make_r1_gp(discr_real_pred, real_batch) * self.gp_coef
fake_loss = F.softplus(discr_fake_pred)
if not self.use_unmasked_for_discr or self.mask_as_fake_target:
# == if masked region should be treated differently
mask = self.interpolate_mask(mask, discr_fake_pred.shape[-2:])
# use_unmasked_for_discr=False only makes sense for fakes;
# for reals there is no difference beetween two regions
fake_loss = fake_loss * mask
if self.mask_as_fake_target:
fake_loss = fake_loss + (1 - mask) * F.softplus(-discr_fake_pred)
sum_discr_loss = real_loss + grad_penalty + fake_loss
metrics = dict(
discr_real_out=discr_real_pred.mean(), discr_fake_out=discr_fake_pred.mean(), discr_real_gp=grad_penalty
)
return sum_discr_loss.mean(), metrics
class BCELoss(BaseAdversarialLoss):
def __init__(self, weight):
self.weight = weight
self.bce_loss = nn.BCEWithLogitsLoss()
def generator_loss(self, discr_fake_pred: torch.Tensor) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
real_mask_gt = torch.zeros(discr_fake_pred.shape).to(discr_fake_pred.device)
fake_loss = self.bce_loss(discr_fake_pred, real_mask_gt) * self.weight
return fake_loss, dict()
def pre_discriminator_step(
self, real_batch: torch.Tensor, fake_batch: torch.Tensor, generator: nn.Module, discriminator: nn.Module
):
real_batch.requires_grad = True
def discriminator_loss(
self, mask: torch.Tensor, discr_real_pred: torch.Tensor, discr_fake_pred: torch.Tensor
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
real_mask_gt = torch.zeros(discr_real_pred.shape).to(discr_real_pred.device)
sum_discr_loss = (self.bce_loss(discr_real_pred, real_mask_gt) + self.bce_loss(discr_fake_pred, mask)) / 2
metrics = dict(discr_real_out=discr_real_pred.mean(), discr_fake_out=discr_fake_pred.mean(), discr_real_gp=0)
return sum_discr_loss, metrics
def make_discrim_loss(kind, **kwargs):
if kind == "r1":
return NonSaturatingWithR1(**kwargs)
elif kind == "bce":
return BCELoss(**kwargs)
raise ValueError(f"Unknown adversarial loss kind {kind}")
@@ -1,154 +0,0 @@
weights = {
"ade20k": [
6.34517766497462,
9.328358208955224,
11.389521640091116,
16.10305958132045,
20.833333333333332,
22.22222222222222,
25.125628140703515,
43.29004329004329,
50.5050505050505,
54.6448087431694,
55.24861878453038,
60.24096385542168,
62.5,
66.2251655629139,
84.74576271186442,
90.90909090909092,
91.74311926605505,
96.15384615384616,
96.15384615384616,
97.08737864077669,
102.04081632653062,
135.13513513513513,
149.2537313432836,
153.84615384615384,
163.93442622950818,
166.66666666666666,
188.67924528301887,
192.30769230769232,
217.3913043478261,
227.27272727272725,
227.27272727272725,
227.27272727272725,
303.03030303030306,
322.5806451612903,
333.3333333333333,
370.3703703703703,
384.61538461538464,
416.6666666666667,
416.6666666666667,
434.7826086956522,
434.7826086956522,
454.5454545454545,
454.5454545454545,
500.0,
526.3157894736842,
526.3157894736842,
555.5555555555555,
555.5555555555555,
555.5555555555555,
555.5555555555555,
555.5555555555555,
555.5555555555555,
555.5555555555555,
588.2352941176471,
588.2352941176471,
588.2352941176471,
588.2352941176471,
588.2352941176471,
666.6666666666666,
666.6666666666666,
666.6666666666666,
666.6666666666666,
714.2857142857143,
714.2857142857143,
714.2857142857143,
714.2857142857143,
714.2857142857143,
769.2307692307693,
769.2307692307693,
769.2307692307693,
833.3333333333334,
833.3333333333334,
833.3333333333334,
833.3333333333334,
909.090909090909,
1000.0,
1111.111111111111,
1111.111111111111,
1111.111111111111,
1111.111111111111,
1111.111111111111,
1250.0,
1250.0,
1250.0,
1250.0,
1250.0,
1428.5714285714287,
1428.5714285714287,
1428.5714285714287,
1428.5714285714287,
1428.5714285714287,
1428.5714285714287,
1428.5714285714287,
1666.6666666666667,
1666.6666666666667,
1666.6666666666667,
1666.6666666666667,
1666.6666666666667,
1666.6666666666667,
1666.6666666666667,
1666.6666666666667,
1666.6666666666667,
1666.6666666666667,
1666.6666666666667,
2000.0,
2000.0,
2000.0,
2000.0,
2000.0,
2000.0,
2000.0,
2000.0,
2000.0,
2000.0,
2000.0,
2000.0,
2000.0,
2000.0,
2000.0,
2000.0,
2000.0,
2500.0,
2500.0,
2500.0,
2500.0,
2500.0,
2500.0,
2500.0,
2500.0,
2500.0,
2500.0,
2500.0,
2500.0,
2500.0,
3333.3333333333335,
3333.3333333333335,
3333.3333333333335,
3333.3333333333335,
3333.3333333333335,
3333.3333333333335,
3333.3333333333335,
3333.3333333333335,
3333.3333333333335,
3333.3333333333335,
3333.3333333333335,
3333.3333333333335,
3333.3333333333335,
5000.0,
5000.0,
5000.0,
]
}
@@ -1,133 +0,0 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision
from .perceptual import IMAGENET_STD, IMAGENET_MEAN
def dummy_distance_weighter(real_img, pred_img, mask):
return mask
def get_gauss_kernel(kernel_size, width_factor=1):
coords = torch.stack(torch.meshgrid(torch.arange(kernel_size), torch.arange(kernel_size)), dim=0).float()
diff = torch.exp(-((coords - kernel_size // 2) ** 2).sum(0) / kernel_size / width_factor)
diff /= diff.sum()
return diff
class BlurMask(nn.Module):
def __init__(self, kernel_size=5, width_factor=1):
super().__init__()
self.filter = nn.Conv2d(1, 1, kernel_size, padding=kernel_size // 2, padding_mode="replicate", bias=False)
self.filter.weight.data.copy_(get_gauss_kernel(kernel_size, width_factor=width_factor))
def forward(self, real_img, pred_img, mask):
with torch.no_grad():
result = self.filter(mask) * mask
return result
class EmulatedEDTMask(nn.Module):
def __init__(self, dilate_kernel_size=5, blur_kernel_size=5, width_factor=1):
super().__init__()
self.dilate_filter = nn.Conv2d(
1, 1, dilate_kernel_size, padding=dilate_kernel_size // 2, padding_mode="replicate", bias=False
)
self.dilate_filter.weight.data.copy_(
torch.ones(1, 1, dilate_kernel_size, dilate_kernel_size, dtype=torch.float)
)
self.blur_filter = nn.Conv2d(
1, 1, blur_kernel_size, padding=blur_kernel_size // 2, padding_mode="replicate", bias=False
)
self.blur_filter.weight.data.copy_(get_gauss_kernel(blur_kernel_size, width_factor=width_factor))
def forward(self, real_img, pred_img, mask):
with torch.no_grad():
known_mask = 1 - mask
dilated_known_mask = (self.dilate_filter(known_mask) > 1).float()
result = self.blur_filter(1 - dilated_known_mask) * mask
return result
class PropagatePerceptualSim(nn.Module):
def __init__(self, level=2, max_iters=10, temperature=500, erode_mask_size=3):
super().__init__()
vgg = torchvision.models.vgg19(pretrained=True).features
vgg_avg_pooling = []
for weights in vgg.parameters():
weights.requires_grad = False
cur_level_i = 0
for module in vgg.modules():
if module.__class__.__name__ == "Sequential":
continue
elif module.__class__.__name__ == "MaxPool2d":
vgg_avg_pooling.append(nn.AvgPool2d(kernel_size=2, stride=2, padding=0))
else:
vgg_avg_pooling.append(module)
if module.__class__.__name__ == "ReLU":
cur_level_i += 1
if cur_level_i == level:
break
self.features = nn.Sequential(*vgg_avg_pooling)
self.max_iters = max_iters
self.temperature = temperature
self.do_erode = erode_mask_size > 0
if self.do_erode:
self.erode_mask = nn.Conv2d(1, 1, erode_mask_size, padding=erode_mask_size // 2, bias=False)
self.erode_mask.weight.data.fill_(1)
def forward(self, real_img, pred_img, mask):
with torch.no_grad():
real_img = (real_img - IMAGENET_MEAN.to(real_img)) / IMAGENET_STD.to(real_img)
real_feats = self.features(real_img)
vertical_sim = torch.exp(
-(real_feats[:, :, 1:] - real_feats[:, :, :-1]).pow(2).sum(1, keepdim=True) / self.temperature
)
horizontal_sim = torch.exp(
-(real_feats[:, :, :, 1:] - real_feats[:, :, :, :-1]).pow(2).sum(1, keepdim=True) / self.temperature
)
mask_scaled = F.interpolate(mask, size=real_feats.shape[-2:], mode="bilinear", align_corners=False)
if self.do_erode:
mask_scaled = (self.erode_mask(mask_scaled) > 1).float()
cur_knowness = 1 - mask_scaled
for iter_i in range(self.max_iters):
new_top_knowness = F.pad(cur_knowness[:, :, :-1] * vertical_sim, (0, 0, 1, 0), mode="replicate")
new_bottom_knowness = F.pad(cur_knowness[:, :, 1:] * vertical_sim, (0, 0, 0, 1), mode="replicate")
new_left_knowness = F.pad(cur_knowness[:, :, :, :-1] * horizontal_sim, (1, 0, 0, 0), mode="replicate")
new_right_knowness = F.pad(cur_knowness[:, :, :, 1:] * horizontal_sim, (0, 1, 0, 0), mode="replicate")
new_knowness = (
torch.stack([new_top_knowness, new_bottom_knowness, new_left_knowness, new_right_knowness], dim=0)
.max(0)
.values
)
cur_knowness = torch.max(cur_knowness, new_knowness)
cur_knowness = F.interpolate(cur_knowness, size=mask.shape[-2:], mode="bilinear")
result = torch.min(mask, 1 - cur_knowness)
return result
def make_mask_distance_weighter(kind="none", **kwargs):
if kind == "none":
return dummy_distance_weighter
if kind == "blur":
return BlurMask(**kwargs)
if kind == "edt":
return EmulatedEDTMask(**kwargs)
if kind == "pps":
return PropagatePerceptualSim(**kwargs)
raise ValueError(f"Unknown mask distance weighter kind {kind}")
@@ -1,33 +0,0 @@
from typing import List
import torch
import torch.nn.functional as F
def masked_l2_loss(pred, target, mask, weight_known, weight_missing):
per_pixel_l2 = F.mse_loss(pred, target, reduction='none')
pixel_weights = mask * weight_missing + (1 - mask) * weight_known
return (pixel_weights * per_pixel_l2).mean()
def masked_l1_loss(pred, target, mask, weight_known, weight_missing):
per_pixel_l1 = F.l1_loss(pred, target, reduction='none')
pixel_weights = mask * weight_missing + (1 - mask) * weight_known
return (pixel_weights * per_pixel_l1).mean()
def feature_matching_loss(fake_features: List[torch.Tensor], target_features: List[torch.Tensor], mask=None):
if mask is None:
res = torch.stack([F.mse_loss(fake_feat, target_feat)
for fake_feat, target_feat in zip(fake_features, target_features)]).mean()
else:
res = 0
norm = 0
for fake_feat, target_feat in zip(fake_features, target_features):
cur_mask = F.interpolate(mask, size=fake_feat.shape[-2:], mode='bilinear', align_corners=False)
error_weights = 1 - cur_mask
cur_val = ((fake_feat - target_feat).pow(2) * error_weights).mean()
res = res + cur_val
norm += 1
res = res / norm
return res
@@ -1,84 +0,0 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision
from ...utils import check_and_warn_input_range
IMAGENET_MEAN = torch.FloatTensor([0.485, 0.456, 0.406])[None, :, None, None]
IMAGENET_STD = torch.FloatTensor([0.229, 0.224, 0.225])[None, :, None, None]
class PerceptualLoss(nn.Module):
def __init__(self, normalize_inputs=True):
super(PerceptualLoss, self).__init__()
self.normalize_inputs = normalize_inputs
self.mean_ = IMAGENET_MEAN
self.std_ = IMAGENET_STD
vgg = torchvision.models.vgg19(pretrained=True).features
vgg_avg_pooling = []
for weights in vgg.parameters():
weights.requires_grad = False
for module in vgg.modules():
if module.__class__.__name__ == "Sequential":
continue
elif module.__class__.__name__ == "MaxPool2d":
vgg_avg_pooling.append(nn.AvgPool2d(kernel_size=2, stride=2, padding=0))
else:
vgg_avg_pooling.append(module)
self.vgg = nn.Sequential(*vgg_avg_pooling)
def do_normalize_inputs(self, x):
return (x - self.mean_.to(x.device)) / self.std_.to(x.device)
def partial_losses(self, input, target, mask=None):
check_and_warn_input_range(target, 0, 1, "PerceptualLoss target in partial_losses")
# we expect input and target to be in [0, 1] range
losses = []
if self.normalize_inputs:
features_input = self.do_normalize_inputs(input)
features_target = self.do_normalize_inputs(target)
else:
features_input = input
features_target = target
for layer in self.vgg[:30]:
features_input = layer(features_input)
features_target = layer(features_target)
if layer.__class__.__name__ == "ReLU":
loss = F.mse_loss(features_input, features_target, reduction="none")
if mask is not None:
cur_mask = F.interpolate(
mask, size=features_input.shape[-2:], mode="bilinear", align_corners=False
)
loss = loss * (1 - cur_mask)
loss = loss.mean(dim=tuple(range(1, len(loss.shape))))
losses.append(loss)
return losses
def forward(self, input, target, mask=None):
losses = self.partial_losses(input, target, mask=mask)
return torch.stack(losses).sum(dim=0)
def get_global_features(self, input):
check_and_warn_input_range(input, 0, 1, "PerceptualLoss input in get_global_features")
if self.normalize_inputs:
features_input = self.do_normalize_inputs(input)
else:
features_input = input
features_input = self.vgg(features_input)
return features_input
@@ -1,43 +0,0 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from .constants import weights as constant_weights
class CrossEntropy2d(nn.Module):
def __init__(self, reduction="mean", ignore_label=255, weights=None, *args, **kwargs):
"""
weight (Tensor, optional): a manual rescaling weight given to each class.
If given, has to be a Tensor of size "nclasses"
"""
super(CrossEntropy2d, self).__init__()
self.reduction = reduction
self.ignore_label = ignore_label
self.weights = weights
if self.weights is not None:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.weights = torch.FloatTensor(constant_weights[weights]).to(device)
def forward(self, predict, target):
"""
Args:
predict:(n, c, h, w)
target:(n, 1, h, w)
"""
target = target.long()
assert not target.requires_grad
assert predict.dim() == 4, "{0}".format(predict.size())
assert target.dim() == 4, "{0}".format(target.size())
assert predict.size(0) == target.size(0), "{0} vs {1} ".format(predict.size(0), target.size(0))
assert target.size(1) == 1, "{0}".format(target.size(1))
assert predict.size(2) == target.size(2), "{0} vs {1} ".format(predict.size(2), target.size(2))
assert predict.size(3) == target.size(3), "{0} vs {1} ".format(predict.size(3), target.size(3))
target = target.squeeze(1)
n, c, h, w = predict.size()
target_mask = (target >= 0) * (target != self.ignore_label)
target = target[target_mask]
predict = predict.transpose(1, 2).transpose(2, 3).contiguous()
predict = predict[target_mask.view(n, h, w, 1).repeat(1, 1, 1, c)].view(-1, c)
loss = F.cross_entropy(predict, target, weight=self.weights, reduction=self.reduction)
return loss
@@ -1,150 +0,0 @@
import torch
import torch.nn as nn
import torchvision.models as models
class PerceptualLoss(nn.Module):
r"""
Perceptual loss, VGG-based
https://arxiv.org/abs/1603.08155
https://github.com/dxyang/StyleTransfer/blob/master/utils.py
"""
def __init__(self, weights=[1.0, 1.0, 1.0, 1.0, 1.0]):
super(PerceptualLoss, self).__init__()
self.add_module("vgg", VGG19())
self.criterion = torch.nn.L1Loss()
self.weights = weights
def __call__(self, x, y):
# Compute features
x_vgg, y_vgg = self.vgg(x), self.vgg(y)
content_loss = 0.0
content_loss += self.weights[0] * self.criterion(x_vgg["relu1_1"], y_vgg["relu1_1"])
content_loss += self.weights[1] * self.criterion(x_vgg["relu2_1"], y_vgg["relu2_1"])
content_loss += self.weights[2] * self.criterion(x_vgg["relu3_1"], y_vgg["relu3_1"])
content_loss += self.weights[3] * self.criterion(x_vgg["relu4_1"], y_vgg["relu4_1"])
content_loss += self.weights[4] * self.criterion(x_vgg["relu5_1"], y_vgg["relu5_1"])
return content_loss
class VGG19(torch.nn.Module):
def __init__(self):
super(VGG19, self).__init__()
features = models.vgg19(pretrained=True).features
self.relu1_1 = torch.nn.Sequential()
self.relu1_2 = torch.nn.Sequential()
self.relu2_1 = torch.nn.Sequential()
self.relu2_2 = torch.nn.Sequential()
self.relu3_1 = torch.nn.Sequential()
self.relu3_2 = torch.nn.Sequential()
self.relu3_3 = torch.nn.Sequential()
self.relu3_4 = torch.nn.Sequential()
self.relu4_1 = torch.nn.Sequential()
self.relu4_2 = torch.nn.Sequential()
self.relu4_3 = torch.nn.Sequential()
self.relu4_4 = torch.nn.Sequential()
self.relu5_1 = torch.nn.Sequential()
self.relu5_2 = torch.nn.Sequential()
self.relu5_3 = torch.nn.Sequential()
self.relu5_4 = torch.nn.Sequential()
for x in range(2):
self.relu1_1.add_module(str(x), features[x])
for x in range(2, 4):
self.relu1_2.add_module(str(x), features[x])
for x in range(4, 7):
self.relu2_1.add_module(str(x), features[x])
for x in range(7, 9):
self.relu2_2.add_module(str(x), features[x])
for x in range(9, 12):
self.relu3_1.add_module(str(x), features[x])
for x in range(12, 14):
self.relu3_2.add_module(str(x), features[x])
for x in range(14, 16):
self.relu3_2.add_module(str(x), features[x])
for x in range(16, 18):
self.relu3_4.add_module(str(x), features[x])
for x in range(18, 21):
self.relu4_1.add_module(str(x), features[x])
for x in range(21, 23):
self.relu4_2.add_module(str(x), features[x])
for x in range(23, 25):
self.relu4_3.add_module(str(x), features[x])
for x in range(25, 27):
self.relu4_4.add_module(str(x), features[x])
for x in range(27, 30):
self.relu5_1.add_module(str(x), features[x])
for x in range(30, 32):
self.relu5_2.add_module(str(x), features[x])
for x in range(32, 34):
self.relu5_3.add_module(str(x), features[x])
for x in range(34, 36):
self.relu5_4.add_module(str(x), features[x])
# don't need the gradients, just want the features
for param in self.parameters():
param.requires_grad = False
def forward(self, x):
relu1_1 = self.relu1_1(x)
relu1_2 = self.relu1_2(relu1_1)
relu2_1 = self.relu2_1(relu1_2)
relu2_2 = self.relu2_2(relu2_1)
relu3_1 = self.relu3_1(relu2_2)
relu3_2 = self.relu3_2(relu3_1)
relu3_3 = self.relu3_3(relu3_2)
relu3_4 = self.relu3_4(relu3_3)
relu4_1 = self.relu4_1(relu3_4)
relu4_2 = self.relu4_2(relu4_1)
relu4_3 = self.relu4_3(relu4_2)
relu4_4 = self.relu4_4(relu4_3)
relu5_1 = self.relu5_1(relu4_4)
relu5_2 = self.relu5_2(relu5_1)
relu5_3 = self.relu5_3(relu5_2)
relu5_4 = self.relu5_4(relu5_3)
out = {
"relu1_1": relu1_1,
"relu1_2": relu1_2,
"relu2_1": relu2_1,
"relu2_2": relu2_2,
"relu3_1": relu3_1,
"relu3_2": relu3_2,
"relu3_3": relu3_3,
"relu3_4": relu3_4,
"relu4_1": relu4_1,
"relu4_2": relu4_2,
"relu4_3": relu4_3,
"relu4_4": relu4_4,
"relu5_1": relu5_1,
"relu5_2": relu5_2,
"relu5_3": relu5_3,
"relu5_4": relu5_4,
}
return out
@@ -1,36 +0,0 @@
import logging
from ..modules.ffc import FFCResNetGenerator
from ..modules.pix2pixhd import (
GlobalGenerator,
MultiDilatedGlobalGenerator,
NLayerDiscriminator,
MultidilatedNLayerDiscriminator,
)
def make_generator(config, kind, **kwargs):
logging.info(f"Make generator {kind}")
if kind == "pix2pixhd_multidilated":
return MultiDilatedGlobalGenerator(**kwargs)
if kind == "pix2pixhd_global":
return GlobalGenerator(**kwargs)
if kind == "ffc_resnet":
return FFCResNetGenerator(**kwargs)
raise ValueError(f"Unknown generator kind {kind}")
def make_discriminator(kind, **kwargs):
logging.info(f"Make discriminator {kind}")
if kind == "pix2pixhd_nlayer_multidilated":
return MultidilatedNLayerDiscriminator(**kwargs)
if kind == "pix2pixhd_nlayer":
return NLayerDiscriminator(**kwargs)
raise ValueError(f"Unknown discriminator kind {kind}")
@@ -1,96 +0,0 @@
import abc
from typing import Tuple, List
import torch
import torch.nn as nn
from .depthwise_sep_conv import DepthWiseSeperableConv
from .multidilated_conv import MultidilatedConv
class BaseDiscriminator(nn.Module):
@abc.abstractmethod
def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, List[torch.Tensor]]:
"""
Predict scores and get intermediate activations. Useful for feature matching loss
:return tuple (scores, list of intermediate activations)
"""
raise NotImplemented()
def get_conv_block_ctor(kind="default"):
if not isinstance(kind, str):
return kind
if kind == "default":
return nn.Conv2d
if kind == "depthwise":
return DepthWiseSeperableConv
if kind == "multidilated":
return MultidilatedConv
raise ValueError(f"Unknown convolutional block kind {kind}")
def get_norm_layer(kind="bn"):
if not isinstance(kind, str):
return kind
if kind == "bn":
return nn.BatchNorm2d
if kind == "in":
return nn.InstanceNorm2d
raise ValueError(f"Unknown norm block kind {kind}")
def get_activation(kind="tanh"):
if kind == "tanh":
return nn.Tanh()
if kind == "sigmoid":
return nn.Sigmoid()
if kind is False:
return nn.Identity()
raise ValueError(f"Unknown activation kind {kind}")
class SimpleMultiStepGenerator(nn.Module):
def __init__(self, steps: List[nn.Module]):
super().__init__()
self.steps = nn.ModuleList(steps)
def forward(self, x):
cur_in = x
outs = []
for step in self.steps:
cur_out = step(cur_in)
outs.append(cur_out)
cur_in = torch.cat((cur_in, cur_out), dim=1)
return torch.cat(outs[::-1], dim=1)
def deconv_factory(kind, ngf, mult, norm_layer, activation, max_features):
if kind == "convtranspose":
return [
nn.ConvTranspose2d(
min(max_features, ngf * mult),
min(max_features, int(ngf * mult / 2)),
kernel_size=3,
stride=2,
padding=1,
output_padding=1,
),
norm_layer(min(max_features, int(ngf * mult / 2))),
activation,
]
elif kind == "bilinear":
return [
nn.Upsample(scale_factor=2, mode="bilinear"),
DepthWiseSeperableConv(
min(max_features, ngf * mult),
min(max_features, int(ngf * mult / 2)),
kernel_size=3,
stride=1,
padding=1,
),
norm_layer(min(max_features, int(ngf * mult / 2))),
activation,
]
else:
raise Exception(f"Invalid deconv kind: {kind}")
@@ -1,18 +0,0 @@
import torch
import torch.nn as nn
class DepthWiseSeperableConv(nn.Module):
def __init__(self, in_dim, out_dim, *args, **kwargs):
super().__init__()
if "groups" in kwargs:
# ignoring groups for Depthwise Sep Conv
del kwargs["groups"]
self.depthwise = nn.Conv2d(in_dim, in_dim, *args, groups=in_dim, **kwargs)
self.pointwise = nn.Conv2d(in_dim, out_dim, kernel_size=1)
def forward(self, x):
out = self.depthwise(x)
out = self.pointwise(out)
return out
@@ -1,50 +0,0 @@
import torch
from kornia.constants import SamplePadding
from kornia.augmentation import RandomAffine, CenterCrop
class FakeFakesGenerator:
def __init__(self, aug_proba=0.5, img_aug_degree=30, img_aug_translate=0.2):
self.grad_aug = RandomAffine(
degrees=360, translate=0.2, padding_mode=SamplePadding.REFLECTION, keepdim=False, p=1
)
self.img_aug = RandomAffine(
degrees=img_aug_degree,
translate=img_aug_translate,
padding_mode=SamplePadding.REFLECTION,
keepdim=True,
p=1,
)
self.aug_proba = aug_proba
def __call__(self, input_images, masks):
blend_masks = self._fill_masks_with_gradient(masks)
blend_target = self._make_blend_target(input_images)
result = input_images * (1 - blend_masks) + blend_target * blend_masks
return result, blend_masks
def _make_blend_target(self, input_images):
batch_size = input_images.shape[0]
permuted = input_images[torch.randperm(batch_size)]
augmented = self.img_aug(input_images)
is_aug = (torch.rand(batch_size, device=input_images.device)[:, None, None, None] < self.aug_proba).float()
result = augmented * is_aug + permuted * (1 - is_aug)
return result
def _fill_masks_with_gradient(self, masks):
batch_size, _, height, width = masks.shape
grad = (
torch.linspace(0, 1, steps=width * 2, device=masks.device, dtype=masks.dtype)
.view(1, 1, 1, -1)
.expand(batch_size, 1, height * 2, width * 2)
)
grad = self.grad_aug(grad)
grad = CenterCrop((height, width))(grad)
grad *= masks
grad_for_min = grad + (1 - masks) * 10
grad -= grad_for_min.view(batch_size, -1).min(-1).values[:, None, None, None]
grad /= grad.view(batch_size, -1).max(-1).values[:, None, None, None] + 1e-6
grad.clamp_(min=0, max=1)
return grad
@@ -1,589 +0,0 @@
# Fast Fourier Convolution NeurIPS 2020
# original implementation https://github.com/pkumivision/FFC/blob/main/model_zoo/ffc.py
# paper https://proceedings.neurips.cc/paper/2020/file/2fd5d41ec6cfab47e32164d5624269b1-Paper.pdf
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from .base import get_activation, BaseDiscriminator
from .spatial_transform import LearnableSpatialTransformWrapper
from .squeeze_excitation import SELayer
class FFCSE_block(nn.Module):
def __init__(self, channels, ratio_g):
super(FFCSE_block, self).__init__()
in_cg = int(channels * ratio_g)
in_cl = channels - in_cg
r = 16
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
self.conv1 = nn.Conv2d(channels, channels // r, kernel_size=1, bias=True)
self.relu1 = nn.ReLU(inplace=True)
self.conv_a2l = None if in_cl == 0 else nn.Conv2d(channels // r, in_cl, kernel_size=1, bias=True)
self.conv_a2g = None if in_cg == 0 else nn.Conv2d(channels // r, in_cg, kernel_size=1, bias=True)
self.sigmoid = nn.Sigmoid()
def forward(self, x):
x = x if type(x) is tuple else (x, 0)
id_l, id_g = x
x = id_l if type(id_g) is int else torch.cat([id_l, id_g], dim=1)
x = self.avgpool(x)
x = self.relu1(self.conv1(x))
x_l = 0 if self.conv_a2l is None else id_l * self.sigmoid(self.conv_a2l(x))
x_g = 0 if self.conv_a2g is None else id_g * self.sigmoid(self.conv_a2g(x))
return x_l, x_g
class FourierUnit(nn.Module):
def __init__(
self,
in_channels,
out_channels,
groups=1,
spatial_scale_factor=None,
spatial_scale_mode="bilinear",
spectral_pos_encoding=False,
use_se=False,
se_kwargs=None,
ffc3d=False,
fft_norm="ortho",
):
# bn_layer not used
super(FourierUnit, self).__init__()
self.groups = groups
self.conv_layer = torch.nn.Conv2d(
in_channels=in_channels * 2 + (2 if spectral_pos_encoding else 0),
out_channels=out_channels * 2,
kernel_size=1,
stride=1,
padding=0,
groups=self.groups,
bias=False,
)
self.bn = torch.nn.BatchNorm2d(out_channels * 2)
self.relu = torch.nn.ReLU(inplace=True)
# squeeze and excitation block
self.use_se = use_se
if use_se:
if se_kwargs is None:
se_kwargs = {}
self.se = SELayer(self.conv_layer.in_channels, **se_kwargs)
self.spatial_scale_factor = spatial_scale_factor
self.spatial_scale_mode = spatial_scale_mode
self.spectral_pos_encoding = spectral_pos_encoding
self.ffc3d = ffc3d
self.fft_norm = fft_norm
def forward(self, x):
batch = x.shape[0]
if self.spatial_scale_factor is not None:
orig_size = x.shape[-2:]
x = F.interpolate(
x, scale_factor=self.spatial_scale_factor, mode=self.spatial_scale_mode, align_corners=False
)
r_size = x.size()
# (batch, c, h, w/2+1, 2)
fft_dim = (-3, -2, -1) if self.ffc3d else (-2, -1)
ffted = torch.fft.rfftn(x, dim=fft_dim, norm=self.fft_norm)
ffted = torch.stack((ffted.real, ffted.imag), dim=-1)
ffted = ffted.permute(0, 1, 4, 2, 3).contiguous() # (batch, c, 2, h, w/2+1)
ffted = ffted.view(
(
batch,
-1,
)
+ ffted.size()[3:]
)
if self.spectral_pos_encoding:
height, width = ffted.shape[-2:]
coords_vert = torch.linspace(0, 1, height)[None, None, :, None].expand(batch, 1, height, width).to(ffted)
coords_hor = torch.linspace(0, 1, width)[None, None, None, :].expand(batch, 1, height, width).to(ffted)
ffted = torch.cat((coords_vert, coords_hor, ffted), dim=1)
if self.use_se:
ffted = self.se(ffted)
ffted = self.conv_layer(ffted) # (batch, c*2, h, w/2+1)
ffted = self.relu(self.bn(ffted))
ffted = (
ffted.view(
(
batch,
-1,
2,
)
+ ffted.size()[2:]
)
.permute(0, 1, 3, 4, 2)
.contiguous()
) # (batch,c, t, h, w/2+1, 2)
ffted = torch.complex(ffted[..., 0], ffted[..., 1])
ifft_shape_slice = x.shape[-3:] if self.ffc3d else x.shape[-2:]
output = torch.fft.irfftn(ffted, s=ifft_shape_slice, dim=fft_dim, norm=self.fft_norm)
if self.spatial_scale_factor is not None:
output = F.interpolate(output, size=orig_size, mode=self.spatial_scale_mode, align_corners=False)
return output
class SpectralTransform(nn.Module):
def __init__(self, in_channels, out_channels, stride=1, groups=1, enable_lfu=True, **fu_kwargs):
# bn_layer not used
super(SpectralTransform, self).__init__()
self.enable_lfu = enable_lfu
if stride == 2:
self.downsample = nn.AvgPool2d(kernel_size=(2, 2), stride=2)
else:
self.downsample = nn.Identity()
self.stride = stride
self.conv1 = nn.Sequential(
nn.Conv2d(in_channels, out_channels // 2, kernel_size=1, groups=groups, bias=False),
nn.BatchNorm2d(out_channels // 2),
nn.ReLU(inplace=True),
)
self.fu = FourierUnit(out_channels // 2, out_channels // 2, groups, **fu_kwargs)
if self.enable_lfu:
self.lfu = FourierUnit(out_channels // 2, out_channels // 2, groups)
self.conv2 = torch.nn.Conv2d(out_channels // 2, out_channels, kernel_size=1, groups=groups, bias=False)
def forward(self, x):
x = self.downsample(x)
x = self.conv1(x)
output = self.fu(x)
if self.enable_lfu:
n, c, h, w = x.shape
split_no = 2
split_s = h // split_no
xs = torch.cat(torch.split(x[:, : c // 4], split_s, dim=-2), dim=1).contiguous()
xs = torch.cat(torch.split(xs, split_s, dim=-1), dim=1).contiguous()
xs = self.lfu(xs)
xs = xs.repeat(1, 1, split_no, split_no).contiguous()
else:
xs = 0
output = self.conv2(x + output + xs)
return output
class FFC(nn.Module):
def __init__(
self,
in_channels,
out_channels,
kernel_size,
ratio_gin,
ratio_gout,
stride=1,
padding=0,
dilation=1,
groups=1,
bias=False,
enable_lfu=True,
padding_type="reflect",
gated=False,
**spectral_kwargs
):
super(FFC, self).__init__()
assert stride == 1 or stride == 2, "Stride should be 1 or 2."
self.stride = stride
in_cg = int(in_channels * ratio_gin)
in_cl = in_channels - in_cg
out_cg = int(out_channels * ratio_gout)
out_cl = out_channels - out_cg
# groups_g = 1 if groups == 1 else int(groups * ratio_gout)
# groups_l = 1 if groups == 1 else groups - groups_g
self.ratio_gin = ratio_gin
self.ratio_gout = ratio_gout
self.global_in_num = in_cg
module = nn.Identity if in_cl == 0 or out_cl == 0 else nn.Conv2d
self.convl2l = module(
in_cl, out_cl, kernel_size, stride, padding, dilation, groups, bias, padding_mode=padding_type
)
module = nn.Identity if in_cl == 0 or out_cg == 0 else nn.Conv2d
self.convl2g = module(
in_cl, out_cg, kernel_size, stride, padding, dilation, groups, bias, padding_mode=padding_type
)
module = nn.Identity if in_cg == 0 or out_cl == 0 else nn.Conv2d
self.convg2l = module(
in_cg, out_cl, kernel_size, stride, padding, dilation, groups, bias, padding_mode=padding_type
)
module = nn.Identity if in_cg == 0 or out_cg == 0 else SpectralTransform
self.convg2g = module(in_cg, out_cg, stride, 1 if groups == 1 else groups // 2, enable_lfu, **spectral_kwargs)
self.gated = gated
module = nn.Identity if in_cg == 0 or out_cl == 0 or not self.gated else nn.Conv2d
self.gate = module(in_channels, 2, 1)
def forward(self, x):
x_l, x_g = x if type(x) is tuple else (x, 0)
out_xl, out_xg = 0, 0
if self.gated:
total_input_parts = [x_l]
if torch.is_tensor(x_g):
total_input_parts.append(x_g)
total_input = torch.cat(total_input_parts, dim=1)
gates = torch.sigmoid(self.gate(total_input))
g2l_gate, l2g_gate = gates.chunk(2, dim=1)
else:
g2l_gate, l2g_gate = 1, 1
if self.ratio_gout != 1:
out_xl = self.convl2l(x_l) + self.convg2l(x_g) * g2l_gate
if self.ratio_gout != 0:
out_xg = self.convl2g(x_l) * l2g_gate + self.convg2g(x_g)
return out_xl, out_xg
class FFC_BN_ACT(nn.Module):
def __init__(
self,
in_channels,
out_channels,
kernel_size,
ratio_gin,
ratio_gout,
stride=1,
padding=0,
dilation=1,
groups=1,
bias=False,
norm_layer=nn.BatchNorm2d,
activation_layer=nn.Identity,
padding_type="reflect",
enable_lfu=True,
**kwargs
):
super(FFC_BN_ACT, self).__init__()
self.ffc = FFC(
in_channels,
out_channels,
kernel_size,
ratio_gin,
ratio_gout,
stride,
padding,
dilation,
groups,
bias,
enable_lfu,
padding_type=padding_type,
**kwargs
)
lnorm = nn.Identity if ratio_gout == 1 else norm_layer
gnorm = nn.Identity if ratio_gout == 0 else norm_layer
global_channels = int(out_channels * ratio_gout)
self.bn_l = lnorm(out_channels - global_channels)
self.bn_g = gnorm(global_channels)
lact = nn.Identity if ratio_gout == 1 else activation_layer
gact = nn.Identity if ratio_gout == 0 else activation_layer
self.act_l = lact(inplace=True)
self.act_g = gact(inplace=True)
def forward(self, x):
x_l, x_g = self.ffc(x)
x_l = self.act_l(self.bn_l(x_l))
x_g = self.act_g(self.bn_g(x_g))
return x_l, x_g
class FFCResnetBlock(nn.Module):
def __init__(
self,
dim,
padding_type,
norm_layer,
activation_layer=nn.ReLU,
dilation=1,
spatial_transform_kwargs=None,
inline=False,
**conv_kwargs
):
super().__init__()
self.conv1 = FFC_BN_ACT(
dim,
dim,
kernel_size=3,
padding=dilation,
dilation=dilation,
norm_layer=norm_layer,
activation_layer=activation_layer,
padding_type=padding_type,
**conv_kwargs
)
self.conv2 = FFC_BN_ACT(
dim,
dim,
kernel_size=3,
padding=dilation,
dilation=dilation,
norm_layer=norm_layer,
activation_layer=activation_layer,
padding_type=padding_type,
**conv_kwargs
)
if spatial_transform_kwargs is not None:
self.conv1 = LearnableSpatialTransformWrapper(self.conv1, **spatial_transform_kwargs)
self.conv2 = LearnableSpatialTransformWrapper(self.conv2, **spatial_transform_kwargs)
self.inline = inline
def forward(self, x):
if self.inline:
x_l, x_g = x[:, : -self.conv1.ffc.global_in_num], x[:, -self.conv1.ffc.global_in_num :]
else:
x_l, x_g = x if type(x) is tuple else (x, 0)
id_l, id_g = x_l, x_g
x_l, x_g = self.conv1((x_l, x_g))
x_l, x_g = self.conv2((x_l, x_g))
x_l, x_g = id_l + x_l, id_g + x_g
out = x_l, x_g
if self.inline:
out = torch.cat(out, dim=1)
return out
class ConcatTupleLayer(nn.Module):
def forward(self, x):
assert isinstance(x, tuple)
x_l, x_g = x
assert torch.is_tensor(x_l) or torch.is_tensor(x_g)
if not torch.is_tensor(x_g):
return x_l
return torch.cat(x, dim=1)
class FFCResNetGenerator(nn.Module):
def __init__(
self,
input_nc,
output_nc,
ngf=64,
n_downsampling=3,
n_blocks=9,
norm_layer=nn.BatchNorm2d,
padding_type="reflect",
activation_layer=nn.ReLU,
up_norm_layer=nn.BatchNorm2d,
up_activation=nn.ReLU(True),
init_conv_kwargs={},
downsample_conv_kwargs={},
resnet_conv_kwargs={},
spatial_transform_layers=None,
spatial_transform_kwargs={},
add_out_act=True,
max_features=1024,
out_ffc=False,
out_ffc_kwargs={},
):
assert n_blocks >= 0
super().__init__()
model = [
nn.ReflectionPad2d(3),
FFC_BN_ACT(
input_nc,
ngf,
kernel_size=7,
padding=0,
norm_layer=norm_layer,
activation_layer=activation_layer,
**init_conv_kwargs
),
]
### downsample
for i in range(n_downsampling):
mult = 2**i
if i == n_downsampling - 1:
cur_conv_kwargs = dict(downsample_conv_kwargs)
cur_conv_kwargs["ratio_gout"] = resnet_conv_kwargs.get("ratio_gin", 0)
else:
cur_conv_kwargs = downsample_conv_kwargs
model += [
FFC_BN_ACT(
min(max_features, ngf * mult),
min(max_features, ngf * mult * 2),
kernel_size=3,
stride=2,
padding=1,
norm_layer=norm_layer,
activation_layer=activation_layer,
**cur_conv_kwargs
)
]
mult = 2**n_downsampling
feats_num_bottleneck = min(max_features, ngf * mult)
### resnet blocks
for i in range(n_blocks):
cur_resblock = FFCResnetBlock(
feats_num_bottleneck,
padding_type=padding_type,
activation_layer=activation_layer,
norm_layer=norm_layer,
**resnet_conv_kwargs
)
if spatial_transform_layers is not None and i in spatial_transform_layers:
cur_resblock = LearnableSpatialTransformWrapper(cur_resblock, **spatial_transform_kwargs)
model += [cur_resblock]
model += [ConcatTupleLayer()]
### upsample
for i in range(n_downsampling):
mult = 2 ** (n_downsampling - i)
model += [
nn.ConvTranspose2d(
min(max_features, ngf * mult),
min(max_features, int(ngf * mult / 2)),
kernel_size=3,
stride=2,
padding=1,
output_padding=1,
),
up_norm_layer(min(max_features, int(ngf * mult / 2))),
up_activation,
]
if out_ffc:
model += [
FFCResnetBlock(
ngf,
padding_type=padding_type,
activation_layer=activation_layer,
norm_layer=norm_layer,
inline=True,
**out_ffc_kwargs
)
]
model += [nn.ReflectionPad2d(3), nn.Conv2d(ngf, output_nc, kernel_size=7, padding=0)]
if add_out_act:
model.append(get_activation("tanh" if add_out_act is True else add_out_act))
self.model = nn.Sequential(*model)
def forward(self, input):
return self.model(input)
class FFCNLayerDiscriminator(BaseDiscriminator):
def __init__(
self,
input_nc,
ndf=64,
n_layers=3,
norm_layer=nn.BatchNorm2d,
max_features=512,
init_conv_kwargs={},
conv_kwargs={},
):
super().__init__()
self.n_layers = n_layers
def _act_ctor(inplace=True):
return nn.LeakyReLU(negative_slope=0.2, inplace=inplace)
kw = 3
padw = int(np.ceil((kw - 1.0) / 2))
sequence = [
[
FFC_BN_ACT(
input_nc,
ndf,
kernel_size=kw,
padding=padw,
norm_layer=norm_layer,
activation_layer=_act_ctor,
**init_conv_kwargs
)
]
]
nf = ndf
for n in range(1, n_layers):
nf_prev = nf
nf = min(nf * 2, max_features)
cur_model = [
FFC_BN_ACT(
nf_prev,
nf,
kernel_size=kw,
stride=2,
padding=padw,
norm_layer=norm_layer,
activation_layer=_act_ctor,
**conv_kwargs
)
]
sequence.append(cur_model)
nf_prev = nf
nf = min(nf * 2, 512)
cur_model = [
FFC_BN_ACT(
nf_prev,
nf,
kernel_size=kw,
stride=1,
padding=padw,
norm_layer=norm_layer,
activation_layer=lambda *args, **kwargs: nn.LeakyReLU(*args, negative_slope=0.2, **kwargs),
**conv_kwargs
),
ConcatTupleLayer(),
]
sequence.append(cur_model)
sequence += [[nn.Conv2d(nf, 1, kernel_size=kw, stride=1, padding=padw)]]
for n in range(len(sequence)):
setattr(self, "model" + str(n), nn.Sequential(*sequence[n]))
def get_all_activations(self, x):
res = [x]
for n in range(self.n_layers + 2):
model = getattr(self, "model" + str(n))
res.append(model(res[-1]))
return res[1:]
def forward(self, x):
act = self.get_all_activations(x)
feats = []
for out in act[:-1]:
if isinstance(out, tuple):
if torch.is_tensor(out[1]):
out = torch.cat(out, dim=1)
else:
out = out[0]
feats.append(out)
return act[-1], feats
@@ -1,117 +0,0 @@
import torch
import torch.nn as nn
import random
from .depthwise_sep_conv import DepthWiseSeperableConv
class MultidilatedConv(nn.Module):
def __init__(
self,
in_dim,
out_dim,
kernel_size,
dilation_num=3,
comb_mode="sum",
equal_dim=True,
shared_weights=False,
padding=1,
min_dilation=1,
shuffle_in_channels=False,
use_depthwise=False,
**kwargs
):
super().__init__()
convs = []
self.equal_dim = equal_dim
assert comb_mode in ("cat_out", "sum", "cat_in", "cat_both"), comb_mode
if comb_mode in ("cat_out", "cat_both"):
self.cat_out = True
if equal_dim:
assert out_dim % dilation_num == 0
out_dims = [out_dim // dilation_num] * dilation_num
self.index = sum(
[[i + j * (out_dims[0]) for j in range(dilation_num)] for i in range(out_dims[0])], []
)
else:
out_dims = [out_dim // 2 ** (i + 1) for i in range(dilation_num - 1)]
out_dims.append(out_dim - sum(out_dims))
index = []
starts = [0] + out_dims[:-1]
lengths = [out_dims[i] // out_dims[-1] for i in range(dilation_num)]
for i in range(out_dims[-1]):
for j in range(dilation_num):
index += list(range(starts[j], starts[j] + lengths[j]))
starts[j] += lengths[j]
self.index = index
assert len(index) == out_dim
self.out_dims = out_dims
else:
self.cat_out = False
self.out_dims = [out_dim] * dilation_num
if comb_mode in ("cat_in", "cat_both"):
if equal_dim:
assert in_dim % dilation_num == 0
in_dims = [in_dim // dilation_num] * dilation_num
else:
in_dims = [in_dim // 2 ** (i + 1) for i in range(dilation_num - 1)]
in_dims.append(in_dim - sum(in_dims))
self.in_dims = in_dims
self.cat_in = True
else:
self.cat_in = False
self.in_dims = [in_dim] * dilation_num
conv_type = DepthWiseSeperableConv if use_depthwise else nn.Conv2d
dilation = min_dilation
for i in range(dilation_num):
if isinstance(padding, int):
cur_padding = padding * dilation
else:
cur_padding = padding[i]
convs.append(
conv_type(
self.in_dims[i], self.out_dims[i], kernel_size, padding=cur_padding, dilation=dilation, **kwargs
)
)
if i > 0 and shared_weights:
convs[-1].weight = convs[0].weight
convs[-1].bias = convs[0].bias
dilation *= 2
self.convs = nn.ModuleList(convs)
self.shuffle_in_channels = shuffle_in_channels
if self.shuffle_in_channels:
# shuffle list as shuffling of tensors is nondeterministic
in_channels_permute = list(range(in_dim))
random.shuffle(in_channels_permute)
# save as buffer so it is saved and loaded with checkpoint
self.register_buffer("in_channels_permute", torch.tensor(in_channels_permute))
def forward(self, x):
if self.shuffle_in_channels:
x = x[:, self.in_channels_permute]
outs = []
if self.cat_in:
if self.equal_dim:
x = x.chunk(len(self.convs), dim=1)
else:
new_x = []
start = 0
for dim in self.in_dims:
new_x.append(x[:, start : start + dim])
start += dim
x = new_x
for i, conv in enumerate(self.convs):
if self.cat_in:
input = x[i]
else:
input = x
outs.append(conv(input))
if self.cat_out:
out = torch.cat(outs, dim=1)[:, self.index]
else:
out = sum(outs)
return out
@@ -1,338 +0,0 @@
from typing import List, Tuple, Union, Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
from .base import get_conv_block_ctor, get_activation
from .pix2pixhd import ResnetBlock
class ResNetHead(nn.Module):
def __init__(
self,
input_nc,
ngf=64,
n_downsampling=3,
n_blocks=9,
norm_layer=nn.BatchNorm2d,
padding_type="reflect",
conv_kind="default",
activation=nn.ReLU(True),
):
assert n_blocks >= 0
super(ResNetHead, self).__init__()
conv_layer = get_conv_block_ctor(conv_kind)
model = [
nn.ReflectionPad2d(3),
conv_layer(input_nc, ngf, kernel_size=7, padding=0),
norm_layer(ngf),
activation,
]
### downsample
for i in range(n_downsampling):
mult = 2**i
model += [
conv_layer(ngf * mult, ngf * mult * 2, kernel_size=3, stride=2, padding=1),
norm_layer(ngf * mult * 2),
activation,
]
mult = 2**n_downsampling
### resnet blocks
for i in range(n_blocks):
model += [
ResnetBlock(
ngf * mult,
padding_type=padding_type,
activation=activation,
norm_layer=norm_layer,
conv_kind=conv_kind,
)
]
self.model = nn.Sequential(*model)
def forward(self, input):
return self.model(input)
class ResNetTail(nn.Module):
def __init__(
self,
output_nc,
ngf=64,
n_downsampling=3,
n_blocks=9,
norm_layer=nn.BatchNorm2d,
padding_type="reflect",
conv_kind="default",
activation=nn.ReLU(True),
up_norm_layer=nn.BatchNorm2d,
up_activation=nn.ReLU(True),
add_out_act=False,
out_extra_layers_n=0,
add_in_proj=None,
):
assert n_blocks >= 0
super(ResNetTail, self).__init__()
mult = 2**n_downsampling
model = []
if add_in_proj is not None:
model.append(nn.Conv2d(add_in_proj, ngf * mult, kernel_size=1))
### resnet blocks
for i in range(n_blocks):
model += [
ResnetBlock(
ngf * mult,
padding_type=padding_type,
activation=activation,
norm_layer=norm_layer,
conv_kind=conv_kind,
)
]
### upsample
for i in range(n_downsampling):
mult = 2 ** (n_downsampling - i)
model += [
nn.ConvTranspose2d(
ngf * mult, int(ngf * mult / 2), kernel_size=3, stride=2, padding=1, output_padding=1
),
up_norm_layer(int(ngf * mult / 2)),
up_activation,
]
self.model = nn.Sequential(*model)
out_layers = []
for _ in range(out_extra_layers_n):
out_layers += [nn.Conv2d(ngf, ngf, kernel_size=1, padding=0), up_norm_layer(ngf), up_activation]
out_layers += [nn.ReflectionPad2d(3), nn.Conv2d(ngf, output_nc, kernel_size=7, padding=0)]
if add_out_act:
out_layers.append(get_activation("tanh" if add_out_act is True else add_out_act))
self.out_proj = nn.Sequential(*out_layers)
def forward(self, input, return_last_act=False):
features = self.model(input)
out = self.out_proj(features)
if return_last_act:
return out, features
else:
return out
class MultiscaleResNet(nn.Module):
def __init__(
self,
input_nc,
output_nc,
ngf=64,
n_downsampling=2,
n_blocks_head=2,
n_blocks_tail=6,
n_scales=3,
norm_layer=nn.BatchNorm2d,
padding_type="reflect",
conv_kind="default",
activation=nn.ReLU(True),
up_norm_layer=nn.BatchNorm2d,
up_activation=nn.ReLU(True),
add_out_act=False,
out_extra_layers_n=0,
out_cumulative=False,
return_only_hr=False,
):
super().__init__()
self.heads = nn.ModuleList(
[
ResNetHead(
input_nc,
ngf=ngf,
n_downsampling=n_downsampling,
n_blocks=n_blocks_head,
norm_layer=norm_layer,
padding_type=padding_type,
conv_kind=conv_kind,
activation=activation,
)
for i in range(n_scales)
]
)
tail_in_feats = ngf * (2**n_downsampling) + ngf
self.tails = nn.ModuleList(
[
ResNetTail(
output_nc,
ngf=ngf,
n_downsampling=n_downsampling,
n_blocks=n_blocks_tail,
norm_layer=norm_layer,
padding_type=padding_type,
conv_kind=conv_kind,
activation=activation,
up_norm_layer=up_norm_layer,
up_activation=up_activation,
add_out_act=add_out_act,
out_extra_layers_n=out_extra_layers_n,
add_in_proj=None if (i == n_scales - 1) else tail_in_feats,
)
for i in range(n_scales)
]
)
self.out_cumulative = out_cumulative
self.return_only_hr = return_only_hr
@property
def num_scales(self):
return len(self.heads)
def forward(
self, ms_inputs: List[torch.Tensor], smallest_scales_num: Optional[int] = None
) -> Union[torch.Tensor, List[torch.Tensor]]:
"""
:param ms_inputs: List of inputs of different resolutions from HR to LR
:param smallest_scales_num: int or None, number of smallest scales to take at input
:return: Depending on return_only_hr:
True: Only the most HR output
False: List of outputs of different resolutions from HR to LR
"""
if smallest_scales_num is None:
assert len(self.heads) == len(ms_inputs), (len(self.heads), len(ms_inputs), smallest_scales_num)
smallest_scales_num = len(self.heads)
else:
assert smallest_scales_num == len(ms_inputs) <= len(self.heads), (
len(self.heads),
len(ms_inputs),
smallest_scales_num,
)
cur_heads = self.heads[-smallest_scales_num:]
ms_features = [cur_head(cur_inp) for cur_head, cur_inp in zip(cur_heads, ms_inputs)]
all_outputs = []
prev_tail_features = None
for i in range(len(ms_features)):
scale_i = -i - 1
cur_tail_input = ms_features[-i - 1]
if prev_tail_features is not None:
if prev_tail_features.shape != cur_tail_input.shape:
prev_tail_features = F.interpolate(
prev_tail_features, size=cur_tail_input.shape[2:], mode="bilinear", align_corners=False
)
cur_tail_input = torch.cat((cur_tail_input, prev_tail_features), dim=1)
cur_out, cur_tail_feats = self.tails[scale_i](cur_tail_input, return_last_act=True)
prev_tail_features = cur_tail_feats
all_outputs.append(cur_out)
if self.out_cumulative:
all_outputs_cum = [all_outputs[0]]
for i in range(1, len(ms_features)):
cur_out = all_outputs[i]
cur_out_cum = cur_out + F.interpolate(
all_outputs_cum[-1], size=cur_out.shape[2:], mode="bilinear", align_corners=False
)
all_outputs_cum.append(cur_out_cum)
all_outputs = all_outputs_cum
if self.return_only_hr:
return all_outputs[-1]
else:
return all_outputs[::-1]
class MultiscaleDiscriminatorSimple(nn.Module):
def __init__(self, ms_impl):
super().__init__()
self.ms_impl = nn.ModuleList(ms_impl)
@property
def num_scales(self):
return len(self.ms_impl)
def forward(
self, ms_inputs: List[torch.Tensor], smallest_scales_num: Optional[int] = None
) -> List[Tuple[torch.Tensor, List[torch.Tensor]]]:
"""
:param ms_inputs: List of inputs of different resolutions from HR to LR
:param smallest_scales_num: int or None, number of smallest scales to take at input
:return: List of pairs (prediction, features) for different resolutions from HR to LR
"""
if smallest_scales_num is None:
assert len(self.ms_impl) == len(ms_inputs), (len(self.ms_impl), len(ms_inputs), smallest_scales_num)
smallest_scales_num = len(self.heads)
else:
assert smallest_scales_num == len(ms_inputs) <= len(self.ms_impl), (
len(self.ms_impl),
len(ms_inputs),
smallest_scales_num,
)
return [cur_discr(cur_input) for cur_discr, cur_input in zip(self.ms_impl[-smallest_scales_num:], ms_inputs)]
class SingleToMultiScaleInputMixin:
def forward(self, x: torch.Tensor) -> List:
orig_height, orig_width = x.shape[2:]
factors = [2**i for i in range(self.num_scales)]
ms_inputs = [
F.interpolate(x, size=(orig_height // f, orig_width // f), mode="bilinear", align_corners=False)
for f in factors
]
return super().forward(ms_inputs)
class GeneratorMultiToSingleOutputMixin:
def forward(self, x):
return super().forward(x)[0]
class DiscriminatorMultiToSingleOutputMixin:
def forward(self, x):
out_feat_tuples = super().forward(x)
return out_feat_tuples[0][0], [f for _, flist in out_feat_tuples for f in flist]
class DiscriminatorMultiToSingleOutputStackedMixin:
def __init__(self, *args, return_feats_only_levels=None, **kwargs):
super().__init__(*args, **kwargs)
self.return_feats_only_levels = return_feats_only_levels
def forward(self, x):
out_feat_tuples = super().forward(x)
outs = [out for out, _ in out_feat_tuples]
scaled_outs = [outs[0]] + [
F.interpolate(cur_out, size=outs[0].shape[-2:], mode="bilinear", align_corners=False)
for cur_out in outs[1:]
]
out = torch.cat(scaled_outs, dim=1)
if self.return_feats_only_levels is not None:
feat_lists = [out_feat_tuples[i][1] for i in self.return_feats_only_levels]
else:
feat_lists = [flist for _, flist in out_feat_tuples]
feats = [f for flist in feat_lists for f in flist]
return out, feats
class MultiscaleDiscrSingleInput(
SingleToMultiScaleInputMixin, DiscriminatorMultiToSingleOutputStackedMixin, MultiscaleDiscriminatorSimple
):
pass
class MultiscaleResNetSingle(GeneratorMultiToSingleOutputMixin, SingleToMultiScaleInputMixin, MultiscaleResNet):
pass
@@ -1,893 +0,0 @@
# original: https://github.com/NVIDIA/pix2pixHD/blob/master/models/networks.py
import collections
from functools import partial
import functools
import logging
from collections import defaultdict
import numpy as np
import torch.nn as nn
from .base import BaseDiscriminator, deconv_factory, get_conv_block_ctor, get_norm_layer, get_activation
from .ffc import FFCResnetBlock
from .multidilated_conv import MultidilatedConv
class DotDict(defaultdict):
# https://stackoverflow.com/questions/2352181/how-to-use-a-dot-to-access-members-of-dictionary
"""dot.notation access to dictionary attributes"""
__getattr__ = defaultdict.get
__setattr__ = defaultdict.__setitem__
__delattr__ = defaultdict.__delitem__
class Identity(nn.Module):
def __init__(self):
super().__init__()
def forward(self, x):
return x
class ResnetBlock(nn.Module):
def __init__(
self,
dim,
padding_type,
norm_layer,
activation=nn.ReLU(True),
use_dropout=False,
conv_kind="default",
dilation=1,
in_dim=None,
groups=1,
second_dilation=None,
):
super(ResnetBlock, self).__init__()
self.in_dim = in_dim
self.dim = dim
if second_dilation is None:
second_dilation = dilation
self.conv_block = self.build_conv_block(
dim,
padding_type,
norm_layer,
activation,
use_dropout,
conv_kind=conv_kind,
dilation=dilation,
in_dim=in_dim,
groups=groups,
second_dilation=second_dilation,
)
if self.in_dim is not None:
self.input_conv = nn.Conv2d(in_dim, dim, 1)
self.out_channnels = dim
def build_conv_block(
self,
dim,
padding_type,
norm_layer,
activation,
use_dropout,
conv_kind="default",
dilation=1,
in_dim=None,
groups=1,
second_dilation=1,
):
conv_layer = get_conv_block_ctor(conv_kind)
conv_block = []
p = 0
if padding_type == "reflect":
conv_block += [nn.ReflectionPad2d(dilation)]
elif padding_type == "replicate":
conv_block += [nn.ReplicationPad2d(dilation)]
elif padding_type == "zero":
p = dilation
else:
raise NotImplementedError("padding [%s] is not implemented" % padding_type)
if in_dim is None:
in_dim = dim
conv_block += [
conv_layer(in_dim, dim, kernel_size=3, padding=p, dilation=dilation),
norm_layer(dim),
activation,
]
if use_dropout:
conv_block += [nn.Dropout(0.5)]
p = 0
if padding_type == "reflect":
conv_block += [nn.ReflectionPad2d(second_dilation)]
elif padding_type == "replicate":
conv_block += [nn.ReplicationPad2d(second_dilation)]
elif padding_type == "zero":
p = second_dilation
else:
raise NotImplementedError("padding [%s] is not implemented" % padding_type)
conv_block += [
conv_layer(dim, dim, kernel_size=3, padding=p, dilation=second_dilation, groups=groups),
norm_layer(dim),
]
return nn.Sequential(*conv_block)
def forward(self, x):
x_before = x
if self.in_dim is not None:
x = self.input_conv(x)
out = x + self.conv_block(x_before)
return out
class ResnetBlock5x5(nn.Module):
def __init__(
self,
dim,
padding_type,
norm_layer,
activation=nn.ReLU(True),
use_dropout=False,
conv_kind="default",
dilation=1,
in_dim=None,
groups=1,
second_dilation=None,
):
super(ResnetBlock5x5, self).__init__()
self.in_dim = in_dim
self.dim = dim
if second_dilation is None:
second_dilation = dilation
self.conv_block = self.build_conv_block(
dim,
padding_type,
norm_layer,
activation,
use_dropout,
conv_kind=conv_kind,
dilation=dilation,
in_dim=in_dim,
groups=groups,
second_dilation=second_dilation,
)
if self.in_dim is not None:
self.input_conv = nn.Conv2d(in_dim, dim, 1)
self.out_channnels = dim
def build_conv_block(
self,
dim,
padding_type,
norm_layer,
activation,
use_dropout,
conv_kind="default",
dilation=1,
in_dim=None,
groups=1,
second_dilation=1,
):
conv_layer = get_conv_block_ctor(conv_kind)
conv_block = []
p = 0
if padding_type == "reflect":
conv_block += [nn.ReflectionPad2d(dilation * 2)]
elif padding_type == "replicate":
conv_block += [nn.ReplicationPad2d(dilation * 2)]
elif padding_type == "zero":
p = dilation * 2
else:
raise NotImplementedError("padding [%s] is not implemented" % padding_type)
if in_dim is None:
in_dim = dim
conv_block += [
conv_layer(in_dim, dim, kernel_size=5, padding=p, dilation=dilation),
norm_layer(dim),
activation,
]
if use_dropout:
conv_block += [nn.Dropout(0.5)]
p = 0
if padding_type == "reflect":
conv_block += [nn.ReflectionPad2d(second_dilation * 2)]
elif padding_type == "replicate":
conv_block += [nn.ReplicationPad2d(second_dilation * 2)]
elif padding_type == "zero":
p = second_dilation * 2
else:
raise NotImplementedError("padding [%s] is not implemented" % padding_type)
conv_block += [
conv_layer(dim, dim, kernel_size=5, padding=p, dilation=second_dilation, groups=groups),
norm_layer(dim),
]
return nn.Sequential(*conv_block)
def forward(self, x):
x_before = x
if self.in_dim is not None:
x = self.input_conv(x)
out = x + self.conv_block(x_before)
return out
class MultidilatedResnetBlock(nn.Module):
def __init__(self, dim, padding_type, conv_layer, norm_layer, activation=nn.ReLU(True), use_dropout=False):
super().__init__()
self.conv_block = self.build_conv_block(dim, padding_type, conv_layer, norm_layer, activation, use_dropout)
def build_conv_block(self, dim, padding_type, conv_layer, norm_layer, activation, use_dropout, dilation=1):
conv_block = []
conv_block += [conv_layer(dim, dim, kernel_size=3, padding_mode=padding_type), norm_layer(dim), activation]
if use_dropout:
conv_block += [nn.Dropout(0.5)]
conv_block += [conv_layer(dim, dim, kernel_size=3, padding_mode=padding_type), norm_layer(dim)]
return nn.Sequential(*conv_block)
def forward(self, x):
out = x + self.conv_block(x)
return out
class MultiDilatedGlobalGenerator(nn.Module):
def __init__(
self,
input_nc,
output_nc,
ngf=64,
n_downsampling=3,
n_blocks=3,
norm_layer=nn.BatchNorm2d,
padding_type="reflect",
conv_kind="default",
deconv_kind="convtranspose",
activation=nn.ReLU(True),
up_norm_layer=nn.BatchNorm2d,
affine=None,
up_activation=nn.ReLU(True),
add_out_act=True,
max_features=1024,
multidilation_kwargs={},
ffc_positions=None,
ffc_kwargs={},
):
assert n_blocks >= 0
super().__init__()
conv_layer = get_conv_block_ctor(conv_kind)
resnet_conv_layer = functools.partial(get_conv_block_ctor("multidilated"), **multidilation_kwargs)
norm_layer = get_norm_layer(norm_layer)
if affine is not None:
norm_layer = partial(norm_layer, affine=affine)
up_norm_layer = get_norm_layer(up_norm_layer)
if affine is not None:
up_norm_layer = partial(up_norm_layer, affine=affine)
model = [
nn.ReflectionPad2d(3),
conv_layer(input_nc, ngf, kernel_size=7, padding=0),
norm_layer(ngf),
activation,
]
identity = Identity()
### downsample
for i in range(n_downsampling):
mult = 2**i
model += [
conv_layer(
min(max_features, ngf * mult),
min(max_features, ngf * mult * 2),
kernel_size=3,
stride=2,
padding=1,
),
norm_layer(min(max_features, ngf * mult * 2)),
activation,
]
mult = 2**n_downsampling
feats_num_bottleneck = min(max_features, ngf * mult)
### resnet blocks
for i in range(n_blocks):
if ffc_positions is not None and i in ffc_positions:
model += [
FFCResnetBlock(
feats_num_bottleneck,
padding_type,
norm_layer,
activation_layer=nn.ReLU,
inline=True,
**ffc_kwargs,
)
]
model += [
MultidilatedResnetBlock(
feats_num_bottleneck,
padding_type=padding_type,
conv_layer=resnet_conv_layer,
activation=activation,
norm_layer=norm_layer,
)
]
### upsample
for i in range(n_downsampling):
mult = 2 ** (n_downsampling - i)
model += deconv_factory(deconv_kind, ngf, mult, up_norm_layer, up_activation, max_features)
model += [nn.ReflectionPad2d(3), nn.Conv2d(ngf, output_nc, kernel_size=7, padding=0)]
if add_out_act:
model.append(get_activation("tanh" if add_out_act is True else add_out_act))
self.model = nn.Sequential(*model)
def forward(self, input):
return self.model(input)
class ConfigGlobalGenerator(nn.Module):
def __init__(
self,
input_nc,
output_nc,
ngf=64,
n_downsampling=3,
n_blocks=3,
norm_layer=nn.BatchNorm2d,
padding_type="reflect",
conv_kind="default",
deconv_kind="convtranspose",
activation=nn.ReLU(True),
up_norm_layer=nn.BatchNorm2d,
affine=None,
up_activation=nn.ReLU(True),
add_out_act=True,
max_features=1024,
manual_block_spec=[],
resnet_block_kind="multidilatedresnetblock",
resnet_conv_kind="multidilated",
resnet_dilation=1,
multidilation_kwargs={},
):
assert n_blocks >= 0
super().__init__()
conv_layer = get_conv_block_ctor(conv_kind)
resnet_conv_layer = functools.partial(get_conv_block_ctor(resnet_conv_kind), **multidilation_kwargs)
norm_layer = get_norm_layer(norm_layer)
if affine is not None:
norm_layer = partial(norm_layer, affine=affine)
up_norm_layer = get_norm_layer(up_norm_layer)
if affine is not None:
up_norm_layer = partial(up_norm_layer, affine=affine)
model = [
nn.ReflectionPad2d(3),
conv_layer(input_nc, ngf, kernel_size=7, padding=0),
norm_layer(ngf),
activation,
]
identity = Identity()
### downsample
for i in range(n_downsampling):
mult = 2**i
model += [
conv_layer(
min(max_features, ngf * mult),
min(max_features, ngf * mult * 2),
kernel_size=3,
stride=2,
padding=1,
),
norm_layer(min(max_features, ngf * mult * 2)),
activation,
]
mult = 2**n_downsampling
feats_num_bottleneck = min(max_features, ngf * mult)
if len(manual_block_spec) == 0:
manual_block_spec = [DotDict(lambda: None, {"n_blocks": n_blocks, "use_default": True})]
### resnet blocks
for block_spec in manual_block_spec:
def make_and_add_blocks(model, block_spec):
block_spec = DotDict(lambda: None, block_spec)
if not block_spec.use_default:
resnet_conv_layer = functools.partial(
get_conv_block_ctor(block_spec.resnet_conv_kind), **block_spec.multidilation_kwargs
)
resnet_conv_kind = block_spec.resnet_conv_kind
resnet_block_kind = block_spec.resnet_block_kind
if block_spec.resnet_dilation is not None:
resnet_dilation = block_spec.resnet_dilation
for i in range(block_spec.n_blocks):
if resnet_block_kind == "multidilatedresnetblock":
model += [
MultidilatedResnetBlock(
feats_num_bottleneck,
padding_type=padding_type,
conv_layer=resnet_conv_layer,
activation=activation,
norm_layer=norm_layer,
)
]
if resnet_block_kind == "resnetblock":
model += [
ResnetBlock(
ngf * mult,
padding_type=padding_type,
activation=activation,
norm_layer=norm_layer,
conv_kind=resnet_conv_kind,
)
]
if resnet_block_kind == "resnetblock5x5":
model += [
ResnetBlock5x5(
ngf * mult,
padding_type=padding_type,
activation=activation,
norm_layer=norm_layer,
conv_kind=resnet_conv_kind,
)
]
if resnet_block_kind == "resnetblockdwdil":
model += [
ResnetBlock(
ngf * mult,
padding_type=padding_type,
activation=activation,
norm_layer=norm_layer,
conv_kind=resnet_conv_kind,
dilation=resnet_dilation,
second_dilation=resnet_dilation,
)
]
make_and_add_blocks(model, block_spec)
### upsample
for i in range(n_downsampling):
mult = 2 ** (n_downsampling - i)
model += deconv_factory(deconv_kind, ngf, mult, up_norm_layer, up_activation, max_features)
model += [nn.ReflectionPad2d(3), nn.Conv2d(ngf, output_nc, kernel_size=7, padding=0)]
if add_out_act:
model.append(get_activation("tanh" if add_out_act is True else add_out_act))
self.model = nn.Sequential(*model)
def forward(self, input):
return self.model(input)
def make_dil_blocks(dilated_blocks_n, dilation_block_kind, dilated_block_kwargs):
blocks = []
for i in range(dilated_blocks_n):
if dilation_block_kind == "simple":
blocks.append(ResnetBlock(**dilated_block_kwargs, dilation=2 ** (i + 1)))
elif dilation_block_kind == "multi":
blocks.append(MultidilatedResnetBlock(**dilated_block_kwargs))
else:
raise ValueError(f'dilation_block_kind could not be "{dilation_block_kind}"')
return blocks
class GlobalGenerator(nn.Module):
def __init__(
self,
input_nc,
output_nc,
ngf=64,
n_downsampling=3,
n_blocks=9,
norm_layer=nn.BatchNorm2d,
padding_type="reflect",
conv_kind="default",
activation=nn.ReLU(True),
up_norm_layer=nn.BatchNorm2d,
affine=None,
up_activation=nn.ReLU(True),
dilated_blocks_n=0,
dilated_blocks_n_start=0,
dilated_blocks_n_middle=0,
add_out_act=True,
max_features=1024,
is_resblock_depthwise=False,
ffc_positions=None,
ffc_kwargs={},
dilation=1,
second_dilation=None,
dilation_block_kind="simple",
multidilation_kwargs={},
):
assert n_blocks >= 0
super().__init__()
conv_layer = get_conv_block_ctor(conv_kind)
norm_layer = get_norm_layer(norm_layer)
if affine is not None:
norm_layer = partial(norm_layer, affine=affine)
up_norm_layer = get_norm_layer(up_norm_layer)
if affine is not None:
up_norm_layer = partial(up_norm_layer, affine=affine)
if ffc_positions is not None:
ffc_positions = collections.Counter(ffc_positions)
model = [
nn.ReflectionPad2d(3),
conv_layer(input_nc, ngf, kernel_size=7, padding=0),
norm_layer(ngf),
activation,
]
identity = Identity()
### downsample
for i in range(n_downsampling):
mult = 2**i
model += [
conv_layer(
min(max_features, ngf * mult),
min(max_features, ngf * mult * 2),
kernel_size=3,
stride=2,
padding=1,
),
norm_layer(min(max_features, ngf * mult * 2)),
activation,
]
mult = 2**n_downsampling
feats_num_bottleneck = min(max_features, ngf * mult)
dilated_block_kwargs = dict(
dim=feats_num_bottleneck, padding_type=padding_type, activation=activation, norm_layer=norm_layer
)
if dilation_block_kind == "simple":
dilated_block_kwargs["conv_kind"] = conv_kind
elif dilation_block_kind == "multi":
dilated_block_kwargs["conv_layer"] = functools.partial(
get_conv_block_ctor("multidilated"), **multidilation_kwargs
)
# dilated blocks at the start of the bottleneck sausage
if dilated_blocks_n_start is not None and dilated_blocks_n_start > 0:
model += make_dil_blocks(dilated_blocks_n_start, dilation_block_kind, dilated_block_kwargs)
# resnet blocks
for i in range(n_blocks):
# dilated blocks at the middle of the bottleneck sausage
if i == n_blocks // 2 and dilated_blocks_n_middle is not None and dilated_blocks_n_middle > 0:
model += make_dil_blocks(dilated_blocks_n_middle, dilation_block_kind, dilated_block_kwargs)
if ffc_positions is not None and i in ffc_positions:
for _ in range(ffc_positions[i]): # same position can occur more than once
model += [
FFCResnetBlock(
feats_num_bottleneck,
padding_type,
norm_layer,
activation_layer=nn.ReLU,
inline=True,
**ffc_kwargs,
)
]
if is_resblock_depthwise:
resblock_groups = feats_num_bottleneck
else:
resblock_groups = 1
model += [
ResnetBlock(
feats_num_bottleneck,
padding_type=padding_type,
activation=activation,
norm_layer=norm_layer,
conv_kind=conv_kind,
groups=resblock_groups,
dilation=dilation,
second_dilation=second_dilation,
)
]
# dilated blocks at the end of the bottleneck sausage
if dilated_blocks_n is not None and dilated_blocks_n > 0:
model += make_dil_blocks(dilated_blocks_n, dilation_block_kind, dilated_block_kwargs)
# upsample
for i in range(n_downsampling):
mult = 2 ** (n_downsampling - i)
model += [
nn.ConvTranspose2d(
min(max_features, ngf * mult),
min(max_features, int(ngf * mult / 2)),
kernel_size=3,
stride=2,
padding=1,
output_padding=1,
),
up_norm_layer(min(max_features, int(ngf * mult / 2))),
up_activation,
]
model += [nn.ReflectionPad2d(3), nn.Conv2d(ngf, output_nc, kernel_size=7, padding=0)]
if add_out_act:
model.append(get_activation("tanh" if add_out_act is True else add_out_act))
self.model = nn.Sequential(*model)
def forward(self, input):
return self.model(input)
class GlobalGeneratorGated(GlobalGenerator):
def __init__(self, *args, **kwargs):
real_kwargs = dict(conv_kind="gated_bn_relu", activation=nn.Identity(), norm_layer=nn.Identity)
real_kwargs.update(kwargs)
super().__init__(*args, **real_kwargs)
class GlobalGeneratorFromSuperChannels(nn.Module):
def __init__(
self,
input_nc,
output_nc,
n_downsampling,
n_blocks,
super_channels,
norm_layer="bn",
padding_type="reflect",
add_out_act=True,
):
super().__init__()
self.n_downsampling = n_downsampling
norm_layer = get_norm_layer(norm_layer)
if type(norm_layer) == functools.partial:
use_bias = norm_layer.func == nn.InstanceNorm2d
else:
use_bias = norm_layer == nn.InstanceNorm2d
channels = self.convert_super_channels(super_channels)
self.channels = channels
model = [
nn.ReflectionPad2d(3),
nn.Conv2d(input_nc, channels[0], kernel_size=7, padding=0, bias=use_bias),
norm_layer(channels[0]),
nn.ReLU(True),
]
for i in range(n_downsampling): # add downsampling layers
mult = 2**i
model += [
nn.Conv2d(channels[0 + i], channels[1 + i], kernel_size=3, stride=2, padding=1, bias=use_bias),
norm_layer(channels[1 + i]),
nn.ReLU(True),
]
mult = 2**n_downsampling
n_blocks1 = n_blocks // 3
n_blocks2 = n_blocks1
n_blocks3 = n_blocks - n_blocks1 - n_blocks2
for i in range(n_blocks1):
c = n_downsampling
dim = channels[c]
model += [ResnetBlock(dim, padding_type=padding_type, norm_layer=norm_layer)]
for i in range(n_blocks2):
c = n_downsampling + 1
dim = channels[c]
kwargs = {}
if i == 0:
kwargs = {"in_dim": channels[c - 1]}
model += [ResnetBlock(dim, padding_type=padding_type, norm_layer=norm_layer, **kwargs)]
for i in range(n_blocks3):
c = n_downsampling + 2
dim = channels[c]
kwargs = {}
if i == 0:
kwargs = {"in_dim": channels[c - 1]}
model += [ResnetBlock(dim, padding_type=padding_type, norm_layer=norm_layer, **kwargs)]
for i in range(n_downsampling): # add upsampling layers
mult = 2 ** (n_downsampling - i)
model += [
nn.ConvTranspose2d(
channels[n_downsampling + 3 + i],
channels[n_downsampling + 3 + i + 1],
kernel_size=3,
stride=2,
padding=1,
output_padding=1,
bias=use_bias,
),
norm_layer(channels[n_downsampling + 3 + i + 1]),
nn.ReLU(True),
]
model += [nn.ReflectionPad2d(3)]
model += [nn.Conv2d(channels[2 * n_downsampling + 3], output_nc, kernel_size=7, padding=0)]
if add_out_act:
model.append(get_activation("tanh" if add_out_act is True else add_out_act))
self.model = nn.Sequential(*model)
def convert_super_channels(self, super_channels):
n_downsampling = self.n_downsampling
result = []
cnt = 0
if n_downsampling == 2:
N1 = 10
elif n_downsampling == 3:
N1 = 13
else:
raise NotImplementedError
for i in range(0, N1):
if i in [1, 4, 7, 10]:
channel = super_channels[cnt] * (2**cnt)
config = {"channel": channel}
result.append(channel)
logging.info(f"Downsample channels {result[-1]}")
cnt += 1
for i in range(3):
for counter, j in enumerate(range(N1 + i * 3, N1 + 3 + i * 3)):
if len(super_channels) == 6:
channel = super_channels[3] * 4
else:
channel = super_channels[i + 3] * 4
config = {"channel": channel}
if counter == 0:
result.append(channel)
logging.info(f"Bottleneck channels {result[-1]}")
cnt = 2
for i in range(N1 + 9, N1 + 21):
if i in [22, 25, 28]:
cnt -= 1
if len(super_channels) == 6:
channel = super_channels[5 - cnt] * (2**cnt)
else:
channel = super_channels[7 - cnt] * (2**cnt)
result.append(int(channel))
logging.info(f"Upsample channels {result[-1]}")
return result
def forward(self, input):
return self.model(input)
# Defines the PatchGAN discriminator with the specified arguments.
class NLayerDiscriminator(BaseDiscriminator):
def __init__(
self,
input_nc,
ndf=64,
n_layers=3,
norm_layer=nn.BatchNorm2d,
):
super().__init__()
self.n_layers = n_layers
kw = 4
padw = int(np.ceil((kw - 1.0) / 2))
sequence = [[nn.Conv2d(input_nc, ndf, kernel_size=kw, stride=2, padding=padw), nn.LeakyReLU(0.2, True)]]
nf = ndf
for n in range(1, n_layers):
nf_prev = nf
nf = min(nf * 2, 512)
cur_model = []
cur_model += [
nn.Conv2d(nf_prev, nf, kernel_size=kw, stride=2, padding=padw),
norm_layer(nf),
nn.LeakyReLU(0.2, True),
]
sequence.append(cur_model)
nf_prev = nf
nf = min(nf * 2, 512)
cur_model = []
cur_model += [
nn.Conv2d(nf_prev, nf, kernel_size=kw, stride=1, padding=padw),
norm_layer(nf),
nn.LeakyReLU(0.2, True),
]
sequence.append(cur_model)
sequence += [[nn.Conv2d(nf, 1, kernel_size=kw, stride=1, padding=padw)]]
for n in range(len(sequence)):
setattr(self, "model" + str(n), nn.Sequential(*sequence[n]))
def get_all_activations(self, x):
res = [x]
for n in range(self.n_layers + 2):
model = getattr(self, "model" + str(n))
res.append(model(res[-1]))
return res[1:]
def forward(self, x):
act = self.get_all_activations(x)
return act[-1], act[:-1]
class MultidilatedNLayerDiscriminator(BaseDiscriminator):
def __init__(self, input_nc, ndf=64, n_layers=3, norm_layer=nn.BatchNorm2d, multidilation_kwargs={}):
super().__init__()
self.n_layers = n_layers
kw = 4
padw = int(np.ceil((kw - 1.0) / 2))
sequence = [[nn.Conv2d(input_nc, ndf, kernel_size=kw, stride=2, padding=padw), nn.LeakyReLU(0.2, True)]]
nf = ndf
for n in range(1, n_layers):
nf_prev = nf
nf = min(nf * 2, 512)
cur_model = []
cur_model += [
MultidilatedConv(nf_prev, nf, kernel_size=kw, stride=2, padding=[2, 3], **multidilation_kwargs),
norm_layer(nf),
nn.LeakyReLU(0.2, True),
]
sequence.append(cur_model)
nf_prev = nf
nf = min(nf * 2, 512)
cur_model = []
cur_model += [
nn.Conv2d(nf_prev, nf, kernel_size=kw, stride=1, padding=padw),
norm_layer(nf),
nn.LeakyReLU(0.2, True),
]
sequence.append(cur_model)
sequence += [[nn.Conv2d(nf, 1, kernel_size=kw, stride=1, padding=padw)]]
for n in range(len(sequence)):
setattr(self, "model" + str(n), nn.Sequential(*sequence[n]))
def get_all_activations(self, x):
res = [x]
for n in range(self.n_layers + 2):
model = getattr(self, "model" + str(n))
res.append(model(res[-1]))
return res[1:]
def forward(self, x):
act = self.get_all_activations(x)
return act[-1], act[:-1]
class NLayerDiscriminatorAsGen(NLayerDiscriminator):
def forward(self, x):
return super().forward(x)[0]
@@ -1,49 +0,0 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from kornia.geometry.transform import rotate
class LearnableSpatialTransformWrapper(nn.Module):
def __init__(self, impl, pad_coef=0.5, angle_init_range=80, train_angle=True):
super().__init__()
self.impl = impl
self.angle = torch.rand(1) * angle_init_range
if train_angle:
self.angle = nn.Parameter(self.angle, requires_grad=True)
self.pad_coef = pad_coef
def forward(self, x):
if torch.is_tensor(x):
return self.inverse_transform(self.impl(self.transform(x)), x)
elif isinstance(x, tuple):
x_trans = tuple(self.transform(elem) for elem in x)
y_trans = self.impl(x_trans)
return tuple(self.inverse_transform(elem, orig_x) for elem, orig_x in zip(y_trans, x))
else:
raise ValueError(f"Unexpected input type {type(x)}")
def transform(self, x):
height, width = x.shape[2:]
pad_h, pad_w = int(height * self.pad_coef), int(width * self.pad_coef)
x_padded = F.pad(x, [pad_w, pad_w, pad_h, pad_h], mode="reflect")
x_padded_rotated = rotate(x_padded, angle=self.angle.to(x_padded))
return x_padded_rotated
def inverse_transform(self, y_padded_rotated, orig_x):
height, width = orig_x.shape[2:]
pad_h, pad_w = int(height * self.pad_coef), int(width * self.pad_coef)
y_padded = rotate(y_padded_rotated, angle=-self.angle.to(y_padded_rotated))
y_height, y_width = y_padded.shape[2:]
y = y_padded[:, :, pad_h : y_height - pad_h, pad_w : y_width - pad_w]
return y
if __name__ == "__main__":
layer = LearnableSpatialTransformWrapper(nn.Identity())
x = torch.arange(2 * 3 * 15 * 15).view(2, 3, 15, 15).float()
y = layer(x)
assert x.shape == y.shape
assert torch.allclose(x[:, :, 1:, 1:][:, :, :-1, :-1], y[:, :, 1:, 1:][:, :, :-1, :-1])
print("all ok")
@@ -1,20 +0,0 @@
import torch.nn as nn
class SELayer(nn.Module):
def __init__(self, channel, reduction=16):
super(SELayer, self).__init__()
self.avg_pool = nn.AdaptiveAvgPool2d(1)
self.fc = nn.Sequential(
nn.Linear(channel, channel // reduction, bias=False),
nn.ReLU(inplace=True),
nn.Linear(channel // reduction, channel, bias=False),
nn.Sigmoid(),
)
def forward(self, x):
b, c, _, _ = x.size()
y = self.avg_pool(x).view(b, c)
y = self.fc(y).view(b, c, 1, 1)
res = x * y.expand_as(x)
return res
@@ -1,31 +0,0 @@
import logging
import torch
from .default import DefaultInpaintingTrainingModule
def get_training_model_class(kind):
if kind == "default":
return DefaultInpaintingTrainingModule
raise ValueError(f"Unknown trainer module {kind}")
def make_training_model(config):
kind = config.training_model.kind
kwargs = dict(config.training_model)
kwargs.pop("kind")
kwargs["use_ddp"] = config.trainer.kwargs.get("accelerator", None) == "ddp"
logging.info(f"Make training model {kind}")
cls = get_training_model_class(kind)
return cls(config, **kwargs)
def load_checkpoint(train_config, path, map_location="cuda", strict=True):
model: torch.nn.Module = make_training_model(train_config)
state = torch.load(path, map_location=map_location)
model.load_state_dict(state["state_dict"], strict=strict)
model.on_load_checkpoint(state)
return model
@@ -1,316 +0,0 @@
import copy
import logging
from typing import Dict, Tuple
import pandas as pd
import pytorch_lightning as ptl
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DistributedSampler
# from saicinpainting.evaluation import make_evaluator
# from saicinpainting.training.data.datasets import make_default_train_dataloader, make_default_val_dataloader
# from saicinpainting.training.losses.adversarial import make_discrim_loss
# from saicinpainting.training.losses.perceptual import PerceptualLoss, ResNetPL
from ..modules import make_generator # , make_discriminator
# from saicinpainting.training.visualizers import make_visualizer
from ...utils import add_prefix_to_keys, average_dicts, set_requires_grad, flatten_dict, get_has_ddp_rank
LOGGER = logging.getLogger(__name__)
def make_optimizer(parameters, kind="adamw", **kwargs):
if kind == "adam":
optimizer_class = torch.optim.Adam
elif kind == "adamw":
optimizer_class = torch.optim.AdamW
else:
raise ValueError(f"Unknown optimizer kind {kind}")
return optimizer_class(parameters, **kwargs)
def update_running_average(result: nn.Module, new_iterate_model: nn.Module, decay=0.999):
with torch.no_grad():
res_params = dict(result.named_parameters())
new_params = dict(new_iterate_model.named_parameters())
for k in res_params.keys():
res_params[k].data.mul_(decay).add_(new_params[k].data, alpha=1 - decay)
def make_multiscale_noise(base_tensor, scales=6, scale_mode="bilinear"):
batch_size, _, height, width = base_tensor.shape
cur_height, cur_width = height, width
result = []
align_corners = False if scale_mode in ("bilinear", "bicubic") else None
for _ in range(scales):
cur_sample = torch.randn(batch_size, 1, cur_height, cur_width, device=base_tensor.device)
cur_sample_scaled = F.interpolate(
cur_sample, size=(height, width), mode=scale_mode, align_corners=align_corners
)
result.append(cur_sample_scaled)
cur_height //= 2
cur_width //= 2
return torch.cat(result, dim=1)
class BaseInpaintingTrainingModule(ptl.LightningModule):
def __init__(
self,
config,
use_ddp,
*args,
predict_only=False,
visualize_each_iters=100,
average_generator=False,
generator_avg_beta=0.999,
average_generator_start_step=30000,
average_generator_period=10,
store_discr_outputs_for_vis=False,
**kwargs,
):
super().__init__(*args, **kwargs)
LOGGER.info("BaseInpaintingTrainingModule init called")
self.config = config
self.generator = make_generator(config, **self.config.generator)
self.use_ddp = use_ddp
# if not get_has_ddp_rank():
# LOGGER.info(f"Generator\n{self.generator}")
# if not predict_only:
# self.save_hyperparameters(self.config)
# self.discriminator = make_discriminator(**self.config.discriminator)
# self.adversarial_loss = make_discrim_loss(**self.config.losses.adversarial)
# self.visualizer = make_visualizer(**self.config.visualizer)
# self.val_evaluator = make_evaluator(**self.config.evaluator)
# self.test_evaluator = make_evaluator(**self.config.evaluator)
# if not get_has_ddp_rank():
# LOGGER.info(f"Discriminator\n{self.discriminator}")
# extra_val = self.config.data.get("extra_val", ())
# if extra_val:
# self.extra_val_titles = list(extra_val)
# self.extra_evaluators = nn.ModuleDict({k: make_evaluator(**self.config.evaluator) for k in extra_val})
# else:
# self.extra_evaluators = {}
# self.average_generator = average_generator
# self.generator_avg_beta = generator_avg_beta
# self.average_generator_start_step = average_generator_start_step
# self.average_generator_period = average_generator_period
# self.generator_average = None
# self.last_generator_averaging_step = -1
# self.store_discr_outputs_for_vis = store_discr_outputs_for_vis
# if self.config.losses.get("l1", {"weight_known": 0})["weight_known"] > 0:
# self.loss_l1 = nn.L1Loss(reduction="none")
# if self.config.losses.get("mse", {"weight": 0})["weight"] > 0:
# self.loss_mse = nn.MSELoss(reduction="none")
# if self.config.losses.perceptual.weight > 0:
# self.loss_pl = PerceptualLoss()
# if self.config.losses.get("resnet_pl", {"weight": 0})["weight"] > 0:
# self.loss_resnet_pl = ResNetPL(**self.config.losses.resnet_pl)
# else:
# self.loss_resnet_pl = None
self.visualize_each_iters = visualize_each_iters
LOGGER.info("BaseInpaintingTrainingModule init done")
def configure_optimizers(self):
discriminator_params = list(self.discriminator.parameters())
return [
dict(optimizer=make_optimizer(self.generator.parameters(), **self.config.optimizers.generator)),
dict(optimizer=make_optimizer(discriminator_params, **self.config.optimizers.discriminator)),
]
def train_dataloader(self):
kwargs = dict(self.config.data.train)
if self.use_ddp:
kwargs["ddp_kwargs"] = dict(
num_replicas=self.trainer.num_nodes * self.trainer.num_processes,
rank=self.trainer.global_rank,
shuffle=True,
)
dataloader = make_default_train_dataloader(**self.config.data.train)
return dataloader
def val_dataloader(self):
res = [make_default_val_dataloader(**self.config.data.val)]
if self.config.data.visual_test is not None:
res = res + [make_default_val_dataloader(**self.config.data.visual_test)]
else:
res = res + res
extra_val = self.config.data.get("extra_val", ())
if extra_val:
res += [make_default_val_dataloader(**extra_val[k]) for k in self.extra_val_titles]
return res
def training_step(self, batch, batch_idx, optimizer_idx=None):
self._is_training_step = True
return self._do_step(batch, batch_idx, mode="train", optimizer_idx=optimizer_idx)
def validation_step(self, batch, batch_idx, dataloader_idx):
extra_val_key = None
if dataloader_idx == 0:
mode = "val"
elif dataloader_idx == 1:
mode = "test"
else:
mode = "extra_val"
extra_val_key = self.extra_val_titles[dataloader_idx - 2]
self._is_training_step = False
return self._do_step(batch, batch_idx, mode=mode, extra_val_key=extra_val_key)
def training_step_end(self, batch_parts_outputs):
if (
self.training
and self.average_generator
and self.global_step >= self.average_generator_start_step
and self.global_step >= self.last_generator_averaging_step + self.average_generator_period
):
if self.generator_average is None:
self.generator_average = copy.deepcopy(self.generator)
else:
update_running_average(self.generator_average, self.generator, decay=self.generator_avg_beta)
self.last_generator_averaging_step = self.global_step
full_loss = (
batch_parts_outputs["loss"].mean()
if torch.is_tensor(batch_parts_outputs["loss"]) # loss is not tensor when no discriminator used
else torch.tensor(batch_parts_outputs["loss"]).float().requires_grad_(True)
)
log_info = {k: v.mean() for k, v in batch_parts_outputs["log_info"].items()}
self.log_dict(log_info, on_step=True, on_epoch=False)
return full_loss
def validation_epoch_end(self, outputs):
outputs = [step_out for out_group in outputs for step_out in out_group]
averaged_logs = average_dicts(step_out["log_info"] for step_out in outputs)
self.log_dict({k: v.mean() for k, v in averaged_logs.items()})
pd.set_option("display.max_columns", 500)
pd.set_option("display.width", 1000)
# standard validation
val_evaluator_states = [s["val_evaluator_state"] for s in outputs if "val_evaluator_state" in s]
val_evaluator_res = self.val_evaluator.evaluation_end(states=val_evaluator_states)
val_evaluator_res_df = pd.DataFrame(val_evaluator_res).stack(1).unstack(0)
val_evaluator_res_df.dropna(axis=1, how="all", inplace=True)
LOGGER.info(
f"Validation metrics after epoch #{self.current_epoch}, "
f"total {self.global_step} iterations:\n{val_evaluator_res_df}"
)
for k, v in flatten_dict(val_evaluator_res).items():
self.log(f"val_{k}", v)
# standard visual test
test_evaluator_states = [s["test_evaluator_state"] for s in outputs if "test_evaluator_state" in s]
test_evaluator_res = self.test_evaluator.evaluation_end(states=test_evaluator_states)
test_evaluator_res_df = pd.DataFrame(test_evaluator_res).stack(1).unstack(0)
test_evaluator_res_df.dropna(axis=1, how="all", inplace=True)
LOGGER.info(
f"Test metrics after epoch #{self.current_epoch}, "
f"total {self.global_step} iterations:\n{test_evaluator_res_df}"
)
for k, v in flatten_dict(test_evaluator_res).items():
self.log(f"test_{k}", v)
# extra validations
if self.extra_evaluators:
for cur_eval_title, cur_evaluator in self.extra_evaluators.items():
cur_state_key = f"extra_val_{cur_eval_title}_evaluator_state"
cur_states = [s[cur_state_key] for s in outputs if cur_state_key in s]
cur_evaluator_res = cur_evaluator.evaluation_end(states=cur_states)
cur_evaluator_res_df = pd.DataFrame(cur_evaluator_res).stack(1).unstack(0)
cur_evaluator_res_df.dropna(axis=1, how="all", inplace=True)
LOGGER.info(
f"Extra val {cur_eval_title} metrics after epoch #{self.current_epoch}, "
f"total {self.global_step} iterations:\n{cur_evaluator_res_df}"
)
for k, v in flatten_dict(cur_evaluator_res).items():
self.log(f"extra_val_{cur_eval_title}_{k}", v)
def _do_step(self, batch, batch_idx, mode="train", optimizer_idx=None, extra_val_key=None):
if optimizer_idx == 0: # step for generator
set_requires_grad(self.generator, True)
set_requires_grad(self.discriminator, False)
elif optimizer_idx == 1: # step for discriminator
set_requires_grad(self.generator, False)
set_requires_grad(self.discriminator, True)
batch = self(batch)
total_loss = 0
metrics = {}
if optimizer_idx is None or optimizer_idx == 0: # step for generator
total_loss, metrics = self.generator_loss(batch)
elif optimizer_idx is None or optimizer_idx == 1: # step for discriminator
if self.config.losses.adversarial.weight > 0:
total_loss, metrics = self.discriminator_loss(batch)
if self.get_ddp_rank() in (None, 0) and (batch_idx % self.visualize_each_iters == 0 or mode == "test"):
if self.config.losses.adversarial.weight > 0:
if self.store_discr_outputs_for_vis:
with torch.no_grad():
self.store_discr_outputs(batch)
vis_suffix = f"_{mode}"
if mode == "extra_val":
vis_suffix += f"_{extra_val_key}"
self.visualizer(self.current_epoch, batch_idx, batch, suffix=vis_suffix)
metrics_prefix = f"{mode}_"
if mode == "extra_val":
metrics_prefix += f"{extra_val_key}_"
result = dict(loss=total_loss, log_info=add_prefix_to_keys(metrics, metrics_prefix))
if mode == "val":
result["val_evaluator_state"] = self.val_evaluator.process_batch(batch)
elif mode == "test":
result["test_evaluator_state"] = self.test_evaluator.process_batch(batch)
elif mode == "extra_val":
result[f"extra_val_{extra_val_key}_evaluator_state"] = self.extra_evaluators[extra_val_key].process_batch(
batch
)
return result
def get_current_generator(self, no_average=False):
if not no_average and not self.training and self.average_generator and self.generator_average is not None:
return self.generator_average
return self.generator
def forward(self, batch: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]:
"""Pass data through generator and obtain at leas 'predicted_image' and 'inpainted' keys"""
raise NotImplementedError()
def generator_loss(self, batch) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
raise NotImplementedError()
def discriminator_loss(self, batch) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
raise NotImplementedError()
def store_discr_outputs(self, batch):
out_size = batch["image"].shape[2:]
discr_real_out, _ = self.discriminator(batch["image"])
discr_fake_out, _ = self.discriminator(batch["predicted_image"])
batch["discr_output_real"] = F.interpolate(discr_real_out, size=out_size, mode="nearest")
batch["discr_output_fake"] = F.interpolate(discr_fake_out, size=out_size, mode="nearest")
batch["discr_output_diff"] = batch["discr_output_real"] - batch["discr_output_fake"]
def get_ddp_rank(self):
return self.trainer.global_rank if (self.trainer.num_nodes * self.trainer.num_processes) > 1 else None
@@ -1,230 +0,0 @@
import logging
import random
import torch
import torch.nn.functional as F
from omegaconf import OmegaConf
from ..losses.distance_weighting import make_mask_distance_weighter
from ..losses.feature_matching import feature_matching_loss, masked_l1_loss
from ..modules.fake_fakes import FakeFakesGenerator
from .base import BaseInpaintingTrainingModule, make_multiscale_noise
from ...utils import add_prefix_to_keys, get_ramp
LOGGER = logging.getLogger(__name__)
def ceil_modulo(x, mod):
if x % mod == 0:
return x
return (x // mod + 1) * mod
def make_constant_area_crop_params(img_height, img_width, min_size=128, max_size=512, area=256 * 256, round_to_mod=16):
min_size = min(img_height, img_width, min_size)
max_size = min(img_height, img_width, max_size)
if random.random() < 0.5:
out_height = min(max_size, ceil_modulo(random.randint(min_size, max_size), round_to_mod))
out_width = min(max_size, ceil_modulo(area // out_height, round_to_mod))
else:
out_width = min(max_size, ceil_modulo(random.randint(min_size, max_size), round_to_mod))
out_height = min(max_size, ceil_modulo(area // out_width, round_to_mod))
start_y = random.randint(0, img_height - out_height)
start_x = random.randint(0, img_width - out_width)
return (start_y, start_x, out_height, out_width)
def make_constant_area_crop_batch(batch, **kwargs):
crop_y, crop_x, crop_height, crop_width = make_constant_area_crop_params(
img_height=batch["image"].shape[2], img_width=batch["image"].shape[3], **kwargs
)
batch["image"] = batch["image"][:, :, crop_y : crop_y + crop_height, crop_x : crop_x + crop_width]
batch["mask"] = batch["mask"][:, :, crop_y : crop_y + crop_height, crop_x : crop_x + crop_width]
return batch
class DefaultInpaintingTrainingModule(BaseInpaintingTrainingModule):
def __init__(
self,
*args,
concat_mask=True,
rescale_scheduler_kwargs=None,
image_to_discriminator="predicted_image",
add_noise_kwargs=None,
noise_fill_hole=False,
const_area_crop_kwargs=None,
distance_weighter_kwargs=None,
distance_weighted_mask_for_discr=False,
fake_fakes_proba=0,
fake_fakes_generator_kwargs=None,
**kwargs
):
super().__init__(*args, **kwargs)
self.concat_mask = concat_mask
self.rescale_size_getter = (
get_ramp(**rescale_scheduler_kwargs) if rescale_scheduler_kwargs is not None else None
)
self.image_to_discriminator = image_to_discriminator
self.add_noise_kwargs = add_noise_kwargs
self.noise_fill_hole = noise_fill_hole
self.const_area_crop_kwargs = const_area_crop_kwargs
self.refine_mask_for_losses = (
make_mask_distance_weighter(**distance_weighter_kwargs) if distance_weighter_kwargs is not None else None
)
self.distance_weighted_mask_for_discr = distance_weighted_mask_for_discr
self.fake_fakes_proba = fake_fakes_proba
if self.fake_fakes_proba > 1e-3:
self.fake_fakes_gen = FakeFakesGenerator(**(fake_fakes_generator_kwargs or {}))
def forward(self, batch):
if self.training and self.rescale_size_getter is not None:
cur_size = self.rescale_size_getter(self.global_step)
batch["image"] = F.interpolate(batch["image"], size=cur_size, mode="bilinear", align_corners=False)
batch["mask"] = F.interpolate(batch["mask"], size=cur_size, mode="nearest")
if self.training and self.const_area_crop_kwargs is not None:
batch = make_constant_area_crop_batch(batch, **self.const_area_crop_kwargs)
img = batch["image"]
mask = batch["mask"]
masked_img = img * (1 - mask)
if self.add_noise_kwargs is not None:
noise = make_multiscale_noise(masked_img, **self.add_noise_kwargs)
if self.noise_fill_hole:
masked_img = masked_img + mask * noise[:, : masked_img.shape[1]]
masked_img = torch.cat([masked_img, noise], dim=1)
if self.concat_mask:
masked_img = torch.cat([masked_img, mask], dim=1)
batch["predicted_image"] = self.generator(masked_img)
batch["inpainted"] = mask * batch["predicted_image"] + (1 - mask) * batch["image"]
if self.fake_fakes_proba > 1e-3:
if self.training and torch.rand(1).item() < self.fake_fakes_proba:
batch["fake_fakes"], batch["fake_fakes_masks"] = self.fake_fakes_gen(img, mask)
batch["use_fake_fakes"] = True
else:
batch["fake_fakes"] = torch.zeros_like(img)
batch["fake_fakes_masks"] = torch.zeros_like(mask)
batch["use_fake_fakes"] = False
batch["mask_for_losses"] = (
self.refine_mask_for_losses(img, batch["predicted_image"], mask)
if self.refine_mask_for_losses is not None and self.training
else mask
)
return batch
def generator_loss(self, batch):
img = batch["image"]
predicted_img = batch[self.image_to_discriminator]
original_mask = batch["mask"]
supervised_mask = batch["mask_for_losses"]
# L1
l1_value = masked_l1_loss(
predicted_img,
img,
supervised_mask,
self.config.losses.l1.weight_known,
self.config.losses.l1.weight_missing,
)
total_loss = l1_value
metrics = dict(gen_l1=l1_value)
# vgg-based perceptual loss
if self.config.losses.perceptual.weight > 0:
pl_value = (
self.loss_pl(predicted_img, img, mask=supervised_mask).sum() * self.config.losses.perceptual.weight
)
total_loss = total_loss + pl_value
metrics["gen_pl"] = pl_value
# discriminator
# adversarial_loss calls backward by itself
mask_for_discr = supervised_mask if self.distance_weighted_mask_for_discr else original_mask
self.adversarial_loss.pre_generator_step(
real_batch=img, fake_batch=predicted_img, generator=self.generator, discriminator=self.discriminator
)
discr_real_pred, discr_real_features = self.discriminator(img)
discr_fake_pred, discr_fake_features = self.discriminator(predicted_img)
adv_gen_loss, adv_metrics = self.adversarial_loss.generator_loss(
real_batch=img,
fake_batch=predicted_img,
discr_real_pred=discr_real_pred,
discr_fake_pred=discr_fake_pred,
mask=mask_for_discr,
)
total_loss = total_loss + adv_gen_loss
metrics["gen_adv"] = adv_gen_loss
metrics.update(add_prefix_to_keys(adv_metrics, "adv_"))
# feature matching
if self.config.losses.feature_matching.weight > 0:
need_mask_in_fm = OmegaConf.to_container(self.config.losses.feature_matching).get("pass_mask", False)
mask_for_fm = supervised_mask if need_mask_in_fm else None
fm_value = (
feature_matching_loss(discr_fake_features, discr_real_features, mask=mask_for_fm)
* self.config.losses.feature_matching.weight
)
total_loss = total_loss + fm_value
metrics["gen_fm"] = fm_value
if self.loss_resnet_pl is not None:
resnet_pl_value = self.loss_resnet_pl(predicted_img, img)
total_loss = total_loss + resnet_pl_value
metrics["gen_resnet_pl"] = resnet_pl_value
return total_loss, metrics
def discriminator_loss(self, batch):
total_loss = 0
metrics = {}
predicted_img = batch[self.image_to_discriminator].detach()
self.adversarial_loss.pre_discriminator_step(
real_batch=batch["image"],
fake_batch=predicted_img,
generator=self.generator,
discriminator=self.discriminator,
)
discr_real_pred, discr_real_features = self.discriminator(batch["image"])
discr_fake_pred, discr_fake_features = self.discriminator(predicted_img)
adv_discr_loss, adv_metrics = self.adversarial_loss.discriminator_loss(
real_batch=batch["image"],
fake_batch=predicted_img,
discr_real_pred=discr_real_pred,
discr_fake_pred=discr_fake_pred,
mask=batch["mask"],
)
total_loss = total_loss + adv_discr_loss
metrics["discr_adv"] = adv_discr_loss
metrics.update(add_prefix_to_keys(adv_metrics, "adv_"))
if batch.get("use_fake_fakes", False):
fake_fakes = batch["fake_fakes"]
self.adversarial_loss.pre_discriminator_step(
real_batch=batch["image"],
fake_batch=fake_fakes,
generator=self.generator,
discriminator=self.discriminator,
)
discr_fake_fakes_pred, _ = self.discriminator(fake_fakes)
fake_fakes_adv_discr_loss, fake_fakes_adv_metrics = self.adversarial_loss.discriminator_loss(
real_batch=batch["image"],
fake_batch=fake_fakes,
discr_real_pred=discr_real_pred,
discr_fake_pred=discr_fake_fakes_pred,
mask=batch["mask"],
)
total_loss = total_loss + fake_fakes_adv_discr_loss
metrics["discr_adv_fake_fakes"] = fake_fakes_adv_discr_loss
metrics.update(add_prefix_to_keys(fake_fakes_adv_metrics, "adv_"))
return total_loss, metrics
@@ -1,15 +0,0 @@
import logging
from saicinpainting.training.visualizers.directory import DirectoryVisualizer
from saicinpainting.training.visualizers.noop import NoopVisualizer
def make_visualizer(kind, **kwargs):
logging.info(f'Make visualizer {kind}')
if kind == 'directory':
return DirectoryVisualizer(**kwargs)
if kind == 'noop':
return NoopVisualizer()
raise ValueError(f'Unknown visualizer kind {kind}')
@@ -1,75 +0,0 @@
import abc
from typing import Dict, List
import numpy as np
import torch
import skimage.color as color
from skimage.segmentation import mark_boundaries
from . import colors
COLORS, _ = colors.generate_colors(151) # 151 - max classes for semantic segmentation
class BaseVisualizer:
@abc.abstractmethod
def __call__(self, epoch_i, batch_i, batch, suffix="", rank=None):
"""
Take a batch, make an image from it and visualize
"""
raise NotImplementedError()
def visualize_mask_and_images(
images_dict: Dict[str, np.ndarray],
keys: List[str],
last_without_mask=True,
rescale_keys=None,
mask_only_first=None,
black_mask=False,
) -> np.ndarray:
mask = images_dict["mask"] > 0.5
result = []
for i, k in enumerate(keys):
img = images_dict[k]
img = np.transpose(img, (1, 2, 0))
if rescale_keys is not None and k in rescale_keys:
img = img - img.min()
img /= img.max() + 1e-5
if len(img.shape) == 2:
img = np.expand_dims(img, 2)
if img.shape[2] == 1:
img = np.repeat(img, 3, axis=2)
elif img.shape[2] > 3:
img_classes = img.argmax(2)
img = color.label2rgb(img_classes, colors=COLORS)
if mask_only_first:
need_mark_boundaries = i == 0
else:
need_mark_boundaries = i < len(keys) - 1 or not last_without_mask
if need_mark_boundaries:
if black_mask:
img = img * (1 - mask[0][..., None])
img = mark_boundaries(img, mask[0], color=(1.0, 0.0, 0.0), outline_color=(1.0, 1.0, 1.0), mode="thick")
result.append(img)
return np.concatenate(result, axis=1)
def visualize_mask_and_images_batch(
batch: Dict[str, torch.Tensor], keys: List[str], max_items=10, last_without_mask=True, rescale_keys=None
) -> np.ndarray:
batch = {k: tens.detach().cpu().numpy() for k, tens in batch.items() if k in keys or k == "mask"}
batch_size = next(iter(batch.values())).shape[0]
items_to_vis = min(batch_size, max_items)
result = []
for i in range(items_to_vis):
cur_dct = {k: tens[i] for k, tens in batch.items()}
result.append(
visualize_mask_and_images(cur_dct, keys, last_without_mask=last_without_mask, rescale_keys=rescale_keys)
)
return np.concatenate(result, axis=0)
@@ -1,95 +0,0 @@
import random
import colorsys
import numpy as np
import matplotlib
matplotlib.use("agg")
import matplotlib.pyplot as plt
from matplotlib.colors import LinearSegmentedColormap
def generate_colors(nlabels, type="bright", first_color_black=False, last_color_black=True, verbose=False):
# https://stackoverflow.com/questions/14720331/how-to-generate-random-colors-in-matplotlib
"""
Creates a random colormap to be used together with matplotlib. Useful for segmentation tasks
:param nlabels: Number of labels (size of colormap)
:param type: 'bright' for strong colors, 'soft' for pastel colors
:param first_color_black: Option to use first color as black, True or False
:param last_color_black: Option to use last color as black, True or False
:param verbose: Prints the number of labels and shows the colormap. True or False
:return: colormap for matplotlib
"""
if type not in ("bright", "soft"):
print('Please choose "bright" or "soft" for type')
return
if verbose:
print("Number of labels: " + str(nlabels))
# Generate color map for bright colors, based on hsv
if type == "bright":
randHSVcolors = [
(
np.random.uniform(low=0.0, high=1),
np.random.uniform(low=0.2, high=1),
np.random.uniform(low=0.9, high=1),
)
for i in range(nlabels)
]
# Convert HSV list to RGB
randRGBcolors = []
for HSVcolor in randHSVcolors:
randRGBcolors.append(colorsys.hsv_to_rgb(HSVcolor[0], HSVcolor[1], HSVcolor[2]))
if first_color_black:
randRGBcolors[0] = [0, 0, 0]
if last_color_black:
randRGBcolors[-1] = [0, 0, 0]
random_colormap = LinearSegmentedColormap.from_list("new_map", randRGBcolors, N=nlabels)
# Generate soft pastel colors, by limiting the RGB spectrum
if type == "soft":
low = 0.6
high = 0.95
randRGBcolors = [
(
np.random.uniform(low=low, high=high),
np.random.uniform(low=low, high=high),
np.random.uniform(low=low, high=high),
)
for i in range(nlabels)
]
if first_color_black:
randRGBcolors[0] = [0, 0, 0]
if last_color_black:
randRGBcolors[-1] = [0, 0, 0]
random_colormap = LinearSegmentedColormap.from_list("new_map", randRGBcolors, N=nlabels)
# Display colorbar
if verbose:
from matplotlib import colors, colorbar
from matplotlib import pyplot as plt
fig, ax = plt.subplots(1, 1, figsize=(15, 0.5))
bounds = np.linspace(0, nlabels, nlabels + 1)
norm = colors.BoundaryNorm(bounds, nlabels)
cb = colorbar.ColorbarBase(
ax,
cmap=random_colormap,
norm=norm,
spacing="proportional",
ticks=None,
boundaries=bounds,
format="%1i",
orientation="horizontal",
)
return randRGBcolors, random_colormap
@@ -1,41 +0,0 @@
import os
import cv2
import numpy as np
from .base import BaseVisualizer, visualize_mask_and_images_batch
from ...utils import check_and_warn_input_range
class DirectoryVisualizer(BaseVisualizer):
DEFAULT_KEY_ORDER = "image predicted_image inpainted".split(" ")
def __init__(
self, outdir, key_order=DEFAULT_KEY_ORDER, max_items_in_batch=10, last_without_mask=True, rescale_keys=None
):
self.outdir = outdir
os.makedirs(self.outdir, exist_ok=True)
self.key_order = key_order
self.max_items_in_batch = max_items_in_batch
self.last_without_mask = last_without_mask
self.rescale_keys = rescale_keys
def __call__(self, epoch_i, batch_i, batch, suffix="", rank=None):
check_and_warn_input_range(batch["image"], 0, 1, "DirectoryVisualizer target image")
vis_img = visualize_mask_and_images_batch(
batch,
self.key_order,
max_items=self.max_items_in_batch,
last_without_mask=self.last_without_mask,
rescale_keys=self.rescale_keys,
)
vis_img = np.clip(vis_img * 255, 0, 255).astype("uint8")
curoutdir = os.path.join(self.outdir, f"epoch{epoch_i:04d}{suffix}")
os.makedirs(curoutdir, exist_ok=True)
rank_suffix = f"_r{rank}" if rank is not None else ""
out_fname = os.path.join(curoutdir, f"batch{batch_i:07d}{rank_suffix}.jpg")
vis_img = cv2.cvtColor(vis_img, cv2.COLOR_RGB2BGR)
cv2.imwrite(out_fname, vis_img)
@@ -1,9 +0,0 @@
from .base import BaseVisualizer
class NoopVisualizer(BaseVisualizer):
def __init__(self, *args, **kwargs):
pass
def __call__(self, epoch_i, batch_i, batch, suffix="", rank=None):
pass
@@ -1,172 +0,0 @@
import bisect
import functools
import logging
import numbers
import os
import sys
import traceback
import warnings
import torch
from pytorch_lightning import seed_everything
LOGGER = logging.getLogger(__name__)
def check_and_warn_input_range(tensor, min_value, max_value, name):
actual_min = tensor.min()
actual_max = tensor.max()
if actual_min < min_value or actual_max > max_value:
warnings.warn(f"{name} must be in {min_value}..{max_value} range, but it ranges {actual_min}..{actual_max}")
def sum_dict_with_prefix(target, cur_dict, prefix, default=0):
for k, v in cur_dict.items():
target_key = prefix + k
target[target_key] = target.get(target_key, default) + v
def average_dicts(dict_list):
result = {}
norm = 1e-3
for dct in dict_list:
sum_dict_with_prefix(result, dct, "")
norm += 1
for k in list(result):
result[k] /= norm
return result
def add_prefix_to_keys(dct, prefix):
return {prefix + k: v for k, v in dct.items()}
def set_requires_grad(module, value):
for param in module.parameters():
param.requires_grad = value
def flatten_dict(dct):
result = {}
for k, v in dct.items():
if isinstance(k, tuple):
k = "_".join(k)
if isinstance(v, dict):
for sub_k, sub_v in flatten_dict(v).items():
result[f"{k}_{sub_k}"] = sub_v
else:
result[k] = v
return result
class LinearRamp:
def __init__(self, start_value=0, end_value=1, start_iter=-1, end_iter=0):
self.start_value = start_value
self.end_value = end_value
self.start_iter = start_iter
self.end_iter = end_iter
def __call__(self, i):
if i < self.start_iter:
return self.start_value
if i >= self.end_iter:
return self.end_value
part = (i - self.start_iter) / (self.end_iter - self.start_iter)
return self.start_value * (1 - part) + self.end_value * part
class LadderRamp:
def __init__(self, start_iters, values):
self.start_iters = start_iters
self.values = values
assert len(values) == len(start_iters) + 1, (len(values), len(start_iters))
def __call__(self, i):
segment_i = bisect.bisect_right(self.start_iters, i)
return self.values[segment_i]
def get_ramp(kind="ladder", **kwargs):
if kind == "linear":
return LinearRamp(**kwargs)
if kind == "ladder":
return LadderRamp(**kwargs)
raise ValueError(f"Unexpected ramp kind: {kind}")
def print_traceback_handler(sig, frame):
LOGGER.warning(f"Received signal {sig}")
bt = "".join(traceback.format_stack())
LOGGER.warning(f"Requested stack trace:\n{bt}")
def handle_deterministic_config(config):
seed = dict(config).get("seed", None)
if seed is None:
return False
seed_everything(seed)
return True
def get_shape(t):
if torch.is_tensor(t):
return tuple(t.shape)
elif isinstance(t, dict):
return {n: get_shape(q) for n, q in t.items()}
elif isinstance(t, (list, tuple)):
return [get_shape(q) for q in t]
elif isinstance(t, numbers.Number):
return type(t)
else:
raise ValueError("unexpected type {}".format(type(t)))
def get_has_ddp_rank():
master_port = os.environ.get("MASTER_PORT", None)
node_rank = os.environ.get("NODE_RANK", None)
local_rank = os.environ.get("LOCAL_RANK", None)
world_size = os.environ.get("WORLD_SIZE", None)
has_rank = master_port is not None or node_rank is not None or local_rank is not None or world_size is not None
return has_rank
def handle_ddp_subprocess():
def main_decorator(main_func):
@functools.wraps(main_func)
def new_main(*args, **kwargs):
# Trainer sets MASTER_PORT, NODE_RANK, LOCAL_RANK, WORLD_SIZE
parent_cwd = os.environ.get("TRAINING_PARENT_WORK_DIR", None)
has_parent = parent_cwd is not None
has_rank = get_has_ddp_rank()
assert has_parent == has_rank, f"Inconsistent state: has_parent={has_parent}, has_rank={has_rank}"
if has_parent:
# we are in the worker
sys.argv.extend(
[
f"hydra.run.dir={parent_cwd}",
# 'hydra/hydra_logging=disabled',
# 'hydra/job_logging=disabled'
]
)
# do nothing if this is a top-level process
# TRAINING_PARENT_WORK_DIR is set in handle_ddp_parent_process after hydra initialization
main_func(*args, **kwargs)
return new_main
return main_decorator
def handle_ddp_parent_process():
parent_cwd = os.environ.get("TRAINING_PARENT_WORK_DIR", None)
has_parent = parent_cwd is not None
has_rank = get_has_ddp_rank()
assert has_parent == has_rank, f"Inconsistent state: has_parent={has_parent}, has_rank={has_rank}"
if parent_cwd is None:
os.environ["TRAINING_PARENT_WORK_DIR"] = os.getcwd()
return has_parent
+82 -37
View File
@@ -5,10 +5,10 @@ from PIL import Image, ImageOps
from typing import Dict
from .sam.nodes import SAMLoader, GetSAMEmbedding, SAMEmbeddingToImage
from .lama import LaMaInpaint
from .lama import LoadLaMaModel, LaMaInpaint
from ..masking import get_crop_region, expand_crop_region
from ..image_utils import ResizeMode, resize_image, flatten_image
from ..image_utils import ResizeMode, resize_image
from ..utils import numpy2pil, tensor2pil, pil2tensor
@@ -21,46 +21,54 @@ class PrepareImageAndMaskForInpaint:
"mask": ("MASK",),
"mask_blur": ("INT", {"default": 4, "min": 0, "max": 64}),
"inpaint_masked": ("BOOLEAN", {"default": False}),
"mask_padding": ("INT", {"default": 32, "min": 0, "max": 256}),
"mask_padding": ("INT", {"default": 32, "min": 0, "max": 1024}),
"width": ("INT", {"default": 0, "min": 0, "max": 2048}),
"height": ("INT", {"default": 0, "min": 0, "max": 2048}),
}
},
"optional": {
"controlnet_image": ("IMAGE",),
},
}
RETURN_TYPES = ("IMAGE", "MASK", "IMAGE", "CROP_REGION")
RETURN_NAMES = ("inpaint_image", "inpaint_mask", "overlay_image", "crop_region")
CATEGORY = "Art Venture/Inpainting"
RETURN_TYPES = ("IMAGE", "MASK", "IMAGE", "CROP_REGION", "IMAGE")
RETURN_NAMES = ("inpaint_image", "inpaint_mask", "overlay_image", "crop_region", "controlnet_image")
CATEGORY = "ArtVenture/Inpainting"
FUNCTION = "prepare"
def prepare(
self,
image: torch.Tensor,
mask: torch.Tensor,
# resize_mode: str,
mask_blur: int,
inpaint_masked: bool,
mask_padding: int,
width: int,
height: int,
controlnet_image: torch.Tensor = None,
):
if image.shape[0] != mask.shape[0]:
raise ValueError("image and mask must have same batch size")
if controlnet_image is not None and image.shape[0] != controlnet_image.shape[0]:
raise ValueError("image and controlnet_image must have same batch size")
if image.shape[1] != mask.shape[1] or image.shape[2] != mask.shape[2]:
raise ValueError("image and mask must have same dimensions")
if width == 0 and height == 0:
height, width = image.shape[1:3]
sourceheight, sourcewidth = image.shape[1:3]
# These are only used if inpaint_masked is True
out_width, out_height = width, height
if inpaint_masked and out_width == 0 and out_height == 0:
out_height, out_width = image.shape[1:3]
source_height, source_width = image.shape[1:3]
masks = []
images = []
overlay_masks = []
masks = []
overlay_images = []
crop_regions = []
processed_controlnet_images = []
for img, msk in zip(image, mask):
for idx, (img, msk) in enumerate(zip(image, mask)):
np_mask: np.ndarray = msk.cpu().numpy()
if mask_blur > 0:
@@ -68,41 +76,76 @@ class PrepareImageAndMaskForInpaint:
np_mask = cv2.GaussianBlur(np_mask, (kernel_size, kernel_size), mask_blur)
pil_mask = numpy2pil(np_mask, "L")
crop_region = None
pil_img = tensor2pil(img)
# --- LOGIC SEPARATION ---
if inpaint_masked:
# --- MODE 1: CROP AND RESIZE ---
crop_region = get_crop_region(np_mask, mask_padding)
crop_region = expand_crop_region(crop_region, width, height, sourcewidth, sourceheight)
# crop mask
overlay_mask = pil_mask
pil_mask = resize_image(pil_mask.crop(crop_region), width, height, ResizeMode.RESIZE_TO_FIT)
pil_mask = pil_mask.convert("L")
crop_region = expand_crop_region(crop_region, out_width, out_height, source_width, source_height)
cropped_img = pil_img.crop(crop_region)
cropped_mask = pil_mask.crop(crop_region)
final_pil_img = resize_image(cropped_img, out_width, out_height, ResizeMode.RESIZE_TO_FIT)
final_pil_mask = resize_image(cropped_mask, out_width, out_height, ResizeMode.RESIZE_TO_FIT).convert(
"L"
)
if controlnet_image is not None:
pil_cimg = tensor2pil(controlnet_image[idx])
cn_source_width, cn_source_height = pil_cimg.size
scale_x = cn_source_width / source_width
scale_y = cn_source_height / source_height
cn_target_width = int(out_width * scale_x)
cn_target_height = int(out_height * scale_y)
x1, y1, x2, y2 = crop_region
cn_crop_region = (int(x1 * scale_x), int(y1 * scale_y), int(x2 * scale_x), int(y2 * scale_y))
cropped_cn_img = pil_cimg.crop(cn_crop_region)
final_cn_img = resize_image(
cropped_cn_img, cn_target_width, cn_target_height, ResizeMode.RESIZE_TO_FIT
)
processed_controlnet_images.append(pil2tensor(final_cn_img))
else:
np_mask = np.clip((np_mask.astype(np.float32)) * 2, 0, 255).astype(np.uint8)
overlay_mask = numpy2pil(np_mask, "L")
# --- MODE 2: PASS-THROUGH (NO RESIZING) ---
final_pil_img = pil_img
final_pil_mask = pil_mask # Already blurred if requested
crop_region = (0, 0, source_width, source_height)
pil_img = tensor2pil(img)
pil_img = flatten_image(pil_img)
if controlnet_image is not None:
# Simply pass the original controlnet image through
final_cn_img = tensor2pil(controlnet_image[idx])
processed_controlnet_images.append(pil2tensor(final_cn_img))
# --- COMMON LOGIC FOR BOTH MODES ---
# The overlay/preview should always be based on the original full-size image
image_masked = Image.new("RGBa", (pil_img.width, pil_img.height))
image_masked.paste(pil_img.convert("RGBA").convert("RGBa"), mask=ImageOps.invert(overlay_mask))
# The mask used here is the potentially blurred one, but before any cropping/resizing
image_masked.paste(pil_img.convert("RGBA").convert("RGBa"), mask=ImageOps.invert(pil_mask))
overlay_images.append(pil2tensor(image_masked.convert("RGBA")))
overlay_masks.append(pil2tensor(overlay_mask))
if crop_region is not None:
pil_img = resize_image(pil_img.crop(crop_region), width, height, ResizeMode.RESIZE_TO_FIT)
else:
crop_region = (0, 0, 0, 0)
images.append(pil2tensor(pil_img))
masks.append(pil2tensor(pil_mask))
images.append(pil2tensor(final_pil_img))
masks.append(pil2tensor(final_pil_mask))
crop_regions.append(torch.tensor(crop_region, dtype=torch.int64))
if processed_controlnet_images:
final_controlnet_tensor = torch.cat(processed_controlnet_images, dim=0)
else:
# If no controlnet image is provided, create a black 64x64 placeholder
batch_size = image.shape[0]
final_controlnet_tensor = torch.zeros((batch_size, 64, 64, 3), dtype=torch.float32, device=image.device)
return (
torch.cat(images, dim=0),
torch.cat(masks, dim=0),
torch.cat(overlay_images, dim=0),
torch.stack(crop_regions),
torch.stack(crop_regions, dim=0),
final_controlnet_tensor,
)
@@ -118,7 +161,7 @@ class OverlayInpaintedLatent:
}
RETURN_TYPES = ("LATENT",)
CATEGORY = "Art Venture/Inpainting"
CATEGORY = "ArtVenture/Inpainting"
FUNCTION = "overlay"
def overlay(self, original: Dict, inpainted: Dict, mask: torch.Tensor):
@@ -162,7 +205,7 @@ class OverlayInpaintedImage:
}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "Art Venture/Inpainting"
CATEGORY = "ArtVenture/Inpainting"
FUNCTION = "overlay"
def overlay(self, inpainted: torch.Tensor, overlay_image: torch.Tensor, crop_region: torch.Tensor):
@@ -198,6 +241,7 @@ NODE_CLASS_MAPPINGS = {
"AV_SAMLoader": SAMLoader,
"GetSAMEmbedding": GetSAMEmbedding,
"SAMEmbeddingToImage": SAMEmbeddingToImage,
"LoadLaMaModel": LoadLaMaModel,
"LaMaInpaint": LaMaInpaint,
"PrepareImageAndMaskForInpaint": PrepareImageAndMaskForInpaint,
"OverlayInpaintedLatent": OverlayInpaintedLatent,
@@ -208,6 +252,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"AV_SAMLoader": "SAM Loader",
"GetSAMEmbedding": "Get SAM Embedding",
"SAMEmbeddingToImage": "SAM Embedding to Image",
"LoadLaMaModel": "LaMa Loader",
"LaMaInpaint": "LaMa Remove Object",
"PrepareImageAndMaskForInpaint": "Prepare Image & Mask for Inpaint",
"OverlayInpaintedLatent": "Overlay Inpainted Latent",
+10 -9
View File
@@ -9,12 +9,13 @@ import comfy.utils
from ...utils import ensure_package, tensor2pil, pil2tensor
folder_paths.folder_names_and_paths["sams"] = (
[
os.path.join(folder_paths.models_dir, "sams"),
],
folder_paths.supported_pt_extensions,
)
if "sams" not in folder_paths.folder_names_and_paths:
folder_paths.folder_names_and_paths["sams"] = (
[
os.path.join(folder_paths.models_dir, "sams"),
],
folder_paths.supported_pt_extensions,
)
gpu = model_management.get_torch_device()
cpu = torch.device("cpu")
@@ -32,7 +33,7 @@ class SAMLoader:
RETURN_TYPES = ("AV_SAM_MODEL",)
RETURN_NAMES = ("sam_model",)
FUNCTION = "load_model"
CATEGORY = "Art Venture/Segmentation"
CATEGORY = "ArtVenture/Segmentation"
def load_model(self, model_name):
modelname = folder_paths.get_full_path("sams", model_name)
@@ -68,7 +69,7 @@ class GetSAMEmbedding:
}
RETURN_TYPES = ("SAM_EMBEDDING",)
CATEGORY = "Art Venture/Segmentation"
CATEGORY = "ArtVenture/Segmentation"
FUNCTION = "get_sam_embedding"
def get_sam_embedding(self, image, sam_model, device_mode="AUTO"):
@@ -102,7 +103,7 @@ class SAMEmbeddingToImage:
}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "Art Venture/Segmentation"
CATEGORY = "ArtVenture/Segmentation"
FUNCTION = "sam_embedding_to_noise_image"
def sam_embedding_to_noise_image(self, embedding: np.ndarray):
+3 -1
View File
@@ -1,14 +1,16 @@
from .blip_node import BlipLoader, BlipCaption
from .blip_node import BlipLoader, BlipCaption, DownloadAndLoadBlip
from .danbooru import DeepDanbooruCaption
NODE_CLASS_MAPPINGS = {
"BLIPLoader": BlipLoader,
"BLIPCaption": BlipCaption,
"DownloadAndLoadBlip": DownloadAndLoadBlip,
"DeepDanbooruCaption": DeepDanbooruCaption,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"BLIPLoader": "BLIP Loader",
"BLIPCaption": "BLIP Caption",
"DownloadAndLoadBlip": "Download and Load BLIP Model",
"DeepDanbooruCaption": "Deep Danbooru Caption",
}
+43 -23
View File
@@ -7,17 +7,25 @@ from torchvision.transforms.functional import InterpolationMode
import folder_paths
from comfy.model_management import text_encoder_device, text_encoder_offload_device, soft_empty_cache
from ..model_utils import download_model
from ..model_utils import download_file
from ..utils import tensor2pil
blips = {}
blip_size = 384
gpu = text_encoder_device()
cpu = text_encoder_offload_device()
model_url = (
"https://storage.googleapis.com/sfr-vision-language-research/BLIP/models/model_base_caption_capfilt_large.pth"
)
model_dir = os.path.join(folder_paths.models_dir, "blip")
models = {
"model_base_caption_capfilt_large.pth": {
"url": "https://storage.googleapis.com/sfr-vision-language-research/BLIP/models/model_base_caption_capfilt_large.pth",
"sha": "96ac8749bd0a568c274ebe302b3a3748ab9be614c737f3d8c529697139174086",
},
"model_base_capfilt_large.pth": {
"url": "https://storage.googleapis.com/sfr-vision-language-research/BLIP/models/model_base_capfilt_large.pth",
"sha": "8f5187458d4d47bb87876faf3038d5947eff17475edf52cf47b62e84da0b235f",
},
}
folder_paths.folder_names_and_paths["blip"] = (
[model_dir],
@@ -90,12 +98,7 @@ def join_caption(caption, prefix, suffix):
def blip_caption(model, image, min_length, max_length):
image = tensor2pil(image)
if "transformers==4.26.1" in packages(True):
print("Using Legacy `transformImaage()`")
tensor = transformImage_legacy(image)
else:
tensor = transformImage(image)
tensor = transformImage(image)
with torch.no_grad():
caption = model.generate(
@@ -119,11 +122,35 @@ class BlipLoader:
RETURN_TYPES = ("BLIP_MODEL",)
FUNCTION = "load_blip"
CATEGORY = "Art Venture/Captioning"
CATEGORY = "ArtVenture/Captioning"
def load_blip(self, model_name):
model = load_blip(model_name)
return (model,)
return (load_blip(model_name),)
class DownloadAndLoadBlip:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model_name": (list(models.keys()),),
},
}
RETURN_TYPES = ("BLIP_MODEL",)
FUNCTION = "download_and_load_blip"
CATEGORY = "ArtVenture/Captioning"
def download_and_load_blip(self, model_name):
if model_name not in folder_paths.get_filename_list("blip"):
model_info = models[model_name]
download_file(
model_info["url"],
os.path.join(model_dir, model_name),
model_info["sha"],
)
return (load_blip(model_name),)
class BlipCaption:
@@ -164,7 +191,7 @@ class BlipCaption:
RETURN_NAMES = ("caption",)
OUTPUT_IS_LIST = (True,)
FUNCTION = "blip_caption"
CATEGORY = "Art Venture/Captioning"
CATEGORY = "ArtVenture/Captioning"
def blip_caption(
self, image, min_length, max_length, device_mode="AUTO", prefix="", suffix="", enabled=True, blip_model=None
@@ -173,15 +200,8 @@ class BlipCaption:
return ([join_caption("", prefix, suffix)],)
if blip_model is None:
ckpts = folder_paths.get_filename_list("blip")
if len(ckpts) == 0:
ckpts = download_model(
model_path=model_dir,
model_url=model_url,
ext_filter=[".pth"],
download_name="model_base_caption_capfilt_large.pth",
)
blip_model = load_blip(ckpts[0])
downloader = DownloadAndLoadBlip()
blip_model = downloader.download_and_load_blip("model_base_caption_capfilt_large.pth")[0]
device = gpu if device_mode != "CPU" else cpu
blip_model = blip_model.to(device)
+9 -12
View File
@@ -7,7 +7,7 @@ import folder_paths
from comfy.model_management import text_encoder_device, text_encoder_offload_device, soft_empty_cache
from ..image_utils import resize_image
from ..model_utils import download_model
from ..model_utils import download_file
from ..utils import is_junction, tensor2pil
from .blip_node import join_caption
@@ -15,28 +15,25 @@ danbooru = None
blip_size = 384
gpu = text_encoder_device()
cpu = text_encoder_offload_device()
model_dir = os.path.join(folder_paths.models_dir, "blip")
model_url = "https://github.com/AUTOMATIC1111/TorchDeepDanbooru/releases/download/v1/model-resnet_custom_v3.pt"
model_sha = "3841542cda4dd037da12a565e854b3347bb2eec8fbcd95ea3941b2c68990a355"
re_special = re.compile(r"([\\()])")
def load_danbooru(device_mode):
global danbooru
if danbooru is None:
blip_dir = os.path.join(folder_paths.models_dir, "blip")
if not os.path.exists(blip_dir) and not is_junction(blip_dir):
os.makedirs(blip_dir, exist_ok=True)
if not os.path.exists(model_dir) and not is_junction(model_dir):
os.makedirs(model_dir, exist_ok=True)
files = download_model(
model_path=blip_dir,
model_url=model_url,
ext_filter=[".pt"],
download_name="model-resnet_custom_v3.pt",
)
model_path = os.path.join(model_dir, "model-resnet_custom_v3.pt")
download_file(model_url, model_path, model_sha)
from .models.deepbooru_model import DeepDanbooruModel
danbooru = DeepDanbooruModel()
danbooru.load_state_dict(torch.load(files[0], map_location="cpu"))
danbooru.load_state_dict(torch.load(model_path, map_location="cpu"))
danbooru.eval()
if device_mode != "CPU":
@@ -82,7 +79,7 @@ class DeepDanbooruCaption:
RETURN_NAMES = ("caption",)
OUTPUT_IS_LIST = (True,)
FUNCTION = "caption"
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
def caption(
self,
+1 -1
View File
@@ -25,10 +25,10 @@ from transformers.modeling_outputs import (
)
from transformers.modeling_utils import (
PreTrainedModel,
apply_chunking_to_forward,
find_pruneable_heads_and_indices,
prune_linear_layer,
)
from transformers.pytorch_utils import apply_chunking_to_forward
from transformers.utils import logging
from transformers.models.bert.configuration_bert import BertConfig
+15 -128
View File
@@ -1,8 +1,5 @@
import os
import json
import torch
from typing import Dict, Tuple, List
from pydantic import BaseModel
from typing import Dict, Tuple
import folder_paths
import comfy.clip_vision
@@ -10,11 +7,10 @@ import comfy.controlnet
import comfy.utils
import comfy.model_management
from .utils import load_module, pil2tensor
from .utility_nodes import load_images_from_url
from .utils import load_module
custom_nodes = folder_paths.get_folder_paths("custom_nodes")
ip_adapter_dir_names = ["IPAdapter", "ComfyUI_IPAdapter_plus"]
ip_adapter_dir_names = ["IPAdapter", "ComfyUI_IPAdapter_plus", "comfyui_ipadapter_plus"]
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
@@ -36,42 +32,11 @@ try:
print("Loaded IPAdapter nodes from", module_path)
nodes: Dict = getattr(module, "NODE_CLASS_MAPPINGS")
IPAdapterUnifiedLoader = nodes.get("IPAdapterUnifiedLoader")
IPAdapterModelLoader = nodes.get("IPAdapterModelLoader")
IPAdapterApply = nodes.get("IPAdapter")
IPAdapterEncoder = nodes.get("IPAdapterEncoder")
IPAdapterEmbeds = nodes.get("IPAdapterEmbeds")
IPAdapterCombineEmbeds = nodes.get("IPAdapterCombineEmbeds")
IPAdapterSimple = nodes.get("IPAdapter")
loader = IPAdapterModelLoader()
unifyLoader = IPAdapterUnifiedLoader()
apply = IPAdapterApply()
encoder = IPAdapterEncoder()
combiner = IPAdapterCombineEmbeds()
embedder = IPAdapterEmbeds()
WEIGHT_TYPES = [
"linear",
"ease in",
"ease out",
"ease in-out",
"reverse in-out",
"weak input",
"weak output",
"weak middle",
"strong middle",
"style transfer (SDXL)",
"composition (SDXL)",
]
PRESETS = [
"LIGHT - SD1.5 only (low strength)",
"STANDARD (medium strength)",
"VIT-G (medium strength)",
"PLUS (high strength)",
"PLUS FACE (portraits)",
"FULL FACE - SD1.5 only (portraits stronger)",
]
apply = IPAdapterSimple()
class AV_IPAdapterPipe:
@classmethod
@@ -85,7 +50,7 @@ try:
RETURN_TYPES = ("IPADAPTER",)
RETURN_NAMES = ("pipeline",)
CATEGORY = "Art Venture/IP Adapter"
CATEGORY = "ArtVenture/IP Adapter"
FUNCTION = "load_ip_adapter"
def load_ip_adapter(self, ip_adapter_name, clip_name):
@@ -97,9 +62,11 @@ try:
pipeline = {"ipadapter": {"model": ip_adapter}, "clipvision": {"model": clip_vision}}
return (pipeline,)
class AV_IPAdapter(IPAdapterModelLoader, IPAdapterApply):
class AV_IPAdapter:
@classmethod
def INPUT_TYPES(cls):
inputs = IPAdapterSimple.INPUT_TYPES()
return {
"required": {
"ip_adapter_name": (["None"] + folder_paths.get_filename_list("ipadapter"),),
@@ -107,7 +74,6 @@ try:
"model": ("MODEL",),
"image": ("IMAGE",),
"weight": ("FLOAT", {"default": 1.0, "min": -1, "max": 3, "step": 0.05}),
"noise": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}),
},
"optional": {
"ip_adapter_opt": ("IPADAPTER",),
@@ -115,17 +81,14 @@ try:
"attn_mask": ("MASK",),
"start_at": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"end_at": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"weight_type": (
["standard", "prompt is more important", "style transfer (SDXL only)"],
{"default": "standard"},
),
"weight_type": inputs["required"]["weight_type"],
"enabled": ("BOOLEAN", {"default": True}),
},
}
RETURN_TYPES = ("MODEL", "IPADAPTER", "CLIP_VISION")
RETURN_NAMES = ("model", "pipeline", "clip_vision")
CATEGORY = "Art Venture/IP Adapter"
CATEGORY = "ArtVenture/IP Adapter"
FUNCTION = "apply_ip_adapter"
def apply_ip_adapter(
@@ -135,7 +98,6 @@ try:
model,
image,
weight,
noise,
ip_adapter_opt=None,
clip_vision_opt=None,
enabled=True,
@@ -169,91 +131,16 @@ try:
return res
class IPAdapterImage(BaseModel):
url: str
weight: float
class IPAdapterData(BaseModel):
images: List[IPAdapterImage]
class AV_StyleApply:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
"preset": (PRESETS,),
"data": (
"STRING",
{
"placeholder": '[{"url": "http://domain/path/image.png", "weight": 1}]',
"multiline": True,
"dynamicPrompts": False,
},
),
"weight": ("FLOAT", {"default": 0.5, "min": -1, "max": 3, "step": 0.05}),
"weight_type": (WEIGHT_TYPES,),
"start_at": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"end_at": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
},
"optional": {
"mask": ("MASK",),
"enabled": ("BOOLEAN", {"default": True}),
},
}
RETURN_TYPES = ("MODEL", "IMAGE")
CATEGORY = "Art Venture/Style"
FUNCTION = "apply_style"
def apply_style(self, model, preset: str, data: str, mask=None, enabled=True, **kwargs):
data = json.loads(data or "[]")
data: IPAdapterData = IPAdapterData(images=data) # validate
if len(data.images) == 0:
images = torch.zeros((1, 64, 64, 3))
return (model, images)
(model, pipeline) = unifyLoader.load_models(model, preset)
urls = [image.url for image in data.images]
pils, _ = load_images_from_url(urls)
embeds_avg = None
neg_embeds_avg = None
images = []
for i, pil in enumerate(pils):
weight = data.images[i].weight
image = pil2tensor(pil)
if i > 0 and image.shape[1:] != images[0].shape[1:]:
image = comfy.utils.common_upscale(
image.movedim(-1, 1), images[0].shape[2], images[0].shape[1], "bilinear", "center"
).movedim(1, -1)
images.append(image)
embeds = encoder.encode(pipeline, image, weight, mask=mask)
if embeds_avg is None:
embeds_avg = embeds[0]
neg_embeds_avg = embeds[1]
else:
embeds_avg = combiner.batch(embeds_avg, method="average", embed2=embeds[0])[0]
neg_embeds_avg = combiner.batch(neg_embeds_avg, method="average", embed2=embeds[1])[0]
images = torch.cat(images)
model = embedder.apply_ipadapter(model, pipeline, embeds_avg, neg_embed=neg_embeds_avg, **kwargs)[0]
return (model, images)
NODE_CLASS_MAPPINGS.update(
{"AV_IPAdapter": AV_IPAdapter, "AV_IPAdapterPipe": AV_IPAdapterPipe, "AV_StyleApply": AV_StyleApply}
{
"AV_IPAdapter": AV_IPAdapter,
"AV_IPAdapterPipe": AV_IPAdapterPipe,
}
)
NODE_DISPLAY_NAME_MAPPINGS.update(
{
"AV_IPAdapter": "IP Adapter Apply",
"AV_IPAdapterPipe": "IP Adapter Pipe",
"AV_StyleApply": "AV Style Apply",
}
)
+3 -1
View File
@@ -1,12 +1,14 @@
from .segmenter import ISNetLoader, ISNetSegment
from .segmenter import ISNetLoader, ISNetSegment, DownloadISNetModel
NODE_CLASS_MAPPINGS = {
"ISNetLoader": ISNetLoader,
"ISNetSegment": ISNetSegment,
"DownloadISNetModel": DownloadISNetModel,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ISNetLoader": "ISNet Loader",
"ISNetSegment": "ISNet Segment",
"DownloadISNetModel": "Download and Load ISNet Model",
}
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+48 -25
View File
@@ -12,17 +12,30 @@ import folder_paths
import comfy.model_management as model_management
import comfy.utils
from ..model_utils import download_model
from ..utils import pil2tensor, tensor2pil, numpy2pil
from ..model_utils import download_file
from ..utils import pil2tensor, tensor2pil
from ..logger import logger
isnets = {}
cache_size = [1024, 1024]
gpu = model_management.get_torch_device()
cpu = torch.device("cpu")
model_dir = os.path.join(folder_paths.models_dir, "isnet")
model_url = "https://huggingface.co/NimaBoscarino/IS-Net_DIS-general-use/resolve/main/isnet-general-use.pth"
cache_size = [1024, 1024]
models = {
"isnet-general-use.pth": {
"url": "https://huggingface.co/NimaBoscarino/IS-Net_DIS-general-use/resolve/main/isnet-general-use.pth",
"sha": "9e1aafea58f0b55d0c35077e0ceade6ba1ba2bce372fd4f8f77215391f3fac13",
},
"isnetis.pth": {
"url": "https://github.com/Sanster/models/releases/download/isnetis/isnetis.pth",
"sha": "90a970badbd99ca7839b4e0beb09a36565d24edba7e4a876de23c761981e79e0",
},
"RMBG-1.4.bin": {
"url": "https://huggingface.co/briaai/RMBG-1.4/resolve/main/pytorch_model.bin",
"sha": "59569acdb281ac9fc9f78f9d33b6f9f17f68e25086b74f9025c35bb5f2848967",
},
}
folder_paths.folder_names_and_paths["isnet"] = (
[model_dir],
@@ -134,23 +147,40 @@ class ISNetLoader:
return {
"required": {
"model_name": (folder_paths.get_filename_list("isnet"),),
"model_override": ("STRING", {"default": "None"}),
},
}
RETURN_TYPES = ("ISNET_MODEL",)
FUNCTION = "load_isnet"
CATEGORY = "Art Venture/Segmentation"
CATEGORY = "ArtVenture/Segmentation"
def load_isnet(self, model_name, model_override="None"):
if model_override != "None":
if model_override not in folder_paths.get_filename_list("isnet"):
logger.warning(f"Model override {model_override} not found. Use {model_name} instead.")
else:
model_name = model_override
def load_isnet(self, model_name):
return (load_isnet_model(model_name),)
model = load_isnet_model(model_name)
return (model,)
class DownloadISNetModel:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model_name": (list(models.keys()),),
},
}
RETURN_TYPES = ("ISNET_MODEL",)
FUNCTION = "download_isnet"
CATEGORY = "ArtVenture/Segmentation"
def download_isnet(self, model_name):
if model_name not in folder_paths.get_filename_list("isnet"):
model_info = models[model_name]
download_file(
model_info["url"],
os.path.join(model_dir, model_name),
model_info["sha"],
)
return (load_isnet_model(model_name),)
class ISNetSegment:
@@ -170,7 +200,7 @@ class ISNetSegment:
RETURN_TYPES = ("IMAGE", "MASK")
RETURN_NAMES = ("segmented", "mask")
CATEGORY = "Art Venture/Segmentation"
CATEGORY = "ArtVenture/Segmentation"
FUNCTION = "segment_isnet"
def segment_isnet(self, images: torch.Tensor, threshold, device_mode="AUTO", enabled=True, isnet_model=None):
@@ -179,15 +209,8 @@ class ISNetSegment:
return (images, masks)
if isnet_model is None:
ckpts = folder_paths.get_filename_list("isnet")
if len(ckpts) == 0:
ckpts = download_model(
model_path=model_dir,
model_url=model_url,
ext_filter=[".pth"],
download_name="isnet-general-use.pth",
)
isnet_model = load_isnet_model(ckpts[0])
downloader = DownloadISNetModel()
isnet_model = downloader.download_isnet("isnet-general-use.pth")[0]
device = gpu if device_mode != "CPU" else cpu
isnet_model = isnet_model.to(device)
@@ -198,7 +221,7 @@ class ISNetSegment:
for image in images:
mask = predict(isnet_model, image, device)
mask_im = tensor2pil(mask.permute(1, 2, 0))
cropped = Image.new("RGBA", mask_im.size, (0,0,0,0))
cropped = Image.new("RGBA", mask_im.size, (0, 0, 0, 0))
cropped.paste(tensor2pil(image), mask=mask_im)
masks.append(mask)
+575 -116
View File
@@ -1,29 +1,114 @@
import os
import base64
import json
import requests
import os
from enum import Enum
from io import BytesIO
from typing import Any, Dict, List, Optional, Union
import requests
import torch
from PIL import Image
from pydantic import BaseModel
from torch import Tensor
from pydantic import BaseModel
from typing import List, Dict, Union, Optional, Any
from ..utils import ensure_package, pil2base64, pil2tensor, tensor2pil
def image_urls_to_tensor(
image_urls: List[str], timeout: int | None = 60
) -> Optional[Tensor]:
tensors: List[Tensor] = []
for url in image_urls:
try:
if isinstance(url, str) and url.startswith("data:"):
comma_index = url.find(",")
if comma_index == -1:
continue
header = url[:comma_index]
data_part = url[comma_index + 1 :]
if ";base64" in header:
raw_bytes = base64.b64decode(data_part)
else:
raw_bytes = data_part.encode("utf-8")
image = Image.open(BytesIO(raw_bytes)).convert("RGB")
else:
resp = requests.get(url, timeout=timeout)
resp.raise_for_status()
image = Image.open(BytesIO(resp.content)).convert("RGB")
tensor = pil2tensor(image)
tensors.append(tensor)
except Exception:
# Silently skip any image that fails to load/parse
continue
if tensors:
return torch.cat(tensors, dim=0)
return None
from ..utils import ensure_package, tensor2pil, pil2base64
gpt_models = [
"gpt-3.5-turbo",
"gpt-3.5-turbo-16k",
"gpt-5.4-mini",
"gpt-5.4",
"gpt-5",
"gpt-5-mini",
"gpt-5-nano",
"gpt-5-chat-latest",
"gpt-4o",
"gpt-4o-mini",
"gpt-4.1",
"gpt-4.1-mini",
"gpt-4.1-nano",
"gpt-4-turbo",
"gpt-4-vision-preview",
"gpt-4-turbo-preview",
"gpt-4-0125-preview",
"gpt-4-1106-preview",
"gpt-4-0613",
"gpt-4",
"gpt-4-vision-preview",
"o1",
"o1-mini",
"o1-preview",
"o1-pro",
"o3",
"o3-mini",
"o3-pro",
"o4-mini",
]
gpt_vision_models = ["gpt-4-turbo", "gpt-4-turbo-preview", "gpt-4-vision-preview"]
claude3_models = ["claude-3-opus-20240229", "claude-3-sonnet-20240229", "claude-3-haiku-20240307"]
claude2_models = ["claude-2.1"]
claude_models = [
"claude-sonnet-4-5-20250929",
"claude-haiku-4-5-20251001",
"claude-sonnet-4-20250514",
"claude-opus-4-1-20250805",
"claude-opus-4-20250514",
"claude-3-7-sonnet-latest",
"claude-3-7-sonnet-20250219",
"claude-3-5-sonnet-latest",
"claude-3-5-sonnet-20241022",
"claude-3-5-haiku-20241022",
"claude-3-opus-latest",
"claude-3-opus-20240229",
"claude-3-sonnet-20240229",
"claude-3-haiku-20240307",
]
nano_banana_models = {
"nano-banana": "gemini-2.5-flash-image",
"nano-banana-2": "gemini-3.1-flash-image-preview",
"nano-banana-pro": "gemini-3-pro-image-preview",
}
gemini_models = [
"gemini-3-flash-preview",
"gemini-3.1-pro-preview",
"gemini-2.5-flash",
"gemini-2.5-flash-lite",
"gemini-2.5-pro",
"gemini-2.0-flash",
"gemini-2.0-flash-lite",
]
aws_regions = [
"us-east-1",
@@ -39,23 +124,31 @@ aws_regions = [
bedrock_anthropic_versions = ["bedrock-2023-05-31"]
bedrock_claude3_models = [
bedrock_claude_models = [
"anthropic.claude-opus-4-20250514-v1:0",
"anthropic.claude-sonnet-4-20250514-v1:0",
"anthropic.claude-3-7-sonnet-20250219-v1:0",
"anthropic.claude-3-5-sonnet-20241022-v2:0",
"anthropic.claude-3-5-haiku-20241022-v1:0",
"anthropic.claude-3-haiku-20240307-v1:0",
"anthropic.claude-3-sonnet-20240229-v1:0",
"anthropic.claude-3-opus-20240229-v1:0",
]
bedrock_claude2_models = [
"anthropic.claude-v2",
"anthropic.claude-v2.1",
]
bedrock_mistral_models = [
"mistral.mistral-7b-instruct-v0:2",
"mistral.mixtral-8x7b-instruct-v0:1",
"mistral.mistral-large-2402-v1:0",
]
all_models = (
gpt_models
+ claude_models
+ gemini_models
+ bedrock_claude_models
+ bedrock_mistral_models
)
default_system_prompt = "You are a useful AI agent."
@@ -65,6 +158,12 @@ class LLMConfig(BaseModel):
temperature: float
class NanoBananaConfig(LLMConfig):
modalities: str = "image+text"
aspect_ratio: str = "auto"
resolution: str = "1K"
class LLMMessageRole(str, Enum):
system = "system"
user = "user"
@@ -74,13 +173,19 @@ class LLMMessageRole(str, Enum):
class LLMMessage(BaseModel):
role: LLMMessageRole = LLMMessageRole.user
text: str
image: Optional[str] = None # base64 enoded image
images: Optional[List[str]] = None # list of base64 encoded images
def to_openai_message(self):
content = [{"type": "text", "text": self.text}]
content: List[Dict[str, Any]] = [{"type": "text", "text": self.text}]
if self.image:
content.insert(0, {"type": "image_url", "image_url": {"url": f"data:image/png;base64,{self.text}"}})
if self.images:
for img in self.images:
content.append(
{
"type": "image_url",
"image_url": {"url": f"data:image/png;base64,{img}"},
}
)
return {
"role": self.role,
@@ -88,22 +193,48 @@ class LLMMessage(BaseModel):
}
def to_claude_message(self):
content = [{"type": "text", "text": self.text}]
content: List[Dict[str, Any]] = [{"type": "text", "text": self.text}]
if self.image:
content.insert(
0,
{
"type": "image",
"source": {"type": "base64", "media_type": "image/png", "data": self.image},
},
)
if self.images:
for img in reversed(self.images):
content.append(
{
"type": "image",
"source": {
"type": "base64",
"media_type": "image/png",
"data": img,
},
}
)
return {
"role": self.role,
"content": content,
}
def to_gemini_message(self):
parts: List[Dict[str, Any]] = [{"text": self.text}]
if self.images:
for img in self.images:
parts.append(
{
"inline_data": {
"mime_type": "image/png",
"data": img,
}
}
)
# Gemini uses "model" and "user" roles instead of "assistant" and "user"
role = "model" if self.role == "assistant" else "user"
return {
"role": role,
"parts": parts,
}
class OpenAIApi(BaseModel):
api_key: str
@@ -111,7 +242,7 @@ class OpenAIApi(BaseModel):
timeout: Optional[int] = 60
def chat(self, messages: List[LLMMessage], config: LLMConfig, seed=None):
if config.model not in gpt_models:
if config.model in all_models and config.model not in gpt_models:
raise Exception(f"Must provide an OpenAI model, got {config.model}")
formated_messages = [m.to_openai_message() for m in messages]
@@ -122,7 +253,6 @@ class OpenAIApi(BaseModel):
"model": config.model,
"max_tokens": config.max_token,
"temperature": config.temperature,
# "seed": seed,
}
headers = {"Authorization": f"Bearer {self.api_key}"}
@@ -130,9 +260,81 @@ class OpenAIApi(BaseModel):
data: Dict = response.json()
if data.get("error", None) is not None:
raise Exception(data.get("error").get("message"))
return (f"OpenAI API error: {data.get('error').get('message')}", None)
return data["choices"][0]["message"]["content"]
text = data["choices"][0]["message"]["content"]
return (text, None)
def complete(self, prompt: str, config: LLMConfig, seed=None):
messages = [LLMMessage(role=LLMMessageRole.user, text=prompt)]
return self.chat(messages, config, seed)
class OpenRouterApi(BaseModel):
api_key: str
endpoint: Optional[str] = "https://openrouter.ai/api/v1"
timeout: Optional[int] = 60
def chat(self, messages: List[LLMMessage], config: LLMConfig, seed=None):
formated_messages = [m.to_openai_message() for m in messages]
url = f"{self.endpoint}/chat/completions"
headers = {"Authorization": f"Bearer {self.api_key}"}
data = {
"messages": formated_messages,
"model": config.model,
"max_tokens": config.max_token,
"temperature": config.temperature,
}
if isinstance(config, NanoBananaConfig):
modalities = (config.modalities or "text+image").split("+")
modalities = [modality.strip().lower() for modality in modalities]
aspectRatio = (
config.aspect_ratio
if (config.aspect_ratio and config.aspect_ratio != "auto")
else None
)
data["modalities"] = modalities
data["image_config"] = {}
if aspectRatio is not None:
data["image_config"]["aspect_ratio"] = aspectRatio
# nano-banana 1 does not support imageSize
if "nano-banana-" in config.model:
data["image_config"]["image_size"] = config.resolution
model = nano_banana_models[config.model]
data["model"] = f"google/{model}"
response = requests.post(url, json=data, headers=headers, timeout=self.timeout)
data: Dict = response.json()
if data.get("error", None) is not None:
return (f"OpenRouter API error: {data.get('error').get('message')}", None)
message = data["choices"][0]["message"]
text = message["content"]
images = None
if message.get("images", None) is not None:
urls: List[str] = []
for m in message["images"]:
if isinstance(m, dict):
if (
"image_url" in m
and isinstance(m["image_url"], dict)
and "url" in m["image_url"]
):
urls.append(m["image_url"]["url"])
images = image_urls_to_tensor(urls, timeout=self.timeout)
return (text, images)
def complete(self, prompt: str, config: LLMConfig, seed=None):
messages = [LLMMessage(role=LLMMessageRole.user, text=prompt)]
@@ -147,8 +349,8 @@ class ClaudeApi(BaseModel):
timeout: Optional[int] = 60
def chat(self, messages: List[LLMMessage], config: LLMConfig, seed=None):
if config.model not in claude3_models:
raise Exception(f"Must provide a Claude v3 model, got {config.model}")
if config.model in all_models and config.model not in claude_models:
raise Exception(f"Must provide a Claude model, got {config.model}")
system_message = [m for m in messages if m.role == "system"]
user_messages = [m for m in messages if m.role != "system"]
@@ -168,30 +370,114 @@ class ClaudeApi(BaseModel):
data: Dict = response.json()
if data.get("error", None) is not None:
raise Exception(data.get("error").get("message"))
return (data.get("error").get("message"), None)
return data["content"][0]["text"]
text = data["content"][0]["text"]
return (text, None)
def complete(self, prompt: str, config: LLMConfig, seed=None):
if config.model not in claude2_models:
raise Exception(f"Must provide a Claude v2 model, got {config.model}")
messages = [LLMMessage(role=LLMMessageRole.user, text=prompt)]
return self.chat(messages, config, seed)
class GeminiApi(BaseModel):
api_key: str
endpoint: Optional[str] = "https://generativelanguage.googleapis.com/v1beta"
timeout: Optional[int] = 60
def chat(self, messages: List[LLMMessage], config: LLMConfig, seed=None):
if config.model in all_models and config.model not in gemini_models:
raise Exception(f"Must provide a Gemini model, got {config.model}")
system_message = [m for m in messages if m.role == "system"]
user_messages = [m for m in messages if m.role != "system"]
if not user_messages:
return (
"Gemini API error: At least one user message is required. System messages alone are not sufficient.",
None,
)
formated_messages = [m.to_gemini_message() for m in user_messages]
url = f"{self.endpoint}/models/{config.model}:generateContent"
headers = {"x-goog-api-key": self.api_key}
prompt = f"\n\nHuman: {prompt}\n\nAssistant:"
url = f"{self.endpoint}/complete"
data = {
"prompt": prompt,
"max_tokens_to_sample": config.max_token,
"temperature": config.temperature,
"contents": formated_messages,
"generationConfig": {
"maxOutputTokens": config.max_token,
"temperature": config.temperature,
},
}
headers = {"x-api-key": self.api_key, "anthropic-version": self.version}
response = requests.post(url, json=data, headers=headers, timeout=self.timeout)
if isinstance(config, NanoBananaConfig):
modalities = (config.modalities or "text+image").split("+")
modalities = [modality.strip().capitalize() for modality in modalities]
aspectRatio = (
config.aspect_ratio
if (config.aspect_ratio and config.aspect_ratio != "auto")
else None
)
data["generationConfig"]["responseModalities"] = modalities
data["generationConfig"]["imageConfig"] = {
"aspectRatio": aspectRatio,
}
# nano-banana 1 does not support imageSize
if "nano-banana-" in config.model:
data["generationConfig"]["imageConfig"]["imageSize"] = config.resolution
model = nano_banana_models[config.model]
url = f"{self.endpoint}/models/{model}:generateContent"
# Add system instruction if provided
if len(system_message) > 0:
data["systemInstruction"] = {"parts": [{"text": system_message[0].text}]}
response = requests.post(url, headers=headers, json=data, timeout=self.timeout)
data: Dict = response.json()
if data.get("error", None) is not None:
raise Exception(data.get("error").get("message"))
error_message = data.get("error").get("message", "Unknown error")
return (f"Gemini API error: {error_message}", None)
return data["completion"]
# Extract text and images from response
if "candidates" in data and len(data["candidates"]) > 0:
candidate = data["candidates"][0]
if "content" in candidate and "parts" in candidate["content"]:
parts = candidate["content"]["parts"]
# Collect text parts
text_parts = []
image_urls = []
for part in parts:
# Handle text
if "text" in part and part["text"]:
text_parts.append(part["text"])
# Handle inline images
elif "inlineData" in part and part["inlineData"]:
image_urls.append(
f"data:{part['inlineData']['mimeType']};base64,{part['inlineData']['data']}"
)
text = (
"".join(text_parts)
if text_parts
else "No text response from Gemini API"
)
images = image_urls_to_tensor(image_urls, timeout=self.timeout)
return (text, images)
return ("No response from Gemini API", None)
def complete(self, prompt: str, config: LLMConfig, seed=None):
messages = [LLMMessage(role=LLMMessageRole.user, text=prompt)]
return self.chat(messages, config, seed)
class AwsBedrockMistralApi(BaseModel):
@@ -205,7 +491,7 @@ class AwsBedrockMistralApi(BaseModel):
def __init__(self, **data):
super().__init__(**data)
ensure_package("boto3", version="1.34.101")
ensure_package("boto3", required_version=">=1.34.101")
import boto3
self.bedrock_runtime = boto3.client(
@@ -230,13 +516,16 @@ class AwsBedrockMistralApi(BaseModel):
"temperature": config.temperature,
}
response = self.bedrock_runtime.invoke_model(body=json.dumps(data), modelId=config.model)
response = self.bedrock_runtime.invoke_model(
body=json.dumps(data), modelId=config.model
)
data: Dict = json.loads(response.get("body").read())
if data.get("error", None) is not None:
raise Exception(data.get("error").get("message"))
return (f"Mistral API error: {data.get('error').get('message')}", None)
return data["outputs"][0]["text"]
text = data["outputs"][0]["text"]
return (text, None)
class AwsBedrockClaudeApi(BaseModel):
@@ -251,7 +540,7 @@ class AwsBedrockClaudeApi(BaseModel):
def __init__(self, **data):
super().__init__(**data)
ensure_package("boto3", version="1.34.101")
ensure_package("boto3", required_version=">=1.34.101")
import boto3
self.bedrock_runtime = boto3.client(
@@ -263,7 +552,7 @@ class AwsBedrockClaudeApi(BaseModel):
)
def chat(self, messages: List[LLMMessage], config: LLMConfig, seed=None):
if config.model not in bedrock_claude3_models:
if config.model not in bedrock_claude_models:
raise Exception(f"Must provide a Claude v3 model, got {config.model}")
system_message = [m for m in messages if m.role == "system"]
@@ -278,35 +567,23 @@ class AwsBedrockClaudeApi(BaseModel):
"system": system_message[0].text if len(system_message) > 0 else None,
}
response = self.bedrock_runtime.invoke_model(body=json.dumps(data), modelId=config.model)
response = self.bedrock_runtime.invoke_model(
body=json.dumps(data), modelId=config.model
)
data: Dict = json.loads(response.get("body").read())
if data.get("error", None) is not None:
raise Exception(data.get("error").get("message"))
return (f"Claude API error: {data.get('error').get('message')}", None)
return data["content"][0]["text"]
text = data["content"][0]["text"]
return (text, None)
def complete(self, prompt: str, config: LLMConfig, seed=None):
if config.model not in bedrock_claude2_models:
raise Exception(f"Must provide a Claude v2 model, got {config.model}")
prompt = f"\n\nHuman: {prompt}\n\nAssistant:"
data = {
"prompt": prompt,
"max_tokens_to_sample": config.max_token,
"temperature": config.temperature,
}
response = self.bedrock_runtime.invoke_model(body=json.dumps(data), modelId=config.model)
data: Dict = json.loads(response.get("body").read())
if data.get("error", None) is not None:
raise Exception(data.get("error").get("message"))
return data["completion"]
messages = [LLMMessage(role=LLMMessageRole.user, text=prompt)]
return self.chat(messages, config, seed)
LLMApi = Union[OpenAIApi, ClaudeApi]
LLMApi = Union[OpenAIApi, OpenRouterApi, ClaudeApi, GeminiApi]
class OpenAIApiNode:
@@ -315,7 +592,10 @@ class OpenAIApiNode:
return {
"required": {
"openai_api_key": ("STRING", {"multiline": False}),
"endpoint": ("STRING", {"multiline": False, "default": "https://api.openai.com/v1"}),
"endpoint": (
"STRING",
{"multiline": False, "default": "https://api.openai.com/v1"},
),
},
}
@@ -332,13 +612,42 @@ class OpenAIApiNode:
return (OpenAIApi(api_key=openai_api_key, endpoint=endpoint),)
class OpenRouterApiNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"openrouter_api_key": ("STRING", {"multiline": False}),
"endpoint": (
"STRING",
{"multiline": False, "default": "https://openrouter.ai/api/v1"},
),
},
}
RETURN_TYPES = ("LLM_API",)
FUNCTION = "create_api"
CATEGORY = "ArtVenture/LLM"
def create_api(self, openrouter_api_key, endpoint):
if not openrouter_api_key or openrouter_api_key == "":
openrouter_api_key = os.environ.get("OPENROUTER_API_KEY")
if not openrouter_api_key:
raise Exception("OpenRouter API key is required.")
return (OpenRouterApi(api_key=openrouter_api_key, endpoint=endpoint),)
class ClaudeApiNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"claude_api_key": ("STRING", {"multiline": False}),
"endpoint": ("STRING", {"multiline": False, "default": "https://api.anthropic.com/v1"}),
"endpoint": (
"STRING",
{"multiline": False, "default": "https://api.anthropic.com/v1"},
),
"version": (["2023-06-01"], {"default": "2023-06-01"}),
},
}
@@ -350,13 +659,47 @@ class ClaudeApiNode:
def create_api(self, claude_api_key, endpoint, version):
if not claude_api_key or claude_api_key == "":
claude_api_key = os.environ.get("CLAUDE_API_KEY")
claude_api_key = os.environ.get(
"ANTHROPIC_API_KEY", os.environ.get("CLAUDE_API_KEY")
)
if not claude_api_key:
raise Exception("Claude API key is required.")
raise Exception("Anthropic API key is required.")
return (ClaudeApi(api_key=claude_api_key, endpoint=endpoint, version=version),)
class GeminiApiNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"gemini_api_key": ("STRING", {"multiline": False}),
"endpoint": (
"STRING",
{
"multiline": False,
"default": "https://generativelanguage.googleapis.com/v1beta",
},
),
},
}
RETURN_TYPES = ("LLM_API",)
RETURN_NAMES = ("llm_api",)
FUNCTION = "create_api"
CATEGORY = "ArtVenture/LLM"
def create_api(self, gemini_api_key, endpoint):
if not gemini_api_key or gemini_api_key == "":
gemini_api_key = os.environ.get(
"GEMINI_API_KEY", os.environ.get("GOOGLE_API_KEY")
)
if not gemini_api_key:
raise Exception("Gemini API key is required.")
return (GeminiApi(api_key=gemini_api_key, endpoint=endpoint),)
class AwsBedrockMistralApiNode:
@classmethod
def INPUT_TYPES(cls):
@@ -374,7 +717,9 @@ class AwsBedrockMistralApiNode:
FUNCTION = "create_api"
CATEGORY = "ArtVenture/LLM"
def create_api(self, aws_access_key_id, aws_secret_access_key, aws_session_token, region):
def create_api(
self, aws_access_key_id, aws_secret_access_key, aws_session_token, region
):
if not aws_access_key_id or aws_access_key_id == "":
aws_access_key_id = os.environ.get("AWS_ACCESS_KEY_ID", None)
if not aws_secret_access_key or aws_secret_access_key == "":
@@ -404,7 +749,10 @@ class AwsBedrockClaudeApiNode:
"aws_secret_access_key": ("STRING", {"multiline": False}),
"aws_session_token": ("STRING", {"multiline": False}),
"region": (aws_regions, {"default": aws_regions[0]}),
"version": (bedrock_anthropic_versions, {"default": bedrock_anthropic_versions[0]}),
"version": (
bedrock_anthropic_versions,
{"default": bedrock_anthropic_versions[0]},
),
},
}
@@ -413,7 +761,14 @@ class AwsBedrockClaudeApiNode:
FUNCTION = "create_api"
CATEGORY = "ArtVenture/LLM"
def create_api(self, aws_access_key_id, aws_secret_access_key, aws_session_token, region, version):
def create_api(
self,
aws_access_key_id,
aws_secret_access_key,
aws_session_token,
region,
version,
):
if not aws_access_key_id or aws_access_key_id == "":
aws_access_key_id = os.environ.get("AWS_ACCESS_KEY_ID", None)
if not aws_secret_access_key or aws_secret_access_key == "":
@@ -441,17 +796,18 @@ class LLMApiConfigNode:
return {
"required": {
"model": (
gpt_models
+ claude3_models
+ claude2_models
+ bedrock_claude3_models
+ bedrock_claude2_models
+ bedrock_mistral_models,
{"default": gpt_vision_models[0]},
all_models,
{"default": gpt_models[0]},
),
"max_token": ("INT", {"default": 1024}),
"temperature": ("FLOAT", {"default": 0, "min": 0, "max": 1.0, "step": 0.001}),
}
"max_token": ("INT", {"default": 1024, "min": 1, "max": 102400}),
"temperature": (
"FLOAT",
{"default": 0.5, "min": 0, "max": 1.0, "step": 0.001},
),
},
"optional": {
"custom_model": ("STRING", {"multiline": False, "default": ""})
},
}
RETURN_TYPES = ("LLM_CONFIG",)
@@ -459,8 +815,77 @@ class LLMApiConfigNode:
FUNCTION = "make_config"
CATEGORY = "ArtVenture/LLM"
def make_config(self, max_token, model, temperature):
return (LLMConfig(model=model, max_token=max_token, temperature=temperature),)
def make_config(self, model, max_token, temperature, custom_model=""):
# Use custom_model if provided, otherwise use the selected model from dropdown
final_model = (
custom_model.strip() if custom_model and custom_model.strip() else model
)
return (
LLMConfig(model=final_model, max_token=max_token, temperature=temperature),
)
class NanoBananaApiConfigNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": (
list(nano_banana_models.keys()),
{"default": "nano-banana"},
),
"temperature": (
"FLOAT",
{"default": 0.5, "min": 0, "max": 1.0, "step": 0.001},
),
"modalities": (
["image+text", "image", "text"],
{"default": "image+text"},
),
"aspect_ratio": (
[
"auto",
"1:1",
"2:3",
"3:2",
"3:4",
"4:3",
"4:5",
"5:4",
"9:16",
"16:9",
"21:9",
],
{"default": "auto"},
),
"resolution": (["1K", "2K", "4K"], {"default": "1K"}),
},
"optional": {},
}
RETURN_TYPES = ("LLM_CONFIG",)
RETURN_NAMES = ("llm_config",)
FUNCTION = "make_config"
CATEGORY = "ArtVenture/LLM"
def make_config(
self,
model,
temperature,
modalities="image+text",
aspect_ratio="auto",
resolution="1K",
):
return (
NanoBananaConfig(
model=model,
max_token=8192,
temperature=temperature,
modalities=modalities,
aspect_ratio=aspect_ratio,
resolution=resolution,
),
)
class LLMMessageNode:
@@ -471,7 +896,13 @@ class LLMMessageNode:
"role": (["system", "user", "assistant"],),
"text": ("STRING", {"multiline": True}),
},
"optional": {"image": ("IMAGE",), "messages": ("LLM_MESSAGE",)},
"optional": {
"messages": ("LLM_MESSAGE",),
"image": ("IMAGE",),
"image_2": ("IMAGE",),
"image_3": ("IMAGE",),
"image_4": ("IMAGE",),
},
}
RETURN_TYPES = ("LLM_MESSAGE",)
@@ -479,21 +910,43 @@ class LLMMessageNode:
FUNCTION = "make_message"
CATEGORY = "ArtVenture/LLM"
def make_message(self, role, text, image: Optional[Tensor] = None, messages: Optional[List[LLMMessage]] = None):
def make_message(
self,
role,
text,
messages: Optional[List[LLMMessage]] = None,
image: Optional[Tensor] = None,
image_2: Optional[Tensor] = None,
image_3: Optional[Tensor] = None,
image_4: Optional[Tensor] = None,
):
messages = [] if messages is None else messages.copy()
if role == "system":
if isinstance(image, Tensor):
raise Exception("System prompt does not support image.")
system_message = [m for m in messages if m.role == "system"]
if len(system_message) > 0:
raise Exception("Only one system prompt is allowed.")
if isinstance(image, Tensor):
pil = tensor2pil(image)
content = pil2base64(pil)
messages.append(LLMMessage(role=role, text=text, image=content))
if any(
isinstance(img, Tensor) for img in [image, image_2, image_3, image_4]
):
raise Exception("System prompt does not support image.")
all_images = []
for img_tensor in [image, image_2, image_3, image_4]:
if isinstance(img_tensor, Tensor):
if len(img_tensor.shape) == 4: # Batch of images
for i in range(img_tensor.shape[0]):
pil = tensor2pil(img_tensor[i])
content = pil2base64(pil)
all_images.append(content)
else: # Single image
pil = tensor2pil(img_tensor)
content = pil2base64(pil)
all_images.append(content)
if all_images:
messages.append(LLMMessage(role=role, text=text, images=all_images))
else:
messages.append(LLMMessage(role=role, text=text))
@@ -512,14 +965,14 @@ class LLMChatNode:
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("response",)
RETURN_TYPES = ("STRING", "IMAGE")
RETURN_NAMES = ("text", "(optional) images")
FUNCTION = "chat"
CATEGORY = "ArtVenture/LLM"
def chat(self, messages: List[LLMMessage], api: LLMApi, config: LLMConfig, seed):
response = api.chat(messages, config, seed)
return (response,)
text, images = api.chat(messages, config, seed)
return (text, images)
class LLMCompletionNode:
@@ -534,22 +987,25 @@ class LLMCompletionNode:
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("response",)
RETURN_TYPES = ("STRING", "IMAGE")
RETURN_NAMES = ("text", "(optional) images")
FUNCTION = "chat"
CATEGORY = "ArtVenture/LLM"
def chat(self, prompt: str, api: LLMApi, config: LLMConfig, seed):
response = api.complete(prompt, config, seed)
return (response,)
text, images = api.complete(prompt, config, seed)
return (text, images)
NODE_CLASS_MAPPINGS = {
"AV_OpenAIApi": OpenAIApiNode,
"AV_OpenRouterApi": OpenRouterApiNode,
"AV_ClaudeApi": ClaudeApiNode,
"AV_GeminiApi": GeminiApiNode,
"AV_AwsBedrockClaudeApi": AwsBedrockClaudeApiNode,
"AV_AwsBedrockMistralApi": AwsBedrockMistralApiNode,
"AV_LLMApiConfig": LLMApiConfigNode,
"AV_NanoBananaApiConfig": NanoBananaApiConfigNode,
"AV_LLMMessage": LLMMessageNode,
"AV_LLMChat": LLMChatNode,
"AV_LLMCompletion": LLMCompletionNode,
@@ -557,10 +1013,13 @@ NODE_CLASS_MAPPINGS = {
NODE_DISPLAY_NAME_MAPPINGS = {
"AV_OpenAIApi": "OpenAI API",
"AV_OpenRouterApi": "OpenRouter API",
"AV_ClaudeApi": "Claude API",
"AV_GeminiApi": "Gemini API",
"AV_AwsBedrockClaudeApi": "AWS Bedrock Claude API",
"AV_AwsBedrockMistralApi": "AWS Bedrock Mistral API",
"AV_LLMApiConfig": "LLM API Config",
"AV_NanoBananaApiConfig": "NanoBanana API Config",
"AV_LLMMessage": "LLM Message",
"AV_LLMChat": "LLM Chat",
"AV_LLMCompletion": "LLM Completion",
+82 -32
View File
@@ -1,7 +1,12 @@
import os
import re
import torch
import hashlib
import urllib.request
import urllib.error
from tqdm import tqdm
from urllib.parse import urlparse
from typing import Dict, Optional
def natural_sort_key(s, regex=re.compile("([0-9]+)")):
@@ -37,7 +42,7 @@ def load_file_from_url(
*,
model_dir: str,
progress: bool = True,
file_name: str | None = None,
file_name: Optional[str] = None,
) -> str:
"""Download a file from `url` into `model_dir`, using the file present if possible.
@@ -56,44 +61,89 @@ def load_file_from_url(
return cached_file
def download_model(
model_path: str,
model_url: str = None,
ext_filter=None,
download_name=None,
ext_blacklist=None,
) -> list:
"""
A one-and done loader to try finding the desired models in specified directories.
def calculate_sha(file: str, force=False) -> Optional[str]:
sha_file = f"{file}.sha"
@param download_name: Specify to download from model_url immediately.
@param model_url: If no other models are found, this will be downloaded on upscale.
@param model_path: The location to store/find models in.
@param ext_filter: An optional list of filename extensions to filter by
@return: A list of paths containing the desired model(s)
# Check if the .sha file exists
if not force and os.path.exists(sha_file):
try:
with open(sha_file, "r") as f:
stored_hash = f.read().strip()
if stored_hash:
return stored_hash
except IOError as e:
print(f"Failed to read hash: {e}")
# Calculate the hash if the .sha file doesn't exist or is empty
try:
with open(file, "rb") as fp:
file_hash = hashlib.sha256()
while chunk := fp.read(8192):
file_hash.update(chunk)
calculated_hash = file_hash.hexdigest()
# Write the calculated hash to the .sha file
try:
with open(sha_file, "w") as f:
f.write(calculated_hash)
except IOError as e:
print(f"Failed to write hash to {sha_file}: {e}")
return calculated_hash
except IOError as e:
print(f"Failed to read file {file}: {e}")
return None
def download_file(url: str, dst: str, sha256sum: Optional[str] = None) -> Dict[str, Optional[str]]:
"""
output = []
Downloads a file from a URL to a destination path, optionally verifying its SHA-256 checksum.
:param url: URL of the file to download
:param dst: Destination path to save the downloaded file
:param sha256sum: Optional SHA-256 checksum to verify the downloaded file
:return: Dictionary with file path, download status, calculated checksum, and checksum match status
"""
# Ensure the directory exists
os.makedirs(os.path.dirname(dst), exist_ok=True)
file_exists = os.path.isfile(dst)
file_checksum = None
checksum_match = None
downloaded = False
try:
for full_path in walk_files(model_path, allowed_extensions=ext_filter):
if os.path.islink(full_path) and not os.path.exists(full_path):
print(f"Skipping broken symlink: {full_path}")
continue
if ext_blacklist is not None and any(full_path.endswith(x) for x in ext_blacklist):
continue
if full_path not in output:
output.append(full_path)
if file_exists:
file_checksum = calculate_sha(dst)
if sha256sum:
checksum_match = file_checksum == sha256sum
if not checksum_match:
os.remove(dst)
if model_url is not None and len(output) == 0:
if download_name is not None:
output.append(load_file_from_url(model_url, model_dir=model_path, file_name=download_name))
else:
output.append(model_url)
if not file_exists or checksum_match == False:
with tqdm(unit="B", unit_scale=True, unit_divisor=1024, miniters=1, desc=dst.split("/")[-1]) as t:
except Exception:
pass
def reporthook(blocknum, blocksize, totalsize):
if t.total is None and totalsize > 0:
t.total = totalsize
read_so_far = blocknum * blocksize
t.update(max(0, read_so_far - t.n))
return output
urllib.request.urlretrieve(url, dst, reporthook=reporthook)
downloaded = True
file_checksum = calculate_sha(dst, force=True)
if sha256sum:
checksum_match = file_checksum == sha256sum
except urllib.error.URLError as ex:
print("Download failed:", ex)
if os.path.isfile(dst):
os.remove(dst)
except Exception as ex:
print("An error occurred:", ex)
finally:
return {"file": dst, "downloaded": downloaded, "sha": file_checksum, "match": checksum_match}
def load_jit_torch_file(model_path: str):
+17 -26
View File
@@ -62,8 +62,6 @@ from .llm import (
NODE_DISPLAY_NAME_MAPPINGS as LLM_NODE_DISPLAY_NAME_MAPPINGS,
)
from .model_utils import load_file_from_url
class AVVAELoader(VAELoader):
@classmethod
@@ -72,7 +70,7 @@ class AVVAELoader(VAELoader):
inputs["optional"] = {"vae_override": ("STRING", {"default": "None"})}
return inputs
CATEGORY = "Art Venture/Loaders"
CATEGORY = "ArtVenture/Loaders"
def load_vae(self, vae_name, vae_override="None"):
if vae_override != "None":
@@ -94,7 +92,7 @@ class AVLoraLoader(LoraLoader):
}
return inputs
CATEGORY = "Art Venture/Loaders"
CATEGORY = "ArtVenture/Loaders"
def load_lora(self, model, clip, lora_name, *args, lora_override="None", enabled=True, **kwargs):
if not enabled:
@@ -121,7 +119,7 @@ class AVLoraListStacker:
RETURN_TYPES = ("LORA_STACK",)
FUNCTION = "load_list_lora"
CATEGORY = "Art Venture/Loaders"
CATEGORY = "ArtVenture/Loaders"
def parse_lora_list(self, data: str):
# data is a list of lora model (lora_name, strength_model, strength_clip, url) in json format
@@ -215,7 +213,7 @@ class AVCheckpointModelsToParametersPipe:
}
RETURN_TYPES = ("PIPE",)
CATEGORY = "Art Venture/Parameters"
CATEGORY = "ArtVenture/Parameters"
FUNCTION = "checkpoint_models_to_parameter_pipe"
def checkpoint_models_to_parameter_pipe(
@@ -257,7 +255,7 @@ class AVPromptsToParametersPipe:
}
RETURN_TYPES = ("PIPE",)
CATEGORY = "Art Venture/Parameters"
CATEGORY = "ArtVenture/Parameters"
FUNCTION = "prompt_to_parameter_pipe"
def prompt_to_parameter_pipe(self, positive, negative, pipe: Dict = {}, image=None, mask=None):
@@ -299,7 +297,7 @@ class AVParametersPipeToCheckpointModels:
"lora_2_name",
"lora_3_name",
)
CATEGORY = "Art Venture/Parameters"
CATEGORY = "ArtVenture/Parameters"
FUNCTION = "parameter_pipe_to_checkpoint_models"
def parameter_pipe_to_checkpoint_models(self, pipe: Dict = {}):
@@ -348,7 +346,7 @@ class AVParametersPipeToPrompts:
"image",
"mask",
)
CATEGORY = "Art Venture/Parameters"
CATEGORY = "ArtVenture/Parameters"
FUNCTION = "parameter_pipe_to_prompt"
def parameter_pipe_to_prompt(self, pipe: Dict = {}):
@@ -374,23 +372,23 @@ class AVCheckpointMerge:
"model1": ("MODEL",),
"model2": ("MODEL",),
"model1_weight": ("FLOAT", {"default": 1.0, "min": -1.0, "max": 1.0, "step": 0.01}),
"model2_weight": ("FLOAT", {"default": 1.0, "min": -1.0, "max": 1.0, "step": 0.01}),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "merge"
CATEGORY = "Art Venture/Model Merging"
CATEGORY = "ArtVenture/Model Merging"
DESCRIPTION = "DEPRECATED: Use ComfyUI's native ModelMergeSimple instead"
def merge(self, model1, model2, model1_weight, model2_weight):
def merge(self, model1, model2, model1_weight):
m = model1.clone()
k1 = model1.get_key_patches("diffusion_model.")
k2 = model2.get_key_patches("diffusion_model.")
for k in k1:
if k in k2:
a = k1[k][0]
b = k2[k][0]
for k in k2:
if k in k1:
a, _ = k1[k][0]
b, _ = k2[k][0]
if a.shape != b.shape and a.shape[0:1] + a.shape[2:] == b.shape[0:1] + b.shape[2:]:
if a.shape[1] == 4 and b.shape[1] == 9:
@@ -402,20 +400,13 @@ class AVCheckpointMerge:
"When merging instruct-pix2pix model with a normal one, model1 must be the instruct-pix2pix model."
)
c = torch.zeros_like(a)
c[:, 0:4, :, :] = b
b = c
m.add_patches({k: (b,)}, model2_weight, model1_weight)
else:
logger.warn(f"Key {k} not found in model2")
m.add_patches({k: k1[k]}, -1.0, 1.0) # zero out
m.add_patches({k: k2[k]}, 1 - model1_weight, model1_weight)
return (m,)
class AVCheckpointSave(CheckpointSave):
CATEGORY = "Art Venture/Model Merging"
CATEGORY = "ArtVenture/Model Merging"
@classmethod
def INPUT_TYPES(s):
@@ -464,7 +455,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"AV_LoraLoader": "Lora Loader",
"AV_LoraListLoader": "Lora List Loader",
"AV_LoraListStacker": "Lora List Stacker",
"AV_CheckpointMerge": "Checkpoint Merge",
"AV_CheckpointMerge": "[Deprecated] Checkpoint Merge",
"AV_CheckpointSave": "Checkpoint Save",
}
+4 -10
View File
@@ -37,17 +37,13 @@ class ColorBlend:
RETURN_TYPES = ("IMAGE",)
FUNCTION = "color_blending_mode"
CATEGORY = "Art Venture/Post Processing"
CATEGORY = "ArtVenture/Post Processing"
def color_blending_mode(self, bw_layer, color_layer):
if bw_layer.shape[0] < color_layer.shape[0]:
bw_layer = bw_layer.repeat(color_layer.shape[0], 1, 1, 1)[
: color_layer.shape[0]
]
bw_layer = bw_layer.repeat(color_layer.shape[0], 1, 1, 1)[: color_layer.shape[0]]
if bw_layer.shape[0] > color_layer.shape[0]:
color_layer = color_layer.repeat(bw_layer.shape[0], 1, 1, 1)[
: bw_layer.shape[0]
]
color_layer = color_layer.repeat(bw_layer.shape[0], 1, 1, 1)[: bw_layer.shape[0]]
batch_size, *_ = bw_layer.shape
tensor_output = torch.empty_like(bw_layer)
@@ -70,8 +66,6 @@ class ColorBlend:
for i in range(batch_size):
blend = color_blend(image1[i], image2[i])
blend = np.stack([blend])
tensor_output[i : i + 1] = (
torch.from_numpy(blend.transpose(0, 3, 1, 2)) / 255.0
).permute(0, 2, 3, 1)
tensor_output[i : i + 1] = (torch.from_numpy(blend.transpose(0, 3, 1, 2)) / 255.0).permute(0, 2, 3, 1)
return (tensor_output,)
+1 -1
View File
@@ -36,7 +36,7 @@ class ColorCorrect:
RETURN_TYPES = ("IMAGE",)
FUNCTION = "color_correct"
CATEGORY = "Art Venture/Post Processing"
CATEGORY = "ArtVenture/Post Processing"
def color_correct(
self,
+239 -130
View File
@@ -5,7 +5,7 @@ import torch
import base64
import random
import requests
from typing import List, Dict, Tuple
from typing import List, Dict, Tuple, Optional
from PIL import Image, ImageOps, ImageFilter
import numpy as np
@@ -20,6 +20,42 @@ from .utils import pil2tensor, tensor2pil, ensure_package, get_dict_attribute
MAX_RESOLUTION = 8192
class AnyType(str):
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
def __ne__(self, __value: object) -> bool:
return False
class FlexibleOptionalInputType(dict):
"""A special class to make flexible nodes that pass data to our python handlers.
Enables both flexible/dynamic input types (like for Any Switch) or a dynamic number of inputs
(like for Any Switch, Context Switch, Context Merge, Power Lora Loader, etc).
Note, for ComfyUI, all that's needed is the `__contains__` override below, which tells ComfyUI
that our node will handle the input, regardless of what it is.
However, with https://github.com/comfyanonymous/ComfyUI/pull/2666 a large change would occur
requiring more details on the input itself. There, we need to return a list/tuple where the first
item is the type. This can be a real type, or use the AnyType for additional flexibility.
This should be forwards compatible unless more changes occur in the PR.
"""
def __init__(self, type):
self.type = type
def __getitem__(self, key):
return (self.type,)
def __contains__(self, key):
return True
any_type = AnyType("*")
def prepare_image_for_preview(image: Image.Image, output_dir: str, prefix=None):
if prefix is None:
prefix = "preview_" + "".join(random.choice("abcdefghijklmnopqrstupvxyz") for x in range(5))
@@ -43,8 +79,8 @@ def prepare_image_for_preview(image: Image.Image, output_dir: str, prefix=None):
def load_images_from_url(urls: List[str], keep_alpha_channel=False):
images = []
masks = []
images: List[Image.Image] = []
masks: List[Optional[Image.Image]] = []
for url in urls:
if url.startswith("data:image/"):
@@ -105,14 +141,8 @@ def load_images_from_url(urls: List[str], keep_alpha_channel=False):
if has_alpha:
mask = i.getchannel("A")
# recreate image to fix weird RGB image
alpha = i.split()[-1]
image = Image.new("RGB", i.size, (0, 0, 0))
image.paste(i, mask=alpha)
image.putalpha(alpha)
if not keep_alpha_channel:
image = image.convert("RGB")
if not keep_alpha_channel:
image = i.convert("RGB")
else:
image = i
@@ -128,11 +158,20 @@ class UtilLoadImageFromUrl:
self.filename_prefix = "TempImageFromUrl"
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {},
"required": {
"image": (
"STRING",
{
"default": "",
"placeholder": "Input image paths or URLS one per line. Eg:\nhttps://example.com/image.png\nfile:///path/to/local/image.jpg\ndata:image/png;base64,...",
"multiline": True,
"dynamicPrompts": False,
},
),
},
"optional": {
"image": ("STRING", {"default": "", "multiline": True, "dynamicPrompts": False}),
"keep_alpha_channel": (
"BOOLEAN",
{"default": False, "label_on": "enabled", "label_off": "disabled"},
@@ -141,71 +180,81 @@ class UtilLoadImageFromUrl:
"BOOLEAN",
{"default": False, "label_on": "list", "label_off": "batch"},
),
"url": ("STRING", {"default": "", "multiline": True, "dynamicPrompts": False}),
},
}
RETURN_TYPES = ("IMAGE", "MASK", "BOOLEAN")
OUTPUT_IS_LIST = (True, True, False)
RETURN_NAMES = ("images", "masks", "has_image")
CATEGORY = "Art Venture/Image"
CATEGORY = "ArtVenture/Image"
FUNCTION = "load_image"
def load_image(self, image="", keep_alpha_channel=False, output_mode=False, url=""):
if not image or image == "":
image = url
def load_image(self, image: str, keep_alpha_channel=False, output_mode=False):
urls = image.strip().split("\n")
images, masks = load_images_from_url(urls, keep_alpha_channel)
if len(images) == 0:
image = torch.zeros((1, 64, 64, 3), dtype=torch.float32, device="cpu")
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
images = [tensor2pil(image)]
masks = [tensor2pil(mask, mode="L")]
pil_images, pil_masks = load_images_from_url(urls, keep_alpha_channel)
has_image = len(pil_images) > 0
if not has_image:
i = torch.zeros((1, 64, 64, 3), dtype=torch.float32, device="cpu")
m = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
pil_images = [tensor2pil(i)]
pil_masks = [tensor2pil(m, mode="L")]
previews = []
np_images = []
np_masks = []
np_images: list[torch.Tensor] = []
np_masks: list[torch.Tensor] = []
for image, mask in zip(images, masks):
# save image to temp folder
preview = prepare_image_for_preview(image, self.output_dir, self.filename_prefix)
image = pil2tensor(image)
if mask:
mask = np.array(mask).astype(np.float32) / 255.0
mask = 1.0 - torch.from_numpy(mask)
for pil_image, pil_mask in zip(pil_images, pil_masks):
if pil_mask is not None:
preview_image = Image.new("RGB", pil_image.size)
preview_image.paste(pil_image, (0, 0))
preview_image.putalpha(pil_mask)
else:
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
preview_image = pil_image
previews.append(preview)
np_images.append(image)
np_masks.append(mask.unsqueeze(0))
previews.append(prepare_image_for_preview(preview_image, self.output_dir, self.filename_prefix))
np_image = pil2tensor(pil_image)
if pil_mask:
np_mask = np.array(pil_mask).astype(np.float32) / 255.0
np_mask = 1.0 - torch.from_numpy(np_mask)
else:
np_mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
np_images.append(np_image)
np_masks.append(np_mask.unsqueeze(0))
if output_mode:
result = (np_images, np_masks, True)
result = (np_images, np_masks, has_image)
else:
has_size_mismatch = False
if len(np_images) > 1:
for image in np_images[1:]:
if image.shape[1] != np_images[0].shape[1] or image.shape[2] != np_images[0].shape[2]:
for np_image in np_images[1:]:
if np_image.shape[1] != np_images[0].shape[1] or np_image.shape[2] != np_images[0].shape[2]:
has_size_mismatch = True
break
if has_size_mismatch:
raise Exception("To output as batch, images must have the same size. Use list output mode instead.")
result = ([torch.cat(np_images)], [torch.cat(np_masks)], True)
result = ([torch.cat(np_images)], [torch.cat(np_masks)], has_image)
return {"ui": {"images": previews}, "result": result}
class UtilLoadImageAsMaskFromUrl(UtilLoadImageFromUrl):
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"image": ("STRING", {"default": "", "multiline": True, "dynamicPrompts": False}),
"image": (
"STRING",
{
"default": "",
"placeholder": "Input image paths or URLS one per line. Eg:\nhttps://example.com/image.png\nfile:///path/to/local/image.jpg\ndata:image/png;base64,...",
"multiline": True,
"dynamicPrompts": False,
},
),
"channel": (["alpha", "red", "green", "blue"],),
},
"optional": {
@@ -225,22 +274,27 @@ class UtilLoadImageAsMaskFromUrl(UtilLoadImageFromUrl):
image = url
urls = image.strip().split("\n")
images, alphas = load_images_from_url(urls, True)
pil_images, pil_alphas = load_images_from_url(urls, True)
masks = []
masks: List[torch.Tensor] = []
for image, alpha in zip(images, alphas):
for img, alpha in zip(pil_images, pil_alphas):
if channel == "alpha":
mask = alpha
elif channel == "red":
mask = image.getchannel("R")
mask = img.getchannel("R")
elif channel == "green":
mask = image.getchannel("G")
mask = img.getchannel("G")
elif channel == "blue":
mask = image.getchannel("B")
mask = img.getchannel("B")
mask = np.array(mask).astype(np.float32) / 255.0
mask = 1.0 - torch.from_numpy(mask)
if mask:
mask = np.array(mask, dtype=np.float32) / 255.0
mask = torch.from_numpy(mask)
if channel == "alpha":
mask = 1.0 - mask
else:
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
masks.append(mask.unsqueeze(0))
@@ -257,7 +311,7 @@ class UtilLoadImageAsMaskFromUrl(UtilLoadImageFromUrl):
class UtilLoadJsonFromText:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"data": (
@@ -268,7 +322,7 @@ class UtilLoadJsonFromText:
}
RETURN_TYPES = ("JSON",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "load_json"
def load_json(self, data: str):
@@ -277,7 +331,7 @@ class UtilLoadJsonFromText:
class UtilLoadJsonFromUrl:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"url": ("STRING", {"default": ""}),
@@ -288,7 +342,7 @@ class UtilLoadJsonFromUrl:
}
RETURN_TYPES = ("JSON",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "load_json"
def load_json(self, url: str, print_to_console=False):
@@ -305,7 +359,7 @@ class UtilLoadJsonFromUrl:
class UtilGetObjectFromJson:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"json": ("JSON",),
@@ -314,7 +368,7 @@ class UtilGetObjectFromJson:
}
RETURN_TYPES = ("JSON",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "get_objects_from_json"
OUTPUT_NODE = True
@@ -324,7 +378,7 @@ class UtilGetObjectFromJson:
class UtilGetTextFromJson:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"json": ("JSON",),
@@ -333,17 +387,17 @@ class UtilGetTextFromJson:
}
RETURN_TYPES = ("STRING",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "get_string_from_json"
OUTPUT_NODE = True
def get_string_from_json(self, json: Dict, key: str):
return (get_dict_attribute(json, key, ""),)
return (str(get_dict_attribute(json, key, "")),)
class UtilGetFloatFromJson:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"json": ("JSON",),
@@ -352,17 +406,17 @@ class UtilGetFloatFromJson:
}
RETURN_TYPES = ("FLOAT",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "get_float_from_json"
OUTPUT_NODE = True
def get_float_from_json(self, json: Dict, key: str):
return (get_dict_attribute(json, key, 0.0),)
return (float(get_dict_attribute(json, key, 0.0)),)
class UtilGetIntFromJson:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"json": ("JSON",),
@@ -371,17 +425,17 @@ class UtilGetIntFromJson:
}
RETURN_TYPES = ("INT",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "get_int_from_json"
OUTPUT_NODE = True
def get_int_from_json(self, json: Dict, key: str):
return (get_dict_attribute(json, key, 0),)
return (int(get_dict_attribute(json, key, 0)),)
class UtilGetBoolFromJson:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"json": ("JSON",),
@@ -390,7 +444,7 @@ class UtilGetBoolFromJson:
}
RETURN_TYPES = ("BOOLEAN",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "get_bool_from_json"
OUTPUT_NODE = True
@@ -400,7 +454,7 @@ class UtilGetBoolFromJson:
class UtilRandomInt:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"min": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
@@ -409,11 +463,11 @@ class UtilRandomInt:
}
RETURN_TYPES = ("INT", "STRING")
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "random_int"
@classmethod
def IS_CHANGED(s, *args, **kwargs):
def IS_CHANGED(cls, *args, **kwargs):
return torch.rand(1).item()
def random_int(self, min: int, max: int):
@@ -423,7 +477,7 @@ class UtilRandomInt:
class UtilRandomFloat:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"min": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 0xFFFFFFFFFFFFFFFF}),
@@ -432,11 +486,11 @@ class UtilRandomFloat:
}
RETURN_TYPES = ("FLOAT", "STRING")
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "random_float"
@classmethod
def IS_CHANGED(s, *args, **kwargs):
def IS_CHANGED(cls, *args, **kwargs):
return torch.rand(1).item()
def random_float(self, min: float, max: float):
@@ -446,13 +500,13 @@ class UtilRandomFloat:
class UtilStringToInt:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {"string": ("STRING", {"default": "0"})},
}
RETURN_TYPES = ("INT",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "string_to_int"
def string_to_int(self, string: str):
@@ -461,7 +515,7 @@ class UtilStringToInt:
class UtilStringToNumber:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"string": ("STRING", {"default": "0"}),
@@ -470,7 +524,7 @@ class UtilStringToNumber:
}
RETURN_TYPES = ("INT", "FLOAT")
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "string_to_numbers"
def string_to_numbers(self, string: str, rounding):
@@ -486,7 +540,7 @@ class UtilStringToNumber:
class UtilNumberScaler:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"min": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 0xFFFFFFFFFFFFFFFF}),
@@ -498,7 +552,7 @@ class UtilNumberScaler:
}
RETURN_TYPES = ("FLOAT",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "scale_number"
def scale_number(self, min: float, max: float, scale_to_min: float, scale_to_max: float, value: float):
@@ -508,7 +562,7 @@ class UtilNumberScaler:
class UtilBooleanPrimitive:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"value": ("BOOLEAN", {"default": False}),
@@ -517,7 +571,7 @@ class UtilBooleanPrimitive:
}
RETURN_TYPES = ("BOOLEAN", "STRING")
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "boolean_primitive"
def boolean_primitive(self, value: bool, reverse: bool):
@@ -527,9 +581,58 @@ class UtilBooleanPrimitive:
return (value, str(value))
class UtilTextSwitchCase:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"switch_cases": (
"STRING",
{
"default": "",
"multiline": True,
"dynamicPrompts": False,
"placeholder": "case_1:output_1\ncase_2:output_2\nthat span multiple lines\ncase_3:output_3",
},
),
"condition": ("STRING", {"default": ""}),
"default_value": ("STRING", {"default": ""}),
"delimiter": ("STRING", {"default": ":"}),
}
}
RETURN_TYPES = ("STRING",)
CATEGORY = "ArtVenture/Utils"
FUNCTION = "text_switch_case"
def text_switch_case(self, switch_cases: str, condition: str, default_value: str, delimiter: str = ":"):
# Split into cases first
cases = switch_cases.split("\n")
current_case = None
current_output = []
for line in cases:
if delimiter in line:
# Process previous case if exists
if current_case is not None and condition == current_case:
return ("\n".join(current_output),)
# Start new case
current_case, output = line.split(delimiter, 1)
current_output = [output]
elif current_case is not None:
current_output.append(line)
# Check last case
if current_case is not None and condition == current_case:
return ("\n".join(current_output),)
return (default_value,)
class UtilImageMuxer:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"image_1": ("IMAGE",),
@@ -540,7 +643,7 @@ class UtilImageMuxer:
}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "image_muxer"
def image_muxer(self, image_1, image_2, input_selector, image_3=None, image_4=None):
@@ -550,7 +653,7 @@ class UtilImageMuxer:
class UtilSDXLAspectRatioSelector:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"aspect_ratio": (
@@ -576,7 +679,7 @@ class UtilSDXLAspectRatioSelector:
RETURN_TYPES = ("STRING", "INT", "INT")
RETURN_NAMES = ("ratio", "width", "height")
FUNCTION = "get_aspect_ratio"
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
def get_aspect_ratio(self, aspect_ratio):
width, height = 1024, 1024
@@ -613,7 +716,7 @@ class UtilSDXLAspectRatioSelector:
class UtilAspectRatioSelector(UtilSDXLAspectRatioSelector):
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"aspect_ratio": (
@@ -643,7 +746,7 @@ class UtilAspectRatioSelector(UtilSDXLAspectRatioSelector):
class UtilDependenciesEdit:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"dependencies": ("DEPENDENCIES",),
@@ -669,7 +772,7 @@ class UtilDependenciesEdit:
}
RETURN_TYPES = ("DEPENDENCIES",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "edit_dependencies"
def edit_dependencies(
@@ -732,7 +835,7 @@ class UtilImageScaleDown:
crop_methods = ["disabled", "center"]
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
@@ -744,12 +847,12 @@ class UtilImageScaleDown:
"INT",
{"default": 512, "min": 1, "max": MAX_RESOLUTION, "step": 1},
),
"crop": (s.crop_methods,),
"crop": (cls.crop_methods,),
}
}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "image_scale_down"
def image_scale_down(self, images, width, height, crop):
@@ -771,7 +874,7 @@ class UtilImageScaleDown:
results = []
for image in s:
img = tensor2pil(image).convert("RGB")
img = img.resize((width, height), Image.LANCZOS)
img = img.resize((width, height), Image.Resampling.LANCZOS)
results.append(pil2tensor(img))
return (torch.cat(results, dim=0),)
@@ -779,7 +882,7 @@ class UtilImageScaleDown:
class UtilImageScaleDownBy(UtilImageScaleDown):
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
@@ -791,7 +894,7 @@ class UtilImageScaleDownBy(UtilImageScaleDown):
}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "image_scale_down_by"
def image_scale_down_by(self, images, scale_by):
@@ -804,7 +907,7 @@ class UtilImageScaleDownBy(UtilImageScaleDown):
class UtilImageScaleDownToSize(UtilImageScaleDownBy):
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
@@ -814,7 +917,7 @@ class UtilImageScaleDownToSize(UtilImageScaleDownBy):
}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "image_scale_down_to_size"
def image_scale_down_to_size(self, images, size, mode):
@@ -830,9 +933,13 @@ class UtilImageScaleDownToSize(UtilImageScaleDownBy):
return self.image_scale_down_by(images, scale_by)
class UtilImageScaleToTotalPixels(UtilImageScaleDownBy, ImageUpscaleWithModel):
class UtilImageScaleToTotalPixels(UtilImageScaleDownBy):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.upscale_model_node = ImageUpscaleWithModel()
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
@@ -844,7 +951,7 @@ class UtilImageScaleToTotalPixels(UtilImageScaleDownBy, ImageUpscaleWithModel):
}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "image_scale_down_to_total_pixels"
def image_scale_up_by(self, images: torch.Tensor, scale_by, upscale_model_opt):
@@ -857,10 +964,10 @@ class UtilImageScaleToTotalPixels(UtilImageScaleDownBy, ImageUpscaleWithModel):
s = s.movedim(1, -1)
return (s,)
else:
s = self.upscale(upscale_model_opt, images)[0]
s = self.upscale_model_node.execute(upscale_model_opt, images)[0]
return self.image_scale_down(s, width, height, "center")
def image_scale_down_to_total_pixels(self, images, megapixels, upscale_model_opt=None):
def image_scale_down_to_total_pixels(self, images, megapixels, *args, upscale_model_opt=None, **kwargs):
width = images.shape[2]
height = images.shape[1]
scale_by = np.sqrt((megapixels * 1024 * 1024) / (width * height))
@@ -873,7 +980,7 @@ class UtilImageScaleToTotalPixels(UtilImageScaleDownBy, ImageUpscaleWithModel):
class UtilImageAlphaComposite:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"image_1": ("IMAGE",),
@@ -882,7 +989,7 @@ class UtilImageAlphaComposite:
}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "image_alpha_composite"
def image_alpha_composite(self, image_1: torch.Tensor, image_2: torch.Tensor):
@@ -905,7 +1012,7 @@ class UtilImageAlphaComposite:
class UtilImageGaussianBlur:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
@@ -914,7 +1021,7 @@ class UtilImageGaussianBlur:
}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "image_gaussian_blur"
def image_gaussian_blur(self, images, radius):
@@ -929,7 +1036,7 @@ class UtilImageGaussianBlur:
class UtilImageExtractChannel:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
@@ -939,7 +1046,7 @@ class UtilImageExtractChannel:
RETURN_TYPES = ("MASK",)
RETURN_NAMES = ("channel_data",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "image_extract_alpha"
def image_extract_alpha(self, images: torch.Tensor, channel):
@@ -959,7 +1066,7 @@ class UtilImageExtractChannel:
class UtilImageApplyChannel:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
@@ -969,7 +1076,7 @@ class UtilImageApplyChannel:
}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "image_apply_channel"
def image_apply_channel(self, images: torch.Tensor, channel_data: torch.Tensor, channel):
@@ -1011,20 +1118,20 @@ class UtillQRCodeGenerator:
RETURN_TYPES = ("IMAGE",)
FUNCTION = "create_qr_code"
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
def create_qr_code(self, text, size, qr_version, error_correction, box_size, border):
ensure_package("qrcode", "qrcode[pil]")
ensure_package("qrcode", install_package_name="qrcode[pil]")
import qrcode
if error_correction == "L":
error_level = qrcode.constants.ERROR_CORRECT_L
error_level = qrcode.ERROR_CORRECT_L
elif error_correction == "M":
error_level = qrcode.constants.ERROR_CORRECT_M
error_level = qrcode.ERROR_CORRECT_M
elif error_correction == "Q":
error_level = qrcode.constants.ERROR_CORRECT_Q
error_level = qrcode.ERROR_CORRECT_Q
else:
error_level = qrcode.constants.ERROR_CORRECT_H
error_level = qrcode.ERROR_CORRECT_H
qr = qrcode.QRCode(version=qr_version, error_correction=error_level, box_size=box_size, border=border)
qr.add_data(text)
@@ -1037,7 +1144,7 @@ class UtillQRCodeGenerator:
class UtilRepeatImages:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
@@ -1046,7 +1153,7 @@ class UtilRepeatImages:
}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "rebatch"
def rebatch(self, images: torch.Tensor, amount):
@@ -1055,7 +1162,7 @@ class UtilRepeatImages:
class UtilSeedSelector:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"mode": ("BOOLEAN", {"default": True, "label_on": "random", "label_off": "fixed"}),
@@ -1069,7 +1176,7 @@ class UtilSeedSelector:
RETURN_TYPES = ("INT",)
RETURN_NAMES = ("seed",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "get_seed"
def get_seed(self, mode, seed, fixed_seed):
@@ -1078,7 +1185,7 @@ class UtilSeedSelector:
class UtilCheckpointSelector:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"ckpt_name": (folder_paths.get_filename_list("checkpoints"),),
@@ -1087,11 +1194,11 @@ class UtilCheckpointSelector:
RETURN_TYPES = (folder_paths.get_filename_list("checkpoints"), "STRING")
RETURN_NAMES = ("ckpt_name", "ckpt_name_str")
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "get_ckpt_name"
@classmethod
def IS_CHANGED(s, *args, **kwargs):
def IS_CHANGED(cls, *args, **kwargs):
return torch.rand(1).item()
def get_ckpt_name(self, ckpt_name):
@@ -1100,7 +1207,7 @@ class UtilCheckpointSelector:
class UtilModelMerge:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"model1": ("MODEL",),
@@ -1110,7 +1217,7 @@ class UtilModelMerge:
}
RETURN_TYPES = ("MODEL",)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "merge_models"
def merge_models(self, model1, model2, ratio=1.0):
@@ -1132,7 +1239,7 @@ class UtilModelMerge:
class UtilTextRandomMultiline:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"multiline": True, "dynamicPrompts": False}),
@@ -1144,7 +1251,7 @@ class UtilTextRandomMultiline:
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("lines",)
OUTPUT_IS_LIST = (True,)
CATEGORY = "Art Venture/Utils"
CATEGORY = "ArtVenture/Utils"
FUNCTION = "random_multiline"
def random_multiline(self, text: str, amount=1, seed=0):
@@ -1190,6 +1297,7 @@ NODE_CLASS_MAPPINGS = {
"NumberScaler": UtilNumberScaler,
"MergeModels": UtilModelMerge,
"TextRandomMultiline": UtilTextRandomMultiline,
"TextSwitchCase": UtilTextSwitchCase,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"LoadImageFromUrl": "Load Image From URL",
@@ -1225,4 +1333,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"NumberScaler": "Number Scaler",
"MergeModels": "Merge Models",
"TextRandomMultiline": "Text Random Multiline",
"TextSwitchCase": "Text Switch Case",
}
+42 -29
View File
@@ -5,9 +5,10 @@ import torch
import base64
import numpy as np
import importlib
import importlib.metadata
import subprocess
import pkg_resources
from pkg_resources import parse_version
from packaging import version
from packaging.specifiers import SpecifierSet
from PIL import Image
from .logger import logger
@@ -21,23 +22,40 @@ class AnyType(str):
any_type = AnyType("*")
def ensure_package(package, version=None, install_package_name=None):
def ensure_package(package, required_version=None, install_package_name=None):
# Try to import the package
try:
module = importlib.import_module(package)
except ImportError:
logger.info(f"Package {package} is not installed. Installing now...")
install_command = _construct_pip_command(install_package_name or package, version)
install_command = _construct_pip_command(install_package_name or package, required_version)
subprocess.check_call(install_command)
else:
# If a specific version is required, check the version
if version:
installed_version = pkg_resources.get_distribution(package).version
if parse_version(installed_version) < parse_version(version):
logger.info(
f"Package {package} is outdated (installed: {installed_version}, required: {version}). Upgrading now..."
)
install_command = _construct_pip_command(install_package_name or package, version)
if required_version:
try:
installed_version = importlib.metadata.version(package)
# Parse version specifier (e.g., ">=1.1.1", "==1.1.1", "<=1.1.1")
if any(op in required_version for op in ['>=', '<=', '==', '!=', '>', '<', '~=']):
spec = SpecifierSet(required_version)
if installed_version not in spec:
logger.info(
f"Package {package} version constraint not satisfied (installed: {installed_version}, required: {required_version}). Installing now..."
)
install_command = _construct_pip_command(install_package_name or package, required_version)
subprocess.check_call(install_command)
else:
# Fallback to simple version comparison for backwards compatibility
if version.parse(installed_version) < version.parse(required_version):
logger.info(
f"Package {package} is outdated (installed: {installed_version}, required: {required_version}). Upgrading now..."
)
install_command = _construct_pip_command(install_package_name or package, required_version)
subprocess.check_call(install_command)
except importlib.metadata.PackageNotFoundError:
logger.info(f"Package {package} version information not found. Installing required version {required_version}...")
install_command = _construct_pip_command(install_package_name or package, required_version)
subprocess.check_call(install_command)
@@ -54,29 +72,24 @@ def _construct_pip_command(package_name, version=None):
return pip_install + [package_name]
# modified from https://stackoverflow.com/questions/22058048/hashing-a-file-in-python
def calculate_file_hash(filename: str, hash_every_n: int = 1):
import hashlib
h = hashlib.sha256()
b = bytearray(10 * 1024 * 1024) # read 10 megabytes at a time
mv = memoryview(b)
with open(filename, "rb", buffering=0) as f:
i = 0
# don't hash entire file, only portions of it if requested
while n := f.readinto(mv):
if i % hash_every_n == 0:
h.update(mv[:n])
i += 1
return h.hexdigest()
def get_dict_attribute(dict_inst: dict, name_string: str, default=None):
nested_keys = name_string.split(".")
value = dict_inst
for key in nested_keys:
value = value.get(key, None)
# Handle array indexing
if key.startswith("[") and key.endswith("]"):
try:
index = int(key[1:-1])
if not isinstance(value, (list, tuple)) or index >= len(value):
return default
value = value[index]
except (ValueError, TypeError):
return default
else:
if not isinstance(value, dict):
return default
value = value.get(key, None)
if value is None:
return default
+3 -3
View File
@@ -52,10 +52,10 @@ try:
@classmethod
def INPUT_TYPES(s):
inputs = LoadVideoPath.INPUT_TYPES()
inputs["required"]["video"] = ("STRING", {"default": "", "multiline": True, "dynamicPrompts": False})
inputs["required"]["video"] = ("STRING", {"default": ""})
return inputs
CATEGORY = "Art Venture/Loaders"
CATEGORY = "ArtVenture/Loaders"
FUNCTION = "load"
RETURN_TYPES = ("IMAGE", "INT", "BOOLEAN")
RETURN_NAMES = ("frames", "frame_count", "has_video")
@@ -164,7 +164,7 @@ try:
from urllib.parse import parse_qs
qs_idx = url.find("?")
qs = parse_qs(url[qs_idx + 1:])
qs = parse_qs(url[qs_idx + 1 :])
filename = qs.get("name", qs.get("filename", None))
if filename is None:
raise Exception(f"Invalid url: {url}")
+18 -6
View File
@@ -1,15 +1,27 @@
[project]
name = "comfyui-art-venture"
description = "Nodes: ImagesConcat, LoadImageFromUrl, AV_UploadImage"
version = "1.0.0"
description = "A comprehensive set of custom nodes for ComfyUI, focusing on utilities for image processing, JSON manipulation, model operations and working with object via URLs"
version = "1.1.7"
license = "LICENSE"
dependencies = ["timm==0.6.13", "transformers", "fairscale", "pycocoevalcap", "opencv-python", "qrcode[pil]", "pytorch_lightning", "kornia", "pydantic", "segment_anything", "omegaconf", "boto3>=1.34.101"]
dependencies = [
"timm==0.6.13",
"transformers",
"fairscale",
"pycocoevalcap",
"opencv-python",
"qrcode[pil]",
"pytorch_lightning",
"kornia",
"pydantic",
"segment_anything",
"boto3>=1.34.101",
]
[project.urls]
Repository = "https://github.com/sipherxyz/comfyui-art-venture"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = ""
DisplayName = "comfyui-art-venture"
Icon = ""
PublisherId = "protogaia"
DisplayName = "ComfyUI ArtVenture"
Icon = "https://cdn.protogaia.com/assets/gaia.png"
+1 -1
View File
@@ -9,4 +9,4 @@ kornia
pydantic
segment_anything
omegaconf
boto3>=1.34.101
boto3
+52
View File
@@ -0,0 +1,52 @@
import { app } from '../../scripts/app.js';
import { ComfyWidgets } from '../../scripts/widgets.js';
import {
addKVState,
chainCallback,
addWidgetChangeCallback,
} from './utils.js';
function addTextSwitchCaseWidget(nodeType) {
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
const dataWidget = this.widgets.find((w) => w.name === 'switch_cases');
const delimiterWidget = this.widgets.find((w) => w.name === 'delimiter');
this.widgets = this.widgets.filter((w) => w.name !== 'condition');
let conditionCombo = null;
const updateConditionCombo = () => {
if (!delimiterWidget.value) return;
const cases = (dataWidget.value ?? '')
.split('\n')
.filter((line) => line.includes(delimiterWidget.value))
.map((line) => line.split(delimiterWidget.value)[0]);
if (!conditionCombo) {
conditionCombo = ComfyWidgets['COMBO'](this, 'condition', [
['__default__', ...(cases ?? [])],
]).widget;
} else {
conditionCombo.options.values = ['__default__', ...cases];
}
};
updateConditionCombo();
dataWidget.inputEl.addEventListener('input', updateConditionCombo);
addWidgetChangeCallback(delimiterWidget, updateConditionCombo);
});
}
app.registerExtension({
name: 'ArtVenture.TextSwitchCase',
async beforeRegisterNodeDef(nodeType, nodeData) {
if (!nodeData) return;
if (nodeData.name !== 'TextSwitchCase') {
return;
}
addKVState(nodeType);
addTextSwitchCaseWidget(nodeType);
},
});
+212 -590
View File
@@ -1,616 +1,244 @@
import { app, ANIM_PREVIEW_WIDGET } from '../../../scripts/app.js';
import { api } from '../../../scripts/api.js';
import { $el } from '../../../scripts/ui.js';
import { createImageHost } from '../../../scripts/ui/imagePreview.js';
import { app } from '../../scripts/app.js';
import { api } from '../../scripts/api.js';
import { $el } from '../../scripts/ui.js';
import { addWidget, DOMWidgetImpl } from '../../scripts/domWidget.js';
import { ComfyWidgets } from '../../scripts/widgets.js'
const style = `
.comfy-img-preview video {
object-fit: contain;
width: var(--comfy-img-preview-width);
height: var(--comfy-img-preview-height);
}
`;
import { chainCallback, addKVState, addWidgetChangeCallback } from './utils.js';
const URL_REGEX = /^((blob:)?https?:\/\/|\/view\?|\/api\/view\?|data:image\/)/
const supportedNodes = ['LoadImageFromUrl', 'LoadImageAsMaskFromUrl'];
const supportedNodes = ['LoadImageFromUrl', 'LoadImageAsMaskFromUrl', 'LoadVideoFromUrl'];
const formatUrl = (url) => {
if (!url) return ""
function chainCallback(object, property, callback) {
if (object == undefined) {
//This should not happen.
console.error('Tried to add callback to non-existant object');
return;
}
if (property in object) {
const callback_orig = object[property];
object[property] = function () {
const r = callback_orig.apply(this, arguments);
callback.apply(this, arguments);
return r;
};
} else {
object[property] = callback;
}
}
if (url.startsWith("http://") || url.startsWith("https://") || url.startsWith("blob:")) return url
if (url.startsWith("/view") || url.startsWith("/api/view")) return url
function injectHidden(widget) {
widget.computeSize = (target_width) => {
if (widget.hidden) {
return [0, -4];
}
return [target_width, 20];
};
widget._type = widget.type;
Object.defineProperty(widget, 'type', {
set: function (value) {
widget._type = value;
},
get: function () {
if (widget.hidden) {
return 'hidden';
}
return widget._type;
},
});
}
function addKVState(nodeType) {
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
chainCallback(this, 'onConfigure', function (info) {
if (!this.widgets) {
//Node has no widgets, there is nothing to restore
return;
}
if (typeof info.widgets_values != 'object') {
//widgets_values is in some unknown inactionable format
return;
}
let widgetDict = info.widgets_values;
if (widgetDict.length == undefined) {
for (let w of this.widgets) {
if (w.name in widgetDict) {
w.value = widgetDict[w.name];
} else {
//attempt to restore default value
let inputs = LiteGraph.getNodeType(this.type).nodeData.input;
let initialValue = null;
if (inputs?.required?.hasOwnProperty(w.name)) {
if (inputs.required[w.name][1]?.hasOwnProperty('default')) {
initialValue = inputs.required[w.name][1].default;
} else if (inputs.required[w.name][0].length) {
initialValue = inputs.required[w.name][0][0];
}
} else if (inputs?.optional?.hasOwnProperty(w.name)) {
if (inputs.optional[w.name][1]?.hasOwnProperty('default')) {
initialValue = inputs.optional[w.name][1].default;
} else if (inputs.optional[w.name][0].length) {
initialValue = inputs.optional[w.name][0][0];
}
}
if (initialValue) {
w.value = initialValue;
}
}
}
}
});
chainCallback(this, 'onSerialize', function (info) {
info.widgets_values = {};
if (!this.widgets) {
//object has no widgets, there is nothing to store
return;
}
for (let w of this.widgets) {
info.widgets_values[w.name] = w.value;
}
});
});
}
function migrateWidget(nodeType, oldWidgetName, newWidgetName) {
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
if (!this.widgets) return;
const oldIndex = this.widgets.findIndex((w) => w.name === oldWidgetName);
if (oldIndex > -1) {
this.widgets.splice(oldIndex, 1);
}
chainCallback(this, 'onConfigure', function (info) {
if (typeof info.widgets_values != 'object') return;
const newWidget = this.widgets.find((w) => w.name === newWidgetName);
if (newWidget && info.widgets_values[oldWidgetName]) {
newWidget.value = info.widgets_values[oldWidgetName];
}
});
});
}
function formatImageUrl(params) {
if (typeof params === "string") {
if (URL_REGEX.test(params)) return params;
const folder_separator = params.lastIndexOf("/");
let subfolder = "";
if (folder_separator > -1) {
subfolder = params.substring(0, folder_separator);
params = params.substring(folder_separator + 1);
}
let type = "input";
if (params.indexOf(" [") > -1) {
type = params.split(" [")[1].split("]")[0];
params = params.split(" [")[0];
}
params = {
filename: params,
type: type,
subfolder: subfolder,
};
let type = "input"
if (url.endsWith(']')) {
const openBracketIndex = url.lastIndexOf('[')
type = url.slice(openBracketIndex + 1, url.length - 1).trim()
url = url.slice(0, openBracketIndex).trim()
}
if (params.url) {
return params.url;
}
const parts = url.split('/')
const filename = parts.pop()
const subfolder = parts.join('/')
params = { ...params };
const params = [
'filename=' + encodeURIComponent(filename),
'type=' + type,
'subfolder=' + subfolder,
app.getRandParam().substring(1)
].join('&')
if (!params.filename && params.name) {
params.filename = params.name;
delete params.name;
}
return api.apiURL("/view?" + new URLSearchParams(params).toString() + app.getPreviewFormatParam());
return api.apiURL(`/view?${params}`)
}
async function uploadFile(file) {
//TODO: Add uploaded file to cache with Cache.put()?
try {
// Wrap file in formdata so it includes filename
const body = new FormData();
const i = file.webkitRelativePath.lastIndexOf('/');
const subfolder = file.webkitRelativePath.slice(0, i + 1);
const new_file = new File([file], file.name, {
type: file.type,
lastModified: file.lastModified,
});
body.append('image', new_file);
if (i > 0) {
body.append('subfolder', subfolder);
}
const resp = await api.fetchApi('/upload/image', {
method: 'POST',
body,
});
// copied from ComfyUI_frontend/src/composables/widgets/useStringWidget.ts
// remove the Object.defineProperty(widget, 'value') part
function addUrlWidget(node, name, options) {
const inputEl = document.createElement('textarea')
inputEl.className = 'comfy-multiline-input'
inputEl.value = options.default
inputEl.placeholder = options.placeholder || name
inputEl.spellcheck = false
if (resp.status === 200 || resp.status === 201) {
return resp.json();
} else {
alert(`Upload failed: ${resp.statusText}`);
}
} catch (error) {
alert(`Upload failed: ${error}`);
}
}
function addVideoCustomSize(nodeType, nodeData, widgetName) {
//Add the extra size widgets now
//This takes some finagling as widget order is defined by key order
const newWidgets = {};
for (let key in nodeData.input.required) {
newWidgets[key] = nodeData.input.required[key];
if (key == widgetName) {
newWidgets[key][0] = newWidgets[key][0].concat(['Custom Width', 'Custom Height', 'Custom']);
newWidgets['custom_width'] = ['INT', { default: 512, min: 8, step: 8 }];
newWidgets['custom_height'] = ['INT', { default: 512, min: 8, step: 8 }];
}
}
nodeData.input.required = newWidgets;
//Add a callback which sets up the actual logic once the node is created
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
const node = this;
const sizeOptionWidget = node.widgets.find((w) => w.name === widgetName);
const widthWidget = node.widgets.find((w) => w.name === 'custom_width');
const heightWidget = node.widgets.find((w) => w.name === 'custom_height');
injectHidden(widthWidget);
widthWidget.options.serialize = false;
injectHidden(heightWidget);
heightWidget.options.serialize = false;
sizeOptionWidget._value = sizeOptionWidget.value;
Object.defineProperty(sizeOptionWidget, 'value', {
set: function (value) {
//TODO: Only modify hidden/reset size when a change occurs
if (value == 'Custom Width') {
widthWidget.hidden = false;
heightWidget.hidden = true;
} else if (value == 'Custom Height') {
widthWidget.hidden = true;
heightWidget.hidden = false;
} else if (value == 'Custom') {
widthWidget.hidden = false;
heightWidget.hidden = false;
} else {
widthWidget.hidden = true;
heightWidget.hidden = true;
}
node.setSize([node.size[0], node.computeSize([node.size[0], node.size[1]])[1]]);
this._value = value;
const widget = new DOMWidgetImpl({
node,
name,
type: 'customtext',
element: inputEl,
options: {
hideOnZoom: true,
getValue() {
return inputEl.value
},
get: function () {
return this._value;
},
});
//Ensure proper visibility/size state for initial value
sizeOptionWidget.value = sizeOptionWidget._value;
setValue(v) {
inputEl.value = v
}
}
})
addWidget(node, widget)
sizeOptionWidget.serializeValue = function () {
if (this.value == 'Custom Width') {
return widthWidget.value + 'x?';
} else if (this.value == 'Custom Height') {
return '?x' + heightWidget;
} else if (this.value == 'Custom') {
return widthWidget.value + 'x' + heightWidget.value;
widget.inputEl = inputEl
widget.options.minNodeSize = [400, 200]
inputEl.addEventListener('input', () => {
widget.value = inputEl.value
widget.callback?.(inputEl.value, true)
})
// Allow middle mouse button panning
inputEl.addEventListener('pointerdown', (event) => {
if (event.button === 1) {
app.canvas.processMouseDown(event)
}
})
inputEl.addEventListener('pointermove', (event) => {
if ((event.buttons & 4) === 4) {
app.canvas.processMouseMove(event)
}
})
inputEl.addEventListener('pointerup', (event) => {
if (event.button === 1) {
app.canvas.processMouseUp(event)
}
})
/** Timer reference. `null` when the timer completes. */
let ignoreEventsTimer = null
/** Total number of events ignored since the timer started. */
let ignoredEvents = 0
// Pass wheel events to the canvas when appropriate
inputEl.addEventListener('wheel', (event) => {
if (!Object.is(event.deltaX, -0)) return
// If the textarea has focus, require more effort to activate pass-through
const multiplier = document.activeElement === inputEl ? 2 : 1
const maxScrollHeight = inputEl.scrollHeight - inputEl.clientHeight
if (
(event.deltaY < 0 && inputEl.scrollTop === 0) ||
(event.deltaY > 0 && inputEl.scrollTop === maxScrollHeight)
) {
// Attempting to scroll past the end of the textarea
if (!ignoreEventsTimer || ignoredEvents > 25 * multiplier) {
app.canvas.processMouseWheel(event)
} else {
return this.value;
ignoredEvents++
}
};
});
} else if (event.deltaY !== 0) {
// Start timer whenever a successful scroll occurs
ignoredEvents = 0
if (ignoreEventsTimer) clearTimeout(ignoreEventsTimer)
ignoreEventsTimer = setTimeout(() => {
ignoreEventsTimer = null
}, 800 * multiplier)
}
})
return widget
}
function addUploadWidget(nodeType, widgetName, type) {
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
this.images = [];
const pathWidget = this.widgets.find((w) => w.name === widgetName);
const supportMultiple = pathWidget.type === 'customtext';
function addImageUploadWidget(nodeType, nodeData, imageInputName) {
const { input } = nodeData ?? {}
const required = input?.required
if (!required) return
if (pathWidget.element) {
pathWidget.options.getMinHeight = () => 50;
pathWidget.options.getMaxHeight = () => 150;
}
const imageOptions = required.image
delete required.image
const fileInput = document.createElement('input');
chainCallback(this, 'onRemoved', () => {
fileInput?.remove();
});
chainCallback(nodeType.prototype, "onNodeCreated", function () {
this.previewMediaType = 'image'
if (type === 'image') {
Object.assign(fileInput, {
type: 'file',
accept: 'image/png,image/jpeg,image/webp',
style: 'display: none',
multiple: supportMultiple,
onchange: async () => {
if (!fileInput.files.length) {
return;
}
const urlWidget = addUrlWidget(this, imageInputName, imageOptions[1])
// move urlWidget to the first position
const widgets = this.widgets.filter(w => w !== urlWidget)
this.widgets = [urlWidget, ...widgets]
let successes = [];
for (const file of fileInput.files) {
const params = await uploadFile(file);
ComfyWidgets.IMAGEUPLOAD(
this,
'upload',
["IMAGEUPLOAD", { "image_upload": true, imageInputName }],
)
if (!!params) {
successes.push(params);
} else {
// Upload failed, but some prior uploads may have succeeded
// Stop future uploads to prevent cascading failures
// and only add to list if an upload has succeeded
if (successes.length) {
break;
} else {
return;
}
}
}
pathWidget.value = successes.map(formatImageUrl).join('\n');
fileInput.value = '';
},
});
} else if (type === 'video') {
Object.assign(fileInput, {
type: 'file',
accept: 'video/webm,video/mp4,video/mkv,image/gif,image/webp',
style: 'display: none',
multiple: supportMultiple,
onchange: async () => {
if (!fileInput.files.length) {
return;
}
let successes = [];
for (const file of fileInput.files) {
const params = await uploadFile(file);
if (!!params) {
successes.push(params);
} else {
// Upload failed, but some prior uploads may have succeeded
// Stop future uploads to prevent cascading failures
// and only add to list if an upload has succeeded
if (successes.length) {
break;
} else {
return;
}
}
}
pathWidget.value = successes.map(formatImageUrl).join('\n');
fileInput.value = '';
},
});
} else {
throw new Error(`Unknown upload type ${type}`);
}
document.body.append(fileInput);
let uploadWidget = this.addWidget('button', 'choose ' + type + ' to upload', 'image', () => {
//clear the active click event
app.canvas.node_widget = null;
fileInput.click();
});
uploadWidget.serialize = false;
// Add handler to check if an image is being dragged over our node
this.onDragOver = function (e) {
if (e.dataTransfer && e.dataTransfer.items) {
const image = [...e.dataTransfer.items].find((f) => f.kind === 'file');
return !!image;
}
return false;
};
// On drop upload files
this.onDragDrop = async function (e) {
let successes = [];
const files = e.dataTransfer.files
.filter((file) => file.type.startsWith('image/'))
.slice(0, supportMultiple ? undefined : 1);
for (const file of files) {
const params = await uploadFile(file);
if (!!params) {
successes.push(params);
}
pathWidget.value = (supportMultiple ? this.images : [])
.concat(...successes.map(formatImageUrl))
.join('\n');
}
return successes.length > 0;
};
this.pasteFile = function (file) {
if (file.type.startsWith('image/')) {
uploadFile(file).then((res) => {
pathWidget.value = (supportMultiple ? this.images : [])
.concat(formatImageUrl(res))
.join('\n');
});
return true;
}
return false;
};
});
}
function patchValueSetter(nodeType, widgetName) {
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
const pathWidget = this.widgets.find((w) => w.name === widgetName);
pathWidget._value = pathWidget.value;
let editing = false;
const setter = (value) => {
if (typeof value !== 'string') value = formatImageUrl(value);
pathWidget._value = value;
this.images = (value ?? '').split('\n').filter(Boolean);
if (pathWidget.type === 'customtext' && !editing) {
pathWidget.inputEl.value = value;
}
delete app.nodeOutputs[this.id]
};
Object.defineProperty(pathWidget, 'value', {
set: setter,
get: () => pathWidget._value,
});
if (pathWidget.type === 'customtext') {
pathWidget.inputEl.addEventListener('focus', (e) => {
editing = true;
});
pathWidget.inputEl.addEventListener('blur', (e) => {
editing = false;
});
pathWidget.inputEl.addEventListener('keyup', (e) => {
setter(e.target.value);
const safeLoadImageFromUrl = (url) => {
return new Promise((resolve, reject) => {
const img = new Image();
img.onload = () => resolve(img);
img.onerror = () => reject(null);
img.src = url;
});
}
pathWidget.callback = setter;
pathWidget.value = pathWidget._value;
});
}
let initialImgs = undefined;
let isUrlImageSet = false
function addVideoPreview(nodeType, widgetName) {
const createVideoNode = (url) => {
return new Promise((cb) => {
const videoEl = document.createElement('video');
Object.defineProperty(videoEl, 'naturalWidth', {
get: () => {
return videoEl.videoWidth;
},
});
Object.defineProperty(videoEl, 'naturalHeight', {
get: () => {
return videoEl.videoHeight;
},
});
videoEl.addEventListener('loadedmetadata', () => {
videoEl.controls = false;
videoEl.loop = true;
videoEl.muted = true;
cb(videoEl);
});
videoEl.addEventListener('error', () => {
cb();
});
videoEl.src = url;
});
};
const createImageNode = (url) => {
return new Promise((cb) => {
const imgEl = document.createElement('img');
imgEl.onload = () => {
cb(imgEl);
};
imgEl.addEventListener('error', () => {
cb();
});
imgEl.src = url;
});
};
nodeType.prototype.onDrawBackground = function (ctx) {
if (this.flags.collapsed) return;
let imageURLs = (this.images ?? []).map(formatImageUrl);
let imagesChanged = false;
if (JSON.stringify(this.displayingImages) !== JSON.stringify(imageURLs)) {
this.displayingImages = imageURLs;
imagesChanged = true;
}
if (!imagesChanged) return;
if (!imageURLs.length) {
this.imgs = null;
this.animatedImages = false;
return;
}
const promises = imageURLs.map((url) => {
if (/^(\/api)?\/view/.test(url)) {
url = window.location.origin + url;
}
let ext = '';
if (url.startsWith('data:')) {
const blob = dataUriToBlob(url);
ext = blob.type.split('/').pop();
url = URL.createObjectURL(blob);
} else {
const u = new URL(url);
const filename =
u.searchParams.get('filename') ||
u.searchParams.get('name') ||
u.pathname.split('/').pop();
ext = filename.split('.').pop();
}
const format = ['gif', 'webp', 'avif'].includes(ext) ? 'image' : 'video';
if (format === 'video') {
return createVideoNode(url);
} else {
return createImageNode(url);
}
});
Promise.all(promises)
.then((imgs) => {
this.imgs = imgs.filter(Boolean);
})
.then(() => {
if (!this.imgs.length) return;
this.animatedImages = true;
const widgetIdx = this.widgets?.findIndex((w) => w.name === ANIM_PREVIEW_WIDGET);
// Instead of using the canvas we'll use a IMG
if (widgetIdx > -1) {
// Replace content
const widget = this.widgets[widgetIdx];
widget.options.host.updateImages(this.imgs);
} else {
const host = createImageHost(this);
this.setSizeForImage(true);
const widget = this.addDOMWidget(ANIM_PREVIEW_WIDGET, 'img', host.el, {
host,
getHeight: host.getHeight,
onDraw: host.onDraw,
hideOnZoom: false,
});
widget.serializeValue = () => ({
height: host.el.clientHeight,
});
// widget.computeSize = (w) => ([w, 220]);
widget.options.host.updateImages(this.imgs);
}
this.imgs.forEach((img) => {
if (img instanceof HTMLVideoElement) {
img.muted = true;
img.autoplay = true;
img.play();
}
});
});
};
patchValueSetter(nodeType, widgetName);
chainCallback(nodeType.prototype, 'onExecuted', function (message) {
if (message?.videos) {
this.images = message?.videos.map(formatImageUrl);
}
});
}
function addImagePreview(nodeType, widgetName) {
const onDrawBackground = nodeType.prototype.onDrawBackground;
nodeType.prototype.onDrawBackground = function (ctx) {
if (this.flags.collapsed) return;
let imageURLs = (this.images ?? []).map(formatImageUrl);
let imagesChanged = false;
if (JSON.stringify(this.displayingImages) !== JSON.stringify(imageURLs)) {
this.displayingImages = imageURLs;
imagesChanged = true;
}
if (imagesChanged) {
const setImagesFromUrl = (value = "") => {
this.imageIndex = null;
if (imageURLs.length > 0) {
Promise.all(
imageURLs.map((src) => {
return new Promise((r) => {
const img = new Image();
img.onload = () => r(img);
img.onerror = () => r(null);
img.src = src;
});
}),
).then((imgs) => {
this.imgs = imgs.filter(Boolean);
this.setSizeForImage?.();
app.graph.setDirtyCanvas(true);
});
} else {
this.imgs = null;
const urls = value.split("\n").filter(Boolean).map(formatUrl);
if (!urls.length) {
this.imgs = undefined;
this.widgets = this.widgets.filter((w) => w.name !== "$$canvas-image-preview");
isUrlImageSet = true;
return
}
return Promise.all(
urls.map(safeLoadImageFromUrl)
).then((imgs) => {
initialImgs = imgs.filter(Boolean);
this.imgs = initialImgs.length > 0 ? initialImgs : undefined;
if (!this.imgs) {
this.widgets = this.widgets.filter((w) => w.name !== "$$canvas-image-preview");
}
app.graph.setDirtyCanvas(true);
// cancel any img change in the next 2 seconds
// to prevent `ComfyUI `overwriting the image
setTimeout(() => {
isUrlImageSet = true;
}, 2000);
return initialImgs;
})
}
onDrawBackground?.call(this, ctx);
};
addWidgetChangeCallback(this, {
name: "imgs",
shouldChange: (value) => {
if (isUrlImageSet) return true;
return value === initialImgs;
}
})
patchValueSetter(nodeType, widgetName);
urlWidget.callback = (value) => {
if (!value) { // from upload
value = urlWidget.value.split("\n").filter(Boolean).map(formatUrl).join("\n")
}
if (Array.isArray(value)) {
value = value.map(formatUrl).join("\n")
}
if (value !== urlWidget.options.getValue()) {
urlWidget.options.setValue(value)
}
setImagesFromUrl(value)
}
this.clipspace = () => {
const widgets = this.widgets
.map(({ type, name, value }) => ({
type,
name,
value,
}))
widgets.push(
{
type: "text",
name: imageInputName,
value: (urlWidget.value || "").split("\n").filter(Boolean)[0],
},
{
type: "text",
name: "url",
value: (urlWidget.value || "").split("\n").filter(Boolean)[0],
}
)
return { widgets, images: undefined }
}
requestAnimationFrame(() => {
setImagesFromUrl(urlWidget.value);
})
})
}
app.registerExtension({
@@ -627,16 +255,10 @@ app.registerExtension({
return;
}
addKVState(nodeType);
if (nodeData.name === 'LoadImageFromUrl' || nodeData.name === 'LoadImageAsMaskFromUrl') {
migrateWidget(nodeType, 'url', 'image');
addUploadWidget(nodeType, 'image', 'image');
addImagePreview(nodeType, 'image');
} else if (nodeData.name == 'LoadVideoFromUrl') {
addVideoCustomSize(nodeType, nodeData, 'force_size');
addUploadWidget(nodeType, 'video', 'video');
addVideoPreview(nodeType, 'video');
addImageUploadWidget(nodeType, nodeData, 'image');
}
addKVState(nodeType);
},
});
+146
View File
@@ -0,0 +1,146 @@
export const CONVERTED_TYPE = "converted-widget"
export function hideWidgetForGood(node, widget, suffix = "") {
widget.origType = widget.type
widget.origComputeSize = widget.computeSize
widget.computeSize = () => [0, -4] // -4 is due to the gap litegraph adds between widgets automatically
widget.type = CONVERTED_TYPE + suffix
// Hide any linked widgets, e.g. seed+seedControl
if (widget.linkedWidgets) {
for (const w of widget.linkedWidgets) {
hideWidgetForGood(node, w, ":" + widget.name)
}
}
}
const doesInputWithNameExist = (node, name) => {
return node.inputs ? node.inputs.some(input => input.name === name) : false
}
const HIDDEN_TAG = "tschide"
const origProps = {}
// Toggle Widget + change size
export function toggleWidget(node, widget, show = false, suffix = "", updateSize = true) {
if (!widget || doesInputWithNameExist(node, widget.name)) return
// Store the original properties of the widget if not already stored
if (!origProps[widget.name]) {
origProps[widget.name] = {
origType: widget.type,
origComputeSize: widget.computeSize,
}
}
const origSize = node.size
// Set the widget type and computeSize based on the show flag
widget.type = show ? origProps[widget.name].origType : HIDDEN_TAG + suffix
widget.computeSize = show ? origProps[widget.name].origComputeSize : () => [0, -4]
// Recursively handle linked widgets if they exist
widget.linkedWidgets?.forEach(w => toggleWidget(node, w, ":" + widget.name, show))
// Calculate the new height for the node based on its computeSize method
if (updateSize) {
const newHeight = node.computeSize()[1]
node.setSize([node.size[0], newHeight])
}
}
export function addWidgetChangeCallback(widget, callback) {
let widgetValue = widget.value
let originalDescriptor = Object.getOwnPropertyDescriptor(widget, "value")
Object.defineProperty(widget, "value", {
get() {
return originalDescriptor && originalDescriptor.get ? originalDescriptor.get.call(widget) : widgetValue
},
set(newVal) {
if (originalDescriptor && originalDescriptor.set) {
originalDescriptor.set.call(widget, newVal)
} else {
widgetValue = newVal
}
callback(newVal)
},
})
}
export function chainCallback(object, property, callback) {
if (object == undefined) {
//This should not happen.
console.error("Tried to add callback to non-existant object")
return
}
if (property in object) {
const callback_orig = object[property]
object[property] = function () {
const r = callback_orig?.apply(this, arguments)
callback.apply(this, arguments)
return r
}
} else {
object[property] = callback
}
}
export function addKVState(nodeType) {
chainCallback(nodeType.prototype, "onNodeCreated", function () {
chainCallback(this, "onConfigure", function (info) {
if (!this.widgets) {
//Node has no widgets, there is nothing to restore
return
}
if (typeof info.widgets_values != "object") {
//widgets_values is in some unknown inactionable format
return
}
let widgetDict = info.widgets_values
if (widgetDict.length == undefined) {
for (let w of this.widgets) {
if (w.name in widgetDict) {
w.value = widgetDict[w.name]
if (w.type !== "button") {
w.callback?.(w.value)
}
} else {
//attempt to restore default value
let inputs = LiteGraph.getNodeType(this.type).nodeData.input
let initialValue = null
if (inputs?.required?.hasOwnProperty(w.name)) {
if (inputs.required[w.name][1]?.hasOwnProperty("default")) {
initialValue = inputs.required[w.name][1].default
} else if (inputs.required[w.name][0].length) {
initialValue = inputs.required[w.name][0][0]
}
} else if (inputs?.optional?.hasOwnProperty(w.name)) {
if (inputs.optional[w.name][1]?.hasOwnProperty("default")) {
initialValue = inputs.optional[w.name][1].default
} else if (inputs.optional[w.name][0].length) {
initialValue = inputs.optional[w.name][0][0]
}
}
if (initialValue) {
w.value = initialValue
if (w.type !== "button") {
w.callback?.(w.value)
}
}
}
}
}
})
chainCallback(this, "onSerialize", function (info) {
info.widgets_values = {}
if (!this.widgets) {
//object has no widgets, there is nothing to store
return
}
for (let w of this.widgets) {
info.widgets_values[w.name] = w.value
}
})
})
}