Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions config_templates/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,9 @@ agent:
train_critic_iters: 80 # Number of iterations to train critic
discount_factor: 0.99
gae_discount_factor: 0.97
# PPO reward mix: r = extrinsic_reward_coeff * r_env + intrinsic_reward_coeff * r_icm
extrinsic_reward_coeff: 1.0
intrinsic_reward_coeff: 1.0
icm:
scaling_factor: 1.0
train_forward_dynamics_iters: 80
Expand Down
6 changes: 6 additions & 0 deletions mineagent/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,10 @@ class PPOConfig:
Discount factor for Generalized Advantage Estimation
focus_loss_coeff : float, optional
Coefficient for the separate REINFORCE loss on the focus/ROI head
extrinsic_reward_coeff : float, optional
Weight λ_ext on environment reward in PPO: r = λ_ext * r_env + λ_icm * r_intrinsic
intrinsic_reward_coeff : float, optional
Weight λ_icm on ICM intrinsic reward in PPO (independent of ICMConfig.scaling_factor)
"""

clip_ratio: float = 0.2
Expand All @@ -60,6 +64,8 @@ class PPOConfig:
discount_factor: float = 0.99
gae_discount_factor: float = 0.97
focus_loss_coeff: float = 0.01
extrinsic_reward_coeff: float = 1.0
intrinsic_reward_coeff: float = 1.0


@dataclass
Expand Down
4 changes: 3 additions & 1 deletion mineagent/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,9 +47,11 @@ def run() -> None:
event_bus.publish(EnvReset(timestamp=datetime.now(), observation=frame))
obs = torch.tensor(frame, dtype=torch.float).unsqueeze(0)
total_return = 0.0
prev_env_reward = 0.0
for _ in range(engine_config.max_steps):
action = agent.act(obs)
action = agent.act(obs, reward=prev_env_reward)
next_frame, reward, terminated, truncated, info = env.step(action)
prev_env_reward = float(reward)
next_obs = torch.tensor(next_frame, dtype=torch.float).unsqueeze(0)
event_bus.publish(
EnvStep(
Expand Down
7 changes: 6 additions & 1 deletion mineagent/learning/ppo.py
Original file line number Diff line number Diff line change
Expand Up @@ -115,6 +115,8 @@ def __init__(self, actor: nn.Module, critic: nn.Module, config: PPOConfig):
self.clip_ratio = config.clip_ratio
self.target_kl = config.target_kl
self.focus_loss_coeff = config.focus_loss_coeff
self.extrinsic_reward_coeff = config.extrinsic_reward_coeff
self.intrinsic_reward_coeff = config.intrinsic_reward_coeff
self.actor_optim = optim.Adam(
self.actor.parameters(),
lr=config.actor_lr,
Expand Down Expand Up @@ -222,7 +224,10 @@ def _finalize_trajectory(self, data: TrajectoryBuffer) -> PPOSample:
# The reward for a_t is at r_{t+1}
env_rewards = torch.tensor(list(data.rewards_buffer)[1:])
intrinsic_rewards = torch.tensor(list(data.intrinsic_rewards_buffer)[1:])
rewards = env_rewards + intrinsic_rewards
rewards = (
self.extrinsic_reward_coeff * env_rewards
+ self.intrinsic_reward_coeff * intrinsic_rewards
)
# Need all values since the final one is used to estimate future reward
values = torch.tensor(list(data.values_buffer))

Expand Down
18 changes: 18 additions & 0 deletions tests/agent/test_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,24 @@ def test_agent_v1_act_single(agent_v1_module: AgentV1):
assert isinstance(float(action["scroll_delta"]), float)


def test_act_stores_extrinsic_reward_from_previous_step() -> None:
"""Mirrors engine timing: reward from env.step is passed into the next act()."""
torch.manual_seed(0)
agent = AgentV1(
AgentConfig(
ppo=PPOConfig(),
icm=ICMConfig(),
td=TDConfig(),
max_buffer_size=100,
),
)
obs = torch.randn((1, 3, 160, 256))
agent.act(obs, reward=0.0)
agent.act(obs, reward=7.5)
assert agent.memory.rewards_buffer[0] == 0.0
assert agent.memory.rewards_buffer[1] == 7.5


def test_agent_v1_params(agent_v1_module: AgentV1):
modules = [
agent_v1_module.vision,
Expand Down
44 changes: 44 additions & 0 deletions tests/learning/test_ppo.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,12 @@
import numpy as np
import pytest
from mineagent.agent.agent import AgentV1
import torch
from mineagent.learning.ppo import PPO
from mineagent.config import AgentConfig, PPOConfig, ICMConfig, TDConfig
from mineagent.memory.trajectory import TrajectoryBuffer
from mineagent.client.protocol import NUM_KEYS
from mineagent.utils import discount_cumsum

ENV_ACTION_DIM = NUM_KEYS + 3 + 3
FOCUS_DIM = 2
Expand Down Expand Up @@ -40,3 +42,45 @@ def test_ppo_update(ppo_module: PPO) -> None:
focus_logp=torch.ones((FOCUS_DIM,), dtype=torch.float),
)
ppo_module.update(trajectory)


def test_ppo_finalize_weighted_rewards() -> None:
"""r = λ_ext * r_env + λ_icm * r_intrinsic before discount_cumsum for returns."""
gamma = 0.99
lam_ext, lam_icm = 2.0, 0.5
agent = AgentV1(
AgentConfig(
ppo=PPOConfig(
train_actor_iters=2,
train_critic_iters=2,
discount_factor=gamma,
extrinsic_reward_coeff=lam_ext,
intrinsic_reward_coeff=lam_icm,
),
icm=ICMConfig(),
td=TDConfig(),
),
)
ppo = agent.ppo

buffer_size = 3
trajectory = TrajectoryBuffer(max_buffer_size=buffer_size)
env_r = [0.0, 4.0, 6.0]
int_r = [0.0, 2.0, 4.0]
for i in range(buffer_size):
trajectory.store(
torch.zeros((EMBED_DIM,), dtype=torch.float),
torch.zeros((ENV_ACTION_DIM,), dtype=torch.float),
env_r[i],
int_r[i],
0.0,
torch.ones((ENV_ACTION_DIM,), dtype=torch.float),
focus=torch.zeros((FOCUS_DIM,), dtype=torch.float),
focus_logp=torch.ones((FOCUS_DIM,), dtype=torch.float),
)

weighted = lam_ext * np.array(env_r[1:]) + lam_icm * np.array(int_r[1:])
expected_returns = discount_cumsum(weighted, gamma)

sample = ppo._finalize_trajectory(trajectory)
assert np.allclose(sample.returns.squeeze().numpy(), expected_returns)
Loading