You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

TensorFlow Dataset是否有类似pandas apply的特征工程等价实现?

问题解答

背景

TensorFlow官方推荐使用tf.data.Dataset实现输入管道,它的批处理、洗牌等操作高效,还能和Keras API无缝集成。但在自定义Python函数生成新列的特征工程任务中,tf.data.Dataset用起来繁琐甚至难以实现,和pandas灵活的apply()函数形成鲜明对比。

用户觉得tf.data.Dataset.map()是最接近apply()的方法,但它不能直接用任意Python函数,得先转成张量操作,这要求用户掌握TensorFlow原生函数(比如用tf.strings.length()替代Python的len()),操作麻烦还容易出维度或类型错误。而用户尝试的tf.py_function也没像预期那样轻松转换Python代码。

问题

TensorFlow的tf.data是不是还没成熟到能像pandas apply()一样处理特征工程?还是用户存在理解误区?

最小可复现示例

用户用以下代码对比pandas DataFrame和TensorFlow Dataset的特征工程效果,目标是生成两个新特征:1. 为字符串特征my_string添加后缀;2. 统计my_string中特定字母的出现次数。第一个操作两者都能轻松实现,但第二个在TensorFlow里很难完成。

from collections import Counter
import numpy as np
import pandas as pd
import tensorflow as tf

# 创建pandas DataFrame
df = pd.DataFrame(
    {'index': list(range(5)),
     'my_string': ['Alondra', 'Spaanbiuk', 'Ibinth', 'Liefelle', 'Yoanda'], 
     'some_other_column': np.random.rand(5),
     }).set_index('index')
print('原始pandas DataFrame:')
print(df, '\n')

# 创建TensorFlow数据集,并定义函数将其转为pandas DataFrame查看
ds = tf.data.Dataset.from_tensors(df.to_dict(orient='list'))
def view_ds(ds):
    data = pd.concat([pd.DataFrame(k) for k in ds.take(5)], axis=0)
    # 将字节字符串转为Python字符串
    object_cols = data.select_dtypes([object])
    data[object_cols.columns] = object_cols.stack().str.decode('utf-8').unstack()
    print('TensorFlow数据集:')
    print(data, '\n')
view_ds(ds)

#### 为`my_string`添加后缀,生成`my_string_with_suffix`
def add_a_suffix(x):
    x['my_string_with_suffix'] = x['my_string']+'.suffix'
    return x

# 应用到pandas DataFrame
print('添加后缀后的DataFrame:')
df = df.apply(add_a_suffix, axis=1)
print(df, '\n')

# 应用到TensorFlow数据集
print('添加后缀后的TensorFlow数据集:')
ds = ds.map(lambda x: add_a_suffix(x))
view_ds(ds)

#### 统计`my_string`中字母`a`的出现次数
def count_letters(x, letter='a'):
    counter = Counter(x['my_string'].lower())
    x[f'{letter}_counts'] = counter[letter]
    return x

# 应用到pandas DataFrame
print('添加字母计数后的DataFrame:')
df = df.apply(count_letters, axis=1)
print(df, '\n')
    
# 如何应用到TensorFlow数据集?
# print('添加字母计数后的TensorFlow数据集:')
# ds = ds.apply(lambda x: count_letters(x))
# view_ds(ds)

代码输出:

原始pandas DataFrame:
       my_string  some_other_column
index                              
0        Alondra           0.209685
1      Spaanbiuk           0.972315
2         Ibinth           0.933700
3       Liefelle           0.186369
4         Yoanda           0.667436 

TensorFlow数据集:
   my_string  some_other_column
0    Alondra           0.209685
1  Spaanbiuk           0.972315
2     Ibinth           0.933700
3   Liefelle           0.186369
4     Yoanda           0.667436 

添加后缀后的DataFrame:
       my_string  some_other_column my_string_with_suffix
index                                                    
0        Alondra           0.209685        Alondra.suffix
1      Spaanbiuk           0.972315      Spaanbiuk.suffix
2         Ibinth           0.933700         Ibinth.suffix
3       Liefelle           0.186369       Liefelle.suffix
4         Yoanda           0.667436         Yoanda.suffix 

添加后缀后的TensorFlow数据集:
TensorFlow数据集:
   my_string  some_other_column my_string_with_suffix
0    Alondra           0.209685        Alondra.suffix
1  Spaanbiuk           0.972315      Spaanbiuk.suffix
2     Ibinth           0.933700         Ibinth.suffix
3   Liefelle           0.186369       Liefelle.suffix
4     Yoanda           0.667436         Yoanda.suffix 

添加字母计数后的DataFrame:
       my_string  some_other_column my_string_with_suffix  a_counts
index                                                                
0        Alondra           0.209685        Alondra.suffix         2
1      Spaanbiuk           0.972315      Spaanbiuk.suffix         2
2         Ibinth           0.933700         Ibinth.suffix         0
3       Liefelle           0.186369       Liefelle.suffix         0
4         Yoanda           0.667436         Yoanda.suffix         2 

解答

这不是tf.data不成熟,而是没抓住它的设计核心——tf.data面向张量操作,追求计算图的可优化性和硬件加速,而pandas的apply()是基于Python对象的逐行处理,两者设计目标完全不同。

要在tf.data里实现类似pandasapply()的自定义逻辑,有两种可行路径:

路径一:用TensorFlow原生函数实现(优先选)

如果能把Python逻辑转成TensorFlow原生操作,就能充分发挥tf.data的性能优势。比如统计字母a的出现次数,用tf.strings模块就能搞定:

def count_letters_tf(x, letter='a'):
    # 转小写
    lower_str = tf.strings.lower(x['my_string'])
    # 拆分每个字符
    chars = tf.strings.unicode_split(lower_str, input_encoding='UTF-8')
    # 匹配目标字母
    matches = tf.equal(chars, letter)
    # 统计匹配数量
    count = tf.reduce_sum(tf.cast(matches, tf.int32))
    # 添加新列
    x[f'{letter}_counts'] = count
    return x

# 应用到数据集
print('添加字母计数后的TensorFlow数据集:')
ds = ds.map(count_letters_tf)
view_ds(ds)

路径二:用tf.py_function封装Python函数

如果必须依赖Python原生逻辑(比如用Counter这类没法转成TensorFlow操作的工具),可以用tf.py_function包装,但得明确输入输出的张量类型和形状——因为TensorFlow需要清楚计算图的结构:

def count_letters_py(x, letter='a'):
    # 把张量转成Python对象
    my_strings = x['my_string'].numpy()
    counts = []
    for s in my_strings:
        # 字节串转普通字符串
        s_str = s.decode('utf-8')
        counter = Counter(s_str.lower())
        counts.append(counter[letter])
    # 添加新列,转回张量
    x[f'{letter}_counts'] = tf.convert_to_tensor(counts, dtype=tf.int32)
    return x

# 包装函数,定义输入输出类型
def wrap_count_letters(x):
    return tf.py_function(
        func=count_letters_py,
        inp=[x],
        Tout=[tf.string, tf.float64, tf.string, tf.int32]  # 对应所有列的类型:my_string, some_other_column, my_string_with_suffix, a_counts
    )

# 把tf.py_function返回的元组转回字典,适配view_ds函数
def tuple_to_dict(t):
    return {
        'my_string': t[0],
        'some_other_column': t[1],
        'my_string_with_suffix': t[2],
        'a_counts': t[3]
    }

print('添加字母计数后的TensorFlow数据集:')
ds = ds.map(wrap_count_letters).map(tuple_to_dict)
view_ds(ds)

核心注意事项

  1. tf.data.Dataset.map()默认要求操作能构建计算图,所以直接跑Python函数会报错,要么转成TensorFlow原生操作,要么用tf.py_function封装。
  2. 用tf.py_function会丢失TensorFlow的部分优化(比如自动微分、硬件加速),能转原生操作的尽量转。
  3. 之前用tf.py_function没成功,大概率是没处理好张量与Python对象的转换,也没明确输出的类型签名。

内容的提问来源于stack exchange,提问作者Tfovid

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.25 08:04:56