fixes
This commit is contained in:
@@ -0,0 +1,118 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
|
||||
|
||||
class QuantileLoss(nn.Module):
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
self.register_buffer('q', torch.tensor(config.quantiles))
|
||||
|
||||
def forward(self, predictions, targets):
|
||||
diff = predictions - targets.unsqueeze(-1).expand(-1, -1, -1, predictions.shape[-1])
|
||||
ql = (1 - self.q) * F.relu(diff) + self.q * F.relu(-diff)
|
||||
losses = ql.view(-1, ql.shape[-1]).mean(0)
|
||||
return losses
|
||||
|
||||
|
||||
def qrisk(pred, tgt, quantiles):
|
||||
diff = pred - tgt
|
||||
ql = (1 - quantiles) * np.clip(diff, 0, float('inf')) + quantiles * np.clip(-diff, 0, float('inf'))
|
||||
losses = ql.reshape(-1, ql.shape[-1])
|
||||
normalizer = np.abs(tgt).mean()
|
||||
risk = 2 * losses / normalizer
|
||||
return risk.mean(0)
|
||||
|
||||
|
||||
def weighted_bce_loss_fn():
|
||||
bce = nn.BCEWithLogitsLoss(reduction="none") # Use reduction="none" to get per-element loss
|
||||
use_class_weights = False
|
||||
time_weight_factor = 1
|
||||
|
||||
def _loss_fn(y_pred, y_true):
|
||||
loss_list = [] # To accumulate per-sample loss values
|
||||
for i in range(len(y_pred)):
|
||||
current_pred = y_pred[i]
|
||||
current_true = y_true[i]
|
||||
loss_raw = bce(current_pred, current_true)
|
||||
|
||||
if use_class_weights:
|
||||
# Calculate class weights using torch operations
|
||||
try:
|
||||
factor = torch.max(current_true) / torch.mean(current_true) / 2
|
||||
non_zero_indices = current_true != 0
|
||||
class_weights = torch.ones_like(current_pred)
|
||||
class_weights[non_zero_indices] = factor
|
||||
except Exception:
|
||||
class_weights = torch.ones_like(current_pred)
|
||||
else:
|
||||
class_weights = torch.ones_like(current_pred)
|
||||
|
||||
# Calculate time weights using torch operations
|
||||
try:
|
||||
# Find indices where the event occurs
|
||||
event_indices = torch.where(current_true == 1)[0]
|
||||
if len(event_indices) > 0:
|
||||
first_event_index = event_indices[0].item()
|
||||
time_weights = torch.ones_like(current_pred)
|
||||
# Weight the loss before the first event higher if needed
|
||||
time_weights[:first_event_index] = time_weight_factor
|
||||
else:
|
||||
time_weights = torch.ones_like(current_pred)
|
||||
except Exception:
|
||||
time_weights = torch.ones_like(current_pred)
|
||||
|
||||
# Apply the weights to the raw loss
|
||||
loss = loss_raw * class_weights * time_weights
|
||||
# Optionally, take the mean over the time dimension
|
||||
loss_list.append(loss.mean())
|
||||
|
||||
# Aggregate the loss for the entire batch
|
||||
return torch.stack(loss_list).mean()
|
||||
|
||||
return _loss_fn
|
||||
|
||||
|
||||
def weighted_mse_prob_loss_fn():
|
||||
mse = nn.MSELoss(reduce=False)
|
||||
after_event_factor = 1
|
||||
use_weighted = False
|
||||
|
||||
def _loss_fn(y_pred, y_true):
|
||||
nonlocal mse, after_event_factor
|
||||
loss_list = list()
|
||||
for i in range(len(y_pred)):
|
||||
current_pred = y_pred[i]
|
||||
current_true = y_true[i]
|
||||
loss_raw = mse(current_pred, current_true)
|
||||
# calculate factor to use for weighting, only works with probability values
|
||||
if use_weighted:
|
||||
try:
|
||||
factor = torch.max(current_true) / torch.mean(current_true) / 2
|
||||
non_zero_indices = current_true != 0
|
||||
class_weights = np.ones(len(current_true))
|
||||
class_weights[non_zero_indices] = factor
|
||||
except:
|
||||
# if there are no non-zero indices, set all weights to 1
|
||||
class_weights = np.ones(len(current_true))
|
||||
else:
|
||||
# if we don't want to weight the loss, set all weights to 1
|
||||
class_weights = np.ones(len(current_true))
|
||||
|
||||
# calculate time weight, here, after the event, the loss should be higher
|
||||
try:
|
||||
last_event_index = np.where(current_true != 0)[-1]
|
||||
time_weights = np.zeros(len(current_true))
|
||||
# here, we want to weight the loss after the 'fertility curve' higher, so that we don't get the signal too late
|
||||
time_weights[last_event_index:] = after_event_factor
|
||||
except:
|
||||
# if there are no non-zero indices, set all weights to 1
|
||||
time_weights = np.ones(len(current_true))
|
||||
|
||||
loss = loss_raw * torch.tensor(class_weights) * torch.tensor(time_weights)
|
||||
loss_list.append(loss.mean())
|
||||
|
||||
return torch.stack(loss_list).mean()
|
||||
|
||||
return _loss_fn
|
||||
Reference in New Issue
Block a user