多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

类不平衡表格数据增强:CTGAN与SMOTE混合过采样策略实战

类不平衡表格数据增强:CTGAN与SMOTE混合过采样策略实战 简介这是一份面向机器学习与数据科学从业者的表格数据合成与质量评估综合项目重点聚焦生成对抗网络CTGAN、TabDiff与经典过采样方法SMOTE、ADA的结合适用于不平衡数据处理、数据增强、隐私保护下的数据共享等场景。资源包含493个文件、约76MB以Python源码41个py和CSV数据集164个csv为核心配套PNG可视化图表、JSON/YAML配置、NPY模型权重及说明文档从模型训练到结果分析均有覆盖。已有65人学习参考。通过该资源可系统掌握CTGAN和TabDiff的表格数据生成逻辑、SMOTE与ADA的过采样实现以及合成数据的统计属性、预测性能与分布一致性评估方法源码、实验数据和图表便于直接复现实验也可作为进一步研究数据质量评估和数据增强的起点。适合希望在真实项目中落地GAN与过采样技术的中高级开发者。1. 表格类不平衡比你想的严重GAN 过采样与经典过采样不是二选一做信贷风险、故障诊断或者医疗预测的应该都有这个经验正样本少得可怜负样本排山倒海。生成对抗网络在表格数据合成上火了这几年CTGAN 这类模型能通过对抗训练造出和原始分布相近的新样本但 SMOTE 这种经典过采样方法在低维小样本上依然有极强性价比。这个项目最让我感兴趣的是它把生成对抗网络CTGAN、TabDiff和经典过采样SMOTE、ADASYN放在同一个流程里对比还配了质量评估。它不是让你二选一而是给你一个混合策略先用经典过采样撑起基线再用深度生成模型逼近列间关系。适合正在做表格数据增广、类不平衡建模、或者被数据隐私问题卡住的人。我会把这套流程的手感、参数和坑逐一写出来。2. 方法拆解与选型CTGAN、TabDiff、SMOTE 各家管哪一段拿到这个项目包先别急着跑训练。里面有两条路线一条是以 CTGAN 为首的生成对抗网络一条是 SMOTE/ADASYN 这类插值型过采样。两条路线的适用场景差得很远硬放在一起比较没有意义。项目把它们并列本质是想让你看清楚什么时候深度生成模型值得付出训练成本什么时候经典的 K 近邻插值已经够用。2.1 CTGAN条件生成与模式归一化解决的是真实表分布CTGAN 是专门为表格数据改造的生成对抗网络。图像 GAN 处理的是连续的像素矩阵而表格数据是混合的有些列是连续浮点数且呈多峰分布有些列是离散类别还有明显的列间关联。CTGAN 做了几件很关键的事。第一是连续列的建模方式。它没有直接对原始数值做 min-max 归一化而是对每个连续列用高斯混合模型估计分布再根据这个分布做变换。说白了一个连续列如果有三个峰值统一归一化会把这些峰值的信息压扁高斯混合能把每个峰都保留下来。这个处理方式我最早在项目中看到时并不觉得特别直到自己跑了一版对比才发现同一份工资收入数据用普通归一化的 CTGAN 生成结果明显缺少两头的长尾。第二是离散列的条件生成机制。表格数据里的离散列往往很不均衡比如欺诈标签只有 5% 是正例。CTGAN 在采样训练批时会额外构造条件向量保证每个类别都能以一定概率被抽到而不是被多数类淹没。生成对抗网络的损失函数里判别器负责区分真实样本和合成样本生成器负责骗过判别器CTGAN 在这个基础上加了梯度惩罚项用 WGAN-GP 的方式稳定训练避免模式坍塌。第三是我个人觉得最容易被低估的一点ctgan 库自带 log_frequency 参数。它让离散列的条件概率按类别频次取对数这样可以防止低频率类别在生成时被彻底忽略。对类不平衡场景来说这比单纯调网络层数更实用。2.2 TabDiff 与扩散模型稳定但更贵的另一条路项目里出现的 TabDiff 在现阶段还不像 CTGAN 那么普及它代表的是扩散模型进入表格数据的方向。和生成对抗网络不一样扩散模型不是让两个网络互相博弈而是对原始数据逐步加噪声再学习如何一步步去噪还原。这个机制的优点是训练过程稳定很多生成质量不容易出现判别器压过生成器导致的塌缩。但代价也很直接训练成本通常比 GAN 高出一截尤其是表格数据这种维度不算大但列类型复杂的场景调试时间往往翻倍。我在实际项目里对 TabDiff 的态度是先看数据规模如果原始样本只有几千条扩散模型很难学出足够丰富的条件分布如果数据量到十万行以上而且列间关系比较复杂这时候才值得把它拉进来和 CTGAN 做对比。它更适合作为备选方案而不是默认首选项。2.3 SMOTE 与 ADASYN经典过采样在合成前先把基线垫高SMOTE 的思路简单粗暴在少数类样本之间沿特征连线方向合成新样本。它不需要训练生成模型先对少数类样本找 K 个近邻再在样本和近邻之间随机插值一个新样本就出来了。ADASYN 是 SMOTE 的变体它的区别在于会根据少数类样本周围多数类的密度动态决定生成数量。周围多数类越多说明这个样本越难学ADASYN 就给这个样本附近多生成一些合成样本。这两种方法在低维、特征稠密的数据上效果立竿见影而且完全可解释。你随时能指出哪个样本是插值来的、插值来自哪两个邻居这在业务审计时有很大优势。但当特征维度很高或者原始少数类样本特别稀疏时SMOTE 合成的样本会大量落在特征空间的空白区域反而引入噪声。项目里把两种路线放在一起我通常会先跑 SMOTE 做基线再跑 CTGAN 看能不能超越它如果深度生成模型连插值基线都打不过那大概率是数据量太少或参数没调到位。方法训练成本可解释性典型适用场景主要风险SMOTE极低高低维稠密小样本高维稀疏时生成噪声ADASYN极低高边界样本被多数类包围对噪声敏感CTGAN较高中混合类型、关系复杂训练不稳定TabDiff很高低大规模复杂表格训练周期长3. 动手复现从原始 CSV 到合成样本的完整操作链一条完整的复现路径大致是准备环境、清洗数据、训练深度生成模型、跑经典过采样、最后做评估。我按项目实际能跑通的方式写一遍命令和参数都是可照抄的。3.1 环境准备与解包先分清两个引擎解压项目包之后建议先建虚拟环境避免把本机的 Python 环境搞乱。代码里主要依赖是 ctgan、imbalanced-learn、pandas、scikit-learn。如果没有 GPUCTGAN 也可以跑 CPU但训练时间会明显拉长样本量大时建议找个带 CUDA 的环境。python -m venv venv source venv/bin/activate # Windows 下用 venv\Scripts\activate pip install ctgan imbalanced-learn pandas scikit-learn逻辑说明这条命令先创建虚拟环境并激活然后安装两套方法需要的基础库。ctgan 里自带 CTGANSynthesizerimbalanced-learn 提供 SMOTE 和 ADASYNpandas 负责数据读写。参数上不需要额外指定镜像源如果网络环境慢可以临时加-i指向国内源但这不是必需项。3.2 数据清洗与类型声明CTGAN 出错往往在这里表格数据建模的第一步不是直接开训而是把每一列的类型定清楚。CTGAN 要求你显式传入离散列名剩下默认按连续列处理。如果连续列里混入了缺失值或者字符串轻则训练报错重则生成结果全是 NaN。我一般先把明显是数值的列转成 float把类别列转成 category再做一次缺失值兜底。import pandas as pd import numpy as np df pd.read_csv(raw.csv) for col in [amount, age, income, duration]: df[col] pd.to_numeric(df[col], errorscoerce) df[label] df[label].astype(category) df df.dropna(subset[label]) numeric_cols df.select_dtypes(include[np.number]).columns df[numeric_cols] df[numeric_cols].fillna(df[numeric_cols].median())逻辑说明这里最关键的是把类别列显式设置为 category让 CTGAN 把它当成离散分布来建模。如果漏掉这一步模型会把类别码当成连续值去拟合生成结果会出现真实数据里不存在的“中间类别”比如性别列生成出 0.5 这种毫无意义的值。缺失值填充用中位数而不是均值是为了尽量避免分布偏向单侧长尾。3.3 训练 CTGAN 并生成样本epochs、batch_size 与 log_frequency数据清洗完成后就可以实例化 CTGAN。这个库的接口很简洁核心参数是 epochs、batch_size、discriminator_steps。epochs 太少学不到分布太多又容易过拟合到训练集。经验值先设 200 到 300观察损失变化再调整。batch_size 一般取 500 或 1000太小会导致离散条件采样不稳定太大则训练太慢。from ctgan import CTGAN ctgan CTGAN( epochs300, batch_size500, discriminator_steps1, log_frequencyTrue, verboseTrue ) ctgan.fit(df, discrete_columns[label, city, occupation]) synthetic ctgan.sample(n_rows10000)逻辑说明fit 传入原始数据和离散列列表sample 根据学到的分布生成新样本。n_rows 不是必须等于原样本量你可以按需生成 1 倍、2 倍甚至更多。log_frequency 建议保持 True它会根据离散列的真实频率调整采样权重对类不平衡问题非常关键。discriminator_steps 表示每训练一步生成器之前先训练几步判别器默认 1 就够用如果损失震荡明显可以改成 5 试试代价是训练时间变长。3.4 用 SMOTE 和 ADASYN 生成等价增强集sampling_strategy 怎么设经典过采样跑起来比 CTGAN 快得多核心参数是 sampling_strategy 和近邻数量。sampling_strategy 这个参数值得仔细说如果是小数表示少数类与多数类数量之比如果是整数表示少数类最终要达到的绝对样本数。不要一上来就设成 1 对 1那样生成的样本量可能过大也会让模型对合成样本过拟合。按 0.5 到 0.8 起步比较稳。from sklearn.model_selection import train_test_split from imblearn.over_sampling import SMOTE, ADASYN X df.drop(columns[label]) y df[label].astype(int) X_train, X_test, y_train, y_test train_test_split( X, y, test_size0.3, stratifyy, random_state42 ) smote SMOTE(sampling_strategy0.8, random_state42, k_neighbors5) X_smote, y_smote smote.fit_resample(X_train, y_train) adasyn ADASYN(sampling_strategy0.8, random_state42, n_neighbors5) X_adasyn, y_adasyn adasyn.fit_resample(X_train, y_train)逻辑说明先把原始数据分成训练集和测试集并且用 stratify 保证测试集和训练集的类别比例一致。SMOTE 和 ADASYN 都只能作用在训练集上绝对不要对测试集做任何过采样否则评估结果会虚高到完全没有参考价值。k_neighbors 和 n_neighbors 在代码里意思相同代表找几个近邻来做插值样本量少时建议降到 3样本量大时 5 到 7 都行。随机种子固定下来这样别人复现时能得到一致结果。4. 集中避坑合成数据翻车的几个常见原因合成数据这条链路翻车点往往不在模型本身而在数据处理和参数设定。我按自己踩过和替别人排过的坑整理几条高频问题。4.1 连续变量被整列当成离散值现象CTGAN 生成的某个连续列只有十几个固定数值明显不是真实分布。原因训练前没有正确声明离散列要么把离散列当连续列要么反过来最常见的是把整数型连续列比如“次数”“人数”默认当成了离散列。解决先用df.select_dtypes查看每列类型整数列如果实际含义是连续量就显式转成 float 再输入 CTGAN。离散列则用 category 类型声明出来两件事分开做不要图省事全交给模型自动推断。4.2 过采样比例超过 1 比 3 后下游反而变差现象把少数类过采样到和多数类一样多模型在训练集上 F1 很高但在测试集上反而比不过采样的基线还差。原因合成样本毕竟不是真实样本比例拉到 1 比 1 之后模型对少数类的决策边界会被大量合成样本推偏真实测试集里的少数类模式根本没有那么多。解决设 sampling_strategy 0.5 到 0.8保留多数类的先验优势。如果确实需要平衡优先尝试调整分类器的 class_weight而不是一味扩大合成样本量。4.3 CTGAN 损失不下降或生成 NaN判别器崩了现象verbose 输出的损失一路震荡最终 sample 出来的 data frame 里有大量 NaN 或者全部是同一个值。原因多半是连续列里有极端离群点或者判别器梯度惩罚参数和 batch_size 不匹配导致 WGAN-GP 训练不稳定。解决先对连续列做分位数裁剪把 99% 分位以外的值拉回边界。再把 batch_size 调小到 256 或 128观察损失是否变得平滑。如果仍然崩把 epochs 降到 100先跑通流程再慢慢加。4.4 SMOTE 在稀疏高维数据上生成重复样本现象合成样本里有大量完全相同的重复行去重后发现有效样本非常少。原因特征维度很高时少数类样本之间距离普遍很大K 近邻找出来的邻居实际上很远插值生成的样本散布在高维空间里很容易跟已有样本重叠。解决先做特征选择或 PCA 降维再跑 SMOTE。或者改用 SMOTE-NC 这类能感知类别特征的变体。另一种思路是直接用 CTGAN 走深度生成路线不再用插值硬刚。4.5 只看单列分布忘掉列间相关性现象合成数据每一列单独看都很像原始数据但两列交叉后明显失真比如年龄和收入的对齐关系消失。原因评估时只画了单变量分布图没有从业务角度验证列与列之间的业务逻辑。CTGAN 虽然能学习列间关联但样本量小时学得不一定完整。解决在评估阶段加上相关性矩阵差异检查重点看业务强相关的几对列比如年龄与工作年限、交易金额与账户余额。发现相关性偏掉优先增加训练轮数再考虑换 TabDiff 这类扩散模型。5. 质量评估用分布距离和下游任务做最后把关合成数据好不好不能只靠肉眼。项目里最有价值的部分不是生成模型本身而是它给出的质量评估思路这比单纯生成几万行假数据更重要。我一般会分三步走。第一步是分布距离检查。对每个连续列比较原始表和合成表的均值、方差、分位数对每个离散列比较类别占比。更深一层是算相关系数矩阵的差值重点关注业务上已知强相关的列对。这一步可以用一个很小的脚本完成比如用 pandas 计算两组数据对应列的相关系数再求绝对值差。第二步是下游任务验证。把原始训练集、SMOTE 增强集、CTGAN 增强集分别送去训练同一个分类器固定随机种子然后看它们在同一个测试集上的 F1 和 AUC。这个方法最直接如果分类器在合成数据增强下没有提升那不管分布图多好看都不能上线。关键点是测试集必须保持原样不能混入任何合成样本。第三步是记录生成配置和随机种子。训练一个 CTGAN 动辄几分钟到几十分钟跑完后把 epochs、batch_size、随机种子、数据版本一起存下来。这样能让结果可复现也能避免同一份代码在不同时间跑出完全不同的样本。很多人忽略这一步等模型上线后想往回查某个版本就彻底没法查了。从那以后我每次跑合成数据实验都会强制走一遍先算原始表和生成表的均值方差与相关矩阵差值再用固定分类器做交叉验证所有模型配置和随机种子入库保存。折腾一圈下来最深的体会是生成对抗网络和经典过采样不是替代关系而是互为兜底——真到业务上线时能解释、能评估、能复现的合成数据才敢放心用。希望帮到你。本文还有配套的精品资源点击获取
返回列表