CTDE로 설계하는 멀티에이전트 강화학습의 중앙 학습과 분산 실행

CTDE 기반 멀티에이전트 강화학습의 중앙 비평가, 가치 분해, 분산 실행 구조와 MADDPG 운영 고려사항을 정리한다.

2026-08-14 · 최초 발행 2024-04-29

학습에는 전체 상황을 보고, 실행은 각자 판단하게 한다

멀티에이전트 강화학습(MARL)에서는 각 에이전트가 일부 관측만 가진 채 서로의 행동에 영향을 준다. 이 상호작용은 학습 대상 환경을 계속 바꾸는 비정상성 문제로 이어지고, 공동 보상 환경에서는 개별 에이전트의 기여를 판단하기도 어렵다.

Centralized Training with Decentralized Execution(CTDE)은 이 문제를 학습과 실행의 정보 범위를 분리해 다룬다. 학습 시에는 글로벌 상태, 다른 에이전트의 행동, 공동 보상 같은 중앙화된 신호를 사용한다. 반면 운영 환경에서는 각 에이전트가 자신이 관측한 정보만으로 독립 정책을 실행한다. 통신 제약이나 지연이 있는 환경에서 실시간 의사결정을 유지하면서 협력·경쟁 과제를 최적화하기 위한 구조다.

학습 단계에서는 중앙화된 가치 추정 또는 가치 분해를 사용해 신용할당 문제를 완화한다. 실행 단계의 정책은 부분 관찰만으로 동작하도록 설계한다.

대표적인 알고리즘 계열은 다음과 같다.

  • 중앙 비평가 계열: MADDPG, COMA, MAPPO 등
  • 가치 분해 계열: VDN, QMIX, QTRAN 등
  • 통신·파라미터 공유 확장: IQL + parameter sharing, FACMAC, G2ANet 등

중앙 신호가 학습을 안정시키는 방식

중앙 비평가(critic)는 학습 시점에 글로벌 상태와 공동 행동을 입력으로 받을 수 있다. 이 구조는 개별 정책만으로는 보기 어려운 상호작용을 가치 추정에 반영하며, 샘플 효율 개선과 신용할당 안정화에 사용된다.

실행 시 각 에이전트는 자신의 관측 o_i만으로 정책 π_i를 수행한다. 따라서 중앙 제어기나 지속적인 에이전트 간 통신에 의존하지 않고도 실행할 수 있으며, 통신 불능 또는 지연 환경에서도 실시간성을 확보할 수 있다.

협력 과제에서는 공동 Q 값을 개별 Q 값과 연결하는 가치 분해 방식도 중요하다. VDN과 QMIX는 공동 Q를 개별 Q로 분해하고, QMIX는 단조성 제약을 통해 학습 안정성을 강화한다. COMA는 counterfactual advantage를 이용해 각 에이전트의 기여도를 추정한다.

비정상성을 다루기 위해서는 joint transition 버퍼, 타깃 네트워크, 중요도 샘플링을 적용할 수 있다. fingerprinting과 정책 지연 업데이트 역시 학습 분산성을 제어하는 장치다. 운영 측면에서는 분산 수집기(worker)와 중앙 학습기(learner)를 분리하고, 체크포인트·재현성(seed)·모니터링 및 평가 파이프라인을 표준화한다.

협력 의사결정이 필요한 환경

군집 로보틱스와 자율 주행 협조에서는 다수 드론이나 AGV의 충돌 회피, 편대 유지, 태스크 할당을 함께 다룰 수 있다. V2V 통신이 불안정하다고 가정하는 경우에도 CTDE로 학습하고 분산 정책으로 실시간 실행을 보장하는 구성이 가능하다.

무선 자원 할당과 엣지 스케줄링에서는 셀 간 간섭을 줄이기 위한 파워·채널 할당이 대상이 된다. 중앙 학습 단계에서 글로벌 간섭 패턴을 학습한 뒤, 실행은 기지국 단위로 분산할 수 있다.

공급망과 창고 피킹에서는 재고 보충, 경로 계획, 스테이션 밸런싱을 동시에 최적화한다. 가치 분해는 공정 지연과 체리피킹을 방지하는 데 활용될 수 있다.

StarCraft II Micromanagement와 Google Research Football은 다중 에이전트 게임·시뮬레이션 벤치마크로 쓰인다. 알고리즘을 비교하고 성능을 검증하는 표준 실험대 역할을 한다.

성능과 운영성 사이의 균형

중앙 신호를 활용하면 환경에 따라 수렴 속도를 1.53배 단축하고, 승률/리워드를 1030%p 향상할 수 있다. 실행은 분산 정책으로 구성하므로 에이전트 수 증가에 대해 선형 근사 확장성을 확보한다.

가치 분해와 카운터팩추얼 이득은 변동성을 줄이고 재학습 비용을 낮추는 데 도움이 된다. 실행 시 통신 실패나 지연이 있어도 정책 성능을 유지할 수 있으며, 장애 격리도 쉬워진다.

학습과 실행의 정보 경로

Conditions & HandlingCentralized TrainingTD targets, advantagesTD targets, advantagespolicy gradientscritic lossYesNoYesNoθ_π only exportedEnv: joint state s_t, obso_t^1..o_t^N, actionsa_t^1..a_t^N, reward r_tJoint Replay BufferCentral Critic Q(s,a_1..a_N) or ValueDecompositionActors π_1(o_1)Actors π_2(o_2)Update θ_π^iUpdate θ_Q / mixerNon-stationarity?Fingerprinting / Target netsProceedPartial observability?LSTM/GRU policiesDecentralized ExecutionAgent i: observe o_iCompute a_i = π_i(o_i)Env step, no global info

학습 입력은 joint transition(s*t, o_t^1..o_t^N, a_t^1..a_t^N, r_t, s*{t+1})다. 중앙 비평가 또는 가치 분해를 학습하고 개별 정책을 업데이트하며, 비정상성과 부분 관찰에 대한 보완을 함께 적용한다. 결과물은 분산 실행에 사용할 수 있는 정책 파라미터 {θ_π^i}다.

CTDE와 다른 학습·실행 구조의 차이

접근 성능(보상/승률) 확장성(에이전트 수) 일관성(재현성/정책 드리프트) 안정성(수렴/분산) 운영 편의(배포/장애 격리)
CTDE 높음, 협력 과제 유리 높음, 실행 분산 높음, 중앙 신호로 안정화 높음, 가치 분해 지원 높음, 현장 통신 의존도 낮음
완전 중앙화 학습·실행 매우 높음(이상적 통신 가정) 낮음, 통신 병목 중간, 단일 실패점 중간, 스케일 시 불안정 낮음, 네트워크 의존
완전 분산 IQL 중간, 협력 한계 매우 높음 중간, 드리프트 위험 중간, 환경 의존 매우 높음

중앙 비평가와 분산 정책을 둔 MADDPG 스켈레톤

전제조건은 Python 3.10, PyTorch 2.3+, PettingZoo MPE(simple_spread_v3) 또는 PettingZoo 최신 버전(최신 정보 확인 필요)이다. CUDA는 선택적이다. 아래 코드는 실험 재현성을 목적으로 간소화했으며, 타깃 네트워크와 노멀라이저 등 학습 안정화 기법을 최소 포함한다.

# pip install pettingzoo==1.24.* supersuit==3.* torch==2.3.*
import numpy as np
import torch
import torch.nn as nn
import torch.optim as optim
from pettingzoo.mpe import simple_spread_v3

torch.manual_seed(7)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

class MLP(nn.Module):
    def __init__(self, in_dim, out_dim, hidden=128):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(in_dim, hidden), nn.ReLU(),
            nn.Linear(hidden, hidden), nn.ReLU(),
            nn.Linear(hidden, out_dim)
        )
    def forward(self, x): return self.net(x)

class Actor(nn.Module):
    def __init__(self, obs_dim, act_dim):
        super().__init__()
        self.body = MLP(obs_dim, act_dim)
    def forward(self, obs):
        return torch.tanh(self.body(obs))  # continuous action clamp

class Critic(nn.Module):
    def __init__(self, joint_obs_dim, joint_act_dim):
        super().__init__()
        self.q = MLP(joint_obs_dim + joint_act_dim, 1)
    def forward(self, joint_obs, joint_act):
        x = torch.cat([joint_obs, joint_act], dim=-1)
        return self.q(x)

def soft_update(targ, src, tau=0.005):
    for tp, sp in zip(targ.parameters(), src.parameters()):
        tp.data.mul_(1 - tau).add_(sp.data, alpha=tau)

def collect_dims(env):
    env.reset()
    obs_dims, act_dims = [], []
    for agent in env.agents:
        obs = env.observe(agent)
        obs_dims.append(obs.shape[0])
        # MPE uses continuous actions with bounds [-1,1] and act_space.n not defined; assume low/high
        act_dims.append(env.action_space(agent).shape[0])
    return obs_dims, act_dims

# Environment
env = simple_spread_v3.env(N=3, local_ratio=0.5, max_cycles=25, continuous_actions=True)
env.reset()

obs_dims, act_dims = collect_dims(env)
n = len(obs_dims)
actors = [Actor(obs_dims[i], act_dims[i]).to(device) for i in range(n)]
targ_actors = [Actor(obs_dims[i], act_dims[i]).to(device) for i in range(n)]
for i in range(n): targ_actors[i].load_state_dict(actors[i].state_dict())

joint_obs_dim = sum(obs_dims)
joint_act_dim = sum(act_dims)
critic = Critic(joint_obs_dim, joint_act_dim).to(device)
targ_critic = Critic(joint_obs_dim, joint_act_dim).to(device)
targ_critic.load_state_dict(critic.state_dict())

opt_actors = [optim.Adam(actors[i].parameters(), lr=1e-3) for i in range(n)]
opt_critic = optim.Adam(critic.parameters(), lr=2e-3)
gamma = 0.95

# Simple replay
from collections import deque
buf = deque(maxlen=50000)
batch_size = 256

def to_tensor(x): return torch.tensor(x, dtype=torch.float32, device=device)

def step_env(env, actors):
    env.reset()
    terminated = {a: False for a in env.agents}
    truncation = False
    ep_ret = 0.0
    last_obs = {a: env.observe(a) for a in env.agents}
    while env.agents:
        actions = {}
        for i, a in enumerate(env.agents):
            o = to_tensor(last_obs[a]).unsqueeze(0)
            with torch.no_grad():
                act = actors[i](o).squeeze(0).cpu().numpy()
            actions[a] = act
        for a, act in actions.items():
            env.step(act)
            # pettingzoo requires acting per active agent order; simple_spread_v3 handles ordering internally
        trans = {}
        joint_obs, joint_act, joint_next_obs, rews, dones = [], [], [], [], []
        for i, a in enumerate(terminated.keys()):
            o = last_obs[a]
            r = env.rewards.get(a, 0.0)
            d = env.terminations.get(a, False) or env.truncations.get(a, False)
            o2 = env.observe(a) if a in env.agents else o
            joint_obs.append(o); joint_next_obs.append(o2)
            joint_act.append(actions.get(a, np.zeros(act_dims[i])))
            rews.append(r); dones.append(d)
            last_obs[a] = o2
        ep_ret += np.mean(rews)
        buf.append((np.concatenate(joint_obs),
                    np.concatenate(joint_act),
                    np.mean(rews),
                    np.concatenate(joint_next_obs),
                    float(any(dones))))
        if any(dones): break
    return ep_ret

def sample_batch():
    idx = np.random.choice(len(buf), size=min(batch_size, len(buf)), replace=False)
    batch = [buf[i] for i in idx]
    obs, act, rew, next_obs, done = map(np.stack, zip(*batch))
    return to_tensor(obs), to_tensor(act), to_tensor(rew).unsqueeze(-1), to_tensor(next_obs), to_tensor(done).unsqueeze(-1)

for ep in range(200):  # demo iterations
    ret = step_env(env, actors)
    if len(buf) < 1000:
        print(f"warmup ep {ep}, ret {ret:.2f}")
        continue
    obs_b, act_b, rew_b, nxt_obs_b, done_b = sample_batch()

    # Target actions from target actors
    # Here we approximate next joint action using current target actors on next per-agent obs slices
    with torch.no_grad():
        offset = 0
        nxt_acts = []
        for i in range(n):
            o_dim = obs_dims[i]; a_dim = act_dims[i]
            o_i = nxt_obs_b[:, offset:offset+o_dim]
            offset += o_dim
            nxt_acts.append(targ_actors[i](o_i))
        nxt_joint_act = torch.cat(nxt_acts, dim=-1)
        y = rew_b + gamma * (1 - done_b) * targ_critic(nxt_obs_b, nxt_joint_act)

    # Critic update
    q = critic(obs_b, act_b)
    critic_loss = nn.MSELoss()(q, y)
    opt_critic.zero_grad(); critic_loss.backward(); opt_critic.step()

    # Actor updates
    offset_o = 0
    act_slices = []
    for i in range(n):
        o_dim = obs_dims[i]; a_dim = act_dims[i]
        o_i = obs_b[:, offset_o:offset_o+o_dim]
        offset_o += o_dim
        act_slices.append(actors[i](o_i))
    joint_act_pred = torch.cat(act_slices, dim=-1)
    # Maximize Q -> minimize -Q
    actor_loss = -critic(obs_b, joint_act_pred).mean()
    for opt in opt_actors: opt.zero_grad()
    actor_loss.backward()
    for i, opt in enumerate(opt_actors): opt.step()

    # Target updates
    soft_update(targ_critic, critic, tau=0.01)
    for i in range(n): soft_update(targ_actors[i], actors[i], tau=0.01)

    if ep % 10 == 0:
        print(f"ep {ep} ret {ret:.2f} | critic {critic_loss.item():.3f} | actor {-actor_loss.item():.3f}")

PettingZoo의 step 호출 방식과 에이전트 순회 규칙은 버전에 따라 상이 가능하므로 최신 정보 확인이 필요하다. 실제 학습에서는 노이즈 탐색(OU), 정책/입력 정규화, 학습률 스케줄, 클리핑, 우선순위 리플레이 적용을 권장한다.

과제 특성에 맞춘 선택과 운영 제어

협력 위주이면서 이산 또는 혼합 행동을 다룬다면 QMIX/VDN을 선택하고, 단조성 제약으로 학습 안정성을 확보할 수 있다. 연속 제어 또는 혼합 협력·경쟁 환경에서는 MADDPG/MAPPO를 선택해 중앙 비평가의 유연성을 활용한다.

파라미터 공유는 샘플 효율과 메모리 절감에 도움이 되지만, 과도하게 공유하면 정책 다양성이 떨어질 수 있다. 분산 수집기와 배치 동기화 주기를 조정하면서 스루풋과 지연의 트레이드오프를 관리한다.

실행 단계에서는 통신 비의존 설계 원칙을 유지하고, 학습 중 통신 채널 드랍아웃을 주입해 로버스트니스를 강화할 수 있다. RNN/Transformer 정책으로 부분 관찰을 보완하며, 지연 허용 범위는 사전에 검증한다.

타깃 네트워크, Double Q, Polyak 평균, gradient clipping은 안정화 기법으로 적용할 수 있다. fingerprinting, 정책 업데이트 지연, mixup 리워드 정규화도 비정상성 완화 관점에서 고려 대상이다.

일반화 평가는 멀티 시드·환경 랜덤화·도메인 랜덤화로 수행한다. 배포에는 정책만 포함하고 중앙 비평가는 제외한다. 롤백, A/B 테스트, 안전 제동 규칙을 병행하는 운영 구성이 필요하다.

멀티에이전트 강화학습CTDEMADDPGQMIX강화학습