From 8773c8f249341510672e198067295191acd0dec6 Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Mon, 9 Feb 2026 16:22:03 -0800 Subject: [PATCH] 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. --- kikotools/tools/empty_latent_batch/node.py | 10 +++++----- tests/unit/tools/test_empty_latent_batch.py | 19 +++++++++++-------- 2 files changed, 16 insertions(+), 13 deletions(-) diff --git a/kikotools/tools/empty_latent_batch/node.py b/kikotools/tools/empty_latent_batch/node.py index 99a0b6e..82411a9 100644 --- a/kikotools/tools/empty_latent_batch/node.py +++ b/kikotools/tools/empty_latent_batch/node.py @@ -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 diff --git a/tests/unit/tools/test_empty_latent_batch.py b/tests/unit/tools/test_empty_latent_batch.py index 1692fc6..5187804 100644 --- a/tests/unit/tools/test_empty_latent_batch.py +++ b/tests/unit/tools/test_empty_latent_batch.py @@ -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)