55 Commits
Author SHA1 Message Date
Michael Poutre f2256932c8 Working monaco 2023-09-12 18:28:40 -07:00
Michael Poutre 94c46fe1ba Working breakpoints 2023-09-12 18:28:40 -07:00
Michael Poutre ca11ea8e82 Working webpack 2023-09-12 18:28:40 -07:00
Michael Poutre fb363aed7a Working TS 2023-09-12 18:28:40 -07:00
Michael Poutre d67d8600ab Merge branch 'func/setDims' 2023-09-05 00:14:13 -07:00
Michael Poutre 6444e2890a feat(func/setDims): Add func 2023-09-04 23:35:24 -07:00
Michael Poutre 28141f2cbe refactor(nodes): Update Build Gif to show preview of Gif 2023-09-04 18:32:53 -07:00
Michael Poutre 70704f5e68 feat(func/scaleDims): Add scaleDims 2023-09-04 18:15:30 -07:00
Michael Poutre e59aaa499d fix(action reg): Don't allow multiple actions with same name 2023-09-04 18:00:35 -07:00
Michael Poutre c98df8d289 feat(fun/average): Add function 2023-08-31 23:50:45 -07:00
Michael Poutre a7c4bbe332 feat(nodes): Update saving method for build_gif 2023-08-31 23:05:41 -07:00
Michael Poutre ce795c52bc refactor(func/mult): Update error messages 2023-08-31 22:51:27 -07:00
Michael Poutre ef9693ec73 feat(nodes): Add frame_duration to build_gif 2023-08-31 22:51:16 -07:00
Michael Poutre 5362ac75fa feat(func/slerp): Add function 2023-08-31 22:50:47 -07:00
Michael Poutre 4fcdee837b fix(func/sum): Allow for passing a single argument 2023-08-31 21:47:22 -07:00
Michael Poutre 6b2bf7485e fix(promptsegment): Put tokens on same device as embeddingmodule weights 2023-08-31 14:19:49 -07:00
Michael Poutre d330663624 fix(promtsegment): Put tokens on the GPU 2023-08-31 14:13:39 -07:00
Michael Poutre 5de8d6821c fix(typing): tuple -> Tuple for <=python3.8 2023-08-31 00:33:58 -07:00
Michael Poutre 0f36544b0e feat(func): Add Mult 2023-08-31 00:29:58 -07:00
Michael Poutre 05920e0392 refactor(nodes): Cleanup some Mypy type errors 2023-08-30 23:42:34 -07:00
Michael Poutre 9ccf5d2583 feat: Use a more generic grammar to allow for easier interoperability 2023-08-30 23:40:39 -07:00
Michael Poutre f185b39f06 feat(func/rand): Add support for defining range for rand 2023-08-30 21:06:23 -07:00
Michael Poutre af761a5620 feat(func): Add rand(<token_length>) 2023-08-30 20:29:46 -07:00
Michael Poutre fee9c56abd refactor: Rename repo to match github 2023-08-30 18:44:11 -07:00
Michael Poutre 6879267391 refactor(Tests): Only run 1 step in sampler 2023-08-30 18:33:59 -07:00
Michael Poutre a40ed34eac fix(CI): Source venv for all python runs 2023-08-30 18:26:11 -07:00
Michael Poutre 94b2347d01 feat(CI): Cache venv instead of pip cache 2023-08-30 18:23:34 -07:00
Michael Poutre bcc083f2fa fix: Use typing.List for backwards compatibility 2023-08-30 18:09:10 -07:00
Michael Poutre f91d3aa233 fix(Typing): Replace | syntax with Union 2023-08-29 16:23:05 -07:00
Michael Poutre b497a81b43 fix(CI): PYTHONBUFFERED to correct step 2023-08-29 16:13:05 -07:00
Michael Poutre 4076d2beb1 fix(CI): Try sleeping for 30s maybe? 2023-08-29 16:03:06 -07:00
Michael Poutre e644f219dd fix(CI): Set PYTHONUNBUFFERED=1 for running ComfyUI server 2023-08-29 15:50:30 -07:00
Michael Poutre 9e7e79a1a1 fix(tests): Read error body, then attempt to load JSON 2023-08-29 15:30:10 -07:00
Michael Poutre f18ef4b292 feat(CI): Archive server.log 2023-08-29 15:18:49 -07:00
Michael Poutre 78c14ce95a feat(CI): Add all supported python version 2023-08-29 15:18:39 -07:00
Michael Poutre 60bc4e5a4d fix(tests): Log error when JSON decode fails 2023-08-29 15:18:22 -07:00
Michael Poutre c11a6c5178 fix(CI): Don't fail fast on multi-python test 2023-08-29 15:11:50 -07:00
Michael Poutre c0adb16e30 feat(CI): Test python 3.9-11 2023-08-29 15:08:27 -07:00
Michael Poutre df98b07ffc fix(ClipTransformer): Fix causal_map change between transformer versions 2023-08-28 22:28:10 -07:00
Michael Poutre 79f17aac60 workflows: Install correct websocket library 2023-08-28 22:09:18 -07:00
Michael Poutre 7804426fa2 workflows: Better output and pass error to actions 2023-08-28 22:01:35 -07:00
Michael Poutre 5be4787a12 fix(ClipTransformer): Update with transformers library 2023-08-28 22:01:16 -07:00
Michael Poutre fd2316c8fa workflows: Fix double ext... 2023-08-28 21:38:44 -07:00
Michael Poutre 137a3f0f24 workflows: Error handling on script 2023-08-28 21:35:08 -07:00
Michael Poutre c422e6000c Move up Debug actino 2023-08-28 21:26:12 -07:00
Michael Poutre 5054bfdb82 actions: Fix again 2023-08-28 21:23:33 -07:00
Michael Poutre 7a854c7fea workflow: Open json relative to script file 2023-08-28 21:19:46 -07:00
Michael Poutre 9f12d2f16d workflow: Don't use cache for SD - To slow 2023-08-28 21:15:00 -07:00
Michael Poutre bd369930bd Fix run_workflow.py 2023-08-28 21:14:39 -07:00
Michael Poutre 9bc7f54bcf Workflow: Sleep longer 2023-08-28 21:08:21 -07:00
Michael Poutre 337dad1cc9 Add more files for workflow testing 2023-08-28 21:00:18 -07:00
Michael Poutre ade09bf806 update(clip_model): Sync with upstream 2023-08-28 20:19:52 -07:00
Michael Poutre 88c3804446 New test node 2023-08-28 19:54:14 -07:00
Michael Poutre acbaf7cefe First one didn't show up for some reason.. 2023-08-28 19:07:49 -07:00
Michael Poutre 1f1e74cd30 Merge branch 'gh-workflow' 2023-08-28 19:05:22 -07:00
317 changed files with 23725 additions and 173 deletions
+85
View File
@@ -0,0 +1,85 @@
name: Run Test Workflow
on: [workflow_dispatch]
jobs:
Test:
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
python-version: [ "3.7", "3.8", "3.9", "3.10", "3.11" ]
steps:
- name: Clone Upstream
uses: actions/checkout@v3
with:
repository: comfyanonymous/ComfyUI
ref: master
fetch-depth: 0
- name: Clone Node
uses: actions/checkout@v3
with:
ref: master
fetch-depth: 0
path: custom_nodes/KepPromptLang
- name: Setup Python
uses: actions/setup-python@v4
with:
# Version range or exact version of Python or PyPy to use, using SemVer's version range syntax. Reads from .python-version if unset.
python-version: ${{ matrix.python-version }}
- name: Cache virtualenv
uses: actions/cache@v3
id: cache-venv
with:
path: ./.venv/
key: ${{ runner.os }}-venv-${{ matrix.python-version }}-${{ hashFiles('**/requirements.txt') }}
restore-keys: |
${{ runner.os }}-venv-${{ matrix.python-version }}-
- name: Install Requirements
if: steps.cache-venv.outputs.cache-hit != 'true'
run: |
python -m venv ./.venv
source ./.venv/bin/activate
pip install torch --index-url https://download.pytorch.org/whl/cpu
pip install -r requirements.txt
pip install -r custom_nodes/KepPromptLang/requirements.txt
pip install huggingface_hub websocket-client
# - name: Cache SD Checkpoint
# uses: actions/cache@v3
# with:
# path: |
# models/checkpoints
# key: ${{ runner.os }}-sd-15-checkpoint
- name: Check and Download Model
run: |
source ./.venv/bin/activate
python custom_nodes/KepPromptLang/test_files/check_and_download_model.py
- name: Run in Background
env:
PYTHONUNBUFFERED: 1
run: |
source ./.venv/bin/activate
python main.py --cpu &> server.log &
sleep 10
# - name: Setup upterm session
# uses: lhotari/action-upterm@v1
- name: Run Workflow
run: |
source ./.venv/bin/activate
python custom_nodes/KepPromptLang/test_files/run_workflow.py
- name: Upload Comfy Server Log
if: always()
uses: actions/upload-artifact@v3
with:
name: comfy-server-log-${{ matrix.python-version }}
path: server.log
+4
View File
@@ -1,9 +1,13 @@
from .nodes import (
BuildGif,
SpecialClipLoader,
MonacoPrompt,
)
NODE_CLASS_MAPPINGS = {
"Build Gif": BuildGif,
"Special CLIP Loader": SpecialClipLoader,
"Monaco Prompt": MonacoPrompt,
}
WEB_DIRECTORY = ("./web/dist", ["app.bundle.js"])
+22
View File
@@ -0,0 +1,22 @@
from custom_nodes.KepPromptLang.lib.actions.avg import AverageAction
from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction
from custom_nodes.KepPromptLang.lib.actions.mult import MultiplyAction
from custom_nodes.KepPromptLang.lib.actions.neg import NegAction
from custom_nodes.KepPromptLang.lib.actions.norm import NormAction
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.parser.registration import register_action
register_action(DiffAction)
register_action(MultiplyAction)
register_action(NegAction)
register_action(NormAction)
register_action(RandAction)
register_action(SumAction)
register_action(SlerpAction)
register_action(AverageAction)
register_action(ScaleDims)
register_action(SetDims)
+35 -18
View File
@@ -1,16 +1,31 @@
from abc import ABC, abstractmethod
from typing import Union
from enum import Enum
from typing import Union, List
from torch import Tensor
from torch.nn import Embedding
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class ActionArity(Enum):
NONE = 0
SINGLE = 1
MULTI = 2
class Action(ABC):
@property
@abstractmethod
def chars(self) -> list[str] | None:
def chars(self) -> Union[List[str], None]:
pass
@property
@abstractmethod
def arity(self) -> ActionArity:
"""
Determines the arity of the action. This is used to determine how many arguments the action supports.
:return:
"""
pass
@property
@@ -36,13 +51,21 @@ class Action(ABC):
pass
@abstractmethod
def get_all_segments(self) -> list[PromptSegment]:
def get_all_segments(self) -> List[PromptSegment]:
"""
Get all segments, including nested segments.
:return:
"""
pass
@abstractmethod
def __init__(self, *args, **kwargs) -> None:
"""
Initialize the action. This is called when the action is parsed from the prompt.
:param args: The arguments for the action.
"""
pass
@abstractmethod
def get_result(self, embedding_module: Embedding) -> Tensor:
"""
@@ -57,7 +80,8 @@ class Action(ABC):
class SingleArgAction(Action, ABC):
def get_all_segments(self) -> list[PromptSegment]:
arity = ActionArity.SINGLE
def get_all_segments(self) -> List[PromptSegment]:
segments = []
for seg_or_action in self.arg:
if isinstance(seg_or_action, Action):
@@ -66,7 +90,7 @@ class SingleArgAction(Action, ABC):
segments.append(seg_or_action)
return segments
def __init__(self, arg: list[PromptSegment | Action]):
def __init__(self, arg: List[Union[PromptSegment, Action]]):
# TODO: Target is a list now... what does this mean for us..
self.arg = arg
@@ -74,15 +98,10 @@ class SingleArgAction(Action, ABC):
return f"{self.name}({self.arg})"
class MultiArgAction(Action, ABC):
def get_all_segments(self) -> list[PromptSegment]:
arity = ActionArity.MULTI
def get_all_segments(self) -> List[PromptSegment]:
segments = []
for seg_or_action in self.base_segment:
if isinstance(seg_or_action, Action):
segments.extend(seg_or_action.get_all_segments())
else:
segments.append(seg_or_action)
for arg in self.args:
for arg in self.all_args:
for seg_or_action in arg:
if isinstance(seg_or_action, Action):
segments.extend(seg_or_action.get_all_segments())
@@ -93,8 +112,6 @@ class MultiArgAction(Action, ABC):
def __init__(
self,
base_segment: list[PromptSegment | Action],
args: list[list[Union[PromptSegment, Action]]],
args: List[List[Union[PromptSegment, Action]]],
):
self.base_segment = base_segment
self.args = args
self.all_args = args
+2 -2
View File
@@ -1,8 +1,8 @@
from torch import Tensor
from torch.nn import Embedding
from custom_nodes.ClipStuff.lib.action.base import Action
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
def get_embedding(seg_or_action: SegOrAction, embedding_module: Embedding) -> Tensor:
+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
+24 -13
View File
@@ -1,8 +1,12 @@
from typing import Union, List
import torch
from torch.nn import Embedding
from custom_nodes.ClipStuff.lib.action.base import Action, MultiArgAction
from custom_nodes.ClipStuff.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.action.base import Action, MultiArgAction
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class DiffAction(MultiArgAction):
@@ -10,20 +14,26 @@ class DiffAction(MultiArgAction):
name = "diff"
chars = ["-", "-"]
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
super().__init__(args)
self.base_arg = args[0]
self.additional_args = args[1:]
def token_length(self) -> int:
# Sum adds to 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_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_segment
for seg_or_action in self.base_arg
]
result = torch.cat(all_base_embeddings, dim=1)
for arg in self.args:
for arg in self.additional_args:
all_arg_embeddings = [
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
]
@@ -44,23 +54,24 @@ class DiffAction(MultiArgAction):
return result
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
# return f"sum(\n\tbase_segment={self.base_segment},\n\tadditional_args={self.additional_args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.args))})"
return f"sum({', '.join(map(str, self.additional_args))})"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_segment, Action):
base_segment_repr = self.base_segment.depth_repr(depth + 1)
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_segment.depth_repr()},\n"
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.args, Action):
target_repr = self.args.depth_repr(depth + 1)
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.args.depth_repr()},\n"
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
+62
View File
@@ -0,0 +1,62 @@
from typing import List
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import (
Action,
MultiArgAction,
)
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class MultiplyAction(MultiArgAction):
grammar = 'mult(" arg+ ")"'
name = "mult"
chars = ["[", "]"]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 2:
raise ValueError("Multiply action should have exactly two arguments")
self.target_arg = args[0]
self._parse_multiplier(args[1])
def _parse_multiplier(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Multiply actions multiplier should have exactly one segment")
multiplier_seg_or_action = arg[0]
if isinstance(multiplier_seg_or_action, Action):
raise ValueError("Multiply actions multiplier must be a number")
try:
self.parsed_multiplier = float(multiplier_seg_or_action.text)
except ValueError:
raise ValueError("Multiply action should have an integer/float as the multiplier")
def token_length(self) -> int:
"""
Mult multiplies the embeddings of the base segment, so the length is the length of the base segment
:return:
"""
total_length = 0
for seg_or_action in self.target_arg:
total_length += seg_or_action.token_length()
return total_length
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
all_embeddings = []
for seg_or_action in self.target_arg:
if isinstance(seg_or_action, Action):
all_embeddings.append(seg_or_action.get_result(embedding_module))
else:
all_embeddings.append(seg_or_action.get_embeddings(embedding_module))
target_embeddings = torch.cat(all_embeddings, dim=1)
return target_embeddings * self.parsed_multiplier
+3 -1
View File
@@ -1,7 +1,8 @@
import torch
from torch.nn import Embedding
from custom_nodes.ClipStuff.lib.action.base import Action, SingleArgAction
from custom_nodes.KepPromptLang.lib.action.base import Action, SingleArgAction
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class NegAction(SingleArgAction):
@@ -30,3 +31,4 @@ class NegAction(SingleArgAction):
target_embeddings = torch.cat(all_embeddings, dim=1)
return target_embeddings * -1
+2 -1
View File
@@ -1,10 +1,11 @@
import torch
from torch.nn import Embedding
from custom_nodes.ClipStuff.lib.action.base import (
from custom_nodes.KepPromptLang.lib.action.base import (
Action,
SingleArgAction,
)
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class NormAction(SingleArgAction):
+87
View File
@@ -0,0 +1,87 @@
from typing import List
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import (
Action,
MultiArgAction,
)
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class RandAction(MultiArgAction):
grammar = 'rand(" arg ")"'
name = "rand"
chars = None
parsed_token_length = 0
range_min = 0
range_max = 1
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 1 and len(args) != 3:
raise ValueError("Random action should have exactly one argument or three arguments")
self._parse_token_length(args[0])
if len(args) == 3:
self._parse_range(args[1], args[2])
def _parse_token_length(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Random action first argument should have exactly one segment")
token_length_seg_or_action = arg[0]
if isinstance(token_length_seg_or_action, Action):
raise ValueError("Random action should not have an action as an argument")
try:
self.parsed_token_length = int(token_length_seg_or_action.text)
except ValueError:
raise ValueError("Random action should have an integer as the first argument")
def _parse_range(self, min_arg: List[SegOrAction], max_arg: List[SegOrAction]) -> None:
if len(min_arg) != 1:
raise ValueError("Random action second argument should have exactly one segment")
if len(max_arg) != 1:
raise ValueError("Random action third argument should have exactly one segment")
min_seg_or_action = min_arg[0]
max_seg_or_action = max_arg[0]
if isinstance(min_seg_or_action, Action):
raise ValueError("Random action should not have an action as an argument")
if isinstance(max_seg_or_action, Action):
raise ValueError("Random action should not have an action as an argument")
try:
self.range_min = int(min_seg_or_action.text)
except ValueError:
raise ValueError("Random action should have an integer as the second argument")
try:
self.range_max = int(max_seg_or_action.text)
except ValueError:
raise ValueError("Random action should have an integer as the third argument")
if self.range_min > self.range_max:
raise ValueError("Random action should have the second argument be less than the third argument")
def token_length(self) -> int:
"""
Random returns a random embedding whose length is the number in the argument
:return:
"""
return self.parsed_token_length
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Create random tensor of size
result = torch.empty(1, self.parsed_token_length, embedding_module.embedding_dim).uniform_(self.range_min, self.range_max)
return result
+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
+23 -40
View File
@@ -1,57 +1,39 @@
from typing import Union
from typing import Union, List
import torch
from torch.nn import Embedding
from custom_nodes.ClipStuff.lib.actions.action_utils import get_embedding
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
from custom_nodes.ClipStuff.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.action.base import Action, MultiArgAction
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class SumAction(Action):
grammar = 'sum(" arg ("|" arg)* ")"'
class SumAction(MultiArgAction):
grammar = 'sum(" arg ("|" arg)+ ")"'
name = "sum"
chars = ["+", "+"]
def __init__(
self,
base_segment: list[PromptSegment | Action],
args: list[list[Union[PromptSegment, Action]]],
):
self.base_segment = base_segment
self.args = args
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
super().__init__(args)
self.base_arg = args[0]
self.additional_args = args[1:]
def token_length(self) -> int:
# Sum adds to 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_segment)
def get_all_segments(self) -> list[PromptSegment]:
segments = []
for seg_or_action in self.base_segment:
if isinstance(seg_or_action, Action):
segments.extend(seg_or_action.get_all_segments())
else:
segments.append(seg_or_action)
for arg in self.args:
for seg_or_action in arg:
if isinstance(seg_or_action, Action):
segments.extend(seg_or_action.get_all_segments())
else:
segments.append(seg_or_action)
return segments
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_segment
for seg_or_action in self.base_arg
]
result = torch.cat(all_base_embeddings, dim=1)
for arg in self.args:
for arg in self.additional_args:
all_arg_embeddings = [
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
]
@@ -76,21 +58,22 @@ class SumAction(Action):
# 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.args))})"
return f"sum({', '.join(map(str, self.additional_args))})"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_segment, Action):
base_segment_repr = self.base_segment.depth_repr(depth + 1)
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_segment.depth_repr()},\n"
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.args, Action):
target_repr = self.args.depth_repr(depth + 1)
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.args.depth_repr()},\n"
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
+2 -2
View File
@@ -1,6 +1,6 @@
from typing import Union
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
from custom_nodes.ClipStuff.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.action.base import Action
SegOrAction = Union[PromptSegment, Action]
+34 -2
View File
@@ -1,6 +1,38 @@
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
from typing import List
def batch_size_info(batch: list[SegOrAction]):
import torch
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
def batch_size_info(batch: List[SegOrAction]):
for segment in batch:
print("Token Len: " + str(segment.token_length()))
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
+19 -10
View File
@@ -1,15 +1,16 @@
import contextlib
import os
from typing import List
import torch
from transformers import CLIPTextConfig, modeling_utils
from comfy import model_management
import comfy.ops
from custom_nodes.ClipStuff.lib.action.base import Action
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
from custom_nodes.ClipStuff.lib.fun_clip_stuff import PromptLangTextModel
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.fun_clip_stuff import PromptLangTextModel
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
# Methods with no comment can be assumed to be the same as comfy.sd1_clip.SD1ClipModel
@@ -23,7 +24,7 @@ class PromptLangClipModel(torch.nn.Module):
def __init__(self, version="openai/clip-vit-large-patch14", device="cpu", max_length=77,
freeze=True, layer="last", layer_idx=None, textmodel_json_config=None,
textmodel_path=None): # clip-vit-base-patch32
textmodel_path=None, dtype=None): # clip-vit-base-patch32
super().__init__()
assert layer in self.LAYERS
self.num_layers = 12
@@ -38,18 +39,22 @@ class PromptLangClipModel(torch.nn.Module):
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "clip_config.json")
config = CLIPTextConfig.from_json_file(textmodel_json_config)
self.num_layers = config.num_hidden_layers
with comfy.ops.use_comfy_ops():
with comfy.ops.use_comfy_ops(device, dtype):
with modeling_utils.no_init_weights():
# Our transformer
self.transformer = PromptLangTextModel(config)
if dtype is not None:
self.transformer.to(dtype)
self.max_length = max_length
if freeze:
self.freeze()
self.layer = layer
self.layer_idx = None
self.empty_tokens = [[49406] + [49407] * 76]
self.text_projection = None
self.text_projection = torch.nn.Parameter(torch.eye(self.transformer.get_input_embeddings().weight.shape[1]))
self.logit_scale = torch.nn.Parameter(torch.tensor(4.6055))
self.layer_norm_hidden_state = True
if layer == "hidden":
assert layer_idx is not None
@@ -75,7 +80,7 @@ class PromptLangClipModel(torch.nn.Module):
self.layer_idx = self.layer_default[1]
# Completely changed to support Segments and actions
def set_up_textual_embeddings(self, tokens: list[list[SegOrAction]], current_embeds):
def set_up_textual_embeddings(self, tokens: List[List[SegOrAction]], current_embeds):
next_new_token = token_dict_size = current_embeds.weight.shape[0] - 1
embedding_weights = []
@@ -165,18 +170,22 @@ class PromptLangClipModel(torch.nn.Module):
pooled_output = outputs.pooler_output
if self.text_projection is not None:
pooled_output = pooled_output.to(self.text_projection.device) @ self.text_projection
pooled_output = pooled_output.float().to(self.text_projection.device) @ self.text_projection.float()
return z.float(), pooled_output.float()
def encode(self, tokens):
return self(tokens)
def load_sd(self, sd):
if "text_projection" in sd:
self.text_projection[:] = sd.pop("text_projection")
if "text_projection.weight" in sd:
self.text_projection[:] = sd.pop("text_projection.weight").transpose(0, 1)
return self.transformer.load_state_dict(sd, strict=False)
# Changed from comfy.sd1_clip.ClipTokenWeightEncoder
# Changed to use PromptSegments
def encode_token_weights(self, prompt_segments: list[list[SegOrAction]]):
def encode_token_weights(self, prompt_segments: List[List[SegOrAction]]):
to_encode = [[PromptSegment(text="_Empty Batch_", tokens=self.empty_tokens[0])]]
for batch in prompt_segments:
to_encode.append(batch)
+24 -12
View File
@@ -1,13 +1,17 @@
from typing import Optional, Tuple, Union
from typing import Optional, Tuple, Union, List
import torch
from transformers import CLIPTextConfig
from transformers.modeling_outputs import BaseModelOutputWithPooling
from transformers.models.clip.modeling_clip import _expand_mask, CLIPTextEmbeddings, CLIPTextTransformer, \
CLIPTextModel
from transformers.models.clip.modeling_clip import (
_expand_mask,
CLIPTextEmbeddings,
CLIPTextTransformer,
CLIPTextModel,
)
from custom_nodes.ClipStuff.lib.action.base import Action
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
def slerp(val, low, high):
low = low.unsqueeze(0)
@@ -25,7 +29,7 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
def forward(
self,
input_dicts: Optional[list[list[SegOrAction]]] = None,
input_dicts: Optional[List[List[SegOrAction]]] = None,
input_ids: Optional[torch.LongTensor] = None,
position_ids: Optional[torch.LongTensor] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
@@ -69,7 +73,7 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
def forward(
self,
input_ids: Optional[list[list[SegOrAction]]] = None,
input_ids: Optional[List[List[SegOrAction]]] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
@@ -97,12 +101,20 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
bsz = len(input_ids)
# TODO: Properly gather this
seq_len = 77
# bsz, seq_len = input_shape
input_shape = torch.Size([bsz, seq_len])
# CLIP's text model uses causal mask, prepare it here.
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(
hidden_states.device
)
## VERSION DIFF ##
# transformers < 4.30.0
if hasattr(self, "_build_causal_attention_mask"):
print("Using transformers < 4.30.0")
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(hidden_states.device)
else:
# transformers >= 4.30.0
print("Using transformers >= 4.30.0")
from transformers.models.clip.modeling_clip import _make_causal_mask
causal_attention_mask = _make_causal_mask(input_shape, hidden_states.dtype, device=hidden_states.device)
# expand attention_mask
if attention_mask is not None:
# [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
@@ -160,7 +172,7 @@ class PromptLangTextModel(CLIPTextModel):
def forward(
self,
input_ids: Optional[list[list[SegOrAction]]] = None,
input_ids: Optional[List[List[SegOrAction]]] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
+4 -12
View File
@@ -3,24 +3,16 @@ grammar = """
item: embedding
| WORD
| function
| generic_function
| QUOTED_STRING
function: sum_function
| neg_function
| norm_function
| diff_function
sum_function: "sum(" arg ("|" arg)* ")"
neg_function: "neg(" arg ")"
norm_function: "norm(" arg ")"
diff_function: "diff(" arg ("|" arg)* ")"
generic_function: FUNC_NAME "(" arg ("|" arg)* ")"
arg: item+
embedding: "embedding:" WORD
WORD: /[A-Za-z0-9,_-]+/
FUNC_NAME: /[A-Za-z_-]+/
WORD: /[A-Za-z0-9,_\.-]+/
QUOTED_STRING: /"([^"\\\]*(\\\.[^"\\\]*)*)"|'([^'\\\]*(\\\.[^'\\\]*)*)'/
%import common.WS
+3 -3
View File
@@ -1,4 +1,4 @@
from typing import Union
from typing import Union, List
import torch
from torch import Tensor
@@ -6,7 +6,7 @@ from torch.nn import Embedding
class PromptSegment:
def __init__(self, text: str, tokens: list[Union[int, Tensor]]):
def __init__(self, text: str, tokens: List[Union[int, Tensor]]):
self.text = text
self.tokens = tokens
@@ -17,7 +17,7 @@ class PromptSegment:
return len(self.tokens)
def get_embeddings(self, embedding_module: Embedding) -> Tensor:
tensors = torch.LongTensor(self.tokens).to(torch.device('cpu'))
tensors = torch.LongTensor(self.tokens).to(embedding_module.weight.device)
unsqueezed_tensors = tensors.unsqueeze(0)
return embedding_module(unsqueezed_tensors)
+20
View File
@@ -0,0 +1,20 @@
from typing import Type, Dict
from custom_nodes.KepPromptLang.lib.action.base import Action
action_registry: Dict[str, Type[Action]] = {}
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
def get_action_by_name(name: str) -> Type[Action]:
if name not in action_registry:
raise ValueError(f"Action {name} not found in registry")
return action_registry[name]
+22 -20
View File
@@ -1,13 +1,17 @@
from typing import List
from lark import Transformer, Token
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.ClipStuff.lib.action.base import Action
from custom_nodes.ClipStuff.lib.actions.diff import DiffAction
from custom_nodes.ClipStuff.lib.parser.utils import build_prompt_segment
from custom_nodes.ClipStuff.lib.actions.neg import NegAction
from custom_nodes.ClipStuff.lib.actions.norm import NormAction
from custom_nodes.ClipStuff.lib.actions.sum import SumAction
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.action.base import Action, ActionArity
from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction
from custom_nodes.KepPromptLang.lib.actions.rand import RandAction
from custom_nodes.KepPromptLang.lib.parser.registration import get_action_by_name
from custom_nodes.KepPromptLang.lib.parser.utils import build_prompt_segment
from custom_nodes.KepPromptLang.lib.actions.neg import NegAction
from custom_nodes.KepPromptLang.lib.actions.norm import NormAction
from custom_nodes.KepPromptLang.lib.actions.sum import SumAction
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class PromptTransformer(Transformer):
@@ -18,7 +22,7 @@ class PromptTransformer(Transformer):
super().__init__()
self.tokenizer = tokenizer
def item(self, items: list[Token]):
def item(self, items: List[Token]):
for item in items:
if isinstance(item, Action):
return item
@@ -47,15 +51,13 @@ class PromptTransformer(Transformer):
def embedding(self, items):
return build_prompt_segment(f'{self.tokenizer.embedding_identifier}{items[0]}', self.tokenizer)
def function(self, items):
for item in items:
if item.data == 'sum_function':
return SumAction(item.children[0][:], item.children[1:][:])
elif item.data == 'neg_function':
return NegAction(item.children[0])
elif item.data == 'norm_function':
return NormAction(item.children[0])
elif item.data == 'diff_function':
return DiffAction(item.children[0][:], item.children[1:][:])
else:
raise Exception("Unknown function type: " + str(item.data))
def generic_function(self, items):
action = get_action_by_name(items[0])
if action.arity == ActionArity.SINGLE:
if len(items) != 2:
raise ValueError(f"Action {action.name} should have exactly one argument")
return action(items[1])
elif action.arity == ActionArity.MULTI:
return action(items[1:][:])
else:
raise ValueError(f"Unknown action arity: {action.arity}")
+1 -1
View File
@@ -1,7 +1,7 @@
from lark import Token
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
def flatten_tree(tree):
+7 -5
View File
@@ -1,11 +1,13 @@
from typing import List
from lark import Tree
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.ClipStuff.lib.parser import PromptParser
from custom_nodes.ClipStuff.lib.parser.transformer import PromptTransformer
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.parser import PromptParser
from custom_nodes.KepPromptLang.lib.parser.transformer import PromptTransformer
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class PromptLangTokenizer(SD1Tokenizer):
def __init__(self, tokenizer_path=None, max_length=77, pad_with_end=True, embedding_directory=None, embedding_size=768, embedding_key='clip_l', special_tokens=None):
@@ -16,7 +18,7 @@ class PromptLangTokenizer(SD1Tokenizer):
Returns batches of segments and actions
:return: List of list(batches) of segments and actions
"""
def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs) -> list[list[SegOrAction]]:
def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs) -> List[List[SegOrAction]]:
if self.pad_with_end:
pad_token = self.end_token
else:
+75 -31
View File
@@ -1,4 +1,6 @@
import random
import os
from typing import List, Tuple, Any
import numpy as np
from PIL import Image
@@ -6,18 +8,37 @@ from PIL import Image
import folder_paths
import comfy.sd
import comfy.ops
from custom_nodes.ClipStuff.lib.clip_model import PromptLangClipModel
from custom_nodes.KepPromptLang.lib.clip_model import PromptLangClipModel
from custom_nodes.ClipStuff.lib.tokenizer import PromptLangTokenizer
from custom_nodes.KepPromptLang.lib.tokenizer import PromptLangTokenizer
class EmptyClass:
pass
class SpecialClipLoader:
class MonacoPrompt:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"clip": ("CLIP",),
"prompt": ("MONACO",),
}
}
RETURN_TYPES = ("CONDITIONING",)
FUNCTION = "do_crap"
OUTPUT_IS_LIST = (False,)
CATEGORY = "conditioning"
@staticmethod
def do_crap(clip, prompt):
return (clip,)
class SpecialClipLoader:
@classmethod
def INPUT_TYPES(cls): # type: ignore
return {
"required": {
"source_clip": ("CLIP",),
@@ -30,7 +51,7 @@ class SpecialClipLoader:
CATEGORY = "conditioning"
@staticmethod
def load_clip(source_clip):
def load_clip(source_clip: comfy.sd.CLIP) -> Tuple[comfy.sd.CLIP]:
clip_target = EmptyClass()
clip_target.params = {}
clip_target.clip = PromptLangClipModel
@@ -43,22 +64,23 @@ class SpecialClipLoader:
return (clip,)
def tensor2img(tensor_img):
def tensor2img(tensor_img) -> Image.Image:
i = 255.0 * tensor_img.cpu().numpy()
i_np_arr = np.clip(i, 0, 255, out=i).astype(np.uint8, copy=False)
return Image.fromarray(i_np_arr)
class BuildGif:
def __init__(self):
def __init__(self) -> None:
self.output_dir = folder_paths.get_output_directory()
pass
@classmethod
def INPUT_TYPES(cls):
def INPUT_TYPES(cls): # type: ignore
return {
"required": {
"images": ("IMAGE",),
"split_every": ("INT", {"default": -1}),
"frame_duration": ("INT", {"default": 125}),
"output_mode": (
["One Per Split", "Big Grid"],
{"default": "Big Grid"},
@@ -67,46 +89,53 @@ class BuildGif:
}
RELOAD_INST = True
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("Gifs",)
RETURN_TYPES = ()
# RETURN_NAMES = ("Gifs",)
INPUT_IS_LIST = True
FUNCTION = "build_gif"
OUTPUT_IS_LIST = (True,)
# OUTPUT_NODE = False
# OUTPUT_IS_LIST = (True,)
OUTPUT_NODE = True
CATEGORY = "List Stuff"
@staticmethod
def build_gif(images: list, split_every: list[int], output_mode: str):
def build_gif(self, images: List[Any], split_every: List[int], frame_duration: List[int], output_mode: List[str]):
print("Build GIF called!")
print(f"{type(images)}")
if len(split_every) > 1:
raise Exception("List input for split every is not supported.")
split_every = split_every[0]
batch_size = images[0].size()[0]
if split_every == -1:
split_chunks = 1
split_every = len(images)
else:
split_chunks = int(len(images) / split_every)
if len(output_mode) > 1:
raise Exception("List input for output_mode is not supported.")
output_mode = output_mode[0]
out = []
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]
batch_size = images[0].size()[0]
if split_every_val == -1:
split_chunks = 1
split_every_val = len(images)
else:
split_chunks = int(len(images) / split_every_val)
num_wide = batch_size
num_tall = split_chunks
chunked_batches = [
images[split_every * chunk_idx : split_every * (chunk_idx + 1)]
images[split_every_val * chunk_idx : split_every_val * (chunk_idx + 1)]
for chunk_idx in range(split_chunks)
]
frames = []
results = list()
if output_mode == "Big Grid":
# For every image in gif
for idx_in_chunk in range(split_every):
for idx_in_chunk in range(split_every_val):
img_shape = images[0][0].shape
img_frame = Image.new(
"RGB", size=(num_wide * img_shape[0], num_tall * img_shape[1])
@@ -121,8 +150,9 @@ class BuildGif:
)
frames.append(img_frame)
file = f"{filename}_{counter:05}_"
save_path = (
f"{folder_paths.get_output_directory()}/{random.randint(1, 100)}"
f"{os.path.join(full_output_folder, file)}"
)
frames[0].save(
f"{save_path}.webp",
@@ -132,15 +162,24 @@ class BuildGif:
save_all=True,
append_images=frames[1:],
optimize=False,
duration=125,
duration=frame_duration,
loop=0,
)
results.append({
"filename": f"{file}.webp",
"subfolder": subfolder,
"type": "output"
})
elif output_mode == "One Per Split":
for split_idx in range(int(split_chunks)):
split_start = split_every * split_idx
split_end = split_every * (split_idx + 1)
split_start = split_every_val * split_idx
split_end = split_every_val * (split_idx + 1)
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)
tensor2img(images[split_start][batch_idx]).save(
f"{save_path}.webp",
@@ -150,7 +189,12 @@ class BuildGif:
for nested_batch in images[split_start + 1 : split_end]
],
optimize=False,
duration=125,
duration=frame_duration,
loop=0,
)
return (out,)
results.append({
"filename": f"{file}.webp",
"subfolder": subfolder,
"type": "output"
})
return { "ui": { "images": results } }
+12
View File
@@ -0,0 +1,12 @@
# If the models/checkpoints folder does not have test.txt, then download the model.
import os
from huggingface_hub import hf_hub_download
FILE = "v1-5-pruned-emaonly.safetensors"
REPO_ID = "runwayml/stable-diffusion-v1-5"
if not os.path.exists(f"models/checkpoints/{FILE}"):
print("Downloading model...")
hf_hub_download(repo_id=REPO_ID, filename=FILE, local_dir="models/checkpoints", local_dir_use_symlinks=False)
else:
print("Model already downloaded.")
+109
View File
@@ -0,0 +1,109 @@
#This is an example that uses the websockets api to know when a prompt execution is done
#Once the prompt execution is done it downloads the images using the /history endpoint
import os
import websocket #NOTE: websocket-client (https://github.com/websocket-client/websocket-client)
import uuid
import json
import urllib.request
import urllib.parse
server_address = "127.0.0.1:8188"
client_id = str(uuid.uuid4())
def queue_prompt(prompt):
p = {"prompt": prompt, "client_id": client_id}
data = json.dumps(p).encode('utf-8')
req = urllib.request.Request("http://{}/prompt".format(server_address), data=data)
try:
response = urllib.request.urlopen(req)
return json.loads(response.read())
except urllib.error.HTTPError as e:
print(f"HTTP Error {e.code}: {e.reason}")
error_body = e.read()
# Attempt to read and print the JSON error body
try:
json_error_body = json.loads(error_body)
print(json_error_body)
raise e
except json.JSONDecodeError:
print("Failed to decode error response as JSON.")
print(error_body)
raise e
except Exception as e:
print(f"An unexpected error occurred: {e}")
raise e
def get_image(filename, subfolder, folder_type):
data = {"filename": filename, "subfolder": subfolder, "type": folder_type}
url_values = urllib.parse.urlencode(data)
with urllib.request.urlopen("http://{}/view?{}".format(server_address, url_values)) as response:
return response.read()
def get_history(prompt_id):
with urllib.request.urlopen("http://{}/history/{}".format(server_address, prompt_id)) as response:
return json.loads(response.read())
def get_images(ws, prompt):
prompt_id = queue_prompt(prompt)['prompt_id']
output_images = {}
while True:
out = ws.recv()
if isinstance(out, str):
message = json.loads(out)
if message['type'] == 'executing':
data = message['data']
if data['node'] is None and data['prompt_id'] == prompt_id:
break #Execution is done
else:
continue #previews are binary data
history = get_history(prompt_id)[prompt_id]
for o in history['outputs']:
for node_id in history['outputs']:
node_output = history['outputs'][node_id]
if 'images' in node_output:
images_output = []
for image in node_output['images']:
image_data = get_image(image['filename'], image['subfolder'], image['type'])
images_output.append(image_data)
output_images[node_id] = images_output
return output_images
# Load json from file relative to this script
prompt = json.load(open(os.path.join(os.path.dirname(os.path.realpath(__file__)), "workflow_api.json")))
#set the text prompt for our positive CLIPTextEncode
# prompt["6"]["inputs"]["text"] = "masterpiece best quality man"
#set the seed for our KSampler node
print(queue_prompt(prompt))
ws = websocket.WebSocket()
ws.connect("ws://{}/ws?clientId={}".format(server_address, client_id))
while True:
out = ws.recv()
if isinstance(out, str):
message = json.loads(out)
# print(message)
if message["type"] == "executing" and message["data"]["node"] is None:
print("Execution is done")
break
if message["type"] == "execution_error":
print("Execution error")
print(json.dumps(message["data"], indent=4))
raise Exception("Execution error")
# images = get_images(ws, prompt)
#Commented out code to display the output images:
# for node_id in images:
# for image_data in images[node_id]:
# from PIL import Image
# import io
# image = Image.open(io.BytesIO(image_data))
# image.show()
+84
View File
@@ -0,0 +1,84 @@
{
"1": {
"inputs": {
"text": "A sum(cat|norm(sum(neg(parrot)|rabbit))) outside",
"clip": [
"2",
0
]
},
"class_type": "CLIPTextEncode"
},
"2": {
"inputs": {
"source_clip": [
"4",
1
]
},
"class_type": "Special CLIP Loader"
},
"3": {
"inputs": {
"seed": 556492279461741,
"steps": 1,
"cfg": 8,
"sampler_name": "euler",
"scheduler": "normal",
"denoise": 1,
"model": [
"4",
0
],
"positive": [
"1",
0
],
"negative": [
"1",
0
],
"latent_image": [
"7",
0
]
},
"class_type": "KSampler"
},
"4": {
"inputs": {
"ckpt_name": "v1-5-pruned-emaonly.safetensors"
},
"class_type": "CheckpointLoaderSimple"
},
"5": {
"inputs": {
"samples": [
"3",
0
],
"vae": [
"4",
2
]
},
"class_type": "VAEDecode"
},
"6": {
"inputs": {
"images": [
"5",
0
]
},
"class_type": "PreviewImage"
},
"7": {
"inputs": {
"width": 512,
"height": 512,
"batch_size": 1
},
"class_type": "EmptyLatentImage"
}
}
Generated Executable
+12963
View File
File diff suppressed because one or more lines are too long
+2047
View File
File diff suppressed because it is too large Load Diff
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.

Some files were not shown because too many files have changed in this diff Show More