프로그래밍

개인 작업 -2-

하바사 2025. 3. 29. 23:51
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}")