Team Ai
Apppublic

Krish280199/davinci-image-processor

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
app.py370 linesDownload Raw Back to root
1# app.py2 3import streamlit as st4import pandas as pd5import os6import requests7from io import BytesIO8from PIL import Image, ImageOps9import numpy as np10from scipy import ndimage11from pathlib import Path12import warnings13 14# --- Suppress specific warnings ---15warnings.filterwarnings("ignore", category=UserWarning, module='torchvision')16 17# --- Dependency Check and Model Setup ---18try:19    from rembg import remove, new_session20    REMBG_AVAILABLE = True21except ImportError:22    REMBG_AVAILABLE = False23    # Define dummy functions if rembg is not available24    def new_session(model_name): return None25    def remove(img, session): return img26 27try:28    from basicsr.archs.rrdbnet_arch import RRDBNet29    from realesrgan import RealESRGANer30    REALESRGAN_AVAILABLE = True31except ImportError:32    REALESRGAN_AVAILABLE = False33 34# --- AI Model Loading (Cached) ---35@st.cache_resource36def get_realesrgan_model():37    """Downloads and loads the Real-ESRGAN model."""38    if not REALESRGAN_AVAILABLE:39        return None40        41    weights_dir = 'weights'42    os.makedirs(weights_dir, exist_ok=True)43    model_path = os.path.join(weights_dir, 'RealESRGAN_x4plus.pth')44    45    if not os.path.exists(model_path):46        model_url = 'https://github.com/xinntao/Real-ESRGAN/releases/download/v0.1.0/RealESRGAN_x4plus.pth'47        st.info(f'Downloading Real-ESRGAN model...')48        try:49            response = requests.get(model_url, stream=True)50            response.raise_for_status()51            with open(model_path, 'wb') as f:52                for chunk in response.iter_content(chunk_size=8192):53                    f.write(chunk)54            st.success('Model downloaded successfully.')55        except Exception as e:56            st.error(f"Failed to download model: {e}")57            return None58    try:59        model = RealESRGANer(60            scale=4, model_path=model_path,61            model=RRDBNet(num_in_ch=3, num_out_ch=3, num_feat=64, num_block=23, num_grow_ch=32),62            tile=0, tile_pad=10, pre_pad=0, half=False63        )64        return model65    except Exception as e:66        st.error(f"Failed to load Real-ESRGAN model: {e}")67        return None68 69# Create a cached function to load the background removal model70@st.cache_resource71def get_bg_removal_session(model_name):72    """Loads a specific rembg model into a session, cached for performance."""73    if not REMBG_AVAILABLE:74        return None75    with st.spinner(f"Loading background model '{model_name}'..."):76        # This will auto-download the model from GitHub on first run77        try:78            session = new_session(model_name)79            return session80        except Exception as e:81            st.error(f"Failed to load model '{model_name}': {e}. Please try another model.")82            return None83 84# --- Core Image Processing Functions ---85 86def load_and_orient_image(image_source) -> Image.Image:87    """Loads an image and corrects its orientation based on EXIF data."""88    img = Image.open(image_source)89    # This line reads the EXIF tag and rotates the image automatically90    img = ImageOps.exif_transpose(img)91    return img92 93def process_image(img: Image.Image, use_sr: bool, use_rembg: bool, bg_model_name: str, **kwargs) -> Image.Image:94    """The main image processing pipeline."""95    96    # 1. (Optional) AI Super-Resolution97    if use_sr and REALESRGAN_AVAILABLE:98        upsampler = get_realesrgan_model()99        if upsampler:100            with st.spinner("Upscaling image with AI..."):101                img_np = np.array(img.convert("RGB"))102                sr_img_np, _ = upsampler.enhance(img_np, outscale=4)103                img = Image.fromarray(sr_img_np)104 105    # 2. (Optional) Background Removal + Smart Cleanup106    if use_rembg and REMBG_AVAILABLE:107        # Get the selected model session108        session = get_bg_removal_session(bg_model_name)109        if session:110            with st.spinner("Removing background..."):111                # Use the session to remove the background112                img_no_bg = remove(img.convert("RGBA"), session=session)113            114            # Smart cleanup115            data = np.array(img_no_bg)116            alpha = data[:, :, 3]117            core_mask = alpha > 245118            protection_zone = ndimage.binary_dilation(core_mask, iterations=2)119            r, g, b = data[:, :, 0], data[:, :, 1], data[:, :, 2]120            is_light_gray_artifact = ((alpha > 0) & (alpha < 230) & (r > 160) & (g > 160) & (b > 160) & (abs(r.astype(np.int16) - g.astype(np.int16)) < 25) & (abs(g.astype(np.int16) - b.astype(np.int16)) < 25))121            pixels_to_delete = is_light_gray_artifact & ~protection_zone122            alpha[pixels_to_delete] = 0123            data[:, :, 3] = alpha124            img_cleaned = Image.fromarray(data)125            bbox = img_cleaned.getbbox()126            if not bbox: return None127            img_for_pipeline = img_cleaned.crop(bbox)128        else:129            st.warning("Background removal session failed to load. Skipping.")130            img_for_pipeline = img.copy()131    else:132        img_for_pipeline = img.copy()133 134    # 3. Final Resizing and Placement (Your Original Logic)135    target_size = kwargs.get('target_size')136    padding = kwargs.get('padding')137    output_format = kwargs.get('output_format')138    max_box = target_size - 2 * padding139    ratio = min(max_box / img_for_pipeline.width, max_box / img_for_pipeline.height) if img_for_pipeline.width > 0 and img_for_pipeline.height > 0 else 1140    new_size = (int(img_for_pipeline.width * ratio), int(img_for_pipeline.height * ratio))141    img_resized = img_for_pipeline.resize(new_size, Image.LANCZOS) if ratio != 1.0 else img_for_pipeline142 143    if not (use_rembg and REMBG_AVAILABLE) or output_format == "JPEG":144        canvas_mode, canvas_color = "RGB", (255, 255, 255)145    else:146        # Use transparent background for PNGs147        canvas_mode, canvas_color = "RGBA", (0, 0, 0, 0)148        149    final_img = Image.new(canvas_mode, (target_size, target_size), canvas_color)150    x_offset, y_offset = (target_size - img_resized.width) // 2, (target_size - img_resized.height) // 2151    x_offset -= 30; y_offset -= 20 # Your Original optical centering152    153    # Ensure image has alpha channel if pasting into RGBA canvas154    if final_img.mode == 'RGBA' and img_resized.mode != 'RGBA':155        img_resized = img_resized.convert('RGBA')156 157    if img_resized.mode == 'RGBA':158        final_img.paste(img_resized, (x_offset, y_offset), img_resized)159    else:160        final_img.paste(img_resized, (x_offset, y_offset))161    return final_img162 163# --- Streamlit UI ---164st.set_page_config(page_title="Davinci Image Processor", page_icon="๐ŸŽจ", layout="wide")165st.title("๐ŸŽจ Davinci - High-Quality Image Processor")166 167if not REMBG_AVAILABLE: st.warning("`rembg` library not found. Background removal is disabled.")168if not REALESRGAN_AVAILABLE: st.warning("`realesrgan` and `basicsr` not found. AI Super-Resolution is disabled.")169 170# --- Sidebar for Configuration ---171with st.sidebar:172    st.header("โš™๏ธ Configuration")173    processing_mode = st.radio("Choose Mode", ["Single Preview", "Batch (Excel)", "Batch (Local Files)"])174 175    st.subheader("Processing Options")176    use_rembg_toggle = st.toggle("Remove Background", value=True, disabled=not REMBG_AVAILABLE)177    178    # --- MODIFICATION ---179    # Removed the selectbox and hard-coded the bria-rmbg model180    bg_model_name = "bria-rmbg"181    st.caption("Using `bria-rmbg` model for background removal.")182    # --- END MODIFICATION ---183    184    output_format = st.radio("Output Format", ["PNG", "JPEG"], index=0)185    jpeg_quality = st.slider("JPEG Quality", 75, 100, 98) if output_format == "JPEG" else 98186    187    st.subheader("AI Tools")188    use_sr_toggle = st.toggle("Enable AI Super-Resolution", value=False, disabled=not REALESRGAN_AVAILABLE, help="Upscales image 4x before processing. Very slow!")189 190    st.subheader("Canvas Settings")191    target_size = st.slider("Target Size (px)", 500, 4000, 1400)192    padding = st.slider("Padding (px)", 0, 500, 100)193    194# Pack all processing kwargs195processing_kwargs = {196    "target_size": target_size, "padding": padding, "use_rembg": use_rembg_toggle,197    "output_format": output_format, "use_sr": use_sr_toggle,198    "bg_model_name": bg_model_name  # Pass the hard-coded model name199}200 201# --- Main App Logic ---202 203## --- Single Image Preview Mode ---204if processing_mode == "Single Preview":205    st.header("Single Image Preview")206    source_option = st.radio("Input type", ["Upload", "URL"], horizontal=True)207    input_image = None208    if source_option == "Upload":209        uploaded_file = st.file_uploader("Choose an image", type=["png", "jpg", "jpeg", "webp"])210        if uploaded_file: input_image = load_and_orient_image(uploaded_file)211    else:212        image_url = st.text_input("Enter image URL")213        if image_url:214            try:215                response = requests.get(image_url, timeout=10); response.raise_for_status()216                input_image = load_and_orient_image(BytesIO(response.content))217            except Exception as e: st.error(f"Failed to load image: {e}")218 219    if input_image:220        st.subheader("Preview")221        col1, col2 = st.columns(2)222        with col1: st.image(input_image, caption="Original", use_column_width=True)223        with st.spinner("Processing..."):224            processed_img = process_image(input_image.copy(), **processing_kwargs)225        with col2:226            if processed_img:227                st.image(processed_img, caption="Processed", use_column_width=True)228                img_byte_arr = BytesIO()229                230                # Convert to RGB if format is JPEG231                save_image_single = processed_img232                if output_format == 'JPEG' and save_image_single.mode == 'RGBA':233                    save_image_single = save_image_single.convert('RGB')234                235                file_ext_single = "jpg" if output_format == "JPEG" else "png"236                save_image_single.save(img_byte_arr, format=output_format, quality=jpeg_quality)237                st.download_button(label=f"Download (.{file_ext_single})", data=img_byte_arr.getvalue(), file_name=f"processed.{file_ext_single}", mime=f"image/{file_ext_single}")238            else: st.error("Processing failed.")239 240## --- Batch Processing Modes (FIXED) ---241else:242    iterable = None243    total = 0244    245    if processing_mode == "Batch (Excel)":246        st.header("Batch Process from Excel")247        uploaded_file = st.file_uploader("Upload Excel file", type=["xlsx"])248        if uploaded_file:249            df = pd.read_excel(uploaded_file); st.dataframe(df.head())250            st.subheader("1. Map Columns")251            image_col = st.selectbox("Image URL Column", df.columns)252            barcode_col = st.selectbox("Barcode Column", df.columns)253            254            # --- FIX: Re-added Batch Preview ---255            st.subheader("2. Preview (Optional)")256            preview_list = ["-- Select to preview --"] + df[barcode_col].dropna().astype(str).unique().tolist()257            selected_barcode = st.selectbox("Choose a barcode to preview:", preview_list)258            259            if selected_barcode != "-- Select to preview --":260                row = df[df[barcode_col].astype(str) == selected_barcode].iloc[0]261                url = row.get(image_col)262                if url and pd.notna(url):263                    try:264                        col1, col2 = st.columns(2)265                        original_preview = load_and_orient_image(BytesIO(requests.get(url).content))266                        with col1:267                            st.image(original_preview, caption=f"Original ({selected_barcode})", use_column_width=True)268                        with col2:269                            with st.spinner("Processing preview..."):270                                processed_preview = process_image(original_preview.copy(), **processing_kwargs)271                                st.image(processed_preview, caption="Processed Preview", use_column_width=True)272                    except Exception as e:273                        st.error(f"Failed to load preview image: {e}")274            275            st.markdown("---")276            st.subheader("3. Start Full Batch Process")277            if st.button("Process All and Save", type="primary"):278                iterable = df.iterrows()279                total = len(df)280    281    elif processing_mode == "Batch (Local Files)":282        st.header("Batch Process from Local Files")283        uploaded_files = st.file_uploader("Upload image files", type=["png", "jpg", "jpeg", "webp"], accept_multiple_files=True)284        if uploaded_files:285            # --- FIX: Re-added Batch Preview ---286            st.subheader("1. Preview (Optional)")287            preview_list = ["-- Select to preview --"] + [f.name for f in uploaded_files]288            selected_file_name = st.selectbox("Choose a file to preview:", preview_list)289 290            if selected_file_name != "-- Select to preview --":291                selected_file = next((f for f in uploaded_files if f.name == selected_file_name), None)292                if selected_file:293                    col1, col2 = st.columns(2)294                    selected_file.seek(0)295                    original_preview = load_and_orient_image(selected_file)296                    with col1:297                        st.image(original_preview, caption=f"Original ({selected_file.name})", use_column_width=True)298                    with col2:299                        with st.spinner("Processing preview..."):300                            processed_preview = process_image(original_preview.copy(), **processing_kwargs)301                            st.image(processed_preview, caption="Processed Preview", use_column_width=True)302            303            st.markdown("---")304            st.subheader("2. Start Full Batch Process")305            if st.button(f"Process All {len(uploaded_files)} Images and Save", type="primary"):306                iterable = uploaded_files # Pass the list of files directly307                total = len(uploaded_files)308 309    # --- ROBUST BATCH EXECUTION LOGIC (FIXED) ---310    if iterable:311        pictures_folder = Path.home() / "Pictures"312        save_dir = pictures_folder / "davinci_output"313        save_dir.mkdir(parents=True, exist_ok=True)314        315        st.info(f"Starting... Images will be saved to: {save_dir}")316        progress_bar = st.progress(0)317        log_area = st.container(height=300)318        processed_count = 0319        320        # We use enumerate for a clean counter321        for i, item in enumerate(iterable):322            progress_bar.progress((i + 1) / total)323            item_id = "" 324            325            try:326                original_img = None327                328                if processing_mode == "Batch (Excel)":329                    row = item[1] # item is (index, row) from iterrows()330                    item_id = str(row.get(barcode_col, f"row_{i+1}")).strip()331                    image_url = row.get(image_col)332                    333                    if not image_url or pd.isna(image_url):334                        log_area.warning(f"โš ๏ธ Skipping {item_id}: Missing image URL.")335                        continue 336                        337                    response = requests.get(image_url, timeout=10); response.raise_for_status()338                    original_img = load_and_orient_image(BytesIO(response.content))339                    340                elif processing_mode == "Batch (Local Files)":341                    file = item # item is an UploadedFile object342                    item_id = Path(file.name).stem343                    original_img = load_and_orient_image(file)344 345                if original_img is None:346                    log_area.error(f"โŒ Error on {item_id}: Could not load image.")347                    continue348 349                final_img = process_image(original_img.copy(), **processing_kwargs)350                351                if final_img:352                    file_ext = "png" if output_format == "PNG" else "jpg"353                    save_path = save_dir / f"{item_id}.{file_ext}"354                    355                    save_image = final_img356                    if output_format == 'JPEG' and save_image.mode == 'RGBA':357                        save_image = save_image.convert('RGB')358                    359                    save_image.save(save_path, format=output_format, quality=jpeg_quality)360                    log_area.success(f"โœ… Saved: {save_path.name}")361                    processed_count += 1362                else:363                    log_area.warning(f"โš ๏ธ Skipping {item_id}: Processing failed (e.g., empty image).")364            365            except Exception as e:366                if not item_id: item_id = f"item {i+1}" # Fallback367                log_area.error(f"โŒ CRITICAL Error on {item_id}: {e}")368        369        progress_bar.empty()370        st.success(f"๐ŸŽ‰ Batch complete! **{processed_count}** images saved to: `{save_dir}`")