CoolFace
Modelpublic

SleepVeryHard/ToriiGate-0.5_GGUF

sourceHugging Facemitupdated 6mo agoView on Hugging Face
8likes1.2kdownloads
gradio_interface.py636 linesDownload Raw Back to scripts
1import json2import base643import io4import requests5from pathlib import Path6from typing import Dict, Any, Optional, Tuple7 8import gradio as gr9from PIL import Image10 11# Import prompt building functions from prompts.py12from prompts import make_user_query, system_prompt, prompts_b13 14# ==================== CONFIGURATION ====================15 16# API settings17API_URL = "http://127.0.0.1:8000/v1/chat/completions"18API_KEY = "not-needed"19 20# Image settings21MAX_PIXELS = 1.0  # Maximum resolution in megapixels (e.g., 4.0 = 4MP)22 23# Request settings24MAX_TOKENS = 409625TEMPERATURE = 0.526REQUEST_TIMEOUT = 5  # Reduced for connection check27WORK_TIMEOUT = 30028 29# Captioning type options (from prompts_b in prompts.py)30CAPTION_TYPES = list(prompts_b.keys())31DEFAULT_C_TYPE = CAPTION_TYPES[0] if CAPTION_TYPES else None32 33if not DEFAULT_C_TYPE:34    raise RuntimeError("No caption types available in prompts_b!")35 36# ==================== END CONFIGURATION ====================37 38 39def check_api_connection(api_url: str) -> Tuple[str, str]:40    """41    Check API connection and return model info.42    Returns (status_message, model_name).43    """44    try:45        # Try to get models endpoint46        base_url = api_url.rstrip('/').split('/v1/')[0]47        models_url = f"{base_url}/v1/models"48        49        response = requests.get(models_url, timeout=REQUEST_TIMEOUT)50        response.raise_for_status()51        52        result = response.json()53        if result and 'data' in result and len(result['data']) > 0:54            model_name = result['data'][0].get('id', 'Unknown')55            return "✅ Connected", model_name56        else:57            return "⚠️ Connected (no model info)", "Unknown"58            59    except requests.exceptions.ConnectionError:60        return "❌ Connection failed", "N/A"61    except requests.exceptions.Timeout:62        return "❌ Timeout", "N/A"63    except Exception as e:64        return f"❌ Error: {str(e)[:50]}", "N/A"65 66 67def encode_image_base64(image: Image.Image, max_pixels: float = MAX_PIXELS) -> str:68    """Encode image to base64 string, resizing if necessary."""69    img = image70    71    if img.mode != 'RGB':72        img = img.convert('RGB')73    74    # Check if resizing needed75    current_pixels = img.width * img.height76    max_pixels_count = max_pixels * 1_000_00077    78    if current_pixels >= max_pixels_count:79        # Calculate new dimensions while preserving aspect ratio80        scale = (max_pixels_count / current_pixels) ** 0.581        new_width = int(img.width * scale)82        new_height = int(img.height * scale)83        84        # Resize with high quality85        img = img.resize((new_width, new_height), Image.Resampling.LANCZOS)86        # No resize needed87    88    # Encode resized image to base6489    buffer = io.BytesIO()90    img.save(buffer, format='JPEG', quality=100)91    return base64.b64encode(buffer.getvalue()).decode("utf-8")92 93 94def call_caption_api(messages: list, api_url: str = API_URL, model_name: str = "toriigate-0.5") -> Optional[str]:95    """Call the captioning API."""96    payload = {97        "model": model_name,98        "messages": messages,99        "max_tokens": MAX_TOKENS,100        "temperature": TEMPERATURE,101        "stream": False102    }103 104    headers = {105        "Content-Type": "application/json",106        "Authorization": f"Bearer {API_KEY}"107    }108 109    try:110        response = requests.post(111            api_url,112            headers=headers,113            json=payload,114            timeout=WORK_TIMEOUT115        )116        response.raise_for_status()117 118        result = response.json()119        content = result['choices'][0]['message']['content']120        return content121 122    except requests.exceptions.RequestException as e:123        return f"API Error: {e}"124    except (KeyError, IndexError) as e:125        return f"Parse Error: {e}"126 127 128def empty_template() -> Dict[str, Any]:129    """Return empty template for missing JSON data."""130    return {131        "tags": [],132        "characters": [],133        "char_p_tags": {"chars": {}, "skins": {}},134        "char_descr": {"chars": {}, "skins": {}}135    }136 137 138def generate_caption(139    image: Image.Image,140    api_url: str,141    model_name: str,142    c_type: str,143    use_names: bool,144    add_tags: bool,145    add_char_list: bool,146    add_chars_tags: bool,147    add_chars_descr: bool,148    tags_text: str,149    characters_text: str,150    char1_name: str,151    char1_tags: str,152    char2_name: str,153    char2_tags: str,154    char3_name: str,155    char3_tags: str,156    char4_name: str,157    char4_tags: str,158    char5_name: str,159    char5_tags: str,160    char_descr1_name: str,161    char_descr1_text: str,162    char_descr2_name: str,163    char_descr2_text: str,164    char_descr3_name: str,165    char_descr3_text: str,166    char_descr4_name: str,167    char_descr4_text: str,168    char_descr5_name: str,169    char_descr5_text: str170) -> str:171    """Generate caption for a single image."""172    if image is None:173        return "Please upload an image first."174 175    # Build item dict from inputs176    item = empty_template()177 178    # Parse tags179    if add_tags and tags_text.strip():180        item["tags"] = [t.strip() for t in tags_text.split(',') if t.strip()]181 182    # Parse characters183    if add_char_list:184        item["characters"] = [c.strip() for c in characters_text.split(',') if c.strip()]185    186    # Auto-populate characters list from char tags/descriptions if not manually specified187    if add_chars_tags or add_chars_descr:188        auto_chars = []189        190        if add_chars_tags:191            char_entries = [192                char1_name, char2_name, char3_name, char4_name, char5_name193            ]194            for name in char_entries:195                if name and name.strip():196                    auto_chars.append(name.strip())197        198        if add_chars_descr:199            descr_entries = [200                char_descr1_name, char_descr2_name, char_descr3_name,201                char_descr4_name, char_descr5_name202            ]203            for name in descr_entries:204                if name and name.strip() and name.strip() not in auto_chars:205                    auto_chars.append(name.strip())206        207        # Only auto-populate if characters list is empty or not manually set208        if auto_chars and (not add_char_list or not item["characters"]):209            item["characters"] = auto_chars210            add_char_list = True211 212    # Parse character tags from structured inputs213    if add_chars_tags:214        chars_dict = {}215        char_entries = [216            (char1_name, char1_tags),217            (char2_name, char2_tags),218            (char3_name, char3_tags),219            (char4_name, char4_tags),220            (char5_name, char5_tags)221        ]222        for name, tags_str in char_entries:223            if name is None:224                continue225            name = name.strip()226            if name:227                tags_list = [t.strip() for t in tags_str.split(',') if t.strip()] if tags_str and tags_str.strip() else []228                chars_dict[name] = tags_list229        230        if chars_dict:231            item["char_p_tags"] = {"chars": chars_dict, "skins": {}}232 233    # Parse character descriptions from structured inputs234    if add_chars_descr:235        descr_dict = {}236        descr_entries = [237            (char_descr1_name, char_descr1_text),238            (char_descr2_name, char_descr2_text),239            (char_descr3_name, char_descr3_text),240            (char_descr4_name, char_descr4_text),241            (char_descr5_name, char_descr5_text)242        ]243        for name, descr in descr_entries:244            if name is None or descr is None:245                continue246            name = name.strip()247            descr = descr.strip()248            if name and descr:249                descr_dict[name] = descr250        251        if descr_dict:252            item["char_descr"] = {"chars": descr_dict, "skins": {}}253 254    # Encode image255    image_data = encode_image_base64(image)256 257    # Prepare messages258    user_query = make_user_query(259        item,260        c_type=c_type,261        use_names=use_names,262        add_tags=add_tags,263        add_characters=add_char_list,264        add_char_tags=add_chars_tags,265        add_description=add_chars_descr,266        underscores_replace=False267    )268 269    messages = [270        {271            "role": "system",272            "content": [{"type": "text", "text": system_prompt}]273        },274        {275            "role": "user",276            "content": [277                {"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{image_data}"}},278                {"type": "text", "text": user_query}279            ]280        }281    ]282 283    # Call API284    return call_caption_api(messages, api_url, model_name)285 286 287def create_ui():288    """Create and return the Gradio interface."""289 290    with gr.Blocks(title="ToriiGate Captioner", theme=gr.themes.Soft()) as app:291        gr.Markdown("# 🖼️ ToriiGate Captioner")292        293        # API URL row with status294        with gr.Row():295            api_url_input = gr.Textbox(296                label="API URL",297                value=API_URL,298                interactive=True,299                scale=4300            )301            api_status = gr.Textbox(302                label="Status",303                value="⏳ Waiting for input...",304                interactive=False,305                scale=1306            )307            model_name_display = gr.Textbox(308                label="Model",309                value="N/A",310                interactive=False,311                scale=1312            )313 314        with gr.Row():315            # Left column - Image input316            with gr.Column(scale=1):317                image_input = gr.Image(318                    label="Upload Image",319                    type="pil",320                    height=400321                )322                323                gr.Markdown("### Configuration")324                325                # Caption type selector326                c_type = gr.Dropdown(327                    choices=CAPTION_TYPES,328                    value=DEFAULT_C_TYPE,329                    label="Caption Type",330                    interactive=True331                )332                333                # Boolean options with conditional text inputs334                with gr.Group():335                    use_names = gr.Checkbox(336                        value=True,337                        label="Use Names (enable character names)"338                    )339                    340                    add_tags = gr.Checkbox(341                        value=False,342                        label="Add Tags"343                    )344                    tags_text = gr.Textbox(345                        label="Tags (comma-separated)",346                        placeholder="e.g., 1girl, blue_hair, school_uniform",347                        interactive=False348                    )349                    350                    add_char_list = gr.Checkbox(351                        value=False,352                        label="Add Character List"353                    )354                    characters_text = gr.Textbox(355                        label="Character Names (comma-separated)",356                        placeholder="e.g., nishizono_mio, hoshimi_miyabi",357                        interactive=False358                    )359 360                    add_chars_tags = gr.Checkbox(361                        value=False,362                        label="Add Character Tags"363                    )364                    365                    with gr.Group(visible=False) as char_tags_group:366                        gr.Markdown("**Add character names and their tags**")367                        368                        with gr.Accordion("Character 1", open=True):369                            char1_name = gr.Textbox(370                                label="Name",371                                placeholder="e.g., albedo",372                                interactive=True373                            )374                            char1_tags = gr.Textbox(375                                label="Tags (comma-separated)",376                                placeholder="e.g., white_hair, green_eyes, horns",377                                interactive=True378                            )379                        380                        with gr.Accordion("Character 2", open=False):381                            char2_name = gr.Textbox(382                                label="Name",383                                placeholder="e.g., hoshimi_miyabi",384                                interactive=True385                            )386                            char2_tags = gr.Textbox(387                                label="Tags (comma-separated)",388                                placeholder="e.g., blue_hair, fox_ears",389                                interactive=True390                            )391                        392                        with gr.Accordion("Character 3", open=False):393                            char3_name = gr.Textbox(394                                label="Name",395                                placeholder="e.g., nishizono_mio",396                                interactive=True397                            )398                            char3_tags = gr.Textbox(399                                label="Tags (comma-separated)",400                                placeholder="e.g., brown_hair, glasses",401                                interactive=True402                            )403                        404                        with gr.Accordion("Character 4", open=False):405                            char4_name = gr.Textbox(406                                label="Name",407                                placeholder="e.g.",408                                interactive=True409                            )410                            char4_tags = gr.Textbox(411                                label="Tags (comma-separated)",412                                placeholder="e.g.",413                                interactive=True414                            )415                        416                        with gr.Accordion("Character 5", open=False):417                            char5_name = gr.Textbox(418                                label="Name",419                                placeholder="e.g.",420                                interactive=True421                            )422                            char5_tags = gr.Textbox(423                                label="Tags (comma-separated)",424                                placeholder="e.g.",425                                interactive=True426                            )427                        428                        char_tags_clear_btn = gr.Button(429                            "🗑️ Clear All",430                            variant="secondary",431                            size="sm"432                        )433 434                    add_chars_descr = gr.Checkbox(435                        value=False,436                        label="Add Character Descriptions"437                    )438                    439                    with gr.Group(visible=False) as char_descr_group:440                        gr.Markdown("**Add character descriptions**")441                        442                        with gr.Accordion("Character 1", open=True):443                            char_descr1_name = gr.Textbox(444                                label="Name",445                                placeholder="e.g., albedo",446                                interactive=True447                            )448                            char_descr1_text = gr.Textbox(449                                label="Description",450                                placeholder="e.g., Albedo is a curvy woman with...",451                                lines=3,452                                interactive=True453                            )454                        455                        with gr.Accordion("Character 2", open=False):456                            char_descr2_name = gr.Textbox(457                                label="Name",458                                placeholder="e.g., hoshimi_miyabi",459                                interactive=True460                            )461                            char_descr2_text = gr.Textbox(462                                label="Description",463                                placeholder="e.g., Miyabi is a calm and collected...",464                                lines=3,465                                interactive=True466                            )467                        468                        with gr.Accordion("Character 3", open=False):469                            char_descr3_name = gr.Textbox(470                                label="Name",471                                placeholder="e.g., nishizono_mio",472                                interactive=True473                            )474                            char_descr3_text = gr.Textbox(475                                label="Description",476                                placeholder="e.g., Mio is a cheerful girl with...",477                                lines=3,478                                interactive=True479                            )480                        481                        with gr.Accordion("Character 4", open=False):482                            char_descr4_name = gr.Textbox(483                                label="Name",484                                placeholder="e.g.",485                                interactive=True486                            )487                            char_descr4_text = gr.Textbox(488                                label="Description",489                                placeholder="e.g.",490                                lines=3,491                                interactive=True492                            )493                        494                        with gr.Accordion("Character 5", open=False):495                            char_descr5_name = gr.Textbox(496                                label="Name",497                                placeholder="e.g.",498                                interactive=True499                            )500                            char_descr5_text = gr.Textbox(501                                label="Description",502                                placeholder="e.g.",503                                lines=3,504                                interactive=True505                            )506                        507                        char_descr_clear_btn = gr.Button(508                            "🗑️ Clear All",509                            variant="secondary",510                            size="sm"511                        )512 513                generate_btn = gr.Button("🚀 Generate Caption", variant="primary", size="lg")514            515            # Right column - Output516            with gr.Column(scale=1):517                output_text = gr.Textbox(518                    label="Caption Output",519                    lines=20,520                    max_lines=50,521                    interactive=False522                )523        524        # Toggle text inputs based on checkbox state525        def toggle_input(is_checked: bool, input_component):526            return gr.update(interactive=is_checked)527 528        add_tags.change(529            lambda x: toggle_input(x, tags_text),530            inputs=add_tags,531            outputs=tags_text532        )533 534        add_char_list.change(535            lambda x: toggle_input(x, characters_text),536            inputs=add_char_list,537            outputs=characters_text538        )539 540        add_chars_tags.change(541            fn=lambda x: gr.update(visible=x),542            inputs=add_chars_tags,543            outputs=char_tags_group544        )545 546        add_chars_descr.change(547            fn=lambda x: gr.update(visible=x),548            inputs=add_chars_descr,549            outputs=char_descr_group550        )551 552        # API URL change handler553        api_url_input.change(554            fn=check_api_connection,555            inputs=api_url_input,556            outputs=[api_status, model_name_display]557        )558 559        # Wire up generate button560        generate_btn.click(561            fn=generate_caption,562            inputs=[563                image_input,564                api_url_input,565                model_name_display,566                c_type,567                use_names,568                add_tags,569                add_char_list,570                add_chars_tags,571                add_chars_descr,572                tags_text,573                characters_text,574                char1_name,575                char1_tags,576                char2_name,577                char2_tags,578                char3_name,579                char3_tags,580                char4_name,581                char4_tags,582                char5_name,583                char5_tags,584                char_descr1_name,585                char_descr1_text,586                char_descr2_name,587                char_descr2_text,588                char_descr3_name,589                char_descr3_text,590                char_descr4_name,591                char_descr4_text,592                char_descr5_name,593                char_descr5_text594            ],595            outputs=output_text596        )597 598        # Clear character tags button handler599        def clear_char_tags():600            return "", "", "", "", "", "", "", "", "", ""601        602        char_tags_clear_btn.click(603            fn=clear_char_tags,604            inputs=[],605            outputs=[606                char1_name, char1_tags,607                char2_name, char2_tags,608                char3_name, char3_tags,609                char4_name, char4_tags,610                char5_name, char5_tags611            ]612        )613 614        # Clear character descriptions button handler615        def clear_char_descr():616            return "", "", "", "", "", "", "", "", "", ""617        618        char_descr_clear_btn.click(619            fn=clear_char_descr,620            inputs=[],621            outputs=[622                char_descr1_name, char_descr1_text,623                char_descr2_name, char_descr2_text,624                char_descr3_name, char_descr3_text,625                char_descr4_name, char_descr4_text,626                char_descr5_name, char_descr5_text627            ]628        )629    630    return app631 632 633if __name__ == "__main__":634    app = create_ui()635    app.launch(server_name="127.0.0.1", server_port=7860)636