evolutionary-policy-optimization 0.3.0__tar.gz → 0.3.5__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.5}/PKG-INFO +3 -3
- evolutionary_policy_optimization-0.3.5/evolutionary_policy_optimization/__init__.py +3 -0
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.5}/evolutionary_policy_optimization/distributed.py +4 -5
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.5}/evolutionary_policy_optimization/env_wrappers.py +3 -3
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.5}/evolutionary_policy_optimization/epo.py +316 -137
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.5}/evolutionary_policy_optimization/experimental.py +5 -5
- evolutionary_policy_optimization-0.3.5/evolutionary_policy_optimization/mock_env.py +128 -0
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.5}/pyproject.toml +3 -3
- evolutionary_policy_optimization-0.3.0/evolutionary_policy_optimization/__init__.py +0 -18
- evolutionary_policy_optimization-0.3.0/evolutionary_policy_optimization/mock_env.py +0 -55
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.5}/.gitignore +0 -0
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.5}/LICENSE +0 -0
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.5}/README.md +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.5
|
|
2
2
|
Name: evolutionary-policy-optimization
|
|
3
|
-
Version: 0.3.
|
|
3
|
+
Version: 0.3.5
|
|
4
4
|
Summary: EPO - Pytorch
|
|
5
5
|
Project-URL: Homepage, https://pypi.org/project/evolutionary-policy-optimization/
|
|
6
6
|
Project-URL: Repository, https://github.com/lucidrains/evolutionary-policy-optimization
|
|
@@ -33,14 +33,14 @@ Classifier: Intended Audience :: Developers
|
|
|
33
33
|
Classifier: License :: OSI Approved :: MIT License
|
|
34
34
|
Classifier: Programming Language :: Python :: 3.8
|
|
35
35
|
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
|
|
36
|
-
Requires-Python: >=3.
|
|
36
|
+
Requires-Python: >=3.10
|
|
37
37
|
Requires-Dist: accelerate>=1.6.0
|
|
38
|
-
Requires-Dist: adam-atan2-pytorch
|
|
39
38
|
Requires-Dist: assoc-scan>=0.0.2
|
|
40
39
|
Requires-Dist: einops>=0.8.1
|
|
41
40
|
Requires-Dist: einx>=0.3.0
|
|
42
41
|
Requires-Dist: ema-pytorch>=0.7.7
|
|
43
42
|
Requires-Dist: hl-gauss-pytorch>=0.2.8
|
|
43
|
+
Requires-Dist: mean-conc-beta>=0.0.7
|
|
44
44
|
Requires-Dist: torch>=2.2
|
|
45
45
|
Requires-Dist: tqdm
|
|
46
46
|
Requires-Dist: x-mlps-pytorch>=0.3.0
|
|
@@ -0,0 +1,3 @@
|
|
|
1
|
+
from evolutionary_policy_optimization.env_wrappers import GymnasiumEnvWrapper, rescale_from_to
|
|
2
|
+
from evolutionary_policy_optimization.epo import EPO, MLP, Actor, Agent, Beta, CategoricalActionDistr, Critic, LatentGenePool, create_agent
|
|
3
|
+
from evolutionary_policy_optimization.mock_env import Env, VectorEnv
|
|
@@ -1,12 +1,11 @@
|
|
|
1
|
+
import einx
|
|
1
2
|
import torch
|
|
2
|
-
|
|
3
|
+
import torch.distributed as dist
|
|
3
4
|
import torch.nn.functional as F
|
|
5
|
+
from einops import rearrange
|
|
6
|
+
from torch import nn
|
|
4
7
|
from torch.autograd import Function
|
|
5
8
|
|
|
6
|
-
import torch.distributed as dist
|
|
7
|
-
|
|
8
|
-
import einx
|
|
9
|
-
from einops import rearrange
|
|
10
9
|
|
|
11
10
|
def exists(val):
|
|
12
11
|
return val is not None
|
|
@@ -5,8 +5,8 @@ from torch.nn import Module
|
|
|
5
5
|
|
|
6
6
|
from evolutionary_policy_optimization.epo import Agent, create_agent, exists
|
|
7
7
|
|
|
8
|
-
def rescale_from_to(x, from_range = (
|
|
9
|
-
# e.g. beta actions on (
|
|
8
|
+
def rescale_from_to(x, from_range = (-1., 1.), to_range = (-1., 1.)):
|
|
9
|
+
# e.g. beta actions on (-1, 1) -> the env's action bounds
|
|
10
10
|
|
|
11
11
|
from_low, from_high = from_range
|
|
12
12
|
to_low, to_high = to_range
|
|
@@ -35,7 +35,7 @@ class GymnasiumEnvWrapper(Module):
|
|
|
35
35
|
return self.env.reset(*args, **kwargs)
|
|
36
36
|
|
|
37
37
|
def step(self, actions, *args, **kwargs):
|
|
38
|
-
#
|
|
38
|
+
# continuous actions are on (-1, 1) - rescale to env bounds at interface
|
|
39
39
|
|
|
40
40
|
if exists(self.rescale_to):
|
|
41
41
|
actions = rescale_from_to(actions, to_range = self.rescale_to)
|
|
@@ -1,50 +1,36 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
|
-
from typing import Callable
|
|
3
2
|
|
|
4
3
|
import math
|
|
5
|
-
from pathlib import Path
|
|
6
|
-
from math import ceil
|
|
7
|
-
from itertools import product
|
|
8
|
-
from functools import partial, wraps
|
|
9
4
|
from collections import namedtuple
|
|
5
|
+
from copy import deepcopy
|
|
6
|
+
from functools import partial, wraps
|
|
7
|
+
from itertools import product
|
|
8
|
+
from math import ceil
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
from typing import Callable
|
|
10
11
|
|
|
12
|
+
import einx
|
|
11
13
|
import numpy as np
|
|
12
|
-
|
|
13
14
|
import torch
|
|
14
|
-
from torch import nn, cat, stack, is_tensor, tensor, from_numpy, Tensor
|
|
15
|
-
import torch.nn.functional as F
|
|
16
15
|
import torch.distributed as dist
|
|
17
|
-
|
|
18
|
-
from
|
|
19
|
-
from torch.
|
|
20
|
-
|
|
21
|
-
import einx
|
|
22
|
-
from einops import rearrange, repeat, reduce, einsum
|
|
23
|
-
from einops.layers.torch import Rearrange
|
|
24
|
-
|
|
25
|
-
from torch.distributions import Beta as _Beta, Categorical, Distribution
|
|
26
|
-
|
|
27
|
-
from x_mlps_pytorch import AttnResidualNormedMLP
|
|
28
|
-
|
|
29
|
-
from evolutionary_policy_optimization.distributed import (
|
|
30
|
-
is_distributed,
|
|
31
|
-
get_world_and_rank,
|
|
32
|
-
maybe_sync_seed,
|
|
33
|
-
all_gather,
|
|
34
|
-
maybe_barrier
|
|
35
|
-
)
|
|
36
|
-
|
|
16
|
+
import torch.nn.functional as F
|
|
17
|
+
from accelerate import Accelerator
|
|
18
|
+
from torch.optim import AdamW
|
|
37
19
|
from assoc_scan import AssocScan
|
|
38
|
-
|
|
39
|
-
from
|
|
40
|
-
|
|
41
|
-
from hl_gauss_pytorch import HLGaussLayer
|
|
42
|
-
|
|
20
|
+
from einops import einsum, rearrange, reduce, repeat
|
|
21
|
+
from einops.layers.torch import Rearrange
|
|
43
22
|
from ema_pytorch import EMA
|
|
44
|
-
|
|
23
|
+
from hl_gauss_pytorch import HLGaussLayer
|
|
24
|
+
from mean_conc_beta import Beta
|
|
25
|
+
from torch import Tensor, cat, from_numpy, is_tensor, nn, stack, tensor
|
|
26
|
+
from torch.distributions import Categorical, Distribution
|
|
27
|
+
from torch.nn import Linear, Module, ModuleList
|
|
28
|
+
from torch.utils._pytree import tree_map
|
|
29
|
+
from torch.utils.data import DataLoader, TensorDataset
|
|
45
30
|
from tqdm import tqdm
|
|
31
|
+
from x_mlps_pytorch import AttnResidualNormedMLP
|
|
46
32
|
|
|
47
|
-
from
|
|
33
|
+
from evolutionary_policy_optimization.distributed import all_gather, get_world_and_rank, is_distributed, maybe_barrier, maybe_sync_seed
|
|
48
34
|
|
|
49
35
|
# helpers
|
|
50
36
|
|
|
@@ -124,9 +110,6 @@ def l2norm(t, dim = -1):
|
|
|
124
110
|
def batch_randperm(shape, device):
|
|
125
111
|
return torch.randn(shape, device = device).argsort(dim = -1)
|
|
126
112
|
|
|
127
|
-
def log(t, eps = 1e-20):
|
|
128
|
-
return t.clamp(min = eps).log()
|
|
129
|
-
|
|
130
113
|
def sum_to_batch(t):
|
|
131
114
|
# fold trailing action dims into one value per state
|
|
132
115
|
|
|
@@ -561,8 +544,8 @@ class DiversityDiscr(Module):
|
|
|
561
544
|
return self.net((state_embed, next_state_embed))
|
|
562
545
|
|
|
563
546
|
# action distributions - the actor always returns a `Distribution`, either
|
|
564
|
-
# categorical (discrete) or beta mean-conc (continuous
|
|
565
|
-
#
|
|
547
|
+
# categorical (discrete) or beta mean-conc (continuous). beta is defined over
|
|
548
|
+
# (-1, 1) - multiplied by some constant usually (e.g. 0.4) for the env action range
|
|
566
549
|
|
|
567
550
|
class CategoricalActionDistr(Module):
|
|
568
551
|
def forward(self, logits, temperature = 1.):
|
|
@@ -571,50 +554,75 @@ class CategoricalActionDistr(Module):
|
|
|
571
554
|
|
|
572
555
|
return Categorical(logits = logits)
|
|
573
556
|
|
|
574
|
-
|
|
557
|
+
# self-predictive representations (SPR) over the actor's hidden embedding
|
|
558
|
+
|
|
559
|
+
class HiddenSpr(Module):
|
|
575
560
|
def __init__(
|
|
576
561
|
self,
|
|
577
|
-
|
|
578
|
-
|
|
579
|
-
|
|
562
|
+
actor,
|
|
563
|
+
dim_hidden,
|
|
564
|
+
num_actions,
|
|
565
|
+
dim_action = 32,
|
|
580
566
|
):
|
|
581
567
|
super().__init__()
|
|
582
|
-
|
|
583
|
-
|
|
584
|
-
self.init_conc = init_conc
|
|
585
|
-
self.min_conc = min_conc
|
|
586
|
-
self.eps = eps
|
|
568
|
+
action_is_continuous = actor.beta_actions
|
|
569
|
+
self.action_is_continuous = action_is_continuous
|
|
587
570
|
|
|
588
|
-
|
|
571
|
+
if action_is_continuous:
|
|
572
|
+
self.action_proj = nn.Linear(num_actions, dim_action)
|
|
573
|
+
else:
|
|
574
|
+
self.action_proj = nn.Embedding(num_actions, dim_action)
|
|
589
575
|
|
|
590
|
-
self.
|
|
576
|
+
self.to_dynamics = nn.Sequential(
|
|
577
|
+
nn.Linear(dim_hidden + dim_action, dim_hidden),
|
|
578
|
+
nn.SiLU(),
|
|
579
|
+
nn.Linear(dim_hidden, dim_hidden)
|
|
580
|
+
)
|
|
591
581
|
|
|
592
|
-
|
|
593
|
-
# the beta mean is exactly (tanh(raw_mean) + 1) / 2 by construction
|
|
582
|
+
self.proj_head = nn.Linear(dim_hidden, dim_hidden, bias = False)
|
|
594
583
|
|
|
595
|
-
|
|
596
|
-
return ((torch.tanh(raw_mean) + 1.) * 0.5).clamp(min = self.eps, max = 1. - self.eps)
|
|
584
|
+
# ema target of the actor body + projection head
|
|
597
585
|
|
|
598
|
-
|
|
599
|
-
|
|
586
|
+
self.target_actor = deepcopy(actor).requires_grad_(False)
|
|
587
|
+
self.target_proj = deepcopy(self.proj_head).requires_grad_(False)
|
|
600
588
|
|
|
601
|
-
|
|
589
|
+
def online_parameters(self):
|
|
590
|
+
return [
|
|
591
|
+
*self.action_proj.parameters(),
|
|
592
|
+
*self.to_dynamics.parameters(),
|
|
593
|
+
*self.proj_head.parameters(),
|
|
594
|
+
]
|
|
602
595
|
|
|
603
|
-
|
|
596
|
+
def predict(self, hidden, action):
|
|
597
|
+
if not self.action_is_continuous:
|
|
598
|
+
action = action.long()
|
|
599
|
+
if action.ndim > 1:
|
|
600
|
+
action = rearrange(action, '... 1 -> ...')
|
|
604
601
|
|
|
605
|
-
|
|
602
|
+
action_proj = self.action_proj(action)
|
|
603
|
+
dynamics = self.to_dynamics(cat((hidden, action_proj), dim = -1))
|
|
604
|
+
return self.proj_head(dynamics + hidden)
|
|
606
605
|
|
|
607
|
-
|
|
606
|
+
@torch.no_grad()
|
|
607
|
+
def target(self, state, latent):
|
|
608
|
+
return self.target_proj(self.target_actor.latent(state, latent))
|
|
608
609
|
|
|
609
|
-
|
|
610
|
+
@torch.no_grad()
|
|
611
|
+
def _ema_update_one_(self, target_module, source_module, decay):
|
|
612
|
+
# if decay is 0, this is a hard copy of the source into the target
|
|
610
613
|
|
|
611
|
-
|
|
612
|
-
|
|
614
|
+
for target_param, source_param in zip(target_module.parameters(), source_module.parameters()):
|
|
615
|
+
target_param.lerp_(source_param, 1. - decay)
|
|
613
616
|
|
|
614
|
-
|
|
615
|
-
|
|
617
|
+
@torch.no_grad()
|
|
618
|
+
def ema_update(self, actor, decay):
|
|
619
|
+
self._ema_update_one_(self.target_actor.init_layer, actor.init_layer, decay)
|
|
620
|
+
self._ema_update_one_(self.target_actor.mlp, actor.mlp, decay)
|
|
621
|
+
self._ema_update_one_(self.target_proj, self.proj_head, decay)
|
|
616
622
|
|
|
617
|
-
|
|
623
|
+
@torch.no_grad()
|
|
624
|
+
def sync_targets_(self, actor):
|
|
625
|
+
self.ema_update(actor, decay = 0.)
|
|
618
626
|
|
|
619
627
|
# actor, critic, and agent (actor + critic)
|
|
620
628
|
# eventually, should just create a separate repo and aggregate all the MLP related architectures
|
|
@@ -629,11 +637,14 @@ class Actor(Module):
|
|
|
629
637
|
state_norm: StateNorm | None = None,
|
|
630
638
|
dim_latent = 0,
|
|
631
639
|
action_is_continuous = False, # continuous control - beta policy
|
|
640
|
+
beta_kwargs: dict = dict(),
|
|
632
641
|
):
|
|
633
642
|
super().__init__()
|
|
634
643
|
|
|
635
644
|
self.state_norm = state_norm
|
|
636
645
|
|
|
646
|
+
self.dim = dim
|
|
647
|
+
self.num_actions = num_actions
|
|
637
648
|
self.dim_latent = dim_latent
|
|
638
649
|
self.beta_actions = action_is_continuous
|
|
639
650
|
|
|
@@ -658,14 +669,16 @@ class Actor(Module):
|
|
|
658
669
|
nn.Linear(dim, num_actions, bias = False),
|
|
659
670
|
)
|
|
660
671
|
|
|
661
|
-
self.action_distr =
|
|
672
|
+
self.action_distr = Beta(**beta_kwargs) if self.beta_actions else CategoricalActionDistr()
|
|
662
673
|
|
|
663
|
-
def
|
|
674
|
+
def latent(
|
|
664
675
|
self,
|
|
665
676
|
state,
|
|
666
|
-
latent
|
|
667
|
-
|
|
668
|
-
|
|
677
|
+
latent
|
|
678
|
+
) -> Tensor:
|
|
679
|
+
# hidden embedding right before the action projection -
|
|
680
|
+
# the dynamics in hidden_spr operate over this
|
|
681
|
+
|
|
669
682
|
if exists(self.state_norm):
|
|
670
683
|
with torch.no_grad():
|
|
671
684
|
self.state_norm.eval()
|
|
@@ -673,8 +686,15 @@ class Actor(Module):
|
|
|
673
686
|
|
|
674
687
|
hidden = self.init_layer(state)
|
|
675
688
|
|
|
676
|
-
|
|
689
|
+
return self.mlp(hidden, latent)
|
|
677
690
|
|
|
691
|
+
def forward(
|
|
692
|
+
self,
|
|
693
|
+
state,
|
|
694
|
+
latent,
|
|
695
|
+
temperature = 1.
|
|
696
|
+
) -> Distribution:
|
|
697
|
+
hidden = self.latent(state, latent)
|
|
678
698
|
return self.action_distr(self.to_out(hidden), temperature = temperature)
|
|
679
699
|
|
|
680
700
|
class Critic(Module):
|
|
@@ -724,9 +744,10 @@ class Critic(Module):
|
|
|
724
744
|
self,
|
|
725
745
|
state,
|
|
726
746
|
latent,
|
|
727
|
-
old_values,
|
|
728
747
|
target,
|
|
729
|
-
|
|
748
|
+
old_values = None,
|
|
749
|
+
clip_value = False,
|
|
750
|
+
eps_clip = 0.8,
|
|
730
751
|
use_improved = True
|
|
731
752
|
):
|
|
732
753
|
|
|
@@ -737,6 +758,9 @@ class Critic(Module):
|
|
|
737
758
|
|
|
738
759
|
logits = self.forward(state, latent, return_logits = True)
|
|
739
760
|
|
|
761
|
+
if not clip_value or not exists(old_values) or not exists(eps_clip):
|
|
762
|
+
return self.loss_fn(logits, target)
|
|
763
|
+
|
|
740
764
|
value = self.maybe_bins_to_value(logits)
|
|
741
765
|
|
|
742
766
|
loss_fn = partial(self.loss_fn, reduction = 'none')
|
|
@@ -866,7 +890,6 @@ class LatentGenePool(Module):
|
|
|
866
890
|
assert (frac_natural_selected + frac_elitism) < 1.
|
|
867
891
|
|
|
868
892
|
self.dim_latent = dim_latent
|
|
869
|
-
self.num_latents = num_latents
|
|
870
893
|
self.num_islands = num_islands
|
|
871
894
|
|
|
872
895
|
latents_per_island = num_latents // num_islands
|
|
@@ -1005,6 +1028,8 @@ class LatentGenePool(Module):
|
|
|
1005
1028
|
should_update_per_island = self.should_run_genetic_algorithm(fitness)
|
|
1006
1029
|
|
|
1007
1030
|
if not should_update_per_island.any():
|
|
1031
|
+
self.advance_step_()
|
|
1032
|
+
|
|
1008
1033
|
if inplace:
|
|
1009
1034
|
return False, None
|
|
1010
1035
|
|
|
@@ -1155,13 +1180,13 @@ class Agent(Module):
|
|
|
1155
1180
|
actor: Actor,
|
|
1156
1181
|
critic: Critic,
|
|
1157
1182
|
latent_gene_pool: LatentGenePool | None,
|
|
1158
|
-
optim_klass =
|
|
1183
|
+
optim_klass = AdamW,
|
|
1159
1184
|
state_norm: StateNorm | None = None,
|
|
1160
1185
|
actor_lr = 8e-4,
|
|
1161
1186
|
critic_lr = 8e-4,
|
|
1162
1187
|
latent_lr = 1e-5,
|
|
1163
|
-
actor_weight_decay =
|
|
1164
|
-
critic_weight_decay =
|
|
1188
|
+
actor_weight_decay = 0.,
|
|
1189
|
+
critic_weight_decay = 0.,
|
|
1165
1190
|
diversity_aux_loss_weight = 0.,
|
|
1166
1191
|
use_critic_ema = True,
|
|
1167
1192
|
critic_ema_beta = 0.95,
|
|
@@ -1177,8 +1202,9 @@ class Agent(Module):
|
|
|
1177
1202
|
entropy_weight = .01,
|
|
1178
1203
|
norm_advantages = True
|
|
1179
1204
|
),
|
|
1205
|
+
clip_value = False,
|
|
1180
1206
|
critic_loss_kwargs: dict = dict(
|
|
1181
|
-
eps_clip = 0.
|
|
1207
|
+
eps_clip = 0.8
|
|
1182
1208
|
),
|
|
1183
1209
|
use_spo = False, # Simple Policy Optimization - Xie et al. https://arxiv.org/abs/2401.16025v9
|
|
1184
1210
|
use_improved_critic_loss = True,
|
|
@@ -1192,6 +1218,11 @@ class Agent(Module):
|
|
|
1192
1218
|
diversity_discr_kwargs: dict = dict(dim = 64, depth = 2),
|
|
1193
1219
|
diversity_discr_lr = 3e-4,
|
|
1194
1220
|
diversity_discr_optim_kwargs: dict = dict(),
|
|
1221
|
+
use_hidden_spr = False, # self-predictive representations over the actor's hidden embedding - https://arxiv.org/abs/2106.04799
|
|
1222
|
+
hidden_spr_dim_action = 32,
|
|
1223
|
+
hidden_spr_lr = 3e-4,
|
|
1224
|
+
hidden_spr_weight = 1.0,
|
|
1225
|
+
hidden_spr_ema_update = 0.99,
|
|
1195
1226
|
get_fitness_scores: Callable[..., Tensor] = get_fitness_scores,
|
|
1196
1227
|
wrap_with_accelerate: bool = True,
|
|
1197
1228
|
accelerate_kwargs: dict = dict(),
|
|
@@ -1290,6 +1321,28 @@ class Agent(Module):
|
|
|
1290
1321
|
self.diversity_discr = None
|
|
1291
1322
|
self.diversity_discr_optim = None
|
|
1292
1323
|
|
|
1324
|
+
self.clip_value = clip_value
|
|
1325
|
+
self.critic_loss_kwargs = critic_loss_kwargs
|
|
1326
|
+
|
|
1327
|
+
# self-predictive representations (SPR)
|
|
1328
|
+
|
|
1329
|
+
self.use_hidden_spr = use_hidden_spr
|
|
1330
|
+
self.hidden_spr_weight = hidden_spr_weight
|
|
1331
|
+
|
|
1332
|
+
if use_hidden_spr:
|
|
1333
|
+
self.hidden_spr = HiddenSpr(
|
|
1334
|
+
actor,
|
|
1335
|
+
dim_hidden = actor.dim,
|
|
1336
|
+
num_actions = actor.num_actions,
|
|
1337
|
+
dim_action = hidden_spr_dim_action,
|
|
1338
|
+
)
|
|
1339
|
+
|
|
1340
|
+
self.hidden_spr_optim = optim_klass(self.hidden_spr.online_parameters(), lr = hidden_spr_lr)
|
|
1341
|
+
self.hidden_spr_ema_update = hidden_spr_ema_update
|
|
1342
|
+
else:
|
|
1343
|
+
self.hidden_spr = None
|
|
1344
|
+
self.hidden_spr_optim = None
|
|
1345
|
+
|
|
1293
1346
|
self.register_buffer('has_diversity_discr_warmed_up', tensor(False))
|
|
1294
1347
|
self.register_buffer('zero', tensor(0.))
|
|
1295
1348
|
|
|
@@ -1341,6 +1394,11 @@ class Agent(Module):
|
|
|
1341
1394
|
self.diversity_discr, self.diversity_discr_optim
|
|
1342
1395
|
)
|
|
1343
1396
|
|
|
1397
|
+
if exists(self.hidden_spr):
|
|
1398
|
+
self.hidden_spr, self.hidden_spr_optim = self.accelerate.prepare(
|
|
1399
|
+
self.hidden_spr, self.hidden_spr_optim
|
|
1400
|
+
)
|
|
1401
|
+
|
|
1344
1402
|
if exists(self.latent_optim):
|
|
1345
1403
|
self.latent_optim = self.accelerate.prepare(self.latent_optim)
|
|
1346
1404
|
|
|
@@ -1373,6 +1431,7 @@ class Agent(Module):
|
|
|
1373
1431
|
|
|
1374
1432
|
def save(self, path, overwrite = False):
|
|
1375
1433
|
path = Path(path)
|
|
1434
|
+
path.parent.mkdir(parents = True, exist_ok = True)
|
|
1376
1435
|
unwrap = self.unwrap_model
|
|
1377
1436
|
unwrap_optim = lambda opt: opt.optimizer if hasattr(opt, 'optimizer') else opt
|
|
1378
1437
|
|
|
@@ -1389,6 +1448,8 @@ class Agent(Module):
|
|
|
1389
1448
|
critic_optim = unwrap_optim(self.critic_optim).state_dict(),
|
|
1390
1449
|
latent_optim = unwrap_optim(self.latent_optim).state_dict() if exists(self.latent_optim) else None,
|
|
1391
1450
|
diversity_discr_optim = unwrap_optim(self.diversity_discr_optim).state_dict() if self.use_diversity_discr else None,
|
|
1451
|
+
hidden_spr = unwrap(self.hidden_spr).state_dict() if self.use_hidden_spr else None,
|
|
1452
|
+
hidden_spr_optim = unwrap_optim(self.hidden_spr_optim).state_dict() if self.use_hidden_spr else None,
|
|
1392
1453
|
)
|
|
1393
1454
|
|
|
1394
1455
|
torch.save(pkg, str(path))
|
|
@@ -1409,7 +1470,7 @@ class Agent(Module):
|
|
|
1409
1470
|
self.critic_ema.load_state_dict(pkg['critic_ema'])
|
|
1410
1471
|
|
|
1411
1472
|
if 'latents' in pkg and exists(pkg['latents']):
|
|
1412
|
-
self.latent_gene_pool.load_state_dict(pkg['latents'])
|
|
1473
|
+
unwrap(self.latent_gene_pool).load_state_dict(pkg['latents'])
|
|
1413
1474
|
|
|
1414
1475
|
if self.use_diversity_discr and 'diversity_discr' in pkg and exists(pkg['diversity_discr']):
|
|
1415
1476
|
unwrap(self.diversity_discr).load_state_dict(pkg['diversity_discr'])
|
|
@@ -1417,6 +1478,11 @@ class Agent(Module):
|
|
|
1417
1478
|
if 'has_diversity_discr_warmed_up' in pkg:
|
|
1418
1479
|
self.has_diversity_discr_warmed_up.copy_(tensor(pkg['has_diversity_discr_warmed_up']))
|
|
1419
1480
|
|
|
1481
|
+
if self.use_hidden_spr and 'hidden_spr' in pkg and exists(pkg['hidden_spr']):
|
|
1482
|
+
unwrap(self.hidden_spr).load_state_dict(pkg['hidden_spr'])
|
|
1483
|
+
elif self.use_hidden_spr:
|
|
1484
|
+
unwrap(self.hidden_spr).sync_targets_(unwrap(self.actor))
|
|
1485
|
+
|
|
1420
1486
|
unwrap_optim(self.actor_optim).load_state_dict(pkg['actor_optim'])
|
|
1421
1487
|
unwrap_optim(self.critic_optim).load_state_dict(pkg['critic_optim'])
|
|
1422
1488
|
|
|
@@ -1426,6 +1492,9 @@ class Agent(Module):
|
|
|
1426
1492
|
if self.use_diversity_discr and 'diversity_discr_optim' in pkg and exists(pkg['diversity_discr_optim']):
|
|
1427
1493
|
unwrap_optim(self.diversity_discr_optim).load_state_dict(pkg['diversity_discr_optim'])
|
|
1428
1494
|
|
|
1495
|
+
if self.use_hidden_spr and 'hidden_spr_optim' in pkg and exists(pkg['hidden_spr_optim']):
|
|
1496
|
+
unwrap_optim(self.hidden_spr_optim).load_state_dict(pkg['hidden_spr_optim'])
|
|
1497
|
+
|
|
1429
1498
|
@move_input_tensors_to_device
|
|
1430
1499
|
def get_actor_distribution(
|
|
1431
1500
|
self,
|
|
@@ -1556,7 +1625,7 @@ class Agent(Module):
|
|
|
1556
1625
|
|
|
1557
1626
|
valid_episode = episode_ids >= 0
|
|
1558
1627
|
|
|
1559
|
-
dataset = TensorDataset(*[t[valid_episode] for t in (advantages, states, next_states, latent_gene_ids, actions, log_probs, values)])
|
|
1628
|
+
dataset = TensorDataset(*[t[valid_episode] for t in (advantages, states, next_states, latent_gene_ids, actions, log_probs, values, dones)])
|
|
1560
1629
|
|
|
1561
1630
|
dataloader = DataLoader(dataset, batch_size = self.batch_size, shuffle = True)
|
|
1562
1631
|
|
|
@@ -1576,7 +1645,8 @@ class Agent(Module):
|
|
|
1576
1645
|
latent_gene_ids,
|
|
1577
1646
|
actions,
|
|
1578
1647
|
log_probs,
|
|
1579
|
-
old_values
|
|
1648
|
+
old_values,
|
|
1649
|
+
dones
|
|
1580
1650
|
) in dataloader:
|
|
1581
1651
|
|
|
1582
1652
|
if self.has_latent_genes:
|
|
@@ -1597,6 +1667,24 @@ class Agent(Module):
|
|
|
1597
1667
|
use_spo = self.use_spo
|
|
1598
1668
|
)
|
|
1599
1669
|
|
|
1670
|
+
# self-predictive representation - predict the next hidden from
|
|
1671
|
+
# the current hidden and action, against the ema target of the
|
|
1672
|
+
# next state - only over non-terminal transitions
|
|
1673
|
+
|
|
1674
|
+
if self.use_hidden_spr:
|
|
1675
|
+
spr = self.unwrap_model(self.hidden_spr)
|
|
1676
|
+
|
|
1677
|
+
hidden = self.unwrap_model(self.actor).latent(states, latents)
|
|
1678
|
+
predicted = spr.predict(hidden, actions)
|
|
1679
|
+
target = spr.target(next_states, latents)
|
|
1680
|
+
|
|
1681
|
+
transitions = ~dones.bool()
|
|
1682
|
+
|
|
1683
|
+
if transitions.any():
|
|
1684
|
+
spr_loss = (2. - F.cosine_similarity(predicted[transitions], target[transitions], dim = -1)).mean()
|
|
1685
|
+
|
|
1686
|
+
actor_loss = actor_loss + self.hidden_spr_weight * spr_loss
|
|
1687
|
+
|
|
1600
1688
|
actor_loss.backward()
|
|
1601
1689
|
|
|
1602
1690
|
if exists(self.has_grad_clip):
|
|
@@ -1605,6 +1693,10 @@ class Agent(Module):
|
|
|
1605
1693
|
self.actor_optim.step()
|
|
1606
1694
|
self.actor_optim.zero_grad()
|
|
1607
1695
|
|
|
1696
|
+
if self.use_hidden_spr:
|
|
1697
|
+
self.hidden_spr_optim.step()
|
|
1698
|
+
self.hidden_spr_optim.zero_grad()
|
|
1699
|
+
|
|
1608
1700
|
# learn critic with maybe classification loss
|
|
1609
1701
|
|
|
1610
1702
|
critic_loss = self.unwrap_model(self.critic).forward_for_loss(
|
|
@@ -1612,6 +1704,7 @@ class Agent(Module):
|
|
|
1612
1704
|
latents,
|
|
1613
1705
|
old_values = old_values,
|
|
1614
1706
|
target = advantages + old_values,
|
|
1707
|
+
clip_value = self.clip_value,
|
|
1615
1708
|
use_improved = self.use_improved_critic_loss,
|
|
1616
1709
|
**self.critic_loss_kwargs
|
|
1617
1710
|
)
|
|
@@ -1691,6 +1784,11 @@ class Agent(Module):
|
|
|
1691
1784
|
if self.use_diversity_discr:
|
|
1692
1785
|
self.has_diversity_discr_warmed_up.copy_(tensor(True))
|
|
1693
1786
|
|
|
1787
|
+
# update the hidden spr ema targets once per learn_from
|
|
1788
|
+
|
|
1789
|
+
if self.use_hidden_spr:
|
|
1790
|
+
self.unwrap_model(self.hidden_spr).ema_update(self.unwrap_model(self.actor), self.hidden_spr_ema_update)
|
|
1791
|
+
|
|
1694
1792
|
# apply evolution
|
|
1695
1793
|
|
|
1696
1794
|
should_update = False
|
|
@@ -1850,6 +1948,61 @@ MemoriesAndCumulativeRewards = namedtuple('MemoriesAndCumulativeRewards', [
|
|
|
1850
1948
|
'cumulative_rewards' # Float['latent episodes']
|
|
1851
1949
|
])
|
|
1852
1950
|
|
|
1951
|
+
Slot = namedtuple('Slot', [
|
|
1952
|
+
'latent_id',
|
|
1953
|
+
'episode_id',
|
|
1954
|
+
'latent',
|
|
1955
|
+
'state',
|
|
1956
|
+
'time',
|
|
1957
|
+
'memories'
|
|
1958
|
+
])
|
|
1959
|
+
|
|
1960
|
+
# rollout of episodes for each latent can be parallelized across the workers
|
|
1961
|
+
# of an env implementing the vectorized interface (`num_envs`, `reset_one`,
|
|
1962
|
+
# `step_batch`) - a gymnasium-style env (reset / step) is adapted to the
|
|
1963
|
+
# vectorized interface with a single worker
|
|
1964
|
+
|
|
1965
|
+
def is_vectorized_env(env):
|
|
1966
|
+
return all(hasattr(env, attr) for attr in ('num_envs', 'reset_one', 'step_batch'))
|
|
1967
|
+
|
|
1968
|
+
class VectorizedEnvAdapter(Module):
|
|
1969
|
+
def __init__(
|
|
1970
|
+
self,
|
|
1971
|
+
env
|
|
1972
|
+
):
|
|
1973
|
+
super().__init__()
|
|
1974
|
+
self.env = env
|
|
1975
|
+
self.num_envs = 1
|
|
1976
|
+
|
|
1977
|
+
def reset_one(
|
|
1978
|
+
self,
|
|
1979
|
+
worker,
|
|
1980
|
+
seed = None
|
|
1981
|
+
):
|
|
1982
|
+
assert worker == 0
|
|
1983
|
+
|
|
1984
|
+
return self.env.reset(seed = seed)
|
|
1985
|
+
|
|
1986
|
+
def step_batch(
|
|
1987
|
+
self,
|
|
1988
|
+
actions,
|
|
1989
|
+
worker_ids = None
|
|
1990
|
+
):
|
|
1991
|
+
assert not exists(worker_ids) or list(worker_ids) == [0]
|
|
1992
|
+
|
|
1993
|
+
action = np.asarray(actions[0])
|
|
1994
|
+
next_state, reward, terminated, truncated, *_ = self.env.step(action)
|
|
1995
|
+
|
|
1996
|
+
return (
|
|
1997
|
+
np.asarray(next_state, dtype = np.float32)[None, ...],
|
|
1998
|
+
np.asarray(reward, dtype = np.float32)[None],
|
|
1999
|
+
np.asarray(terminated, dtype = bool)[None],
|
|
2000
|
+
np.asarray(truncated, dtype = bool)[None]
|
|
2001
|
+
)
|
|
2002
|
+
|
|
2003
|
+
def to_vectorized_env(env):
|
|
2004
|
+
return env if is_vectorized_env(env) else VectorizedEnvAdapter(env)
|
|
2005
|
+
|
|
1853
2006
|
class EPO(Module):
|
|
1854
2007
|
|
|
1855
2008
|
def __init__(
|
|
@@ -1930,110 +2083,136 @@ class EPO(Module):
|
|
|
1930
2083
|
memories: list[Memory] | None = None,
|
|
1931
2084
|
fix_environ_across_latents = None
|
|
1932
2085
|
) -> MemoriesAndCumulativeRewards:
|
|
2086
|
+
"""rollout episodes for each latent - parallelized across the workers
|
|
2087
|
+
of a vectorized env, or one episode at a time on a single worker for a
|
|
2088
|
+
gymnasium-style env. when an episode ends, the next one from the
|
|
2089
|
+
rollout generator is loaded into that worker"""
|
|
1933
2090
|
|
|
1934
2091
|
fix_environ_across_latents = default(fix_environ_across_latents, self.fix_environ_across_latents)
|
|
1935
2092
|
|
|
2093
|
+
env = to_vectorized_env(env)
|
|
2094
|
+
num_envs = env.num_envs
|
|
2095
|
+
|
|
1936
2096
|
self.agent.eval()
|
|
1937
2097
|
|
|
1938
|
-
invalid_episode = tensor(-1) #
|
|
2098
|
+
invalid_episode = tensor(-1) # bootstrap transitions carry this id, to be discarded when learning
|
|
2099
|
+
num_episodes = self.num_latents * self.episodes_per_latent
|
|
1939
2100
|
|
|
1940
|
-
if
|
|
1941
|
-
memories = []
|
|
2101
|
+
memories = memories if exists(memories) else []
|
|
1942
2102
|
|
|
1943
2103
|
rewards_per_latent_episode = torch.zeros((self.num_latents, self.episodes_per_latent), device = self.device)
|
|
1944
2104
|
|
|
1945
|
-
rollout_gen = self.rollouts_for_machine(fix_environ_across_latents)
|
|
1946
|
-
|
|
1947
|
-
for latent_id, episode_id, maybe_seed in tqdm(rollout_gen, desc = 'rollout', disable = self.agent.quiet or not self.agent.is_main_process):
|
|
1948
|
-
|
|
1949
|
-
time = 0
|
|
2105
|
+
rollout_gen = iter(self.rollouts_for_machine(fix_environ_across_latents))
|
|
1950
2106
|
|
|
1951
|
-
|
|
2107
|
+
# slots - the episode each worker is currently rolling out
|
|
1952
2108
|
|
|
1953
|
-
|
|
2109
|
+
slots: list[Slot | None] = [None] * num_envs
|
|
1954
2110
|
|
|
1955
|
-
|
|
1956
|
-
|
|
1957
|
-
|
|
1958
|
-
|
|
2111
|
+
def fill_slot(i):
|
|
2112
|
+
rollout = next(rollout_gen, None)
|
|
2113
|
+
if rollout is None:
|
|
2114
|
+
return
|
|
1959
2115
|
|
|
1960
|
-
|
|
2116
|
+
latent_id, episode_id, maybe_seed = rollout
|
|
1961
2117
|
|
|
1962
2118
|
latent = self.agent.unwrapped_latent_gene_pool(latent_id = latent_id) if self.agent.has_latent_genes else None
|
|
1963
2119
|
|
|
1964
|
-
|
|
2120
|
+
seed = maybe_seed if fix_environ_across_latents else None
|
|
2121
|
+
state, _ = interface_torch_numpy(env.reset_one, device = self.device)(i, seed = seed)
|
|
1965
2122
|
|
|
1966
|
-
|
|
2123
|
+
slots[i] = Slot(latent_id, episode_id, latent, state, 0, [])
|
|
1967
2124
|
|
|
1968
|
-
|
|
2125
|
+
for i in range(num_envs):
|
|
2126
|
+
fill_slot(i)
|
|
1969
2127
|
|
|
1970
|
-
|
|
2128
|
+
pbar = tqdm(total = num_episodes, desc = 'rollout', disable = self.agent.quiet or not self.agent.is_main_process)
|
|
1971
2129
|
|
|
1972
|
-
|
|
2130
|
+
# each iteration, batch all active slots through the actor and critic,
|
|
2131
|
+
# then step the env - slots that finish are refilled and the batch
|
|
2132
|
+
# stays full until all episodes have been rolled out
|
|
1973
2133
|
|
|
1974
|
-
|
|
2134
|
+
while num_episodes > 0:
|
|
2135
|
+
active = [i for i, slot in enumerate(slots) if exists(slot)]
|
|
1975
2136
|
|
|
1976
|
-
|
|
2137
|
+
obs = stack([slots[i].state for i in active])
|
|
2138
|
+
latents = stack([slots[i].latent for i in active]) if self.agent.has_latent_genes else None
|
|
1977
2139
|
|
|
1978
|
-
|
|
2140
|
+
actions, log_probs = self.agent.get_actor_actions(obs, latent = latents, sample = True, temperature = self.action_sample_temperature, use_unwrapped_model = True)
|
|
2141
|
+
values = self.agent.get_critic_values(obs, latent = latents, use_ema_if_available = True, use_unwrapped_model = True)
|
|
1979
2142
|
|
|
1980
|
-
|
|
2143
|
+
next_obs, rewards, terminated, truncated = env.step_batch(actions.cpu().numpy(), worker_ids = active)[:4]
|
|
1981
2144
|
|
|
1982
|
-
|
|
2145
|
+
rewards = from_numpy(np.asarray(rewards)).float().to(self.device)
|
|
2146
|
+
terminated = from_numpy(np.asarray(terminated)).to(self.device)
|
|
2147
|
+
truncated = from_numpy(np.asarray(truncated)).to(self.device)
|
|
1983
2148
|
|
|
1984
|
-
|
|
1985
|
-
|
|
1986
|
-
diversity_discr = self.agent.unwrap_model(self.agent.diversity_discr)
|
|
1987
|
-
diversity_discr.eval()
|
|
2149
|
+
for k, worker in enumerate(active):
|
|
2150
|
+
slot = slots[worker]
|
|
1988
2151
|
|
|
1989
|
-
|
|
1990
|
-
log_probs = logits.log_softmax(dim=-1)
|
|
2152
|
+
latent_id, episode_id, latent, state = slot.latent_id, slot.episode_id, slot.latent, slot.state
|
|
1991
2153
|
|
|
1992
|
-
|
|
2154
|
+
next_state = from_numpy(np.array(next_obs[k])).float().to(self.device)
|
|
1993
2155
|
|
|
1994
|
-
|
|
2156
|
+
reward, terminated_at, truncated_at = rewards[k], terminated[k], truncated[k]
|
|
1995
2157
|
|
|
1996
|
-
|
|
2158
|
+
# maybe diversity reward for the latent
|
|
1997
2159
|
|
|
1998
|
-
|
|
2160
|
+
if self.agent.use_diversity_discr and self.agent.has_diversity_discr_warmed_up.item():
|
|
2161
|
+
with torch.no_grad():
|
|
2162
|
+
logits = self.agent.unwrap_model(self.agent.diversity_discr)(rearrange(state, '... -> 1 ...'), rearrange(next_state, '... -> 1 ...'))
|
|
1999
2163
|
|
|
2000
|
-
|
|
2164
|
+
latent_log_probs = logits.log_softmax(dim = -1)
|
|
2165
|
+
diversity_reward = latent_log_probs[0, latent_id] + math.log(self.agent.num_latents)
|
|
2166
|
+
|
|
2167
|
+
reward = reward + self.diversity_reward_weight * diversity_reward.item()
|
|
2001
2168
|
|
|
2002
|
-
|
|
2169
|
+
done = truncated_at or terminated_at
|
|
2003
2170
|
|
|
2004
2171
|
memory = Memory(
|
|
2005
2172
|
tensor(episode_id),
|
|
2006
2173
|
state,
|
|
2007
2174
|
next_state,
|
|
2008
2175
|
tensor(latent_id),
|
|
2009
|
-
|
|
2010
|
-
|
|
2176
|
+
actions[k],
|
|
2177
|
+
log_probs[k],
|
|
2011
2178
|
reward,
|
|
2012
|
-
|
|
2013
|
-
|
|
2179
|
+
values[k],
|
|
2180
|
+
terminated_at
|
|
2014
2181
|
)
|
|
2015
2182
|
|
|
2016
2183
|
memory = Memory(*tuple(t.cpu() for t in memory))
|
|
2017
2184
|
|
|
2018
|
-
memories.append(memory)
|
|
2185
|
+
slot.memories.append(memory)
|
|
2019
2186
|
|
|
2020
|
-
|
|
2187
|
+
rewards_per_latent_episode[latent_id, episode_id] += reward
|
|
2021
2188
|
|
|
2022
|
-
|
|
2189
|
+
# episode is done - bootstrap value if truncated, then refill
|
|
2023
2190
|
|
|
2024
|
-
|
|
2025
|
-
|
|
2191
|
+
if not done and slot.time + 1 < self.max_episode_length:
|
|
2192
|
+
slots[worker] = slot._replace(state = next_state, time = slot.time + 1)
|
|
2193
|
+
continue
|
|
2026
2194
|
|
|
2027
|
-
|
|
2195
|
+
if not terminated_at:
|
|
2196
|
+
next_value = temp_batch_dim(self.agent.get_critic_values)(next_state, latent = latent, use_ema_if_available = True, use_unwrapped_model = True)
|
|
2028
2197
|
|
|
2029
|
-
|
|
2030
|
-
|
|
2031
|
-
|
|
2032
|
-
|
|
2033
|
-
|
|
2034
|
-
|
|
2198
|
+
memory = memory._replace(
|
|
2199
|
+
episode_id = invalid_episode,
|
|
2200
|
+
reward = next_value.cpu(),
|
|
2201
|
+
value = next_value.cpu(),
|
|
2202
|
+
done = tensor(True)
|
|
2203
|
+
)
|
|
2204
|
+
|
|
2205
|
+
slot.memories.append(memory)
|
|
2206
|
+
|
|
2207
|
+
memories.extend(slot.memories)
|
|
2208
|
+
|
|
2209
|
+
num_episodes -= 1
|
|
2210
|
+
pbar.update(1)
|
|
2211
|
+
|
|
2212
|
+
slots[worker] = None
|
|
2213
|
+
fill_slot(worker)
|
|
2035
2214
|
|
|
2036
|
-
|
|
2215
|
+
pbar.close()
|
|
2037
2216
|
|
|
2038
2217
|
return MemoriesAndCumulativeRewards(
|
|
2039
2218
|
memories = memories,
|
|
@@ -1,14 +1,14 @@
|
|
|
1
|
-
from random import uniform
|
|
2
1
|
from copy import deepcopy
|
|
2
|
+
from random import uniform
|
|
3
3
|
|
|
4
|
+
import einx
|
|
4
5
|
import torch
|
|
5
|
-
from torch import Tensor
|
|
6
6
|
import torch.nn.functional as F
|
|
7
|
-
from
|
|
7
|
+
from einops import rearrange, reduce, repeat
|
|
8
|
+
from torch import Tensor
|
|
9
|
+
from torch.func import functional_call, vmap
|
|
8
10
|
from torch.nn import Module, ParameterList
|
|
9
11
|
|
|
10
|
-
import einx
|
|
11
|
-
from einops import rearrange, reduce, repeat
|
|
12
12
|
|
|
13
13
|
def exists(v):
|
|
14
14
|
return v is not None
|
|
@@ -0,0 +1,128 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from random import choice
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import torch
|
|
7
|
+
from torch import randint, randn, tensor
|
|
8
|
+
from torch.nn import Module
|
|
9
|
+
|
|
10
|
+
# helpers
|
|
11
|
+
|
|
12
|
+
def cast_tuple(v):
|
|
13
|
+
return v if isinstance(v, tuple) else (v,)
|
|
14
|
+
|
|
15
|
+
def rng_or_none(seed):
|
|
16
|
+
return None if seed is None else int(seed)
|
|
17
|
+
|
|
18
|
+
# mock env
|
|
19
|
+
|
|
20
|
+
class Env(Module):
|
|
21
|
+
def __init__(
|
|
22
|
+
self,
|
|
23
|
+
state_shape: int | tuple[int, ...],
|
|
24
|
+
can_terminate_after = 2
|
|
25
|
+
):
|
|
26
|
+
super().__init__()
|
|
27
|
+
self.state_shape = cast_tuple(state_shape)
|
|
28
|
+
|
|
29
|
+
self.can_terminate_after = can_terminate_after
|
|
30
|
+
self.register_buffer('_step', tensor(0))
|
|
31
|
+
|
|
32
|
+
@property
|
|
33
|
+
def device(self):
|
|
34
|
+
return self._step.device
|
|
35
|
+
|
|
36
|
+
def reset(
|
|
37
|
+
self,
|
|
38
|
+
seed = None
|
|
39
|
+
):
|
|
40
|
+
state = randn(self.state_shape, device = self.device)
|
|
41
|
+
self._step.zero_()
|
|
42
|
+
return state.numpy(), None
|
|
43
|
+
|
|
44
|
+
def step(
|
|
45
|
+
self,
|
|
46
|
+
actions,
|
|
47
|
+
):
|
|
48
|
+
state = randn(self.state_shape, device = self.device)
|
|
49
|
+
reward = randint(0, 5, (), device = self.device).float()
|
|
50
|
+
|
|
51
|
+
if self._step > self.can_terminate_after:
|
|
52
|
+
truncated = tensor(choice((True, False)), device = self.device)
|
|
53
|
+
terminated = tensor(choice((True, False)), device = self.device)
|
|
54
|
+
else:
|
|
55
|
+
truncated = terminated = tensor(False, device = self.device)
|
|
56
|
+
|
|
57
|
+
self._step.add_(1)
|
|
58
|
+
|
|
59
|
+
out = (state, reward, terminated, truncated)
|
|
60
|
+
return (*tuple(t.numpy() for t in out), None)
|
|
61
|
+
|
|
62
|
+
# vectorized mock env - `num_envs` independent workers in one object, where
|
|
63
|
+
# each worker's trajectory is deterministic given its seed. satisfies the
|
|
64
|
+
# vectorized interface used by EPO for parallel rollout (`num_envs`,
|
|
65
|
+
# `reset_one`, `step_batch`), stepping only the listed workers
|
|
66
|
+
|
|
67
|
+
class VectorEnv(Module):
|
|
68
|
+
def __init__(
|
|
69
|
+
self,
|
|
70
|
+
state_shape: int | tuple[int, ...],
|
|
71
|
+
num_envs = 4,
|
|
72
|
+
can_terminate_after = 2
|
|
73
|
+
):
|
|
74
|
+
super().__init__()
|
|
75
|
+
self.state_shape = cast_tuple(state_shape)
|
|
76
|
+
self.num_envs = num_envs
|
|
77
|
+
self.can_terminate_after = can_terminate_after
|
|
78
|
+
|
|
79
|
+
# per-worker rng and step count, derived from the worker index
|
|
80
|
+
|
|
81
|
+
self._worker_steps = np.zeros(num_envs, dtype = np.int64)
|
|
82
|
+
self._worker_rngs = [np.random.default_rng(i) for i in range(num_envs)]
|
|
83
|
+
|
|
84
|
+
def _next_transition(self, i):
|
|
85
|
+
rng = self._worker_rngs[i]
|
|
86
|
+
state = rng.standard_normal(self.state_shape).astype(np.float32)
|
|
87
|
+
reward = float(rng.integers(0, 5))
|
|
88
|
+
|
|
89
|
+
if self._worker_steps[i] > self.can_terminate_after:
|
|
90
|
+
terminated = bool(rng.choice((True, False)))
|
|
91
|
+
truncated = bool(rng.choice((True, False)))
|
|
92
|
+
else:
|
|
93
|
+
terminated = truncated = False
|
|
94
|
+
|
|
95
|
+
self._worker_steps[i] += 1
|
|
96
|
+
return state, reward, terminated, truncated
|
|
97
|
+
|
|
98
|
+
def reset_one(
|
|
99
|
+
self,
|
|
100
|
+
i,
|
|
101
|
+
seed = None
|
|
102
|
+
):
|
|
103
|
+
self._worker_rngs[i] = np.random.default_rng(rng_or_none(seed))
|
|
104
|
+
self._worker_steps[i] = 0
|
|
105
|
+
|
|
106
|
+
state = self._worker_rngs[i].standard_normal(self.state_shape).astype(np.float32)
|
|
107
|
+
return state, None
|
|
108
|
+
|
|
109
|
+
def step_batch(
|
|
110
|
+
self,
|
|
111
|
+
actions,
|
|
112
|
+
worker_ids = None
|
|
113
|
+
):
|
|
114
|
+
worker_ids = np.arange(self.num_envs) if worker_ids is None else np.asarray(worker_ids)
|
|
115
|
+
actions = np.asarray(actions)
|
|
116
|
+
|
|
117
|
+
states = np.zeros((len(worker_ids),) + self.state_shape, dtype = np.float32)
|
|
118
|
+
rewards = np.zeros(len(worker_ids), dtype = np.float32)
|
|
119
|
+
terminated = np.zeros(len(worker_ids), dtype = bool)
|
|
120
|
+
truncated = np.zeros(len(worker_ids), dtype = bool)
|
|
121
|
+
|
|
122
|
+
for k, i in enumerate(worker_ids):
|
|
123
|
+
states[k], rewards[k], terminated[k], truncated[k] = self._next_transition(int(i))
|
|
124
|
+
|
|
125
|
+
return states, rewards, terminated, truncated
|
|
126
|
+
|
|
127
|
+
def close(self):
|
|
128
|
+
pass
|
{evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.5}/pyproject.toml
RENAMED
|
@@ -1,12 +1,12 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "evolutionary-policy-optimization"
|
|
3
|
-
version = "0.3.
|
|
3
|
+
version = "0.3.5"
|
|
4
4
|
description = "EPO - Pytorch"
|
|
5
5
|
authors = [
|
|
6
6
|
{ name = "Phil Wang", email = "lucidrains@gmail.com" }
|
|
7
7
|
]
|
|
8
8
|
readme = "README.md"
|
|
9
|
-
requires-python = ">= 3.
|
|
9
|
+
requires-python = ">= 3.10"
|
|
10
10
|
license = { file = "LICENSE" }
|
|
11
11
|
keywords = [
|
|
12
12
|
'artificial intelligence',
|
|
@@ -25,12 +25,12 @@ classifiers = [
|
|
|
25
25
|
|
|
26
26
|
dependencies = [
|
|
27
27
|
"accelerate>=1.6.0",
|
|
28
|
-
"adam-atan2-pytorch",
|
|
29
28
|
'assoc-scan>=0.0.2',
|
|
30
29
|
'einx>=0.3.0',
|
|
31
30
|
'einops>=0.8.1',
|
|
32
31
|
'ema-pytorch>=0.7.7',
|
|
33
32
|
'hl-gauss-pytorch>=0.2.8',
|
|
33
|
+
'mean-conc-beta>=0.0.7',
|
|
34
34
|
'torch>=2.2',
|
|
35
35
|
'tqdm',
|
|
36
36
|
"x-mlps-pytorch>=0.3.0",
|
|
@@ -1,18 +0,0 @@
|
|
|
1
|
-
from evolutionary_policy_optimization.epo import (
|
|
2
|
-
MLP,
|
|
3
|
-
Actor,
|
|
4
|
-
BetaActionDistr,
|
|
5
|
-
CategoricalActionDistr,
|
|
6
|
-
Critic,
|
|
7
|
-
create_agent,
|
|
8
|
-
Agent,
|
|
9
|
-
LatentGenePool,
|
|
10
|
-
EPO
|
|
11
|
-
)
|
|
12
|
-
|
|
13
|
-
from evolutionary_policy_optimization.mock_env import Env
|
|
14
|
-
|
|
15
|
-
from evolutionary_policy_optimization.env_wrappers import (
|
|
16
|
-
GymnasiumEnvWrapper,
|
|
17
|
-
rescale_from_to
|
|
18
|
-
)
|
|
@@ -1,55 +0,0 @@
|
|
|
1
|
-
from __future__ import annotations
|
|
2
|
-
from random import choice
|
|
3
|
-
|
|
4
|
-
import torch
|
|
5
|
-
from torch import tensor, randn, randint
|
|
6
|
-
from torch.nn import Module
|
|
7
|
-
|
|
8
|
-
# functions
|
|
9
|
-
|
|
10
|
-
def cast_tuple(v):
|
|
11
|
-
return v if isinstance(v, tuple) else (v,)
|
|
12
|
-
|
|
13
|
-
# mock env
|
|
14
|
-
|
|
15
|
-
class Env(Module):
|
|
16
|
-
def __init__(
|
|
17
|
-
self,
|
|
18
|
-
state_shape: int | tuple[int, ...],
|
|
19
|
-
can_terminate_after = 2
|
|
20
|
-
):
|
|
21
|
-
super().__init__()
|
|
22
|
-
self.state_shape = cast_tuple(state_shape)
|
|
23
|
-
|
|
24
|
-
self.can_terminate_after = can_terminate_after
|
|
25
|
-
self.register_buffer('_step', tensor(0))
|
|
26
|
-
|
|
27
|
-
@property
|
|
28
|
-
def device(self):
|
|
29
|
-
return self._step.device
|
|
30
|
-
|
|
31
|
-
def reset(
|
|
32
|
-
self,
|
|
33
|
-
seed = None
|
|
34
|
-
):
|
|
35
|
-
state = randn(self.state_shape, device = self.device)
|
|
36
|
-
self._step.zero_()
|
|
37
|
-
return state.numpy(), None
|
|
38
|
-
|
|
39
|
-
def step(
|
|
40
|
-
self,
|
|
41
|
-
actions,
|
|
42
|
-
):
|
|
43
|
-
state = randn(self.state_shape, device = self.device)
|
|
44
|
-
reward = randint(0, 5, (), device = self.device).float()
|
|
45
|
-
|
|
46
|
-
if self._step > self.can_terminate_after:
|
|
47
|
-
truncated = tensor(choice((True, False)), device =self.device)
|
|
48
|
-
terminated = tensor(choice((True, False)), device =self.device)
|
|
49
|
-
else:
|
|
50
|
-
truncated = terminated = tensor(False, device = self.device)
|
|
51
|
-
|
|
52
|
-
self._step.add_(1)
|
|
53
|
-
|
|
54
|
-
out = (state, reward, terminated, truncated)
|
|
55
|
-
return (*tuple(t.numpy() for t in out), None)
|
{evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.5}/.gitignore
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|