|
|
| from collections import deque |
| import random |
| import numpy as np |
|
|
| import torch |
| import torch.nn as nn |
| import torch.optim as optim |
|
|
| class DQN(nn.Module): |
| def __init__(self): |
| super().__init__() |
| self.net = nn.Sequential( |
| nn.Linear(1, 32), |
| nn.ReLU(), |
| nn.Linear(32, 32), |
| nn.ReLU(), |
| nn.Linear(32, 2) |
| ) |
|
|
| def forward(self, x): |
| return self.net(x) |
|
|
| |
| model = DQN() |
| optimizer = optim.Adam(model.parameters(), lr=0.01) |
|
|
| |
| target_model = DQN() |
| target_model.load_state_dict(model.state_dict()) |
|
|
| epsilon = 0.2 |
|
|
| def dqn_choose_action(state): |
| if np.random.rand() < epsilon: |
| return np.random.randint(0, 2) |
|
|
| state_tensor = torch.tensor([[float(state)]]) |
| |
|
|
| q_values = model(state_tensor) |
| |
|
|
| action = torch.argmax(q_values).item() |
| return action |
|
|
| |
| replay_buffer = deque(maxlen=1000) |
| batch_size = 32 |
| gamma = 0.9 |
|
|
|
|
| |
| def train_from_buffer(replay_buffer, model, target_model, optimizer, gamma, batch_size=32): |
| if len(replay_buffer) < batch_size: |
| return None |
|
|
| batch = random.sample(replay_buffer, batch_size) |
|
|
| states, actions, rewards, next_states, dones = zip(*batch) |
|
|
| states = torch.tensor(states, dtype=torch.float32).unsqueeze(1) |
| actions = torch.tensor(actions) |
| rewards = torch.tensor(rewards, dtype=torch.float32) |
| next_states = torch.tensor(next_states, dtype=torch.float32).unsqueeze(1) |
| dones = torch.tensor(dones, dtype=torch.float32) |
| |
| |
| |
| q_values = model(states) |
| q_values = q_values.gather(1, actions.unsqueeze(1)).squeeze() |
|
|
| |
| with torch.no_grad(): |
| next_q_values = target_model(next_states) |
| next_max = next_q_values.max(1)[0] |
|
|
| target = rewards + gamma * next_max * (1 - dones) |
|
|
| loss = ((q_values - target) ** 2).mean() |
|
|
| optimizer.zero_grad() |
| loss.backward() |
| optimizer.step() |
|
|
| return loss.item() |
|
|
| import asyncio |
| from client import MyEnvClient |
| from models import MyAction |
|
|
| async def train_one_step(): |
| client = MyEnvClient("http://127.0.0.1:8000") |
| obs = await client.reset() |
| state = obs["observation"]["state"] |
|
|
| obs = await client.step(MyAction(move=1)) |
| |
| |
| next_state = obs["observation"]["state"] |
| reward = obs["reward"] |
| done = obs["done"] |
| truncated = False |
| print(f"State: {state}, Action: {1}, Next State: {next_state}, Reward: {reward}, Done: {done}, Truncated: {truncated}") |
|
|
| async def train(episodes=100): |
| client = MyEnvClient("http://127.0.0.1:8000") |
|
|
| episode_rewards = [] |
|
|
| for episode in range(episodes): |
|
|
| obs = await client.reset() |
| state = obs["observation"]["state"] |
|
|
| total_reward = 0 |
|
|
| while True: |
|
|
| |
| action = dqn_choose_action(state) |
|
|
| |
| obs = await client.step(MyAction(move=action)) |
|
|
| next_state = obs["observation"]["state"] |
| reward = obs["reward"] |
| done = obs["done"] |
| truncated = False |
|
|
| |
| replay_buffer.append((state, action, reward, next_state, done)) |
|
|
| |
| loss = train_from_buffer( |
| replay_buffer, model, target_model, optimizer, gamma, batch_size |
| ) |
|
|
| |
| state = next_state |
| total_reward += reward |
|
|
| if done: |
| break |
|
|
| episode_rewards.append(total_reward) |
|
|
| |
| if episode % 10 == 0: |
| target_model.load_state_dict(model.state_dict()) |
|
|
| print(f"Episode {episode}, Reward: {total_reward}, Loss: {loss}") |
| |
| torch.save(model.state_dict(), "dqn.pt") |
|
|
|
|
| if __name__ == "__main__": |
| asyncio.run(train()) |
|
|