"""Stable Layers — decompose an image into editable back-to-front RGBA layers. LoRA: StabilityLabs/Stable-Layers (PEFT adapter, subfolder `model/`) Base: Qwen/Qwen-Image-Layered (QwenImageLayeredPipeline, 20.4B DiT) Inference follows the authors' reference script (decompose.py) exactly: Heun 2nd-order sampler, 50 steps, CFG 1.0 (off), 640 px max dim, 4 layers. The pipeline's own __call__ (Euler + true_cfg 4.0) is the *base model* recipe and garbles this LoRA, so the denoise loop is reimplemented here. """ import spaces # MUST be the first import (before torch / diffusers / peft) import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") import base64 import csv import inspect import io import random import tempfile import time import zipfile from pathlib import Path # gradio reads its cached-example log with a bare csv.reader, whose 128 KB per-field # default chokes on the inline layer-viewer markup ("field larger than field limit"). csv.field_size_limit(2**31 - 1) import gradio as gr # noqa: E402 import numpy as np import torch from PIL import Image from tqdm.auto import tqdm # --------------------------------------------------------------------------- # Constants (from the model card's "Recommended Inference Settings") # --------------------------------------------------------------------------- BASE_MODEL = "Qwen/Qwen-Image-Layered" LORA_REPO = "StabilityLabs/Stable-Layers" LORA_SUBFOLDER = "model" AOTI_REPO = "multimodalart/stable-layers-aoti" # precompiled QwenImageTransformerBlock STEPS = 50 # locked — fewer steps garbles the decomposition GUIDANCE = 1.0 # locked — CFG off RESOLUTION = 640 # locked — max dim; higher resolution garbles DEFAULT_PROMPT = "a clean, well composed image" PROMPT_TEMPLATE = ( "<|im_start|>system\nDescribe the image by detailing the color, shape, size, " "texture, quantity, text, spatial relationships of the objects and background:" "<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" ) PROMPT_TEMPLATE_DROP_IDX = 34 VAE_SCALE_FACTOR = 8 MAX_SEED = np.iinfo(np.int32).max # --------------------------------------------------------------------------- # Latent plumbing — ported 1:1 from the reference decompose.py # --------------------------------------------------------------------------- def rgb_to_rgba(image): b, _, h, w = image.shape alpha = torch.ones(b, 1, h, w, device=image.device, dtype=image.dtype) return torch.cat([image, alpha], dim=1) def normalize_latents(latents, vae): mean = torch.tensor(vae.config.latents_mean).view(1, -1, 1, 1, 1).to(latents.device, latents.dtype) std = torch.tensor(vae.config.latents_std).view(1, -1, 1, 1, 1).to(latents.device, latents.dtype) return (latents - mean) / std def denormalize_latents(latents, vae): mean = torch.tensor(vae.config.latents_mean).view(1, -1, 1, 1, 1).to(latents.device, latents.dtype) inv_std = (1.0 / torch.tensor(vae.config.latents_std)).view(1, -1, 1, 1, 1).to(latents.device, latents.dtype) return latents / inv_std + mean def pack_latents(latents, batch_size, num_channels, height, width, num_frames): latents = latents.view(batch_size, num_frames, num_channels, height // 2, 2, width // 2, 2) latents = latents.permute(0, 1, 3, 5, 2, 4, 6) return latents.reshape(batch_size, num_frames * (height // 2) * (width // 2), num_channels * 4) def unpack_latents(latents, height, width, num_layers, vae_scale_factor=VAE_SCALE_FACTOR): batch_size, _, channels = latents.shape frames = num_layers + 1 h = 2 * (int(height) // (vae_scale_factor * 2)) w = 2 * (int(width) // (vae_scale_factor * 2)) latents = latents.view(batch_size, frames, h // 2, w // 2, channels // 4, 2, 2) latents = latents.permute(0, 1, 4, 2, 5, 3, 6) latents = latents.reshape(batch_size, frames, channels // 4, h, w) return latents.permute(0, 2, 1, 3, 4) @torch.no_grad() def encode_condition_image(vae, image): """RGB image tensor in [-1,1], (B,3,H,W) -> packed condition latents.""" b = image.shape[0] vae_input_ch = getattr(vae.config, "input_channels", 3) if vae_input_ch == 4 and image.shape[1] == 3: image = rgb_to_rgba(image) elif vae_input_ch == 3 and image.shape[1] == 4: image = image[:, :3] image_5d = image.unsqueeze(2).to(dtype=vae.dtype) dist = vae.encode(image_5d) if hasattr(dist, "latent_dist"): latents = dist.latent_dist.mode() elif hasattr(dist, "mode"): latents = dist.mode() else: latents = dist latents = normalize_latents(latents, vae) z_dim, lh, lw = latents.shape[1], latents.shape[3], latents.shape[4] latents = latents.permute(0, 2, 1, 3, 4) return pack_latents(latents, b, z_dim, lh, lw, 1).to(dtype=torch.bfloat16) @torch.no_grad() def decode_layers(vae, latents, height, width, num_layers): b = latents.shape[0] unpacked = unpack_latents(latents, height, width, num_layers) unpacked = denormalize_latents(unpacked, vae) layer_latents = unpacked[:, :, 1:].permute(0, 2, 1, 3, 4) # drop the base frame _, _, c, h, w = layer_latents.shape decoded = vae.decode( layer_latents.reshape(b * num_layers, c, 1, h, w).to(dtype=vae.dtype), return_dict=False )[0] decoded = decoded.squeeze(2).float().clamp(-1.0, 1.0) _, c_out, h_out, w_out = decoded.shape return decoded.reshape(b, num_layers, c_out, h_out, w_out) def composite_layers(layers): """Porter-Duff 'over', back-to-front.""" if layers.shape[2] == 3: return layers.mean(dim=1) result = torch.zeros_like(layers[:, 0, :3]) for i in range(layers.shape[1]): rgb = layers[:, i, :3] alpha = (layers[:, i, 3:4] + 1.0) / 2.0 result = rgb * alpha + result * (1.0 - alpha) return result def tensor_to_pil(t): arr = ((t.clamp(-1, 1) + 1) / 2 * 255).byte().permute(1, 2, 0).cpu().numpy() return Image.fromarray(arr, mode="RGB") def layer_to_pil_rgba(layer): layer = layer.clamp(-1, 1).float() rgba = torch.cat([(layer[:3] + 1) / 2, (layer[3:4] + 1) / 2], dim=0) arr = (rgba * 255).byte().permute(1, 2, 0).cpu().numpy() return Image.fromarray(arr, mode="RGBA") def compute_aspect_resize(orig_w, orig_h, max_size): """Keep aspect ratio, max dim = max_size, both dims multiples of 16.""" scale = max_size / max(orig_w, orig_h) return ( max(int(round(orig_w * scale / 16)) * 16, 16), max(int(round(orig_h * scale / 16)) * 16, 16), ) # --------------------------------------------------------------------------- # Heun denoising — ported 1:1 from decompose.py # --------------------------------------------------------------------------- def transformer_forward(transformer, hidden_states, timestep, encoder_hidden_states, img_shapes, encoder_hidden_states_mask=None, additional_t_cond=None): return transformer( hidden_states=hidden_states, timestep=timestep / 1000, encoder_hidden_states=encoder_hidden_states, encoder_hidden_states_mask=encoder_hidden_states_mask, img_shapes=img_shapes, guidance=None, additional_t_cond=additional_t_cond, return_dict=False, )[0] def build_img_shapes(num_layers, packed_h, packed_w): # (num_layers + 1) generated frames + 1 condition frame return [[*[(1, packed_h, packed_w) for _ in range(num_layers + 1)], (1, packed_h, packed_w)]] @torch.no_grad() def denoise(transformer, scheduler, latents, timesteps, prompt_embeds, prompt_mask, img_shapes, condition_latents, additional_t_cond): """Heun's method (2nd order) over the scheduler's sigmas. CFG is off (1.0).""" gen_seq_len = latents.shape[1] def velocity(lat, t_val): ts = t_val.expand(1).to(lat.dtype) model_input = torch.cat([lat, condition_latents], dim=1) try: out = transformer_forward( transformer, model_input, ts, prompt_embeds, img_shapes, prompt_mask, additional_t_cond, ) except Exception as exc: # noqa: BLE001 # A stale AoTI package (torch / diffusers drift) blows up on the first # forward. Swap the compiled blocks out and carry on eagerly rather than # failing the request — costs one wasted forward, not a whole rerun. if not AOTI_ACTIVE: raise print(f"[gpu] AoTI forward failed ({exc!r}) -> eager", flush=True) _drop_aoti() out = transformer_forward( transformer, model_input, ts, prompt_embeds, img_shapes, prompt_mask, additional_t_cond, ) return out[:, :gen_seq_len] sigmas = scheduler.sigmas for i, t in enumerate(tqdm(timesteps, desc="Heun steps")): sigma = sigmas[i] sigma_next = sigmas[i + 1] if i + 1 < len(sigmas) else torch.tensor(0.0, device=t.device) dt = sigma_next - sigma v1 = velocity(latents, t) latents_mid = latents + dt * v1 if sigma_next > 0: # Heun corrector v2 = velocity(latents_mid, sigma_next * 1000) latents = latents + dt * 0.5 * (v1 + v2) else: latents = latents_mid return latents # --------------------------------------------------------------------------- # Load base pipeline + fuse the LoRA, at module scope (ZeroGPU rule 2) # --------------------------------------------------------------------------- print(f"[load] base pipeline {BASE_MODEL}", flush=True) from diffusers import QwenImageLayeredPipeline # noqa: E402 pipe = QwenImageLayeredPipeline.from_pretrained(BASE_MODEL, torch_dtype=torch.bfloat16) # The adapter is a raw PEFT checkpoint on QwenImageTransformer2DModel: its keys have no # `transformer.` prefix, so diffusers' load_lora_weights finds nothing and silently # no-ops. PEFT + merge_and_unload is the path that actually applies it — and merging # also keeps module FQNs identical to the plain transformer, which the AoTI package # below depends on. torch_device="cpu" stops PEFT's infer_device() from poking the # hijacked cuda device at import time. print(f"[load] LoRA {LORA_REPO}/{LORA_SUBFOLDER} (fusing on CPU)", flush=True) from peft import PeftModel # noqa: E402 _peft = PeftModel.from_pretrained( pipe.transformer, LORA_REPO, subfolder=LORA_SUBFOLDER, torch_device="cpu" ) pipe.transformer = _peft.merge_and_unload() del _peft pipe.transformer.eval().requires_grad_(False) pipe.vae.eval().requires_grad_(False) pipe.text_encoder.eval().requires_grad_(False) IN_CHANNELS = getattr(pipe.transformer.config, "in_channels", 64) SUPPORTS_SIGMAS = "sigmas" in set(inspect.signature(pipe.scheduler.set_timesteps).parameters) pipe.to("cuda") print("[load] pipeline on cuda (hijacked)", flush=True) # Ahead-of-time-compiled QwenImageTransformerBlock (weights stay runtime inputs, so the # same graph serves the LoRA-fused weights). Falls back to eager if the artifact can't # be loaded; a cheap canary in the GPU fn below catches a package that loads but can't run. AOTI_ACTIVE = False try: spaces.aoti_blocks_load(pipe.transformer, AOTI_REPO) AOTI_ACTIVE = True print(f"[load] AoTI blocks loaded from {AOTI_REPO}", flush=True) except Exception as exc: # noqa: BLE001 print(f"[load] AoTI unavailable ({exc!r}); running eager", flush=True) def _drop_aoti(): """Best-effort revert to eager forwards (aoti patches instance attributes).""" global AOTI_ACTIVE for blk in pipe.transformer.transformer_blocks: blk.__dict__.pop("forward", None) AOTI_ACTIVE = False @torch.no_grad() def encode_prompt(text): """Qwen2.5-VL prompt encoding with the template prefix dropped (as in decompose.py).""" tok = pipe.tokenizer([PROMPT_TEMPLATE.format(text)], padding=True, return_tensors="pt") tok = tok.to(pipe.text_encoder.device) out = pipe.text_encoder( input_ids=tok.input_ids, attention_mask=tok.attention_mask, output_hidden_states=True ) hidden = out.hidden_states[-1] mask_bool = tok.attention_mask.bool() lengths = mask_bool.sum(dim=1) chunks = [c[PROMPT_TEMPLATE_DROP_IDX:] for c in torch.split(hidden[mask_bool], lengths.tolist(), dim=0)] max_len = max(c.size(0) for c in chunks) embeds = torch.stack([torch.cat([c, c.new_zeros(max_len - c.size(0), c.size(1))]) for c in chunks]) mask = torch.stack([ torch.cat([ torch.ones(c.size(0), dtype=torch.long, device=c.device), torch.zeros(max_len - c.size(0), dtype=torch.long, device=c.device), ]) for c in chunks ]) return embeds.to(dtype=torch.bfloat16, device="cuda"), mask.to("cuda") print("[load] ready", flush=True) # --------------------------------------------------------------------------- # Interactive layer-stack widget (behaviour lives in PANEL_JS, below) # --------------------------------------------------------------------------- def _data_uri(pil, max_side=None): """Inline an RGBA layer for the viewer. WebP keeps alpha at roughly a tenth of PNG's size, which matters when four layers ride inline in one HTML payload.""" img = pil if max_side: img = pil.copy() img.thumbnail((max_side, max_side), Image.LANCZOS) buf = io.BytesIO() img.save(buf, format="WEBP", quality=90, method=4) return "data:image/webp;base64," + base64.b64encode(buf.getvalue()).decode("ascii") PLACEHOLDER_HTML = """
🗂️
Your layer stack will appear here.
Background + object layers, each a real RGBA image. Hide them, drag them around, and watch the composite rebuild itself.
""" def build_layer_panel(layer_pils, coverage, width, height): """Return the HTML for the stacked/toggleable/draggable layer viewer.""" n = len(layer_pils) # Pre-select the front-most layer that actually has content — the model often # leaves the top slots blank, and selecting a blank one makes drag look broken. sel_i = next((i for i in range(n - 1, 0, -1) if coverage[i] >= 0.5), 0) stack, rows = [], [] for i, pil in enumerate(layer_pils): stack.append( f'layer {i}' ) for i in range(n - 1, -1, -1): # layers panel: front-most on top name = "Background" if i == 0 else f"Layer {i}" pct = coverage[i] empty = pct < 0.5 tag = "empty" if empty else f"{pct:.0f}% cover" cls = " sl-sel" if i == sel_i else "" cls += " sl-blank" if empty else "" rows.append( f'
' f'' f'' f'{name}{tag}
' ) return ( f'
' f'
' f'
{"".join(stack)}
' f'
Click a layer to select it, then drag on the canvas to move it. ' f'Toggle visible to delete it from the composite.
' f'
' f'
Layers
{"".join(rows)}' f'
' f'
' ) # --------------------------------------------------------------------------- # Inference # --------------------------------------------------------------------------- def _estimate_duration(*args, **kwargs): """Measured on ZeroGPU xlarge at 640 px with the AoTI blocks: L=4 (6 frames) 113 s, L=6 (8 frames) 178 s. Attention makes it superlinear in the frame count, so fit 8.6·f + 1.72·f² and add the ~16 s encode/decode/fork overhead. """ n = 4 if len(args) > 1: try: n = int(args[1]) except (TypeError, ValueError): n = 4 f = max(2, min(6, n)) + 2 est = 16.0 + 8.6 * f + 1.72 * f * f if not AOTI_ACTIVE: est *= 1.6 # the eager fallback is materially slower return int(est * 1.15) + 1 @spaces.GPU(duration=_estimate_duration, size="xlarge") def decompose(input_image, num_layers=4, seed=42, randomize_seed=False, prompt="", progress=gr.Progress(track_tqdm=True)): """Decompose an image into back-to-front editable RGBA layers. Args: input_image: the source image to split into layers. num_layers: how many layers to decompose into (4 is the recommended default). seed: RNG seed for the initial noise. randomize_seed: pick a fresh random seed for this run. prompt: optional caption nudge; the default works well for most images. Returns: The interactive layer-stack HTML, a gallery of every layer in order, a gallery of the layers, a ZIP of the RGBA PNGs, and the seed that was actually used. """ if input_image is None: raise gr.Error("Please upload an image first.") t0 = time.perf_counter() num_layers = int(max(2, min(6, int(num_layers)))) seed = random.randint(0, MAX_SEED) if randomize_seed else int(seed) pil = Image.fromarray(input_image) if isinstance(input_image, np.ndarray) else input_image if isinstance(pil, str): pil = Image.open(pil) image = pil.convert("RGB") tw, th = compute_aspect_resize(*image.size, RESOLUTION) source = image.resize((tw, th), Image.LANCZOS) packed_h, packed_w = (th // VAE_SCALE_FACTOR) // 2, (tw // VAE_SCALE_FACTOR) // 2 img_shapes = build_img_shapes(num_layers, packed_h, packed_w) additional_t_cond = torch.zeros(1, dtype=torch.long, device="cuda") # 0 = RGB input img_t = torch.from_numpy(np.array(source)).permute(2, 0, 1).float() / 127.5 - 1.0 condition_latents = encode_condition_image(pipe.vae, img_t.unsqueeze(0).to("cuda")) p_embeds, p_mask = encode_prompt(prompt.strip() if prompt and prompt.strip() else DEFAULT_PROMPT) # Fresh scheduler per call: set_timesteps mutates state and handlers run concurrently. scheduler = pipe.scheduler.__class__.from_config(pipe.scheduler.config) mu = (condition_latents.shape[1] / (256 * 256 / 16 / 16)) ** 0.5 if SUPPORTS_SIGMAS: scheduler.set_timesteps(sigmas=np.linspace(1.0, 0, STEPS + 1)[:-1], device="cuda", mu=mu) else: scheduler.set_timesteps(num_inference_steps=STEPS, device="cuda") gen = torch.Generator(device="cuda").manual_seed(seed) seq_len = (num_layers + 1) * packed_h * packed_w latents = torch.randn(1, seq_len, IN_CHANNELS, device="cuda", dtype=torch.bfloat16, generator=gen) t_denoise = time.perf_counter() latents = denoise( pipe.transformer, scheduler, latents, scheduler.timesteps, p_embeds, p_mask, img_shapes, condition_latents, additional_t_cond, ) print(f"[gpu] denoise {time.perf_counter() - t_denoise:.1f}s " f"(L={num_layers}, {tw}x{th}, aoti={AOTI_ACTIVE})", flush=True) decoded = decode_layers(pipe.vae, latents, th, tw, num_layers) composite = tensor_to_pil(composite_layers(decoded)[0]) layers = [layer_to_pil_rgba(decoded[0, i]) for i in range(num_layers)] coverage = [ float((np.asarray(l)[..., 3] > 12).mean() * 100.0) for l in layers ] # ZIP of the real RGBA PNGs + the recomposite (per the reference script's layout) workdir = Path(tempfile.mkdtemp(prefix="stable_layers_")) zip_path = workdir / f"stable_layers_seed{seed}.zip" with zipfile.ZipFile(zip_path, "w", zipfile.ZIP_DEFLATED) as zf: for name, im in [("source.png", source), ("composite.png", composite)] + [ (f"layer_{i}.png", l) for i, l in enumerate(layers) ]: p = workdir / name im.save(p) zf.write(p, name) # One entry per layer, in back-to-front order, for the browsable gallery # sitting directly below the custom layer viewer. per_layer_gallery = [ (l, ("background" if i == 0 else f"layer {i}") + (" · empty" if coverage[i] < 0.5 else f" · {coverage[i]:.0f}% cover")) for i, l in enumerate(layers) ] gallery = [(composite, "composite (all layers)")] + per_layer_gallery print(f"[gpu] total {time.perf_counter() - t0:.1f}s", flush=True) return build_layer_panel(layers, coverage, tw, th), per_layer_gallery, gallery, str(zip_path), seed # --------------------------------------------------------------------------- # UI # --------------------------------------------------------------------------- CSS = """ #col-container { max-width: 1360px; margin: 0 auto; } .dark .gradio-container { color: var(--body-text-color); } .sl-panel { display: flex; gap: 14px; align-items: flex-start; flex-wrap: wrap; } .sl-panel .sl-stagewrap { flex: 1 1 380px; min-width: 300px; } .sl-panel .sl-stage { position: relative; width: 100%; overflow: hidden; border-radius: 10px; border: 1px solid var(--border-color-primary); touch-action: none; cursor: grab; background-color: #fff; background-image: linear-gradient(45deg, #e3e3e3 25%, transparent 25%, transparent 75%, #e3e3e3 75%), linear-gradient(45deg, #e3e3e3 25%, transparent 25%, transparent 75%, #e3e3e3 75%); background-size: 22px 22px; background-position: 0 0, 11px 11px; } .sl-panel .sl-stage:active { cursor: grabbing; } .sl-panel .sl-lyr { position: absolute; inset: 0; width: 100%; height: 100%; object-fit: fill; user-select: none; -webkit-user-drag: none; transition: opacity .12s ease; } .sl-panel .sl-lyr.sl-off { opacity: 0; pointer-events: none; } .sl-panel .sl-hint { font-size: 12px; opacity: .65; margin-top: 6px; line-height: 1.4; } .sl-panel .sl-side { flex: 0 0 226px; display: flex; flex-direction: column; gap: 6px; } .sl-panel .sl-sidehead { font-size: 11px; letter-spacing: .09em; text-transform: uppercase; opacity: .6; margin-bottom: 2px; } .sl-panel .sl-row { display: flex; align-items: center; gap: 8px; padding: 5px 7px; border-radius: 9px; border: 1px solid var(--border-color-primary); cursor: pointer; background: var(--background-fill-secondary); } .sl-panel .sl-row.sl-sel { border-color: var(--color-accent); box-shadow: 0 0 0 1px var(--color-accent) inset; } .sl-panel .sl-row.sl-dim, .sl-panel .sl-row.sl-blank { opacity: .42; } .sl-panel .sl-th { width: 44px; height: 44px; object-fit: contain; border-radius: 6px; flex: 0 0 44px; background-color: #fff; background-image: linear-gradient(45deg, #e3e3e3 25%, transparent 25%, transparent 75%, #e3e3e3 75%), linear-gradient(45deg, #e3e3e3 25%, transparent 25%, transparent 75%, #e3e3e3 75%); background-size: 12px 12px; background-position: 0 0, 6px 6px; } .sl-panel .sl-name { font-size: 13px; line-height: 1.25; display: flex; flex-direction: column; } .sl-panel .sl-name em { font-style: normal; font-size: 11px; opacity: .55; } .sl-panel .sl-eye, .sl-panel .sl-reset { font-size: 0; border: 1px solid var(--border-color-primary); border-radius: 6px; background: var(--background-fill-primary); cursor: pointer; padding: 0; } .sl-panel .sl-eye { width: 22px; height: 22px; flex: 0 0 22px; box-sizing: border-box; font-size: 12px; line-height: 20px; padding-top: 0; padding-left: 8px; overflow: hidden; white-space: nowrap; text-align: left; } .sl-panel .sl-eye.sl-hid { opacity: .5; } .sl-panel .sl-reset { font-size: 12px; padding: 6px 8px; margin-top: 4px; } .sl-empty { text-align: center; padding: 46px 18px; opacity: .7; line-height: 1.6; } .sl-empty-icon { font-size: 34px; margin-bottom: 6px; } .sl-empty-sub { font-size: 12px; max-width: 330px; margin: 4px auto 0; } """ # Behaviour for the layer panel. gr.HTML renders via innerHTML, which never executes # inline