CoolFace
Apppublic

OnMoon/character_testing_ui

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py154 linesDownload Raw Back to root
1import os2import json3import base644import requests5import gradio as gr6import matplotlib.pyplot as plt7 8from PIL import Image9from io import BytesIO10from dotenv import load_dotenv11from huggingface_hub import HfFileSystem, hf_hub_download12 13def request_to_endpoint (request:dict, endpoint_id='s4i4um7xakaq37'):14    api_key = os.getenv("RUNPOD_API_KEY")15    url = f"https://api.runpod.ai/v2/{endpoint_id}/runsync"16    headers = {17        "accept": "application/json",18        "authorization": api_key,19        "content-type": "application/json"20    }21    response = requests.post(url, headers=headers, data=json.dumps(request)) 22 23    return response24 25# Собирает только sdxl версии26def prepare_characters(hf_path="OnMoon/loras"):27    fs = HfFileSystem()28 29    files = fs.ls(hf_path, detail=False)30    character_names = []31    character_configs = {}32    for file in files:33        if file.endswith(".safetensors") and file.startswith(f"{hf_path}/sdxl_"):34            character = file[len(f"{hf_path}/sdxl_"): -len(".safetensors")]35            character_names.append(character)36        elif file.endswith(".json") and file.startswith(f"{hf_path}/sdxl_"):37            character = file[len(f"{hf_path}/sdxl_"): -len(".json")]38            with fs.open(file, 'r', encoding='utf-8') as file_json:39                character_configs[character] = json.load(file_json)  40        41    return character_names, character_configs42 43character_names, character_configs = prepare_characters()44 45 46def process_input (name, scale, triggers, tech_prompt, tech_negative_prompt, prompts_count, *args):    47    # Если хотим задать свои триггерные слова или технический промпт48    triggers = triggers if triggers != "" else ",".join(character_configs[name]["trigger_words"])49    tech_prompt = tech_prompt if tech_prompt != "" else character_configs[name]["tech_prompt"]50    tech_negative_prompt = tech_negative_prompt if tech_negative_prompt != "" else character_configs[name]["tech_negative_prompt"]51 52    model = character_configs[name]["model"]53    model['loras'] = {name: scale}54    params = character_configs[name]["params"]55    params["cross_attention_kwargs"]['scale'] = scale56 57    boxes = list(args)58    prompts = []59    negative_prompts = []60    for n in range(prompts_count):61        prompts.append(f'{triggers}, {tech_prompt}, {boxes[n]}')62        negative_prompts.append(f"{tech_negative_prompt}")63    64    request_data = {65        "input": {66            "model": model,67            "params": params,68            "prompt": prompts,69            "negative_prompt": negative_prompts,70            "height": 1216,71            "width": 832,72        }73    }74    75    response = request_to_endpoint(request_data)76 77    images = []78    for base64_string in response.json()['output']['images']:79        img = Image.open(BytesIO(base64.b64decode(base64_string)))80        images.append(img)81 82    gallery = [[images[i], f"{prompts[i]}"] for i in range(prompts_count)]83 84    return gallery85 86 87 88 89 90 91                            ######################################################92                            #   ____               _ _         _                 #93                            #  / ___|_ __ __ _  __| (_) ___   / \   _ __  _ __   #94                            # | |  _| '__/ _` |/ _` | |/ _ \ / _ \ | '_ \| '_ \  #95                            # | |_| | | | (_| | (_| | | (_) / ___ \| |_) | |_) | #96                            #  \____|_|  \__,_|\__,_|_|\___/_/   \_\ .__/| .__/  #97                            #                                      |_|   |_|     #98############################################################################################################99with gr.Blocks() as demo:100    with gr.Group():101        name = gr.Radio(102            character_names,103            label="Select character:",104            interactive=True,105            visible=True,106        )107 108        scale = gr.Slider(109            minimum=0, 110            maximum=2.0,111            value=0.75, 112            step=0.01, 113            label="Selected LoRA scale:", 114            interactive=True,115        )116 117        with gr.Accordion(open=False):118            triggers = gr.Textbox(label=f"Trigger words:")119            tech_prompt = gr.Textbox(label=f"Technical prompt:")120            tech_negative_prompt = gr.Textbox(label=f"Negative technical prompt:")121 122    prompts_count = gr.State(1)123 124    with gr.Group():125        add_btn = gr.Button("Add prompt")126        del_btn = gr.Button("Delete prompt")127            128        add_btn.click(lambda x: x + 1, prompts_count, prompts_count)129        del_btn.click(lambda x: x - 1, prompts_count, prompts_count)130 131        @gr.render(inputs=prompts_count)132        def render_count(count):133            boxes = []134            for i in range(count):135                with gr.Group():136                    prompt = gr.Textbox(key=str(i), label=f"Prompt {i+1}")137                boxes.append(prompt)138            139            generate_btn.click(140                process_input, 141                [name, scale, triggers, tech_prompt, tech_negative_prompt, prompts_count]+boxes,142                output143            )144 145    generate_btn = gr.Button("Generate!")146 147    output = gr.Gallery(148        label="Generation results:",149        object_fit="contain", 150        height="auto",151    )152 153demo.launch()154############################################################################################################