如何在TorchRL中定义依赖环境状态的Q-learning策略?
问题:TorchRL中适配动态合法动作子集的Q-learning策略定义
我正在训练一款扑克变体卡牌游戏的模型。每个玩家回合的合法动作是其手牌的子集(可空)。我将当前手牌和动作编码为52维布尔向量:
hand = torch.zeros(self.batch_size + (52,), dtype=torch.bool)
动作规格定义如下:
self.action_spec = BinaryDiscreteTensorSpec(52, dtype=torch.bool)
玩家的合法动作列表可从自定义环境中获取,比如下面是简化错误检查后的rand_action方法:
def rand_action(self, tensordict: TensorDictBase | None = None) -> TensorDictBase: if tensordict is not None: shape = tensordict.shape elif not self.batch_size: shape = torch.Size([]) moves = [] for i in range(self.total_batch_size()): game = self.games[i] game_moves = game.getMoves() game_move = game_moves[np.random.randint(len(game_moves))] moves.append([bool(i in game_move) for i in range(52)]) if len(self.batch_size) == 0: r = TensorDict({"action": torch.tensor(moves[0], dtype=torch.bool)}, batch_size=self.batch_size) else: r = TensorDict({"action": torch.tensor(moves, dtype=torch.bool)}, batch_size=self.batch_size) if tensordict is None: return r tensordict.update(r) return tensordict
我需要定义一个State × Vect(52) → Values的Q函数,策略要从环境返回的合法动作列表中选择最优动作。但标准的QValueActor和QValueModule类似乎是从预定义动作列表中做argmax选最优动作,TorchRL的文档和示例也都依赖预定义动作空间,找不到针对“动作空间是固定集合的动态子集”的实现示例。请问如何在TorchRL中定义这类Q-learning策略?
解决方案
针对动态合法动作子集的场景,核心思路是先计算所有合法动作的Q值,再从中选取最优动作,以下是具体实现步骤:
1. 自定义QValueModule
继承TorchRL的QValueModule,适配52维布尔动作的输入格式,实现状态与动作的特征融合及Q值输出:
from torchrl.modules import QValueModule import torch.nn as nn import torch class CustomQValueModule(QValueModule): def __init__(self, state_dim=52, action_dim=52): super().__init__() # 状态编码器:将手牌(52维布尔向量)编码为特征 self.state_encoder = nn.Sequential( nn.Flatten(), nn.Linear(state_dim, 128), nn.ReLU(), nn.Linear(128, 64) ) # 动作编码器:处理52维布尔动作向量 self.action_encoder = nn.Sequential( nn.Flatten(), nn.Linear(action_dim, 64), nn.ReLU() ) # 融合特征并输出Q值 self.q_head = nn.Sequential( nn.Linear(64 + 64, 32), nn.ReLU(), nn.Linear(32, 1) ) def forward(self, tensordict): # 提取状态和动作并转浮点型 state = tensordict["hand"].float() action = tensordict["action"].float() state_feat = self.state_encoder(state) action_feat = self.action_encoder(action) combined_feat = torch.cat([state_feat, action_feat], dim=-1) q_value = self.q_head(combined_feat) return tensordict.set("q_value", q_value)
2. 自定义QValueActor
继承QValueActor,重写逻辑实现“仅在合法动作中选最优”:
from torchrl.modules import QValueActor class CustomQValueActor(QValueActor): def __init__(self, qvalue_module, action_spec): super().__init__(qvalue_module, action_spec) def forward(self, tensordict): # 从tensordict中取出合法动作列表,shape=(batch_size, num_valid_moves, 52) valid_actions = tensordict["valid_actions"].float() batch_size, num_valid = valid_actions.shape[:2] # 复制状态到每个合法动作对应的维度,shape=(batch_size, num_valid, 52) state = tensordict["hand"].unsqueeze(1).repeat(1, num_valid, 1).float() # 构造临时张量字典,计算所有合法动作的Q值 temp_td = tensordict.clone() temp_td["hand"] = state temp_td["action"] = valid_actions q_values = self.qvalue_module(temp_td)["q_value"].squeeze(-1) # shape=(batch_size, num_valid) # 选出每个样本中Q值最大的动作索引 max_idx = torch.argmax(q_values, dim=-1) # 根据索引提取最优动作 best_actions = torch.gather( valid_actions, 1, max_idx.unsqueeze(-1).unsqueeze(-1).repeat(1, 1, 52) ).squeeze(1).bool() tensordict["action"] = best_actions return tensordict
3. 环境适配:传递合法动作到TensorDict
修改自定义环境,在step或reset时将合法动作列表以张量形式存入tensordict:
# 以step方法为例,添加合法动作存储逻辑 def step(self, tensordict): # ... 原有环境逻辑 ... # 初始化合法动作张量容器(需匹配batch维度) valid_actions_list = [] for i in range(self.total_batch_size()): game = self.games[i] valid_moves = game.getMoves() # 将合法动作转换为(num_valid, 52)的布尔张量 move_tensor = torch.tensor( [[bool(card in move) for card in range(52)] for move in valid_moves], dtype=torch.bool ) valid_actions_list.append(move_tensor) # 拼接为(batch_size, num_valid, 52)的张量(注意处理不同样本的num_valid差异,可补0对齐) valid_actions = torch.nn.utils.rnn.pad_sequence(valid_actions_list, batch_first=True) tensordict["valid_actions"] = valid_actions # ... 其他环境逻辑 ... return tensordict
4. 训练流程集成
将自定义模块与TorchRL的DQN训练框架结合:
# 初始化模块与Actor q_module = CustomQValueModule() actor = CustomQValueActor(q_module, env.action_spec) # 定义DQN损失函数 from torchrl.objectives import DQNLoss loss_fn = DQNLoss(qvalue_module=q_module, action_space=env.action_spec) # 训练循环(简化版) optimizer = torch.optim.Adam(q_module.parameters(), lr=1e-3) num_epochs = 1000 for _ in range(num_epochs): # 收集训练数据 data = env.collect(actor, max_steps=1000) # 计算损失并更新参数 loss = loss_fn(data) optimizer.zero_grad() loss.backward() optimizer.step()
内容的提问来源于stack exchange,提问作者Matt Cushman
相关产品推荐
相关产品推荐

