feat(blocks): key normalization and weight-string parsing
This commit is contained in:
@@ -0,0 +1,36 @@
|
||||
"""
|
||||
Pure logic for block-wise LoRA weighting (PM Block Selector + model block nodes).
|
||||
|
||||
No ComfyUI or intra-package imports so this module is unit-testable in isolation.
|
||||
"""
|
||||
import logging
|
||||
import re
|
||||
from typing import Any, Dict, Iterable, List, Optional
|
||||
|
||||
|
||||
def normalize_key(key: Any) -> str:
|
||||
"""Return the string form of a LoRA layer key (ComfyUI keys may be tuples)."""
|
||||
if isinstance(key, tuple):
|
||||
return str(key[0])
|
||||
return str(key)
|
||||
|
||||
|
||||
def parse_weight_list(s: Optional[str], default: float = 1.0) -> List[float]:
|
||||
"""Parse a comma-separated per-group weight string into a list of floats.
|
||||
|
||||
Empty or unparseable tokens fall back to ``default``. ``None`` yields ``[]``.
|
||||
"""
|
||||
if s is None:
|
||||
return []
|
||||
out: List[float] = []
|
||||
for tok in str(s).split(","):
|
||||
tok = tok.strip()
|
||||
if tok == "":
|
||||
out.append(default)
|
||||
continue
|
||||
try:
|
||||
out.append(float(tok))
|
||||
except ValueError:
|
||||
logging.warning(f"[PM Blocks] Could not parse weight '{tok}', using {default}")
|
||||
out.append(default)
|
||||
return out
|
||||
@@ -0,0 +1,31 @@
|
||||
from blocks import normalize_key, parse_weight_list
|
||||
|
||||
|
||||
class TestNormalizeKey:
|
||||
def test_string_key_unchanged(self):
|
||||
assert normalize_key("diffusion_model.blocks.0.attn.wq.weight") == \
|
||||
"diffusion_model.blocks.0.attn.wq.weight"
|
||||
|
||||
def test_tuple_key_uses_first_element(self):
|
||||
assert normalize_key(("diffusion_model.single_blocks.0.linear1.weight", (0, 0, 9))) == \
|
||||
"diffusion_model.single_blocks.0.linear1.weight"
|
||||
|
||||
|
||||
class TestParseWeightList:
|
||||
def test_simple_list(self):
|
||||
assert parse_weight_list("1,1,0.8,0.5,0") == [1.0, 1.0, 0.8, 0.5, 0.0]
|
||||
|
||||
def test_whitespace_tolerated(self):
|
||||
assert parse_weight_list(" 1.0 , 0.5 ") == [1.0, 0.5]
|
||||
|
||||
def test_empty_token_uses_default(self):
|
||||
assert parse_weight_list("1,,0.5", default=1.0) == [1.0, 1.0, 0.5]
|
||||
|
||||
def test_empty_string_returns_single_default(self):
|
||||
assert parse_weight_list("", default=1.0) == [1.0]
|
||||
|
||||
def test_invalid_token_uses_default(self):
|
||||
assert parse_weight_list("1,abc,0.5", default=1.0) == [1.0, 1.0, 0.5]
|
||||
|
||||
def test_none_returns_empty(self):
|
||||
assert parse_weight_list(None) == []
|
||||
Reference in New Issue
Block a user