Spaces:
Sleeping
Sleeping
File size: 2,914 Bytes
64ea2b1 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 | import random
import torch
import torch.nn as nn
import numpy as np
from collections import deque
from env.environment import AttentionEnv
from agents.dqn_agent import QNetwork, featurize
# -------- Hyperparameters --------
EPISODES = 2000
BATCH_SIZE = 64
GAMMA = 0.95
LR = 0.001
EPSILON_START = 1.0
EPSILON_END = 0.05
EPSILON_DECAY = 0.995
TARGET_UPDATE = 50
BUFFER_SIZE = 5000
# -------- Replay Buffer --------
buffer = deque(maxlen=BUFFER_SIZE)
def sample_batch():
batch = random.sample(buffer, BATCH_SIZE)
states, targets = [], []
for state, item, reward, next_state, done in batch:
x = torch.FloatTensor(featurize(state, item))
if done:
target = reward
else:
next_qs = []
for next_item in next_state.items:
x_next = torch.FloatTensor(featurize(next_state, next_item))
next_qs.append(target_model(x_next).item())
target = reward + GAMMA * max(next_qs)
states.append(x)
targets.append([target])
return torch.stack(states), torch.FloatTensor(targets)
# -------- Init --------
env = AttentionEnv()
sample_state = env.reset()
sample_item = sample_state.items[0]
input_dim = len(featurize(sample_state, sample_item))
model = QNetwork(input_dim)
target_model = QNetwork(input_dim)
target_model.load_state_dict(model.state_dict())
optimizer = torch.optim.Adam(model.parameters(), lr=LR)
loss_fn = nn.MSELoss()
epsilon = EPSILON_START
rewards_log = []
# -------- Training Loop --------
for episode in range(EPISODES):
state = env.reset()
done = False
total_reward = 0
while not done:
# Epsilon-greedy
if random.random() < epsilon:
item = random.choice(state.items)
else:
qs = []
for i in state.items:
x = torch.FloatTensor(featurize(state, i))
qs.append(model(x).item())
item = state.items[np.argmax(qs)]
next_state, reward, done, _ = env.step(type("A", (), {"item_id": item.id})())
buffer.append((state, item, reward.value, next_state, done))
state = next_state
total_reward += reward.value
# Train
if len(buffer) > BATCH_SIZE:
states, targets = sample_batch()
preds = model(states)
loss = loss_fn(preds, targets)
optimizer.zero_grad()
loss.backward()
optimizer.step()
rewards_log.append(total_reward)
# Update target network
if episode % TARGET_UPDATE == 0:
target_model.load_state_dict(model.state_dict())
epsilon = max(EPSILON_END, epsilon * EPSILON_DECAY)
if episode % 100 == 0:
print(f"Episode {episode}, Reward: {total_reward:.2f}, Epsilon: {epsilon:.2f}")
# Save model
torch.save(model.state_dict(), "dqn_model.pth")
# Save rewards
np.save("rewards.npy", rewards_log) |