Files
aszc-dev-ComfyUI-CoreMLSuite/tests/test_chunks.py
T
2024-06-28 15:52:53 +02:00

64 lines
1.8 KiB
Python

import pytest
import torch
from comfy.model_management import get_torch_device
from coreml_suite.latents import chunk_batch, merge_chunks
from coreml_suite.controlnet import chunk_control
@pytest.mark.parametrize("batch_size", [2, 4, 5, 9])
def test_batch_chunking(batch_size):
latent_image = torch.randn(batch_size, 4, 64, 64).to(get_torch_device())
target_shape = (4, 4, 64, 64)
chunked = chunk_batch(latent_image, target_shape)
for chunk in chunked:
assert chunk.shape == target_shape
if batch_size % target_shape[0] != 0:
assert chunked[-1][batch_size % target_shape[0] :].sum() == 0
@pytest.mark.parametrize("batch_size", [2, 4, 5, 9])
def test_merge_chunks(batch_size):
input_tensor = torch.randn(batch_size, 4, 64, 64).to(get_torch_device())
target_shape = (4, 4, 64, 64)
chunked = chunk_batch(input_tensor, target_shape)
merged = merge_chunks(chunked, input_tensor.shape)
assert merged.shape == input_tensor.shape
assert torch.equal(input_tensor, merged)
def test_chunking_controlnet():
cn = {
"output": [
torch.randn(4, 4, 64, 64).to(get_torch_device()),
torch.randn(4, 4, 128, 128).to(get_torch_device()),
],
"middle": [
torch.randn(4, 4, 256, 256).to(get_torch_device()),
],
}
target_batch_size = 2
num_chunks = cn["output"][0].shape[0] // target_batch_size
chunked = chunk_control(cn, num_chunks)
for chunk in chunked:
assert chunk["output"][0].shape == (target_batch_size, 4, 64, 64)
assert chunk["output"][1].shape == (target_batch_size, 4, 128, 128)
assert chunk["middle"][0].shape == (target_batch_size, 4, 256, 256)
def test_chunking_no_control():
cn = None
num_chunks = 2
chunked = chunk_control(cn, num_chunks)
assert chunked == [None, None]