alonsosilva/embeddings
1
1import solara2 3import numpy as np4import pandas as pd5from sentence_transformers import SentenceTransformer6from huggingface_hub import snapshot_download7from umap import UMAP8from annoy import AnnoyIndex9from cluestar import plot_text10 11news = pd.read_csv('https://raw.githubusercontent.com/alonsosilvaallende/fake-and-real-news-titles/main/example.csv')12texts = list(news["title"].values)13texts = [str(text) for text in texts if str(text) != 'nan']14 15sentences = ["This is an example sentence", "Each sentence is converted"]16model_path = snapshot_download(17 repo_id="TaylorAI/gte-tiny", allow_patterns=["*.json", "pytorch_model.bin"]18)19 20embedder2 = SentenceTransformer(model_path)21embeddings2 = [embedder2.encode(str(texts[i])) for i in range(500)]22 23reducer = UMAP()24X2 = reducer.fit_transform(embeddings2)25 26f = len(embeddings2[0])27t = AnnoyIndex(f, 'angular')28for i, embedded_text in enumerate(embeddings2):29 t.add_item(i, embedded_text)30t.build(1000)31 32query = solara.reactive("What did Nancy Pelosi said about Obamacare?")33@solara.component34def Page():35 with solara.Column(margin=10):36 solara.Markdown("#Embeddings")37 solara.InputText("Enter some query:", query, continuous_update=True)38 if query.value != "":39 embedded_query = embedder2.encode(query.value)40 idx, distances = t.get_nns_by_vector(embedded_query, 10, include_distances=True)41 df_neighbors = pd.DataFrame()42 df_neighbors["neighbors"]=[texts[i] for i in idx]43 df_neighbors["distances"] = distances44 x = reducer.transform([embedded_query])45 color_array = ["texts" if i not in idx else "neighbors" for i in range(len(texts[:500]))]+["query"]46 solara.AltairChart(plot_text(np.vstack((X2,x)), texts[:500]+[query.value], color_array=color_array).configure_range(47 category=['#0000ff', '#ff0000', '#a0aab4']48 ))49 solara.DataFrame(df_neighbors, items_per_page=10)50 solara.Markdown("Dataset: 'Fake and real news' from [kaggle](https://www.kaggle.com/datasets/clmentbisaillon/fake-and-real-news-dataset)")51 else:52 color_array = ["texts" for _ in range(500)]53 solara.AltairChart(plot_text(X2, texts[:500], color_array=color_array))54 