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
3 changes: 2 additions & 1 deletion dcor/_dcor_internals_numba.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
from numba.types import Array, Tuple

from ._dcor_internals import _dcov_from_terms
from ._utils import FS_CACHE

if TYPE_CHECKING:
import numpy.typing
Expand Down Expand Up @@ -36,7 +37,7 @@
int64,
boolean,
),
cache=True,
cache=FS_CACHE,
)(_dcov_from_terms)


Expand Down
14 changes: 7 additions & 7 deletions dcor/_fast_dcov_avl.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@
NumbaVectorReadOnlyNonContiguous,
_generate_distance_covariance_sqr_from_terms_impl,
)
from ._utils import CompileMode, _transform_to_1d
from ._utils import CompileMode, _transform_to_1d, FS_CACHE

if TYPE_CHECKING:
NumpyArrayType = np.typing.NDArray[np.number[Any]]
Expand Down Expand Up @@ -141,7 +141,7 @@ def _dyad_update_compiled_version(
NumbaVector,
NumbaIntVectorReadOnly,
),
cache=True,
cache=FS_CACHE,
)(
_dyad_update_compiled_version,
)
Expand Down Expand Up @@ -193,7 +193,7 @@ def _partial_sum_2d(
NumbaIntVectorReadOnly,
NumbaVector,
),
cache=True,
cache=FS_CACHE,
)(
_generate_partial_sum_2d(compiled=True),
)
Expand Down Expand Up @@ -303,7 +303,7 @@ def _get_impl_args(
NumbaIntVectorReadOnly,
NumbaMatrix,
))(NumbaVectorReadOnlyNonContiguous, NumbaVectorReadOnlyNonContiguous),
cache=True,
cache=FS_CACHE,
)(_get_impl_args)


Expand Down Expand Up @@ -429,7 +429,7 @@ def _distance_covariance_sqr_terms_avl_impl(
numba.optional(float64),
numba.optional(float64),
))(NumbaVectorReadOnlyNonContiguous, NumbaVectorReadOnlyNonContiguous, boolean),
cache=True,
cache=FS_CACHE,
)(
_generate_distance_covariance_sqr_terms_avl_impl(compiled=True),
)
Expand All @@ -446,7 +446,7 @@ def _distance_covariance_sqr_terms_avl_impl(
NumbaVectorReadOnlyNonContiguous,
boolean,
),
cache=True,
cache=FS_CACHE,
)(
_generate_distance_covariance_sqr_from_terms_impl(
compiled=True,
Expand Down Expand Up @@ -582,7 +582,7 @@ def _rowwise_distance_covariance_sqr_avl_generic_internal(
)],
'(n),(n),()->()',
nopython=True,
cache=True,
cache=FS_CACHE,
target=target,
)(_rowwise_distance_covariance_sqr_avl_generic_internal)

Expand Down
12 changes: 6 additions & 6 deletions dcor/_fast_dcov_mergesort.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@
NumbaVectorReadOnlyNonContiguous,
_generate_distance_covariance_sqr_from_terms_impl,
)
from ._utils import CompileMode, _transform_to_1d
from ._utils import CompileMode, _transform_to_1d, FS_CACHE

if TYPE_CHECKING:
NumpyArrayType = np.typing.NDArray[np.number[Any]]
Expand Down Expand Up @@ -132,7 +132,7 @@ def _compute_weight_sums(

_compute_weight_sums_compiled = numba.njit(
NumbaMatrix(NumbaVectorReadOnly, NumbaMatrixReadOnly),
cache=True,
cache=FS_CACHE,
)(_compute_weight_sums)


Expand Down Expand Up @@ -179,7 +179,7 @@ def _compute_aijbij_term(
_compute_aijbij_term = _generate_compute_aijbij_term(compiled=False)
_compute_aijbij_term_compiled = numba.njit(
float64(NumbaVectorReadOnly, NumbaVectorReadOnly),
cache=True,
cache=FS_CACHE,
)(
_generate_compute_aijbij_term(
compiled=True,
Expand All @@ -205,7 +205,7 @@ def _compute_row_sums(

_compute_row_sums_compiled = numba.njit(
NumbaVector(NumbaVectorReadOnly),
cache=True)(_compute_row_sums)
cache=FS_CACHE)(_compute_row_sums)


def _generate_distance_covariance_sqr_terms_mergesort_impl(
Expand Down Expand Up @@ -285,7 +285,7 @@ def _distance_covariance_sqr_terms_mergesort_impl(
NumbaVectorReadOnlyNonContiguous,
boolean,
),
cache=True,
cache=FS_CACHE,
)(
_generate_distance_covariance_sqr_terms_mergesort_impl(
compiled=True,
Expand All @@ -306,7 +306,7 @@ def _distance_covariance_sqr_terms_mergesort_impl(
NumbaVectorReadOnlyNonContiguous,
boolean,
),
cache=True,
cache=FS_CACHE,
)(
_generate_distance_covariance_sqr_from_terms_impl(
compiled=True,
Expand Down
3 changes: 3 additions & 0 deletions dcor/_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
import enum
import warnings
from typing import TYPE_CHECKING, Any, Iterable, TypeVar, Union
import os

import numpy as np
from array_api_compat import (
Expand All @@ -26,6 +27,8 @@
None,
]

FS_CACHE = False if os.environ.get("DCOR_DISABLE_FS_CACHE") else True


class CompileMode(enum.Enum):
"""Compilation mode of the algorithm."""
Expand Down
Loading