nanankawa/extrasneo-CodeSandBox
0
1from functools import wraps2from flask import (3 Flask,4 jsonify,5 request,6 Response,7 render_template_string,8 abort,9 send_from_directory,10 send_file,11)12from flask_cors import CORS13from flask_compress import Compress14import markdown15import argparse16from transformers import AutoTokenizer, AutoProcessor, pipeline17from transformers import AutoModelForCausalLM, AutoModelForSeq2SeqLM18from transformers import BlipForConditionalGeneration19import unicodedata20import torch21import time22import os23import gc24import sys25import secrets26from PIL import Image27import base6428from io import BytesIO29from random import randint30import webuiapi31import hashlib32from constants import *33from colorama import Fore, Style, init as colorama_init34 35colorama_init()36 37if sys.hexversion < 0x030b0000:38 print(f"{Fore.BLUE}{Style.BRIGHT}Python 3.11 or newer is recommended to run this program.{Style.RESET_ALL}")39 time.sleep(2)40 41class SplitArgs(argparse.Action):42 def __call__(self, parser, namespace, values, option_string=None):43 setattr(44 namespace, self.dest, values.replace('"', "").replace("'", "").split(",")45 )46 47#Setting Root Folders for Silero Generations so it is compatible with STSL, should not effect regular runs. - Rolyat48parent_dir = os.path.dirname(os.path.abspath(__file__))49SILERO_SAMPLES_PATH = os.path.join(parent_dir, "tts_samples")50SILERO_SAMPLE_TEXT = os.path.join(parent_dir)51 52# Create directories if they don't exist53if not os.path.exists(SILERO_SAMPLES_PATH):54 os.makedirs(SILERO_SAMPLES_PATH)55if not os.path.exists(SILERO_SAMPLE_TEXT):56 os.makedirs(SILERO_SAMPLE_TEXT)57 58# Script arguments59parser = argparse.ArgumentParser(60 prog="SillyTavern Extras", description="Web API for transformers models"61)62parser.add_argument(63 "--port", type=int, help="Specify the port on which the application is hosted"64)65parser.add_argument(66 "--listen", action="store_true", help="Host the app on the local network"67)68parser.add_argument(69 "--share", action="store_true", help="Share the app on CloudFlare tunnel"70)71parser.add_argument("--cpu", action="store_true", help="Run the models on the CPU")72parser.add_argument("--cuda", action="store_false", dest="cpu", help="Run the models on the GPU")73parser.add_argument("--cuda-device", help="Specify the CUDA device to use")74parser.add_argument("--mps", "--apple", "--m1", "--m2", action="store_false", dest="cpu", help="Run the models on Apple Silicon")75parser.set_defaults(cpu=True)76parser.add_argument("--summarization-model", help="Load a custom summarization model")77parser.add_argument(78 "--classification-model", help="Load a custom text classification model"79)80parser.add_argument("--captioning-model", help="Load a custom captioning model")81parser.add_argument("--embedding-model", help="Load a custom text embedding model")82parser.add_argument("--chroma-host", help="Host IP for a remote ChromaDB instance")83parser.add_argument("--chroma-port", help="HTTP port for a remote ChromaDB instance (defaults to 8000)")84parser.add_argument("--chroma-folder", help="Path for chromadb persistence folder", default='.chroma_db')85parser.add_argument('--chroma-persist', help="ChromaDB persistence", default=True, action=argparse.BooleanOptionalAction)86parser.add_argument(87 "--secure", action="store_true", help="Enforces the use of an API key"88)89parser.add_argument("--talkinghead-gpu", action="store_true", help="Run the talkinghead animation on the GPU (CPU is default)")90 91parser.add_argument("--coqui-gpu", action="store_true", help="Run the voice models on the GPU (CPU is default)")92parser.add_argument("--coqui-models", help="Install given Coqui-api TTS model at launch (comma separated list, last one will be loaded at start)")93 94parser.add_argument("--max-content-length", help="Set the max")95parser.add_argument("--rvc-save-file", action="store_true", help="Save the last rvc input/output audio file into data/tmp/ folder (for research)")96 97parser.add_argument("--stt-vosk-model-path", help="Load a custom vosk speech-to-text model")98parser.add_argument("--stt-whisper-model-path", help="Load a custom vosk speech-to-text model")99sd_group = parser.add_mutually_exclusive_group()100 101local_sd = parser.add_argument_group("sd-local")102local_sd.add_argument("--sd-model", help="Load a custom SD image generation model")103local_sd.add_argument("--sd-cpu", help="Force the SD pipeline to run on the CPU", action="store_true")104 105remote_sd = parser.add_argument_group("sd-remote")106remote_sd.add_argument(107 "--sd-remote", action="store_true", help="Use a remote backend for SD"108)109remote_sd.add_argument(110 "--sd-remote-host", type=str, help="Specify the host of the remote SD backend"111)112remote_sd.add_argument(113 "--sd-remote-port", type=int, help="Specify the port of the remote SD backend"114)115remote_sd.add_argument(116 "--sd-remote-ssl", action="store_true", help="Use SSL for the remote SD backend"117)118remote_sd.add_argument(119 "--sd-remote-auth",120 type=str,121 help="Specify the username:password for the remote SD backend (if required)",122)123 124parser.add_argument(125 "--enable-modules",126 action=SplitArgs,127 default=[],128 help="Override a list of enabled modules",129)130 131args = parser.parse_args()132# [HF, Huggingface] Set port to 7860, set host to remote. 133port = 7860134host = "0.0.0.0"135summarization_model = (136 args.summarization_model137 if args.summarization_model138 else DEFAULT_SUMMARIZATION_MODEL139)140classification_model = (141 args.classification_model142 if args.classification_model143 else DEFAULT_CLASSIFICATION_MODEL144)145captioning_model = (146 args.captioning_model if args.captioning_model else DEFAULT_CAPTIONING_MODEL147)148embedding_model = (149 args.embedding_model if args.embedding_model else DEFAULT_EMBEDDING_MODEL150)151 152sd_use_remote = False if args.sd_model else True153sd_model = args.sd_model if args.sd_model else DEFAULT_SD_MODEL154sd_remote_host = args.sd_remote_host if args.sd_remote_host else DEFAULT_REMOTE_SD_HOST155sd_remote_port = args.sd_remote_port if args.sd_remote_port else DEFAULT_REMOTE_SD_PORT156sd_remote_ssl = args.sd_remote_ssl157sd_remote_auth = args.sd_remote_auth158 159modules = (160 args.enable_modules if args.enable_modules and len(args.enable_modules) > 0 else []161)162 163if len(modules) == 0:164 print(165 f"{Fore.RED}{Style.BRIGHT}You did not select any modules to run! Choose them by adding an --enable-modules option"166 )167 print(f"Example: --enable-modules=caption,summarize{Style.RESET_ALL}")168 169# Models init170cuda_device = DEFAULT_CUDA_DEVICE if not args.cuda_device else args.cuda_device171device_string = cuda_device if torch.cuda.is_available() and not args.cpu else 'mps' if torch.backends.mps.is_available() and not args.cpu else 'cpu'172device = torch.device(device_string)173torch_dtype = torch.float32 if device_string != cuda_device else torch.float16174 175if not torch.cuda.is_available() and not args.cpu:176 print(f"{Fore.YELLOW}{Style.BRIGHT}torch-cuda is not supported on this device.{Style.RESET_ALL}")177 if not torch.backends.mps.is_available() and not args.cpu:178 print(f"{Fore.YELLOW}{Style.BRIGHT}torch-mps is not supported on this device.{Style.RESET_ALL}")179 180 181print(f"{Fore.GREEN}{Style.BRIGHT}Using torch device: {device_string}{Style.RESET_ALL}")182 183if "talkinghead" in modules:184 import sys185 import threading186 mode = "cuda" if args.talkinghead_gpu else "cpu"187 print("Initializing talkinghead pipeline in " + mode + " mode....")188 talkinghead_path = os.path.abspath(os.path.join(os.getcwd(), "talkinghead"))189 sys.path.append(talkinghead_path) # Add the path to the 'tha3' module to the sys.path list190 191 try:192 import talkinghead.tha3.app.app as talkinghead193 from talkinghead import *194 def launch_talkinghead_gui():195 talkinghead.launch_gui(mode, "separable_float")196 #choices=['standard_float', 'separable_float', 'standard_half', 'separable_half'],197 #choices='The device to use for PyTorch ("cuda" for GPU, "cpu" for CPU).'198 talkinghead_thread = threading.Thread(target=launch_talkinghead_gui)199 talkinghead_thread.daemon = True # Set the thread as a daemon thread200 talkinghead_thread.start()201 202 except ModuleNotFoundError:203 print("Error: Could not import the 'talkinghead' module.")204 205if "caption" in modules:206 print("Initializing an image captioning model...")207 captioning_processor = AutoProcessor.from_pretrained(captioning_model)208 if "blip" in captioning_model:209 captioning_transformer = BlipForConditionalGeneration.from_pretrained(210 captioning_model, torch_dtype=torch_dtype211 ).to(device)212 else:213 captioning_transformer = AutoModelForCausalLM.from_pretrained(214 captioning_model, torch_dtype=torch_dtype215 ).to(device)216 217if "summarize" in modules:218 print("Initializing a text summarization model...")219 summarization_tokenizer = AutoTokenizer.from_pretrained(summarization_model)220 summarization_transformer = AutoModelForSeq2SeqLM.from_pretrained(221 summarization_model, torch_dtype=torch_dtype222 ).to(device)223 224if "classify" in modules:225 print("Initializing a sentiment classification pipeline...")226 classification_pipe = pipeline(227 "text-classification",228 model=classification_model,229 top_k=None,230 device=device,231 torch_dtype=torch_dtype,232 )233 234if "sd" in modules and not sd_use_remote:235 from diffusers import StableDiffusionPipeline236 from diffusers import EulerAncestralDiscreteScheduler237 238 print("Initializing Stable Diffusion pipeline...")239 sd_device_string = cuda_device if torch.cuda.is_available() else 'mps' if torch.backends.mps.is_available() else 'cpu'240 sd_device = torch.device(sd_device_string)241 sd_torch_dtype = torch.float32 if sd_device_string != cuda_device else torch.float16242 sd_pipe = StableDiffusionPipeline.from_pretrained(243 sd_model, custom_pipeline="lpw_stable_diffusion", torch_dtype=sd_torch_dtype244 ).to(sd_device)245 sd_pipe.safety_checker = lambda images, clip_input: (images, False)246 sd_pipe.enable_attention_slicing()247 # pipe.scheduler = KarrasVeScheduler.from_config(pipe.scheduler.config)248 sd_pipe.scheduler = EulerAncestralDiscreteScheduler.from_config(249 sd_pipe.scheduler.config250 )251elif "sd" in modules and sd_use_remote:252 print("Initializing Stable Diffusion connection")253 try:254 sd_remote = webuiapi.WebUIApi(255 host=sd_remote_host, port=sd_remote_port, use_https=sd_remote_ssl256 )257 if sd_remote_auth:258 username, password = sd_remote_auth.split(":")259 sd_remote.set_auth(username, password)260 sd_remote.util_wait_for_ready()261 except Exception as e:262 # remote sd from modules263 print(264 f"{Fore.RED}{Style.BRIGHT}Could not connect to remote SD backend at http{'s' if sd_remote_ssl else ''}://{sd_remote_host}:{sd_remote_port}! Disabling SD module...{Style.RESET_ALL}"265 )266 modules.remove("sd")267 268if "tts" in modules:269 print("tts module is deprecated. Please use silero-tts instead.")270 modules.remove("tts")271 modules.append("silero-tts")272 273 274if "silero-tts" in modules:275 if not os.path.exists(SILERO_SAMPLES_PATH):276 os.makedirs(SILERO_SAMPLES_PATH)277 print("Initializing Silero TTS server")278 from silero_api_server import tts279 280 tts_service = tts.SileroTtsService(SILERO_SAMPLES_PATH)281 if len(os.listdir(SILERO_SAMPLES_PATH)) == 0:282 print("Generating Silero TTS samples...")283 tts_service.update_sample_text(SILERO_SAMPLE_TEXT)284 tts_service.generate_samples()285 286if "edge-tts" in modules:287 print("Initializing Edge TTS client")288 import tts_edge as edge289 290 291if "chromadb" in modules:292 print("Initializing ChromaDB")293 import chromadb294 import posthog295 from chromadb.config import Settings296 from sentence_transformers import SentenceTransformer297 298 # Assume that the user wants in-memory unless a host is specified299 # Also disable chromadb telemetry300 posthog.capture = lambda *args, **kwargs: None301 if args.chroma_host is None:302 if args.chroma_persist:303 chromadb_client = chromadb.PersistentClient(path=args.chroma_folder, settings=Settings(anonymized_telemetry=False))304 print(f"ChromaDB is running in-memory with persistence. Persistence is stored in {args.chroma_folder}. Can be cleared by deleting the folder or purging db.")305 else:306 chromadb_client = chromadb.EphemeralClient(Settings(anonymized_telemetry=False))307 print(f"ChromaDB is running in-memory without persistence.")308 else:309 chroma_port=(310 args.chroma_port if args.chroma_port else DEFAULT_CHROMA_PORT311 )312 chromadb_client = chromadb.HttpClient(host=args.chroma_host, port=chroma_port, settings=Settings(anonymized_telemetry=False))313 print(f"ChromaDB is remotely configured at {args.chroma_host}:{chroma_port}")314 315 chromadb_embedder = SentenceTransformer(embedding_model, device=device_string)316 chromadb_embed_fn = lambda *args, **kwargs: chromadb_embedder.encode(*args, **kwargs).tolist()317 318 # Check if the db is connected and running, otherwise tell the user319 try:320 chromadb_client.heartbeat()321 print("Successfully pinged ChromaDB! Your client is successfully connected.")322 except:323 print("Could not ping ChromaDB! If you are running remotely, please check your host and port!")324 325# Flask init326app = Flask(__name__)327CORS(app) # allow cross-domain requests328Compress(app) # compress responses329app.config["MAX_CONTENT_LENGTH"] = 100 * 1024 * 1024330 331max_content_length = (332 args.max_content_length333 if args.max_content_length334 else None)335 336if max_content_length is not None:337 print("Setting MAX_CONTENT_LENGTH to",max_content_length,"Mb")338 app.config["MAX_CONTENT_LENGTH"] = int(max_content_length) * 1024 * 1024339 340if "vosk-stt" in modules:341 print("Initializing Vosk speech-recognition (from ST request file)")342 vosk_model_path = (343 args.stt_vosk_model_path344 if args.stt_vosk_model_path345 else None)346 347 import modules.speech_recognition.vosk_module as vosk_module348 349 vosk_module.model = vosk_module.load_model(file_path=vosk_model_path)350 app.add_url_rule("/api/speech-recognition/vosk/process-audio", view_func=vosk_module.process_audio, methods=["POST"])351 352if "whisper-stt" in modules:353 print("Initializing Whisper speech-recognition (from ST request file)")354 whisper_model_path = (355 args.stt_whisper_model_path356 if args.stt_whisper_model_path357 else None)358 359 import modules.speech_recognition.whisper_module as whisper_module360 361 whisper_module.model = whisper_module.load_model(file_path=whisper_model_path)362 app.add_url_rule("/api/speech-recognition/whisper/process-audio", view_func=whisper_module.process_audio, methods=["POST"])363 364if "streaming-stt" in modules:365 print("Initializing vosk/whisper speech-recognition (from extras server microphone)")366 whisper_model_path = (367 args.stt_whisper_model_path368 if args.stt_whisper_model_path369 else None)370 371 import modules.speech_recognition.streaming_module as streaming_module372 373 streaming_module.whisper_model, streaming_module.vosk_model = streaming_module.load_model(file_path=whisper_model_path)374 app.add_url_rule("/api/speech-recognition/streaming/record-and-transcript", view_func=streaming_module.record_and_transcript, methods=["POST"])375 376if "rvc" in modules:377 print("Initializing RVC voice conversion (from ST request file)")378 print("Increasing server upload limit")379 rvc_save_file = (380 args.rvc_save_file381 if args.rvc_save_file382 else False)383 384 if rvc_save_file:385 print("RVC saving file option detected, input/output audio will be savec into data/tmp/ folder")386 387 import sys388 sys.path.insert(0,'modules/voice_conversion')389 390 import modules.voice_conversion.rvc_module as rvc_module391 rvc_module.save_file = rvc_save_file392 rvc_module.fix_model_install()393 app.add_url_rule("/api/voice-conversion/rvc/get-models-list", view_func=rvc_module.rvc_get_models_list, methods=["POST"])394 app.add_url_rule("/api/voice-conversion/rvc/upload-models", view_func=rvc_module.rvc_upload_models, methods=["POST"])395 app.add_url_rule("/api/voice-conversion/rvc/process-audio", view_func=rvc_module.rvc_process_audio, methods=["POST"])396 397 398if "coqui-tts" in modules:399 mode = "GPU" if args.coqui_gpu else "CPU"400 print("Initializing Coqui TTS client in " + mode + " mode")401 import modules.text_to_speech.coqui.coqui_module as coqui_module402 403 if mode == "GPU":404 coqui_module.gpu_mode = True405 406 coqui_models = (407 args.coqui_models408 if args.coqui_models409 else None410 )411 412 if coqui_models is not None:413 coqui_models = coqui_models.split(",")414 for i in coqui_models:415 if not coqui_module.install_model(i):416 raise ValueError("Coqui model loading failed, most likely a wrong model name in --coqui-models argument, check log above to see which one")417 418 # Coqui-api models419 app.add_url_rule("/api/text-to-speech/coqui/coqui-api/check-model-state", view_func=coqui_module.coqui_check_model_state, methods=["POST"])420 app.add_url_rule("/api/text-to-speech/coqui/coqui-api/install-model", view_func=coqui_module.coqui_install_model, methods=["POST"])421 422 # Users models423 app.add_url_rule("/api/text-to-speech/coqui/local/get-models", view_func=coqui_module.coqui_get_local_models, methods=["POST"])424 425 # Handle both coqui-api/users models426 app.add_url_rule("/api/text-to-speech/coqui/generate-tts", view_func=coqui_module.coqui_generate_tts, methods=["POST"])427 428def require_module(name):429 def wrapper(fn):430 @wraps(fn)431 def decorated_view(*args, **kwargs):432 if name not in modules:433 abort(403, "Module is disabled by config")434 return fn(*args, **kwargs)435 436 return decorated_view437 438 return wrapper439 440 441# AI stuff442def classify_text(text: str) -> list:443 output = classification_pipe(444 text,445 truncation=True,446 max_length=classification_pipe.model.config.max_position_embeddings,447 )[0]448 return sorted(output, key=lambda x: x["score"], reverse=True)449 450 451def caption_image(raw_image: Image, max_new_tokens: int = 20) -> str:452 inputs = captioning_processor(raw_image.convert("RGB"), return_tensors="pt").to(453 device, torch_dtype454 )455 outputs = captioning_transformer.generate(**inputs, max_new_tokens=max_new_tokens)456 caption = captioning_processor.decode(outputs[0], skip_special_tokens=True)457 return caption458 459 460def summarize_chunks(text: str, params: dict) -> str:461 try:462 return summarize(text, params)463 except IndexError:464 print(465 "Sequence length too large for model, cutting text in half and calling again"466 )467 new_params = params.copy()468 new_params["max_length"] = new_params["max_length"] // 2469 new_params["min_length"] = new_params["min_length"] // 2470 return summarize_chunks(471 text[: (len(text) // 2)], new_params472 ) + summarize_chunks(text[(len(text) // 2) :], new_params)473 474 475def summarize(text: str, params: dict) -> str:476 # Tokenize input477 inputs = summarization_tokenizer(text, return_tensors="pt").to(device)478 token_count = len(inputs[0])479 480 bad_words_ids = [481 summarization_tokenizer(bad_word, add_special_tokens=False).input_ids482 for bad_word in params["bad_words"]483 ]484 summary_ids = summarization_transformer.generate(485 inputs["input_ids"],486 num_beams=2,487 max_new_tokens=max(token_count, int(params["max_length"])),488 min_new_tokens=min(token_count, int(params["min_length"])),489 repetition_penalty=float(params["repetition_penalty"]),490 temperature=float(params["temperature"]),491 length_penalty=float(params["length_penalty"]),492 bad_words_ids=bad_words_ids,493 )494 summary = summarization_tokenizer.batch_decode(495 summary_ids, skip_special_tokens=True, clean_up_tokenization_spaces=True496 )[0]497 summary = normalize_string(summary)498 return summary499 500 501def normalize_string(input: str) -> str:502 output = " ".join(unicodedata.normalize("NFKC", input).strip().split())503 return output504 505 506def generate_image(data: dict) -> Image:507 prompt = normalize_string(f'{data["prompt_prefix"]} {data["prompt"]}')508 509 if sd_use_remote:510 image = sd_remote.txt2img(511 prompt=prompt,512 negative_prompt=data["negative_prompt"],513 sampler_name=data["sampler"],514 steps=data["steps"],515 cfg_scale=data["scale"],516 width=data["width"],517 height=data["height"],518 restore_faces=data["restore_faces"],519 enable_hr=data["enable_hr"],520 save_images=True,521 send_images=True,522 do_not_save_grid=False,523 do_not_save_samples=False,524 ).image525 else:526 image = sd_pipe(527 prompt=prompt,528 negative_prompt=data["negative_prompt"],529 num_inference_steps=data["steps"],530 guidance_scale=data["scale"],531 width=data["width"],532 height=data["height"],533 ).images[0]534 535 image.save("./debug.png")536 return image537 538 539def image_to_base64(image: Image, quality: int = 75) -> str:540 buffer = BytesIO()541 image.convert("RGB")542 image.save(buffer, format="JPEG", quality=quality)543 img_str = base64.b64encode(buffer.getvalue()).decode("utf-8")544 return img_str545 546ignore_auth = [] 547# [HF, Huggingface] Get password instead of text file.548api_key = os.environ.get("password")549 550def is_authorize_ignored(request):551 view_func = app.view_functions.get(request.endpoint)552 553 if view_func is not None:554 if view_func in ignore_auth:555 return True556 return False557 558@app.before_request559def before_request():560 # Request time measuring561 request.start_time = time.time()562 563 # Checks if an API key is present and valid, otherwise return unauthorized564 # The options check is required so CORS doesn't get angry565 try:566 if request.method != 'OPTIONS' and is_authorize_ignored(request) == False and getattr(request.authorization, 'token', '') != api_key:567 print(f"WARNING: Unauthorized API key access from {request.remote_addr}")568 if request.method == 'POST':569 print(f"Incoming POST request with {request.headers.get('Authorization')}")570 response = jsonify({ 'error': '401: Invalid API key' })571 response.status_code = 401572 return "https://(hf_name)-(space_name).hf.space/"573 except Exception as e:574 print(f"API key check error: {e}")575 return "https://(hf_name)-(space_name).hf.space/"576 577 578@app.after_request579def after_request(response):580 duration = time.time() - request.start_time581 response.headers["X-Request-Duration"] = str(duration)582 return response583 584 585@app.route("/", methods=["GET"])586def index():587 with open("./README.md", "r", encoding="utf8") as f:588 content = f.read()589 return render_template_string(markdown.markdown(content, extensions=["tables"]))590 591 592@app.route("/api/extensions", methods=["GET"])593def get_extensions():594 extensions = dict(595 {596 "extensions": [597 {598 "name": "not-supported",599 "metadata": {600 "display_name": """<span style="white-space:break-spaces;">Extensions serving using Extensions API is no longer supported. Please update the mod from: <a href="https://github.com/Cohee1207/SillyTavern">https://github.com/Cohee1207/SillyTavern</a></span>""",601 "requires": [],602 "assets": [],603 },604 }605 ]606 }607 )608 return jsonify(extensions)609 610 611@app.route("/api/caption", methods=["POST"])612@require_module("caption")613def api_caption():614 data = request.get_json()615 616 if "image" not in data or not isinstance(data["image"], str):617 abort(400, '"image" is required')618 619 image = Image.open(BytesIO(base64.b64decode(data["image"])))620 image = image.convert("RGB")621 image.thumbnail((512, 512))622 caption = caption_image(image)623 thumbnail = image_to_base64(image)624 print("Caption:", caption, sep="\n")625 gc.collect()626 return jsonify({"caption": caption, "thumbnail": thumbnail})627 628 629@app.route("/api/summarize", methods=["POST"])630@require_module("summarize")631def api_summarize():632 data = request.get_json()633 634 if "text" not in data or not isinstance(data["text"], str):635 abort(400, '"text" is required')636 637 params = DEFAULT_SUMMARIZE_PARAMS.copy()638 639 if "params" in data and isinstance(data["params"], dict):640 params.update(data["params"])641 642 print("Summary input:", data["text"], sep="\n")643 summary = summarize_chunks(data["text"], params)644 print("Summary output:", summary, sep="\n")645 gc.collect()646 return jsonify({"summary": summary})647 648 649@app.route("/api/classify", methods=["POST"])650@require_module("classify")651def api_classify():652 data = request.get_json()653 654 if "text" not in data or not isinstance(data["text"], str):655 abort(400, '"text" is required')656 657 print("Classification input:", data["text"], sep="\n")658 classification = classify_text(data["text"])659 print("Classification output:", classification, sep="\n")660 gc.collect()661 if "talkinghead" in modules: #send emotion to talkinghead662 talkinghead.setEmotion(classification)663 return jsonify({"classification": classification})664 665 666@app.route("/api/classify/labels", methods=["GET"])667@require_module("classify")668def api_classify_labels():669 classification = classify_text("")670 labels = [x["label"] for x in classification]671 if "talkinghead" in modules:672 labels.append('talkinghead') # Add 'talkinghead' to the labels list673 return jsonify({"labels": labels})674 675@app.route("/api/talkinghead/load", methods=["POST"])676def live_load():677 file = request.files['file']678 # convert stream to bytes and pass to talkinghead_load679 return talkinghead.talkinghead_load_file(file.stream)680 681@app.route('/api/talkinghead/unload')682def live_unload():683 return talkinghead.unload()684 685@app.route('/api/talkinghead/start_talking')686def start_talking():687 return talkinghead.start_talking()688 689@app.route('/api/talkinghead/stop_talking')690def stop_talking():691 return talkinghead.stop_talking()692 693@app.route('/api/talkinghead/result_feed')694def result_feed():695 return talkinghead.result_feed()696 697@app.route("/api/image", methods=["POST"])698@require_module("sd")699def api_image():700 required_fields = {701 "prompt": str,702 }703 704 optional_fields = {705 "steps": 30,706 "scale": 6,707 "sampler": "DDIM",708 "width": 512,709 "height": 512,710 "restore_faces": False,711 "enable_hr": False,712 "prompt_prefix": PROMPT_PREFIX,713 "negative_prompt": NEGATIVE_PROMPT,714 }715 716 data = request.get_json()717 718 # Check required fields719 for field, field_type in required_fields.items():720 if field not in data or not isinstance(data[field], field_type):721 abort(400, f'"{field}" is required')722 723 # Set optional fields to default values if not provided724 for field, default_value in optional_fields.items():725 type_match = (726 (int, float)727 if isinstance(default_value, (int, float))728 else type(default_value)729 )730 if field not in data or not isinstance(data[field], type_match):731 data[field] = default_value732 733 try:734 print("SD inputs:", data, sep="\n")735 image = generate_image(data)736 base64image = image_to_base64(image, quality=90)737 return jsonify({"image": base64image})738 except RuntimeError as e:739 abort(400, str(e))740 741 742@app.route("/api/image/model", methods=["POST"])743@require_module("sd")744def api_image_model_set():745 data = request.get_json()746 747 if not sd_use_remote:748 abort(400, "Changing model for local sd is not supported.")749 if "model" not in data or not isinstance(data["model"], str):750 abort(400, '"model" is required')751 752 old_model = sd_remote.util_get_current_model()753 sd_remote.util_set_model(data["model"], find_closest=False)754 # sd_remote.util_set_model(data['model'])755 sd_remote.util_wait_for_ready()756 new_model = sd_remote.util_get_current_model()757 758 return jsonify({"previous_model": old_model, "current_model": new_model})759 760 761@app.route("/api/image/model", methods=["GET"])762@require_module("sd")763def api_image_model_get():764 model = sd_model765 766 if sd_use_remote:767 model = sd_remote.util_get_current_model()768 769 return jsonify({"model": model})770 771 772@app.route("/api/image/models", methods=["GET"])773@require_module("sd")774def api_image_models():775 models = [sd_model]776 777 if sd_use_remote:778 models = sd_remote.util_get_model_names()779 780 return jsonify({"models": models})781 782 783@app.route("/api/image/samplers", methods=["GET"])784@require_module("sd")785def api_image_samplers():786 samplers = ["Euler a"]787 788 if sd_use_remote:789 samplers = [sampler["name"] for sampler in sd_remote.get_samplers()]790 791 return jsonify({"samplers": samplers})792 793 794@app.route("/api/modules", methods=["GET"])795def get_modules():796 return jsonify({"modules": modules})797 798 799@app.route("/api/tts/speakers", methods=["GET"])800@require_module("silero-tts")801def tts_speakers():802 voices = [803 {804 "name": speaker,805 "voice_id": speaker,806 "preview_url": f"{str(request.url_root)}api/tts/sample/{speaker}",807 }808 for speaker in tts_service.get_speakers()809 ]810 return jsonify(voices)811 812# Added fix for Silero not working as new files were unable to be created if one already existed. - Rolyat 7/7/23813@app.route("/api/tts/generate", methods=["POST"])814@require_module("silero-tts")815def tts_generate():816 voice = request.get_json()817 if "text" not in voice or not isinstance(voice["text"], str):818 abort(400, '"text" is required')819 if "speaker" not in voice or not isinstance(voice["speaker"], str):820 abort(400, '"speaker" is required')821 # Remove asterisks822 voice["text"] = voice["text"].replace("*", "")823 try:824 # Remove the destination file if it already exists825 if os.path.exists('test.wav'):826 os.remove('test.wav')827 828 audio = tts_service.generate(voice["speaker"], voice["text"])829 audio_file_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), os.path.basename(audio))830 831 os.rename(audio, audio_file_path)832 return send_file(audio_file_path, mimetype="audio/x-wav")833 except Exception as e:834 print(e)835 abort(500, voice["speaker"])836 837 838@app.route("/api/tts/sample/<speaker>", methods=["GET"])839@require_module("silero-tts")840def tts_play_sample(speaker: str):841 return send_from_directory(SILERO_SAMPLES_PATH, f"{speaker}.wav")842 843 844@app.route("/api/edge-tts/list", methods=["GET"])845@require_module("edge-tts")846def edge_tts_list():847 voices = edge.get_voices()848 return jsonify(voices)849 850 851@app.route("/api/edge-tts/generate", methods=["POST"])852@require_module("edge-tts")853def edge_tts_generate():854 data = request.get_json()855 if "text" not in data or not isinstance(data["text"], str):856 abort(400, '"text" is required')857 if "voice" not in data or not isinstance(data["voice"], str):858 abort(400, '"voice" is required')859 if "rate" in data and isinstance(data['rate'], int):860 rate = data['rate']861 else:862 rate = 0863 # Remove asterisks864 data["text"] = data["text"].replace("*", "")865 try:866 audio = edge.generate_audio(text=data["text"], voice=data["voice"], rate=rate)867 return Response(audio, mimetype="audio/mpeg")868 except Exception as e:869 print(e)870 abort(500, data["voice"])871 872 873@app.route("/api/chromadb", methods=["POST"])874@require_module("chromadb")875def chromadb_add_messages():876 data = request.get_json()877 if "chat_id" not in data or not isinstance(data["chat_id"], str):878 abort(400, '"chat_id" is required')879 if "messages" not in data or not isinstance(data["messages"], list):880 abort(400, '"messages" is required')881 882 chat_id_md5 = hashlib.md5(data["chat_id"].encode()).hexdigest()883 collection = chromadb_client.get_or_create_collection(884 name=f"chat-{chat_id_md5}", embedding_function=chromadb_embed_fn885 )886 887 documents = [m["content"] for m in data["messages"]]888 ids = [m["id"] for m in data["messages"]]889 metadatas = [890 {"role": m["role"], "date": m["date"], "meta": m.get("meta", "")}891 for m in data["messages"]892 ]893 894 collection.upsert(895 ids=ids,896 documents=documents,897 metadatas=metadatas,898 )899 900 return jsonify({"count": len(ids)})901 902 903@app.route("/api/chromadb/purge", methods=["POST"])904@require_module("chromadb")905def chromadb_purge():906 data = request.get_json()907 if "chat_id" not in data or not isinstance(data["chat_id"], str):908 abort(400, '"chat_id" is required')909 910 chat_id_md5 = hashlib.md5(data["chat_id"].encode()).hexdigest()911 collection = chromadb_client.get_or_create_collection(912 name=f"chat-{chat_id_md5}", embedding_function=chromadb_embed_fn913 )914 915 count = collection.count()916 collection.delete()917 print("ChromaDB embeddings deleted", count)918 return 'Ok', 200919 920 921@app.route("/api/chromadb/query", methods=["POST"])922@require_module("chromadb")923def chromadb_query():924 data = request.get_json()925 if "chat_id" not in data or not isinstance(data["chat_id"], str):926 abort(400, '"chat_id" is required')927 if "query" not in data or not isinstance(data["query"], str):928 abort(400, '"query" is required')929 930 if "n_results" not in data or not isinstance(data["n_results"], int):931 n_results = 1932 else:933 n_results = data["n_results"]934 935 chat_id_md5 = hashlib.md5(data["chat_id"].encode()).hexdigest()936 collection = chromadb_client.get_or_create_collection(937 name=f"chat-{chat_id_md5}", embedding_function=chromadb_embed_fn938 )939 940 if collection.count() == 0:941 print(f"Queried empty/missing collection for {repr(data['chat_id'])}.")942 return jsonify([])943 944 945 n_results = min(collection.count(), n_results)946 query_result = collection.query(947 query_texts=[data["query"]],948 n_results=n_results,949 )950 951 documents = query_result["documents"][0]952 ids = query_result["ids"][0]953 metadatas = query_result["metadatas"][0]954 distances = query_result["distances"][0]955 956 messages = [957 {958 "id": ids[i],959 "date": metadatas[i]["date"],960 "role": metadatas[i]["role"],961 "meta": metadatas[i]["meta"],962 "content": documents[i],963 "distance": distances[i],964 }965 for i in range(len(ids))966 ]967 968 return jsonify(messages)969 970@app.route("/api/chromadb/multiquery", methods=["POST"])971@require_module("chromadb")972def chromadb_multiquery():973 data = request.get_json()974 if "chat_list" not in data or not isinstance(data["chat_list"], list):975 abort(400, '"chat_list" is required and should be a list')976 if "query" not in data or not isinstance(data["query"], str):977 abort(400, '"query" is required')978 979 if "n_results" not in data or not isinstance(data["n_results"], int):980 n_results = 1981 else:982 n_results = data["n_results"]983 984 messages = []985 986 for chat_id in data["chat_list"]:987 if not isinstance(chat_id, str):988 continue989 990 try:991 chat_id_md5 = hashlib.md5(chat_id.encode()).hexdigest()992 collection = chromadb_client.get_collection(993 name=f"chat-{chat_id_md5}", embedding_function=chromadb_embed_fn994 )995 996 # Skip this chat if the collection is empty997 if collection.count() == 0:998 continue999 1000 n_results_per_chat = min(collection.count(), n_results)1001 query_result = collection.query(1002 query_texts=[data["query"]],1003 n_results=n_results_per_chat,1004 )1005 documents = query_result["documents"][0]1006 ids = query_result["ids"][0]1007 metadatas = query_result["metadatas"][0]1008 distances = query_result["distances"][0]1009 1010 chat_messages = [1011 {1012 "id": ids[i],1013 "date": metadatas[i]["date"],1014 "role": metadatas[i]["role"],1015 "meta": metadatas[i]["meta"],1016 "content": documents[i],1017 "distance": distances[i],1018 }1019 for i in range(len(ids))1020 ]1021 1022 messages.extend(chat_messages)1023 except Exception as e:1024 print(e)1025 1026 #remove duplicate msgs, filter down to the right number1027 seen = set()1028 messages = [d for d in messages if not (d['content'] in seen or seen.add(d['content']))]1029 messages = sorted(messages, key=lambda x: x['distance'])[0:n_results]1030 1031 return jsonify(messages)1032 1033 1034@app.route("/api/chromadb/export", methods=["POST"])1035@require_module("chromadb")1036def chromadb_export():1037 data = request.get_json()1038 if "chat_id" not in data or not isinstance(data["chat_id"], str):1039 abort(400, '"chat_id" is required')1040 1041 chat_id_md5 = hashlib.md5(data["chat_id"].encode()).hexdigest()1042 try:1043 collection = chromadb_client.get_collection(1044 name=f"chat-{chat_id_md5}", embedding_function=chromadb_embed_fn1045 )1046 except Exception as e:1047 print(e)1048 abort(400, "Chat collection not found in chromadb")1049 1050 collection_content = collection.get()1051 documents = collection_content.get('documents', [])1052 ids = collection_content.get('ids', [])1053 metadatas = collection_content.get('metadatas', [])1054 1055 unsorted_content = [1056 {1057 "id": ids[i],1058 "metadata": metadatas[i],1059 "document": documents[i],1060 }1061 for i in range(len(ids))1062 ]1063 1064 sorted_content = sorted(unsorted_content, key=lambda x: x['metadata']['date'])1065 1066 export = {1067 "chat_id": data["chat_id"],1068 "content": sorted_content1069 }1070 1071 return jsonify(export)1072 1073@app.route("/api/chromadb/import", methods=["POST"])1074@require_module("chromadb")1075def chromadb_import():1076 data = request.get_json()1077 content = data['content']1078 if "chat_id" not in data or not isinstance(data["chat_id"], str):1079 abort(400, '"chat_id" is required')1080 1081 chat_id_md5 = hashlib.md5(data["chat_id"].encode()).hexdigest()1082 collection = chromadb_client.get_or_create_collection(1083 name=f"chat-{chat_id_md5}", embedding_function=chromadb_embed_fn1084 )1085 1086 documents = [item['document'] for item in content]1087 metadatas = [item['metadata'] for item in content]1088 ids = [item['id'] for item in content]1089 1090 1091 collection.upsert(documents=documents, metadatas=metadatas, ids=ids)1092 print(f"Imported {len(ids)} (total {collection.count()}) content entries into {repr(data['chat_id'])}")1093 1094 return jsonify({"count": len(ids)})1095 1096 1097if args.share:1098 from flask_cloudflared import _run_cloudflared1099 import inspect1100 1101 sig = inspect.signature(_run_cloudflared)1102 sum = sum(1103 11104 for param in sig.parameters.values()1105 if param.kind == param.POSITIONAL_OR_KEYWORD1106 )1107 if sum > 1:1108 metrics_port = randint(8100, 9000)1109 cloudflare = _run_cloudflared(port, metrics_port)1110 else:1111 cloudflare = _run_cloudflared(port)1112 print("\x1b[32mRunning on", cloudflare + "\x1b[0m")1113 1114ignore_auth.append(tts_play_sample)1115app.run(host=host, port=port)1116 