Uh oh!
There was an error while loading. Please reload this page.
Increase number of FP8 tensors per GEMM - #22
Conversation
ptrendx
commented
Oct 28, 2022
@vasunvidia Please sign your commits (see CONTRIBUTING.rst) |
ptrendx
commented
Oct 28, 2022
/blossom-ci |
vasunvidia
commented
Oct 29, 2022
Signed the commit. Thanks. |
ptrendx
commented
Nov 28, 2022
@ksivaman Could you review this? Thanks :-)! |
Uh oh!
There was an error while loading. Please reload this page.
| CUBLASLT_MATMUL_DESC_AMAX_D_POINTER, | ||
| &D_amax, | ||
| sizeof(D_amax))); | ||
| NVTE_CHECK_CUBLAS(cublasLtMatrixLayoutCreate(&Cdesc, bias_type, m, n, ldd)); |
There was a problem hiding this comment.
What if C desc is same as D desc?
There was a problem hiding this comment.
For FP8 output, C cannot be FP8. So C type should be same as bias_type for FP8 output type and D_type for other output types.
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.
| return_output = True | ||
| out_dtype = tex.DType.kFloat32 if fp32_output else TE_DType[out_dtype] | ||
| bias_dtype = output_dtype if bias is None else TE_DType[bias.dtype] |
There was a problem hiding this comment.
SHould not require bias_type in C api.
There was a problem hiding this comment.
Check with cublas team
There was a problem hiding this comment.
@ksivaman Could you remind why bias_type should not be required?
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.
There was a problem hiding this comment.
You are leaking memory here since Cdesc was already created.
There was a problem hiding this comment.
Thanks for pointing it out. Will fix this.
There was a problem hiding this comment.
TBH I fail to see why do you have to set C descriptor - it is there for the beta=1 case, right? So it should always be of the same type as D?
There was a problem hiding this comment.
This breaks beta=1 case that we use for wgrad accumulation, no?
There was a problem hiding this comment.
Got it. Will address this.
Uh oh!
There was an error while loading. Please reload this page.
There was a problem hiding this comment.
I just copied the usage for other scale factors such as
at::Tensor A_scale_inverse_arg = A_scale_inverse.clone();
Is it unnecessary?
There was a problem hiding this comment.
Yes this shouldn't be cloned for 2 reasons:
- Unnecessary malloc,
- and more importantly, we want to populate the original amax tensor received in the arguments, which would remain unchanged here.
There was a problem hiding this comment.
Thanks. I'll make the change.
There was a problem hiding this comment.
@ksivaman That begs the question why are there the other clones for scale inverses? I believe those are unnecessary as well and create additional (albeit pretty small)D2D copies before every FP8 gemm.
There was a problem hiding this comment.
Yes, I believe all 7 clones in transformer_engine/pytorch/csrc/ts_fp8_op.cpp can be removed. The scale inverse clones also seem unused? @asfiyab-nvidia Could you comment on why these were added to begin with?
There was a problem hiding this comment.
@ksivaman These aren't necessary. I'll create a PR with a fix shortly. Thanks for pointing it out
Uh oh!
There was an error while loading. Please reload this page.
ptrendx
commented
Jan 31, 2023
/te-ci |
Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com>
Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com>
Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com>
Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com>
Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com>
Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com>
Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com>
Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com>
Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com>
Signed-off-by: Vasudevan Rengasamy <vrengasamy@nvidia.com>
vasunvidia
commented
Feb 2, 2023
/te-ci |
1 similar comment
ptrendx
commented
Feb 2, 2023
/te-ci |
Signed-off-by: Vasudevan Rengasamy <vrengasamy@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>
* add flash attention to TransformerLayer Signed-off-by: Charlene Yang <charleney@nvidia.com> * Add docs for FP8 calibration (#61) Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> Signed-off-by: Charlene Yang <charleney@nvidia.com> * Fix the integer overflow in fused softmax (#60) Signed-off-by: Przemek Tredak <ptredak@nvidia.com> Signed-off-by: Charlene Yang <charleney@nvidia.com> * prefix flash attn env var with NVTE_ Signed-off-by: Charlene Yang <charleney@nvidia.com> * Address steady memory increase and bloated checkpoints (#63) Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> Signed-off-by: Charlene Yang <charleney@nvidia.com> * fix env var logic Signed-off-by: cyanguwa <cyang.uwa@gmail.com> Signed-off-by: Charlene Yang <charleney@nvidia.com> * fix flash attn env var logic again Signed-off-by: cyanguwa <cyang.uwa@gmail.com> Signed-off-by: Charlene Yang <charleney@nvidia.com> * remove d2d copies (#64) * 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 (#22) * 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 (#65) * 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> * replace rearrange with transpose Signed-off-by: cyanguwa <cyang.uwa@gmail.com> Signed-off-by: Charlene Yang <charleney@nvidia.com> * QKV parameters unfused path fixes and optimization (#66) * 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> * Fix gradients when using AMP (#70) retain grad related attrs while casting Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> Signed-off-by: Charlene Yang <charleney@nvidia.com> * fix pylint violations fixed pyline violations such as trailing white spaces and too long lines Signed-off-by: cyanguwa <cyang.uwa@gmail.com> * fix pylint violation on line 264 with R1719 Signed-off-by: cyanguwa <cyang.uwa@gmail.com> * fix two more pylint violations Signed-off-by: cyanguwa <cyang.uwa@gmail.com> * DotProductAttention API Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Add docs for attention Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * fix assert always true Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * check for correct flash-attn version Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * address review comments Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * lint+build fixes, correct settings for default flash-attn Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * correct version Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * review comments and fixes Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * fix onnx and disable flash-attn export test Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * remove einops dependency Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * cleanup internal API; rm duplication Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * fix Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * only install TE wheel (exclude flash-attn to rm conflicts) Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * forgot to change install wheel path Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * next round review comments Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * fix flash_attn output Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * fix QK layer scaling Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * update docs Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * review comments and fixes to selective checkpointing Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> --------- Signed-off-by: Charlene Yang <charleney@nvidia.com> Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> Signed-off-by: cyanguwa <cyang.uwa@gmail.com> Co-authored-by: Charlene Yang <charleney@nvidia.com> Co-authored-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
- Remove the flag_gems.use_gems() context to avoid context-switching overhead - Call flag_gems.xxx directly wherever possible.

No description provided.