Skip to content
Open
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
8 changes: 5 additions & 3 deletions fvcore/nn/jit_handles.py
Original file line number Diff line number Diff line change
Expand Up @@ -221,11 +221,13 @@ def matmul_flop_jit(inputs: List[Any], outputs: List[Any]) -> Number:
Count flops for matmul.
"""
# Inputs should be a list of length 2.
# Inputs contains the shapes of two matrices.
input_shapes = [get_shape(v) for v in inputs]
assert len(input_shapes) == 2, input_shapes
assert input_shapes[0][-1] == input_shapes[1][-2], input_shapes
flop = prod(input_shapes[0]) * input_shapes[-1][-1]
# A vector on the right has no penultimate dimension. Each output
# element still requires a dot product over the shared input dimension.
right_contract_dim = -1 if len(input_shapes[1]) == 1 else -2
assert input_shapes[0][-1] == input_shapes[1][right_contract_dim], input_shapes
flop = prod(get_shape(outputs[0])) * input_shapes[0][-1]
return flop


Expand Down
16 changes: 16 additions & 0 deletions tests/test_flop_count.py
Original file line number Diff line number Diff line change
Expand Up @@ -601,6 +601,22 @@ def test_matmul(self) -> None:
flop_dict, gt_dict, "Matmul operation failed to pass the flop count test."
)

def test_matmul_vectors(self) -> None:
"""Matmul counts vector operands, including batched matrices."""
cases = [
((10,), (10,), 10),
((20, 10), (10,), 200),
((2, 20, 10), (10,), 400),
((10,), (10, 20), 200),
((10,), (2, 10, 20), 400),
]
for left, right, expected in cases:
with self.subTest(left=left, right=right):
inputs = (torch.randn(*left), torch.randn(*right))
self.assertEqual(
FlopCountAnalysis(MatmulNet(), inputs).total(), expected
)

def test_matmul_broadcast(self) -> None:
"""
Test flop count for operation matmul.
Expand Down