| import os |
| os.environ.setdefault("MPLBACKEND", "Agg") |
| import pandas as pd |
| import yfinance as yf |
| from IPython.display import Markdown, display |
| import backtrader as bt |
| import tempfile, os |
| import base64 |
| from io import BytesIO |
| import numpy as np |
| from io import BytesIO |
| import matplotlib |
| matplotlib.use('Agg') |
| import matplotlib.pyplot as plt |
| from datetime import datetime, timedelta |
|
|
|
|
| |
| CACHE_DIR = "data_cache" |
| os.makedirs(CACHE_DIR, exist_ok=True) |
|
|
|
|
| class TransactionRecorder(bt.Analyzer): |
| """ Records all buy/sell transactions with details.""" |
| def __init__(self): |
| self.records = [] |
|
|
| def notify_order(self, order): |
| if order.status != order.Completed: |
| return |
|
|
| self.records.append({ |
| 'datetime': self.strategy.data.datetime.datetime(), |
| 'type': 'BUY' if order.isbuy() else 'SELL', |
| 'price': order.executed.price, |
| 'size': order.executed.size, |
| 'value': order.executed.value, |
| 'commission': order.executed.comm, |
| }) |
|
|
| def get_analysis(self): |
| return pd.DataFrame(self.records) |
|
|
| class TradeRecorder(bt.Analyzer): |
| """ Records detailed trade information including entry/exit prices and PnL.""" |
| def __init__(self): |
| self.records = [] |
| self.trade_id = 0 |
|
|
| def notify_trade(self, trade): |
| |
| if trade.isopen: |
| self.current_entry_price = trade.price |
| self.current_size = trade.size |
| return |
|
|
| |
| if trade.isclosed: |
| self.trade_id += 1 |
|
|
| |
| size = self.current_size |
|
|
| |
| |
| if size is not None and size != 0: |
| exit_price = self.current_entry_price + (trade.pnl / size) |
| else: |
| exit_price = None |
|
|
| self.records.append({ |
| 'trade_id': self.trade_id, |
| 'size': size, |
| 'entry_price': self.current_entry_price, |
| 'exit_price': exit_price, |
| 'bruto_profitloss': trade.pnl, |
| 'neto_profitloss': trade.pnlcomm, |
| 'date_open': self._format_dt(trade.dtopen), |
| 'date_close': self._format_dt(trade.dtclose), |
| }) |
|
|
| |
| self.current_entry_price = None |
| self.current_size = None |
|
|
| def get_analysis(self): |
| return pd.DataFrame(self.records) |
|
|
| def _format_dt(self, val): |
| """Convert a backtrader/matplotlib numeric datetime to a readable string.""" |
| if val is None: |
| return None |
| try: |
| ord_day = int(val) |
| frac = val - ord_day |
| dt = datetime.fromordinal(ord_day) + timedelta(days=frac) |
| return dt.strftime('%Y-%m-%d %H:%M:%S') |
| except Exception as e: |
| return "Error formatting dt: " + str(e) |
| |
|
|
| def _get_cache_filename(tckr_symbl, interval, start_date, end_date, adjust_prices): |
| """Generate a cache filename based on ticker, interval, date range, and adjust_prices. |
| Format: ticker_interval_start_to_end_adjust.csv""" |
| start_str = start_date.replace('-', '_') |
| end_str = end_date.replace('-', '_') |
| adjust_str = "adj" if adjust_prices else "raw" |
| return os.path.join(CACHE_DIR, f"{tckr_symbl}_{interval}_{start_str}_to_{end_str}_{adjust_str}.csv") |
|
|
| def _is_valid_cache(df): |
| """Simple check if cached dataframe is valid and usable.""" |
| if df.empty or len(df) < 2: |
| return False |
| if not isinstance(df.index, pd.DatetimeIndex): |
| return False |
| |
| required_cols = ['Open', 'High', 'Low', 'Close', 'Volume'] |
| if not all(col in df.columns for col in required_cols): |
| return False |
| return True |
|
|
| def _load_cached_data(cache_file): |
| """Load cached data from CSV if it exists.""" |
| if os.path.exists(cache_file): |
| try: |
| |
| df = pd.read_csv(cache_file, index_col=0, parse_dates=True, header=0) |
| |
| |
| if isinstance(df.columns, pd.MultiIndex): |
| df.columns = df.columns.get_level_values(0) |
| |
| |
| if not isinstance(df.index, pd.DatetimeIndex): |
| try: |
| df.index = pd.to_datetime(df.index) |
| except Exception as e: |
| print(f"Could not parse dates in cache: {e}, will re-download") |
| return None |
| |
| |
| expected_cols = ['Open', 'High', 'Low', 'Close', 'Volume'] |
| actual_cols = list(df.columns) |
| |
| if not any(col in actual_cols for col in expected_cols): |
| print(f"Unexpected columns in cache: {actual_cols}, will re-download") |
| return None |
| |
| print(f"Loaded cached data: {len(df)} rows, columns: {list(df.columns)} from {cache_file}") |
| return df |
| except Exception as e: |
| print(f"Error loading cache: {e}, will re-download") |
| import traceback |
| traceback.print_exc() |
| return None |
| return None |
|
|
| def _save_cached_data(df, cache_file): |
| """Save dataframe to CSV cache.""" |
| try: |
| df.to_csv(cache_file) |
| print(f"Cached data saved: {cache_file}") |
| except Exception as e: |
| print(f"Error saving cache: {e}") |
|
|
| def _fig_to_numpy(fig, dpi=150): |
| buf = BytesIO() |
| fig.savefig(buf, format="png", dpi=dpi, bbox_inches="tight") |
| buf.seek(0) |
| img = plt.imread(buf, format="png") |
| buf.close() |
| return img |
|
|
| |
| def plot_bt(figs, symbol, market_name, save_img=True): |
| """ |
| Plot backtrader results with market name in title |
| and readable date ticks on all x-axes. |
| """ |
| images = [] |
|
|
| plt.ioff() |
| for i, fig_list in enumerate(figs): |
| for j, fig in enumerate(fig_list): |
| |
| if fig.axes: |
| main_ax = fig.axes[0] |
| current_title = main_ax.get_title() |
| new_title = f"{market_name} ({symbol.upper()}) - {current_title or 'Price Chart'}" |
| main_ax.set_title(new_title, fontsize=8, fontweight="bold") |
| fig.suptitle( |
| f"Trading Strategy Analysis: {market_name}", |
| fontsize=10, |
| fontweight="bold", |
| y=0.98, |
| ) |
|
|
| fig.set_size_inches(12, 6) |
| fig.autofmt_xdate() |
| fig.tight_layout(rect=[0, 0, 1, 0.95]) |
| |
| img_data = _fig_to_numpy(fig) |
| images.append(img_data) |
|
|
| if save_img: |
| filename = f"plot_{symbol}_{i}_{j}.png" |
| fig.savefig(filename, dpi=300, bbox_inches="tight") |
| print(f"Chart saved as: {filename}") |
| |
| plt.close(fig) |
| return images |
|
|
| def _download_data(tckr_symbl, interval, date, adjust_prices, auto_period=True, period='60d'): |
| |
| try: |
| print("Interval: ", interval) |
| start_dt = datetime.strptime(date["start"], "%Y-%m-%d") |
| end_dt = datetime.strptime(date["end"], "%Y-%m-%d") |
| |
| |
| cache_file = _get_cache_filename(tckr_symbl, interval, date["start"], date["end"], adjust_prices) |
| |
| |
| cached_df = _load_cached_data(cache_file) |
| if cached_df is not None and _is_valid_cache(cached_df): |
| |
| df = cached_df |
| print(f"Using cached data: {len(df)} rows (no download needed)") |
| else: |
| |
| print("No valid cache, downloading...") |
| |
| if interval in ['1m', '2m', '5m', '15m', '30m', '60m', '1h'] and auto_period: |
| |
| if interval in ['1m']: |
| max_days = 7 |
| elif interval in ['2m', '5m', '15m', '30m']: |
| max_days = 60 |
| else: |
| max_days = 730 |
| desired_days = max(1, (end_dt - start_dt).days or 1) |
| clamped_days = min(desired_days, max_days) |
| period = f"{clamped_days}d" |
| df = yf.download(tckr_symbl, period=period, interval=interval, auto_adjust=adjust_prices) |
| print(f"Downloaded {interval} data for {period}") |
| else: |
| |
| df = yf.download(tckr_symbl, start=date["start"], end=date["end"], interval=interval, auto_adjust=adjust_prices) |
| print(f"Downloaded data from {date['start']} to {date['end']} with {interval} interval") |
| |
| |
| if isinstance(df.columns, pd.MultiIndex): |
| df.columns = df.columns.get_level_values(0) |
| |
| |
| _save_cached_data(df, cache_file) |
| |
| |
| if isinstance(df.columns, pd.MultiIndex): |
| df.columns = df.columns.get_level_values(0) |
| |
| |
| if df.empty: |
| raise ValueError("No data available for the specified parameters!") |
| |
| if df.index.tz is not None: |
| df.index = df.index.tz_localize(None) |
| |
| print(f"Data points: {len(df)}") |
| return df |
| |
| except Exception as e: |
| raise ValueError(f"Error downloading data: {e}") |
| |
|
|
| |
|
|
| def run_bt(cerebro, |
| date={'start':'1990-01-01', 'end':'2024-12-31'}, |
| tckr_symbl="SPY", |
| save_img=False, |
| interval='1d', |
| auto_period =True, |
| period='60d', |
| market_name = "Complete Market Name here", |
| initial_capital=10000.0, |
| commission=0.001, |
| slippage_percent=0.01, |
| adjust_prices=True): |
| """ |
| Run backtrader strategy with enhanced plotting |
| |
| Args: |
| strategy: Backtrader strategy class already init |
| date: Dictionary with start and end dates |
| tckr_symbl: Stock ticker symbol |
| save_img: Whether to save plot images |
| interval: Data interval ('1d', '1h', '30m', etc.) |
| period: Period for intraday data (e.g., '60d') |
| auto_period: Whether to auto adjust period based on interval |
| initial_capital: Starting capital for the broker |
| commission: Commission per share |
| slippage_percent: Percent (e.g., 0.01 for 0.01%) applied as slippage |
| adjust_prices: Whether to pull adjusted (dividend/split) prices |
| """ |
|
|
| print(f"Running strategy on: {market_name} ({tckr_symbl.upper()})") |
| print("-" * 50) |
| print(f"Initial capital: {initial_capital}, Commission: {commission}, Slippage%: {slippage_percent}, Adjusted prices: {adjust_prices}") |
| |
| df = _download_data(tckr_symbl, interval, date, adjust_prices, auto_period, period) |
| |
| |
| data = bt.feeds.PandasData(dataname=df) |
| cerebro.adddata(data) |
| |
| |
| initial_cash = float(initial_capital) |
| cerebro.broker.setcash(initial_cash) |
| cerebro.broker.setcommission(commission=float(commission)) |
| slippage_decimal = float(slippage_percent) / 100.0 |
| cerebro.broker.set_slippage_perc(slippage_decimal) |
| |
| |
| cerebro.addanalyzer(bt.analyzers.Returns, _name='returns') |
| cerebro.addanalyzer(bt.analyzers.SharpeRatio, _name='sharpe') |
| cerebro.addanalyzer(bt.analyzers.DrawDown, _name='drawdown') |
| cerebro.addanalyzer(TransactionRecorder, _name='transactions') |
| cerebro.addanalyzer(TradeRecorder, _name='trades') |
| |
| |
| print(f'Starting Portfolio Value: ${cerebro.broker.getvalue():,.2f}') |
| |
| |
| results = cerebro.run() |
| |
| |
| final_value = cerebro.broker.getvalue() |
| total_return = (final_value - initial_cash) / initial_cash * 100 |
| |
| print(f'Final Portfolio Value: ${final_value:,.2f}') |
| print(f'Total Return: {total_return:.2f}%') |
| |
| |
| strat = results[0] |
| |
| try: |
| sharpe = strat.analyzers.sharpe.get_analysis().get('sharperatio', 'N/A') |
| if sharpe != 'N/A': |
| print(f'Sharpe Ratio: {sharpe:.3f}') |
| else: |
| print('Sharpe Ratio: N/A') |
| except: |
| print('Sharpe Ratio: N/A') |
| |
| try: |
| max_dd = strat.analyzers.drawdown.get_analysis()['max']['drawdown'] |
| print(f'Max Drawdown: {max_dd:.2f}%') |
| except: |
| print('Max Drawdown: N/A') |
| |
| print("-" * 50) |
| |
| df_transactions = strat.analyzers.transactions.get_analysis() |
| df_trades = strat.analyzers.trades.get_analysis() |
| if not df_transactions.empty: |
| print("Transactions logs generated") |
| else: |
| print("No transactions recorded.") |
| if not df_trades.empty: |
| print("Trades logs generated") |
| else: |
| print("No trades recorded.") |
| |
| |
| figs = cerebro.plot(style='candlestick', |
| barstyle='candlestick', |
| rowsmajor=True, |
| dpi=120, |
| |
| subplot=True, |
| plotabove=False, |
| |
| downsample=False, |
| plotstyle='multiple', |
| iplot=False, |
| show=False |
| ) |
| print ("How many charts where created: " , len(figs)) |
| fig = plot_bt(figs, market_name, tckr_symbl, save_img) |
| |
| return final_value, total_return, fig, df_trades, df_transactions |