diff --git a/fvcore/nn/jit_handles.py b/fvcore/nn/jit_handles.py index a498f82..1cae663 100644 --- a/fvcore/nn/jit_handles.py +++ b/fvcore/nn/jit_handles.py @@ -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 diff --git a/tests/test_flop_count.py b/tests/test_flop_count.py index a5edbb9..383d723 100644 --- a/tests/test_flop_count.py +++ b/tests/test_flop_count.py @@ -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.