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

1"""Base indicator class. 

2 

3Abstract base class for all trading indicators. 

4""" 

5 

6from __future__ import annotations 

7 

8from abc import ABC, abstractmethod 

9from dataclasses import dataclass, field 

10from typing import Any 

11 

12import numpy as np 

13import pandas as pd 

14 

15from monte_neo.utils.logger import get_logger 

16 

17logger = get_logger(__name__) 

18 

19 

20@dataclass 

21class IndicatorConfig: 

22 """Indicator configuration.""" 

23 

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) 

28 

29 

30class BaseIndicator(ABC): 

31 """Abstract base class for trading indicators.""" 

32 

33 def __init__(self, config: IndicatorConfig | None = None) -> None: 

34 """Initialize indicator. 

35 

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) 

41 

42 @property 

43 def name(self) -> str: 

44 """Get indicator name.""" 

45 return self.config.name 

46 

47 def get_parameters(self) -> dict[str, Any]: 

48 """Get all parameters. 

49 

50 Returns: 

51 Dictionary of parameter names to values. 

52 """ 

53 return dict(self._parameters) 

54 

55 def set_parameter(self, name: str, value: Any) -> None: 

56 """Set a parameter value. 

57 

58 Args: 

59 name: Parameter name. 

60 value: Parameter value. 

61 """ 

62 self._parameters[name] = value 

63 logger.debug(f"Set {name}={value}") 

64 

65 def set_parameters(self, params: dict[str, Any]) -> None: 

66 """Set multiple parameters. 

67 

68 Args: 

69 params: Dictionary of parameters. 

70 """ 

71 for name, value in params.items(): 

72 self.set_parameter(name, value) 

73 

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}" 

79 

80 @abstractmethod 

81 def calculate(self, data: pd.DataFrame) -> pd.DataFrame: 

82 """Calculate indicator values. 

83 

84 Args: 

85 data: OHLCV DataFrame. 

86 

87 Returns: 

88 DataFrame with indicator values. 

89 """ 

90 pass 

91 

92 @abstractmethod 

93 def generate_signals(self, data: pd.DataFrame) -> pd.DataFrame: 

94 """Generate trading signals. 

95 

96 Args: 

97 data: OHLCV DataFrame. 

98 

99 Returns: 

100 DataFrame with 'signal' column (1=buy, -1=sell, 0=hold). 

101 """ 

102 pass 

103 

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 

107 

108 def to_mlx_representation(self) -> Any | None: 

109 """Convert to MLX representation for GPU execution. 

110  

111 Returns: 

112 MLX Strategy object or None if not supported. 

113 """ 

114 return None 

115 

116 def generate_signals_fast(self, data: pd.DataFrame | np.ndarray) -> np.ndarray: 

117 """Fast version of signal generation returning numpy array. 

118  

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) 

127 

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) 

133 

134 def get_formula(self) -> str: 

135 """Get the formula or logic of the indicator.""" 

136 return self.name 

137 

138 def validate_data(self, data: pd.DataFrame) -> bool: 

139 """Validate input data. 

140 

141 Args: 

142 data: OHLCV DataFrame. 

143 

144 Returns: 

145 True if data is valid. 

146 """ 

147 required_cols = ["open", "high", "low", "close", "volume"] 

148 

149 for col in required_cols: 

150 if col not in data.columns: 

151 logger.warning(f"Missing column: {col}") 

152 return False 

153 

154 if len(data) < 10: 

155 logger.warning("Insufficient data points") 

156 return False 

157 

158 return True 

159 

160 def get_min_periods(self) -> int: 

161 """Get minimum periods required for calculation. 

162 

163 Returns: 

164 Minimum number of periods. 

165 """ 

166 return 1 

167 

168 def to_dict(self) -> dict: 

169 """Convert indicator to dictionary. 

170 

171 Returns: 

172 Serializable dictionary. 

173 """ 

174 return { 

175 "name": self.name, 

176 "class": self.__class__.__name__, 

177 "parameters": self.get_parameters(), 

178 } 

179 

180 @classmethod 

181 def from_dict(cls, data: dict) -> BaseIndicator: 

182 """Create indicator from dictionary. 

183 

184 Args: 

185 data: Dictionary representation. 

186 

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