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.
This commit is contained in:
Vito Sansevero
2026-02-09 17:11:01 -08:00
parent 3bf8391e88
commit c69281e795
2 changed files with 76 additions and 4 deletions
+30 -4
View File
@@ -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
@@ -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."""