简介:围绕生成对抗网络与过采样技术的综合性机器学习项目包,聚焦CTGAN、TabDiff与SMOTE、ADA的联合建模,实现表格数据合成及质量评估。面向数据科学研究者、机器学习开发者,尤其适用于处理不平衡数据集、数据稀缺或隐私保护场景;项目代码、数据集与评估结果一并提供,可复现实验或二次开发。压缩包共493个文件,以164个csv数据表、41个py源码脚本和213个png可视化图表为主体,辅以json配置、npy数组结果等,整体大小75.99MB,目录结构清晰。目前已有65人学习下载。读者可获得完整的合成数据生成流程、过采样方法实现对比,以及基于随机森林、决策树、逻辑回归的SHAP特征重要性分析和多种指标下的合成数据质量验证结果,为数据增强与GAN应用研究提供可直接上手的实战参考。
1. 表格数据合成:当少数类样本不够时,GAN 和过采样哪个先上场
先给一个反直觉的结论:在多数表格数据场景里,SMOTE 这类传统过采样往往比 CTGAN 更容易帮你拿到“能用的模型”,但生成对抗网络(也就是对抗生成网络,GAN 家族)能覆盖 SMOTE 碰不到的那类需求——比如要从原始数据分布里生成全新的样本,而不是在已知样本之间插值。这个综合性项目把 CTGAN、TabDiff 和 SMOTE/ADASYN 放在同一个框架里对比,本质上是在回答一个问题:缺样本的时候,你是要做“复制变异”还是“按分布重画”。适合谁?做分类任务被不均衡数据折磨的从业者,搞数据脱敏需要替身数据集的工程师,以及想给模型做鲁棒性验证的算法团队。本文按“选型逻辑 → 建模参数 → 评估方法 → 踩坑记录 → 验证技巧”的顺序把这套流程讲透。
2. SMOTE 与 ADA 的边界:为什么插值法先赢,然后撞墙
2.1 SMOTE 的核心机制与适用前提
SMOTE(Synthetic Minority Over-sampling Technique)的做法一句话能说清:在少数类的 K 近邻之间连线,在连线上随机取点生成新样本。它假设的是“少数类样本在特征空间里分布是连续的,两点之间依然属于这个类别”。这个假设在低维、特征相关性不强的表格数据上通常是成立的,所以它几乎零成本地就能让多数场景下的 F1 值涨上一截。
我一般会在拿不定主意的时候先跑一版 SMOTE 作为基线,因为它快、稳定、可解释。代码实现也不复杂:
from imblearn.over_sampling import SMOTE from sklearn.ensemble import RandomForestClassifier from sklearn.model_selection import cross_val_score smote = SMOTE(random_state=42, k_neighbors=5) X_res, y_res = smote.fit_resample(X_train, y_train) clf = RandomForestClassifier(n_estimators=200, max_depth=10, random_state=42) scores = cross_val_score(clf, X_res, y_res, cv=5, scoring='f1_macro') print(f"SMOTE + RF F1: {scores.mean():.4f}")逻辑说明:首先用fit_resample同时完成拟合和重采样,k_neighbors=5控制生成样本时参考的近邻数量,偏小容易过拟合局部噪声,偏大则生成的样本会向其他类别方向“漂移”。交叉验证分数用于快速判断这个方案的上限——如果 SMOTE 连 0.6 都上不去,后面换 GAN 大概率也只是微调。
参数说明:random_state必须固定,否则每次跑出来的实验数据都不一样,后面做对比时你分不清是模型差异还是随机性造成的。另一个容易忽略的是,SMOTE 只能作用于数值特征,如果你的表格里有多分类或高基数类别特征,需要先用编码器处理或者改用 SMOTE-NC 变体。
2.2 ADA 与边界样本的博弈
ADASYN(Adaptive Synthetic Sampling)是 SMOTE 的改进版,核心思路是:对每个少数类样本,根据它周围多数类样本的密度决定生成数量——周围多数类越多,生成的样本就越多。它把生成火力集中在“最容易混淆”的边界区域,这也是它名字里 Adaptive 的来源。
听起来比 SMOTE 聪明,但在真实项目里 ADA 翻车的频率不低。原因是边界区域本身噪声就大,过度在边界合成样本,等于变相放大了分类器对边界噪声的敏感度,有时候 F1 涨了,精确率掉得一塌糊涂。
from imblearn.over_sampling import ADASYN ada = ADASYN(random_state=42, n_neighbors=5, sampling_strategy='auto') X_ada, y_ada = ada.fit_resample(X_train, y_train) from sklearn.metrics import precision_score, recall_score, f1_score # 训练完成后检查各类别指标分布 print(f"Precision: {precision_score(y_val, y_pred):.4f}") print(f"Recall: {recall_score(y_val, y_pred):.4f}")逻辑说明:sampling_strategy='auto'表示把所有少数类都提升到与多数类同等数量,如果你希望控制合成比例,可以传一个字典,比如{1: 2000},表示把类别 1 合成到 2000 条。代码里单独打印精确率和召回率是为了检查 ADA 是否“只追召回不计精度”。
参数说明:n_neighbors在 ADA 中不仅影响边界密度估计,还影响合成样本的生成位置。默认 5 在小数据集上经常不够稳定,建议在 3-10 之间做网格搜索。
2.3 插值法的三个硬边界
第一个硬边界是特征共线性。SMOTE 类方法生成的样本沿近邻连线分布,当特征之间高度相关时,新样本可能落在原始分布的子空间外。第二个硬边界是类别特征的失真——对独热编码后的类别特征做插值,结果可能是“半男半女”这种现实中不存在的样本。第三个硬边界是数据量级太小,比如少数类只有几十条,K 近邻本身就不可靠,生成的样本就是在噪声之间反复插值。
这就是为什么项目标题里要把 GAN 和过采样放在一起。SMOTE 负责快速止血,GAN 负责处理插值法解决不了的问题——尤其是需要生成“看起来真实但不重复”的新样本时。
3. 从 SMOTE 到 CTGAN:当插值法失效时,对抗生成网络怎么接手表格数据
3.1 表格数据生成为什么不能用普通 GAN
很多从图像 GAN 转过来的人第一次用 DCGAN 生成表格数据,生成的样本要么模式坍塌成几条重复数据,要么在离散特征上给出完全不合理的组合。原因是图像 GAN 假设数据是连续像素分布,而表格数据是混合类型的——连续列有偏态分布,离散列有稀疏类别。普通 GAN 的生成器输出是一个连续向量,没法直接表达“这个类别出现的概率分布”。
所以 CTGAN 引入了一个关键设计:对离散列做 one-hot 编码后,用 Gumbel softmax 让生成器能输出离散分布;对连续列,先用高斯混合模型(VGM)估计每个连续列的分布,再做条件归一化。这一步相当于把表格数据“翻译”成 GAN 能理解的格式。
3.2 CTGAN 的条件生成与训练策略
CTGAN 的另一个核心机制是条件生成器。它在每个训练 step 随机选一个离散列的一个类别作为条件,让生成器在这个条件下生成样本,并用一个判别器判断样本是否同时满足“真实性”和“条件匹配性”。这个设计直接解决了类别不平衡带来的模式坍塌问题——原来少数类在整体数据里占比低,GAN 可能直接忽略它们,现在强制生成器每次都要生成指定类别的样本。
from ctgan import CTGAN ctgan = CTGAN( epochs=300, batch_size=500, log_frequency=True, verbose=True, generator_dim=(256, 256), discriminator_dim=(256, 256), generator_lr=2e-4, discriminator_lr=2e-4, discriminator_steps=1, ) ctgan.fit(train_data, discrete_columns=['education', 'marital_status', 'loan_status']) samples = ctgan.sample(5000) print(samples.head())逻辑说明:discrete_columns必须把数据里所有类别特征列名都传进去,漏掉一个都会让模型把这个离散列当成连续列处理,生成出来的组合会非常怪异。log_frequency=True是最重要的一个参数,它让条件生成时按照频率的对数来采样,避免高频类别主导训练。
参数说明:generator_dim=(256, 256)和discriminator_dim=(256, 256)是双隐藏层的维度配置,表格数据一般不需要像图像那样堆到 512 或 1024 的宽度,因为特征维度通常几十维,太宽的网络反而容易让判别器过早收敛。discriminator_steps=1表示每训练生成器一次,判别器训练一次;如果训练不稳定可以把判别器步数提高到 5,但也会增加模式坍塌风险。epochs=300在数据集几万行时够用,如果数据量上百万,可以适当降低到 100-200。
3.3 TabDiff 为什么能成为 CTGAN 的互补方案
TabDiff 走的是扩散模型路线,核心思想是对原始数据逐步加高斯噪声直到完全变成随机噪声,然后学习一个反向过程,从纯噪声里逐步去噪还原出数据。它的优势在于生成多样性比 GAN 强,因为扩散模型的训练目标是拟合完整数据分布,而不是像 GAN 那样在生成器和判别器的博弈中找均衡点。
但 TabDiff 的问题是采样速度慢,生成一批样本需要跑几十步去噪过程,不像 CTGAN 的生成器是前向一遍就能出结果。这个项目把 TabDiff 也纳进来,合理的用法是把它当成 CTRGAN 的效果上限参考——如果 TabDiff 生成的数据在质量评估指标上也不比 CTGAN 好到哪里去,那说明任务难度主要在下游分类器而不是生成模型。
4. 跑通完整流程:从数据预处理到质量评估的最小复现路径
4.1 实验设计:用一个分类任务统一三种方法
为了让对比公平,我会定义一个统一的评估流程:原始不平衡数据 → 分别用 SMOTE、ADASYN、CTGAN、TabDiff 生成或合成训练集 → 在同一个分类器(比如 XGBoost)上训练 → 在同一个固定的测试集上评估 F1、AUC、精确率。这里必须保证测试集是原始数据,不能掺任何合成样本,否则评估结果是虚高的。
4.2 数据预处理的标准动作
第一步是划分训练集和测试集,并且在划分之后再做任何合成操作。如果先做合成再划分,测试集里可能有合成样本的“影子”,导致评估失真。第二步是统一做特征对齐,数值列做标准化,离散列做标签编码或独热编码。CTGAN 内部有自己的离散列处理逻辑,所以给它喂原始编码即可,但 SMOTE 需要单独处理。
import pandas as pd from sklearn.model_selection import train_test_split from sklearn.preprocessing import StandardScaler df = pd.read_csv('credit_risk.csv') X = df.drop(columns='default') y = df['default'] X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.2, random_state=42, stratify=y ) num_cols = X_train.select_dtypes(include=['float64', 'int64']).columns scaler = StandardScaler() X_train[num_cols] = scaler.fit_transform(X_train[num_cols]) X_test[num_cols] = scaler.transform(X_test[num_cols])逻辑说明:stratify=y保证训练集和测试集中的正负样本比例一致,这是不均衡分类的基础操作,不设置的话划分出来的测试集可能极度失衡,导致评估指标完全失真。fit_transform和transform分开用是防止测试集的统计量泄漏进训练过程。
4.3 四种方案统一评估的脚本框架
from xgboost import XGBClassifier from sklearn.metrics import roc_auc_score, f1_score def evaluate(y_true, y_pred_proba, threshold=0.5): y_pred = (y_pred_proba >= threshold).astype(int) return { 'auc': roc_auc_score(y_true, y_pred_proba), 'f1': f1_score(y_true, y_pred), } results = {} # SMOTE 方案 smote = SMOTE(random_state=42) X_sm, y_sm = smote.fit_resample(X_train, y_train) clf = XGBClassifier(n_estimators=200, max_depth=6, learning_rate=0.1, random_state=42) clf.fit(X_sm, y_sm) results['smote'] = evaluate(y_test, clf.predict_proba(X_test)[:, 1]) # CTGAN 方案 ctgan = CTGAN(epochs=300, batch_size=500, verbose=False) ctgan.fit(X_train, discrete_columns=['marital_status']) X_ctgan = ctgan.sample(len(X_train)) # 注意:这里只取少数类样本,与原始多数类样本拼接 # 确保训练集包含全部原始多数类 + 生成的少数类 X_ctgan_combined = pd.concat([X_train[y_train == 0], X_ctgan[y_train == 1]], axis=0) # 重训模型,评估同上逻辑说明:CTGAN 合成后用条件生成的方式生成少数类样本(通过 CTGAN 的采样接口无法直接控制标签,常见做法是先按标签分组,每个组单独训练一个 CTGAN,或者在这里用原始多数类拼接合成少数类的方式)。这个拼接逻辑是整个流程中最容易被忽略的一步——直接拿 CTGAN 生成全部训练集,导致原始数据的信息丢失,效果反而不如 SMOTE。
参数说明:threshold=0.5是默认分类阈值,在不均衡数据中可以根据验证集调整到 0.3 或 0.6,但要确保所有方案用同一个阈值对比才有意义。
4.4 质量评估的三个维度
第一维是分布距离指标,常见的是 Wasserstein 距离或 KL 散度,用来衡量合成数据的整体分布与原始数据有多接近。第二维是特征相关性保留——计算原始数据和合成数据的相关系数矩阵,看差异有多大。第三维是下游任务效用,就是前面统一评估的 F1/AUC,这个是最有说服力的指标,因为最终目的是让模型好用,而不是让数据看起来像。
| 评估维度 | 代表指标 | 说明 |
|---|---|---|
| 分布距离 | Wasserstein Distance | 值越小越好,但注意高维下计算不稳定 |
| 相关性保留 | Correlation Difference | 比较各特征对的相关系数绝对值差 |
| 下游效用 | F1 / AUC | 最终决策依据 |
5. 避坑与排查:表格数据合成最容易翻车的五个细节
5.1 现象:CTGAN 训练完成后生成的数据全是重复行
原因分析:这是模式坍塌的典型表现。可能诱因是判别器能力过强导致生成器梯度消失,或者batch_size相对于数据量过小,让生成器只学会了一条容易被判别器放行的样本形态。解决思路:降低判别器学习率(比如从 2e-4 降到 1e-4),调大batch_size,同时把epochs减半观察训练前期的生成多样性变化。另一个从工程侧有效的办法是把连续列做分位变换,比如用QuantileTransformer映射到正态分布,降低原始偏态分布给生成器带来的拟合压力。
5.2 现象:SMOTE 生成的数据在离散列上出现“不可能的组合”
原因分析:直接对标签编码后的离散列做 K 近邻插值,两个近邻在线段中点取值,编码值四舍五入后可能得到一个在原始数据里根本不存在的类别。解决思路:改用 SMOTE-NC,并且在调用时传入categorical_features参数,让算法对离散列采用众数投票而不是插值。这一步不需要改其他流程,替换效果立竿见影。
5.3 现象:CTGAN 评估时 AUC 比原始数据还低
原因分析:这不是生成模型的问题,很大概率是评估流程出了问题。最常见的是测试集里混入了合成数据,或者分类器用了默认参数且没有做交叉验证。排查顺序:先检查测试集的y分布与原始测试集是否一致,再确认训练集里原始数据与合成数据的比例是否合理。CTGAN 生成的少数类样本量一般建议与原始多数类样本量 1:1 拼接,如果少数类占比过低,会导致模型对少数类的学习不够充分。
5.4 现象:TabDiff 生成一次要几个小时,无法接受
原因分析:扩散模型的反向采样步骤默认设置偏高。TabDiff 的采样步数(通常是 50-1000 步)直接决定耗时,但表格数据的特征维度低,50 步已经能获得不错的效果。把采样步数从默认的 1000 降到 100,生成时间能缩短 90%。另一个优化手段是对连续特征做降维,比如先用 PCA 压到 20 维再训练,生成后再用逆变换回到原特征空间,模型训练和采样都快得多。
5.5 现象:合成数据在分布距离指标上很好,但下游 F1 没提升
原因分析:分布距离近不等于分类信息保留得好。GAN 可能学到了边缘分布,但特征之间的联合交互没学全,而 XGBoost 这类模型恰恰依赖特征交叉来分类。解决思路:在质量评估中加入“特征对相关性差异”指标,或者做一个快速验证——用合成数据训练逻辑回归,如果逻辑回归在测试集上的表现也接近原始数据训练的模型,说明线性可分的联合信息也被保留了。
6. 生成质量自检技巧:用异常检测器给合成数据做压力测试
最后一个分享一个百试不爽的验证技巧:训练一个异常检测模型(比如 Isolation Forest),把原始数据标记为正常,把合成数据标记为异常,看检测器能不能把两者分开。
具体做法是这样。用原始训练集拟合 Isolation Forest,然后对合成数据计算异常分数。如果合成数据与原始数据同分布,异常分数不会显著高于原始数据的异常分数;如果生成模型学偏了,异常分数会出现明显的双峰分布——这个现象比任何分布距离指标都直观。我习惯把这个检测放在 CTGAN 训练中期的检查点执行,每 50 个 epoch 生成一批样本做一次测试,如果某个检查点开始异常分数飙升,说明训练已经开始过拟合到训练集噪声了,应该早停。
另一个技巧是把合成数据当成增强集与原始数据混合后训练,然后只在原始数据的留出集上测试。分别记录合成数据的增减对 F1 的影响曲线,如果增加合成数据反而让 F1 下降,说明数据质量有问题,不是量的问题。这个测试也能帮助判断到底该用 SMOTE 还是 CTGAN——如果 SMOTE 增强后 F1 已经接近上限,CTGAN 的边际收益很低,不值得花时间调参。
这个方向我最深的感受是:表格数据合成不是“模型越复杂越好”,而是“够用且可验证才是真好”。CTGAN 和 TabDiff 有它们不可替代的场景,但先用 SMOTE 跑通基线、再用 GAN 补盲区、最后用异常检测做压力测试,这条路几乎不会出错。希望这套流程能帮你少走我当年走过的弯路。
本文还有配套的精品资源,点击获取