Skip to content

Fix matmul flop counting for vector operands - #159

Open
harshitgavita-07 wants to merge 2 commits into
facebookresearch:mainfrom
harshitgavita-07:fix/matmul-vector-flop-count
Open

harshitgavita-07 wants to merge 2 commits into
facebookresearch:mainfrom
harshitgavita-07:fix/matmul-vector-flop-count

Conversation

@harshitgavita-07

Copy link
Copy Markdown

Summary

Fixes #130.

matmul_flop_jit() assumed both operands are matrices. With a vector on the right (vector-vector or matrix-vector matmul), the shape assertion indexes [-2] on a 1-D shape and raises IndexError. With a vector on the left against a batched matrix, the formula prod(input_shapes[0]) * input_shapes[1][-1] counts only one batch and undercounts.

The handler now checks the shared contraction dimension (-1 for a vector right operand, -2 otherwise) and computes prod(output_shape) * contraction_dim, which matches the existing behavior for plain matrix-matrix inputs and covers the vector cases.

Tests

Added test_matmul_vectors covering vector-vector, matrix-vector, batched-matrix-vector, vector-matrix, and vector-batched-matrix. 4 of the 5 cases fail on current main (3 IndexError, 1 wrong count); all pass with the fix. Full tests/test_flop_count.py suite passes.

@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. label Sep 30, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed.

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Flop counter for matmul does not support matrix-vector product.

1 participant