Ë
    þÍ:jg"  ã                   óf  — d dl mZ d dlZd dlmZ d dlmZ d dlmZ d dlm	Z	 d dl
mZ 	 	 	 dded	ed
eed      ded   deed      ddfd„Z	 ddededed	eded   deeeef   fd„Z	 	 	 ddeded
eed      deed      dee   defd„Z	 	 	 	 ddededed	ed
eed      ded   deed      defd„Zy)é    )ÚOptionalN)ÚTensor)ÚLiteral)Ú_segmentation_inputs_format)Úrank_zero_warn)Ú_safe_divideÚnum_classesÚinclude_backgroundÚaverage©ÚmicroÚmacroÚweightedÚnoneÚinput_format©úone-hotÚindexÚmixedÚaggregation_level©Ú
samplewiseÚglobalÚreturnc                 ó  — t        | t        «      r| dk  rt        d| › d�«      ‚t        |t        «      st        d|› d�«      ‚g d¢}|�||vrt        d|› d|› d�«      ‚|d	vrt        d
|› d�«      ‚|dvrt        d|› �«      ‚y)z%Validate the arguments of the metric.r   zDExpected argument `num_classes` must be a positive integer, but got ú.zBExpected argument `include_background` must be a boolean, but got r   Nz)Expected argument `average` to be one of z or None, but got r   zSExpected argument `input_format` to be one of 'one-hot', 'index', 'mixed', but got r   zSExpected argument `aggregation_level` to be one of `samplewise`, `global`, but got )Ú
isinstanceÚintÚ
ValueErrorÚbool)r	   r
   r   r   r   Úallowed_averages         ú~/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torchmetrics/functional/segmentation/dice.pyÚ_dice_score_validate_argsr#      sÐ   € ô �k¤3Ô'¨;¸!Ò+;ÜÐ_Ð`kÐ_lÐlmÐnÓoÐoÜÐ(¬$Ô/ÜÐ]Ð^pÐ]qÐqrÐsÓtÐtÚ<€OØÐ˜w¨oÑ=ÜÐDÀ_ÐDUÐUgÐhoÐgpÐpqÐrÓsÐsØÐ8Ñ8ÜØaÐbnÐaoÐopÐqó
ð 	
ð Ð 8Ñ8ÜØaÐbsÐatÐuó
ð 	
ð 9ó    ÚpredsÚtargetc                 ó  — t        | ||||«      \  } }t        t        d|j                  «      «      }t	        j
                  | |z  |¬«      }t	        j
                  ||¬«      }t	        j
                  | |¬«      }d|z  }	||z   }
|}|	|
|fS )z8Update the state with the current prediction and target.é   ©Údim)r   ÚlistÚrangeÚndimÚtorchÚsum)r%   r&   r	   r
   r   Úreduce_axisÚintersectionÚ
target_sumÚpred_sumÚ	numeratorÚdenominatorÚsupports               r"   Ú_dice_score_updater7   2   sŒ   € ô 0°°vÐ?QÐS^Ð`lÓm�M€Eˆ6ä”u˜Q §¡Ó,Ó-€KÜ—9‘9˜U V™^°Ô=€LÜ—‘˜6 {Ô3€JÜ�y‰y˜ KÔ0€Hà�LÑ €IØ˜ZÑ'€KØ€GØ�k 7Ð*Ð*r$   r4   r5   r6   c                 ó.  — |dk(  rgt        j                  | d¬«      j                  d«      } t        j                  |d¬«      j                  d«      }|�t        j                  |d¬«      nd}|dk(  r<t        j                  | d¬«      } t        j                  |d¬«      }t        | |d¬«      S t        | |d¬«      }|d	k(  rt        j                  |d¬«      S |d
k(  r¥t        |t         j                  «      st        dt        |«      › d�«      ‚t        |t        j                  |dd¬«      d¬«      }|j                  «       j                  d¬«      }t        j                  ||z  d¬«      }t         j                  ||<   |S |dv r|S t        d|› d�«      ‚)z:Compute the Dice score from the numerator and denominator.r   r   r)   Nr   éÿÿÿÿÚnan)Úzero_divisionr   r   z1Expected argument `support` to be a tensor, got: r   T)r*   Úkeepdim)r   NzInvalid value for `average`: )r.   r/   Ú	unsqueezer   Únanmeanr   r   r   ÚtypeÚisnanÚallÚnansumr:   )r4   r5   r   r   r6   ÚdiceÚweightsÚnan_masks           r"   Ú_dice_score_computerF   G   sh  € ð ˜HÒ$Ü—I‘I˜i¨QÔ/×9Ñ9¸!Ó<ˆ	Ü—i‘i °Ô3×=Ñ=¸aÓ@ˆØ/6Ð/B”%—)‘)˜G¨Õ+Èˆà�'ÒÜ—I‘I˜i¨RÔ0ˆ	Ü—i‘i °Ô4ˆÜ˜I {À%ÔHÐHä˜	 ;¸eÔD€DØ�'ÒÜ�}‰}˜T rÔ*Ð*Ø�*ÒÜ˜'¤5§<¡<Ô0ÜÐPÔQUÐV]ÓQ^ÐP_Ð_`ÐaÓbÐbÜ˜w¬¯	©	°'¸rÈ4Ô(PÐ`eÔfˆØ—:‘:“<×#Ñ#¨Ð#Ó+ˆÜ�|‰|˜D 7™N°Ô3ˆÜŸ™ˆˆX‰ØˆØ�.Ñ ØˆÜ
Ð4°W°I¸QÐ?Ó
@Ð@r$   c                 ó�   — |dk(  rt        dt        «       t        |||||«       t        | ||||«      \  }}}	t	        |||||	¬«      S )aØ  Compute the Dice score for semantic segmentation.

    Args:
        preds: Predictions from model
        target: Ground truth values
        num_classes: Number of classes
        include_background: Whether to include the background class in the computation
        average: The method to average the dice score. Options are ``"micro"``, ``"macro"``, ``"weighted"``, ``"none"``
            or ``None``. This determines how to average the dice score across different classes.
        input_format: What kind of input the function receives.
            Choose between ``"one-hot"`` for one-hot encoded tensors, ``"index"`` for index tensors
            or ``"mixed"`` for one one-hot encoded and one index tensor
        aggregation_level: The level at which to aggregate the dice score. Options are ``"samplewise"`` or ``"global"``.
            For ``"samplewise"`` the dice score is computed for each sample and then averaged. For ``"global"`` the dice
            score is computed globally over all samples.

    Returns:
        The Dice score.

    Example (with one-hot encoded tensors):
        >>> from torch import randint
        >>> from torchmetrics.functional.segmentation import dice_score
        >>> _ = torch.manual_seed(42)
        >>> preds = randint(0, 2, (4, 5, 16, 16))  # 4 samples, 5 classes, 16x16 prediction
        >>> target = randint(0, 2, (4, 5, 16, 16))  # 4 samples, 5 classes, 16x16 target
        >>> # dice score micro averaged over all classes
        >>> dice_score(preds, target, num_classes=5, average="micro")
        tensor([0.4842, 0.4968, 0.5053, 0.4902])
        >>> # dice score per sample and class
        >>> dice_score(preds, target, num_classes=5, average="none")
        tensor([[0.4724, 0.5185, 0.4710, 0.5062, 0.4500],
                [0.4571, 0.4980, 0.5191, 0.4380, 0.5649],
                [0.5428, 0.4904, 0.5358, 0.4830, 0.4724],
                [0.4715, 0.4925, 0.4797, 0.5267, 0.4788]])
        >>> # global dice score over all samples with macro averaging
        >>> dice_score(preds, target, num_classes=5, average="macro", aggregation_level="global")
        tensor([0.4942])

    Example (with index tensors):
        >>> from torch import randint
        >>> from torchmetrics.functional.segmentation import dice_score
        >>> _ = torch.manual_seed(42)
        >>> preds = randint(0, 5, (4, 16, 16))  # 4 samples, 5 classes, 16x16 prediction
        >>> target = randint(0, 5, (4, 16, 16))  # 4 samples, 5 classes, 16x16 target
        >>> # dice score micro averaged over all classes
        >>> dice_score(preds, target, num_classes=5, average="micro", input_format="index")
        tensor([0.2031, 0.1914, 0.2266, 0.1641])
        >>> # dice score per sample and class
        >>> dice_score(preds, target, num_classes=5, average="none", input_format="index")
        tensor([[0.1731, 0.1667, 0.2400, 0.2424, 0.1947],
                [0.2245, 0.2247, 0.2321, 0.1132, 0.1682],
                [0.2500, 0.2476, 0.1887, 0.1818, 0.2718],
                [0.1308, 0.1800, 0.1980, 0.1607, 0.1522]])
        >>> # global dice score over all samples with macro averaging
        >>> dice_score(preds, target, num_classes=5, average="macro", aggregation_level="global", input_format="index")
        tensor([0.1965])

    r   zµdice_score metric currently defaults to `average=micro`, but will change to`average=macro` in the v1.9 release. If you've explicitly set this parameter, you can ignore this warning.)r   r6   )r   ÚUserWarningr#   r7   rF   )
r%   r&   r	   r
   r   r   r   r4   r5   r6   s
             r"   Ú
dice_scorerI   i   sd   € ðF �'ÒÜðUô ô		
ô ˜kÐ+=¸wÈÐVgÔhÜ&8¸ÀÈÐUgÐiuÓ&vÑ#€Iˆ{˜GÜ˜y¨+°wÐRcÐmtÔuÐur$   )r   r   r   )r   )r   r   N)Tr   r   r   )Útypingr   r.   r   Útyping_extensionsr   Ú*torchmetrics.functional.segmentation.utilsr   Útorchmetrics.utilitiesr   Útorchmetrics.utilities.computer   r   r    r#   Útupler7   rF   rI   © r$   r"   ú<module>rQ      sÑ  ðõ ã Ý Ý %å RÝ 1Ý 7ð HOØ9BØCOñ
Øð
àð
ð �gÐBÑCÑDð
ð Ð5Ñ6ð	
ð
   Ð(>Ñ ?Ñ@ð
ð 
ó
ð< :Cñ+Øð+àð+ð ð+ð ð	+ð
 Ð5Ñ6ð+ð ˆ6�6˜6Ð!Ñ"ó+ð0 HOØCOØ $ñAØðAàðAð �gÐBÑCÑDðAð   Ð(>Ñ ?Ñ@ð	Að
 �fÑðAð óAðL  $ØGNØ9BØCOñLvØðLvàðLvð ðLvð ð	Lvð
 �gÐBÑCÑDðLvð Ð5Ñ6ðLvð   Ð(>Ñ ?Ñ@ðLvð ôLvr$   