MCPcopy Create free account
hub / github.com/chenbq/CAVDN / td_lambda_target

Function td_lambda_target

common/utils.py:45–85  ·  view source on GitHub ↗
(batch, max_episode_len, q_targets, args)

Source from the content-addressed store, hash-verified

43
44
45def td_lambda_target(batch, max_episode_len, q_targets, args):
46 # batch.shep = (episode_num, max_episode_len, n_agents,n_actions)
47 # q_targets.shape = (episode_num, max_episode_len, n_agents)
48 episode_num = batch['o'].shape[0]
49 mask = (1 - batch["padded"].float()).repeat(1, 1, args.n_agents)
50 terminated = (1 - batch["terminated"].float()).repeat(1, 1, args.n_agents)
51 r = batch['r'].repeat((1, 1, args.n_agents))
52 # --------------------------------------------------n_step_return---------------------------------------------------
53 '''
54 1. 每条经验都有若干个n_step_return,所以给一个最大的max_episode_len维度用来装n_step_return
55 最后一维,第n个数代表 n+1 step。
56 2. 因为batch中各个episode的长度不一样,所以需要用mask将多出的n-step return置为0,
57 否则的话会影响后面的lambda return。第t条经验的lambda return是和它后面的所有n-step return有关的,
58 如果没有置0,在计算td-error后再置0是来不及的
59 3. terminated用来将超出当前episode长度的q_targets和r置为0
60 '''
61 n_step_return = torch.zeros((episode_num, max_episode_len, args.n_agents, max_episode_len))
62 for transition_idx in range(max_episode_len - 1, -1, -1):
63 # 最后计算1 step return
64 n_step_return[:, transition_idx, :, 0] = (r[:, transition_idx] + args.gamma * q_targets[:, transition_idx] * terminated[:, transition_idx]) * mask[:, transition_idx] # 经验transition_idx上的obs有max_episode_len - transition_idx个return, 分别计算每种step return
65 # 同时要注意n step return对应的index为n-1
66 for n in range(1, max_episode_len - transition_idx):
67 # t时刻的n step return =r + gamma * (t + 1 时刻的 n-1 step return)
68 # n=1除外, 1 step return =r + gamma * (t + 1 时刻的 Q)
69 n_step_return[:, transition_idx, :, n] = (r[:, transition_idx] + args.gamma * n_step_return[:, transition_idx + 1, :, n - 1]) * mask[:, transition_idx]
70 # --------------------------------------------------n_step_return---------------------------------------------------
71
72 # --------------------------------------------------lambda return---------------------------------------------------
73 '''
74 lambda_return.shape = (episode_num, max_episode_len,n_agents)
75 '''
76 lambda_return = torch.zeros((episode_num, max_episode_len, args.n_agents))
77 for transition_idx in range(max_episode_len):
78 returns = torch.zeros((episode_num, args.n_agents))
79 for n in range(1, max_episode_len - transition_idx):
80 returns += pow(args.td_lambda, n - 1) * n_step_return[:, transition_idx, :, n - 1]
81 lambda_return[:, transition_idx] = (1 - args.td_lambda) * returns + \
82 pow(args.td_lambda, max_episode_len - transition_idx - 1) * \
83 n_step_return[:, transition_idx, :, max_episode_len - transition_idx - 1]
84 # --------------------------------------------------lambda return---------------------------------------------------
85 return lambda_return

Callers 1

_train_criticMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected