Skip to content

Commit 12af02f

Browse files
authored
Fix NVTE_FRAMEWORK=all installation (#1850)
* Fix NVTE_FRAMEWORK=all Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * fix Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Fix Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Workflow tests and fixes Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Fix jax install Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Fix Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Update dep Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Add numpy Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> * Add dep Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com> --------- Signed-off-by: Kirthi Shankar Sivamani <ksivamani@nvidia.com>
1 parent 97e493f commit 12af02f

29 files changed

Lines changed: 50 additions & 31 deletions

‎.github/workflows/build.yml‎

Lines changed: 22 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -43,7 +43,7 @@ jobs:
4343
run: |
4444
apt-get update
4545
apt-get install -y git python3.9 pip ninja-build cudnn9-cuda-12
46-
pip install cmake torch pydantic importlib-metadata>=1.0 packaging pybind11
46+
pip install cmake torch pydantic importlib-metadata>=1.0 packaging pybind11 numpy einops
4747
- name: 'Checkout'
4848
uses: actions/checkout@v3
4949
with:
@@ -54,7 +54,6 @@ jobs:
5454
NVTE_FRAMEWORK: pytorch
5555
MAX_JOBS: 1
5656
- name: 'Sanity check'
57-
if: false # Sanity import test requires Flash Attention
5857
run: python3 tests/pytorch/test_sanity_import.py
5958
jax:
6059
name: 'JAX'
@@ -73,4 +72,24 @@ jobs:
7372
NVTE_FRAMEWORK: jax
7473
MAX_JOBS: 1
7574
- name: 'Sanity check'
76-
run: python tests/jax/test_sanity_import.py
75+
run: python3 tests/jax/test_sanity_import.py
76+
all:
77+
name: 'All'
78+
runs-on: ubuntu-latest
79+
container:
80+
image: ghcr.io/nvidia/jax:jax
81+
options: --user root
82+
steps:
83+
- name: 'Dependencies'
84+
run: pip install torch pybind11[global] einops
85+
- name: 'Checkout'
86+
uses: actions/checkout@v3
87+
with:
88+
submodules: recursive
89+
- name: 'Build'
90+
run: pip install --no-build-isolation . -v
91+
env:
92+
NVTE_FRAMEWORK: all
93+
MAX_JOBS: 1
94+
- name: 'Sanity check'
95+
run: python3 tests/pytorch/test_sanity_import.py && python3 tests/jax/test_sanity_import.py

‎transformer_engine/jax/csrc/extensions/activation.cpp‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77

88
#include <cuda_runtime.h>
99

10-
#include "extensions.h"
10+
#include "../extensions.h"
1111
#include "transformer_engine/cast.h"
1212
#include "xla/ffi/api/c_api.h"
1313

‎transformer_engine/jax/csrc/extensions/attention.cpp‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
* See LICENSE for license information.
55
************************************************************************/
66

7-
#include "extensions.h"
7+
#include "../extensions.h"
88
#include "transformer_engine/fused_attn.h"
99
#include "transformer_engine/transformer_engine.h"
1010

‎transformer_engine/jax/csrc/extensions/cublas.cpp‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
* See LICENSE for license information.
55
************************************************************************/
66

7-
#include "extensions.h"
7+
#include "../extensions.h"
88
#include "transformer_engine/gemm.h"
99
#include "xla/ffi/api/c_api.h"
1010

‎transformer_engine/jax/csrc/extensions/cudnn.cpp‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@
66

77
#include "transformer_engine/cudnn.h"
88

9-
#include "extensions.h"
9+
#include "../extensions.h"
1010
#include "xla/ffi/api/c_api.h"
1111

1212
namespace transformer_engine {

‎transformer_engine/jax/csrc/extensions/gemm.cpp‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,9 +7,9 @@
77

88
#include <memory>
99

10+
#include "../extensions.h"
1011
#include "common/util/cuda_runtime.h"
1112
#include "common/util/system.h"
12-
#include "extensions.h"
1313
#include "xla/ffi/api/c_api.h"
1414

1515
namespace transformer_engine {

‎transformer_engine/jax/csrc/extensions/misc.cpp‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
* See LICENSE for license information.
55
************************************************************************/
66

7-
#include "extensions.h"
7+
#include "../extensions.h"
88

99
namespace transformer_engine {
1010
namespace jax {

‎transformer_engine/jax/csrc/extensions/normalization.cpp‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77

88
#include <cuda_runtime.h>
99

10-
#include "extensions.h"
10+
#include "../extensions.h"
1111

1212
namespace transformer_engine {
1313
namespace jax {

‎transformer_engine/jax/csrc/extensions/pybind.cpp‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
* See LICENSE for license information.
55
************************************************************************/
66

7-
#include "extensions.h"
7+
#include "../extensions.h"
88

99
namespace transformer_engine {
1010
namespace jax {

‎transformer_engine/jax/csrc/extensions/quantization.cpp‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
************************************************************************/
66
#include <cuda_runtime.h>
77

8-
#include "extensions.h"
8+
#include "../extensions.h"
99
#include "transformer_engine/cast.h"
1010
#include "transformer_engine/recipe.h"
1111
#include "xla/ffi/api/c_api.h"

0 commit comments

Comments
 (0)