22 Commits
Author SHA1 Message Date
Michael Poutre a0b3695800 refactor(Docs): Add table of all actions, and tiny cleanup 2023-11-19 00:26:57 -08:00
Michael Poutre c299354c9c feat(Docs): Add script for generating docs 2023-11-19 00:26:33 -08:00
Michael Poutre e002d0ad64 refactor(Action): Update property names to allow automatic doc gen 2023-11-19 00:26:19 -08:00
Michael Poutre b5a728a997 feat(Action): Add experimental pooledAvg, pooler actions(_exp- prefix) 2023-11-18 23:41:27 -08:00
Michael Poutre 72f46ad938 feat(Action): Add PostPost action 2023-11-18 23:07:01 -08:00
Michael Poutre f74728267d feat: Add call to process_with_transformers for custom actions 2023-11-18 23:05:42 -08:00
Michael Poutre 0de2ab3f9a fix: Properly determine seq_len for TextTransformer processing 2023-11-18 23:04:56 -08:00
Michael Poutre 2fa1b95045 fix: Handle EOT as well as __PAD__ segments for pooler output EOT calc 2023-11-18 23:04:06 -08:00
Michael Poutre dba41a2c4d fix: Drop invalid embeddings during prompt parsing to avoid it later 2023-11-18 22:48:10 -08:00
Michael Poutre 3cbfc8b78c feat(Actions): Add bypass_pos_embed PostModifier 2023-11-18 22:46:39 -08:00
Michael Poutre fe74c556f5 refactor(PromptTransformer): Simple cleanup and consolidation 2023-11-18 21:36:28 -08:00
Michael Poutre d042fc572f fix: Update importlib.metadata import 2023-11-13 19:54:15 -08:00
Michael Poutre 46ae926911 fix: Update to work with transformers changes to attention masks 2023-11-13 19:47:56 -08:00
Michael Poutre 2823a80078 feat(PosScale: Add node 2023-11-13 19:47:56 -08:00
Michael Poutre fccffdbdf0 feat: Add support for post embedding modifiers 2023-11-13 19:47:56 -08:00
Michael Poutre 7e138268f9 chore: Cleanup imports 2023-11-13 19:47:56 -08:00
Michael Poutre ef3f66ef91 refactor/fix: More fixes to align with clip changes in base 2023-11-13 19:47:56 -08:00
Michael Poutre 2f115352ec refactor/fix(SD1): Update for changes to clip handling
Fix end token calculation for long prompts
2023-11-13 19:47:56 -08:00
Michael Poutre a0ebcdaf4d fix(Clip): Fix bug in EOT token detection 2023-11-13 19:47:56 -08:00
Michael Poutre 55cc09b0be feat(SDXL): Initial stab at SDXL 2023-11-13 19:47:56 -08:00
Michael Poutre 2e731a480d refactor(Node): Add checks for clip model type 2023-11-13 19:47:56 -08:00
Michael Poutre 65bc2c0327 refactor(Node): Use new comfy.supported_models_base.ClipTarget 2023-11-13 19:47:56 -08:00
24 changed files with 486 additions and 69 deletions
+23 -25
View File
@@ -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> |
+7
View File
@@ -3,7 +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.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
@@ -22,3 +26,6 @@ register_action(AverageAction)
register_action(ScaleDims) register_action(ScaleDims)
register_action(SetDims) register_action(SetDims)
register_action(PosScaleAction) register_action(PosScaleAction)
register_action(PoolerAction)
register_action(PostPosAction)
register_action(PooledAvgAction)
+36 -2
View File
@@ -4,6 +4,7 @@ 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
@@ -18,6 +19,7 @@ class PostModifiers(TypedDict):
A dictionary of post modifiers for an action result. A dictionary of post modifiers for an action result.
""" """
position_embed_scale: Union[float, None] position_embed_scale: Union[float, None]
bypass_pos_embed: Union[bool, None]
class Action(ABC): class Action(ABC):
@@ -37,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:
@@ -82,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()
@@ -97,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
@@ -117,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]]],
+9
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
+68
View File
@@ -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?"
)
+67
View File
@@ -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?"
)
+7 -1
View File
@@ -13,9 +13,15 @@ from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
class PosScaleAction(MultiArgAction): class PosScaleAction(MultiArgAction):
grammar = 'posScale(" arg+ ")"' grammar = 'posScale(" arg+ ")"'
name = "posScale"
chars = ["[", "]"] 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: def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args) super().__init__(args)
if len(args) != 2: if len(args) != 2:
+47
View File
@@ -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)
+8
View File
@@ -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
+7 -3
View File
@@ -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)
+7 -3
View File
@@ -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)
+7 -1
View File
@@ -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
View File
@@ -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)
+3 -4
View File
@@ -108,12 +108,11 @@ class PromptLangSDClipModel(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
+43 -13
View File
@@ -1,5 +1,5 @@
from typing import Optional, Tuple, Union, List, TypedDict from typing import Optional, Tuple, Union, List, TypedDict, TYPE_CHECKING
from importlib_metadata import version as import_version from importlib.metadata import version as import_version
from packaging import version from packaging import version
import torch import torch
@@ -14,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)
@@ -46,11 +51,10 @@ 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]] = [] pos_modifiers: List[List[PosModifier]] = []
for batch_idx, batch in enumerate(input_dicts): for batch_idx, batch in enumerate(input_dicts):
results = [] results = []
@@ -58,15 +62,23 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
token_idx = 0 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):
action_result = 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): if isinstance(action_result, tuple):
result, post_modifiers = action_result result, post_modifiers = action_result
if post_modifiers["position_embed_scale"] is not None: if post_modifiers.get("position_embed_scale", None) is not None:
post_modifiers["start_idx"] = token_idx post_modifiers["start_idx"] = token_idx
post_modifiers["end_idx"] = ( post_modifiers["end_idx"] = (
token_idx + seg_or_action.token_length() token_idx + seg_or_action.token_length()
) )
batch_pos_modifiers.append(post_modifiers)
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: else:
result = action_result result = action_result
else: else:
@@ -88,14 +100,26 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
else: else:
embeds.append(torch.cat(batch, dim=-2)) embeds.append(torch.cat(batch, dim=-2))
# Iterate over the batches and apply the pos modifiers to the position embeddings then add them to the embeddings
for idx, batch_pos_modifiers in enumerate(pos_modifiers): for idx, batch_pos_modifiers in enumerate(pos_modifiers):
position_embeddings = self.position_embedding(position_ids) position_embeddings = self.position_embedding(position_ids)
if len(batch_pos_modifiers) > 0: if len(batch_pos_modifiers) > 0:
print(f"Found {len(batch_pos_modifiers)} pos modifiers for batch {idx}") 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: for post_modifier in batch_pos_modifiers:
position_embeddings[ if post_modifier.get("bypass_pos_embed", False):
0, post_modifier["start_idx"] : post_modifier["end_idx"] position_embeddings[
] *= post_modifier["position_embed_scale"] 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 embeds[idx] = embeds[idx] + position_embeddings
embeddings = torch.cat(embeds, dim=0) embeddings = torch.cat(embeds, dim=0)
@@ -166,11 +190,17 @@ 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
causal_attention_mask, attention_mask = self.process_attention_mask(hidden_states, attention_mask, bsz, seq_len) causal_attention_mask, attention_mask = self.process_attention_mask(hidden_states, attention_mask, bsz, seq_len)
@@ -195,7 +225,7 @@ 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 # Is a segment, and isn't the pad segment
+3 -3
View File
@@ -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:
+6 -5
View File
@@ -10,6 +10,10 @@ 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
@@ -19,10 +23,7 @@ class PromptTransformer(Transformer):
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":
@@ -36,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))
+6 -3
View File
@@ -24,10 +24,13 @@ def build_prompt_segment(text: str, tokenizer: SDTokenizer) -> 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
+83
View File
@@ -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