| 43 | |
| 44 | |
| 45 | def 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 |