TensorFlow 2.0张量切片与更新:PyTorch转译报错求助
PyTorch转TensorFlow 2.0切片赋值问题解决
原PyTorch代码
def __getitem__(self, idx): mask = torch.from_numpy(self.mask[idx]) input_seq = torch.zeros(self.dataset[idx].shape, dtype=torch.float32) input_seq[1:, :] = torch.from_numpy(self.dataset[idx, :-1, :]) target = torch.from_numpy(self.dataset[idx]) return (input_seq, target, mask)
出错的TensorFlow实现及错误
def __getitem__(self, idx): mask = tf.convert_to_tensor(self.mask[idx]) input_seq = tf.zeros(self.dataset[idx].shape, dtype=tf.float32) input_seq[1:, :] = tf.convert_to_tensor(self.dataset[idx, :-1, :]) target = tf.convert_to_tensor(self.dataset[idx]) return (input_seq, target, mask)
错误信息:
'tensorflow.python.framework.ops.EagerTensor' object does not support item assignment
解决方法
TensorFlow中的EagerTensor是不可变对象,无法像PyTorch张量那样直接进行切片赋值。可以通过tf.concat直接构造符合要求的input_seq,替代赋值逻辑:
def __getitem__(self, idx): mask = tf.convert_to_tensor(self.mask[idx]) # 获取当前样本的形状 seq_len, feat_dim = self.dataset[idx].shape # 构造第一行全0张量 zero_row = tf.zeros((1, feat_dim), dtype=tf.float32) # 转换数据集切片为TensorFlow张量 dataset_slice = tf.convert_to_tensor(self.dataset[idx, :-1, :], dtype=tf.float32) # 拼接得到input_seq input_seq = tf.concat([zero_row, dataset_slice], axis=0) target = tf.convert_to_tensor(self.dataset[idx]) return (input_seq, target, mask)
原理说明
原PyTorch逻辑是初始化全0张量后,将第1行及以后的部分替换为数据集的前seq_len-1行。用tf.concat可以直接把全0的首行和数据集切片拼接,一步得到目标张量,避免了对不可变张量的赋值操作,同时逻辑更直观高效。
内容的提问来源于stack exchange,提问作者Aryan Vijaywargia
相关产品推荐
相关产品推荐

