Coverage for src / monte_neo / indicators / code_gen.py: 100%

39 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-01-28 16:27 +0200

1"""Dynamic code generation module.""" 

2 

3from __future__ import annotations 

4 

5import numpy as np 

6 

7 

8class CodeGenerator: 

9 """Generates random Python code for DynamicIndicator.""" 

10 

11 def __init__(self, rng: np.random.Generator | None = None): 

12 """Initialize generator. 

13 

14 Args: 

15 rng: Random number generator. 

16 """ 

17 self.rng = rng or np.random.default_rng() 

18 

19 def generate_code(self, depth: int = 0) -> str: 

20 """Generate a random valid Python expression for an indicator.""" 

21 # Operands 

22 operands = [ 

23 "data['close']", 

24 "data['open']", 

25 "data['high']", 

26 "data['low']", 

27 "data['volume']", 

28 ] 

29 

30 # Terminal condition (max depth or random stop) 

31 if depth >= 3 or (depth > 0 and self.rng.random() < 0.3): 

32 return self.rng.choice(operands) 

33 

34 # Operators / Functions 

35 # 0: Binary Op, 1: Unary/Func 

36 op_type = self.rng.integers(0, 2) 

37 

38 if op_type == 0: 

39 # Binary 

40 # For simplicity, let's stick to arithmetic and let DynamicIndicator handle >0 logic 

41 # UNLESS we explicitly want boolean signals. 

42 # The current DynamicIndicator maps >0 to 1, <0 to -1. 

43 # So (Close - MA) is good. 

44 

45 op = self.rng.choice(["+", "-", "*", "/"]) 

46 left = self.generate_code(depth + 1) 

47 right = self.generate_code(depth + 1) 

48 return f"({left} {op} {right})" 

49 

50 else: 

51 # Functions 

52 # rolling_mean, diff, shift, rsi, bbands, macd 

53 

54 func_type = self.rng.choice(["mean", "max", "min", "std", "diff", "shift", "rsi", "bbands", "macd"]) 

55 period = int(self.rng.integers(3, 50)) 

56 inner = self.generate_code(depth + 1) 

57 

58 if func_type == "mean": 

59 return f"{inner}.rolling({period}).mean()" 

60 elif func_type == "max": 

61 return f"{inner}.rolling({period}).max()" 

62 elif func_type == "min": 

63 return f"{inner}.rolling({period}).min()" 

64 elif func_type == "std": 

65 return f"{inner}.rolling({period}).std()" 

66 elif func_type == "diff": 

67 return f"{inner}.diff()" 

68 elif func_type == "shift": 

69 return f"{inner}.shift({period})" 

70 elif func_type == "rsi": 

71 return f"rsi({inner}, {period})" 

72 elif func_type == "bbands": 

73 return f"({inner} - {inner}.rolling({period}).mean()) / {inner}.rolling({period}).std()" 

74 elif func_type == "macd": 

75 fast = period 

76 slow = int(period * 2.2) 

77 return f"({inner}.ewm(span={fast}).mean() - {inner}.ewm(span={slow}).mean())" 

78 

79 return "data['close']" # Fallback