强化学习场景下,如何无复制通过多切片原地更新PyTorch Tensor
强化学习状态更新优化问题
我正在开展强化学习相关工作:游戏状态由一组车辆构成,每辆车存在若干漏洞,每个漏洞包含三个特征:攻击成功率、攻击严重程度、攻击成功标记。批量状态的Tensor形状为[batch, vehicle, vuln, feature]。攻击者的动作是选择每个状态中的部分车辆进行攻击,状态转换时需对选中车辆的每个漏洞进行攻击成功判定。
目前我已实现批量状态与动作的处理,但采用了临时的workaround,想请教是否存在更优方案,无需复制Tensor,通过多切片操作原地更新得到next_states Tensor?
原实现代码
import torch num_batches = 4 MAX_VEHICLES = 3 MAX_VULNS = 3 MAX_ATTACK = 2 # 批量状态示例定义 prob_dist = torch.distributions.Normal(loc=torch.as_tensor(0.5, dtype=torch.float32), scale=torch.as_tensor(0.25, dtype=torch.float32)) sev_dist = torch.distributions.Normal(loc=torch.as_tensor(2, dtype=torch.float32),scale=torch.as_tensor(1, dtype=torch.float32)) states = torch.zeros((num_batches, MAX_VEHICLES, MAX_VULNS, 3), dtype=torch.float32) states[:,:,:,0] = prob_dist.sample(torch.Size((num_batches, MAX_VEHICLES, MAX_VULNS))).clamp(0,1) states[:,:,:,1] = sev_dist.sample(torch.Size((num_batches, MAX_VEHICLES, MAX_VULNS))).round().clamp(1,5) print("states", states.shape, states) # 批量动作示例定义:攻击者从每个状态中选择MAX_ATTACK辆车 priority = (states[:,:,:,0] * states[:,:,:,1] * (1-states[:,:,:,2])).sum(dim=-1) attack = priority.topk(MAX_ATTACK).indices attack_mask = torch.zeros((states.shape[0], MAX_VEHICLES), dtype=torch.float32).to(states.device).scatter_(1, attack, 1).bool() print("attack mask", attack_mask) # 对被攻击车辆的每个漏洞进行攻击成功判定 next_states = states.clone() attack = next_states[attack_mask] probs = torch.rand(attack.shape[:-1], dtype=torch.float32, device=states.device) success_mask = (probs > 1-attack[:,:,0]) print("success mask", success_mask) # 此方法失效:高级切片会克隆Tensor # next_states[attack_mask][success_mask][:,2] = 1 # 临时workaround:必须重建Tensor success = attack[success_mask] success[:,2] = 1 attack[success_mask] = success next_states[attack_mask] = attack print("next states", next_states)
优化方案
核心思路是构建全局维度的掩码,直接定位到需要更新的位置,避免多次克隆子Tensor,实现原地更新。
优化后代码
import torch num_batches = 4 MAX_VEHICLES = 3 MAX_VULNS = 3 MAX_ATTACK = 2 # 生成批量状态示例 prob_dist = torch.distributions.Normal(loc=torch.as_tensor(0.5, dtype=torch.float32), scale=torch.as_tensor(0.25, dtype=torch.float32)) sev_dist = torch.distributions.Normal(loc=torch.as_tensor(2, dtype=torch.float32),scale=torch.as_tensor(1, dtype=torch.float32)) states = torch.zeros((num_batches, MAX_VEHICLES, MAX_VULNS, 3), dtype=torch.float32) states[:,:,:,0] = prob_dist.sample(torch.Size((num_batches, MAX_VEHICLES, MAX_VULNS))).clamp(0,1) states[:,:,:,1] = sev_dist.sample(torch.Size((num_batches, MAX_VEHICLES, MAX_VULNS))).round().clamp(1,5) print("states", states.shape, states) # 生成批量动作示例:每个状态选择MAX_ATTACK辆车攻击 priority = (states[:,:,:,0] * states[:,:,:,1] * (1-states[:,:,:,2])).sum(dim=-1) attack = priority.topk(MAX_ATTACK).indices attack_mask = torch.zeros((states.shape[0], MAX_VEHICLES), dtype=torch.float32).to(states.device).scatter_(1, attack, 1).bool() print("attack mask", attack_mask) # 状态转换:直接原地更新next_states,避免多余克隆 next_states = states.clone() # 仅克隆原状态一次 # 将attack_mask扩展到vuln维度,匹配next_states的前三维结构 attack_mask_expanded = attack_mask.unsqueeze(-1).expand(-1, -1, MAX_VULNS) # 生成全局的攻击成功概率判定 probs = torch.rand(next_states.shape[0], next_states.shape[1], next_states.shape[2], device=states.device) # 构建全局成功掩码:选中攻击的车辆 且 漏洞攻击成功 global_success_mask = attack_mask_expanded & (probs > 1 - next_states[:,:,:,0]) # 直接更新对应位置的攻击成功标记(第3个特征,索引为2) next_states[global_success_mask, 2] = 1 print("next states", next_states)
优化说明
- 减少Tensor拷贝:原方案中多次克隆子Tensor并来回赋值,优化后仅需克隆一次原状态,后续通过全局掩码直接修改,无额外子Tensor复制操作。
- 原地更新实现:通过
global_success_mask直接定位到next_states中需要修改的(batch, vehicle, vuln)位置,直接修改该位置的攻击成功标记,无需中间变量传递。 - 逻辑更简洁:从全局维度构建掩码,避免了多层索引导致的Tensor拷贝问题,代码可读性和执行效率都更优。
示例输出
states torch.Size([2, 3, 3, 3]) tensor([[[[0.3672, 2.0000, 0.0000], [0.2386, 2.0000, 0.0000], [1.0000, 1.0000, 0.0000]], [[0.5873, 2.0000, 0.0000], [0.7048, 3.0000, 0.0000], [0.5552, 1.0000, 0.0000]], [[0.4456, 3.0000, 0.0000], [0.3886, 2.0000, 0.0000], [0.3745, 1.0000, 0.0000]]], [[[0.0000, 3.0000, 0.0000], [0.3182, 3.0000, 0.0000], [0.0357, 2.0000, 0.0000]], [[0.6474, 1.0000, 0.0000], [0.7376, 2.0000, 0.0000], [0.6590, 2.0000, 0.0000]], [[0.6119, 1.0000, 0.0000], [0.7148, 3.0000, 0.0000], [0.6434, 2.0000, 0.0000]]]]) attack mask tensor([[False, True, True], [False, True, True]]) next states tensor([[[[0.3672, 2.0000, 0.0000], [0.2386, 2.0000, 0.0000], [1.0000, 1.0000, 0.0000]], [[0.5873, 2.0000, 1.0000], [0.7048, 3.0000, 1.0000], [0.5552, 1.0000, 0.0000]], [[0.4456, 3.0000, 0.0000], [0.3886, 2.0000, 1.0000], [0.3745, 1.0000, 0.0000]]], [[[0.0000, 3.0000, 0.0000], [0.3182, 3.0000, 0.0000], [0.0357, 2.0000, 0.0000]], [[0.6474, 1.0000, 1.0000], [0.7376, 2.0000, 1.0000], [0.6590, 2.0000, 1.0000]], [[0.6119, 1.0000, 1.0000], [0.7148, 3.0000, 1.0000], [0.6434, 2.0000, 1.0000]]]])
内容的提问来源于stack exchange,提问作者TeamDman
相关产品推荐
相关产品推荐

