refactor(phase3): split pure logic into coreml_suite.core

Phase 3 of the modernization plan: move the framework-free math out of
the comfy-coupled modules so Tier-0 tests can run on plain Linux without
ComfyUI, coremltools, or python_coreml_stable_diffusion.

New pure-core package (no comfy / coreml / mps imports):
- coreml_suite.core.latents: chunk_batch, merge_chunks
- coreml_suite.core.controlnet: expand_inputs, no_control,
  extract_residual_kwargs, chunk_control
- coreml_suite.core.inputs: CoreMLInputs (chunks + coreml_kwargs)
- coreml_suite.core.sdxl: is_sdxl / is_sdxl_base / is_sdxl_refiner,
  build_sdxl_time_ids (base len 6, refiner len 5), build_sdxl_text_embeds,
  sdxl_model_function_wrapper
- coreml_suite.core.naming: compose_out_name, lora_names_from_params

Thin adapters keep the public import paths:
- coreml_suite.latents / coreml_suite.controlnet: re-export from core
- coreml_suite.models: CoreMLModelWrapper, CoreMLModelWrapperLCM,
  add_sdxl_model_options (now uses the pure builders from core.sdxl),
  get_latent_image, get_model_patcher remain framework-coupled
- coreml_suite.nodes: CoreMLConverter.convert now delegates the out_name
  composition to core.naming.compose_out_name

Test infra:
- tests/unit/* re-pointed at coreml_suite.core.*
- test_chunks.py dropped `from comfy.model_management import ...` and
  the dead `model_config` fixture (Phase 1 left it broken; Phase 3
  removes it entirely)
- test_characterization_sdxl_options now targets the pure builders
  directly via inspect.getclosurevars on the wrapper closure
- test_characterization_out_name now calls compose_out_name without the
  heavy CoreMLConverter monkey-patching that Phase 2 needed
- tests/unit/test_tier0_purity.py: new gate that fails if comfy /
  coremltools / etc leak into sys.modules during a pure `-m unit` run
  (skipped in mixed runs where m2 / integration legitimately import them)
- tests/__init__.py + top-level conftest.py + pyproject addopts
  `--import-mode=importlib --confcutdir=tests` together stop pytest from
  importing the repo-root `__init__.py` (the ComfyUI custom-node entry
  pulls in comfy)
- tests/conftest.py adds tier-aware collect_ignore so `-m unit` skips
  tests/m2 + tests/integration at collection time

Verification:
- `pytest -m unit tests/` → 88 passed in ~2s; deterministic across runs
- Tier-0 purity gate confirms no comfy/coreml/etc in sys.modules
- m2 golden image (Phase 2 anchor) still hashes identical → refactor
  produced bit-for-bit unchanged output
- `git diff main -- __init__.py coreml_suite/nodes.py` shows zero churn
  to NODE_CLASS_MAPPINGS keys or INPUT_TYPES field names (public
  workflow contract intact)
This commit is contained in:
aszc-dev
2026-05-22 16:08:01 +02:00
parent 04911d0052
commit 5dafd261b7
21 changed files with 737 additions and 575 deletions
+4
View File
@@ -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"]
+13 -61
View File
@@ -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",
]
+10
View File
@@ -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.
"""
+67
View File
@@ -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
+114
View File
@@ -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,
)
]
+42
View File
@@ -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]]
+49
View File
@@ -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])]
+91
View File
@@ -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
+3 -35
View File
@@ -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"]
+40 -188
View File
@@ -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
+9 -16
View File
@@ -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}")
+4
View File
@@ -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"]
View File
+25
View File
@@ -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:
@@ -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,
+1 -1
View File
@@ -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)
+1 -1
View File
@@ -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)
+96 -126
View File
@@ -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([]) == []
+100 -120
View File
@@ -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)
+24 -26
View File
@@ -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),
],
}
+43
View File
@@ -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."
)