Quantiles PR nit fix

This commit is contained in:
misha-chertushkin
2025-02-05 13:41:13 +00:00
parent 402ebf52a5
commit 67e0eaaa3d
2 changed files with 4 additions and 4 deletions
+2 -2
View File
@@ -13,7 +13,7 @@ import torch.distributed as dist
import torch.nn as nn
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, Dataset
from timesfm.pytorch_patched_decoder import _create_quantiles
from timesfm.pytorch_patched_decoder import create_quantiles
import wandb
@@ -267,7 +267,7 @@ class TimesFMFinetuner:
loss = self.loss_fn(last_patch_pred, x_future.squeeze(-1))
if self.config.use_quantile_loss:
quantiles = self.config.quantiles or _create_quantiles()
quantiles = self.config.quantiles or create_quantiles()
for i, quantile in enumerate(quantiles):
last_patch_quantile = predictions[:, -1, :, i + 1]
loss += torch.mean(
+2 -2
View File
@@ -21,7 +21,7 @@ from torch import nn
import torch.nn.functional as F
def _create_quantiles() -> list[float]:
def create_quantiles() -> list[float]:
return [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9]
@@ -48,7 +48,7 @@ class TimesFMConfig:
# Horizon length
horizon_len: int = 128
# quantiles
quantiles: List[float] = dataclasses.field(default_factory=_create_quantiles)
quantiles: List[float] = dataclasses.field(default_factory=create_quantiles)
# Padding value
pad_val: float = 1123581321.0
# Tolerance