12 Commits
11 changed files with 550 additions and 18 deletions
+10
View File
@@ -1,8 +1,13 @@
from custom_nodes.KepPromptLang.lib.actions.avg import AverageAction
from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction
from custom_nodes.KepPromptLang.lib.actions.mult import MultiplyAction from custom_nodes.KepPromptLang.lib.actions.mult import MultiplyAction
from custom_nodes.KepPromptLang.lib.actions.neg import NegAction from custom_nodes.KepPromptLang.lib.actions.neg import NegAction
from custom_nodes.KepPromptLang.lib.actions.norm import NormAction from custom_nodes.KepPromptLang.lib.actions.norm import NormAction
from custom_nodes.KepPromptLang.lib.actions.project import ProjectAction
from custom_nodes.KepPromptLang.lib.actions.rand import RandAction from custom_nodes.KepPromptLang.lib.actions.rand import RandAction
from custom_nodes.KepPromptLang.lib.actions.scale_dims import ScaleDims
from custom_nodes.KepPromptLang.lib.actions.set_dims import SetDims
from custom_nodes.KepPromptLang.lib.actions.slerp import SlerpAction
from custom_nodes.KepPromptLang.lib.actions.sum import SumAction from custom_nodes.KepPromptLang.lib.actions.sum import SumAction
from custom_nodes.KepPromptLang.lib.parser.registration import register_action from custom_nodes.KepPromptLang.lib.parser.registration import register_action
@@ -12,3 +17,8 @@ register_action(NegAction)
register_action(NormAction) register_action(NormAction)
register_action(RandAction) register_action(RandAction)
register_action(SumAction) register_action(SumAction)
register_action(SlerpAction)
register_action(AverageAction)
register_action(ScaleDims)
register_action(SetDims)
register_action(ProjectAction)
+9
View File
@@ -1,3 +1,6 @@
from typing import List
import torch
from torch import Tensor from torch import Tensor
from torch.nn import Embedding from torch.nn import Embedding
@@ -9,3 +12,9 @@ def get_embedding(seg_or_action: SegOrAction, embedding_module: Embedding) -> Te
if isinstance(seg_or_action, Action): if isinstance(seg_or_action, Action):
return seg_or_action.get_result(embedding_module) return seg_or_action.get_result(embedding_module)
return seg_or_action.get_embeddings(embedding_module) return seg_or_action.get_embeddings(embedding_module)
def get_embedding_for_segments(
segments: List[SegOrAction], embedding_module: Embedding
) -> Tensor:
return torch.cat([get_embedding(segment, embedding_module) for segment in segments], dim=1)
+100
View File
@@ -0,0 +1,100 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import MultiArgAction, Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.actions.utils import slerp
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class AverageAction(MultiArgAction):
grammar = 'avg(" arg "|" arg "|" arg ")"'
name = "avg"
chars = ["+", "+"]
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
super().__init__(args)
if len(args) != 3:
raise ValueError("Average action should have exactly three arguments(2 vectors and a weight)")
self.first_arg = args[0]
self.second_arg = args[1]
self._parse_weight(args[2])
self._validate_args()
def _parse_weight(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Average weight should have exactly one segment")
weight_seg_or_action = arg[0]
if isinstance(weight_seg_or_action, Action):
raise ValueError("Average weight should not have an action as an argument")
try:
self.parsed_weight = float(weight_seg_or_action.text)
except ValueError:
raise ValueError("Average should have an integer/float as the weight")
def _validate_args(self) -> None:
first_arg_token_length = sum(seg_or_action.token_length() for seg_or_action in self.first_arg)
second_arg_token_length = sum(seg_or_action.token_length() for seg_or_action in self.second_arg)
if first_arg_token_length != second_arg_token_length:
raise ValueError(f"Average start and end arguments should have the same length. Got {start_arg_token_length} and {end_arg_token_length}")
if self.parsed_weight < 0 or self.parsed_weight > 1:
print(f"WARNING: Average weight should be between 0 and 1. Got {self.parsed_weight}")
def token_length(self) -> int:
# Average interpolates between the embeddings of the start and end segments, so the length is the length of the start segment
return sum(seg_or_action.token_length() for seg_or_action in self.first_arg)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the start segment
all_start_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.first_arg
]
start_embedding = torch.cat(all_start_embeddings, dim=1)
# Calculate the embeddings for the end segment
all_end_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.second_arg
]
end_embedding = torch.cat(all_end_embeddings, dim=1)
# Perform the weighted average
result = start_embedding * (1 - self.parsed_weight) + end_embedding * self.parsed_weight
return result
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.additional_args))})"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_arg, Action):
base_segment_repr = self.base_arg.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.additional_args, Action):
target_repr = self.additional_args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
+3 -3
View File
@@ -26,17 +26,17 @@ class MultiplyAction(MultiArgAction):
def _parse_multiplier(self, arg: List[SegOrAction]) -> None: def _parse_multiplier(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1: if len(arg) != 1:
raise ValueError("Multiply action first argument should have exactly one segment") raise ValueError("Multiply actions multiplier should have exactly one segment")
multiplier_seg_or_action = arg[0] multiplier_seg_or_action = arg[0]
if isinstance(multiplier_seg_or_action, Action): if isinstance(multiplier_seg_or_action, Action):
raise ValueError("Multiply action should not have an action as an argument") raise ValueError("Multiply actions multiplier must be a number")
try: try:
self.parsed_multiplier = float(multiplier_seg_or_action.text) self.parsed_multiplier = float(multiplier_seg_or_action.text)
except ValueError: except ValueError:
raise ValueError("Multiply action should have an integer/float as the first argument") raise ValueError("Multiply action should have an integer/float as the multiplier")
def token_length(self) -> int: def token_length(self) -> int:
""" """
+109
View File
@@ -0,0 +1,109 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import MultiArgAction, Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import (
get_embedding,
get_embedding_for_segments,
)
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.actions.utils import slerp
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class ProjectAction(MultiArgAction):
grammar = 'project(" arg "|" arg ("|" arg ")?)"'
name = "project"
chars = ["+", "+"]
weight = 1.0
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
super().__init__(args)
num_args = len(args)
if num_args != 3 and num_args != 2:
raise ValueError(
"Project action should have exactly three arguments(2 vectors and a weight)"
)
self.source_argument = args[0]
self.source_argument_token_length = sum(
seg_or_action.token_length() for seg_or_action in self.source_argument
)
self.onto_argument = args[1]
self.onto_argument_token_length = sum(
seg_or_action.token_length() for seg_or_action in self.onto_argument
)
if num_args == 3:
self._parse_weight(args[2])
self._validate_args()
def _parse_weight(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Project weight should have exactly one segment")
weight_seg_or_action = arg[0]
if isinstance(weight_seg_or_action, Action):
raise ValueError("Project weight should not have an action as an argument")
try:
self.weight = float(weight_seg_or_action.text)
except ValueError:
raise ValueError("Project should have an integer/float as the weight")
def _validate_args(self) -> None:
if (
self.source_argument_token_length != self.onto_argument_token_length
and self.onto_argument_token_length != 1
):
raise ValueError(
f"Project source and target arguments should have the same token lengths, or target should be one token. Got {self.source_argument_token_length} source tokens and {self.onto_argument_token_length} target tokens"
)
def token_length(self) -> int:
# Project projects the source onto the target, so the length of the result is the length of source
return self.source_argument_token_length
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the start segment
source_embedding = get_embedding_for_segments(
self.source_argument, embedding_module
).to(dtype=torch.float32)
onto_embedding = get_embedding_for_segments(
self.onto_argument, embedding_module
).to(dtype=torch.float32)
# Perform the projection
return torch.mul(
torch.mul(source_embedding, onto_embedding)
/ torch.mul(onto_embedding, onto_embedding),
onto_embedding,
).to(dtype=torch.float16)
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.additional_args))})"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_arg, Action):
base_segment_repr = self.base_arg.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.additional_args, Action):
target_repr = self.additional_args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
+72
View File
@@ -0,0 +1,72 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import MultiArgAction, Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class ScaleDims(MultiArgAction):
grammar = 'scaleDims(" arg ("|" arg)* ")"'
name = "scaleDims"
description = "Scales the specified dimensions of the input embeddings by the specified amount"
example = "'The scaleDims(cat|4,1.5|76,1.2) is happy' scales the 4th dimension by 1.5 and the 76th dimension by 1.2 for the word 'cat'"
chars = ["-", "-"]
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
super().__init__(args)
self.base_arg = args[0]
self._parse_scale_args(args[1:])
def _parse_scale_args(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
# scaleDims scale args should have format of "<dim>,<scale>" where dim is the dimension to scale and scale is the amount to scale it by
# scaleDims(some words|4,1.5|76,1.2)
self.scale_args = []
for arg in args:
if isinstance(arg, Action):
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got an action")
if len(arg) != 1:
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got multiple segments")
extracted_arg = arg[0]
assert isinstance(extracted_arg, PromptSegment)
if "," not in extracted_arg.text:
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got a segment with no comma: " + extracted_arg.text)
# Split prompt segment into text and scale args
dim, scale = extracted_arg.text.split(",")
try:
# TODO: Check that dim is within the bounds of the embedding
parsed_dim = int(dim)
except ValueError:
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got a segment with a non-integer dim: " + str(dim))
try:
parsed_scale = float(scale)
except ValueError:
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got a segment with a non-float scale: " + str(scale))
self.scale_args.append((parsed_dim, parsed_scale))
def token_length(self) -> int:
# scaleDims modifies the embeddings of the base segment, so the length is the length of the base segment
return sum(seg_or_action.token_length() for seg_or_action in self.base_arg)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the base segment
all_base_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.base_arg
]
base_embeddings = torch.cat(all_base_embeddings, dim=1)
for dim, scale in self.scale_args:
base_embeddings[0, :, dim] *= scale
return base_embeddings
+72
View File
@@ -0,0 +1,72 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import MultiArgAction, Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class SetDims(MultiArgAction):
grammar = 'setDims(" arg ("|" arg)* ")"'
name = "setDims"
description = "Sets the specified dimensions of the input embeddings to the specified value"
example = "'The scaleDims(cat|4,1.5|76,1.2) is happy' scales the 4th dimension by 1.5 and the 76th dimension by 1.2 for the word 'cat'"
chars = ["-", "-"]
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
super().__init__(args)
self.base_arg = args[0]
self._parse_value_args(args[1:])
def _parse_value_args(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
# setDims args should have format of "<dim>,<value>" where dim is the dimension to set and value is the value to set it to
# setDims(some words|4,-0.01254|76,1.2)
self.value_args = []
for arg in args:
if isinstance(arg, Action):
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got an action")
if len(arg) != 1:
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got multiple segments")
extracted_arg = arg[0]
assert isinstance(extracted_arg, PromptSegment)
if "," not in extracted_arg.text:
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got a segment with no comma: " + extracted_arg.text)
# Split prompt segment into text and value args
dim, value = extracted_arg.text.split(",")
try:
# TODO: Check that dim is within the bounds of the embedding
parsed_dim = int(dim)
except ValueError:
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got a segment with a non-integer dim: " + str(dim))
try:
parsed_value = float(value)
except ValueError:
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got a segment with a non-float scale: " + str(value))
self.value_args.append((parsed_dim, parsed_value))
def token_length(self) -> int:
# setDims modifies the embeddings of the base segment, so the length is the length of the base segment
return sum(seg_or_action.token_length() for seg_or_action in self.base_arg)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the base segment
all_base_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.base_arg
]
base_embeddings = torch.cat(all_base_embeddings, dim=1)
for dim, value in self.value_args:
base_embeddings[0, :, dim] = value
return base_embeddings
+100
View File
@@ -0,0 +1,100 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import MultiArgAction, Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.actions.utils import slerp
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class SlerpAction(MultiArgAction):
grammar = 'slerp(" arg "|" arg "|" arg ")"'
name = "slerp"
chars = ["+", "+"]
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
super().__init__(args)
if len(args) != 3:
raise ValueError("Slerp action should have exactly three arguments(2 vectors and a weight)")
self.start_argument = args[0]
self.end_argument = args[1]
self._parse_weight(args[2])
self._validate_args()
def _parse_weight(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Slerp weight should have exactly one segment")
weight_seg_or_action = arg[0]
if isinstance(weight_seg_or_action, Action):
raise ValueError("Slerp weight should not have an action as an argument")
try:
self.parsed_weight = float(weight_seg_or_action.text)
except ValueError:
raise ValueError("Slerp should have an integer/float as the weight")
def _validate_args(self) -> None:
start_arg_token_length = sum(seg_or_action.token_length() for seg_or_action in self.start_argument)
end_arg_token_length = sum(seg_or_action.token_length() for seg_or_action in self.end_argument)
if start_arg_token_length != end_arg_token_length:
raise ValueError(f"Slerp start and end arguments should have the same length. Got {start_arg_token_length} and {end_arg_token_length}")
if self.parsed_weight < 0 or self.parsed_weight > 1:
print(f"WARNING: Slerp weight should be between 0 and 1. Got {self.parsed_weight}")
def token_length(self) -> int:
# Slerp interpolates between the embeddings of the start and end segments, so the length is the length of the start segment
return sum(seg_or_action.token_length() for seg_or_action in self.start_argument)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the start segment
all_start_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.start_argument
]
start_embedding = torch.cat(all_start_embeddings, dim=1)
# Calculate the embeddings for the end segment
all_end_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.end_argument
]
end_embedding = torch.cat(all_end_embeddings, dim=1)
# Perform the slerp
result = slerp(self.parsed_weight, start_embedding, end_embedding)
return result
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.additional_args))})"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_arg, Action):
base_segment_repr = self.base_arg.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.additional_args, Action):
target_repr = self.additional_args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
+30
View File
@@ -1,8 +1,38 @@
from typing import List from typing import List
import torch
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
def batch_size_info(batch: List[SegOrAction]): def batch_size_info(batch: List[SegOrAction]):
for segment in batch: for segment in batch:
print("Token Len: " + str(segment.token_length())) print("Token Len: " + str(segment.token_length()))
print(segment.depth_repr()) print(segment.depth_repr())
def slerp(val: float, low: torch.Tensor, high: torch.Tensor, epsilon=1e-5):
# Convert val to tensor and clamp between 0 and 1
val = torch.tensor(val, dtype=torch.float32).clamp(0, 1)
# Normalize the vectors
low_norm = low / torch.norm(low, dim=-1, keepdim=True)
high_norm = high / torch.norm(high, dim=-1, keepdim=True)
# Calculate the cosine of the angle between the vectors
dot = (low_norm * high_norm).sum(-1, keepdim=True)
# Clamp to prevent numerical errors
dot = torch.clamp(dot, -1, 1)
omega = torch.acos(dot)
# Slerp formula
sin_omega = torch.sin(omega)
scale_0 = torch.sin((1.0 - val) * omega) / (sin_omega + epsilon)
scale_1 = torch.sin(val * omega) / (sin_omega + epsilon)
# Handle the case where omega is small (the vectors are close)
close_condition = sin_omega < epsilon
scale_0 = torch.where(close_condition, 1.0 - val, scale_0)
scale_1 = torch.where(close_condition, val, scale_1)
return scale_0 * low + scale_1 * high
+6
View File
@@ -5,6 +5,12 @@ from custom_nodes.KepPromptLang.lib.action.base import Action
action_registry: Dict[str, Type[Action]] = {} action_registry: Dict[str, Type[Action]] = {}
def register_action(action: Type[Action]) -> None: def register_action(action: Type[Action]) -> None:
"""
:rtype: object
"""
if action.name in action_registry:
raise ValueError(f"Action {action.name} already registered")
action_registry[str(action.name)] = action action_registry[str(action.name)] = action
def get_action_by_name(name: str) -> Type[Action]: def get_action_by_name(name: str) -> Type[Action]:
+39 -15
View File
@@ -1,5 +1,6 @@
import random import random
from typing import List, Tuple import os
from typing import List, Tuple, Any
import numpy as np import numpy as np
from PIL import Image from PIL import Image
@@ -49,9 +50,9 @@ def tensor2img(tensor_img) -> Image.Image:
i_np_arr = np.clip(i, 0, 255, out=i).astype(np.uint8, copy=False) i_np_arr = np.clip(i, 0, 255, out=i).astype(np.uint8, copy=False)
return Image.fromarray(i_np_arr) return Image.fromarray(i_np_arr)
class BuildGif: class BuildGif:
def __init__(self) -> None: def __init__(self) -> None:
self.output_dir = folder_paths.get_output_directory()
pass pass
@classmethod @classmethod
@@ -60,6 +61,7 @@ class BuildGif:
"required": { "required": {
"images": ("IMAGE",), "images": ("IMAGE",),
"split_every": ("INT", {"default": -1}), "split_every": ("INT", {"default": -1}),
"frame_duration": ("INT", {"default": 125}),
"output_mode": ( "output_mode": (
["One Per Split", "Big Grid"], ["One Per Split", "Big Grid"],
{"default": "Big Grid"}, {"default": "Big Grid"},
@@ -68,23 +70,31 @@ class BuildGif:
} }
RELOAD_INST = True RELOAD_INST = True
RETURN_TYPES = ("IMAGE",) RETURN_TYPES = ()
RETURN_NAMES = ("Gifs",) # RETURN_NAMES = ("Gifs",)
INPUT_IS_LIST = True INPUT_IS_LIST = True
FUNCTION = "build_gif" FUNCTION = "build_gif"
OUTPUT_IS_LIST = (True,) # OUTPUT_IS_LIST = (True,)
# OUTPUT_NODE = False OUTPUT_NODE = True
CATEGORY = "List Stuff" CATEGORY = "List Stuff"
@staticmethod def build_gif(self, images: List[Any], split_every: List[int], frame_duration: List[int], output_mode: List[str]):
def build_gif(images: list, split_every: List[int], output_mode: str):
print("Build GIF called!") print("Build GIF called!")
print(f"{type(images)}") print(f"{type(images)}")
if len(split_every) > 1: if len(split_every) > 1:
raise Exception("List input for split every is not supported.") raise Exception("List input for split every is not supported.")
if len(output_mode) > 1:
raise Exception("List input for output_mode is not supported.")
output_mode = output_mode[0]
if len(frame_duration) > 1:
raise Exception("List input for frame_duration is not supported.")
frame_duration = frame_duration[0]
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix="Gif", output_dir=self.output_dir, image_width=0, image_height=0)
split_every_val = split_every[0] split_every_val = split_every[0]
batch_size = images[0].size()[0] batch_size = images[0].size()[0]
if split_every_val == -1: if split_every_val == -1:
@@ -93,8 +103,6 @@ class BuildGif:
else: else:
split_chunks = int(len(images) / split_every_val) split_chunks = int(len(images) / split_every_val)
out = []
num_wide = batch_size num_wide = batch_size
num_tall = split_chunks num_tall = split_chunks
@@ -104,6 +112,7 @@ class BuildGif:
] ]
frames = [] frames = []
results = list()
if output_mode == "Big Grid": if output_mode == "Big Grid":
# For every image in gif # For every image in gif
@@ -122,8 +131,9 @@ class BuildGif:
) )
frames.append(img_frame) frames.append(img_frame)
file = f"{filename}_{counter:05}_"
save_path = ( save_path = (
f"{folder_paths.get_output_directory()}/{random.randint(1, 100)}" f"{os.path.join(full_output_folder, file)}"
) )
frames[0].save( frames[0].save(
f"{save_path}.webp", f"{save_path}.webp",
@@ -133,15 +143,24 @@ class BuildGif:
save_all=True, save_all=True,
append_images=frames[1:], append_images=frames[1:],
optimize=False, optimize=False,
duration=125, duration=frame_duration,
loop=0, loop=0,
) )
results.append({
"filename": f"{file}.webp",
"subfolder": subfolder,
"type": "output"
})
elif output_mode == "One Per Split": elif output_mode == "One Per Split":
for split_idx in range(int(split_chunks)): for split_idx in range(int(split_chunks)):
split_start = split_every_val * split_idx split_start = split_every_val * split_idx
split_end = split_every_val * (split_idx + 1) split_end = split_every_val * (split_idx + 1)
for batch_idx in range(batch_size): for batch_idx in range(batch_size):
save_path = f"{folder_paths.get_output_directory()}/-{batch_idx}-{random.randint(1, 100)}" file = f"{filename}_{counter:05}_"
save_path = (
f"{os.path.join(full_output_folder, file)}"
)
counter += 1
print(save_path) print(save_path)
tensor2img(images[split_start][batch_idx]).save( tensor2img(images[split_start][batch_idx]).save(
f"{save_path}.webp", f"{save_path}.webp",
@@ -151,7 +170,12 @@ class BuildGif:
for nested_batch in images[split_start + 1 : split_end] for nested_batch in images[split_start + 1 : split_end]
], ],
optimize=False, optimize=False,
duration=125, duration=frame_duration,
loop=0, loop=0,
) )
return (out,) results.append({
"filename": f"{file}.webp",
"subfolder": subfolder,
"type": "output"
})
return { "ui": { "images": results } }