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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: evolutionary-policy-optimization
3
- Version: 0.2.22
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.1.19
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
- actions = actor(state, latent)
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(agent, env, num_learning_cycles = 5)
149
+ epo(env, num_learning_cycles = 5)
149
150
 
150
151
  # saving and loading
151
152
 
@@ -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
- actions = actor(state, latent)
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(agent, env, num_learning_cycles = 5)
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, pack
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, (np.ndarray, np.float64)):
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 gumbel_noise(t):
124
- return -log(-log(torch.rand_like(t)))
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
- if not is_greedy:
130
- t = (t / temperature) + gumbel_noise(t)
133
+ return reduce(t, 'b ... -> b', 'sum')
131
134
 
132
- return t.argmax(dim = -1)
135
+ def mode_action(distr):
136
+ # greedy action - argmax for categorical, mean for beta
133
137
 
134
- def calc_entropy(logits):
135
- prob = logits.softmax(dim = -1)
136
- return -(prob * log(prob)).sum(dim = -1)
138
+ if isinstance(distr, Categorical):
139
+ return distr.probs.argmax(dim = -1)
137
140
 
138
- def gather_log_prob(
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.to_out = nn.Sequential(
594
- nn.RMSNorm(dim),
595
- nn.Linear(dim, num_actions, bias = False),
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 = -10.,
625
- max_value = 10.,
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
- pred_kwargs = dict(return_logits = return_logits) if not self.use_regression else dict()
719
- return self.to_pred(hidden, **pred_kwargs)
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 = 1.5 # not sure what the value is
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
- # equation (3)
809
+ # eq (3) - fire when the fitness spread is a meaningful fraction of the level
733
810
 
734
- # max(fitness) - min(fitness) > gamma * median(fitness)
735
- # however, this equation does not make much sense to me if fitness increases unbounded
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 (fitnesses.amax(dim = -1) - fitnesses.amin(dim = -1)) > (self.gamma * torch.median(fitnesses, dim = -1).values)
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 = 1.5,
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 get_actor_actions(
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
- logits = maybe_unwrap(self.actor)(state, latent)
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 logits
1464
+ return mode_action(distr)
1369
1465
 
1370
- actions = gumbel_sample(logits, temperature = temperature)
1466
+ # temperature <= 0 is greedy - mode works for both discrete and continuous
1371
1467
 
1372
- log_probs = gather_log_prob(logits, actions)
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
- logits = self.actor(states, latents)
1593
+ distr = self.actor(states, latents)
1496
1594
 
1497
- actor_loss = self.actor_loss(logits, log_probs, actions, advantages, use_spo = self.use_spo)
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
- logits, # Float[b l]
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 = logits.shape[0]
1738
+ batch = advantages.shape[0]
1634
1739
 
1635
- log_probs = gather_log_prob(logits, actions)
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, truncated, terminated, _ = interface_torch_numpy(env.step, device = self.device)(action)
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')
@@ -51,5 +51,5 @@ class Env(Module):
51
51
 
52
52
  self._step.add_(1)
53
53
 
54
- out = (state, reward, truncated, terminated)
54
+ out = (state, reward, terminated, truncated)
55
55
  return (*tuple(t.numpy() for t in out), None)
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "evolutionary-policy-optimization"
3
- version = "0.2.22"
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.1.19',
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
- )