Files
aszc-dev-ComfyUI-CoreMLSuite/coreml_suite/core/latents.py
T
aszc 02b6e8ece3 feat: modernize toolchain, refactor core, add tiered CI and opt-in quantization
Modernizes ComfyUI-CoreMLSuite onto Python 3.12 / torch 2.7 / coremltools 9
with a characterization-test safety net. The default conversion path is
unchanged; existing saved workflows produce identical output.

- Toolchain bump (Python 3.12, torch 2.7, coremltools 9, numpy <2) with the
  blocking upstream pins overridden.
- Framework-free logic moved into coreml_suite/core/ (no comfy/coremltools
  imports); old module paths re-export from there.
- Opt-in quantize_nbits dropdown (none|8|6|4) for k-means weight
  palettization; default none is byte-for-byte identical to before.
- Tiered CI: Tier 0 (Linux unit), Tier 1 (macOS-ARM smoke), Tier 2
  (self-hosted Apple Silicon golden-image check on the ANE).
2026-05-25 19:11:49 +02:00

43 lines
1.3 KiB
Python

"""Pure batch-chunking helpers for Core ML's fixed-shape UNet inputs.
Re-exported by coreml_suite.latents. 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]]