OnMoon/character_testing_ui
0
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############################################################################################################