Weight syntax: - Grammar: `(arg:SIGNED_NUMBER)` parses to a WeightedGroup - emph(text|w) is a function-style alias for the same - Nested weights multiply: ((cat:1.2):0.5) -> 0.6, matching A1111 - Weights are emitted in the (token, weight) tuples that comfy's stock encode_token_weights already applies post-transformer To make per-position weight indexing work, rows are now always exactly max_length entries: a multi-slot Action emits as one (action, w) entry followed by (ACTION_CONTINUATION, w) placeholders that process_tokens drops before handing to super().process_tokens. This replaces the previous variable-length-row + padding-math. PromptLang Inspect: - New node: CLIP + text -> STRING report of per-slot weight, L2 norm, and top-k nearest vocab tokens for the resolved embedding - lib/inspect.py uses the wrapper's stored attr name (cond.clip / tok.clip) rather than hardcoding 'clip_l', and uses the tokenizer's inv_vocab for decode Fix: parse_numeric_arg now rejects WeightedGroup (was only Action). SegOrAction widened to include WeightedGroup.
66 lines
2.4 KiB
Python
66 lines
2.4 KiB
Python
from typing import List
|
|
|
|
from lark import Token, Transformer
|
|
|
|
from comfy.sd1_clip import SDTokenizer
|
|
|
|
from ..actions.action_utils import parse_numeric_arg
|
|
from ..actions.base import Action, ActionArity
|
|
from ..actions.weighted import WeightedGroup
|
|
from .prompt_segment import PromptSegment
|
|
from .registration import get_action_by_name
|
|
from .utils import build_prompt_segment
|
|
|
|
|
|
class PromptTransformer(Transformer):
|
|
"""Maps the Lark parse tree into a flat list of PromptSegments and Actions."""
|
|
|
|
def __init__(self, tokenizer: SDTokenizer):
|
|
super().__init__()
|
|
self.tokenizer = tokenizer
|
|
|
|
def item(self, items: List[Token]):
|
|
for item in items:
|
|
if isinstance(item, (Action, PromptSegment, WeightedGroup)):
|
|
return item
|
|
|
|
if item.type == "WORD":
|
|
return build_prompt_segment(str(item), self.tokenizer)
|
|
if item.type == "QUOTED_STRING":
|
|
# Strip surrounding quotes, unescape \" and \'.
|
|
unquoted = item[1:-1]
|
|
unescaped = unquoted.replace('\\"', '"').replace("\\'", "'")
|
|
return build_prompt_segment(unescaped, self.tokenizer)
|
|
raise ValueError(f"Unknown item type: {item.type}")
|
|
|
|
def arg(self, items):
|
|
return items
|
|
|
|
def weighted(self, items):
|
|
arg_items, weight_token = items
|
|
return WeightedGroup(arg_items, float(weight_token))
|
|
|
|
def embedding(self, items):
|
|
return build_prompt_segment(
|
|
f"{self.tokenizer.embedding_identifier}{items[0]}",
|
|
self.tokenizer,
|
|
)
|
|
|
|
def generic_function(self, items):
|
|
# `emph(text|w)` is sugar for `(text:w)`; handled here so it doesn't need
|
|
# to fit the Action ABC (it changes weights, not embeddings).
|
|
if str(items[0]) == "emph":
|
|
if len(items) != 3:
|
|
raise ValueError("emph expects exactly two arguments: emph(text|weight)")
|
|
weight = parse_numeric_arg(items[2], action_name="emph", role="weight", cast=float)
|
|
return WeightedGroup(items[1], weight)
|
|
|
|
action = get_action_by_name(items[0])
|
|
if action.arity == ActionArity.SINGLE:
|
|
if len(items) != 2:
|
|
raise ValueError(f"Action {action.action_name} expects exactly one argument")
|
|
return action(items[1])
|
|
if action.arity == ActionArity.MULTI:
|
|
return action(items[1:])
|
|
raise ValueError(f"Unknown action arity: {action.arity}")
|