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

强化学习场景下,如何无复制通过多切片原地更新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)

优化说明

  1. 减少Tensor拷贝:原方案中多次克隆子Tensor并来回赋值,优化后仅需克隆一次原状态,后续通过全局掩码直接修改,无额外子Tensor复制操作。
  2. 原地更新实现:通过global_success_mask直接定位到next_states中需要修改的(batch, vehicle, vuln)位置,直接修改该位置的攻击成功标记,无需中间变量传递。
  3. 逻辑更简洁:从全局维度构建掩码,避免了多层索引导致的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 14:12:05