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
74 changes: 58 additions & 16 deletions conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,33 +5,75 @@
# running tests from inside VS code.
# See https://stackoverflow.com/a/34520971

import json
import os
from importlib.metadata import PackageNotFoundError, version

import pytest

from stumpy import rng


def get_specs():
"""
Find and return all package versions
"""
pkgs = [
# Alphabetical Order
"black",
"coverage",
"dask",
"distributed",
"flake8",
"isort",
"numba",
"numpy",
"pandas",
"polars",
"pytest",
"python",
"ray",
"scipy",
]

specs = []
for pkg in pkgs:
try: # pragma: no cover
pkg_version = version(pkg)
specs.append(f"--spec {pkg}={pkg_version}")
except PackageNotFoundError:
pass

return " ".join(specs)


def get_env_vars():
"""
Find and return all environment variables
"""
keys = [
"NUMBA_DISABLE_JIT",
"NUMBA_ENABLE_CUDASIM",
]

env_vars = []
for key in keys:
value = os.getenv(key)
if value is not None:
env_vars.append(f"{key}={value}")

return " ".join(env_vars)


def pytest_configure(config):
"""
Called after command line options have been parsed
and all plugins and initial conftest files been loaded.
"""
state = rng.STATE
state_str = json.dumps(
(
state[0],
state[1].tolist(),
state[2],
state[3],
state[4],
)
)

# Store details of starting random state in case of failure
pytest.STUMPY_MSG = (
f"\n\nSTUMPY_STATE='{state_str}' pixi run tests custom {config.args[0]}"
)
# Store details of starting random seed in case of failure
env_vars = get_env_vars()
specs = get_specs()
pytest.STUMPY_MSG = f"\n\nSTUMPY_SEED={rng.SEED} {env_vars} "
pytest.STUMPY_MSG += f"pixi exec {specs} ./test.sh custom 1 {config.args[0]}"


def pytest_sessionfinish(session, exitstatus):
Expand Down
22 changes: 2 additions & 20 deletions stumpy/rng.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,4 @@
import json
import os
import warnings
from contextlib import contextmanager

import numpy as np
Expand All @@ -9,28 +7,12 @@
# in order to account for unit testing
if os.getenv("STUMPY_SEED") is not None: # pragma: no cover
SEED = int(os.getenv("STUMPY_SEED"))
if SEED == 0:
raise ValueError("STUMPY_SEED must be greater than zero!")
else:
SEED = np.random.randint(1, 4_294_967_296, dtype=np.uint32)
RNG = np.random.RandomState(seed=SEED)

if os.getenv("STUMPY_STATE") is not None: # pragma: no cover
if os.getenv("STUMPY_SEED") is not None: # pragma: no cover
warnings.warn("STUMPY_SEED was ignored in lieu of STUMPY_STATE")
state_str = os.getenv("STUMPY_STATE")
state = json.loads(state_str)
STATE = (
state[0],
np.array(state[1], dtype=np.uint32),
state[2],
state[3],
state[4],
)
RNG.set_state(STATE)
else:
STATE = RNG.get_state()

# seed = RNG.get_state()[1][0]


@contextmanager
def fix_seed(seed):
Expand Down
Loading