finalize tf basic op functions
This commit is contained in:
@@ -93,18 +93,21 @@ class TFLongformerSelfAttention(tf.keras.layers.Layer):
|
||||
0.0000, 0.0000, -0.7584, 0.4206, -0.0405, 0.1599, 0.0000
|
||||
0.0000, 0.0000, 0.0000, 2.0514, -1.1600, 0.5372, 0.2629 ]
|
||||
"""
|
||||
total_num_heads, num_chunks, window_overlap, hidden_dim = chunked_hidden_states.size()
|
||||
chunked_hidden_states = F.pad(
|
||||
chunked_hidden_states, (0, window_overlap + 1)
|
||||
total_num_heads, num_chunks, window_overlap, hidden_dim = shape_list(chunked_hidden_states)
|
||||
|
||||
paddings = tf.constant([[0, 0], [0, 0], [0, 0], [0, window_overlap + 1]])
|
||||
chunked_hidden_states = tf.pad(
|
||||
chunked_hidden_states, paddings
|
||||
) # total_num_heads x num_chunks x window_overlap x (hidden_dim+window_overlap+1). Padding value is not important because it'll be overwritten
|
||||
chunked_hidden_states = chunked_hidden_states.view(
|
||||
total_num_heads, num_chunks, -1
|
||||
|
||||
chunked_hidden_states = tf.reshape(
|
||||
chunked_hidden_states, (total_num_heads, num_chunks, -1)
|
||||
) # total_num_heads x num_chunks x window_overlapL+window_overlapwindow_overlap+window_overlap
|
||||
chunked_hidden_states = chunked_hidden_states[
|
||||
:, :, :-window_overlap
|
||||
] # total_num_heads x num_chunks x window_overlapL+window_overlapwindow_overlap
|
||||
chunked_hidden_states = chunked_hidden_states.view(
|
||||
total_num_heads, num_chunks, window_overlap, window_overlap + hidden_dim
|
||||
chunked_hidden_states = tf.reshape(
|
||||
chunked_hidden_states, (total_num_heads, num_chunks, window_overlap, window_overlap + hidden_dim)
|
||||
) # total_num_heads x num_chunks, window_overlap x hidden_dim+window_overlap
|
||||
chunked_hidden_states = chunked_hidden_states[:, :, :, :-1]
|
||||
return chunked_hidden_states
|
||||
@@ -112,19 +115,26 @@ class TFLongformerSelfAttention(tf.keras.layers.Layer):
|
||||
@staticmethod
|
||||
def _chunk(hidden_states, window_overlap):
|
||||
"""convert into overlapping chunkings. Chunk size = 2w, overlap size = w"""
|
||||
batch_size, seq_length, hidden_dim = shape_list(hidden_states)
|
||||
num_output_chunks = 2 * (seq_length // (2 * window_overlap)) - 1
|
||||
|
||||
# non-overlapping chunks of size = 2w
|
||||
hidden_states = hidden_states.view(
|
||||
hidden_states.size(0),
|
||||
hidden_states.size(1) // (window_overlap * 2),
|
||||
window_overlap * 2,
|
||||
hidden_states.size(2),
|
||||
# define frame size and frame stride (similar to convolution)
|
||||
frame_hop_size = window_overlap * hidden_dim
|
||||
frame_size = 2 * frame_hop_size
|
||||
|
||||
hidden_states = tf.reshape(hidden_states, (batch_size, seq_length * hidden_dim))
|
||||
|
||||
# chunk with overlap
|
||||
chunked_hidden_states = tf.signal.frame(hidden_states, frame_size, frame_hop_size)
|
||||
|
||||
assert shape_list(chunked_hidden_states) == [
|
||||
batch_size,
|
||||
num_output_chunks,
|
||||
frame_size,
|
||||
], f"Make sure chunking is correctly applied. `Chunked hidden states should have output dimension {[batch_size, frame_size, num_output_chunks]}, but got {shape_list(chunked_hidden_states)}."
|
||||
|
||||
chunked_hidden_states = tf.reshape(
|
||||
chunked_hidden_states, (batch_size, num_output_chunks, 2 * window_overlap, hidden_dim)
|
||||
)
|
||||
|
||||
# use `as_strided` to make the chunks overlap with an overlap size = window_overlap
|
||||
chunk_size = list(hidden_states.size())
|
||||
chunk_size[1] = chunk_size[1] * 2 - 1
|
||||
|
||||
chunk_stride = list(hidden_states.stride())
|
||||
chunk_stride[1] = chunk_stride[1] // 2
|
||||
return hidden_states.as_strided(size=chunk_size, stride=chunk_stride)
|
||||
return chunked_hidden_states
|
||||
|
||||
@@ -367,34 +367,27 @@ class TFLongformerModelIntegrationTest(unittest.TestCase):
|
||||
dtype=tf.float32,
|
||||
)
|
||||
|
||||
def _test_diagonalize(self):
|
||||
def test_diagonalize(self):
|
||||
hidden_states = self._get_hidden_states()
|
||||
hidden_states = hidden_states.reshape((1, 8, 4)) # set seq length = 8, hidden dim = 4
|
||||
hidden_states = tf.reshape(hidden_states, (1, 8, 4)) # set seq length = 8, hidden dim = 4
|
||||
chunked_hidden_states = TFLongformerSelfAttention._chunk(hidden_states, window_overlap=2)
|
||||
window_overlap_size = chunked_hidden_states.shape[2]
|
||||
window_overlap_size = shape_list(chunked_hidden_states)[2]
|
||||
self.assertTrue(window_overlap_size == 4)
|
||||
|
||||
padded_hidden_states = TFLongformerSelfAttention._pad_and_diagonalize(chunked_hidden_states)
|
||||
|
||||
self.assertTrue(padded_hidden_states.shape[-1] == chunked_hidden_states.shape[-1] + window_overlap_size - 1)
|
||||
self.assertTrue(
|
||||
shape_list(padded_hidden_states)[-1] == shape_list(chunked_hidden_states)[-1] + window_overlap_size - 1
|
||||
)
|
||||
|
||||
# first row => [0.4983, 2.6918, -0.0071, 1.0492, 0.0000, 0.0000, 0.0000]
|
||||
self.assertTrue(torch.allclose(padded_hidden_states[0, 0, 0, :4], chunked_hidden_states[0, 0, 0], atol=1e-3))
|
||||
self.assertTrue(
|
||||
torch.allclose(
|
||||
padded_hidden_states[0, 0, 0, 4:],
|
||||
torch.zeros((3,), device=torch_device, dtype=torch.float32),
|
||||
atol=1e-3,
|
||||
)
|
||||
)
|
||||
tf.debugging.assert_near(padded_hidden_states[0, 0, 0, :4], chunked_hidden_states[0, 0, 0], rtol=1e-3)
|
||||
tf.debugging.assert_near(padded_hidden_states[0, 0, 0, 4:], tf.zeros((3,), dtype=tf.dtypes.float32), rtol=1e-3)
|
||||
|
||||
# last row => [0.0000, 0.0000, 0.0000, 2.0514, -1.1600, 0.5372, 0.2629]
|
||||
self.assertTrue(torch.allclose(padded_hidden_states[0, 0, -1, 3:], chunked_hidden_states[0, 0, -1], atol=1e-3))
|
||||
self.assertTrue(
|
||||
torch.allclose(
|
||||
padded_hidden_states[0, 0, -1, :3],
|
||||
torch.zeros((3,), device=torch_device, dtype=torch.float32),
|
||||
atol=1e-3,
|
||||
)
|
||||
tf.debugging.assert_near(padded_hidden_states[0, 0, -1, 3:], chunked_hidden_states[0, 0, -1], rtol=1e-3)
|
||||
tf.debugging.assert_near(
|
||||
padded_hidden_states[0, 0, -1, :3], tf.zeros((3,), dtype=tf.dtypes.float32), rtol=1e-3
|
||||
)
|
||||
|
||||
def test_pad_and_transpose_last_two_dims(self):
|
||||
@@ -413,7 +406,7 @@ class TFLongformerModelIntegrationTest(unittest.TestCase):
|
||||
hidden_states[0, -1, :], tf.reshape(padded_hidden_states, (1, -1))[0, 24:32], rtol=1e-6
|
||||
)
|
||||
|
||||
def _test_chunk(self):
|
||||
def test_chunk(self):
|
||||
hidden_states = self._get_hidden_states()
|
||||
batch_size = 1
|
||||
seq_length = 8
|
||||
@@ -423,16 +416,12 @@ class TFLongformerModelIntegrationTest(unittest.TestCase):
|
||||
chunked_hidden_states = TFLongformerSelfAttention._chunk(hidden_states, window_overlap=2)
|
||||
|
||||
# expected slices across chunk and seq length dim
|
||||
expected_slice_along_seq_length = torch.tensor(
|
||||
[0.4983, -0.7584, -1.6944], device=torch_device, dtype=torch.float32
|
||||
)
|
||||
expected_slice_along_chunk = torch.tensor(
|
||||
[0.4983, -1.8348, -0.7584, 2.0514], device=torch_device, dtype=torch.float32
|
||||
)
|
||||
expected_slice_along_seq_length = tf.convert_to_tensor([0.4983, -0.7584, -1.6944], dtype=tf.dtypes.float32)
|
||||
expected_slice_along_chunk = tf.convert_to_tensor([0.4983, -1.8348, -0.7584, 2.0514], dtype=tf.dtypes.float32)
|
||||
|
||||
self.assertTrue(torch.allclose(chunked_hidden_states[0, :, 0, 0], expected_slice_along_seq_length, atol=1e-3))
|
||||
self.assertTrue(torch.allclose(chunked_hidden_states[0, 0, :, 0], expected_slice_along_chunk, atol=1e-3))
|
||||
self.assertTrue(chunked_hidden_states.shape, (1, 3, 4, 4))
|
||||
self.assertTrue(shape_list(chunked_hidden_states) == [1, 3, 4, 4])
|
||||
tf.debugging.assert_near(chunked_hidden_states[0, :, 0, 0], expected_slice_along_seq_length, rtol=1e-3)
|
||||
tf.debugging.assert_near(chunked_hidden_states[0, 0, :, 0], expected_slice_along_chunk, rtol=1e-3)
|
||||
|
||||
def _test_layer_local_attn(self):
|
||||
model = TFLongformerModel.from_pretrained("patrickvonplaten/longformer-random-tiny")
|
||||
|
||||
Reference in New Issue
Block a user