feat(latent): add batch_size as 4th output from EmptyLatentBatchNode

Exposes batch_size as an output so downstream nodes can reference it
directly without needing a separate input.
This commit is contained in:
Vito Sansevero
2026-02-09 16:22:03 -08:00
parent 0a1cbe4990
commit 8773c8f249
2 changed files with 16 additions and 13 deletions
+5 -5
View File
@@ -93,14 +93,14 @@ class EmptyLatentBatchNode(ComfyAssetsBaseNode):
}
}
RETURN_TYPES = ("LATENT", "INT", "INT")
RETURN_NAMES = ("latent", "width", "height")
RETURN_TYPES = ("LATENT", "INT", "INT", "INT")
RETURN_NAMES = ("latent", "width", "height", "batch_size")
FUNCTION = "create_empty_latent"
CATEGORY = "🫶 ComfyAssets/📦 Latents"
def create_empty_latent(
self, preset: str, width: int, height: int, batch_size: int
) -> Tuple[Dict[str, torch.Tensor], int, int]:
) -> Tuple[Dict[str, torch.Tensor], int, int, int]:
"""
Create empty latent tensor with specified dimensions and batch size.
@@ -111,7 +111,7 @@ class EmptyLatentBatchNode(ComfyAssetsBaseNode):
batch_size: Number of latents in the batch
Returns:
Tuple containing (latent dictionary with 'samples' tensor, width, height)
Tuple containing (latent dict, width, height, batch_size)
"""
try:
# Extract original preset name from formatted string if needed
@@ -160,7 +160,7 @@ class EmptyLatentBatchNode(ComfyAssetsBaseNode):
f"(pixel dims: {final_width}×{final_height})"
)
return (latent_dict, final_width, final_height)
return (latent_dict, final_width, final_height, batch_size)
except Exception as e:
# Handle any unexpected errors gracefully
+11 -8
View File
@@ -131,8 +131,8 @@ class TestEmptyLatentBatchNode:
def test_node_attributes(self):
"""Test node class attributes."""
assert EmptyLatentBatchNode.RETURN_TYPES == ("LATENT", "INT", "INT")
assert EmptyLatentBatchNode.RETURN_NAMES == ("latent", "width", "height")
assert EmptyLatentBatchNode.RETURN_TYPES == ("LATENT", "INT", "INT", "INT")
assert EmptyLatentBatchNode.RETURN_NAMES == ("latent", "width", "height", "batch_size")
assert EmptyLatentBatchNode.FUNCTION == "create_empty_latent"
assert EmptyLatentBatchNode.CATEGORY == "🫶 ComfyAssets/📦 Latents"
@@ -141,13 +141,14 @@ class TestEmptyLatentBatchNode:
result = self.node.create_empty_latent("custom", 512, 512, 1)
assert isinstance(result, tuple)
assert len(result) == 3 # Now returns (latent, width, height)
assert len(result) == 4 # Returns (latent, width, height, batch_size)
latent_dict, width, height = result
latent_dict, width, height, batch_size = result
assert isinstance(latent_dict, dict)
assert "samples" in latent_dict
assert width == 512
assert height == 512
assert batch_size == 1
samples = latent_dict["samples"]
assert isinstance(samples, torch.Tensor)
@@ -155,12 +156,13 @@ class TestEmptyLatentBatchNode:
def test_create_empty_latent_with_batch(self):
"""Test empty latent creation with batch size."""
batch_size = 3
result = self.node.create_empty_latent("custom", 1024, 768, batch_size)
input_batch_size = 3
result = self.node.create_empty_latent("custom", 1024, 768, input_batch_size)
latent_dict, width, height = result
latent_dict, width, height, batch_size = result
assert width == 1024
assert height == 768
assert batch_size == 3
samples = latent_dict["samples"]
assert samples.shape == (3, 4, 96, 128) # batch=3, 768/8=96, 1024/8=128
@@ -169,10 +171,11 @@ class TestEmptyLatentBatchNode:
# Input dimensions not divisible by 8
result = self.node.create_empty_latent("custom", 513, 515, 1)
latent_dict, width, height = result
latent_dict, width, height, batch_size = result
# Dimensions should be rounded UP to nearest multiple of 8
assert width == 520 # 513 -> 520
assert height == 520 # 515 -> 520
assert batch_size == 1
samples = latent_dict["samples"]
# Should be adjusted to 520x520 -> 65x65 latent
assert samples.shape == (1, 4, 65, 65)