Files
M1kep-KepPromptLang/lib/actions/base.py
T
Claude 3dddbe1671 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.
2026-04-12 06:02:00 +00:00

74 lines
1.8 KiB
Python

from abc import ABC, abstractmethod
from dataclasses import dataclass
from enum import Enum
from typing import List, Optional, Tuple, Union
from torch import Tensor
from torch.nn import Embedding
class ActionArity(Enum):
NONE = 0
SINGLE = 1
MULTI = 2
@dataclass
class PostModifiers:
"""Optional position-embedding tweaks an action can request for its token range.
`start_idx` / `end_idx` are filled in by the encoder once the action's position in
the final token stream is known.
"""
position_embed_scale: Optional[float] = None
bypass_pos_embed: bool = False
start_idx: int = 0
end_idx: int = 0
ActionResult = Union[Tensor, Tuple[Tensor, PostModifiers]]
# Tokenizer placeholder for the 2nd..Nth slots of a multi-token Action, so each
# row stays exactly max_length entries (required for comfy's per-position weight
# indexing). process_tokens drops these; the Action's tensor fills the slots.
ACTION_CONTINUATION = object()
class Action(ABC):
arity: ActionArity = ActionArity.NONE
display_name: str = ""
action_name: str = ""
description: str = ""
grammar: str = ""
usage_examples: List[str] = []
@abstractmethod
def __init__(self, *args, **kwargs) -> None: ...
@abstractmethod
def token_length(self) -> int: ...
@abstractmethod
def get_result(self, embedding_module: Embedding) -> ActionResult: ...
class SingleArgAction(Action, ABC):
arity = ActionArity.SINGLE
def __init__(self, arg: List):
self.arg = arg
def __repr__(self) -> str:
return f"{self.action_name}({self.arg})"
class MultiArgAction(Action, ABC):
arity = ActionArity.MULTI
def __init__(self, args: List[List]):
self.all_args = args
def __repr__(self) -> str:
joined = " | ".join(str(a) for a in self.all_args)
return f"{self.action_name}({joined})"