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

如何对含数组型目标列的数据集进行重采样以实现均匀分布

如何对含数组型目标列的数据集进行重采样以实现均匀分布

嘿,看起来你有个用Polars构建的数据集,target列还是数组类型的,想要通过重采样让数据分布变得均匀对吧?我来给你一步步拆解怎么搞定这个事儿~

首先得明确:数组型的目标列没法直接用来分组重采样,咱们得先把数组里的关键信息提取出来,转换成能用来划分类别的标签,之后再进行重采样操作。从你给的例子来看,target数组的前两个元素都是1.0,只有第三个元素在变化,那咱们就以这个元素为核心来处理。

步骤1:提取数组中的关键特征作为分组依据

先把target数组的第三个元素(索引为2)提取出来,生成一个新列target_key:

import polars as pl

# 你的原始数据集
df = pl.DataFrame(
    {
        "target": [
            [1.0, 1.0, 0.0],
            [1.0, 1.0, 0.1],
            [1.0, 1.0, 0.2],
            [1.0, 1.0, 0.8],
            [1.0, 1.0, 0.9],
            [1.0, 1.0, 1.0],
        ],
        "feature": ["a", "b", "c", "d", "e", "f"],
    },
    schema={"target": pl.List(pl.Float64), "feature": pl.Utf8}
)

# 提取target数组的第三个元素作为分组键
df = df.with_columns(
    pl.col("target").list.get(2).alias("target_key")
)

步骤2:划分类别区间

接下来把target_key转换成明确的类别标签,方便后续分组。比如咱们把它分成「低区间(0-0.5)」和「高区间(0.5-1.0)」两类:

# 划分区间并生成类别标签
df = df.with_columns(
    pl.col("target_key").cut(bins=[0.0, 0.5, 1.0], labels=["low", "high"]).alias("target_label")
)

步骤3:重采样实现均匀分布

这里分两种常见场景,你可以根据自己的数据情况选:

场景1:欠采样(减少多数类样本)

如果某一类样本数量远多于其他类,咱们可以对多数类进行随机采样,把所有类别样本数降到和少数类一致:

# 统计每个类别的样本数量
class_counts = df.group_by("target_label").len().sort("len")
min_count = class_counts["len"][0]  # 取样本最少的类别数量

# 对每个类别采样到min_count个样本
resampled_df = df.group_by("target_label").apply(
    lambda group: group.sample(n=min_count, seed=42)  # seed固定保证结果可复现
).drop(["target_key", "target_label"])

场景2:过采样(增加少数类样本)

如果某一类样本太少,咱们可以对少数类进行重复采样,让所有类别样本数和多数类一致:

# 统计每个类别的样本数量
class_counts = df.group_by("target_label").len().sort("len", descending=True)
max_count = class_counts["len"][0]  # 取样本最多的类别数量

# 对每个类别采样到max_count个样本,少数类用重复采样补充
resampled_df = df.group_by("target_label").apply(
    lambda group: group.sample(n=max_count, seed=42, replace=True)
).drop(["target_key", "target_label"])

步骤4:验证重采样结果

最后可以检查一下重采样后的分布是不是均匀了:

# 重新提取target_key来验证分布
resampled_df = resampled_df.with_columns(
    pl.col("target").list.get(2).alias("target_key")
)
# 统计各区间的样本数
print(resampled_df.group_by(pl.col("target_key").cut(bins=[0.0,0.5,1.0])).len())

备注:内容来源于stack exchange,提问作者DJDuque

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.13 15:59:31