quantex 0.5.0__tar.gz → 0.5.2__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.3
2
2
  Name: quantex
3
- Version: 0.5.0
3
+ Version: 0.5.2
4
4
  Summary: A simple quant strategy creation and backtesting package.
5
5
  License: MIT
6
6
  Author: Daniel Green
@@ -12,6 +12,8 @@ Classifier: Programming Language :: Python :: 3.11
12
12
  Classifier: Programming Language :: Python :: 3.12
13
13
  Classifier: Programming Language :: Python :: 3.13
14
14
  Requires-Dist: fastparquet (>=2024.11.0,<2025.0.0)
15
+ Requires-Dist: matplotlib (>=3.10.8,<4.0.0)
16
+ Requires-Dist: mplfinance (>=0.12.10b0,<0.13.0)
15
17
  Requires-Dist: numpy (>=2.4.3,<3.0.0)
16
18
  Requires-Dist: optuna (>=4.8.0,<5.0.0)
17
19
  Requires-Dist: pandas (>=2.3.0,<3.0.0)
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "quantex"
3
- version = "0.5.0"
3
+ version = "0.5.2"
4
4
  description = "A simple quant strategy creation and backtesting package."
5
5
  authors = [
6
6
  {name = "Daniel Green",email = "dangreen07@outlook.com"}
@@ -15,6 +15,8 @@ dependencies = [
15
15
  "tqdm (>=4.67.1,<5.0.0)",
16
16
  "numpy (>=2.4.3,<3.0.0)",
17
17
  "optuna (>=4.8.0,<5.0.0)",
18
+ "matplotlib (>=3.10.8,<4.0.0)",
19
+ "mplfinance (>=0.12.10b0,<0.13.0)",
18
20
  ]
19
21
 
20
22
  [tool.poetry]
@@ -8,6 +8,8 @@ import numpy as np
8
8
  import pandas as pd
9
9
  from tqdm import tqdm
10
10
 
11
+ from ..backtester.walk_forward import WalkForwardResult
12
+
11
13
  from ..broker import Order
12
14
  from ..strategy import Strategy
13
15
 
@@ -153,7 +155,7 @@ class SimpleBacktester:
153
155
  broker.margin_call = self.margin_call
154
156
  broker.leverage = self.leverage
155
157
  broker.commision = np.float64(self.commission)
156
- broker.commision_type = self.commission_type
158
+ broker.commision_type = self.commission_type # type: ignore[assignment]
157
159
 
158
160
  self.strategy.init()
159
161
 
@@ -182,7 +184,8 @@ class SimpleBacktester:
182
184
  tradeRecord.extend(trades)
183
185
  orders.extend(val.complete_orders)
184
186
 
185
- index = list(self.strategy.positions.values())[0].source.data['Close'].index
187
+ source = list(self.strategy.positions.values())[0].source
188
+ index = source.data['Close'].index
186
189
  return BacktestReport(
187
190
  starting_cash=np.float64(self.cash),
188
191
  final_cash=self.PnLRecord[-1],
@@ -193,7 +196,8 @@ class SimpleBacktester:
193
196
  event
194
197
  for broker in self.strategy.positions.values()
195
198
  for event in getattr(broker, "margin_call_events", [])
196
- ] or None)
199
+ ] or None,
200
+ data=source.data.copy())
197
201
 
198
202
  def optimize(
199
203
  self,
@@ -1536,3 +1540,171 @@ class SimpleBacktester:
1536
1540
  seed=seed,
1537
1541
  progress_bar=progress_bar,
1538
1542
  )
1543
+
1544
+ def walk_forward_analyze(
1545
+ self,
1546
+ params: dict[str, Any],
1547
+ train_periods: int,
1548
+ test_periods: int,
1549
+ step_periods: int | None = None,
1550
+ constraint: Callable[[dict[str, Any]], bool] | None = None,
1551
+ objective: str = "sharpe",
1552
+ risk_tolerance: dict[str, float] | None = None,
1553
+ min_train_periods: int = 30,
1554
+ min_test_periods: int = 10,
1555
+ progress_bar: bool = True,
1556
+ **optimizer_kwargs: Any,
1557
+ ) -> "WalkForwardResult":
1558
+ """
1559
+ Perform walk-forward analysis on the strategy.
1560
+
1561
+ Walk-forward analysis is a rigorous method for evaluating trading strategies
1562
+ that simulates real-world deployment conditions. It uses rolling windows to:
1563
+ 1. Train: Optimize parameters on historical data
1564
+ 2. Test: Evaluate best parameters on unseen future data
1565
+
1566
+ This method provides a convenient way to run walk-forward analysis directly
1567
+ on the backtester instance without needing to create a separate WalkForwardAnalyzer.
1568
+
1569
+ Args:
1570
+ params (dict[str, range]): Dictionary mapping strategy attribute names
1571
+ to iterables of candidate values. For example:
1572
+ ```python
1573
+ {
1574
+ 'fast_period': [5, 10, 15, 20],
1575
+ 'slow_period': [20, 30, 40, 50],
1576
+ 'threshold': [0.01, 0.02, 0.05]
1577
+ }
1578
+ ```
1579
+ train_periods (int): Number of periods for each training window.
1580
+ This is the lookback window used for parameter optimization.
1581
+ test_periods (int): Number of periods for each test window.
1582
+ This is the forward-looking window for out-of-sample evaluation.
1583
+ step_periods (int | None, optional): Number of periods to step forward
1584
+ between windows. If None, uses test_periods (non-overlapping windows).
1585
+ Defaults to None.
1586
+ constraint (Callable[[dict[str, Any]], bool] | None, optional): Optional
1587
+ callable that takes a parameter dict and returns True to evaluate
1588
+ or False to skip. Useful for enforcing logical constraints.
1589
+ Defaults to None.
1590
+ objective (str, optional): Metric to optimize. Defaults to "sharpe".
1591
+ Supports: "final_cash", "total_return", "sharpe", "max_drawdown", "trades".
1592
+ risk_tolerance (dict[str, float] | None, optional): Optional maximum
1593
+ allowed values for risk metrics. Defaults to None.
1594
+ min_train_periods (int, optional): Minimum required training periods.
1595
+ Defaults to 30.
1596
+ min_test_periods (int, optional): Minimum required test periods.
1597
+ Defaults to 10.
1598
+ progress_bar (bool, optional): Whether to show progress bar.
1599
+ Defaults to True.
1600
+ **optimizer_kwargs: Additional keyword arguments passed to the
1601
+ optimizer. Can include:
1602
+ - workers (int): For parallel optimization
1603
+ - n_trials (int): For Optuna optimization
1604
+ - timeout (int): For Optuna timeout
1605
+ - random_seed (int): For reproducibility
1606
+
1607
+ Returns:
1608
+ WalkForwardResult: Object containing:
1609
+ - n_windows: Total number of walk-forward windows
1610
+ - train_periods: Periods in each training window
1611
+ - test_periods: Periods in each test window
1612
+ - window_results: List of WalkForwardWindow objects
1613
+ - aggregated_metrics: Aggregated statistics across windows
1614
+ - all_windows_results_df: DataFrame with all results
1615
+ - plot(): Visualization method for results
1616
+
1617
+ Example:
1618
+ >>> # Run walk-forward analysis with grid search
1619
+ >>> result = bt.walk_forward_analyze(
1620
+ ... params={'fast': [5, 10, 20], 'slow': [20, 50, 100]},
1621
+ ... train_periods=252, # 1 year training
1622
+ ... test_periods=63, # 3 months testing
1623
+ ... objective='sharpe'
1624
+ ... )
1625
+ >>> print(result)
1626
+ >>> print(f"Average OOS Sharpe: {result.aggregated_metrics['out_of_sample_sharpe_mean']:.2f}")
1627
+ >>> result.plot()
1628
+
1629
+ >>> # Run with parallel optimization
1630
+ >>> result = bt.walk_forward_analyze(
1631
+ ... params={'period': range(5, 50, 5)},
1632
+ ... train_periods=252,
1633
+ ... test_periods=63,
1634
+ ... workers=4 # Use parallel optimization
1635
+ ... )
1636
+
1637
+ >>> # Run with Optuna optimization
1638
+ >>> result = bt.walk_forward_analyze(
1639
+ ... params={'period': (5, 50)},
1640
+ ... train_periods=252,
1641
+ ... test_periods=63,
1642
+ ... n_trials=100 # Optuna-specific
1643
+ ... )
1644
+
1645
+ Note:
1646
+ - Walk-forward analysis helps detect overfitting by evaluating parameter
1647
+ stability and out-of-sample performance over multiple time windows
1648
+ - A stability ratio (OOS/IS Sharpe) close to 1.0 indicates robust parameters
1649
+ - Use step_periods < test_periods for overlapping windows with more samples
1650
+ """
1651
+ from .walk_forward import (
1652
+ WalkForwardAnalyzer,
1653
+ WalkForwardResult,
1654
+ )
1655
+
1656
+ # Handle step_periods - use test_periods if None
1657
+ actual_step = test_periods if step_periods is None else step_periods
1658
+
1659
+ # Determine which optimizer to use based on kwargs
1660
+ # If workers > 1, use optimize_parallel; if n_trials provided, use optimize_optuna
1661
+ if "n_trials" in optimizer_kwargs:
1662
+ # Use Optuna optimizer
1663
+ optimizer = lambda bt, params, **kw: bt.optimize_optuna(
1664
+ param_space=params,
1665
+ n_trials=kw.get("n_trials", 100),
1666
+ objective=kw.get("objective", objective),
1667
+ risk_tolerance=kw.get("risk_tolerance", risk_tolerance),
1668
+ constraint=kw.get("constraint", constraint),
1669
+ timeout=kw.get("timeout", None),
1670
+ random_seed=kw.get("random_seed", None),
1671
+ workers=kw.get("workers", None),
1672
+ progress_bar=kw.get("progress_bar", False),
1673
+ )
1674
+ elif optimizer_kwargs.get("workers", 1) > 1:
1675
+ # Use parallel optimizer
1676
+ optimizer = lambda bt, params, **kw: bt.optimize_parallel(
1677
+ params=params,
1678
+ constraint=kw.get("constraint", constraint),
1679
+ objective=kw.get("objective", objective),
1680
+ risk_tolerance=kw.get("risk_tolerance", risk_tolerance),
1681
+ workers=kw.get("workers", None),
1682
+ chunksize=kw.get("chunksize", "auto"),
1683
+ )
1684
+ else:
1685
+ # Use sequential optimizer
1686
+ optimizer = lambda bt, params, **kw: bt.optimize(
1687
+ params=params,
1688
+ constraint=kw.get("constraint", constraint),
1689
+ objective=kw.get("objective", objective),
1690
+ risk_tolerance=kw.get("risk_tolerance", risk_tolerance),
1691
+ )
1692
+
1693
+ # Create analyzer and run analysis
1694
+ analyzer = WalkForwardAnalyzer(
1695
+ backtester=self,
1696
+ train_periods=train_periods,
1697
+ test_periods=test_periods,
1698
+ min_train_periods=min_train_periods,
1699
+ min_test_periods=min_test_periods,
1700
+ selection_criterion=objective,
1701
+ )
1702
+
1703
+ return analyzer.analyze(
1704
+ optimizer=optimizer,
1705
+ params=params,
1706
+ constraint=constraint,
1707
+ objective=objective,
1708
+ risk_tolerance=risk_tolerance,
1709
+ progress_bar=progress_bar,
1710
+ )
@@ -1,10 +1,13 @@
1
- from dataclasses import dataclass
1
+ from dataclasses import dataclass, field
2
2
  from typing import Any
3
3
 
4
+ import mplfinance as mpf
4
5
  import numpy as np
5
6
  import pandas as pd
6
7
  from matplotlib import pyplot as plt
7
8
 
9
+ from ..broker.types import Order
10
+
8
11
 
9
12
  @dataclass
10
13
  class OptimizationResult:
@@ -50,14 +53,17 @@ class BacktestReport:
50
53
  final_cash (np.float64): Final cash amount at end of backtest.
51
54
  PnlRecord (pd.Series): Time series of P&L values throughout the backtest.
52
55
  orders (list[Order]): List of all orders executed during the backtest.
56
+ tradeRecord (list[np.float64]): List of individual trade P&L values.
53
57
  margin_call_events (list[dict]): Margin call events triggered during the run.
58
+ data (pd.DataFrame): OHLC price data used in the backtest.
54
59
  """
55
60
  starting_cash: np.float64
56
61
  final_cash: np.float64
57
62
  PnlRecord: pd.Series
58
- orders: list
63
+ orders: list[Order]
59
64
  tradeRecord: list[np.float64]
60
65
  margin_call_events: list[dict] | None = None
66
+ data: pd.DataFrame = field(default_factory=lambda: pd.DataFrame())
61
67
 
62
68
  @property
63
69
  def annual_rf(self):
@@ -161,6 +167,128 @@ class BacktestReport:
161
167
  plt.tight_layout()
162
168
  plt.show()
163
169
 
170
+ def plot_trades(
171
+ self,
172
+ figsize: tuple = (16, 8),
173
+ style: str = "yahoo",
174
+ title: str | None = None,
175
+ start_date: str | None = None,
176
+ end_date: str | None = None,
177
+ volume: bool = True,
178
+ ) -> None:
179
+ """
180
+ Plot the price chart with trade entry and exit markers.
181
+
182
+ Creates a candlestick chart showing all trades with:
183
+ - Green triangles (^) for buy entries
184
+ - Red triangles (v) for sell exits
185
+ - Position is closed when an order in the opposite direction is executed
186
+
187
+ Args:
188
+ figsize (tuple, optional): Figure size as (width, height) in inches.
189
+ Defaults to (16, 8).
190
+ style (str, optional): mplfinance style name. Defaults to "yahoo".
191
+ title (str, optional): Chart title. Defaults to None (uses default title).
192
+ start_date (str, optional): Start date filter (e.g., '2026-03-01').
193
+ Defaults to None (show all).
194
+ end_date (str, optional): End date filter (e.g., '2026-03-31').
195
+ Defaults to None (show all).
196
+ volume (bool, optional): Whether to show volume subplot. Defaults to True.
197
+
198
+ Raises:
199
+ ValueError: If no OHLC data is available in the backtest report.
200
+ ValueError: If no orders are available.
201
+ """
202
+ from ..broker.types import OrderSide
203
+
204
+ if self.data.empty:
205
+ raise ValueError(
206
+ "No OHLC data available in backtest report. "
207
+ "Ensure the backtest was run with data storage enabled."
208
+ )
209
+
210
+ if not self.orders:
211
+ raise ValueError("No orders available in backtest report.")
212
+
213
+ # Filter data by date range if specified
214
+ data = self.data.copy()
215
+ if start_date is not None:
216
+ data = data[data.index > start_date]
217
+ if end_date is not None:
218
+ data = data[data.index <= end_date]
219
+
220
+ # Create marker series with the same index as the data
221
+ buy_signal = pd.Series(index=data.index, dtype=float)
222
+ sell_signal = pd.Series(index=data.index, dtype=float)
223
+
224
+ for order in self.orders:
225
+ timestamp = order.timestamp
226
+
227
+ # Only set signal if timestamp exists in the data index
228
+ if timestamp not in data.index:
229
+ continue
230
+
231
+ # Set the signal at the order timestamp
232
+ if order.side == OrderSide.BUY:
233
+ buy_signal.loc[timestamp] = data.loc[timestamp, 'Close'] # type: ignore
234
+ else: # OrderSide.SELL
235
+ sell_signal.loc[timestamp] = data.loc[timestamp, 'Close'] # type: ignore
236
+
237
+ # If no trades to plot and no data, raise an error
238
+ if data.empty:
239
+ raise ValueError("No data available to plot. Check your date range.")
240
+
241
+ # Check if there are actual trade markers (non-NaN values)
242
+ has_buy_signals = not buy_signal.dropna().empty
243
+ has_sell_signals = not sell_signal.dropna().empty
244
+
245
+ # Build addplot list with markers (only if we have data points)
246
+ apds = []
247
+ if has_buy_signals:
248
+ apds.append(
249
+ mpf.make_addplot(
250
+ buy_signal,
251
+ type="scatter",
252
+ marker="^",
253
+ markersize=120,
254
+ color="green",
255
+ )
256
+ )
257
+ if has_sell_signals:
258
+ apds.append(
259
+ mpf.make_addplot(
260
+ sell_signal,
261
+ type="scatter",
262
+ marker="v",
263
+ markersize=120,
264
+ color="red",
265
+ )
266
+ )
267
+
268
+ # If no trades found in the date range, warn the user
269
+ if not apds:
270
+ import warnings
271
+ warnings.warn(
272
+ "No trades found in the specified date range. "
273
+ "Plotting price data without trade markers."
274
+ )
275
+
276
+ # Generate default title if not provided
277
+ if title is None:
278
+ total_trades = len(self.orders)
279
+ title = f"Price chart with trade entries/exits ({total_trades} orders)"
280
+
281
+ mpf.plot(
282
+ data,
283
+ type="candle",
284
+ volume=volume,
285
+ addplot=apds,
286
+ figsize=figsize,
287
+ style=style,
288
+ title=title,
289
+ warn_too_much_data=len(data) + 5
290
+ )
291
+
164
292
  def __str__(self) -> str:
165
293
  """
166
294
  Generate a formatted string summary of backtest results.
@@ -7,13 +7,18 @@ optimization uses rolling windows to test parameter stability over time.
7
7
  """
8
8
 
9
9
  from dataclasses import dataclass, field
10
- from typing import Any, Callable, Protocol
10
+ from datetime import timedelta
11
+ from typing import Any, Callable, Protocol, Union
11
12
 
12
13
  import numpy as np
13
14
  import pandas as pd
14
15
  from tqdm import tqdm
15
16
 
16
17
 
18
+ # Type alias for period specification (int or timedelta)
19
+ PeriodSpec = Union[int, timedelta, pd.Timedelta]
20
+
21
+
17
22
  @dataclass
18
23
  class WalkForwardWindow:
19
24
  """
@@ -27,6 +32,10 @@ class WalkForwardWindow:
27
32
  test_end (int): Ending index for test data.
28
33
  train_periods (int): Number of periods in training.
29
34
  test_periods (int): Number of periods in testing.
35
+ train_periods_spec (PeriodSpec): Original specification for train periods
36
+ (can be int or timedelta).
37
+ test_periods_spec (PeriodSpec): Original specification for test periods
38
+ (can be int or timedelta).
30
39
  best_params (dict): Best parameters found during training.
31
40
  train_metrics (dict): Metrics computed on training data.
32
41
  test_metrics (dict): Metrics computed on test (out-of-sample) data.
@@ -40,9 +49,11 @@ class WalkForwardWindow:
40
49
  test_end: int
41
50
  train_periods: int
42
51
  test_periods: int
43
- best_params: dict
44
- train_metrics: dict
45
- test_metrics: dict
52
+ train_periods_spec: PeriodSpec | None = None
53
+ test_periods_spec: PeriodSpec | None = None
54
+ best_params: dict = field(default_factory=dict)
55
+ train_metrics: dict = field(default_factory=dict)
56
+ test_metrics: dict = field(default_factory=dict)
46
57
  train_report: Any = None
47
58
  test_report: Any = None
48
59
 
@@ -57,8 +68,12 @@ class WalkForwardResult:
57
68
 
58
69
  Attributes:
59
70
  n_windows (int): Total number of walk-forward windows.
60
- train_periods (int): Number of periods in each training window.
61
- test_periods (int): Number of periods in each test window.
71
+ train_periods (int): Number of periods in each training window (computed value).
72
+ test_periods (int): Number of periods in each test window (computed value).
73
+ train_periods_spec (PeriodSpec): Original specification for train periods
74
+ (can be int or timedelta).
75
+ test_periods_spec (PeriodSpec): Original specification for test periods
76
+ (can be int or timedelta).
62
77
  window_results (list[WalkForwardWindow]): Results for each window.
63
78
  aggregated_metrics (dict): Aggregated statistics across all windows.
64
79
  all_windows_results_df (pd.DataFrame): DataFrame with results from all windows.
@@ -66,6 +81,8 @@ class WalkForwardResult:
66
81
  n_windows: int
67
82
  train_periods: int
68
83
  test_periods: int
84
+ train_periods_spec: PeriodSpec | None = None
85
+ test_periods_spec: PeriodSpec | None = None
69
86
  window_results: list[WalkForwardWindow] = field(default_factory=list)
70
87
  aggregated_metrics: dict = field(default_factory=dict)
71
88
  all_windows_results_df: pd.DataFrame = field(default_factory=pd.DataFrame)
@@ -269,7 +286,6 @@ class OptimizerProtocol(Protocol):
269
286
  def __call__(self, *args: Any, **kwargs: Any) -> Any: ...
270
287
 
271
288
 
272
- @dataclass
273
289
  class WalkForwardAnalyzer:
274
290
  """
275
291
  Analyzer for walk-forward optimization.
@@ -286,10 +302,14 @@ class WalkForwardAnalyzer:
286
302
 
287
303
  Attributes:
288
304
  backtester (SimpleBacktester): The backtester instance to use.
289
- train_periods (int): Number of periods for each training window.
290
- test_periods (int): Number of periods for each test window.
291
- step_periods (int): Number of periods to step forward between windows.
292
- If None, uses test_periods (non-overlapping windows).
305
+ train_periods (PeriodSpec): Number of periods or timedelta for each
306
+ training window. Can be an int (number of periods) or a timedelta
307
+ (e.g., pd.Timedelta('252 days') or timedelta(days=365)).
308
+ test_periods (PeriodSpec): Number of periods or timedelta for each
309
+ test window. Can be an int or timedelta.
310
+ step_periods (PeriodSpec): Number of periods or timedelta to step
311
+ forward between windows. If None, uses test_periods
312
+ (non-overlapping windows).
293
313
  min_train_periods (int): Minimum required training periods.
294
314
  min_test_periods (int): Minimum required test periods.
295
315
  selection_criterion (str): Metric to use for selecting best parameters.
@@ -297,13 +317,22 @@ class WalkForwardAnalyzer:
297
317
  Example:
298
318
  >>> from quantex.backtester.walk_forward import WalkForwardAnalyzer
299
319
  >>> from quantex import SimpleBacktester
320
+ >>> import pandas as pd
321
+ >>>
322
+ >>> # Create analyzer with grid search optimizer (using periods)
323
+ >>> analyzer = WalkForwardAnalyzer(
324
+ ... backtester=bt,
325
+ ... train_periods=252, # 1 year training (252 trading days)
326
+ ... test_periods=63, # 3 months testing
327
+ ... step_periods=63 # Move forward 3 months each window
328
+ ... )
300
329
  >>>
301
- >>> # Create analyzer with grid search optimizer
330
+ >>> # Or using timedelta (data frequency is inferred from the data)
302
331
  >>> analyzer = WalkForwardAnalyzer(
303
332
  ... backtester=bt,
304
- ... train_periods=252, # 1 year training
305
- ... test_periods=63, # 3 months testing
306
- ... step_periods=63 # Move forward 3 months each window
333
+ ... train_periods=pd.Timedelta('252 days'), # 1 year training
334
+ ... test_periods=pd.Timedelta('90 days'), # 3 months testing
335
+ ... step_periods=pd.Timedelta('90 days') # Move forward 3 months
307
336
  ... )
308
337
  >>>
309
338
  >>> # Run with grid search
@@ -316,37 +345,28 @@ class WalkForwardAnalyzer:
316
345
  >>> print(result)
317
346
  >>> print(f"Average out-of-sample Sharpe: {result.aggregated_metrics['out_of_sample_sharpe_mean']:.2f}")
318
347
  """
319
- backtester: Any
320
- train_periods: int
321
- test_periods: int
322
- step_periods: int = 0 # Will be set to test_periods in __post_init__ if 0
323
- min_train_periods: int = 30
324
- min_test_periods: int = 10
325
- selection_criterion: str = "sharpe"
326
-
327
- def __post_init__(self):
328
- """Validate and set default values."""
329
- # Handle default step_periods
330
- if self.step_periods <= 0:
331
- self.step_periods = self.test_periods
332
-
333
- if self.train_periods < self.min_train_periods:
334
- raise ValueError(
335
- f"train_periods must be at least {self.min_train_periods}, "
336
- f"got {self.train_periods}"
337
- )
338
-
339
- if self.test_periods < self.min_test_periods:
340
- raise ValueError(
341
- f"test_periods must be at least {self.min_test_periods}, "
342
- f"got {self.test_periods}"
343
- )
344
-
345
- if self.step_periods <= 0:
346
- raise ValueError("step_periods must be positive")
347
-
348
- # Get data length from the strategy's data source
349
- positions = self.backtester.strategy.positions
348
+
349
+ def __init__(
350
+ self,
351
+ backtester: Any,
352
+ train_periods: PeriodSpec,
353
+ test_periods: PeriodSpec,
354
+ step_periods: PeriodSpec | None = None,
355
+ min_train_periods: int = 30,
356
+ min_test_periods: int = 10,
357
+ selection_criterion: str = "sharpe",
358
+ ):
359
+ """Initialize the WalkForwardAnalyzer."""
360
+ self.backtester = backtester
361
+ self._train_periods_spec = train_periods
362
+ self._test_periods_spec = test_periods
363
+ self._step_periods_spec = step_periods
364
+ self.min_train_periods = min_train_periods
365
+ self.min_test_periods = min_test_periods
366
+ self.selection_criterion = selection_criterion
367
+
368
+ # Get data source for frequency information
369
+ positions = backtester.strategy.positions
350
370
  if not positions:
351
371
  raise ValueError(
352
372
  "Strategy must have at least one data source registered. "
@@ -354,32 +374,139 @@ class WalkForwardAnalyzer:
354
374
  )
355
375
  source = positions[next(iter(positions))].source
356
376
  self.data_length = len(source.data)
377
+ self.data_index = source.data.index
378
+
379
+ # Convert timedelta specifications to period counts
380
+ self._train_periods_int = self._resolve_periods(train_periods)
381
+ self._test_periods_int = self._resolve_periods(test_periods)
382
+
383
+ if step_periods is None:
384
+ self._step_periods_int = self._test_periods_int
385
+ else:
386
+ self._step_periods_int = self._resolve_periods(step_periods)
387
+
388
+ # Validate period counts
389
+ if self._train_periods_int < self.min_train_periods:
390
+ raise ValueError(
391
+ f"train_periods ({train_periods}) resolves to "
392
+ f"{self._train_periods_int} periods, which is less than "
393
+ f"the minimum of {self.min_train_periods}"
394
+ )
395
+
396
+ if self._test_periods_int < self.min_test_periods:
397
+ raise ValueError(
398
+ f"test_periods ({test_periods}) resolves to "
399
+ f"{self._test_periods_int} periods, which is less than "
400
+ f"the minimum of {self.min_test_periods}"
401
+ )
402
+
403
+ if self._step_periods_int <= 0:
404
+ raise ValueError(
405
+ f"step_periods ({step_periods}) resolves to "
406
+ f"{self._step_periods_int}, which must be positive"
407
+ )
408
+
409
+ def _resolve_periods(self, value: PeriodSpec) -> int:
410
+ """
411
+ Convert a period specification to number of periods.
412
+
413
+ Args:
414
+ value: Either an int (number of periods) or a timedelta.
415
+
416
+ Returns:
417
+ int: Number of periods.
418
+ """
419
+ if isinstance(value, (timedelta, pd.Timedelta)):
420
+ td = pd.Timedelta(value)
421
+ # Calculate number of periods based on data frequency
422
+ if len(self.data_index) < 2:
423
+ raise ValueError(
424
+ "Cannot convert timedelta to periods: "
425
+ "data source has fewer than 2 data points"
426
+ )
427
+
428
+ # Determine the frequency of the data
429
+ # Compute from the first two points (most reliable method)
430
+ freq_delta = self.data_index[1] - self.data_index[0]
431
+ freq_delta = pd.Timedelta(freq_delta)
432
+
433
+ if freq_delta.total_seconds() <= 0:
434
+ raise ValueError(
435
+ "Cannot convert timedelta to periods: "
436
+ "data frequency could not be determined"
437
+ )
438
+
439
+ # Calculate number of periods covered by the timedelta
440
+ periods = int(round(td / freq_delta))
441
+
442
+ if periods <= 0:
443
+ raise ValueError(
444
+ f"timedelta {td} is less than one data period "
445
+ f"(frequency: {freq_delta})"
446
+ )
447
+
448
+ return periods
449
+
450
+ # It's already an int
451
+ return int(value)
452
+
453
+ @property
454
+ def train_periods_spec(self) -> PeriodSpec:
455
+ """Original train_periods specification (int or timedelta)."""
456
+ return self._train_periods_spec
457
+
458
+ @property
459
+ def test_periods_spec(self) -> PeriodSpec:
460
+ """Original test_periods specification (int or timedelta)."""
461
+ return self._test_periods_spec
462
+
463
+ @property
464
+ def step_periods_spec(self) -> PeriodSpec | None:
465
+ """Original step_periods specification (int or timedelta)."""
466
+ return self._step_periods_spec
467
+
468
+ @property
469
+ def train_periods(self) -> int:
470
+ """Resolved train_periods as number of periods (int)."""
471
+ return self._train_periods_int
472
+
473
+ @property
474
+ def test_periods(self) -> int:
475
+ """Resolved test_periods as number of periods (int)."""
476
+ return self._test_periods_int
477
+
478
+ @property
479
+ def step_periods(self) -> int:
480
+ """Resolved step_periods as number of periods (int)."""
481
+ return self._step_periods_int
357
482
 
358
483
  def _create_window_splits(self) -> list[tuple[int, int, int, int]]:
359
484
  """
360
485
  Create the train/test splits for all walk-forward windows.
361
-
486
+
362
487
  Returns:
363
488
  List of tuples: (train_start, train_end, test_start, test_end)
364
489
  """
365
490
  splits = []
366
491
  train_start = 0
367
- step = self.step_periods # Already validated to be non-None in __post_init__
368
-
492
+ train_periods = self._train_periods_int
493
+ test_periods = self._test_periods_int
494
+ step = self._step_periods_int
495
+
369
496
  while True:
370
- train_end = train_start + self.train_periods
497
+ train_end = train_start + train_periods
371
498
  test_start = train_end
372
- test_end = test_start + self.test_periods
373
-
499
+ test_end = test_start + test_periods
500
+
374
501
  # Check if we have enough data for test period
375
502
  if test_end > self.data_length:
376
503
  break
377
-
504
+
378
505
  splits.append((train_start, train_end, test_start, test_end))
379
-
506
+
380
507
  # Move forward
381
508
  train_start += step
382
-
509
+
383
510
  return splits
384
511
 
385
512
  def _slice_strategy_for_window(
@@ -525,8 +652,8 @@ class WalkForwardAnalyzer:
525
652
  if n_windows == 0:
526
653
  raise ValueError(
527
654
  f"Data length ({self.data_length}) is too short for the configured "
528
- f"train_periods ({self.train_periods}) and test_periods ({self.test_periods}). "
529
- f"Need at least {self.train_periods + self.test_periods} periods."
655
+ f"train_periods ({self._train_periods_int}) and test_periods ({self._test_periods_int}). "
656
+ f"Need at least {self._train_periods_int + self._test_periods_int} periods."
530
657
  )
531
658
 
532
659
  window_results: list[WalkForwardWindow] = []
@@ -608,6 +735,8 @@ class WalkForwardAnalyzer:
608
735
  test_end=test_end,
609
736
  train_periods=train_end - train_start,
610
737
  test_periods=test_end - test_start,
738
+ train_periods_spec=self._train_periods_spec,
739
+ test_periods_spec=self._test_periods_spec,
611
740
  best_params=best_params,
612
741
  train_metrics=train_metrics,
613
742
  test_metrics=test_metrics,
@@ -637,8 +766,10 @@ class WalkForwardAnalyzer:
637
766
 
638
767
  return WalkForwardResult(
639
768
  n_windows=n_windows,
640
- train_periods=self.train_periods,
641
- test_periods=self.test_periods,
769
+ train_periods=self._train_periods_int,
770
+ test_periods=self._test_periods_int,
771
+ train_periods_spec=self._train_periods_spec,
772
+ test_periods_spec=self._test_periods_spec,
642
773
  window_results=window_results,
643
774
  aggregated_metrics=aggregated,
644
775
  all_windows_results_df=results_df,
@@ -713,9 +844,9 @@ def walk_forward_analyze(
713
844
  backtester: Any,
714
845
  optimizer: OptimizerProtocol,
715
846
  params: dict[str, Any],
716
- train_periods: int,
717
- test_periods: int,
718
- step_periods: int | None = None,
847
+ train_periods: PeriodSpec,
848
+ test_periods: PeriodSpec,
849
+ step_periods: PeriodSpec | None = None,
719
850
  constraint: Callable[[dict], bool] | None = None,
720
851
  objective: str = "sharpe",
721
852
  risk_tolerance: dict[str, float] | None = None,
@@ -734,10 +865,12 @@ def walk_forward_analyze(
734
865
  backtester (SimpleBacktester): The backtester instance to use.
735
866
  optimizer (OptimizerProtocol): Optimizer function to use.
736
867
  params (dict[str, Any]): Parameter space for optimization.
737
- train_periods (int): Number of periods for each training window.
738
- test_periods (int): Number of periods for each test window.
739
- step_periods (int | None, optional): Periods to step between windows.
740
- If None, uses test_periods. Defaults to None.
868
+ train_periods (PeriodSpec): Number of periods or timedelta for each
869
+ training window. Can be an int or timedelta.
870
+ test_periods (PeriodSpec): Number of periods or timedelta for each
871
+ test window. Can be an int or timedelta.
872
+ step_periods (PeriodSpec | None, optional): Periods or timedelta to
873
+ step between windows. If None, uses test_periods. Defaults to None.
741
874
  constraint (Callable[[dict], bool] | None, optional): Constraint function.
742
875
  objective (str, optional): Metric to optimize. Defaults to "sharpe".
743
876
  risk_tolerance (dict[str, float] | None, optional): Risk tolerance.
@@ -752,6 +885,7 @@ def walk_forward_analyze(
752
885
  WalkForwardResult: Walk-forward analysis results.
753
886
 
754
887
  Example:
888
+ Using periods:
755
889
  >>> from quantex.backtester.walk_forward import walk_forward_analyze
756
890
  >>> result = walk_forward_analyze(
757
891
  ... backtester=bt,
@@ -761,16 +895,25 @@ def walk_forward_analyze(
761
895
  ... test_periods=63,
762
896
  ... objective='sharpe'
763
897
  ... )
898
+
899
+ Using timedelta:
900
+ >>> import pandas as pd
901
+ >>> result = walk_forward_analyze(
902
+ ... backtester=bt,
903
+ ... optimizer=lambda bt, params: bt.optimize(params),
904
+ ... params={'fast': [5, 10, 20], 'slow': [20, 50, 100]},
905
+ ... train_periods=pd.Timedelta('365 days'),
906
+ ... test_periods=pd.Timedelta('90 days'),
907
+ ... objective='sharpe'
908
+ ... )
909
+
764
910
  >>> print(f"Average OOS Sharpe: {result.aggregated_metrics['out_of_sample_sharpe_mean']:.2f}")
765
911
  """
766
- # Handle step_periods - use test_periods if None
767
- actual_step = test_periods if step_periods is None else step_periods
768
-
769
912
  analyzer = WalkForwardAnalyzer(
770
913
  backtester=backtester,
771
914
  train_periods=train_periods,
772
915
  test_periods=test_periods,
773
- step_periods=actual_step,
916
+ step_periods=step_periods,
774
917
  min_train_periods=min_train_periods,
775
918
  min_test_periods=min_test_periods,
776
919
  selection_criterion=objective,
File without changes
File without changes
File without changes
File without changes
File without changes