如何实现基于字符串键的TensorFlow Dataset分组操作?
解决TensorFlow字符串键分组问题
方法一:字符串转整数映射适配group_by_window
由于group_by_window要求key_func必须返回int64类型张量,核心思路是将字符串键映射为唯一整数ID,再用该ID作为分组键。
步骤1:生成字符串到整数的映射表
# 获取所有唯一的person字符串 unique_persons = tf.unique(person_tensor)[0] # 创建静态哈希表,为每个person分配唯一整数ID person_to_id = tf.lookup.StaticHashTable( tf.lookup.KeyValueTensorInitializer(unique_persons, tf.range(tf.size(unique_persons))), default_value=-1 # 处理未知键,当前场景可忽略 )
步骤2:修改分组逻辑使用映射后的整数ID
window_size = 10 grouped_ds = ds.group_by_window( key_func=lambda row: person_to_id.lookup(row['person']), window_size=window_size, reduce_func=lambda key, rows: rows.batch(window_size) ) # 验证分组结果 for batch in grouped_ds.as_numpy_iterator(): print(batch)
该方法全程在TensorFlow图模式下运行,适配大数据集,性能稳定。如果数据集是动态生成的(无法提前获取所有唯一键),可改用tf.lookup.experimental.DynamicHashTable动态添加映射关系。
方法二:小数据集内存分组(适合数据量不大的场景)
若数据集规模较小,可先将数据提取到内存中完成分组,再转回TensorFlow数据集:
import numpy as np # 提取所有数据到numpy数组 persons = person_tensor.numpy() values = value_tensor.numpy() # 按person字符串手动分组 groups = {} for p, v in zip(persons, values): groups.setdefault(p, []).append(v) # 将分组后的数据转回tf.data.Dataset grouped_ds = tf.data.Dataset.from_generator( lambda: ((k, v) for k, vs in groups.items() for v in vs), output_signature=( tf.TensorSpec(shape=(), dtype=tf.string), tf.TensorSpec(shape=(), dtype=tf.int32) ) ).map(lambda p, v: {'person': p, 'value': v}) # 可选:如需批量处理,仍需映射整数ID后使用group_by_window grouped_ds = grouped_ds.group_by_window( key_func=lambda row: person_to_id.lookup(row['person']), window_size=window_size, reduce_func=lambda key, rows: rows.batch(window_size) )
这种方法简单直接,但需要将所有数据加载到内存,不适合超大数据集。
内容的提问来源于stack exchange,提问作者leleogere
相关产品推荐
相关产品推荐

