Compare commits
42
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
61f84e06ed | ||
|
|
4fcdee837b | ||
|
|
6b2bf7485e | ||
|
|
d330663624 | ||
|
|
5de8d6821c | ||
|
|
0f36544b0e | ||
|
|
05920e0392 | ||
|
|
9ccf5d2583 | ||
|
|
f185b39f06 | ||
|
|
af761a5620 | ||
|
|
fee9c56abd | ||
|
|
6879267391 | ||
|
|
a40ed34eac | ||
|
|
94b2347d01 | ||
|
|
bcc083f2fa | ||
|
|
f91d3aa233 | ||
|
|
b497a81b43 | ||
|
|
4076d2beb1 | ||
|
|
e644f219dd | ||
|
|
9e7e79a1a1 | ||
|
|
f18ef4b292 | ||
|
|
78c14ce95a | ||
|
|
60bc4e5a4d | ||
|
|
c11a6c5178 | ||
|
|
c0adb16e30 | ||
|
|
df98b07ffc | ||
|
|
79f17aac60 | ||
|
|
7804426fa2 | ||
|
|
5be4787a12 | ||
|
|
fd2316c8fa | ||
|
|
137a3f0f24 | ||
|
|
c422e6000c | ||
|
|
5054bfdb82 | ||
|
|
7a854c7fea | ||
|
|
9f12d2f16d | ||
|
|
bd369930bd | ||
|
|
9bc7f54bcf | ||
|
|
337dad1cc9 | ||
|
|
ade09bf806 | ||
|
|
88c3804446 | ||
|
|
acbaf7cefe | ||
|
|
1f1e74cd30 |
@@ -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
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
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.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)
|
||||
|
||||
+35
-18
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
+24
-13
@@ -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
|
||||
|
||||
|
||||
@@ -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 action first argument should have exactly one segment")
|
||||
|
||||
multiplier_seg_or_action = arg[0]
|
||||
|
||||
if isinstance(multiplier_seg_or_action, Action):
|
||||
raise ValueError("Multiply action should not have an action as an argument")
|
||||
|
||||
try:
|
||||
self.parsed_multiplier = float(multiplier_seg_or_action.text)
|
||||
except ValueError:
|
||||
raise ValueError("Multiply action should have an integer/float as the first argument")
|
||||
|
||||
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
@@ -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
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
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:
|
||||
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
@@ -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
@@ -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
@@ -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:
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import random
|
||||
from typing import List, Tuple
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
@@ -6,9 +7,9 @@ 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:
|
||||
@@ -17,7 +18,7 @@ class EmptyClass:
|
||||
|
||||
class SpecialClipLoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls): # type: ignore
|
||||
return {
|
||||
"required": {
|
||||
"source_clip": ("CLIP",),
|
||||
@@ -30,7 +31,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,18 +44,18 @@ 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:
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
def INPUT_TYPES(cls): # type: ignore
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
@@ -77,20 +78,20 @@ class BuildGif:
|
||||
CATEGORY = "List Stuff"
|
||||
|
||||
@staticmethod
|
||||
def build_gif(images: list, split_every: list[int], output_mode: str):
|
||||
def build_gif(images: list, split_every: List[int], output_mode: 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]
|
||||
split_every_val = split_every[0]
|
||||
batch_size = images[0].size()[0]
|
||||
if split_every == -1:
|
||||
if split_every_val == -1:
|
||||
split_chunks = 1
|
||||
split_every = len(images)
|
||||
split_every_val = len(images)
|
||||
else:
|
||||
split_chunks = int(len(images) / split_every)
|
||||
split_chunks = int(len(images) / split_every_val)
|
||||
|
||||
out = []
|
||||
|
||||
@@ -98,7 +99,7 @@ class BuildGif:
|
||||
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)
|
||||
]
|
||||
|
||||
@@ -106,7 +107,7 @@ class BuildGif:
|
||||
|
||||
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])
|
||||
@@ -137,8 +138,8 @@ class BuildGif:
|
||||
)
|
||||
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)}"
|
||||
print(save_path)
|
||||
|
||||
@@ -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.")
|
||||
@@ -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()
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user