evolutionary-policy-optimization 0.2.20__tar.gz → 0.2.22__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
- Metadata-Version: 2.4
1
+ Metadata-Version: 2.5
2
2
  Name: evolutionary-policy-optimization
3
- Version: 0.2.20
3
+ Version: 0.2.22
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
@@ -363,20 +363,21 @@ That's it
363
363
  year = {2018},
364
364
  volume = {abs/1802.06070},
365
365
  url = {https://arxiv.org/abs/1802.06070}
366
- url = {https://arxiv.org/abs/2602.16863},
366
+ }
367
367
  ```
368
368
 
369
369
  ```bibtex
370
370
  @article{Zhu2025Hybrid,
371
- author = {Zhu, Zimo and Yu, Chuanqiang and Wang, Junti},
372
- title = {A Hybrid Genetic Algorithm and Proximal Policy Optimization System for Efficient Multi-Agent Task Allocation},
373
- journal = {Systems},
374
- volume = {13},
375
- year = {2025},
376
- number = {6},
371
+ author = {Zhu, Zimo and Yu, Chuanqiang and Wang, Junti},
372
+ title = {A Hybrid Genetic Algorithm and Proximal Policy Optimization System for Efficient Multi-Agent Task Allocation},
373
+ journal = {Systems},
374
+ volume = {13},
375
+ year = {2025},
376
+ number = {6},
377
377
  article-number = {453},
378
- url = {https://www.mdpi.com/2079-8954/13/6/453},
379
- issn = {2079-8954}
378
+ url = {https://www.mdpi.com/2079-8954/13/6/453},
379
+ issn = {2079-8954}
380
+ }
380
381
  ```
381
382
 
382
383
  ```bibtex
@@ -301,20 +301,21 @@ That's it
301
301
  year = {2018},
302
302
  volume = {abs/1802.06070},
303
303
  url = {https://arxiv.org/abs/1802.06070}
304
- url = {https://arxiv.org/abs/2602.16863},
304
+ }
305
305
  ```
306
306
 
307
307
  ```bibtex
308
308
  @article{Zhu2025Hybrid,
309
- author = {Zhu, Zimo and Yu, Chuanqiang and Wang, Junti},
310
- title = {A Hybrid Genetic Algorithm and Proximal Policy Optimization System for Efficient Multi-Agent Task Allocation},
311
- journal = {Systems},
312
- volume = {13},
313
- year = {2025},
314
- number = {6},
309
+ author = {Zhu, Zimo and Yu, Chuanqiang and Wang, Junti},
310
+ title = {A Hybrid Genetic Algorithm and Proximal Policy Optimization System for Efficient Multi-Agent Task Allocation},
311
+ journal = {Systems},
312
+ volume = {13},
313
+ year = {2025},
314
+ number = {6},
315
315
  article-number = {453},
316
- url = {https://www.mdpi.com/2079-8954/13/6/453},
317
- issn = {2079-8954}
316
+ url = {https://www.mdpi.com/2079-8954/13/6/453},
317
+ issn = {2079-8954}
318
+ }
318
319
  ```
319
320
 
320
321
  ```bibtex
@@ -362,7 +362,7 @@ class StateNorm(Module):
362
362
  self,
363
363
  state
364
364
  ):
365
- assert state.shape[-1] == self.dim, f'expected feature dimension of {self.dim} but received {x.shape[-1]}'
365
+ assert state.shape[-1] == self.dim, f'expected feature dimension of {self.dim} but received {state.shape[-1]}'
366
366
 
367
367
  time = self.step.item()
368
368
  mean = self.running_mean
@@ -512,7 +512,7 @@ class MLP(Module):
512
512
  if latent.ndim == 1:
513
513
  latent = repeat(latent, 'd -> b d', b = batch)
514
514
 
515
- assert latent.shape[0] == x.shape[0], f'received state with batch size {x.shape[0]} but latent ids received had batch size {latent_id.shape[0]}'
515
+ assert latent.shape[0] == x.shape[0], f'received state with batch size {x.shape[0]} but latent ids received had batch size {latent.shape[0]}'
516
516
 
517
517
  # layers
518
518
 
@@ -540,13 +540,11 @@ class DiversityDiscr(Module):
540
540
  def __init__(
541
541
  self,
542
542
  dim_state,
543
- num_actions,
544
543
  num_latents,
545
544
  dim = 64,
546
545
  depth = 2
547
546
  ):
548
547
  super().__init__()
549
- self.action_embed = nn.Embedding(num_actions, dim)
550
548
  self.state_proj = nn.Linear(dim_state, dim)
551
549
 
552
550
  self.net = AttnResidualNormedMLP(
@@ -558,13 +556,13 @@ class DiversityDiscr(Module):
558
556
 
559
557
  def reset_parameters(self):
560
558
  for module in self.modules():
561
- if isinstance(module, (nn.Linear, nn.LayerNorm, nn.RMSNorm, nn.Embedding)):
559
+ if isinstance(module, (nn.Linear, nn.LayerNorm, nn.RMSNorm)):
562
560
  module.reset_parameters()
563
561
 
564
- def forward(self, state, action):
562
+ def forward(self, state, next_state):
565
563
  state_embed = self.state_proj(state)
566
- action_embed = self.action_embed(action)
567
- return self.net((state_embed, action_embed))
564
+ next_state_embed = self.state_proj(next_state)
565
+ return self.net((state_embed, next_state_embed))
568
566
 
569
567
  # actor, critic, and agent (actor + critic)
570
568
  # eventually, should just create a separate repo and aggregate all the MLP related architectures
@@ -1201,11 +1199,9 @@ class Agent(Module):
1201
1199
  if use_diversity_discr:
1202
1200
  assert exists(latent_gene_pool), 'latent_gene_pool must be present to use DIAYN'
1203
1201
  dim_state = actor.init_layer[0].in_features
1204
- num_actions = actor.to_out[1].out_features
1205
1202
 
1206
1203
  self.diversity_discr = DiversityDiscr(
1207
1204
  dim_state = dim_state,
1208
- num_actions = num_actions,
1209
1205
  num_latents = latent_gene_pool.num_latents,
1210
1206
  **diversity_discr_kwargs
1211
1207
  )
@@ -1282,6 +1278,10 @@ class Agent(Module):
1282
1278
  def device(self):
1283
1279
  return self.step.device
1284
1280
 
1281
+ @property
1282
+ def is_main_process(self):
1283
+ return not self.wrap_with_accelerate or self.accelerate.is_main_process
1284
+
1285
1285
  @property
1286
1286
  def unwrapped_latent_gene_pool(self):
1287
1287
  return self.unwrap_model(self.latent_gene_pool)
@@ -1435,6 +1435,7 @@ class Agent(Module):
1435
1435
  (
1436
1436
  episode_ids,
1437
1437
  states,
1438
+ next_states,
1438
1439
  latent_gene_ids,
1439
1440
  actions,
1440
1441
  log_probs,
@@ -1457,7 +1458,7 @@ class Agent(Module):
1457
1458
 
1458
1459
  valid_episode = episode_ids >= 0
1459
1460
 
1460
- dataset = TensorDataset(*[t[valid_episode] for t in (advantages, states, latent_gene_ids, actions, log_probs, values)])
1461
+ dataset = TensorDataset(*[t[valid_episode] for t in (advantages, states, next_states, latent_gene_ids, actions, log_probs, values)])
1461
1462
 
1462
1463
  dataloader = DataLoader(dataset, batch_size = self.batch_size, shuffle = True)
1463
1464
 
@@ -1469,10 +1470,11 @@ class Agent(Module):
1469
1470
  self.actor.train()
1470
1471
  self.critic.train()
1471
1472
 
1472
- for _ in tqdm(range(epochs), desc = 'learning actor/critic epoch'):
1473
+ for _ in tqdm(range(epochs), desc = 'learning actor/critic epoch', disable = not self.is_main_process):
1473
1474
  for (
1474
1475
  advantages,
1475
1476
  states,
1477
+ next_states,
1476
1478
  latent_gene_ids,
1477
1479
  actions,
1478
1480
  log_probs,
@@ -1526,7 +1528,7 @@ class Agent(Module):
1526
1528
  diversity_discr_loss = self.zero
1527
1529
 
1528
1530
  if self.use_diversity_discr:
1529
- diversity_discr_logits = self.diversity_discr(states, actions)
1531
+ diversity_discr_logits = self.diversity_discr(states, next_states)
1530
1532
  diversity_discr_loss = F.cross_entropy(diversity_discr_logits, latent_gene_ids)
1531
1533
 
1532
1534
  diversity_discr_loss.backward()
@@ -1580,7 +1582,7 @@ class Agent(Module):
1580
1582
  if exists(self.state_norm):
1581
1583
  self.state_norm.train()
1582
1584
 
1583
- for _, states, *_ in tqdm(dataloader, desc = 'state norm learning'):
1585
+ for _, states, *_ in tqdm(dataloader, desc = 'state norm learning', disable = not self.is_main_process):
1584
1586
  self.state_norm(states)
1585
1587
 
1586
1588
  # update warmed up state for discriminator
@@ -1728,6 +1730,7 @@ def create_agent(
1728
1730
  Memory = namedtuple('Memory', [
1729
1731
  'episode_id',
1730
1732
  'state',
1733
+ 'next_state',
1731
1734
  'latent_gene_id',
1732
1735
  'action',
1733
1736
  'log_prob',
@@ -1835,7 +1838,7 @@ class EPO(Module):
1835
1838
 
1836
1839
  rollout_gen = self.rollouts_for_machine(fix_environ_across_latents)
1837
1840
 
1838
- for latent_id, episode_id, maybe_seed in tqdm(rollout_gen, desc = 'rollout'):
1841
+ for latent_id, episode_id, maybe_seed in tqdm(rollout_gen, desc = 'rollout', disable = not self.agent.is_main_process):
1839
1842
 
1840
1843
  time = 0
1841
1844
 
@@ -1868,7 +1871,7 @@ class EPO(Module):
1868
1871
 
1869
1872
  # get the next state, action, and reward
1870
1873
 
1871
- state, reward, truncated, terminated, _ = interface_torch_numpy(env.step, device = self.device)(action)
1874
+ next_state, reward, truncated, terminated, _ = interface_torch_numpy(env.step, device = self.device)(action)
1872
1875
 
1873
1876
  # diversity reward
1874
1877
 
@@ -1877,7 +1880,7 @@ class EPO(Module):
1877
1880
  diversity_discr = self.agent.unwrap_model(self.agent.diversity_discr)
1878
1881
  diversity_discr.eval()
1879
1882
 
1880
- logits = diversity_discr(rearrange(state, '... -> 1 ...'), rearrange(action, '... -> 1 ...'))
1883
+ logits = diversity_discr(rearrange(state, '... -> 1 ...'), rearrange(next_state, '... -> 1 ...'))
1881
1884
  log_probs = logits.log_softmax(dim=-1)
1882
1885
 
1883
1886
  diversity_reward = log_probs[0, latent_id] + log(tensor(self.agent.num_latents))
@@ -1895,6 +1898,7 @@ class EPO(Module):
1895
1898
  memory = Memory(
1896
1899
  tensor(episode_id),
1897
1900
  state,
1901
+ next_state,
1898
1902
  tensor(latent_id),
1899
1903
  action,
1900
1904
  log_prob,
@@ -1907,15 +1911,18 @@ class EPO(Module):
1907
1911
 
1908
1912
  memories.append(memory)
1909
1913
 
1914
+ state = next_state
1915
+
1910
1916
  time += 1
1911
1917
 
1912
1918
  if not terminated:
1913
1919
  # add bootstrap value if truncated
1914
1920
 
1915
- next_value = temp_batch_dim(self.agent.get_critic_values)(state, latent = latent, use_ema_if_available = True, use_unwrapped_model = True)
1921
+ next_value = temp_batch_dim(self.agent.get_critic_values)(next_state, latent = latent, use_ema_if_available = True, use_unwrapped_model = True)
1916
1922
 
1917
1923
  memory_for_gae = memory._replace(
1918
1924
  episode_id = invalid_episode,
1925
+ reward = next_value.cpu(),
1919
1926
  value = next_value.cpu(),
1920
1927
  done = tensor(True)
1921
1928
  )
@@ -1939,7 +1946,7 @@ class EPO(Module):
1939
1946
  torch.manual_seed(seed)
1940
1947
  np.random.seed(seed)
1941
1948
 
1942
- for _ in tqdm(range(num_learning_cycles), desc = 'learning cycle'):
1949
+ for _ in tqdm(range(num_learning_cycles), desc = 'learning cycle', disable = not self.agent.is_main_process):
1943
1950
 
1944
1951
  memories = self.gather_experience_from(env)
1945
1952
 
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "evolutionary-policy-optimization"
3
- version = "0.2.20"
3
+ version = "0.2.22"
4
4
  description = "EPO - Pytorch"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }