Retrieval-augmented-Generation

Pre-trained neural language models can learn a substantial amount of in-depth knowledge from data without any access to an external memory, as a parameterized implicit knowledge base.

However, the downsides are that they cannot easily expand or revise their memory, cannot straightforwardly provide insight into their predictions and may produce hallucinations.

RAG architecture

RAG architecture
Endow pretrained parametric-memory generation models with a non-parametric memory through a general-purpose fine-tuning approch: retrieval-augmented generation (RAG).

The models leverage two components:

  1. a retriever \(p_{\eta}(z|x)\) with parameters \(\eta\) that returns (top-\(K\) truncated) distributions over text passages given a query \(x\),
  2. a generator \(p_{\theta}(y_i|x, z, y_{<i})\) parameterized by \(\theta\) that auto-regressively generates output token \(y_i\) given previous tokens, the query \(x\) and retrieved passages \(z\).

RAG-Sequence Model: use the same retrieved document to generate the complete sequence.

\[p_{\text{RAG-Sequence}}(y|x)\approx \sum_{z\in \text{top-K}(p(\cdot|x))}p_\eta(z|x)p_\theta(y|x,z)=\sum_{z\in \text{top-K}(p(\cdot|x))}p_\eta(z|x)\prod_{i=1}^{|y|}p_\theta(y_i|x,z,y_{<i}).\]
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
class RAGSequence(nn.Module):
def __init__(self, retriever: DenseRetriever, generator: RAGGenerator, pad_id: int):
super().__init__()
self.retriever = retriever
self.generator = generator
self.pad_id = pad_id

def forward(
self,
query_ids: torch.Tensor,
query_mask: torch.Tensor,
answer_ids: torch.Tensor,
doc_ids: torch.Tensor,
doc_mask: torch.Tensor,
topk: int = 3,
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
retrieval = self.retriever(query_ids, query_mask, doc_ids, doc_mask, topk=topk)
doc_log_probs = F.log_softmax(retrieval["top_scores"], dim=-1)

context_ids, context_mask = build_contexts(
query_ids.detach().cpu(),
query_mask.detach().cpu(),
doc_ids.detach().cpu(),
retrieval["top_indices"].detach().cpu(),
self.pad_id,
)
context_ids = context_ids.to(query_ids.device)
context_mask = context_mask.to(query_ids.device)

decoder_input_ids = answer_ids[:, :-1]
labels = answer_ids[:, 1:]
batch_size = query_ids.shape[0]

repeated_decoder_input = decoder_input_ids.repeat_interleave(topk, dim=0)
repeated_labels = labels.repeat_interleave(topk, dim=0)

logits = self.generator(context_ids, context_mask, repeated_decoder_input)
doc_sequence_log_probs = sequence_log_prob(logits, repeated_labels, self.pad_id)
doc_sequence_log_probs = doc_sequence_log_probs.view(batch_size, topk)

marginal_log_prob = torch.logsumexp(doc_log_probs + doc_sequence_log_probs, dim=-1)
loss = -marginal_log_prob.mean()
info = {
**retrieval,
"doc_log_probs": doc_log_probs,
"doc_sequence_log_probs": doc_sequence_log_probs,
"marginal_log_prob": marginal_log_prob,
}
return loss, info

RAG-Token Model: marginalize over the retrieved documents at each generation step.

\[p_{\text{RAG-Token}}(y|x)\approx \prod_{i=1}^{|y|}\sum_{z\in \text{top-K}(p(\cdot|x))}p_\eta(z|x)p_\theta(y_i|x,z,y_{<i}).\]
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
def rag_token_nll(
doc_log_probs: torch.Tensor,
logits: torch.Tensor,
labels: torch.Tensor,
pad_id: int,
batch_size: int,
topk: int,
) -> torch.Tensor:
vocab_log_probs = F.log_softmax(logits, dim=-1)
seq_len = labels.shape[1]
vocab_log_probs = vocab_log_probs.view(batch_size, topk, seq_len, -1)
labels = labels.view(batch_size, topk, seq_len)
token_log_probs = vocab_log_probs.gather(dim=-1, index=labels.unsqueeze(-1)).squeeze(-1)
mixed_token_log_probs = torch.logsumexp(doc_log_probs[:, :, None] + token_log_probs, dim=1)
mask = labels[:, 0, :].ne(pad_id)
example_log_probs = (mixed_token_log_probs * mask.float()).sum(dim=-1)
return -example_log_probs.mean()


class RAGToken(RAGSequence):
def forward(
self,
query_ids: torch.Tensor,
query_mask: torch.Tensor,
answer_ids: torch.Tensor,
doc_ids: torch.Tensor,
doc_mask: torch.Tensor,
topk: int = 3,
) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
retrieval = self.retriever(query_ids, query_mask, doc_ids, doc_mask, topk=topk)
doc_log_probs = F.log_softmax(retrieval["top_scores"], dim=-1)

context_ids, context_mask = build_contexts(
query_ids.detach().cpu(),
query_mask.detach().cpu(),
doc_ids.detach().cpu(),
retrieval["top_indices"].detach().cpu(),
self.pad_id,
)
context_ids = context_ids.to(query_ids.device)
context_mask = context_mask.to(query_ids.device)

topk = retrieval["top_indices"].shape[1]
batch_size = query_ids.shape[0]
decoder_input_ids = answer_ids[:, :-1]
labels = answer_ids[:, 1:]
repeated_decoder_input = decoder_input_ids.repeat_interleave(topk, dim=0)
repeated_labels = labels.repeat_interleave(topk, dim=0)

logits = self.generator(context_ids, context_mask, repeated_decoder_input)
loss = rag_token_nll(doc_log_probs, logits, repeated_labels, self.pad_id, batch_size, topk)
return loss, {**retrieval, "doc_log_probs": doc_log_probs}

Retriever: Dense Passage Retriever (DPR).
\(p_\eta(z|x)\) is based on DPR, following a bi-encoder architecture:

\[p_\eta(z|x)\propto \exp(d(z)^T q(x)),\]

where \(d(z)=\text{BERT}_d(z)\) and \(q(x)=\text{BERT}_q(x)\) are the dense vector representations of the passage \(z\) and query \(x\), respectively. Calculating top-k of \(p_\eta(\cdot|x)\) is the maximum inner product search (MIPS) problem, which can be solved in sub-linear time (using approximation algorithms like clustering, HNSW or hashing ).

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
class MeanTextEncoder(nn.Module):
"""
Mean pooling text encoder. Used for displaying the DPR retriever. In practice, we can use any encoder architecture (e.g., BERT, RoBERTa, etc.) to encode the text.
"""
def __init__(self, vocab_size, d_model, pad_id):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model, padding_idx=pad_id)
self.d_model = d_model
self.pad_id = pad_id
self.proj = nn.Linear(d_model, d_model)
self.norm = nn.LayerNorm(d_model)
def forward(self, token_ids, mask):
if mask is None:
mask = token_ids.ne(self.pad_id)
embedded = self.embedding(token_ids) # embedded = [batch_size, seq_len, d_model]
mask = mask.unsqueeze(-1).float() # mask = [batch_size, seq_len, 1]
summed = (embedded * mask).sum(dim=1) # summed = [batch_size, d_model]
demon = mask.sum(dim=1).clamp_min(1.0) # demon = [batch_size, 1]
pooled = summed / demon # pooled = [batch_size, d_model]
return self.norm(torch.tanh(self.proj(pooled)))

class DenseRetriever(nn.Module):
"""
Dense retriever based on DPR.
"""
def __init__(self, vocab_size: int, d_model: int, pad_id: int, freeze_doc_encoder: bool = True):
super().__init__()
self.query_encoder = MeanTextEncoder(vocab_size, d_model, pad_id)
self.doc_encoder = MeanTextEncoder(vocab_size, d_model, pad_id)
if freeze_doc_encoder:
for param in self.doc_encoder.parameters():
param.requires_grad = False
def encode_docs(self, doc_ids: torch.Tensor, doc_mask: torch.Tensor) -> torch.Tensor:
return self.doc_encoder(doc_ids, doc_mask)
def forward(
self,
query_ids: torch.Tensor,
query_mask: torch.Tensor,
doc_ids: torch.Tensor,
doc_mask: torch.Tensor,
topk: int,
) -> Dict[str, torch.Tensor]:
query_vecs = self.query_encoder(query_ids, query_mask)
doc_vecs = self.encode_docs(doc_ids, doc_mask)
scores = query_vecs @ doc_vecs.T # use FAISS for large-scale retrieval
top_scores, top_indices = scores.topk(k=topk, dim=-1)
return {
"query_vecs": query_vecs,
"doc_vecs": doc_vecs,
"scores": scores,
"top_scores": top_scores,
"top_indices": top_indices,
}

retriever = DenseRetriever(vocab_size, d_model=64, pad_id=tokenizer.pad_token_id).to(device)
with torch.no_grad():
# inference time: disable gradient calculation for retrieval
retrieval = retriever(
question_ids.to(device),
question_mask.to(device),
docs_ids.to(device),
docs_mask.to(device),
topk=5,
)
for row, ex in enumerate(examples[:3]):
hits = retrieval["top_indices"][row].cpu().tolist()
print(ex["question"])
for rank, doc_i in enumerate(hits, start=1):
print(f" {rank}. {documents[doc_i]['title']}")

Generator: BART. The generator \(p_\theta(y_i|x, z, y_{<i})\) is based on BART, a denoising autoencoder for pretraining sequence-to-sequence models. It uses a standard Transformer-based encoder-decoder architecture with a bidirectional encoder and an autoregressive decoder.

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
# Context = [question tokens] + [document tokens]
# decoder input = previous answer tokens
# output = next-token logits

class RAGGenerator(nn.Module):
"""
RAG generator based on BART.
"""
def __init__(self, model_name: str):
super().__init__()
self.model = BartForConditionalGeneration.from_pretrained(model_name)
def forward(
self,
context_ids: torch.Tensor,
context_mask: torch.Tensor,
decoder_input_ids: torch.Tensor,
decoder_attention_mask: torch.Tensor,
) -> torch.Tensor:
return self.model(
input_ids=context_ids,
attention_mask=context_mask,
decoder_input_ids=decoder_input_ids,
decoder_attention_mask=decoder_attention_mask,
).logits
# returns the logits of the next token prediction
# log-probabilities can be obtained by applying log_softmax to the logits
def strip_pad(row: torch.Tensor, pad_id: int) -> List[int]:
"""
Strip padding tokens from a row of token IDs.
"""
return row[row.ne(pad_id)].tolist()

def pad_sequences(sequences: List[List[int]], pad_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
max_len = max(len(seq) for seq in sequences)
ids = torch.full((len(sequences), max_len), pad_id, dtype=torch.long)
mask = torch.zeros((len(sequences), max_len), dtype=torch.bool)
for row, seq in enumerate(sequences):
ids[row, : len(seq)] = torch.tensor(seq, dtype=torch.long)
mask[row, : len(seq)] = True
return ids, mask

def build_contexts(
query_ids: torch.Tensor,
query_mask: torch.Tensor,
all_doc_ids: torch.Tensor,
doc_indices: torch.Tensor,
pad_id: int,
max_len: int = 512,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Build contexts by concatenating query tokens and retrieved document tokens.
"""
contextx = []
batch_size, topk = doc_indices.size()
for b in range(batch_size):
q = strip_pad(query_ids[b], pad_id)
for k in range(topk):
doc_idx = doc_indices[b, k].item()
d = strip_pad(all_doc_ids[doc_idx], pad_id)
context = q + d
if len(context) > max_len:
context = context[:max_len]
contextx.append(context)
return pad_sequences(contextx, pad_id)

def seq_log_prob(logits: torch.Tensor, labels: torch.Tensor, pad_id: int) -> torch.Tensor:
"""
Compute the log-probabilities of the labels given the logits.
"""
log_probs = F.log_softmax(logits, dim=-1)
label_log_probs = log_probs.gather(dim=-1, index=labels.unsqueeze(-1)).squeeze(-1)
mask = labels.ne(pad_id)
return (label_log_probs * mask.float()).sum(dim=-1) / mask.float().sum(dim=-1)
Discussion

Comments

Sign in with GitHub to join the conversation.