From db9f94fb5d8d73dee5fdf0818626eac5eeae75fb Mon Sep 17 00:00:00 2001 From: larsupb Date: Sun, 19 Jul 2026 00:34:19 +0200 Subject: [PATCH] feat(blocks): key normalization and weight-string parsing --- src/blocks.py | 36 ++++++++++++++++++++++++++++++++++++ tests/test_blocks.py | 31 +++++++++++++++++++++++++++++++ 2 files changed, 67 insertions(+) create mode 100644 src/blocks.py create mode 100644 tests/test_blocks.py diff --git a/src/blocks.py b/src/blocks.py new file mode 100644 index 0000000..6ce60e9 --- /dev/null +++ b/src/blocks.py @@ -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 diff --git a/tests/test_blocks.py b/tests/test_blocks.py new file mode 100644 index 0000000..514eb9c --- /dev/null +++ b/tests/test_blocks.py @@ -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) == []