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)
115 lines
4.1 KiB
Python
115 lines
4.1 KiB
Python
"""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,
|
|
)
|
|
]
|