""" Feature Importance Tracker — track which features drive planner decisions. Enables: - Understanding which market features matter most - Identifying overfitting to specific features - Guiding feature engineering - Explaining decision rationale """ from __future__ import annotations import time from collections import defaultdict from dataclasses import dataclass, field from typing import Any, Dict, List, Mapping, Optional, Tuple from malkhut.state import MarketWorldState from malkhut.actions import FulfilmentAction from malkhut.features import DefaultFeatureExtractor, FeatureExtractor @dataclass(frozen=True, slots=True) class FeatureImportance: """Importance score for a feature in a specific context.""" feature_name: str importance: float regime: str action_type: str sample_count: int class FeatureImportanceTracker: """ Track which features drive planner decisions. Uses a simple attribution method: - When a decision is made, record which features were above/below thresholds - Aggregate across decisions to compute importance scores """ def __init__(self, feature_extractor: Optional[FeatureExtractor] = None) -> None: self._extractor = feature_extractor or DefaultFeatureExtractor() self._feature_counts: Dict[str, Dict[str, int]] = defaultdict(lambda: defaultdict(int)) self._feature_values: Dict[str, List[float]] = defaultdict(list) self._total_decisions = 0 def record_decision( self, state: MarketWorldState, action: FulfilmentAction, regime: str = "unknown", ) -> None: """Record which features were relevant for this decision.""" self._total_decisions += 1 fv = self._extractor.extract(state).values # Track which features were "active" (non-zero or above threshold) for name, value in fv.items(): if abs(value) > 1e-6: # non-zero self._feature_counts[name][regime] += 1 self._feature_counts[name]["_total"] += 1 # Track value distribution self._feature_values[name].append(value) def get_importance( self, top_n: int = 10, regime: Optional[str] = None, ) -> List[FeatureImportance]: """Get top N most important features.""" scores = [] for name, regime_counts in self._feature_counts.items(): total = regime_counts.get("_total", 0) if regime: count = regime_counts.get(regime, 0) else: count = total importance = count / max(self._total_decisions, 1) scores.append(FeatureImportance( feature_name=name, importance=importance, regime=regime or "all", action_type="all", sample_count=count, )) scores.sort(key=lambda s: s.importance, reverse=True) return scores[:top_n] def get_feature_stats(self, feature_name: str) -> Dict[str, float]: """Get statistics for a specific feature.""" values = self._feature_values.get(feature_name, []) if not values: return {} return { "mean": sum(values) / len(values), "min": min(values), "max": max(values), "count": len(values), } @property def total_decisions(self) -> int: return self._total_decisions @property def feature_count(self) -> int: return len(self._feature_counts)