Compare commits
10
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9df0382735 | ||
|
|
239113501a | ||
|
|
c3d45e306e | ||
|
|
611ee06761 | ||
|
|
97f48d0029 | ||
|
|
1059daf5e9 | ||
|
|
2a4b724d59 | ||
|
|
55918edc54 | ||
|
|
21ba1299d6 | ||
|
|
407724a64c |
@@ -29,10 +29,33 @@ See example workflow in examples folder.
|
||||
- Example: `"Hello World"`, `'It\'s a sunny day'`
|
||||
- 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:
|
||||
- Each function takes one or more arguments.
|
||||
- An argument (`arg`) can be an embedding, multiple words, another function, or a quoted string.
|
||||
- An argument (`arg`) can be an embedding, a word, another function, or a quoted string.
|
||||
- For functions that accept multiple arguments, they are separated by the `|` symbol.
|
||||
|
||||
## Examples
|
||||
@@ -55,25 +78,4 @@ See example workflow in examples folder.
|
||||
```
|
||||
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> |
|
||||
|
||||
```
|
||||
|
||||
@@ -3,11 +3,7 @@ from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction
|
||||
from custom_nodes.KepPromptLang.lib.actions.mult import MultiplyAction
|
||||
from custom_nodes.KepPromptLang.lib.actions.neg import NegAction
|
||||
from custom_nodes.KepPromptLang.lib.actions.norm import NormAction
|
||||
|
||||
from custom_nodes.KepPromptLang.lib.actions.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.scale_dims import ScaleDims
|
||||
from custom_nodes.KepPromptLang.lib.actions.set_dims import SetDims
|
||||
@@ -26,6 +22,3 @@ register_action(AverageAction)
|
||||
register_action(ScaleDims)
|
||||
register_action(SetDims)
|
||||
register_action(PosScaleAction)
|
||||
register_action(PoolerAction)
|
||||
register_action(PostPosAction)
|
||||
register_action(PooledAvgAction)
|
||||
|
||||
+2
-36
@@ -4,7 +4,6 @@ from typing import Union, List, TypedDict, Tuple
|
||||
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
from transformers.models.clip.modeling_clip import CLIPTextTransformer
|
||||
|
||||
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
@@ -19,7 +18,6 @@ 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):
|
||||
@@ -39,23 +37,9 @@ class Action(ABC):
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def display_name(self) -> str:
|
||||
def name(self) -> str:
|
||||
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
|
||||
@abstractmethod
|
||||
def grammar(self) -> str:
|
||||
@@ -98,13 +82,6 @@ class Action(ABC):
|
||||
"""
|
||||
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:
|
||||
raise NotImplementedError()
|
||||
|
||||
@@ -120,17 +97,12 @@ class SingleArgAction(Action, ABC):
|
||||
segments.append(seg_or_action)
|
||||
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]]):
|
||||
# TODO: Target is a list now... what does this mean for us..
|
||||
self.arg = arg
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.display_name}({self.arg})"
|
||||
return f"{self.name}({self.arg})"
|
||||
|
||||
class MultiArgAction(Action, ABC):
|
||||
arity = ActionArity.MULTI
|
||||
@@ -145,12 +117,6 @@ class MultiArgAction(Action, ABC):
|
||||
|
||||
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__(
|
||||
self,
|
||||
args: List[List[Union[PromptSegment, Action]]],
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
from typing import List
|
||||
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
|
||||
@@ -11,10 +9,3 @@ def get_embedding(seg_or_action: SegOrAction, embedding_module: Embedding) -> Te
|
||||
if isinstance(seg_or_action, Action):
|
||||
return seg_or_action.get_result(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
|
||||
|
||||
+1
-9
@@ -12,17 +12,9 @@ from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
class AverageAction(MultiArgAction):
|
||||
grammar = 'avg(" arg "|" arg "|" arg ")"'
|
||||
name = "avg"
|
||||
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:
|
||||
super().__init__(args)
|
||||
|
||||
|
||||
+1
-9
@@ -11,17 +11,9 @@ from custom_nodes.KepPromptLang.lib.parser.registration import register_action
|
||||
|
||||
class DiffAction(MultiArgAction):
|
||||
grammar = 'diff(" arg ("|" arg)* ")"'
|
||||
name = "diff"
|
||||
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]]]):
|
||||
super().__init__(args)
|
||||
self.base_arg = args[0]
|
||||
|
||||
+1
-8
@@ -13,16 +13,9 @@ from custom_nodes.KepPromptLang.lib.parser.registration import register_action
|
||||
|
||||
class MultiplyAction(MultiArgAction):
|
||||
grammar = 'mult(" arg+ ")"'
|
||||
name = "mult"
|
||||
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:
|
||||
super().__init__(args)
|
||||
if len(args) != 2:
|
||||
|
||||
+1
-8
@@ -7,16 +7,9 @@ from custom_nodes.KepPromptLang.lib.parser.registration import register_action
|
||||
|
||||
class NegAction(SingleArgAction):
|
||||
grammar = 'neg(" arg+ ")"'
|
||||
name = "neg"
|
||||
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:
|
||||
"""
|
||||
Neg negates the embeddings of the base segment, so the length is the length of the base segment
|
||||
|
||||
+1
-8
@@ -10,16 +10,9 @@ from custom_nodes.KepPromptLang.lib.parser.registration import register_action
|
||||
|
||||
class NormAction(SingleArgAction):
|
||||
grammar = 'norm(" arg+ ")"'
|
||||
name = "norm"
|
||||
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:
|
||||
"""
|
||||
Norm normalizes the embeddings of the base segment, so the length is the length of the base segment
|
||||
|
||||
@@ -1,68 +0,0 @@
|
||||
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?"
|
||||
)
|
||||
@@ -1,67 +0,0 @@
|
||||
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?"
|
||||
)
|
||||
@@ -13,15 +13,9 @@ from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
||||
|
||||
class PosScaleAction(MultiArgAction):
|
||||
grammar = 'posScale(" arg+ ")"'
|
||||
name = "posScale"
|
||||
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:
|
||||
|
||||
@@ -1,47 +0,0 @@
|
||||
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,14 +17,6 @@ class RandAction(MultiArgAction):
|
||||
name = "rand"
|
||||
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
|
||||
range_min = 0
|
||||
range_max = 1
|
||||
|
||||
@@ -10,14 +10,10 @@ from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
class ScaleDims(MultiArgAction):
|
||||
grammar = 'scaleDims(" arg ("|" arg)* ")"'
|
||||
chars = ["-", "-"]
|
||||
|
||||
display_name = "Scale Dimensions"
|
||||
action_name = "scaleDims"
|
||||
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",
|
||||
]
|
||||
example = "'The scaleDims(cat|4,1.5|76,1.2) is happy' scales the 4th dimension by 1.5 and the 76th dimension by 1.2 for the word 'cat'"
|
||||
chars = ["-", "-"]
|
||||
|
||||
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
|
||||
super().__init__(args)
|
||||
|
||||
@@ -10,14 +10,10 @@ from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
class SetDims(MultiArgAction):
|
||||
grammar = 'setDims(" arg ("|" arg)* ")"'
|
||||
chars = ["-", "-"]
|
||||
|
||||
display_name = "Set Dimensions"
|
||||
action_name = "setDims"
|
||||
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"
|
||||
]
|
||||
example = "'The scaleDims(cat|4,1.5|76,1.2) is happy' scales the 4th dimension by 1.5 and the 76th dimension by 1.2 for the word 'cat'"
|
||||
chars = ["-", "-"]
|
||||
|
||||
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
|
||||
super().__init__(args)
|
||||
|
||||
@@ -12,15 +12,9 @@ from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
class SlerpAction(MultiArgAction):
|
||||
grammar = 'slerp(" arg "|" arg "|" arg ")"'
|
||||
name = "slerp"
|
||||
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:
|
||||
super().__init__(args)
|
||||
|
||||
|
||||
+1
-7
@@ -11,15 +11,9 @@ from custom_nodes.KepPromptLang.lib.parser.registration import register_action
|
||||
|
||||
class SumAction(MultiArgAction):
|
||||
grammar = 'sum(" arg ("|" arg)+ ")"'
|
||||
name = "sum"
|
||||
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:
|
||||
super().__init__(args)
|
||||
|
||||
|
||||
+4
-3
@@ -108,11 +108,12 @@ class PromptLangSDClipModel(torch.nn.Module):
|
||||
tokens_temp += [next_new_token]
|
||||
next_new_token += 1
|
||||
else:
|
||||
raise Exception("WARNING: shape mismatch when trying to apply embedding. Should have been caught during tokenization.",
|
||||
print("WARNING: shape mismatch when trying to apply embedding, embedding will be ignored",
|
||||
tid_or_tensor.shape[0], current_embeds.weight.shape[1])
|
||||
if len(tokens_temp) < segment_length:
|
||||
# This should never happen...
|
||||
raise Exception("Segment size mismatch. Please submit an issue on Github.")
|
||||
# Pretty sure this is only needed if the embedding is not the same size as the CLIP embedding
|
||||
print("WARNING: segment length mismatch, padding with EOS token")
|
||||
tokens_temp.extend([self.empty_tokens[0][-1] * (segment_length - len(tokens_temp))])
|
||||
segment.tokens = tokens_temp
|
||||
|
||||
n = token_dict_size
|
||||
|
||||
+13
-43
@@ -1,5 +1,5 @@
|
||||
from typing import Optional, Tuple, Union, List, TypedDict, TYPE_CHECKING
|
||||
from importlib.metadata import version as import_version
|
||||
from typing import Optional, Tuple, Union, List, TypedDict
|
||||
from importlib_metadata import version as import_version
|
||||
from packaging import version
|
||||
|
||||
import torch
|
||||
@@ -14,11 +14,6 @@ from transformers.models.clip.modeling_clip import (
|
||||
from custom_nodes.KepPromptLang.lib.action.base import Action
|
||||
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):
|
||||
low = low.unsqueeze(0)
|
||||
high = high.unsqueeze(0)
|
||||
@@ -51,10 +46,11 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||
) -> torch.Tensor:
|
||||
|
||||
if input_dicts is None:
|
||||
raise ValueError("You have to specify input_dicts")
|
||||
|
||||
batches: List[List[Tensor | Tuple[Tensor, PostModifiers] | Action]] = []
|
||||
batches = []
|
||||
pos_modifiers: List[List[PosModifier]] = []
|
||||
for batch_idx, batch in enumerate(input_dicts):
|
||||
results = []
|
||||
@@ -62,23 +58,15 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||
token_idx = 0
|
||||
for seg_or_action in batch:
|
||||
if isinstance(seg_or_action, Action):
|
||||
action_result: Union[
|
||||
Tensor, Tuple[Tensor, PostModifiers]
|
||||
] = seg_or_action.get_result(self.token_embedding)
|
||||
action_result = 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:
|
||||
if post_modifiers["position_embed_scale"] 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)
|
||||
batch_pos_modifiers.append(post_modifiers)
|
||||
else:
|
||||
result = action_result
|
||||
else:
|
||||
@@ -100,26 +88,14 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||
else:
|
||||
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):
|
||||
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
|
||||
position_embeddings[
|
||||
0, post_modifier["start_idx"] : post_modifier["end_idx"]
|
||||
] *= post_modifier["position_embed_scale"]
|
||||
embeds[idx] = embeds[idx] + position_embeddings
|
||||
embeddings = torch.cat(embeds, dim=0)
|
||||
|
||||
@@ -190,17 +166,11 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
|
||||
# input_shape = input_ids.size()
|
||||
# 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)
|
||||
|
||||
bsz = len(input_ids)
|
||||
seq_len = hidden_states.shape[1]
|
||||
# TODO: Properly gather this
|
||||
seq_len = 77
|
||||
|
||||
causal_attention_mask, attention_mask = self.process_attention_mask(hidden_states, attention_mask, bsz, seq_len)
|
||||
|
||||
@@ -225,7 +195,7 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
|
||||
if isinstance(seg_or_action, Action):
|
||||
idx += seg_or_action.token_length()
|
||||
else:
|
||||
if seg_or_action.text == "__PAD__" or seg_or_action.text == "[EOT]":
|
||||
if seg_or_action.text == '__PAD__':
|
||||
break
|
||||
|
||||
# Is a segment, and isn't the pad segment
|
||||
|
||||
@@ -9,9 +9,9 @@ def register_action(action: Type[Action]) -> None:
|
||||
|
||||
:rtype: object
|
||||
"""
|
||||
if action.action_name in action_registry:
|
||||
raise ValueError(f"Action {action.action_name} already registered")
|
||||
action_registry[str(action.action_name)] = action
|
||||
if action.name in action_registry:
|
||||
raise ValueError(f"Action {action.name} already registered")
|
||||
action_registry[str(action.name)] = action
|
||||
|
||||
def get_action_by_name(name: str) -> Type[Action]:
|
||||
if name not in action_registry:
|
||||
|
||||
@@ -10,10 +10,6 @@ from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
|
||||
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):
|
||||
# return items
|
||||
|
||||
@@ -23,7 +19,10 @@ class PromptTransformer(Transformer):
|
||||
|
||||
def item(self, items: List[Token]):
|
||||
for item in items:
|
||||
if isinstance(item, (Action, PromptSegment)):
|
||||
if isinstance(item, Action):
|
||||
return item
|
||||
|
||||
if isinstance(item, PromptSegment):
|
||||
return item
|
||||
|
||||
if item.type == "WORD":
|
||||
@@ -37,7 +36,7 @@ class PromptTransformer(Transformer):
|
||||
elif item.type == "embedding":
|
||||
return build_prompt_segment(item, self.tokenizer)
|
||||
elif item.type == "function":
|
||||
raise Exception("Unexpected type in prompt transformer: function. Please report this issue on GitHub.")
|
||||
return item
|
||||
else:
|
||||
raise Exception("Unknown item type: " + str(item.type))
|
||||
|
||||
|
||||
+3
-6
@@ -24,13 +24,10 @@ def build_prompt_segment(text: str, tokenizer: SDTokenizer) -> PromptSegment:
|
||||
if embedding is None:
|
||||
print(f"warning, embedding:{embedding_name} does not exist, ignoring")
|
||||
else:
|
||||
if embedding.shape[1] != tokenizer.embedding_size:
|
||||
print(f"warning, embedding:{embedding_name} has size {embedding.shape[1]}, expected {tokenizer.embedding_size}, ignoring")
|
||||
if len(embedding.shape) == 1:
|
||||
tokens.append(embedding)
|
||||
else:
|
||||
if len(embedding.shape) == 1:
|
||||
tokens.append(embedding)
|
||||
else:
|
||||
tokens.extend(embedding)
|
||||
tokens.extend(embedding)
|
||||
|
||||
if leftover != "":
|
||||
word = leftover
|
||||
|
||||
@@ -1,83 +0,0 @@
|
||||
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
|
||||
Reference in New Issue
Block a user