File size: 2,490 Bytes
4f03424
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
"""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)