Files
sentiment-engine/MALKHUT/malkhut/training/trajectory.py

105 lines
3.1 KiB
Python
Raw Normal View History

"""
Trajectory Persistence — store CWM trajectories to ClickHouse.
Enables:
- Post-hoc analysis of decision quality
- Replay verification against stored trajectories
- Training data for learned leaf values
- Audit trail for every decision
"""
from __future__ import annotations
import json
import time
from dataclasses import dataclass, field
from typing import Any, List, Mapping, Optional, Sequence
from malkhut.state import MarketWorldState
from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision
from malkhut.storage.ch_store import MalkhutCHStore
@dataclass(frozen=True, slots=True)
class TrajectoryStep:
"""One step in a persisted trajectory."""
step_index: int
ts_ns: int
symbol: str
state_hash: str
action_kind: str
action_side: Optional[str]
action_price: Optional[float]
action_qty: float
next_state_hash: str
pnl_bps: float
reward: float
entropy: float
diagnostics: Mapping[str, Any] = field(default_factory=dict)
@dataclass(frozen=True, slots=True)
class TrajectoryRecord:
"""Complete trajectory for one episode."""
trajectory_id: str
policy_version: str
scenario_id: str
seed: int
steps: Tuple[TrajectoryStep, ...]
total_pnl_bps: float
max_drawdown_bps: float
fill_count: int
cancel_count: int
noop_count: int
duration_ns: int
created_ts_ns: int
class TrajectoryPersister:
"""
Persist CWM trajectories to ClickHouse.
Enables post-hoc analysis and replay verification.
"""
def __init__(self, store: MalkhutCHStore) -> None:
self._store = store
def persist(self, record: TrajectoryRecord) -> None:
"""Persist a trajectory record to CH."""
self._store.store_episode(
policy_version=record.policy_version,
scenario_id=record.scenario_id,
seed=record.seed,
pnl_bps=record.total_pnl_bps,
max_drawdown_bps=record.max_drawdown_bps,
fill_ratio=record.fill_count / max(record.fill_count + record.cancel_count + record.noop_count, 1),
adverse_fill_ratio=0.0,
avg_slippage_bps=0.0,
liq_near_misses=0,
cancel_count=record.cancel_count,
diagnostics=json.dumps({
"trajectory_id": record.trajectory_id,
"steps": len(record.steps),
"fill_count": record.fill_count,
"noop_count": record.noop_count,
"duration_ns": record.duration_ns,
}),
)
def query_trajectories(
self,
policy_version: Optional[str] = None,
scenario_id: Optional[str] = None,
limit: int = 100,
) -> str:
"""Query stored trajectories."""
where_parts = []
if policy_version:
where_parts.append(f"policy_version = '{policy_version}'")
if scenario_id:
where_parts.append(f"scenario_id = '{scenario_id}'")
where = " AND ".join(where_parts) if where_parts else "1=1"
return self._store.query(
f"SELECT * FROM self_play_episodes WHERE {where} ORDER BY ts_ns DESC LIMIT {limit}"
)