Team Ai
Apppublic

nanankawa/extrasneo-CodeSandBox

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
server.py1116 linesDownload Raw Back to root
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