HF-test-lab/bulk_embeddings
0
1import gradio as gr2 3from data import download_dataset, tokenize_dataset, load_tokenized_dataset4from infer import get_model_and_tokenizer, batch_embed5 6# TODO: add instructor models7# "hkunlp/instructor-xl",8# "hkunlp/instructor-large",9# "hkunlp/instructor-base",10 11# model ids and hidden sizes12models_and_hidden_sizes = [13 ("intfloat/e5-small-v2", 384),14 ("intfloat/e5-base-v2", 768),15 ("intfloat/e5-large-v2", 1024),16 ("intfloat/multilingual-e5-small", 384),17 ("intfloat/multilingual-e5-base", 768),18 ("intfloat/multilingual-e5-large", 1024),19 ("sentence-transformers/all-MiniLM-L6-v2", 384),20 ("sentence-transformers/all-MiniLM-L12-v2", 384),21 ("sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2", 384),22]23 24model_options = [25 f"{model_name} (hidden_size = {hidden_size})"26 for model_name, hidden_size in models_and_hidden_sizes27]28 29 30opt2desc = {31 "O2": "Most precise, slowest (O2: basic and extended general optimizations, transformers-specific fusions)",32 "O3": "Less precise, faster (O3: O2 + gelu approx)",33 "O4": "Least precise, fastest (O4: O3 + fp16/bf16)",34}35 36desc2opt = {v: k for k, v in opt2desc.items()}37 38 39optimization_options = list(opt2desc.values())40 41 42def download_and_tokenize(43 ds_name,44 ds_config,45 column_name,46 ds_split,47 model_choice,48 opt_desc,49 num2skip,50 num2embed,51 progress=gr.Progress(track_tqdm=True),52):53 num_samples = download_dataset(ds_name, ds_config, ds_split)54 55 opt_level = desc2opt[opt_desc]56 57 model_name = model_choice.split()[0]58 59 tokenize_dataset(60 ds_name=ds_name,61 ds_config=ds_config,62 model_name=model_name,63 opt_level=opt_level,64 column_name=column_name,65 num2skip=num2skip,66 num2embed=num2embed,67 )68 69 return f"Downloaded! It has {len(num_samples)} docs."70 71 72def embed(73 ds_name,74 ds_config,75 column_name,76 ds_split,77 model_choice,78 opt_desc,79 new_dataset_id,80 num2skip,81 num2embed,82 progress=gr.Progress(track_tqdm=True),83):84 ds = load_tokenized_dataset(ds_name, ds_config, ds_split)85 86 opt_level = desc2opt[opt_desc]87 88 model_name = model_choice.split()[0]89 90 if progress is not None:91 progress(0.2, "Downloading model and tokenizer...")92 model, tokenizer = get_model_and_tokenizer(model_name, opt_level, progress)93 94 doc_count, seconds_taken = batch_embed(95 ds,96 model,97 tokenizer,98 model_name=model_name,99 column_name=column_name,100 new_dataset_id=new_dataset_id,101 opt_level=opt_level,102 num2skip=num2skip,103 num2embed=num2embed,104 progress=progress,105 )106 107 return f"Embedded {doc_count} docs in {seconds_taken/60:.2f} minutes ({doc_count/seconds_taken:.1f} docs/sec)"108 109 110with gr.Blocks(title="Bulk embeddings") as demo:111 gr.Markdown(112 """113 # Bulk Embeddings114 115 116 This Space allows you to embed a large dataset easily. For instance, this can easily create vectors for Wikipedia \117 articles -- taking about __ hours and costing approximately $__. 118 This utilizes state-of-the-art open-source embedding models, \119 and optimizes them for inference using Hugging Face [optimum](https://github.com/huggingface/optimum). There are various \120 levels of optimizations that can be applied - the quality of the embeddings will degrade as the optimizations increase. 121 Currently available options: O2/O3/O4 on T4/A10 GPUs using onnx runtime. 122 Future options: 123 - OpenVino for CPU inference124 - TensorRT for GPU inference125 - Quantized models126 - Instructor models127 - Text splitting options128 - More control about which rows to embed (skip some, stop early)129 - Dynamic padding130 131 ## Steps132 1. Upload the dataset to the Hugging Face Hub.133 2. Enter dataset details into the form below.134 3. Choose a model. These are taken from the top of the [MTEB leaderboard](https://huggingface.co/spaces/mteb/leaderboard).135 4. Enter optimization level. See [here](https://huggingface.co/docs/optimum/onnxruntime/usage_guides/optimization#optimization-configuration) for details.136 5. Choose a name for the new dataset.137 6. Hit run!138 139 ### Note:140 If you have short documents, O3 will be faster than O4. If you have long documents, O4 will be faster than O3. \141 O4 requires the tokenized documents to be padded to max length.142 """143 )144 145 with gr.Row():146 ds_name = gr.Textbox(147 lines=1,148 label="Dataset to load from Hugging Face Hub",149 value="wikipedia",150 )151 ds_config = gr.Textbox(152 lines=1,153 label="Dataset config (leave blank to use default)",154 value="20220301.en",155 )156 157 column_name = gr.Textbox(lines=1, label="Enter column to embed", value="text")158 ds_split = gr.Dropdown(159 choices=["train", "validation", "test"],160 label="Dataset split",161 value="train",162 )163 # TODO: idx column164 # TODO: text splitting options165 166 with gr.Row():167 model_choice = gr.Dropdown(168 choices=model_options, label="Embedding model", value=model_options[0]169 )170 opt_desc = gr.Dropdown(171 choices=optimization_options,172 label="Optimization level",173 value=optimization_options[0],174 )175 176 with gr.Row():177 new_dataset_id = gr.Textbox(178 lines=1,179 label="New dataset name, including username",180 value="wiki-embeds",181 )182 183 num2skip = gr.Slider(184 value=0,185 minimum=0,186 maximum=100_000_000,187 step=1,188 label="Number of rows to skip",189 )190 191 num2embed = gr.Slider(192 value=30000,193 minimum=-1,194 maximum=100_000_000,195 step=1,196 label="Number of rows to embed (-1 = all)",197 )198 199 num2upload = gr.Slider(200 value=10000,201 minimum=1000,202 maximum=100000,203 step=1000,204 label="Chunk size for uploading",205 )206 207 with gr.Row():208 download_btn = gr.Button(value="Download and tokenize dataset!")209 embed_btn = gr.Button(value="Embed texts!")210 211 last = gr.Textbox(value="")212 213 download_btn.click(214 fn=download_and_tokenize,215 inputs=[216 ds_name,217 ds_config,218 column_name,219 ds_split,220 model_choice,221 opt_desc,222 num2skip,223 num2embed,224 ],225 outputs=last,226 )227 228 embed_btn.click(229 fn=embed,230 inputs=[231 ds_name,232 ds_config,233 column_name,234 ds_split,235 model_choice,236 opt_desc,237 new_dataset_id,238 num2skip,239 num2embed,240 ],241 outputs=last,242 )243 244 245if __name__ == "__main__":246 demo.queue(concurrency_count=20).launch(show_error=True, debug=True)247 