johnnyclem/tool_retriever
0
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 