From c69281e79565b608edfd02efe18a9166ed7a47c6 Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Mon, 9 Feb 2026 17:11:01 -0800 Subject: [PATCH] fix(batch_list_converter): preserve all latent dict keys during split/join Previously split_latent_batch and join_latent_batch only handled the "samples" key, silently dropping metadata like noise_mask or batch_index. This broke inpaint and masked workflows after a round-trip conversion. --- kikotools/tools/batch_list_converter/logic.py | 34 ++++++++++++-- tests/unit/tools/test_batch_list_converter.py | 46 +++++++++++++++++++ 2 files changed, 76 insertions(+), 4 deletions(-) diff --git a/kikotools/tools/batch_list_converter/logic.py b/kikotools/tools/batch_list_converter/logic.py index 5b28993..61539fe 100644 --- a/kikotools/tools/batch_list_converter/logic.py +++ b/kikotools/tools/batch_list_converter/logic.py @@ -17,13 +17,39 @@ def join_image_batch(image_list: List[torch.Tensor]) -> torch.Tensor: def split_latent_batch( latent: Dict[str, torch.Tensor], ) -> List[Dict[str, torch.Tensor]]: - """Split latent dict into list of single-item latent dicts.""" + """Split latent dict into list of single-item latent dicts. + + Preserves all keys (e.g. noise_mask, batch_index). Tensor values whose + first dimension matches the batch size of ``samples`` are sliced along + dim-0; all other values are copied as-is to every item. + """ samples = latent["samples"] - return [{"samples": samples[i : i + 1]} for i in range(samples.shape[0])] + batch_size = samples.shape[0] + result: List[Dict[str, torch.Tensor]] = [] + for i in range(batch_size): + item: Dict[str, torch.Tensor] = {} + for key, value in latent.items(): + if isinstance(value, torch.Tensor) and value.shape[0] == batch_size: + item[key] = value[i : i + 1] + else: + item[key] = value + result.append(item) + return result def join_latent_batch( latent_list: List[Dict[str, torch.Tensor]], ) -> Dict[str, torch.Tensor]: - """Join list of latent dicts into single batched latent dict.""" - return {"samples": torch.cat([lat["samples"] for lat in latent_list], dim=0)} + """Join list of latent dicts into single batched latent dict. + + Tensor values that were sliced during split are concatenated along dim-0. + Non-tensor values are taken from the first item. + """ + result: Dict[str, torch.Tensor] = {} + first = latent_list[0] + for key in first: + if isinstance(first[key], torch.Tensor): + result[key] = torch.cat([lat[key] for lat in latent_list], dim=0) + else: + result[key] = first[key] + return result diff --git a/tests/unit/tools/test_batch_list_converter.py b/tests/unit/tools/test_batch_list_converter.py index 5c529e7..3eba35d 100644 --- a/tests/unit/tools/test_batch_list_converter.py +++ b/tests/unit/tools/test_batch_list_converter.py @@ -107,6 +107,52 @@ class TestBatchListConverterLogic: reconstructed = join_latent_batch(split_latent_batch(original)) assert torch.equal(original["samples"], reconstructed["samples"]) + def test_split_latent_preserves_noise_mask(self): + """noise_mask is sliced alongside samples.""" + latent = { + "samples": torch.rand(3, 4, 32, 32), + "noise_mask": torch.rand(3, 1, 32, 32), + } + result = split_latent_batch(latent) + assert len(result) == 3 + for i, item in enumerate(result): + assert "noise_mask" in item + assert item["noise_mask"].shape == (1, 1, 32, 32) + assert torch.equal(item["noise_mask"], latent["noise_mask"][i : i + 1]) + + def test_join_latent_preserves_noise_mask(self): + """noise_mask is concatenated alongside samples.""" + latent_list = [ + { + "samples": torch.rand(1, 4, 32, 32), + "noise_mask": torch.rand(1, 1, 32, 32), + } + for _ in range(3) + ] + result = join_latent_batch(latent_list) + assert "noise_mask" in result + assert result["noise_mask"].shape == (3, 1, 32, 32) + + def test_latent_roundtrip_with_extra_keys(self): + """Round-trip preserves all tensor keys.""" + original = { + "samples": torch.rand(4, 4, 64, 64), + "noise_mask": torch.rand(4, 1, 64, 64), + } + reconstructed = join_latent_batch(split_latent_batch(original)) + assert torch.equal(original["samples"], reconstructed["samples"]) + assert torch.equal(original["noise_mask"], reconstructed["noise_mask"]) + + def test_split_latent_copies_non_tensor_values(self): + """Non-tensor metadata is copied to each item.""" + latent = { + "samples": torch.rand(2, 4, 32, 32), + "some_flag": "preserve_me", + } + result = split_latent_batch(latent) + for item in result: + assert item["some_flag"] == "preserve_me" + class TestBatchListConverterNodes: """Test ComfyUI node classes."""