+
+ if lookahead_delta is not None:
+ r = rewards
+ print(f"{r.size()=} {lookahead_delta=}")
+ u = F.pad(r, (0, lookahead_delta - 1)).as_strided(
+ (r.size(0), r.size(1), lookahead_delta),
+ (r.size(1) + lookahead_delta - 1, 1, 1),
+ )
+ a = u.min(dim=-1).values
+ b = u.max(dim=-1).values
+ s = (a < 0).long() * a + (a >= 0).long() * b
+ lookahead_rewards = (1 + s[:, :, None]) + first_lookahead_rewards_code
+
+ r = rewards[:, :, None]
+ rewards = (r + 1) + first_rewards_code