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

满足维度元素条件时反转张量指定维度的实现方法

张量维度反转处理方案

问题说明

需要对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)

关键逻辑说明

  1. 掩码生成:通过比较目标元素是否为0,生成与原张量维度匹配的布尔掩码,用于标记需要反转的样本;
  2. 张量反转:使用框架原生的反转API,针对目标维度(2维的feature、3维的horizon)执行反转;
  3. 条件选择:通过where操作,根据掩码对原张量和反转张量进行元素级选择,得到最终结果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 03:54:15