CoolFace
Apppublic

johnnyclem/tool_retriever

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py157 linesDownload Raw Back to root
1import os2import sys3import faiss4import numpy as np5import streamlit as st6import pandas as pd7from text2vec import SentenceModel8from src.jsonl_Indexer import JSONLIndexer9 10def get_cli_args():11    args = {}12    argv = sys.argv[2:] if len(sys.argv) > 2 else []13    for arg in argv:14        if '=' in arg:15            key, value = arg.split('=', 1)16            args[key.strip()] = value.strip()17    return args18 19cli_args = get_cli_args()20 21DEFAULT_CONFIG = {22    'model_path': 'BAAI/bge-base-en-v1.5',23    'dataset_path': 'tool-embedding.jsonl',24    'vector_size': 768,25    'embedding_field': 'embedding',26    'id_field': 'id'27}28 29config = DEFAULT_CONFIG.copy()30config.update(cli_args)31config['vector_size'] = int(config['vector_size'])32 33# ---------------------------34# 缓存数据集加载函数(避免每次运行时重复下载数据)35# ---------------------------36@st.cache_data37def load_tools_datasets():38    from datasets import load_dataset, concatenate_datasets39    ds1 = load_dataset("mangopy/ToolRet-Tools", "code")40    ds2 = load_dataset("mangopy/ToolRet-Tools", "customized")41    ds3 = load_dataset("mangopy/ToolRet-Tools", "web")42    ds = concatenate_datasets([ds1['tools'], ds2['tools'], ds3['tools']])43    # 重命名'id'字段为'tool'44    ds = ds.rename_columns({'id': 'tool'})45    return ds46 47ds = load_tools_datasets()48df2 = ds.to_pandas()49# 如果数据量较大,可以通过设置索引加速后续的合并操作50df2.set_index('tool', inplace=True)51 52# ---------------------------53# 缓存模型加载函数54# ---------------------------55@st.cache_resource56def get_model(model_path: str = config['model_path']):57    return SentenceModel(model_path)58 59# 缓存检索器创建函数60@st.cache_resource61def create_retriever(vector_sz: int, dataset_path: str, embedding_field: str, id_field: str, _model):62    retriever = JSONLIndexer(vector_sz=vector_sz, model=_model)63    retriever.load_jsonl(dataset_path, embedding_field=embedding_field, id_field=id_field)64    return retriever65 66# ---------------------------67# 侧边栏配置68# ---------------------------69st.sidebar.markdown("<div style='text-align: center;'><h3>📄 Model Configuration</h3></div>", unsafe_allow_html=True)70model_options = ["BAAI/bge-base-en-v1.5"]71selected_model = st.sidebar.selectbox("Select Model", model_options)72st.sidebar.write("Selected model:", selected_model)73st.sidebar.write("Embedding length: 768")74 75# 使用下拉框选中的模型(避免重复加载)76model = get_model(selected_model)77retriever = create_retriever(78    config['vector_size'],79    config['dataset_path'],80    config['embedding_field'],81    config['id_field'],82    _model=model83)84 85# ---------------------------86# 界面样式设置87# ---------------------------88st.markdown("""89    <style>90    .search-container {91        display: flex;92        justify-content: center;93        align-items: center;94        gap: 10px;95        margin-top: 20px;96    }97    .search-box input {98        width: 500px !important;99        height: 45px;100        font-size: 16px;101        border-radius: 25px;102        padding-left: 15px;103    }104    .search-btn button {105        height: 45px;106        font-size: 16px;107        border-radius: 25px;108    }109    </style>110""", unsafe_allow_html=True)111 112st.markdown("<h1 style='text-align: center;'>🔍 Tool Retrieval</h1>", unsafe_allow_html=True)113 114# ---------------------------115# 主体检索区域116# ---------------------------117col1, col2 = st.columns([4, 1])118with col1:119    query = st.text_input("", placeholder="Enter your search query...", key="search_query", label_visibility="collapsed")120with col2:121    search_clicked = st.button("🔎 Search", use_container_width=True)122 123top_k = st.slider("Top-K tools", 1, 100, 50, help="Choose the number of results to display")124 125if search_clicked and query:126    rec_ids, scores = retriever.search_return_id(query, top_k)127    # 构建检索结果 DataFrame128    df1 = pd.DataFrame({"relevance": scores, "tool": rec_ids})129    # 使用 join 加速合并(前提是 df2 已设置好索引)130    results_df = df1.join(df2, on='tool', how='left').reset_index(drop=False)131    132    st.subheader("🗂️ Retrieval results")133    134    styled_results = results_df.style.apply(135        lambda x: [136            "background-color: #F7F7F7" if i % 2 == 0 else "background-color: #FFFFFF"137            for i in range(len(x))138        ],139        axis=0,140    ).format({"relevance": "{:.4f}"})141    142    st.dataframe(143        styled_results,144        column_config={145            "relevance": st.column_config.ProgressColumn(146                "relevance",147                help="记录与查询的匹配程度",148                format="%.4f",149                min_value=0,150                max_value=float(max(scores)) if len(scores) > 0 else 1,151            ),152            "tool": st.column_config.TextColumn("tool", help="tool help text", width="medium")153        },154        hide_index=True,155        use_container_width=True,156    )157