[JAX] Cleanup the MLP warning for TE GEMM + TP (#2054)
* fix pspec check Signed-off-by:Phuong Nguyen <phuonguyen@nvidia.com> * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * cleaning Signed-off-by:
Phuong Nguyen <phuonguyen@nvidia.com> * add docstring Signed-off-by:
Phuong Nguyen <phuonguyen@nvidia.com> * use dict.get() Signed-off-by:
Phuong Nguyen <phuonguyen@nvidia.com> * fix lint Signed-off-by:
Phuong Nguyen <phuonguyen@nvidia.com> --------- Signed-off-by:
Phuong Nguyen <phuonguyen@nvidia.com> Co-authored-by:
pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Showing
Please register or sign in to comment