my_env / train.py
Sreeharan's picture
Upload folder using huggingface_hub
115d05e verified
Raw
History Blame Contribute Delete
4.11 kB
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), # state -> hidden
nn.ReLU(),
nn.Linear(32, 32), # hidden -> hidden
nn.ReLU(),
nn.Linear(32, 2) # hidden -> Q_left, Q_right
)
def forward(self, x):
return self.net(x)
# Model used for predicting the next move
model = DQN()
optimizer = optim.Adam(model.parameters(), lr=0.01)
# Model used as base reference (like a teacher model)
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) # random
state_tensor = torch.tensor([[float(state)]])
# e.g., state=2 → [[2.0]]
q_values = model(state_tensor)
# e.g., [[0.3, 0.7]]
action = torch.argmax(q_values).item()
return action
# replay buffer
replay_buffer = deque(maxlen=1000)
batch_size = 32
gamma = 0.9
# training function
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)
# current Q values
q_values = model(states) # shape: [batch, 2]
q_values = q_values.gather(1, actions.unsqueeze(1)).squeeze()
# next Q values (target network)
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:
# choose action
action = dqn_choose_action(state)
# take step
obs = await client.step(MyAction(move=action))
next_state = obs["observation"]["state"]
reward = obs["reward"]
done = obs["done"]
truncated = False
# store experience
replay_buffer.append((state, action, reward, next_state, done))
# train
loss = train_from_buffer(
replay_buffer, model, target_model, optimizer, gamma, batch_size
)
# move forward
state = next_state
total_reward += reward
if done:
break
episode_rewards.append(total_reward)
# update target model
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())