Compare commits
84
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
210dc072b1 | ||
|
|
9bd02bd62b | ||
|
|
a445c82b2b | ||
|
|
91eceb7d57 | ||
|
|
090418eb8d | ||
|
|
f1830ba85b | ||
|
|
51ed4a0bc4 | ||
|
|
c62b79b6d9 | ||
|
|
1138a1f4d9 | ||
|
|
f37981bffb | ||
|
|
e40e9244b3 | ||
|
|
c01aafb508 | ||
|
|
61171a2a87 | ||
|
|
91b4761689 | ||
|
|
75b47d8eb4 | ||
|
|
c6b58caca6 | ||
|
|
8965cfb3c3 | ||
|
|
6df1bc2298 | ||
|
|
6107619468 | ||
|
|
7dbb8fb35c | ||
|
|
746c1cefbd | ||
|
|
e5e027e25c | ||
|
|
dd6673b4d5 | ||
|
|
42c116dbb0 | ||
|
|
12ee9ebe16 | ||
|
|
12aa820fdc | ||
|
|
beae36b7ed | ||
|
|
b886928282 | ||
|
|
1721ff7a70 | ||
|
|
8e83109a4e | ||
|
|
e4510faffb | ||
|
|
27c0905ba4 | ||
|
|
896a59a294 | ||
|
|
b3c5a98603 | ||
|
|
52d1b5c874 | ||
|
|
e3b2493b1f | ||
|
|
20952fd8c1 | ||
|
|
43461905b2 | ||
|
|
a2aaa32fbc | ||
|
|
707517937a | ||
|
|
71722e4c7f | ||
|
|
0d7bcc5e23 | ||
|
|
2503bc2cf3 | ||
|
|
3a6a0c2b52 | ||
|
|
efaa2dcd6c | ||
|
|
c77a9b386e | ||
|
|
c3bacdc0c4 | ||
|
|
4e97ff8c4a | ||
|
|
64fa05980d | ||
|
|
d78b709e31 | ||
|
|
4d6caa301b | ||
|
|
ad57177bcd | ||
|
|
fc00f4a094 | ||
|
|
633352bf5d | ||
|
|
3e97c544f2 | ||
|
|
3bf0cfa0cc | ||
|
|
50abaace75 | ||
|
|
d0cf2abdaf | ||
|
|
36ed264742 | ||
|
|
fb8bc917ba | ||
|
|
8d538c9678 | ||
|
|
8034af6478 | ||
|
|
713f6de761 | ||
|
|
83c3732201 | ||
|
|
de986044b3 | ||
|
|
bb6994f677 | ||
|
|
4405434b9c | ||
|
|
e7e9e58c66 | ||
|
|
51dd8fcb7c | ||
|
|
4be544aa9e | ||
|
|
ba06d209d4 | ||
|
|
5d22cae422 | ||
|
|
b9dc7e59cb | ||
|
|
135e58a6e9 | ||
|
|
539864865b | ||
|
|
c467bbe54f | ||
|
|
08fa873f2c | ||
|
|
e1139c55c3 | ||
|
|
a8ceae60ea | ||
|
|
ee47479097 | ||
|
|
ee55653500 | ||
|
|
af630185c3 | ||
|
|
6bf9ad1f3d | ||
|
|
be0a668549 |
@@ -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 }}
|
||||
@@ -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.
|
||||
@@ -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.
|
||||
|
||||

|
||||
|
||||
### 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.
|
||||
|
||||

|
||||
|
||||
### 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
|
||||
|
||||

|
||||
|
||||
#### 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.
|
||||
|
||||

|
||||
|
||||
### 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)
|
||||
|
||||

|
||||
|
||||
#### 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.
|
||||
|
||||

|
||||
|
||||
### 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.
|
||||
|
||||

|
||||
|
||||

|
||||
|
||||
# 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.
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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, **_):
|
||||
|
||||
@@ -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)
|
||||
@@ -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,)
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
{
|
||||
|
||||
@@ -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"},
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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",
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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",
|
||||
}
|
||||
|
||||
|
||||
@@ -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,)
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -9,4 +9,4 @@ kornia
|
||||
pydantic
|
||||
segment_anything
|
||||
omegaconf
|
||||
boto3>=1.34.101
|
||||
boto3
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
})
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user