Uh oh!
There was an error while loading. Please reload this page.
flash-attn integration - #62
Conversation
ksivaman
commented
Feb 4, 2023
@cyanguwa Could you sign off the commits? This guide explains how to do it. |
Signed-off-by: Charlene Yang <charleney@nvidia.com>
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> Signed-off-by: Charlene Yang <charleney@nvidia.com>
Signed-off-by: Przemek Tredak <ptredak@nvidia.com> Signed-off-by: Charlene Yang <charleney@nvidia.com>
Signed-off-by: Charlene Yang <charleney@nvidia.com>
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> Signed-off-by: Charlene Yang <charleney@nvidia.com>
Signed-off-by: cyanguwa <cyang.uwa@gmail.com> Signed-off-by: Charlene Yang <charleney@nvidia.com>
Signed-off-by: cyanguwa <cyang.uwa@gmail.com> Signed-off-by: Charlene Yang <charleney@nvidia.com>
* remove d2d copies Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * cleanup Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> --------- Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> Signed-off-by: Charlene Yang <charleney@nvidia.com>
* Increase number of FP8 tensors per GEMM Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com> * Enable FP8 output tensor for fp8_gemm Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com> * [BERT FP8] Initial TE review comments Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com> * Temporary fix for cuda graph non convergence Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com> * Address review comments-2 Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com> * Review comments-3 Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com> * Cleanup Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com> * Change for New API Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com> * Remove unnecessary clone for D_scale, D_amax Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com> * Avoid Roll for AMAX history size = 1 Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com> * Update onnx_te_gemm API Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com> * Fix Lint errors Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com> --------- Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com> Signed-off-by: Charlene Yang <charleney@nvidia.com>
* Bug fixes from PR 22 Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Add FP8 tests to ci Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * bundle unittests for ci Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> --------- Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> Signed-off-by: Charlene Yang <charleney@nvidia.com>
Signed-off-by: cyanguwa <cyang.uwa@gmail.com> Signed-off-by: Charlene Yang <charleney@nvidia.com>
* Bug fixes from PR 22 Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Add FP8 tests to ci Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Better QKV parameter fusion Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * small fix Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * keep original param for unfused case to retain externally set attrs Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * lint fix Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Fix ONNX exports Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * improve arg naming Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * No need to set data pointers Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * lint Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Assert memory loc in NoopCat Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Handle case of different memory in param and buffer Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * fix assert always true Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Reassign params memory to avoid more concats Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> --------- Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> Signed-off-by: Charlene Yang <charleney@nvidia.com>
retain grad related attrs while casting Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> Signed-off-by: Charlene Yang <charleney@nvidia.com>
Signed-off-by: cyanguwa <cyang.uwa@gmail.com>
fixed pyline violations such as trailing white spaces and too long lines Signed-off-by: cyanguwa <cyang.uwa@gmail.com>
Signed-off-by: cyanguwa <cyang.uwa@gmail.com>
Signed-off-by: cyanguwa <cyang.uwa@gmail.com>
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
ksivaman
commented
Feb 16, 2023
/te-ci |
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
ksivaman
commented
Feb 16, 2023
/te-ci |
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
ksivaman
commented
Feb 17, 2023
/te-ci |
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
| softmax_scale=1.0/self.norm_factor, causal=self.attn_causal_mask | ||
| ) | ||
| # [b, sq, np, hn] |
There was a problem hiding this comment.
The dimension comments are to indicate shape before not after
There was a problem hiding this comment.
Well, then the code is wrong, since to get that you would need to do transpose, not view, right?
There was a problem hiding this comment.
You are actually right, everything is wrong here; comment (input shape) and the view. I couldn't catch in my numerical tests since I'm running with mbs=1, sigh
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>
ksivaman
commented
Feb 17, 2023
/te-ci |
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
ksivaman
commented
Feb 22, 2023
/te-ci |
ksivaman
commented
Feb 22, 2023
ci error is unrelated ONNX failure due to tolerance |
| ), 'FlashAttention currently only supports CUDA tensors.' | ||
| assert ( | ||
| attention_mask is None | ||
| ), 'FlashAttention currently does not support attention mask.' |
There was a problem hiding this comment.
FA calculates the causal attention mask inside itself so maybe here we meant 'FlashAttention doesn't support external attention mask'?
| pip install pytest==6.2.5 onnxruntime==1.13.1 | ||
| pytest -v -s $TE_PATH/tests/*.py | ||
| pytest -v -s $TE_PATH/tests/test_transformerengine.py $TE_PATH/tests/test_fp8.py | ||
| NVTE_FLASH_ATTN=0 pytest -v -s $TE_PATH/tests/test_onnx_export.py |
There was a problem hiding this comment.
We don't test NVTE_FLASH_ATTN=1 I guess for ONNX?
| self.attention_dropout_ctx = attention_dropout_ctx | ||
| self.attention_dropout = attention_dropout | ||
| self.layer_number = layer_number | ||
| self.apply_query_key_layer_scaling = apply_query_key_layer_scaling |
There was a problem hiding this comment.
Do we need L231-232 since we don't use layer_number and apply_query_key_layer_scaling in FA? I guess you're trying to keep the interface consistent between Unfused and Flash?
There was a problem hiding this comment.
Yes the second point is correct, just for consistency
| if use_flash_attention: | ||
| if checkpoint_core_attention: | ||
| return self._checkpointed_attention_forward(query_layer, key_layer, value_layer) |
There was a problem hiding this comment.
Did we forget to pass in the attention_func here? Same for L457.
There was a problem hiding this comment.
Yes there are some bugs in the selective activation checkpointing due to the MLM integration that doesn't actually run this path when I test for it, thus I wrongly assumed it works, fixing it now. Also will add standalone tests for this separately. Good catch :)
| custom_forward, | ||
| False, | ||
| self.get_rng_state_tracker, | ||
| self.tp_group, |
There was a problem hiding this comment.
Does self have a tp_group member?
Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
ksivaman
commented
Feb 22, 2023
/te-ci |
This PR is to provide support for flash attention in TE for self-attention calculations.