Update examples with new feedback
This commit is contained in:
@@ -44,18 +44,27 @@
|
||||
"class TimeSeriesDataset(Dataset):\n",
|
||||
" \"\"\"Dataset for time series data compatible with TimesFM.\"\"\"\n",
|
||||
"\n",
|
||||
" def __init__(self, series: np.ndarray, context_length: int, horizon_length: int):\n",
|
||||
" def __init__(self,\n",
|
||||
" series: np.ndarray,\n",
|
||||
" context_length: int,\n",
|
||||
" horizon_length: int,\n",
|
||||
" freq_type: int = 0):\n",
|
||||
" \"\"\"\n",
|
||||
" Initialize dataset.\n",
|
||||
" Initialize dataset.\n",
|
||||
"\n",
|
||||
" Args:\n",
|
||||
" series: Time series data\n",
|
||||
" context_length: Number of past timesteps to use as input\n",
|
||||
" horizon_length: Number of future timesteps to predict\n",
|
||||
" freq_type: Frequency type (0, 1, or 2)\n",
|
||||
" \"\"\"\n",
|
||||
" if freq_type not in [0, 1, 2]:\n",
|
||||
" raise ValueError(\"freq_type must be 0, 1, or 2\")\n",
|
||||
"\n",
|
||||
" Args:\n",
|
||||
" series: Time series data\n",
|
||||
" context_length: Number of past timesteps to use as input\n",
|
||||
" horizon_length: Number of future timesteps to predict\n",
|
||||
" \"\"\"\n",
|
||||
" self.series = series\n",
|
||||
" self.context_length = context_length\n",
|
||||
" self.horizon_length = horizon_length\n",
|
||||
" self.freq_type = freq_type\n",
|
||||
" self._prepare_samples()\n",
|
||||
"\n",
|
||||
" def _prepare_samples(self) -> None:\n",
|
||||
@@ -66,47 +75,57 @@
|
||||
" for start_idx in range(0, len(self.series) - total_length + 1):\n",
|
||||
" end_idx = start_idx + self.context_length\n",
|
||||
" x_context = self.series[start_idx:end_idx]\n",
|
||||
" x_future = self.series[end_idx : end_idx + self.horizon_length]\n",
|
||||
" x_future = self.series[end_idx:end_idx + self.horizon_length]\n",
|
||||
" self.samples.append((x_context, x_future))\n",
|
||||
"\n",
|
||||
" def __len__(self) -> int:\n",
|
||||
" return len(self.samples)\n",
|
||||
"\n",
|
||||
" def __getitem__(self, index: int) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:\n",
|
||||
" def __getitem__(\n",
|
||||
" self, index: int\n",
|
||||
" ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:\n",
|
||||
" x_context, x_future = self.samples[index]\n",
|
||||
"\n",
|
||||
" x_context = torch.tensor(x_context, dtype=torch.float32)\n",
|
||||
" x_future = torch.tensor(x_future, dtype=torch.float32)\n",
|
||||
"\n",
|
||||
" input_padding = torch.zeros_like(x_context)\n",
|
||||
" freq = torch.zeros(1, dtype=torch.long)\n",
|
||||
" freq = torch.tensor([self.freq_type], dtype=torch.long)\n",
|
||||
"\n",
|
||||
" return x_context, input_padding, freq, x_future\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def prepare_datasets(\n",
|
||||
" series: np.ndarray, context_length: int, horizon_length: int, train_split: float = 0.8\n",
|
||||
") -> Tuple[Dataset, Dataset]:\n",
|
||||
"def prepare_datasets(series: np.ndarray,\n",
|
||||
" context_length: int,\n",
|
||||
" horizon_length: int,\n",
|
||||
" freq_type: int = 0,\n",
|
||||
" train_split: float = 0.8) -> Tuple[Dataset, Dataset]:\n",
|
||||
" \"\"\"\n",
|
||||
" Prepare training and validation datasets from time series data.\n",
|
||||
" Prepare training and validation datasets from time series data.\n",
|
||||
"\n",
|
||||
" Args:\n",
|
||||
" series: Input time series data\n",
|
||||
" context_length: Number of past timesteps to use\n",
|
||||
" horizon_length: Number of future timesteps to predict\n",
|
||||
" train_split: Fraction of data to use for training\n",
|
||||
" Args:\n",
|
||||
" series: Input time series data\n",
|
||||
" context_length: Number of past timesteps to use\n",
|
||||
" horizon_length: Number of future timesteps to predict\n",
|
||||
" freq_type: Frequency type (0, 1, or 2)\n",
|
||||
" train_split: Fraction of data to use for training\n",
|
||||
"\n",
|
||||
" Returns:\n",
|
||||
" Tuple of (train_dataset, val_dataset)\n",
|
||||
" \"\"\"\n",
|
||||
" Returns:\n",
|
||||
" Tuple of (train_dataset, val_dataset)\n",
|
||||
" \"\"\"\n",
|
||||
" train_size = int(len(series) * train_split)\n",
|
||||
" train_data = series[:train_size]\n",
|
||||
" val_data = series[train_size:]\n",
|
||||
"\n",
|
||||
" # Create datasets\n",
|
||||
" train_dataset = TimeSeriesDataset(train_data, context_length=context_length, horizon_length=horizon_length)\n",
|
||||
" # Create datasets with specified frequency type\n",
|
||||
" train_dataset = TimeSeriesDataset(train_data,\n",
|
||||
" context_length=context_length,\n",
|
||||
" horizon_length=horizon_length,\n",
|
||||
" freq_type=freq_type)\n",
|
||||
"\n",
|
||||
" val_dataset = TimeSeriesDataset(val_data, context_length=context_length, horizon_length=horizon_length)\n",
|
||||
" val_dataset = TimeSeriesDataset(val_data,\n",
|
||||
" context_length=context_length,\n",
|
||||
" horizon_length=horizon_length,\n",
|
||||
" freq_type=freq_type)\n",
|
||||
"\n",
|
||||
" return train_dataset, val_dataset\n"
|
||||
]
|
||||
@@ -128,14 +147,16 @@
|
||||
" device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n",
|
||||
" repo_id = \"google/timesfm-2.0-500m-pytorch\"\n",
|
||||
" hparams = TimesFmHparams(\n",
|
||||
" backend=device,\n",
|
||||
" per_core_batch_size=32,\n",
|
||||
" horizon_len=128,\n",
|
||||
" num_layers=50,\n",
|
||||
" use_positional_embedding=False,\n",
|
||||
" context_len=192,\n",
|
||||
" backend=device,\n",
|
||||
" per_core_batch_size=32,\n",
|
||||
" horizon_len=128,\n",
|
||||
" num_layers=50,\n",
|
||||
" use_positional_embedding=False,\n",
|
||||
" context_len=\n",
|
||||
" 192, # Context length can be anything up to 2048 in multiples of 32\n",
|
||||
" )\n",
|
||||
" tfm = TimesFm(hparams=hparams, checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id))\n",
|
||||
" tfm = TimesFm(hparams=hparams,\n",
|
||||
" checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id))\n",
|
||||
"\n",
|
||||
" model = PatchedTimeSeriesDecoder(tfm._model_config)\n",
|
||||
" if load_weights:\n",
|
||||
@@ -152,18 +173,18 @@
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def plot_predictions(\n",
|
||||
" model: TimesFm,\n",
|
||||
" val_dataset: Dataset,\n",
|
||||
" save_path: Optional[str] = \"predictions.png\",\n",
|
||||
" model: TimesFm,\n",
|
||||
" val_dataset: Dataset,\n",
|
||||
" save_path: Optional[str] = \"predictions.png\",\n",
|
||||
") -> None:\n",
|
||||
" \"\"\"\n",
|
||||
" Plot model predictions against ground truth for a batch of validation data.\n",
|
||||
" Plot model predictions against ground truth for a batch of validation data.\n",
|
||||
"\n",
|
||||
" Args:\n",
|
||||
" model: Trained TimesFM model\n",
|
||||
" val_dataset: Validation dataset\n",
|
||||
" save_path: Path to save the plot\n",
|
||||
" \"\"\"\n",
|
||||
" Args:\n",
|
||||
" model: Trained TimesFM model\n",
|
||||
" val_dataset: Validation dataset\n",
|
||||
" save_path: Path to save the plot\n",
|
||||
" \"\"\"\n",
|
||||
" import matplotlib.pyplot as plt\n",
|
||||
"\n",
|
||||
" model.eval()\n",
|
||||
@@ -194,18 +215,26 @@
|
||||
"\n",
|
||||
" plt.figure(figsize=(12, 6))\n",
|
||||
"\n",
|
||||
" plt.plot(range(context_len), context_vals, label=\"Historical Data\", color=\"blue\", linewidth=2)\n",
|
||||
" plt.plot(range(context_len),\n",
|
||||
" context_vals,\n",
|
||||
" label=\"Historical Data\",\n",
|
||||
" color=\"blue\",\n",
|
||||
" linewidth=2)\n",
|
||||
"\n",
|
||||
" plt.plot(\n",
|
||||
" range(context_len, context_len + horizon_len),\n",
|
||||
" future_vals,\n",
|
||||
" label=\"Ground Truth\",\n",
|
||||
" color=\"green\",\n",
|
||||
" linestyle=\"--\",\n",
|
||||
" linewidth=2,\n",
|
||||
" range(context_len, context_len + horizon_len),\n",
|
||||
" future_vals,\n",
|
||||
" label=\"Ground Truth\",\n",
|
||||
" color=\"green\",\n",
|
||||
" linestyle=\"--\",\n",
|
||||
" linewidth=2,\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" plt.plot(range(context_len, context_len + horizon_len), pred_vals, label=\"Prediction\", color=\"red\", linewidth=2)\n",
|
||||
" plt.plot(range(context_len, context_len + horizon_len),\n",
|
||||
" pred_vals,\n",
|
||||
" label=\"Prediction\",\n",
|
||||
" color=\"red\",\n",
|
||||
" linewidth=2)\n",
|
||||
"\n",
|
||||
" plt.xlabel(\"Time Step\")\n",
|
||||
" plt.ylabel(\"Value\")\n",
|
||||
@@ -218,43 +247,44 @@
|
||||
" print(f\"Plot saved to {save_path}\")\n",
|
||||
"\n",
|
||||
" plt.close()\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def get_data(context_len: int, horizon_len: int) -> Tuple[Dataset, Dataset]:\n",
|
||||
" df = yf.download(\"AAPL\", start=\"2010-01-01\", end=\"2019-01-01\")\n",
|
||||
" time_series = df[\"Close\"].values\n",
|
||||
"\n",
|
||||
" train_dataset, val_dataset = prepare_datasets(\n",
|
||||
" series=time_series,\n",
|
||||
" context_length=context_len,\n",
|
||||
" horizon_length=horizon_len,\n",
|
||||
" train_split=0.8,\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" print(f\"Created datasets:\")\n",
|
||||
" print(f\"- Training samples: {len(train_dataset)}\")\n",
|
||||
" print(f\"- Validation samples: {len(val_dataset)}\")\n",
|
||||
" return train_dataset, val_dataset\n",
|
||||
"\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"\n",
|
||||
"def single_gpu_example():\n",
|
||||
" \"\"\"Basic example of finetuning TimesFM on stock data.\"\"\"\n",
|
||||
" model, hparams, tfm_config = get_model(load_weights=True)\n",
|
||||
" config = FinetuningConfig(batch_size=256, num_epochs=5, learning_rate=1e-4, use_wandb=True)\n",
|
||||
" config = FinetuningConfig(batch_size=256,\n",
|
||||
" num_epochs=5,\n",
|
||||
" learning_rate=1e-4,\n",
|
||||
" use_wandb=True,\n",
|
||||
" freq_type=1,\n",
|
||||
" log_every_n_steps=10,\n",
|
||||
" val_check_interval=0.5,\n",
|
||||
" use_quantile_loss=True)\n",
|
||||
"\n",
|
||||
" train_dataset, val_dataset = get_data(128, tfm_config.horizon_len)\n",
|
||||
" train_dataset, val_dataset = get_data(128,\n",
|
||||
" tfm_config.horizon_len,\n",
|
||||
" freq_type=config.freq_type)\n",
|
||||
" finetuner = TimesFMFinetuner(model, config)\n",
|
||||
"\n",
|
||||
" print(\"\\nStarting finetuning...\")\n",
|
||||
" results = finetuner.finetune(train_dataset=train_dataset, val_dataset=val_dataset)\n",
|
||||
" results = finetuner.finetune(train_dataset=train_dataset,\n",
|
||||
" val_dataset=val_dataset)\n",
|
||||
"\n",
|
||||
" print(\"\\nFinetuning completed!\")\n",
|
||||
" print(f\"Training history: {len(results['history']['train_loss'])} epochs\")\n",
|
||||
"\n",
|
||||
" plot_predictions(\n",
|
||||
" model=model,\n",
|
||||
" val_dataset=val_dataset,\n",
|
||||
" save_path=\"timesfm_predictions.png\",\n",
|
||||
" model=model,\n",
|
||||
" val_dataset=val_dataset,\n",
|
||||
" save_path=\"timesfm_predictions.png\",\n",
|
||||
" )\n"
|
||||
]
|
||||
},
|
||||
|
||||
Reference in New Issue
Block a user