From 9a41f177f9f31f77ec7d8d36960ce3e9bca99f91 Mon Sep 17 00:00:00 2001 From: harshitgavita-07 Date: Thu, 1 Oct 2026 02:11:43 +0530 Subject: [PATCH 1/2] Fix matmul flop counting for vector operands --- fvcore/nn/jit_handles.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) 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 From 2b70b93195b10daf7967856a7cd840cdc686fe73 Mon Sep 17 00:00:00 2001 From: harshitgavita-07 Date: Thu, 1 Oct 2026 02:12:10 +0530 Subject: [PATCH 2/2] Add regression tests for matmul vector flop counting --- tests/test_flop_count.py | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) 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.