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:
Claude
2026-04-12 06:02:00 +00:00
parent c6910ab775
commit 3dddbe1671
8 changed files with 313 additions and 4 deletions
+7
View File
@@ -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
+20
View File
@@ -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
View File
@@ -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:
+45
View File
@@ -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)
+34
View File
@@ -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
+72
View File
@@ -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))
+37
View File
@@ -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)
+96
View File
@@ -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