feat(blocks): key normalization and weight-string parsing

This commit is contained in:
larsupb
2026-07-19 00:34:19 +02:00
parent 3c73c551f6
commit db9f94fb5d
2 changed files with 67 additions and 0 deletions
+36
View File
@@ -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
+31
View File
@@ -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) == []