111 lines
3.5 KiB
Python
111 lines
3.5 KiB
Python
|
|
"""
|
||
|
|
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)
|