如何根据事件布尔标记对时序张量事件前时间步生成掩码
时序事件掩码实现
需求说明
- 输入时序数据形状为
[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
相关产品推荐
相关产品推荐

