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:
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user