Keras 3自定义Loss函数适配JAX/Torch后端:NaN掩码问题求助
Keras 3跨后端自定义掩码损失函数适配问题
我在Keras 3中构建可适配JAX或PyTorch后端的自定义Loss函数,需求是将y_true中特定值对应的y_pred和y_true索引进行掩码,仅将有效数据传入指定损失函数。但两种实现方式在JAX和Torch后端运行时均报错,报错提示无法处理张量索引或掩码操作。此前基于TensorFlow的实现可正常运行,以下是相关代码与报错信息:
实现代码
import keras from keras import Loss, ops class NanValueLossA(Loss): def __init__( self, loss_to_use=None, nan_value=None, name="nan_value_loss", **kwargs, ): self.nan_value = nan_value self.loss_to_use=loss_to_use super().__init__(name=name, **kwargs) def call(self, y_true, y_pred): valid_mask = ops.not_equal(y_true, self.nan_value) return self.loss_to_use(y_true[valid_mask], y_pred[valid_mask]) class NanValueLossB(Loss): def __init__( self, loss_to_use=None, nan_value=None, name="nan_value_loss", **kwargs, ): self.nan_value = nan_value self.loss_to_use=loss_to_use super().__init__(name=name, **kwargs) def call(self, y_true, y_pred): valid_mask = ops.not_equal(y_true, self.nan_value) valid_indices = ops.where(valid_mask) masked_y_pred = ops.take(y_pred,valid_indices) masked_y_true = ops.take(y_true,valid_indices) return self.loss_to_use(masked_y_true, masked_y_pred)
报错信息
NanValueLossA
- PyTorch后端:
File "c:\....\Lib\site-packages\keras\src\backend\torch\core.py", line 162, in convert_to_tensor x = x.to(device) ^^^^^^^^^^^^ NotImplementedError: Cannot copy out of meta tensor; no data!
- JAX后端:
File "c:....\Lib\site-packages\jax\_src\numpy\lax_numpy.py", line 6976, in _expand_bool_indices raise errors.NonConcreteBooleanIndexError(abstract_i) jax.errors.NonConcreteBooleanIndexError: Array boolean indices must be concrete; got ShapedArray(bool[32,1,128,128,1])
NanValueLossB
- PyTorch后端:
File "c:\....\Lib\site-packages\keras\src\backend\torch\core.py", line 162, in convert_to_tensor x = x.to(device) ^^^^^^^^^^^^ NotImplementedError: Cannot copy out of meta tensor; no data!
- JAX后端:
File "C:....\advanced_losses.py", line 651, in call valid_indices = ops.where(valid_mask) ^^^^^^^^^^^^^^^^^^^^^ File "....\Lib\site-packages\jax\_src\numpy\lax_numpy.py", line 1946, in where return nonzero(condition, size=size, fill_value=fill_value) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File ".....\Lib\site-packages\jax\_src\numpy\lax_numpy.py", line 2378, in nonzero calculated_size = core.concrete_dim_or_error(calculated_size, ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ jax.errors.ConcretizationTypeError: Abstract tracer value encountered where concrete value is expected: traced array with shape int32[]. The size argument of jnp.nonzero must be statically specified to use jnp.nonzero within JAX transformations.
原TensorFlow可运行实现
import numpy as np import tensorflow as tf from tensorflow.keras import backend as K def nan_mean_squared_error_loss(nan_value=np.nan): # Create a loss function def loss(y_true, y_pred): indices = tf.where(tf.not_equal(y_true, nan_value)) return tf.keras.losses.mean_squared_error( tf.gather_nd(y_true, indices), tf.gather_nd(y_pred, indices) ) # Return a function return loss
解决方案
问题核心在于JAX和PyTorch后端对动态张量索引/掩码的处理逻辑与TensorFlow不同,不能直接通过索引或ops.where提取动态长度的有效样本(JAX要求静态形状,PyTorch的Meta张量阶段无法执行数据提取)。正确的做法是保留原张量形状,通过掩码将无效区域的损失置0,最后求加权平均,避免动态改变张量形状。
修改后的跨后端兼容实现:
import keras from keras import Loss, ops class MaskedLoss(Loss): def __init__( self, loss_to_use=None, mask_value=None, name="masked_loss", **kwargs, ): self.mask_value = mask_value self.loss_to_use = loss_to_use super().__init__(name=name, **kwargs) def call(self, y_true, y_pred): # 生成有效掩码 valid_mask = ops.not_equal(y_true, self.mask_value) # 计算逐元素损失 elementwise_loss = self.loss_to_use(y_true, y_pred) # 掩码无效区域的损失为0 masked_loss = ops.where(valid_mask, elementwise_loss, 0.0) # 计算有效样本的损失均值(避免除以0) valid_count = ops.maximum(ops.sum(ops.cast(valid_mask, dtype="float32")), 1.0) return ops.sum(masked_loss) / valid_count
关键说明
- 不修改张量形状:全程保持
y_true和y_pred的原始形状,避免JAX的静态形状限制和PyTorch的Meta张量问题。 - 逐元素损失计算:先计算所有元素的损失,再通过掩码过滤无效值。
- 加权平均:用有效样本的数量做分母,确保损失计算的准确性,同时避免除以0的情况。
这样实现可以无缝适配Keras 3的JAX、PyTorch和TensorFlow后端,逻辑与原TensorFlow版本一致,但兼容跨后端要求。
内容的提问来源于stack exchange,提问作者João Santos
相关产品推荐
相关产品推荐

