Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a0b3695800 | ||
|
|
c299354c9c | ||
|
|
e002d0ad64 | ||
|
|
b5a728a997 | ||
|
|
72f46ad938 | ||
|
|
f74728267d | ||
|
|
0de2ab3f9a | ||
|
|
2fa1b95045 | ||
|
|
dba41a2c4d | ||
|
|
3cbfc8b78c | ||
|
|
fe74c556f5 | ||
|
|
d042fc572f | ||
|
|
46ae926911 | ||
|
|
2823a80078 | ||
|
|
fccffdbdf0 | ||
|
|
7e138268f9 | ||
|
|
ef3f66ef91 | ||
|
|
2f115352ec | ||
|
|
a0ebcdaf4d | ||
|
|
55cc09b0be | ||
|
|
2e731a480d | ||
|
|
65bc2c0327 |
@@ -29,33 +29,10 @@ See example workflow in examples folder.
|
|||||||
- Example: `"Hello World"`, `'It\'s a sunny day'`
|
- Example: `"Hello World"`, `'It\'s a sunny day'`
|
||||||
- Represents string literals.
|
- Represents string literals.
|
||||||
|
|
||||||
## Functions
|
|
||||||
|
|
||||||
Here are the available functions and their usage:
|
|
||||||
|
|
||||||
1. **Sum Function**:
|
|
||||||
- Syntax: `sum(arg1 | arg2 | ... | argN)`
|
|
||||||
- Adds together multiple embeddings.
|
|
||||||
- Example: `sum(embedding:face1 | dog)`
|
|
||||||
|
|
||||||
2. **Negation Function**:
|
|
||||||
- Syntax: `neg(arg)`
|
|
||||||
- Negates the output.
|
|
||||||
- Example: `neg(A embedding:happycats outside)`
|
|
||||||
|
|
||||||
3. **Normalization Function**:
|
|
||||||
- Syntax: `norm(arg)`
|
|
||||||
- Normalizes the given vector embedding.
|
|
||||||
- Example: `norm(sum(embedding:face1 | embedding:face2))`
|
|
||||||
|
|
||||||
4. **Difference Function**:
|
|
||||||
- Syntax: `diff(arg1 | arg2 | ... | argN)`
|
|
||||||
- Computes the difference between multiple vector embeddings.
|
|
||||||
- Example: `diff(embedding:face1 | embedding:face2)`
|
|
||||||
|
|
||||||
### Notes on Arguments:
|
### Notes on Arguments:
|
||||||
- Each function takes one or more arguments.
|
- Each function takes one or more arguments.
|
||||||
- An argument (`arg`) can be an embedding, a word, another function, or a quoted string.
|
- An argument (`arg`) can be an embedding, multiple words, another function, or a quoted string.
|
||||||
- For functions that accept multiple arguments, they are separated by the `|` symbol.
|
- For functions that accept multiple arguments, they are separated by the `|` symbol.
|
||||||
|
|
||||||
## Examples
|
## Examples
|
||||||
@@ -78,4 +55,25 @@ Here are the available functions and their usage:
|
|||||||
```
|
```
|
||||||
sum(king|neg(man)|woman)
|
sum(king|neg(man)|woman)
|
||||||
```
|
```
|
||||||
```
|
|
||||||
|
## Functions
|
||||||
|
|
||||||
|
Here are the available functions and their usage:
|
||||||
|
|
||||||
|
| Display Name | Action Name | Description | Usage Examples |
|
||||||
|
| --- | --- | --- | --- |
|
||||||
|
| Multiply | mult | Multiplies the provided segments or actions by the multiplier. | <ul><li>mult(The cat is\|2.5)</li><li>mult(Cat\|-1)</li></ul> |
|
||||||
|
| Set Dimensions | setDims | Sets the specified dimensions of the input embeddings to the specified value | <ul><li>The setDims(cat\|4, -0.01253\|76, 1.2) is happy</li></ul> |
|
||||||
|
| Negate | neg | Negates the provided segments or actions. | <ul><li>neg(cat)</li><li>sum(king\|neg(man)\|women)</li></ul> |
|
||||||
|
| Normalize | norm | Normalizes the provided segments or actions. | <ul><li>norm(cat)</li><li>sum(cat\|norm(sum(tiger\|fish)))</li></ul> |
|
||||||
|
| Positional Embedding Scale | posScale | Scales(Multiplies) the positional embeddings of the provided segments or actions by the multiplier. | <ul><li>A posScale(cat\|1.5) on a rainy day</li></ul> |
|
||||||
|
| Random Embedding | rand | Returns a random embedding of the specified token length, with the values optionally bounded by the second and third arguments. | <ul><li>A rand(1) cat</li><li>A rand(1\|-1\|1) cat</li></ul> |
|
||||||
|
| Difference | diff | Subtracts the segments in the order they are given. The first segment is subtracted from the second, then the third from the result, and so on. | <ul><li>diff(The cat is\|The dog is)</li><li>diff(Cat\|Dog)</li><li>sum(diff(king\|man)\|woman)</li></ul> |
|
||||||
|
| Slerp | slerp | Performs a slerp(Interpolation) between two segments or actions, with the given weight. The recommended weight is 0 - 1 | <ul><li>The slerp(cat\|dog\|0.5) is happy</li></ul> |
|
||||||
|
| Pooled Average(Experimental) | _exp-pooledAvg | Processes the provided segments or actions fully through CLIP and creates a pooled average of the last hidden state by averaging the last hidden state of each token. | <ul><li>A cat on a _exp-pooledAvg(beautiful sunny day)</li><li>A _exp-pooledAvg(broken glass) bottle</li></ul> |
|
||||||
|
| Sum | sum | Adds the embeddings of the provided segments or actions. | <ul><li>A happy sum(cat\|dog\|shark)</li></ul> |
|
||||||
|
| Pooler Output(Experimental) | _exp-pooler | Processes the provided segments or actions fully through CLIP and returns the pooler_output from the transformer | <ul><li>A cat on a _exp-pooler(beautiful sunny day)</li><li>A _exp-pooler(broken glass) bottle</li></ul> |
|
||||||
|
| Scale Dimensions | scaleDims | Scales the specified dimensions of the input embeddings by the specified amount | <ul><li>The scaleDims(cat\|4,1.5\|76,1.2) is happy</li></ul> |
|
||||||
|
| Ignore Positional Embeddings | postPos | Prevents positional embeddings from being applied to the provided segments or actions. | <ul><li>A postPos(cat) on a rainy day</li></ul> |
|
||||||
|
| Average | avg | Performs a weighted average between two segments or actions. The recommended weight is 0 - 1. | <ul><li>avg(The cat is\|The dog is\|0.5)</li><li>avg(Cat\|Dog\|0.5)</li></ul> |
|
||||||
|
|
||||||
|
|||||||
@@ -1,13 +1,9 @@
|
|||||||
from .nodes import (
|
from .nodes import (
|
||||||
BuildGif,
|
BuildGif,
|
||||||
SpecialClipLoader,
|
SpecialClipLoader,
|
||||||
MonacoPrompt,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"Build Gif": BuildGif,
|
"Build Gif": BuildGif,
|
||||||
"Special CLIP Loader": SpecialClipLoader,
|
"Special CLIP Loader": SpecialClipLoader,
|
||||||
"Monaco Prompt": MonacoPrompt,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
WEB_DIRECTORY = ("./web/dist", ["app.bundle.js"])
|
|
||||||
|
|||||||
@@ -3,6 +3,11 @@ from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction
|
|||||||
from custom_nodes.KepPromptLang.lib.actions.mult import MultiplyAction
|
from custom_nodes.KepPromptLang.lib.actions.mult import MultiplyAction
|
||||||
from custom_nodes.KepPromptLang.lib.actions.neg import NegAction
|
from custom_nodes.KepPromptLang.lib.actions.neg import NegAction
|
||||||
from custom_nodes.KepPromptLang.lib.actions.norm import NormAction
|
from custom_nodes.KepPromptLang.lib.actions.norm import NormAction
|
||||||
|
|
||||||
|
from custom_nodes.KepPromptLang.lib.actions.pooled_avg import PooledAvgAction
|
||||||
|
from custom_nodes.KepPromptLang.lib.actions.pooler import PoolerAction
|
||||||
|
from custom_nodes.KepPromptLang.lib.actions.pos_scale import PosScaleAction
|
||||||
|
from custom_nodes.KepPromptLang.lib.actions.post_pos import PostPosAction
|
||||||
from custom_nodes.KepPromptLang.lib.actions.rand import RandAction
|
from custom_nodes.KepPromptLang.lib.actions.rand import RandAction
|
||||||
from custom_nodes.KepPromptLang.lib.actions.scale_dims import ScaleDims
|
from custom_nodes.KepPromptLang.lib.actions.scale_dims import ScaleDims
|
||||||
from custom_nodes.KepPromptLang.lib.actions.set_dims import SetDims
|
from custom_nodes.KepPromptLang.lib.actions.set_dims import SetDims
|
||||||
@@ -20,3 +25,7 @@ register_action(SlerpAction)
|
|||||||
register_action(AverageAction)
|
register_action(AverageAction)
|
||||||
register_action(ScaleDims)
|
register_action(ScaleDims)
|
||||||
register_action(SetDims)
|
register_action(SetDims)
|
||||||
|
register_action(PosScaleAction)
|
||||||
|
register_action(PoolerAction)
|
||||||
|
register_action(PostPosAction)
|
||||||
|
register_action(PooledAvgAction)
|
||||||
|
|||||||
+45
-4
@@ -1,9 +1,10 @@
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Union, List
|
from typing import Union, List, TypedDict, Tuple
|
||||||
|
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
from torch.nn import Embedding
|
from torch.nn import Embedding
|
||||||
|
from transformers.models.clip.modeling_clip import CLIPTextTransformer
|
||||||
|
|
||||||
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||||
|
|
||||||
@@ -13,6 +14,14 @@ class ActionArity(Enum):
|
|||||||
SINGLE = 1
|
SINGLE = 1
|
||||||
MULTI = 2
|
MULTI = 2
|
||||||
|
|
||||||
|
class PostModifiers(TypedDict):
|
||||||
|
"""
|
||||||
|
A dictionary of post modifiers for an action result.
|
||||||
|
"""
|
||||||
|
position_embed_scale: Union[float, None]
|
||||||
|
bypass_pos_embed: Union[bool, None]
|
||||||
|
|
||||||
|
|
||||||
class Action(ABC):
|
class Action(ABC):
|
||||||
@property
|
@property
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
@@ -30,9 +39,23 @@ class Action(ABC):
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def name(self) -> str:
|
def display_name(self) -> str:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
@property
|
||||||
|
@abstractmethod
|
||||||
|
def action_name(self) -> str:
|
||||||
|
pass
|
||||||
|
|
||||||
|
@property
|
||||||
|
@abstractmethod
|
||||||
|
def description(self) -> str:
|
||||||
|
pass
|
||||||
|
|
||||||
|
@property
|
||||||
|
def usage_examples(self) -> List[str]:
|
||||||
|
return []
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def grammar(self) -> str:
|
def grammar(self) -> str:
|
||||||
@@ -67,7 +90,7 @@ class Action(ABC):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
def get_result(self, embedding_module: Embedding) -> Union[Tensor, Tuple[Tensor, PostModifiers]]:
|
||||||
"""
|
"""
|
||||||
Get the result of this action. This is called when the embeddings are being calculated.
|
Get the result of this action. This is called when the embeddings are being calculated.
|
||||||
:param embedding_module: The embedding module to use to get the base embeddings for tokens.
|
:param embedding_module: The embedding module to use to get the base embeddings for tokens.
|
||||||
@@ -75,6 +98,13 @@ class Action(ABC):
|
|||||||
"""
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
def process_with_transformer(self, transformer: CLIPTextTransformer, embedding_module: Embedding) -> None:
|
||||||
|
"""
|
||||||
|
For actions that need access to the TextTransformer, this method is called. Results are expected to be returned via get_result still.
|
||||||
|
:param transformer: An instance of CLIPTextTransformer
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
def depth_repr(self, depth: int = 1) -> str:
|
def depth_repr(self, depth: int = 1) -> str:
|
||||||
raise NotImplementedError()
|
raise NotImplementedError()
|
||||||
|
|
||||||
@@ -90,12 +120,17 @@ class SingleArgAction(Action, ABC):
|
|||||||
segments.append(seg_or_action)
|
segments.append(seg_or_action)
|
||||||
return segments
|
return segments
|
||||||
|
|
||||||
|
def process_with_transformer(self, transformer: CLIPTextTransformer, embedding_module: Embedding) -> None:
|
||||||
|
for seg_or_action in self.arg:
|
||||||
|
if isinstance(seg_or_action, Action):
|
||||||
|
seg_or_action.process_with_transformer(transformer, embedding_module)
|
||||||
|
|
||||||
def __init__(self, arg: List[Union[PromptSegment, Action]]):
|
def __init__(self, arg: List[Union[PromptSegment, Action]]):
|
||||||
# TODO: Target is a list now... what does this mean for us..
|
# TODO: Target is a list now... what does this mean for us..
|
||||||
self.arg = arg
|
self.arg = arg
|
||||||
|
|
||||||
def __repr__(self) -> str:
|
def __repr__(self) -> str:
|
||||||
return f"{self.name}({self.arg})"
|
return f"{self.display_name}({self.arg})"
|
||||||
|
|
||||||
class MultiArgAction(Action, ABC):
|
class MultiArgAction(Action, ABC):
|
||||||
arity = ActionArity.MULTI
|
arity = ActionArity.MULTI
|
||||||
@@ -110,6 +145,12 @@ class MultiArgAction(Action, ABC):
|
|||||||
|
|
||||||
return segments
|
return segments
|
||||||
|
|
||||||
|
def process_with_transformer(self, transformer: CLIPTextTransformer, embedding_module: Embedding) -> None:
|
||||||
|
for arg in self.all_args:
|
||||||
|
for seg_or_action in arg:
|
||||||
|
if isinstance(seg_or_action, Action):
|
||||||
|
seg_or_action.process_with_transformer(transformer, embedding_module)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
args: List[List[Union[PromptSegment, Action]]],
|
args: List[List[Union[PromptSegment, Action]]],
|
||||||
|
|||||||
@@ -1,3 +1,5 @@
|
|||||||
|
from typing import List
|
||||||
|
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
from torch.nn import Embedding
|
from torch.nn import Embedding
|
||||||
|
|
||||||
@@ -9,3 +11,10 @@ def get_embedding(seg_or_action: SegOrAction, embedding_module: Embedding) -> Te
|
|||||||
if isinstance(seg_or_action, Action):
|
if isinstance(seg_or_action, Action):
|
||||||
return seg_or_action.get_result(embedding_module)
|
return seg_or_action.get_result(embedding_module)
|
||||||
return seg_or_action.get_embeddings(embedding_module)
|
return seg_or_action.get_embeddings(embedding_module)
|
||||||
|
|
||||||
|
def get_total_length(args: List[SegOrAction]) -> int:
|
||||||
|
total_length = 0
|
||||||
|
for seg_or_action in args:
|
||||||
|
total_length += seg_or_action.token_length()
|
||||||
|
|
||||||
|
return total_length
|
||||||
|
|||||||
+9
-1
@@ -12,9 +12,17 @@ from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
|||||||
|
|
||||||
class AverageAction(MultiArgAction):
|
class AverageAction(MultiArgAction):
|
||||||
grammar = 'avg(" arg "|" arg "|" arg ")"'
|
grammar = 'avg(" arg "|" arg "|" arg ")"'
|
||||||
name = "avg"
|
|
||||||
chars = ["+", "+"]
|
chars = ["+", "+"]
|
||||||
|
|
||||||
|
display_name = "Average"
|
||||||
|
action_name = "avg"
|
||||||
|
description = "Performs a weighted average between two segments or actions. The recommended weight is 0 - 1."
|
||||||
|
usage_examples = [
|
||||||
|
"avg(The cat is|The dog is|0.5)",
|
||||||
|
"avg(Cat|Dog|0.5)",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
|
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
|
||||||
super().__init__(args)
|
super().__init__(args)
|
||||||
|
|
||||||
|
|||||||
+9
-1
@@ -11,9 +11,17 @@ from custom_nodes.KepPromptLang.lib.parser.registration import register_action
|
|||||||
|
|
||||||
class DiffAction(MultiArgAction):
|
class DiffAction(MultiArgAction):
|
||||||
grammar = 'diff(" arg ("|" arg)* ")"'
|
grammar = 'diff(" arg ("|" arg)* ")"'
|
||||||
name = "diff"
|
|
||||||
chars = ["-", "-"]
|
chars = ["-", "-"]
|
||||||
|
|
||||||
|
display_name = "Difference"
|
||||||
|
action_name = "diff"
|
||||||
|
description = "Subtracts the segments in the order they are given. The first segment is subtracted from the second, then the third from the result, and so on."
|
||||||
|
usage_examples = [
|
||||||
|
"diff(The cat is|The dog is)",
|
||||||
|
"diff(Cat|Dog)",
|
||||||
|
"sum(diff(king|man)|woman)",
|
||||||
|
]
|
||||||
|
|
||||||
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
|
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
|
||||||
super().__init__(args)
|
super().__init__(args)
|
||||||
self.base_arg = args[0]
|
self.base_arg = args[0]
|
||||||
|
|||||||
+8
-1
@@ -13,9 +13,16 @@ from custom_nodes.KepPromptLang.lib.parser.registration import register_action
|
|||||||
|
|
||||||
class MultiplyAction(MultiArgAction):
|
class MultiplyAction(MultiArgAction):
|
||||||
grammar = 'mult(" arg+ ")"'
|
grammar = 'mult(" arg+ ")"'
|
||||||
name = "mult"
|
|
||||||
chars = ["[", "]"]
|
chars = ["[", "]"]
|
||||||
|
|
||||||
|
display_name = "Multiply"
|
||||||
|
action_name = "mult"
|
||||||
|
description = "Multiplies the provided segments or actions by the multiplier."
|
||||||
|
usage_examples = [
|
||||||
|
"mult(The cat is|2.5)",
|
||||||
|
"mult(Cat|-1)",
|
||||||
|
]
|
||||||
|
|
||||||
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
||||||
super().__init__(args)
|
super().__init__(args)
|
||||||
if len(args) != 2:
|
if len(args) != 2:
|
||||||
|
|||||||
+8
-1
@@ -7,9 +7,16 @@ from custom_nodes.KepPromptLang.lib.parser.registration import register_action
|
|||||||
|
|
||||||
class NegAction(SingleArgAction):
|
class NegAction(SingleArgAction):
|
||||||
grammar = 'neg(" arg+ ")"'
|
grammar = 'neg(" arg+ ")"'
|
||||||
name = "neg"
|
|
||||||
chars = ["[", "]"]
|
chars = ["[", "]"]
|
||||||
|
|
||||||
|
display_name = "Negate"
|
||||||
|
action_name = "neg"
|
||||||
|
description = "Negates the provided segments or actions."
|
||||||
|
usage_examples = [
|
||||||
|
"neg(cat)",
|
||||||
|
"sum(king|neg(man)|women)",
|
||||||
|
]
|
||||||
|
|
||||||
def token_length(self) -> int:
|
def token_length(self) -> int:
|
||||||
"""
|
"""
|
||||||
Neg negates the embeddings of the base segment, so the length is the length of the base segment
|
Neg negates the embeddings of the base segment, so the length is the length of the base segment
|
||||||
|
|||||||
+8
-1
@@ -10,9 +10,16 @@ from custom_nodes.KepPromptLang.lib.parser.registration import register_action
|
|||||||
|
|
||||||
class NormAction(SingleArgAction):
|
class NormAction(SingleArgAction):
|
||||||
grammar = 'norm(" arg+ ")"'
|
grammar = 'norm(" arg+ ")"'
|
||||||
name = "norm"
|
|
||||||
chars = None
|
chars = None
|
||||||
|
|
||||||
|
display_name = "Normalize"
|
||||||
|
action_name = "norm"
|
||||||
|
description = "Normalizes the provided segments or actions."
|
||||||
|
usage_examples = [
|
||||||
|
"norm(cat)",
|
||||||
|
"sum(cat|norm(sum(tiger|fish)))",
|
||||||
|
]
|
||||||
|
|
||||||
def token_length(self) -> int:
|
def token_length(self) -> int:
|
||||||
"""
|
"""
|
||||||
Norm normalizes the embeddings of the base segment, so the length is the length of the base segment
|
Norm normalizes the embeddings of the base segment, so the length is the length of the base segment
|
||||||
|
|||||||
@@ -0,0 +1,68 @@
|
|||||||
|
import torch
|
||||||
|
from torch.nn import Embedding
|
||||||
|
from transformers.modeling_outputs import BaseModelOutputWithPooling
|
||||||
|
from transformers.models.clip.modeling_clip import CLIPTextTransformer
|
||||||
|
|
||||||
|
from custom_nodes.KepPromptLang.lib.action.base import Action, SingleArgAction
|
||||||
|
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_total_length
|
||||||
|
from custom_nodes.KepPromptLang.lib.fun_clip_stuff import (
|
||||||
|
PromptLangCLIPTextEmbeddings,
|
||||||
|
PrompLangCLIPTextTransformer,
|
||||||
|
)
|
||||||
|
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||||
|
|
||||||
|
|
||||||
|
class PooledAvgAction(SingleArgAction):
|
||||||
|
grammar = 'pooledAvg(" arg+ ")"'
|
||||||
|
chars = ["[", "]"]
|
||||||
|
|
||||||
|
display_name = "Pooled Average(Experimental)"
|
||||||
|
action_name = "_exp-pooledAvg"
|
||||||
|
description = "Processes the provided segments or actions fully through CLIP and creates a pooled average of the last hidden state by averaging the last hidden state of each token."
|
||||||
|
usage_examples = [
|
||||||
|
"A cat on a _exp-pooledAvg(beautiful sunny day)",
|
||||||
|
"A _exp-pooledAvg(broken glass) bottle",
|
||||||
|
]
|
||||||
|
|
||||||
|
def __init__(self, args):
|
||||||
|
super().__init__(args)
|
||||||
|
self.result = None
|
||||||
|
|
||||||
|
def token_length(self) -> int:
|
||||||
|
"""
|
||||||
|
PooledAvg returns the average of the last hidden state, so the length is 1
|
||||||
|
"""
|
||||||
|
return 1
|
||||||
|
|
||||||
|
def process_with_transformer(
|
||||||
|
self, transformer: CLIPTextTransformer, embedding_module: Embedding
|
||||||
|
) -> None:
|
||||||
|
""" """
|
||||||
|
# SOT + tokens + EOT
|
||||||
|
eot_token = embedding_module.num_embeddings - 1
|
||||||
|
print("Using EOT token", eot_token)
|
||||||
|
#TODO: Play with impact of padding on pooled output
|
||||||
|
arg_length = get_total_length(self.arg)
|
||||||
|
|
||||||
|
# SOT + arg length + EOT
|
||||||
|
empty_tokens = [[49406] + [eot_token] * (arg_length + 1)]
|
||||||
|
|
||||||
|
transformer_results: BaseModelOutputWithPooling = transformer(
|
||||||
|
[
|
||||||
|
[PromptSegment(text="_Empty Batch_", tokens=empty_tokens[0])],
|
||||||
|
[PromptSegment(text="[SOT]", tokens=[49406])]
|
||||||
|
+ self.arg
|
||||||
|
+ [PromptSegment(text="[EOT]", tokens=[eot_token])],
|
||||||
|
]
|
||||||
|
)
|
||||||
|
self.result = (
|
||||||
|
transformer_results.last_hidden_state[1, 1:-1, :].mean(dim=0).unsqueeze(0).unsqueeze(0)
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
|
||||||
|
if self.result is not None:
|
||||||
|
return self.result
|
||||||
|
|
||||||
|
raise Exception(
|
||||||
|
"PooledAvg action result is not set. Did you forget to call process_with_transformer?"
|
||||||
|
)
|
||||||
@@ -0,0 +1,67 @@
|
|||||||
|
import torch
|
||||||
|
from torch.nn import Embedding
|
||||||
|
from transformers.modeling_outputs import BaseModelOutputWithPooling
|
||||||
|
from transformers.models.clip.modeling_clip import CLIPTextTransformer
|
||||||
|
|
||||||
|
from custom_nodes.KepPromptLang.lib.action.base import Action, SingleArgAction
|
||||||
|
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_total_length
|
||||||
|
from custom_nodes.KepPromptLang.lib.fun_clip_stuff import (
|
||||||
|
PromptLangCLIPTextEmbeddings,
|
||||||
|
PrompLangCLIPTextTransformer,
|
||||||
|
)
|
||||||
|
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||||
|
|
||||||
|
|
||||||
|
class PoolerAction(SingleArgAction):
|
||||||
|
grammar = 'pooler(" arg+ ")"'
|
||||||
|
chars = ["[", "]"]
|
||||||
|
|
||||||
|
display_name = "Pooler Output(Experimental)"
|
||||||
|
action_name = "_exp-pooler"
|
||||||
|
description = "Processes the provided segments or actions fully through CLIP and returns the pooler_output from the transformer"
|
||||||
|
usage_examples = [
|
||||||
|
"A cat on a _exp-pooler(beautiful sunny day)",
|
||||||
|
"A _exp-pooler(broken glass) bottle",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def __init__(self, args):
|
||||||
|
super().__init__(args)
|
||||||
|
self.result = None
|
||||||
|
|
||||||
|
def token_length(self) -> int:
|
||||||
|
"""
|
||||||
|
Pooler returns the embedding of the EOT token, so the length is 1
|
||||||
|
"""
|
||||||
|
return 1
|
||||||
|
|
||||||
|
def process_with_transformer(
|
||||||
|
self, transformer: CLIPTextTransformer, embedding_module: Embedding
|
||||||
|
) -> None:
|
||||||
|
""" """
|
||||||
|
# SOT + tokens + EOT
|
||||||
|
eot_token = embedding_module.num_embeddings - 1
|
||||||
|
print("Using EOT token", eot_token)
|
||||||
|
|
||||||
|
arg_length = get_total_length(self.arg)
|
||||||
|
|
||||||
|
# SOT + arg length + EOT
|
||||||
|
empty_tokens = [[49406] + [eot_token] * (arg_length + 1)]
|
||||||
|
|
||||||
|
transformer_results: BaseModelOutputWithPooling = transformer(
|
||||||
|
[
|
||||||
|
[PromptSegment(text="_Empty Batch_", tokens=empty_tokens[0])],
|
||||||
|
[PromptSegment(text="[SOT]", tokens=[49406])]
|
||||||
|
+ self.arg
|
||||||
|
+ [PromptSegment(text="[EOT]", tokens=[eot_token])],
|
||||||
|
]
|
||||||
|
)
|
||||||
|
self.result = transformer_results.pooler_output[1].unsqueeze(0).unsqueeze(0)
|
||||||
|
|
||||||
|
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
|
||||||
|
if self.result is not None:
|
||||||
|
return self.result
|
||||||
|
|
||||||
|
raise Exception(
|
||||||
|
"Pooled action result is not set. Did you forget to call process_with_transformer?"
|
||||||
|
)
|
||||||
@@ -0,0 +1,71 @@
|
|||||||
|
from typing import Tuple, List
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch.nn import Embedding
|
||||||
|
|
||||||
|
from custom_nodes.KepPromptLang.lib.action.base import (
|
||||||
|
Action,
|
||||||
|
PostModifiers,
|
||||||
|
MultiArgAction,
|
||||||
|
)
|
||||||
|
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
||||||
|
|
||||||
|
|
||||||
|
class PosScaleAction(MultiArgAction):
|
||||||
|
grammar = 'posScale(" arg+ ")"'
|
||||||
|
chars = ["[", "]"]
|
||||||
|
|
||||||
|
display_name = "Positional Embedding Scale"
|
||||||
|
action_name = "posScale"
|
||||||
|
description = "Scales(Multiplies) the positional embeddings of the provided segments or actions by the multiplier."
|
||||||
|
usage_examples = [
|
||||||
|
"A posScale(cat|1.5) on a rainy day",
|
||||||
|
]
|
||||||
|
|
||||||
|
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
||||||
|
super().__init__(args)
|
||||||
|
if len(args) != 2:
|
||||||
|
raise ValueError("PosScale 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(
|
||||||
|
"PosScale actions multiplier should have exactly one segment"
|
||||||
|
)
|
||||||
|
|
||||||
|
multiplier_seg_or_action = arg[0]
|
||||||
|
|
||||||
|
if isinstance(multiplier_seg_or_action, Action):
|
||||||
|
raise ValueError("PosScale actions multiplier must be a number")
|
||||||
|
|
||||||
|
try:
|
||||||
|
self.parsed_multiplier = float(multiplier_seg_or_action.text)
|
||||||
|
except ValueError:
|
||||||
|
raise ValueError(
|
||||||
|
"PosScale action should have an integer/float as the multiplier"
|
||||||
|
)
|
||||||
|
|
||||||
|
def token_length(self) -> int:
|
||||||
|
"""
|
||||||
|
PosScale modifies the posional 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) -> Tuple[torch.Tensor, PostModifiers]:
|
||||||
|
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, {"position_embed_scale": self.parsed_multiplier}
|
||||||
@@ -0,0 +1,47 @@
|
|||||||
|
from typing import Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import Tensor
|
||||||
|
from torch.nn import Embedding
|
||||||
|
|
||||||
|
from custom_nodes.KepPromptLang.lib.action.base import (
|
||||||
|
SingleArgAction,
|
||||||
|
Action,
|
||||||
|
PostModifiers,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class PostPosAction(SingleArgAction):
|
||||||
|
grammar = 'postPos(" arg+ ")"'
|
||||||
|
chars = ["[", "]"]
|
||||||
|
|
||||||
|
display_name = "Ignore Positional Embeddings"
|
||||||
|
action_name = "postPos"
|
||||||
|
description = "Prevents positional embeddings from being applied to the provided segments or actions."
|
||||||
|
usage_examples = [
|
||||||
|
"A postPos(cat) on a rainy day",
|
||||||
|
]
|
||||||
|
|
||||||
|
def __init__(self, args):
|
||||||
|
super().__init__(args)
|
||||||
|
self.result = None
|
||||||
|
|
||||||
|
def token_length(self) -> int:
|
||||||
|
"""
|
||||||
|
PostPos returns the results of the wrapped action, so the length is the length of the wrapped action
|
||||||
|
"""
|
||||||
|
total_length = 0
|
||||||
|
for seg_or_action in self.arg:
|
||||||
|
total_length += seg_or_action.token_length()
|
||||||
|
|
||||||
|
return total_length
|
||||||
|
|
||||||
|
def get_result(self, embedding_module: Embedding) -> Tuple[Tensor, PostModifiers]:
|
||||||
|
all_embeddings = []
|
||||||
|
for seg_or_action in self.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))
|
||||||
|
|
||||||
|
return torch.cat(all_embeddings, dim=1), PostModifiers(position_embed_scale=None, bypass_pos_embed=True)
|
||||||
@@ -17,6 +17,14 @@ class RandAction(MultiArgAction):
|
|||||||
name = "rand"
|
name = "rand"
|
||||||
chars = None
|
chars = None
|
||||||
|
|
||||||
|
display_name = "Random Embedding"
|
||||||
|
action_name = "rand"
|
||||||
|
description = "Returns a random embedding of the specified token length, with the values optionally bounded by the second and third arguments."
|
||||||
|
usage_examples = [
|
||||||
|
"A rand(1) cat",
|
||||||
|
"A rand(1|-1|1) cat",
|
||||||
|
]
|
||||||
|
|
||||||
parsed_token_length = 0
|
parsed_token_length = 0
|
||||||
range_min = 0
|
range_min = 0
|
||||||
range_max = 1
|
range_max = 1
|
||||||
|
|||||||
@@ -10,11 +10,15 @@ from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
|||||||
|
|
||||||
class ScaleDims(MultiArgAction):
|
class ScaleDims(MultiArgAction):
|
||||||
grammar = 'scaleDims(" arg ("|" arg)* ")"'
|
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 = ["-", "-"]
|
chars = ["-", "-"]
|
||||||
|
|
||||||
|
display_name = "Scale Dimensions"
|
||||||
|
action_name = "scaleDims"
|
||||||
|
description = "Scales the specified dimensions of the input embeddings by the specified amount"
|
||||||
|
usage_examples = [
|
||||||
|
"The scaleDims(cat|4,1.5|76,1.2) is happy",
|
||||||
|
]
|
||||||
|
|
||||||
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
|
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
|
||||||
super().__init__(args)
|
super().__init__(args)
|
||||||
|
|
||||||
|
|||||||
@@ -10,11 +10,15 @@ from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
|||||||
|
|
||||||
class SetDims(MultiArgAction):
|
class SetDims(MultiArgAction):
|
||||||
grammar = 'setDims(" arg ("|" arg)* ")"'
|
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 = ["-", "-"]
|
chars = ["-", "-"]
|
||||||
|
|
||||||
|
display_name = "Set Dimensions"
|
||||||
|
action_name = "setDims"
|
||||||
|
description = "Sets the specified dimensions of the input embeddings to the specified value"
|
||||||
|
usage_examples = [
|
||||||
|
"The setDims(cat|4, -0.01253|76, 1.2) is happy"
|
||||||
|
]
|
||||||
|
|
||||||
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
|
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
|
||||||
super().__init__(args)
|
super().__init__(args)
|
||||||
|
|
||||||
|
|||||||
@@ -12,9 +12,15 @@ from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
|||||||
|
|
||||||
class SlerpAction(MultiArgAction):
|
class SlerpAction(MultiArgAction):
|
||||||
grammar = 'slerp(" arg "|" arg "|" arg ")"'
|
grammar = 'slerp(" arg "|" arg "|" arg ")"'
|
||||||
name = "slerp"
|
|
||||||
chars = ["+", "+"]
|
chars = ["+", "+"]
|
||||||
|
|
||||||
|
display_name = "Slerp"
|
||||||
|
action_name = "slerp"
|
||||||
|
description = "Performs a slerp(Interpolation) between two segments or actions, with the given weight. The recommended weight is 0 - 1"
|
||||||
|
usage_examples = [
|
||||||
|
"The slerp(cat|dog|0.5) is happy",
|
||||||
|
]
|
||||||
|
|
||||||
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
|
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
|
||||||
super().__init__(args)
|
super().__init__(args)
|
||||||
|
|
||||||
|
|||||||
+7
-1
@@ -11,9 +11,15 @@ from custom_nodes.KepPromptLang.lib.parser.registration import register_action
|
|||||||
|
|
||||||
class SumAction(MultiArgAction):
|
class SumAction(MultiArgAction):
|
||||||
grammar = 'sum(" arg ("|" arg)+ ")"'
|
grammar = 'sum(" arg ("|" arg)+ ")"'
|
||||||
name = "sum"
|
|
||||||
chars = ["+", "+"]
|
chars = ["+", "+"]
|
||||||
|
|
||||||
|
display_name = "Sum"
|
||||||
|
action_name = "sum"
|
||||||
|
description = "Adds the embeddings of the provided segments or actions."
|
||||||
|
usage_examples = [
|
||||||
|
"A happy sum(cat|dog|shark)",
|
||||||
|
]
|
||||||
|
|
||||||
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
|
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
|
||||||
super().__init__(args)
|
super().__init__(args)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,23 @@
|
|||||||
|
{
|
||||||
|
"architectures": [
|
||||||
|
"CLIPTextModel"
|
||||||
|
],
|
||||||
|
"attention_dropout": 0.0,
|
||||||
|
"bos_token_id": 0,
|
||||||
|
"dropout": 0.0,
|
||||||
|
"eos_token_id": 2,
|
||||||
|
"hidden_act": "gelu",
|
||||||
|
"hidden_size": 1280,
|
||||||
|
"initializer_factor": 1.0,
|
||||||
|
"initializer_range": 0.02,
|
||||||
|
"intermediate_size": 5120,
|
||||||
|
"layer_norm_eps": 1e-05,
|
||||||
|
"max_position_embeddings": 77,
|
||||||
|
"model_type": "clip_text_model",
|
||||||
|
"num_attention_heads": 20,
|
||||||
|
"num_hidden_layers": 32,
|
||||||
|
"pad_token_id": 1,
|
||||||
|
"projection_dim": 1280,
|
||||||
|
"torch_dtype": "float32",
|
||||||
|
"vocab_size": 49408
|
||||||
|
}
|
||||||
+36
-5
@@ -7,6 +7,8 @@ from transformers import CLIPTextConfig, modeling_utils
|
|||||||
|
|
||||||
from comfy import model_management
|
from comfy import model_management
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
|
from comfy.sd1_clip import SD1ClipModel
|
||||||
|
from comfy.sdxl_clip import SDXLClipModel
|
||||||
from custom_nodes.KepPromptLang.lib.action.base import Action
|
from custom_nodes.KepPromptLang.lib.action.base import Action
|
||||||
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
||||||
from custom_nodes.KepPromptLang.lib.fun_clip_stuff import PromptLangTextModel
|
from custom_nodes.KepPromptLang.lib.fun_clip_stuff import PromptLangTextModel
|
||||||
@@ -14,7 +16,7 @@ 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
|
# Methods with no comment can be assumed to be the same as comfy.sd1_clip.SD1ClipModel
|
||||||
class PromptLangClipModel(torch.nn.Module):
|
class PromptLangSDClipModel(torch.nn.Module):
|
||||||
"""Uses the CLIP transformer encoder for text (from huggingface)"""
|
"""Uses the CLIP transformer encoder for text (from huggingface)"""
|
||||||
LAYERS = [
|
LAYERS = [
|
||||||
"last",
|
"last",
|
||||||
@@ -106,12 +108,11 @@ class PromptLangClipModel(torch.nn.Module):
|
|||||||
tokens_temp += [next_new_token]
|
tokens_temp += [next_new_token]
|
||||||
next_new_token += 1
|
next_new_token += 1
|
||||||
else:
|
else:
|
||||||
print("WARNING: shape mismatch when trying to apply embedding, embedding will be ignored",
|
raise Exception("WARNING: shape mismatch when trying to apply embedding. Should have been caught during tokenization.",
|
||||||
tid_or_tensor.shape[0], current_embeds.weight.shape[1])
|
tid_or_tensor.shape[0], current_embeds.weight.shape[1])
|
||||||
if len(tokens_temp) < segment_length:
|
if len(tokens_temp) < segment_length:
|
||||||
# Pretty sure this is only needed if the embedding is not the same size as the CLIP embedding
|
# This should never happen...
|
||||||
print("WARNING: segment length mismatch, padding with EOS token")
|
raise Exception("Segment size mismatch. Please submit an issue on Github.")
|
||||||
tokens_temp.extend([self.empty_tokens[0][-1] * (segment_length - len(tokens_temp))])
|
|
||||||
segment.tokens = tokens_temp
|
segment.tokens = tokens_temp
|
||||||
|
|
||||||
n = token_dict_size
|
n = token_dict_size
|
||||||
@@ -209,3 +210,33 @@ class PromptLangClipModel(torch.nn.Module):
|
|||||||
if (len(output) == 0):
|
if (len(output) == 0):
|
||||||
return z_empty.cpu(), first_pooled.cpu()
|
return z_empty.cpu(), first_pooled.cpu()
|
||||||
return torch.cat(output, dim=-2).cpu(), first_pooled.cpu()
|
return torch.cat(output, dim=-2).cpu(), first_pooled.cpu()
|
||||||
|
|
||||||
|
class PromptLangSD1ClipModel(SD1ClipModel):
|
||||||
|
def __init__(self, device="cpu", dtype=None, clip_name="l", clip_model=PromptLangSDClipModel):
|
||||||
|
super().__init__()
|
||||||
|
self.clip_name = clip_name
|
||||||
|
self.clip = "clip_{}".format(self.clip_name)
|
||||||
|
setattr(self, self.clip, clip_model(device=device, dtype=dtype))
|
||||||
|
|
||||||
|
|
||||||
|
class PromptLangSDXLClipModel(SDXLClipModel):
|
||||||
|
def __init__(self, device="cpu", dtype=None) -> None:
|
||||||
|
# Skip SDXLClipModel's init
|
||||||
|
super(SDXLClipModel, self).__init__()
|
||||||
|
self.clip_l = PromptLangSDClipModel(layer="hidden", layer_idx=11, device=device, dtype=dtype)
|
||||||
|
self.clip_l.layer_norm_hidden_state = False
|
||||||
|
self.clip_g = PromptLangSDXLClipG(device, dtype)
|
||||||
|
|
||||||
|
class PromptLangSDXLClipG(PromptLangSDClipModel):
|
||||||
|
def __init__(self, device="cpu", max_length=77, freeze=True, layer="penultimate", layer_idx=None, textmodel_path=None, dtype=None):
|
||||||
|
if layer == "penultimate":
|
||||||
|
layer="hidden"
|
||||||
|
layer_idx=-2
|
||||||
|
|
||||||
|
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "clip_config_bigg.json")
|
||||||
|
super().__init__(device=device, freeze=freeze, layer=layer, layer_idx=layer_idx, textmodel_json_config=textmodel_json_config, textmodel_path=textmodel_path, dtype=dtype)
|
||||||
|
self.empty_tokens = [[49406] + [49407] + [0] * 75]
|
||||||
|
self.layer_norm_hidden_state = False
|
||||||
|
|
||||||
|
def load_sd(self, sd):
|
||||||
|
return super().load_sd(sd)
|
||||||
|
|||||||
+115
-28
@@ -1,10 +1,11 @@
|
|||||||
from typing import Optional, Tuple, Union, List
|
from typing import Optional, Tuple, Union, List, TypedDict, TYPE_CHECKING
|
||||||
|
from importlib.metadata import version as import_version
|
||||||
|
from packaging import version
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from transformers import CLIPTextConfig
|
from transformers import CLIPTextConfig
|
||||||
from transformers.modeling_outputs import BaseModelOutputWithPooling
|
from transformers.modeling_outputs import BaseModelOutputWithPooling
|
||||||
from transformers.models.clip.modeling_clip import (
|
from transformers.models.clip.modeling_clip import (
|
||||||
_expand_mask,
|
|
||||||
CLIPTextEmbeddings,
|
CLIPTextEmbeddings,
|
||||||
CLIPTextTransformer,
|
CLIPTextTransformer,
|
||||||
CLIPTextModel,
|
CLIPTextModel,
|
||||||
@@ -13,6 +14,11 @@ from transformers.models.clip.modeling_clip import (
|
|||||||
from custom_nodes.KepPromptLang.lib.action.base import Action
|
from custom_nodes.KepPromptLang.lib.action.base import Action
|
||||||
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from torch import Tensor
|
||||||
|
from custom_nodes.KepPromptLang.lib.action.base import PostModifiers
|
||||||
|
|
||||||
|
|
||||||
def slerp(val, low, high):
|
def slerp(val, low, high):
|
||||||
low = low.unsqueeze(0)
|
low = low.unsqueeze(0)
|
||||||
high = high.unsqueeze(0)
|
high = high.unsqueeze(0)
|
||||||
@@ -23,6 +29,17 @@ def slerp(val, low, high):
|
|||||||
res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high
|
res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high
|
||||||
return res
|
return res
|
||||||
|
|
||||||
|
|
||||||
|
class PosModifier(TypedDict):
|
||||||
|
"""
|
||||||
|
A dictionary of post modifiers for an action result.
|
||||||
|
"""
|
||||||
|
|
||||||
|
position_embed_scale: Union[float]
|
||||||
|
start_idx: Union[int]
|
||||||
|
end_idx: Union[int]
|
||||||
|
|
||||||
|
|
||||||
class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||||
def __init__(self, config: CLIPTextConfig):
|
def __init__(self, config: CLIPTextConfig):
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
@@ -34,19 +51,42 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
|||||||
position_ids: Optional[torch.LongTensor] = None,
|
position_ids: Optional[torch.LongTensor] = None,
|
||||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
|
||||||
if input_dicts is None:
|
if input_dicts is None:
|
||||||
raise ValueError("You have to specify input_dicts")
|
raise ValueError("You have to specify input_dicts")
|
||||||
|
|
||||||
batches = []
|
batches: List[List[Tensor | Tuple[Tensor, PostModifiers] | Action]] = []
|
||||||
|
pos_modifiers: List[List[PosModifier]] = []
|
||||||
for batch_idx, batch in enumerate(input_dicts):
|
for batch_idx, batch in enumerate(input_dicts):
|
||||||
results = []
|
results = []
|
||||||
|
batch_pos_modifiers = []
|
||||||
|
token_idx = 0
|
||||||
for seg_or_action in batch:
|
for seg_or_action in batch:
|
||||||
if isinstance(seg_or_action, Action):
|
if isinstance(seg_or_action, Action):
|
||||||
results.append(seg_or_action.get_result(self.token_embedding))
|
action_result: Union[
|
||||||
|
Tensor, Tuple[Tensor, PostModifiers]
|
||||||
|
] = seg_or_action.get_result(self.token_embedding)
|
||||||
|
if isinstance(action_result, tuple):
|
||||||
|
result, post_modifiers = action_result
|
||||||
|
if post_modifiers.get("position_embed_scale", None) is not None:
|
||||||
|
post_modifiers["start_idx"] = token_idx
|
||||||
|
post_modifiers["end_idx"] = (
|
||||||
|
token_idx + seg_or_action.token_length()
|
||||||
|
)
|
||||||
|
|
||||||
|
if post_modifiers.get("bypass_pos_embed", False):
|
||||||
|
post_modifiers["start_idx"] = token_idx
|
||||||
|
post_modifiers["end_idx"] = (
|
||||||
|
token_idx + seg_or_action.token_length()
|
||||||
|
)
|
||||||
|
batch_pos_modifiers.append(post_modifiers)
|
||||||
|
else:
|
||||||
|
result = action_result
|
||||||
else:
|
else:
|
||||||
results.append(seg_or_action.get_embeddings(self.token_embedding))
|
result = seg_or_action.get_embeddings(self.token_embedding)
|
||||||
|
results.append(result)
|
||||||
|
token_idx += seg_or_action.token_length()
|
||||||
batches.append(results)
|
batches.append(results)
|
||||||
|
pos_modifiers.append(batch_pos_modifiers)
|
||||||
|
|
||||||
seq_length = batches[0][0].shape[-2]
|
seq_length = batches[0][0].shape[-2]
|
||||||
|
|
||||||
@@ -60,8 +100,28 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
|||||||
else:
|
else:
|
||||||
embeds.append(torch.cat(batch, dim=-2))
|
embeds.append(torch.cat(batch, dim=-2))
|
||||||
|
|
||||||
position_embeddings = self.position_embedding(position_ids)
|
# Iterate over the batches and apply the pos modifiers to the position embeddings then add them to the embeddings
|
||||||
embeddings = torch.cat(embeds, dim=0) + position_embeddings
|
for idx, batch_pos_modifiers in enumerate(pos_modifiers):
|
||||||
|
position_embeddings = self.position_embedding(position_ids)
|
||||||
|
if len(batch_pos_modifiers) > 0:
|
||||||
|
print(f"Found {len(batch_pos_modifiers)} pos modifiers for batch {idx}")
|
||||||
|
# Apply each pos modifier to the position embeddings at the specified indices
|
||||||
|
for post_modifier in batch_pos_modifiers:
|
||||||
|
if post_modifier.get("bypass_pos_embed", False):
|
||||||
|
position_embeddings[
|
||||||
|
0, post_modifier["start_idx"] : post_modifier["end_idx"]
|
||||||
|
] = 0
|
||||||
|
elif post_modifier["position_embed_scale"] is not None:
|
||||||
|
position_embeddings[
|
||||||
|
0, post_modifier["start_idx"] : post_modifier["end_idx"]
|
||||||
|
] *= post_modifier["position_embed_scale"]
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
"Pos modifier must have a scale or bypass_pos_embed"
|
||||||
|
)
|
||||||
|
# Add the possibly modified position embeddings to the embeddings
|
||||||
|
embeds[idx] = embeds[idx] + position_embeddings
|
||||||
|
embeddings = torch.cat(embeds, dim=0)
|
||||||
|
|
||||||
return embeddings
|
return embeddings
|
||||||
|
|
||||||
@@ -70,6 +130,40 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
|
|||||||
def __init__(self, config: CLIPTextConfig):
|
def __init__(self, config: CLIPTextConfig):
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
self.embeddings = PromptLangCLIPTextEmbeddings(config)
|
self.embeddings = PromptLangCLIPTextEmbeddings(config)
|
||||||
|
self.transformers_version = version.parse(import_version('transformers'))
|
||||||
|
|
||||||
|
def process_attention_mask(self, hidden_states, attention_mask, bsz, seq_len):
|
||||||
|
# Parse the transformer version
|
||||||
|
input_shape = torch.Size([bsz, seq_len])
|
||||||
|
|
||||||
|
v4_30 = version.parse('4.30.0')
|
||||||
|
v4_35 = version.parse('4.35')
|
||||||
|
if self.transformers_version < v4_30:
|
||||||
|
print("Using transformers < 4.30.0")
|
||||||
|
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(
|
||||||
|
hidden_states.device)
|
||||||
|
elif v4_30 <= self.transformers_version < v4_35:
|
||||||
|
print("Using transformers >= 4.30.0 and <= 4.34.*")
|
||||||
|
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)
|
||||||
|
else:
|
||||||
|
print("Using transformers >= 4.35")
|
||||||
|
from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask
|
||||||
|
causal_attention_mask = _create_4d_causal_attention_mask(
|
||||||
|
input_shape, hidden_states.dtype, device=hidden_states.device
|
||||||
|
)
|
||||||
|
|
||||||
|
# Expand attention_mask if it exists
|
||||||
|
if attention_mask is not None:
|
||||||
|
# Import _expand_mask or _prepare_4d_attention_mask based on version
|
||||||
|
if self.transformers_version < v4_35:
|
||||||
|
from transformers.models.clip.modeling_clip import _expand_mask
|
||||||
|
attention_mask = _expand_mask(attention_mask, hidden_states.dtype)
|
||||||
|
else:
|
||||||
|
from transformers.modeling_attn_mask_utils import _prepare_4d_attention_mask
|
||||||
|
attention_mask = _prepare_4d_attention_mask(attention_mask, hidden_states.dtype)
|
||||||
|
|
||||||
|
return causal_attention_mask, attention_mask
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -96,29 +190,19 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
|
|||||||
# input_shape = input_ids.size()
|
# input_shape = input_ids.size()
|
||||||
# input_ids = input_ids.view(-1, input_shape[-1])
|
# input_ids = input_ids.view(-1, input_shape[-1])
|
||||||
|
|
||||||
|
for batch_idx, batch in enumerate(input_ids):
|
||||||
|
for seg_or_action in batch:
|
||||||
|
if isinstance(seg_or_action, Action):
|
||||||
|
seg_or_action.process_with_transformer(
|
||||||
|
self, self.embeddings.token_embedding
|
||||||
|
)
|
||||||
|
|
||||||
hidden_states = self.embeddings(input_dicts=input_ids)
|
hidden_states = self.embeddings(input_dicts=input_ids)
|
||||||
|
|
||||||
bsz = len(input_ids)
|
bsz = len(input_ids)
|
||||||
# TODO: Properly gather this
|
seq_len = hidden_states.shape[1]
|
||||||
seq_len = 77
|
|
||||||
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
|
|
||||||
## 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
|
causal_attention_mask, attention_mask = self.process_attention_mask(hidden_states, attention_mask, bsz, seq_len)
|
||||||
if attention_mask is not None:
|
|
||||||
# [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
|
|
||||||
attention_mask = _expand_mask(attention_mask, hidden_states.dtype)
|
|
||||||
|
|
||||||
encoder_outputs = self.encoder(
|
encoder_outputs = self.encoder(
|
||||||
inputs_embeds=hidden_states,
|
inputs_embeds=hidden_states,
|
||||||
@@ -141,8 +225,11 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
|
|||||||
if isinstance(seg_or_action, Action):
|
if isinstance(seg_or_action, Action):
|
||||||
idx += seg_or_action.token_length()
|
idx += seg_or_action.token_length()
|
||||||
else:
|
else:
|
||||||
if seg_or_action.text == '__PAD__':
|
if seg_or_action.text == "__PAD__" or seg_or_action.text == "[EOT]":
|
||||||
break
|
break
|
||||||
|
|
||||||
|
# Is a segment, and isn't the pad segment
|
||||||
|
idx += seg_or_action.token_length()
|
||||||
eot_idx.append(idx)
|
eot_idx.append(idx)
|
||||||
# text_embeds.shape = [batch_size, sequence_length, transformer.width]
|
# text_embeds.shape = [batch_size, sequence_length, transformer.width]
|
||||||
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
||||||
|
|||||||
@@ -9,9 +9,9 @@ def register_action(action: Type[Action]) -> None:
|
|||||||
|
|
||||||
:rtype: object
|
:rtype: object
|
||||||
"""
|
"""
|
||||||
if action.name in action_registry:
|
if action.action_name in action_registry:
|
||||||
raise ValueError(f"Action {action.name} already registered")
|
raise ValueError(f"Action {action.action_name} already registered")
|
||||||
action_registry[str(action.name)] = action
|
action_registry[str(action.action_name)] = action
|
||||||
|
|
||||||
def get_action_by_name(name: str) -> Type[Action]:
|
def get_action_by_name(name: str) -> Type[Action]:
|
||||||
if name not in action_registry:
|
if name not in action_registry:
|
||||||
|
|||||||
@@ -2,32 +2,28 @@ from typing import List
|
|||||||
|
|
||||||
from lark import Transformer, Token
|
from lark import Transformer, Token
|
||||||
|
|
||||||
from comfy.sd1_clip import SD1Tokenizer
|
from comfy.sd1_clip import SDTokenizer
|
||||||
from custom_nodes.KepPromptLang.lib.action.base import Action, ActionArity
|
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.registration import get_action_by_name
|
||||||
from custom_nodes.KepPromptLang.lib.parser.utils import build_prompt_segment
|
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
|
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||||
|
|
||||||
|
|
||||||
class PromptTransformer(Transformer):
|
class PromptTransformer(Transformer):
|
||||||
|
"""
|
||||||
|
Transforms the parsed prompt into a list of segments and actions
|
||||||
|
Types from the grammar are mapped to the methods in this class
|
||||||
|
"""
|
||||||
# def WORD(self, items):
|
# def WORD(self, items):
|
||||||
# return items
|
# return items
|
||||||
|
|
||||||
def __init__(self, tokenizer: SD1Tokenizer):
|
def __init__(self, tokenizer: SDTokenizer):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.tokenizer = tokenizer
|
self.tokenizer = tokenizer
|
||||||
|
|
||||||
def item(self, items: List[Token]):
|
def item(self, items: List[Token]):
|
||||||
for item in items:
|
for item in items:
|
||||||
if isinstance(item, Action):
|
if isinstance(item, (Action, PromptSegment)):
|
||||||
return item
|
|
||||||
|
|
||||||
if isinstance(item, PromptSegment):
|
|
||||||
return item
|
return item
|
||||||
|
|
||||||
if item.type == "WORD":
|
if item.type == "WORD":
|
||||||
@@ -41,7 +37,7 @@ class PromptTransformer(Transformer):
|
|||||||
elif item.type == "embedding":
|
elif item.type == "embedding":
|
||||||
return build_prompt_segment(item, self.tokenizer)
|
return build_prompt_segment(item, self.tokenizer)
|
||||||
elif item.type == "function":
|
elif item.type == "function":
|
||||||
return item
|
raise Exception("Unexpected type in prompt transformer: function. Please report this issue on GitHub.")
|
||||||
else:
|
else:
|
||||||
raise Exception("Unknown item type: " + str(item.type))
|
raise Exception("Unknown item type: " + str(item.type))
|
||||||
|
|
||||||
|
|||||||
+8
-5
@@ -1,6 +1,6 @@
|
|||||||
from lark import Token
|
from lark import Token
|
||||||
|
|
||||||
from comfy.sd1_clip import SD1Tokenizer
|
from comfy.sd1_clip import SDTokenizer
|
||||||
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||||
|
|
||||||
|
|
||||||
@@ -11,7 +11,7 @@ def flatten_tree(tree):
|
|||||||
return [str(tree.data)] + sum([flatten_tree(child) for child in tree.children], [])
|
return [str(tree.data)] + sum([flatten_tree(child) for child in tree.children], [])
|
||||||
|
|
||||||
|
|
||||||
def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment:
|
def build_prompt_segment(text: str, tokenizer: SDTokenizer) -> PromptSegment:
|
||||||
split_text = text.split(" ")
|
split_text = text.split(" ")
|
||||||
tokens = []
|
tokens = []
|
||||||
for word in split_text:
|
for word in split_text:
|
||||||
@@ -24,10 +24,13 @@ def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment:
|
|||||||
if embedding is None:
|
if embedding is None:
|
||||||
print(f"warning, embedding:{embedding_name} does not exist, ignoring")
|
print(f"warning, embedding:{embedding_name} does not exist, ignoring")
|
||||||
else:
|
else:
|
||||||
if len(embedding.shape) == 1:
|
if embedding.shape[1] != tokenizer.embedding_size:
|
||||||
tokens.append(embedding)
|
print(f"warning, embedding:{embedding_name} has size {embedding.shape[1]}, expected {tokenizer.embedding_size}, ignoring")
|
||||||
else:
|
else:
|
||||||
tokens.extend(embedding)
|
if len(embedding.shape) == 1:
|
||||||
|
tokens.append(embedding)
|
||||||
|
else:
|
||||||
|
tokens.extend(embedding)
|
||||||
|
|
||||||
if leftover != "":
|
if leftover != "":
|
||||||
word = leftover
|
word = leftover
|
||||||
|
|||||||
+27
-7
@@ -1,18 +1,18 @@
|
|||||||
from typing import List
|
from typing import List, Dict
|
||||||
|
|
||||||
from lark import Tree
|
from lark import Tree
|
||||||
|
|
||||||
from comfy.sd1_clip import SD1Tokenizer
|
from comfy.sd1_clip import SD1Tokenizer, SDTokenizer
|
||||||
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
||||||
|
|
||||||
from custom_nodes.KepPromptLang.lib.parser import PromptParser
|
from custom_nodes.KepPromptLang.lib.parser import PromptParser
|
||||||
from custom_nodes.KepPromptLang.lib.parser.transformer import PromptTransformer
|
from custom_nodes.KepPromptLang.lib.parser.transformer import PromptTransformer
|
||||||
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
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):
|
|
||||||
super().__init__(tokenizer_path, max_length, pad_with_end, embedding_directory, embedding_size, embedding_key)
|
|
||||||
|
|
||||||
|
class PromptLangSDTokenizer(SDTokenizer):
|
||||||
|
def __init__(self, tokenizer_path=None, max_length=77, pad_with_end=True, embedding_directory=None, embedding_size=768, embedding_key='clip_l'):
|
||||||
|
super().__init__(tokenizer_path, max_length, pad_with_end, embedding_directory, embedding_size, embedding_key)
|
||||||
"""
|
"""
|
||||||
Doesn't actually tokenize...
|
Doesn't actually tokenize...
|
||||||
Returns batches of segments and actions
|
Returns batches of segments and actions
|
||||||
@@ -43,9 +43,9 @@ class PromptLangTokenizer(SD1Tokenizer):
|
|||||||
|
|
||||||
# If the segment is too large to fit in a single batch, pad the current batch and start a new one
|
# If the segment is too large to fit in a single batch, pad the current batch and start a new one
|
||||||
if num_tokens + batch_size > self.max_length - 1:
|
if num_tokens + batch_size > self.max_length - 1:
|
||||||
remaining_length = self.max_length - batch_size - 1 # -1 for end token
|
remaining_length = self.max_length - batch_size
|
||||||
# Pad batch
|
# Pad batch
|
||||||
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length - 1))
|
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * (remaining_length - 1))) # -1 for end token
|
||||||
batched_segments.append(batch)
|
batched_segments.append(batch)
|
||||||
|
|
||||||
# start new batch
|
# start new batch
|
||||||
@@ -66,3 +66,23 @@ class PromptLangTokenizer(SD1Tokenizer):
|
|||||||
# batch_size_info(batch)
|
# batch_size_info(batch)
|
||||||
|
|
||||||
return batched_segments
|
return batched_segments
|
||||||
|
|
||||||
|
class PromptLangSD1Tokenizer(SD1Tokenizer):
|
||||||
|
def __init__(self, embedding_directory=None, clip_name='l', tokenizer=PromptLangSDTokenizer) -> None:
|
||||||
|
super().__init__(embedding_directory, clip_name, tokenizer)
|
||||||
|
|
||||||
|
|
||||||
|
class PromptLangSDXLClipGTokenizer(PromptLangSDTokenizer):
|
||||||
|
def __init__(self, tokenizer_path=None, embedding_directory=None):
|
||||||
|
super().__init__(tokenizer_path, pad_with_end=False, embedding_directory=embedding_directory, embedding_size=1280, embedding_key='clip_g')
|
||||||
|
|
||||||
|
class PromptLangSDXLTokenizer(SD1Tokenizer):
|
||||||
|
def __init__(self, embedding_directory=None) -> None:
|
||||||
|
self.clip_l = PromptLangSDTokenizer(embedding_directory=embedding_directory)
|
||||||
|
self.clip_g = PromptLangSDXLClipGTokenizer(embedding_directory=embedding_directory)
|
||||||
|
|
||||||
|
def tokenize_with_weights(self, text:str, return_word_ids=False) -> Dict[str, List[List[SegOrAction]]]:
|
||||||
|
out = {}
|
||||||
|
out["g"] = self.clip_g.tokenize_with_weights(text, return_word_ids)
|
||||||
|
out["l"] = self.clip_l.tokenize_with_weights(text, return_word_ids)
|
||||||
|
return out
|
||||||
|
|||||||
@@ -1,4 +1,3 @@
|
|||||||
import random
|
|
||||||
import os
|
import os
|
||||||
from typing import List, Tuple, Any
|
from typing import List, Tuple, Any
|
||||||
|
|
||||||
@@ -8,34 +7,24 @@ from PIL import Image
|
|||||||
import folder_paths
|
import folder_paths
|
||||||
import comfy.sd
|
import comfy.sd
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
from custom_nodes.KepPromptLang.lib.clip_model import PromptLangClipModel
|
from comfy.sd2_clip import SD2ClipModel
|
||||||
|
from comfy.sdxl_clip import SDXLClipModel
|
||||||
|
from comfy.supported_models_base import ClipTarget
|
||||||
|
from custom_nodes.KepPromptLang.lib.clip_model import (
|
||||||
|
PromptLangSDXLClipModel,
|
||||||
|
PromptLangSD1ClipModel,
|
||||||
|
)
|
||||||
|
|
||||||
from custom_nodes.KepPromptLang.lib.tokenizer import PromptLangTokenizer
|
from custom_nodes.KepPromptLang.lib.tokenizer import (
|
||||||
|
PromptLangSDXLTokenizer,
|
||||||
|
PromptLangSD1Tokenizer,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class EmptyClass:
|
class EmptyClass:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
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:
|
class SpecialClipLoader:
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls): # type: ignore
|
def INPUT_TYPES(cls): # type: ignore
|
||||||
@@ -52,15 +41,22 @@ class SpecialClipLoader:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def load_clip(source_clip: comfy.sd.CLIP) -> Tuple[comfy.sd.CLIP]:
|
def load_clip(source_clip: comfy.sd.CLIP) -> Tuple[comfy.sd.CLIP]:
|
||||||
clip_target = EmptyClass()
|
|
||||||
clip_target.params = {}
|
|
||||||
clip_target.clip = PromptLangClipModel
|
|
||||||
clip_target.tokenizer = PromptLangTokenizer
|
|
||||||
|
|
||||||
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.embedding_directory)
|
if isinstance(source_clip.cond_stage_model, SDXLClipModel):
|
||||||
comfy.sd.load_clip_weights(
|
clip_target = ClipTarget(PromptLangSDXLTokenizer, PromptLangSDXLClipModel)
|
||||||
clip.cond_stage_model, source_clip.cond_stage_model.state_dict()
|
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.clip_g.embedding_directory)
|
||||||
)
|
comfy.sd.load_clip_weights(clip.cond_stage_model.clip_g,source_clip.cond_stage_model.clip_g.state_dict())
|
||||||
|
comfy.sd.load_clip_weights(
|
||||||
|
clip.cond_stage_model.clip_l, source_clip.cond_stage_model.clip_l.state_dict()
|
||||||
|
)
|
||||||
|
elif isinstance(source_clip, SD2ClipModel):
|
||||||
|
raise ValueError("SD2 Clip model is not supported.")
|
||||||
|
else:
|
||||||
|
clip_target = ClipTarget(PromptLangSD1Tokenizer, PromptLangSD1ClipModel)
|
||||||
|
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.clip_l.embedding_directory)
|
||||||
|
comfy.sd.load_clip_weights(
|
||||||
|
clip.cond_stage_model, source_clip.cond_stage_model.state_dict()
|
||||||
|
)
|
||||||
return (clip,)
|
return (clip,)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1 +1,2 @@
|
|||||||
lark
|
lark
|
||||||
|
packaging
|
||||||
|
|||||||
@@ -0,0 +1,83 @@
|
|||||||
|
import importlib
|
||||||
|
import inspect
|
||||||
|
import os
|
||||||
|
from typing import List, Type
|
||||||
|
|
||||||
|
from custom_nodes.KepPromptLang.lib.action.base import Action
|
||||||
|
|
||||||
|
|
||||||
|
EXCLUDED_MODULES = ["utils.py", "action_utils.py", "types.py"]
|
||||||
|
def import_module_from_path(path: str):
|
||||||
|
module_name = path.replace("/", ".")[:-3]
|
||||||
|
return importlib.import_module(module_name, package="custom_nodes.KepPromptLang.lib.actions")
|
||||||
|
|
||||||
|
|
||||||
|
# Function to find and import action classes
|
||||||
|
def find_action_classes(directory: str) -> List[Type[Action]]:
|
||||||
|
action_classes = []
|
||||||
|
for filename in os.listdir(directory):
|
||||||
|
if filename.endswith(".py") and not filename.startswith("__") and filename not in EXCLUDED_MODULES:
|
||||||
|
module_path = os.path.join(directory, filename)
|
||||||
|
module = import_module_from_path(module_path)
|
||||||
|
for name, obj in inspect.getmembers(module, inspect.isclass):
|
||||||
|
if issubclass(obj, Action) and obj is not Action and obj.__name__ != "MultiArgAction" and obj.__name__ != "SingleArgAction":
|
||||||
|
action_classes.append(obj)
|
||||||
|
return action_classes
|
||||||
|
|
||||||
|
|
||||||
|
# Function to extract info from an action class
|
||||||
|
def extract_class_info(cls: Type[Action]) -> dict:
|
||||||
|
class_info = {
|
||||||
|
'class_name': cls.__name__,
|
||||||
|
'properties': {
|
||||||
|
'display_name': getattr(cls, 'display_name', None),
|
||||||
|
'action_name': getattr(cls, 'action_name', None),
|
||||||
|
'description': getattr(cls, 'description', None),
|
||||||
|
'usage_examples': getattr(cls, 'usage_examples', None)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return class_info
|
||||||
|
|
||||||
|
def escape_pipes(text: str) -> str:
|
||||||
|
return text.replace('|', '\\|')
|
||||||
|
|
||||||
|
def generate_markdown_documentation(classes_info: List[dict]) -> str:
|
||||||
|
documentation = "# Actions Documentation\n\n"
|
||||||
|
|
||||||
|
# Define table columns
|
||||||
|
# columns = ["Class", "Display Name", "Action Name", "Description", "Usage Examples"]
|
||||||
|
columns = ["Display Name", "Action Name", "Description", "Usage Examples"]
|
||||||
|
documentation += "| " + " | ".join(columns) + " |\n"
|
||||||
|
documentation += "| --- " * len(columns) + "|\n"
|
||||||
|
|
||||||
|
for cls_info in classes_info:
|
||||||
|
# row = [cls_info['class_name']]
|
||||||
|
row = []
|
||||||
|
# Iterate over properties in a predefined order
|
||||||
|
for prop in ["display_name", "action_name", "description", "usage_examples"]:
|
||||||
|
prop_doc = cls_info['properties'].get(prop, 'N/A')
|
||||||
|
|
||||||
|
# Format and escape usage examples
|
||||||
|
if isinstance(prop_doc, list):
|
||||||
|
escaped_examples = [escape_pipes(example) for example in prop_doc]
|
||||||
|
prop_doc = "<ul>" + "".join([f"<li>{example}</li>" for example in escaped_examples]) + "</ul>"
|
||||||
|
else:
|
||||||
|
prop_doc = escape_pipes(prop_doc)
|
||||||
|
|
||||||
|
row.append(prop_doc)
|
||||||
|
|
||||||
|
documentation += "| " + " | ".join(row) + " |\n"
|
||||||
|
|
||||||
|
return documentation
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
# Main execution
|
||||||
|
if __name__ == "__main__":
|
||||||
|
actions_directory = (
|
||||||
|
"../lib/actions" # Update this path as per your project structure
|
||||||
|
)
|
||||||
|
action_classes = find_action_classes(actions_directory)
|
||||||
|
class_infos = [extract_class_info(cls) for cls in action_classes]
|
||||||
|
docs = generate_markdown_documentation(class_infos)
|
||||||
|
print(docs) # Or write to a file
|
||||||
Generated
-12963
File diff suppressed because one or more lines are too long
Generated
-2047
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.
BIN
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Vendored
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
Binary file not shown.
BIN
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.
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
Reference in New Issue
Block a user