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)
92 lines
2.7 KiB
Python
92 lines
2.7 KiB
Python
"""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
|