Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 5 additions & 5 deletions .github/workflows/intel-a770.yml
Original file line number Diff line number Diff line change
Expand Up @@ -81,21 +81,21 @@ jobs:
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=1 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run pytest on test files
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
continue-on-error: true
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run pytest on varlen test files
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
continue-on-error: true
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=0 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }} || \
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }} || \
echo "Varlen tests failed (non-critical)"

- name: Test full compiling on all test files
Expand All @@ -110,12 +110,12 @@ jobs:
continue-on-error: true
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run full pytest on varlen test files
if: false && github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') && steps.check_skip.outputs.skip_tests == 'false'
continue-on-error: true
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=0 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }} || \
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }} || \
echo "Varlen tests failed (non-critical)"
20 changes: 10 additions & 10 deletions .github/workflows/nvidia-4090.yml
Original file line number Diff line number Diff line change
Expand Up @@ -86,19 +86,19 @@ jobs:
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=1 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run pytest on test files
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run pytest on varlen test files
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=0 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }} || \
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }} || \
echo "Varlen tests failed (non-critical)"

- name: Test full compiling on all test files
Expand All @@ -111,13 +111,13 @@ jobs:
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run full pytest on varlen test files
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=0 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }} || \
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }} || \
echo "Varlen tests failed (non-critical)"

test-models:
Expand Down Expand Up @@ -169,19 +169,19 @@ jobs:
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=1 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run pytest on test files
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run pytest on varlen test files
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=0 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }} || \
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }} || \
echo "Varlen tests failed (non-critical)"

- name: Test full compiling on all test files
Expand All @@ -194,11 +194,11 @@ jobs:
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run full pytest on varlen test files
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=0 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }} || \
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }} || \
echo "Varlen tests failed (non-critical)"
20 changes: 10 additions & 10 deletions .github/workflows/nvidia-a100.yml
Original file line number Diff line number Diff line change
Expand Up @@ -87,19 +87,19 @@ jobs:
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=1 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run pytest on test files
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run pytest on varlen test files
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=0 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }} || \
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }} || \
echo "Varlen tests failed (non-critical)"

- name: Test full compiling on all test files
Expand All @@ -112,13 +112,13 @@ jobs:
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run full pytest on varlen test files
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=0 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }} || \
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }} || \
echo "Varlen tests failed (non-critical)"

test-models:
Expand Down Expand Up @@ -170,19 +170,19 @@ jobs:
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=1 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run pytest on test files
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run pytest on varlen test files
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=0 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }} || \
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }} || \
echo "Varlen tests failed (non-critical)"

- name: Test full compiling on all test files
Expand All @@ -195,11 +195,11 @@ jobs:
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run full pytest on varlen test files
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=0 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }} || \
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }} || \
echo "Varlen tests failed (non-critical)"
20 changes: 10 additions & 10 deletions .github/workflows/nvidia-h100.yml
Original file line number Diff line number Diff line change
Expand Up @@ -87,19 +87,19 @@ jobs:
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=1 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run pytest on test files
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run pytest on varlen test files
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=0 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }} || \
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }} || \
echo "Varlen tests failed (non-critical)"

- name: Test full compiling on all test files
Expand All @@ -112,13 +112,13 @@ jobs:
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run full pytest on varlen test files
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=0 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }} || \
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }} || \
echo "Varlen tests failed (non-critical)"

test-models:
Expand Down Expand Up @@ -170,19 +170,19 @@ jobs:
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=1 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run pytest on test files
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run pytest on varlen test files
if: steps.find-dependent-tests.outputs.test_files && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=0 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }} || \
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }} || \
echo "Varlen tests failed (non-critical)"

- name: Test full compiling on all test files
Expand All @@ -195,11 +195,11 @@ jobs:
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=1 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }}
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }}

- name: Run full pytest on varlen test files
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags/') && steps.check_skip.outputs.skip_tests == 'false'
run: |
FLA_COMPILER_MODE=0 TRITON_PRINT_AUTOTUNING=0 SKIP_TEST_CHUNK_VARLEN=0 \
pytest ${{ steps.find-dependent-tests.outputs.test_files }} || \
pytest -s -v ${{ steps.find-dependent-tests.outputs.test_files }} || \
echo "Varlen tests failed (non-critical)"
54 changes: 38 additions & 16 deletions tests/models/test_modeling_abc.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,25 +11,47 @@
# ===================================================================================
# Test for Modeling (Forward/Backward Pass)
# ===================================================================================
@pytest.mark.parametrize("L", [4])
@pytest.mark.parametrize("B", [4])
@pytest.mark.parametrize("T", [1024])
@pytest.mark.parametrize("H", [4])
@pytest.mark.parametrize("D", [64, 128])
@pytest.mark.parametrize("dtype", [torch.bfloat16])
@pytest.mark.parametrize("use_l2warp", [True, False])
def test_modeling(L, B, T, H, D, dtype, use_l2warp):
run_test_model_forward_backward(L, B, T, H, D, ABCConfig, dtype, use_l2warp)
@pytest.mark.parametrize(
['L', 'B', 'T', 'H', 'D', 'dtype', 'use_l2warp'],
[
pytest.param(*test, id="L{}-B{}-T{}-H{}-D{}-use_l2warp{}-{}".format(*test))
for test in [
(4, 4, 1024, 4, 64, True, torch.bfloat16),
(4, 4, 1024, 4, 64, False, torch.bfloat16),
(4, 4, 1024, 4, 128, False, torch.bfloat16),
]
]
)
def test_modeling(
L: int,
B: int,
T: int,
H: int,
D: int,
use_l2warp: bool,
dtype: torch.dtype,
):
run_test_model_forward_backward(L, B, T, H, D, ABCConfig, use_l2warp=use_l2warp, dtype=dtype)


# ===================================================================================
# Test for Generation
# ===================================================================================
@pytest.mark.parametrize("L", [2])
@pytest.mark.parametrize("B", [4])
@pytest.mark.parametrize("T", [4000])
@pytest.mark.parametrize("H", [8])
@pytest.mark.parametrize("D", [64])
@pytest.mark.parametrize("dtype", [torch.float16])
def test_generation(L, B, T, H, D, dtype):
@pytest.mark.parametrize(
['L', 'B', 'T', 'H', 'D', 'dtype'],
[
pytest.param(*test, id="L{}-B{}-T{}-H{}-D{}-{}".format(*test))
for test in [
(2, 4, 2000, 8, 64, torch.float16),
]
]
)
def test_generation(
L: int,
B: int,
T: int,
H: int,
D: int,
dtype: torch.dtype,
):
run_test_generation(L, B, T, H, D, ABCConfig, dtype)
6 changes: 3 additions & 3 deletions tests/models/test_modeling_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,8 +31,8 @@ def run_test_model_forward_backward(
H: int,
D: int,
config_class: type,
dtype: torch.dtype,
use_l2warp: bool,
dtype: torch.dtype,
):
"""
A foundational test for the forward and backward passes of a model.
Expand All @@ -44,7 +44,7 @@ def run_test_model_forward_backward(
if config_class.__name__ in NOT_READY_FOR_TESTING:
pytest.skip(f"{config_class.__name__} is not yet ready for testing.")

model, config = create_model_and_config(config_class, L, H, D, dtype, use_l2warp=use_l2warp)
model, config = create_model_and_config(config_class, L, H, D, use_l2warp=use_l2warp, dtype=dtype)
input_ids = torch.randint(low=0, high=config.vocab_size, size=(B, T), device=device)
output_fixed = model(input_ids, output_hidden_states=True).hidden_states[-1]
assert output_fixed.shape == (B, T, config.hidden_size)
Expand Down Expand Up @@ -87,7 +87,7 @@ def run_test_generation(
pytest.skip(f"{config_class.__name__} is not yet ready for testing.")

if model is None:
model, config = create_model_and_config(config_class, L, H, D, dtype, use_l2warp=use_l2warp)
model, config = create_model_and_config(config_class, L, H, D, use_l2warp=use_l2warp, dtype=dtype)
model.eval()
model = model.to(dtype).to(device)

Expand Down
Loading
Loading