Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 6 additions & 18 deletions tests/models/flava/test_flava.py
Original file line number Diff line number Diff line change
Expand Up @@ -216,9 +216,7 @@ def test_forward_image_text(self, image_encoder, text_encoder, flava, inputs):
image, _, text, _ = inputs
actual = flava(image, text)
expected_image = image_encoder(image)
expected_text = text_encoder(
text, return_attn_weights=True, return_hidden_states=True
)
expected_text = text_encoder(text, return_hidden_states=True)
assert actual.text_masked == TransformerOutput()
assert actual.multimodal_masked == TransformerOutput()
assert actual.multimodal == TransformerOutput()
Expand All @@ -244,12 +242,8 @@ def test_forward_masked_image_and_text(
)
expected_image = image_encoder(image)
expected_image_masked = image_encoder(image, masked_image)
expected_text = text_encoder(
text, return_attn_weights=True, return_hidden_states=True
)
expected_text_masked = text_encoder(
masked_text, return_attn_weights=True, return_hidden_states=True
)
expected_text = text_encoder(text, return_hidden_states=True)
expected_text_masked = text_encoder(masked_text, return_hidden_states=True)
assert actual.multimodal == TransformerOutput()
assert_expected(actual.text_masked, expected_text_masked)
assert_expected(
Expand Down Expand Up @@ -277,9 +271,7 @@ def test_forward_masked_text(self, text_encoder, flava, inputs):
text = torch.ones(2, 3, dtype=torch.int32)
masked_text = torch.ones(2, 3, dtype=torch.int32)
actual = flava(text=text, text_masked=masked_text)
expected_text = text_encoder(
text, return_attn_weights=True, return_hidden_states=True
)
expected_text = text_encoder(text, return_hidden_states=True)

assert actual.multimodal_masked == TransformerOutput()
assert actual.multimodal == TransformerOutput()
Expand All @@ -289,9 +281,7 @@ def test_forward_masked_text(self, text_encoder, flava, inputs):
assert_expected(actual.text, expected_text)
assert_expected(
actual.text_masked,
text_encoder(
masked_text, return_attn_weights=True, return_hidden_states=True
),
text_encoder(masked_text, return_hidden_states=True),
)
assert_expected(
actual.projected_text_embeddings, expected_text.last_hidden_state[:, 0, :]
Expand All @@ -300,9 +290,7 @@ def test_forward_masked_text(self, text_encoder, flava, inputs):
def test_forward_text(self, text_encoder, flava, inputs):
_, _, text, _ = inputs
actual = flava(text=text)
expected_text = text_encoder(
text, return_attn_weights=True, return_hidden_states=True
)
expected_text = text_encoder(text, return_hidden_states=True)
assert actual.multimodal_masked == TransformerOutput()
assert actual.multimodal == TransformerOutput()
assert actual.image == TransformerOutput()
Expand Down
29 changes: 0 additions & 29 deletions tests/models/flava/test_image_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -139,32 +139,3 @@ def test_image_encoder(self, image_encoder_components, input):
atol=1e-4,
rtol=0,
)
assert_expected(
out.attentions,
(
torch.Tensor(
[
[
[
[0.2000, 0.2000, 0.2000, 0.2000, 0.2000],
[0.1999, 0.2000, 0.2000, 0.2000, 0.2000],
[0.1999, 0.2000, 0.2000, 0.2000, 0.2000],
[0.1999, 0.2000, 0.2000, 0.2000, 0.2000],
[0.1999, 0.2000, 0.2000, 0.2000, 0.2000],
]
],
[
[
[0.2000, 0.2000, 0.2000, 0.2000, 0.2000],
[0.1999, 0.2000, 0.2000, 0.2000, 0.2000],
[0.1999, 0.2000, 0.2000, 0.2000, 0.2000],
[0.1999, 0.2000, 0.2000, 0.2000, 0.2000],
[0.1999, 0.2000, 0.2000, 0.2000, 0.2000],
]
],
]
),
),
atol=1e-4,
rtol=0,
)
5 changes: 0 additions & 5 deletions tests/models/flava/test_text_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,7 +77,6 @@ def test_text_transformer(self, text_encoder_components, input_ids):
text_encoder, _ = text_encoder_components
out = text_encoder(
input_ids,
return_attn_weights=True,
return_hidden_states=True,
)

Expand All @@ -95,16 +94,13 @@ def test_text_transformer(self, text_encoder_components, input_ids):
rtol=0.0,
)

assert_expected(out.attentions, (torch.Tensor([[[[0, 1.0], [0.0, 1.0]]]]),))

def test_text_transformer_attn_mask(
self, text_encoder_components, input_ids, attn_mask
):
text_encoder, _ = text_encoder_components
out = text_encoder(
input_ids,
attention_mask=attn_mask,
return_attn_weights=True,
return_hidden_states=True,
)

Expand All @@ -123,4 +119,3 @@ def test_text_transformer_attn_mask(
)

assert_expected(out.pooler_output, torch.Tensor([[[1.0, -1.0], [-1.0, 1.0]]]))
assert_expected(out.attentions, (torch.Tensor([[[[1.0, 0], [1.0, 0]]]]),))
66 changes: 2 additions & 64 deletions tests/models/flava/test_transformer.py
Original file line number Diff line number Diff line change
Expand Up @@ -109,11 +109,10 @@ def inputs_ln(self):
return torch.rand((2, 3, 4))

def test_forward(self, inputs, encoder):
output = encoder(inputs, return_hidden_states=True, return_attn_weights=True)
output = encoder(inputs, return_hidden_states=True)

actual_last_hidden_state = output.last_hidden_state
actual_hidden_states = torch.sum(torch.stack(output.hidden_states), dim=0)
actual_attentions = torch.sum(torch.stack(output.attentions), dim=0)

expected_last_hidden_state = torch.Tensor(
[
Expand All @@ -127,42 +126,13 @@ def test_forward(self, inputs, encoder):
[[5.1976, 1.9218], [3.8499, 2.2402], [3.1757, -0.1730]],
]
)
expected_attentions = torch.Tensor(
[
[
[
[0.8520, 0.5740, 0.5740],
[0.6232, 0.6884, 0.6884],
[0.6232, 0.6884, 0.6884],
],
[
[0.5859, 0.7071, 0.7071],
[0.6515, 0.6742, 0.6742],
[0.6515, 0.6742, 0.6742],
],
],
[
[
[0.7392, 0.5216, 0.7392],
[0.6434, 0.7132, 0.6434],
[0.7392, 0.5216, 0.7392],
],
[
[0.6207, 0.7586, 0.6207],
[0.6589, 0.6822, 0.6589],
[0.6207, 0.7586, 0.6207],
],
],
]
)

assert_expected(
actual_last_hidden_state, expected_last_hidden_state, rtol=0.0, atol=1e-4
)
assert_expected(
actual_hidden_states, expected_hidden_states, rtol=0.0, atol=1e-4
)
assert_expected(actual_attentions, expected_attentions, rtol=0.0, atol=1e-4)

# set flags to false
output = encoder(inputs)
Expand All @@ -172,13 +142,10 @@ def test_forward(self, inputs, encoder):
)

def test_forward_ln(self, inputs_ln, encoder_ln):
output = encoder_ln(
inputs_ln, return_hidden_states=True, return_attn_weights=True
)
output = encoder_ln(inputs_ln, return_hidden_states=True)

actual_last_hidden_state = output.last_hidden_state
actual_hidden_states = torch.sum(torch.stack(output.hidden_states), dim=0)
actual_attentions = torch.sum(torch.stack(output.attentions), dim=0)

expected_last_hidden_state = torch.Tensor(
[
Expand Down Expand Up @@ -208,42 +175,13 @@ def test_forward_ln(self, inputs_ln, encoder_ln):
],
]
)
expected_attentions = torch.Tensor(
[
[
[
[0.6653, 0.6376, 0.6971],
[0.7078, 0.5621, 0.7302],
[0.6506, 0.6943, 0.6551],
],
[
[0.6333, 0.7897, 0.5770],
[0.7207, 0.7019, 0.5774],
[0.7285, 0.7195, 0.5520],
],
],
[
[
[0.6919, 0.7021, 0.6060],
[0.6274, 0.7462, 0.6264],
[0.7025, 0.7090, 0.5885],
],
[
[0.5826, 0.6227, 0.7947],
[0.6855, 0.6174, 0.6971],
[0.7317, 0.6057, 0.6625],
],
],
]
)

assert_expected(
actual_last_hidden_state, expected_last_hidden_state, rtol=0.0, atol=1e-4
)
assert_expected(
actual_hidden_states, expected_hidden_states, rtol=0.0, atol=1e-4
)
assert_expected(actual_attentions, expected_attentions, rtol=0.0, atol=1e-4)

# set flags to false
output = encoder_ln(inputs_ln)
Expand Down
1 change: 0 additions & 1 deletion tests/modules/layers/test_transformer.py
Original file line number Diff line number Diff line change
Expand Up @@ -139,7 +139,6 @@ def test_forward(
):
assert_expected(state_1, state_2)

assert actual.attentions == expected_output.attentions
assert_expected(
actual.last_hidden_state,
expected_output.last_hidden_state,
Expand Down
5 changes: 1 addition & 4 deletions torchmultimodal/models/albef/multimodal_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,9 +82,7 @@ def __init__(
def _self_attention_block(
self, hidden_states: Tensor, attention_mask: Optional[Tensor] = None
) -> Tensor:
output = self.attention(
hidden_states, attention_mask=attention_mask, return_attn_weights=False
)
output = self.attention(hidden_states, attention_mask=attention_mask)
output = self.attention_dropout(output)
return output

Expand All @@ -98,7 +96,6 @@ def _cross_attention_block(
hidden_states,
encoder_hidden_states,
attention_mask=cross_attention_mask,
return_attn_weights=False,
)
output = self.cross_attention_dropout(output)
return output
Expand Down
3 changes: 0 additions & 3 deletions torchmultimodal/models/flava/image_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -217,7 +217,6 @@ def forward(
encoder_output = self.encoder(
embedding_output,
attention_mask=attention_mask,
return_attn_weights=True,
return_hidden_states=True,
)
sequence_output = encoder_output.last_hidden_state
Expand All @@ -230,7 +229,6 @@ def forward(
last_hidden_state=sequence_output,
pooler_output=pooled_output,
hidden_states=encoder_output.hidden_states,
attentions=encoder_output.attentions,
)


Expand Down Expand Up @@ -308,5 +306,4 @@ def forward(
last_hidden_state=output.last_hidden_state,
pooler_output=output.pooler_output,
hidden_states=output.hidden_states,
attentions=output.attentions,
)
1 change: 0 additions & 1 deletion torchmultimodal/models/flava/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -254,7 +254,6 @@ def encode_text(
encoded_text = self.text_encoder(
input_ids=text,
attention_mask=text_mask,
return_attn_weights=True,
return_hidden_states=True,
)
if projection:
Expand Down
Loading
Loading