Compare commits

..

10 Commits

Author SHA1 Message Date
freedakgmail 6041e3ff27 Add stock prediction experiments: AAPL forecast + 4-version backtest comparison
Python package build / build (push) Has been cancelled
- stock_forecast.py: Basic AAPL price forecast
- stock_backtest.py: V1 original price prediction backtest
- stock_backtest_optimized.py: V2 log-return + ensemble optimization
- stock_backtest_xreg.py: V3 XReg covariates (volume/RSI/SPY)
- stock_backtest_trick.py: V4 SPY trend guidance + volatility adjustment
- Visualization charts for all versions
2026-07-01 21:01:59 +08:00
Rajat Sen 8a22ca28a0 Update README.md 2026-06-08 11:27:29 -07:00
Rajat Sen 6ed1d8a7a2 Update global-temperature example to use TimesFM 2.5 API 2026-06-08 17:31:07 +00:00
Rajat Sen e56854bc9e Bump version to 2.0.1 2026-06-08 17:17:36 +00:00
Rajat Sen 2f1625c208 Fix model loading issues and forecast_naive slicing bug in TimesFM 2.5
- Allow model wrapper constructors (__init__) to accept and ignore extra keyword arguments (e.g. proxies) passed by huggingface_hub during from_pretrained.
- Implement load_checkpoint for TimesFM_2p5_200M_torch and TimesFM_2p5_200M_flax to restore weights from local paths.
- Fix slicing bug in PyTorch's forecast_naive to correctly slice the time/horizon dimension ([:, :horizon, :]) instead of quantiles.
- Add unit tests in tests/test_model_loading.py covering local checkpoint loading, hub compatibility, and prediction shape correctness.
2026-06-07 17:28:32 +00:00
Yichen Zhou b3d0dec6ec Update README.md 2026-06-05 12:42:18 -07:00
Yichen Zhou ace12a8a94 Update README.md
timesfm=2.0.0
2026-06-05 12:41:22 -07:00
Yichen Zhou d720daa678 Update README.md 2026-04-14 22:59:09 -07:00
Yichen Zhou eacf761c32 Merge pull request #398 from darkpowerxo/feat/peft-finetuning-pipeline-2.5
feat: PEFT fine-tuning pipeline (LoRA/DoRA, multi-GPU) for TimesFM 2.5
2026-04-14 22:49:22 -07:00
darkpowerxo 6ae67d41d8 revert: drop PR #393 (xreg batch behavior) and PR #390 (SKILL.md link) per maintainer feedback 2026-04-09 23:00:36 -04:00
22 changed files with 1177 additions and 265 deletions
+1
View File
@@ -8,3 +8,4 @@ datasets/
results/
uv.lock
development_setup.md
debug.log
+25 -6
View File
@@ -9,8 +9,10 @@ model developed by Google Research for time-series forecasting.
* All checkpoints:
[TimesFM Hugging Face Collection](https://huggingface.co/collections/google/timesfm-release-66e4be5fdb56e960c1e482a6).
* [Google Research blog](https://research.google/blog/a-decoder-only-foundation-model-for-time-series-forecasting/).
* [TimesFM in BigQuery](https://cloud.google.com/bigquery/docs/timesfm-model):
an official Google product.
* TimesFM in Google 1P Products:
* [BigQuery ML](https://cloud.google.com/bigquery/docs/timesfm-model): Enterprise level SQL queries for scalability and reliability.
* [Google Sheets](https://workspaceupdates.googleblog.com/2026/02/forecast-data-in-connected-sheets-BigQueryML-TimesFM.html): For your daily spreadsheet.
* [Vertex Model Garden](https://pantheon.corp.google.com/vertex-ai/publishers/google/model-garden/timesfm): Dockerized endpoint for agentic calling.
This open version is not an officially supported Google product.
@@ -21,17 +23,21 @@ This open version is not an officially supported Google product.
- 1.0 and 2.0: relevant code archived in the sub directory `v1`. You can `pip
install timesfm==1.3.0` to install an older version of this package to load
them.
## Update - June 5, 2026
Updated PyPI to `timesfm=2.0.0`. See [Install](https://github.com/google-research/timesfm#from-pypi).
## Update - Apr. 9, 2026
Added fine-tuning example using HuggingFace Transformers + PEFT (LoRA) — see
[`timesfm-forecasting/examples/finetuning/`](timesfm-forecasting/examples/finetuning/).
Also added unit tests (`tests/`), fixed per-input ridge regression in XReg to
prevent data leakage, and incorporated several community fixes.
Also added unit tests (`tests/`) and incorporated several community fixes.
Shoutout to [@kashif](https://github.com/kashif) and [@darkpowerxo](https://github.com/darkpowerxo).
## Update - Mar. 19, 2026
Huge shoutout to [@borealBytes](https://github.com/borealBytes) for adding the support for [AGENTS](https://github.com/google-research/timesfm/blob/master/AGENTS.md)! TimesFM [SKILL.md](https://github.com/google-research/timesfm/blob/master/timesfm-forecasting/SKILL.md) is out.
Huge shoutout to [@borealBytes](https://github.com/borealBytes) for adding the support for [AGENTS](https://github.com/google-research/timesfm/blob/master/AGENTS.md)! TimesFM [SKILL.md](https://github.com/google-research/timesfm/tree/master/timesfm-forecasting) is out.
## Update - Oct. 29, 2025
@@ -61,6 +67,19 @@ Since the Sept. 2025 launch, the following improvements have been completed:
### Install
#### From `PyPI`
```shell
# Install the package with torch
pip install timesfm[torch]
# Or with Flax
pip install timesfm[flax]
# And when XReg is needed
pip install timesfm[xreg]
```
#### Local Install
1. Clone the repository:
```shell
git clone https://github.com/google-research/timesfm.git
@@ -79,7 +98,7 @@ Since the Sept. 2025 launch, the following improvements have been completed:
uv pip install -e .[torch]
# Or with flax
uv pip install -e .[flax]
# Or XReg is needed
# And when XReg is needed
uv pip install -e .[xreg]
```
BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 116 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 143 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 142 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 147 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 108 KiB

+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "timesfm"
version = "2.0.0"
version = "2.0.1"
description = "A time series foundation model."
authors = [
{name = "Rajat Sen", email = "senrajat@google.com"},
+16 -4
View File
@@ -447,6 +447,21 @@ class TimesFM_2p5_200M_flax(timesfm_2p5_base.TimesFM_2p5):
model: nnx.Module = TimesFM_2p5_200M_flax_module()
def __init__(self, **kwargs):
self.model = TimesFM_2p5_200M_flax_module()
def load_checkpoint(self, path: str):
"""Loads a TimesFM model from a checkpoint."""
if os.path.isdir(path):
model_file_path = path
else:
model_file_path = os.path.dirname(path)
checkpointer = ocp.StandardCheckpointer()
graph, state = nnx.split(self.model)
state = checkpointer.restore(model_file_path, state)
self.model = nnx.merge(graph, state)
@classmethod
def from_pretrained(
cls,
@@ -485,10 +500,7 @@ class TimesFM_2p5_200M_flax(timesfm_2p5_base.TimesFM_2p5):
)
logging.info("Loading checkpoint from: %s", model_file_path)
checkpointer = ocp.StandardCheckpointer()
graph, state = nnx.split(instance.model)
state = checkpointer.restore(model_file_path, state)
instance.model = nnx.merge(graph, state)
instance.load_checkpoint(model_file_path)
return instance
def compile(
+16 -2
View File
@@ -257,7 +257,7 @@ class TimesFM_2p5_200M_torch_module(nn.Module):
to_concat = [t_pf[:, -1, ...]]
if t_ar is not None:
to_concat.append(t_ar.reshape(1, -1, self.q))
torch_forecast = torch.cat(to_concat, dim=1)[..., :horizon]
torch_forecast = torch.cat(to_concat, dim=1)[:, :horizon, :]
torch_forecast = torch_forecast.squeeze(0)
outputs.append(torch_forecast.detach().cpu().numpy())
return outputs
@@ -283,12 +283,26 @@ class TimesFM_2p5_200M_torch(
self,
torch_compile: bool = True,
config: Optional[dict] = None,
**kwargs,
):
self.model = TimesFM_2p5_200M_torch_module()
self.torch_compile = torch_compile
if config is not None:
self._hub_mixin_config = config
def load_checkpoint(self, path: str, **kwargs):
"""Loads a TimesFM model from a checkpoint directory or file."""
if os.path.isdir(path):
model_file_path = os.path.join(path, self.WEIGHTS_FILENAME)
if not os.path.exists(model_file_path):
raise FileNotFoundError(
f"{self.WEIGHTS_FILENAME} not found in directory {path}"
)
else:
model_file_path = path
self.model.load_checkpoint(model_file_path, **kwargs)
@classmethod
def _from_pretrained(
cls,
@@ -333,7 +347,7 @@ class TimesFM_2p5_200M_torch(
logging.info("Loading checkpoint from: %s", model_file_path)
# Load the weights into the model.
instance.model.load_checkpoint(
instance.load_checkpoint(
model_file_path, torch_compile=instance.torch_compile
)
return instance
+51 -59
View File
@@ -370,20 +370,11 @@ class BatchedInContextXRegBase:
x_train = np.concatenate(x_train, axis=1)
x_test = np.concatenate(x_test, axis=1)
# Normalize per-input for robustness (batch-wide normalization
# would make each input's result depend on batch composition).
train_splits = np.cumsum(self.train_lens)[:-1]
test_splits = np.cumsum(self.test_lens)[:-1]
train_parts = np.split(x_train, train_splits, axis=0)
test_parts = np.split(x_test, test_splits, axis=0)
norm_train, norm_test = [], []
for tr, te in zip(train_parts, test_parts):
m = np.mean(tr, axis=0, keepdims=True)
s = np.where((w := np.std(tr, axis=0, keepdims=True)) > _TOL, w, 1.0)
norm_train.append((tr - m) / s)
norm_test.append((te - m) / s)
x_train = [np.concatenate(norm_train, axis=0)]
x_test = [np.concatenate(norm_test, axis=0)]
# Normalize for robustness.
x_mean = np.mean(x_train, axis=0, keepdims=True)
x_std = np.where((w := np.std(x_train, axis=0, keepdims=True)) > _TOL, w, 1.0)
x_train = [(x_train - x_mean) / x_std]
x_test = [(x_test - x_mean) / x_std]
# Categorical features. Encode one by one.
one_hot_encoder = preprocessing.OneHotEncoder(
@@ -472,57 +463,58 @@ class BatchedInContextXRegLinear(BatchedInContextXRegBase):
assert_covariate_shapes=assert_covariate_shapes,
)
x_train = x_train_raw.copy()
if max_rows_per_col:
nrows, ncols = x_train.shape
if nrows > (w := ncols * max_rows_per_col):
subsample = jax.random.choice(
jax.random.PRNGKey(max_rows_per_col_sample_seed),
nrows,
(w,),
replace=False,
)
x_train = x_train[subsample]
flat_targets = flat_targets[subsample]
device = jax.devices("cpu")[0] if force_on_cpu else None
# Runs jitted version of the solvers which are quicker at the cost of
# running jitting during the first time calling. Re-jitting happens whenever
# new (padded) shapes are encountered.
# Ocassionally it helps with the speed and the accuracy if we force single
# thread execution on cpu for accelerator machines:
# 1. Avoid moving data to accelarator memory.
# 2. Avoid precision loss if any.
with jax.default_device(device):
x_train_raw = _to_padded_jax_array(x_train_raw)
x_train = _to_padded_jax_array(x_train)
flat_targets = _to_padded_jax_array(flat_targets)
x_test = _to_padded_jax_array(x_test)
beta_hat = (
jnp.linalg.pinv(
x_train.T @ x_train + ridge * jnp.eye(x_train.shape[1]),
hermitian=True,
)
@ x_train.T
@ flat_targets
)
y_hat = x_test @ beta_hat
y_hat_context = x_train_raw @ beta_hat if debug_info else None
outputs = []
outputs_context = []
train_idx, test_idx = 0, 0
with jax.default_device(device):
for trl, tel in zip(self.train_lens, self.test_lens):
x_tr = x_train_raw[train_idx : train_idx + trl]
x_te = x_test[test_idx : test_idx + tel]
y_tr = flat_targets[train_idx : train_idx + trl]
x_tr_fit = x_tr.copy()
if max_rows_per_col:
nrows, ncols = x_tr_fit.shape
if nrows > (w := ncols * max_rows_per_col):
subsample = jax.random.choice(
jax.random.PRNGKey(max_rows_per_col_sample_seed),
nrows,
(w,),
replace=False,
)
x_tr_fit = x_tr_fit[subsample]
y_tr = y_tr[subsample]
x_tr_raw_j = _to_padded_jax_array(x_tr)
x_tr_j = _to_padded_jax_array(x_tr_fit)
y_tr_j = _to_padded_jax_array(y_tr)
x_te_j = _to_padded_jax_array(x_te)
beta_hat = (
jnp.linalg.pinv(
x_tr_j.T @ x_tr_j + ridge * jnp.eye(x_tr_j.shape[1]),
hermitian=True,
)
@ x_tr_j.T
@ y_tr_j
# Reconstruct the ragged 2-dim batched forecasts from flattened linear fits.
train_index, test_index = 0, 0
for train_index_delta, test_index_delta in zip(self.train_lens, self.test_lens):
outputs.append(np.array(y_hat[test_index : (test_index + test_index_delta)]))
if debug_info:
outputs_context.append(
np.array(y_hat_context[train_index : (train_index + train_index_delta)])
)
outputs.append(np.array(x_te_j @ beta_hat)[:tel])
if debug_info:
outputs_context.append(np.array(x_tr_raw_j @ beta_hat)[:trl])
train_idx += trl
test_idx += tel
train_index += train_index_delta
test_index += test_index_delta
if debug_info:
return (
outputs,
outputs_context,
_to_padded_jax_array(flat_targets),
_to_padded_jax_array(x_train_raw),
_to_padded_jax_array(x_test),
)
return outputs, outputs_context, flat_targets, x_train, x_test
else:
return outputs
+115
View File
@@ -0,0 +1,115 @@
import yfinance as yf
import numpy as np
import torch
import timesfm
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from datetime import datetime, timedelta
# 1. 获取 AAPL 数据(到今天为止)
print("下载 AAPL 股票数据...", flush=True)
end = datetime.now()
start = end - timedelta(days=365)
df = yf.download("AAPL", start=start.strftime("%Y-%m-%d"), end=end.strftime("%Y-%m-%d"), progress=False)
close = df["Close"].values.flatten().astype(np.float32)
dates = df.index
# 2. 分割:6/15 之前为训练集,6/16 之后为实际值
split_date = "2026-06-15"
split_idx = None
for i, d in enumerate(dates):
if str(d.date()) <= split_date:
split_idx = i
# split_idx 是 <= 6/15 的最后一个索引
train_data = close[: split_idx + 1]
train_dates = dates[: split_idx + 1]
actual_data = close[split_idx + 1 :]
actual_dates = dates[split_idx + 1 :]
HORIZON = len(actual_data)
print(f"训练数据: {len(train_data)} 天(截至 {train_dates[-1].date()}", flush=True)
print(f"实际数据: {HORIZON} 天({actual_dates[0].date()} ~ {actual_dates[-1].date()}", flush=True)
print(f"分割日收盘价: {train_data[-1]:.2f}", flush=True)
# 3. 加载模型
print()
print("加载 TimesFM 模型...", flush=True)
torch.set_float32_matmul_precision("high")
model = timesfm.TimesFM_2p5_200M_torch.from_pretrained(
"google/timesfm-2.5-200m-pytorch", torch_compile=False
)
print("模型加载完成", flush=True)
# 4. 编译并预测
model.compile(
timesfm.ForecastConfig(
max_context=512,
max_horizon=128,
normalize_inputs=True,
use_continuous_quantile_head=True,
force_flip_invariance=False,
infer_is_positive=True,
fix_quantile_crossing=True,
)
)
print(f"预测未来 {HORIZON} 个交易日...", flush=True)
pf, qf = model.forecast(horizon=HORIZON, inputs=[train_data])
print("预测完成!", flush=True)
# 5. 计算误差指标
actual = actual_data
pred = pf[0][:HORIZON]
mae = np.mean(np.abs(actual - pred))
rmse = np.sqrt(np.mean((actual - pred) ** 2))
mape = np.mean(np.abs((actual - pred) / actual)) * 100
# 方向准确率
actual_dir = np.diff(actual)
pred_dir = np.diff(pred)
dir_acc = np.mean(actual_dir * pred_dir > 0) * 100
# 置信区间覆盖率
in_80 = np.mean((actual >= qf[0, :HORIZON, 1]) & (actual <= qf[0, :HORIZON, 9])) * 100
in_40 = np.mean((actual >= qf[0, :HORIZON, 3]) & (actual <= qf[0, :HORIZON, 7])) * 100
print()
print("=== 回测结果:6/16 ~ 6/30 预测 vs 实际 ===", flush=True)
print()
print(" 日期 实际价 预测价 误差 误差%", flush=True)
for i in range(HORIZON):
err = pred[i] - actual[i]
err_pct = err / actual[i] * 100
print(
f" {actual_dates[i].date()} {actual[i]:7.2f} {pred[i]:7.2f} {err:+7.2f} {err_pct:+6.2f}%",
flush=True,
)
print()
print("=== 误差指标 ===", flush=True)
print(f" MAE (平均绝对误差): ${mae:.2f}", flush=True)
print(f" RMSE (均方根误差): ${rmse:.2f}", flush=True)
print(f" MAPE (平均绝对百分比误差): {mape:.2f}%", flush=True)
print(f" 方向准确率: {dir_acc:.1f}%", flush=True)
print(f" 80%置信区间覆盖率: {in_80:.1f}%", flush=True)
print(f" 40%置信区间覆盖率: {in_40:.1f}%", flush=True)
# 6. 画图
fig, ax = plt.subplots(figsize=(14, 6))
show_n = min(60, len(train_data))
ax.plot(range(show_n), train_data[-show_n:], label="Historical Close", color="steelblue", linewidth=1.5)
x_actual = range(show_n, show_n + HORIZON)
ax.plot(x_actual, actual, label="Actual", color="forestgreen", linewidth=2, marker="o", markersize=3)
ax.plot(x_actual, pred, label="Forecast", color="tomato", linewidth=2, linestyle="--")
ax.fill_between(x_actual, qf[0, :HORIZON, 1], qf[0, :HORIZON, 9], alpha=0.15, color="tomato", label="80% CI")
ax.fill_between(x_actual, qf[0, :HORIZON, 3], qf[0, :HORIZON, 7], alpha=0.3, color="tomato", label="40% CI")
ax.axvline(x=show_n - 1, color="gray", linestyle="--", alpha=0.5, label="Forecast Start")
ax.set_title(f"AAPL Backtest: Forecast vs Actual (MAPE={mape:.2f}%, Dir.Acc={dir_acc:.0f}%)", fontsize=14)
ax.set_xlabel("Trading Days")
ax.set_ylabel("Price (USD)")
ax.legend(loc="upper left")
plt.tight_layout()
plt.savefig("aapl_backtest.png", dpi=150)
print()
print("Chart saved: aapl_backtest.png", flush=True)
print("Done!", flush=True)
+184
View File
@@ -0,0 +1,184 @@
import yfinance as yf
import numpy as np
import torch
import timesfm
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from datetime import datetime, timedelta
# ---------- 技术指标计算 ----------
def calc_rsi(prices, period=14):
deltas = np.diff(prices)
gains = np.where(deltas > 0, deltas, 0.0)
losses = np.where(deltas < 0, -deltas, 0.0)
avg_gain = np.convolve(gains, np.ones(period) / period, mode="valid")
avg_loss = np.convolve(losses, np.ones(period) / period, mode="valid")
avg_loss = np.where(avg_loss == 0, 1e-10, avg_loss)
rs = avg_gain / avg_loss
rsi = 100.0 - (100.0 / (1.0 + rs))
# pad front to match length
return np.concatenate([np.full(period, 50.0), rsi])
def calc_sma(prices, period):
sma = np.convolve(prices, np.ones(period) / period, mode="valid")
return np.concatenate([np.full(period - 1, sma[0] if len(sma) > 0 else 0.0), sma])
# ---------- 1. 获取数据 ----------
print("下载 AAPL 股票数据...", flush=True)
end = datetime.now()
start = end - timedelta(days=365)
df = yf.download("AAPL", start=start.strftime("%Y-%m-%d"), end=end.strftime("%Y-%m-%d"), progress=False)
close = df["Close"].values.flatten().astype(np.float32)
volume = df["Volume"].values.flatten().astype(np.float32)
dates = df.index
# 分割
split_date = "2026-06-15"
split_idx = None
for i, d in enumerate(dates):
if str(d.date()) <= split_date:
split_idx = i
train_close = close[: split_idx + 1]
train_vol = volume[: split_idx + 1]
train_dates = dates[: split_idx + 1]
actual_close = close[split_idx + 1 :]
actual_dates = dates[split_idx + 1 :]
HORIZON = len(actual_close)
print(f"训练数据: {len(train_close)} 天(截至 {train_dates[-1].date()}", flush=True)
print(f"实际数据: {HORIZON} 天({actual_dates[0].date()} ~ {actual_dates[-1].date()}", flush=True)
# ---------- 2. 计算对数收益率 ----------
# log returns = ln(P_t / P_{t-1}),更平稳
train_logret = np.diff(np.log(train_close)).astype(np.float32) # length = N-1
actual_logret = np.diff(np.log(np.concatenate([train_close[-1:], actual_close]))).astype(np.float32)
print(f"对数收益率: 均值={train_logret.mean():.6f}, 标准差={train_logret.std():.6f}", flush=True)
# ---------- 3. 加载模型 ----------
print()
print("加载 TimesFM 模型...", flush=True)
torch.set_float32_matmul_precision("high")
model = timesfm.TimesFM_2p5_200M_torch.from_pretrained(
"google/timesfm-2.5-200m-pytorch", torch_compile=False
)
print("模型加载完成", flush=True)
# ---------- 4. 多窗口集成预测(对数收益率) ----------
CONTEXTS = [128, 256, 512] # 不同 context 长度
all_preds = []
all_quantiles = []
for ctx_len in CONTEXTS:
ctx_data = train_logret[-ctx_len:] if len(train_logret) >= ctx_len else train_logret
actual_ctx = min(ctx_len, len(ctx_data))
# 对收益率: infer_is_positive=False (可负), normalize_inputs=True
model.compile(
timesfm.ForecastConfig(
max_context=actual_ctx,
max_horizon=128,
normalize_inputs=True,
use_continuous_quantile_head=True,
force_flip_invariance=False,
infer_is_positive=False, # 收益率可正可负
fix_quantile_crossing=True,
)
)
pf, qf = model.forecast(horizon=HORIZON, inputs=[ctx_data])
all_preds.append(pf[0][:HORIZON])
all_quantiles.append(qf[0][:HORIZON])
print(f" context={actual_ctx} 预测完成", flush=True)
# 集成: 取平均
ensemble_pred_logret = np.mean(all_preds, axis=0)
ensemble_q_logret = np.mean(all_quantiles, axis=0)
# ---------- 5. 转换回价格 ----------
# P_t = P_{t-1} * exp(r_t)
last_price = train_close[-1]
pred_prices = []
for i in range(HORIZON):
last_price = last_price * np.exp(ensemble_pred_logret[i])
pred_prices.append(last_price)
pred_prices = np.array(pred_prices)
# 分位数价格
q_prices = np.zeros((HORIZON, 10))
for qi in range(10):
p = train_close[-1]
for i in range(HORIZON):
p = p * np.exp(ensemble_q_logret[i, qi])
q_prices[i, qi] = p
# ---------- 6. 计算误差 ----------
actual = actual_close
pred = pred_prices
mae = np.mean(np.abs(actual - pred))
rmse = np.sqrt(np.mean((actual - pred) ** 2))
mape = np.mean(np.abs((actual - pred) / actual)) * 100
actual_dir = np.diff(actual)
pred_dir = np.diff(pred)
dir_acc = np.mean(actual_dir * pred_dir > 0) * 100
in_80 = np.mean((actual >= q_prices[:, 1]) & (actual <= q_prices[:, 9])) * 100
in_40 = np.mean((actual >= q_prices[:, 3]) & (actual <= q_prices[:, 7])) * 100
print()
print("=== 优化版回测结果:6/16 ~ 6/30 ===", flush=True)
print()
print(" 日期 实际价 预测价 误差 误差%", flush=True)
for i in range(HORIZON):
err = pred[i] - actual[i]
err_pct = err / actual[i] * 100
print(
f" {actual_dates[i].date()} {actual[i]:7.2f} {pred[i]:7.2f} {err:+7.2f} {err_pct:+6.2f}%",
flush=True,
)
print()
print("=== 误差指标(优化版 vs 原始版)===", flush=True)
print(f" MAE: ${mae:.2f} (原始: $8.46)", flush=True)
print(f" RMSE: ${rmse:.2f} (原始: $11.29)", flush=True)
print(f" MAPE: {mape:.2f}% (原始: 2.98%)", flush=True)
print(f" 方向准确率: {dir_acc:.1f}% (原始: 55.6%)", flush=True)
print(f" 80%CI覆盖率: {in_80:.1f}% (原始: 90.0%)", flush=True)
print(f" 40%CI覆盖率: {in_40:.1f}% (原始: 60.0%)", flush=True)
# ---------- 7. 画对比图 ----------
fig, axes = plt.subplots(2, 1, figsize=(14, 10), sharex=False)
# 上图: 价格对比
ax = axes[0]
show_n = min(60, len(train_close))
ax.plot(range(show_n), train_close[-show_n:], label="Historical", color="steelblue", linewidth=1.5)
x_actual = range(show_n, show_n + HORIZON)
ax.plot(x_actual, actual, label="Actual", color="forestgreen", linewidth=2, marker="o", markersize=4)
ax.plot(x_actual, pred, label="Optimized Forecast", color="tomato", linewidth=2, linestyle="--")
ax.fill_between(x_actual, q_prices[:, 1], q_prices[:, 9], alpha=0.15, color="tomato", label="80% CI")
ax.fill_between(x_actual, q_prices[:, 3], q_prices[:, 7], alpha=0.3, color="tomato", label="40% CI")
ax.axvline(x=show_n - 1, color="gray", linestyle="--", alpha=0.5)
ax.set_title(f"Optimized: Log-Return + Ensemble (MAPE={mape:.2f}%, Dir={dir_acc:.0f}%)", fontsize=13)
ax.set_ylabel("Price (USD)")
ax.legend(loc="upper left")
# 下图: 预测误差对比
ax2 = axes[1]
errors = pred - actual
ax2.bar(range(HORIZON), errors, color=["tomato" if e > 0 else "steelblue" for e in errors], alpha=0.7)
ax2.axhline(y=0, color="black", linewidth=0.8)
ax2.set_title("Forecast Error per Day (Optimized)", fontsize=13)
ax2.set_xlabel("Trading Days After Split")
ax2.set_ylabel("Error (USD)")
ax2.set_xticks(range(HORIZON))
ax2.set_xticklabels([str(d.date()) for d in actual_dates], rotation=45, fontsize=8)
plt.tight_layout()
plt.savefig("aapl_backtest_optimized.png", dpi=150)
print()
print("Chart saved: aapl_backtest_optimized.png", flush=True)
print("Done!", flush=True)
+217
View File
@@ -0,0 +1,217 @@
"""TimesFM 纯技巧优化:
1. 对数收益率预测(更平稳)
2. 多窗口集成(128/256/512
3. SPY 大盘走势引导:先预测 SPY,用 SPY 预测的趋势辅助判断 AAPL 方向
4. 波动率调整:用近期波动率缩放置信区间
5. infer_is_positive=False(收益率可负)
"""
import yfinance as yf
import numpy as np
import torch
import timesfm
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from datetime import datetime, timedelta
# 1. 获取数据
print("下载 AAPL + SPY 数据...", flush=True)
end = datetime.now()
start = end - timedelta(days=365)
df_aapl = yf.download("AAPL", start=start.strftime("%Y-%m-%d"), end=end.strftime("%Y-%m-%d"), progress=False)
df_spy = yf.download("SPY", start=start.strftime("%Y-%m-%d"), end=end.strftime("%Y-%m-%d"), progress=False)
close = df_aapl["Close"].values.flatten().astype(np.float32)
spy_close = df_spy["Close"].values.flatten().astype(np.float32)
dates = df_aapl.index
min_len = min(len(close), len(spy_close))
close = close[-min_len:]
spy_close = spy_close[-min_len:]
dates = dates[-min_len:]
# 分割
split_date = "2026-06-15"
split_idx = None
for i, d in enumerate(dates):
if str(d.date()) <= split_date:
split_idx = i
train_close = close[: split_idx + 1]
train_spy = spy_close[: split_idx + 1]
train_dates = dates[: split_idx + 1]
actual_close = close[split_idx + 1 :]
actual_spy = spy_close[split_idx + 1 :]
actual_dates = dates[split_idx + 1 :]
HORIZON = len(actual_close)
print(f"训练数据: {len(train_close)} 天 | 预测: {HORIZON}", flush=True)
# 2. 对数收益率
train_logret = np.diff(np.log(train_close)).astype(np.float32)
train_spy_logret = np.diff(np.log(train_spy)).astype(np.float32)
actual_logret = np.diff(np.log(np.concatenate([train_close[-1:], actual_close]))).astype(np.float32)
# 3. 加载模型
print()
print("加载 TimesFM 模型...", flush=True)
torch.set_float32_matmul_precision("high")
model = timesfm.TimesFM_2p5_200M_torch.from_pretrained(
"google/timesfm-2.5-200m-pytorch", torch_compile=False
)
print("模型加载完成", flush=True)
# 4. 多窗口集成预测 AAPL 收益率
CONTEXTS = [128, 256, 512]
all_preds = []
all_quantiles = []
for ctx_len in CONTEXTS:
ctx_data = train_logret[-ctx_len:] if len(train_logret) >= ctx_len else train_logret
actual_ctx = min(ctx_len, len(ctx_data))
# 向上取整到 32 的倍数
actual_ctx = ((actual_ctx + 31) // 32) * 32
ctx_data = train_logret[-actual_ctx:]
model.compile(
timesfm.ForecastConfig(
max_context=actual_ctx,
max_horizon=128,
normalize_inputs=True,
use_continuous_quantile_head=True,
force_flip_invariance=False,
infer_is_positive=False,
fix_quantile_crossing=True,
)
)
pf, qf = model.forecast(horizon=HORIZON, inputs=[ctx_data])
all_preds.append(pf[0][:HORIZON])
all_quantiles.append(qf[0][:HORIZON])
print(f" AAPL context={actual_ctx} done", flush=True)
# 5. 预测 SPY 收益率(大盘趋势引导)
spy_ctx = train_spy_logret[-256:]
spy_ctx_len = ((len(spy_ctx) + 31) // 32) * 32
spy_ctx = train_spy_logret[-spy_ctx_len:]
model.compile(
timesfm.ForecastConfig(
max_context=spy_ctx_len,
max_horizon=128,
normalize_inputs=True,
use_continuous_quantile_head=True,
force_flip_invariance=False,
infer_is_positive=False,
fix_quantile_crossing=True,
)
)
spy_pf, _ = model.forecast(horizon=HORIZON, inputs=[spy_ctx])
print(f" SPY context={spy_ctx_len} done", flush=True)
# 6. 集成 + SPY 趋势调整
ensemble_pred_logret = np.mean(all_preds, axis=0)
ensemble_q_logret = np.mean(all_quantiles, axis=0)
spy_pred_logret = spy_pf[0][:HORIZON]
# SPY 趋势调整:如果 SPY 预测下跌,对 AAPL 预测施加向下的调整
# 计算 AAPL 对 SPY 的 beta(敏感度)
beta = np.corrcoef(train_logret[-60:], train_spy_logret[-60:])[0, 1]
print(f" AAPL-SPY 60日相关系数: {beta:.3f}", flush=True)
# 调整:将 SPY 预测的偏离均值部分 * beta 加到 AAPL 预测上
spy_mean = np.mean(train_spy_logret[-60:])
spy_deviation = spy_pred_logret - spy_mean # SPY 偏离其均值的部分
adjustment = beta * spy_deviation * 0.3 # 0.3 是调整强度,避免过度修正
adjusted_pred_logret = ensemble_pred_logret + adjustment
# 7. 波动率调整置信区间
recent_vol = np.std(train_logret[-20:])
long_vol = np.std(train_logret[-60:])
print(f" 近20日波动率: {recent_vol:.5f} | 近60日波动率: {long_vol:.5f}", flush=True)
# 如果近期波动率高于长期,扩大置信区间
vol_ratio = recent_vol / max(long_vol, 1e-8)
vol_scale = max(vol_ratio, 1.0) # 只扩大不缩小
adjusted_q_logret = ensemble_q_logret.copy()
median_idx = 5
for qi in range(10):
if qi != median_idx:
adjusted_q_logret[:, qi] = ensemble_q_logret[:, median_idx] + (
ensemble_q_logret[:, qi] - ensemble_q_logret[:, median_idx]
) * vol_scale
# 8. 转换回价格
last_price = train_close[-1]
pred_prices = []
for i in range(HORIZON):
last_price = last_price * np.exp(adjusted_pred_logret[i])
pred_prices.append(last_price)
pred_prices = np.array(pred_prices)
q_prices = np.zeros((HORIZON, 10))
for qi in range(10):
p = train_close[-1]
for i in range(HORIZON):
p = p * np.exp(adjusted_q_logret[i, qi])
q_prices[i, qi] = p
# 9. 误差指标
actual = actual_close
pred = pred_prices
mae = np.mean(np.abs(actual - pred))
rmse = np.sqrt(np.mean((actual - pred) ** 2))
mape = np.mean(np.abs((actual - pred) / actual)) * 100
actual_dir = np.diff(actual)
pred_dir = np.diff(pred)
dir_acc = np.mean(actual_dir * pred_dir > 0) * 100
in_80 = np.mean((actual >= q_prices[:, 1]) & (actual <= q_prices[:, 9])) * 100
in_40 = np.mean((actual >= q_prices[:, 3]) & (actual <= q_prices[:, 7])) * 100
print()
print("=== 纯技巧优化版回测结果 ===", flush=True)
print()
print(" 日期 实际价 预测价 误差 误差%", flush=True)
for i in range(HORIZON):
err = pred[i] - actual[i]
err_pct = err / actual[i] * 100
print(f" {actual_dates[i].date()} {actual[i]:7.2f} {pred[i]:7.2f} {err:+7.2f} {err_pct:+6.2f}%", flush=True)
print()
print("=== 误差指标(四版对比)===", flush=True)
print(f" MAE: ${mae:.2f} (原始: $8.46, 对数收益率: $8.00, XReg: $9.21)", flush=True)
print(f" RMSE: ${rmse:.2f} (原始: $11.29, 对数收益率: $11.09, XReg: $12.47)", flush=True)
print(f" MAPE: {mape:.2f}% (原始: 2.98%, 对数收益率: 2.82%, XReg: 3.25%)", flush=True)
print(f" 方向准确率: {dir_acc:.1f}% (原始: 55.6%, 对数收益率: 44.4%, XReg: 44.4%)", flush=True)
print(f" 80%CI覆盖率: {in_80:.1f}% (原始: 90.0%, 对数收益率: 100.0%, XReg: 70.0%)", flush=True)
print(f" 40%CI覆盖率: {in_40:.1f}% (原始: 60.0%, 对数收益率: 80.0%, XReg: 40.0%)", flush=True)
# 10. 画图
fig, axes = plt.subplots(2, 1, figsize=(14, 10))
show_n = min(60, len(train_close))
ax = axes[0]
ax.plot(range(show_n), train_close[-show_n:], label="Historical", color="steelblue", linewidth=1.5)
x_actual = range(show_n, show_n + HORIZON)
ax.plot(x_actual, actual, label="Actual", color="forestgreen", linewidth=2, marker="o", markersize=4)
ax.plot(x_actual, pred, label="Trick Forecast", color="purple", linewidth=2, linestyle="--")
ax.fill_between(x_actual, q_prices[:, 1], q_prices[:, 9], alpha=0.15, color="purple", label="80% CI")
ax.fill_between(x_actual, q_prices[:, 3], q_prices[:, 7], alpha=0.3, color="purple", label="40% CI")
ax.axvline(x=show_n - 1, color="gray", linestyle="--", alpha=0.5)
ax.set_title(f"Trick: LogRet + Ensemble + SPY + VolAdj (MAPE={mape:.2f}%, Dir={dir_acc:.0f}%)", fontsize=13)
ax.set_ylabel("Price (USD)")
ax.legend(loc="upper left")
ax2 = axes[1]
errors = pred - actual
ax2.bar(range(HORIZON), errors, color=["purple" if e > 0 else "steelblue" for e in errors], alpha=0.7)
ax2.axhline(y=0, color="black", linewidth=0.8)
ax2.set_title("Trick Forecast Error per Day", fontsize=13)
ax2.set_xlabel("Trading Days After Split")
ax2.set_ylabel("Error (USD)")
ax2.set_xticks(range(HORIZON))
ax2.set_xticklabels([str(d.date()) for d in actual_dates], rotation=45, fontsize=8)
plt.tight_layout()
plt.savefig("aapl_backtest_trick.png", dpi=150)
print()
print("Chart saved: aapl_backtest_trick.png", flush=True)
print("Done!", flush=True)
+203
View File
@@ -0,0 +1,203 @@
"""TimesFM + XReg 协变量预测:加入成交量、RSI、SPY 大盘指数"""
import yfinance as yf
import numpy as np
import torch
import timesfm
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from datetime import datetime, timedelta
def calc_rsi(prices, period=14):
deltas = np.diff(prices)
gains = np.where(deltas > 0, deltas, 0.0)
losses = np.where(deltas < 0, -deltas, 0.0)
avg_gain = np.convolve(gains, np.ones(period) / period, mode="valid")
avg_loss = np.convolve(losses, np.ones(period) / period, mode="valid")
avg_loss = np.where(avg_loss == 0, 1e-10, avg_loss)
rs = avg_gain / avg_loss
rsi = 100.0 - (100.0 / (1.0 + rs))
return np.concatenate([np.full(period, 50.0), rsi])
# 1. 获取数据
print("下载 AAPL + SPY 数据...", flush=True)
end = datetime.now()
start = end - timedelta(days=365)
df_aapl = yf.download("AAPL", start=start.strftime("%Y-%m-%d"), end=end.strftime("%Y-%m-%d"), progress=False)
df_spy = yf.download("SPY", start=start.strftime("%Y-%m-%d"), end=end.strftime("%Y-%m-%d"), progress=False)
close = df_aapl["Close"].values.flatten().astype(np.float64)
volume = df_aapl["Volume"].values.flatten().astype(np.float64)
spy_close = df_spy["Close"].values.flatten().astype(np.float64)
dates = df_aapl.index
# 对齐长度
min_len = min(len(close), len(spy_close))
close = close[-min_len:]
volume = volume[-min_len:]
spy_close = spy_close[-min_len:]
dates = dates[-min_len:]
# 计算 RSI
rsi = calc_rsi(close)
# 分割
split_date = "2026-06-15"
split_idx = None
for i, d in enumerate(dates):
if str(d.date()) <= split_date:
split_idx = i
train_close = close[: split_idx + 1]
train_vol = volume[: split_idx + 1]
train_spy = spy_close[: split_idx + 1]
train_rsi = rsi[: split_idx + 1]
train_dates = dates[: split_idx + 1]
actual_close = close[split_idx + 1 :]
actual_dates = dates[split_idx + 1 :]
HORIZON = len(actual_close)
print(f"训练数据: {len(train_close)}", flush=True)
print(f"预测目标: {HORIZON}", flush=True)
# 2. 准备协变量
# 动态数值协变量需要覆盖 context + horizon 的完整长度
# 对每个序列: train 部分用实际值,test 部分需要"未来值"
# 对于 RSI 和 Volume,我们没有未来值,用最后一个值填充
# 对于 SPY,用训练集最后一个值填充(因为我们无法预知未来 SPY)
# 协变量需要: 每个协变量是一个 list,每个元素对应一个输入序列的完整长度(context+horizon)
full_len = len(train_close) + HORIZON
# 成交量: train 用实际值,future 用最近 5 日均值
vol_future = np.mean(train_vol[-5:])
vol_full = np.concatenate([train_vol, np.full(HORIZON, vol_future)])
# RSI: train 用实际值,future 用 50(中性)
rsi_full = np.concatenate([train_rsi, np.full(HORIZON, 50.0)])
# SPY: train 用实际值,future 用最后一个值
spy_future = train_spy[-1]
spy_full = np.concatenate([train_spy, np.full(HORIZON, spy_future)])
# 3. 加载模型
print()
print("加载 TimesFM 模型...", flush=True)
torch.set_float32_matmul_precision("high")
model = timesfm.TimesFM_2p5_200M_torch.from_pretrained(
"google/timesfm-2.5-200m-pytorch", torch_compile=False
)
print("模型加载完成", flush=True)
# 4. 编译(return_backcast=True 是 XReg 必需的)
model.compile(
timesfm.ForecastConfig(
max_context=512,
max_horizon=128,
normalize_inputs=True,
use_continuous_quantile_head=True,
force_flip_invariance=False,
infer_is_positive=True,
fix_quantile_crossing=True,
return_backcast=True, # XReg 需要
)
)
print("编译完成", flush=True)
# 5. 用 XReg 预测
print("运行 XReg 协变量预测...", flush=True)
# 动态数值协变量: dict[str, list[list[float]]]
# 每个协变量是一个 list,其中每个元素是一个序列(对应一个输入)
# 这里只有一个输入序列
dynamic_num_covs = {
"volume": [vol_full],
"rsi": [rsi_full],
"spy_close": [spy_full],
}
# 静态数值协变量
static_num_covs = {
"avg_volume": [np.mean(train_vol)],
}
point_outputs, quantile_outputs = model.forecast_with_covariates(
inputs=[train_close],
dynamic_numerical_covariates=dynamic_num_covs,
static_numerical_covariates=static_num_covs,
xreg_mode="xreg + timesfm", # 先回归再预测残差
normalize_xreg_target_per_input=True,
ridge=1.0,
)
print("XReg 预测完成!", flush=True)
# 6. 提取结果
pred = np.array(point_outputs[0][:HORIZON])
q = np.array(quantile_outputs[0]) # (horizon, 10) or (full, 10)
# quantile_outputs 可能包含 backcast,取最后 HORIZON 个
if q.shape[0] > HORIZON:
q = q[-HORIZON:]
# 7. 计算误差
actual = actual_close
mae = np.mean(np.abs(actual - pred))
rmse = np.sqrt(np.mean((actual - pred) ** 2))
mape = np.mean(np.abs((actual - pred) / actual)) * 100
actual_dir = np.diff(actual)
pred_dir = np.diff(pred)
dir_acc = np.mean(actual_dir * pred_dir > 0) * 100
in_80 = np.mean((actual >= q[:, 1]) & (actual <= q[:, 9])) * 100
in_40 = np.mean((actual >= q[:, 3]) & (actual <= q[:, 7])) * 100
print()
print("=== XReg 协变量版回测结果 ===", flush=True)
print()
print(" 日期 实际价 预测价 误差 误差%", flush=True)
for i in range(HORIZON):
err = pred[i] - actual[i]
err_pct = err / actual[i] * 100
print(f" {actual_dates[i].date()} {actual[i]:7.2f} {pred[i]:7.2f} {err:+7.2f} {err_pct:+6.2f}%", flush=True)
print()
print("=== 误差指标(三版对比)===", flush=True)
print(f" MAE: ${mae:.2f} (原始: $8.46, 优化: $8.00)", flush=True)
print(f" RMSE: ${rmse:.2f} (原始: $11.29, 优化: $11.09)", flush=True)
print(f" MAPE: {mape:.2f}% (原始: 2.98%, 优化: 2.82%)", flush=True)
print(f" 方向准确率: {dir_acc:.1f}% (原始: 55.6%, 优化: 44.4%)", flush=True)
print(f" 80%CI覆盖率: {in_80:.1f}% (原始: 90.0%, 优化: 100.0%)", flush=True)
print(f" 40%CI覆盖率: {in_40:.1f}% (原始: 60.0%, 优化: 80.0%)", flush=True)
# 8. 画图
fig, axes = plt.subplots(2, 1, figsize=(14, 10))
show_n = min(60, len(train_close))
ax = axes[0]
ax.plot(range(show_n), train_close[-show_n:], label="Historical", color="steelblue", linewidth=1.5)
x_actual = range(show_n, show_n + HORIZON)
ax.plot(x_actual, actual, label="Actual", color="forestgreen", linewidth=2, marker="o", markersize=4)
ax.plot(x_actual, pred, label="XReg Forecast", color="darkorange", linewidth=2, linestyle="--")
ax.fill_between(x_actual, q[:, 1], q[:, 9], alpha=0.15, color="darkorange", label="80% CI")
ax.fill_between(x_actual, q[:, 3], q[:, 7], alpha=0.3, color="darkorange", label="40% CI")
ax.axvline(x=show_n - 1, color="gray", linestyle="--", alpha=0.5)
ax.set_title(f"XReg: AAPL + Volume/RSI/SPY (MAPE={mape:.2f}%, Dir={dir_acc:.0f}%)", fontsize=13)
ax.set_ylabel("Price (USD)")
ax.legend(loc="upper left")
# 误差柱状图
ax2 = axes[1]
errors = pred - actual
ax2.bar(range(HORIZON), errors, color=["darkorange" if e > 0 else "steelblue" for e in errors], alpha=0.7)
ax2.axhline(y=0, color="black", linewidth=0.8)
ax2.set_title("XReg Forecast Error per Day", fontsize=13)
ax2.set_xlabel("Trading Days After Split")
ax2.set_ylabel("Error (USD)")
ax2.set_xticks(range(HORIZON))
ax2.set_xticklabels([str(d.date()) for d in actual_dates], rotation=45, fontsize=8)
plt.tight_layout()
plt.savefig("aapl_backtest_xreg.png", dpi=150)
print()
print("Chart saved: aapl_backtest_xreg.png", flush=True)
print("Done!", flush=True)
+84
View File
@@ -0,0 +1,84 @@
import yfinance as yf
import numpy as np
import torch
import timesfm
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from datetime import datetime, timedelta
# 1. 获取 AAPL 近 1 年收盘价
print("下载 AAPL 股票数据...", flush=True)
end = datetime.now()
start = end - timedelta(days=365)
df = yf.download("AAPL", start=start.strftime("%Y-%m-%d"), end=end.strftime("%Y-%m-%d"), progress=False)
close = df["Close"].values.flatten().astype(np.float32)
dates = df.index
print(f"获取到 {len(close)} 个交易日", flush=True)
print(f"价格范围: {close.min():.2f} ~ {close.max():.2f}", flush=True)
print(f"最近 5 日收盘价: {close[-5:]}", flush=True)
# 2. 加载 TimesFM 模型
print()
print("加载 TimesFM 模型...", flush=True)
torch.set_float32_matmul_precision("high")
model = timesfm.TimesFM_2p5_200M_torch.from_pretrained(
"google/timesfm-2.5-200m-pytorch", torch_compile=False
)
print("模型加载完成", flush=True)
# 3. 编译并预测未来 20 个交易日
HORIZON = 20
model.compile(
timesfm.ForecastConfig(
max_context=512,
max_horizon=128,
normalize_inputs=True,
use_continuous_quantile_head=True,
force_flip_invariance=False,
infer_is_positive=True,
fix_quantile_crossing=True,
)
)
print("编译完成,开始预测...", flush=True)
pf, qf = model.forecast(horizon=HORIZON, inputs=[close])
print("预测完成!", flush=True)
# 4. 输出结果
print()
print("=== AAPL 未来 20 个交易日预测 ===", flush=True)
print(f"当前价格: {close[-1]:.2f}", flush=True)
print()
print(" 日期(估计) 点预测 q10(低) q90(高)", flush=True)
last_date = dates[-1]
for i in range(HORIZON):
est_date = last_date + timedelta(days=i + 1)
print(
f" {est_date.strftime('%Y-%m-%d')} {pf[0][i]:7.2f} {qf[0][i,1]:7.2f} {qf[0][i,9]:7.2f}",
flush=True,
)
print()
print(f"预测均价: {pf[0].mean():.2f}", flush=True)
print(f"预测涨跌: {(pf[0][-1] - close[-1]) / close[-1] * 100:+.2f}%", flush=True)
print(f"80%置信区间: {qf[0,-1,1]:.2f} ~ {qf[0,-1,9]:.2f}", flush=True)
# 5. 画图
fig, ax = plt.subplots(figsize=(14, 6))
show_n = min(60, len(close))
ax.plot(range(show_n), close[-show_n:], label="历史收盘价", color="steelblue", linewidth=1.5)
x_fc = range(show_n, show_n + HORIZON)
ax.plot(x_fc, pf[0], label="点预测(中位数)", color="tomato", linewidth=2)
ax.fill_between(x_fc, qf[0, :, 1], qf[0, :, 9], alpha=0.2, color="tomato", label="80% 置信区间")
ax.fill_between(x_fc, qf[0, :, 3], qf[0, :, 7], alpha=0.35, color="tomato", label="40% 置信区间")
ax.axvline(x=show_n - 1, color="gray", linestyle="--", alpha=0.5, label="预测起点")
ax.set_title("AAPL 收盘价预测 (TimesFM 2.5)", fontsize=14)
ax.set_xlabel("交易日")
ax.set_ylabel("价格 (USD)")
ax.legend(loc="upper left")
plt.tight_layout()
plt.savefig("aapl_forecast.png", dpi=150)
print()
print("图表已保存: aapl_forecast.png", flush=True)
print("✅ 完成!", flush=True)
+68
View File
@@ -0,0 +1,68 @@
# Copyright 2025 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Tests for loading TimesFM 2.5 models."""
import os
import tempfile
from timesfm.timesfm_2p5.timesfm_2p5_torch import TimesFM_2p5_200M_torch
from timesfm.timesfm_2p5.timesfm_2p5_flax import TimesFM_2p5_200M_flax
class TestModelLoading:
"""Tests to verify model instantiation, loading, and compatibility."""
def test_torch_load_checkpoint_and_from_pretrained_local(self):
"""Verifies that PyTorch load_checkpoint and from_pretrained work locally."""
# 1. Instantiate the model wrapper with compilation disabled
tfm = TimesFM_2p5_200M_torch(torch_compile=False)
with tempfile.TemporaryDirectory() as tmpdir:
# 2. Save the model's randomly-initialized weights
tfm._save_pretrained(tmpdir)
# Verify weights file is written
weights_path = os.path.join(tmpdir, "model.safetensors")
assert os.path.exists(weights_path)
# 3. Verify that load_checkpoint works from the temp directory path
tfm2 = TimesFM_2p5_200M_torch(torch_compile=False)
tfm2.load_checkpoint(tmpdir, torch_compile=False)
# 4. Verify that from_pretrained works with a local directory path
# and accepts/ignores extra kwargs (like proxies) without raising TypeError
tfm3 = TimesFM_2p5_200M_torch.from_pretrained(
tmpdir,
torch_compile=False,
proxies={"http": "http://dummy.proxy"},
custom_kwarg="dummy_value",
)
assert tfm3 is not None
assert not tfm3.torch_compile
# 5. Run a simple prediction step to verify the loaded model performs forward pass
import numpy as np
inputs = [np.random.randn(32)]
forecasts = tfm3.model.forecast_naive(horizon=10, inputs=inputs)
assert len(forecasts) == 1
assert forecasts[0].shape == (10, 10)
def test_flax_model_init_kwargs(self):
"""Verifies that Flax model wrapper constructor accepts arbitrary kwargs."""
tfm = TimesFM_2p5_200M_flax(
proxies={"http": "http://dummy.proxy"},
custom_kwarg="dummy_value",
)
assert tfm is not None
@@ -1,13 +1,13 @@
date,point_forecast,q10,q20,q30,q40,q50,q60,q70,q80,q90,q99
2025-01-01,1.2593384,1.248188,1.140702,1.1880752,1.2137158,1.2394564,1.2593384,1.2767732,1.297132,1.32396,1.367888
2025-02-01,1.2856668,1.2773758,1.1406044,1.1960833,1.2322671,1.2593892,1.2856668,1.3110137,1.3400218,1.3751202,1.4253658
2025-03-01,1.2950127,1.2869918,1.126852,1.1876173,1.234988,1.2675052,1.2950127,1.328448,1.354729,1.4035482,1.4642649
2025-04-01,1.2207624,1.2084007,1.0352504,1.1041918,1.151865,1.1853008,1.2207624,1.256663,1.2898555,1.3310349,1.4016538
2025-05-01,1.1702554,1.153313,0.9691495,1.0431063,1.0932612,1.1276176,1.1702554,1.201966,1.2390311,1.2891905,1.3632389
2025-06-01,1.1455553,1.1275499,0.94203794,1.0110554,1.0658777,1.1061188,1.1455553,1.1806211,1.2180579,1.2702757,1.345366
2025-07-01,1.1702348,1.1510556,0.9503718,1.0347577,1.0847733,1.1287677,1.1702348,1.2114835,1.2482276,1.2997853,1.3807325
2025-08-01,1.2026825,1.1859496,0.9709255,1.0594383,1.1106675,1.1579902,1.2026825,1.2399211,1.2842004,1.3408126,1.419526
2025-09-01,1.1909748,1.1784849,0.95943713,1.0403702,1.103606,1.1511956,1.1909748,1.2390201,1.2832941,1.3354731,1.416972
2025-10-01,1.1490841,1.1264795,0.9079477,0.99529266,1.0548235,1.1052223,1.1490841,1.1897774,1.240414,1.2868769,1.3775467
2025-11-01,1.0804785,1.0624356,0.8361266,0.9259792,0.9882403,1.0386353,1.0804785,1.1281581,1.1759715,1.228377,1.3122478
2025-12-01,1.0613453,1.0366092,0.80220693,0.89521873,0.9593707,1.0152239,1.0613453,1.1032857,1.15315,1.216908,1.2959521
date,point_forecast,mean,q10,q20,q30,q40,q50,q60,q70,q80,q90
2025-01-01,1.2223774,1.2215943,1.1230627,1.1613995,1.1832488,1.2030286,1.2223774,1.240989,1.2626991,1.2927711,1.339603
2025-02-01,1.2563584,1.2502017,1.1482248,1.1891649,1.2134029,1.2355918,1.2563584,1.2787917,1.3061507,1.3358022,1.3880283
2025-03-01,1.286477,1.2816916,1.1694773,1.2141376,1.2430503,1.2636719,1.286477,1.3101407,1.3365126,1.3725822,1.4265702
2025-04-01,1.240488,1.2406754,1.119298,1.1689228,1.1957527,1.2162383,1.240488,1.265359,1.2880774,1.3237736,1.3806401
2025-05-01,1.2026378,1.1969143,1.0776879,1.1280149,1.1554759,1.1801765,1.2026378,1.2274072,1.2546852,1.2890899,1.3469675
2025-06-01,1.21002,1.1963896,1.0811386,1.1352499,1.1687461,1.1853771,1.21002,1.2333429,1.255994,1.2938039,1.3532466
2025-07-01,1.2253109,1.2151253,1.0917634,1.1474797,1.176413,1.2029413,1.2253109,1.250956,1.2764856,1.3113455,1.3728349
2025-08-01,1.2421811,1.2292916,1.1043029,1.1598811,1.1929151,1.2164187,1.2421811,1.266998,1.2924675,1.3304923,1.394975
2025-09-01,1.269735,1.2603163,1.1239274,1.1872942,1.2162873,1.242951,1.269735,1.2976494,1.3236798,1.3580028,1.4252181
2025-10-01,1.2496669,1.2436218,1.0962446,1.1630417,1.1938806,1.221712,1.2496669,1.2757387,1.3041382,1.3376691,1.409816
2025-11-01,1.2135266,1.2031629,1.0545524,1.1223581,1.157211,1.1864623,1.2135266,1.2447275,1.2706528,1.3091388,1.3801678
2025-12-01,1.2034141,1.1867243,1.0412307,1.1113251,1.1433371,1.1752325,1.2034141,1.2308054,1.2621957,1.2912495,1.370078
1 date point_forecast mean q10 q20 q30 q40 q50 q60 q70 q80 q99 q90
2 2025-01-01 1.2593384 1.2223774 1.2215943 1.248188 1.1230627 1.140702 1.1613995 1.1880752 1.1832488 1.2137158 1.2030286 1.2394564 1.2223774 1.2593384 1.240989 1.2767732 1.2626991 1.297132 1.2927711 1.367888 1.32396 1.339603
3 2025-02-01 1.2856668 1.2563584 1.2502017 1.2773758 1.1482248 1.1406044 1.1891649 1.1960833 1.2134029 1.2322671 1.2355918 1.2593892 1.2563584 1.2856668 1.2787917 1.3110137 1.3061507 1.3400218 1.3358022 1.4253658 1.3751202 1.3880283
4 2025-03-01 1.2950127 1.286477 1.2816916 1.2869918 1.1694773 1.126852 1.2141376 1.1876173 1.2430503 1.234988 1.2636719 1.2675052 1.286477 1.2950127 1.3101407 1.328448 1.3365126 1.354729 1.3725822 1.4642649 1.4035482 1.4265702
5 2025-04-01 1.2207624 1.240488 1.2406754 1.2084007 1.119298 1.0352504 1.1689228 1.1041918 1.1957527 1.151865 1.2162383 1.1853008 1.240488 1.2207624 1.265359 1.256663 1.2880774 1.2898555 1.3237736 1.4016538 1.3310349 1.3806401
6 2025-05-01 1.1702554 1.2026378 1.1969143 1.153313 1.0776879 0.9691495 1.1280149 1.0431063 1.1554759 1.0932612 1.1801765 1.1276176 1.2026378 1.1702554 1.2274072 1.201966 1.2546852 1.2390311 1.2890899 1.3632389 1.2891905 1.3469675
7 2025-06-01 1.1455553 1.21002 1.1963896 1.1275499 1.0811386 0.94203794 1.1352499 1.0110554 1.1687461 1.0658777 1.1853771 1.1061188 1.21002 1.1455553 1.2333429 1.1806211 1.255994 1.2180579 1.2938039 1.345366 1.2702757 1.3532466
8 2025-07-01 1.1702348 1.2253109 1.2151253 1.1510556 1.0917634 0.9503718 1.1474797 1.0347577 1.176413 1.0847733 1.2029413 1.1287677 1.2253109 1.1702348 1.250956 1.2114835 1.2764856 1.2482276 1.3113455 1.3807325 1.2997853 1.3728349
9 2025-08-01 1.2026825 1.2421811 1.2292916 1.1859496 1.1043029 0.9709255 1.1598811 1.0594383 1.1929151 1.1106675 1.2164187 1.1579902 1.2421811 1.2026825 1.266998 1.2399211 1.2924675 1.2842004 1.3304923 1.419526 1.3408126 1.394975
10 2025-09-01 1.1909748 1.269735 1.2603163 1.1784849 1.1239274 0.95943713 1.1872942 1.0403702 1.2162873 1.103606 1.242951 1.1511956 1.269735 1.1909748 1.2976494 1.2390201 1.3236798 1.2832941 1.3580028 1.416972 1.3354731 1.4252181
11 2025-10-01 1.1490841 1.2496669 1.2436218 1.1264795 1.0962446 0.9079477 1.1630417 0.99529266 1.1938806 1.0548235 1.221712 1.1052223 1.2496669 1.1490841 1.2757387 1.1897774 1.3041382 1.240414 1.3376691 1.3775467 1.2868769 1.409816
12 2025-11-01 1.0804785 1.2135266 1.2031629 1.0624356 1.0545524 0.8361266 1.1223581 0.9259792 1.157211 0.9882403 1.1864623 1.0386353 1.2135266 1.0804785 1.2447275 1.1281581 1.2706528 1.1759715 1.3091388 1.3122478 1.228377 1.3801678
13 2025-12-01 1.0613453 1.2034141 1.1867243 1.0366092 1.0412307 0.80220693 1.1113251 0.89521873 1.1433371 0.9593707 1.1752325 1.0152239 1.2034141 1.0613453 1.2308054 1.1032857 1.2621957 1.15315 1.2912495 1.2959521 1.216908 1.370078
@@ -1,5 +1,5 @@
{
"model": "TimesFM 1.0 (200M) PyTorch",
"model": "TimesFM 2.5 (200M) PyTorch",
"input": {
"source": "NOAA GISTEMP Global Temperature Anomaly",
"n_observations": 36,
@@ -23,166 +23,166 @@
"2025-12"
],
"point": [
1.25933837890625,
1.285666823387146,
1.2950127124786377,
1.2207623720169067,
1.170255422592163,
1.1455552577972412,
1.1702347993850708,
1.2026824951171875,
1.1909748315811157,
1.1490840911865234,
1.080478549003601,
1.0613453388214111
1.2223774194717407,
1.2563583850860596,
1.286476969718933,
1.240488052368164,
1.202637791633606,
1.2100199460983276,
1.2253109216690063,
1.2421810626983643,
1.2697349786758423,
1.2496669292449951,
1.2135266065597534,
1.2034140825271606
],
"quantiles": {
"mean": [
1.2215943336486816,
1.25020170211792,
1.281691551208496,
1.240675449371338,
1.1969143152236938,
1.1963895559310913,
1.215125322341919,
1.229291558265686,
1.260316252708435,
1.243621826171875,
1.2031629085540771,
1.186724305152893
],
"10%": [
1.2481880187988281,
1.2773758172988892,
1.286991834640503,
1.2084007263183594,
1.1533130407333374,
1.1275498867034912,
1.1510555744171143,
1.1859495639801025,
1.1784849166870117,
1.1264795064926147,
1.0624356269836426,
1.036609172821045
1.1230627298355103,
1.1482248306274414,
1.1694773435592651,
1.119297981262207,
1.0776878595352173,
1.0811386108398438,
1.0917633771896362,
1.1043028831481934,
1.123927354812622,
1.0962445735931396,
1.054552435874939,
1.0412306785583496
],
"20%": [
1.1407020092010498,
1.1406043767929077,
1.126852035522461,
1.0352504253387451,
0.9691494703292847,
0.9420379400253296,
0.9503718018531799,
0.970925509929657,
0.9594371318817139,
0.9079477190971375,
0.8361266255378723,
0.8022069334983826
1.161399483680725,
1.1891648769378662,
1.2141375541687012,
1.168922781944275,
1.1280149221420288,
1.1352498531341553,
1.1474796533584595,
1.1598811149597168,
1.1872942447662354,
1.1630417108535767,
1.1223580837249756,
1.1113251447677612
],
"30%": [
1.1880751848220825,
1.1960833072662354,
1.187617301940918,
1.104191780090332,
1.0431063175201416,
1.01105535030365,
1.0347577333450317,
1.0594383478164673,
1.040370225906372,
0.9952926635742188,
0.9259791970252991,
0.8952187299728394
1.18324875831604,
1.2134028673171997,
1.2430503368377686,
1.195752739906311,
1.1554758548736572,
1.1687461137771606,
1.1764130592346191,
1.1929150819778442,
1.2162872552871704,
1.193880558013916,
1.1572109460830688,
1.1433371305465698
],
"40%": [
1.2137157917022705,
1.232267141342163,
1.2349879741668701,
1.151865005493164,
1.0932612419128418,
1.0658776760101318,
1.084773302078247,
1.1106674671173096,
1.1036059856414795,
1.0548235177993774,
0.9882403016090393,
0.9593706727027893
1.2030285596847534,
1.2355917692184448,
1.263671875,
1.216238260269165,
1.1801764965057373,
1.1853771209716797,
1.2029412984848022,
1.216418743133545,
1.2429510354995728,
1.2217119932174683,
1.1864622831344604,
1.1752325296401978
],
"50%": [
1.2394564151763916,
1.2593891620635986,
1.267505168914795,
1.1853008270263672,
1.127617597579956,
1.1061187982559204,
1.128767728805542,
1.1579902172088623,
1.1511956453323364,
1.1052223443984985,
1.03863525390625,
1.0152238607406616
1.2223774194717407,
1.2563583850860596,
1.286476969718933,
1.240488052368164,
1.202637791633606,
1.2100199460983276,
1.2253109216690063,
1.2421810626983643,
1.2697349786758423,
1.2496669292449951,
1.2135266065597534,
1.2034140825271606
],
"60%": [
1.25933837890625,
1.285666823387146,
1.2950127124786377,
1.2207623720169067,
1.170255422592163,
1.1455552577972412,
1.1702347993850708,
1.2026824951171875,
1.1909748315811157,
1.1490840911865234,
1.080478549003601,
1.0613453388214111
1.2409889698028564,
1.2787916660308838,
1.3101407289505005,
1.2653590440750122,
1.2274072170257568,
1.2333428859710693,
1.2509560585021973,
1.266998052597046,
1.2976493835449219,
1.2757387161254883,
1.2447274923324585,
1.2308053970336914
],
"70%": [
1.27677321434021,
1.3110136985778809,
1.3284480571746826,
1.2566629648208618,
1.2019660472869873,
1.1806211471557617,
1.2114834785461426,
1.2399210929870605,
1.2390201091766357,
1.1897773742675781,
1.1281580924987793,
1.1032856702804565
1.2626991271972656,
1.3061506748199463,
1.336512565612793,
1.2880773544311523,
1.2546851634979248,
1.2559939622879028,
1.276485562324524,
1.292467474937439,
1.323679804801941,
1.30413818359375,
1.2706527709960938,
1.2621957063674927
],
"80%": [
1.2971320152282715,
1.3400218486785889,
1.3547290563583374,
1.2898554801940918,
1.2390310764312744,
1.2180578708648682,
1.248227596282959,
1.2842004299163818,
1.2832940816879272,
1.240414023399353,
1.175971508026123,
1.153149962425232
1.2927711009979248,
1.3358021974563599,
1.372582197189331,
1.3237736225128174,
1.2890899181365967,
1.2938039302825928,
1.3113454580307007,
1.3304922580718994,
1.358002781867981,
1.3376691341400146,
1.3091387748718262,
1.2912495136260986
],
"90%": [
1.3239599466323853,
1.3751201629638672,
1.403548240661621,
1.3310348987579346,
1.2891905307769775,
1.2702757120132446,
1.2997852563858032,
1.3408125638961792,
1.3354730606079102,
1.286876916885376,
1.2283769845962524,
1.2169079780578613
],
"99%": [
1.3678879737854004,
1.4253658056259155,
1.4642648696899414,
1.40165376663208,
1.3632389307022095,
1.3453660011291504,
1.380732536315918,
1.4195259809494019,
1.416972041130066,
1.3775466680526733,
1.3122477531433105,
1.2959520816802979
1.3396029472351074,
1.3880282640457153,
1.426570177078247,
1.3806401491165161,
1.3469674587249756,
1.3532465696334839,
1.3728349208831787,
1.394974946975708,
1.425218105316162,
1.409816026687622,
1.380167841911316,
1.3700779676437378
]
}
},
"summary": {
"forecast_mean_c": 1.186,
"forecast_max_c": 1.295,
"forecast_min_c": 1.061,
"vs_last_year_mean": -0.067
"forecast_mean_c": 1.235,
"forecast_max_c": 1.286,
"forecast_min_c": 1.203,
"vs_last_year_mean": -0.017
}
}
Binary file not shown.

Before

Width:  |  Height:  |  Size: 153 KiB

After

Width:  |  Height:  |  Size: 147 KiB

@@ -11,6 +11,7 @@ from pathlib import Path
import numpy as np
import pandas as pd
import timesfm
# Preflight check
print("=" * 60)
@@ -35,27 +36,28 @@ print(
# TimesFM expects a list of 1D numpy arrays
input_series = df["anomaly_c"].values.astype(np.float32)
# Load TimesFM 1.0 (PyTorch)
# NOTE: TimesFM 2.5 PyTorch checkpoint has a file format issue at time of writing.
# The model.safetensors file is not loadable via torch.load().
# Using TimesFM 1.0 PyTorch which works correctly.
print("\n🤖 Loading TimesFM 1.0 (200M) PyTorch...")
import timesfm
# Load TimesFM 2.5 (PyTorch)
print("\n🤖 Loading TimesFM 2.5 (200M) PyTorch...")
hparams = timesfm.TimesFmHparams(horizon_len=12)
checkpoint = timesfm.TimesFmCheckpoint(
huggingface_repo_id="google/timesfm-1.0-200m-pytorch"
model = timesfm.TimesFM_2p5_200M_torch.from_pretrained(
"google/timesfm-2.5-200m-pytorch",
torch_compile=False,
)
model = timesfm.TimesFm(hparams=hparams, checkpoint=checkpoint)
model.compile(timesfm.ForecastConfig(
max_context=512,
max_horizon=12,
normalize_inputs=True,
use_continuous_quantile_head=True,
fix_quantile_crossing=True,
))
# Forecast
print("\n📈 Running forecast (12 months ahead)...")
forecast_input = [input_series]
frequency_input = [0] # Monthly data
point_forecast, experimental_quantile_forecast = model.forecast(
forecast_input,
freq=frequency_input,
horizon=12,
inputs=forecast_input,
)
print(f" Point forecast shape: {point_forecast.shape}")
@@ -65,9 +67,8 @@ print(f" Quantile forecast shape: {experimental_quantile_forecast.shape}")
point = point_forecast[0] # Shape: (horizon,)
quantiles = experimental_quantile_forecast[0] # Shape: (horizon, num_quantiles)
# TimesFM quantiles: [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 0.99]
# Index mapping: 0=10%, 1=20%, ..., 4=50% (median), ..., 9=99%
quantile_labels = ["10%", "20%", "30%", "40%", "50%", "60%", "70%", "80%", "90%", "99%"]
# TimesFM 2.5 columns: 0=mean, 1=10%, 2=20%, ..., 5=50% (median), ..., 9=90%
quantile_labels = ["mean", "10%", "20%", "30%", "40%", "50%", "60%", "70%", "80%", "90%"]
# Create forecast dates (2025 monthly)
last_date = df["date"].max()
@@ -80,16 +81,16 @@ output_df = pd.DataFrame(
{
"date": forecast_dates.strftime("%Y-%m-%d"),
"point_forecast": point,
"q10": quantiles[:, 0],
"q20": quantiles[:, 1],
"q30": quantiles[:, 2],
"q40": quantiles[:, 3],
"q50": quantiles[:, 4], # Median
"q60": quantiles[:, 5],
"q70": quantiles[:, 6],
"q80": quantiles[:, 7],
"q90": quantiles[:, 8],
"q99": quantiles[:, 9],
"mean": quantiles[:, 0],
"q10": quantiles[:, 1],
"q20": quantiles[:, 2],
"q30": quantiles[:, 3],
"q40": quantiles[:, 4],
"q50": quantiles[:, 5], # Median
"q60": quantiles[:, 6],
"q70": quantiles[:, 7],
"q80": quantiles[:, 8],
"q90": quantiles[:, 9],
}
)
@@ -100,7 +101,7 @@ output_df.to_csv(output_dir / "forecast_output.csv", index=False)
# JSON output for the report
output_json = {
"model": "TimesFM 1.0 (200M) PyTorch",
"model": "TimesFM 2.5 (200M) PyTorch",
"input": {
"source": "NOAA GISTEMP Global Temperature Anomaly",
"n_observations": len(df),
@@ -135,24 +136,24 @@ print("=" * 60)
print(
f"\n📅 Forecast period: {forecast_dates[0].strftime('%Y-%m')} to {forecast_dates[-1].strftime('%Y-%m')}"
)
print(f"\n🌡️ Temperature Anomaly Forecast (°C above 1951-1980 baseline):")
print(f"\n {'Month':<10} {'Point':>8} {'80% CI':>15} {'90% CI':>15}")
print("\n🌡️ Temperature Anomaly Forecast (°C above 1951-1980 baseline):")
print(f"\n {'Month':<10} {'Point':>8} {'60% CI':>15} {'80% CI':>15}")
print(f" {'-' * 10} {'-' * 8} {'-' * 15} {'-' * 15}")
for i, (date, pt, q10, q90, q05, q95) in enumerate(
for i, (date, pt, q20, q80, q10, q90) in enumerate(
zip(
forecast_dates.strftime("%Y-%m"),
point,
quantiles[:, 1], # 20%
quantiles[:, 7], # 80%
quantiles[:, 0], # 10%
quantiles[:, 8], # 90%
quantiles[:, 2], # 20%
quantiles[:, 8], # 80%
quantiles[:, 1], # 10%
quantiles[:, 9], # 90%
)
):
print(
f" {date:<10} {pt:>8.3f} [{q10:>6.3f}, {q90:>6.3f}] [{q05:>6.3f}, {q95:>6.3f}]"
f" {date:<10} {pt:>8.3f} [{q20:>6.3f}, {q80:>6.3f}] [{q10:>6.3f}, {q90:>6.3f}]"
)
print(f"\n📊 Summary Statistics:")
print("\n📊 Summary Statistics:")
print(f" Mean forecast: {point.mean():.3f}°C")
print(
f" Max forecast: {point.max():.3f}°C (Month: {forecast_dates[point.argmax()].strftime('%Y-%m')})"
@@ -162,6 +163,6 @@ print(
)
print(f" vs 2024 mean: {point.mean() - df['anomaly_c'].iloc[-12:].mean():+.3f}°C")
print(f"\n✅ Output saved to:")
print("\n✅ Output saved to:")
print(f" {output_dir / 'forecast_output.csv'}")
print(f" {output_dir / 'forecast_output.json'}")
@@ -16,6 +16,8 @@ from __future__ import annotations
import json
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
@@ -57,11 +59,11 @@ def main() -> None:
label="Historical (NOAA GISTEMP)",
)
# Plot 90% CI (outer band)
ax.fill_between(dates, q10, q90, alpha=0.2, color="#dc2626", label="90% CI")
# Plot 80% CI (outer band)
ax.fill_between(dates, q10, q90, alpha=0.2, color="#dc2626", label="80% CI")
# Plot 80% CI (inner band)
ax.fill_between(dates, q20, q80, alpha=0.3, color="#dc2626", label="80% CI")
# Plot 60% CI (inner band)
ax.fill_between(dates, q20, q80, alpha=0.3, color="#dc2626", label="60% CI")
# Plot point forecast
ax.plot(