You need to agree to share your contact information to access this model

This repository is publicly accessible, but you have to accept the conditions to access its files and content.

Log in or Sign Up to review the conditions and access this model content.

YAML Metadata Warning:empty or missing yaml metadata in repo card

Check out the documentation for more information.

lora效果运行:

from pathlib import Path
from typing import Optional

import torch
import gradio as gr
from PIL import Image
from diffusers import QwenImageEditPlusPipeline


MODEL_PATH = "path/FireRed-Image-Edit-1.1"
LORA_PATH = "path/rotate60-1534"

DEVICE = "cuda:0"
DTYPE = torch.bfloat16

pipe = None


def resolve_lora_path(lora_path: str) -> Path:
    path = Path(lora_path).expanduser().resolve()
    if not path.exists():
        raise FileNotFoundError(f"LoRA path not found: {path}")

    if path.is_dir() and (path / "adapter").is_dir():
        return path / "adapter"

    return path


def attach_lora(
    pipeline: QwenImageEditPlusPipeline,
    lora_path: str,
) -> QwenImageEditPlusPipeline:
    resolved = resolve_lora_path(lora_path)

    if resolved.is_file():
        pipeline.load_lora_weights(
            str(resolved.parent),
            weight_name=resolved.name,
        )
        print(f"Loaded LoRA weights from file: {resolved}")
        return pipeline

    peft_weight_files = (
        "adapter_model.safetensors",
        "adapter_model.bin",
    )

    diffusers_weight_files = (
        "pytorch_lora_weights.safetensors",
        "pytorch_lora_weights.bin",
    )

    if any((resolved / name).exists() for name in peft_weight_files):
        from peft import PeftModel

        pipeline.transformer = PeftModel.from_pretrained(
            pipeline.transformer,
            str(resolved),
            is_trainable=False,
        )
        pipeline.transformer.eval()

        print(f"Loaded PEFT LoRA adapter from directory: {resolved}")
        return pipeline

    if any((resolved / name).exists() for name in diffusers_weight_files):
        pipeline.load_lora_weights(str(resolved))
        print(f"Loaded diffusers-format LoRA weights from directory: {resolved}")
        return pipeline

    raise FileNotFoundError(
        f"Unsupported LoRA directory layout: {resolved}"
    )


def load_pipeline(
    model_path: str = MODEL_PATH,
    lora_path: Optional[str] = LORA_PATH,
) -> QwenImageEditPlusPipeline:
    global pipe

    if pipe is not None:
        return pipe

    print("Loading FireRed-Image-Edit pipeline...")

    pipeline = QwenImageEditPlusPipeline.from_pretrained(
        model_path,
        torch_dtype=DTYPE,
    )

    if lora_path:
        pipeline = attach_lora(pipeline, lora_path)

    pipeline.to(DEVICE)
    pipeline.set_progress_bar_config(disable=None)

    pipe = pipeline

    print("Pipeline loaded.")
    return pipe


def infer(
    input_image: Image.Image,
    prompt: str,
    seed: int,
    true_cfg_scale: float,
    num_inference_steps: int,
):
    if input_image is None:
        raise gr.Error("Please upload an image.")

    if not prompt or not prompt.strip():
        raise gr.Error("Please input an editing prompt.")

    pipeline = load_pipeline()

    image = input_image.convert("RGB")
    print(f"Image size: {image.size[0]}x{image.size[1]}")

    generator = torch.Generator(device=DEVICE).manual_seed(int(seed))

    with torch.inference_mode():
        result = pipeline(
            image=image,
            prompt=prompt,
            generator=generator,
            true_cfg_scale=true_cfg_scale,
            negative_prompt=" ",
            num_inference_steps=int(num_inference_steps),
            num_images_per_prompt=1,
        )

    return result.images[0]


def build_demo():
    with gr.Blocks(title="FireRed-Image-Edit Single-Image Demo") as demo:
        gr.Markdown("# FireRed-Image-Edit Single-Image Demo")

        with gr.Row():
            with gr.Column():
                input_image = gr.Image(
                    label="Input Image",
                    type="pil",
                )

                prompt = gr.Textbox(
                    label="Editing Prompt",
                    lines=6,
                    value='把图中的物体按照俯视视角顺时针旋转60度',
                )

                seed = gr.Number(
                    label="Seed",
                    value=49,
                    precision=0,
                )

                true_cfg_scale = gr.Slider(
                    label="True CFG Scale",
                    minimum=1.0,
                    maximum=10.0,
                    value=4.0,
                    step=0.1,
                )

                num_inference_steps = gr.Slider(
                    label="Inference Steps",
                    minimum=1,
                    maximum=80,
                    value=40,
                    step=1,
                )

                run_btn = gr.Button("Run Inference")

            with gr.Column():
                output_image = gr.Image(
                    label="Edited Image",
                    type="pil",
                )

        run_btn.click(
            fn=infer,
            inputs=[
                input_image,
                prompt,
                seed,
                true_cfg_scale,
                num_inference_steps,
            ],
            outputs=output_image,
        )

    return demo

if __name__ == "__main__":
    load_pipeline()
    demo = build_demo()
    demo.launch(
        server_name="0.0.0.0",
        server_port=7860,
        share=False,
    )

firered原生效果运行:

"""FireRed-Image-Edit Gradio demo."""

import argparse
from pathlib import Path
from typing import Optional, Tuple

import gradio as gr
import torch
from PIL import Image
from diffusers import QwenImageEditPlusPipeline

from utils.fast_pipeline import load_fast_pipeline


MODEL_PATH = "path/FireRed-Image-Edit-1.1"
DEVICE = "cuda:2"
DTYPE = torch.bfloat16

LIGHTNING_LORA_DIR = Path(
    "path/FireRed-Image-Edit-1.0-Lightning"
)
DEFAULT_LIGHTNING_WEIGHT = (
    "FireRed-Image-Edit-1.0-Lightning-8steps-v1.1.safetensors"
)
LIGHTNING_STEPS = 8
LIGHTNING_CFG_SCALE = 1.0

pipe: Optional[QwenImageEditPlusPipeline] = None
pipe_cache_key: Optional[tuple] = None


def attach_lightning_lora(
    pipeline: QwenImageEditPlusPipeline,
    weight_name: str = DEFAULT_LIGHTNING_WEIGHT,
) -> QwenImageEditPlusPipeline:
    weight_path = LIGHTNING_LORA_DIR / weight_name
    if not weight_path.exists():
        raise FileNotFoundError(f"Lightning LoRA not found: {weight_path}")

    pipeline.load_lora_weights(str(LIGHTNING_LORA_DIR), weight_name=weight_name)
    print(f"Loaded Lightning LoRA: {weight_name}")
    return pipeline


def load_pipeline(
    model_path: str = MODEL_PATH,
    optimized: bool = False,
    lightning: bool = False,
    lightning_weight: str = DEFAULT_LIGHTNING_WEIGHT,
) -> QwenImageEditPlusPipeline:
    global pipe, pipe_cache_key

    cache_key = (model_path, optimized, lightning, lightning_weight)
    if pipe is not None and pipe_cache_key == cache_key:
        return pipe

    if optimized and lightning:
        raise ValueError("Lightning 模式与 Int8/Compile 优化模式不能同时启用。")

    print("Loading FireRed-Image-Edit pipeline...")
    if optimized:
        pipeline = load_fast_pipeline(model_path, device=DEVICE)
    else:
        pipeline = QwenImageEditPlusPipeline.from_pretrained(
            model_path,
            torch_dtype=DTYPE,
        )
        pipeline.to(DEVICE)
        pipeline.set_progress_bar_config(disable=None)

    if lightning:
        attach_lightning_lora(pipeline, lightning_weight)

    pipe = pipeline
    pipe_cache_key = cache_key
    print("Pipeline loaded.")
    return pipe


def on_lightning_toggle(enabled: bool):
    if enabled:
        return (
            gr.update(value=LIGHTNING_STEPS),
            gr.update(value=LIGHTNING_CFG_SCALE),
            gr.update(interactive=False),
        )
    return gr.update(), gr.update(), gr.update(interactive=True)


def infer(
    image: Optional[Image.Image],
    prompt: str,
    seed: int,
    true_cfg_scale: float,
    num_inference_steps: int,
    optimized: bool,
    lightning: bool,
    lightning_weight: str,
    model_path: str,
) -> Tuple[Image.Image, str]:
    if image is None:
        raise gr.Error("请上传输入图像。")

    if not prompt or not prompt.strip():
        raise gr.Error("请输入编辑提示词。")

    if optimized and lightning:
        raise gr.Error("Lightning 加速与 Int8/Compile 优化不能同时启用,请只选其一。")

    prompt = prompt.strip()
    images = [image.convert("RGB")]
    status_lines = ["已加载 1 张图像。"]

    pipeline = load_pipeline(
        model_path,
        optimized=optimized,
        lightning=lightning,
        lightning_weight=lightning_weight,
    )

    if lightning:
        status_lines.append(
            f"Lightning 加速:{lightning_weight},"
            f"{int(num_inference_steps)} 步,CFG={true_cfg_scale}"
        )
    elif optimized:
        status_lines.append(
            "已启用优化模式:首次推理在编译后可能需要 1-2 分钟。"
        )

    generator = torch.Generator(device=DEVICE).manual_seed(int(seed))

    inputs = {
        "image": images,
        "prompt": prompt,
        "generator": generator,
        "true_cfg_scale": true_cfg_scale,
        "negative_prompt": " ",
        "num_inference_steps": int(num_inference_steps),
        "num_images_per_prompt": 1,
    }

    with torch.inference_mode():
        result = pipeline(**inputs)

    status_lines.append("推理完成。")
    return result.images[0], "\n".join(status_lines)


def build_demo() -> gr.Blocks:
    with gr.Blocks(title="FireRed-Image-Edit") as demo:
        gr.Markdown(
            "# FireRed-Image-Edit\n"
            "上传图像并输入编辑提示词,生成编辑结果。\n\n"
            "勾选 **Lightning 8步加速** 可加载蒸馏 LoRA,约 8 步完成推理(推荐 CFG=1.0)。"
        )

        with gr.Row():
            with gr.Column():
                image = gr.Image(label="输入图像", type="pil")

                prompt = gr.Textbox(
                    label="编辑提示词",
                    lines=4,
                    value='"大动力机芯"修改为"High-Power Motor"',
                )

                with gr.Accordion("高级参数", open=False):
                    model_path = gr.Textbox(
                        label="模型路径",
                        value=MODEL_PATH,
                    )
                    lightning = gr.Checkbox(
                        label="Lightning 8步加速(推荐)",
                        value=True,
                    )
                    lightning_weight = gr.Dropdown(
                        label="Lightning LoRA 权重",
                        choices=sorted(
                            p.name for p in LIGHTNING_LORA_DIR.glob("*.safetensors")
                        ),
                        value=DEFAULT_LIGHTNING_WEIGHT,
                    )
                    seed = gr.Number(label="随机种子", value=49, precision=0)
                    true_cfg_scale = gr.Slider(
                        label="True CFG Scale",
                        minimum=1.0,
                        maximum=10.0,
                        value=LIGHTNING_CFG_SCALE,
                        step=0.1,
                    )
                    num_inference_steps = gr.Slider(
                        label="推理步数",
                        minimum=1,
                        maximum=80,
                        value=LIGHTNING_STEPS,
                        step=1,
                    )
                    optimized = gr.Checkbox(
                        label="启用优化模式(Int8 / Cache / Compile,与 Lightning 互斥)",
                        value=False,
                    )

                run_btn = gr.Button("开始编辑", variant="primary")

            with gr.Column():
                output_image = gr.Image(label="编辑结果", type="pil")
                status = gr.Textbox(label="状态", lines=8, interactive=False)

        lightning.change(
            fn=on_lightning_toggle,
            inputs=lightning,
            outputs=[num_inference_steps, true_cfg_scale, optimized],
        )

        run_btn.click(
            fn=infer,
            inputs=[
                image,
                prompt,
                seed,
                true_cfg_scale,
                num_inference_steps,
                optimized,
                lightning,
                lightning_weight,
                model_path,
            ],
            outputs=[output_image, status],
        )

    return demo


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="FireRed-Image-Edit Gradio demo")
    parser.add_argument("--model_path", type=str, default=MODEL_PATH)
    parser.add_argument("--optimized", action="store_true", default=False)
    parser.add_argument(
        "--lightning",
        action=argparse.BooleanOptionalAction,
        default=True,
    )
    parser.add_argument(
        "--lightning_weight",
        type=str,
        default=DEFAULT_LIGHTNING_WEIGHT,
    )
    parser.add_argument("--server_name", type=str, default="0.0.0.0")
    parser.add_argument("--server_port", type=int, default=7860)
    parser.add_argument("--share", action="store_true", default=False)
    return parser.parse_args()


if __name__ == "__main__":
    args = parse_args()
    load_pipeline(
        args.model_path,
        optimized=args.optimized,
        lightning=args.lightning,
        lightning_weight=args.lightning_weight,
    )
    demo = build_demo()
    demo.launch(
        server_name=args.server_name,
        server_port=args.server_port,
        share=args.share,
    )
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support