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

124 lines
4.0 KiB
Python
Raw Normal View History

"""
VBT Post-Simulation Analysis — vectorized trade analysis and metrics.
Uses VectorBT for analyzing MALKHUT's trade records after the engine
produces results. This is an ANALYSIS tool, not an engine component.
VBT is used here for:
- Trade record → performance metrics (Sharpe, Sortino, VaR, CVaR)
- Multi-asset equity curves and comparisons
- Parameter sensitivity visualization
- Portfolio-level risk analysis
Usage:
from malkhut.training.vbt_analysis import analyze_episodes
metrics = analyze_episodes(results)
"""
from __future__ import annotations
from typing import Any, Dict, List, Optional, Sequence
import numpy as np
from malkhut.training.cma_trainer import EpisodeResult
def episodes_to_pnl_array(
results: Sequence[EpisodeResult],
) -> np.ndarray:
"""Convert episode results to a flat PnL array."""
return np.array([r.pnl_bps for r in results], dtype=np.float64)
def episodes_to_metrics(results: Sequence[EpisodeResult]) -> Dict[str, float]:
"""Compute standard performance metrics from episode results.
Returns dict with: total_pnl, mean_pnl, max_drawdown, sharpe, sortino,
win_rate, profit_factor, fill_ratio, avg_slippage.
"""
if not results:
return {}
pnls = episodes_to_pnl_array(results)
n = len(pnls)
total_pnl = float(np.sum(pnls))
mean_pnl = float(np.mean(pnls))
std_pnl = float(np.std(pnls, ddof=1)) if n > 1 else 0.0
# Sharpe (annualized, assuming ~30 trading days)
sharpe = (mean_pnl / std_pnl * np.sqrt(30)) if std_pnl > 0 else 0.0
# Sortino (downside deviation)
downside = pnls[pnls < 0]
downside_std = float(np.std(downside, ddof=1)) if len(downside) > 1 else 0.0
sortino = (mean_pnl / downside_std * np.sqrt(30)) if downside_std > 0 else 0.0
# Max drawdown (cumulative)
cum_pnl = np.cumsum(pnls)
peak = np.maximum.accumulate(cum_pnl)
drawdowns = peak - cum_pnl
max_dd = float(np.max(drawdowns)) if len(drawdowns) > 0 else 0.0
# Win rate
wins = np.sum(pnls > 0)
win_rate = float(wins / n) if n > 0 else 0.0
# Profit factor
gross_profit = float(np.sum(pnls[pnls > 0])) if np.any(pnls > 0) else 0.0
gross_loss = float(np.abs(np.sum(pnls[pnls < 0]))) if np.any(pnls < 0) else 1e-12
profit_factor = gross_profit / gross_loss if gross_loss > 0 else float("inf")
# Fill ratio
fill_ratios = [r.fill_ratio for r in results]
avg_fill_ratio = sum(fill_ratios) / len(fill_ratios) if fill_ratios else 0.0
# Tail risk
tail_losses = [r.tail_loss_bps for r in results if r.tail_loss_bps < 0]
worst_tail = min(tail_losses) if tail_losses else 0.0
return {
"total_pnl_bps": total_pnl,
"mean_pnl_bps": mean_pnl,
"std_pnl_bps": std_pnl,
"sharpe_ratio": sharpe,
"sortino_ratio": sortino,
"max_drawdown_bps": max_dd,
"win_rate": win_rate,
"profit_factor": profit_factor,
"avg_fill_ratio": avg_fill_ratio,
"worst_tail_bps": worst_tail,
"n_episodes": n,
}
def cross_asset_comparison(
results_by_asset: Dict[str, Sequence[EpisodeResult]],
) -> Dict[str, Dict[str, float]]:
"""Compare performance across assets."""
comparisons = {}
for asset, results in results_by_asset.items():
comparisons[asset] = episodes_to_metrics(results)
return comparisons
def parameter_sensitivity(
results_by_param: Dict[str, Sequence[EpisodeResult]],
) -> Dict[str, float]:
"""Compare performance across parameter variations."""
sensitivities = {}
for param_key, results in results_by_param.items():
metrics = episodes_to_metrics(results)
sensitivities[param_key] = metrics.get("mean_pnl_bps", 0.0)
return sensitivities
def format_metrics(metrics: Dict[str, float]) -> str:
"""Pretty-print metrics dict."""
lines = []
for key, val in sorted(metrics.items()):
if isinstance(val, float):
lines.append(f" {key:25s} {val:>12.2f}")
else:
lines.append(f" {key:25s} {val:>12}")
return "\n".join(lines)