150 lines
6.8 KiB
Python
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],
|
|
}
|