Chunk and pad batches
This commit is contained in:
+28
-15
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user