Chunk and pad batches

This commit is contained in:
aszc-dev
2023-10-30 16:17:18 +01:00
parent 6319d2aedb
commit 8a3e9332e1
5 changed files with 118 additions and 55 deletions
+28 -15
View File
@@ -1,20 +1,33 @@
import torch
from torchvision.transforms.functional import resize
from coreml_suite.logger import logger
def reshape_latent_image(latent_image, target_shape):
if latent_image is None:
logger.warning("No latent image provided, using zeros.")
return {"samples": torch.zeros(target_shape)}
def chunk_batch(latent_image, target_shape):
if latent_image.shape == target_shape:
return [latent_image]
if latent_image["samples"].shape == target_shape:
return latent_image
batch_size = latent_image.shape[0]
target_batch_size = target_shape[0]
logger.warning(
"Latent image shape does not match model input shape,"
" resizing to match models expected input shape."
)
resized = resize(latent_image["samples"], target_shape[-2:])
return {"samples": resized}
num_chunks = batch_size // target_batch_size
if num_chunks == 0:
padding = torch.zeros(target_batch_size - batch_size, *target_shape[1:])
return [torch.cat((latent_image, padding), dim=0)]
mod = batch_size % target_batch_size
if mod != 0:
chunks = list(torch.chunk(latent_image[:-mod], num_chunks))
padding = torch.zeros(target_batch_size - mod, *target_shape[1:])
padded = torch.cat((latent_image[-mod:], padding), dim=0)
chunks.append(padded)
return [chunk for chunk in chunks]
chunks = list(torch.chunk(latent_image, num_chunks))
return [chunk for chunk in chunks]
def merge_chunks(chunks, orig_shape):
merged = torch.cat(chunks, dim=0)
if merged.shape == orig_shape:
return merged
return merged[: orig_shape[0]]
+35 -4
View File
@@ -7,6 +7,7 @@ from comfy.latent_formats import SD15
from comfy.model_base import BaseModel
from coreml_suite.controlnet import expand_inputs, extract_residual_kwargs
from coreml_suite.latents import chunk_batch, merge_chunks
def get_model_config():
@@ -38,6 +39,36 @@ class CoreMLModelWrapper(BaseModel):
c_adm=None,
control=None,
transformer_options={},
):
sample_shape = self.diffusion_model.expected_inputs["sample"]["shape"]
chunked_x = chunk_batch(x, sample_shape)
ts = t.chunk(len(chunked_x), dim=0)
chunked_context = c_crossattn.chunk(len(chunked_x), dim=0)
chunked_out = [
self._apply_model(
x, t, c_concat, c_crossattn, c_adm, control, transformer_options
)
for x, t, c_crossattn in zip(chunked_x, ts, chunked_context)
]
merged_out = merge_chunks(chunked_out, x.shape)
return merged_out
def get_dtype(self):
# Hardcoding torch-compatible dtype (used for memory allocation)
return torch.float16
def _apply_model(
self,
x,
t,
c_concat=None,
c_crossattn=None,
c_adm=None,
control=None,
transformer_options={},
):
sample = x.cpu().numpy().astype(np.float16)
@@ -53,11 +84,11 @@ class CoreMLModelWrapper(BaseModel):
}
residual_kwargs = extract_residual_kwargs(self.diffusion_model, control)
model_input_kwargs |= residual_kwargs
model_input_kwargs = expand_inputs(model_input_kwargs)
# model_input_kwargs = expand_inputs(model_input_kwargs)
np_out = self.diffusion_model(**model_input_kwargs)["noise_pred"]
return torch.from_numpy(np_out).to(x.device)
def get_dtype(self):
# Hardcoding torch-compatible dtype (used for memory allocation)
return torch.float16
@property
def expected_inputs(self):
return self.diffusion_model.expected_inputs
+24 -17
View File
@@ -1,5 +1,6 @@
import os
import torch
from coremltools import ComputeUnit
from python_coreml_stable_diffusion.coreml_model import CoreMLModel
@@ -7,7 +8,6 @@ import folder_paths
from comfy.model_management import get_torch_device
from comfy.model_patcher import ModelPatcher
from coreml_suite.logger import logger
from coreml_suite.latents import reshape_latent_image
from nodes import KSampler
from coreml_suite.models import CoreMLModelWrapper, get_model_config
@@ -28,26 +28,33 @@ class CoreMLSampler(KSampler):
CATEGORY = "Core ML Suite"
def sample(
self,
coreml_model,
seed,
steps,
cfg,
sampler_name,
scheduler,
positive,
negative,
latent_image=None,
denoise=1.0,
self,
coreml_model,
seed,
steps,
cfg,
sampler_name,
scheduler,
positive,
negative,
latent_image=None,
denoise=1.0,
):
sample_shape = coreml_model.expected_inputs["sample"]["shape"]
latent_image = reshape_latent_image(latent_image, sample_shape)
latent_image["samples"] = latent_image["samples"][0:1]
model_config = get_model_config()
wrapped_model = CoreMLModelAdapter(model_config, coreml_model)
wrapped_model = CoreMLModelWrapper(model_config, coreml_model)
model = ModelPatcher(wrapped_model, get_torch_device(), None)
if latent_image is None:
logger.warning("No latent image provided, using empty tensor.")
expected = coreml_model.expected_inputs["sample"]["shape"]
latent_image = {"samples": torch.zeros(expected[0] // 2, *expected[1:])}
batch_size = latent_image["samples"].shape[0]
if batch_size != coreml_model.expected_inputs["sample"]["shape"][0]:
logger.warning(
"Batch size is different from expected input size. Chunking and/or padding will be applied."
)
return super().sample(
model,
seed,
+31
View File
@@ -0,0 +1,31 @@
import pytest
import torch
from coreml_suite.latents import chunk_batch, merge_chunks
@pytest.mark.parametrize("batch_size", [2, 4, 5, 9])
def test_batch_chunking(batch_size):
latent_image = torch.randn(batch_size, 4, 64, 64)
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)
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)
-19
View File
@@ -1,19 +0,0 @@
import pytest
import torch
from coreml_suite.latents import reshape_latent_image
def test_fix_latents_no_latent_image():
reshaped = reshape_latent_image(None, (2, 4, 64, 64))
assert reshaped["samples"].shape == (2, 4, 64, 64)
@pytest.mark.parametrize(
"latent_shape", [(2, 4, 64, 64), (2, 4, 128, 128), (2, 4, 32, 32), (2, 4, 128, 64)]
)
def test_reshape_latents(latent_shape):
latent_image = {"samples": torch.zeros(latent_shape)}
reshaped = reshape_latent_image(latent_image, (2, 4, 64, 64))
assert reshaped["samples"].shape == (2, 4, 64, 64)