简介:这是一份面向机器学习与数据科学从业者的表格数据合成与质量评估综合项目,重点聚焦生成对抗网络(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。如果没有 GPU,CTGAN 也可以跑 CPU,但训练时间会明显拉长,样本量大时建议找个带 CUDA 的环境。
python -m venv venv source venv/bin/activate # Windows 下用 venv\Scripts\activate pip install ctgan imbalanced-learn pandas scikit-learn逻辑说明:这条命令先创建虚拟环境并激活,然后安装两套方法需要的基础库。ctgan 里自带 CTGANSynthesizer,imbalanced-learn 提供 SMOTE 和 ADASYN,pandas 负责数据读写。参数上不需要额外指定镜像源,如果网络环境慢可以临时加-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], errors="coerce") 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( epochs=300, batch_size=500, discriminator_steps=1, log_frequency=True, verbose=True ) ctgan.fit(df, discrete_columns=["label", "city", "occupation"]) synthetic = ctgan.sample(n_rows=10000)逻辑说明: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_size=0.3, stratify=y, random_state=42 ) smote = SMOTE(sampling_strategy=0.8, random_state=42, k_neighbors=5) X_smote, y_smote = smote.fit_resample(X_train, y_train) adasyn = ADASYN(sampling_strategy=0.8, random_state=42, n_neighbors=5) 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、随机种子、数据版本一起存下来。这样能让结果可复现,也能避免同一份代码在不同时间跑出完全不同的样本。很多人忽略这一步,等模型上线后想往回查某个版本就彻底没法查了。
从那以后,我每次跑合成数据实验都会强制走一遍:先算原始表和生成表的均值方差与相关矩阵差值,再用固定分类器做交叉验证,所有模型配置和随机种子入库保存。折腾一圈下来最深的体会是,生成对抗网络和经典过采样不是替代关系,而是互为兜底——真到业务上线时,能解释、能评估、能复现的合成数据才敢放心用。希望帮到你。
本文还有配套的精品资源,点击获取