Coverage for src / monte_neo / indicators / base.py: 61%
71 statements
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-28 16:27 +0200
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-28 16:27 +0200
1"""Base indicator class.
3Abstract base class for all trading indicators.
4"""
6from __future__ import annotations
8from abc import ABC, abstractmethod
9from dataclasses import dataclass, field
10from typing import Any
12import numpy as np
13import pandas as pd
15from monte_neo.utils.logger import get_logger
17logger = get_logger(__name__)
20@dataclass
21class IndicatorConfig:
22 """Indicator configuration."""
24 name: str
25 parameters: dict[str, Any] = field(default_factory=dict)
26 entry_conditions: list[str] = field(default_factory=list)
27 exit_conditions: list[str] = field(default_factory=list)
30class BaseIndicator(ABC):
31 """Abstract base class for trading indicators."""
33 def __init__(self, config: IndicatorConfig | None = None) -> None:
34 """Initialize indicator.
36 Args:
37 config: Indicator configuration.
38 """
39 self.config = config or IndicatorConfig(name=self.__class__.__name__)
40 self._parameters: dict[str, Any] = dict(self.config.parameters)
42 @property
43 def name(self) -> str:
44 """Get indicator name."""
45 return self.config.name
47 def get_parameters(self) -> dict[str, Any]:
48 """Get all parameters.
50 Returns:
51 Dictionary of parameter names to values.
52 """
53 return dict(self._parameters)
55 def set_parameter(self, name: str, value: Any) -> None:
56 """Set a parameter value.
58 Args:
59 name: Parameter name.
60 value: Parameter value.
61 """
62 self._parameters[name] = value
63 logger.debug(f"Set {name}={value}")
65 def set_parameters(self, params: dict[str, Any]) -> None:
66 """Set multiple parameters.
68 Args:
69 params: Dictionary of parameters.
70 """
71 for name, value in params.items():
72 self.set_parameter(name, value)
74 def get_id(self) -> str:
75 """Get unique identifier for this indicator instance."""
76 import json
77 params_str = json.dumps(self._parameters, sort_keys=True)
78 return f"{self.__class__.__name__}_{params_str}"
80 @abstractmethod
81 def calculate(self, data: pd.DataFrame) -> pd.DataFrame:
82 """Calculate indicator values.
84 Args:
85 data: OHLCV DataFrame.
87 Returns:
88 DataFrame with indicator values.
89 """
90 pass
92 @abstractmethod
93 def generate_signals(self, data: pd.DataFrame) -> pd.DataFrame:
94 """Generate trading signals.
96 Args:
97 data: OHLCV DataFrame.
99 Returns:
100 DataFrame with 'signal' column (1=buy, -1=sell, 0=hold).
101 """
102 pass
104 def get_metal_params(self, commission_bps: float = 0.0, slippage_bps: float = 0.0) -> list[float] | None:
105 """Return parameters for native Metal kernel (5 floats)."""
106 return None
108 def to_mlx_representation(self) -> Any | None:
109 """Convert to MLX representation for GPU execution.
111 Returns:
112 MLX Strategy object or None if not supported.
113 """
114 return None
116 def generate_signals_fast(self, data: pd.DataFrame | np.ndarray) -> np.ndarray:
117 """Fast version of signal generation returning numpy array.
119 Default implementation calls generate_signals and extracts the array.
120 Subclasses should override this for better performance.
121 """
122 if isinstance(data, pd.DataFrame):
123 sigs = self.generate_signals(data)
124 if isinstance(sigs, pd.DataFrame):
125 return sigs["signal"].to_numpy(dtype=np.float32)
126 return np.asarray(sigs, dtype=np.float32)
128 # If it's already a numpy array, we might need a dummy DataFrame
129 # but this is exactly what we want to avoid.
130 # Subclasses MUST override this if they want to support pure numpy paths.
131 dummy_df = pd.DataFrame({"close": data[:, 3] if data.ndim > 1 else data})
132 return self.generate_signals(dummy_df)["signal"].to_numpy(dtype=np.float32)
134 def get_formula(self) -> str:
135 """Get the formula or logic of the indicator."""
136 return self.name
138 def validate_data(self, data: pd.DataFrame) -> bool:
139 """Validate input data.
141 Args:
142 data: OHLCV DataFrame.
144 Returns:
145 True if data is valid.
146 """
147 required_cols = ["open", "high", "low", "close", "volume"]
149 for col in required_cols:
150 if col not in data.columns:
151 logger.warning(f"Missing column: {col}")
152 return False
154 if len(data) < 10:
155 logger.warning("Insufficient data points")
156 return False
158 return True
160 def get_min_periods(self) -> int:
161 """Get minimum periods required for calculation.
163 Returns:
164 Minimum number of periods.
165 """
166 return 1
168 def to_dict(self) -> dict:
169 """Convert indicator to dictionary.
171 Returns:
172 Serializable dictionary.
173 """
174 return {
175 "name": self.name,
176 "class": self.__class__.__name__,
177 "parameters": self.get_parameters(),
178 }
180 @classmethod
181 def from_dict(cls, data: dict) -> BaseIndicator:
182 """Create indicator from dictionary.
184 Args:
185 data: Dictionary representation.
187 Returns:
188 Indicator instance.
189 """
190 config = IndicatorConfig(
191 name=data.get("name", cls.__name__),
192 parameters=data.get("parameters", {}),
193 )
194 instance = cls(config)
195 instance.set_parameters(data.get("parameters", {}))
196 return instance