malkhut(T4): Strategy DSL v2 + generator + supporting modules
Strategy DSL v2 (dsl.py): 40+ action primitives, 40+ market sensors, 12 comparison operators, 16 builtins, full parser. Strategy Generator (generator.py): genetic programming evolution — crossover, mutation, tournament selection, pool management. Supporting: discrepancy tracking, execution quality, hooks, feature importance, observability, parallel eval, auto-rollback, stress testing, structured observations, trajectory recording.
This commit is contained in:
110
MALKHUT/malkhut/training/importance.py
Normal file
110
MALKHUT/malkhut/training/importance.py
Normal file
@@ -0,0 +1,110 @@
|
||||
"""
|
||||
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)
|
||||
Reference in New Issue
Block a user