Skip to content

Add device-side workspace prep kernel#4

Closed
Micky774 wants to merge 3 commits into
ipanfilo/te_gfx1250from
zain/ck/device-prep
Closed

Add device-side workspace prep kernel#4
Micky774 wants to merge 3 commits into
ipanfilo/te_gfx1250from
zain/ck/device-prep

Conversation

@Micky774

Copy link
Copy Markdown
Collaborator

Motivation

Removes host-side memory management requirements by computing metadata directly on device.

Technical Details

Test Plan

Test Result

Submission Checklist

@Micky774
Micky774 requested a review from wangye805 as a code owner July 10, 2026 17:58

@wangye805 wangye805 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.

Overall looks good. Two issues:
1). With host pinned memory moved to device, we will need to increase the device workspace size for the previous host pinned memory, right? But I didn't see where are they
2). How about the performance of this host->device change? Can we test several configs, maybe starting from ck example c++ api first.

If the performance and correctness looks good, we can file PR to CK first. Then we will also need to validate this change in tridao's flash-attn

@Micky774

Micky774 commented Jul 10, 2026

Copy link
Copy Markdown
Collaborator Author

1). With host pinned memory moved to device, we will need to increase the device workspace size for the previous host pinned memory, right? But I didn't see where are they

The (very small) host-side memory was simply for metadata for the kernel launch since it needed to be graph safe. With the on-device kernel the memory cost is eliminated entirely.

2). How about the performance of this host->device change? Can we test several configs, maybe starting from ck example c++ api first.

Sure, I'll post a follow-up with comparisons.

@Micky774

Copy link
Copy Markdown
Collaborator Author
Host-side kernel
=============
  causal(batch)    B=2 S=  512 H=16 D=128: mean=   0.256ms  median=   0.257ms  min=   0.244ms
  padcausal(group) B=2 S=  512 H=16 D=128: mean=   0.241ms  median=   0.240ms  min=   0.224ms
  causal(batch)    B=2 S= 1024 H=16 D=128: mean=   0.303ms  median=   0.302ms  min=   0.297ms
  padcausal(group) B=2 S= 1024 H=16 D=128: mean=   0.280ms  median=   0.280ms  min=   0.270ms
  causal(batch)    B=4 S= 2048 H=16 D=128: mean=   0.820ms  median=   0.818ms  min=   0.800ms
  padcausal(group) B=4 S= 2048 H=16 D=128: mean=   0.694ms  median=   0.694ms  min=   0.683ms
  causal(batch)    B=1 S= 4096 H=32 D=128: mean=   1.346ms  median=   1.346ms  min=   1.334ms
  padcausal(group) B=1 S= 4096 H=32 D=128: mean=   0.994ms  median=   0.995ms  min=   0.978ms
  causal(batch)    B=2 S= 8192 H= 8 D=128: mean=   2.394ms  median=   2.391ms  min=   2.361ms
  padcausal(group) B=2 S= 8192 H= 8 D=128: mean=   1.672ms  median=   1.671ms  min=   1.649ms



Device-side kernel
============
  causal(batch)    B=2 S=  512 H=16 D=128: mean=   0.217ms  median=   0.217ms  min=   0.212ms
  padcausal(group) B=2 S=  512 H=16 D=128: mean=   0.206ms  median=   0.206ms  min=   0.197ms
  causal(batch)    B=2 S= 1024 H=16 D=128: mean=   0.288ms  median=   0.287ms  min=   0.282ms
  padcausal(group) B=2 S= 1024 H=16 D=128: mean=   0.257ms  median=   0.258ms  min=   0.249ms
  causal(batch)    B=4 S= 2048 H=16 D=128: mean=   0.788ms  median=   0.787ms  min=   0.778ms
  padcausal(group) B=4 S= 2048 H=16 D=128: mean=   0.670ms  median=   0.670ms  min=   0.661ms
  causal(batch)    B=1 S= 4096 H=32 D=128: mean=   1.325ms  median=   1.325ms  min=   1.314ms
  padcausal(group) B=1 S= 4096 H=32 D=128: mean=   0.970ms  median=   0.971ms  min=   0.955ms
  causal(batch)    B=2 S= 8192 H= 8 D=128: mean=   2.376ms  median=   2.375ms  min=   2.361ms
  padcausal(group) B=2 S= 8192 H= 8 D=128: mean=   1.651ms  median=   1.649ms  min=   1.631ms

The runtimes are comparable with no indication of a meaningful regression.

@wangye805

Copy link
Copy Markdown
Collaborator
Host-side kernel
=============
  causal(batch)    B=2 S=  512 H=16 D=128: mean=   0.256ms  median=   0.257ms  min=   0.244ms
  padcausal(group) B=2 S=  512 H=16 D=128: mean=   0.241ms  median=   0.240ms  min=   0.224ms
  causal(batch)    B=2 S= 1024 H=16 D=128: mean=   0.303ms  median=   0.302ms  min=   0.297ms
  padcausal(group) B=2 S= 1024 H=16 D=128: mean=   0.280ms  median=   0.280ms  min=   0.270ms
  causal(batch)    B=4 S= 2048 H=16 D=128: mean=   0.820ms  median=   0.818ms  min=   0.800ms
  padcausal(group) B=4 S= 2048 H=16 D=128: mean=   0.694ms  median=   0.694ms  min=   0.683ms
  causal(batch)    B=1 S= 4096 H=32 D=128: mean=   1.346ms  median=   1.346ms  min=   1.334ms
  padcausal(group) B=1 S= 4096 H=32 D=128: mean=   0.994ms  median=   0.995ms  min=   0.978ms
  causal(batch)    B=2 S= 8192 H= 8 D=128: mean=   2.394ms  median=   2.391ms  min=   2.361ms
  padcausal(group) B=2 S= 8192 H= 8 D=128: mean=   1.672ms  median=   1.671ms  min=   1.649ms



Device-side kernel
============
  causal(batch)    B=2 S=  512 H=16 D=128: mean=   0.217ms  median=   0.217ms  min=   0.212ms
  padcausal(group) B=2 S=  512 H=16 D=128: mean=   0.206ms  median=   0.206ms  min=   0.197ms
  causal(batch)    B=2 S= 1024 H=16 D=128: mean=   0.288ms  median=   0.287ms  min=   0.282ms
  padcausal(group) B=2 S= 1024 H=16 D=128: mean=   0.257ms  median=   0.258ms  min=   0.249ms
  causal(batch)    B=4 S= 2048 H=16 D=128: mean=   0.788ms  median=   0.787ms  min=   0.778ms
  padcausal(group) B=4 S= 2048 H=16 D=128: mean=   0.670ms  median=   0.670ms  min=   0.661ms
  causal(batch)    B=1 S= 4096 H=32 D=128: mean=   1.325ms  median=   1.325ms  min=   1.314ms
  padcausal(group) B=1 S= 4096 H=32 D=128: mean=   0.970ms  median=   0.971ms  min=   0.955ms
  causal(batch)    B=2 S= 8192 H= 8 D=128: mean=   2.376ms  median=   2.375ms  min=   2.361ms
  padcausal(group) B=2 S= 8192 H= 8 D=128: mean=   1.651ms  median=   1.649ms  min=   1.631ms

The runtimes are comparable with no indication of a meaningful regression.

Actually, the runtime looks consistently better. Try to have hd64 comparison as well. If the device side is always better, we should suggest to CK to remove the host side kernel

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.

2 participants