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/ results/
uv.lock uv.lock
development_setup.md development_setup.md
debug.log
+25 -6
View File
@@ -9,8 +9,10 @@ model developed by Google Research for time-series forecasting.
* All checkpoints: * All checkpoints:
[TimesFM Hugging Face Collection](https://huggingface.co/collections/google/timesfm-release-66e4be5fdb56e960c1e482a6). [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/). * [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): * TimesFM in Google 1P Products:
an official Google product. * [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. 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 - 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 install timesfm==1.3.0` to install an older version of this package to load
them. 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 ## Update - Apr. 9, 2026
Added fine-tuning example using HuggingFace Transformers + PEFT (LoRA) — see Added fine-tuning example using HuggingFace Transformers + PEFT (LoRA) — see
[`timesfm-forecasting/examples/finetuning/`](timesfm-forecasting/examples/finetuning/). [`timesfm-forecasting/examples/finetuning/`](timesfm-forecasting/examples/finetuning/).
Also added unit tests (`tests/`), fixed per-input ridge regression in XReg to Also added unit tests (`tests/`) and incorporated several community fixes.
prevent data leakage, and incorporated several community fixes.
Shoutout to [@kashif](https://github.com/kashif) and [@darkpowerxo](https://github.com/darkpowerxo).
## Update - Mar. 19, 2026 ## 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 ## Update - Oct. 29, 2025
@@ -61,6 +67,19 @@ Since the Sept. 2025 launch, the following improvements have been completed:
### Install ### 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: 1. Clone the repository:
```shell ```shell
git clone https://github.com/google-research/timesfm.git 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] uv pip install -e .[torch]
# Or with flax # Or with flax
uv pip install -e .[flax] uv pip install -e .[flax]
# Or XReg is needed # And when XReg is needed
uv pip install -e .[xreg] 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] [project]
name = "timesfm" name = "timesfm"
version = "2.0.0" version = "2.0.1"
description = "A time series foundation model." description = "A time series foundation model."
authors = [ authors = [
{name = "Rajat Sen", email = "senrajat@google.com"}, {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() 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 @classmethod
def from_pretrained( def from_pretrained(
cls, cls,
@@ -485,10 +500,7 @@ class TimesFM_2p5_200M_flax(timesfm_2p5_base.TimesFM_2p5):
) )
logging.info("Loading checkpoint from: %s", model_file_path) logging.info("Loading checkpoint from: %s", model_file_path)
checkpointer = ocp.StandardCheckpointer() instance.load_checkpoint(model_file_path)
graph, state = nnx.split(instance.model)
state = checkpointer.restore(model_file_path, state)
instance.model = nnx.merge(graph, state)
return instance return instance
def compile( def compile(
+16 -2
View File
@@ -257,7 +257,7 @@ class TimesFM_2p5_200M_torch_module(nn.Module):
to_concat = [t_pf[:, -1, ...]] to_concat = [t_pf[:, -1, ...]]
if t_ar is not None: if t_ar is not None:
to_concat.append(t_ar.reshape(1, -1, self.q)) 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) torch_forecast = torch_forecast.squeeze(0)
outputs.append(torch_forecast.detach().cpu().numpy()) outputs.append(torch_forecast.detach().cpu().numpy())
return outputs return outputs
@@ -283,12 +283,26 @@ class TimesFM_2p5_200M_torch(
self, self,
torch_compile: bool = True, torch_compile: bool = True,
config: Optional[dict] = None, config: Optional[dict] = None,
**kwargs,
): ):
self.model = TimesFM_2p5_200M_torch_module() self.model = TimesFM_2p5_200M_torch_module()
self.torch_compile = torch_compile self.torch_compile = torch_compile
if config is not None: if config is not None:
self._hub_mixin_config = config 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 @classmethod
def _from_pretrained( def _from_pretrained(
cls, cls,
@@ -333,7 +347,7 @@ class TimesFM_2p5_200M_torch(
logging.info("Loading checkpoint from: %s", model_file_path) logging.info("Loading checkpoint from: %s", model_file_path)
# Load the weights into the model. # Load the weights into the model.
instance.model.load_checkpoint( instance.load_checkpoint(
model_file_path, torch_compile=instance.torch_compile model_file_path, torch_compile=instance.torch_compile
) )
return instance return instance
+41 -49
View File
@@ -370,20 +370,11 @@ class BatchedInContextXRegBase:
x_train = np.concatenate(x_train, axis=1) x_train = np.concatenate(x_train, axis=1)
x_test = np.concatenate(x_test, axis=1) x_test = np.concatenate(x_test, axis=1)
# Normalize per-input for robustness (batch-wide normalization # Normalize for robustness.
# would make each input's result depend on batch composition). x_mean = np.mean(x_train, axis=0, keepdims=True)
train_splits = np.cumsum(self.train_lens)[:-1] x_std = np.where((w := np.std(x_train, axis=0, keepdims=True)) > _TOL, w, 1.0)
test_splits = np.cumsum(self.test_lens)[:-1] x_train = [(x_train - x_mean) / x_std]
train_parts = np.split(x_train, train_splits, axis=0) x_test = [(x_test - x_mean) / x_std]
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)]
# Categorical features. Encode one by one. # Categorical features. Encode one by one.
one_hot_encoder = preprocessing.OneHotEncoder( one_hot_encoder = preprocessing.OneHotEncoder(
@@ -472,20 +463,9 @@ class BatchedInContextXRegLinear(BatchedInContextXRegBase):
assert_covariate_shapes=assert_covariate_shapes, assert_covariate_shapes=assert_covariate_shapes,
) )
device = jax.devices("cpu")[0] if force_on_cpu else None x_train = x_train_raw.copy()
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: if max_rows_per_col:
nrows, ncols = x_tr_fit.shape nrows, ncols = x_train.shape
if nrows > (w := ncols * max_rows_per_col): if nrows > (w := ncols * max_rows_per_col):
subsample = jax.random.choice( subsample = jax.random.choice(
jax.random.PRNGKey(max_rows_per_col_sample_seed), jax.random.PRNGKey(max_rows_per_col_sample_seed),
@@ -493,36 +473,48 @@ class BatchedInContextXRegLinear(BatchedInContextXRegBase):
(w,), (w,),
replace=False, replace=False,
) )
x_tr_fit = x_tr_fit[subsample] x_train = x_train[subsample]
y_tr = y_tr[subsample] flat_targets = flat_targets[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)
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 = ( beta_hat = (
jnp.linalg.pinv( jnp.linalg.pinv(
x_tr_j.T @ x_tr_j + ridge * jnp.eye(x_tr_j.shape[1]), x_train.T @ x_train + ridge * jnp.eye(x_train.shape[1]),
hermitian=True, hermitian=True,
) )
@ x_tr_j.T @ x_train.T
@ y_tr_j @ flat_targets
) )
outputs.append(np.array(x_te_j @ beta_hat)[:tel]) y_hat = x_test @ beta_hat
if debug_info: y_hat_context = x_train_raw @ beta_hat if debug_info else None
outputs_context.append(np.array(x_tr_raw_j @ beta_hat)[:trl])
train_idx += trl outputs = []
test_idx += tel outputs_context = []
# 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)])
)
train_index += train_index_delta
test_index += test_index_delta
if debug_info: if debug_info:
return ( return outputs, outputs_context, flat_targets, x_train, x_test
outputs,
outputs_context,
_to_padded_jax_array(flat_targets),
_to_padded_jax_array(x_train_raw),
_to_padded_jax_array(x_test),
)
else: else:
return outputs 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 date,point_forecast,mean,q10,q20,q30,q40,q50,q60,q70,q80,q90
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-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.2856668,1.2773758,1.1406044,1.1960833,1.2322671,1.2593892,1.2856668,1.3110137,1.3400218,1.3751202,1.4253658 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.2950127,1.2869918,1.126852,1.1876173,1.234988,1.2675052,1.2950127,1.328448,1.354729,1.4035482,1.4642649 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.2207624,1.2084007,1.0352504,1.1041918,1.151865,1.1853008,1.2207624,1.256663,1.2898555,1.3310349,1.4016538 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.1702554,1.153313,0.9691495,1.0431063,1.0932612,1.1276176,1.1702554,1.201966,1.2390311,1.2891905,1.3632389 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.1455553,1.1275499,0.94203794,1.0110554,1.0658777,1.1061188,1.1455553,1.1806211,1.2180579,1.2702757,1.345366 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.1702348,1.1510556,0.9503718,1.0347577,1.0847733,1.1287677,1.1702348,1.2114835,1.2482276,1.2997853,1.3807325 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.2026825,1.1859496,0.9709255,1.0594383,1.1106675,1.1579902,1.2026825,1.2399211,1.2842004,1.3408126,1.419526 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.1909748,1.1784849,0.95943713,1.0403702,1.103606,1.1511956,1.1909748,1.2390201,1.2832941,1.3354731,1.416972 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.1490841,1.1264795,0.9079477,0.99529266,1.0548235,1.1052223,1.1490841,1.1897774,1.240414,1.2868769,1.3775467 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.0804785,1.0624356,0.8361266,0.9259792,0.9882403,1.0386353,1.0804785,1.1281581,1.1759715,1.228377,1.3122478 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.0613453,1.0366092,0.80220693,0.89521873,0.9593707,1.0152239,1.0613453,1.1032857,1.15315,1.216908,1.2959521 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": { "input": {
"source": "NOAA GISTEMP Global Temperature Anomaly", "source": "NOAA GISTEMP Global Temperature Anomaly",
"n_observations": 36, "n_observations": 36,
@@ -23,166 +23,166 @@
"2025-12" "2025-12"
], ],
"point": [ "point": [
1.25933837890625, 1.2223774194717407,
1.285666823387146, 1.2563583850860596,
1.2950127124786377, 1.286476969718933,
1.2207623720169067, 1.240488052368164,
1.170255422592163, 1.202637791633606,
1.1455552577972412, 1.2100199460983276,
1.1702347993850708, 1.2253109216690063,
1.2026824951171875, 1.2421810626983643,
1.1909748315811157, 1.2697349786758423,
1.1490840911865234, 1.2496669292449951,
1.080478549003601, 1.2135266065597534,
1.0613453388214111 1.2034140825271606
], ],
"quantiles": { "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%": [ "10%": [
1.2481880187988281, 1.1230627298355103,
1.2773758172988892, 1.1482248306274414,
1.286991834640503, 1.1694773435592651,
1.2084007263183594, 1.119297981262207,
1.1533130407333374, 1.0776878595352173,
1.1275498867034912, 1.0811386108398438,
1.1510555744171143, 1.0917633771896362,
1.1859495639801025, 1.1043028831481934,
1.1784849166870117, 1.123927354812622,
1.1264795064926147, 1.0962445735931396,
1.0624356269836426, 1.054552435874939,
1.036609172821045 1.0412306785583496
], ],
"20%": [ "20%": [
1.1407020092010498, 1.161399483680725,
1.1406043767929077, 1.1891648769378662,
1.126852035522461, 1.2141375541687012,
1.0352504253387451, 1.168922781944275,
0.9691494703292847, 1.1280149221420288,
0.9420379400253296, 1.1352498531341553,
0.9503718018531799, 1.1474796533584595,
0.970925509929657, 1.1598811149597168,
0.9594371318817139, 1.1872942447662354,
0.9079477190971375, 1.1630417108535767,
0.8361266255378723, 1.1223580837249756,
0.8022069334983826 1.1113251447677612
], ],
"30%": [ "30%": [
1.1880751848220825, 1.18324875831604,
1.1960833072662354, 1.2134028673171997,
1.187617301940918, 1.2430503368377686,
1.104191780090332, 1.195752739906311,
1.0431063175201416, 1.1554758548736572,
1.01105535030365, 1.1687461137771606,
1.0347577333450317, 1.1764130592346191,
1.0594383478164673, 1.1929150819778442,
1.040370225906372, 1.2162872552871704,
0.9952926635742188, 1.193880558013916,
0.9259791970252991, 1.1572109460830688,
0.8952187299728394 1.1433371305465698
], ],
"40%": [ "40%": [
1.2137157917022705, 1.2030285596847534,
1.232267141342163, 1.2355917692184448,
1.2349879741668701, 1.263671875,
1.151865005493164, 1.216238260269165,
1.0932612419128418, 1.1801764965057373,
1.0658776760101318, 1.1853771209716797,
1.084773302078247, 1.2029412984848022,
1.1106674671173096, 1.216418743133545,
1.1036059856414795, 1.2429510354995728,
1.0548235177993774, 1.2217119932174683,
0.9882403016090393, 1.1864622831344604,
0.9593706727027893 1.1752325296401978
], ],
"50%": [ "50%": [
1.2394564151763916, 1.2223774194717407,
1.2593891620635986, 1.2563583850860596,
1.267505168914795, 1.286476969718933,
1.1853008270263672, 1.240488052368164,
1.127617597579956, 1.202637791633606,
1.1061187982559204, 1.2100199460983276,
1.128767728805542, 1.2253109216690063,
1.1579902172088623, 1.2421810626983643,
1.1511956453323364, 1.2697349786758423,
1.1052223443984985, 1.2496669292449951,
1.03863525390625, 1.2135266065597534,
1.0152238607406616 1.2034140825271606
], ],
"60%": [ "60%": [
1.25933837890625, 1.2409889698028564,
1.285666823387146, 1.2787916660308838,
1.2950127124786377, 1.3101407289505005,
1.2207623720169067, 1.2653590440750122,
1.170255422592163, 1.2274072170257568,
1.1455552577972412, 1.2333428859710693,
1.1702347993850708, 1.2509560585021973,
1.2026824951171875, 1.266998052597046,
1.1909748315811157, 1.2976493835449219,
1.1490840911865234, 1.2757387161254883,
1.080478549003601, 1.2447274923324585,
1.0613453388214111 1.2308053970336914
], ],
"70%": [ "70%": [
1.27677321434021, 1.2626991271972656,
1.3110136985778809, 1.3061506748199463,
1.3284480571746826, 1.336512565612793,
1.2566629648208618, 1.2880773544311523,
1.2019660472869873, 1.2546851634979248,
1.1806211471557617, 1.2559939622879028,
1.2114834785461426, 1.276485562324524,
1.2399210929870605, 1.292467474937439,
1.2390201091766357, 1.323679804801941,
1.1897773742675781, 1.30413818359375,
1.1281580924987793, 1.2706527709960938,
1.1032856702804565 1.2621957063674927
], ],
"80%": [ "80%": [
1.2971320152282715, 1.2927711009979248,
1.3400218486785889, 1.3358021974563599,
1.3547290563583374, 1.372582197189331,
1.2898554801940918, 1.3237736225128174,
1.2390310764312744, 1.2890899181365967,
1.2180578708648682, 1.2938039302825928,
1.248227596282959, 1.3113454580307007,
1.2842004299163818, 1.3304922580718994,
1.2832940816879272, 1.358002781867981,
1.240414023399353, 1.3376691341400146,
1.175971508026123, 1.3091387748718262,
1.153149962425232 1.2912495136260986
], ],
"90%": [ "90%": [
1.3239599466323853, 1.3396029472351074,
1.3751201629638672, 1.3880282640457153,
1.403548240661621, 1.426570177078247,
1.3310348987579346, 1.3806401491165161,
1.2891905307769775, 1.3469674587249756,
1.2702757120132446, 1.3532465696334839,
1.2997852563858032, 1.3728349208831787,
1.3408125638961792, 1.394974946975708,
1.3354730606079102, 1.425218105316162,
1.286876916885376, 1.409816026687622,
1.2283769845962524, 1.380167841911316,
1.2169079780578613 1.3700779676437378
],
"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
] ]
} }
}, },
"summary": { "summary": {
"forecast_mean_c": 1.186, "forecast_mean_c": 1.235,
"forecast_max_c": 1.295, "forecast_max_c": 1.286,
"forecast_min_c": 1.061, "forecast_min_c": 1.203,
"vs_last_year_mean": -0.067 "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 numpy as np
import pandas as pd import pandas as pd
import timesfm
# Preflight check # Preflight check
print("=" * 60) print("=" * 60)
@@ -35,27 +36,28 @@ print(
# TimesFM expects a list of 1D numpy arrays # TimesFM expects a list of 1D numpy arrays
input_series = df["anomaly_c"].values.astype(np.float32) input_series = df["anomaly_c"].values.astype(np.float32)
# Load TimesFM 1.0 (PyTorch) # Load TimesFM 2.5 (PyTorch)
# NOTE: TimesFM 2.5 PyTorch checkpoint has a file format issue at time of writing. print("\n🤖 Loading TimesFM 2.5 (200M) PyTorch...")
# 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
hparams = timesfm.TimesFmHparams(horizon_len=12) model = timesfm.TimesFM_2p5_200M_torch.from_pretrained(
checkpoint = timesfm.TimesFmCheckpoint( "google/timesfm-2.5-200m-pytorch",
huggingface_repo_id="google/timesfm-1.0-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 # Forecast
print("\n📈 Running forecast (12 months ahead)...") print("\n📈 Running forecast (12 months ahead)...")
forecast_input = [input_series] forecast_input = [input_series]
frequency_input = [0] # Monthly data
point_forecast, experimental_quantile_forecast = model.forecast( point_forecast, experimental_quantile_forecast = model.forecast(
forecast_input, horizon=12,
freq=frequency_input, inputs=forecast_input,
) )
print(f" Point forecast shape: {point_forecast.shape}") 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,) point = point_forecast[0] # Shape: (horizon,)
quantiles = experimental_quantile_forecast[0] # Shape: (horizon, num_quantiles) 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] # TimesFM 2.5 columns: 0=mean, 1=10%, 2=20%, ..., 5=50% (median), ..., 9=90%
# Index mapping: 0=10%, 1=20%, ..., 4=50% (median), ..., 9=99% quantile_labels = ["mean", "10%", "20%", "30%", "40%", "50%", "60%", "70%", "80%", "90%"]
quantile_labels = ["10%", "20%", "30%", "40%", "50%", "60%", "70%", "80%", "90%", "99%"]
# Create forecast dates (2025 monthly) # Create forecast dates (2025 monthly)
last_date = df["date"].max() last_date = df["date"].max()
@@ -80,16 +81,16 @@ output_df = pd.DataFrame(
{ {
"date": forecast_dates.strftime("%Y-%m-%d"), "date": forecast_dates.strftime("%Y-%m-%d"),
"point_forecast": point, "point_forecast": point,
"q10": quantiles[:, 0], "mean": quantiles[:, 0],
"q20": quantiles[:, 1], "q10": quantiles[:, 1],
"q30": quantiles[:, 2], "q20": quantiles[:, 2],
"q40": quantiles[:, 3], "q30": quantiles[:, 3],
"q50": quantiles[:, 4], # Median "q40": quantiles[:, 4],
"q60": quantiles[:, 5], "q50": quantiles[:, 5], # Median
"q70": quantiles[:, 6], "q60": quantiles[:, 6],
"q80": quantiles[:, 7], "q70": quantiles[:, 7],
"q90": quantiles[:, 8], "q80": quantiles[:, 8],
"q99": quantiles[:, 9], "q90": quantiles[:, 9],
} }
) )
@@ -100,7 +101,7 @@ output_df.to_csv(output_dir / "forecast_output.csv", index=False)
# JSON output for the report # JSON output for the report
output_json = { output_json = {
"model": "TimesFM 1.0 (200M) PyTorch", "model": "TimesFM 2.5 (200M) PyTorch",
"input": { "input": {
"source": "NOAA GISTEMP Global Temperature Anomaly", "source": "NOAA GISTEMP Global Temperature Anomaly",
"n_observations": len(df), "n_observations": len(df),
@@ -135,24 +136,24 @@ print("=" * 60)
print( print(
f"\n📅 Forecast period: {forecast_dates[0].strftime('%Y-%m')} to {forecast_dates[-1].strftime('%Y-%m')}" 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("\n🌡️ Temperature Anomaly Forecast (°C above 1951-1980 baseline):")
print(f"\n {'Month':<10} {'Point':>8} {'80% CI':>15} {'90% CI':>15}") print(f"\n {'Month':<10} {'Point':>8} {'60% CI':>15} {'80% CI':>15}")
print(f" {'-' * 10} {'-' * 8} {'-' * 15} {'-' * 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( zip(
forecast_dates.strftime("%Y-%m"), forecast_dates.strftime("%Y-%m"),
point, point,
quantiles[:, 1], # 20% quantiles[:, 2], # 20%
quantiles[:, 7], # 80% quantiles[:, 8], # 80%
quantiles[:, 0], # 10% quantiles[:, 1], # 10%
quantiles[:, 8], # 90% quantiles[:, 9], # 90%
) )
): ):
print( 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" Mean forecast: {point.mean():.3f}°C")
print( print(
f" Max forecast: {point.max():.3f}°C (Month: {forecast_dates[point.argmax()].strftime('%Y-%m')})" 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" 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.csv'}")
print(f" {output_dir / 'forecast_output.json'}") print(f" {output_dir / 'forecast_output.json'}")
@@ -16,6 +16,8 @@ from __future__ import annotations
import json import json
from pathlib import Path from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt import matplotlib.pyplot as plt
import numpy as np import numpy as np
import pandas as pd import pandas as pd
@@ -57,11 +59,11 @@ def main() -> None:
label="Historical (NOAA GISTEMP)", label="Historical (NOAA GISTEMP)",
) )
# Plot 90% CI (outer band) # Plot 80% CI (outer band)
ax.fill_between(dates, q10, q90, alpha=0.2, color="#dc2626", label="90% CI") ax.fill_between(dates, q10, q90, alpha=0.2, color="#dc2626", label="80% CI")
# Plot 80% CI (inner band) # Plot 60% CI (inner band)
ax.fill_between(dates, q20, q80, alpha=0.3, color="#dc2626", label="80% CI") ax.fill_between(dates, q20, q80, alpha=0.3, color="#dc2626", label="60% CI")
# Plot point forecast # Plot point forecast
ax.plot( ax.plot(