File size: 2,084 Bytes
28ecc36
 
0d93e34
69d8bc3
5ed432c
69d8bc3
 
28ecc36
 
 
 
 
 
 
 
 
69d8bc3
f020416
 
69d8bc3
 
 
 
28ecc36
 
5ed432c
69d8bc3
f020416
 
69d8bc3
f020416
69d8bc3
 
 
 
 
f020416
69d8bc3
 
 
5ed432c
69d8bc3
 
 
 
 
28ecc36
 
 
 
 
 
69d8bc3
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
50
51
52
53
54
# streamlit_app.py
import os
import streamlit as st
from PIL import Image
import torch
from diffusers import PaintByExamplePipeline

# 👉 Forcer les dossiers de cache dans /app/.cache (où tu as les droits)
os.environ["HF_HOME"] = "/app/.cache/huggingface"
os.environ["TRANSFORMERS_CACHE"] = "/app/.cache/transformers"
os.environ["DIFFUSERS_CACHE"] = "/app/.cache/diffusers"

os.makedirs("/app/.cache/huggingface", exist_ok=True)
os.makedirs("/app/.cache/transformers", exist_ok=True)
os.makedirs("/app/.cache/diffusers", exist_ok=True)

st.set_page_config(page_title="Paint-By-Example", layout="centered")

@st.cache_resource(show_spinner=False)
def load_pipe():
    dtype = torch.float16 if torch.cuda.is_available() else torch.float32
    pipe = PaintByExamplePipeline.from_pretrained(
        "Fantasy-Studio/Paint-by-Example",
        torch_dtype=dtype,
        cache_dir=os.environ["DIFFUSERS_CACHE"]  # 👉 important : on précise le cache
    )
    pipe = pipe.to("cuda" if torch.cuda.is_available() else "cpu")
    return pipe

pipe = load_pipe()

st.title("🧩 Paint-By-Example (édition par image de référence)")
img_file = st.file_uploader("Image d’entrée", type=["png","jpg","jpeg"])
mask_file = st.file_uploader("Masque (blanc = à modifier)", type=["png","jpg","jpeg"])
ref_file  = st.file_uploader("Image de référence (style/objet à copier)", type=["png","jpg","jpeg"])
steps = st.slider("Steps", 1, 100, 50)

if st.button("Appliquer"):
    if not (img_file and mask_file and ref_file):
        st.warning("Charge l’image, le masque et l’image de référence.")
    else:
        image = Image.open(img_file).convert("RGB")
        mask = Image.open(mask_file).convert("L")   # masque en niveaux de gris
        ref = Image.open(ref_file).convert("RGB")

        with st.spinner("Génération…"):
            out = pipe(
                image=image,
                mask_image=mask,
                example_image=ref,
                num_inference_steps=steps
            ).images[0]
        st.image(out, caption="Résultat", use_column_width=True)