"""FeatureDQNPolicy — action-as-feature DQN (Phase 2·3). 고정 슬롯 Q(s)→[11개] 대신 ScoreNet(상태벡터 + 카드임베딩) → 스칼라 점수. 결정 시 가용 카드 풀을 순회 채점해 argmax → 카드 추가/삭제/새 카드(zero-shot)에 구조 변화 없음. 협력사 특징은 상태벡터에 포함(feature_builder) → '협력사를 입력으로' 달성. 가변 행동 학습: replay 에 다음 상태의 '가용 카드 임베딩들'을 함께 저장, target = r + γ · max_{c'∈next_avail} Q(s', c') · (1-done) """ import math import random from collections import deque from typing import Dict, List, Optional, Tuple import numpy as np import torch import torch.nn as nn class ScoreNet(nn.Module): """(상태 + 카드임베딩) → 스칼라 점수.""" def __init__(self, state_dim: int, card_dim: int, hidden: int = 128): super().__init__() self.net = nn.Sequential( nn.Linear(state_dim + card_dim, hidden), nn.ReLU(), nn.Linear(hidden, hidden), nn.ReLU(), nn.Linear(hidden, 1), ) def forward(self, x: torch.Tensor) -> torch.Tensor: # x: [B, state+card] return self.net(x).squeeze(-1) # [B] class FeatureDQNPolicy: name = "feature_dqn" def __init__(self, state_dim: int, card_dim: int, device: str = "cpu", lr: float = 1e-3, gamma: float = 0.95, hidden: int = 128, eps_start: float = 1.0, eps_end: float = 0.05, eps_decay: int = 6000, buffer_size: int = 50_000, batch_size: int = 64, target_sync: int = 500): self.device = device self.gamma = gamma self.batch_size = batch_size self.target_sync = target_sync self.q = ScoreNet(state_dim, card_dim, hidden).to(device) self.tgt = ScoreNet(state_dim, card_dim, hidden).to(device) self.tgt.load_state_dict(self.q.state_dict()) self.opt = torch.optim.Adam(self.q.parameters(), lr=lr) self.buf: deque = deque(maxlen=buffer_size) self.eps_start, self.eps_end, self.eps_decay = eps_start, eps_end, eps_decay self.steps = 0 self.greedy = False # 평가 모드(탐색 끔) # ---- 탐색 스케줄 ---------------------------------------------------- def eps(self) -> float: if self.greedy: return 0.0 return self.eps_end + (self.eps_start - self.eps_end) * math.exp(-self.steps / self.eps_decay) # ---- 채점/선택 ------------------------------------------------------- def scores(self, state_feat: np.ndarray, card_embs: np.ndarray) -> np.ndarray: """가용 카드 K개 일괄 채점. card_embs: [K, card_dim] → [K].""" k = card_embs.shape[0] x = np.concatenate([np.repeat(state_feat[None, :], k, axis=0), card_embs], axis=1) with torch.no_grad(): return self.q(torch.tensor(x, device=self.device)).cpu().numpy() def select(self, state_feat: np.ndarray, card_embs: np.ndarray) -> Tuple[int, float, float]: """(선택 인덱스, propensity, 선택 점수). 인덱스는 card_embs 행 기준.""" k = card_embs.shape[0] sc = self.scores(state_feat, card_embs) e = self.eps() if random.random() < e: i = random.randrange(k) prop = e / k else: i = int(sc.argmax()) prop = (1.0 - e) + e / k return i, prop, float(sc[i]) # ---- 경험/학습 ------------------------------------------------------- def remember(self, state_feat: np.ndarray, card_emb: np.ndarray, reward: float, next_state_feat: Optional[np.ndarray], next_card_embs: Optional[np.ndarray], done: bool): self.buf.append((state_feat, card_emb, reward, next_state_feat, next_card_embs, done)) def train_step(self) -> Optional[float]: if len(self.buf) < self.batch_size: return None batch = random.sample(self.buf, self.batch_size) # Q(s, a_chosen) xs = np.stack([np.concatenate([s, c]) for s, c, *_ in batch]) q_sa = self.q(torch.tensor(xs, device=self.device)) # target = r + γ·max_{c'} Q_tgt(s', c') — 가변 후보라 후보 전체를 한 번에 forward 후 세그먼트 max rewards = torch.tensor([b[2] for b in batch], device=self.device, dtype=torch.float32) dones = torch.tensor([float(b[5]) for b in batch], device=self.device) next_rows, owner = [], [] for bi, (_, _, _, s2, cands, done) in enumerate(batch): if done or s2 is None or cands is None or len(cands) == 0: continue for c in cands: next_rows.append(np.concatenate([s2, c])) owner.append(bi) q_next_max = torch.zeros(self.batch_size, device=self.device) if next_rows: with torch.no_grad(): q_all = self.tgt(torch.tensor(np.stack(next_rows), device=self.device)) owner_t = torch.tensor(owner, device=self.device) q_next_max = q_next_max.index_reduce_(0, owner_t, q_all, "amax", include_self=False) target = rewards + self.gamma * q_next_max * (1.0 - dones) loss = nn.functional.smooth_l1_loss(q_sa, target) self.opt.zero_grad() loss.backward() self.opt.step() self.steps += 1 if self.steps % self.target_sync == 0: self.tgt.load_state_dict(self.q.state_dict()) return float(loss) # ---- 저장/로드 ------------------------------------------------------- def save(self, path: str): torch.save(self.q.state_dict(), path) def load(self, path: str): sd = torch.load(path, map_location=self.device) self.q.load_state_dict(sd) self.tgt.load_state_dict(sd)