Team Ai
Apppublic

renderfy/FitStudioAI

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
streamlit_app.py498 linesDownload Raw Back to root
1# streamlit_app.py — Fit Studio AI (Fashion & Apparel) v0.6.5 UI2# Fix: "Video (loop)" duplicate — single placeholder (st.empty) updated in-place.3 4import os, io, zipfile, base64, requests, streamlit as st, uuid, tempfile5from pathlib import Path6from PIL import Image7 8# ================= Env & API =================9def _env(k): return (os.getenv(k) or "").strip().strip("'\"")10 11API_BASE = (12    _env("FITSTUDIO_API")13    or _env("AI_LIGHTBOX_API")14    or _env("LUXFIT_API")15).rstrip("/") if (16    _env("FITSTUDIO_API") or _env("AI_LIGHTBOX_API") or _env("LUXFIT_API")17) else ""18 19HF_TOKEN = _env("FITSTUDIO_TOKEN") or _env("AI_LIGHTBOX_TOKEN") or _env("LUXFIT_TOKEN")20HEADERS = {"Authorization": f"Bearer {HF_TOKEN}"} if HF_TOKEN else {}21 22# === Tile geometry (strict 9:16) ===23TILE_H = int(os.getenv("TILE_H", "640"))24TILE_W = max(200, int(round(TILE_H * 9 / 16)))25PREVIEW_BG = (255, 255, 255)26 27GIF_MAX_W = int(os.getenv("GIF_MAX_W", "512"))28GIF_FPS   = int(os.getenv("GIF_FPS", "12"))29GIF_MAX_FRAMES = GIF_FPS * 1230 31# ================= Session State =================32defaults = {"results": None, "job_counter": 0, "duration": "6", "rand": uuid.uuid4().hex[:8]}33for k, v in defaults.items():34    if k not in st.session_state:35        st.session_state[k] = v36 37# ================= CSS =================38st.markdown(f"""39<style>40h4.tile-title {{ margin: 0 0 8px 0; font-weight: 600; }}41div.tile-box {{42  width: {TILE_W}px; height: {TILE_H}px;43  display: flex; align-items: center; justify-content: center;44  background: #ffffff; border-radius: 12px; overflow: hidden;45  box-shadow: 0 1px 6px rgba(0,0,0,.08);46  margin: 0 0 6px 0;47}}48img.tile-media, video.tile-media {{49  width: 100%; height: 100%; object-fit: contain; background: #ffffff;50  display: block;51}}52.small-cap {{ font-size: 12px; color: #6b7280; margin: 4px 0 8px 2px; }}53</style>54""", unsafe_allow_html=True)55 56# ================= Helpers =================57def _needs_auth(url: str) -> bool:58    return (API_BASE and (url.startswith(API_BASE) or url.startswith("outputs/")))59 60@st.cache_data(ttl=300, show_spinner=False)61def fetch_bytes(url_or_path: str | None):62    if not url_or_path:63        return None64    try:65        p = Path(url_or_path)66        if p.is_file():67            return p.read_bytes()68        url = f"{API_BASE}/{url_or_path}" if url_or_path.startswith("outputs/") else url_or_path69        headers = (HEADERS if (_needs_auth(url) and HF_TOKEN) else None)70        r = requests.get(url, headers=headers, timeout=180)71        r.raise_for_status()72        return r.content73    except Exception:74        return None75 76@st.cache_data(ttl=300, show_spinner=False)77def image_to_9_16_canvas(img_bytes: bytes, w: int = TILE_W, h: int = TILE_H, bg_rgb=PREVIEW_BG) -> bytes:78    im = Image.open(io.BytesIO(img_bytes)).convert("RGBA")79    iw, ih = im.size80    scale = min(w / iw, h / ih)81    nw, nh = max(1, int(iw*scale)), max(1, int(ih*scale))82    im_resized = im.resize((nw, nh), Image.LANCZOS)83    canvas = Image.new("RGBA", (w, h), (*bg_rgb, 255))84    off = ((w - nw)//2, (h - nh)//2)85    canvas.paste(im_resized, off, im_resized)86    out = io.BytesIO(); canvas.save(out, format="PNG"); return out.getvalue()87 88@st.cache_data(ttl=300, show_spinner=False)89def image_size_from_bytes(b: bytes):90    try:91        im = Image.open(io.BytesIO(b)); return im.size92    except Exception:93        return None94 95def infer_ext(data: bytes) -> str:96    try:97        im = Image.open(io.BytesIO(data))98        fmt = (im.format or "JPEG").lower()99        return "." + {"jpeg":"jpg","jpg":"jpg","png":"png","webp":"webp"}.get(fmt, "jpg")100    except Exception:101        return ".jpg"102 103def unique_key(prefix: str) -> str:104    return f"{prefix}_{st.session_state['job_counter']}_{st.session_state['rand']}_{uuid.uuid4().hex[:6]}"105 106def _b64(data: bytes, mime: str) -> str:107    return f"data:{mime};base64," + base64.b64encode(data).decode("ascii")108 109def show_image_tile(col, title: str, src: str|bytes|None, filename_stub="image", key_prefix=""):110    col.markdown(f"<h4 class='tile-title'>{title}</h4>", unsafe_allow_html=True)111    if not src:112        placeholder = Image.new("RGBA", (TILE_W, TILE_H), (240,240,240,255))113        buf = io.BytesIO(); placeholder.save(buf, "PNG")114        html = f"<div class='tile-box'><img class='tile-media' src='{_b64(buf.getvalue(),'image/png')}'/></div>"115        col.markdown(html, unsafe_allow_html=True)116        col.markdown("<div class='small-cap'>—</div>", unsafe_allow_html=True)117        return None, None118 119    b = src if isinstance(src, (bytes, bytearray)) else fetch_bytes(src)120    if not b:121        col.warning("Image could not be loaded"); return None, None122 123    tile_png = image_to_9_16_canvas(b, w=TILE_W, h=TILE_H, bg_rgb=PREVIEW_BG)124    html = f"<div class='tile-box'><img class='tile-media' src='{_b64(tile_png,'image/png')}'/></div>"125    col.markdown(html, unsafe_allow_html=True)126 127    sz = image_size_from_bytes(b)128    if sz:129        col.markdown(f"<div class='small-cap'>Source: {sz[0]}×{sz[1]} px</div>", unsafe_allow_html=True)130 131    ext = infer_ext(b)132    mime = "image/jpeg" if ext in [".jpg",".jpeg"] else ("image/png" if ext==".png" else "image/webp")133    col.download_button(134        "Download (full quality)",135        data=b,136        file_name=f"{filename_stub}{ext}",137        mime=mime,138        key=unique_key(f"{key_prefix}_{filename_stub}_dl")139    )140    return src, b141 142def _video_html_src(src: str) -> str:143    return f"<div class='tile-box'><video class='tile-media' src='{src}' autoplay muted loop playsinline controls></video></div>"144 145# --- Video rendering in a single placeholder (prevents duplicate heading) ---146def render_video_placeholder(slot, video_url: str|None, key_prefix=""):147    with slot.container():148        st.markdown("<h4 class='tile-title'>Video (loop)</h4>", unsafe_allow_html=True)149        if not video_url:150            st.markdown("<div class='small-cap'>Not generated yet.</div>", unsafe_allow_html=True)151            return None152        vb = fetch_bytes(video_url)153        vsrc = _b64(vb, "video/mp4") if vb else video_url154        st.markdown(_video_html_src(vsrc), unsafe_allow_html=True)155        if vb:156            st.download_button(157                "Download video (MP4)",158                data=vb,159                file_name="fitstudio_video.mp4",160                mime="video/mp4",161                key=unique_key(f"{key_prefix}_video_dl")162            )163        return vb164 165# ---- MP4 → GIF ----166@st.cache_data(show_spinner=False)167def mp4_to_gif_bytes(mp4_bytes: bytes, target_w: int = GIF_MAX_W, fps: int = GIF_FPS, max_frames: int = GIF_MAX_FRAMES) -> bytes:168    try:169        import imageio.v3 as iio170        from PIL import Image171    except Exception as e:172        raise RuntimeError("imageio.v3 + PIL required for GIF conversion") from e173 174    with tempfile.NamedTemporaryFile(delete=False, suffix=".mp4") as f:175        f.write(mp4_bytes); mp4_path = f.name176 177    try:178        import imageio.v3 as iio2  # same pkg; just to be explicit179        meta = {}180        try: meta = iio2.immeta(mp4_path)181        except Exception: meta = {}182        src_fps = meta.get("fps", fps)183        step = max(1, int(round(src_fps / fps))) if isinstance(src_fps, (int, float)) and src_fps > 0 else 1184    except Exception:185        step = 1186 187    frames = []188    try:189        for idx, frame in enumerate(iio.imiter(mp4_path)):190            if idx % step != 0: continue191            im = Image.fromarray(frame)192            if im.width > target_w:193                new_h = max(1, int(im.height * target_w / im.width))194                im = im.resize((target_w, new_h), Image.LANCZOS)195            frames.append(im.convert("P", palette=Image.ADAPTIVE))196            if len(frames) >= max_frames: break197    except Exception as e:198        raise RuntimeError("Failed to decode frames for GIF. ffmpeg/pyav may be missing.") from e199    finally:200        try: Path(mp4_path).unlink(missing_ok=True)201        except Exception: pass202 203    if not frames: raise RuntimeError("No frames decoded for GIF.")204    out = io.BytesIO()205    frames[0].save(out, format="GIF", save_all=True, append_images=frames[1:],206                   loop=0, duration=max(10, int(1000 / fps)), disposal=2)207    return out.getvalue()208 209def make_zip(named_bytes):210    buf = io.BytesIO()211    with zipfile.ZipFile(buf, "w", compression=zipfile.ZIP_DEFLATED) as z:212        for fname, b in named_bytes:213            if b: z.writestr(fname, b)214    buf.seek(0); return buf.read()215 216def backend_ok() -> bool:217    if not API_BASE: return False218    try:219        r = requests.get(f"{API_BASE}/health", headers=HEADERS if HF_TOKEN else None, timeout=10)220        r.raise_for_status(); return True221    except Exception:222        return False223 224def post_image_edit(data: dict, files_payload):225    r = requests.post(f"{API_BASE}/v1/image/edit", data=data, files=files_payload or None,226                      headers=HEADERS if HF_TOKEN else None, timeout=600)227    r.raise_for_status(); return r.json()228 229def post_chain(data: dict, files_payload):230    payload = {**data, "to_video": "false"}231    r = requests.post(f"{API_BASE}/v1/tryon/chain", data=payload, files=files_payload or None,232                      headers=HEADERS if HF_TOKEN else None, timeout=600)233    r.raise_for_status(); return r.json()234 235def post_video(image_url: str, gender="auto", age_group="auto"):236    payload = {"image_url": image_url, "gender": gender, "age_group": age_group, "frame_aspect": "9:16"}237    r = requests.post(f"{API_BASE}/v1/video/from-image", data=payload,238                      headers=HEADERS if HF_TOKEN else None, timeout=600)239    r.raise_for_status(); return r.json()240 241# ================= Page =================242st.set_page_config(page_title="Fit Studio AI — by Vahit Feryad", layout="wide", page_icon="🧵")243 244left, right = st.columns([0.7, 0.3], vertical_alignment="center")245with left:246    st.title("🧵 Fit Studio AI — Fashion & Apparel")247    st.caption("Built by **Vahit Feryad** — Demo for Groove Jones")248with right:249    badge = "🟢 Online" if backend_ok() else "🔴 Offline"250    st.metric("Backend", badge)251 252st.markdown("**9:16** commercial stills (VTO) → optional runway-style video (looping).  \nNo persistence: outputs are remote URLs only. Gender & Age controls steer prompts server-side.")253 254if not API_BASE:255    st.error("API_BASE not set. Define env var FITSTUDIO_API (or AI_LIGHTBOX_API / LUXFIT_API).")256    st.stop()257 258# ================= Controls =================259c0, c_gender, c_age, c2, c3, c4 = st.columns([1.0, 1.0, 1.2, 1.8, 1.4, 1.6])260with c0:       category = st.selectbox("Category", ["general","underwear"], index=0, key="category")261with c_gender: gender   = st.selectbox("Gender", ["auto","female","male","unisex"], index=0, key="gender")262with c_age:    age_group= st.selectbox("Age group", ["auto","teen","young_adult","adult","mature"], index=0, key="age_group")263with c2:       file_list= st.file_uploader("Garment image (upload)", type=["jpg","jpeg","png","webp"], accept_multiple_files=True, key="file_upl")264with c3:       image_url= st.text_input("or Garment URL", key="image_url")265with c4:266    num_images    = st.selectbox("Number of looks", [1,2,3,4], index=0, key="num_images")267    custom_prompt = st.text_area("Custom prompt (blank → default)", value="", height=90, key="custom_prompt")268 269bA, bB = st.columns([1.0, 1.0])270with bA: st.info(f"API: {API_BASE}", icon="🔌")271with bB:272    run_image = st.button("Run (Image only)", type="primary", use_container_width=True, key="run_img_btn")273    run_chain = st.button("Chain (Image → prep Video)", use_container_width=True, key="run_chain_btn")274 275# ================= Demo preview =================276st.markdown("---"); st.subheader("Demo preview")277dp_in, dp_sp, dp_v = st.columns([1,1,1])278show_image_tile(dp_in, "Input", fetch_bytes("input.jpg"), filename_stub="input", key_prefix="demo_in")279show_image_tile(dp_sp, "Sample Output", fetch_bytes("tryon.jpg"), filename_stub="tryon", key_prefix="demo_out")280# keep separate label to avoid confusion with "Video (loop)"281demo_vbytes = fetch_bytes("fitstudio_video.mp4")282if demo_vbytes:283    dp_v.markdown("<h4 class='tile-title'>Sample Video (loop)</h4>", unsafe_allow_html=True)284    dp_v.markdown(f"<div class='tile-box'><video class='tile-media' src='{_b64(demo_vbytes,'video/mp4')}' autoplay muted loop playsinline controls></video></div>", unsafe_allow_html=True)285 286# ================= Input preview =================287st.markdown("---")288g1, _, _, _ = st.columns([1,1,1,1])289input_preview = None290if file_list:291    try: input_preview = file_list[0].getvalue()292    except Exception: input_preview = None293elif image_url:294    input_preview = fetch_bytes(image_url)295show_image_tile(g1, "Input (preview)", input_preview, filename_stub="preview", key_prefix="preview")296 297# ================ Run helpers ================298def _build_files_payload():299    files_payload = []300    if file_list:301        for f in file_list:302            files_payload.append(("files", (f.name, f.getvalue(), f.type or "image/jpeg")))303    return files_payload304 305def _common_form_data():306    data = {307        "category": st.session_state["category"],308        "gender": st.session_state["gender"],309        "age_group": st.session_state["age_group"],310        "frame_aspect": "9:16",311        "num_images": str(st.session_state["num_images"]),312    }313    if st.session_state.get("custom_prompt","").strip():314        data["prompt"] = st.session_state["custom_prompt"].strip()315    if image_url:316        data["image_urls"] = image_url.strip()317    return data318 319if run_image or run_chain:320    if not backend_ok():321        st.error("Backend not reachable. Check API_BASE / token.")322    elif not (file_list or image_url):323        st.error("Provide at least one garment image (upload or URL).")324    else:325        try:326            files_payload = _build_files_payload()327            form = _common_form_data()328            with st.spinner("Running…"):329                if run_image:330                    out = post_image_edit(form, files_payload)331                    image_json = out332                else:333                    out = post_chain(form, files_payload)   # to_video=false334                    image_json = out.get("image_step") or out335 336            res_block = (image_json.get("result") or {})337            imgs_norm = res_block.get("images_9_16") or []338            imgs_raw  = res_block.get("images") or []339 340            urls_norm, urls_raw = [], []341            for it in imgs_norm:342                u = it.get("url"); 343                if u: urls_norm.append(u)344            for it in imgs_raw:345                u = it.get("url"); 346                if u: urls_raw.append(u)347 348            st.session_state["job_counter"] += 1349            st.session_state["results"] = {350                "job_id": (image_json.get("job_id") or out.get("job_id")),351                "schema_version": (image_json.get("schema_version") or out.get("schema_version")),352                "image_urls_norm": urls_norm,353                "image_urls_raw": urls_raw,354                "video_url": None,355            }356 357        except requests.HTTPError as e:358            st.error(f"HTTP {e.response.status_code}")359        except Exception as e:360            st.error(f"Error: {e}")361 362# ===== Helper: dual download (normalized + native) =====363def show_dual_image_tile(col, title: str, norm_src: str|bytes|None, native_src: str|bytes|None,364                         filename_stub="image", key_prefix=""):365    preview_src = norm_src or native_src366    _, preview_bytes = show_image_tile(col, title, preview_src, filename_stub=filename_stub, key_prefix=key_prefix)367 368    norm_bytes = None; native_bytes = None369    if isinstance(norm_src, (bytes, bytearray)): norm_bytes = bytes(norm_src)370    elif isinstance(norm_src, str): norm_bytes = fetch_bytes(norm_src)371    if isinstance(native_src, (bytes, bytearray)): native_bytes = bytes(native_src)372    elif isinstance(native_src, str): native_bytes = fetch_bytes(native_src)373 374    if norm_bytes:375        extn = infer_ext(norm_bytes)376        mime = "image/jpeg" if extn in [".jpg",".jpeg"] else ("image/png" if extn==".png" else "image/webp")377        col.download_button(378            "Download 9:16 (normalized)",379            data=norm_bytes,380            file_name=f"{filename_stub}_9x16{extn}",381            mime=mime,382            key=unique_key(f"{key_prefix}_{filename_stub}_norm_dl")383        )384    if native_bytes:385        ext = infer_ext(native_bytes)386        mime = "image/jpeg" if ext in [".jpg",".jpeg"] else ("image/png" if ext==".png" else "image/webp")387        col.download_button(388            "Download native",389            data=native_bytes,390            file_name=f"{filename_stub}_native{ext}",391            mime=mime,392            key=unique_key(f"{key_prefix}_{filename_stub}_raw_dl")393        )394    return preview_bytes, norm_bytes, native_bytes395 396# ================ Results ================397res = st.session_state.get("results") or {}398if res:399    key_prefix = f"job{st.session_state['job_counter']}"400    img_urls_norm = res.get("image_urls_norm") or []401    img_urls_raw  = res.get("image_urls_raw") or []402    video_url = res.get("video_url")403 404    st.markdown("### Results")405    c1, c2, c3, c4 = st.columns([1,1,1,1])406    named = []407 408    def pair_at(i):409        norm = img_urls_norm[i] if i < len(img_urls_norm) else None410        raw  = img_urls_raw[i]  if i < len(img_urls_raw)  else None411        return norm, raw412 413    for idx, col in enumerate([c1, c2, c3]):414        norm_u, raw_u = pair_at(idx)415        if norm_u or raw_u:416            title = "Look" if idx == 0 else f"Look #{idx+1}"417            _, b_norm, b_raw = show_dual_image_tile(418                col, title, norm_u, raw_u, filename_stub=f"look_{idx+1}", key_prefix=f"{key_prefix}_{idx}"419            )420            if b_norm: named.append((f"look_{idx+1}_9x16{infer_ext(b_norm)}", b_norm))421            if b_raw:  named.append((f"look_{idx+1}_native{infer_ext(b_raw)}", b_raw))422        else:423            show_image_tile(col, f"Look #{idx+1}", None, key_prefix=f"{key_prefix}_{idx}")424 425    # --- Single, persistent video slot ---426    video_slot = c4.empty()427    vbytes = render_video_placeholder(video_slot, video_url, key_prefix=key_prefix)428 429    # Controls430    st.markdown("---")431    colv1, colv2 = st.columns([1.2, 1.2])432    with colv1:433        gen_video = st.checkbox("Generate video from first look", value=False, key="video_toggle")434    with colv2:435        go_video = st.button("Create Video", use_container_width=True, key="video_btn")436 437    if gen_video and go_video:438        src_for_video = (img_urls_norm[0] if img_urls_norm else (img_urls_raw[0] if img_urls_raw else None))439        if not src_for_video:440            st.warning("No source image to create video.")441        else:442            with st.spinner("Generating video…"):443                v = post_video(src_for_video, gender=st.session_state.get("gender","auto"),444                               age_group=st.session_state.get("age_group","auto"))445            vurl = (v.get("result") or {}).get("video", {}).get("url")446            if vurl:447                st.session_state["results"]["video_url"] = vurl448                # Update same placeholder in-place → no duplicate heading449                vbytes = render_video_placeholder(video_slot, vurl, key_prefix=key_prefix)450            else:451                st.info("Video URL not returned.")452 453    # ZIP + GIF454    if named or vbytes:455        if vbytes:456            named.append(("video.mp4", vbytes))457        zip_bytes = make_zip(named)458        st.download_button(459            "Download All (ZIP)",460            data=zip_bytes,461            file_name="fitstudio_outputs.zip",462            mime="application/zip",463            key=unique_key(f"{key_prefix}_zip_dl")464        )465        if vbytes:466            try:467                with st.spinner("Preparing GIF…"):468                    gif_bytes = mp4_to_gif_bytes(vbytes, target_w=GIF_MAX_W, fps=GIF_FPS, max_frames=GIF_MAX_FRAMES)469                st.download_button(470                    f"Download GIF ({GIF_MAX_W}px, {GIF_FPS}fps)",471                    data=gif_bytes,472                    file_name="fitstudio_video.gif",473                    mime="image/gif",474                    key=unique_key(f"{key_prefix}_gif_dl")475                )476            except Exception as e:477                st.info(f"GIF conversion not available: {e}. Try installing ffmpeg / imageio-ffmpeg.")478 479# ================= Sidebar =================480with st.sidebar:481    st.caption(f"Resolved API_BASE: {API_BASE}")482    st.caption("Auth header: " + ("ON" if HF_TOKEN else "OFF"))483    st.markdown("---")484    st.markdown("**About this demo**")485    st.caption("Production-style, no local persistence. Built to show scalable image→video chaining with guardrails for consistent 9:16 outputs.")486    if st.button("Ping /health", key=unique_key("ping_health")):487        try:488            r = requests.get(f"{API_BASE}/health", headers=HEADERS if HF_TOKEN else None, timeout=10)489            st.write(r.status_code)490            try: st.json(r.json())491            except Exception: st.code((r.text or "")[:1000])492        except Exception as e:493            st.error(f"Health error: {e}")494    if st.button("Clear cache", key=unique_key("clear_cache")):495        st.cache_data.clear(); st.success("Cache cleared.")496    if st.button("Clear results", key=unique_key("clear_results")):497        st.session_state["results"] = None; st.success("Results cleared.")498