Skip to content

[Example] Add two Flashattention accelerator example - #615

Open
RuizeYu05 wants to merge 9 commits into
cornell-zhang:mainfrom
RuizeYu05:two_fa
Open

RuizeYu05 wants to merge 9 commits into
cornell-zhang:mainfrom
RuizeYu05:two_fa

Conversation

@RuizeYu05

@RuizeYu05 RuizeYu05 commented Sep 30, 2026 •

Copy link
Copy Markdown

Description

I merged the two PR (#564 and #581) submitted before about Flashattention examples

Problems

Add example for Allo

Proposed Solutions

Implement Flashattention accelerator in allo

Examples

It's an example

Checklist

Please make sure to review and check all of these items:

  • PR's title starts with a category (e.g. [Bugfix], [IR], [Builder], etc)
  • All changes have test coverage (It would be good to provide ~2 different test cases to test the robustness of your code)
  • Pass the formatting check locally
  • Code is well-documented

@RuizeYu05 RuizeYu05 changed the title Add two Flashattention accelerator example [Example] Add two Flashattention accelerator example Sep 30, 2026
@Fangtangtang

Copy link
Copy Markdown
Collaborator

Thanks! Does this include PRs #564 and #581? If so, I'll close them.

@RuizeYu05

RuizeYu05 commented Sep 30, 2026 via email

Copy link
Copy Markdown
Author

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot review overview

🟡 Changes recommended

Invalid shape handling, softmax initialization, and silently skipped systolic correctness testing can produce incorrect results or ineffective CI coverage.

Review effort: Balanced
Findings: 3 High severity · 1 Medium severity

Open (4)
What changed in this PR

Adds two FPGA attention accelerator examples: tiled FlashAttention and a quantized systolic MHA design.

Changes:

  • Adds both accelerator implementations and numerical/HLS test drivers.
  • Documents their algorithms and hardware organization.
  • Integrates checks into standard and weekly FPGA workflows.
File Description
examples/​attention/​README.md Documents both attention accelerators.
examples/​attention/​fused_MHA_systolic/​fused_MHA_systolic.py Implements quantized systolic MHA.
examples/​attention/​fused_MHA_systolic/​test_systolic.py Adds reference and HLS testing.
examples/​attention/​fused_MHA_systolic/​__init__.py Initializes the example package.
examples/​attention/​flashattention/​flash_Atten.py Implements tiled FlashAttention.
examples/​attention/​flashattention/​test_flash.py Adds LLVM and HLS testing.
examples/​attention/​flashattention/​__init__.py Initializes the example package.
.github/​workflows/​config.yml Adds both examples to PR CI.
.github/​workflows/​fpga_weekly.yml Adds weekly synthesis coverage.

💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

NUM_HEADS: int,
BLOCK_T: int = 4,
):
HEAD_DIM = HIDDEN_SIZE // NUM_HEADS
Comment on lines +19 to +22
HEAD_DIM = HIDDEN_SIZE // NUM_HEADS
NUM_TC = CONTEXT_LENGTH // BLOCK_T

assert NUM_TC == BLOCK_T, "This design requires NUM_TC == BLOCK_T"
Comment on lines +240 to +244
m_new: float32 = m_cur
if x > m_cur:
m_new = x
ep: float32 = allo.exp(m_cur - m_new)
ex: float32 = allo.exp(x - m_new)
Comment thread examples/attention/fused_mha_systolic/test_fused_mha_systolic.py

@Fangtangtang Fangtangtang left a comment •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could you please clean up some of the structural issues first? and check copilot's suggestions

Comment thread .github/workflows/config.yml
Comment thread examples/attention/flash_attention/__init__.py
Comment thread examples/attention/flash_attention/flash_attention.py
Comment thread examples/attention/fused_mha_systolic/__init__.py
Comment thread examples/attention/fused_mha_systolic/fused_mha_systolic.py
Comment thread examples/attention/flash_attention/test_flash_attention.py
Comment thread examples/attention/fused_mha_systolic/test_fused_mha_systolic.py
Comment thread examples/attention/fused_mha_systolic/test_fused_mha_systolic.py
@RuizeYu05

Copy link
Copy Markdown
Author

Could you please clean up some of the structural issues first? and check copilot's suggestions

All done as you wish, my lord.

@Fangtangtang Fangtangtang left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for your patience! I have a few additional suggestions.

Comment thread .github/workflows/fpga_weekly.yml Outdated
Comment thread examples/attention/flash_attention/test_flash_attention.py
def test_flashattention():
run_test_with_params(
BATCH_SIZE=4, CONTEXT_LENGTH=16, HIDDEN_SIZE=64, NUM_HEADS=4, BLOCK_T=4
)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This wrapper seems to exist only for pytest test discovery?
We can make run_test_with_params the pytest test directly and pass these arguments using @pytest.mark.parametrize. Something like

@pytest.mark.parametrize(
    "BATCH_SIZE, CONTEXT_LENGTH, HIDDEN_SIZE, NUM_HEADS, BLOCK_T",
    [
        (4, 16, 64, 4, 4),
    ],
)
def test_flashattention(
    BATCH_SIZE,
    CONTEXT_LENGTH,
    HIDDEN_SIZE,
    NUM_HEADS,
    BLOCK_T,
):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Actually I was thinking of removing the wrapper: rename run_test_with_params to test_flashattention so it can be triggered by pytest. The goal is to avoid introducing unnecessary functions.

import allo.dataflow as df

int8 = Int(8)
int32 = Int(32)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

why not import int8, int32directly from allo.ir.types

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We are testing quantization method then. It's easier for us to change data type.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

If you'd like in this way, I will modify it.

@Fangtangtang Fangtangtang Oct 8, 2026 •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I see, but I'm not sure how this makes changing data types easier. I'd prefer directly importing the predefined types

Comment on lines +84 to +87
def test_fused_MHA_systolic():
run_test_with_params(
BATCH_SIZE=4, CONTEXT_LENGTH=16, HIDDEN_SIZE=16, NUM_HEADS=4, BLOCK_T=4
)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Similar to test_flash_attention.py above.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

you may need to update file names in this README. It would be better to make them relative links to the corresponding files.
Also, shall we simplify this and make it more technical? I think a short description of each design plus a concrete result table would be clearer.

Comment thread examples/attention/flash_attention/test_flash_attention.py Outdated
def test_flashattention():
run_test_with_params(
BATCH_SIZE=4, CONTEXT_LENGTH=16, HIDDEN_SIZE=64, NUM_HEADS=4, BLOCK_T=4
)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Actually I was thinking of removing the wrapper: rename run_test_with_params to test_flashattention so it can be triggered by pytest. The goal is to avoid introducing unnecessary functions.

Comment on lines +84 to +103
@pytest.mark.parametrize(
"BATCH_SIZE, CONTEXT_LENGTH, HIDDEN_SIZE, NUM_HEADS, BLOCK_T",
[
(4, 16, 16, 4, 4),
],
)
def test_fused_MHA_systolic(
BATCH_SIZE,
CONTEXT_LENGTH,
HIDDEN_SIZE,
NUM_HEADS,
BLOCK_T,
):
run_test_with_params(
BATCH_SIZE=BATCH_SIZE,
CONTEXT_LENGTH=CONTEXT_LENGTH,
HIDDEN_SIZE=HIDDEN_SIZE,
NUM_HEADS=NUM_HEADS,
BLOCK_T=BLOCK_T,
)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Similar to test_flash_attention.py above. Please avoid introducing unnecessary functions.

Comment thread examples/attention/README.md Outdated
Comment thread examples/attention/README.md Outdated
@RuizeYu05

Copy link
Copy Markdown
Author

All changes applied now.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants