使用CTGAN库生成样本时遇ValueError的解决方法咨询
CTGAN生成模拟数据时触发ValueError问题排查与解决
问题背景
在Colab笔记本中使用CTGAN库处理含一个分类特征的表格数据集,模型训练无报错,但生成模拟数据时出现ValueError。
可复现代码
import pandas as pd import numpy as np import seaborn as sns from ctgan import CTGAN iris = sns.load_dataset('iris') iris.head() from sklearn import preprocessing le = preprocessing.LabelEncoder() le.fit(iris['species'].unique()) iris['species'] = pd.DataFrame(le.transform(iris['species'])) data = iris.copy() ctgan_model = CTGAN(epochs=2,batch_size=50,verbose = True) ctgan_model.fit(data) n_ctgan_generated_data = 2000 synthetic_data = ctgan.sample(n_ctgan_generated_data)
完整错误信息
--------------------------------------------------------------------------- ValueError Traceback (most recent call last) <ipython-input-17-199b6dc04389> in <module> 1 n_ctgan_generated_data = 2000 ----> 2 synthetic_data = ctgan.sample(n_ctgan_generated_data) 6 frames /usr/local/lib/python3.8/dist-packages/ctgan/synthesizers/base.py in wrapper(self, *args, **kwargs) 48 def wrapper(self, *args, **kwargs): 49 if self.random_states is None: ---> 50 return function(self, *args, **kwargs) 51 52 else: /usr/local/lib/python3.8/dist-packages/ctgan/synthesizers/ctgan.py in sample(self, n, condition_column, condition_value) 475 data = data[:n] 476 ---> 477 return self._transformer.inverse_transform(data) 478 479 def set_device(self, device): /usr/local/lib/python3.8/dist-packages/ctgan/data_transformer.py in inverse_transform(self, data, sigmas) 211 column_data = data[:, st:st + dim] 212 if column_transform_info.column_type == 'continuous': ---> 213 recovered_column_data = self._inverse_transform_continuous( 214 column_transform_info, column_data, sigmas, st) 215 else: /usr/local/lib/python3.8/dist-packages/ctgan/data_transformer.py in _inverse_transform_continuous(self, column_transform_info, column_data, sigmas, st) 185 def _inverse_transform_continuous(self, column_transform_info, column_data, sigmas, st): 186 gm = column_transform_info.transform ---> 187 data = pd.DataFrame(column_data[:, :2], columns=list(gm.get_output_sdtypes())) 188 data.iloc[:, 1] = np.argmax(column_data[:, 1:], axis=1) 189 if sigmas is not None: /usr/local/lib/python3.8/dist-packages/pandas/core/frame.py in __init__(self, data, index, columns, dtype, copy) 670 ) 671 else: ---> 672 mgr = ndarray_to_mgr( 673 data, 674 index, /usr/local/lib/python3.8/dist-packages/pandas/core/internals/construction.py in ndarray_to_mgr(values, index, columns, dtype, copy, typ) 322 ) 323 ---> 324 _check_values_indices_shape_match(values, index, columns) 325 326 if typ == "array": /usr/local/lib/python3.8/dist-packages/pandas/core/internals/construction.py in _check_values_indices_shape_match(values, index, columns) 391 passed = values.shape 392 implied = (len(index), len(columns)-1) ---> 393 raise ValueError(f"Shape of passed values is {passed}, indices imply {implied}") 394 395 ValueError: Shape of passed values is (2000, 2), indices imply (2000, 3)
问题分析与解决
错误根源
这个错误不是CTGAN库本身的问题,而是因为你手动用LabelEncoder将分类特征species转换为整数类型后,CTGAN默认将该列识别为连续特征,但连续特征的逆变换逻辑需要匹配特定维度,最终导致维度不匹配报错。
解决方案(无需修改源码)
CTGAN内置了分类特征的处理逻辑,不需要手动进行LabelEncoder编码,只需在训练时明确指定categorical_features参数即可:
import pandas as pd import seaborn as sns from ctgan import CTGAN iris = sns.load_dataset('iris') data = iris.copy() # 明确告知CTGAN哪些列是分类特征 ctgan_model = CTGAN(epochs=2, batch_size=50, verbose=True) ctgan_model.fit(data, categorical_features=['species']) n_ctgan_generated_data = 2000 synthetic_data = ctgan_model.sample(n_ctgan_generated_data)
补充说明
如果坚持要手动编码分类特征,需将编码后的列转换为字符串类型,让CTGAN识别为分类特征,但这种方式冗余且容易出错,更推荐使用上述官方推荐的方法。
内容的提问来源于stack exchange,提问作者Arav
相关产品推荐
相关产品推荐

