|
|
|
@@ -350,41 +350,48 @@ class LSHSelfAttention(nn.Module, EfficientAttentionMixin):
|
|
|
|
|
value_vectors.shape[-1], self.attention_head_size
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# set `num_buckets` on the fly, recommended way to do it
|
|
|
|
|
if self.num_buckets is None:
|
|
|
|
|
self._set_num_buckets(sequence_length)
|
|
|
|
|
# LSH attention only makes sense if chunked attention should be performed
|
|
|
|
|
if self.chunk_length < sequence_length:
|
|
|
|
|
# set `num_buckets` on the fly, recommended way to do it
|
|
|
|
|
if self.num_buckets is None:
|
|
|
|
|
self._set_num_buckets(sequence_length)
|
|
|
|
|
|
|
|
|
|
# use cached buckets for backprop only
|
|
|
|
|
if buckets is None:
|
|
|
|
|
# hash query key vectors into buckets
|
|
|
|
|
buckets = self._hash_vectors(query_key_vectors, num_hashes)
|
|
|
|
|
# use cached buckets for backprop only
|
|
|
|
|
if buckets is None:
|
|
|
|
|
# hash query key vectors into buckets
|
|
|
|
|
buckets = self._hash_vectors(query_key_vectors, num_hashes)
|
|
|
|
|
|
|
|
|
|
assert (
|
|
|
|
|
int(buckets.shape[-1]) == num_hashes * sequence_length
|
|
|
|
|
), "last dim of buckets is {}, but should be {}".format(buckets.shape[-1], num_hashes * sequence_length)
|
|
|
|
|
|
|
|
|
|
sorted_bucket_idx, undo_sorted_bucket_idx = self._get_sorted_bucket_idx_and_undo_sorted_bucket_idx(
|
|
|
|
|
sequence_length, buckets, num_hashes
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# make sure bucket idx is not longer then sequence length
|
|
|
|
|
sorted_bucket_idx = sorted_bucket_idx % sequence_length
|
|
|
|
|
|
|
|
|
|
# cluster query key value vectors according to hashed buckets
|
|
|
|
|
query_key_vectors = self._gather_by_expansion(query_key_vectors, sorted_bucket_idx, num_hashes)
|
|
|
|
|
value_vectors = self._gather_by_expansion(value_vectors, sorted_bucket_idx, num_hashes)
|
|
|
|
|
|
|
|
|
|
query_key_vectors = self._split_seq_length_dim_to(
|
|
|
|
|
query_key_vectors, -1, self.chunk_length, self.num_attention_heads, self.attention_head_size,
|
|
|
|
|
)
|
|
|
|
|
value_vectors = self._split_seq_length_dim_to(
|
|
|
|
|
value_vectors, -1, self.chunk_length, self.num_attention_heads, self.attention_head_size,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
if self.chunk_length is None:
|
|
|
|
|
assert (
|
|
|
|
|
self.num_chunks_before == 0 and self.num_chunks_after == 0
|
|
|
|
|
), "If `config.chunk_length` is `None`, make sure `config.num_chunks_after` and `config.num_chunks_before` are set to 0."
|
|
|
|
|
int(buckets.shape[-1]) == num_hashes * sequence_length
|
|
|
|
|
), "last dim of buckets is {}, but should be {}".format(buckets.shape[-1], num_hashes * sequence_length)
|
|
|
|
|
|
|
|
|
|
sorted_bucket_idx, undo_sorted_bucket_idx = self._get_sorted_bucket_idx_and_undo_sorted_bucket_idx(
|
|
|
|
|
sequence_length, buckets, num_hashes
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# make sure bucket idx is not longer then sequence length
|
|
|
|
|
sorted_bucket_idx = sorted_bucket_idx % sequence_length
|
|
|
|
|
|
|
|
|
|
# cluster query key value vectors according to hashed buckets
|
|
|
|
|
query_key_vectors = self._gather_by_expansion(query_key_vectors, sorted_bucket_idx, num_hashes)
|
|
|
|
|
value_vectors = self._gather_by_expansion(value_vectors, sorted_bucket_idx, num_hashes)
|
|
|
|
|
|
|
|
|
|
query_key_vectors = self._split_seq_length_dim_to(
|
|
|
|
|
query_key_vectors, -1, self.chunk_length, self.num_attention_heads, self.attention_head_size,
|
|
|
|
|
)
|
|
|
|
|
value_vectors = self._split_seq_length_dim_to(
|
|
|
|
|
value_vectors, -1, self.chunk_length, self.num_attention_heads, self.attention_head_size,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
if self.chunk_length is None:
|
|
|
|
|
assert (
|
|
|
|
|
self.num_chunks_before == 0 and self.num_chunks_after == 0
|
|
|
|
|
), "If `config.chunk_length` is `None`, make sure `config.num_chunks_after` and `config.num_chunks_before` are set to 0."
|
|
|
|
|
else:
|
|
|
|
|
# get sequence length indices
|
|
|
|
|
sorted_bucket_idx = torch.arange(sequence_length, device=query_key_vectors.device).repeat(
|
|
|
|
|
batch_size, self.num_attention_heads, 1
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# scale key vectors
|
|
|
|
|
key_vectors = self._len_and_dim_norm(query_key_vectors)
|
|
|
|
@@ -397,31 +404,35 @@ class LSHSelfAttention(nn.Module, EfficientAttentionMixin):
|
|
|
|
|
sorted_bucket_idx=sorted_bucket_idx,
|
|
|
|
|
attention_mask=attention_mask,
|
|
|
|
|
head_mask=head_mask,
|
|
|
|
|
sequence_length=sequence_length,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# free memory
|
|
|
|
|
del query_key_vectors, key_vectors, value_vectors
|
|
|
|
|
|
|
|
|
|
# sort clusters back to correct ordering
|
|
|
|
|
out_vectors, logits = ReverseSort.apply(
|
|
|
|
|
out_vectors, logits, sorted_bucket_idx, undo_sorted_bucket_idx, self.num_hashes
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# sum up all hash rounds
|
|
|
|
|
if num_hashes > 1:
|
|
|
|
|
out_vectors = self._split_seq_length_dim_to(
|
|
|
|
|
out_vectors, num_hashes, sequence_length, self.num_attention_heads, self.attention_head_size,
|
|
|
|
|
# re-order out_vectors and logits
|
|
|
|
|
if self.chunk_length < sequence_length:
|
|
|
|
|
# sort clusters back to correct ordering
|
|
|
|
|
out_vectors, logits = ReverseSort.apply(
|
|
|
|
|
out_vectors, logits, sorted_bucket_idx, undo_sorted_bucket_idx, self.num_hashes
|
|
|
|
|
)
|
|
|
|
|
logits = self._split_seq_length_dim_to(
|
|
|
|
|
logits, num_hashes, sequence_length, self.num_attention_heads, self.attention_head_size,
|
|
|
|
|
).unsqueeze(-1)
|
|
|
|
|
|
|
|
|
|
probs_vectors = torch.exp(logits - torch.logsumexp(logits, dim=2, keepdim=True))
|
|
|
|
|
out_vectors = torch.sum(out_vectors * probs_vectors, dim=2)
|
|
|
|
|
# sum up all hash rounds
|
|
|
|
|
if num_hashes > 1:
|
|
|
|
|
out_vectors = self._split_seq_length_dim_to(
|
|
|
|
|
out_vectors, num_hashes, sequence_length, self.num_attention_heads, self.attention_head_size,
|
|
|
|
|
)
|
|
|
|
|
logits = self._split_seq_length_dim_to(
|
|
|
|
|
logits, num_hashes, sequence_length, self.num_attention_heads, self.attention_head_size,
|
|
|
|
|
).unsqueeze(-1)
|
|
|
|
|
|
|
|
|
|
probs_vectors = torch.exp(logits - torch.logsumexp(logits, dim=2, keepdim=True))
|
|
|
|
|
out_vectors = torch.sum(out_vectors * probs_vectors, dim=2)
|
|
|
|
|
# free memory
|
|
|
|
|
del probs_vectors
|
|
|
|
|
|
|
|
|
|
# free memory
|
|
|
|
|
del probs_vectors
|
|
|
|
|
|
|
|
|
|
# free memory
|
|
|
|
|
del logits
|
|
|
|
|
del logits
|
|
|
|
|
|
|
|
|
|
assert out_vectors.shape == (
|
|
|
|
|
batch_size,
|
|
|
|
@@ -553,10 +564,13 @@ class LSHSelfAttention(nn.Module, EfficientAttentionMixin):
|
|
|
|
|
self.num_buckets = num_buckets
|
|
|
|
|
|
|
|
|
|
def _attend(
|
|
|
|
|
self, query_vectors, key_vectors, value_vectors, sorted_bucket_idx, attention_mask, head_mask,
|
|
|
|
|
self, query_vectors, key_vectors, value_vectors, sorted_bucket_idx, attention_mask, head_mask, sequence_length
|
|
|
|
|
):
|
|
|
|
|
key_vectors = self._look_adjacent(key_vectors, self.num_chunks_before, self.num_chunks_after)
|
|
|
|
|
value_vectors = self._look_adjacent(value_vectors, self.num_chunks_before, self.num_chunks_after)
|
|
|
|
|
|
|
|
|
|
# look at previous and following chunks if chunked attention
|
|
|
|
|
if self.chunk_length < sequence_length:
|
|
|
|
|
key_vectors = self._look_adjacent(key_vectors, self.num_chunks_before, self.num_chunks_after)
|
|
|
|
|
value_vectors = self._look_adjacent(value_vectors, self.num_chunks_before, self.num_chunks_after)
|
|
|
|
|
|
|
|
|
|
# get logits and dots
|
|
|
|
|
query_key_dots = torch.matmul(query_vectors, key_vectors.transpose(-1, -2))
|
|
|
|
@@ -564,10 +578,14 @@ class LSHSelfAttention(nn.Module, EfficientAttentionMixin):
|
|
|
|
|
# free memory
|
|
|
|
|
del query_vectors, key_vectors
|
|
|
|
|
|
|
|
|
|
query_bucket_idx = self._split_seq_length_dim_to(
|
|
|
|
|
sorted_bucket_idx, -1, self.chunk_length, self.num_attention_heads
|
|
|
|
|
)
|
|
|
|
|
key_value_bucket_idx = self._look_adjacent(query_bucket_idx, self.num_chunks_before, self.num_chunks_after)
|
|
|
|
|
# if chunked attention split bucket idxs to query and key
|
|
|
|
|
if self.chunk_length < sequence_length:
|
|
|
|
|
query_bucket_idx = self._split_seq_length_dim_to(
|
|
|
|
|
sorted_bucket_idx, -1, self.chunk_length, self.num_attention_heads
|
|
|
|
|
)
|
|
|
|
|
key_value_bucket_idx = self._look_adjacent(query_bucket_idx, self.num_chunks_before, self.num_chunks_after)
|
|
|
|
|
else:
|
|
|
|
|
query_bucket_idx = key_value_bucket_idx = sorted_bucket_idx
|
|
|
|
|
|
|
|
|
|
# get correct mask values depending on precision
|
|
|
|
|
if query_key_dots.dtype == torch.float16:
|
|
|
|
@@ -577,7 +595,7 @@ class LSHSelfAttention(nn.Module, EfficientAttentionMixin):
|
|
|
|
|
self_mask_value = self.self_mask_value_float32
|
|
|
|
|
mask_value = self.mask_value_float32
|
|
|
|
|
|
|
|
|
|
mask = self._compute_attn_mask(query_bucket_idx, key_value_bucket_idx, attention_mask)
|
|
|
|
|
mask = self._compute_attn_mask(query_bucket_idx, key_value_bucket_idx, attention_mask, sequence_length)
|
|
|
|
|
|
|
|
|
|
if mask is not None:
|
|
|
|
|
query_key_dots = torch.where(mask, query_key_dots, mask_value)
|
|
|
|
@@ -626,12 +644,13 @@ class LSHSelfAttention(nn.Module, EfficientAttentionMixin):
|
|
|
|
|
del value_vectors
|
|
|
|
|
|
|
|
|
|
# merge chunk length
|
|
|
|
|
logits = logits.flatten(start_dim=2, end_dim=3).squeeze(-1)
|
|
|
|
|
out_vectors = out_vectors.flatten(start_dim=2, end_dim=3)
|
|
|
|
|
if self.chunk_length < sequence_length:
|
|
|
|
|
logits = logits.flatten(start_dim=2, end_dim=3).squeeze(-1)
|
|
|
|
|
out_vectors = out_vectors.flatten(start_dim=2, end_dim=3)
|
|
|
|
|
|
|
|
|
|
return out_vectors, logits, attention_probs
|
|
|
|
|
|
|
|
|
|
def _compute_attn_mask(self, query_indices, key_indices, attention_mask):
|
|
|
|
|
def _compute_attn_mask(self, query_indices, key_indices, attention_mask, sequence_length):
|
|
|
|
|
mask = None
|
|
|
|
|
|
|
|
|
|
# Causal mask
|
|
|
|
@@ -641,15 +660,27 @@ class LSHSelfAttention(nn.Module, EfficientAttentionMixin):
|
|
|
|
|
# Attention mask: chunk, look up correct mask value from key_value_bucket_idx
|
|
|
|
|
# IMPORTANT: official trax code does not use a mask for LSH Atttention. Not sure why.
|
|
|
|
|
if attention_mask is not None:
|
|
|
|
|
attention_mask = attention_mask.to(torch.uint8)[:, None, None, :]
|
|
|
|
|
# expand attn_mask to fit with key_value_bucket_idx shape
|
|
|
|
|
attention_mask = attention_mask.expand(query_indices.shape[:-1] + (-1,))
|
|
|
|
|
key_attn_mask = torch.gather(attention_mask, -1, key_indices)
|
|
|
|
|
query_attn_mask = torch.gather(attention_mask, -1, query_indices)
|
|
|
|
|
# expand to query_key_dots shape: duplicate along query axis since key sorting is the same for each query position in chunk
|
|
|
|
|
attn_mask = query_attn_mask.unsqueeze(-1) * key_attn_mask.unsqueeze(-2)
|
|
|
|
|
# if chunked attention, the attention mask has to correspond to LSH order
|
|
|
|
|
if sequence_length > self.chunk_length:
|
|
|
|
|
attention_mask = attention_mask.to(torch.uint8)[:, None, None, :]
|
|
|
|
|
# expand attn_mask to fit with key_value_bucket_idx shape
|
|
|
|
|
attention_mask = attention_mask.expand(query_indices.shape[:-1] + (-1,))
|
|
|
|
|
key_attn_mask = torch.gather(attention_mask, -1, key_indices)
|
|
|
|
|
query_attn_mask = torch.gather(attention_mask, -1, query_indices)
|
|
|
|
|
# expand to query_key_dots shape: duplicate along query axis since key sorting is the same for each query position in chunk
|
|
|
|
|
attn_mask = query_attn_mask.unsqueeze(-1) * key_attn_mask.unsqueeze(-2)
|
|
|
|
|
|
|
|
|
|
# free memory
|
|
|
|
|
del query_attn_mask, key_attn_mask
|
|
|
|
|
else:
|
|
|
|
|
# usual attention mask creation
|
|
|
|
|
attention_mask = attention_mask.to(torch.uint8)[:, None, :]
|
|
|
|
|
attn_mask = (attention_mask.unsqueeze(-1) * attention_mask.unsqueeze(-2)).expand(
|
|
|
|
|
query_indices.shape + attention_mask.shape[-1:]
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# free memory
|
|
|
|
|
del query_attn_mask, key_attn_mask, attention_mask
|
|
|
|
|
del attention_mask
|
|
|
|
|
|
|
|
|
|
# multiply by casaul mask if necessary
|
|
|
|
|
if mask is not None:
|
|
|
|
@@ -809,36 +840,45 @@ class LocalSelfAttention(nn.Module, EfficientAttentionMixin):
|
|
|
|
|
torch.tensor(self.attention_head_size, device=key_vectors.device, dtype=key_vectors.dtype)
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# chunk vectors
|
|
|
|
|
# B x Num_Attn_Head x Seq_Len // chunk_len x chunk_len x attn_head_size
|
|
|
|
|
query_vectors = self._split_seq_length_dim_to(
|
|
|
|
|
query_vectors, -1, self.chunk_length, self.num_attention_heads, self.attention_head_size,
|
|
|
|
|
)
|
|
|
|
|
key_vectors = self._split_seq_length_dim_to(
|
|
|
|
|
key_vectors, -1, self.chunk_length, self.num_attention_heads, self.attention_head_size,
|
|
|
|
|
)
|
|
|
|
|
value_vectors = self._split_seq_length_dim_to(
|
|
|
|
|
value_vectors, -1, self.chunk_length, self.num_attention_heads, self.attention_head_size,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# chunk indices
|
|
|
|
|
# get sequence length indices
|
|
|
|
|
indices = torch.arange(sequence_length, device=query_vectors.device).repeat(
|
|
|
|
|
batch_size, self.num_attention_heads, 1
|
|
|
|
|
)
|
|
|
|
|
query_indices = self._split_seq_length_dim_to(indices, -1, self.chunk_length, self.num_attention_heads)
|
|
|
|
|
key_indices = self._split_seq_length_dim_to(indices, -1, self.chunk_length, self.num_attention_heads)
|
|
|
|
|
|
|
|
|
|
# append chunks before and after
|
|
|
|
|
key_vectors = self._look_adjacent(key_vectors, self.num_chunks_before, self.num_chunks_after)
|
|
|
|
|
value_vectors = self._look_adjacent(value_vectors, self.num_chunks_before, self.num_chunks_after)
|
|
|
|
|
key_indices = self._look_adjacent(key_indices, self.num_chunks_before, self.num_chunks_after)
|
|
|
|
|
# if input should be chunked
|
|
|
|
|
if self.chunk_length < sequence_length:
|
|
|
|
|
# chunk vectors
|
|
|
|
|
# B x Num_Attn_Head x Seq_Len // chunk_len x chunk_len x attn_head_size
|
|
|
|
|
query_vectors = self._split_seq_length_dim_to(
|
|
|
|
|
query_vectors, -1, self.chunk_length, self.num_attention_heads, self.attention_head_size,
|
|
|
|
|
)
|
|
|
|
|
key_vectors = self._split_seq_length_dim_to(
|
|
|
|
|
key_vectors, -1, self.chunk_length, self.num_attention_heads, self.attention_head_size,
|
|
|
|
|
)
|
|
|
|
|
value_vectors = self._split_seq_length_dim_to(
|
|
|
|
|
value_vectors, -1, self.chunk_length, self.num_attention_heads, self.attention_head_size,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
# chunk indices
|
|
|
|
|
query_indices = self._split_seq_length_dim_to(indices, -1, self.chunk_length, self.num_attention_heads)
|
|
|
|
|
key_indices = self._split_seq_length_dim_to(indices, -1, self.chunk_length, self.num_attention_heads)
|
|
|
|
|
|
|
|
|
|
# append chunks before and after
|
|
|
|
|
key_vectors = self._look_adjacent(key_vectors, self.num_chunks_before, self.num_chunks_after)
|
|
|
|
|
value_vectors = self._look_adjacent(value_vectors, self.num_chunks_before, self.num_chunks_after)
|
|
|
|
|
key_indices = self._look_adjacent(key_indices, self.num_chunks_before, self.num_chunks_after)
|
|
|
|
|
else:
|
|
|
|
|
query_indices = key_indices = indices
|
|
|
|
|
|
|
|
|
|
# query-key matmul: QK^T
|
|
|
|
|
query_key_dots = torch.matmul(query_vectors, key_vectors.transpose(-1, -2))
|
|
|
|
|
|
|
|
|
|
# free memory
|
|
|
|
|
del query_vectors, key_vectors
|
|
|
|
|
|
|
|
|
|
mask = self._compute_attn_mask(query_indices, key_indices, attention_mask, query_key_dots.shape)
|
|
|
|
|
mask = self._compute_attn_mask(
|
|
|
|
|
query_indices, key_indices, attention_mask, query_key_dots.shape, sequence_length
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
if mask is not None:
|
|
|
|
|
# get mask tensor depending on half precision or not
|
|
|
|
@@ -873,7 +913,8 @@ class LocalSelfAttention(nn.Module, EfficientAttentionMixin):
|
|
|
|
|
del value_vectors
|
|
|
|
|
|
|
|
|
|
# merge chunk length
|
|
|
|
|
out_vectors = out_vectors.flatten(start_dim=2, end_dim=3)
|
|
|
|
|
if self.chunk_length < sequence_length:
|
|
|
|
|
out_vectors = out_vectors.flatten(start_dim=2, end_dim=3)
|
|
|
|
|
|
|
|
|
|
assert out_vectors.shape == (batch_size, self.num_attention_heads, sequence_length, self.attention_head_size,)
|
|
|
|
|
|
|
|
|
@@ -884,14 +925,18 @@ class LocalSelfAttention(nn.Module, EfficientAttentionMixin):
|
|
|
|
|
|
|
|
|
|
return LocalSelfAttentionOutput(hidden_states=out_vectors, attention_probs=attention_probs)
|
|
|
|
|
|
|
|
|
|
def _compute_attn_mask(self, query_indices, key_indices, attention_mask, query_key_dots_shape):
|
|
|
|
|
def _compute_attn_mask(self, query_indices, key_indices, attention_mask, query_key_dots_shape, sequence_length):
|
|
|
|
|
mask = None
|
|
|
|
|
|
|
|
|
|
# chunk attention mask and look before and after
|
|
|
|
|
if attention_mask is not None:
|
|
|
|
|
attention_mask = attention_mask.to(torch.uint8)[:, None, :]
|
|
|
|
|
attention_mask = self._split_seq_length_dim_to(attention_mask, -1, self.chunk_length, 1)
|
|
|
|
|
attention_mask_key = self._look_adjacent(attention_mask, self.num_chunks_before, self.num_chunks_after)
|
|
|
|
|
|
|
|
|
|
if self.chunk_length < sequence_length:
|
|
|
|
|
attention_mask = self._split_seq_length_dim_to(attention_mask, -1, self.chunk_length, 1)
|
|
|
|
|
attention_mask_key = self._look_adjacent(attention_mask, self.num_chunks_before, self.num_chunks_after)
|
|
|
|
|
else:
|
|
|
|
|
attention_mask_key = attention_mask
|
|
|
|
|
|
|
|
|
|
# Causal mask
|
|
|
|
|
if self.is_decoder is True:
|
|
|
|
@@ -1564,7 +1609,9 @@ class ReformerModel(ReformerPreTrainedModel):
|
|
|
|
|
|
|
|
|
|
# if needs padding
|
|
|
|
|
least_common_mult_chunk_length = _get_least_common_mult_chunk_len(self.config)
|
|
|
|
|
must_pad_to_match_chunk_length = input_shape[-1] % least_common_mult_chunk_length != 0
|
|
|
|
|
must_pad_to_match_chunk_length = (
|
|
|
|
|
input_shape[-1] % least_common_mult_chunk_length != 0 and input_shape[-1] > least_common_mult_chunk_length
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
if must_pad_to_match_chunk_length:
|
|
|
|
|
padding_length = least_common_mult_chunk_length - input_shape[-1] % least_common_mult_chunk_length
|
|
|
|
|