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

如何在自定义知识蒸馏Keras模型的train_step中应用class_weight

在Keras自定义知识蒸馏模型中处理类别不平衡问题

一、解决class_weight传入报错的问题

自定义模型的train_step不会自动处理class_weight参数,你需要手动将类别权重转换为样本权重,再在损失计算中应用。具体步骤如下:

  1. 计算类别权重对应的样本权重
    你可以根据标签分布手动计算,比如0类(多数类)权重设为1,1类(少数类)权重为9024/842≈10.72;也可以用sklearn工具自动计算平衡权重:

    from sklearn.utils.class_weight import compute_class_weight
    import numpy as np
    
    class_weights = compute_class_weight('balanced', classes=np.unique(y_train), y=y_train)
    class_weight_dict = {0: class_weights[0], 1: class_weights[1]}
    
  2. 在自定义Distiller的train_step中集成样本权重
    修改train_step方法,要么内部根据标签生成样本权重,要么直接接收外部传入的sample_weight,并将其应用到损失计算和指标更新中。示例代码如下:

    class Distiller(keras.Model):
        def __init__(self, student, teacher):
            super().__init__()
            self.teacher = teacher
            self.student = student
            self.class_weight = None  # 用于存储类别权重
    
        def compile(
            self,
            optimizer,
            metrics,
            student_loss_fn,
            distillation_loss_fn,
            alpha=0.1,
            temperature=3,
            class_weight=None
        ):
            super().compile(optimizer=optimizer, metrics=metrics)
            self.student_loss_fn = student_loss_fn
            self.distillation_loss_fn = distillation_loss_fn
            self.alpha = alpha
            self.temperature = temperature
            self.class_weight = class_weight  # 保存类别权重
    
        def train_step(self, data):
            x, y = data
            sample_weight = None
            # 根据标签生成样本权重
            if self.class_weight is not None:
                sample_weight = tf.gather(list(self.class_weight.values()), tf.cast(y, tf.int32))
    
            # 教师模型预测(不训练)
            teacher_predictions = self.teacher(x, training=False)
    
            with tf.GradientTape() as tape:
                student_predictions = self.student(x, training=True)
                # 计算基础损失与蒸馏损失
                student_loss = self.student_loss_fn(y, student_predictions)
                distillation_loss = self.distillation_loss_fn(
                    tf.nn.softmax(teacher_predictions / self.temperature, axis=1),
                    tf.nn.softmax(student_predictions / self.temperature, axis=1),
                )
                total_loss = self.alpha * student_loss + (1 - self.alpha) * distillation_loss
    
                # 应用样本权重到总损失
                if sample_weight is not None:
                    total_loss = tf.reduce_mean(total_loss * sample_weight)
    
            # 更新学生模型参数
            trainable_vars = self.student.trainable_variables
            gradients = tape.gradient(total_loss, trainable_vars)
            self.optimizer.apply_gradients(zip(gradients, trainable_vars))
    
            # 带权重更新评估指标
            self.compiled_metrics.update_state(y, student_predictions, sample_weight=sample_weight)
            return {m.name: m.result() for m in self.metrics}
    

    调用时直接传入class_weight即可:

    distiller.compile(
        optimizer='adam',
        metrics=['accuracy'],
        student_loss_fn=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
        distillation_loss_fn=keras.losses.KLDivergence(),
        alpha=0.1,
        temperature=3,
        class_weight=class_weight_dict
    )
    distiller.fit(x_train, y_train, epochs=10)
    

    或者提前转换为sample_weight传入fit,效果一致:

    sample_weights = np.array([class_weight_dict[label] for label in y_train])
    distiller.fit(x_train, y_train, sample_weight=sample_weights, epochs=10)
    

二、class_weight加权与正样本10倍采样的区别

这两种方式并非完全等价,但核心目标都是让模型重视少数类:

  • class_weight加权:不改变样本数量,仅在损失计算时给少数类的错误赋予更高权重。优势是不会增加数据集大小,避免过拟合少数类,计算效率更高;缺点是如果少数类样本质量差,加权可能放大噪声影响。
  • 正样本10倍采样:通过复制少数类样本增加其数量,让模型多次学习少数类样本。优势是直观,适合少数类样本质量高的场景;缺点是会增加训练数据量,容易导致模型过拟合少数类,尤其是当少数类样本本身存在噪声时。

你的场景中,class_weight计算出的少数类权重约为10.72,和10倍采样的效果接近,但本质不同——加权是损失层面的调整,采样是数据层面的调整,实际训练效果可能因数据集特性略有差异。

内容的提问来源于stack exchange,提问作者Jaeyoung Park

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 20:35:18