概念引入
问题背景
多智能体的通信问题
架构与公式

关键公式:
简要介绍
RIAL 算法通信信息也作为一个离散的动作空间来考虑,因此,RIAL需要学习两个Q网络,一个用于输出动作Q值,另一个用于输出通信动作Q值。这里的通信动作是离散的,因此梯度无法从一个智能体传递给另一个智能体。
部分代码复现与解析
这段代码用于构建RNN网络,具体为GRU层,输入为opt.model_rnn_size,大小为128,层数为opt.model_rnn_layer=2,若dropout不为0则增加dropout层,之后的连接输出为128的BN层,Relu激活层和一层输入为128输出为opt.game_action_space_toal的全连接层,输出的参数视具体的实验而定。
#RIAL算法/增强智能体间学习
class SwitchCNet(nn,MOdule):
def __init__(SwitchCNet,self):
super(SwitchCNet,self).__init__()
self.opt =opt
droupt_rate = opt.model_rnn_droupt_rate or 0
self.rnn = nn.GRU(input_size = opt,model_rnn_size,hidden_size = opt.model_rnn_size,num_layers = opt.model_rnn_layers,droupt = droupt_rate,batch_first = True)
self.outputs = nn.Sequential()
if droupt_rate > 0:
self.outputs.add_moudle('droupt1',nn.Droupt(droupt_rate))
self.outputs.add_moudle('liner1',nn.liner(opt.model_rnn_size,opt.model_rnn_size))
if opt.model_bn:
self.outputs.add_moudle('batchnorm1',nn.batchnorm1d(opt.model_rnn_size))
self.outputs.add_moudle('relu1',nn.Relu(inplace = True))
self.outputs.add_moudle('liner2',nn.Linear(opt.model_rnn_size,opt.game_action_space_toal))
o_t为观测值,messages,hidden,prev_action,agent_index为智能体的索引agent_index, o_t , prev_action通过查找表传递,分别为 z_a =self.agent_lookup(agent_index), z_o= self.state_lookup(o_t)和 z_u =self.prev_action_lookup(prev_action).messages通过一个一层MLP,z_m =self.messages_mlp(messages.view(-1,self.comm_size)).输出大小为128,将 z_a, z_o, z_u 和 z_m求和得到z,将Z和内部状态hidden输入RNN网络,输出得到动作outputs和内部状态h_out.
def forward(self,o_t,messages.hidden,prev_action,agent_index):
opt = self.opt
o_t = Variable(o_t)
hidden = Variable(hidden)
prev_message = None
if opt.model_dial:
if opt.model_action_aware:
prev_action = Variable(prev_action)
else:
if opt,model_action_aware:
prev_action,prev_message = prev_action
prev_action = Variable(prev_action)
prev_message = Variable(messages)
agent_index = Variable(agent_index)
z_a,z_o,z_u,z_m = [0]*4
z_a = self.agent_index_lookup(agent_index)
z_o = self.state_lookup(o_t)
if opt.model_action_aware:
z_u = self.prev_action_lookup(prev_action)
if prev_message is not None:
z_u += self.prev_message_lookup(prev_message)
z_m = self.messages_mlp(messages.view(-1,self.comm_size))
z = z_a+z_o+z_u+z_m
z = z.unsqueeze(1)
rnn_out,h_out = self.rnn(z,hidden)
outputs = self.outputs(rnn_out[:,-1,:].squeeze())
return h_out,outputs
版权声明:本文为CSDNXXCQ原创文章,遵循CC 4.0 BY-SA版权协议,转载请附上原文出处链接和本声明。