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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: evolutionary-policy-optimization
3
- Version: 0.3.0
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.9
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
- from torch import nn
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 = (0., 1.), to_range = (-1., 1.)):
9
- # e.g. beta actions on (0, 1) -> the env's action bounds
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
- # beta lives on (0, 1) - rescale to the env's bounds at the interface
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
- from torch.nn import Linear, Module, ModuleList
18
- from torch.utils.data import TensorDataset, DataLoader
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
-
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 adam_atan2_pytorch import AdoptAtan2
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 accelerate import Accelerator
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, bounded to (0, 1) by
565
- # construction - scale to the env's action range at the interface)
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
- class BetaActionDistr(Module):
557
+ # self-predictive representations (SPR) over the actor's hidden embedding
558
+
559
+ class HiddenSpr(Module):
575
560
  def __init__(
576
561
  self,
577
- init_conc = 2.,
578
- min_conc = 0.,
579
- eps = 1e-5
562
+ actor,
563
+ dim_hidden,
564
+ num_actions,
565
+ dim_action = 32,
580
566
  ):
581
567
  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
568
+ action_is_continuous = actor.beta_actions
569
+ self.action_is_continuous = action_is_continuous
587
570
 
588
- # softplus offset so the concentration at raw_conc = 0 is exactly init_conc
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.raw_init_conc = math.log(math.expm1(init_conc - min_conc))
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
- def mean(self, params):
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
- raw_mean, _ = params.unbind(dim = -1)
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
- def forward(self, params, temperature = 1.):
599
- _, raw_conc = params.unbind(dim = -1)
586
+ self.target_actor = deepcopy(actor).requires_grad_(False)
587
+ self.target_proj = deepcopy(self.proj_head).requires_grad_(False)
600
588
 
601
- mean = self.mean(params)
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
- conc = F.softplus(raw_conc + self.raw_init_conc) + self.min_conc
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
- # concentration floor - unimodal (alpha > 1 and beta > 1), mean kept exact
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
- conc = conc + 1. / torch.minimum(mean, 1. - mean).clamp(min = self.eps)
606
+ @torch.no_grad()
607
+ def target(self, state, latent):
608
+ return self.target_proj(self.target_actor.latent(state, latent))
608
609
 
609
- # temperature scales the concentration - lower temperature, sharper policy
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
- if temperature > 0. and temperature != 1.:
612
- conc = conc / temperature
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
- alpha = mean * conc
615
- beta = (1. - mean) * conc
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
- return _Beta(alpha, beta)
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 = BetaActionDistr() if self.beta_actions else CategoricalActionDistr()
672
+ self.action_distr = Beta(**beta_kwargs) if self.beta_actions else CategoricalActionDistr()
662
673
 
663
- def forward(
674
+ def latent(
664
675
  self,
665
676
  state,
666
- latent,
667
- temperature = 1.
668
- ) -> Distribution:
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
- hidden = self.mlp(hidden, latent)
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
- eps_clip = 0.4,
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 = AdoptAtan2,
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 = 5e-4,
1164
- critic_weight_decay = 5e-4,
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.4
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) # will use `episode_id` value of `-1` for the `next_value`, needed for not discarding last reward for generalized advantage estimate
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 not exists(memories):
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
- # initial state
2107
+ # slots - the episode each worker is currently rolling out
1952
2108
 
1953
- reset_kwargs = dict()
2109
+ slots: list[Slot | None] = [None] * num_envs
1954
2110
 
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)
2111
+ def fill_slot(i):
2112
+ rollout = next(rollout_gen, None)
2113
+ if rollout is None:
2114
+ return
1959
2115
 
1960
- # get latent from pool
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
- # until maximum episode length
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
- done = tensor(False)
2123
+ slots[i] = Slot(latent_id, episode_id, latent, state, 0, [])
1967
2124
 
1968
- while time < self.max_episode_length and not done:
2125
+ for i in range(num_envs):
2126
+ fill_slot(i)
1969
2127
 
1970
- # sample action
2128
+ pbar = tqdm(total = num_episodes, desc = 'rollout', disable = self.agent.quiet or not self.agent.is_main_process)
1971
2129
 
1972
- action, log_prob = temp_batch_dim(self.agent.get_actor_actions)(state, latent = latent, sample = True, temperature = self.action_sample_temperature, use_unwrapped_model = True)
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
- # values
2134
+ while num_episodes > 0:
2135
+ active = [i for i, slot in enumerate(slots) if exists(slot)]
1975
2136
 
1976
- value = temp_batch_dim(self.agent.get_critic_values)(state, latent = latent, use_ema_if_available = True, use_unwrapped_model = True)
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
- # get the next state, action, and reward
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
- next_state, reward, terminated, truncated, _ = interface_torch_numpy(env.step, device = self.device)(action)
2143
+ next_obs, rewards, terminated, truncated = env.step_batch(actions.cpu().numpy(), worker_ids = active)[:4]
1981
2144
 
1982
- # diversity reward
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
- if self.agent.use_diversity_discr and self.agent.has_diversity_discr_warmed_up.item():
1985
- with torch.no_grad():
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
- logits = diversity_discr(rearrange(state, '... -> 1 ...'), rearrange(next_state, '... -> 1 ...'))
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
- diversity_reward = log_probs[0, latent_id] + log(tensor(self.agent.num_latents))
2154
+ next_state = from_numpy(np.array(next_obs[k])).float().to(self.device)
1993
2155
 
1994
- reward = reward + diversity_reward * self.diversity_reward_weight
2156
+ reward, terminated_at, truncated_at = rewards[k], terminated[k], truncated[k]
1995
2157
 
1996
- done = truncated or terminated
2158
+ # maybe diversity reward for the latent
1997
2159
 
1998
- # update cumulative rewards per latent, to be used as default fitness score
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
- rewards_per_latent_episode[latent_id, episode_id] += reward
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
- # store memories
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
- action,
2010
- log_prob,
2176
+ actions[k],
2177
+ log_probs[k],
2011
2178
  reward,
2012
- value,
2013
- terminated
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
- state = next_state
2187
+ rewards_per_latent_episode[latent_id, episode_id] += reward
2021
2188
 
2022
- time += 1
2189
+ # episode is done - bootstrap value if truncated, then refill
2023
2190
 
2024
- if not terminated:
2025
- # add bootstrap value if truncated
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
- next_value = temp_batch_dim(self.agent.get_critic_values)(next_state, latent = latent, use_ema_if_available = True, use_unwrapped_model = True)
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
- memory_for_gae = memory._replace(
2030
- episode_id = invalid_episode,
2031
- reward = next_value.cpu(),
2032
- value = next_value.cpu(),
2033
- done = tensor(True)
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
- memories.append(memory_for_gae)
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 torch.func import vmap, functional_call
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
@@ -1,12 +1,12 @@
1
1
  [project]
2
2
  name = "evolutionary-policy-optimization"
3
- version = "0.3.0"
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"
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)