class_weight参数通过加权损失使模型重视少数类,应使用sklearn推荐公式计算权重,适用于标准fit流程但不适用于自定义训练循环;过采样需在tf.data前完成且注意数据类型;多分类不可用weighted_cross_entropy;验证时需用pr曲线、精准率与召回率评估,并对推理概率做后校准。

用 class_weight 参数让模型正视少数类
TensorFlow 的 model.fit() 支持直接传入 class_weight,这是最轻量、最常用的处理方式。它不会改动数据,只在损失计算时给少数类样本更高的权重,相当于告诉模型“分错这个类别代价更大”。
常见错误是手动算错权重比例——别用 1 / class_count 简单倒数,而应按 sklearn 推荐公式:class_weight = total_samples / (n_classes * class_count),这样能平衡各类总权重贡献。
- 用
sklearn.utils.class_weight.compute_class_weight()生成字典,再转成{0: w0, 1: w1, ...}格式传给fit() - 二分类时,若正样本仅占 5%,
class_weight大致为{0: 1.0, 1: 19.0},不是 20.0(注意分母是n_classes * class_count) - 该方式对
tf.keras.Model和tf.keras.Sequential都有效,但不适用于自定义训练循环(tf.GradientTape),需手动加权损失
过采样慎用 imblearn + tf.data.Dataset
用 imblearn.over_sampling.SMOTE 或 RandomOverSampler 在 fit 前扩增少数类,看似直观,但和 TensorFlow 的 tf.data.Dataset 流水线容易冲突:SMOTE 要求输入是 numpy 数组,而 Dataset 常含 tf.Tensor 或动态 batch,直接喂会报 ValueError: Expected 2D array, got 1D array instead。
图片提示词生成器?不止如此。 马甲系统 —— 把脑海中的画面,翻译成AI能理解的专业表达。 用得越多,它越懂你:首次需要多问几句确认方向,用久了几乎一说就懂。 用得越多,它越快:缓存机制让后续对话越来越省。 RAG进化:成功案例持续入库,越跑越聪明。 输入「新手指南」查看完整功能介绍
- 必须在构建
Dataset之前完成重采样:先as_numpy_iterator()拿出全部数据(小心内存!),再用SMOTE.fit_resample(X, y) - 重采样后标签类型可能变(如从
int64变float64),训练时报InvalidArgumentError: labels must be an int32 or int64,需显式y_resampled = y_resampled.astype(np.int32) - 时间序列或图像等结构化数据慎用 SMOTE——它对特征做线性插值,会生成无意义的中间样本
自定义损失函数时,tf.nn.weighted_cross_entropy_with_logits 不适用多分类
有人看到“weighted”就直奔 tf.nn.weighted_cross_entropy_with_logits,但它只支持二分类(logits 形状为 [batch, 1])。多分类要用 tf.keras.losses.CategoricalCrossentropy 或 SparseCategoricalCrossentropy 的 sample_weight 参数,或自己写带类权重的 loss。
- 多分类下正确做法:在
model.compile()中设loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),然后在fit()里传sample_weight数组(长度 = batch_size,每个值对应样本的类权重) - 若用自定义
@tf.function训练循环,需在 loss 计算中手动乘权重:weighted_loss = tf.reduce_mean(loss_per_sample * sample_weights) -
tf.nn.weighted_cross_entropy_with_logits的pos_weight是标量,只能调正负样本比,无法支持三类及以上
验证阶段别只看 accuracy,tf.keras.metrics.AUC 默认不处理类别偏斜
accuracy 在 95% 负样本场景下可能高达 0.95,但完全没意义。AUC 看起来更鲁棒,但默认 tf.keras.metrics.AUC 用的是全量样本的 ROC 曲线,当负样本爆炸时,FPR 计算会被淹没,曲线贴边,AUC 值虚高。
- 改用
tf.keras.metrics.AUC(curve='PR')(PR 曲线),它关注查准率/查全率,在不平衡场景下更敏感 - 务必同时监控
tf.keras.metrics.Precision(class_id=1)和tf.keras.metrics.Recall(class_id=1),尤其后者——召回率低说明模型根本不敢预测少数类 - 验证集本身也要保持和训练集一致的分布,不要用
stratify重划分后又忽略 class_weight,否则评估失真
实际中最容易被忽略的是:class_weight 只影响训练损失,不改变模型输出 logits 的尺度,因此推理时 softmax 后的概率依然偏向多数类。如果业务需要校准概率(比如风控阈值决策),得额外用 Platt scaling 或 isotonic regression 对输出做后处理——这一步不在 TensorFlow 内置流程里,得自己接 sklearn。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!










