Coverage for book/marimo/notebooks/Experiment4.py: 100%
47 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-31 10:05 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-31 10:05 +0000
1# /// script
2# requires-python = ">=3.12"
3# dependencies = [
4# "marimo==0.23.15",
5# "numpy==2.4.6",
6# "plotly==6.9.0",
7# "polars==1.43.1",
8# "jquantstats==0.9.7",
9# "tinycta==0.14.0"
10# ]
11# ///
13"""Experiment 4: CTA strategy with optimization and risk scaling.
15This module demonstrates a more advanced trend-following strategy that
16incorporates portfolio optimization techniques and risk scaling to
17improve performance and risk-adjusted returns.
18"""
20import marimo
22__generated_with = "0.23.1"
23app = marimo.App()
25with app.setup:
26 import sys
27 from pathlib import Path
29 import marimo as mo
30 import numpy as np
31 import polars as pl
32 from jquantstats import Portfolio
33 from tinycta.osc import osc
34 from tinycta.util import vol_adj
36 sys.path.insert(0, str(Path(__file__).parent))
38 from preamble import date_col, load_prices
40 prices = load_prices(__file__)
41 prices_only = prices.drop(date_col)
42 assets = prices_only.columns
45@app.cell(hide_code=True)
46def _():
47 mo.md(r"""# CTA 4.0 - Optimization 1.0""")
48 return
51@app.function
52def f(price: "pl.Expr", fast: int = 32, slow: int = 96, vola: int = 32, clip: float = 4.2) -> "pl.Expr":
53 """Return the tanh oscillator of vol-adjusted cumulative price."""
54 return osc(vol_adj(price, vola=vola, clip=clip, min_samples=300).cum_sum(), fast=fast, slow=slow).tanh()
57@app.cell
58def _():
59 fast = mo.ui.slider(4, 192, step=4, value=32, label="Fast Moving Average")
60 slow = mo.ui.slider(4, 192, step=4, value=96, label="Slow Moving Average")
61 vola = mo.ui.slider(4, 192, step=4, value=32, label="Volatility")
62 winsor = mo.ui.slider(1.0, 6.0, step=0.1, value=4.2, label="Winsorizing")
64 mo.vstack([fast, slow, vola, winsor])
66 return fast, slow, vola, winsor
69@app.cell
70def _(fast, slow, vola, winsor):
71 mu_np = prices_only.select(
72 f(pl.all(), fast=fast.value, slow=slow.value, vola=vola.value, clip=winsor.value)
73 ).to_numpy()
74 volax_np = prices_only.select(
75 pl.all().fill_nan(None).pct_change().ewm_std(com=vola.value, min_samples=vola.value)
76 ).to_numpy()
77 euclid_norm = np.sqrt(np.nansum(mu_np**2, axis=1, keepdims=True))
78 euclid_norm[euclid_norm == 0] = np.nan
79 risk_scaled_np = mu_np / euclid_norm
81 pos_np = np.nan_to_num(5e5 * risk_scaled_np / volax_np, nan=0.0)
82 portfolio = Portfolio.from_cash_position(
83 prices=prices,
84 cash_position=pl.concat(
85 [prices.select(date_col), pl.from_numpy(pos_np, schema=dict.fromkeys(assets, pl.Float64))],
86 how="horizontal_extend",
87 ),
88 aum=1e8,
89 )
90 return (portfolio,)
93@app.cell
94def _(portfolio):
95 print(portfolio.stats.sharpe())
98@app.cell
99def _(portfolio):
100 fig = portfolio.plots.snapshot()
101 fig
102 return
105if __name__ == "__main__":
106 app.run()