Skip to content
Open
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
49 changes: 41 additions & 8 deletions src/kamae/keras/tensorflow/layers/list_max.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,11 @@
from kamae.keras.core.backend import TENSORFLOW_ONLY
from kamae.keras.core.base import BaseLayer
from kamae.keras.core.utils.input_utils import allow_single_or_multiple_tensor_input
from kamae.keras.tensorflow.utils.list_utils import get_top_n, segmented_operation
from kamae.keras.tensorflow.utils.list_utils import (
get_top_n,
min_filter_mask,
segmented_operation,
)
from kamae.keras.tensorflow.utils.transform_utils import map_fn_w_axis


Expand Down Expand Up @@ -108,6 +112,10 @@ def compatible_dtypes(self) -> Optional[List[str]]:
"float16",
"float32",
"float64",
"int8",
"int16",
"int32",
"int64",
"string",
]

Expand Down Expand Up @@ -145,11 +153,12 @@ def _call(self, inputs: Iterable[KerasTensor], **kwargs: Any) -> KerasTensor:

# Apply the mask to filter out elements less than or equal to the threshold
if self.min_filter_value is not None:
mask = tf.greater_equal(val_tensor, self.min_filter_value)
neg_inf = val_tensor.dtype.min
val_tensor = tf.where(mask, val_tensor, neg_inf)
else:
val_tensor = val_tensor
mask = min_filter_mask(val_tensor, self.min_filter_value)
# The dtype minimum is only a neutral element for the reduction here: a
# real value equal to it still wins the max, so substituting it for the
# filtered entries cannot change the result.
val_tensor = tf.where(mask, val_tensor, val_tensor.dtype.min)
kept = tf.cast(mask, tf.int32)

# Apply segmented calculation
if self.with_segment:
Expand All @@ -167,8 +176,32 @@ def _call(self, inputs: Iterable[KerasTensor], **kwargs: Any) -> KerasTensor:
listwise_max = tf.broadcast_to(listwise_max, output_shape)

if self.min_filter_value is not None:
fill_val = tf.constant(self.nan_fill_value, dtype=listwise_max.dtype)
listwise_max = tf.where(listwise_max != neg_inf, listwise_max, fill_val)
# Whether anything survived the filter is read from the mask rather than
# by testing the result against the sentinel. dtype.min is a legitimate
# value for the narrow integer dtypes (-128 for int8), so a sentinel test
# would silently overwrite real data with nan_fill_value.
if self.with_segment:
any_kept = map_fn_w_axis(
elems=[kept, segment_tensor],
fn=lambda x: segmented_operation(x, tf.math.unsorted_segment_max),
axis=self.axis,
fn_output_signature=tf.TensorSpec(
shape=kept.shape[self.axis :], dtype=kept.dtype
),
)
any_kept = tf.ensure_shape(any_kept, kept.shape)
else:
any_kept = tf.reduce_max(kept, axis=self.axis, keepdims=True)
any_kept = tf.broadcast_to(any_kept, output_shape)

# nan_fill_value is a Python float, which tf.constant cannot convert
# directly to an integer dtype. Narrowing via numpy first handles the
# integer dtypes while preserving full precision for the float ones.
fill_val = tf.constant(
listwise_max.dtype.as_numpy_dtype(self.nan_fill_value),
dtype=listwise_max.dtype,
)
listwise_max = tf.where(any_kept > 0, listwise_max, fill_val)

return listwise_max

Expand Down
7 changes: 6 additions & 1 deletion src/kamae/keras/tensorflow/utils/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,5 +37,10 @@
datetime_year,
unix_timestamp_to_datetime,
)
from .list_utils import get_top_n, listify_tensors, segmented_operation # noqa: F401
from .list_utils import ( # noqa: F401
get_top_n,
listify_tensors,
min_filter_mask,
segmented_operation,
)
from .transform_utils import map_fn_w_axis # noqa: F401
26 changes: 26 additions & 0 deletions src/kamae/keras/tensorflow/utils/list_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import math
from typing import Any, Callable, List, Union

import numpy as np
Expand Down Expand Up @@ -70,6 +71,31 @@ def get_top_n(
)


def min_filter_mask(val_tensor: Tensor, min_filter_value: float) -> Tensor:
"""
Build the mask of the entries that meet a minimum value threshold.

TensorFlow will not compare an integer tensor against a float threshold, so for
those the threshold is rounded up, which is equivalent under >= on integers.
A threshold beyond the bounds of the dtype cannot be narrowed at all, since the
cast would wrap around silently, so those two cases are answered directly.
Floats need neither adjustment because they saturate to +/-inf rather than wrap.

:param val_tensor: Value tensor to filter.
:param min_filter_value: Minimum value an entry must meet to be kept.
:returns: Boolean tensor that is True wherever the entry is kept.
"""
if not val_tensor.dtype.is_integer:
return tf.greater_equal(val_tensor, min_filter_value)

threshold = math.ceil(min_filter_value)
if threshold <= val_tensor.dtype.min:
return tf.ones_like(val_tensor, dtype=tf.bool)
if threshold > val_tensor.dtype.max:
return tf.zeros_like(val_tensor, dtype=tf.bool)
return tf.greater_equal(val_tensor, tf.cast(threshold, val_tensor.dtype))


def listify_tensors(x: Union[tf.Tensor, np.ndarray, List[Any]]) -> List[Any]:
"""
Converts any tensors or numpy arrays to lists for config serialization.
Expand Down
15 changes: 14 additions & 1 deletion src/kamae/spark/transformers/list_max.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,16 @@
import tensorflow as tf
from pyspark import keyword_only
from pyspark.sql import DataFrame
from pyspark.sql.types import DataType, DoubleType, FloatType, StringType
from pyspark.sql.types import (
ByteType,
DataType,
DoubleType,
FloatType,
IntegerType,
LongType,
ShortType,
StringType,
)

from kamae.keras.core.backend import TENSORFLOW_ONLY
from kamae.keras.tensorflow.layers import ListMaxLayer
Expand Down Expand Up @@ -124,6 +133,10 @@ def compatible_dtypes(self) -> Optional[List[DataType]]:
return [
FloatType(),
DoubleType(),
ByteType(),
ShortType(),
IntegerType(),
LongType(),
StringType(),
]

Expand Down
Loading
Loading