TensorFlow中shape为()的字符串张量无法拆分该如何解决?
原因解释
张量shape为()的原因
你当前使用的是0维(标量)字符串张量,内部仅存储1个独立的字符串元素,你看到的多个短语是这个单个字符串的文本内容,而非张量的独立元素,因此张量的形状为空元组()。
set_shape报错的原因
set_shape仅用于补充张量的静态形状推断信息,不会修改张量的实际存储结构。你的张量实际是0维标量,强行指定为1维形状会触发不兼容报错。如果需要修改张量的实际形状可以调用tf.reshape,但该方法仅能做维度的拆分合并,无法改变元素总数,你当前张量的元素总数仅为1,因此reshape也无法直接得到包含4个元素的数组,必须先解析字符串内容。
拆分操作实现步骤
首先将标量内的字符串内容解析为包含4个独立字符串元素的1维张量,再按索引拆分即可:
- 解析标量字符串为1维字符串张量
import tensorflow as tf import ast # 假设你的原始张量为raw_tensor raw_tensor = tf.constant(b'["Can we make this quick? Roxanne Korrine and Andrew Barrett are having an incredibly horrendous public break- up on the quad. Again.",\n"Well, I thought we\'d start with pronunciation, if that\'s okay with you.",\n"Not the hacking and gagging and spitting part. Please.",\n"Okay... then how \'bout we try out some French cuisine. Saturday? Night?"]') # 场景1:Eager执行模式、非tf.data流水线场景 # 先转成Python字符串,解析为列表后重新构造张量 str_content = raw_tensor.numpy().decode('utf-8') str_list = ast.literal_eval(str_content) full_tensor = tf.constant(str_list, dtype=tf.string) # 场景2:静态图模式、tf.data流水线场景 # 纯TensorFlow算子实现解析 processed = tf.strings.regex_replace(raw_tensor, r'^b\[|\]$', '') splits = tf.strings.split(processed, sep=',\n') splits = tf.strings.regex_replace(splits, r'^[\'"]|[\'"]$', '') full_tensor = splits
- 拆分得到张量A和张量B
# 索引从0开始计数,第1、3条短语对应索引0、2,第2、4条对应索引1、3 tensor_a = tf.gather(full_tensor, [0, 2]) tensor_b = tf.gather(full_tensor, [1, 3])
内容的提问来源于stack exchange,提问作者OK 400
相关产品推荐
相关产品推荐

