Uh oh!
There was an error while loading. Please reload this page.
Faster split of QKV for FlashAttention - #166
Conversation
Signed-off-by: Przemek Tredak <ptredak@nvidia.com>
ptrendx
commented
Apr 22, 2023
/te-ci |
ptrendx
commented
Apr 22, 2023
/te-ci |
ptrendx
commented
Apr 22, 2023
/te-ci |
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
| if (q.scalar_type() == at::ScalarType::Half) { | ||
| using dtype = at::Half; | ||
| flash_attention::prepare_kernel_bwd<dtype><<<grid, threads, 0, | ||
| at::cuda::getCurrentCUDAStream()>>>( | ||
| q.data_ptr<dtype>(), | ||
| k.data_ptr<dtype>(), | ||
| v.data_ptr<dtype>(), | ||
| qkv.data_ptr<dtype>(), | ||
| q.size(0), | ||
| q.size(1), | ||
| q.size(2), | ||
| q.size(3)); | ||
| } else { | ||
| using dtype = at::BFloat16; | ||
| flash_attention::prepare_kernel_bwd<dtype><<<grid, threads, 0, | ||
| at::cuda::getCurrentCUDAStream()>>>( | ||
| q.data_ptr<dtype>(), | ||
| k.data_ptr<dtype>(), | ||
| v.data_ptr<dtype>(), | ||
| qkv.data_ptr<dtype>(), | ||
| q.size(0), | ||
| q.size(1), | ||
| q.size(2), | ||
| q.size(3)); | ||
| } |
There was a problem hiding this comment.
https://github.com/pytorch/pytorch/blob/main/aten/src/ATen/Dispatch.h/#L284 seems like a cleaner method
There was a problem hiding this comment.
Will look into it in a follow up generalization PR (since in principle this code should work for any type, not just FP16 and BF16).
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Co-authored-by: Przemyslaw Tredak <ptredak@nvidia.com> Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Uh oh!
There was an error while loading. Please reload this page.
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
ksivaman
commented
Apr 24, 2023
/te-ci |
* Faster split of QKV for FlashAttention Signed-off-by: Przemek Tredak <ptredak@nvidia.com> * CI fixes Signed-off-by: Przemek Tredak <ptredak@nvidia.com> * Fix Signed-off-by: Przemek Tredak <ptredak@nvidia.com> * review comments Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Message with assert Co-authored-by: Przemyslaw Tredak <ptredak@nvidia.com> Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Review comments Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * review Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * fix misalignment error Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * make clarifying comment and check strides Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> --------- Signed-off-by: Przemek Tredak <ptredak@nvidia.com> Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
* Faster split of QKV for FlashAttention Signed-off-by: Przemek Tredak <ptredak@nvidia.com> * CI fixes Signed-off-by: Przemek Tredak <ptredak@nvidia.com> * Fix Signed-off-by: Przemek Tredak <ptredak@nvidia.com> * review comments Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Message with assert Co-authored-by: Przemyslaw Tredak <ptredak@nvidia.com> Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Review comments Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * review Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * fix misalignment error Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * make clarifying comment and check strides Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> --------- Signed-off-by: Przemek Tredak <ptredak@nvidia.com> Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
No description provided.