Datasets 文档
搜索索引
开始使用
教程
操作指南
概览
通用
加载过程流式传输配合 PyTorch 使用配合 TensorFlow 使用配合 NumPy 使用配合 JAX 使用配合 Pandas 使用配合 Polars 使用配合 PyArrow 使用配合 Spark 使用缓存管理云存储搜索索引CLI故障排除
音频
视觉
文本
表格
数据集仓库
概念指南
参考
加入 Hugging Face 社区
并获得增强的文档体验
开始使用
搜索索引
FAISS 和 Elasticsearch 允许在数据集内搜索示例。当你想要从数据集中检索与你的 NLP 任务相关的特定示例时,这会非常有用。例如,如果你正在处理一个开放域问答(Open Domain Question Answering)任务,你可能只想返回与回答你的问题相关的示例。
本指南将向你展示如何为你的数据集构建一个允许搜索的索引。
FAISS
FAISS 根据向量表示的相似度来检索文档。在本示例中,你将使用 DPR 模型生成向量表示。
- 从 🤗 Transformers 下载 DPR 模型
>>> from transformers import DPRContextEncoder, DPRContextEncoderTokenizer
>>> import torch
>>> torch.set_grad_enabled(False)
>>> ctx_encoder = DPRContextEncoder.from_pretrained("facebook/dpr-ctx_encoder-single-nq-base")
>>> ctx_tokenizer = DPRContextEncoderTokenizer.from_pretrained("facebook/dpr-ctx_encoder-single-nq-base")- 加载你的数据集并计算向量表示
>>> from datasets import load_dataset
>>> ds = load_dataset('community-datasets/crime_and_punish', split='train[:100]')
>>> ds_with_embeddings = ds.map(lambda example: {'embeddings': ctx_encoder(**ctx_tokenizer(example["line"], return_tensors="pt"))[0][0].numpy()})- 使用 Dataset.add_faiss_index() 创建索引
>>> ds_with_embeddings.add_faiss_index(column='embeddings')- 现在你可以使用
embeddings索引来查询你的数据集。加载 DPR 问题编码器(Question Encoder),并使用 Dataset.get_nearest_examples() 搜索问题
>>> from transformers import DPRQuestionEncoder, DPRQuestionEncoderTokenizer
>>> q_encoder = DPRQuestionEncoder.from_pretrained("facebook/dpr-question_encoder-single-nq-base")
>>> q_tokenizer = DPRQuestionEncoderTokenizer.from_pretrained("facebook/dpr-question_encoder-single-nq-base")
>>> question = "Is it serious ?"
>>> question_embedding = q_encoder(**q_tokenizer(question, return_tensors="pt"))[0][0].numpy()
>>> scores, retrieved_examples = ds_with_embeddings.get_nearest_examples('embeddings', question_embedding, k=10)
>>> retrieved_examples["line"][0]
'_that_ serious? It is not serious at all. It’s simply a fantasy to amuse\r\n'- 你可以通过 Dataset.get_index() 访问索引,并将其用于特殊操作,例如使用
range_search进行查询
>>> faiss_index = ds_with_embeddings.get_index('embeddings').faiss_index
>>> limits, distances, indices = faiss_index.range_search(x=question_embedding.reshape(1, -1), thresh=0.95)- 查询完成后,使用 Dataset.save_faiss_index() 将索引保存到磁盘
>>> ds_with_embeddings.save_faiss_index('embeddings', 'my_index.faiss')- 稍后可以使用 Dataset.load_faiss_index() 重新加载它
>>> ds = load_dataset('community-datasets/crime_and_punish', split='train[:100]')
>>> ds.load_faiss_index('embeddings', 'my_index.faiss')Elasticsearch
与 FAISS 不同,Elasticsearch 根据精确匹配来检索文档。
在你的机器上启动 Elasticsearch,如果你还没有安装,请参阅 Elasticsearch 安装指南。
- 加载你想要建立索引的数据集
>>> from datasets import load_dataset
>>> squad = load_dataset('rajpurkar/squad', split='validation')>>> squad.add_elasticsearch_index("context", host="localhost", port="9200")- 然后你可以使用 Dataset.get_nearest_examples() 查询
context索引
>>> query = "machine"
>>> scores, retrieved_examples = squad.get_nearest_examples("context", query, k=10)
>>> retrieved_examples["title"][0]
'Computational_complexity_theory'- 如果你想重复使用该索引,请在构建索引时定义
es_index_name参数
>>> from datasets import load_dataset
>>> squad = load_dataset('rajpurkar/squad', split='validation')
>>> squad.add_elasticsearch_index("context", host="localhost", port="9200", es_index_name="hf_squad_val_context")
>>> squad.get_index("context").es_index_name
hf_squad_val_context- 稍后在调用 Dataset.load_elasticsearch_index() 时使用索引名称重新加载它
>>> from datasets import load_dataset
>>> squad = load_dataset('rajpurkar/squad', split='validation')
>>> squad.load_elasticsearch_index("context", host="localhost", port="9200", es_index_name="hf_squad_val_context")
>>> query = "machine"
>>> scores, retrieved_examples = squad.get_nearest_examples("context", query, k=10)对于更高级的 Elasticsearch 用法,你可以通过自定义设置指定你自己的配置
>>> import elasticsearch as es
>>> import elasticsearch.helpers
>>> from elasticsearch import Elasticsearch
>>> es_client = Elasticsearch([{"host": "localhost", "port": "9200"}]) # default client
>>> es_config = {
... "settings": {
... "number_of_shards": 1,
... "analysis": {"analyzer": {"stop_standard": {"type": "standard", " stopwords": "_english_"}}},
... },
... "mappings": {"properties": {"text": {"type": "text", "analyzer": "standard", "similarity": "BM25"}}},
... } # default config
>>> es_index_name = "hf_squad_context" # name of the index in Elasticsearch
>>> squad.add_elasticsearch_index("context", es_client=es_client, es_config=es_config, es_index_name=es_index_name)