import torch
import torch.nn as nn
import torch.optim as optim
import numpy as np
# 기본 설정
NUM_PATTERNS = 100
PHASE_ENCODING_SIZE = 4
STATE_SIZE = NUM_PATTERNS + PHASE_ENCODING_SIZE
# 임시 데이터 생성: (상태 벡터, 선택된 패턴 index)
# 상태 = [패턴 선택 히스토리 (100차원) + 시점정보 (4차원)]
dummy_state1 = np.zeros(104)
dummy_state1[0] = 1 # 첫 번째 패턴이 선택되었음을 의미
dummy_state1[100] = 1 # 초반 시점 one-hot
dummy_state2 = np.zeros(104)
dummy_state2[10] = 1
dummy_state2[101] = 1 # 중반 시점 one-hot
training_data = [
(dummy_state1, 20), # 상태 1, 선택된 패턴 인덱스 = 20
(dummy_state2, 33) # 상태 2, 선택된 패턴 인덱스 = 33
]
# 정책 네트워크 정의
class PolicyNetwork(nn.Module):
def __init__(self, state_size, action_size):
super(PolicyNetwork, self).__init__()
self.fc = nn.Sequential(
nn.Linear(state_size, 128),
nn.ReLU(),
nn.Linear(128, action_size)
)
def forward(self, x):
return self.fc(x)
model = PolicyNetwork(STATE_SIZE, NUM_PATTERNS)
optimizer = optim.Adam(model.parameters(), lr=0.001)
loss_fn = nn.CrossEntropyLoss()
for epoch in range(50):
total_loss = 0
for state, action_index in training_data:
pred = model(torch.FloatTensor(state))
loss = loss_fn(pred.unsqueeze(0), torch.LongTensor([action_index]))
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += loss.item()
print(f"Epoch {epoch+1}, Loss: {total_loss}")