o2o-negosium-original/agent/router/v1/learning/learning.py

150 lines
6.8 KiB
Python

"""learning/qtable 관리 API (Chat_server 14개 API 보존, tenant 스코프).
엔드포인트: /q-table/versions·switch·current, /experience-logs, /reset-learning, /reset-all,
/invalidate-session, /train. 모두 X-Tenant-ID 헤더로 company_id 격리.
"""
from typing import Optional
from fastapi import APIRouter, Depends
from pydantic import BaseModel
from common.enums import ErrorType
from negotiation.policies.base import Transition
from negotiation.policies.qtable_policy import UCBQTablePolicy
from negotiation.policy.model_store import QTablePolicyStore
from negotiation.qtable.domain.model.q_table import QTable
from negotiation.qtable.infra.repository.learning_repository import LearningRepository
from router.deps import get_tenant_engine
from tenancy.registry import TenantEngine
router = APIRouter(prefix="/v1", tags=["Learning / Q-Table"], responses={404: {"description": "Not found"}})
def _repo(engine: TenantEngine) -> LearningRepository:
return LearningRepository(engine.company_id)
@router.get("/q-table/versions", summary="Q-Table 버전 목록")
async def list_versions(engine: TenantEngine = Depends(get_tenant_engine)):
repo = _repo(engine)
err, versions = await repo.read(lambda s: repo.list_versions(s))
return {"company_id": engine.company_id, "versions": [
{"version_name": v.version_name, "is_active": v.is_active, "scope": v.scope,
"state_space_size": v.state_space_size, "action_space_size": v.action_space_size,
"epochs": v.epochs, "created_at": str(v.created_at)} for v in (versions or [])]}
class SwitchReq(BaseModel):
version_name: str
@router.post("/q-table/switch", summary="활성 Q-Table 버전 전환")
async def switch_version(req: SwitchReq, engine: TenantEngine = Depends(get_tenant_engine)):
err = await _repo(engine).set_active_version(req.version_name)
ok = err == ErrorType.SUCCESS
return {"success": ok, "active_version": req.version_name if ok else None,
"desc": err.name}
@router.get("/q-table/current", summary="현재 활성 Q-Table 상태")
async def current_qtable(engine: TenantEngine = Depends(get_tenant_engine)):
repo = _repo(engine)
err, active = await repo.read(lambda s: repo.get_active_version(s))
if not active:
return {"company_id": engine.company_id, "active_version": None, "q_value_rows": 0}
rows = await repo.qvalue_count(active.version_id)
return {"company_id": engine.company_id, "active_version": active.version_name,
"state_space_size": active.state_space_size, "action_space_size": active.action_space_size,
"q_value_rows": rows}
@router.get("/experience-logs", summary="최근 경험 로그")
async def experience_logs(limit: int = 50, engine: TenantEngine = Depends(get_tenant_engine)):
repo = _repo(engine)
err, total = await repo.read(lambda s: repo.count_experience(s))
logs = await repo.recent_experience(limit=min(limit, 500))
return {"company_id": engine.company_id, "total": total, "logs": logs}
@router.post("/reset-learning", summary="[관리] 학습 데이터 초기화(자사 스코프)")
async def reset_learning(engine: TenantEngine = Depends(get_tenant_engine)):
err = await _repo(engine).reset_learning()
return {"success": err == ErrorType.SUCCESS, "company_id": engine.company_id, "desc": err.name}
@router.post("/reset-all", summary="[관리] 완전 초기화(자사 스코프, 타테넌트 무영향)")
async def reset_all(engine: TenantEngine = Depends(get_tenant_engine)):
err = await _repo(engine).reset_all()
return {"success": err == ErrorType.SUCCESS, "company_id": engine.company_id, "desc": err.name}
class InvalidateReq(BaseModel):
session_id: str
reason: Optional[str] = "invalidated"
@router.post("/invalidate-session", summary="세션 학습 데이터 무효화")
async def invalidate_session(req: InvalidateReq, engine: TenantEngine = Depends(get_tenant_engine)):
repo = _repo(engine)
err = await repo.write_one(lambda s: repo.invalidate_by_session(s, req.session_id, req.reason))
return {"success": err == ErrorType.SUCCESS, "session_id": req.session_id, "desc": err.name}
class TrainReq(BaseModel):
epochs: int = 1
@router.post("/train", summary="오프라인 Q-Learning 학습(경험로그 기반)")
async def train(req: TrainReq, engine: TenantEngine = Depends(get_tenant_engine)):
repo = _repo(engine)
transitions = await repo.load_transitions()
if not transitions:
return {"success": True, "trained_transitions": 0, "note": "경험 로그 없음 (먼저 /chat 진행)"}
S, A = engine.state_space_size, engine.action_space_size
qt = QTable(S, A, learning_rate=engine.config.policy.learning_rate, discount_factor=engine.config.policy.gamma)
# 활성 버전에서 현재 Q 적재
version_id = await repo.get_or_create_active_version(
state_space_size=S, action_space_size=A,
learning_rate=engine.config.policy.learning_rate, discount_factor=engine.config.policy.gamma)
qrows, vrows = await repo.load_cells(version_id)
for st, a, q in qrows:
if 0 <= st < S and 0 <= a < A:
qt.q[st, a] = q
policy = UCBQTablePolicy(qt, mark_visits=False)
updates = 0
touched = set()
for _ in range(max(1, req.epochs)):
for st, a, r, ns, done in transitions:
if st is None or a is None or not (0 <= st < S and 0 <= a < A):
continue
policy.update(Transition(state_index=st, action_id=a, reward=r, next_state_index=ns, done=bool(done)))
touched.add((st, a))
updates += 1
# 갱신된 셀 영속화
for (st, a) in touched:
await QTablePolicyStore.persist_cell(repo, version_id, policy, st, a)
return {"success": True, "trained_transitions": len(transitions), "epochs": req.epochs,
"updates": updates, "cells_persisted": len(touched), "company_id": engine.company_id}
@router.get("/verification-report", summary="Q-Table 검증 리포트(요약)")
async def verification_report(engine: TenantEngine = Depends(get_tenant_engine)):
repo = _repo(engine)
err, active = await repo.read(lambda s: repo.get_active_version(s))
err, total_exp = await repo.read(lambda s: repo.count_experience(s))
if not active:
return {"company_id": engine.company_id, "active_version": None, "experience_total": total_exp}
qrows, vrows = await repo.load_cells(active.version_id)
nonzero = [(s, a, q) for s, a, q in qrows if q != 0]
top = sorted(nonzero, key=lambda x: -x[2])[:10]
return {
"company_id": engine.company_id, "active_version": active.version_name,
"experience_total": total_exp, "q_value_rows": len(qrows), "nonzero_q": len(nonzero),
"visited_cells": len(vrows),
"top_q": [{"state_index": s, "action_id": a, "q_value": round(q, 4)} for s, a, q in top],
}