SleepVeryHard/ToriiGate-0.5_GGUF
81.2k
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 