|
215 | 215 | "outputs": [], |
216 | 216 | "source": [ |
217 | 217 | "from torch.utils.data import DataLoader\n", |
| 218 | + "\n", |
218 | 219 | "from s4casting.data.utils import collate_single_interval\n", |
219 | 220 | "\n", |
220 | 221 | "N = len(batcher.train)\n", |
221 | 222 | "print(f\"The training dataset has {N} samples\")\n", |
222 | 223 | "first, _ = next(\n", |
223 | 224 | " iter(\n", |
224 | | - " DataLoader(batcher.train, batch_size=1, num_workers=0, shuffle=True, pin_memory=False, persistent_workers=False, collate_fn=collate_single_interval)\n", |
| 225 | + " DataLoader(\n", |
| 226 | + " batcher.train,\n", |
| 227 | + " batch_size=1,\n", |
| 228 | + " num_workers=0,\n", |
| 229 | + " shuffle=True,\n", |
| 230 | + " pin_memory=False,\n", |
| 231 | + " persistent_workers=False,\n", |
| 232 | + " collate_fn=collate_single_interval,\n", |
| 233 | + " )\n", |
225 | 234 | " )\n", |
226 | 235 | ")\n", |
227 | 236 | "X, Xm, Y, Ym = (t for t in first)\n", |
|
284 | 293 | "import numpy as np\n", |
285 | 294 | "\n", |
286 | 295 | "predict_width_samples = (config.model.predict_width * 24 * 60) // config.model.base_sample_interval_minutes\n", |
287 | | - "input_width_samples = ((config.model.context_window[0] * 24 * 60) // config.model.base_sample_interval_minutes) - predict_width_samples\n", |
| 296 | + "input_width_samples = (\n", |
| 297 | + " (config.model.context_window[0] * 24 * 60) // config.model.base_sample_interval_minutes\n", |
| 298 | + ") - predict_width_samples\n", |
288 | 299 | "days_per_sample = config.model.context_window[0]\n", |
289 | 300 | "samples_per_day = 96 # 24 hours x 4 samples/hour\n", |
290 | 301 | "\n", |
|
369 | 380 | " Xm = Xm.to(context.machine.torch_device).float()\n", |
370 | 381 | " Ym = Ym.to(context.machine.torch_device).float()\n", |
371 | 382 | "\n", |
372 | | - " input_interval = batch_cfg.sample_interval_minutes.to(context.machine.torch_device)\n", |
| 383 | + " input_interval = batch_cfg.sample_interval_minutes.to(context.machine.torch_device)\n", |
373 | 384 | " output_interval = select_rate(input_interval, config.model.output_sample_intervals_minutes)\n", |
374 | 385 | "\n", |
375 | 386 | " # input data to model\n", |
376 | | - " _, loss = context.model_container.model(\n", |
377 | | - " X, Xm, input_interval, output_interval, Y, Ym\n", |
378 | | - " )\n", |
| 387 | + " _, loss = context.model_container.model(X, Xm, input_interval, output_interval, Y, Ym)\n", |
379 | 388 | " loss.backward()\n", |
380 | 389 | " total_loss += loss.item()\n", |
381 | 390 | "\n", |
|
0 commit comments