evolutionary-policy-optimization 0.3.0__tar.gz → 0.3.1__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.1}/PKG-INFO +1 -2
- evolutionary_policy_optimization-0.3.1/evolutionary_policy_optimization/__init__.py +3 -0
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.1}/evolutionary_policy_optimization/distributed.py +4 -5
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.1}/evolutionary_policy_optimization/env_wrappers.py +1 -0
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.1}/evolutionary_policy_optimization/epo.py +334 -106
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.1}/evolutionary_policy_optimization/experimental.py +5 -5
- evolutionary_policy_optimization-0.3.1/evolutionary_policy_optimization/mock_env.py +128 -0
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.1}/pyproject.toml +1 -2
- 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.1}/.gitignore +0 -0
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.1}/LICENSE +0 -0
- {evolutionary_policy_optimization-0.3.0 → evolutionary_policy_optimization-0.3.1}/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.1
|
|
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
|
|
@@ -35,7 +35,6 @@ Classifier: Programming Language :: Python :: 3.8
|
|
|
35
35
|
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
|
|
36
36
|
Requires-Python: >=3.9
|
|
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
|
|
@@ -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, BetaActionDistr, 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
|
|
@@ -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.utils._pytree import tree_map
|
|
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
|
-
|
|
37
|
-
from assoc_scan import AssocScan
|
|
38
|
-
|
|
16
|
+
import torch.nn.functional as F
|
|
17
|
+
from accelerate import Accelerator
|
|
39
18
|
from adam_atan2_pytorch import AdoptAtan2
|
|
40
|
-
|
|
41
|
-
from
|
|
42
|
-
|
|
19
|
+
from assoc_scan import AssocScan
|
|
20
|
+
from einops import einsum, rearrange, reduce, repeat
|
|
21
|
+
from einops.layers.torch import Rearrange
|
|
43
22
|
from ema_pytorch import EMA
|
|
23
|
+
from hl_gauss_pytorch import HLGaussLayer
|
|
24
|
+
from torch import Tensor, cat, from_numpy, is_tensor, nn, stack, tensor
|
|
25
|
+
from torch.distributions import Beta, Categorical, Distribution
|
|
44
26
|
|
|
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
|
+
# (0, 1) - rescale to the env's action range at the interface
|
|
566
549
|
|
|
567
550
|
class CategoricalActionDistr(Module):
|
|
568
551
|
def forward(self, logits, temperature = 1.):
|
|
@@ -595,6 +578,13 @@ class BetaActionDistr(Module):
|
|
|
595
578
|
raw_mean, _ = params.unbind(dim = -1)
|
|
596
579
|
return ((torch.tanh(raw_mean) + 1.) * 0.5).clamp(min = self.eps, max = 1. - self.eps)
|
|
597
580
|
|
|
581
|
+
def entropy(self, params, temperature = 1.):
|
|
582
|
+
# entropy is over the rescaled (2x - 1) in [-1, 1], with log 2 jacobian
|
|
583
|
+
# adjustment for the differential entropy of an affine transform
|
|
584
|
+
|
|
585
|
+
distr = self.forward(params, temperature = temperature)
|
|
586
|
+
return distr.entropy() + math.log(2.)
|
|
587
|
+
|
|
598
588
|
def forward(self, params, temperature = 1.):
|
|
599
589
|
_, raw_conc = params.unbind(dim = -1)
|
|
600
590
|
|
|
@@ -614,7 +604,77 @@ class BetaActionDistr(Module):
|
|
|
614
604
|
alpha = mean * conc
|
|
615
605
|
beta = (1. - mean) * conc
|
|
616
606
|
|
|
617
|
-
return
|
|
607
|
+
return Beta(alpha, beta)
|
|
608
|
+
|
|
609
|
+
# self-predictive representations (SPR) over the actor's hidden embedding
|
|
610
|
+
|
|
611
|
+
class HiddenSpr(Module):
|
|
612
|
+
def __init__(
|
|
613
|
+
self,
|
|
614
|
+
actor,
|
|
615
|
+
dim_hidden,
|
|
616
|
+
num_actions,
|
|
617
|
+
dim_action = 32,
|
|
618
|
+
):
|
|
619
|
+
super().__init__()
|
|
620
|
+
action_is_continuous = actor.beta_actions
|
|
621
|
+
self.action_is_continuous = action_is_continuous
|
|
622
|
+
|
|
623
|
+
if action_is_continuous:
|
|
624
|
+
self.action_proj = nn.Linear(num_actions, dim_action)
|
|
625
|
+
else:
|
|
626
|
+
self.action_proj = nn.Embedding(num_actions, dim_action)
|
|
627
|
+
|
|
628
|
+
self.to_dynamics = nn.Sequential(
|
|
629
|
+
nn.Linear(dim_hidden + dim_action, dim_hidden),
|
|
630
|
+
nn.SiLU(),
|
|
631
|
+
nn.Linear(dim_hidden, dim_hidden)
|
|
632
|
+
)
|
|
633
|
+
|
|
634
|
+
self.proj_head = nn.Linear(dim_hidden, dim_hidden, bias = False)
|
|
635
|
+
|
|
636
|
+
# ema target of the actor body + projection head
|
|
637
|
+
|
|
638
|
+
self.target_actor = deepcopy(actor).requires_grad_(False)
|
|
639
|
+
self.target_proj = deepcopy(self.proj_head).requires_grad_(False)
|
|
640
|
+
|
|
641
|
+
def online_parameters(self):
|
|
642
|
+
return [
|
|
643
|
+
*self.action_proj.parameters(),
|
|
644
|
+
*self.to_dynamics.parameters(),
|
|
645
|
+
*self.proj_head.parameters(),
|
|
646
|
+
]
|
|
647
|
+
|
|
648
|
+
def predict(self, hidden, action):
|
|
649
|
+
if not self.action_is_continuous:
|
|
650
|
+
action = action.long()
|
|
651
|
+
if action.ndim > 1:
|
|
652
|
+
action = rearrange(action, '... 1 -> ...')
|
|
653
|
+
|
|
654
|
+
action_proj = self.action_proj(action)
|
|
655
|
+
dynamics = self.to_dynamics(cat((hidden, action_proj), dim = -1))
|
|
656
|
+
return self.proj_head(dynamics + hidden)
|
|
657
|
+
|
|
658
|
+
@torch.no_grad()
|
|
659
|
+
def target(self, state, latent):
|
|
660
|
+
return self.target_proj(self.target_actor.latent(state, latent))
|
|
661
|
+
|
|
662
|
+
@torch.no_grad()
|
|
663
|
+
def _ema_update_one_(self, target_module, source_module, decay):
|
|
664
|
+
# if decay is 0, this is a hard copy of the source into the target
|
|
665
|
+
|
|
666
|
+
for target_param, source_param in zip(target_module.parameters(), source_module.parameters()):
|
|
667
|
+
target_param.lerp_(source_param, 1. - decay)
|
|
668
|
+
|
|
669
|
+
@torch.no_grad()
|
|
670
|
+
def ema_update(self, actor, decay):
|
|
671
|
+
self._ema_update_one_(self.target_actor.init_layer, actor.init_layer, decay)
|
|
672
|
+
self._ema_update_one_(self.target_actor.mlp, actor.mlp, decay)
|
|
673
|
+
self._ema_update_one_(self.target_proj, self.proj_head, decay)
|
|
674
|
+
|
|
675
|
+
@torch.no_grad()
|
|
676
|
+
def sync_targets_(self, actor):
|
|
677
|
+
self.ema_update(actor, decay = 0.)
|
|
618
678
|
|
|
619
679
|
# actor, critic, and agent (actor + critic)
|
|
620
680
|
# eventually, should just create a separate repo and aggregate all the MLP related architectures
|
|
@@ -634,6 +694,8 @@ class Actor(Module):
|
|
|
634
694
|
|
|
635
695
|
self.state_norm = state_norm
|
|
636
696
|
|
|
697
|
+
self.dim = dim
|
|
698
|
+
self.num_actions = num_actions
|
|
637
699
|
self.dim_latent = dim_latent
|
|
638
700
|
self.beta_actions = action_is_continuous
|
|
639
701
|
|
|
@@ -660,12 +722,14 @@ class Actor(Module):
|
|
|
660
722
|
|
|
661
723
|
self.action_distr = BetaActionDistr() if self.beta_actions else CategoricalActionDistr()
|
|
662
724
|
|
|
663
|
-
def
|
|
725
|
+
def latent(
|
|
664
726
|
self,
|
|
665
727
|
state,
|
|
666
|
-
latent
|
|
667
|
-
|
|
668
|
-
|
|
728
|
+
latent
|
|
729
|
+
) -> Tensor:
|
|
730
|
+
# hidden embedding right before the action projection -
|
|
731
|
+
# the dynamics in hidden_spr operate over this
|
|
732
|
+
|
|
669
733
|
if exists(self.state_norm):
|
|
670
734
|
with torch.no_grad():
|
|
671
735
|
self.state_norm.eval()
|
|
@@ -673,8 +737,15 @@ class Actor(Module):
|
|
|
673
737
|
|
|
674
738
|
hidden = self.init_layer(state)
|
|
675
739
|
|
|
676
|
-
|
|
740
|
+
return self.mlp(hidden, latent)
|
|
677
741
|
|
|
742
|
+
def forward(
|
|
743
|
+
self,
|
|
744
|
+
state,
|
|
745
|
+
latent,
|
|
746
|
+
temperature = 1.
|
|
747
|
+
) -> Distribution:
|
|
748
|
+
hidden = self.latent(state, latent)
|
|
678
749
|
return self.action_distr(self.to_out(hidden), temperature = temperature)
|
|
679
750
|
|
|
680
751
|
class Critic(Module):
|
|
@@ -724,9 +795,10 @@ class Critic(Module):
|
|
|
724
795
|
self,
|
|
725
796
|
state,
|
|
726
797
|
latent,
|
|
727
|
-
old_values,
|
|
728
798
|
target,
|
|
729
|
-
|
|
799
|
+
old_values = None,
|
|
800
|
+
clip_value = False,
|
|
801
|
+
eps_clip = 0.8,
|
|
730
802
|
use_improved = True
|
|
731
803
|
):
|
|
732
804
|
|
|
@@ -737,6 +809,9 @@ class Critic(Module):
|
|
|
737
809
|
|
|
738
810
|
logits = self.forward(state, latent, return_logits = True)
|
|
739
811
|
|
|
812
|
+
if not clip_value or not exists(old_values) or not exists(eps_clip):
|
|
813
|
+
return self.loss_fn(logits, target)
|
|
814
|
+
|
|
740
815
|
value = self.maybe_bins_to_value(logits)
|
|
741
816
|
|
|
742
817
|
loss_fn = partial(self.loss_fn, reduction = 'none')
|
|
@@ -866,7 +941,6 @@ class LatentGenePool(Module):
|
|
|
866
941
|
assert (frac_natural_selected + frac_elitism) < 1.
|
|
867
942
|
|
|
868
943
|
self.dim_latent = dim_latent
|
|
869
|
-
self.num_latents = num_latents
|
|
870
944
|
self.num_islands = num_islands
|
|
871
945
|
|
|
872
946
|
latents_per_island = num_latents // num_islands
|
|
@@ -1177,8 +1251,9 @@ class Agent(Module):
|
|
|
1177
1251
|
entropy_weight = .01,
|
|
1178
1252
|
norm_advantages = True
|
|
1179
1253
|
),
|
|
1254
|
+
clip_value = False,
|
|
1180
1255
|
critic_loss_kwargs: dict = dict(
|
|
1181
|
-
eps_clip = 0.
|
|
1256
|
+
eps_clip = 0.8
|
|
1182
1257
|
),
|
|
1183
1258
|
use_spo = False, # Simple Policy Optimization - Xie et al. https://arxiv.org/abs/2401.16025v9
|
|
1184
1259
|
use_improved_critic_loss = True,
|
|
@@ -1192,6 +1267,11 @@ class Agent(Module):
|
|
|
1192
1267
|
diversity_discr_kwargs: dict = dict(dim = 64, depth = 2),
|
|
1193
1268
|
diversity_discr_lr = 3e-4,
|
|
1194
1269
|
diversity_discr_optim_kwargs: dict = dict(),
|
|
1270
|
+
use_hidden_spr = False, # self-predictive representations over the actor's hidden embedding - https://arxiv.org/abs/2106.04799
|
|
1271
|
+
hidden_spr_dim_action = 32,
|
|
1272
|
+
hidden_spr_lr = 3e-4,
|
|
1273
|
+
hidden_spr_weight = 1.0,
|
|
1274
|
+
hidden_spr_ema_update = 0.99,
|
|
1195
1275
|
get_fitness_scores: Callable[..., Tensor] = get_fitness_scores,
|
|
1196
1276
|
wrap_with_accelerate: bool = True,
|
|
1197
1277
|
accelerate_kwargs: dict = dict(),
|
|
@@ -1290,6 +1370,28 @@ class Agent(Module):
|
|
|
1290
1370
|
self.diversity_discr = None
|
|
1291
1371
|
self.diversity_discr_optim = None
|
|
1292
1372
|
|
|
1373
|
+
self.clip_value = clip_value
|
|
1374
|
+
self.critic_loss_kwargs = critic_loss_kwargs
|
|
1375
|
+
|
|
1376
|
+
# self-predictive representations (SPR)
|
|
1377
|
+
|
|
1378
|
+
self.use_hidden_spr = use_hidden_spr
|
|
1379
|
+
self.hidden_spr_weight = hidden_spr_weight
|
|
1380
|
+
|
|
1381
|
+
if use_hidden_spr:
|
|
1382
|
+
self.hidden_spr = HiddenSpr(
|
|
1383
|
+
actor,
|
|
1384
|
+
dim_hidden = actor.dim,
|
|
1385
|
+
num_actions = actor.num_actions,
|
|
1386
|
+
dim_action = hidden_spr_dim_action,
|
|
1387
|
+
)
|
|
1388
|
+
|
|
1389
|
+
self.hidden_spr_optim = optim_klass(self.hidden_spr.online_parameters(), lr = hidden_spr_lr)
|
|
1390
|
+
self.hidden_spr_ema_update = hidden_spr_ema_update
|
|
1391
|
+
else:
|
|
1392
|
+
self.hidden_spr = None
|
|
1393
|
+
self.hidden_spr_optim = None
|
|
1394
|
+
|
|
1293
1395
|
self.register_buffer('has_diversity_discr_warmed_up', tensor(False))
|
|
1294
1396
|
self.register_buffer('zero', tensor(0.))
|
|
1295
1397
|
|
|
@@ -1341,6 +1443,11 @@ class Agent(Module):
|
|
|
1341
1443
|
self.diversity_discr, self.diversity_discr_optim
|
|
1342
1444
|
)
|
|
1343
1445
|
|
|
1446
|
+
if exists(self.hidden_spr):
|
|
1447
|
+
self.hidden_spr, self.hidden_spr_optim = self.accelerate.prepare(
|
|
1448
|
+
self.hidden_spr, self.hidden_spr_optim
|
|
1449
|
+
)
|
|
1450
|
+
|
|
1344
1451
|
if exists(self.latent_optim):
|
|
1345
1452
|
self.latent_optim = self.accelerate.prepare(self.latent_optim)
|
|
1346
1453
|
|
|
@@ -1373,6 +1480,7 @@ class Agent(Module):
|
|
|
1373
1480
|
|
|
1374
1481
|
def save(self, path, overwrite = False):
|
|
1375
1482
|
path = Path(path)
|
|
1483
|
+
path.parent.mkdir(parents = True, exist_ok = True)
|
|
1376
1484
|
unwrap = self.unwrap_model
|
|
1377
1485
|
unwrap_optim = lambda opt: opt.optimizer if hasattr(opt, 'optimizer') else opt
|
|
1378
1486
|
|
|
@@ -1389,6 +1497,8 @@ class Agent(Module):
|
|
|
1389
1497
|
critic_optim = unwrap_optim(self.critic_optim).state_dict(),
|
|
1390
1498
|
latent_optim = unwrap_optim(self.latent_optim).state_dict() if exists(self.latent_optim) else None,
|
|
1391
1499
|
diversity_discr_optim = unwrap_optim(self.diversity_discr_optim).state_dict() if self.use_diversity_discr else None,
|
|
1500
|
+
hidden_spr = unwrap(self.hidden_spr).state_dict() if self.use_hidden_spr else None,
|
|
1501
|
+
hidden_spr_optim = unwrap_optim(self.hidden_spr_optim).state_dict() if self.use_hidden_spr else None,
|
|
1392
1502
|
)
|
|
1393
1503
|
|
|
1394
1504
|
torch.save(pkg, str(path))
|
|
@@ -1409,7 +1519,7 @@ class Agent(Module):
|
|
|
1409
1519
|
self.critic_ema.load_state_dict(pkg['critic_ema'])
|
|
1410
1520
|
|
|
1411
1521
|
if 'latents' in pkg and exists(pkg['latents']):
|
|
1412
|
-
self.latent_gene_pool.load_state_dict(pkg['latents'])
|
|
1522
|
+
unwrap(self.latent_gene_pool).load_state_dict(pkg['latents'])
|
|
1413
1523
|
|
|
1414
1524
|
if self.use_diversity_discr and 'diversity_discr' in pkg and exists(pkg['diversity_discr']):
|
|
1415
1525
|
unwrap(self.diversity_discr).load_state_dict(pkg['diversity_discr'])
|
|
@@ -1417,6 +1527,11 @@ class Agent(Module):
|
|
|
1417
1527
|
if 'has_diversity_discr_warmed_up' in pkg:
|
|
1418
1528
|
self.has_diversity_discr_warmed_up.copy_(tensor(pkg['has_diversity_discr_warmed_up']))
|
|
1419
1529
|
|
|
1530
|
+
if self.use_hidden_spr and 'hidden_spr' in pkg and exists(pkg['hidden_spr']):
|
|
1531
|
+
unwrap(self.hidden_spr).load_state_dict(pkg['hidden_spr'])
|
|
1532
|
+
elif self.use_hidden_spr:
|
|
1533
|
+
unwrap(self.hidden_spr).sync_targets_(unwrap(self.actor))
|
|
1534
|
+
|
|
1420
1535
|
unwrap_optim(self.actor_optim).load_state_dict(pkg['actor_optim'])
|
|
1421
1536
|
unwrap_optim(self.critic_optim).load_state_dict(pkg['critic_optim'])
|
|
1422
1537
|
|
|
@@ -1426,6 +1541,9 @@ class Agent(Module):
|
|
|
1426
1541
|
if self.use_diversity_discr and 'diversity_discr_optim' in pkg and exists(pkg['diversity_discr_optim']):
|
|
1427
1542
|
unwrap_optim(self.diversity_discr_optim).load_state_dict(pkg['diversity_discr_optim'])
|
|
1428
1543
|
|
|
1544
|
+
if self.use_hidden_spr and 'hidden_spr_optim' in pkg and exists(pkg['hidden_spr_optim']):
|
|
1545
|
+
unwrap_optim(self.hidden_spr_optim).load_state_dict(pkg['hidden_spr_optim'])
|
|
1546
|
+
|
|
1429
1547
|
@move_input_tensors_to_device
|
|
1430
1548
|
def get_actor_distribution(
|
|
1431
1549
|
self,
|
|
@@ -1556,7 +1674,7 @@ class Agent(Module):
|
|
|
1556
1674
|
|
|
1557
1675
|
valid_episode = episode_ids >= 0
|
|
1558
1676
|
|
|
1559
|
-
dataset = TensorDataset(*[t[valid_episode] for t in (advantages, states, next_states, latent_gene_ids, actions, log_probs, values)])
|
|
1677
|
+
dataset = TensorDataset(*[t[valid_episode] for t in (advantages, states, next_states, latent_gene_ids, actions, log_probs, values, dones)])
|
|
1560
1678
|
|
|
1561
1679
|
dataloader = DataLoader(dataset, batch_size = self.batch_size, shuffle = True)
|
|
1562
1680
|
|
|
@@ -1576,7 +1694,8 @@ class Agent(Module):
|
|
|
1576
1694
|
latent_gene_ids,
|
|
1577
1695
|
actions,
|
|
1578
1696
|
log_probs,
|
|
1579
|
-
old_values
|
|
1697
|
+
old_values,
|
|
1698
|
+
dones
|
|
1580
1699
|
) in dataloader:
|
|
1581
1700
|
|
|
1582
1701
|
if self.has_latent_genes:
|
|
@@ -1597,6 +1716,24 @@ class Agent(Module):
|
|
|
1597
1716
|
use_spo = self.use_spo
|
|
1598
1717
|
)
|
|
1599
1718
|
|
|
1719
|
+
# self-predictive representation - predict the next hidden from
|
|
1720
|
+
# the current hidden and action, against the ema target of the
|
|
1721
|
+
# next state - only over non-terminal transitions
|
|
1722
|
+
|
|
1723
|
+
if self.use_hidden_spr:
|
|
1724
|
+
spr = self.unwrap_model(self.hidden_spr)
|
|
1725
|
+
|
|
1726
|
+
hidden = self.unwrap_model(self.actor).latent(states, latents)
|
|
1727
|
+
predicted = spr.predict(hidden, actions)
|
|
1728
|
+
target = spr.target(next_states, latents)
|
|
1729
|
+
|
|
1730
|
+
transitions = ~dones.bool()
|
|
1731
|
+
|
|
1732
|
+
if transitions.any():
|
|
1733
|
+
spr_loss = (2. - F.cosine_similarity(predicted[transitions], target[transitions], dim = -1)).mean()
|
|
1734
|
+
|
|
1735
|
+
actor_loss = actor_loss + self.hidden_spr_weight * spr_loss
|
|
1736
|
+
|
|
1600
1737
|
actor_loss.backward()
|
|
1601
1738
|
|
|
1602
1739
|
if exists(self.has_grad_clip):
|
|
@@ -1605,6 +1742,10 @@ class Agent(Module):
|
|
|
1605
1742
|
self.actor_optim.step()
|
|
1606
1743
|
self.actor_optim.zero_grad()
|
|
1607
1744
|
|
|
1745
|
+
if self.use_hidden_spr:
|
|
1746
|
+
self.hidden_spr_optim.step()
|
|
1747
|
+
self.hidden_spr_optim.zero_grad()
|
|
1748
|
+
|
|
1608
1749
|
# learn critic with maybe classification loss
|
|
1609
1750
|
|
|
1610
1751
|
critic_loss = self.unwrap_model(self.critic).forward_for_loss(
|
|
@@ -1612,6 +1753,7 @@ class Agent(Module):
|
|
|
1612
1753
|
latents,
|
|
1613
1754
|
old_values = old_values,
|
|
1614
1755
|
target = advantages + old_values,
|
|
1756
|
+
clip_value = self.clip_value,
|
|
1615
1757
|
use_improved = self.use_improved_critic_loss,
|
|
1616
1758
|
**self.critic_loss_kwargs
|
|
1617
1759
|
)
|
|
@@ -1691,6 +1833,11 @@ class Agent(Module):
|
|
|
1691
1833
|
if self.use_diversity_discr:
|
|
1692
1834
|
self.has_diversity_discr_warmed_up.copy_(tensor(True))
|
|
1693
1835
|
|
|
1836
|
+
# update the hidden spr ema targets once per learn_from
|
|
1837
|
+
|
|
1838
|
+
if self.use_hidden_spr:
|
|
1839
|
+
self.unwrap_model(self.hidden_spr).ema_update(self.unwrap_model(self.actor), self.hidden_spr_ema_update)
|
|
1840
|
+
|
|
1694
1841
|
# apply evolution
|
|
1695
1842
|
|
|
1696
1843
|
should_update = False
|
|
@@ -1850,6 +1997,61 @@ MemoriesAndCumulativeRewards = namedtuple('MemoriesAndCumulativeRewards', [
|
|
|
1850
1997
|
'cumulative_rewards' # Float['latent episodes']
|
|
1851
1998
|
])
|
|
1852
1999
|
|
|
2000
|
+
Slot = namedtuple('Slot', [
|
|
2001
|
+
'latent_id',
|
|
2002
|
+
'episode_id',
|
|
2003
|
+
'latent',
|
|
2004
|
+
'state',
|
|
2005
|
+
'time',
|
|
2006
|
+
'memories'
|
|
2007
|
+
])
|
|
2008
|
+
|
|
2009
|
+
# rollout of episodes for each latent can be parallelized across the workers
|
|
2010
|
+
# of an env implementing the vectorized interface (`num_envs`, `reset_one`,
|
|
2011
|
+
# `step_batch`) - a gymnasium-style env (reset / step) is adapted to the
|
|
2012
|
+
# vectorized interface with a single worker
|
|
2013
|
+
|
|
2014
|
+
def is_vectorized_env(env):
|
|
2015
|
+
return all(hasattr(env, attr) for attr in ('num_envs', 'reset_one', 'step_batch'))
|
|
2016
|
+
|
|
2017
|
+
class VectorizedEnvAdapter(Module):
|
|
2018
|
+
def __init__(
|
|
2019
|
+
self,
|
|
2020
|
+
env
|
|
2021
|
+
):
|
|
2022
|
+
super().__init__()
|
|
2023
|
+
self.env = env
|
|
2024
|
+
self.num_envs = 1
|
|
2025
|
+
|
|
2026
|
+
def reset_one(
|
|
2027
|
+
self,
|
|
2028
|
+
worker,
|
|
2029
|
+
seed = None
|
|
2030
|
+
):
|
|
2031
|
+
assert worker == 0
|
|
2032
|
+
|
|
2033
|
+
return self.env.reset(seed = seed)
|
|
2034
|
+
|
|
2035
|
+
def step_batch(
|
|
2036
|
+
self,
|
|
2037
|
+
actions,
|
|
2038
|
+
worker_ids = None
|
|
2039
|
+
):
|
|
2040
|
+
assert not exists(worker_ids) or list(worker_ids) == [0]
|
|
2041
|
+
|
|
2042
|
+
action = np.asarray(actions[0])
|
|
2043
|
+
next_state, reward, terminated, truncated, *_ = self.env.step(action)
|
|
2044
|
+
|
|
2045
|
+
return (
|
|
2046
|
+
np.asarray(next_state, dtype = np.float32)[None, ...],
|
|
2047
|
+
np.asarray(reward, dtype = np.float32)[None],
|
|
2048
|
+
np.asarray(terminated, dtype = bool)[None],
|
|
2049
|
+
np.asarray(truncated, dtype = bool)[None]
|
|
2050
|
+
)
|
|
2051
|
+
|
|
2052
|
+
def to_vectorized_env(env):
|
|
2053
|
+
return env if is_vectorized_env(env) else VectorizedEnvAdapter(env)
|
|
2054
|
+
|
|
1853
2055
|
class EPO(Module):
|
|
1854
2056
|
|
|
1855
2057
|
def __init__(
|
|
@@ -1930,110 +2132,136 @@ class EPO(Module):
|
|
|
1930
2132
|
memories: list[Memory] | None = None,
|
|
1931
2133
|
fix_environ_across_latents = None
|
|
1932
2134
|
) -> MemoriesAndCumulativeRewards:
|
|
2135
|
+
"""rollout episodes for each latent - parallelized across the workers
|
|
2136
|
+
of a vectorized env, or one episode at a time on a single worker for a
|
|
2137
|
+
gymnasium-style env. when an episode ends, the next one from the
|
|
2138
|
+
rollout generator is loaded into that worker"""
|
|
1933
2139
|
|
|
1934
2140
|
fix_environ_across_latents = default(fix_environ_across_latents, self.fix_environ_across_latents)
|
|
1935
2141
|
|
|
2142
|
+
env = to_vectorized_env(env)
|
|
2143
|
+
num_envs = env.num_envs
|
|
2144
|
+
|
|
1936
2145
|
self.agent.eval()
|
|
1937
2146
|
|
|
1938
|
-
invalid_episode = tensor(-1) #
|
|
2147
|
+
invalid_episode = tensor(-1) # bootstrap transitions carry this id, to be discarded when learning
|
|
2148
|
+
num_episodes = self.num_latents * self.episodes_per_latent
|
|
1939
2149
|
|
|
1940
|
-
if
|
|
1941
|
-
memories = []
|
|
2150
|
+
memories = memories if exists(memories) else []
|
|
1942
2151
|
|
|
1943
2152
|
rewards_per_latent_episode = torch.zeros((self.num_latents, self.episodes_per_latent), device = self.device)
|
|
1944
2153
|
|
|
1945
|
-
rollout_gen = self.rollouts_for_machine(fix_environ_across_latents)
|
|
2154
|
+
rollout_gen = iter(self.rollouts_for_machine(fix_environ_across_latents))
|
|
1946
2155
|
|
|
1947
|
-
|
|
2156
|
+
# slots - the episode each worker is currently rolling out
|
|
1948
2157
|
|
|
1949
|
-
|
|
2158
|
+
slots: list[Slot | None] = [None] * num_envs
|
|
1950
2159
|
|
|
1951
|
-
|
|
2160
|
+
def fill_slot(i):
|
|
2161
|
+
rollout = next(rollout_gen, None)
|
|
2162
|
+
if rollout is None:
|
|
2163
|
+
return
|
|
1952
2164
|
|
|
1953
|
-
|
|
1954
|
-
|
|
1955
|
-
if fix_environ_across_latents:
|
|
1956
|
-
reset_kwargs.update(seed = maybe_seed)
|
|
1957
|
-
|
|
1958
|
-
state, _ = interface_torch_numpy(env.reset, device = self.device)(**reset_kwargs)
|
|
1959
|
-
|
|
1960
|
-
# get latent from pool
|
|
2165
|
+
latent_id, episode_id, maybe_seed = rollout
|
|
1961
2166
|
|
|
1962
2167
|
latent = self.agent.unwrapped_latent_gene_pool(latent_id = latent_id) if self.agent.has_latent_genes else None
|
|
1963
2168
|
|
|
1964
|
-
|
|
2169
|
+
seed = maybe_seed if fix_environ_across_latents else None
|
|
2170
|
+
state, _ = interface_torch_numpy(env.reset_one, device = self.device)(i, seed = seed)
|
|
1965
2171
|
|
|
1966
|
-
|
|
2172
|
+
slots[i] = Slot(latent_id, episode_id, latent, state, 0, [])
|
|
1967
2173
|
|
|
1968
|
-
|
|
2174
|
+
for i in range(num_envs):
|
|
2175
|
+
fill_slot(i)
|
|
1969
2176
|
|
|
1970
|
-
|
|
2177
|
+
pbar = tqdm(total = num_episodes, desc = 'rollout', disable = self.agent.quiet or not self.agent.is_main_process)
|
|
1971
2178
|
|
|
1972
|
-
|
|
2179
|
+
# each iteration, batch all active slots through the actor and critic,
|
|
2180
|
+
# then step the env - slots that finish are refilled and the batch
|
|
2181
|
+
# stays full until all episodes have been rolled out
|
|
1973
2182
|
|
|
1974
|
-
|
|
2183
|
+
while num_episodes > 0:
|
|
2184
|
+
active = [i for i, slot in enumerate(slots) if exists(slot)]
|
|
1975
2185
|
|
|
1976
|
-
|
|
2186
|
+
obs = stack([slots[i].state for i in active])
|
|
2187
|
+
latents = stack([slots[i].latent for i in active]) if self.agent.has_latent_genes else None
|
|
1977
2188
|
|
|
1978
|
-
|
|
2189
|
+
actions, log_probs = self.agent.get_actor_actions(obs, latent = latents, sample = True, temperature = self.action_sample_temperature, use_unwrapped_model = True)
|
|
2190
|
+
values = self.agent.get_critic_values(obs, latent = latents, use_ema_if_available = True, use_unwrapped_model = True)
|
|
1979
2191
|
|
|
1980
|
-
|
|
2192
|
+
next_obs, rewards, terminated, truncated = env.step_batch(actions.cpu().numpy(), worker_ids = active)[:4]
|
|
1981
2193
|
|
|
1982
|
-
|
|
2194
|
+
rewards = from_numpy(np.asarray(rewards)).float().to(self.device)
|
|
2195
|
+
terminated = from_numpy(np.asarray(terminated)).to(self.device)
|
|
2196
|
+
truncated = from_numpy(np.asarray(truncated)).to(self.device)
|
|
1983
2197
|
|
|
1984
|
-
|
|
1985
|
-
|
|
1986
|
-
diversity_discr = self.agent.unwrap_model(self.agent.diversity_discr)
|
|
1987
|
-
diversity_discr.eval()
|
|
2198
|
+
for k, worker in enumerate(active):
|
|
2199
|
+
slot = slots[worker]
|
|
1988
2200
|
|
|
1989
|
-
|
|
1990
|
-
log_probs = logits.log_softmax(dim=-1)
|
|
2201
|
+
latent_id, episode_id, latent, state = slot.latent_id, slot.episode_id, slot.latent, slot.state
|
|
1991
2202
|
|
|
1992
|
-
|
|
2203
|
+
next_state = from_numpy(np.array(next_obs[k])).float().to(self.device)
|
|
1993
2204
|
|
|
1994
|
-
|
|
2205
|
+
reward, terminated_at, truncated_at = rewards[k], terminated[k], truncated[k]
|
|
1995
2206
|
|
|
1996
|
-
|
|
2207
|
+
# maybe diversity reward for the latent
|
|
1997
2208
|
|
|
1998
|
-
|
|
2209
|
+
if self.agent.use_diversity_discr and self.agent.has_diversity_discr_warmed_up.item():
|
|
2210
|
+
with torch.no_grad():
|
|
2211
|
+
logits = self.agent.unwrap_model(self.agent.diversity_discr)(rearrange(state, '... -> 1 ...'), rearrange(next_state, '... -> 1 ...'))
|
|
1999
2212
|
|
|
2000
|
-
|
|
2213
|
+
latent_log_probs = logits.log_softmax(dim = -1)
|
|
2214
|
+
diversity_reward = latent_log_probs[0, latent_id] + math.log(self.agent.num_latents)
|
|
2001
2215
|
|
|
2002
|
-
|
|
2216
|
+
reward = reward + self.diversity_reward_weight * diversity_reward.item()
|
|
2217
|
+
|
|
2218
|
+
done = truncated_at or terminated_at
|
|
2003
2219
|
|
|
2004
2220
|
memory = Memory(
|
|
2005
2221
|
tensor(episode_id),
|
|
2006
2222
|
state,
|
|
2007
2223
|
next_state,
|
|
2008
2224
|
tensor(latent_id),
|
|
2009
|
-
|
|
2010
|
-
|
|
2225
|
+
actions[k],
|
|
2226
|
+
log_probs[k],
|
|
2011
2227
|
reward,
|
|
2012
|
-
|
|
2013
|
-
|
|
2228
|
+
values[k],
|
|
2229
|
+
terminated_at
|
|
2014
2230
|
)
|
|
2015
2231
|
|
|
2016
2232
|
memory = Memory(*tuple(t.cpu() for t in memory))
|
|
2017
2233
|
|
|
2018
|
-
memories.append(memory)
|
|
2234
|
+
slot.memories.append(memory)
|
|
2019
2235
|
|
|
2020
|
-
|
|
2236
|
+
rewards_per_latent_episode[latent_id, episode_id] += reward
|
|
2021
2237
|
|
|
2022
|
-
|
|
2238
|
+
# episode is done - bootstrap value if truncated, then refill
|
|
2023
2239
|
|
|
2024
|
-
|
|
2025
|
-
|
|
2240
|
+
if not done and slot.time + 1 < self.max_episode_length:
|
|
2241
|
+
slots[worker] = slot._replace(state = next_state, time = slot.time + 1)
|
|
2242
|
+
continue
|
|
2026
2243
|
|
|
2027
|
-
|
|
2244
|
+
if not terminated_at:
|
|
2245
|
+
next_value = temp_batch_dim(self.agent.get_critic_values)(next_state, latent = latent, use_ema_if_available = True, use_unwrapped_model = True)
|
|
2028
2246
|
|
|
2029
|
-
|
|
2030
|
-
|
|
2031
|
-
|
|
2032
|
-
|
|
2033
|
-
|
|
2034
|
-
|
|
2247
|
+
memory = memory._replace(
|
|
2248
|
+
episode_id = invalid_episode,
|
|
2249
|
+
reward = next_value.cpu(),
|
|
2250
|
+
value = next_value.cpu(),
|
|
2251
|
+
done = tensor(True)
|
|
2252
|
+
)
|
|
2253
|
+
|
|
2254
|
+
slot.memories.append(memory)
|
|
2255
|
+
|
|
2256
|
+
memories.extend(slot.memories)
|
|
2257
|
+
|
|
2258
|
+
num_episodes -= 1
|
|
2259
|
+
pbar.update(1)
|
|
2260
|
+
|
|
2261
|
+
slots[worker] = None
|
|
2262
|
+
fill_slot(worker)
|
|
2035
2263
|
|
|
2036
|
-
|
|
2264
|
+
pbar.close()
|
|
2037
2265
|
|
|
2038
2266
|
return MemoriesAndCumulativeRewards(
|
|
2039
2267
|
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.1}/pyproject.toml
RENAMED
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "evolutionary-policy-optimization"
|
|
3
|
-
version = "0.3.
|
|
3
|
+
version = "0.3.1"
|
|
4
4
|
description = "EPO - Pytorch"
|
|
5
5
|
authors = [
|
|
6
6
|
{ name = "Phil Wang", email = "lucidrains@gmail.com" }
|
|
@@ -25,7 +25,6 @@ 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',
|
|
@@ -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.1}/.gitignore
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|