| |
| |
|
|
| import argparse |
| import os |
| import threading |
|
|
| import gradio as gr |
| import numpy as np |
|
|
| from kimodo.model import resolve_target |
|
|
| from .gradio_theme import get_gradio_theme |
|
|
| os.environ["HF_ENABLE_PARALLEL_LOADING"] = "YES" |
| DEFAULT_TEXT = "A person walks and falls to the ground." |
| DEFAULT_SERVER_NAME = "0.0.0.0" |
| DEFAULT_SERVER_PORT = 9550 |
| DEFAULT_TMP_FOLDER = "/tmp/text_encoder/" |
| DEFAULT_TEXT_ENCODER = "llm2vec" |
| TEXT_ENCODER_PRESETS = { |
| "llm2vec": { |
| "target": "kimodo.model.LLM2VecEncoder", |
| "kwargs": { |
| "base_model_name_or_path": "McGill-NLP/LLM2Vec-Meta-Llama-3-8B-Instruct-mntp", |
| "peft_model_name_or_path": "McGill-NLP/LLM2Vec-Meta-Llama-3-8B-Instruct-mntp-supervised", |
| "dtype": "bfloat16", |
| "llm_dim": 4096, |
| "device": "auto", |
| }, |
| "display_name": "LLM2Vec", |
| } |
| } |
|
|
|
|
| class DemoWrapper: |
| def __init__(self, builder, tmp_folder): |
| |
| |
| |
| self._builder = builder |
| self._encoder = None |
| self._lock = threading.Lock() |
| self.tmp_folder = tmp_folder |
|
|
| @property |
| def text_encoder(self): |
| if self._encoder is None: |
| with self._lock: |
| if self._encoder is None: |
| print("[text-encoder] lazy-loading model on first request...", flush=True) |
| self._encoder = self._builder() |
| print("[text-encoder] model loaded", flush=True) |
| return self._encoder |
|
|
| def __call__(self, text, filename, progress=gr.Progress()): |
| |
| tensor, length = self.text_encoder(text) |
| embedding = tensor[:length] |
| embedding = embedding.cpu().numpy() |
|
|
| |
| path = os.path.join(self.tmp_folder, filename) |
| np.save(path, embedding) |
|
|
| output_title = gr.Markdown(visible=True) |
| output_text = gr.Markdown(visible=True, value=f"Text: {text}") |
| download = gr.DownloadButton(visible=True, value=path) |
| return download, output_title, output_text |
|
|
|
|
| def _get_env(name: str, default): |
| return os.getenv(name, default) |
|
|
|
|
| def _build_text_encoder(name: str, fp32: bool = False): |
| if name not in TEXT_ENCODER_PRESETS: |
| available = ", ".join(sorted(TEXT_ENCODER_PRESETS)) |
| raise ValueError(f"Unknown TEXT_ENCODER='{name}'. Available: {available}") |
| preset = TEXT_ENCODER_PRESETS[name] |
| target_cls = resolve_target(preset["target"]) |
| if fp32: |
| preset["kwargs"]["dtype"] = "float32" |
| return target_cls(**preset["kwargs"]) |
|
|
|
|
| def parse_args(): |
| parser = argparse.ArgumentParser(description="Run text encoder Gradio server.") |
| parser.add_argument( |
| "--text-encoder", |
| default=_get_env("TEXT_ENCODER", DEFAULT_TEXT_ENCODER), |
| choices=sorted(TEXT_ENCODER_PRESETS.keys()), |
| help="Text encoder preset.", |
| ) |
| parser.add_argument( |
| "--tmp-folder", |
| default=_get_env("TEXT_ENCODER_TMP_FOLDER", DEFAULT_TMP_FOLDER), |
| ) |
| parser.add_argument( |
| "--fp32", |
| action="store_true", |
| help="Uses fp32 for the text encoder rather than default bfloat16.", |
| ) |
| return parser.parse_args() |
|
|
|
|
| def main(): |
| args = parse_args() |
| server_name = _get_env("GRADIO_SERVER_NAME", DEFAULT_SERVER_NAME) |
| server_port = int(_get_env("GRADIO_SERVER_PORT", DEFAULT_SERVER_PORT)) |
| theme, css = get_gradio_theme() |
| os.makedirs(args.tmp_folder, exist_ok=True) |
| display_name = TEXT_ENCODER_PRESETS[args.text_encoder]["display_name"] |
| |
| demo_wrapper_fn = DemoWrapper(lambda: _build_text_encoder(args.text_encoder, args.fp32), args.tmp_folder) |
|
|
| with gr.Blocks(title="Text encoder", css=css, theme=theme) as demo: |
| gr.Markdown(f"# Text encoder: {display_name}") |
| gr.Markdown("## Description") |
| gr.Markdown("Get a embeddings from a text.") |
|
|
| gr.Markdown("## Inputs") |
| with gr.Row(): |
| text = gr.Textbox( |
| placeholder="Type the motion you want to generate with a sentence", |
| show_label=True, |
| label="Text prompt", |
| value=DEFAULT_TEXT, |
| type="text", |
| ) |
| with gr.Row(scale=3): |
| with gr.Column(scale=1): |
| btn = gr.Button("Encode", variant="primary") |
| with gr.Column(scale=1): |
| clear = gr.Button("Clear", variant="secondary") |
| with gr.Column(scale=3): |
| pass |
|
|
| output_title = gr.Markdown("## Outputs", visible=False) |
| output_text = gr.Markdown("", visible=False) |
| with gr.Row(scale=3): |
| with gr.Column(scale=1): |
| download = gr.DownloadButton("Download", variant="primary", visible=False) |
| with gr.Column(scale=4): |
| pass |
|
|
| filename = gr.Textbox( |
| visible=False, |
| value="embedding.npy", |
| ) |
|
|
| def clear_fn(): |
| return [ |
| gr.DownloadButton(visible=False), |
| gr.Markdown(visible=False), |
| gr.Markdown(visible=False), |
| ] |
|
|
| outputs = [download, output_title, output_text] |
|
|
| gr.on( |
| triggers=[text.submit, btn.click], |
| fn=clear_fn, |
| inputs=None, |
| outputs=outputs, |
| ).then( |
| fn=demo_wrapper_fn, |
| inputs=[text, filename], |
| outputs=outputs, |
| ) |
|
|
| def download_file(): |
| return gr.DownloadButton() |
|
|
| download.click( |
| fn=download_file, |
| inputs=None, |
| outputs=[download], |
| ) |
| clear.click(fn=clear_fn, inputs=None, outputs=outputs) |
|
|
| demo.launch(server_name=server_name, server_port=server_port) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|