20 Commits
Author SHA1 Message Date
Kijai ac5bb32bb8 editor linking progress 2024-05-07 16:51:07 +03:00
Kijai 43dded2f42 Add PreviewAnimation -node 2024-05-07 14:07:24 +03:00
Kijai a93b8687ac Update spline_editor.js 2024-05-07 12:19:08 +03:00
kijai eb2d3762a5 Merge branch 'main' into develop 2024-05-06 21:43:25 +03:00
Kijai 6afc4c38e7 continue 2024-05-06 17:37:51 +03:00
Kijai b1a061fe47 continue later 2024-05-06 16:44:35 +03:00
Kijai 2c262f722b Update spline_editor.js 2024-05-06 15:55:33 +03:00
kijai 894d9e5a0b Initial work on SplineEditor linking 2024-05-06 01:30:28 +03:00
kijai 48736ca845 Merge branch 'main' into develop 2024-05-06 00:29:03 +03:00
kijai a33961fdab Update spline_editor.js 2024-05-05 23:59:15 +03:00
kijai 62a704fc14 Merge branch 'main' into develop 2024-05-05 23:57:57 +03:00
kijai d5e37ff797 Update spline_editor.js 2024-04-20 21:22:29 +03:00
kijai bdade789ee spline editor fixes 2024-04-20 21:21:34 +03:00
kijai 1aef2b8e43 spline editor updates 2024-04-20 18:20:02 +03:00
kijai cb1e98abf1 Update spline_editor.js 2024-04-20 11:27:50 +03:00
kijai 891daeb438 Merge branch 'main' into develop 2024-04-19 18:40:44 +03:00
kijai d9712b0e04 spline editor work 2024-04-18 01:45:32 +03:00
Kijai 92711dc762 Update spline_editor.js 2024-04-17 19:43:49 +03:00
Kijai 6f256423af Update spline_editor.js 2024-04-17 19:21:32 +03:00
Kijai 47c23d5a19 reworking spline editor (not functional yet) 2024-04-16 19:19:16 +03:00
31 changed files with 1233 additions and 11158 deletions
-2
View File
@@ -1,2 +0,0 @@
github: [kijai]
custom: ["https://www.paypal.me/kijaidesign"]
-25
View File
@@ -1,25 +0,0 @@
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 == 'kijai' }}
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 }}
+1 -4
View File
@@ -1,11 +1,8 @@
__pycache__
/venv
*.code-workspace
.history
.vscode
*.ckpt
*.pth
types
models
jsconfig.json
custom_dimensions.json
jsconfig.json
+1 -1
View File
@@ -17,7 +17,7 @@ This is still work in progress, like everything else.
## Javascript
### browserstatus.js
Sets the favicon to green circle when not processing anything, sets it to red when processing and shows progress percentage and the length of your queue.
Sets the favicon to green circle when not processing anything, sets it to red when processing and shows progress percentage and the lenghth of your queue.
Default off, needs to be enabled from options, overrides Custom-Scripts favicon when enabled.
## Nodes:
+10 -107
View File
@@ -5,11 +5,8 @@ from .nodes.audioscheduler_nodes import *
from .nodes.image_nodes import *
from .nodes.intrinsic_lora_nodes import *
from .nodes.mask_nodes import *
from .nodes.model_optimization_nodes import *
from .nodes.lora_nodes import *
NODE_CONFIG = {
#constants
"BOOLConstant": {"class": BOOLConstant, "name": "BOOL Constant"},
"INTConstant": {"class": INTConstant, "name": "INT Constant"},
"FloatConstant": {"class": FloatConstant, "name": "Float Constant"},
"StringConstant": {"class": StringConstant, "name": "String Constant"},
@@ -22,7 +19,6 @@ NODE_CONFIG = {
"ConditioningSetMaskAndCombine5": {"class": ConditioningSetMaskAndCombine5, "name": "ConditioningSetMaskAndCombine5"},
"CondPassThrough": {"class": CondPassThrough},
#masking
"DownloadAndLoadCLIPSeg": {"class": DownloadAndLoadCLIPSeg, "name": "(Down)load CLIPSeg"},
"BatchCLIPSeg": {"class": BatchCLIPSeg, "name": "Batch CLIPSeg"},
"ColorToMask": {"class": ColorToMask, "name": "Color To Mask"},
"CreateGradientMask": {"class": CreateGradientMask, "name": "Create Gradient Mask"},
@@ -41,64 +37,32 @@ NODE_CONFIG = {
"RemapMaskRange": {"class": RemapMaskRange, "name": "Remap Mask Range"},
"ResizeMask": {"class": ResizeMask, "name": "Resize Mask"},
"RoundMask": {"class": RoundMask, "name": "Round Mask"},
"SeparateMasks": {"class": SeparateMasks, "name": "Separate Masks"},
#images
"AddLabel": {"class": AddLabel, "name": "Add Label"},
"ColorMatch": {"class": ColorMatch, "name": "Color Match"},
"ImageTensorList": {"class": ImageTensorList, "name": "Image Tensor List"},
"CrossFadeImages": {"class": CrossFadeImages, "name": "Cross Fade Images"},
"CrossFadeImagesMulti": {"class": CrossFadeImagesMulti, "name": "Cross Fade Images Multi"},
"GetImagesFromBatchIndexed": {"class": GetImagesFromBatchIndexed, "name": "Get Images From Batch Indexed"},
"GetImageRangeFromBatch": {"class": GetImageRangeFromBatch, "name": "Get Image or Mask Range From Batch"},
"GetLatentRangeFromBatch": {"class": GetLatentRangeFromBatch, "name": "Get Latent Range From Batch"},
"GetLatentSizeAndCount": {"class": GetLatentSizeAndCount, "name": "Get Latent Size & Count"},
"GetImageRangeFromBatch": {"class": GetImageRangeFromBatch, "name": "Get Image Range From Batch"},
"GetImageSizeAndCount": {"class": GetImageSizeAndCount, "name": "Get Image Size & Count"},
"FastPreview": {"class": FastPreview, "name": "Fast Preview"},
"ImageBatchFilter": {"class": ImageBatchFilter, "name": "Image Batch Filter"},
"ImageAndMaskPreview": {"class": ImageAndMaskPreview},
"ImageAddMulti": {"class": ImageAddMulti, "name": "Image Add Multi"},
"ImageBatchJoinWithTransition": {"class": ImageBatchJoinWithTransition, "name": "Image Batch Join With Transition"},
"ImageBatchMulti": {"class": ImageBatchMulti, "name": "Image Batch Multi"},
"ImageBatchRepeatInterleaving": {"class": ImageBatchRepeatInterleaving},
"ImageBatchTestPattern": {"class": ImageBatchTestPattern, "name": "Image Batch Test Pattern"},
"ImageConcanate": {"class": ImageConcanate, "name": "Image Concatenate"},
"ImageConcatFromBatch": {"class": ImageConcatFromBatch, "name": "Image Concatenate From Batch"},
"ImageConcatMulti": {"class": ImageConcatMulti, "name": "Image Concatenate Multi"},
"ImageCropByMask": {"class": ImageCropByMask, "name": "Image Crop By Mask"},
"ImageCropByMaskAndResize": {"class": ImageCropByMaskAndResize, "name": "Image Crop By Mask And Resize"},
"ImageCropByMaskBatch": {"class": ImageCropByMaskBatch, "name": "Image Crop By Mask Batch"},
"ImageUncropByMask": {"class": ImageUncropByMask, "name": "Image Uncrop By Mask"},
"ImageGrabPIL": {"class": ImageGrabPIL, "name": "Image Grab PIL"},
"ImageGridComposite2x2": {"class": ImageGridComposite2x2, "name": "Image Grid Composite 2x2"},
"ImageGridComposite3x3": {"class": ImageGridComposite3x3, "name": "Image Grid Composite 3x3"},
"ImageGridtoBatch": {"class": ImageGridtoBatch, "name": "Image Grid To Batch"},
"ImageNoiseAugmentation": {"class": ImageNoiseAugmentation, "name": "Image Noise Augmentation"},
"ImageNormalize_Neg1_To_1": {"class": ImageNormalize_Neg1_To_1, "name": "Image Normalize -1 to 1"},
"ImagePass": {"class": ImagePass},
"ImagePadKJ": {"class": ImagePadKJ, "name": "ImagePad KJ"},
"ImagePadForOutpaintMasked": {"class": ImagePadForOutpaintMasked, "name": "Image Pad For Outpaint Masked"},
"ImagePadForOutpaintTargetSize": {"class": ImagePadForOutpaintTargetSize, "name": "Image Pad For Outpaint Target Size"},
"ImagePrepForICLora": {"class": ImagePrepForICLora, "name": "Image Prep For ICLora"},
"ImageResizeKJ": {"class": ImageResizeKJ, "name": "Resize Image (deprecated)"},
"ImageResizeKJv2": {"class": ImageResizeKJv2, "name": "Resize Image v2"},
"ImageUpscaleWithModelBatched": {"class": ImageUpscaleWithModelBatched, "name": "Image Upscale With Model Batched"},
"InsertImagesToBatchIndexed": {"class": InsertImagesToBatchIndexed, "name": "Insert Images To Batch Indexed"},
"InsertLatentToIndexed": {"class": InsertLatentToIndex, "name": "Insert Latent To Index"},
"LoadAndResizeImage": {"class": LoadAndResizeImage, "name": "Load & Resize Image"},
"LoadImagesFromFolderKJ": {"class": LoadImagesFromFolderKJ, "name": "Load Images From Folder (KJ)"},
"LoadVideosFromFolder": {"class": LoadVideosFromFolder, "name": "Load Videos From Folder"},
"MergeImageChannels": {"class": MergeImageChannels, "name": "Merge Image Channels"},
"PadImageBatchInterleaved": {"class": PadImageBatchInterleaved, "name": "Pad Image Batch Interleaved"},
"PreviewAnimation": {"class": PreviewAnimation, "name": "Preview Animation"},
"RemapImageRange": {"class": RemapImageRange, "name": "Remap Image Range"},
"ReverseImageBatch": {"class": ReverseImageBatch, "name": "Reverse Image Batch"},
"ReplaceImagesInBatch": {"class": ReplaceImagesInBatch, "name": "Replace Images In Batch"},
"SaveImageWithAlpha": {"class": SaveImageWithAlpha, "name": "Save Image With Alpha"},
"SaveImageKJ": {"class": SaveImageKJ, "name": "Save Image KJ"},
"ShuffleImageBatch": {"class": ShuffleImageBatch, "name": "Shuffle Image Batch"},
"SplitImageChannels": {"class": SplitImageChannels, "name": "Split Image Channels"},
"TransitionImagesMulti": {"class": TransitionImagesMulti, "name": "Transition Images Multi"},
"TransitionImagesInBatch": {"class": TransitionImagesInBatch, "name": "Transition Images In Batch"},
"SplitImageChannels": {"class": SplitImageChannels, "name": "Split Image Channels"},
#batch cropping
"BatchCropFromMask": {"class": BatchCropFromMask, "name": "Batch Crop From Mask"},
"BatchCropFromMaskAdvanced": {"class": BatchCropFromMaskAdvanced, "name": "Batch Crop From Mask Advanced"},
@@ -115,52 +79,34 @@ NODE_CONFIG = {
"InjectNoiseToLatent": {"class": InjectNoiseToLatent, "name": "Inject Noise To Latent"},
"CustomSigmas": {"class": CustomSigmas, "name": "Custom Sigmas"},
#utility
"StringToFloatList": {"class": StringToFloatList, "name": "String to Float List"},
"WidgetToString": {"class": WidgetToString, "name": "Widget To String"},
"SaveStringKJ": {"class": SaveStringKJ, "name": "Save String KJ"},
"DummyOut": {"class": DummyOut, "name": "Dummy Out"},
"DummyLatentOut": {"class": DummyLatentOut, "name": "Dummy Latent Out"},
"GetLatentsFromBatchIndexed": {"class": GetLatentsFromBatchIndexed, "name": "Get Latents From Batch Indexed"},
"ScaleBatchPromptSchedule": {"class": ScaleBatchPromptSchedule, "name": "Scale Batch Prompt Schedule"},
"CameraPoseVisualizer": {"class": CameraPoseVisualizer, "name": "Camera Pose Visualizer"},
"AppendStringsToList": {"class": AppendStringsToList, "name": "Append Strings To List"},
"JoinStrings": {"class": JoinStrings, "name": "Join Strings"},
"JoinStringMulti": {"class": JoinStringMulti, "name": "Join String Multi"},
"SomethingToString": {"class": SomethingToString, "name": "Something To String"},
"Sleep": {"class": Sleep, "name": "Sleep"},
"VRAM_Debug": {"class": VRAM_Debug, "name": "VRAM Debug"},
"SomethingToString": {"class": SomethingToString, "name": "Something To String"},
"EmptyLatentImagePresets": {"class": EmptyLatentImagePresets, "name": "Empty Latent Image Presets"},
"EmptyLatentImageCustomPresets": {"class": EmptyLatentImageCustomPresets, "name": "Empty Latent Image Custom Presets"},
"ModelPassThrough": {"class": ModelPassThrough, "name": "ModelPass"},
"ModelSaveKJ": {"class": ModelSaveKJ, "name": "Model Save KJ"},
"SetShakkerLabsUnionControlNetType": {"class": SetShakkerLabsUnionControlNetType, "name": "Set Shakker Labs Union ControlNet Type"},
"StyleModelApplyAdvanced": {"class": StyleModelApplyAdvanced, "name": "Style Model Apply Advanced"},
"DiffusionModelSelector": {"class": DiffusionModelSelector, "name": "Diffusion Model Selector"},
"LazySwitchKJ": {"class": LazySwitchKJ, "name": "Lazy Switch KJ"},
#audioscheduler stuff
"NormalizedAmplitudeToMask": {"class": NormalizedAmplitudeToMask},
"NormalizedAmplitudeToFloatList": {"class": NormalizedAmplitudeToFloatList},
"OffsetMaskByNormalizedAmplitude": {"class": OffsetMaskByNormalizedAmplitude},
"ImageTransformByNormalizedAmplitude": {"class": ImageTransformByNormalizedAmplitude},
"AudioConcatenate": {"class": AudioConcatenate},
#curve nodes
"SplineEditor": {"class": SplineEditor, "name": "Spline Editor"},
"CreateShapeImageOnPath": {"class": CreateShapeImageOnPath, "name": "Create Shape Image On Path"},
"CreateShapeMaskOnPath": {"class": CreateShapeMaskOnPath, "name": "Create Shape Mask On Path"},
"CreateTextOnPath": {"class": CreateTextOnPath, "name": "Create Text On Path"},
"CreateGradientFromCoords": {"class": CreateGradientFromCoords, "name": "Create Gradient From Coords"},
"CutAndDragOnPath": {"class": CutAndDragOnPath, "name": "Cut And Drag On Path"},
"GradientToFloat": {"class": GradientToFloat, "name": "Gradient To Float"},
"WeightScheduleExtend": {"class": WeightScheduleExtend, "name": "Weight Schedule Extend"},
"MaskOrImageToWeight": {"class": MaskOrImageToWeight, "name": "Mask Or Image To Weight"},
"WeightScheduleConvert": {"class": WeightScheduleConvert, "name": "Weight Schedule Convert"},
"FloatToMask": {"class": FloatToMask, "name": "Float To Mask"},
"FloatToSigmas": {"class": FloatToSigmas, "name": "Float To Sigmas"},
"SigmasToFloat": {"class": SigmasToFloat, "name": "Sigmas To Float"},
"PlotCoordinates": {"class": PlotCoordinates, "name": "Plot Coordinates"},
"InterpolateCoords": {"class": InterpolateCoords, "name": "Interpolate Coords"},
"PointsEditor": {"class": PointsEditor, "name": "Points Editor"},
#experimental
"StabilityAPI_SD3": {"class": StabilityAPI_SD3, "name": "Stability API SD3"},
"SoundReactive": {"class": SoundReactive, "name": "Sound Reactive"},
"StableZero123_BatchSchedule": {"class": StableZero123_BatchSchedule, "name": "Stable Zero123 Batch Schedule"},
"SV3D_BatchSchedule": {"class": SV3D_BatchSchedule, "name": "SV3D Batch Schedule"},
@@ -168,51 +114,10 @@ NODE_CONFIG = {
"Superprompt": {"class": Superprompt, "name": "Superprompt"},
"GLIGENTextBoxApplyBatchCoords": {"class": GLIGENTextBoxApplyBatchCoords},
"Intrinsic_lora_sampling": {"class": Intrinsic_lora_sampling, "name": "Intrinsic Lora Sampling"},
"CheckpointPerturbWeights": {"class": CheckpointPerturbWeights, "name": "CheckpointPerturbWeights"},
"Screencap_mss": {"class": Screencap_mss, "name": "Screencap mss"},
"WebcamCaptureCV2": {"class": WebcamCaptureCV2, "name": "Webcam Capture CV2"},
"DifferentialDiffusionAdvanced": {"class": DifferentialDiffusionAdvanced, "name": "Differential Diffusion Advanced"},
"DiTBlockLoraLoader": {"class": DiTBlockLoraLoader, "name": "DiT Block Lora Loader"},
"FluxBlockLoraSelect": {"class": FluxBlockLoraSelect, "name": "Flux Block Lora Select"},
"HunyuanVideoBlockLoraSelect": {"class": HunyuanVideoBlockLoraSelect, "name": "Hunyuan Video Block Lora Select"},
"Wan21BlockLoraSelect": {"class": Wan21BlockLoraSelect, "name": "Wan21 Block Lora Select"},
"CustomControlNetWeightsFluxFromList": {"class": CustomControlNetWeightsFluxFromList, "name": "Custom ControlNet Weights Flux From List"},
"CheckpointLoaderKJ": {"class": CheckpointLoaderKJ, "name": "CheckpointLoaderKJ"},
"DiffusionModelLoaderKJ": {"class": DiffusionModelLoaderKJ, "name": "Diffusion Model Loader KJ"},
"TorchCompileModelFluxAdvanced": {"class": TorchCompileModelFluxAdvanced, "name": "TorchCompileModelFluxAdvanced"},
"TorchCompileModelFluxAdvancedV2": {"class": TorchCompileModelFluxAdvancedV2, "name": "TorchCompileModelFluxAdvancedV2"},
"TorchCompileModelHyVideo": {"class": TorchCompileModelHyVideo, "name": "TorchCompileModelHyVideo"},
"TorchCompileVAE": {"class": TorchCompileVAE, "name": "TorchCompileVAE"},
"TorchCompileControlNet": {"class": TorchCompileControlNet, "name": "TorchCompileControlNet"},
"PatchModelPatcherOrder": {"class": PatchModelPatcherOrder, "name": "Patch Model Patcher Order"},
"TorchCompileLTXModel": {"class": TorchCompileLTXModel, "name": "TorchCompileLTXModel"},
"TorchCompileCosmosModel": {"class": TorchCompileCosmosModel, "name": "TorchCompileCosmosModel"},
"TorchCompileModelQwenImage": {"class": TorchCompileModelQwenImage, "name": "TorchCompileModelQwenImage"},
"TorchCompileModelWanVideo": {"class": TorchCompileModelWanVideo, "name": "TorchCompileModelWanVideo"},
"TorchCompileModelWanVideoV2": {"class": TorchCompileModelWanVideoV2, "name": "TorchCompileModelWanVideoV2"},
"PathchSageAttentionKJ": {"class": PathchSageAttentionKJ, "name": "Patch Sage Attention KJ"},
"LeapfusionHunyuanI2VPatcher": {"class": LeapfusionHunyuanI2V, "name": "Leapfusion Hunyuan I2V Patcher"},
"VAELoaderKJ": {"class": VAELoaderKJ, "name": "VAELoader KJ"},
"ScheduledCFGGuidance": {"class": ScheduledCFGGuidance, "name": "Scheduled CFG Guidance"},
"ApplyRifleXRoPE_HunuyanVideo": {"class": ApplyRifleXRoPE_HunuyanVideo, "name": "Apply RifleXRoPE HunuyanVideo"},
"ApplyRifleXRoPE_WanVideo": {"class": ApplyRifleXRoPE_WanVideo, "name": "Apply RifleXRoPE WanVideo"},
"WanVideoTeaCacheKJ": {"class": WanVideoTeaCacheKJ, "name": "WanVideo Tea Cache (native)"},
"WanVideoEnhanceAVideoKJ": {"class": WanVideoEnhanceAVideoKJ, "name": "WanVideo Enhance A Video (native)"},
"SkipLayerGuidanceWanVideo": {"class": SkipLayerGuidanceWanVideo, "name": "Skip Layer Guidance WanVideo"},
"TimerNodeKJ": {"class": TimerNodeKJ, "name": "Timer Node KJ"},
"HunyuanVideoEncodeKeyframesToCond": {"class": HunyuanVideoEncodeKeyframesToCond, "name": "HunyuanVideo Encode Keyframes To Cond"},
"CFGZeroStarAndInit": {"class": CFGZeroStarAndInit, "name": "CFG Zero Star/Init"},
"ModelPatchTorchSettings": {"class": ModelPatchTorchSettings, "name": "Model Patch Torch Settings"},
"WanVideoNAG": {"class": WanVideoNAG, "name": "WanVideoNAG"},
#instance diffusion
"CreateInstanceDiffusionTracking": {"class": CreateInstanceDiffusionTracking},
"AppendInstanceDiffusionTracking": {"class": AppendInstanceDiffusionTracking},
"DrawInstanceDiffusionTracking": {"class": DrawInstanceDiffusionTracking},
#lora
"LoraExtractKJ": {"class": LoraExtractKJ, "name": "LoraExtractKJ"},
"LoraReduceRankKJ": {"class": LoraReduceRank, "name": "LoraReduceRank"}
}
def generate_node_mappings(node_config):
@@ -236,11 +141,9 @@ from server import PromptServer
from pathlib import Path
if hasattr(PromptServer, "instance"):
try:
# NOTE: we add an extra static path to avoid comfy mechanism
# that loads every script in web.
PromptServer.instance.app.add_routes(
[web.static("/kjweb_async", (Path(__file__).parent.absolute() / "kjweb_async").as_posix())]
)
except:
pass
# NOTE: we add an extra static path to avoid comfy mechanism
# that loads every script in web.
PromptServer.instance.app.add_routes(
[web.static("/kjweb_async", (Path(__file__).parent.absolute() / "kjweb_async").as_posix())]
)
+3
View File
@@ -0,0 +1,3 @@
{
"sai_api_key": "your_api_key_here"
}
-22
View File
@@ -1,22 +0,0 @@
[
{
"label": "SD",
"value": "512x512"
},
{
"label": "HD",
"value": "768x768"
},
{
"label": "Full HD",
"value": "1024x1024"
},
{
"label": "4k",
"value": "2048x2048"
},
{
"label": "SVD",
"value": "1024x576"
}
]
File diff suppressed because it is too large Load Diff
+12 -38
View File
@@ -694,7 +694,6 @@ class BboxVisualize:
"images": ("IMAGE",),
"bboxes": ("BBOX",),
"line_width": ("INT", {"default": 1,"min": 1, "max": 10, "step": 1}),
"bbox_format": (["xywh", "xyxy"], {"default": "xywh"}),
},
}
@@ -707,56 +706,31 @@ Visualizes the specified bbox on the image.
CATEGORY = "KJNodes/masking"
def visualizebbox(self, bboxes, images, line_width, bbox_format):
def visualizebbox(self, bboxes, images, line_width):
image_list = []
for image, bbox in zip(images, bboxes):
if bbox_format == "xywh":
x_min, y_min, width, height = bbox
elif bbox_format == "xyxy":
x_min, y_min, x_max, y_max = bbox
width = x_max - x_min
height = y_max - y_min
else:
raise ValueError(f"Unknown bbox_format: {bbox_format}")
# Ensure bbox coordinates are integers
x_min = int(x_min)
y_min = int(y_min)
width = int(width)
height = int(height)
# Permute the image dimensions
x_min, y_min, width, height = bbox
image = image.permute(2, 0, 1)
# Clone the image to draw bounding boxes
img_with_bbox = image.clone()
# Define the color for the bbox, e.g., red
color = torch.tensor([1, 0, 0], dtype=torch.float32)
# Ensure color tensor matches the image channels
if color.shape[0] != img_with_bbox.shape[0]:
color = color.unsqueeze(1).expand(-1, line_width)
# Draw lines for each side of the bbox with the specified line width
for lw in range(line_width):
# Top horizontal line
if y_min + lw < img_with_bbox.shape[1]:
img_with_bbox[:, y_min + lw, x_min:x_min + width] = color[:, None]
img_with_bbox[:, y_min + lw, x_min:x_min + width] = color[:, None]
# Bottom horizontal line
if y_min + height - lw < img_with_bbox.shape[1]:
img_with_bbox[:, y_min + height - lw, x_min:x_min + width] = color[:, None]
img_with_bbox[:, y_min + height - lw, x_min:x_min + width] = color[:, None]
# Left vertical line
if x_min + lw < img_with_bbox.shape[2]:
img_with_bbox[:, y_min:y_min + height, x_min + lw] = color[:, None]
img_with_bbox[:, y_min:y_min + height, x_min + lw] = color[:, None]
# Right vertical line
if x_min + width - lw < img_with_bbox.shape[2]:
img_with_bbox[:, y_min:y_min + height, x_min + width - lw] = color[:, None]
# Permute the image dimensions back
img_with_bbox[:, y_min:y_min + height, x_min + width - lw] = color[:, None]
img_with_bbox = img_with_bbox.permute(1, 2, 0).unsqueeze(0)
image_list.append(img_with_bbox)
+43 -726
View File
@@ -1,49 +1,10 @@
import torch
from torchvision import transforms
import json
from PIL import Image, ImageDraw, ImageFont, ImageColor, ImageFilter, ImageChops
from PIL import Image, ImageDraw, ImageFont
import numpy as np
from ..utility.utility import pil2tensor, tensor2pil
from ..utility.utility import pil2tensor
import folder_paths
import io
import base64
from comfy.utils import common_upscale
def parse_color(color):
if isinstance(color, str) and ',' in color:
return tuple(int(c.strip()) for c in color.split(','))
return color
def parse_json_tracks(tracks):
tracks_data = []
try:
# If tracks is a string, try to parse it as JSON
if isinstance(tracks, str):
parsed = json.loads(tracks.replace("'", '"'))
tracks_data.extend(parsed)
else:
# If tracks is a list of strings, parse each one
for track_str in tracks:
parsed = json.loads(track_str.replace("'", '"'))
tracks_data.append(parsed)
# Check if we have a single track (dict with x,y) or a list of tracks
if tracks_data and isinstance(tracks_data[0], dict) and 'x' in tracks_data[0]:
# Single track detected, wrap it in a list
tracks_data = [tracks_data]
elif tracks_data and isinstance(tracks_data[0], list) and tracks_data[0] and isinstance(tracks_data[0][0], dict) and 'x' in tracks_data[0][0]:
# Already a list of tracks, nothing to do
pass
else:
# Unexpected format
print(f"Warning: Unexpected track format: {type(tracks_data[0])}")
except json.JSONDecodeError as e:
print(f"Error parsing tracks JSON: {e}")
tracks_data = []
return tracks_data
def plot_coordinates_to_tensor(coordinates, height, width, bbox_height, bbox_width, size_multiplier, prompt):
import matplotlib
@@ -129,10 +90,8 @@ Plots coordinates to sequence of images using Matplotlib.
def append(self, coordinates, text, width, height, bbox_width, bbox_height, size_multiplier=[1.0]):
coordinates = json.loads(coordinates.replace("'", '"'))
coordinates = [(coord['x'], coord['y']) for coord in coordinates]
batch_size = len(coordinates)
if not size_multiplier or len(size_multiplier) != batch_size:
size_multiplier = [0] * batch_size
else:
batch_size = len(coordinates)
if len(size_multiplier) != batch_size:
size_multiplier = size_multiplier * (batch_size // len(size_multiplier)) + size_multiplier[:batch_size % len(size_multiplier)]
plot_image_tensor = plot_coordinates_to_tensor(coordinates, height, width, bbox_height, bbox_width, size_multiplier, text)
@@ -154,8 +113,6 @@ class SplineEditor:
[
'path',
'time',
'controlpoints',
'speed'
],
{
"default": 'time'
@@ -189,12 +146,12 @@ class SplineEditor:
"optional": {
"min_value": ("FLOAT", {"default": 0.0, "min": -10000.0, "max": 10000.0, "step": 0.01}),
"max_value": ("FLOAT", {"default": 1.0, "min": -10000.0, "max": 10000.0, "step": 0.01}),
"bg_image": ("IMAGE", ),
"editor_link": ("EDITORLINK",),
}
}
RETURN_TYPES = ("MASK", "STRING", "FLOAT", "INT", "STRING",)
RETURN_NAMES = ("mask", "coord_str", "float", "count", "normalized_str",)
RETURN_TYPES = ("MASK", "STRING", "FLOAT", "INT", "EDITORLINK",)
RETURN_NAMES = ("mask", "coord_str", "float", "count", "editor_link",)
FUNCTION = "splinedata"
CATEGORY = "KJNodes/weights"
DESCRIPTION = """
@@ -212,16 +169,6 @@ guaranteed!!
Note that you can't delete from start/end.
Right click on canvas for context menu:
NEW!:
- Add new spline
- Creates a new spline on same canvas, currently these paths are only outputed
as coordinates.
- Add single point
- Creates a single point that only returns it's current position coords
- Delete spline
- Deletes the currently selected spline, you can select a spline by clicking on
it's path, or cycle through them with the 'Next spline' -option.
These are purely visual options, doesn't affect the output:
- Toggle handles visibility
- Display sample points: display the points to be returned.
@@ -232,7 +179,6 @@ actual control points, so the interpolation type matters.
sampling_method:
- time: samples along the time axis, used for schedules
- path: samples along the path itself, useful for coordinates
- controlpoints: samples only the control points themselves
output types:
- mask batch
@@ -247,77 +193,36 @@ output types:
"""
def splinedata(self, mask_width, mask_height, coordinates, float_output_type, interpolation,
points_to_sample, sampling_method, points_store, tension, repeat_output,
min_value=0.0, max_value=1.0, bg_image=None):
points_to_sample, sampling_method, points_store, tension, repeat_output, min_value=0.0, max_value=1.0, editor_link=None):
coordinates = json.loads(coordinates)
# Handle nested list structure if present
all_normalized = []
all_normalized_y_values = []
# Check if we have a nested list structure
if isinstance(coordinates, list) and len(coordinates) > 0 and isinstance(coordinates[0], list):
# Process each list of coordinates in the nested structure
coordinate_sets = coordinates
else:
# If not nested, treat as a single list of coordinates
coordinate_sets = [coordinates]
# Process each set of coordinates
for coord_set in coordinate_sets:
normalized = []
normalized_y_values = []
for coord in coordinates:
coord['x'] = int(round(coord['x']))
coord['y'] = int(round(coord['y']))
for coord in coord_set:
coord['x'] = int(round(coord['x']))
coord['y'] = int(round(coord['y']))
norm_x = (1.0 - (coord['x'] / mask_height) - 0.0) * (max_value - min_value) + min_value
norm_y = (1.0 - (coord['y'] / mask_height) - 0.0) * (max_value - min_value) + min_value
normalized_y_values.append(norm_y)
normalized.append({'x':norm_x, 'y':norm_y})
all_normalized.extend(normalized)
all_normalized_y_values.extend(normalized_y_values)
# Use the combined normalized values for output
normalized_y_values = [
(1.0 - (point['y'] / mask_height) - 0.0) * (max_value - min_value) + min_value
for point in coordinates
]
if float_output_type == 'list':
out_floats = all_normalized_y_values * repeat_output
out_floats = normalized_y_values * repeat_output
elif float_output_type == 'pandas series':
try:
import pandas as pd
except:
raise Exception("MaskOrImageToWeight: pandas is not installed. Please install pandas to use this output_type")
out_floats = pd.Series(all_normalized_y_values * repeat_output),
out_floats = pd.Series(normalized_y_values * repeat_output),
elif float_output_type == 'tensor':
out_floats = torch.tensor(all_normalized_y_values * repeat_output, dtype=torch.float32)
out_floats = torch.tensor(normalized_y_values * repeat_output, dtype=torch.float32)
# Create a color map for grayscale intensities
color_map = lambda y: torch.full((mask_height, mask_width, 3), y, dtype=torch.float32)
# Create image tensors for each normalized y value
mask_tensors = [color_map(y) for y in all_normalized_y_values]
mask_tensors = [color_map(y) for y in normalized_y_values]
masks_out = torch.stack(mask_tensors)
masks_out = masks_out.repeat(repeat_output, 1, 1, 1)
masks_out = masks_out.mean(dim=-1)
if bg_image is None:
return (masks_out, json.dumps(coordinates if len(coordinates) > 1 else coordinates[0]), out_floats, len(out_floats), json.dumps(all_normalized))
else:
transform = transforms.ToPILImage()
image = transform(bg_image[0].permute(2, 0, 1))
buffered = io.BytesIO()
image.save(buffered, format="JPEG", quality=75)
# Encode the image bytes to a Base64 string
img_bytes = buffered.getvalue()
img_base64 = base64.b64encode(img_bytes).decode('utf-8')
return {
"ui": {"bg_image": [img_base64]},
"result": (masks_out, json.dumps(coordinates if len(coordinates) > 1 else coordinates[0]), out_floats, len(out_floats), json.dumps(all_normalized))
}
return (masks_out, str(coordinates), out_floats, len(out_floats), editor_link,)
class CreateShapeMaskOnPath:
@@ -328,8 +233,8 @@ class CreateShapeMaskOnPath:
DESCRIPTION = """
Creates a mask or batch of masks with the specified shape.
Locations are center locations.
Grow value is the amount to grow the shape on each frame, creating animated masks.
"""
DEPRECATED = True
@classmethod
def INPUT_TYPES(s):
@@ -362,9 +267,7 @@ Locations are center locations.
batch_size = len(coordinates)
out = []
color = "white"
if not size_multiplier or len(size_multiplier) != batch_size:
size_multiplier = [0] * batch_size
else:
if len(size_multiplier) != batch_size:
size_multiplier = size_multiplier * (batch_size // len(size_multiplier)) + size_multiplier[:batch_size % len(size_multiplier)]
for i, coord in enumerate(coordinates):
image = Image.new("RGB", (frame_width, frame_height), "black")
@@ -400,320 +303,6 @@ Locations are center locations.
out.append(mask)
outstack = torch.cat(out, dim=0)
return (outstack, 1.0 - outstack,)
class CreateShapeImageOnPath:
RETURN_TYPES = ("IMAGE", "MASK",)
RETURN_NAMES = ("image","mask", )
FUNCTION = "createshapemask"
CATEGORY = "KJNodes/image"
DESCRIPTION = """
Creates an image or batch of images with the specified shape.
Locations are center locations.
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"shape": (
[ 'circle',
'square',
'triangle',
],
{
"default": 'circle'
}),
"coordinates": ("STRING", {"forceInput": True}),
"frame_width": ("INT", {"default": 512,"min": 16, "max": 4096, "step": 1}),
"frame_height": ("INT", {"default": 512,"min": 16, "max": 4096, "step": 1}),
"shape_width": ("INT", {"default": 128,"min": 2, "max": 4096, "step": 1}),
"shape_height": ("INT", {"default": 128,"min": 2, "max": 4096, "step": 1}),
"shape_color": ("STRING", {"default": 'white'}),
"bg_color": ("STRING", {"default": 'black'}),
"blur_radius": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100, "step": 0.1}),
"intensity": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 100.0, "step": 0.01}),
},
"optional": {
"size_multiplier": ("FLOAT", {"default": [1.0], "forceInput": True}),
"trailing": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
"border_width": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1}),
"border_color": ("STRING", {"default": 'black'}),
}
}
def createshapemask(self, coordinates, frame_width, frame_height, shape_width, shape_height, shape_color,
bg_color, blur_radius, shape, intensity, size_multiplier=[1.0], trailing=1.0, border_width=0, border_color='black'):
shape_color = parse_color(shape_color)
border_color = parse_color(border_color)
bg_color = parse_color(bg_color)
coords_list = parse_json_tracks(coordinates)
batch_size = len(coords_list[0])
images_list = []
masks_list = []
if not size_multiplier or len(size_multiplier) != batch_size:
size_multiplier = [1] * batch_size
else:
size_multiplier = size_multiplier * (batch_size // len(size_multiplier)) + size_multiplier[:batch_size % len(size_multiplier)]
previous_output = None
for i in range(batch_size):
image = Image.new("RGB", (frame_width, frame_height), bg_color)
draw = ImageDraw.Draw(image)
# Calculate the size for this frame and ensure it's not less than 0
current_width = shape_width * size_multiplier[i]
current_height = shape_height * size_multiplier[i]
for coords in coords_list:
location_x = coords[i]['x']
location_y = coords[i]['y']
if shape == 'circle' or shape == 'square':
# Define the bounding box for the shape
left_up_point = (location_x - current_width // 2, location_y - current_height // 2)
right_down_point = (location_x + current_width // 2, location_y + current_height // 2)
two_points = [left_up_point, right_down_point]
if shape == 'circle':
if border_width > 0:
draw.ellipse(two_points, fill=shape_color, outline=border_color, width=border_width)
else:
draw.ellipse(two_points, fill=shape_color)
elif shape == 'square':
if border_width > 0:
draw.rectangle(two_points, fill=shape_color, outline=border_color, width=border_width)
else:
draw.rectangle(two_points, fill=shape_color)
elif shape == 'triangle':
# Define the points for the triangle
left_up_point = (location_x - current_width // 2, location_y + current_height // 2) # bottom left
right_down_point = (location_x + current_width // 2, location_y + current_height // 2) # bottom right
top_point = (location_x, location_y - current_height // 2) # top point
if border_width > 0:
draw.polygon([top_point, left_up_point, right_down_point], fill=shape_color, outline=border_color, width=border_width)
else:
draw.polygon([top_point, left_up_point, right_down_point], fill=shape_color)
if blur_radius != 0:
image = image.filter(ImageFilter.GaussianBlur(blur_radius))
# Blend the current image with the accumulated image
image = pil2tensor(image)
if trailing != 1.0 and previous_output is not None:
# Add the decayed previous output to the current frame
image += trailing * previous_output
image = image / image.max()
previous_output = image
image = image * intensity
mask = image[:, :, :, 0]
masks_list.append(mask)
images_list.append(image)
out_images = torch.cat(images_list, dim=0).cpu().float()
out_masks = torch.cat(masks_list, dim=0)
return (out_images, out_masks)
class CreateTextOnPath:
RETURN_TYPES = ("IMAGE", "MASK", "MASK",)
RETURN_NAMES = ("image", "mask", "mask_inverted",)
FUNCTION = "createtextmask"
CATEGORY = "KJNodes/masking/generate"
DESCRIPTION = """
Creates a mask or batch of masks with the specified text.
Locations are center locations.
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"coordinates": ("STRING", {"forceInput": True}),
"text": ("STRING", {"default": 'text', "multiline": True}),
"frame_width": ("INT", {"default": 512,"min": 16, "max": 4096, "step": 1}),
"frame_height": ("INT", {"default": 512,"min": 16, "max": 4096, "step": 1}),
"font": (folder_paths.get_filename_list("kjnodes_fonts"), ),
"font_size": ("INT", {"default": 42}),
"alignment": (
[ 'left',
'center',
'right'
],
{"default": 'center'}
),
"text_color": ("STRING", {"default": 'white'}),
},
"optional": {
"size_multiplier": ("FLOAT", {"default": [1.0], "forceInput": True}),
}
}
def createtextmask(self, coordinates, frame_width, frame_height, font, font_size, text, text_color, alignment, size_multiplier=[1.0]):
coordinates = coordinates.replace("'", '"')
coordinates = json.loads(coordinates)
batch_size = len(coordinates)
mask_list = []
image_list = []
color = parse_color(text_color)
font_path = folder_paths.get_full_path("kjnodes_fonts", font)
if len(size_multiplier) != batch_size:
size_multiplier = size_multiplier * (batch_size // len(size_multiplier)) + size_multiplier[:batch_size % len(size_multiplier)]
for i, coord in enumerate(coordinates):
image = Image.new("RGB", (frame_width, frame_height), "black")
draw = ImageDraw.Draw(image)
lines = text.split('\n') # Split the text into lines
# Apply the size multiplier to the font size for this iteration
current_font_size = int(font_size * size_multiplier[i])
current_font = ImageFont.truetype(font_path, current_font_size)
line_heights = [current_font.getbbox(line)[3] for line in lines] # List of line heights
total_text_height = sum(line_heights) # Total height of text block
# Calculate the starting Y position to center the block of text
start_y = coord['y'] - total_text_height // 2
for j, line in enumerate(lines):
text_width, text_height = current_font.getbbox(line)[2], line_heights[j]
if alignment == 'left':
location_x = coord['x']
elif alignment == 'center':
location_x = int(coord['x'] - text_width // 2)
elif alignment == 'right':
location_x = int(coord['x'] - text_width)
location_y = int(start_y + sum(line_heights[:j]))
text_position = (location_x, location_y)
# Draw the text
try:
draw.text(text_position, line, fill=color, font=current_font, features=['-liga'])
except:
draw.text(text_position, line, fill=color, font=current_font)
image = pil2tensor(image)
non_black_pixels = (image > 0).any(dim=-1)
mask = non_black_pixels.to(image.dtype)
mask_list.append(mask)
image_list.append(image)
out_images = torch.cat(image_list, dim=0).cpu().float()
out_masks = torch.cat(mask_list, dim=0)
return (out_images, out_masks, 1.0 - out_masks,)
class CreateGradientFromCoords:
RETURN_TYPES = ("IMAGE", )
RETURN_NAMES = ("image", )
FUNCTION = "generate"
CATEGORY = "KJNodes/image"
DESCRIPTION = """
Creates a gradient image from coordinates.
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"coordinates": ("STRING", {"forceInput": True}),
"frame_width": ("INT", {"default": 512,"min": 16, "max": 4096, "step": 1}),
"frame_height": ("INT", {"default": 512,"min": 16, "max": 4096, "step": 1}),
"start_color": ("STRING", {"default": 'white'}),
"end_color": ("STRING", {"default": 'black'}),
"multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 100.0, "step": 0.01}),
},
}
def generate(self, coordinates, frame_width, frame_height, start_color, end_color, multiplier):
# Parse the coordinates
coordinates = json.loads(coordinates.replace("'", '"'))
# Create an image
image = Image.new("RGB", (frame_width, frame_height))
draw = ImageDraw.Draw(image)
# Extract start and end points for the gradient
start_coord = coordinates[0]
end_coord = coordinates[1]
start_color = parse_color(start_color)
end_color = parse_color(end_color)
# Calculate the gradient direction (vector)
gradient_direction = (end_coord['x'] - start_coord['x'], end_coord['y'] - start_coord['y'])
gradient_length = (gradient_direction[0] ** 2 + gradient_direction[1] ** 2) ** 0.5
# Iterate over each pixel in the image
for y in range(frame_height):
for x in range(frame_width):
# Calculate the projection of the point on the gradient line
point_vector = (x - start_coord['x'], y - start_coord['y'])
projection = (point_vector[0] * gradient_direction[0] + point_vector[1] * gradient_direction[1]) / gradient_length
projection = max(min(projection, gradient_length), 0) # Clamp the projection value
# Calculate the blend factor for the current pixel
blend = projection * multiplier / gradient_length
# Determine the color of the current pixel
color = (
int(start_color[0] + (end_color[0] - start_color[0]) * blend),
int(start_color[1] + (end_color[1] - start_color[1]) * blend),
int(start_color[2] + (end_color[2] - start_color[2]) * blend)
)
# Set the pixel color
draw.point((x, y), fill=color)
# Convert the PIL image to a tensor (assuming such a function exists in your context)
image_tensor = pil2tensor(image)
return (image_tensor,)
class GradientToFloat:
RETURN_TYPES = ("FLOAT", "FLOAT",)
RETURN_NAMES = ("float_x", "float_y", )
FUNCTION = "sample"
CATEGORY = "KJNodes/image"
DESCRIPTION = """
Calculates list of floats from image.
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE", ),
"steps": ("INT", {"default": 10, "min": 2, "max": 10000, "step": 1}),
},
}
def sample(self, image, steps):
# Assuming image is a tensor with shape [B, H, W, C]
B, H, W, C = image.shape
# Sample along the width axis (W)
w_intervals = torch.linspace(0, W - 1, steps=steps, dtype=torch.int64)
# Assuming we're sampling from the first batch and the first channel
w_sampled = image[0, :, w_intervals, 0]
# Sample along the height axis (H)
h_intervals = torch.linspace(0, H - 1, steps=steps, dtype=torch.int64)
# Assuming we're sampling from the first batch and the first channel
h_sampled = image[0, h_intervals, :, 0]
# Taking the mean across the height for width sampling, and across the width for height sampling
w_values = w_sampled.mean(dim=0).tolist()
h_values = h_sampled.mean(dim=1).tolist()
return (w_values, h_values)
class MaskOrImageToWeight:
@@ -759,7 +348,7 @@ and returns that as the selected output type.
# Convert mean_values to the specified output_type
if output_type == 'list':
out = mean_values
out = mean_values,
elif output_type == 'pandas series':
try:
import pandas as pd
@@ -1016,25 +605,6 @@ Creates a sigmas tensor from list of float values.
def customsigmas(self, float_list):
return torch.tensor(float_list, dtype=torch.float32),
class SigmasToFloat:
@classmethod
def INPUT_TYPES(s):
return {"required":
{
"sigmas": ("SIGMAS",),
}
}
RETURN_TYPES = ("FLOAT",)
RETURN_NAMES = ("float",)
CATEGORY = "KJNodes/noise"
FUNCTION = "customsigmas"
DESCRIPTION = """
Creates a float list from sigmas tensors.
"""
def customsigmas(self, sigmas):
return sigmas.tolist(),
class GLIGENTextBoxApplyBatchCoords:
@classmethod
def INPUT_TYPES(s):
@@ -1163,10 +733,8 @@ for example:
batch_size = len(coordinates)
# Initialize a list to hold the coordinates for the current ID
id_coordinates = []
if not size_multiplier or len(size_multiplier) != batch_size:
size_multiplier = [0] * batch_size
else:
size_multiplier = size_multiplier * (batch_size // len(size_multiplier)) + size_multiplier[:batch_size % len(size_multiplier)]
if len(size_multiplier) != batch_size:
size_multiplier = size_multiplier * (batch_size // len(size_multiplier)) + size_multiplier[:batch_size % len(size_multiplier)]
for i, coord in enumerate(coordinates):
x = coord['x']
y = coord['y']
@@ -1184,13 +752,6 @@ for example:
top_left_y = max(0, top_left_y)
bottom_right_x = min(width, bottom_right_x)
bottom_right_y = min(height, bottom_right_y)
# Ensure width and height are positive
adjusted_bbox_width = max(1, bottom_right_x - top_left_x)
adjusted_bbox_height = max(1, bottom_right_y - top_left_y)
# Update the coordinates with the new width and height
bottom_right_x = top_left_x + adjusted_bbox_width
bottom_right_y = top_left_y + adjusted_bbox_height
# Append the top left and bottom right coordinates to the list for the current ID
id_coordinates.append([top_left_x, top_left_y, bottom_right_x, bottom_right_y, width, height])
@@ -1270,51 +831,48 @@ Interpolates coordinates based on a curve.
}
def interpolate(self, coordinates, interpolation_curve):
# Parse the JSON string to get the list of coordinates
# Parse the JSON string to get the list of coordinates
coordinates = json.loads(coordinates.replace("'", '"'))
# Convert the list of dictionaries to a list of (x, y) tuples for easier processing
coordinates = [(coord['x'], coord['y']) for coord in coordinates]
# Calculate the total length of the original path
path_length = sum(np.linalg.norm(np.array(coordinates[i]) - np.array(coordinates[i-1]))
for i in range(1, len(coordinates)))
path_length = sum(np.linalg.norm(np.array(coordinates[i]) - np.array(coordinates[i-1])) for i in range(1, len(coordinates)))
# Normalize the interpolation curve
normalized_curve = [x / path_length for x in interpolation_curve]
# Initialize variables for interpolation
interpolated_coords = []
current_length = 0
current_index = 0
current_index = 1
# Iterate over the normalized curve
for normalized_length in interpolation_curve:
target_length = normalized_length * path_length # Convert to the original scale
while current_index < len(coordinates) - 1:
segment_start, segment_end = np.array(coordinates[current_index]), np.array(coordinates[current_index + 1])
segment_length = np.linalg.norm(segment_end - segment_start)
if current_length + segment_length >= target_length:
break
for target_length in normalized_curve:
target_length *= path_length # Convert back to the original scale
while current_length < target_length and current_index < len(coordinates):
segment_length = np.linalg.norm(np.array(coordinates[current_index]) - np.array(coordinates[current_index-1]))
current_length += segment_length
current_index += 1
# Interpolate between the last two points
if current_index < len(coordinates) - 1:
p1, p2 = np.array(coordinates[current_index]), np.array(coordinates[current_index + 1])
if current_index == 1:
interpolated_coords.append(coordinates[0])
else:
p1, p2 = np.array(coordinates[current_index-2]), np.array(coordinates[current_index-1])
segment_length = np.linalg.norm(p2 - p1)
if segment_length > 0:
t = (target_length - current_length) / segment_length
t = (target_length - (current_length - segment_length)) / segment_length
interpolated_point = p1 + t * (p2 - p1)
interpolated_coords.append(interpolated_point.tolist())
else:
interpolated_coords.append(p1.tolist())
else:
# If the target_length is at or beyond the end of the path, add the last coordinate
interpolated_coords.append(coordinates[-1])
# Convert back to string format if necessary
interpolated_coords_str = "[" + ", ".join([f"{{'x': {round(coord[0])}, 'y': {round(coord[1])}}}" for coord in interpolated_coords]) + "]"
print(interpolated_coords_str)
return (interpolated_coords_str,)
return (interpolated_coords_str, )
class DrawInstanceDiffusionTracking:
@@ -1392,245 +950,4 @@ CreateInstanceDiffusionTracking -node.
# Stack the modified images back into a batch
image_tensor_batch = torch.stack(modified_images).cpu().float()
return image_tensor_batch,
class PointsEditor:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"points_store": ("STRING", {"multiline": False}),
"coordinates": ("STRING", {"multiline": False}),
"neg_coordinates": ("STRING", {"multiline": False}),
"bbox_store": ("STRING", {"multiline": False}),
"bboxes": ("STRING", {"multiline": False}),
"bbox_format": (
[
'xyxy',
'xywh',
],
),
"width": ("INT", {"default": 512, "min": 8, "max": 4096, "step": 8}),
"height": ("INT", {"default": 512, "min": 8, "max": 4096, "step": 8}),
"normalize": ("BOOLEAN", {"default": False}),
},
"optional": {
"bg_image": ("IMAGE", ),
},
}
RETURN_TYPES = ("STRING", "STRING", "BBOX", "MASK", "IMAGE")
RETURN_NAMES = ("positive_coords", "negative_coords", "bbox", "bbox_mask", "cropped_image")
FUNCTION = "pointdata"
CATEGORY = "KJNodes/experimental"
DESCRIPTION = """
# WORK IN PROGRESS
Do not count on this as part of your workflow yet,
probably contains lots of bugs and stability is not
guaranteed!!
## Graphical editor to create coordinates
**Shift + click** to add a positive (green) point.
**Shift + right click** to add a negative (red) point.
**Ctrl + click** to draw a box.
**Right click on a point** to delete it.
Note that you can't delete from start/end of the points array.
To add an image select the node and copy/paste or drag in the image.
Or from the bg_image input on queue (first frame of the batch).
**THE IMAGE IS SAVED TO THE NODE AND WORKFLOW METADATA**
you can clear the image from the context menu by right clicking on the canvas
"""
def pointdata(self, points_store, bbox_store, width, height, coordinates, neg_coordinates, normalize, bboxes, bbox_format="xyxy", bg_image=None):
coordinates = json.loads(coordinates)
pos_coordinates = []
for coord in coordinates:
coord['x'] = int(round(coord['x']))
coord['y'] = int(round(coord['y']))
if normalize:
norm_x = coord['x'] / width
norm_y = coord['y'] / height
pos_coordinates.append({'x': norm_x, 'y': norm_y})
else:
pos_coordinates.append({'x': coord['x'], 'y': coord['y']})
if neg_coordinates:
coordinates = json.loads(neg_coordinates)
neg_coordinates = []
for coord in coordinates:
coord['x'] = int(round(coord['x']))
coord['y'] = int(round(coord['y']))
if normalize:
norm_x = coord['x'] / width
norm_y = coord['y'] / height
neg_coordinates.append({'x': norm_x, 'y': norm_y})
else:
neg_coordinates.append({'x': coord['x'], 'y': coord['y']})
# Create a blank mask
mask = np.zeros((height, width), dtype=np.uint8)
bboxes = json.loads(bboxes)
print(bboxes)
valid_bboxes = []
for bbox in bboxes:
if (bbox.get("startX") is None or
bbox.get("startY") is None or
bbox.get("endX") is None or
bbox.get("endY") is None):
continue # Skip this bounding box if any value is None
else:
# Ensure that endX and endY are greater than startX and startY
x_min = min(int(bbox["startX"]), int(bbox["endX"]))
y_min = min(int(bbox["startY"]), int(bbox["endY"]))
x_max = max(int(bbox["startX"]), int(bbox["endX"]))
y_max = max(int(bbox["startY"]), int(bbox["endY"]))
valid_bboxes.append((x_min, y_min, x_max, y_max))
bboxes_xyxy = []
for bbox in valid_bboxes:
x_min, y_min, x_max, y_max = bbox
bboxes_xyxy.append((x_min, y_min, x_max, y_max))
mask[y_min:y_max, x_min:x_max] = 1 # Fill the bounding box area with 1s
if bbox_format == "xywh":
bboxes_xywh = []
for bbox in valid_bboxes:
x_min, y_min, x_max, y_max = bbox
width = x_max - x_min
height = y_max - y_min
bboxes_xywh.append((x_min, y_min, width, height))
bboxes = bboxes_xywh
else:
bboxes = bboxes_xyxy
mask_tensor = torch.from_numpy(mask)
mask_tensor = mask_tensor.unsqueeze(0).float().cpu()
if bg_image is not None and len(valid_bboxes) > 0:
x_min, y_min, x_max, y_max = bboxes[0]
cropped_image = bg_image[:, y_min:y_max, x_min:x_max, :]
elif bg_image is not None:
cropped_image = bg_image
if bg_image is None:
return (json.dumps(pos_coordinates), json.dumps(neg_coordinates), bboxes, mask_tensor)
else:
transform = transforms.ToPILImage()
image = transform(bg_image[0].permute(2, 0, 1))
buffered = io.BytesIO()
image.save(buffered, format="JPEG", quality=75)
# Step 3: Encode the image bytes to a Base64 string
img_bytes = buffered.getvalue()
img_base64 = base64.b64encode(img_bytes).decode('utf-8')
return {
"ui": {"bg_image": [img_base64]},
"result": (json.dumps(pos_coordinates), json.dumps(neg_coordinates), bboxes, mask_tensor, cropped_image)
}
class CutAndDragOnPath:
RETURN_TYPES = ("IMAGE", "MASK",)
RETURN_NAMES = ("image","mask", )
FUNCTION = "cutanddrag"
CATEGORY = "KJNodes/image"
DESCRIPTION = """
Cuts the masked area from the image, and drags it along the path. If inpaint is enabled, and no bg_image is provided, the cut area is filled using cv2 TELEA algorithm.
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"coordinates": ("STRING", {"forceInput": True}),
"mask": ("MASK",),
"frame_width": ("INT", {"default": 512,"min": 16, "max": 4096, "step": 1}),
"frame_height": ("INT", {"default": 512,"min": 16, "max": 4096, "step": 1}),
"inpaint": ("BOOLEAN", {"default": True}),
},
"optional": {
"bg_image": ("IMAGE",),
}
}
def cutanddrag(self, image, coordinates, mask, frame_width, frame_height, inpaint, bg_image=None):
# Parse coordinates
coords_list = parse_json_tracks(coordinates)
batch_size = len(coords_list[0])
images_list = []
masks_list = []
# Convert input image and mask to PIL
input_image = tensor2pil(image)[0]
input_mask = tensor2pil(mask)[0]
# Find masked region bounds
mask_array = np.array(input_mask)
y_indices, x_indices = np.where(mask_array > 0)
if len(x_indices) == 0 or len(y_indices) == 0:
return (image, mask)
x_min, x_max = x_indices.min(), x_indices.max()
y_min, y_max = y_indices.min(), y_indices.max()
# Cut out the masked region
cut_width = x_max - x_min
cut_height = y_max - y_min
cut_image = input_image.crop((x_min, y_min, x_max, y_max))
cut_mask = input_mask.crop((x_min, y_min, x_max, y_max))
# Create inpainted background
if bg_image is None:
background = input_image.copy()
# Inpaint the cut area
if inpaint:
import cv2
border = 5 # Create small border around cut area for better inpainting
fill_mask = Image.new("L", background.size, 0)
draw = ImageDraw.Draw(fill_mask)
draw.rectangle([x_min-border, y_min-border, x_max+border, y_max+border], fill=255)
background = cv2.inpaint(
np.array(background),
np.array(fill_mask),
inpaintRadius=3,
flags=cv2.INPAINT_TELEA
)
background = Image.fromarray(background)
else:
background = tensor2pil(bg_image)[0]
# Create batch of images with cut region at different positions
for i in range(batch_size):
# Create new image
new_image = background.copy()
new_mask = Image.new("L", (frame_width, frame_height), 0)
# Get target position from coordinates
for coords in coords_list:
target_x = int(coords[i]['x'] - cut_width/2)
target_y = int(coords[i]['y'] - cut_height/2)
# Paste cut region at new position
new_image.paste(cut_image, (target_x, target_y), cut_mask)
new_mask.paste(cut_mask, (target_x, target_y))
# Convert to tensor and append
image_tensor = pil2tensor(new_image)
mask_tensor = pil2tensor(new_mask)
images_list.append(image_tensor)
masks_list.append(mask_tensor)
# Stack tensors into batches
out_images = torch.cat(images_list, dim=0).cpu().float()
out_masks = torch.cat(masks_list, dim=0)
return (out_images, out_masks)
return image_tensor_batch,
+173 -2810
View File
File diff suppressed because it is too large Load Diff
+3 -3
View File
@@ -7,7 +7,7 @@ import comfy.sample
from nodes import CLIPTextEncode
script_directory = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
folder_paths.add_model_folder_path("intrinsic_loras", os.path.join(script_directory, "intrinsic_loras"))
folder_paths.add_model_folder_path("intristic_loras", os.path.join(script_directory, "intristic_loras"))
class Intrinsic_lora_sampling:
def __init__(self):
@@ -16,7 +16,7 @@ class Intrinsic_lora_sampling:
@classmethod
def INPUT_TYPES(s):
return {"required": { "model": ("MODEL",),
"lora_name": (folder_paths.get_filename_list("intrinsic_loras"), ),
"lora_name": (folder_paths.get_filename_list("intristic_loras"), ),
"task": (
[
'depth map',
@@ -81,7 +81,7 @@ with this node pack.
#load lora
model_clone = model.clone()
lora_path = folder_paths.get_full_path("intrinsic_loras", lora_name)
lora_path = folder_paths.get_full_path("intristic_loras", lora_name)
lora = load_torch_file(lora_path, safe_load=True)
self.loaded_lora = (lora_path, lora)
-552
View File
@@ -1,552 +0,0 @@
import torch
import comfy.model_management
import comfy.utils
import folder_paths
import os
import logging
from tqdm import tqdm
import numpy as np
device = comfy.model_management.get_torch_device()
CLAMP_QUANTILE = 0.99
def extract_lora(diff, key, rank, algorithm, lora_type, lowrank_iters=7, adaptive_param=1.0, clamp_quantile=True):
"""
Extracts LoRA weights from a weight difference tensor using SVD.
"""
conv2d = (len(diff.shape) == 4)
kernel_size = None if not conv2d else diff.size()[2:4]
conv2d_3x3 = conv2d and kernel_size != (1, 1)
out_dim, in_dim = diff.size()[0:2]
if conv2d:
if conv2d_3x3:
diff = diff.flatten(start_dim=1)
else:
diff = diff.squeeze()
diff_float = diff.float()
if algorithm == "svd_lowrank":
U, S, V = torch.svd_lowrank(diff_float, q=min(rank, in_dim, out_dim), niter=lowrank_iters)
U = U @ torch.diag(S)
Vh = V.t()
else:
#torch.linalg.svdvals()
U, S, Vh = torch.linalg.svd(diff_float)
# Flexible rank selection logic like locon: https://github.com/KohakuBlueleaf/LyCORIS/blob/main/tools/extract_locon.py
if "adaptive" in lora_type:
if lora_type == "adaptive_ratio":
min_s = torch.max(S) * adaptive_param
lora_rank = torch.sum(S > min_s).item()
elif lora_type == "adaptive_energy":
energy = torch.cumsum(S**2, dim=0)
total_energy = torch.sum(S**2)
threshold = adaptive_param * total_energy # e.g., adaptive_param=0.95 for 95%
lora_rank = torch.sum(energy < threshold).item() + 1
elif lora_type == "adaptive_quantile":
s_cum = torch.cumsum(S, dim=0)
min_cum_sum = adaptive_param * torch.sum(S)
lora_rank = torch.sum(s_cum < min_cum_sum).item()
print(f"{key} Extracted LoRA rank: {lora_rank}")
else:
lora_rank = rank
lora_rank = max(1, lora_rank)
lora_rank = min(out_dim, in_dim, lora_rank)
U = U[:, :lora_rank]
S = S[:lora_rank]
U = U @ torch.diag(S)
Vh = Vh[:lora_rank, :]
if clamp_quantile:
dist = torch.cat([U.flatten(), Vh.flatten()])
if dist.numel() > 100_000:
# Sample 100,000 elements for quantile estimation
idx = torch.randperm(dist.numel(), device=dist.device)[:100_000]
dist_sample = dist[idx]
hi_val = torch.quantile(dist_sample, CLAMP_QUANTILE)
else:
hi_val = torch.quantile(dist, CLAMP_QUANTILE)
low_val = -hi_val
U = U.clamp(low_val, hi_val)
Vh = Vh.clamp(low_val, hi_val)
if conv2d:
U = U.reshape(out_dim, lora_rank, 1, 1)
Vh = Vh.reshape(lora_rank, in_dim, kernel_size[0], kernel_size[1])
return (U, Vh)
def calc_lora_model(model_diff, rank, prefix_model, prefix_lora, output_sd, lora_type, algorithm, lowrank_iters, out_dtype, bias_diff=False, adaptive_param=1.0, clamp_quantile=True):
comfy.model_management.load_models_gpu([model_diff], force_patch_weights=True)
model_diff.model.diffusion_model.cpu()
sd = model_diff.model_state_dict(filter_prefix=prefix_model)
del model_diff
comfy.model_management.soft_empty_cache()
for k, v in sd.items():
if isinstance(v, torch.Tensor):
sd[k] = v.cpu()
# Get total number of keys to process for progress bar
total_keys = len([k for k in sd if k.endswith(".weight") or (bias_diff and k.endswith(".bias"))])
# Create progress bar
progress_bar = tqdm(total=total_keys, desc=f"Extracting LoRA ({prefix_lora.strip('.')})")
comfy_pbar = comfy.utils.ProgressBar(total_keys)
for k in sd:
if k.endswith(".weight"):
weight_diff = sd[k]
if weight_diff.ndim == 5:
logging.info(f"Skipping 5D tensor for key {k}") #skip patch embed
progress_bar.update(1)
comfy_pbar.update(1)
continue
if lora_type != "full":
if weight_diff.ndim < 2:
if bias_diff:
output_sd["{}{}.diff".format(prefix_lora, k[len(prefix_model):-7])] = weight_diff.contiguous().to(out_dtype).cpu()
progress_bar.update(1)
comfy_pbar.update(1)
continue
try:
out = extract_lora(weight_diff.to(device), k, rank, algorithm, lora_type, lowrank_iters=lowrank_iters, adaptive_param=adaptive_param, clamp_quantile=clamp_quantile)
output_sd["{}{}.lora_up.weight".format(prefix_lora, k[len(prefix_model):-7])] = out[0].contiguous().to(out_dtype).cpu()
output_sd["{}{}.lora_down.weight".format(prefix_lora, k[len(prefix_model):-7])] = out[1].contiguous().to(out_dtype).cpu()
except Exception as e:
logging.warning(f"Could not generate lora weights for key {k}, error {e}")
else:
output_sd["{}{}.diff".format(prefix_lora, k[len(prefix_model):-7])] = weight_diff.contiguous().to(out_dtype).cpu()
progress_bar.update(1)
comfy_pbar.update(1)
elif bias_diff and k.endswith(".bias"):
output_sd["{}{}.diff_b".format(prefix_lora, k[len(prefix_model):-5])] = sd[k].contiguous().to(out_dtype).cpu()
progress_bar.update(1)
comfy_pbar.update(1)
progress_bar.close()
return output_sd
class LoraExtractKJ:
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
@classmethod
def INPUT_TYPES(s):
return {"required":
{
"finetuned_model": ("MODEL",),
"original_model": ("MODEL",),
"filename_prefix": ("STRING", {"default": "loras/ComfyUI_extracted_lora"}),
"rank": ("INT", {"default": 8, "min": 1, "max": 4096, "step": 1}),
"lora_type": (["standard", "full", "adaptive_ratio", "adaptive_quantile", "adaptive_energy"],),
"algorithm": (["svd_linalg", "svd_lowrank"], {"default": "svd_linalg", "tooltip": "SVD algorithm to use, svd_lowrank is faster but less accurate."}),
"lowrank_iters": ("INT", {"default": 7, "min": 1, "max": 100, "step": 1, "tooltip": "The number of subspace iterations for lowrank SVD algorithm."}),
"output_dtype": (["fp16", "bf16", "fp32"], {"default": "fp16"}),
"bias_diff": ("BOOLEAN", {"default": True}),
"adaptive_param": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "For ratio mode, this is the ratio of the maximum singular value. For quantile mode, this is the quantile of the singular values."}),
"clamp_quantile": ("BOOLEAN", {"default": True}),
},
}
RETURN_TYPES = ()
FUNCTION = "save"
OUTPUT_NODE = True
CATEGORY = "KJNodes/lora"
def save(self, finetuned_model, original_model, filename_prefix, rank, lora_type, algorithm, lowrank_iters, output_dtype, bias_diff, adaptive_param, clamp_quantile):
if algorithm == "svd_lowrank" and lora_type != "standard":
raise ValueError("svd_lowrank algorithm is only supported for standard LoRA extraction.")
dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[output_dtype]
m = finetuned_model.clone()
kp = original_model.get_key_patches("diffusion_model.")
for k in kp:
m.add_patches({k: kp[k]}, - 1.0, 1.0)
model_diff = m
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir)
output_sd = {}
if model_diff is not None:
output_sd = calc_lora_model(model_diff, rank, "diffusion_model.", "diffusion_model.", output_sd, lora_type, algorithm, lowrank_iters, dtype, bias_diff=bias_diff, adaptive_param=adaptive_param, clamp_quantile=clamp_quantile)
if "adaptive" in lora_type:
rank_str = f"{lora_type}_{adaptive_param:.2f}"
else:
rank_str = rank
output_checkpoint = f"{filename}_rank_{rank_str}_{output_dtype}_{counter:05}_.safetensors"
output_checkpoint = os.path.join(full_output_folder, output_checkpoint)
comfy.utils.save_torch_file(output_sd, output_checkpoint, metadata=None)
return {}
NODE_CLASS_MAPPINGS = {
"LoraExtractKJ": LoraExtractKJ
}
NODE_DISPLAY_NAME_MAPPINGS = {
"LoraExtractKJ": "LoraExtractKJ"
}
class LoraReduceRank:
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
@classmethod
def INPUT_TYPES(s):
return {"required":
{
"lora_name": (folder_paths.get_filename_list("loras"), {"tooltip": "The name of the LoRA."}),
"new_rank": ("INT", {"default": 8, "min": 1, "max": 4096, "step": 1, "tooltip": "The new rank to resize the LoRA. Acts as max rank when using dynamic_method."}),
"dynamic_method": (["disabled", "sv_ratio", "sv_cumulative", "sv_fro"], {"default": "disabled", "tooltip": "Method to use for dynamically determining new alphas and dims"}),
"dynamic_param": ("FLOAT", {"default": 0.2, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Method to use for dynamically determining new alphas and dims"}),
"output_dtype": (["match_original", "fp16", "bf16", "fp32"], {"default": "match_original", "tooltip": "Data type to save the LoRA as."}),
"verbose": ("BOOLEAN", {"default": True}),
},
}
RETURN_TYPES = ()
FUNCTION = "save"
OUTPUT_NODE = True
EXPERIMENTAL = True
DESCRIPTION = "Resize a LoRA model by reducing it's rank. Based on kohya's sd-scripts: https://github.com/kohya-ss/sd-scripts/blob/main/networks/resize_lora.py"
CATEGORY = "KJNodes/lora"
def save(self, lora_name, new_rank, output_dtype, dynamic_method, dynamic_param, verbose):
lora_path = folder_paths.get_full_path("loras", lora_name)
lora_sd, metadata = comfy.utils.load_torch_file(lora_path, return_metadata=True)
if output_dtype == "fp16":
save_dtype = torch.float16
elif output_dtype == "bf16":
save_dtype = torch.bfloat16
elif output_dtype == "fp32":
save_dtype = torch.float32
elif output_dtype == "match_original":
first_weight_key = next(k for k in lora_sd if k.endswith(".weight") and isinstance(lora_sd[k], torch.Tensor))
save_dtype = lora_sd[first_weight_key].dtype
new_lora_sd = {}
for k, v in lora_sd.items():
new_lora_sd[k.replace(".default", "")] = v
del lora_sd
print("Resizing Lora...")
output_sd, old_dim, new_alpha, rank_list = resize_lora_model(new_lora_sd, new_rank, save_dtype, device, dynamic_method, dynamic_param, verbose)
# update metadata
if metadata is None:
metadata = {}
comment = metadata.get("ss_training_comment", "")
if dynamic_method == "disabled":
metadata["ss_training_comment"] = f"dimension is resized from {old_dim} to {new_rank}; {comment}"
metadata["ss_network_dim"] = str(new_rank)
metadata["ss_network_alpha"] = str(new_alpha)
else:
metadata["ss_training_comment"] = f"Dynamic resize with {dynamic_method}: {dynamic_param} from {old_dim}; {comment}"
metadata["ss_network_dim"] = "Dynamic"
metadata["ss_network_alpha"] = "Dynamic"
# cast to save_dtype before calculating hashes
for key in list(output_sd.keys()):
value = output_sd[key]
if type(value) == torch.Tensor and value.dtype.is_floating_point and value.dtype != save_dtype:
output_sd[key] = value.to(save_dtype)
output_filename_prefix = "loras/" + lora_name
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(output_filename_prefix, self.output_dir)
output_dtype_str = f"_{output_dtype}" if output_dtype != "match_original" else ""
average_rank = str(int(np.mean(rank_list)))
rank_str = new_rank if dynamic_method == "disabled" else f"dynamic_{average_rank}"
output_checkpoint = f"{filename.replace('.safetensors', '')}_resized_from_{old_dim}_to_{rank_str}{output_dtype_str}_{counter:05}_.safetensors"
output_checkpoint = os.path.join(full_output_folder, output_checkpoint)
print(f"Saving resized LoRA to {output_checkpoint}")
comfy.utils.save_torch_file(output_sd, output_checkpoint, metadata=metadata)
return {}
NODE_CLASS_MAPPINGS = {
"LoraExtractKJ": LoraExtractKJ
}
NODE_DISPLAY_NAME_MAPPINGS = {
"LoraExtractKJ": "LoraExtractKJ"
}
# Convert LoRA to different rank approximation (should only be used to go to lower rank)
# This code is based off the extract_lora_from_models.py file which is based on https://github.com/cloneofsimo/lora/blob/develop/lora_diffusion/cli_svd.py
# Thanks to cloneofsimo
# This version is based on
# https://github.com/kohya-ss/sd-scripts/blob/main/networks/resize_lora.py
MIN_SV = 1e-6
LORA_DOWN_UP_FORMATS = [
("lora_down", "lora_up"), # sd-scripts LoRA
("lora_A", "lora_B"), # PEFT LoRA
("down", "up"), # ControlLoRA
]
# Indexing functions
def index_sv_cumulative(S, target):
original_sum = float(torch.sum(S))
cumulative_sums = torch.cumsum(S, dim=0) / original_sum
index = int(torch.searchsorted(cumulative_sums, target)) + 1
index = max(1, min(index, len(S) - 1))
return index
def index_sv_fro(S, target):
S_squared = S.pow(2)
S_fro_sq = float(torch.sum(S_squared))
sum_S_squared = torch.cumsum(S_squared, dim=0) / S_fro_sq
index = int(torch.searchsorted(sum_S_squared, target**2)) + 1
index = max(1, min(index, len(S) - 1))
return index
def index_sv_ratio(S, target):
max_sv = S[0]
min_sv = max_sv / target
index = int(torch.sum(S > min_sv).item())
index = max(1, min(index, len(S) - 1))
return index
# Modified from Kohaku-blueleaf's extract/merge functions
def extract_conv(weight, lora_rank, dynamic_method, dynamic_param, device, scale=1):
out_size, in_size, kernel_size, _ = weight.size()
if weight.dtype != torch.float32:
weight = weight.to(torch.float32)
U, S, Vh = torch.linalg.svd(weight.reshape(out_size, -1).to(device))
param_dict = rank_resize(S, lora_rank, dynamic_method, dynamic_param, scale)
lora_rank = param_dict["new_rank"]
U = U[:, :lora_rank]
S = S[:lora_rank]
U = U @ torch.diag(S)
Vh = Vh[:lora_rank, :]
param_dict["lora_down"] = Vh.reshape(lora_rank, in_size, kernel_size, kernel_size).cpu()
param_dict["lora_up"] = U.reshape(out_size, lora_rank, 1, 1).cpu()
del U, S, Vh, weight
return param_dict
def extract_linear(weight, lora_rank, dynamic_method, dynamic_param, device, scale=1):
out_size, in_size = weight.size()
if weight.dtype != torch.float32:
weight = weight.to(torch.float32)
U, S, Vh = torch.linalg.svd(weight.to(device))
param_dict = rank_resize(S, lora_rank, dynamic_method, dynamic_param, scale)
lora_rank = param_dict["new_rank"]
U = U[:, :lora_rank]
S = S[:lora_rank]
U = U @ torch.diag(S)
Vh = Vh[:lora_rank, :]
param_dict["lora_down"] = Vh.reshape(lora_rank, in_size).cpu()
param_dict["lora_up"] = U.reshape(out_size, lora_rank).cpu()
del U, S, Vh, weight
return param_dict
def merge_conv(lora_down, lora_up, device):
in_rank, in_size, kernel_size, k_ = lora_down.shape
out_size, out_rank, _, _ = lora_up.shape
assert in_rank == out_rank and kernel_size == k_, f"rank {in_rank} {out_rank} or kernel {kernel_size} {k_} mismatch"
lora_down = lora_down.to(device)
lora_up = lora_up.to(device)
merged = lora_up.reshape(out_size, -1) @ lora_down.reshape(in_rank, -1)
weight = merged.reshape(out_size, in_size, kernel_size, kernel_size)
del lora_up, lora_down
return weight
def merge_linear(lora_down, lora_up, device):
in_rank, in_size = lora_down.shape
out_size, out_rank = lora_up.shape
assert in_rank == out_rank, f"rank {in_rank} {out_rank} mismatch"
lora_down = lora_down.to(device)
lora_up = lora_up.to(device)
weight = lora_up @ lora_down
del lora_up, lora_down
return weight
# Calculate new rank
def rank_resize(S, rank, dynamic_method, dynamic_param, scale=1):
param_dict = {}
if dynamic_method == "sv_ratio":
# Calculate new dim and alpha based off ratio
new_rank = index_sv_ratio(S, dynamic_param) + 1
new_alpha = float(scale * new_rank)
elif dynamic_method == "sv_cumulative":
# Calculate new dim and alpha based off cumulative sum
new_rank = index_sv_cumulative(S, dynamic_param) + 1
new_alpha = float(scale * new_rank)
elif dynamic_method == "sv_fro":
# Calculate new dim and alpha based off sqrt sum of squares
new_rank = index_sv_fro(S, dynamic_param) + 1
new_alpha = float(scale * new_rank)
else:
new_rank = rank
new_alpha = float(scale * new_rank)
if S[0] <= MIN_SV: # Zero matrix, set dim to 1
new_rank = 1
new_alpha = float(scale * new_rank)
elif new_rank > rank: # cap max rank at rank
new_rank = rank
new_alpha = float(scale * new_rank)
# Calculate resize info
s_sum = torch.sum(torch.abs(S))
s_rank = torch.sum(torch.abs(S[:new_rank]))
S_squared = S.pow(2)
s_fro = torch.sqrt(torch.sum(S_squared))
s_red_fro = torch.sqrt(torch.sum(S_squared[:new_rank]))
fro_percent = float(s_red_fro / s_fro)
param_dict["new_rank"] = new_rank
param_dict["new_alpha"] = new_alpha
param_dict["sum_retained"] = (s_rank) / s_sum
param_dict["fro_retained"] = fro_percent
param_dict["max_ratio"] = S[0] / S[new_rank - 1]
return param_dict
def resize_lora_model(lora_sd, new_rank, save_dtype, device, dynamic_method, dynamic_param, verbose):
max_old_rank = None
new_alpha = None
verbose_str = "\n"
fro_list = []
rank_list = []
if dynamic_method:
print(f"Dynamically determining new alphas and dims based off {dynamic_method}: {dynamic_param}, max rank is {new_rank}")
lora_down_weight = None
lora_up_weight = None
o_lora_sd = lora_sd.copy()
block_down_name = None
block_up_name = None
total_keys = len([k for k in lora_sd if k.endswith(".weight")])
pbar = comfy.utils.ProgressBar(total_keys)
for key, value in tqdm(lora_sd.items()):
key_parts = key.split(".")
block_down_name = None
for _format in LORA_DOWN_UP_FORMATS:
# Currently we only match lora_down_name in the last two parts of key
# because ("down", "up") are general words and may appear in block_down_name
if len(key_parts) >= 2 and _format[0] == key_parts[-2]:
block_down_name = ".".join(key_parts[:-2])
lora_down_name = "." + _format[0]
lora_up_name = "." + _format[1]
weight_name = "." + key_parts[-1]
break
if len(key_parts) >= 1 and _format[0] == key_parts[-1]:
block_down_name = ".".join(key_parts[:-1])
lora_down_name = "." + _format[0]
lora_up_name = "." + _format[1]
weight_name = ""
break
if block_down_name is None:
# This parameter is not lora_down
continue
# Now weight_name can be ".weight" or ""
# Find corresponding lora_up and alpha
block_up_name = block_down_name
lora_down_weight = value
lora_up_weight = lora_sd.get(block_up_name + lora_up_name + weight_name, None)
lora_alpha = lora_sd.get(block_down_name + ".alpha", None)
weights_loaded = lora_down_weight is not None and lora_up_weight is not None
if weights_loaded:
conv2d = len(lora_down_weight.size()) == 4
old_rank = lora_down_weight.size()[0]
max_old_rank = max(max_old_rank or 0, old_rank)
if lora_alpha is None:
scale = 1.0
else:
scale = lora_alpha / old_rank
if conv2d:
full_weight_matrix = merge_conv(lora_down_weight, lora_up_weight, device)
param_dict = extract_conv(full_weight_matrix, new_rank, dynamic_method, dynamic_param, device, scale)
else:
full_weight_matrix = merge_linear(lora_down_weight, lora_up_weight, device)
param_dict = extract_linear(full_weight_matrix, new_rank, dynamic_method, dynamic_param, device, scale)
if verbose:
max_ratio = param_dict["max_ratio"]
sum_retained = param_dict["sum_retained"]
fro_retained = param_dict["fro_retained"]
if not np.isnan(fro_retained):
fro_list.append(float(fro_retained))
verbose_str += f"{block_down_name:75} | "
verbose_str += f"sum(S) retained: {sum_retained:.1%}, fro retained: {fro_retained:.1%}, max(S) ratio: {max_ratio:0.1f}"
print(verbose_str)
if verbose and dynamic_method:
verbose_str += f", dynamic | dim: {param_dict['new_rank']}, alpha: {param_dict['new_alpha']}\n"
else:
verbose_str += "\n"
new_alpha = param_dict["new_alpha"]
o_lora_sd[block_down_name + lora_down_name + weight_name] = param_dict["lora_down"].to(save_dtype).contiguous()
o_lora_sd[block_up_name + lora_up_name + weight_name] = param_dict["lora_up"].to(save_dtype).contiguous()
o_lora_sd[block_down_name + ".alpha"] = torch.tensor(param_dict["new_alpha"]).to(save_dtype)
block_down_name = None
block_up_name = None
lora_down_weight = None
lora_up_weight = None
weights_loaded = False
rank_list.append(param_dict["new_rank"])
del param_dict
pbar.update(1)
if verbose:
print(verbose_str)
print(f"Average Frobenius norm retention: {np.mean(fro_list):.2%} | std: {np.std(fro_list):0.3f}")
return o_lora_sd, max_old_rank, new_alpha, rank_list
+74 -335
View File
@@ -4,12 +4,13 @@ from torchvision.transforms import functional as TF
from PIL import Image, ImageDraw, ImageFilter, ImageFont
import scipy.ndimage
import numpy as np
import matplotlib.pyplot as plt
from contextlib import nullcontext
import os
from comfy import model_management
import model_management
from comfy.utils import ProgressBar
from comfy.utils import common_upscale
from nodes import MAX_RESOLUTION
import folder_paths
@@ -30,155 +31,73 @@ class BatchCLIPSeg:
{
"images": ("IMAGE",),
"text": ("STRING", {"multiline": False}),
"threshold": ("FLOAT", {"default": 0.5,"min": 0.0, "max": 10.0, "step": 0.001}),
"threshold": ("FLOAT", {"default": 0.1,"min": 0.0, "max": 10.0, "step": 0.001}),
"binary_mask": ("BOOLEAN", {"default": True}),
"combine_mask": ("BOOLEAN", {"default": False}),
"use_cuda": ("BOOLEAN", {"default": True}),
},
"optional":
{
"blur_sigma": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100.0, "step": 0.1}),
"opt_model": ("CLIPSEGMODEL", ),
"prev_mask": ("MASK", {"default": None}),
"image_bg_level": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
"invert": ("BOOLEAN", {"default": False}),
}
}
CATEGORY = "KJNodes/masking"
RETURN_TYPES = ("MASK", "IMAGE", )
RETURN_NAMES = ("Mask", "Image", )
RETURN_TYPES = ("MASK",)
RETURN_NAMES = ("Mask",)
FUNCTION = "segment_image"
DESCRIPTION = """
Segments an image or batch of images using CLIPSeg.
"""
def segment_image(self, images, text, threshold, binary_mask, combine_mask, use_cuda, blur_sigma=0.0, opt_model=None, prev_mask=None, invert= False, image_bg_level=0.5):
def segment_image(self, images, text, threshold, binary_mask, combine_mask, use_cuda):
from transformers import CLIPSegProcessor, CLIPSegForImageSegmentation
import torchvision.transforms as transforms
offload_device = model_management.unet_offload_device()
device = model_management.get_torch_device()
if not use_cuda:
out = []
height, width, _ = images[0].shape
if use_cuda and torch.cuda.is_available():
device = torch.device("cuda")
else:
device = torch.device("cpu")
dtype = model_management.unet_dtype()
if opt_model is None:
checkpoint_path = os.path.join(folder_paths.models_dir,'clip_seg', 'clipseg-rd64-refined-fp16')
if not hasattr(self, "model"):
try:
if not os.path.exists(checkpoint_path):
from huggingface_hub import snapshot_download
snapshot_download(repo_id="Kijai/clipseg-rd64-refined-fp16", local_dir=checkpoint_path, local_dir_use_symlinks=False)
self.model = CLIPSegForImageSegmentation.from_pretrained(checkpoint_path)
except:
checkpoint_path = "CIDAS/clipseg-rd64-refined"
self.model = CLIPSegForImageSegmentation.from_pretrained(checkpoint_path)
processor = CLIPSegProcessor.from_pretrained(checkpoint_path)
else:
self.model = opt_model['model']
processor = opt_model['processor']
self.model.to(dtype).to(device)
B, H, W, C = images.shape
model = CLIPSegForImageSegmentation.from_pretrained("CIDAS/clipseg-rd64-refined")
model.to(dtype)
model.to(device)
images = images.to(device)
processor = CLIPSegProcessor.from_pretrained("CIDAS/clipseg-rd64-refined")
pbar = ProgressBar(images.shape[0])
autocast_condition = (dtype != torch.float32) and not model_management.is_device_mps(device)
with torch.autocast(model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
for image in images:
image = (image* 255).type(torch.uint8)
prompt = text
input_prc = processor(text=prompt, images=image, return_tensors="pt")
# Move the processed input to the device
for key in input_prc:
input_prc[key] = input_prc[key].to(device)
outputs = model(**input_prc)
tensor = torch.sigmoid(outputs[0])
tensor_thresholded = torch.where(tensor > threshold, tensor, torch.tensor(0, dtype=torch.float))
tensor_normalized = (tensor_thresholded - tensor_thresholded.min()) / (tensor_thresholded.max() - tensor_thresholded.min())
tensor = tensor_normalized
PIL_images = [Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)) for image in images ]
prompt = [text] * len(images)
input_prc = processor(text=prompt, images=PIL_images, return_tensors="pt")
# Resize the mask
if len(tensor.shape) == 3:
tensor = tensor.unsqueeze(0)
resized_tensor = F.interpolate(tensor, size=(height, width), mode='nearest')
for key in input_prc:
input_prc[key] = input_prc[key].to(device)
outputs = self.model(**input_prc)
mask_tensor = torch.sigmoid(outputs.logits)
mask_tensor = (mask_tensor - mask_tensor.min()) / (mask_tensor.max() - mask_tensor.min())
mask_tensor = torch.where(mask_tensor > (threshold), mask_tensor, torch.tensor(0, dtype=torch.float))
print(mask_tensor.shape)
if len(mask_tensor.shape) == 2:
mask_tensor = mask_tensor.unsqueeze(0)
mask_tensor = F.interpolate(mask_tensor.unsqueeze(1), size=(H, W), mode='nearest')
mask_tensor = mask_tensor.squeeze(1)
self.model.to(offload_device)
# Remove the extra dimensions
resized_tensor = resized_tensor[0, 0, :, :]
pbar.update(1)
out.append(resized_tensor)
results = torch.stack(out).cpu().float()
if binary_mask:
mask_tensor = (mask_tensor > 0).float()
if blur_sigma > 0:
kernel_size = int(6 * int(blur_sigma) + 1)
blur = transforms.GaussianBlur(kernel_size=(kernel_size, kernel_size), sigma=(blur_sigma, blur_sigma))
mask_tensor = blur(mask_tensor)
if combine_mask:
mask_tensor = torch.max(mask_tensor, dim=0)[0]
mask_tensor = mask_tensor.unsqueeze(0).repeat(len(images),1,1)
combined_results = torch.max(results, dim=0)[0]
results = combined_results.unsqueeze(0).repeat(len(images),1,1)
del outputs
model_management.soft_empty_cache()
if prev_mask is not None:
if prev_mask.shape != mask_tensor.shape:
prev_mask = F.interpolate(prev_mask.unsqueeze(1), size=(H, W), mode='nearest')
mask_tensor = mask_tensor + prev_mask.to(device)
torch.clamp(mask_tensor, min=0.0, max=1.0)
if invert:
mask_tensor = 1 - mask_tensor
image_tensor = images * mask_tensor.unsqueeze(-1) + (1 - mask_tensor.unsqueeze(-1)) * image_bg_level
image_tensor = torch.clamp(image_tensor, min=0.0, max=1.0).cpu().float()
mask_tensor = mask_tensor.cpu().float()
return mask_tensor, image_tensor,
class DownloadAndLoadCLIPSeg:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {"required":
{
"model": (
[ 'Kijai/clipseg-rd64-refined-fp16',
'CIDAS/clipseg-rd64-refined',
],
),
},
}
CATEGORY = "KJNodes/masking"
RETURN_TYPES = ("CLIPSEGMODEL",)
RETURN_NAMES = ("clipseg_model",)
FUNCTION = "segment_image"
DESCRIPTION = """
Downloads and loads CLIPSeg model with huggingface_hub,
to ComfyUI/models/clip_seg
"""
def segment_image(self, model):
from transformers import CLIPSegProcessor, CLIPSegForImageSegmentation
checkpoint_path = os.path.join(folder_paths.models_dir,'clip_seg', os.path.basename(model))
if not hasattr(self, "model"):
if not os.path.exists(checkpoint_path):
from huggingface_hub import snapshot_download
snapshot_download(repo_id=model, local_dir=checkpoint_path, local_dir_use_symlinks=False)
self.model = CLIPSegForImageSegmentation.from_pretrained(checkpoint_path)
processor = CLIPSegProcessor.from_pretrained(checkpoint_path)
clipseg_model = {}
clipseg_model['model'] = self.model
clipseg_model['processor'] = processor
return clipseg_model,
if binary_mask:
results = results.round()
return results,
class CreateTextMask:
@@ -358,7 +277,7 @@ class CreateFluidMask:
return {
"required": {
"invert": ("BOOLEAN", {"default": False}),
"frames": ("INT", {"default": 1,"min": 1, "max": 4096, "step": 1}),
"frames": ("INT", {"default": 0,"min": 0, "max": 255, "step": 1}),
"width": ("INT", {"default": 256,"min": 16, "max": 4096, "step": 1}),
"height": ("INT", {"default": 256,"min": 16, "max": 4096, "step": 1}),
"inflow_count": ("INT", {"default": 3,"min": 0, "max": 255, "step": 1}),
@@ -371,10 +290,7 @@ class CreateFluidMask:
#using code from https://github.com/GregTJ/stable-fluids
def createfluidmask(self, frames, width, height, invert, inflow_count, inflow_velocity, inflow_radius, inflow_padding, inflow_duration):
from ..utility.fluid import Fluid
try:
from scipy.special import erf
except:
from scipy.spatial import erf
from scipy.spatial import erf
out = []
masks = []
RESOLUTION = width, height
@@ -528,7 +444,7 @@ class CreateFadeMask:
return {
"required": {
"invert": ("BOOLEAN", {"default": False}),
"frames": ("INT", {"default": 2,"min": 2, "max": 10000, "step": 1}),
"frames": ("INT", {"default": 2,"min": 2, "max": 255, "step": 1}),
"width": ("INT", {"default": 256,"min": 16, "max": 4096, "step": 1}),
"height": ("INT", {"default": 256,"min": 16, "max": 4096, "step": 1}),
"interpolation": (["linear", "ease_in", "ease_out", "ease_in_out"],),
@@ -615,10 +531,10 @@ and interpolating from that to fully black at the 16th frame.
"required": {
"points_string": ("STRING", {"default": "0:(0.0),\n7:(1.0),\n15:(0.0)\n", "multiline": True}),
"invert": ("BOOLEAN", {"default": False}),
"frames": ("INT", {"default": 16,"min": 2, "max": 10000, "step": 1}),
"frames": ("INT", {"default": 16,"min": 2, "max": 255, "step": 1}),
"width": ("INT", {"default": 512,"min": 1, "max": 4096, "step": 1}),
"height": ("INT", {"default": 512,"min": 1, "max": 4096, "step": 1}),
"interpolation": (["linear", "ease_in", "ease_out", "ease_in_out", "none", "default_to_black"],),
"interpolation": (["linear", "ease_in", "ease_out", "ease_in_out"],),
},
}
@@ -642,7 +558,7 @@ and interpolating from that to fully black at the 16th frame.
points.append((frame, color))
# Check if the last frame is already in the points
if (interpolation != "default_to_black") and (len(points) == 0 or points[-1][0] != frames - 1):
if len(points) == 0 or points[-1][0] != frames - 1:
# If not, add it with the color of the last specified frame
points.append((frames - 1, points[-1][1] if points else 0))
@@ -662,39 +578,17 @@ and interpolating from that to fully black at the 16th frame.
# Interpolate between the previous point and the next point
prev_point = next_point - 1
t = (i - points[prev_point][0]) / (points[next_point][0] - points[prev_point][0])
if interpolation == "ease_in":
t = ease_in(t)
elif interpolation == "ease_out":
t = ease_out(t)
elif interpolation == "ease_in_out":
t = ease_in_out(t)
elif interpolation == "linear":
pass # No need to modify `t` for linear interpolation
if interpolation == "none":
exact_match = False
for p in points:
if p[0] == i: # Exact frame match
color = p[1]
exact_match = True
break
if not exact_match:
color = points[prev_point][1]
elif interpolation == "default_to_black":
exact_match = False
for p in points:
if p[0] == i: # Exact frame match
color = p[1]
exact_match = True
break
if not exact_match:
color = 0
else:
t = (i - points[prev_point][0]) / (points[next_point][0] - points[prev_point][0])
if interpolation == "ease_in":
t = ease_in(t)
elif interpolation == "ease_out":
t = ease_out(t)
elif interpolation == "ease_in_out":
t = ease_in_out(t)
elif interpolation == "linear":
pass # No need to modify `t` for linear interpolation
color = points[prev_point][1] - t * (points[prev_point][1] - points[next_point][1])
color = points[prev_point][1] - t * (points[prev_point][1] - points[next_point][1])
color = np.clip(color, 0, 255)
image = np.full((height, width), color, dtype=np.float32)
image_batch[i] = image
@@ -730,7 +624,6 @@ class CreateMagicMask:
def createmagicmask(self, frames, transitions, depth, distortion, seed, frame_width, frame_height):
from ..utility.magictex import coordinate_grid, random_transform, magic
import matplotlib.pyplot as plt
rng = np.random.default_rng(seed)
out = []
coords = coordinate_grid((frame_width, frame_height))
@@ -1010,7 +903,7 @@ class GrowMaskWithBlur:
previous_output = None
current_expand = expand
for m in growmask:
output = m.numpy().astype(np.float32)
output = m.numpy()
for _ in range(abs(round(current_expand))):
if current_expand < 0:
output = scipy.ndimage.grey_erosion(output, footprint=kernel)
@@ -1202,17 +1095,14 @@ Rounds the mask or batch of masks to a binary mask.
return (mask,)
class ResizeMask:
upscale_methods = ["nearest-exact", "bilinear", "area", "bicubic", "lanczos"]
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mask": ("MASK",),
"width": ("INT", { "default": 512, "min": 0, "max": MAX_RESOLUTION, "step": 1, "display": "number" }),
"height": ("INT", { "default": 512, "min": 0, "max": MAX_RESOLUTION, "step": 1, "display": "number" }),
"width": ("INT", { "default": 512, "min": 0, "max": MAX_RESOLUTION, "step": 8, "display": "number" }),
"height": ("INT", { "default": 512, "min": 0, "max": MAX_RESOLUTION, "step": 8, "display": "number" }),
"keep_proportions": ("BOOLEAN", { "default": False }),
"upscale_method": (s.upscale_methods,),
"crop": (["disabled","center"],),
}
}
@@ -1224,21 +1114,20 @@ class ResizeMask:
Resizes the mask or batch of masks to the specified width and height.
"""
def resize(self, mask, width, height, keep_proportions, upscale_method,crop):
def resize(self, mask, width, height, keep_proportions):
if keep_proportions:
_, oh, ow = mask.shape
_, oh, ow, _ = mask.shape
width = ow if width == 0 else width
height = oh if height == 0 else height
ratio = min(width / ow, height / oh)
width = round(ow*ratio)
height = round(oh*ratio)
outputs = mask.unsqueeze(0) # Add an extra dimension for batch size
outputs = F.interpolate(outputs, size=(height, width), mode="nearest")
outputs = outputs.squeeze(0) # Remove the extra dimension after interpolation
if upscale_method == "lanczos":
out_mask = common_upscale(mask.unsqueeze(1).repeat(1, 3, 1, 1), width, height, upscale_method, crop=crop).movedim(1,-1)[:, :, :, 0]
else:
out_mask = common_upscale(mask.unsqueeze(1), width, height, upscale_method, crop=crop).squeeze(1)
return(out_mask, out_mask.shape[2], out_mask.shape[1],)
return(outputs, outputs.shape[2], outputs.shape[1],)
class RemapMaskRange:
@classmethod
@@ -1274,154 +1163,4 @@ Sets new min and max values for the mask.
# Clamp the values to ensure they are within [0.0, 1.0]
scaled_mask = torch.clamp(scaled_mask, min=0.0, max=1.0)
return (scaled_mask, )
def get_mask_polygon(self, mask_np):
import cv2
"""Helper function to get polygon points from mask"""
# Find contours
contours, _ = cv2.findContours(mask_np, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
if not contours:
return None
# Get the largest contour
largest_contour = max(contours, key=cv2.contourArea)
# Approximate polygon
epsilon = 0.02 * cv2.arcLength(largest_contour, True)
polygon = cv2.approxPolyDP(largest_contour, epsilon, True)
return polygon.squeeze()
import cv2
class SeparateMasks:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"mask": ("MASK", ),
"size_threshold_width" : ("INT", {"default": 256, "min": 0.0, "max": 4096, "step": 1}),
"size_threshold_height" : ("INT", {"default": 256, "min": 0.0, "max": 4096, "step": 1}),
"mode": (["convex_polygons", "area", "box"],),
"max_poly_points": ("INT", {"default": 8, "min": 3, "max": 32, "step": 1}),
},
}
RETURN_TYPES = ("MASK",)
RETURN_NAMES = ("mask",)
FUNCTION = "separate"
CATEGORY = "KJNodes/masking"
OUTPUT_NODE = True
DESCRIPTION = "Separates a mask into multiple masks based on the size of the connected components."
def polygon_to_mask(self, polygon, shape):
mask = np.zeros((shape[0], shape[1]), dtype=np.uint8) # Fixed shape handling
if len(polygon.shape) == 2: # Check if polygon points are valid
polygon = polygon.astype(np.int32)
cv2.fillPoly(mask, [polygon], 1)
return mask
def get_mask_polygon(self, mask_np, max_points):
contours, _ = cv2.findContours(mask_np, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
if not contours:
return None
largest_contour = max(contours, key=cv2.contourArea)
hull = cv2.convexHull(largest_contour)
# Initialize with smaller epsilon for more points
perimeter = cv2.arcLength(hull, True)
epsilon = perimeter * 0.01 # Start smaller
min_eps = perimeter * 0.001 # Much smaller minimum
max_eps = perimeter * 0.2 # Smaller maximum
best_approx = None
best_diff = float('inf')
max_iterations = 20
#print(f"Target points: {max_points}, Perimeter: {perimeter}")
for i in range(max_iterations):
curr_eps = (min_eps + max_eps) / 2
approx = cv2.approxPolyDP(hull, curr_eps, True)
points_diff = len(approx) - max_points
#print(f"Iteration {i}: points={len(approx)}, eps={curr_eps:.4f}")
if abs(points_diff) < best_diff:
best_approx = approx
best_diff = abs(points_diff)
if len(approx) > max_points:
min_eps = curr_eps * 1.1 # More gradual adjustment
elif len(approx) < max_points:
max_eps = curr_eps * 0.9 # More gradual adjustment
else:
return approx.squeeze()
if abs(max_eps - min_eps) < perimeter * 0.0001: # Relative tolerance
break
# If we didn't find exact match, return best approximation
return best_approx.squeeze() if best_approx is not None else hull.squeeze()
def separate(self, mask: torch.Tensor, size_threshold_width: int, size_threshold_height: int, max_poly_points: int, mode: str):
from scipy.ndimage import label, center_of_mass
import numpy as np
B, H, W = mask.shape
separated = []
mask = mask.round()
for b in range(B):
mask_np = mask[b].cpu().numpy().astype(np.uint8)
structure = np.ones((3, 3), dtype=np.int8)
labeled, ncomponents = label(mask_np, structure=structure)
pbar = ProgressBar(ncomponents)
for component in range(1, ncomponents + 1):
component_mask_np = (labeled == component).astype(np.uint8)
rows = np.any(component_mask_np, axis=1)
cols = np.any(component_mask_np, axis=0)
y_min, y_max = np.where(rows)[0][[0, -1]]
x_min, x_max = np.where(cols)[0][[0, -1]]
width = x_max - x_min + 1
height = y_max - y_min + 1
centroid_x = (x_min + x_max) / 2 # Calculate x centroid
print(f"Component {component}: width={width}, height={height}, x_pos={centroid_x}")
if width >= size_threshold_width and height >= size_threshold_height:
if mode == "convex_polygons":
polygon = self.get_mask_polygon(component_mask_np, max_poly_points)
if polygon is not None:
poly_mask = self.polygon_to_mask(polygon, (H, W))
poly_mask = torch.tensor(poly_mask, device=mask.device)
separated.append((centroid_x, poly_mask))
elif mode == "box":
# Create bounding box mask
box_mask = np.zeros((H, W), dtype=np.uint8)
box_mask[y_min:y_max+1, x_min:x_max+1] = 1
box_mask = torch.tensor(box_mask, device=mask.device)
separated.append((centroid_x, box_mask))
else:
area_mask = torch.tensor(component_mask_np, device=mask.device)
separated.append((centroid_x, area_mask))
pbar.update(1)
if len(separated) > 0:
# Sort by x position and extract only the masks
separated.sort(key=lambda x: x[0])
separated = [x[1] for x in separated]
out_masks = torch.stack(separated, dim=0)
return out_masks,
else:
return torch.empty((1, 64, 64), device=mask.device),
return (scaled_mask, )
File diff suppressed because it is too large Load Diff
+230 -1223
View File
File diff suppressed because it is too large Load Diff
-15
View File
@@ -1,15 +0,0 @@
[project]
name = "comfyui-kjnodes"
description = "Various quality of life -nodes for ComfyUI, mostly just visual stuff to improve usability."
version = "1.1.4"
license = {file = "LICENSE"}
dependencies = ["librosa", "numpy", "pillow>=10.3.0", "scipy", "color-matcher", "matplotlib", "huggingface_hub"]
[project.urls]
Repository = "https://github.com/kijai/ComfyUI-KJNodes"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "kijai"
DisplayName = "ComfyUI-KJNodes"
Icon = "https://avatars.githubusercontent.com/u/40791699"
+2 -3
View File
@@ -1,7 +1,6 @@
librosa
numpy
pillow>=10.3.0
scipy
color-matcher
matplotlib
huggingface_hub
mss
opencv-python
+1 -3
View File
@@ -47,9 +47,7 @@ app.registerExtension({
)
if (pythongossFeed) {
console.warn("KJNodes - Overriding pysssss.FaviconStatus")
pythongossFeed.setup = function() {
console.warn("Disabled by KJNodes")
};
app.extensions = app.extensions.filter(item => item !== pythongossFeed);
}
},
});
+53 -48
View File
@@ -48,100 +48,105 @@ app.registerExtension({
}
},
async setup(app) {
const updateSlots = (value) => {
const valuesToAddToIn = ["GetNode"];
const valuesToAddToOut = ["SetNode"];
// Remove entries if they exist
for (const arr of Object.values(LiteGraph.slot_types_default_in)) {
for (const valueToAdd of valuesToAddToIn) {
const idx = arr.indexOf(valueToAdd);
if (idx !== -1) {
arr.splice(idx, 1);
}
}
}
const onChange = (value) => {
if (value) {
const valuesToAddToIn = ["GetNode"];
const valuesToAddToOut = ["SetNode"];
for (const arr of Object.values(LiteGraph.slot_types_default_out)) {
for (const valueToAdd of valuesToAddToOut) {
const idx = arr.indexOf(valueToAdd);
if (idx !== -1) {
arr.splice(idx, 1);
}
}
}
if (value!="disabled") {
for (const arr of Object.values(LiteGraph.slot_types_default_in)) {
for (const valueToAdd of valuesToAddToIn) {
const idx = arr.indexOf(valueToAdd);
if (idx !== -1) {
if (idx !== 0) {
arr.splice(idx, 1);
}
if (value === "top") {
arr.unshift(valueToAdd);
} else {
arr.push(valueToAdd);
}
arr.unshift(valueToAdd);
}
}
for (const arr of Object.values(LiteGraph.slot_types_default_out)) {
for (const valueToAdd of valuesToAddToOut) {
const idx = arr.indexOf(valueToAdd);
if (idx !== -1) {
if (idx !== 0) {
arr.splice(idx, 1);
}
if (value === "top") {
arr.unshift(valueToAdd);
} else {
arr.push(valueToAdd);
}
arr.unshift(valueToAdd);
}
}
}
};
app.ui.settings.addSetting({
id: "KJNodes.SetGetMenu",
name: "KJNodes: Make Set/Get -nodes defaults",
tooltip: 'Adds Set/Get nodes to the top or bottom of the list of available node suggestions.',
options: ['disabled', 'top', 'bottom'],
defaultValue: 'disabled',
type: "combo",
onChange: updateSlots,
id: "🦛 KJNodes.SetGetMenu",
name: "🦛 KJNodes: Make Set/Get -nodes defaults (turn off and reload to disable)",
defaultValue: false,
type: "boolean",
options: (value) => [
{
value: true,
text: "On",
selected: value === true,
},
{
value: false,
text: "Off",
selected: value === false,
},
],
onChange: onChange,
});
app.ui.settings.addSetting({
id: "KJNodes.MiddleClickDefault",
name: "KJNodes: Middle click default node adding",
id: "KJNodes.DisableMiddleClickDefault",
name: "🦛 KJNodes: Middle click default node adding",
defaultValue: false,
type: "boolean",
options: (value) => [
{ value: true, text: "On", selected: value === true },
{ value: false, text: "Off", selected: value === false },
],
onChange: (value) => {
LiteGraph.middle_click_slot_add_default_node = value;
},
});
app.ui.settings.addSetting({
id: "KJNodes.nodeAutoColor",
name: "KJNodes: Automatically set node colors",
type: "boolean",
name: "🦛 KJNodes: Automatically set node colors",
defaultValue: true,
type: "boolean",
options: (value) => [
{ value: true, text: "On", selected: value === true },
{ value: false, text: "Off", selected: value === false },
],
});
app.ui.settings.addSetting({
id: "KJNodes.helpPopup",
name: "KJNodes: Help popups",
name: "🦛 KJNodes: Help popups",
defaultValue: true,
type: "boolean",
options: (value) => [
{ value: true, text: "On", selected: value === true },
{ value: false, text: "Off", selected: value === false },
],
});
app.ui.settings.addSetting({
id: "KJNodes.disablePrefix",
name: "KJNodes: Disable automatic Set_ and Get_ prefix",
defaultValue: true,
name: "🦛 KJNodes: Disable automatic Set_ and Get_ prefix",
defaultValue: false,
type: "boolean",
options: (value) => [
{ value: true, text: "On", selected: value === true },
{ value: false, text: "Off", selected: value === false },
],
});
app.ui.settings.addSetting({
id: "KJNodes.browserStatus",
name: "KJNodes: 🟢 Stoplight browser status icon 🔴",
name: "🦛 KJNodes: 🟢 Stoplight browser status icon 🔴",
defaultValue: false,
type: "boolean",
options: (value) => [
{ value: true, text: "On", selected: value === true },
{ value: false, text: "Off", selected: value === false },
],
});
}
});
-95
View File
@@ -1,95 +0,0 @@
import { app } from '../../../scripts/app.js'
//from melmass
export function makeUUID() {
let dt = new Date().getTime()
const uuid = 'xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx'.replace(/[xy]/g, (c) => {
const r = ((dt + Math.random() * 16) % 16) | 0
dt = Math.floor(dt / 16)
return (c === 'x' ? r : (r & 0x3) | 0x8).toString(16)
})
return uuid
}
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;
}
}
app.registerExtension({
name: 'KJNodes.FastPreview',
async beforeRegisterNodeDef(nodeType, nodeData) {
if (nodeData?.name === 'FastPreview') {
chainCallback(nodeType.prototype, "onNodeCreated", function () {
var element = document.createElement("div");
this.uuid = makeUUID()
element.id = `fast-preview-${this.uuid}`
this.previewWidget = this.addDOMWidget(nodeData.name, "FastPreviewWidget", element, {
serialize: false,
hideOnZoom: false,
});
this.previewer = new Previewer(this);
this.setSize([550, 550]);
this.resizable = false;
this.previewWidget.parentEl = document.createElement("div");
this.previewWidget.parentEl.className = "fast-preview";
this.previewWidget.parentEl.id = `fast-preview-${this.uuid}`
element.appendChild(this.previewWidget.parentEl);
chainCallback(this, "onExecuted", function (message) {
let bg_image = message["bg_image"];
this.properties.imgData = {
name: "bg_image",
base64: bg_image
};
this.previewer.refreshBackgroundImage(this);
});
}); // onAfterGraphConfigured
}//node created
} //before register
})//register
class Previewer {
constructor(context) {
this.node = context;
this.previousWidth = null;
this.previousHeight = null;
}
refreshBackgroundImage = () => {
const imgData = this.node?.properties?.imgData;
if (imgData?.base64) {
const base64String = imgData.base64;
const imageUrl = `data:${imgData.type};base64,${base64String}`;
const img = new Image();
img.src = imageUrl;
img.onload = () => {
const { width, height } = img;
if (width !== this.previousWidth || height !== this.previousHeight) {
this.node.setSize([width, height]);
this.previousWidth = width;
this.previousHeight = height;
}
this.node.previewWidget.element.style.backgroundImage = `url(${imageUrl})`;
};
}
};
}
+5 -6
View File
@@ -38,14 +38,14 @@ export const loadScript = (
})
}
loadScript('kjweb_async/marked.min.js').catch((e) => {
loadScript('/kjweb_async/marked.min.js').catch((e) => {
console.log(e)
})
loadScript('kjweb_async/purify.min.js').catch((e) => {
loadScript('/kjweb_async/purify.min.js').catch((e) => {
console.log(e)
})
const categories = ["KJNodes", "SUPIR", "VoiceCraft", "Marigold", "IC-Light", "WanVideoWrapper"];
const categories = ["KJNodes", "SUPIR", "VoiceCraft", "Marigold"];
app.registerExtension({
name: "KJNodes.HelpPopup",
async beforeRegisterNodeDef(nodeType, nodeData) {
@@ -257,13 +257,12 @@ const create_documentation_stylesheet = () => {
const scale = new DOMMatrix()
.scaleSelf(transform.a, transform.d);
const bcr = app.canvas.canvas.getBoundingClientRect()
const styleObject = {
transformOrigin: '0 0',
transform: scale,
left: `${transform.a + bcr.x + transform.e}px`,
top: `${transform.d + bcr.y + transform.f}px`,
left: `${transform.a + transform.e}px`,
top: `${transform.d + transform.f}px`,
};
Object.assign(docElement.style, styleObject);
}
+73 -192
View File
@@ -1,5 +1,4 @@
import { app } from "../../../scripts/app.js";
import { applyTextReplacements } from "../../../scripts/utils.js";
app.registerExtension({
name: "KJNodes.jsnodes",
@@ -10,158 +9,87 @@ app.registerExtension({
switch (nodeData.name) {
case "ConditioningMultiCombine":
nodeType.prototype.onNodeCreated = function () {
this._type = "CONDITIONING"
this.cond_type = "CONDITIONING"
this.inputs_offset = nodeData.name.includes("selective")?1:0
this.addWidget("button", "Update inputs", null, () => {
if (!this.inputs) {
this.inputs = [];
}
const target_number_of_inputs = this.widgets.find(w => w.name === "inputcount")["value"];
const num_inputs = this.inputs.filter(input => input.type === this._type).length
if(target_number_of_inputs===num_inputs)return; // already set, do nothing
if(target_number_of_inputs===this.inputs.length)return; // already set, do nothing
if(target_number_of_inputs < num_inputs){
const inputs_to_remove = num_inputs - target_number_of_inputs;
for(let i = 0; i < inputs_to_remove; i++) {
this.removeInput(this.inputs.length - 1);
}
}
else{
for(let i = num_inputs+1; i <= target_number_of_inputs; ++i)
this.addInput(`conditioning_${i}`, this._type)
}
if(target_number_of_inputs < this.inputs.length){
for(let i = this.inputs.length; i>=this.inputs_offset+target_number_of_inputs; i--)
this.removeInput(i)
}
else{
for(let i = this.inputs.length+1-this.inputs_offset; i <= target_number_of_inputs; ++i)
this.addInput(`conditioning_${i}`, this.cond_type)
}
});
}
break;
case "ImageBatchMulti":
case "ImageAddMulti":
case "ImageConcatMulti":
case "CrossFadeImagesMulti":
case "TransitionImagesMulti":
nodeType.prototype.onNodeCreated = function () {
this._type = "IMAGE"
this.inputs_offset = nodeData.name.includes("selective")?1:0
this.addWidget("button", "Update inputs", null, () => {
if (!this.inputs) {
this.inputs = [];
}
const target_number_of_inputs = this.widgets.find(w => w.name === "inputcount")["value"];
const num_inputs = this.inputs.filter(input => input.type === this._type).length
if(target_number_of_inputs===num_inputs)return; // already set, do nothing
if(target_number_of_inputs===this.inputs.length)return; // already set, do nothing
if(target_number_of_inputs < num_inputs){
const inputs_to_remove = num_inputs - target_number_of_inputs;
for(let i = 0; i < inputs_to_remove; i++) {
this.removeInput(this.inputs.length - 1);
if(target_number_of_inputs < this.inputs.length){
for(let i = this.inputs.length; i>=this.inputs_offset+target_number_of_inputs; i--)
this.removeInput(i)
}
else{
for(let i = this.inputs.length+1-this.inputs_offset; i <= target_number_of_inputs; ++i)
this.addInput(`image_${i}`, this._type)
}
}
else{
for(let i = num_inputs+1; i <= target_number_of_inputs; ++i)
this.addInput(`image_${i}`, this._type, {shape: 7});
}
});
}
break;
case "MaskBatchMulti":
nodeType.prototype.onNodeCreated = function () {
this._type = "MASK"
this.inputs_offset = nodeData.name.includes("selective")?1:0
this.addWidget("button", "Update inputs", null, () => {
if (!this.inputs) {
this.inputs = [];
}
const target_number_of_inputs = this.widgets.find(w => w.name === "inputcount")["value"];
const num_inputs = this.inputs.filter(input => input.type === this._type).length
if(target_number_of_inputs===num_inputs)return; // already set, do nothing
if(target_number_of_inputs===this.inputs.length)return; // already set, do nothing
if(target_number_of_inputs < num_inputs){
const inputs_to_remove = num_inputs - target_number_of_inputs;
for(let i = 0; i < inputs_to_remove; i++) {
this.removeInput(this.inputs.length - 1);
if(target_number_of_inputs < this.inputs.length){
for(let i = this.inputs.length; i>=this.inputs_offset+target_number_of_inputs; i--)
this.removeInput(i)
}
}
else{
for(let i = num_inputs+1; i <= target_number_of_inputs; ++i)
this.addInput(`mask_${i}`, this._type)
}
});
else{
for(let i = this.inputs.length+1-this.inputs_offset; i <= target_number_of_inputs; ++i)
this.addInput(`mask_${i}`, this._type)
}
});
}
break;
case "FluxBlockLoraSelect":
case "HunyuanVideoBlockLoraSelect":
case "Wan21BlockLoraSelect":
nodeType.prototype.onNodeCreated = function () {
this.addWidget("button", "Set all", null, () => {
const userInput = prompt("Enter the values to set for widgets (e.g., s0,1,2-7=2.0, d0,1,2-7=2.0, or 1.0):", "");
if (userInput) {
const regex = /([sd])?(\d+(?:,\d+|-?\d+)*?)?=(\d+(\.\d+)?)/;
const match = userInput.match(regex);
if (match) {
const type = match[1];
const indicesPart = match[2];
const value = parseFloat(match[3]);
let targetWidgets = [];
if (type === 's') {
targetWidgets = this.widgets.filter(widget => widget.name.includes("single"));
} else if (type === 'd') {
targetWidgets = this.widgets.filter(widget => widget.name.includes("double"));
} else {
targetWidgets = this.widgets; // No type specified, all widgets
}
if (indicesPart) {
const indices = indicesPart.split(',').flatMap(part => {
if (part.includes('-')) {
const [start, end] = part.split('-').map(Number);
return Array.from({ length: end - start + 1 }, (_, i) => start + i);
}
return Number(part);
});
for (const index of indices) {
if (index < targetWidgets.length) {
targetWidgets[index].value = value;
}
}
} else {
// No indices provided, set value for all target widgets
for (const widget of targetWidgets) {
widget.value = value;
}
}
} else if (!isNaN(parseFloat(userInput))) {
// Single value provided, set it for all widgets
const value = parseFloat(userInput);
for (const widget of this.widgets) {
widget.value = value;
}
} else {
alert("Invalid input format. Please use the format s0,1,2-7=2.0, d0,1,2-7=2.0, or 1.0");
}
} else {
alert("Invalid input. Please enter a value.");
}
});
};
break;
case "GetMaskSizeAndCount":
const onGetMaskSizeConnectInput = nodeType.prototype.onConnectInput;
nodeType.prototype.onConnectInput = function (targetSlot, type, output, originNode, originSlot) {
const v = onGetMaskSizeConnectInput? onGetMaskSizeConnectInput.apply(this, arguments): undefined
this.outputs[1]["label"] = "width"
this.outputs[2]["label"] = "height"
this.outputs[3]["label"] = "count"
const v = onGetMaskSizeConnectInput?.(this, arguments);
targetSlot.outputs[1]["name"] = "width"
targetSlot.outputs[2]["name"] = "height"
targetSlot.outputs[3]["name"] = "count"
return v;
}
const onGetMaskSizeExecuted = nodeType.prototype.onAfterExecuteNode;
const onGetMaskSizeExecuted = nodeType.prototype.onExecuted;
nodeType.prototype.onExecuted = function(message) {
const r = onGetMaskSizeExecuted? onGetMaskSizeExecuted.apply(this,arguments): undefined
let values = message["text"].toString().split('x').map(Number);
this.outputs[1]["label"] = values[1] + " width"
this.outputs[2]["label"] = values[2] + " height"
this.outputs[3]["label"] = values[0] + " count"
this.outputs[1]["name"] = values[1] + " width"
this.outputs[2]["name"] = values[2] + " height"
this.outputs[3]["name"] = values[0] + " count"
return r
}
break;
@@ -169,51 +97,19 @@ app.registerExtension({
case "GetImageSizeAndCount":
const onGetImageSizeConnectInput = nodeType.prototype.onConnectInput;
nodeType.prototype.onConnectInput = function (targetSlot, type, output, originNode, originSlot) {
console.log(this)
const v = onGetImageSizeConnectInput? onGetImageSizeConnectInput.apply(this, arguments): undefined
//console.log(this)
this.outputs[1]["label"] = "width"
this.outputs[2]["label"] = "height"
this.outputs[3]["label"] = "count"
const v = onGetImageSizeConnectInput?.(this, arguments);
targetSlot.outputs[1]["name"] = "width"
targetSlot.outputs[2]["name"] = "height"
targetSlot.outputs[3]["name"] = "count"
return v;
}
//const onGetImageSizeExecuted = nodeType.prototype.onExecuted;
const onGetImageSizeExecuted = nodeType.prototype.onAfterExecuteNode;
const onGetImageSizeExecuted = nodeType.prototype.onExecuted;
nodeType.prototype.onExecuted = function(message) {
console.log(this)
const r = onGetImageSizeExecuted? onGetImageSizeExecuted.apply(this,arguments): undefined
let values = message["text"].toString().split('x').map(Number);
console.log(values)
this.outputs[1]["label"] = values[1] + " width"
this.outputs[2]["label"] = values[2] + " height"
this.outputs[3]["label"] = values[0] + " count"
return r
}
break;
case "GetLatentSizeAndCount":
const onGetLatentConnectInput = nodeType.prototype.onConnectInput;
nodeType.prototype.onConnectInput = function (targetSlot, type, output, originNode, originSlot) {
console.log(this)
const v = onGetLatentConnectInput? onGetLatentConnectInput.apply(this, arguments): undefined
//console.log(this)
this.outputs[1]["label"] = "width"
this.outputs[2]["label"] = "height"
this.outputs[3]["label"] = "count"
return v;
}
//const onGetImageSizeExecuted = nodeType.prototype.onExecuted;
const onGetLatentSizeExecuted = nodeType.prototype.onAfterExecuteNode;
nodeType.prototype.onExecuted = function(message) {
console.log(this)
const r = onGetLatentSizeExecuted? onGetLatentSizeExecuted.apply(this,arguments): undefined
let values = message["text"].toString().split('x').map(Number);
console.log(values)
this.outputs[1]["label"] = values[0] + " batch"
this.outputs[2]["label"] = values[1] + " channels"
this.outputs[3]["label"] = values[2] + " frames"
this.outputs[4]["label"] = values[3] + " height"
this.outputs[5]["label"] = values[4] + " width"
this.outputs[1]["name"] = values[1] + " width"
this.outputs[2]["name"] = values[2] + " height"
this.outputs[3]["name"] = values[0] + " count"
return r
}
break;
@@ -221,14 +117,15 @@ app.registerExtension({
case "PreviewAnimation":
const onPreviewAnimationConnectInput = nodeType.prototype.onConnectInput;
nodeType.prototype.onConnectInput = function (targetSlot, type, output, originNode, originSlot) {
const v = onPreviewAnimationConnectInput? onPreviewAnimationConnectInput.apply(this, arguments): undefined
this.title = "Preview Animation"
const v = onPreviewAnimationConnectInput?.(this, arguments);
targetSlot.title = "Preview Animation"
return v;
}
const onPreviewAnimationExecuted = nodeType.prototype.onAfterExecuteNode;
const onPreviewAnimationExecuted = nodeType.prototype.onExecuted;
nodeType.prototype.onExecuted = function(message) {
const r = onPreviewAnimationExecuted? onPreviewAnimationExecuted.apply(this,arguments): undefined
let values = message["text"].toString();
console.log(this)
this.title = "Preview Animation " + values
return r
}
@@ -237,48 +134,43 @@ app.registerExtension({
case "VRAM_Debug":
const onVRAM_DebugConnectInput = nodeType.prototype.onConnectInput;
nodeType.prototype.onConnectInput = function (targetSlot, type, output, originNode, originSlot) {
const v = onVRAM_DebugConnectInput? onVRAM_DebugConnectInput.apply(this, arguments): undefined
this.outputs[3]["label"] = "freemem_before"
this.outputs[4]["label"] = "freemem_after"
const v = onVRAM_DebugConnectInput?.(this, arguments);
targetSlot.outputs[3]["name"] = "freemem_before"
targetSlot.outputs[4]["name"] = "freemem_after"
return v;
}
const onVRAM_DebugExecuted = nodeType.prototype.onAfterExecuteNode;
const onVRAM_DebugExecuted = nodeType.prototype.onExecuted;
nodeType.prototype.onExecuted = function(message) {
const r = onVRAM_DebugExecuted? onVRAM_DebugExecuted.apply(this,arguments): undefined
let values = message["text"].toString().split('x');
this.outputs[3]["label"] = values[0] + " freemem_before"
this.outputs[4]["label"] = values[1] + " freemem_after"
this.outputs[3]["name"] = values[0] + " freemem_before"
this.outputs[4]["name"] = values[1] + " freemem_after"
return r
}
break;
case "JoinStringMulti":
const originalOnNodeCreated = nodeType.prototype.onNodeCreated || function() {};
nodeType.prototype.onNodeCreated = function () {
originalOnNodeCreated.apply(this, arguments);
this._type = "STRING";
this.addWidget("button", "Update inputs", null, () => {
if (!this.inputs) {
this.inputs = [];
}
const target_number_of_inputs = this.widgets.find(w => w.name === "inputcount")["value"];
const num_inputs = this.inputs.filter(input => input.name && input.name.toLowerCase().includes("string_")).length
if (target_number_of_inputs === num_inputs) return; // already set, do nothing
if(target_number_of_inputs < num_inputs){
const inputs_to_remove = num_inputs - target_number_of_inputs;
for(let i = 0; i < inputs_to_remove; i++) {
this.removeInput(this.inputs.length - 1);
}
this._type = "STRING"
this.inputs_offset = nodeData.name.includes("selective")?1:0
this.addWidget("button", "Update inputs", null, () => {
if (!this.inputs) {
this.inputs = [];
}
const target_number_of_inputs = this.widgets.find(w => w.name === "inputcount")["value"];
if(target_number_of_inputs===this.inputs.length)return; // already set, do nothing
if(target_number_of_inputs < this.inputs.length){
for(let i = this.inputs.length; i>=this.inputs_offset+target_number_of_inputs; i--)
this.removeInput(i)
}
else{
for(let i = num_inputs+1; i <= target_number_of_inputs; ++i)
this.addInput(`string_${i}`, this._type, {shape: 7});
}
});
}
break;
for(let i = this.inputs.length+1-this.inputs_offset; i <= target_number_of_inputs; ++i)
this.addInput(`string_${i}`, this._type)
}
});
}
break;
case "SoundReactive":
nodeType.prototype.onNodeCreated = function () {
let audioContext;
@@ -381,17 +273,6 @@ app.registerExtension({
this.addWidget("button", "Stop mic capture", null, stopMicrophoneCapture);
};
break;
case "SaveImageKJ":
const onNodeCreated = nodeType.prototype.onNodeCreated;
nodeType.prototype.onNodeCreated = function() {
const r = onNodeCreated ? onNodeCreated.apply(this, arguments) : void 0;
const widget = this.widgets.find((w) => w.name === "filename_prefix");
widget.serializeValue = () => {
return applyTextReplacements(app, widget.value);
};
return r;
};
break;
}
-734
View File
@@ -1,734 +0,0 @@
import { app } from '../../../scripts/app.js'
//from melmass
export function makeUUID() {
let dt = new Date().getTime()
const uuid = 'xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx'.replace(/[xy]/g, (c) => {
const r = ((dt + Math.random() * 16) % 16) | 0
dt = Math.floor(dt / 16)
return (c === 'x' ? r : (r & 0x3) | 0x8).toString(16)
})
return uuid
}
export const loadScript = (
FILE_URL,
async = true,
type = 'text/javascript',
) => {
return new Promise((resolve, reject) => {
try {
// Check if the script already exists
const existingScript = document.querySelector(`script[src="${FILE_URL}"]`)
if (existingScript) {
resolve({ status: true, message: 'Script already loaded' })
return
}
const scriptEle = document.createElement('script')
scriptEle.type = type
scriptEle.async = async
scriptEle.src = FILE_URL
scriptEle.addEventListener('load', (ev) => {
resolve({ status: true })
})
scriptEle.addEventListener('error', (ev) => {
reject({
status: false,
message: `Failed to load the script ${FILE_URL}`,
})
})
document.body.appendChild(scriptEle)
} catch (error) {
reject(error)
}
})
}
const create_documentation_stylesheet = () => {
const tag = 'kj-pointseditor-stylesheet'
let styleTag = document.head.querySelector(tag)
if (!styleTag) {
styleTag = document.createElement('style')
styleTag.type = 'text/css'
styleTag.id = tag
styleTag.innerHTML = `
.points-editor {
position: absolute;
font: 12px monospace;
line-height: 1.5em;
padding: 10px;
z-index: 0;
overflow: hidden;
}
`
document.head.appendChild(styleTag)
}
}
loadScript('kjweb_async/svg-path-properties.min.js').catch((e) => {
console.log(e)
})
loadScript('kjweb_async/protovis.min.js').catch((e) => {
console.log(e)
})
create_documentation_stylesheet()
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;
}
}
app.registerExtension({
name: 'KJNodes.PointEditor',
async beforeRegisterNodeDef(nodeType, nodeData) {
if (nodeData?.name === 'PointsEditor') {
chainCallback(nodeType.prototype, "onNodeCreated", function () {
hideWidgetForGood(this, this.widgets.find(w => w.name === "coordinates"))
hideWidgetForGood(this, this.widgets.find(w => w.name === "neg_coordinates"))
hideWidgetForGood(this, this.widgets.find(w => w.name === "bboxes"))
var element = document.createElement("div");
this.uuid = makeUUID()
element.id = `points-editor-${this.uuid}`
this.previewMediaType = 'image'
this.pointsEditor = this.addDOMWidget(nodeData.name, "PointsEditorWidget", element, {
serialize: false,
hideOnZoom: false,
});
// context menu
this.contextMenu = document.createElement("div");
this.contextMenu.id = "context-menu";
this.contextMenu.style.display = "none";
this.contextMenu.style.position = "absolute";
this.contextMenu.style.backgroundColor = "#202020";
this.contextMenu.style.minWidth = "100px";
this.contextMenu.style.boxShadow = "0px 8px 16px 0px rgba(0,0,0,0.2)";
this.contextMenu.style.zIndex = "100";
this.contextMenu.style.padding = "5px";
function styleMenuItem(menuItem) {
menuItem.style.display = "block";
menuItem.style.padding = "5px";
menuItem.style.color = "#FFF";
menuItem.style.fontFamily = "Arial, sans-serif";
menuItem.style.fontSize = "16px";
menuItem.style.textDecoration = "none";
menuItem.style.marginBottom = "5px";
}
function createMenuItem(id, textContent) {
let menuItem = document.createElement("a");
menuItem.href = "#";
menuItem.id = `menu-item-${id}`;
menuItem.textContent = textContent;
styleMenuItem(menuItem);
return menuItem;
}
// Create an array of menu items using the createMenuItem function
this.menuItems = [
createMenuItem(0, "Load Image"),
createMenuItem(1, "Clear Image"),
];
// Add mouseover and mouseout event listeners to each menu item for styling
this.menuItems.forEach(menuItem => {
menuItem.addEventListener('mouseover', function () {
this.style.backgroundColor = "gray";
});
menuItem.addEventListener('mouseout', function () {
this.style.backgroundColor = "#202020";
});
});
// Append each menu item to the context menu
this.menuItems.forEach(menuItem => {
this.contextMenu.appendChild(menuItem);
});
document.body.appendChild(this.contextMenu);
this.addWidget("button", "New canvas", null, () => {
if (!this.properties || !("points" in this.properties)) {
this.editor = new PointsEditor(this);
this.addProperty("points", this.constructor.type, "string");
this.addProperty("neg_points", this.constructor.type, "string");
}
else {
this.editor = new PointsEditor(this, true);
}
});
this.setSize([550, 550]);
this.resizable = false;
this.pointsEditor.parentEl = document.createElement("div");
this.pointsEditor.parentEl.className = "points-editor";
this.pointsEditor.parentEl.id = `points-editor-${this.uuid}`
element.appendChild(this.pointsEditor.parentEl);
chainCallback(this, "onConfigure", function () {
try {
this.editor = new PointsEditor(this);
} catch (error) {
console.error("An error occurred while configuring the editor:", error);
}
});
chainCallback(this, "onExecuted", function (message) {
let bg_image = message["bg_image"];
this.properties.imgData = {
name: "bg_image",
base64: bg_image
};
this.editor.refreshBackgroundImage(this);
});
}); // onAfterGraphConfigured
}//node created
} //before register
})//register
class PointsEditor {
constructor(context, reset = false) {
this.node = context;
this.reset = reset;
const self = this; // Keep a reference to the main class context
console.log("creatingPointEditor")
this.node.pasteFile = (file) => {
if (file.type.startsWith("image/")) {
this.handleImageFile(file);
return true;
}
return false;
};
this.node.onDragOver = function (e) {
if (e.dataTransfer && e.dataTransfer.items) {
return [...e.dataTransfer.items].some(f => f.kind === "file" && f.type.startsWith("image/"));
}
return false;
};
// On drop upload files
this.node.onDragDrop = (e) => {
console.log("onDragDrop called");
let handled = false;
for (const file of e.dataTransfer.files) {
if (file.type.startsWith("image/")) {
this.handleImageFile(file);
handled = true;
}
}
return handled;
};
// context menu
this.createContextMenu();
if (reset && context.pointsEditor.element) {
context.pointsEditor.element.innerHTML = ''; // Clear the container
}
this.pos_coordWidget = context.widgets.find(w => w.name === "coordinates");
this.neg_coordWidget = context.widgets.find(w => w.name === "neg_coordinates");
this.pointsStoreWidget = context.widgets.find(w => w.name === "points_store");
this.widthWidget = context.widgets.find(w => w.name === "width");
this.heightWidget = context.widgets.find(w => w.name === "height");
this.bboxStoreWidget = context.widgets.find(w => w.name === "bbox_store");
this.bboxWidget = context.widgets.find(w => w.name === "bboxes");
//widget callbacks
this.widthWidget.callback = () => {
this.width = this.widthWidget.value;
if (this.width > 256) {
context.setSize([this.width + 45, context.size[1]]);
}
this.vis.width(this.width);
this.updateData();
}
this.heightWidget.callback = () => {
this.height = this.heightWidget.value
this.vis.height(this.height)
context.setSize([context.size[0], this.height + 300]);
this.updateData();
}
this.pointsStoreWidget.callback = () => {
this.points = JSON.parse(pointsStoreWidget.value).positive;
this.neg_points = JSON.parse(pointsStoreWidget.value).negative;
this.updateData();
}
this.bboxStoreWidget.callback = () => {
this.bbox = JSON.parse(bboxStoreWidget.value)
this.updateData();
}
this.width = this.widthWidget.value;
this.height = this.heightWidget.value;
var i = 3;
this.points = [];
this.neg_points = [];
this.bbox = [{}];
var drawing = false;
// Initialize or reset points array
if (!reset && this.pointsStoreWidget.value != "") {
this.points = JSON.parse(this.pointsStoreWidget.value).positive;
this.neg_points = JSON.parse(this.pointsStoreWidget.value).negative;
this.bbox = JSON.parse(this.bboxStoreWidget.value);
console.log(this.bbox)
} else {
this.points = [
{
x: this.width / 2, // Middle point horizontally centered
y: this.height / 2 // Middle point vertically centered
}
];
this.neg_points = [
{
x: 0, // Middle point horizontally centered
y: 0 // Middle point vertically centered
}
];
const combinedPoints = {
positive: this.points,
negative: this.neg_points,
};
this.pointsStoreWidget.value = JSON.stringify(combinedPoints);
this.bboxStoreWidget.value = JSON.stringify(this.bbox);
}
//create main canvas panel
this.vis = new pv.Panel()
.width(this.width)
.height(this.height)
.fillStyle("#222")
.strokeStyle("gray")
.lineWidth(2)
.antialias(false)
.margin(10)
.event("mousedown", function () {
if (pv.event.shiftKey && pv.event.button === 2) { // Use pv.event to access the event object
let scaledMouse = {
x: this.mouse().x / app.canvas.ds.scale,
y: this.mouse().y / app.canvas.ds.scale
};
i = self.neg_points.push(scaledMouse) - 1;
self.updateData();
return this;
}
else if (pv.event.shiftKey) {
let scaledMouse = {
x: this.mouse().x / app.canvas.ds.scale,
y: this.mouse().y / app.canvas.ds.scale
};
i = self.points.push(scaledMouse) - 1;
self.updateData();
return this;
}
else if (pv.event.ctrlKey) {
console.log("start drawing at " + this.mouse().x / app.canvas.ds.scale + ", " + this.mouse().y / app.canvas.ds.scale);
drawing = true;
self.bbox[0].startX = this.mouse().x / app.canvas.ds.scale;
self.bbox[0].startY = this.mouse().y / app.canvas.ds.scale;
}
else if (pv.event.button === 2) {
self.node.contextMenu.style.display = 'block';
self.node.contextMenu.style.left = `${pv.event.clientX}px`;
self.node.contextMenu.style.top = `${pv.event.clientY}px`;
}
})
.event("mousemove", function () {
if (drawing) {
self.bbox[0].endX = this.mouse().x / app.canvas.ds.scale;
self.bbox[0].endY = this.mouse().y / app.canvas.ds.scale;
self.vis.render();
}
})
.event("mouseup", function () {
console.log("end drawing at " + this.mouse().x / app.canvas.ds.scale + ", " + this.mouse().y / app.canvas.ds.scale);
drawing = false;
self.updateData();
});
this.backgroundImage = this.vis.add(pv.Image).visible(false)
//create bounding box
this.bounding_box = this.vis.add(pv.Area)
.data(function () {
if (drawing || (self.bbox && self.bbox[0] && Object.keys(self.bbox[0]).length > 0)) {
return [self.bbox[0].startX, self.bbox[0].endX];
} else {
return [];
}
})
.bottom(function () {return self.height - Math.max(self.bbox[0].startY, self.bbox[0].endY); })
.left(function (d) {return d; })
.height(function () {return Math.abs(self.bbox[0].startY - self.bbox[0].endY);})
.fillStyle("rgba(70, 130, 180, 0.5)")
.strokeStyle("steelblue")
.visible(function () {return drawing || Object.keys(self.bbox[0]).length > 0; })
.add(pv.Dot)
.visible(function () {return drawing || Object.keys(self.bbox[0]).length > 0; })
.data(() => {
if (self.bbox && Object.keys(self.bbox[0]).length > 0) {
return [{
x: self.bbox[0].endX,
y: self.bbox[0].endY
}];
} else {
return [];
}
})
.left(d => d.x)
.top(d => d.y)
.radius(Math.log(Math.min(self.width, self.height)) * 1)
.shape("square")
.cursor("move")
.strokeStyle("steelblue")
.lineWidth(2)
.fillStyle(function () { return "rgba(100, 100, 100, 0.6)"; })
.event("mousedown", pv.Behavior.drag())
.event("drag", function () {
let adjustedX = this.mouse().x / app.canvas.ds.scale; // Adjust the new position by the inverse of the scale factor
let adjustedY = this.mouse().y / app.canvas.ds.scale;
// Adjust the new position if it would place the dot outside the bounds of the vis.Panel
adjustedX = Math.max(0, Math.min(self.vis.width(), adjustedX));
adjustedY = Math.max(0, Math.min(self.vis.height(), adjustedY));
self.bbox[0].endX = this.mouse().x / app.canvas.ds.scale;
self.bbox[0].endY = this.mouse().y / app.canvas.ds.scale;
self.vis.render();
})
.event("dragend", function () {
self.updateData();
});
//create positive points
this.vis.add(pv.Dot)
.data(() => this.points)
.left(d => d.x)
.top(d => d.y)
.radius(Math.log(Math.min(self.width, self.height)) * 4)
.shape("circle")
.cursor("move")
.strokeStyle(function () { return i == this.index ? "#07f907" : "#139613"; })
.lineWidth(4)
.fillStyle(function () { return "rgba(100, 100, 100, 0.6)"; })
.event("mousedown", pv.Behavior.drag())
.event("dragstart", function () {
i = this.index;
})
.event("dragend", function () {
if (pv.event.button === 2 && i !== 0 && i !== self.points.length - 1) {
this.index = i;
self.points.splice(i--, 1);
}
self.updateData();
})
.event("drag", function () {
let adjustedX = this.mouse().x / app.canvas.ds.scale; // Adjust the new X position by the inverse of the scale factor
let adjustedY = this.mouse().y / app.canvas.ds.scale; // Adjust the new Y position by the inverse of the scale factor
// Determine the bounds of the vis.Panel
const panelWidth = self.vis.width();
const panelHeight = self.vis.height();
// Adjust the new position if it would place the dot outside the bounds of the vis.Panel
adjustedX = Math.max(0, Math.min(panelWidth, adjustedX));
adjustedY = Math.max(0, Math.min(panelHeight, adjustedY));
self.points[this.index] = { x: adjustedX, y: adjustedY }; // Update the point's position
self.vis.render(); // Re-render the visualization to reflect the new position
})
.anchor("center")
.add(pv.Label)
.left(d => d.x < this.width / 2 ? d.x + 30 : d.x - 35) // Shift label to right if on left half, otherwise shift to left
.top(d => d.y < this.height / 2 ? d.y + 25 : d.y - 25) // Shift label down if on top half, otherwise shift up
.font(25 + "px sans-serif")
.text(d => {return this.points.indexOf(d); })
.textStyle("#139613")
.textShadow("2px 2px 2px black")
.add(pv.Dot) // Add smaller point in the center
.data(() => this.points)
.left(d => d.x)
.top(d => d.y)
.radius(2) // Smaller radius for the center point
.shape("circle")
.fillStyle("red") // Color for the center point
.lineWidth(1); // Stroke thickness for the center point
//create negative points
this.vis.add(pv.Dot)
.data(() => this.neg_points)
.left(d => d.x)
.top(d => d.y)
.radius(Math.log(Math.min(self.width, self.height)) * 4)
.shape("circle")
.cursor("move")
.strokeStyle(function () { return i == this.index ? "#f91111" : "#891616"; })
.lineWidth(4)
.fillStyle(function () { return "rgba(100, 100, 100, 0.6)"; })
.event("mousedown", pv.Behavior.drag())
.event("dragstart", function () {
i = this.index;
})
.event("dragend", function () {
if (pv.event.button === 2 && i !== 0 && i !== self.neg_points.length - 1) {
this.index = i;
self.neg_points.splice(i--, 1);
}
self.updateData();
})
.event("drag", function () {
let adjustedX = this.mouse().x / app.canvas.ds.scale; // Adjust the new X position by the inverse of the scale factor
let adjustedY = this.mouse().y / app.canvas.ds.scale; // Adjust the new Y position by the inverse of the scale factor
// Determine the bounds of the vis.Panel
const panelWidth = self.vis.width();
const panelHeight = self.vis.height();
// Adjust the new position if it would place the dot outside the bounds of the vis.Panel
adjustedX = Math.max(0, Math.min(panelWidth, adjustedX));
adjustedY = Math.max(0, Math.min(panelHeight, adjustedY));
self.neg_points[this.index] = { x: adjustedX, y: adjustedY }; // Update the point's position
self.vis.render(); // Re-render the visualization to reflect the new position
})
.anchor("center")
.add(pv.Label)
.left(d => d.x < this.width / 2 ? d.x + 30 : d.x - 35) // Shift label to right if on left half, otherwise shift to left
.top(d => d.y < this.height / 2 ? d.y + 25 : d.y - 25) // Shift label down if on top half, otherwise shift up
.font(25 + "px sans-serif")
.text(d => {return this.neg_points.indexOf(d); })
.textStyle("red")
.textShadow("2px 2px 2px black")
.add(pv.Dot) // Add smaller point in the center
.data(() => this.neg_points)
.left(d => d.x)
.top(d => d.y)
.radius(2) // Smaller radius for the center point
.shape("circle")
.fillStyle("red") // Color for the center point
.lineWidth(1); // Stroke thickness for the center point
if (this.points.length != 0) {
this.vis.render();
}
var svgElement = this.vis.canvas();
svgElement.style['zIndex'] = "2"
svgElement.style['position'] = "relative"
this.node.pointsEditor.element.appendChild(svgElement);
if (this.width > 256) {
this.node.setSize([this.width + 45, this.node.size[1]]);
}
this.node.setSize([this.node.size[0], this.height + 300]);
this.updateData();
this.refreshBackgroundImage();
}//end constructor
updateData = () => {
if (!this.points || this.points.length === 0) {
console.log("no points");
return;
}
const combinedPoints = {
positive: this.points,
negative: this.neg_points,
};
this.pointsStoreWidget.value = JSON.stringify(combinedPoints);
this.pos_coordWidget.value = JSON.stringify(this.points);
this.neg_coordWidget.value = JSON.stringify(this.neg_points);
if (this.bbox.length != 0) {
let bboxString = JSON.stringify(this.bbox);
this.bboxStoreWidget.value = bboxString;
this.bboxWidget.value = bboxString;
}
this.vis.render();
};
handleImageLoad = (img, file, base64String) => {
console.log(img.width, img.height); // Access width and height here
this.widthWidget.value = img.width;
this.heightWidget.value = img.height;
if (img.width != this.vis.width() || img.height != this.vis.height()) {
if (img.width > 256) {
this.node.setSize([img.width + 45, this.node.size[1]]);
}
this.node.setSize([this.node.size[0], img.height + 300]);
this.vis.width(img.width);
this.vis.height(img.height);
this.height = img.height;
this.width = img.width;
this.updateData();
}
this.backgroundImage.url(file ? URL.createObjectURL(file) : `data:${this.node.properties.imgData.type};base64,${base64String}`).visible(true).root.render();
};
processImage = (img, file) => {
const canvas = document.createElement('canvas');
const ctx = canvas.getContext('2d');
const maxWidth = 800; // maximum width
const maxHeight = 600; // maximum height
let width = img.width;
let height = img.height;
// Calculate the new dimensions while preserving the aspect ratio
if (width > height) {
if (width > maxWidth) {
height *= maxWidth / width;
width = maxWidth;
}
} else {
if (height > maxHeight) {
width *= maxHeight / height;
height = maxHeight;
}
}
canvas.width = width;
canvas.height = height;
ctx.drawImage(img, 0, 0, width, height);
// Get the compressed image data as a Base64 string
const base64String = canvas.toDataURL('image/jpeg', 0.5).replace('data:', '').replace(/^.+,/, ''); // 0.5 is the quality from 0 to 1
this.node.properties.imgData = {
name: file.name,
lastModified: file.lastModified,
size: file.size,
type: file.type,
base64: base64String
};
handleImageLoad(img, file, base64String);
};
handleImageFile = (file) => {
const reader = new FileReader();
reader.onloadend = () => {
const img = new Image();
img.src = reader.result;
img.onload = () => processImage(img, file);
};
reader.readAsDataURL(file);
const imageUrl = URL.createObjectURL(file);
const img = new Image();
img.src = imageUrl;
img.onload = () => this.handleImageLoad(img, file, null);
};
refreshBackgroundImage = () => {
if (this.node.properties.imgData && this.node.properties.imgData.base64) {
const base64String = this.node.properties.imgData.base64;
const imageUrl = `data:${this.node.properties.imgData.type};base64,${base64String}`;
const img = new Image();
img.src = imageUrl;
img.onload = () => this.handleImageLoad(img, null, base64String);
}
};
createContextMenu = () => {
self = this;
document.addEventListener('contextmenu', function (e) {
e.preventDefault();
});
document.addEventListener('click', function (e) {
if (!self.node.contextMenu.contains(e.target)) {
self.node.contextMenu.style.display = 'none';
}
});
this.node.menuItems.forEach((menuItem, index) => {
self = this;
menuItem.addEventListener('click', function (e) {
e.preventDefault();
switch (index) {
case 0:
// Create file input element
const fileInput = document.createElement('input');
fileInput.type = 'file';
fileInput.accept = 'image/*'; // Accept only image files
// Listen for file selection
fileInput.addEventListener('change', function (event) {
const file = event.target.files[0]; // Get the selected file
if (file) {
const imageUrl = URL.createObjectURL(file);
let img = new Image();
img.src = imageUrl;
img.onload = () => self.handleImageLoad(img, file, null);
}
});
fileInput.click();
self.node.contextMenu.style.display = 'none';
break;
case 1:
self.backgroundImage.visible(false).root.render();
self.node.properties.imgData = null;
self.node.contextMenu.style.display = 'none';
break;
}
});
});
}//end createContextMenu
}//end class
//from melmass
export function hideWidgetForGood(node, widget, suffix = '') {
widget.origType = widget.type
widget.origComputeSize = widget.computeSize
widget.origSerializeValue = widget.serializeValue
widget.computeSize = () => [0, -4] // -4 is due to the gap litegraph adds between widgets automatically
widget.type = "converted-widget" + suffix
// widget.serializeValue = () => {
// // Prevent serializing the widget if we have no input linked
// const w = node.inputs?.find((i) => i.widget?.name === widget.name);
// if (w?.link == null) {
// return undefined;
// }
// return widget.origSerializeValue ? widget.origSerializeValue() : widget.value;
// };
// Hide any linked widgets, e.g. seed+seedControl
if (widget.linkedWidgets) {
for (const w of widget.linkedWidgets) {
hideWidgetForGood(node, w, ':' + widget.name)
}
}
}
+23 -45
View File
@@ -12,13 +12,8 @@ function setColorAndBgColor(type) {
"IMAGE": LGraphCanvas.node_colors.pale_blue,
"CLIP": LGraphCanvas.node_colors.yellow,
"FLOAT": LGraphCanvas.node_colors.green,
"MASK": { color: "#1c5715", bgcolor: "#1f401b"},
"MASK": LGraphCanvas.node_colors.cyan,
"INT": { color: "#1b4669", bgcolor: "#29699c"},
"CONTROL_NET": { color: "#156653", bgcolor: "#1c453b"},
"NOISE": { color: "#2e2e2e", bgcolor: "#242121"},
"GUIDER": { color: "#3c7878", bgcolor: "#1c453b"},
"SAMPLER": { color: "#614a4a", bgcolor: "#3b2c2c"},
"SIGMAS": { color: "#485248", bgcolor: "#272e27"},
};
@@ -28,21 +23,20 @@ function setColorAndBgColor(type) {
this.bgcolor = colors.bgcolor;
}
}
let isAlertShown = false;
let disablePrefix = app.ui.settings.getSettingValue("KJNodes.disablePrefix")
const LGraphNode = LiteGraph.LGraphNode
function showAlert(message) {
app.extensionManager.toast.add({
severity: 'warn',
summary: "KJ Get/Set",
detail: `${message}. Most likely you're missing custom nodes`,
life: 5000,
})
function showAlertWithThrottle(message, delay) {
if (!isAlertShown) {
isAlertShown = true;
alert(message);
setTimeout(() => isAlertShown = false, delay);
}
}
app.registerExtension({
name: "SetNode",
registerCustomNodes() {
class SetNode extends LGraphNode {
class SetNode {
defaultVisibility = true;
serialize_widgets = true;
drawConnection = false;
@@ -51,8 +45,7 @@ app.registerExtension({
canvas = app.canvas;
menuEntry = "Show connections";
constructor(title) {
super(title)
constructor() {
if (!this.properties) {
this.properties = {
"previousName": ""
@@ -96,11 +89,10 @@ app.registerExtension({
}
}
if (slotType == 2 && !isChangeConnect) {
if (this.outputs && this.outputs[slot]) {
this.outputs[slot].type = '*';
this.outputs[slot].name = '*';
}
}
this.outputs[slot].type = '*';
this.outputs[slot].name = '*';
}
//On Connect
if (link_info && node.graph && slotType == 1 && isChangeConnect) {
const fromNode = node.graph._nodes.find((otherNode) => otherNode.id == link_info.origin_id);
@@ -123,7 +115,7 @@ app.registerExtension({
setColorAndBgColor.call(this, type);
}
} else {
showAlert("node input undefined.")
alert("Error: Set node input undefined. Most likely you're missing custom nodes");
}
}
if (link_info && node.graph && slotType == 2 && isChangeConnect) {
@@ -135,7 +127,7 @@ app.registerExtension({
this.outputs[0].type = type;
this.outputs[0].name = type;
} else {
showAlert('node output undefined');
alert("Error: Get Set node output undefined. Most likely you're missing custom nodes");
}
}
@@ -328,8 +320,6 @@ app.registerExtension({
];
}
// Provide a default link object with necessary properties, to avoid errors as link can't be null anymore
const defaultLink = { type: 'default', color: this.slotColor };
for (const getter of this.currentGetters) {
if (!this.flags.collapsed) {
@@ -350,7 +340,7 @@ app.registerExtension({
ctx,
start_node_slotpos,
end_node_slotpos,
defaultLink,
null,
false,
null,
this.slotColor,
@@ -375,7 +365,7 @@ app.registerExtension({
app.registerExtension({
name: "GetNode",
registerCustomNodes() {
class GetNode extends LGraphNode {
class GetNode {
defaultVisibility = true;
serialize_widgets = true;
@@ -384,8 +374,7 @@ app.registerExtension({
currentSetter = null;
canvas = app.canvas;
constructor(title) {
super(title)
constructor() {
if (!this.properties) {
this.properties = {};
}
@@ -450,7 +439,7 @@ app.registerExtension({
if (this.outputs[0].type !== '*' && this.outputs[0].links) {
this.outputs[0].links.filter(linkId => {
const link = node.graph.links[linkId];
return link && (!link.type.split(",").includes(this.outputs[0].type) && link.type !== '*');
return link && (link.type !== this.outputs[0].type && link.type !== '*');
}).forEach(linkId => {
node.graph.removeLink(linkId);
});
@@ -481,9 +470,6 @@ app.registerExtension({
getInputLink(slot) {
const setter = this.findSetter(this.graph);
if (this.mode !== 0) {
return null;
}
if (setter) {
const slotInfo = setter.inputs[slot];
@@ -491,8 +477,8 @@ app.registerExtension({
return link;
} else {
const errorMessage = "No SetNode found for " + this.widgets[0].value + "(" + this.type + ")";
showAlert(errorMessage);
//throw new Error(errorMessage);
showAlertWithThrottle(errorMessage, 5000);
throw new Error(errorMessage);
}
}
onAdded(graph) {
@@ -523,11 +509,6 @@ app.registerExtension({
}
onDrawForeground(ctx, lGraphCanvas) {
if (this.mode === 4) {
console.log(`Mode is ${this.mode}, setting to disabled`)
this.mode = 2;
return null;
}
if (this.drawConnection) {
this._drawVirtualLink(lGraphCanvas, ctx);
}
@@ -539,9 +520,6 @@ app.registerExtension({
// }
_drawVirtualLink(lGraphCanvas, ctx) {
if (!this.currentSetter) return;
// Provide a default link object with necessary properties, to avoid errors as link can't be null anymore
const defaultLink = { type: 'default', color: this.slotColor };
let start_node_slotpos = this.currentSetter.getConnectionPos(false, 0);
start_node_slotpos = [
@@ -553,7 +531,7 @@ app.registerExtension({
ctx,
start_node_slotpos,
end_node_slotpos,
defaultLink,
null,
false,
null,
this.slotColor
+526 -1096
View File
File diff suppressed because it is too large Load Diff