|
|
| import torch
|
| import torch.nn as nn
|
|
|
| __all__ = [
|
| "DynamicBetaSmoothL1Loss",
|
| "DynamicBetaSmoothL1LossMultiStep",
|
| "DynamicMultiStepHLoss",
|
| "DynamicMultiStepQLoss",
|
| ]
|
|
|
| import torch
|
| import torch.nn as nn
|
| from scipy.ndimage import gaussian_filter1d
|
|
|
| class DynamicMultiStepHLoss(nn.Module):
|
| def __init__(self, init_beta=1.0, alpha=0.01,
|
| horizon_decay=0.9, reduction='mean',
|
| min_beta=0.1, update_freq=50,
|
| smooth_sigma=0.0):
|
| """
|
| Dynamic Multi-step Huber Loss with horizon decay + optional Gaussian smoothing
|
|
|
| Args:
|
| init_beta: initial huber beta
|
| alpha: update rate for beta
|
| horizon_decay: decay factor for weighting steps
|
| reduction: 'mean', 'sum', 'none'
|
| min_beta: lower bound for beta
|
| update_freq: frequency (steps) to update beta
|
| smooth_sigma: if > 0, apply gaussian_filter1d to errors along horizon
|
| """
|
| super().__init__()
|
| self.register_buffer('beta', torch.tensor(float(init_beta)))
|
| self.alpha = alpha
|
| self.horizon_decay = horizon_decay
|
| self.reduction = reduction
|
| self.min_beta = min_beta
|
| self.update_freq = update_freq
|
| self.register_buffer('step_count', torch.tensor(0))
|
| self.smooth_sigma = smooth_sigma
|
|
|
| def _align_shapes(self, pred, true):
|
| """
|
| Align shapes for pred and true.
|
| 1. Truncate horizon to min(pred, true)
|
| 2. Match batch size by slicing/repeating target if needed
|
| 3. Handle extra dimensions
|
| """
|
|
|
| if pred.dim() == 3 and pred.shape[2] == 1:
|
| pred = pred.squeeze(-1)
|
| if true.dim() == 3 and true.shape[2] == 1:
|
| true = true.squeeze(-1)
|
| if pred.dim() == 3 and pred.shape[1] == 1:
|
| pred = pred.squeeze(1)
|
| if true.dim() == 3 and true.shape[1] == 1:
|
| true = true.squeeze(1)
|
|
|
|
|
| if pred.dim() != 2:
|
| pred = pred.view(pred.shape[0], -1)
|
| if true.dim() != 2:
|
| true = true.view(true.shape[0], -1)
|
|
|
| Bp, Hp = pred.shape
|
| Bt, Ht = true.shape
|
|
|
|
|
| H = min(Hp, Ht)
|
| pred = pred[:, :H]
|
| true = true[:, :H]
|
|
|
|
|
| if Bt < Bp:
|
|
|
| reps = (Bp + Bt - 1) // Bt
|
| true = true.repeat(reps, 1)[:Bp, :]
|
| elif Bt > Bp:
|
|
|
| true = true[:Bp, :]
|
|
|
| return pred, true
|
|
|
|
|
| def _apply_smoothing(self, diff):
|
| """Apply gaussian_filter1d along horizon axis (dim=1)"""
|
| if self.smooth_sigma and self.smooth_sigma > 0:
|
| diff_np = diff.detach().cpu().numpy()
|
| smoothed = gaussian_filter1d(diff_np, sigma=self.smooth_sigma, axis=1)
|
| return torch.tensor(smoothed, dtype=diff.dtype, device=diff.device)
|
| return diff
|
|
|
| def forward(self, pred, true):
|
| self.step_count += 1
|
| pred, true = self._align_shapes(pred, true)
|
|
|
|
|
| diff = torch.abs(pred - true)
|
|
|
|
|
| diff = self._apply_smoothing(diff)
|
|
|
| B, H = diff.shape
|
|
|
|
|
| with torch.no_grad():
|
| step_idx = torch.arange(H, device=diff.device, dtype=torch.float32)
|
| weights = torch.pow(self.horizon_decay, step_idx)
|
| weights = weights / weights.sum()
|
| weights = weights.unsqueeze(0)
|
|
|
|
|
| beta = max(self.beta.item(), self.min_beta)
|
| loss = torch.where(
|
| diff < beta,
|
| 0.5 * (diff ** 2) / beta,
|
| diff - 0.5 * beta
|
| )
|
|
|
|
|
| weighted_loss = loss * weights
|
|
|
|
|
| if self.step_count % self.update_freq == 0:
|
| with torch.no_grad():
|
| median_error = torch.median(diff).item()
|
| new_beta = (1 - self.alpha) * self.beta.item() + self.alpha * median_error
|
| self.beta.fill_(max(new_beta, self.min_beta))
|
|
|
| if self.reduction == 'mean':
|
| return weighted_loss.mean()
|
| elif self.reduction == 'sum':
|
| return weighted_loss.sum()
|
| elif self.reduction == 'none':
|
| return weighted_loss
|
| else:
|
| return weighted_loss.mean()
|
|
|
| def get_beta(self):
|
| return self.beta.item()
|
|
|
|
|
| class DynamicBetaSmoothL1Loss(nn.Module):
|
| def __init__(self, init_beta=0.5, alpha=0.1, reduction='mean'):
|
| super(DynamicBetaSmoothL1Loss, self).__init__()
|
| self.beta = init_beta
|
| self.alpha = alpha
|
| self.reduction = reduction
|
|
|
| def forward(self, pred, target):
|
|
|
| pred = torch.as_tensor(pred, device=target.device, dtype=target.dtype)
|
| target = torch.as_tensor(target, device=target.device, dtype=target.dtype)
|
|
|
| diff = torch.abs(pred - target)
|
|
|
|
|
| current_beta = torch.median(diff.detach())
|
| self.beta = (1 - self.alpha) * self.beta + self.alpha * current_beta.item()
|
|
|
| loss = torch.where(
|
| diff < self.beta,
|
| 0.5 * (diff ** 2) / self.beta,
|
| diff - 0.5 * self.beta
|
| )
|
|
|
| if self.reduction == 'mean':
|
| return loss.mean()
|
| elif self.reduction == 'sum':
|
| return loss.sum()
|
| return loss
|
|
|
|
|
| class DynamicBetaSmoothL1LossMultiStep(nn.Module):
|
| def __init__(self, init_beta=0.5, alpha=0.1, reduction='mean', eps=1e-6):
|
| super().__init__()
|
| self.register_buffer('beta', torch.tensor(init_beta))
|
| self.alpha = alpha
|
| self.reduction = reduction
|
| self.eps = eps
|
|
|
| def forward(self, pred, target):
|
|
|
| diff = torch.abs(pred - target)
|
|
|
|
|
| batch_beta = diff.mean(dim=(0,1)).detach() + self.eps
|
| self.beta = self.alpha * self.beta + (1 - self.alpha) * batch_beta
|
| beta = self.beta
|
|
|
|
|
| loss = torch.where(diff < beta, 0.5 * diff**2 / beta, diff - 0.5*beta)
|
|
|
| if self.reduction=='mean':
|
| return loss.mean()
|
| elif self.reduction=='sum':
|
| return loss.sum()
|
| else:
|
| return loss
|
|
|
| class DynamicMultiStepQLoss(nn.Module):
|
| def __init__(self, quantiles=(0.5,), weight=1.0, step_decay=0.9, reduction='mean'):
|
| """
|
| Multi-step Quantile Loss (Pinball loss) พร้อม step weighting
|
| รองรับ target/pred หลาย step
|
|
|
| Args:
|
| quantiles: list/tuple ของ quantiles (เช่น [0.1, 0.5, 0.9])
|
| weight: global weight factor
|
| step_decay: decay สำหรับ step หลัง ๆ
|
| reduction: 'mean', 'sum', หรือ None
|
| """
|
| super().__init__()
|
| self.quantiles = torch.tensor(quantiles, dtype=torch.float32)
|
| self.weight = float(weight)
|
| self.step_decay = float(step_decay)
|
| self.reduction = reduction
|
|
|
| @staticmethod
|
| def _to_bh(x):
|
| """บังคับ tensor ให้อยู่รูป (B,H)"""
|
| if x.dim() == 1:
|
| x = x.unsqueeze(0)
|
|
|
| if x.dim() == 2:
|
| return x
|
|
|
| if x.dim() == 3:
|
|
|
| if x.size(-1) == 1:
|
| return x.squeeze(-1)
|
|
|
| if x.size(1) == 1:
|
| return x.squeeze(1)
|
|
|
| if x.size(-1) == 2:
|
| return x[..., 0]
|
|
|
| if x.dim() >= 3:
|
| B = x.size(0)
|
| return x.reshape(B, -1)
|
|
|
| return x
|
|
|
| def forward(self, pred, true):
|
| """
|
| pred: (B,H,Q) หรือ (B,H) ถ้า Q=1
|
| true: (B,H)
|
| """
|
| pred = torch.as_tensor(pred)
|
| true = torch.as_tensor(true)
|
|
|
|
|
| true = self._to_bh(true)
|
|
|
| if pred.dim() == 2:
|
|
|
| pred = pred.unsqueeze(-1)
|
|
|
| if pred.dim() == 3:
|
| Bp, Hp, Qp = pred.shape
|
| Bt, Ht = true.shape
|
| if Bp != Bt:
|
| raise ValueError(f"[DynamicMultiStepQLoss] batch mismatch: pred {pred.shape}, true {true.shape}")
|
| H = min(Hp, Ht)
|
| pred = pred[:, :H, :]
|
| true = true[:, :H]
|
| else:
|
| raise ValueError(f"[DynamicMultiStepQLoss] pred shape must be (B,H) or (B,H,Q), got {pred.shape}")
|
|
|
|
|
| true = true.unsqueeze(-1).expand_as(pred)
|
|
|
|
|
| H = pred.size(1)
|
| step_weights = torch.pow(
|
| torch.tensor(self.step_decay, dtype=pred.dtype, device=pred.device),
|
| torch.arange(H, dtype=pred.dtype, device=pred.device)
|
| ).view(1, H, 1)
|
|
|
|
|
| quantiles = self.quantiles.to(pred.device).view(1, 1, -1)
|
|
|
|
|
| diff = true - pred
|
| loss = torch.max(quantiles * diff, (quantiles - 1) * diff)
|
|
|
|
|
| loss = loss * step_weights * self.weight
|
|
|
| if self.reduction == 'mean':
|
| return loss.mean()
|
| elif self.reduction == 'sum':
|
| return loss.sum()
|
| else:
|
| return loss |