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.
@@ -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.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
- 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,6 +5,7 @@ from torch.nn import Module
5
5
 
6
6
  from evolutionary_policy_optimization.epo import Agent, create_agent, exists
7
7
 
8
+
8
9
  def rescale_from_to(x, from_range = (0., 1.), to_range = (-1., 1.)):
9
10
  # e.g. beta actions on (0, 1) -> the env's action bounds
10
11
 
@@ -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
-
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 hl_gauss_pytorch import HLGaussLayer
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 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
+ # (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 _Beta(alpha, beta)
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 forward(
725
+ def latent(
664
726
  self,
665
727
  state,
666
- latent,
667
- temperature = 1.
668
- ) -> Distribution:
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
- hidden = self.mlp(hidden, latent)
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
- eps_clip = 0.4,
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.4
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) # will use `episode_id` value of `-1` for the `next_value`, needed for not discarding last reward for generalized advantage estimate
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 not exists(memories):
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
- for latent_id, episode_id, maybe_seed in tqdm(rollout_gen, desc = 'rollout', disable = self.agent.quiet or not self.agent.is_main_process):
2156
+ # slots - the episode each worker is currently rolling out
1948
2157
 
1949
- time = 0
2158
+ slots: list[Slot | None] = [None] * num_envs
1950
2159
 
1951
- # initial state
2160
+ def fill_slot(i):
2161
+ rollout = next(rollout_gen, None)
2162
+ if rollout is None:
2163
+ return
1952
2164
 
1953
- reset_kwargs = dict()
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
- # until maximum episode length
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
- done = tensor(False)
2172
+ slots[i] = Slot(latent_id, episode_id, latent, state, 0, [])
1967
2173
 
1968
- while time < self.max_episode_length and not done:
2174
+ for i in range(num_envs):
2175
+ fill_slot(i)
1969
2176
 
1970
- # sample action
2177
+ pbar = tqdm(total = num_episodes, desc = 'rollout', disable = self.agent.quiet or not self.agent.is_main_process)
1971
2178
 
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)
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
- # values
2183
+ while num_episodes > 0:
2184
+ active = [i for i, slot in enumerate(slots) if exists(slot)]
1975
2185
 
1976
- value = temp_batch_dim(self.agent.get_critic_values)(state, latent = latent, use_ema_if_available = True, use_unwrapped_model = True)
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
- # get the next state, action, and reward
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
- next_state, reward, terminated, truncated, _ = interface_torch_numpy(env.step, device = self.device)(action)
2192
+ next_obs, rewards, terminated, truncated = env.step_batch(actions.cpu().numpy(), worker_ids = active)[:4]
1981
2193
 
1982
- # diversity reward
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
- 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()
2198
+ for k, worker in enumerate(active):
2199
+ slot = slots[worker]
1988
2200
 
1989
- logits = diversity_discr(rearrange(state, '... -> 1 ...'), rearrange(next_state, '... -> 1 ...'))
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
- diversity_reward = log_probs[0, latent_id] + log(tensor(self.agent.num_latents))
2203
+ next_state = from_numpy(np.array(next_obs[k])).float().to(self.device)
1993
2204
 
1994
- reward = reward + diversity_reward * self.diversity_reward_weight
2205
+ reward, terminated_at, truncated_at = rewards[k], terminated[k], truncated[k]
1995
2206
 
1996
- done = truncated or terminated
2207
+ # maybe diversity reward for the latent
1997
2208
 
1998
- # update cumulative rewards per latent, to be used as default fitness score
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
- rewards_per_latent_episode[latent_id, episode_id] += reward
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
- # store memories
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
- action,
2010
- log_prob,
2225
+ actions[k],
2226
+ log_probs[k],
2011
2227
  reward,
2012
- value,
2013
- terminated
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
- state = next_state
2236
+ rewards_per_latent_episode[latent_id, episode_id] += reward
2021
2237
 
2022
- time += 1
2238
+ # episode is done - bootstrap value if truncated, then refill
2023
2239
 
2024
- if not terminated:
2025
- # add bootstrap value if truncated
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
- next_value = temp_batch_dim(self.agent.get_critic_values)(next_state, latent = latent, use_ema_if_available = True, use_unwrapped_model = True)
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
- 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
- )
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
- memories.append(memory_for_gae)
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 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,6 +1,6 @@
1
1
  [project]
2
2
  name = "evolutionary-policy-optimization"
3
- version = "0.3.0"
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)