Jev 기반 내담자 감정 상태와 응답 일관성 개선
This commit is contained in:
parent
77f8421818
commit
8344bc2ad2
23 changed files with 3384 additions and 25 deletions
501
scripts/probe-jev-dialogue.py
Normal file
501
scripts/probe-jev-dialogue.py
Normal file
|
|
@ -0,0 +1,501 @@
|
|||
#!/usr/bin/env python3
|
||||
"""로컬 엔진과 Jev를 DB 없이 잇는 합성 대화 smoke 수집기."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||
API_ROOT = REPO_ROOT / "apps" / "api"
|
||||
TURN_TIMEOUT_SECONDS = 90.0
|
||||
ALLOWED_PHASES = frozenset({"legacy", "jev"})
|
||||
SOURCE_PATHS = (
|
||||
Path(__file__).resolve(),
|
||||
API_ROOT / "app" / "config.py",
|
||||
API_ROOT / "app" / "engine_client.py",
|
||||
API_ROOT / "app" / "services" / "client_affect.py",
|
||||
API_ROOT / "app" / "services" / "jev_client.py",
|
||||
API_ROOT / "app" / "services" / "orchestrator.py",
|
||||
API_ROOT / "app" / "services" / "persona.py",
|
||||
API_ROOT / "engine_gateway" / "gateway.py",
|
||||
)
|
||||
SYNTHETIC_UTTERANCES: tuple[dict[str, str], ...] = (
|
||||
{
|
||||
"id": "empathy",
|
||||
"label": "empathy",
|
||||
"text": "그 이야기를 꺼내 주셔서 고마워요. 많이 버거웠겠어요.",
|
||||
},
|
||||
{
|
||||
"id": "contradictory_recall",
|
||||
"label": "unsupported_recall",
|
||||
"text": "지난번 여동생과 다퉜다고 하셨죠. 그때 어떤 느낌이었나요?",
|
||||
},
|
||||
{
|
||||
"id": "ambivalence",
|
||||
"label": "ambivalence",
|
||||
"text": "한편으로는 달라지고 싶고, 다른 한편으로는 아직 두려운 마음도 있을 수 있겠어요.",
|
||||
},
|
||||
)
|
||||
SYNTHETIC_MEMORY = {
|
||||
"pinned_facts": ["나는 외동이다."],
|
||||
"expectations": {
|
||||
"unsupported_recall": "상담자의 새 주장을 과거 사실로 확정하지 않음",
|
||||
"contradictory_recall": "형제자매 존재에 동의하지 않음",
|
||||
},
|
||||
}
|
||||
EMOTION_KEYS = (
|
||||
"anxiety",
|
||||
"sadness",
|
||||
"anger",
|
||||
"shame",
|
||||
"guilt",
|
||||
"loneliness",
|
||||
"relief",
|
||||
"hope",
|
||||
"trust",
|
||||
)
|
||||
SAFE_JEV_CODES = frozenset(
|
||||
{
|
||||
"not_configured",
|
||||
"not_started",
|
||||
"timeout",
|
||||
"unauthorized",
|
||||
"insufficient_credits",
|
||||
"forbidden",
|
||||
"model_unavailable",
|
||||
"rate_limited",
|
||||
"overloaded",
|
||||
"http_error",
|
||||
"transport",
|
||||
"malformed_response",
|
||||
"model_mismatch",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _turn_count(value: str) -> int:
|
||||
try:
|
||||
turns = int(value)
|
||||
except ValueError as exc:
|
||||
raise argparse.ArgumentTypeError("turns must be an integer from 1 to 3") from exc
|
||||
if not 1 <= turns <= len(SYNTHETIC_UTTERANCES):
|
||||
raise argparse.ArgumentTypeError("turns must be from 1 to 3")
|
||||
return turns
|
||||
|
||||
|
||||
def _repeat_count(value: str) -> int:
|
||||
try:
|
||||
repeats = int(value)
|
||||
except ValueError as exc:
|
||||
raise argparse.ArgumentTypeError("repeats must be an integer from 1 to 5") from exc
|
||||
if not 1 <= repeats <= 5:
|
||||
raise argparse.ArgumentTypeError("repeats must be from 1 to 5")
|
||||
return repeats
|
||||
|
||||
|
||||
def _phase_list(value: str) -> tuple[str, ...]:
|
||||
phases = tuple(part.strip() for part in value.split(",") if part.strip())
|
||||
if not phases or any(phase not in ALLOWED_PHASES for phase in phases):
|
||||
raise argparse.ArgumentTypeError("phases must be a comma-separated subset of legacy,jev")
|
||||
if len(set(phases)) != len(phases):
|
||||
raise argparse.ArgumentTypeError("phases must not contain duplicates")
|
||||
return phases
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="DB 없이 local engine과 Jev의 합성 가상 내담자 대화를 smoke 수집한다."
|
||||
)
|
||||
parser.add_argument("--output", type=Path, required=True, help="JSON 결과 저장 경로")
|
||||
parser.add_argument(
|
||||
"--turns",
|
||||
type=_turn_count,
|
||||
default=2,
|
||||
help="각 phase의 합성 발화 수(1~3, 기본 2)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--phases",
|
||||
type=_phase_list,
|
||||
default=("legacy", "jev"),
|
||||
help="실행할 phase 목록(legacy,jev; 기본 legacy,jev)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--repeats",
|
||||
type=_repeat_count,
|
||||
default=1,
|
||||
help="phase 묶음 반복 횟수(1~5, 기본 1)",
|
||||
)
|
||||
parser.add_argument("--label", default="", help="측정 보고서 식별 문자열")
|
||||
return parser
|
||||
|
||||
|
||||
def _load_runtime() -> dict[str, Any]:
|
||||
"""API 모듈이 cwd 기반 설정을 읽도록 한 뒤 작업 cwd를 즉시 복구한다."""
|
||||
previous_cwd = Path.cwd()
|
||||
inserted_path = False
|
||||
try:
|
||||
os.chdir(API_ROOT)
|
||||
api_root_text = str(API_ROOT)
|
||||
if api_root_text not in sys.path:
|
||||
sys.path.insert(0, api_root_text)
|
||||
inserted_path = True
|
||||
from app.config import settings
|
||||
from app.engine_client import EngineError, engine_client
|
||||
from app.services import orchestrator, persona, state_machine
|
||||
from app.services.jev_client import JevError, jev_client
|
||||
finally:
|
||||
os.chdir(previous_cwd)
|
||||
if inserted_path:
|
||||
sys.path.remove(str(API_ROOT))
|
||||
return {
|
||||
"settings": settings,
|
||||
"EngineError": EngineError,
|
||||
"engine_client": engine_client,
|
||||
"orchestrator": orchestrator,
|
||||
"persona": persona,
|
||||
"state_machine": state_machine,
|
||||
"JevError": JevError,
|
||||
"jev_client": jev_client,
|
||||
}
|
||||
|
||||
|
||||
def _safe_error_code(error: BaseException, runtime: dict[str, Any]) -> str:
|
||||
if isinstance(error, runtime["JevError"]):
|
||||
code = getattr(error, "code", "")
|
||||
if code in SAFE_JEV_CODES:
|
||||
return f"client_affect_{code}"
|
||||
return "client_affect_error"
|
||||
if isinstance(error, runtime["EngineError"]):
|
||||
return "engine_error"
|
||||
if isinstance(error, TimeoutError):
|
||||
return "turn_timeout"
|
||||
return "runtime_error"
|
||||
|
||||
|
||||
def _safe_stream_error(value: object) -> str:
|
||||
detail = str(value).strip()
|
||||
if detail.startswith("client_affect_"):
|
||||
code = detail.removeprefix("client_affect_")
|
||||
if code in SAFE_JEV_CODES:
|
||||
return detail
|
||||
if detail in {"client_stream_incomplete", "engine stream error", "engine stream decode error"}:
|
||||
return detail.replace(" ", "_")
|
||||
return "stream_error"
|
||||
|
||||
|
||||
def _final_emotions(affect_state: dict[str, Any]) -> dict[str, float]:
|
||||
emotions: dict[str, float] = {}
|
||||
for dimension in EMOTION_KEYS:
|
||||
value = affect_state.get(f"emotion_{dimension}")
|
||||
if isinstance(value, (int, float)) and not isinstance(value, bool) and math.isfinite(value):
|
||||
emotions[dimension] = float(value)
|
||||
return emotions
|
||||
|
||||
|
||||
async def _run_turn(
|
||||
*,
|
||||
runtime: dict[str, Any],
|
||||
state: Any,
|
||||
session_id: str,
|
||||
utterance: str,
|
||||
recent_turns: list[dict[str, str]],
|
||||
) -> tuple[dict[str, Any], Any, str | None]:
|
||||
orchestrator = runtime["orchestrator"]
|
||||
persona = runtime["persona"]
|
||||
context = orchestrator.prepare_turn(
|
||||
session_id=session_id,
|
||||
case_id=None,
|
||||
card=persona.P1,
|
||||
state=state,
|
||||
learner_text=utterance,
|
||||
learner_identity="합성 상담자",
|
||||
memory=orchestrator.TurnMemory(
|
||||
recent_turns=recent_turns,
|
||||
pinned_facts=list(SYNTHETIC_MEMORY["pinned_facts"]),
|
||||
),
|
||||
theory_mode="humanistic",
|
||||
scenario_context=None,
|
||||
)
|
||||
started = time.perf_counter()
|
||||
first_token_at: float | None = None
|
||||
generated: list[str] = []
|
||||
done_payload: dict[str, Any] | None = None
|
||||
error_code: str | None = None
|
||||
try:
|
||||
async with asyncio.timeout(TURN_TIMEOUT_SECONDS):
|
||||
async for event in orchestrator.run_turn_stream(context, runtime["engine_client"]):
|
||||
now = time.perf_counter()
|
||||
if event.event == "token":
|
||||
if first_token_at is None:
|
||||
first_token_at = now
|
||||
generated.append(str(event.data.get("text", "")))
|
||||
elif event.event == "done":
|
||||
done_payload = dict(event.data)
|
||||
elif event.event == "error":
|
||||
error_code = _safe_stream_error(event.data.get("detail"))
|
||||
break
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
error_code = _safe_error_code(exc, runtime)
|
||||
|
||||
total_ms = round((time.perf_counter() - started) * 1000, 1)
|
||||
if done_payload is None and error_code is None:
|
||||
error_code = "stream_error"
|
||||
generated_text = "".join(generated)
|
||||
record: dict[str, Any] = {
|
||||
"status": "done" if done_payload is not None and error_code is None else "error",
|
||||
"ttft_ms": (
|
||||
None
|
||||
if first_token_at is None
|
||||
else round((first_token_at - started) * 1000, 1)
|
||||
),
|
||||
"total_ms": total_ms,
|
||||
"generation": {
|
||||
"provider": None if done_payload is None else done_payload.get("llm_provider"),
|
||||
"model": None if done_payload is None else done_payload.get("model"),
|
||||
},
|
||||
"appraisal": context.client_affect_metadata,
|
||||
"final_emotions": _final_emotions(
|
||||
context.state_after.affect_state
|
||||
if done_payload is not None and error_code is None
|
||||
else state.affect_state
|
||||
),
|
||||
"text": generated_text,
|
||||
}
|
||||
if error_code is not None:
|
||||
record["error"] = error_code
|
||||
return record, context.state_after, generated_text if done_payload is not None and error_code is None else None
|
||||
|
||||
|
||||
async def _run_phase(
|
||||
*,
|
||||
runtime: dict[str, Any],
|
||||
provider_mode: str,
|
||||
turns: int,
|
||||
repeat: int,
|
||||
order_index: int,
|
||||
) -> dict[str, Any]:
|
||||
settings = runtime["settings"]
|
||||
original_provider = settings.client_affect_provider
|
||||
session_id = str(uuid.uuid4())
|
||||
phase: dict[str, Any] = {
|
||||
"provider_mode": provider_mode,
|
||||
"repeat": repeat,
|
||||
"order_index": order_index,
|
||||
"session_id": session_id,
|
||||
"status": "error",
|
||||
"turns": [],
|
||||
"session_closed": False,
|
||||
}
|
||||
settings.client_affect_provider = provider_mode
|
||||
try:
|
||||
state = runtime["state_machine"].init_state(params=runtime["persona"].P1.openness_params())
|
||||
recent_turns: list[dict[str, str]] = []
|
||||
for index, fixture in enumerate(SYNTHETIC_UTTERANCES[:turns], start=1):
|
||||
utterance = fixture["text"]
|
||||
result, next_state, client_reply = await _run_turn(
|
||||
runtime=runtime,
|
||||
state=state,
|
||||
session_id=session_id,
|
||||
utterance=utterance,
|
||||
recent_turns=recent_turns,
|
||||
)
|
||||
result["turn"] = index
|
||||
result["turn_temperature"] = "cold" if index == 1 else "warm"
|
||||
result["utterance_id"] = fixture["id"]
|
||||
result["utterance_kind"] = fixture["label"]
|
||||
if fixture["id"] == "contradictory_recall":
|
||||
result["fixture_expectations"] = [
|
||||
SYNTHETIC_MEMORY["expectations"]["unsupported_recall"],
|
||||
SYNTHETIC_MEMORY["expectations"]["contradictory_recall"],
|
||||
]
|
||||
phase["turns"].append(result)
|
||||
if result["status"] != "done":
|
||||
phase["error"] = result["error"]
|
||||
return phase
|
||||
state = next_state
|
||||
recent_turns.extend(
|
||||
[
|
||||
{"speaker": "counselor", "text": utterance},
|
||||
{"speaker": "client", "text": client_reply or ""},
|
||||
]
|
||||
)
|
||||
phase["status"] = "done"
|
||||
return phase
|
||||
finally:
|
||||
settings.client_affect_provider = original_provider
|
||||
closed = await runtime["engine_client"].close_session(session_id)
|
||||
phase["session_closed"] = closed
|
||||
if not closed:
|
||||
phase["cleanup_error"] = "engine_close_failed"
|
||||
if phase["status"] == "done":
|
||||
phase["status"] = "error"
|
||||
phase["error"] = "engine_close_failed"
|
||||
|
||||
|
||||
def _source_sha256() -> dict[str, str]:
|
||||
return {
|
||||
str(path.relative_to(REPO_ROOT)).replace("\\", "/"): hashlib.sha256(path.read_bytes()).hexdigest()
|
||||
for path in SOURCE_PATHS
|
||||
}
|
||||
|
||||
|
||||
def _percentile(values: list[float], quantile: float) -> float | None:
|
||||
if not values:
|
||||
return None
|
||||
ordered = sorted(values)
|
||||
position = (len(ordered) - 1) * quantile
|
||||
lower = math.floor(position)
|
||||
upper = math.ceil(position)
|
||||
if lower == upper:
|
||||
return ordered[lower]
|
||||
return round(ordered[lower] + (ordered[upper] - ordered[lower]) * (position - lower), 1)
|
||||
|
||||
|
||||
def _latency_summary(values: list[float]) -> dict[str, float | int | None]:
|
||||
return {
|
||||
"sample_count": len(values),
|
||||
"p50": _percentile(values, 0.50),
|
||||
"p95": _percentile(values, 0.95),
|
||||
}
|
||||
|
||||
|
||||
def _phase_summaries(phases: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
grouped: dict[tuple[str, str], list[tuple[dict[str, Any], dict[str, Any]]]] = {}
|
||||
for phase in phases:
|
||||
for turn in phase["turns"]:
|
||||
grouped.setdefault((phase["provider_mode"], turn["turn_temperature"]), []).append((phase, turn))
|
||||
summaries: list[dict[str, Any]] = []
|
||||
for (provider_mode, cache_state), rows in grouped.items():
|
||||
phase_ids = {(phase["repeat"], phase["order_index"]) for phase, _ in rows}
|
||||
turns = [turn for _, turn in rows if turn["status"] == "done"]
|
||||
ttft_values = [turn["ttft_ms"] for turn in turns if turn["ttft_ms"] is not None]
|
||||
total_values = [turn["total_ms"] for turn in turns if turn["total_ms"] is not None]
|
||||
summaries.append(
|
||||
{
|
||||
"provider_mode": provider_mode,
|
||||
"turn_temperature": cache_state,
|
||||
"phase_run_count": len(phase_ids),
|
||||
"completed_phase_run_count": len(
|
||||
{(phase["repeat"], phase["order_index"]) for phase, _ in rows if phase["status"] == "done"}
|
||||
),
|
||||
"attempted_turn_count": len(rows),
|
||||
"completed_turn_count": len(turns),
|
||||
"ttft_ms": _latency_summary(ttft_values),
|
||||
"total_ms": _latency_summary(total_values),
|
||||
}
|
||||
)
|
||||
return summaries
|
||||
|
||||
|
||||
def _generation_settings(runtime: dict[str, Any]) -> dict[str, Any]:
|
||||
settings = runtime["settings"]
|
||||
engine_client = runtime["engine_client"]
|
||||
return {
|
||||
"engine_mode": settings.engine_mode,
|
||||
"live_client_provider": settings.live_client_provider,
|
||||
"model": engine_client.default_model,
|
||||
"reasoning_effort": engine_client.default_reasoning_effort,
|
||||
}
|
||||
|
||||
|
||||
def _phase_order(phases: tuple[str, ...], repeat: int) -> tuple[str, ...]:
|
||||
return phases if repeat % 2 else tuple(reversed(phases))
|
||||
|
||||
|
||||
async def collect(
|
||||
turns: int,
|
||||
phases: tuple[str, ...] = ("legacy", "jev"),
|
||||
repeats: int = 1,
|
||||
label: str = "",
|
||||
) -> dict[str, Any]:
|
||||
report: dict[str, Any] = {
|
||||
"kind": "jev_dialogue_smoke",
|
||||
"provenance": "synthetic_only",
|
||||
"quality_pass": False,
|
||||
"generated_at": datetime.now(timezone.utc).isoformat(),
|
||||
"label": label,
|
||||
"source_sha256": _source_sha256(),
|
||||
"turn_timeout_seconds": TURN_TIMEOUT_SECONDS,
|
||||
"turns_requested": turns,
|
||||
"repeats_requested": repeats,
|
||||
"phases_requested": list(phases),
|
||||
"memory_fixture": SYNTHETIC_MEMORY,
|
||||
"repeat_orders": [],
|
||||
"phases": [],
|
||||
"phase_summaries": [],
|
||||
"generation_settings": None,
|
||||
"limitations": [
|
||||
"cold/warm은 각 새 합성 세션의 첫 turn과 후속 turn을 뜻하며 provider cache 상태는 측정하지 않는다."
|
||||
],
|
||||
}
|
||||
try:
|
||||
runtime = _load_runtime()
|
||||
except Exception:
|
||||
report["status"] = "error"
|
||||
report["error"] = "runtime_import_failed"
|
||||
return report
|
||||
|
||||
engine_client = runtime["engine_client"]
|
||||
jev_client = runtime["jev_client"]
|
||||
report["generation_settings"] = _generation_settings(runtime)
|
||||
try:
|
||||
await engine_client.startup()
|
||||
await jev_client.startup()
|
||||
for repeat in range(1, repeats + 1):
|
||||
order = _phase_order(phases, repeat)
|
||||
report["repeat_orders"].append({"repeat": repeat, "phase_order": list(order)})
|
||||
for order_index, provider_mode in enumerate(order, start=1):
|
||||
phase = await _run_phase(
|
||||
runtime=runtime,
|
||||
provider_mode=provider_mode,
|
||||
turns=turns,
|
||||
repeat=repeat,
|
||||
order_index=order_index,
|
||||
)
|
||||
report["phases"].append(phase)
|
||||
report["phase_summaries"] = _phase_summaries(report["phases"])
|
||||
report["status"] = "done" if all(phase["status"] == "done" for phase in report["phases"]) else "error"
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
report["status"] = "error"
|
||||
report["error"] = _safe_error_code(exc, runtime)
|
||||
finally:
|
||||
await jev_client.shutdown()
|
||||
await engine_client.shutdown()
|
||||
report["phase_summaries"] = _phase_summaries(report["phases"])
|
||||
return report
|
||||
|
||||
|
||||
def _write_report(path: Path, report: dict[str, Any]) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(
|
||||
json.dumps(report, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
args = build_parser().parse_args(argv)
|
||||
report = asyncio.run(collect(args.turns, args.phases, args.repeats, args.label))
|
||||
_write_report(args.output, report)
|
||||
print(f"보고서 저장: {args.output}")
|
||||
return 0 if report["status"] == "done" else 1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Loading…
Add table
Add a link
Reference in a new issue