Skip to content
Merged
Show file tree
Hide file tree
Changes from all 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
37 changes: 37 additions & 0 deletions tests/e2e/nightly/single_node/ops/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,43 @@
from datetime import datetime
import pytest

DURATION_THRESHOLD = 120

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

The Pull Request title and summary do not adhere to the repository's style guide regarding format and required sections.

Suggested PR Title:

[CI][Misc] Add timeout check for custom op CI and optimize test parameters

Suggested PR Summary:

### What this PR does / why we need it?

This PR introduces a mechanism to track test duration in `conftest.py` and skip subsequent tests in a file if a certain number of tests exceed a timeout threshold. This is intended to prevent CI hangs or long-running nightly tests. Additionally, it reduces the parameter space for `test_fused_qkvzba_split_reshape_cat.py` to further optimize CI runtime.

### Does this PR introduce _any_ user-facing change?

no

### How was this patch tested?

nightly
References
  1. The PR title and summary must follow the specific format defined in the Repository Style Guide, including the [Branch][Module][Action] prefix for the title and specific headers for the summary. (link)

SLOW_COUNT_LIMIT = 5


_per_file_slow_cases = {}
_current_file = None


def pytest_runtest_setup(item):
item.start_time = time.time()


def pytest_runtest_teardown(item, nextitem):
global _current_file

file_path = item.fspath
duration = time.time() - item.start_time


if file_path not in _per_file_slow_cases:
_per_file_slow_cases[file_path] = 0

if duration > DURATION_THRESHOLD:
_per_file_slow_cases[file_path] += 1
cnt = _per_file_slow_cases[file_path]
print(f" Detected that the test case took too long, ({cnt}/{SLOW_COUNT_LIMIT}):{duration:.2f}s")

if cnt >= SLOW_COUNT_LIMIT:
print(f"\n The number of timeout test cases {file_path} ≥{SLOW_COUNT_LIMIT}\n")
_current_file = file_path


def pytest_runtest_call(item):
if _current_file == item.fspath:
print(f"CASE SKIP:{item.nodeid}")
pytest.skip(f"The use case takes too long.")
Comment on lines +9 to +40

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

The current implementation uses a single global variable _current_file to track which file should be skipped. This logic is flawed if multiple test files are executed: once a new file exceeds the SLOW_COUNT_LIMIT, _current_file is overwritten, and any remaining tests in previously "timed out" files will no longer be skipped.

I suggest using the _per_file_slow_cases dictionary directly to check the skip condition for each file, which is more robust and removes the need for the _current_file global variable.

_per_file_slow_cases = {}


def pytest_runtest_setup(item):
    item.start_time = time.time()


def pytest_runtest_teardown(item, nextitem):
    file_path = item.fspath
    duration = time.time() - item.start_time

    if duration > DURATION_THRESHOLD:
        _per_file_slow_cases[file_path] = _per_file_slow_cases.get(file_path, 0) + 1
        cnt = _per_file_slow_cases[file_path]
        print(f" Detected that the test case took too long, ({cnt}/{SLOW_COUNT_LIMIT}):{duration:.2f}s")

        if cnt >= SLOW_COUNT_LIMIT:
            print(f"\n The number of timeout test cases  {file_path}   ≥{SLOW_COUNT_LIMIT}\n")


def pytest_runtest_call(item):
    if _per_file_slow_cases.get(item.fspath, 0) >= SLOW_COUNT_LIMIT:
        print(f"CASE SKIP:{item.nodeid}")
        pytest.skip("The use case takes too long.")



@pytest.hookimpl(tryfirst=True, hookwrapper=True)
def pytest_runtest_makereport(item, call):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -38,11 +38,11 @@ def validate_cmp(y_cal, y_ref, dtype, device='npu'):
'Invalid parameter \"dtype\" is found : {}'.format(dtype))


@pytest.mark.parametrize("seq_len", [1, 16, 64, 128, 256, 1024, 2048, 3567])
@pytest.mark.parametrize("num_heads_qk", [2, 4, 8, 16])
@pytest.mark.parametrize("num_heads_v", [2, 4, 8])
@pytest.mark.parametrize("head_qk_dim", [64, 128, 256])
@pytest.mark.parametrize("head_v_dim", [64, 128])
@pytest.mark.parametrize("seq_len", [1, 64, 1024, 2048])
@pytest.mark.parametrize("num_heads_qk", [2, 4, 8])
@pytest.mark.parametrize("num_heads_v", [8])
@pytest.mark.parametrize("head_qk_dim", [256])
@pytest.mark.parametrize("head_v_dim", [128])
@pytest.mark.parametrize("dtype",
[torch.float32, torch.float16, torch.bfloat16])
def test_fused_qkvzba_split_reshape_cat(
Expand Down
Loading