FLUX.1-schnell-OpenVINO-INT4 / inference_int4_flux.py
HelloSun's picture
Add inference_int4_flux.py
7fd6408 verified
Raw History Blame
2.03 kB
"""Minimal inference with the OpenVINO INT4 FLUX.1-schnell pipeline.
Usage:
python inference_int4_flux.py --prompt "a cat" --output out.png
"""
import argparse
import time
import torch
from optimum.intel import OVFluxPipeline
MODEL_PATH = "/home/user/app/flux-schnell-ov-int4"
def load_pipeline(model_path: str = MODEL_PATH, device: str = "CPU"):
t0 = time.perf_counter()
pipe = OVFluxPipeline.from_pretrained(model_path, compile=True, device=device)
load_s = time.perf_counter() - t0
return pipe, load_s
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--model_path", type=str, default=MODEL_PATH)
parser.add_argument("--prompt", type=str, default="A cinematic photo of a mountain lake at sunrise")
parser.add_argument("--negative_prompt", type=str, default="")
parser.add_argument("--output", type=str, default="output.png")
parser.add_argument("--width", type=int, default=1024)
parser.add_argument("--height", type=int, default=1024)
parser.add_argument("--steps", type=int, default=4)
parser.add_argument("--guidance_scale", type=float, default=0.0)
parser.add_argument("--max_sequence_length", type=int, default=256)
parser.add_argument("--seed", type=int, default=42)
args = parser.parse_args()
pipe, load_s = load_pipeline(args.model_path)
print(f"load+compile: {load_s:.2f}s")
generator = torch.Generator(device="cpu").manual_seed(args.seed)
t0 = time.perf_counter()
result = pipe(
prompt=args.prompt,
negative_prompt=args.negative_prompt,
width=args.width,
height=args.height,
num_inference_steps=args.steps,
guidance_scale=args.guidance_scale,
max_sequence_length=args.max_sequence_length,
generator=generator,
)
elapsed = time.perf_counter() - t0
result.images[0].save(args.output)
print(f"generated in {elapsed:.2f}s ({elapsed / args.steps:.2f}s/step) -> {args.output}")
if __name__ == "__main__":
main()