满足维度元素条件时反转张量指定维度的实现方法
张量维度反转处理方案
问题说明
需要对3维张量(batch, horizon, feature)执行以下转换规则(示例简化为2维(batch, feature)展示效果):
- 遍历每个batch样本,若该样本horizon维度的第0个元素(示例中对应2维样本的第0个特征值)为0,则反转该样本的整个horizon维度;
- 若元素为1,则保持样本原样。
示例输入(2维简化版)
import torch input_tensor = torch.tensor([ [1., 1., 1., 1.], [1., 1., 1., 0.], [1., 1., 0., 0.], [1., 0., 0., 0.], [0., 0., 0., 0.], [0., 0., 0., 1.], [0., 0., 1., 1.], [0., 1., 1., 1.] ])
示例输出(2维简化版)
output_tensor = torch.tensor([ [1., 1., 1., 1.], [1., 1., 1., 0.], [1., 1., 0., 0.], [1., 0., 0., 0.], [0., 0., 0., 0.], [1., 0., 0., 0.], [1., 1., 0., 0.], [1., 1., 1., 0.] ])
实现代码
PyTorch 版本
通过掩码判断+张量反转实现,支持2维和3维张量:
import torch def transform_tensor(input_tensor): # 生成掩码:标记需要反转的样本(1=反转,0=保持) if input_tensor.dim() == 2: mask = (input_tensor[:, 0] == 0).unsqueeze(1) elif input_tensor.dim() == 3: mask = (input_tensor[:, 0, 0] == 0).unsqueeze(1).unsqueeze(2) # 反转对应维度:2维反转feature维度,3维反转horizon维度 reversed_dims = [1] reversed_tensor = input_tensor.flip(dims=reversed_dims) # 根据掩码选择原张量或反转后的张量 output_tensor = torch.where(mask, reversed_tensor, input_tensor) return output_tensor # 测试2维示例 input_2d = torch.tensor([ [1., 1., 1., 1.], [1., 1., 1., 0.], [1., 1., 0., 0.], [1., 0., 0., 0.], [0., 0., 0., 0.], [0., 0., 0., 1.], [0., 0., 1., 1.], [0., 1., 1., 1.] ]) output_2d = transform_tensor(input_2d) print("2维输出:\n", output_2d) # 测试3维示例 input_3d = torch.randint(0, 2, (2, 4, 3)).float() # shape: (batch=2, horizon=4, feature=3) print("\n3维输入:\n", input_3d) output_3d = transform_tensor(input_3d) print("\n3维输出:\n", output_3d)
TensorFlow 版本
逻辑与PyTorch一致,使用原生API实现:
import tensorflow as tf def transform_tensor(input_tensor): # 生成掩码 if input_tensor.shape.rank == 2: mask = tf.expand_dims(tf.equal(input_tensor[:, 0], 0), axis=1) elif input_tensor.shape.rank == 3: mask = tf.expand_dims(tf.expand_dims(tf.equal(input_tensor[:, 0, 0], 0), axis=1), axis=2) # 反转对应维度 reversed_axis = [1] reversed_tensor = tf.reverse(input_tensor, axis=reversed_axis) # 选择结果 output_tensor = tf.where(mask, reversed_tensor, input_tensor) return output_tensor # 测试2维示例 input_2d = tf.constant([ [1., 1., 1., 1.], [1., 1., 1., 0.], [1., 1., 0., 0.], [1., 0., 0., 0.], [0., 0., 0., 0.], [0., 0., 0., 1.], [0., 0., 1., 1.], [0., 1., 1., 1.] ]) output_2d = transform_tensor(input_2d) tf.print("2维输出:\n", output_2d)
关键逻辑说明
- 掩码生成:通过比较目标元素是否为0,生成与原张量维度匹配的布尔掩码,用于标记需要反转的样本;
- 张量反转:使用框架原生的反转API,针对目标维度(2维的feature、3维的horizon)执行反转;
- 条件选择:通过
where操作,根据掩码对原张量和反转张量进行元素级选择,得到最终结果。
内容的提问来源于stack exchange,提问作者SnakeWasTheNameTheyGaveMe
相关产品推荐
相关产品推荐

