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

如何根据事件布尔标记对时序张量事件前时间步生成掩码

时序事件掩码实现

需求说明

  • 输入时序数据形状为 [batch_size, horizon, feature],事件不定期触发,使用和数据同形状的meta张量标记事件位置:事件发生位置值为1,其余位置填0。
  • 掩码规则:
    • 时间窗口(horizon维度)默认掩码全为1,代表所有时间步可访问;
    • 如果窗口内存在事件,禁止模型访问最后一次事件发生前的所有数据,对应位置掩码置0;
    • 规则仅以窗口内最后一次事件为判定基准,更早的事件不影响最后一次事件的位置判定。
  • 一维场景映射示例:
[0, 0, 1, 0] -> [0, 0, 1, 1]
[0, 0, 0, 1] -> [0, 0, 0, 1]
[1, 0, 1, 0] -> [0, 0, 1, 1]
[1, 0, 0, 0] -> [1, 1, 1, 1]
[0, 0, 0, 0] -> [1, 1, 1, 1]

实现代码(PyTorch)

核心逻辑是沿时间轴反向做累积求和,快速定位最后一个事件的位置,无需遍历循环,支持批量并行计算:

import torch

def build_event_mask(meta: torch.Tensor) -> torch.Tensor:
    """
    生成符合规则的时序访问掩码
    :param meta: 事件标记张量,形状[batch_size, horizon, feature],事件位置为1,其余为0
    :return: 访问掩码,形状和meta一致,1代表对应位置可访问,0代表不可访问
    """
    # 沿时间轴反向翻转后做累积求和,最后一个事件及之后位置求和结果>=1,之前为0
    reversed_cumsum = torch.flip(meta, dims=[1]).cumsum(dim=1)
    mask = torch.flip(reversed_cumsum >= 1, dims=[1]).to(meta.dtype)
    # 窗口内无事件时,掩码全置为1
    no_event_flag = meta.sum(dim=1, keepdim=True) == 0
    mask = mask + no_event_flag.to(meta.dtype)
    return mask

效果验证

用给出的一维样例测试,所有结果均符合预期:

if __name__ == "__main__":
    test_samples = [
        ([0,0,1,0], [0,0,1,1]),
        ([0,0,0,1], [0,0,0,1]),
        ([1,0,1,0], [0,0,1,1]),
        ([1,0,0,0], [1,1,1,1]),
        ([0,0,0,0], [1,1,1,1]),
    ]
    for meta_input, expected_output in test_samples:
        # 升维适配[batch_size, horizon, feature]格式,batch=1、feature=1
        meta_tensor = torch.tensor(meta_input, dtype=torch.int64).reshape(1, 4, 1)
        output = build_event_mask(meta_tensor).flatten().tolist()
        print(f"输入:{meta_input},输出:{output},匹配预期:{output == expected_output}")

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 08:45:33