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

TensorFlow类别样本数与基数不符及类别权重设置咨询

问题解答

一、关于数据集样本数显示异常的推测验证

你的推测不正确,TensorFlow本身不会默认假设数据集是均衡的,日志中样本数与实际不符的问题,大概率出在数据加载、标签解析或统计逻辑环节:

  • 日志显示训练集Approved和Rejected均为36个,但实际Approved有1456个,说明统计代码只识别到了36个Approved样本,可能是标签匹配错误(比如标签字符串大小写、格式不一致)、数据过滤逻辑误删了大量样本,或是FastFile模式下数据集读取的路径/格式配置问题。
  • 数据集基数(1492、328)和实际样本总数一致(1456+36=1492,319+9=328),说明所有样本都被加载了,但标签统计环节出现错误,和TensorFlow的默认均衡假设无关。

二、为SageMaker TensorFlow Estimator设置类别权重

SageMaker的Estimator本身没有直接配置类别权重的参数,需要在**训练脚本(transfer_learning.py)**中完成权重计算并传入模型训练步骤,具体操作如下:

1. 计算类别权重

根据样本数量手动计算,核心逻辑是:类别权重 = 总样本数 / (类别数量 × 该类样本数),示例代码:

# 直接使用已知的训练集样本数量
total_train_samples = 1492
class_counts = {'Approved': 1456, 'Rejected': 36}
class_weights = {}
for cls, count in class_counts.items():
    class_weights[cls] = total_train_samples / (2 * count)  # 二分类场景,类别数为2
# 计算结果:Approved≈0.51,Rejected≈20.72

2. 在模型训练时传入类别权重

在model.fit()中添加class_weight参数,示例:

# 假设train_dataset、val_dataset已完成加载和预处理
model.fit(
    train_dataset,
    validation_data=val_dataset,
    epochs=hyperparameters['epochs'],
    class_weight=class_weights
)

注意事项

  • 如果你的标签是数字格式(比如0对应Approved,1对应Rejected),需要将class_weights的键改为对应数字。
  • 确保在训练脚本中使用实际的样本类别数量,避免用日志中错误的统计值计算权重。

内容的提问来源于stack exchange,提问作者Sam Baker

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 07:52:20