如何使用TensorFlow Datasets(tfds)获取鸢尾花数据集的特征名称?
解决方法
你可以通过提取数据集元数据中的特征名称,再拆分DataFrame中的ndarray列,来实现具体特征名称与数值的对应:
1. 获取特征名称
iris数据集的特征名称存储在ds_info.features['features']的元数据中,执行以下代码即可获取:
feature_names = ds_info.features['features'].metadata['feature_names'] # 输出结果:['sepal_length', 'sepal_width', 'petal_length', 'petal_width']
2. 拆分DataFrame的特征列
拿到名称后,将包含ndarray的features列拆分为独立的特征列,替换原列:
import pandas as pd # 生成初始DataFrame df = tfds.as_dataframe(ds.take(10), ds_info) # 将ndarray列拆分为对应名称的多列 df[feature_names] = pd.DataFrame(df['features'].tolist(), index=df.index) # 删除原有的features列 df = df.drop('features', axis=1) # 查看处理后的DataFrame print(df)
替代方案:加载时直接格式化特征
也可以在数据集加载阶段就将张量格式的特征拆分为独立字段,再转成DataFrame:
def split_features(example): # 将features张量按索引拆分到对应字段 example['sepal_length'] = example['features'][0] example['sepal_width'] = example['features'][1] example['petal_length'] = example['features'][2] example['petal_width'] = example['features'][3] del example['features'] return example # 应用格式化函数 ds_formatted = ds.map(split_features) # 转换为带具体特征名的DataFrame df = tfds.as_dataframe(ds_formatted.take(10), ds_info) print(df)
内容的提问来源于stack exchange,提问作者Cyrille
相关产品推荐
相关产品推荐

