Spaces:
Running on Zero
Running on Zero
| import os | |
| import random | |
| import gradio as gr | |
| import numpy as np | |
| import spaces | |
| import torch | |
| from diffusers import Flux2KleinPipeline | |
| from PIL import Image | |
| BASE_MODEL = "black-forest-labs/FLUX.2-klein-9B" | |
| LORA_REPO = "dx8152/Flux2-Klein-9B-Migration" | |
| LORA_FILE = "Klein-Migration.safetensors" | |
| MAX_SEED = np.iinfo(np.int32).max | |
| PROMPT_PRESETS = { | |
| "Light migration": "将图1角色的光线替换为图2的光线", | |
| "Style migration": "将图1变为图2的画风", | |
| } | |
| DEFAULT_PRESET = "Light migration" | |
| dtype = torch.bfloat16 | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| pipe = Flux2KleinPipeline.from_pretrained( | |
| BASE_MODEL, | |
| torch_dtype=dtype, | |
| token=os.environ.get("HF_TOKEN"), | |
| ).to(device) | |
| pipe.load_lora_weights(LORA_REPO, weight_name=LORA_FILE, adapter_name="migration") | |
| pipe.set_adapters(["migration"], adapter_weights=[1.0]) | |
| pipe.fuse_lora(adapter_names=["migration"], lora_scale=1.0) | |
| pipe.unload_lora_weights() | |
| def infer( | |
| image, | |
| reference, | |
| prompt, | |
| seed=0, | |
| randomize_seed=True, | |
| num_inference_steps=4, | |
| width=0, | |
| height=0, | |
| progress=gr.Progress(track_tqdm=True), | |
| ): | |
| """ | |
| Migrate lighting or art style from a reference image (Image 2) | |
| onto a target image (Image 1) using FLUX.2 Klein 9B with the | |
| dx8152/Flux2-Klein-9B-Migration LoRA. | |
| Args: | |
| image: Target image (Image 1). | |
| reference: Reference image (Image 2) providing the lighting | |
| or style to migrate. | |
| prompt: Instruction describing the migration (Chinese prompts work | |
| best, e.g. 将图1角色的光线替换为图2的光线). | |
| seed: Random seed. Ignored when randomize_seed is True. | |
| randomize_seed: Pick a random seed per call. | |
| num_inference_steps: Denoising steps (Klein is distilled to 4). | |
| width: Output width. 0 = derive from Image 1. | |
| height: Output height. 0 = derive from Image 1. | |
| Returns: | |
| The edited image and the seed used. | |
| """ | |
| if image is None: | |
| raise gr.Error("Please upload a target image (Image 1).") | |
| if reference is None: | |
| raise gr.Error("Please upload a reference image (Image 2).") | |
| source = image.convert("RGB") | |
| ref = reference.convert("RGB") | |
| if not prompt or not prompt.strip(): | |
| raise gr.Error("Please provide a prompt.") | |
| if randomize_seed: | |
| seed = random.randint(0, MAX_SEED) | |
| generator = torch.Generator(device=device).manual_seed(seed) | |
| result = pipe( | |
| image=[source, ref], | |
| prompt=prompt, | |
| width=width if width else None, | |
| height=height if height else None, | |
| num_inference_steps=num_inference_steps, | |
| guidance_scale=1.0, | |
| generator=generator, | |
| ).images[0] | |
| return result, seed | |
| def apply_preset(preset): | |
| return PROMPT_PRESETS.get(preset, "") | |
| def update_dimensions(image): | |
| if image is None: | |
| return 0, 0 | |
| w, h = image.size | |
| if w * h > 1024 * 1024: | |
| scale = (1024 * 1024 / (w * h)) ** 0.5 | |
| w, h = int(w * scale), int(h * scale) | |
| return (w // 16) * 16, (h // 16) * 16 | |
| css = """ | |
| #col-container { max-width: 1100px; margin: 0 auto; } | |
| .dark .progress-text { color: white !important } | |
| #examples { max-width: 1100px; margin: 0 auto; } | |
| """ | |
| with gr.Blocks() as demo: | |
| with gr.Column(elem_id="col-container"): | |
| gr.Markdown("## 🪄 FLUX.2 Klein — Anything Migration") | |
| gr.Markdown(""" | |
| Migrate **lighting** or **art style** from a reference image onto your target image ⚡ | |
| Using [dx8152's Flux2-Klein-9B-Migration LoRA](https://huggingface.co/dx8152/Flux2-Klein-9B-Migration) | |
| on [FLUX.2 Klein 9B](https://huggingface.co/black-forest-labs/FLUX.2-klein-9B) — 4-step distilled inference 💨 | |
| """) | |
| with gr.Row(): | |
| with gr.Column(): | |
| with gr.Row(): | |
| image = gr.Image( | |
| label="Image 1 (Target)", | |
| type="pil", | |
| ) | |
| reference = gr.Image( | |
| label="Image 2 (Reference)", | |
| type="pil", | |
| ) | |
| preset = gr.Radio( | |
| label="Migration type", | |
| choices=list(PROMPT_PRESETS.keys()), | |
| value=DEFAULT_PRESET, | |
| ) | |
| prompt = gr.Textbox( | |
| label="Prompt", | |
| value=PROMPT_PRESETS[DEFAULT_PRESET], | |
| lines=2, | |
| info="Editable — Chinese prompts work best. 图1 = Image 1, 图2 = Image 2.", | |
| ) | |
| run_btn = gr.Button("🪄 Migrate", variant="primary", size="lg") | |
| with gr.Accordion("Advanced Settings", open=False): | |
| seed = gr.Slider(label="Seed", minimum=0, maximum=MAX_SEED, step=1, value=0) | |
| randomize_seed = gr.Checkbox(label="Randomize Seed", value=True) | |
| num_inference_steps = gr.Slider( | |
| label="Inference Steps", minimum=1, maximum=16, step=1, value=4 | |
| ) | |
| width = gr.Slider( | |
| label="Width (0 = auto from Image 1)", | |
| minimum=0, maximum=2048, step=16, value=0, | |
| ) | |
| height = gr.Slider( | |
| label="Height (0 = auto from Image 1)", | |
| minimum=0, maximum=2048, step=16, value=0, | |
| ) | |
| with gr.Column(): | |
| result = gr.Image(label="Output Image", interactive=False) | |
| seed_used = gr.Number(label="Seed Used", interactive=False) | |
| gr.Examples( | |
| examples=[ | |
| ["character_1.png", "light_1.png", PROMPT_PRESETS["Light migration"]], | |
| ["character_2.png", "light_4.png", PROMPT_PRESETS["Light migration"]], | |
| ["place_1.png", "light_6.png", PROMPT_PRESETS["Light migration"]], | |
| ["character_1.png", "metropolis.jpg", PROMPT_PRESETS["Style migration"]], | |
| ], | |
| inputs=[image, reference, prompt], | |
| outputs=[result, seed_used], | |
| fn=infer, | |
| cache_examples=True, | |
| cache_mode="lazy", | |
| elem_id="examples", | |
| ) | |
| preset.change(fn=apply_preset, inputs=[preset], outputs=[prompt]) | |
| image.change(fn=update_dimensions, inputs=[image], outputs=[width, height]) | |
| run_btn.click( | |
| fn=infer, | |
| inputs=[image, reference, prompt, seed, randomize_seed, num_inference_steps, width, height], | |
| outputs=[result, seed_used], | |
| ) | |
| demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=css, footer_links=["api", "gradio", "settings"]) | |