credit.losses.crps#
Ring-reduce ensemble CRPS loss for distributed training.
One ensemble member per data-parallel rank. Member diversity comes from the per-dp-rank RNG seed (see train_gen2: seed + data_rank) acting on stochastic model components (dropout, SKEBS, IC perturbations); every dp rank must receive the SAME batch, so ring-crps training disables dp dataset sharding.
Two entry points share the same ring exchange:
flat criterion:
loss: {training_loss: ring-crps}— RingCRPSLoss scores the whole normalizedy_predtensor at once.Gen 2 BaseLoss univariate:
loss: {type: base, args: {training_loss: ring-crps, ...}}— each variable is scored elementwise in physical units and combined with BaseLoss’s per-variable weighting (K-1 exchanges per variable per step). The launch contract is identical; train_gen2 resolves the nested name viacredit.losses.effective_loss_name.
Attributes#
Classes#
Criterion-compatible wrapper around the ring CRPS. |
Functions#
|
Elementwise fair-CRPS contribution via ring communication. |
|
Fair CRPS via ring communication, reduced to a scalar. |
Module Contents#
- credit.losses.crps.logger#
- credit.losses.crps.ring_crps_elementwise(y_pred, y, group=None)#
Elementwise fair-CRPS contribution via ring communication.
Each dp rank holds 1 ensemble member of the same sample. K-1 ring shifts pass one member buffer at a time so rank r accumulates
|x_r - x_j|for every j != r without ever materialising the full K-member ensemble on a single device.Gradient correctness (no cross-rank backward needed):
d(CRPS)/d(x_r) = sign(x_r-y)/K - sum_{j!=r} sign(x_r-x_j) / (K*(K-1))
Both terms are computed entirely from rank r’s local graph. DDP averaging (1/K) then gives the correct model-parameter gradient.
Local elementwise contribution (scaled so DDP avg = d(CRPS)/d(params)):
loss_r = [ |y_pred_r - y| - sum_{j!=r} |y_pred_r - x_j| / (K-1) ] / K
The scaling survives any linear reduction applied downstream — a plain
.mean()(ring_crps_loss) or a weighted mean such as BaseLoss’s per-variable latitude/variance weighting.- Parameters:
y_pred – Local member prediction, requires_grad. Same shape on every rank in the group.
y – Truth tensor, broadcastable against y_pred.
group – Process group holding one member per rank (the dp group). None falls back to WORLD. With K == 1 (or torch.distributed not initialized) the spread term vanishes and the loss reduces to elementwise MAE.
- Returns:
Elementwise tensor, same shape as
y_pred; backward populatesy_pred.gradcorrectly.
- credit.losses.crps.ring_crps_loss(y_pred, y, group=None)#
Fair CRPS via ring communication, reduced to a scalar.
ring_crps_elementwise(...).mean()— see that function for the member exchange, the local formula, and the gradient-correctness argument.- Returns:
Scalar loss; backward populates
y_pred.gradcorrectly.
- class credit.losses.crps.RingCRPSLoss(reduction='mean', group=None)#
Bases:
torch.nn.ModuleCriterion-compatible wrapper around the ring CRPS.
Matches the trainer call convention criterion(y, y_pred): target first, prediction second.
reductionfollows the PyTorch convention:"mean"(default) returns a scalar;"none"returns the elementwise contribution, which is what the config loaders pass and what BaseLoss requires — the trainer’s.mean()(or BaseLoss’s weighted per-variable mean) then reduces it to the same scalar, with identical gradients (the ring scaling is linear, seering_crps_elementwise).Usable inside BaseLoss as a per-variable univariate loss (
loss: {type: base, args: {training_loss: ring-crps, ...}}): the ensemble dimension lives across data-parallel ranks, not in the tensor, so each variable is scored elementwise locally. Requires the same launch contract as flat ring-crps training (one member per dp rank, every dp rank fed the SAME batch — train_gen2 arranges this when it detects ring-crps). Note this exchanges K-1 buffers per scored variable per step. For validation, prefer a deterministicvalidation_loss(e.g.mae): validation keeps dp dataset sharding, so ranks hold different samples and a cross-rank spread term would be meaningless.The process group is resolved lazily at first forward (the dp group is registered by distributed_model_wrapper_gen2, which runs before load_loss).
- supports_elementwise = True#
- reduction = 'mean'#
- forward(target, pred)#