FastInstanceSegmentationRegressionLoss#
- class FastInstanceSegmentationRegressionLoss(cost_mask: float = 1.0, cost_dice: float = 1.0, cost_class: float = 0.0, num_points: int = 0, ignore_index: int = -1, loss_weight_focal: float = 1.0, loss_weight_dice: float = 1.0, cls_weight_matched: float = 2.0, cls_weight_noobj: float = 0.1, focal_alpha: float = 0.25, focal_gamma: float = 2.0, truth_label: str = 'segment', aux_loss_weight: float = 1.0, regression_targets=None)[source]#
Bases:
ModuleInstance segmentation loss plus configurable per-instance regressions.
Extends
FastInstanceSegmentationLoss(same Hungarian matching, focal mask, Dice, and classification terms) with an arbitrary set of continuous regression heads supervised on matched query/instance pairs. Each entry inregression_targetsnames a prediction key (defaultpred_<name>), aninput_dicttarget key (default<name>), a per-instance aggregation ("mean"or"first"), a per-targetcriterion(defaultSmoothL1RegressionLoss), and aloss_weight. Deep supervision overaux_outputsis summed and scaled byaux_loss_weight.forward(pred, input_dict)returns(loss, components)with the mask/ class diagnostics plus one scalar per regression target. Registered asFastInstanceSegmentationRegressionLoss.- Parameters:
cost_mask (float) – Matcher weight on the focal mask cost. Defaults to
1.0.cost_dice (float) – Matcher weight on the Dice cost. Defaults to
1.0.cost_class (float) – Matcher weight on the classification cost. Defaults to
0.0.num_points (int) – Sampled points for the mask cost (
0uses all). Defaults to0.ignore_index (int) – Instance/segment id treated as void. Defaults to
-1.loss_weight_focal (float) – Weight on the focal mask loss. Defaults to
1.0.loss_weight_dice (float) – Weight on the Dice loss. Defaults to
1.0.cls_weight_matched (float) – Weight on the matched-query CE term. Defaults to
2.0.cls_weight_noobj (float) – Weight on the no-object CE term. Defaults to
0.1.focal_alpha (float) – Focal-loss positive-class weight. Defaults to
0.25.focal_gamma (float) – Focal-loss focusing exponent. Defaults to
2.0.truth_label (str) –
input_dictkey holding per-point instance ids. Defaults to"segment".aux_loss_weight (float) – Scale on the summed auxiliary loss. Defaults to
1.0.regression_targets – List (or name-keyed dict) of regression-target configs as described above. Defaults to
None.
Example
>>> import torch >>> from pimm.models.losses.instance_regression_fast import ( ... FastInstanceSegmentationRegressionLoss) >>> crit = FastInstanceSegmentationRegressionLoss( ... truth_label="instance", ... regression_targets=[dict(name="momentum", loss_weight=1.0)]) >>> # pred needs pred_masks plus pred_<name> for each regression target >>> pred = {"pred_masks": [torch.randn(8, 50)], ... "pred_momentum": [torch.randn(8, 1)]} >>> inst = torch.zeros(50, dtype=torch.long) >>> inst[10:30] = 1; inst[30:50] = 2 >>> input_dict = {"instance": inst.view(-1, 1), "momentum": torch.randn(50, 1)} >>> loss, comp = crit(pred, input_dict) # scalar loss + diagnostics dict >>> "momentum" in comp # one scalar per regression target True