You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

关键说明

  1. 不修改张量形状:全程保持y_true和y_pred的原始形状,避免JAX的静态形状限制和PyTorch的Meta张量问题。
  2. 逐元素损失计算:先计算所有元素的损失,再通过掩码过滤无效值。
  3. 加权平均:用有效样本的数量做分母,确保损失计算的准确性,同时避免除以0的情况。

这样实现可以无缝适配Keras 3的JAX、PyTorch和TensorFlow后端,逻辑与原TensorFlow版本一致,但兼容跨后端要求。

内容的提问来源于stack exchange,提问作者João Santos

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.21 19:55:54