在拟合TensorFlow模型时,如何对验证数据使用 `class_weights`?
我正在训练一个TensorFlow模型,想在评估训练批次和验证批次时使用 class_weight。然而,我只能在训练数据上使用权重。下面的代码演示了这个问题:
import tensorflow as tf
class MyAccuracy(tf.keras.metrics.Accuracy):
def update_state(self, y_true, y_pred, sample_weight=None):
tf.print("SAMPLE_WEIGHT", sample_weight)
super().update_state(y_true, y_pred, sample_weight)
input_dim = 2
output_dim = 3
inputs = tf.keras.layers.Input(shape=(input_dim,))
outputs = tf.keras.layers.Dense(output_dim)(inputs)
model = tf.keras.Model(inputs=inputs,
outputs=outputs)
model.compile(loss=tf.keras.losses.CategoricalCrossentropy(),
weighted_metrics=[MyAccuracy,])
model.fit(x=tf.keras.random.uniform((10, input_dim)),
y=tf.keras.random.uniform((10, output_dim)),
validation_split=0.5,
class_weight={0: .2,
1: .7,
2: .1},
verbose=2)
当我在TensorFlow 2.21和 Python 3.12.3下执行这段代码时,fit() 调用会先使用 class_weight 创建 sample_weight。随后输出如下:
SAMPLE_WEIGHT [0.7 0.7 0.2 0.2 0.7]
SAMPLE_WEIGHT None
1/1 - 1s - 759ms/step - accuracy: 0.0000e+00 - loss: 1.1852 - val_accuracy: 0.0000e+00 - val_loss: 9.0971
第一行输出显示了为训练批次提供的权重。第二行表明对验证批次并未提供权重。
如何让 fit() 在评估验证数据时提供权重?
解决方案
根据 tf.keras.Model.fit 的文档,class_weight 和 sample_weight 仅在训练阶段应用。如果你想在验证阶段应用权重,需要显式使用 validation_data 参数,而不是 validation_split,并在其中传入权重。
完整示例:
import tensorflow as tf
class MyAccuracy(tf.keras.metrics.SparseCategoricalAccuracy):
def update_state(self, y_true, y_pred, sample_weight=None):
tf.print("SAMPLE_WEIGHT", sample_weight)
super().update_state(y_true, y_pred, sample_weight)
input_dim = 2
output_dim = 3
inputs = tf.keras.layers.Input(shape=(input_dim,))
outputs = tf.keras.layers.Dense(output_dim)(inputs)
model = tf.keras.Model(inputs=inputs, outputs=outputs)
model.compile(
loss=tf.keras.losses.SparseCategoricalCrossentropy(),
weighted_metrics=[MyAccuracy()],
)
x = tf.random.uniform((10, input_dim))
y = tf.random.uniform((10, 1), 0, output_dim, dtype=tf.int32)
class_weight = {0: 0.2, 1: 0.7, 2: 0.1}
class_weight_tensor = tf.constant([class_weight[i] for i in range(output_dim)], dtype=tf.float32)
sample_weight = tf.gather(class_weight_tensor, tf.squeeze(y, axis=-1))
print(sample_weight.numpy()) # [0.7 0.2 0.2 0.2 0.1 0.1 0.7 0.1 0.7 0.1]
def split(x: tf.Tensor, validation_split: float) -> tuple[tf.Tensor, tf.Tensor]:
i = int(x.shape[0] * (1 - validation_split))
return x[:i], x[i:]
validation_split = 0.3
(x_train, x_val), (y_train, y_val), (weights_train, weights_val) = (
split(tensor, validation_split) for tensor in (x, y, sample_weight)
)
model.fit(
x=x_train,
y=y_train,
sample_weight=weights_train,
validation_data=(x_val, y_val, weights_val),
verbose=2,
)
站内所有文章版权归属LeftHeroAI导航站,无授权禁止任何主体转载、抄袭、复制内容,亦不得私自架设镜像站点。一经侵权,本站将通过法律途径追责。