diff --git a/conftest.py b/conftest.py new file mode 100644 index 0000000..9d2e1f4 --- /dev/null +++ b/conftest.py @@ -0,0 +1,4 @@ +"""Top-level conftest: prevent pytest from importing the repo-root +__init__.py (the ComfyUI custom-node entry point pulls in comfy + nodes, +which breaks the Tier-0 'no-framework' promise).""" +collect_ignore = ["__init__.py"] diff --git a/coreml_suite/controlnet.py b/coreml_suite/controlnet.py index a984bfb..141b471 100644 --- a/coreml_suite/controlnet.py +++ b/coreml_suite/controlnet.py @@ -1,62 +1,14 @@ -from itertools import chain -from math import ceil +"""Phase 3 compatibility shim — re-exports from coreml_suite.core.controlnet.""" +from coreml_suite.core.controlnet import ( + chunk_control, + expand_inputs, + extract_residual_kwargs, + no_control, +) -import numpy as np -import torch - -from coreml_suite.latents import chunk_batch - - -def expand_inputs(inputs): - expanded = inputs.copy() - for k, v in inputs.items(): - if isinstance(v, np.ndarray): - expanded[k] = np.concatenate([v] * 2) if v.shape[0] == 1 else v - elif isinstance(v, torch.Tensor): - expanded[k] = torch.cat([v] * 2) if v.shape[0] == 1 else v - elif isinstance(v, list): - expanded[k] = v * 2 if len(v) == 1 else v - elif isinstance(v, dict): - expand_inputs(v) - return expanded - - -def extract_residual_kwargs(expected_inputs, control): - if "additional_residual_0" not in expected_inputs.keys(): - return {} - if control is None: - return no_control(expected_inputs) - - residual_kwargs = { - "additional_residual_{}".format(i): r.cpu().numpy().astype(np.float16) - for i, r in enumerate(chain(control["output"], control["middle"])) - } - return residual_kwargs - - -def no_control(expected_inputs): - shapes_dict = { - k: v["shape"] for k, v in expected_inputs.items() if k.startswith("additional") - } - residual_kwargs = { - k: torch.zeros(*shape).cpu().numpy().astype(dtype=np.float16) - for k, shape in shapes_dict.items() - } - return residual_kwargs - - -def chunk_control(cn, target_size): - if cn is None: - return [None] * target_size - - num_chunks = ceil(cn["output"][0].shape[0] / target_size) - - out = [{"output": [], "middle": []} for _ in range(num_chunks)] - - for k, v in cn.items(): - for i, x in enumerate(v): - chunks = chunk_batch(x, (target_size, *x.shape[1:])) - for j, chunk in enumerate(chunks): - out[j][k].append(chunk) - - return out +__all__ = [ + "chunk_control", + "expand_inputs", + "extract_residual_kwargs", + "no_control", +] diff --git a/coreml_suite/core/__init__.py b/coreml_suite/core/__init__.py new file mode 100644 index 0000000..064138f --- /dev/null +++ b/coreml_suite/core/__init__.py @@ -0,0 +1,10 @@ +"""Framework-free pure-logic core of ComfyUI-CoreMLSuite. + +Modules under this package must NOT import `comfy`, `coremltools`, +`python_coreml_stable_diffusion`, `folder_paths`, `nodes`, or any other +ComfyUI / Apple runtime. Only `numpy` and `torch` are allowed. + +The thin adapters in `coreml_suite.{latents,controlnet,models}` keep the +old public import paths working so `coreml_suite/nodes.py` and downstream +ComfyUI workflows are unchanged. +""" diff --git a/coreml_suite/core/controlnet.py b/coreml_suite/core/controlnet.py new file mode 100644 index 0000000..0f063d3 --- /dev/null +++ b/coreml_suite/core/controlnet.py @@ -0,0 +1,67 @@ +"""Pure helpers around the ControlNet residual inputs of the Core ML UNet. + +Moved from coreml_suite.controlnet in Phase 3. Behavior unchanged — Phase 2 +characterization tests cover shapes, dtype (fp16), and zero-fill fallback. +""" +from itertools import chain +from math import ceil + +import numpy as np +import torch + +from coreml_suite.core.latents import chunk_batch + + +def expand_inputs(inputs): + expanded = inputs.copy() + for k, v in inputs.items(): + if isinstance(v, np.ndarray): + expanded[k] = np.concatenate([v] * 2) if v.shape[0] == 1 else v + elif isinstance(v, torch.Tensor): + expanded[k] = torch.cat([v] * 2) if v.shape[0] == 1 else v + elif isinstance(v, list): + expanded[k] = v * 2 if len(v) == 1 else v + elif isinstance(v, dict): + expand_inputs(v) + return expanded + + +def extract_residual_kwargs(expected_inputs, control): + if "additional_residual_0" not in expected_inputs.keys(): + return {} + if control is None: + return no_control(expected_inputs) + + residual_kwargs = { + "additional_residual_{}".format(i): r.cpu().numpy().astype(np.float16) + for i, r in enumerate(chain(control["output"], control["middle"])) + } + return residual_kwargs + + +def no_control(expected_inputs): + shapes_dict = { + k: v["shape"] for k, v in expected_inputs.items() if k.startswith("additional") + } + residual_kwargs = { + k: torch.zeros(*shape).cpu().numpy().astype(dtype=np.float16) + for k, shape in shapes_dict.items() + } + return residual_kwargs + + +def chunk_control(cn, target_size): + if cn is None: + return [None] * target_size + + num_chunks = ceil(cn["output"][0].shape[0] / target_size) + + out = [{"output": [], "middle": []} for _ in range(num_chunks)] + + for k, v in cn.items(): + for i, x in enumerate(v): + chunks = chunk_batch(x, (target_size, *x.shape[1:])) + for j, chunk in enumerate(chunks): + out[j][k].append(chunk) + + return out diff --git a/coreml_suite/core/inputs.py b/coreml_suite/core/inputs.py new file mode 100644 index 0000000..9263a47 --- /dev/null +++ b/coreml_suite/core/inputs.py @@ -0,0 +1,114 @@ +"""Pure transform from torch sampler inputs to Core ML UNet kwargs. + +Moved from coreml_suite.models in Phase 3. Behavior unchanged — Phase 2 +characterization tests cover SD1.5 / SDXL base / SDXL refiner / LCM +variants and the chunked-batch fan-out. +""" +import numpy as np +import torch + +from coreml_suite.core.controlnet import extract_residual_kwargs, chunk_control +from coreml_suite.core.latents import chunk_batch + + +class CoreMLInputs: + def __init__(self, x, t, context, control, **kwargs): + self.x = x + self.t = t + self.context = context + self.control = control + self.time_ids = kwargs.get("time_ids") + self.text_embeds = kwargs.get("text_embeds") + self.ts_cond = kwargs.get("timestep_cond") + + def coreml_kwargs(self, expected_inputs): + sample = self.x.cpu().numpy().astype(np.float16) + + context = self.context.cpu().numpy().astype(np.float16) + context = context.transpose(0, 2, 1)[:, :, None, :] + + t = self.t.cpu().numpy().astype(np.float16) + + model_input_kwargs = { + "sample": sample, + "encoder_hidden_states": context, + "timestep": t, + } + residual_kwargs = extract_residual_kwargs(expected_inputs, self.control) + model_input_kwargs |= residual_kwargs + + # LCM + if self.ts_cond is not None: + model_input_kwargs["timestep_cond"] = ( + self.ts_cond.cpu().numpy().astype(np.float16) + ) + + # SDXL + if "text_embeds" in expected_inputs: + model_input_kwargs["text_embeds"] = ( + self.text_embeds.cpu().numpy().astype(np.float16) + ) + if "time_ids" in expected_inputs: + model_input_kwargs["time_ids"] = ( + self.time_ids.cpu().numpy().astype(np.float16) + ) + + return model_input_kwargs + + def chunks(self, expected_inputs): + sample_shape = expected_inputs["sample"]["shape"] + timestep_shape = expected_inputs["timestep"]["shape"] + hidden_shape = expected_inputs["encoder_hidden_states"]["shape"] + context_shape = (hidden_shape[0], hidden_shape[3], hidden_shape[1]) + + chunked_x = chunk_batch(self.x, sample_shape) + ts = list(torch.full((len(chunked_x), timestep_shape[0]), self.t[0])) + chunked_context = chunk_batch(self.context, context_shape) + + chunked_control = [None] * len(chunked_x) + if self.control is not None: + chunked_control = chunk_control(self.control, sample_shape[0]) + + chunked_ts_cond = [None] * len(chunked_x) + if self.ts_cond is not None: + ts_cond_shape = expected_inputs["timestep_cond"]["shape"] + chunked_ts_cond = chunk_batch(self.ts_cond, ts_cond_shape) + + chunked_time_ids = [None] * len(chunked_x) + if expected_inputs.get("time_ids") is not None: + time_ids_shape = expected_inputs["time_ids"]["shape"] + if self.time_ids is None: + self.time_ids = torch.zeros(len(chunked_x), *time_ids_shape[1:]).to( + self.x.device + ) + chunked_time_ids = chunk_batch(self.time_ids, time_ids_shape) + + chunked_text_embeds = [None] * len(chunked_x) + if expected_inputs.get("text_embeds") is not None: + text_embeds_shape = expected_inputs["text_embeds"]["shape"] + if self.text_embeds is None: + self.text_embeds = torch.zeros( + len(chunked_x), *text_embeds_shape[1:] + ).to(self.x.device) + chunked_text_embeds = chunk_batch(self.text_embeds, text_embeds_shape) + + return [ + CoreMLInputs( + x, + t, + context, + control, + timestep_cond=ts_cond, + time_ids=time_ids, + text_embeds=text_embeds, + ) + for x, t, context, control, ts_cond, time_ids, text_embeds in zip( + chunked_x, + ts, + chunked_context, + chunked_control, + chunked_ts_cond, + chunked_time_ids, + chunked_text_embeds, + ) + ] diff --git a/coreml_suite/core/latents.py b/coreml_suite/core/latents.py new file mode 100644 index 0000000..58f72ce --- /dev/null +++ b/coreml_suite/core/latents.py @@ -0,0 +1,42 @@ +"""Pure batch-chunking helpers for Core ML's fixed-shape UNet inputs. + +Moved from coreml_suite.latents in Phase 3. Behavior unchanged — Phase 2 +characterization tests cover the contract (padding-zero regions, truncation +in merge_chunks, identity-passthrough when shape already matches). +""" +import torch + + +def chunk_batch(input_tensor, target_shape): + if input_tensor.shape == target_shape: + return [input_tensor] + + batch_size = input_tensor.shape[0] + target_batch_size = target_shape[0] + + num_chunks = batch_size // target_batch_size + if num_chunks == 0: + padding = torch.zeros(target_batch_size - batch_size, *target_shape[1:]).to( + input_tensor.device + ) + return [torch.cat((input_tensor, padding), dim=0)] + + mod = batch_size % target_batch_size + if mod != 0: + chunks = list(torch.chunk(input_tensor[:-mod], num_chunks)) + padding = torch.zeros(target_batch_size - mod, *target_shape[1:]).to( + input_tensor.device + ) + padded = torch.cat((input_tensor[-mod:], padding), dim=0) + chunks.append(padded) + return chunks + + chunks = list(torch.chunk(input_tensor, num_chunks)) + return 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]] diff --git a/coreml_suite/core/naming.py b/coreml_suite/core/naming.py new file mode 100644 index 0000000..9f193d3 --- /dev/null +++ b/coreml_suite/core/naming.py @@ -0,0 +1,49 @@ +"""Pure out_name composition for the Core ML UNet artifact. + +Extracted from CoreMLConverter.convert in Phase 3 so the filename contract +can be tested + reused without instantiating the node. The string is the +cache key: every workflow that references a converted .mlpackage depends +on it staying byte-for-byte identical. +""" +from typing import Iterable, Tuple + +ATTN_SUFFIX = { + "SPLIT_EINSUM": "se", + "SPLIT_EINSUM_V2": "se2", + "ORIGINAL": "orig", +} + + +def compose_out_name( + *, + ckpt_name: str, + batch_size: int, + width: int, + height: int, + controlnet_support: bool, + attention_implementation: str, + lora_names: Iterable[str] = (), +) -> str: + """Build the .mlpackage stem from convert() parameters. + + Locked behaviour (Phase 2 characterization tests): + - first '.' in ckpt_name wins (`a.b.c.safetensors` -> `a`) + - spaces collapse to underscores + - LoRA names are taken stem-only, sorted, joined with '_' and + prefixed with '_' when present (caller is expected to pass a + sorted list; we sort defensively) + - controlnet adds `_cn` + - attn suffix is `_se` | `_se2` | `_orig` + """ + stem = ckpt_name.split(".")[0] + sorted_names = sorted(lora_names) + lora_str = "_" + "_".join(name.split(".")[0] for name in sorted_names) if sorted_names else "" + cn_suffix = "_cn" if controlnet_support else "" + attn_suffix = "_" + ATTN_SUFFIX[attention_implementation] + out_name = f"{stem}{lora_str}_{batch_size}x{width}x{height}{cn_suffix}{attn_suffix}" + return out_name.replace(" ", "_") + + +def lora_names_from_params(lora_params: Iterable[Tuple[str, float]]) -> list[str]: + """Mirror the sort applied inside CoreMLConverter.convert.""" + return [name for name, _ in sorted(lora_params, key=lambda pair: pair[0])] diff --git a/coreml_suite/core/sdxl.py b/coreml_suite/core/sdxl.py new file mode 100644 index 0000000..7ea16eb --- /dev/null +++ b/coreml_suite/core/sdxl.py @@ -0,0 +1,91 @@ +"""Pure SDXL detection + time_ids/text_embeds assembly. + +Moved from coreml_suite.models in Phase 3. The framework-coupled adapter +`add_sdxl_model_options` lives in models.py and now delegates the math +here. Phase 2 characterization tests cover base (len 6) vs refiner +(len 5) and the closure free-vars produced by `sdxl_model_function_wrapper`. +""" +import torch + + +def is_sdxl(coreml_model): + return ( + "time_ids" in coreml_model.expected_inputs + and "text_embeds" in coreml_model.expected_inputs + ) + + +def is_sdxl_base(coreml_model): + return ( + is_sdxl(coreml_model) + and coreml_model.expected_inputs["time_ids"]["shape"][1] == 6 + ) + + +def is_sdxl_refiner(coreml_model): + return ( + is_sdxl(coreml_model) + and coreml_model.expected_inputs["time_ids"]["shape"][1] == 5 + ) + + +def build_sdxl_time_ids(pos_dict, neg_dict, *, is_base: bool, is_refiner: bool): + """Compose the (2, N) time_ids tensor for the SDXL Core ML UNet. + + - base: N=6 -> [h, w, crop_h, crop_w, target_h, target_w] + - refiner: N=5 -> [h, w, crop_h, crop_w, aesthetic_score] + - neither: N=4 -> [h, w, crop_h, crop_w] (edge case kept for parity) + """ + pos_time_ids = [ + pos_dict.get("height", 768), + pos_dict.get("width", 768), + pos_dict.get("crop_h", 0), + pos_dict.get("crop_w", 0), + ] + neg_time_ids = [ + neg_dict.get("height", 768), + neg_dict.get("width", 768), + neg_dict.get("crop_h", 0), + neg_dict.get("crop_w", 0), + ] + + if is_base: + pos_time_ids += [ + pos_dict.get("target_height", 768), + pos_dict.get("target_width", 768), + ] + neg_time_ids += [ + neg_dict.get("target_height", 768), + neg_dict.get("target_width", 768), + ] + + if is_refiner: + pos_time_ids += [pos_dict.get("aesthetic_score", 6)] + neg_time_ids += [neg_dict.get("aesthetic_score", 2.5)] + + return torch.tensor([pos_time_ids, neg_time_ids]) + + +def build_sdxl_text_embeds(pos_pooled, neg_pooled): + """Concat pos then neg along the batch dim. Locked contract.""" + return torch.cat((pos_pooled, neg_pooled)) + + +def sdxl_model_function_wrapper(time_ids, text_embeds, refiner=False): + def wrapper(model_function, params): + x = params["input"] + t = params["timestep"] + c = params["c"] + + context = c.get("c_crossattn") + + if context is None: + return torch.zeros_like(x) + + if refiner and context is not None: + # converted refiner accepts only g clip + c["c_crossattn"] = context[:, :, 768:] + + return model_function(x, t, **c, time_ids=time_ids, text_embeds=text_embeds) + + return wrapper diff --git a/coreml_suite/latents.py b/coreml_suite/latents.py index 7b44b30..d2896df 100644 --- a/coreml_suite/latents.py +++ b/coreml_suite/latents.py @@ -1,36 +1,4 @@ -import torch +"""Phase 3 compatibility shim — re-exports from coreml_suite.core.latents.""" +from coreml_suite.core.latents import chunk_batch, merge_chunks - -def chunk_batch(input_tensor, target_shape): - if input_tensor.shape == target_shape: - return [input_tensor] - - batch_size = input_tensor.shape[0] - target_batch_size = target_shape[0] - - num_chunks = batch_size // target_batch_size - if num_chunks == 0: - padding = torch.zeros(target_batch_size - batch_size, *target_shape[1:]).to( - input_tensor.device - ) - return [torch.cat((input_tensor, padding), dim=0)] - - mod = batch_size % target_batch_size - if mod != 0: - chunks = list(torch.chunk(input_tensor[:-mod], num_chunks)) - padding = torch.zeros(target_batch_size - mod, *target_shape[1:]).to( - input_tensor.device - ) - padded = torch.cat((input_tensor[-mod:], padding), dim=0) - chunks.append(padded) - return chunks - - chunks = list(torch.chunk(input_tensor, num_chunks)) - return 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]] +__all__ = ["chunk_batch", "merge_chunks"] diff --git a/coreml_suite/models.py b/coreml_suite/models.py index 05a506d..f4bc360 100644 --- a/coreml_suite/models.py +++ b/coreml_suite/models.py @@ -1,15 +1,44 @@ -import numpy as np +"""Framework-coupled glue between Core ML UNets and ComfyUI's sampler stack. + +Pure math (CoreMLInputs, SDXL detection, time_ids/text_embeds assembly, +sdxl_model_function_wrapper) lives in coreml_suite.core.* after Phase 3. +This module is what touches comfy.*: model_base, ModelPatcher, the +diffusion_model wrapper, and the maintainer-facing add_sdxl_model_options +adapter. +""" import torch from comfy import model_base from comfy.model_management import get_torch_device from comfy.model_patcher import ModelPatcher + from coreml_suite.config import get_model_config, ModelVersion -from coreml_suite.controlnet import extract_residual_kwargs, chunk_control -from coreml_suite.latents import chunk_batch, merge_chunks +from coreml_suite.core.inputs import CoreMLInputs +from coreml_suite.core.latents import merge_chunks +from coreml_suite.core.sdxl import ( + build_sdxl_text_embeds, + build_sdxl_time_ids, + is_sdxl, + is_sdxl_base, + is_sdxl_refiner, + sdxl_model_function_wrapper, +) from coreml_suite.lcm.utils import is_lcm from coreml_suite.logger import logger +__all__ = [ + "CoreMLInputs", + "CoreMLModelWrapper", + "CoreMLModelWrapperLCM", + "add_sdxl_model_options", + "get_latent_image", + "get_model_patcher", + "is_sdxl", + "is_sdxl_base", + "is_sdxl_refiner", + "sdxl_model_function_wrapper", +] + class CoreMLModelWrapper: def __init__(self, coreml_model): @@ -68,204 +97,27 @@ class CoreMLModelWrapperLCM(CoreMLModelWrapper): self.config = None -class CoreMLInputs: - def __init__(self, x, t, context, control, **kwargs): - self.x = x - self.t = t - self.context = context - self.control = control - self.time_ids = kwargs.get("time_ids") - self.text_embeds = kwargs.get("text_embeds") - self.ts_cond = kwargs.get("timestep_cond") - - def coreml_kwargs(self, expected_inputs): - sample = self.x.cpu().numpy().astype(np.float16) - - context = self.context.cpu().numpy().astype(np.float16) - context = context.transpose(0, 2, 1)[:, :, None, :] - - t = self.t.cpu().numpy().astype(np.float16) - - model_input_kwargs = { - "sample": sample, - "encoder_hidden_states": context, - "timestep": t, - } - residual_kwargs = extract_residual_kwargs(expected_inputs, self.control) - model_input_kwargs |= residual_kwargs - - # LCM - if self.ts_cond is not None: - model_input_kwargs["timestep_cond"] = ( - self.ts_cond.cpu().numpy().astype(np.float16) - ) - - # SDXL - if "text_embeds" in expected_inputs: - model_input_kwargs["text_embeds"] = ( - self.text_embeds.cpu().numpy().astype(np.float16) - ) - if "time_ids" in expected_inputs: - model_input_kwargs["time_ids"] = ( - self.time_ids.cpu().numpy().astype(np.float16) - ) - - return model_input_kwargs - - def chunks(self, expected_inputs): - sample_shape = expected_inputs["sample"]["shape"] - timestep_shape = expected_inputs["timestep"]["shape"] - hidden_shape = expected_inputs["encoder_hidden_states"]["shape"] - context_shape = (hidden_shape[0], hidden_shape[3], hidden_shape[1]) - - chunked_x = chunk_batch(self.x, sample_shape) - ts = list(torch.full((len(chunked_x), timestep_shape[0]), self.t[0])) - chunked_context = chunk_batch(self.context, context_shape) - - chunked_control = [None] * len(chunked_x) - if self.control is not None: - chunked_control = chunk_control(self.control, sample_shape[0]) - - chunked_ts_cond = [None] * len(chunked_x) - if self.ts_cond is not None: - ts_cond_shape = expected_inputs["timestep_cond"]["shape"] - chunked_ts_cond = chunk_batch(self.ts_cond, ts_cond_shape) - - chunked_time_ids = [None] * len(chunked_x) - if expected_inputs.get("time_ids") is not None: - time_ids_shape = expected_inputs["time_ids"]["shape"] - if self.time_ids is None: - self.time_ids = torch.zeros(len(chunked_x), *time_ids_shape[1:]).to( - self.x.device - ) - chunked_time_ids = chunk_batch(self.time_ids, time_ids_shape) - - chunked_text_embeds = [None] * len(chunked_x) - if expected_inputs.get("text_embeds") is not None: - text_embeds_shape = expected_inputs["text_embeds"]["shape"] - if self.text_embeds is None: - self.text_embeds = torch.zeros( - len(chunked_x), *text_embeds_shape[1:] - ).to(self.x.device) - chunked_text_embeds = chunk_batch(self.text_embeds, text_embeds_shape) - - return [ - CoreMLInputs( - x, - t, - context, - control, - timestep_cond=ts_cond, - time_ids=time_ids, - text_embeds=text_embeds, - ) - for x, t, context, control, ts_cond, time_ids, text_embeds in zip( - chunked_x, - ts, - chunked_context, - chunked_control, - chunked_ts_cond, - chunked_time_ids, - chunked_text_embeds, - ) - ] - - -def is_sdxl(coreml_model): - return ( - "time_ids" in coreml_model.expected_inputs - and "text_embeds" in coreml_model.expected_inputs - ) - - -def is_sdxl_base(coreml_model): - return ( - is_sdxl(coreml_model) - and coreml_model.expected_inputs["time_ids"]["shape"][1] == 6 - ) - - -def is_sdxl_refiner(coreml_model): - return ( - is_sdxl(coreml_model) - and coreml_model.expected_inputs["time_ids"]["shape"][1] == 5 - ) - - -def sdxl_model_function_wrapper(time_ids, text_embeds, refiner=False): - def wrapper(model_function, params): - x = params["input"] - t = params["timestep"] - c = params["c"] - - context = c.get("c_crossattn") - - if context is None: - return torch.zeros_like(x) - - if refiner and context is not None: - # converted refiner accepts only g clip - c["c_crossattn"] = context[:, :, 768:] - - return model_function(x, t, **c, time_ids=time_ids, text_embeds=text_embeds) - - return wrapper - - def add_sdxl_model_options(model_patcher, positive, negative): mp = model_patcher.clone() pos_dict = positive[0][1] neg_dict = negative[0][1] - pos_pooled = pos_dict["pooled_output"] - neg_pooled = neg_dict["pooled_output"] - - pos_time_ids = [ - pos_dict.get("height", 768), - pos_dict.get("width", 768), - pos_dict.get("crop_h", 0), - pos_dict.get("crop_w", 0), - ] - - neg_time_ids = [ - neg_dict.get("height", 768), - neg_dict.get("width", 768), - neg_dict.get("crop_h", 0), - neg_dict.get("crop_w", 0), - ] - - if model_patcher.model.diffusion_model.is_sdxl_base: - pos_time_ids += [ - pos_dict.get("target_height", 768), - pos_dict.get("target_width", 768), - ] - - neg_time_ids += [ - neg_dict.get("target_height", 768), - neg_dict.get("target_width", 768), - ] - + is_base = model_patcher.model.diffusion_model.is_sdxl_base is_refiner = model_patcher.model.diffusion_model.is_sdxl_refiner - if is_refiner: - pos_time_ids += [ - pos_dict.get("aesthetic_score", 6), - ] - neg_time_ids += [ - neg_dict.get("aesthetic_score", 2.5), - ] + time_ids = build_sdxl_time_ids( + pos_dict, neg_dict, is_base=is_base, is_refiner=is_refiner + ) + text_embeds = build_sdxl_text_embeds( + pos_dict["pooled_output"], neg_dict["pooled_output"] + ) - time_ids = torch.tensor([pos_time_ids, neg_time_ids]) - text_embeds = torch.cat((pos_pooled, neg_pooled)) - - model_options = { + mp.model_options |= { "model_function_wrapper": sdxl_model_function_wrapper( time_ids, text_embeds, is_refiner ), } - mp.model_options |= model_options - return mp diff --git a/coreml_suite/nodes.py b/coreml_suite/nodes.py index 97dca41..e71df1a 100644 --- a/coreml_suite/nodes.py +++ b/coreml_suite/nodes.py @@ -8,6 +8,7 @@ import folder_paths from coreml_suite import COREML_NODE from coreml_suite import converter from coreml_suite.config import ModelVersion +from coreml_suite.core.naming import compose_out_name, lora_names_from_params from coreml_suite.lcm.utils import add_lcm_model_options, lcm_patch, is_lcm from coreml_suite.logger import logger from nodes import KSampler, LoraLoader, KSamplerAdvanced @@ -288,24 +289,16 @@ class CoreMLConverter(COREML_NODE): h = height w = width sample_size = (h // 8, w // 8) - batch_size = batch_size - cn_support_str = "_cn" if controlnet_support else "" - lora_str = ( - "_" + "_".join(lora_param[0].split(".")[0] for lora_param in lora_params) - if lora_params - else "" + out_name = compose_out_name( + ckpt_name=ckpt_name, + batch_size=batch_size, + width=w, + height=h, + controlnet_support=controlnet_support, + attention_implementation=attention_implementation, + lora_names=lora_names_from_params(lora_params), ) - attn_str = ( - "_" - + {"SPLIT_EINSUM": "se", "SPLIT_EINSUM_V2": "se2", "ORIGINAL": "orig"}[ - attention_implementation - ] - ) - - out_name = f"{ckpt_name.split('.')[0]}{lora_str}_{batch_size}x{w}x{h}{cn_support_str}{attn_str}" - out_name = out_name.replace(" ", "_") - logger.info(f"Converting {ckpt_name} to {out_name}") logger.info(f"Batch size: {batch_size}") logger.info(f"Width: {w}, Height: {h}") diff --git a/pyproject.toml b/pyproject.toml index b35c783..66e49a8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -53,3 +53,7 @@ markers = [ "smoke: macOS-ARM smoke test on a synthetic micro-model (Tier 1)", ] testpaths = ["tests"] +# importlib mode keeps pytest from importing the repo-root __init__.py +# (which is the ComfyUI custom-node entry and pulls in comfy + nodes). +# Without this Tier-0 leaks the entire ComfyUI runtime on collection. +addopts = ["--import-mode=importlib", "--confcutdir=tests"] diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/conftest.py b/tests/conftest.py index 0754a0c..1587338 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -40,6 +40,31 @@ _TIER_BY_DIR = { "tests/smoke": "smoke", } +# When the user asks for a single tier (-m unit / -m m2), skip the other +# directories at collection time. Tier-0 cannot afford to import tests/m2 +# files because they pull in PIL + ComfyUI runtime which Linux CI won't have. +_TIER_DIRS = { + "unit": ("/tests/unit/",), + "m2": ("/tests/m2/", "/tests/integration/"), + "smoke": ("/tests/smoke/",), +} + + +def pytest_ignore_collect(collection_path, config): + expr = config.option.markexpr + if expr not in _TIER_DIRS: + return None + allowed = _TIER_DIRS[expr] + rel = str(collection_path).replace("\\", "/") + if "/tests/" not in rel: + return None + # Always allow tests/ root + the tier's own dirs. + if rel.endswith("/tests"): + return None + if any(frag in rel + "/" for frag in allowed): + return None + return True + def pytest_collection_modifyitems(config, items): for item in items: diff --git a/tests/unit/test_characterization_controlnet.py b/tests/unit/test_characterization_controlnet.py index 1c20756..bb46258 100644 --- a/tests/unit/test_characterization_controlnet.py +++ b/tests/unit/test_characterization_controlnet.py @@ -9,7 +9,7 @@ import numpy as np import pytest import torch -from coreml_suite.controlnet import ( +from coreml_suite.core.controlnet import ( chunk_control, expand_inputs, extract_residual_kwargs, diff --git a/tests/unit/test_characterization_inputs.py b/tests/unit/test_characterization_inputs.py index c1769ca..8209b1d 100644 --- a/tests/unit/test_characterization_inputs.py +++ b/tests/unit/test_characterization_inputs.py @@ -11,7 +11,7 @@ import numpy as np import pytest import torch -from coreml_suite.models import CoreMLInputs +from coreml_suite.core.inputs import CoreMLInputs @pytest.fixture(autouse=True) diff --git a/tests/unit/test_characterization_latents.py b/tests/unit/test_characterization_latents.py index ff6f136..24e3df7 100644 --- a/tests/unit/test_characterization_latents.py +++ b/tests/unit/test_characterization_latents.py @@ -7,7 +7,7 @@ cannot silently shift either contract. import pytest import torch -from coreml_suite.latents import chunk_batch, merge_chunks +from coreml_suite.core.latents import chunk_batch, merge_chunks @pytest.fixture(autouse=True) diff --git a/tests/unit/test_characterization_out_name.py b/tests/unit/test_characterization_out_name.py index cb2f22d..6a27220 100644 --- a/tests/unit/test_characterization_out_name.py +++ b/tests/unit/test_characterization_out_name.py @@ -1,72 +1,16 @@ -"""Phase 2 characterization tests for CoreMLConverter.out_name composition. +"""Phase 2 characterization tests, Phase 3 re-pointed. -out_name is encoded into the .mlpackage filename and therefore drives the -"have we converted this combo already?" cache check. Drift here silently -invalidates user caches and breaks workflow node references. - -Tests intercept converter.get_out_path to capture the composed string, -and stub out the heavy conversion + Core ML model load. +After Phase 3 the .mlpackage filename composition is the pure +coreml_suite.core.naming.compose_out_name function. CoreMLConverter.convert +calls it; the previous Phase 2 test had to monkey-patch heavy converter +internals just to capture the string, which made the test framework-coupled. """ import pytest -import folder_paths -from coreml_suite import converter -from coreml_suite import nodes as nodes_mod -from coreml_suite.nodes import CoreMLConverter +from coreml_suite.core.naming import compose_out_name, lora_names_from_params -@pytest.fixture -def capture_out_name(monkeypatch): - captured = {} - - def fake_get_out_path(submodule, name): - captured["submodule"] = submodule - captured["out_name"] = name - return f"/tmp/fake/{name}.mlpackage" - - monkeypatch.setattr(converter, "get_out_path", fake_get_out_path) - monkeypatch.setattr(converter, "convert", lambda **kwargs: None) - monkeypatch.setattr( - converter, - "compile_model", - lambda out_path, out_name, submodule_name: f"/tmp/fake/{out_name}_{submodule_name}.mlmodelc", - ) - monkeypatch.setattr( - folder_paths, - "get_full_path", - lambda kind, name: f"/tmp/ckpts/{name}" if kind == "checkpoints" else None, - ) - monkeypatch.setattr(nodes_mod, "CoreMLModel", lambda *a, **kw: object()) - return captured - - -def _convert( - *, - ckpt_name="dreamshaper_8.safetensors", - model_version="SD15", - height=512, - width=512, - batch_size=1, - attention_implementation="SPLIT_EINSUM", - compute_unit="CPU_AND_NE", - controlnet_support=False, - lora_params=None, -): - node = CoreMLConverter() - node.convert( - ckpt_name=ckpt_name, - model_version=model_version, - height=height, - width=width, - batch_size=batch_size, - attention_implementation=attention_implementation, - compute_unit=compute_unit, - controlnet_support=controlnet_support, - lora_params=lora_params, - ) - - -# ---------- basic golden strings -------------------------------------------- +# ---------- attention suffixes ---------------------------------------------- @pytest.mark.parametrize( @@ -77,99 +21,125 @@ def _convert( ("ORIGINAL", "orig"), ], ) -def test_out_name_attention_suffix(capture_out_name, attn_name, suffix): - _convert(attention_implementation=attn_name) - assert capture_out_name["out_name"] == f"dreamshaper_8_1x512x512_{suffix}" +def test_attention_suffix(attn_name, suffix): + out = compose_out_name( + ckpt_name="dreamshaper_8.safetensors", + batch_size=1, width=512, height=512, + controlnet_support=False, + attention_implementation=attn_name, + ) + assert out == f"dreamshaper_8_1x512x512_{suffix}" -def test_out_name_includes_batch_and_size(capture_out_name): - _convert(batch_size=4, width=768, height=1024) - assert capture_out_name["out_name"] == "dreamshaper_8_4x768x1024_se" +# ---------- batch / size ---------------------------------------------------- -def test_out_name_appends_cn_suffix_when_controlnet_support_true(capture_out_name): - _convert(controlnet_support=True) - assert capture_out_name["out_name"] == "dreamshaper_8_1x512x512_cn_se" +def test_includes_batch_and_size(): + out = compose_out_name( + ckpt_name="dreamshaper_8.safetensors", + batch_size=4, width=768, height=1024, + controlnet_support=False, + attention_implementation="SPLIT_EINSUM", + ) + assert out == "dreamshaper_8_4x768x1024_se" -def test_out_name_drops_dot_extension_only_at_first_period(capture_out_name): - """`ckpt_name.split('.')[0]` — first '.' wins; locked behaviour.""" - _convert(ckpt_name="my.checkpoint.v2.safetensors") - assert capture_out_name["out_name"] == "my_1x512x512_se" +# ---------- ControlNet ------------------------------------------------------ -def test_out_name_replaces_spaces_with_underscores(capture_out_name): - _convert(ckpt_name="dream shaper 8.safetensors") - assert capture_out_name["out_name"] == "dream_shaper_8_1x512x512_se" +def test_appends_cn_suffix_when_controlnet_support_true(): + out = compose_out_name( + ckpt_name="dreamshaper_8.safetensors", + batch_size=1, width=512, height=512, + controlnet_support=True, + attention_implementation="SPLIT_EINSUM", + ) + assert out == "dreamshaper_8_1x512x512_cn_se" + + +# ---------- ckpt name massage ----------------------------------------------- + + +def test_drops_extension_at_first_period(): + out = compose_out_name( + ckpt_name="my.checkpoint.v2.safetensors", + batch_size=1, width=512, height=512, + controlnet_support=False, + attention_implementation="SPLIT_EINSUM", + ) + assert out == "my_1x512x512_se" + + +def test_replaces_spaces_with_underscores(): + out = compose_out_name( + ckpt_name="dream shaper 8.safetensors", + batch_size=1, width=512, height=512, + controlnet_support=False, + attention_implementation="SPLIT_EINSUM", + ) + assert out == "dream_shaper_8_1x512x512_se" # ---------- LoRA suffixes --------------------------------------------------- -def test_out_name_with_single_lora(capture_out_name, monkeypatch): - monkeypatch.setattr( - folder_paths, - "get_full_path", - lambda kind, name: f"/tmp/{kind}/{name}", - ) - _convert(lora_params={"epi_noiseoffset.safetensors": (0.8,)}) - assert ( - capture_out_name["out_name"] - == "dreamshaper_8_epi_noiseoffset_1x512x512_se" +def test_single_lora(): + out = compose_out_name( + ckpt_name="dreamshaper_8.safetensors", + batch_size=1, width=512, height=512, + controlnet_support=False, + attention_implementation="SPLIT_EINSUM", + lora_names=["epi_noiseoffset.safetensors"], ) + assert out == "dreamshaper_8_epi_noiseoffset_1x512x512_se" -def test_out_name_with_multiple_loras_sorted(capture_out_name, monkeypatch): - """LoRAs are sorted by name then joined with '_' — locks the order.""" - monkeypatch.setattr( - folder_paths, - "get_full_path", - lambda kind, name: f"/tmp/{kind}/{name}", - ) - _convert( - lora_params={ - "zoom.safetensors": (1.0,), - "alpha.safetensors": (0.5,), - "moody.safetensors": (0.3,), - } - ) - assert ( - capture_out_name["out_name"] - == "dreamshaper_8_alpha_moody_zoom_1x512x512_se" +def test_multiple_loras_sorted(): + out = compose_out_name( + ckpt_name="dreamshaper_8.safetensors", + batch_size=1, width=512, height=512, + controlnet_support=False, + attention_implementation="SPLIT_EINSUM", + lora_names=["zoom.safetensors", "alpha.safetensors", "moody.safetensors"], ) + assert out == "dreamshaper_8_alpha_moody_zoom_1x512x512_se" -def test_out_name_lora_plus_controlnet(capture_out_name, monkeypatch): - monkeypatch.setattr( - folder_paths, - "get_full_path", - lambda kind, name: f"/tmp/{kind}/{name}", - ) - _convert( - lora_params={"a.safetensors": (1.0,)}, +def test_lora_plus_controlnet(): + out = compose_out_name( + ckpt_name="dreamshaper_8.safetensors", + batch_size=1, width=512, height=512, controlnet_support=True, + attention_implementation="SPLIT_EINSUM", + lora_names=["a.safetensors"], ) - assert capture_out_name["out_name"] == "dreamshaper_8_a_1x512x512_cn_se" + assert out == "dreamshaper_8_a_1x512x512_cn_se" # ---------- sdxl combinations ----------------------------------------------- -def test_out_name_sdxl_1024(capture_out_name): - _convert( +def test_sdxl_1024_original_gpu(): + out = compose_out_name( ckpt_name="sd_xl_base_1.0.safetensors", - model_version="SDXL", - width=1024, - height=1024, + batch_size=1, width=1024, height=1024, + controlnet_support=False, attention_implementation="ORIGINAL", - compute_unit="CPU_AND_GPU", ) - assert capture_out_name["out_name"] == "sd_xl_base_1_1x1024x1024_orig" + assert out == "sd_xl_base_1_1x1024x1024_orig" -# ---------- submodule path -------------------------------------------------- +# ---------- lora_names_from_params helper ---------------------------------- -def test_get_out_path_invoked_with_unet_submodule(capture_out_name): - _convert() - assert capture_out_name["submodule"] == "unet" +def test_lora_names_from_params_sorts_by_name(): + names = lora_names_from_params([ + ("zebra.safetensors", 1.0), + ("apple.safetensors", 0.5), + ("mango.safetensors", 0.7), + ]) + assert names == ["apple.safetensors", "mango.safetensors", "zebra.safetensors"] + + +def test_lora_names_from_params_empty_list(): + assert lora_names_from_params([]) == [] diff --git a/tests/unit/test_characterization_sdxl_options.py b/tests/unit/test_characterization_sdxl_options.py index 7b5f9f0..c51c374 100644 --- a/tests/unit/test_characterization_sdxl_options.py +++ b/tests/unit/test_characterization_sdxl_options.py @@ -1,22 +1,20 @@ -"""Phase 2 characterization tests for add_sdxl_model_options. +"""Phase 2 characterization tests, Phase 3 re-pointed. -Locks the SDXL time_ids / text_embeds assembly: base produces a (2, 6) -time_ids vector (h, w, crop_h, crop_w, target_h, target_w), refiner -produces (2, 5) (h, w, crop_h, crop_w, aesthetic_score). Both stack pos -then neg along the batch dim, and text_embeds is cat(pos_pooled, -neg_pooled). - -Uses a SimpleNamespace-shaped fake ModelPatcher because exercising the -real comfy.model_patcher.ModelPatcher here is overkill — only three -attribute paths are read by the SUT. +After Phase 3 the SDXL time_ids / text_embeds math lives in +coreml_suite.core.sdxl as pure builders. The framework adapter +add_sdxl_model_options (in models.py) is exercised separately by the m2 +golden image test; here we just lock the pure math. """ import inspect -from types import SimpleNamespace import pytest import torch -from coreml_suite.models import add_sdxl_model_options +from coreml_suite.core.sdxl import ( + build_sdxl_text_embeds, + build_sdxl_time_ids, + sdxl_model_function_wrapper, +) @pytest.fixture(autouse=True) @@ -24,124 +22,106 @@ def _deterministic_seed(): torch.manual_seed(0) -def _fake_patcher(is_base: bool, is_refiner: bool): - diffusion = SimpleNamespace(is_sdxl_base=is_base, is_sdxl_refiner=is_refiner) - model = SimpleNamespace(diffusion_model=diffusion) - patcher = SimpleNamespace(model=model, model_options={}) - patcher.clone = lambda: patcher # in-place: simplest mock that matches contract - return patcher +# ---------- build_sdxl_time_ids: base (len 6) ------------------------------- -def _cond(pooled, **overrides): - base = {"pooled_output": pooled} - base.update(overrides) - return [(None, base)] - - -def _closure_vars(wrapper): - return inspect.getclosurevars(wrapper).nonlocals - - -# ---------- base (len 6) ----------------------------------------------------- - - -def test_add_sdxl_model_options_base_produces_len6_time_ids_with_defaults(): - pos_pooled = torch.randn(1, 1280) - neg_pooled = torch.randn(1, 1280) - patcher = _fake_patcher(is_base=True, is_refiner=False) - out = add_sdxl_model_options(patcher, _cond(pos_pooled), _cond(neg_pooled)) - wrapper = out.model_options["model_function_wrapper"] - closure = _closure_vars(wrapper) - assert closure["time_ids"].shape == (2, 6) - # Defaults: 768 height/width, 0 crop, 768 target. +def test_build_time_ids_base_defaults(): + out = build_sdxl_time_ids({}, {}, is_base=True, is_refiner=False) expected = torch.tensor([[768, 768, 0, 0, 768, 768], [768, 768, 0, 0, 768, 768]]) - assert torch.equal(closure["time_ids"], expected) + assert out.shape == (2, 6) + assert torch.equal(out, expected) + + +def test_build_time_ids_base_respects_overrides(): + pos = {"height": 1024, "width": 512, "crop_h": 8, "crop_w": 4, + "target_height": 1024, "target_width": 1024} + neg = {"height": 256, "width": 256, "crop_h": 0, "crop_w": 0, + "target_height": 256, "target_width": 256} + out = build_sdxl_time_ids(pos, neg, is_base=True, is_refiner=False) + expected = torch.tensor([[1024, 512, 8, 4, 1024, 1024], [256, 256, 0, 0, 256, 256]]) + assert torch.equal(out, expected) + + +# ---------- build_sdxl_time_ids: refiner (len 5) ---------------------------- + + +def test_build_time_ids_refiner_defaults(): + out = build_sdxl_time_ids({}, {}, is_base=False, is_refiner=True) + expected = torch.tensor([[768, 768, 0, 0, 6.0], [768, 768, 0, 0, 2.5]]) + assert out.shape == (2, 5) + assert torch.equal(out, expected) + + +def test_build_time_ids_refiner_respects_aesthetic_score(): + pos = {"aesthetic_score": 8.5} + neg = {"aesthetic_score": 1.5} + out = build_sdxl_time_ids(pos, neg, is_base=False, is_refiner=True) + expected = torch.tensor([[768, 768, 0, 0, 8.5], [768, 768, 0, 0, 1.5]]) + assert torch.equal(out, expected) + + +# ---------- build_sdxl_time_ids: edge case ---------------------------------- + + +def test_build_time_ids_neither_base_nor_refiner_returns_len4(): + out = build_sdxl_time_ids({}, {}, is_base=False, is_refiner=False) + assert out.shape == (2, 4) + + +# ---------- build_sdxl_text_embeds ------------------------------------------ + + +def test_text_embeds_concat_pos_then_neg(): + pos = torch.full((1, 1280), 1.0) + neg = torch.full((1, 1280), -1.0) + out = build_sdxl_text_embeds(pos, neg) + assert out.shape == (2, 1280) + assert torch.equal(out[0], pos[0]) + assert torch.equal(out[1], neg[0]) + + +# ---------- sdxl_model_function_wrapper closure ----------------------------- + + +def test_wrapper_captures_time_ids_text_embeds_refiner_via_closure(): + time_ids = torch.zeros(2, 6) + text_embeds = torch.zeros(2, 1280) + wrapper = sdxl_model_function_wrapper(time_ids, text_embeds, refiner=False) + closure = inspect.getclosurevars(wrapper).nonlocals + assert closure["time_ids"] is time_ids + assert closure["text_embeds"] is text_embeds assert closure["refiner"] is False -def test_add_sdxl_model_options_base_respects_overrides(): - pos_pooled = torch.randn(1, 1280) - neg_pooled = torch.randn(1, 1280) - patcher = _fake_patcher(is_base=True, is_refiner=False) - out = add_sdxl_model_options( - patcher, - _cond(pos_pooled, height=1024, width=512, crop_h=8, crop_w=4, target_height=1024, target_width=1024), - _cond(neg_pooled, height=256, width=256, crop_h=0, crop_w=0, target_height=256, target_width=256), +def test_wrapper_returns_zero_when_context_missing(): + """When c_crossattn is None the wrapper short-circuits to zeros_like(x). + Locked here because Phase 3 mustn't change this default.""" + wrapper = sdxl_model_function_wrapper(torch.zeros(2, 6), torch.zeros(2, 1280)) + x = torch.randn(2, 4, 16, 16) + out = wrapper( + model_function=lambda *a, **kw: pytest.fail("model_function must not run"), + params={"input": x, "timestep": torch.zeros(2), "c": {}}, ) - closure = _closure_vars(out.model_options["model_function_wrapper"]) - expected = torch.tensor([[1024, 512, 8, 4, 1024, 1024], [256, 256, 0, 0, 256, 256]]) - assert torch.equal(closure["time_ids"], expected) + assert torch.equal(out, torch.zeros_like(x)) -# ---------- refiner (len 5) ------------------------------------------------- +def test_wrapper_refiner_truncates_context_to_g_clip(): + """refiner=True slices c_crossattn[:, :, 768:] before forwarding.""" + captured = {} + def fake_model(x, t, **c): + captured["context_shape"] = c["c_crossattn"].shape + captured["time_ids_shape"] = c["time_ids"].shape + return x -def test_add_sdxl_model_options_refiner_produces_len5_time_ids(): - pos_pooled = torch.randn(1, 1280) - neg_pooled = torch.randn(1, 1280) - patcher = _fake_patcher(is_base=False, is_refiner=True) - out = add_sdxl_model_options(patcher, _cond(pos_pooled), _cond(neg_pooled)) - closure = _closure_vars(out.model_options["model_function_wrapper"]) - assert closure["time_ids"].shape == (2, 5) - # Defaults: pos aesthetic_score=6, neg aesthetic_score=2.5. - expected = torch.tensor( - [[768, 768, 0, 0, 6.0], [768, 768, 0, 0, 2.5]], + wrapper = sdxl_model_function_wrapper( + torch.zeros(2, 5), torch.zeros(2, 1280), refiner=True ) - assert torch.equal(closure["time_ids"], expected) - assert closure["refiner"] is True - - -def test_add_sdxl_model_options_refiner_respects_aesthetic_score_overrides(): - pos_pooled = torch.randn(1, 1280) - neg_pooled = torch.randn(1, 1280) - patcher = _fake_patcher(is_base=False, is_refiner=True) - out = add_sdxl_model_options( - patcher, - _cond(pos_pooled, aesthetic_score=8.5), - _cond(neg_pooled, aesthetic_score=1.5), + x = torch.randn(2, 4, 16, 16) + context = torch.randn(2, 77, 2048) # 768 + 1280 dims + wrapper( + model_function=fake_model, + params={"input": x, "timestep": torch.zeros(2), "c": {"c_crossattn": context}}, ) - closure = _closure_vars(out.model_options["model_function_wrapper"]) - expected = torch.tensor([[768, 768, 0, 0, 8.5], [768, 768, 0, 0, 1.5]]) - assert torch.equal(closure["time_ids"], expected) - - -# ---------- text_embeds ---------------------------------------------------- - - -def test_text_embeds_is_concat_pos_then_neg_along_batch(): - pos_pooled = torch.full((1, 1280), 1.0) - neg_pooled = torch.full((1, 1280), -1.0) - patcher = _fake_patcher(is_base=True, is_refiner=False) - out = add_sdxl_model_options(patcher, _cond(pos_pooled), _cond(neg_pooled)) - closure = _closure_vars(out.model_options["model_function_wrapper"]) - embeds = closure["text_embeds"] - assert embeds.shape == (2, 1280) - assert torch.equal(embeds[0], pos_pooled[0]) - assert torch.equal(embeds[1], neg_pooled[0]) - - -def test_neither_base_nor_refiner_yields_len4_time_ids(): - """Locked edge case: if both is_sdxl_base and is_sdxl_refiner are False, - no extra entries are appended -> time_ids is only the 4 shared fields.""" - pos_pooled = torch.randn(1, 1280) - neg_pooled = torch.randn(1, 1280) - patcher = _fake_patcher(is_base=False, is_refiner=False) - out = add_sdxl_model_options(patcher, _cond(pos_pooled), _cond(neg_pooled)) - closure = _closure_vars(out.model_options["model_function_wrapper"]) - assert closure["time_ids"].shape == (2, 4) - - -# ---------- patcher contract ------------------------------------------------ - - -def test_model_options_dict_is_merged_into_clone_not_original(): - """Verify SUT writes to the clone's model_options. Our fake reuses the - same instance, so the merge should still leave the model_function_wrapper - key present after the call.""" - pos_pooled = torch.randn(1, 1280) - neg_pooled = torch.randn(1, 1280) - patcher = _fake_patcher(is_base=True, is_refiner=False) - patcher.model_options["pre_existing"] = "kept" - out = add_sdxl_model_options(patcher, _cond(pos_pooled), _cond(neg_pooled)) - assert "model_function_wrapper" in out.model_options - assert out.model_options["pre_existing"] == "kept" + assert captured["context_shape"] == (2, 77, 1280) + assert captured["time_ids_shape"] == (2, 5) diff --git a/tests/unit/test_chunks.py b/tests/unit/test_chunks.py index 318c720..8ead7ed 100644 --- a/tests/unit/test_chunks.py +++ b/tests/unit/test_chunks.py @@ -1,19 +1,23 @@ -import pytest +"""Original smoke tests, re-pointed at coreml_suite.core.* in Phase 3. +Drops the `from comfy.model_management import get_torch_device` import +(replaced with torch.device('cpu') so Tier 0 runs without ComfyUI) and the +dead `model_config` fixture (defined in Phase 1 but never consumed). +""" +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 -from coreml_suite.models import ( - CoreMLInputs, -) -from coreml_suite.config import ModelVersion, get_model_config +from coreml_suite.core.controlnet import chunk_control +from coreml_suite.core.inputs import CoreMLInputs +from coreml_suite.core.latents import chunk_batch, merge_chunks + + +CPU = torch.device("cpu") @pytest.fixture def expected_inputs(): - expected = { + return { "sample": {"shape": (2, 4, 64, 64)}, "timestep": {"shape": (2,)}, "timestep_cond": {"shape": (2, 256)}, @@ -21,17 +25,11 @@ def expected_inputs(): "additional_residual_0": {"shape": (2, 320, 64, 64)}, "additional_residual_1": {"shape": (2, 640, 32, 32)}, } - return expected - - -@pytest.fixture -def model_config(): - return get_model_config(ModelVersion.SD15) @pytest.mark.parametrize("batch_size", [1, 2, 4, 5, 9]) def test_batch_chunking(batch_size): - latent_image = torch.randn(batch_size, 4, 64, 64).to(get_torch_device()) + latent_image = torch.randn(batch_size, 4, 64, 64).to(CPU) target_shape = (4, 4, 64, 64) chunked = chunk_batch(latent_image, target_shape) @@ -45,7 +43,7 @@ def test_batch_chunking(batch_size): @pytest.mark.parametrize("batch_size", [1, 2, 4, 5, 9]) def test_merge_chunks(batch_size): - input_tensor = torch.randn(batch_size, 4, 64, 64).to(get_torch_device()) + input_tensor = torch.randn(batch_size, 4, 64, 64).to(CPU) target_shape = (4, 4, 64, 64) chunked = chunk_batch(input_tensor, target_shape) @@ -57,16 +55,16 @@ def test_merge_chunks(batch_size): @pytest.fixture def inputs(): - x = torch.randn(1, 4, 64, 64).to(get_torch_device()) - t = torch.randn([1]).to(get_torch_device()) - c_crossattn = torch.randn(1, 77, 768).to(get_torch_device()) + x = torch.randn(1, 4, 64, 64).to(CPU) + t = torch.randn([1]).to(CPU) + c_crossattn = torch.randn(1, 77, 768).to(CPU) control = { "output": [ - torch.randn(1, 320, 64, 64).to(get_torch_device()), - torch.randn(1, 640, 32, 32).to(get_torch_device()), + torch.randn(1, 320, 64, 64).to(CPU), + torch.randn(1, 640, 32, 32).to(CPU), ], } - timestep_cond = torch.randn(1, 256).to(get_torch_device()) + timestep_cond = torch.randn(1, 256).to(CPU) return CoreMLInputs(x, t, c_crossattn, control, timestep_cond=timestep_cond) @@ -86,11 +84,11 @@ def inputs(): def test_chunking_controlnet(b, target_size, num_chunks): cn = { "output": [ - torch.randn(b, 320, 64, 64).to(get_torch_device()), - torch.randn(b, 640, 32, 32).to(get_torch_device()), + torch.randn(b, 320, 64, 64).to(CPU), + torch.randn(b, 640, 32, 32).to(CPU), ], "middle": [ - torch.randn(b, 1280, 8, 8).to(get_torch_device()), + torch.randn(b, 1280, 8, 8).to(CPU), ], } diff --git a/tests/unit/test_tier0_purity.py b/tests/unit/test_tier0_purity.py new file mode 100644 index 0000000..ffa0690 --- /dev/null +++ b/tests/unit/test_tier0_purity.py @@ -0,0 +1,43 @@ +"""Phase 3 gate: prove the Tier-0 lane is framework-free. + +In a pure `pytest -m unit` run, none of the banned runtime modules +(comfy, coremltools, python_coreml_stable_diffusion, folder_paths, +nodes, comfy_extras, diffusers, diffusionkit) may be in sys.modules +after collection. If they are, a tests/unit/ file is transitively +pulling them in and the Tier-0 promise — "runs on Linux with no Mac +stack" — is broken. + +When other tiers (m2 / integration) are also collected, comfy is +expected in sys.modules (integration imports it deliberately), so the +check is skipped in mixed runs — Tier-0 purity is only meaningful when +nothing else is loaded. +""" +import sys + +import pytest + +BANNED_ROOTS = { + "comfy", + "comfy_extras", + "coremltools", + "python_coreml_stable_diffusion", + "folder_paths", + "nodes", + "diffusers", + "diffusionkit", +} + + +def test_no_framework_modules_loaded_by_unit_tier(request): + markexpr = request.config.option.markexpr + if markexpr != "unit": + pytest.skip( + "purity gate only meaningful in a pure `-m unit` run " + f"(got markexpr={markexpr!r}); other tiers are expected to " + "import comfy/coremltools." + ) + loaded = {name for name in sys.modules if name.split(".")[0] in BANNED_ROOTS} + assert not loaded, ( + f"Tier-0 leakage: these framework modules are in sys.modules after " + f"collecting tests/unit/: {sorted(loaded)}. Pure-core promise broken." + )