Image21-MLX-8bit / scripts /runtime.py
ixim's picture
Add files using upload-large-folder tool
4f03424 verified
Raw History Blame Contribute Delete
2.49 kB
"""Pinned MLX runtime adapter preserving the source model's RGBA output."""
from pathlib import Path
import mlx.core as mx
import numpy as np
from PIL import Image
from scripts.mlx_pipeline import QwenImagePipeline
def enable_progress():
"""Materialize at the existing Euler boundary and show progress every ten steps."""
from scripts import mlx_pipeline as module
from mlx_vlm.models.qwen_image.scheduler import FlowMatchEulerDiscreteScheduler
import time
class ProgressScheduler(FlowMatchEulerDiscreteScheduler):
def step(self,*,noise,step_index,latents):
value=super().step(noise=noise,step_index=step_index,latents=latents)
mx.eval(value)
if step_index==0 or (step_index+1)%10==0:
print(f' denoise step {step_index+1} at {time.strftime("%H:%M:%S")}',flush=True)
return value
module.FlowMatchEulerDiscreteScheduler=ProgressScheduler
def load(model,phase_offload=True):
return QwenImagePipeline.from_pretrained(model_path=Path(model),download=False,phase_offload=phase_offload)
def generate(pipe,prompt,*,seed=42,steps=40,width=1024,height=1024,inputs=None,
source_seed=None,resolution=1024):
if steps<1 or width<256 or height<256 or width%32 or height%32:
raise ValueError('Use positive steps and dimensions >=256 divisible by 32')
if inputs:
if source_seed is not None and seed==source_seed:
raise ValueError('Editing must use a different seed from the source image')
if len(inputs)>10: raise ValueError('At most ten references are supported')
result=pipe.edit_array(prompt,inputs,seed=seed,steps=steps,width=width,height=height,
guidance=1.0,output_resolution=resolution,use_kv_cache=True)
else:
# The pinned upstream generate_array slices RGB. Call its unchanged RGBA
# sampler so transparent generation keeps all four native output channels.
print('Encoding prompt',flush=True)
pipe.activate('text_encoder')
emb=pipe.text_encoder.encode(prompt).astype(mx.bfloat16)
mx.eval(emb)
print('Denoising',flush=True)
result=pipe._sample(emb,None,seed=seed,steps=steps,width=width,height=height,guidance=1.0,use_kv_cache=True)
mx.eval(result)
pixels=np.asarray(result)
pipe.release()
if pixels.shape!=(height,width,4): raise ValueError(f'Unexpected output: {pixels.shape}')
return Image.fromarray(pixels)