HelloSun commited on
Commit
7fd6408
·
verified ·
1 Parent(s): d2213f0

Add inference_int4_flux.py

Browse files
Files changed (1) hide show
  1. inference_int4_flux.py +60 -0
inference_int4_flux.py ADDED
@@ -0,0 +1,60 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Minimal inference with the OpenVINO INT4 FLUX.1-schnell pipeline.
2
+
3
+ Usage:
4
+ python inference_int4_flux.py --prompt "a cat" --output out.png
5
+ """
6
+
7
+ import argparse
8
+ import time
9
+
10
+ import torch
11
+ from optimum.intel import OVFluxPipeline
12
+
13
+ MODEL_PATH = "/home/user/app/flux-schnell-ov-int4"
14
+
15
+
16
+ def load_pipeline(model_path: str = MODEL_PATH, device: str = "CPU"):
17
+ t0 = time.perf_counter()
18
+ pipe = OVFluxPipeline.from_pretrained(model_path, compile=True, device=device)
19
+ load_s = time.perf_counter() - t0
20
+ return pipe, load_s
21
+
22
+
23
+ def main() -> None:
24
+ parser = argparse.ArgumentParser()
25
+ parser.add_argument("--model_path", type=str, default=MODEL_PATH)
26
+ parser.add_argument("--prompt", type=str, default="A cinematic photo of a mountain lake at sunrise")
27
+ parser.add_argument("--negative_prompt", type=str, default="")
28
+ parser.add_argument("--output", type=str, default="output.png")
29
+ parser.add_argument("--width", type=int, default=1024)
30
+ parser.add_argument("--height", type=int, default=1024)
31
+ parser.add_argument("--steps", type=int, default=4)
32
+ parser.add_argument("--guidance_scale", type=float, default=0.0)
33
+ parser.add_argument("--max_sequence_length", type=int, default=256)
34
+ parser.add_argument("--seed", type=int, default=42)
35
+ args = parser.parse_args()
36
+
37
+ pipe, load_s = load_pipeline(args.model_path)
38
+ print(f"load+compile: {load_s:.2f}s")
39
+
40
+ generator = torch.Generator(device="cpu").manual_seed(args.seed)
41
+
42
+ t0 = time.perf_counter()
43
+ result = pipe(
44
+ prompt=args.prompt,
45
+ negative_prompt=args.negative_prompt,
46
+ width=args.width,
47
+ height=args.height,
48
+ num_inference_steps=args.steps,
49
+ guidance_scale=args.guidance_scale,
50
+ max_sequence_length=args.max_sequence_length,
51
+ generator=generator,
52
+ )
53
+ elapsed = time.perf_counter() - t0
54
+
55
+ result.images[0].save(args.output)
56
+ print(f"generated in {elapsed:.2f}s ({elapsed / args.steps:.2f}s/step) -> {args.output}")
57
+
58
+
59
+ if __name__ == "__main__":
60
+ main()