Add proj, reject, renorm, noise, nearest actions and lerp alias
proj(a|b) - project a onto the (mean, unit) direction of b reject(a|b) - a minus that projection (orthogonal component) renorm(a|ref)- rescale a so each token's L2 norm matches ref's mean L2 noise(a|std) - a + N(0, std) nearest(e|k) - snap mean(e) to its k nearest vocab tokens (cosine sim) lerp(a|b|t) - alias for avg(a|b|t) All follow the existing MultiArgAction pattern; nearest reuses the embedding-weight cosine machinery the Inspect node uses. Nine new tests cover the math invariants (proj+reject reconstructs input, reject is orthogonal to b, renorm matches ref norm, noise(_, 0) is identity, nearest of a single token returns that token). Relax MultiArgAction/SingleArgAction __init__ type hints to List since SegOrAction now includes WeightedGroup and the Union isn't importable in base.py without a cycle.
This commit is contained in:
@@ -56,16 +56,23 @@ Arguments inside a function are separated by `|`. Each arg can itself be plain t
|
||||
| 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> |
|
||||
| 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> |
|
||||
| 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> |
|
||||
| Nearest Vocab | nearest | Snaps a computed vector to the k nearest real vocabulary tokens (by cosine similarity), returning their embeddings concatenated. The input is mean-pooled before lookup. | <ul><li>nearest(sum(diff(king\|man)\|woman))</li><li>nearest(sum(red\|blue)\|3)</li></ul> |
|
||||
| Negate | neg | Negates the provided segments or actions. | <ul><li>neg(cat)</li><li>sum(king\|neg(man)\|women)</li></ul> |
|
||||
| Noise | noise | Adds Gaussian noise (mean 0, given std) to the embeddings of the first argument. | <ul><li>A noise(cat\|0.05) on a sunny day</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> |
|
||||
| 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> |
|
||||
| Project | proj | Projects the first argument onto the direction of the second (mean, unit-normalized). | <ul><li>proj(king\|gender)</li><li>diff(style\|proj(style\|photorealistic))</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> |
|
||||
| Reject | reject | Removes the component of the first argument along the direction of the second (a - proj(a\|b)). | <ul><li>reject(anime girl\|anime)</li></ul> |
|
||||
| Renormalize | renorm | Rescales the first argument so each token's L2 norm matches the (mean) L2 norm of the reference. | <ul><li>renorm(sum(king\|neg(man)\|woman)\|queen)</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> |
|
||||
| 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> |
|
||||
| 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> |
|
||||
| Sum | sum | Adds the embeddings of the provided segments or actions. | <ul><li>A happy sum(cat\|dog\|shark)</li></ul> |
|
||||
|
||||
`lerp(a|b|t)` is also accepted as an alias for `avg(a|b|t)`.
|
||||
|
||||
Regenerate the table with `python tools/build_docs.py`.
|
||||
|
||||
## Development
|
||||
|
||||
@@ -2,11 +2,15 @@ from ..parser.registration import register_action
|
||||
from .avg import AverageAction
|
||||
from .diff import DiffAction
|
||||
from .mult import MultiplyAction
|
||||
from .nearest import NearestAction
|
||||
from .neg import NegAction
|
||||
from .noise import NoiseAction
|
||||
from .norm import NormAction
|
||||
from .pos_scale import PosScaleAction
|
||||
from .post_pos import PostPosAction
|
||||
from .project import ProjectAction, RejectAction
|
||||
from .rand import RandAction
|
||||
from .renorm import RenormAction
|
||||
from .scale_dims import ScaleDims
|
||||
from .set_dims import SetDims
|
||||
from .slerp import SlerpAction
|
||||
@@ -16,14 +20,30 @@ for _action in [
|
||||
AverageAction,
|
||||
DiffAction,
|
||||
MultiplyAction,
|
||||
NearestAction,
|
||||
NegAction,
|
||||
NoiseAction,
|
||||
NormAction,
|
||||
PosScaleAction,
|
||||
PostPosAction,
|
||||
ProjectAction,
|
||||
RandAction,
|
||||
RejectAction,
|
||||
RenormAction,
|
||||
ScaleDims,
|
||||
SetDims,
|
||||
SlerpAction,
|
||||
SumAction,
|
||||
]:
|
||||
register_action(_action)
|
||||
|
||||
|
||||
class _LerpAlias(AverageAction):
|
||||
"""`lerp(a|b|t)` is sugar for `avg(a|b|t)`."""
|
||||
|
||||
display_name = "Lerp"
|
||||
action_name = "lerp"
|
||||
usage_examples = ["lerp(cat|dog|0.5)"]
|
||||
|
||||
|
||||
register_action(_LerpAlias)
|
||||
|
||||
+2
-4
@@ -6,8 +6,6 @@ from typing import List, Optional, Tuple, Union
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
|
||||
from ..parser.prompt_segment import PromptSegment
|
||||
|
||||
|
||||
class ActionArity(Enum):
|
||||
NONE = 0
|
||||
@@ -57,7 +55,7 @@ class Action(ABC):
|
||||
class SingleArgAction(Action, ABC):
|
||||
arity = ActionArity.SINGLE
|
||||
|
||||
def __init__(self, arg: List[Union[PromptSegment, "Action"]]):
|
||||
def __init__(self, arg: List):
|
||||
self.arg = arg
|
||||
|
||||
def __repr__(self) -> str:
|
||||
@@ -67,7 +65,7 @@ class SingleArgAction(Action, ABC):
|
||||
class MultiArgAction(Action, ABC):
|
||||
arity = ActionArity.MULTI
|
||||
|
||||
def __init__(self, args: List[List[Union[PromptSegment, "Action"]]]):
|
||||
def __init__(self, args: List[List]):
|
||||
self.all_args = args
|
||||
|
||||
def __repr__(self) -> str:
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
from typing import List
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
|
||||
from .action_utils import concat_embeddings, parse_numeric_arg
|
||||
from .base import MultiArgAction
|
||||
from .types import SegOrAction
|
||||
|
||||
|
||||
class NearestAction(MultiArgAction):
|
||||
grammar = 'nearest(" arg ("|" arg)? ")"'
|
||||
|
||||
display_name = "Nearest Vocab"
|
||||
action_name = "nearest"
|
||||
description = (
|
||||
"Snaps a computed vector to the k nearest real vocabulary tokens (by cosine similarity), "
|
||||
"returning their embeddings concatenated. The input is mean-pooled before lookup."
|
||||
)
|
||||
usage_examples = [
|
||||
"nearest(sum(diff(king|man)|woman))",
|
||||
"nearest(sum(red|blue)|3)",
|
||||
]
|
||||
|
||||
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
||||
super().__init__(args)
|
||||
if len(args) not in (1, 2):
|
||||
raise ValueError("nearest expects one or two arguments: nearest(expr) or nearest(expr|k)")
|
||||
self.expr_arg = args[0]
|
||||
self.k = parse_numeric_arg(args[1], action_name="nearest", role="k", cast=int) if len(args) == 2 else 1
|
||||
|
||||
def token_length(self) -> int:
|
||||
return self.k
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
||||
weight = embedding_module.weight.to(torch.float32)
|
||||
weight_norm = torch.nn.functional.normalize(weight, dim=-1)
|
||||
|
||||
expr = concat_embeddings(self.expr_arg, embedding_module).to(torch.float32)
|
||||
query = torch.nn.functional.normalize(expr.mean(dim=1), dim=-1)
|
||||
|
||||
sims = query @ weight_norm.T
|
||||
top_ids = sims.topk(self.k, dim=-1).indices.squeeze(0)
|
||||
return weight[top_ids].unsqueeze(0)
|
||||
@@ -0,0 +1,34 @@
|
||||
from typing import List
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
|
||||
from .action_utils import concat_embeddings, get_total_length, parse_numeric_arg
|
||||
from .base import MultiArgAction
|
||||
from .types import SegOrAction
|
||||
|
||||
|
||||
class NoiseAction(MultiArgAction):
|
||||
grammar = 'noise(" arg "|" arg ")"'
|
||||
|
||||
display_name = "Noise"
|
||||
action_name = "noise"
|
||||
description = "Adds Gaussian noise (mean 0, given std) to the embeddings of the first argument."
|
||||
usage_examples = [
|
||||
"A noise(cat|0.05) on a sunny day",
|
||||
]
|
||||
|
||||
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
||||
super().__init__(args)
|
||||
if len(args) != 2:
|
||||
raise ValueError("noise expects exactly two arguments: noise(a|std)")
|
||||
self.a_arg = args[0]
|
||||
self.std = parse_numeric_arg(args[1], action_name="noise", role="std", cast=float)
|
||||
|
||||
def token_length(self) -> int:
|
||||
return get_total_length(self.a_arg)
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
||||
a = concat_embeddings(self.a_arg, embedding_module)
|
||||
return a + torch.randn_like(a) * self.std
|
||||
@@ -0,0 +1,72 @@
|
||||
from typing import List
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
|
||||
from .action_utils import concat_embeddings, get_total_length
|
||||
from .base import MultiArgAction
|
||||
from .types import SegOrAction
|
||||
|
||||
|
||||
def _direction(args: List[SegOrAction], embedding_module: Embedding) -> Tensor:
|
||||
"""Mean unit direction of an arg's embeddings: [1, 1, hidden]."""
|
||||
emb = concat_embeddings(args, embedding_module)
|
||||
mean = emb.mean(dim=1, keepdim=True)
|
||||
return torch.nn.functional.normalize(mean, dim=-1)
|
||||
|
||||
|
||||
def _project(a: Tensor, b_hat: Tensor) -> Tensor:
|
||||
coeff = (a * b_hat).sum(dim=-1, keepdim=True)
|
||||
return coeff * b_hat
|
||||
|
||||
|
||||
class ProjectAction(MultiArgAction):
|
||||
grammar = 'proj(" arg "|" arg ")"'
|
||||
|
||||
display_name = "Project"
|
||||
action_name = "proj"
|
||||
description = "Projects the first argument onto the direction of the second (mean, unit-normalized)."
|
||||
usage_examples = [
|
||||
"proj(king|gender)",
|
||||
"diff(style|proj(style|photorealistic))",
|
||||
]
|
||||
|
||||
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
||||
super().__init__(args)
|
||||
if len(args) != 2:
|
||||
raise ValueError("proj expects exactly two arguments: proj(a|b)")
|
||||
self.a_arg = args[0]
|
||||
self.b_arg = args[1]
|
||||
|
||||
def token_length(self) -> int:
|
||||
return get_total_length(self.a_arg)
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
||||
a = concat_embeddings(self.a_arg, embedding_module)
|
||||
return _project(a, _direction(self.b_arg, embedding_module))
|
||||
|
||||
|
||||
class RejectAction(MultiArgAction):
|
||||
grammar = 'reject(" arg "|" arg ")"'
|
||||
|
||||
display_name = "Reject"
|
||||
action_name = "reject"
|
||||
description = "Removes the component of the first argument along the direction of the second (a - proj(a|b))."
|
||||
usage_examples = [
|
||||
"reject(anime girl|anime)",
|
||||
]
|
||||
|
||||
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
||||
super().__init__(args)
|
||||
if len(args) != 2:
|
||||
raise ValueError("reject expects exactly two arguments: reject(a|b)")
|
||||
self.a_arg = args[0]
|
||||
self.b_arg = args[1]
|
||||
|
||||
def token_length(self) -> int:
|
||||
return get_total_length(self.a_arg)
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
||||
a = concat_embeddings(self.a_arg, embedding_module)
|
||||
return a - _project(a, _direction(self.b_arg, embedding_module))
|
||||
@@ -0,0 +1,37 @@
|
||||
from typing import List
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
|
||||
from .action_utils import concat_embeddings, get_total_length
|
||||
from .base import MultiArgAction
|
||||
from .types import SegOrAction
|
||||
|
||||
|
||||
class RenormAction(MultiArgAction):
|
||||
grammar = 'renorm(" arg "|" arg ")"'
|
||||
|
||||
display_name = "Renormalize"
|
||||
action_name = "renorm"
|
||||
description = "Rescales the first argument so each token's L2 norm matches the (mean) L2 norm of the reference."
|
||||
usage_examples = [
|
||||
"renorm(sum(king|neg(man)|woman)|queen)",
|
||||
]
|
||||
|
||||
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
||||
super().__init__(args)
|
||||
if len(args) != 2:
|
||||
raise ValueError("renorm expects exactly two arguments: renorm(a|ref)")
|
||||
self.a_arg = args[0]
|
||||
self.ref_arg = args[1]
|
||||
|
||||
def token_length(self) -> int:
|
||||
return get_total_length(self.a_arg)
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
||||
a = concat_embeddings(self.a_arg, embedding_module)
|
||||
ref = concat_embeddings(self.ref_arg, embedding_module)
|
||||
a_norm = torch.norm(a, dim=-1, keepdim=True).clamp(min=1e-8)
|
||||
ref_norm = torch.norm(ref, dim=-1, keepdim=True).mean()
|
||||
return a * (ref_norm / a_norm)
|
||||
@@ -0,0 +1,96 @@
|
||||
import pytest
|
||||
|
||||
torch = pytest.importorskip("torch")
|
||||
|
||||
from KepPromptLang.lib.actions.nearest import NearestAction
|
||||
from KepPromptLang.lib.actions.noise import NoiseAction
|
||||
from KepPromptLang.lib.actions.project import ProjectAction, RejectAction
|
||||
from KepPromptLang.lib.actions.renorm import RenormAction
|
||||
from KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||
from KepPromptLang.lib.parser.registration import get_action_by_name
|
||||
|
||||
EMBED_DIM = 4
|
||||
VOCAB = 50
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def embedding():
|
||||
torch.manual_seed(0)
|
||||
return torch.nn.Embedding(VOCAB, EMBED_DIM)
|
||||
|
||||
|
||||
def seg(*token_ids):
|
||||
return PromptSegment(text="x", tokens=list(token_ids))
|
||||
|
||||
|
||||
def test_proj_plus_reject_reconstructs_input(embedding):
|
||||
a, b = seg(1, 2), seg(3)
|
||||
proj = ProjectAction([[a], [b]]).get_result(embedding)
|
||||
rej = RejectAction([[a], [b]]).get_result(embedding)
|
||||
a_emb = embedding(torch.LongTensor([[1, 2]]))
|
||||
assert torch.allclose(proj + rej, a_emb, atol=1e-5)
|
||||
|
||||
|
||||
def test_reject_is_orthogonal_to_b(embedding):
|
||||
a, b = seg(1, 2), seg(3)
|
||||
rej = RejectAction([[a], [b]]).get_result(embedding)
|
||||
b_dir = torch.nn.functional.normalize(
|
||||
embedding(torch.LongTensor([[3]])).mean(dim=1, keepdim=True), dim=-1
|
||||
)
|
||||
dots = (rej * b_dir).sum(dim=-1)
|
||||
assert torch.allclose(dots, torch.zeros_like(dots), atol=1e-5)
|
||||
|
||||
|
||||
def test_renorm_matches_ref_norm(embedding):
|
||||
a, ref = seg(1, 2), seg(3)
|
||||
out = RenormAction([[a], [ref]]).get_result(embedding)
|
||||
ref_norm = torch.norm(embedding(torch.LongTensor([[3]])), dim=-1).mean()
|
||||
out_norms = torch.norm(out, dim=-1)
|
||||
assert torch.allclose(out_norms, ref_norm.expand_as(out_norms), atol=1e-5)
|
||||
|
||||
|
||||
def test_noise_shape_and_mean(embedding):
|
||||
std = PromptSegment(text="0.01", tokens=[1])
|
||||
out = NoiseAction([[seg(1, 2)], [std]]).get_result(embedding)
|
||||
base = embedding(torch.LongTensor([[1, 2]]))
|
||||
assert out.shape == base.shape
|
||||
# Perturbation magnitude bounded (5σ with margin); std=0.01, EMBED_DIM=4.
|
||||
assert (out - base).abs().max() < 0.2
|
||||
|
||||
|
||||
def test_noise_zero_std_is_identity(embedding):
|
||||
std = PromptSegment(text="0.0", tokens=[1])
|
||||
out = NoiseAction([[seg(1, 2)], [std]]).get_result(embedding)
|
||||
base = embedding(torch.LongTensor([[1, 2]]))
|
||||
assert torch.allclose(out, base)
|
||||
|
||||
|
||||
def test_nearest_returns_exact_token_for_that_token(embedding):
|
||||
out = NearestAction([[seg(7)]]).get_result(embedding)
|
||||
assert out.shape == (1, 1, EMBED_DIM)
|
||||
assert torch.allclose(out[0, 0], embedding.weight[7])
|
||||
|
||||
|
||||
def test_nearest_k_tokens(embedding):
|
||||
k = PromptSegment(text="3", tokens=[1])
|
||||
action = NearestAction([[seg(7)], [k]])
|
||||
assert action.token_length() == 3
|
||||
out = action.get_result(embedding)
|
||||
assert out.shape == (1, 3, EMBED_DIM)
|
||||
# First match should be the token itself.
|
||||
assert torch.allclose(out[0, 0], embedding.weight[7])
|
||||
|
||||
|
||||
def test_lerp_is_registered_as_avg_alias():
|
||||
lerp_cls = get_action_by_name("lerp")
|
||||
avg_cls = get_action_by_name("avg")
|
||||
assert issubclass(lerp_cls, avg_cls)
|
||||
|
||||
|
||||
def test_token_lengths():
|
||||
a, b = seg(1, 2, 3), seg(4)
|
||||
assert ProjectAction([[a], [b]]).token_length() == 3
|
||||
assert RejectAction([[a], [b]]).token_length() == 3
|
||||
assert RenormAction([[a], [b]]).token_length() == 3
|
||||
std = PromptSegment(text="0.1", tokens=[1])
|
||||
assert NoiseAction([[a], [std]]).token_length() == 3
|
||||
Reference in New Issue
Block a user