evolutionary-policy-optimization 0.2.22__tar.gz → 0.3.0__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.2.22 → evolutionary_policy_optimization-0.3.0}/PKG-INFO +5 -4
- {evolutionary_policy_optimization-0.2.22 → evolutionary_policy_optimization-0.3.0}/README.md +3 -2
- {evolutionary_policy_optimization-0.2.22 → evolutionary_policy_optimization-0.3.0}/evolutionary_policy_optimization/__init__.py +4 -1
- evolutionary_policy_optimization-0.3.0/evolutionary_policy_optimization/env_wrappers.py +70 -0
- {evolutionary_policy_optimization-0.2.22 → evolutionary_policy_optimization-0.3.0}/evolutionary_policy_optimization/epo.py +169 -64
- {evolutionary_policy_optimization-0.2.22 → evolutionary_policy_optimization-0.3.0}/evolutionary_policy_optimization/mock_env.py +1 -1
- {evolutionary_policy_optimization-0.2.22 → evolutionary_policy_optimization-0.3.0}/pyproject.toml +2 -2
- evolutionary_policy_optimization-0.2.22/evolutionary_policy_optimization/env_wrappers.py +0 -36
- {evolutionary_policy_optimization-0.2.22 → evolutionary_policy_optimization-0.3.0}/.gitignore +0 -0
- {evolutionary_policy_optimization-0.2.22 → evolutionary_policy_optimization-0.3.0}/LICENSE +0 -0
- {evolutionary_policy_optimization-0.2.22 → evolutionary_policy_optimization-0.3.0}/evolutionary_policy_optimization/distributed.py +0 -0
- {evolutionary_policy_optimization-0.2.22 → evolutionary_policy_optimization-0.3.0}/evolutionary_policy_optimization/experimental.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.5
|
|
2
2
|
Name: evolutionary-policy-optimization
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 0.3.0
|
|
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
|
|
@@ -40,7 +40,7 @@ Requires-Dist: assoc-scan>=0.0.2
|
|
|
40
40
|
Requires-Dist: einops>=0.8.1
|
|
41
41
|
Requires-Dist: einx>=0.3.0
|
|
42
42
|
Requires-Dist: ema-pytorch>=0.7.7
|
|
43
|
-
Requires-Dist: hl-gauss-pytorch>=0.
|
|
43
|
+
Requires-Dist: hl-gauss-pytorch>=0.2.8
|
|
44
44
|
Requires-Dist: torch>=2.2
|
|
45
45
|
Requires-Dist: tqdm
|
|
46
46
|
Requires-Dist: x-mlps-pytorch>=0.3.0
|
|
@@ -99,7 +99,8 @@ critic = Critic(dim_state = 32, dim = 256, mlp_depth = 3, dim_latent = 32)
|
|
|
99
99
|
|
|
100
100
|
latent = latent_pool(latent_id = 2)
|
|
101
101
|
|
|
102
|
-
|
|
102
|
+
action_distr = actor(state, latent)
|
|
103
|
+
actions = action_distr.sample()
|
|
103
104
|
value = critic(state, latent)
|
|
104
105
|
|
|
105
106
|
# interact with environment and receive rewards, termination etc
|
|
@@ -145,7 +146,7 @@ epo = EPO(
|
|
|
145
146
|
|
|
146
147
|
env = Env((512,))
|
|
147
148
|
|
|
148
|
-
epo(
|
|
149
|
+
epo(env, num_learning_cycles = 5)
|
|
149
150
|
|
|
150
151
|
# saving and loading
|
|
151
152
|
|
{evolutionary_policy_optimization-0.2.22 → evolutionary_policy_optimization-0.3.0}/README.md
RENAMED
|
@@ -37,7 +37,8 @@ critic = Critic(dim_state = 32, dim = 256, mlp_depth = 3, dim_latent = 32)
|
|
|
37
37
|
|
|
38
38
|
latent = latent_pool(latent_id = 2)
|
|
39
39
|
|
|
40
|
-
|
|
40
|
+
action_distr = actor(state, latent)
|
|
41
|
+
actions = action_distr.sample()
|
|
41
42
|
value = critic(state, latent)
|
|
42
43
|
|
|
43
44
|
# interact with environment and receive rewards, termination etc
|
|
@@ -83,7 +84,7 @@ epo = EPO(
|
|
|
83
84
|
|
|
84
85
|
env = Env((512,))
|
|
85
86
|
|
|
86
|
-
epo(
|
|
87
|
+
epo(env, num_learning_cycles = 5)
|
|
87
88
|
|
|
88
89
|
# saving and loading
|
|
89
90
|
|
|
@@ -1,6 +1,8 @@
|
|
|
1
1
|
from evolutionary_policy_optimization.epo import (
|
|
2
2
|
MLP,
|
|
3
3
|
Actor,
|
|
4
|
+
BetaActionDistr,
|
|
5
|
+
CategoricalActionDistr,
|
|
4
6
|
Critic,
|
|
5
7
|
create_agent,
|
|
6
8
|
Agent,
|
|
@@ -11,5 +13,6 @@ from evolutionary_policy_optimization.epo import (
|
|
|
11
13
|
from evolutionary_policy_optimization.mock_env import Env
|
|
12
14
|
|
|
13
15
|
from evolutionary_policy_optimization.env_wrappers import (
|
|
14
|
-
GymnasiumEnvWrapper
|
|
16
|
+
GymnasiumEnvWrapper,
|
|
17
|
+
rescale_from_to
|
|
15
18
|
)
|
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
from math import prod
|
|
2
|
+
|
|
3
|
+
import torch
|
|
4
|
+
from torch.nn import Module
|
|
5
|
+
|
|
6
|
+
from evolutionary_policy_optimization.epo import Agent, create_agent, exists
|
|
7
|
+
|
|
8
|
+
def rescale_from_to(x, from_range = (0., 1.), to_range = (-1., 1.)):
|
|
9
|
+
# e.g. beta actions on (0, 1) -> the env's action bounds
|
|
10
|
+
|
|
11
|
+
from_low, from_high = from_range
|
|
12
|
+
to_low, to_high = to_range
|
|
13
|
+
|
|
14
|
+
if torch.is_tensor(x):
|
|
15
|
+
dd_kwargs = dict(device = x.device, dtype = x.dtype)
|
|
16
|
+
from_low, from_high, to_low, to_high = [torch.as_tensor(t, **dd_kwargs) for t in (from_low, from_high, to_low, to_high)]
|
|
17
|
+
|
|
18
|
+
return to_low + (to_high - to_low) * (x - from_low) / (from_high - from_low)
|
|
19
|
+
|
|
20
|
+
class GymnasiumEnvWrapper(Module):
|
|
21
|
+
def __init__(
|
|
22
|
+
self,
|
|
23
|
+
env,
|
|
24
|
+
rescale_to = None # e.g. (-2., 2.) for pendulum's torque bounds
|
|
25
|
+
):
|
|
26
|
+
super().__init__()
|
|
27
|
+
self.env = env
|
|
28
|
+
|
|
29
|
+
if not exists(rescale_to) and not hasattr(env.action_space, 'n'):
|
|
30
|
+
rescale_to = (env.action_space.low, env.action_space.high)
|
|
31
|
+
|
|
32
|
+
self.rescale_to = rescale_to
|
|
33
|
+
|
|
34
|
+
def reset(self, *args, **kwargs):
|
|
35
|
+
return self.env.reset(*args, **kwargs)
|
|
36
|
+
|
|
37
|
+
def step(self, actions, *args, **kwargs):
|
|
38
|
+
# beta lives on (0, 1) - rescale to the env's bounds at the interface
|
|
39
|
+
|
|
40
|
+
if exists(self.rescale_to):
|
|
41
|
+
actions = rescale_from_to(actions, to_range = self.rescale_to)
|
|
42
|
+
|
|
43
|
+
return self.env.step(actions, *args, **kwargs)
|
|
44
|
+
|
|
45
|
+
def close(self, *args, **kwargs):
|
|
46
|
+
return self.env.close(*args, **kwargs)
|
|
47
|
+
|
|
48
|
+
def to_agent_hparams(self):
|
|
49
|
+
action_space = self.env.action_space
|
|
50
|
+
is_continuous = not hasattr(action_space, 'n')
|
|
51
|
+
|
|
52
|
+
num_actions = action_space.n if not is_continuous else prod(action_space.shape)
|
|
53
|
+
|
|
54
|
+
return dict(
|
|
55
|
+
dim_state = self.env.observation_space.shape[0],
|
|
56
|
+
actor_num_actions = num_actions,
|
|
57
|
+
action_is_continuous = is_continuous
|
|
58
|
+
)
|
|
59
|
+
|
|
60
|
+
def to_epo_agent(
|
|
61
|
+
self,
|
|
62
|
+
*args,
|
|
63
|
+
**kwargs
|
|
64
|
+
) -> Agent:
|
|
65
|
+
|
|
66
|
+
return create_agent(
|
|
67
|
+
*args,
|
|
68
|
+
**self.to_agent_hparams(),
|
|
69
|
+
**kwargs
|
|
70
|
+
)
|
|
@@ -1,12 +1,12 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
2
|
from typing import Callable
|
|
3
3
|
|
|
4
|
+
import math
|
|
4
5
|
from pathlib import Path
|
|
5
6
|
from math import ceil
|
|
6
7
|
from itertools import product
|
|
7
8
|
from functools import partial, wraps
|
|
8
9
|
from collections import namedtuple
|
|
9
|
-
from random import randrange
|
|
10
10
|
|
|
11
11
|
import numpy as np
|
|
12
12
|
|
|
@@ -19,9 +19,11 @@ from torch.utils.data import TensorDataset, DataLoader
|
|
|
19
19
|
from torch.utils._pytree import tree_map
|
|
20
20
|
|
|
21
21
|
import einx
|
|
22
|
-
from einops import rearrange, repeat, reduce, einsum
|
|
22
|
+
from einops import rearrange, repeat, reduce, einsum
|
|
23
23
|
from einops.layers.torch import Rearrange
|
|
24
24
|
|
|
25
|
+
from torch.distributions import Beta as _Beta, Categorical, Distribution
|
|
26
|
+
|
|
25
27
|
from x_mlps_pytorch import AttnResidualNormedMLP
|
|
26
28
|
|
|
27
29
|
from evolutionary_policy_optimization.distributed import (
|
|
@@ -80,11 +82,16 @@ def interface_torch_numpy(fn, device):
|
|
|
80
82
|
|
|
81
83
|
@maybe
|
|
82
84
|
def to_torch_tensor(t):
|
|
83
|
-
if isinstance(t,
|
|
85
|
+
if isinstance(t, np.ndarray):
|
|
84
86
|
t = from_numpy(np.array(t))
|
|
87
|
+
elif isinstance(t, np.generic):
|
|
88
|
+
t = tensor(t.item())
|
|
85
89
|
elif isinstance(t, (float, int, bool)):
|
|
86
90
|
t = tensor(t)
|
|
87
91
|
|
|
92
|
+
if is_tensor(t) and t.is_floating_point():
|
|
93
|
+
t = t.float()
|
|
94
|
+
|
|
88
95
|
return t.to(device)
|
|
89
96
|
|
|
90
97
|
@wraps(fn)
|
|
@@ -120,29 +127,18 @@ def batch_randperm(shape, device):
|
|
|
120
127
|
def log(t, eps = 1e-20):
|
|
121
128
|
return t.clamp(min = eps).log()
|
|
122
129
|
|
|
123
|
-
def
|
|
124
|
-
|
|
125
|
-
|
|
126
|
-
def gumbel_sample(t, temperature = 1.):
|
|
127
|
-
is_greedy = temperature <= 0.
|
|
130
|
+
def sum_to_batch(t):
|
|
131
|
+
# fold trailing action dims into one value per state
|
|
128
132
|
|
|
129
|
-
|
|
130
|
-
t = (t / temperature) + gumbel_noise(t)
|
|
133
|
+
return reduce(t, 'b ... -> b', 'sum')
|
|
131
134
|
|
|
132
|
-
|
|
135
|
+
def mode_action(distr):
|
|
136
|
+
# greedy action - argmax for categorical, mean for beta
|
|
133
137
|
|
|
134
|
-
|
|
135
|
-
|
|
136
|
-
return -(prob * log(prob)).sum(dim = -1)
|
|
138
|
+
if isinstance(distr, Categorical):
|
|
139
|
+
return distr.probs.argmax(dim = -1)
|
|
137
140
|
|
|
138
|
-
|
|
139
|
-
logits, # Float[b l]
|
|
140
|
-
indices # Int[b]
|
|
141
|
-
): # Float[b]
|
|
142
|
-
indices = rearrange(indices, '... -> ... 1')
|
|
143
|
-
log_probs = logits.log_softmax(dim = -1)
|
|
144
|
-
log_prob = log_probs.gather(-1, indices)
|
|
145
|
-
return rearrange(log_prob, '... 1 -> ...')
|
|
141
|
+
return distr.mean
|
|
146
142
|
|
|
147
143
|
def temp_batch_dim(fn):
|
|
148
144
|
|
|
@@ -564,6 +560,62 @@ class DiversityDiscr(Module):
|
|
|
564
560
|
next_state_embed = self.state_proj(next_state)
|
|
565
561
|
return self.net((state_embed, next_state_embed))
|
|
566
562
|
|
|
563
|
+
# action distributions - the actor always returns a `Distribution`, either
|
|
564
|
+
# categorical (discrete) or beta mean-conc (continuous, bounded to (0, 1) by
|
|
565
|
+
# construction - scale to the env's action range at the interface)
|
|
566
|
+
|
|
567
|
+
class CategoricalActionDistr(Module):
|
|
568
|
+
def forward(self, logits, temperature = 1.):
|
|
569
|
+
if temperature > 0. and temperature != 1.:
|
|
570
|
+
logits = logits / temperature
|
|
571
|
+
|
|
572
|
+
return Categorical(logits = logits)
|
|
573
|
+
|
|
574
|
+
class BetaActionDistr(Module):
|
|
575
|
+
def __init__(
|
|
576
|
+
self,
|
|
577
|
+
init_conc = 2.,
|
|
578
|
+
min_conc = 0.,
|
|
579
|
+
eps = 1e-5
|
|
580
|
+
):
|
|
581
|
+
super().__init__()
|
|
582
|
+
assert init_conc > min_conc
|
|
583
|
+
|
|
584
|
+
self.init_conc = init_conc
|
|
585
|
+
self.min_conc = min_conc
|
|
586
|
+
self.eps = eps
|
|
587
|
+
|
|
588
|
+
# softplus offset so the concentration at raw_conc = 0 is exactly init_conc
|
|
589
|
+
|
|
590
|
+
self.raw_init_conc = math.log(math.expm1(init_conc - min_conc))
|
|
591
|
+
|
|
592
|
+
def mean(self, params):
|
|
593
|
+
# the beta mean is exactly (tanh(raw_mean) + 1) / 2 by construction
|
|
594
|
+
|
|
595
|
+
raw_mean, _ = params.unbind(dim = -1)
|
|
596
|
+
return ((torch.tanh(raw_mean) + 1.) * 0.5).clamp(min = self.eps, max = 1. - self.eps)
|
|
597
|
+
|
|
598
|
+
def forward(self, params, temperature = 1.):
|
|
599
|
+
_, raw_conc = params.unbind(dim = -1)
|
|
600
|
+
|
|
601
|
+
mean = self.mean(params)
|
|
602
|
+
|
|
603
|
+
conc = F.softplus(raw_conc + self.raw_init_conc) + self.min_conc
|
|
604
|
+
|
|
605
|
+
# concentration floor - unimodal (alpha > 1 and beta > 1), mean kept exact
|
|
606
|
+
|
|
607
|
+
conc = conc + 1. / torch.minimum(mean, 1. - mean).clamp(min = self.eps)
|
|
608
|
+
|
|
609
|
+
# temperature scales the concentration - lower temperature, sharper policy
|
|
610
|
+
|
|
611
|
+
if temperature > 0. and temperature != 1.:
|
|
612
|
+
conc = conc / temperature
|
|
613
|
+
|
|
614
|
+
alpha = mean * conc
|
|
615
|
+
beta = (1. - mean) * conc
|
|
616
|
+
|
|
617
|
+
return _Beta(alpha, beta)
|
|
618
|
+
|
|
567
619
|
# actor, critic, and agent (actor + critic)
|
|
568
620
|
# eventually, should just create a separate repo and aggregate all the MLP related architectures
|
|
569
621
|
|
|
@@ -576,12 +628,14 @@ class Actor(Module):
|
|
|
576
628
|
mlp_depth,
|
|
577
629
|
state_norm: StateNorm | None = None,
|
|
578
630
|
dim_latent = 0,
|
|
631
|
+
action_is_continuous = False, # continuous control - beta policy
|
|
579
632
|
):
|
|
580
633
|
super().__init__()
|
|
581
634
|
|
|
582
635
|
self.state_norm = state_norm
|
|
583
636
|
|
|
584
637
|
self.dim_latent = dim_latent
|
|
638
|
+
self.beta_actions = action_is_continuous
|
|
585
639
|
|
|
586
640
|
self.init_layer = nn.Sequential(
|
|
587
641
|
nn.Linear(dim_state, dim),
|
|
@@ -590,16 +644,28 @@ class Actor(Module):
|
|
|
590
644
|
|
|
591
645
|
self.mlp = MLP(dim = dim, depth = mlp_depth, dim_latent = dim_latent)
|
|
592
646
|
|
|
593
|
-
self.
|
|
594
|
-
|
|
595
|
-
|
|
596
|
-
|
|
647
|
+
if self.beta_actions:
|
|
648
|
+
# beta head - (raw mean, raw concentration) per action dim
|
|
649
|
+
|
|
650
|
+
self.to_out = nn.Sequential(
|
|
651
|
+
nn.RMSNorm(dim),
|
|
652
|
+
nn.Linear(dim, num_actions * 2, bias = False),
|
|
653
|
+
Rearrange('... (d params) -> ... d params', params = 2),
|
|
654
|
+
)
|
|
655
|
+
else:
|
|
656
|
+
self.to_out = nn.Sequential(
|
|
657
|
+
nn.RMSNorm(dim),
|
|
658
|
+
nn.Linear(dim, num_actions, bias = False),
|
|
659
|
+
)
|
|
660
|
+
|
|
661
|
+
self.action_distr = BetaActionDistr() if self.beta_actions else CategoricalActionDistr()
|
|
597
662
|
|
|
598
663
|
def forward(
|
|
599
664
|
self,
|
|
600
665
|
state,
|
|
601
|
-
latent
|
|
602
|
-
|
|
666
|
+
latent,
|
|
667
|
+
temperature = 1.
|
|
668
|
+
) -> Distribution:
|
|
603
669
|
if exists(self.state_norm):
|
|
604
670
|
with torch.no_grad():
|
|
605
671
|
self.state_norm.eval()
|
|
@@ -609,7 +675,7 @@ class Actor(Module):
|
|
|
609
675
|
|
|
610
676
|
hidden = self.mlp(hidden, latent)
|
|
611
677
|
|
|
612
|
-
return self.to_out(hidden)
|
|
678
|
+
return self.action_distr(self.to_out(hidden), temperature = temperature)
|
|
613
679
|
|
|
614
680
|
class Critic(Module):
|
|
615
681
|
def __init__(
|
|
@@ -621,8 +687,8 @@ class Critic(Module):
|
|
|
621
687
|
use_regression = False,
|
|
622
688
|
state_norm: StateNorm | None = None,
|
|
623
689
|
hl_gauss_loss_kwargs: dict = dict(
|
|
624
|
-
min_value =
|
|
625
|
-
max_value =
|
|
690
|
+
min_value = 0.,
|
|
691
|
+
max_value = 500.,
|
|
626
692
|
num_bins = 250
|
|
627
693
|
)
|
|
628
694
|
):
|
|
@@ -715,27 +781,37 @@ class Critic(Module):
|
|
|
715
781
|
|
|
716
782
|
hidden = self.final_norm(hidden)
|
|
717
783
|
|
|
718
|
-
|
|
719
|
-
|
|
784
|
+
if self.use_regression:
|
|
785
|
+
return self.to_pred(hidden)
|
|
786
|
+
|
|
787
|
+
logits = self.to_pred(hidden, return_logits = True)
|
|
788
|
+
|
|
789
|
+
if return_logits:
|
|
790
|
+
return logits
|
|
791
|
+
|
|
792
|
+
value = self.maybe_bins_to_value(logits)
|
|
793
|
+
|
|
794
|
+
return value
|
|
720
795
|
|
|
721
796
|
# criteria for running genetic algorithm
|
|
722
797
|
|
|
723
798
|
class ShouldRunGeneticAlgorithm(Module):
|
|
724
799
|
def __init__(
|
|
725
800
|
self,
|
|
726
|
-
gamma =
|
|
801
|
+
gamma = 0.25, # fire when the spread exceeds this fraction of the fitness level
|
|
802
|
+
min_spread = 1.0, # absolute spread floor
|
|
727
803
|
):
|
|
728
804
|
super().__init__()
|
|
729
805
|
self.gamma = gamma
|
|
806
|
+
self.min_spread = min_spread
|
|
730
807
|
|
|
731
808
|
def forward(self, fitnesses):
|
|
732
|
-
#
|
|
809
|
+
# eq (3) - fire when the fitness spread is a meaningful fraction of the level
|
|
733
810
|
|
|
734
|
-
|
|
735
|
-
|
|
736
|
-
# just let it be customizable, and offer a variant where mean and variance is over some threshold (could account for skew too)
|
|
811
|
+
spread = fitnesses.amax(dim = -1) - fitnesses.amin(dim = -1)
|
|
812
|
+
scale = torch.abs(fitnesses.median(dim = -1).values)
|
|
737
813
|
|
|
738
|
-
return
|
|
814
|
+
return spread > (self.gamma * scale + self.min_spread)
|
|
739
815
|
|
|
740
816
|
# classes
|
|
741
817
|
|
|
@@ -756,7 +832,7 @@ class LatentGenePool(Module):
|
|
|
756
832
|
fast_genetic_algorithm = False,
|
|
757
833
|
fast_ga_values = torch.linspace(1, 5, 10),
|
|
758
834
|
should_run_genetic_algorithm: Module | None = None, # eq (3) in paper
|
|
759
|
-
default_should_run_ga_gamma =
|
|
835
|
+
default_should_run_ga_gamma = 0.25,
|
|
760
836
|
migrate_every = 100, # how many steps before a migration between islands
|
|
761
837
|
apply_genetic_algorithm_every = 2, # how many steps before crossover + mutation happens for genes
|
|
762
838
|
init_latent_fn: Callable | None = None
|
|
@@ -1120,9 +1196,12 @@ class Agent(Module):
|
|
|
1120
1196
|
wrap_with_accelerate: bool = True,
|
|
1121
1197
|
accelerate_kwargs: dict = dict(),
|
|
1122
1198
|
accelerator = None,
|
|
1199
|
+
quiet: bool = False,
|
|
1123
1200
|
):
|
|
1124
1201
|
super().__init__()
|
|
1125
1202
|
|
|
1203
|
+
self.quiet = quiet
|
|
1204
|
+
|
|
1126
1205
|
# hf accelerate
|
|
1127
1206
|
|
|
1128
1207
|
self.wrap_with_accelerate = wrap_with_accelerate
|
|
@@ -1348,28 +1427,47 @@ class Agent(Module):
|
|
|
1348
1427
|
unwrap_optim(self.diversity_discr_optim).load_state_dict(pkg['diversity_discr_optim'])
|
|
1349
1428
|
|
|
1350
1429
|
@move_input_tensors_to_device
|
|
1351
|
-
def
|
|
1430
|
+
def get_actor_distribution(
|
|
1352
1431
|
self,
|
|
1353
1432
|
state,
|
|
1354
1433
|
latent_id = None,
|
|
1355
1434
|
latent = None,
|
|
1356
|
-
sample = False,
|
|
1357
1435
|
temperature = 1.,
|
|
1358
1436
|
use_unwrapped_model = False
|
|
1359
|
-
):
|
|
1437
|
+
) -> Distribution:
|
|
1360
1438
|
maybe_unwrap = identity if not use_unwrapped_model else self.unwrap_model
|
|
1361
1439
|
|
|
1362
|
-
if not exists(latent) and exists(latent_id):
|
|
1440
|
+
if not exists(latent) and exists(latent_id) and exists(self.latent_gene_pool):
|
|
1363
1441
|
latent = maybe_unwrap(self.latent_gene_pool)(latent_id = latent_id)
|
|
1364
1442
|
|
|
1365
|
-
|
|
1443
|
+
return maybe_unwrap(self.actor)(state, latent, temperature = temperature)
|
|
1444
|
+
|
|
1445
|
+
@move_input_tensors_to_device
|
|
1446
|
+
def get_actor_actions(
|
|
1447
|
+
self,
|
|
1448
|
+
state,
|
|
1449
|
+
latent_id = None,
|
|
1450
|
+
latent = None,
|
|
1451
|
+
sample = False,
|
|
1452
|
+
temperature = 1.,
|
|
1453
|
+
use_unwrapped_model = False
|
|
1454
|
+
):
|
|
1455
|
+
distr = self.get_actor_distribution(
|
|
1456
|
+
state,
|
|
1457
|
+
latent_id = latent_id,
|
|
1458
|
+
latent = latent,
|
|
1459
|
+
temperature = temperature if sample else 1.,
|
|
1460
|
+
use_unwrapped_model = use_unwrapped_model
|
|
1461
|
+
)
|
|
1366
1462
|
|
|
1367
1463
|
if not sample:
|
|
1368
|
-
return
|
|
1464
|
+
return mode_action(distr)
|
|
1369
1465
|
|
|
1370
|
-
|
|
1466
|
+
# temperature <= 0 is greedy - mode works for both discrete and continuous
|
|
1371
1467
|
|
|
1372
|
-
|
|
1468
|
+
actions = mode_action(distr) if temperature <= 0. else distr.sample()
|
|
1469
|
+
|
|
1470
|
+
log_probs = sum_to_batch(distr.log_prob(actions))
|
|
1373
1471
|
|
|
1374
1472
|
return actions, log_probs
|
|
1375
1473
|
|
|
@@ -1385,7 +1483,7 @@ class Agent(Module):
|
|
|
1385
1483
|
|
|
1386
1484
|
maybe_unwrap = identity if not use_unwrapped_model else self.unwrap_model
|
|
1387
1485
|
|
|
1388
|
-
if not exists(latent) and exists(latent_id):
|
|
1486
|
+
if not exists(latent) and exists(latent_id) and exists(self.latent_gene_pool):
|
|
1389
1487
|
latent = maybe_unwrap(self.latent_gene_pool)(latent_id = latent_id)
|
|
1390
1488
|
|
|
1391
1489
|
critic_forward = maybe_unwrap(self.critic)
|
|
@@ -1470,7 +1568,7 @@ class Agent(Module):
|
|
|
1470
1568
|
self.actor.train()
|
|
1471
1569
|
self.critic.train()
|
|
1472
1570
|
|
|
1473
|
-
for _ in tqdm(range(epochs), desc = 'learning actor/critic epoch', disable = not self.is_main_process):
|
|
1571
|
+
for _ in tqdm(range(epochs), desc = 'learning actor/critic epoch', disable = self.quiet or not self.is_main_process):
|
|
1474
1572
|
for (
|
|
1475
1573
|
advantages,
|
|
1476
1574
|
states,
|
|
@@ -1492,9 +1590,12 @@ class Agent(Module):
|
|
|
1492
1590
|
|
|
1493
1591
|
# learn actor
|
|
1494
1592
|
|
|
1495
|
-
|
|
1593
|
+
distr = self.actor(states, latents)
|
|
1496
1594
|
|
|
1497
|
-
actor_loss = self.actor_loss(
|
|
1595
|
+
actor_loss = self.actor_loss(
|
|
1596
|
+
distr, log_probs, actions, advantages,
|
|
1597
|
+
use_spo = self.use_spo
|
|
1598
|
+
)
|
|
1498
1599
|
|
|
1499
1600
|
actor_loss.backward()
|
|
1500
1601
|
|
|
@@ -1582,7 +1683,7 @@ class Agent(Module):
|
|
|
1582
1683
|
if exists(self.state_norm):
|
|
1583
1684
|
self.state_norm.train()
|
|
1584
1685
|
|
|
1585
|
-
for _, states, *_ in tqdm(dataloader, desc = 'state norm learning', disable = not self.is_main_process):
|
|
1686
|
+
for _, states, *_ in tqdm(dataloader, desc = 'state norm learning', disable = self.quiet or not self.is_main_process):
|
|
1586
1687
|
self.state_norm(states)
|
|
1587
1688
|
|
|
1588
1689
|
# update warmed up state for discriminator
|
|
@@ -1592,6 +1693,8 @@ class Agent(Module):
|
|
|
1592
1693
|
|
|
1593
1694
|
# apply evolution
|
|
1594
1695
|
|
|
1696
|
+
should_update = False
|
|
1697
|
+
|
|
1595
1698
|
if self.has_latent_genes:
|
|
1596
1699
|
should_update, _ = self.latent_gene_pool.genetic_algorithm_step(fitness_scores)
|
|
1597
1700
|
|
|
@@ -1617,12 +1720,14 @@ class Agent(Module):
|
|
|
1617
1720
|
|
|
1618
1721
|
self.step.add_(1)
|
|
1619
1722
|
|
|
1723
|
+
return should_update, fitness_scores
|
|
1724
|
+
|
|
1620
1725
|
# reinforcement learning related - ppo
|
|
1621
1726
|
|
|
1622
1727
|
def actor_loss(
|
|
1623
|
-
|
|
1728
|
+
distr,
|
|
1624
1729
|
old_log_probs, # Float[b]
|
|
1625
|
-
actions, # Int[b]
|
|
1730
|
+
actions, # Int[b], or Float[b l] for beta
|
|
1626
1731
|
advantages, # Float[b]
|
|
1627
1732
|
eps_clip = 0.2,
|
|
1628
1733
|
entropy_weight = .01,
|
|
@@ -1630,9 +1735,10 @@ def actor_loss(
|
|
|
1630
1735
|
norm_advantages = True,
|
|
1631
1736
|
use_spo = False
|
|
1632
1737
|
):
|
|
1633
|
-
batch =
|
|
1738
|
+
batch = advantages.shape[0]
|
|
1634
1739
|
|
|
1635
|
-
log_probs =
|
|
1740
|
+
log_probs = sum_to_batch(distr.log_prob(actions))
|
|
1741
|
+
entropy = sum_to_batch(distr.entropy())
|
|
1636
1742
|
|
|
1637
1743
|
ratio = (log_probs - old_log_probs).exp()
|
|
1638
1744
|
|
|
@@ -1655,8 +1761,6 @@ def actor_loss(
|
|
|
1655
1761
|
|
|
1656
1762
|
# add entropy loss for exploration
|
|
1657
1763
|
|
|
1658
|
-
entropy = calc_entropy(logits)
|
|
1659
|
-
|
|
1660
1764
|
entropy_aux_loss = -entropy_weight * entropy
|
|
1661
1765
|
|
|
1662
1766
|
return (actor_loss + entropy_aux_loss).mean()
|
|
@@ -1675,6 +1779,7 @@ def create_agent(
|
|
|
1675
1779
|
critic_mlp_depth,
|
|
1676
1780
|
use_critic_ema = True,
|
|
1677
1781
|
use_state_norm = False,
|
|
1782
|
+
action_is_continuous = False, # continuous control - beta policy
|
|
1678
1783
|
latent_gene_pool_kwargs: dict = dict(),
|
|
1679
1784
|
actor_kwargs: dict = dict(),
|
|
1680
1785
|
critic_kwargs: dict = dict(),
|
|
@@ -1701,6 +1806,7 @@ def create_agent(
|
|
|
1701
1806
|
dim = actor_dim,
|
|
1702
1807
|
mlp_depth = actor_mlp_depth,
|
|
1703
1808
|
state_norm = state_norm,
|
|
1809
|
+
action_is_continuous = action_is_continuous,
|
|
1704
1810
|
**actor_kwargs
|
|
1705
1811
|
)
|
|
1706
1812
|
|
|
@@ -1838,7 +1944,7 @@ class EPO(Module):
|
|
|
1838
1944
|
|
|
1839
1945
|
rollout_gen = self.rollouts_for_machine(fix_environ_across_latents)
|
|
1840
1946
|
|
|
1841
|
-
for latent_id, episode_id, maybe_seed in tqdm(rollout_gen, desc = 'rollout', disable = not self.agent.is_main_process):
|
|
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):
|
|
1842
1948
|
|
|
1843
1949
|
time = 0
|
|
1844
1950
|
|
|
@@ -1871,7 +1977,7 @@ class EPO(Module):
|
|
|
1871
1977
|
|
|
1872
1978
|
# get the next state, action, and reward
|
|
1873
1979
|
|
|
1874
|
-
next_state, reward,
|
|
1980
|
+
next_state, reward, terminated, truncated, _ = interface_torch_numpy(env.step, device = self.device)(action)
|
|
1875
1981
|
|
|
1876
1982
|
# diversity reward
|
|
1877
1983
|
|
|
@@ -1936,7 +2042,6 @@ class EPO(Module):
|
|
|
1936
2042
|
|
|
1937
2043
|
def forward(
|
|
1938
2044
|
self,
|
|
1939
|
-
agent: Agent,
|
|
1940
2045
|
env,
|
|
1941
2046
|
num_learning_cycles,
|
|
1942
2047
|
seed = None
|
|
@@ -1946,10 +2051,10 @@ class EPO(Module):
|
|
|
1946
2051
|
torch.manual_seed(seed)
|
|
1947
2052
|
np.random.seed(seed)
|
|
1948
2053
|
|
|
1949
|
-
for _ in tqdm(range(num_learning_cycles), desc = 'learning cycle', disable = not self.agent.is_main_process):
|
|
2054
|
+
for _ in tqdm(range(num_learning_cycles), desc = 'learning cycle', disable = self.agent.quiet or not self.agent.is_main_process):
|
|
1950
2055
|
|
|
1951
2056
|
memories = self.gather_experience_from(env)
|
|
1952
2057
|
|
|
1953
|
-
agent.learn_from(memories)
|
|
2058
|
+
self.agent.learn_from(memories)
|
|
1954
2059
|
|
|
1955
2060
|
print('training complete')
|
{evolutionary_policy_optimization-0.2.22 → evolutionary_policy_optimization-0.3.0}/pyproject.toml
RENAMED
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "evolutionary-policy-optimization"
|
|
3
|
-
version = "0.
|
|
3
|
+
version = "0.3.0"
|
|
4
4
|
description = "EPO - Pytorch"
|
|
5
5
|
authors = [
|
|
6
6
|
{ name = "Phil Wang", email = "lucidrains@gmail.com" }
|
|
@@ -30,7 +30,7 @@ dependencies = [
|
|
|
30
30
|
'einx>=0.3.0',
|
|
31
31
|
'einops>=0.8.1',
|
|
32
32
|
'ema-pytorch>=0.7.7',
|
|
33
|
-
'hl-gauss-pytorch>=0.
|
|
33
|
+
'hl-gauss-pytorch>=0.2.8',
|
|
34
34
|
'torch>=2.2',
|
|
35
35
|
'tqdm',
|
|
36
36
|
"x-mlps-pytorch>=0.3.0",
|
|
@@ -1,36 +0,0 @@
|
|
|
1
|
-
import torch
|
|
2
|
-
from torch.nn import Module
|
|
3
|
-
|
|
4
|
-
from evolutionary_policy_optimization.epo import create_agent, Agent
|
|
5
|
-
|
|
6
|
-
class GymnasiumEnvWrapper(Module):
|
|
7
|
-
def __init__(
|
|
8
|
-
self,
|
|
9
|
-
env
|
|
10
|
-
):
|
|
11
|
-
super().__init__()
|
|
12
|
-
self.env = env
|
|
13
|
-
|
|
14
|
-
def reset(self, *args, **kwargs):
|
|
15
|
-
return self.env.reset(*args, **kwargs)
|
|
16
|
-
|
|
17
|
-
def step(self, *args, **kwargs):
|
|
18
|
-
return self.env.step(*args, **kwargs)
|
|
19
|
-
|
|
20
|
-
def to_agent_hparams(self):
|
|
21
|
-
return dict(
|
|
22
|
-
dim_state = self.env.observation_space.shape[0],
|
|
23
|
-
actor_num_actions = self.env.action_space.n
|
|
24
|
-
)
|
|
25
|
-
|
|
26
|
-
def to_epo_agent(
|
|
27
|
-
self,
|
|
28
|
-
*args,
|
|
29
|
-
**kwargs
|
|
30
|
-
) -> Agent:
|
|
31
|
-
|
|
32
|
-
return create_agent(
|
|
33
|
-
*args,
|
|
34
|
-
**self.to_agent_hparams(),
|
|
35
|
-
**kwargs
|
|
36
|
-
)
|
{evolutionary_policy_optimization-0.2.22 → evolutionary_policy_optimization-0.3.0}/.gitignore
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|