BAAI/LLARA-passage
BAAI
Similitud de oraciones
En este proyecto, introducimos LLaRA: EBAE: Embedding-Based Auto-Encoding. EBAR: Embedding-Based Auto-Regression. LLaRA es un modelo transformador de oraciones diseñado para la extracción de características y para la similitud de oraciones. Utiliza diversas bibliotecas y tiene compatibilidad con AutoTrain y Endpoints de inferencia.
Como usar
import torch
from transformers import AutoModel, AutoTokenizer, LlamaModel
def get_query_inputs(queries, tokenizer, max_length=512):
prefix = '"'
suffix = '", predict the following passage within eight words: '
prefix_ids = tokenizer(prefix, return_tensors=None)['input_ids']
suffix_ids = tokenizer(suffix, return_tensors=None)['input_ids'][1:]
queries_inputs = []
for query in queries:
inputs = tokenizer(query,
return_tensors=None,
max_length=max_length,
truncation=True,
add_special_tokens=False)
inputs['input_ids'] = prefix_ids + inputs['input_ids'] + suffix_ids
inputs['attention_mask'] = [1] * len(inputs['input_ids'])
queries_inputs.append(inputs)
return tokenizer.pad(
queries_inputs,
padding=True,
max_length=max_length,
pad_to_multiple_of=8,
return_tensors='pt',
)
def get_passage_inputs(passages, tokenizer, max_length=512):
prefix = '"'
suffix = '", summarize the above passage within eight words: '
prefix_ids = tokenizer(prefix, return_tensors=None)['input_ids']
suffix_ids = tokenizer(suffix, return_tensors=None)['input_ids'][1:]
passages_inputs = []
for passage in passages:
inputs = tokenizer(passage,
return_tensors=None,
max_length=max_length,
truncation=True,
add_special_tokens=False)
inputs['input_ids'] = prefix_ids + inputs['input_ids'] + suffix_ids
inputs['attention_mask'] = [1] * len(inputs['input_ids'])
passages_inputs.append(inputs)
return tokenizer.pad(
passages_inputs,
padding=True,
max_length=max_length,
pad_to_multiple_of=8,
return_tensors='pt',
)
# Cargar el tokenizador y el modelo
tokenizer = AutoTokenizer.from_pretrained('BAAI/LLARA-passage')
model = AutoModel.from_pretrained('BAAI/LLARA-passage')
# Definir entradas de consulta y pasaje
query = "What is llama?"
title = "Llama"
passage = "The llama is a domesticated South American camelid, widely used as a meat and pack animal by Andean cultures since the pre-Columbian era."
query_input = get_query_inputs([query], tokenizer)
passage_input = get_passage_inputs([passage], tokenizer)
with torch.no_grad():
# Computar la incrustación de la consulta
query_outputs = model(**query_input, return_dict=True, output_hidden_states=True)
query_embedding = query_outputs.hidden_states[-1][:, -8:, :]
query_embedding = torch.mean(query_embedding, dim=1)
query_embedding = torch.nn.functional.normalize(query_embedding, dim=-1)
# Computar la incrustación del pasaje
passage_outputs = model(**passage_input, return_dict=True, output_hidden_states=True)
passage_embeddings = passage_outputs.hidden_states[-1][:, -8:, :]
passage_embeddings = torch.mean(passage_embeddings, dim=1)
passage_embeddings = torch.nn.functional.normalize(passage_embeddings, dim=-1)
# Computar el puntaje de similitud
score = query_embedding @ passage_embeddings.T
print(score)
Funcionalidades
- Similitud de oraciones
- Transformadores de oraciones
- Safetensors
- Llama
- Extracción de características
- Compatible con AutoTrain
- Compatible con Endpoints de inferencia
Casos de uso
- Búsqueda semántica
- Comparación de textos
- Recuperación de información densa