
      i#n                        U d dl mZ d dlmZmZmZmZmZmZ d dl	Z	d dl
mZ d dlmZmZ erd dlmZ d dlmZ eeeeef   f   Z G d d	      Z G d
 de      Z	 dde	j4                  j6                  deg e	j8                  f   deee	j8                  ge	j8                  f      defdZi de	j<                  de	j>                  ddde	j@                  de	jB                  de	jD                  dide	j<                  de	j>                  ddde	j@                  de	jB                  de	jD                  dide	j<                  de	j>                  ddd e	j@                  d!e	jB                  d!e	jD                  d"id#e	j<                  d$e	j>                  d%dd&e	j@                  d e	jB                  d e	jD                  d'id(e	j<                  d)e	j>                  d*dd+e	j@                  d,e	jB                  d,e	jD                  d-id.e	j>                  d/dd/e	j@                  d/e	jB                  d/e	jD                  d0d1d2id3e	j>                  d4dd4e	j@                  d4e	jB                  d4e	jD                  d5d1d6id7e	j>                  d8dd8e	j@                  d8e	jB                  d8e	jD                  d9d1d:id;e	j>                  d<dde	j@                  d=e	jB                  d=e	jD                  d>d1d?id@e	j>                  dAddAe	j@                  dBe	jB                  dBe	jD                  dCd1dDidEe	j<                  dFe	j>                  dGddHe	j@                  dIe	jB                  dIe	jD                  dJidKe	j>                  dLddMe	j@                  dLe	jB                  dLe	jD                  dNd1dOidPe	j>                  dQddRe	j@                  dQe	jB                  dQe	jD                  dSd1dTidUe	j>                  dVddWe	j@                  dXe	jB                  dXe	jD                  dYd1dZid[e	j>                  d\dd\e	j@                  d\e	jB                  d\e	jD                  d]d1d^id_e	j>                  d`dd`e	j@                  d`e	jB                  d`e	jD                  dad1dbidce	j>                  ddddde	j@                  dde	jB                  dde	jD                  ded1dfie	j>                  dgddge	j@                  dge	jB                  dge	jD                  dhd1diie	j>                  djddje	j@                  dje	jB                  dje	jD                  dkd1dlie	j>                  dme	jB                  dne	jD                  dod1dpie	j>                  dqe	jB                  drie	j>                  dqe	jB                  dse	jD                  dtd1duie	j>                  dve	jB                  dwe	jD                  dxd1dyie	j>                  dze	jB                  d{e	jD                  d|d1d}ie	j>                  d~e	jB                  de	jD                  dd1die	j>                  de	jB                  de	jD                  dd1die	j<                  de	j>                  de	jB                  dXie	j<                  de	j>                  de	jB                  die	j<                  de	j>                  de	jB                  doidZ#eeeeee	jH                  f   ef   f   e%d<   dddddZ&de	jN                  dee	jH                  ef   dee   fdZ(ddde	jH                  fdZ) ede      Z* G d de+e*         Z,y)    )deque)TYPE_CHECKINGAnyCallableOptionalTypeVarUnionN)override)rank_zero_onlyrank_zero_warn)Fabric)	Precisionc                   ~    e Zd ZdZ	 ddee   dedededdf
dZddd	d
edededee   dee   ddfdZ	de
fdZddZy)
Throughputa  Computes throughput.

    +------------------------+-------------------------------------------------------------------------------------+
    | Key                    | Value                                                                               |
    +========================+=====================================================================================+
    | batches_per_sec        | Rolling average (over ``window_size`` most recent updates) of the number of batches |
    |                        | processed per second                                                                |
    +--------------------------+-----------------------------------------------------------------------------------+
    | samples_per_sec        | Rolling average (over ``window_size`` most recent updates) of the number of samples |
    |                        | processed per second                                                                |
    +--------------------------+-----------------------------------------------------------------------------------+
    | items_per_sec          | Rolling average (over ``window_size`` most recent updates) of the number of items   |
    |                        | processed per second                                                                |
    +--------------------------+-----------------------------------------------------------------------------------+
    | flpps_per_sec          | Rolling average (over ``window_size`` most recent updates) of the number of flops   |
    |                        | processed per second                                                                |
    +--------------------------+-----------------------------------------------------------------------------------+
    | device/batches_per_sec | batches_per_sec divided by world size                                               |
    +--------------------------+-----------------------------------------------------------------------------------+
    | device/samples_per_sec | samples_per_sec divided by world size                                               |
    +--------------------------+-----------------------------------------------------------------------------------+
    | device/items_per_sec   | items_per_sec divided by world size. This may include padding depending on the data |
    +--------------------------+-----------------------------------------------------------------------------------+
    | device/flops_per_sec   | flops_per_sec divided by world size.                                                |
    +--------------------------+-----------------------------------------------------------------------------------+
    | device/mfu             | device/flops_per_sec divided by world size.                                         |
    +--------------------------+-----------------------------------------------------------------------------------+
    | time                   | Total elapsed time                                                                  |
    +--------------------------+-----------------------------------------------------------------------------------+
    | batches                | Total batches seen                                                                  |
    +--------------------------+-----------------------------------------------------------------------------------+
    | samples                | Total samples seen                                                                  |
    +--------------------------+-----------------------------------------------------------------------------------+
    | lengths                | Total items seen                                                                    |
    +--------------------------+-----------------------------------------------------------------------------------+

    Example::

        throughput = Throughput()
        t0 = time()
        for i in range(1000):
            do_work()
            if torch.cuda.is_available(): torch.cuda.synchronize()  # required or else time() won't be correct
            throughput.update(time=time() - t0, samples=i)
            if i % 10 == 0:
                print(throughput.compute())

    Notes:
        - The implementation assumes that devices FLOPs are all the same as it normalizes by the world size and only
          takes a single ``available_flops`` value.
        - items_per_sec, flops_per_sec and MFU do not account for padding if present. We suggest using
          samples_per_sec or batches_per_sec to measure throughput under this circumstance.

    Args:
        available_flops: Number of theoretical flops available for a single device.
        world_size: Number of devices available across hosts. Global metrics are not included if the world size is 1.
        window_size: Number of batches to use for a rolling average.
        separator: Key separator to use when creating per-device and global metrics.

    Navailable_flops
world_sizewindow_size	separatorreturnc                     || _         || _        |dkD  sJ || _        |dkD  sJ t        |      | _        t        |      | _        t        |      | _        t        |      | _        t        |      | _	        y )Nr      )maxlen)
r   r   r   _MonotonicWindow_time_batches_samples_lengthsr   _flops)selfr   r   r   r   s        z/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/lightning_fabric/utilities/throughput.py__init__zThroughput.__init__^   sr      /"A~~$ Q /?k.R
/?{/S/?{/S/?{/S"'{";    )lengthsflopstimebatchessamplesr#   r$   c                X   | j                   j                  |       ||k  rt        d| d| d      | j                  j                  |       | j                  j                  |       |||k  rt        d| d| d      | j
                  j                  |       t        | j                        t        | j
                        k7  r8t        dt        | j
                         dt        | j                         d      |)| j                  j                  || j                  z         yy)	ad  Update throughput metrics.

        Args:
            time: Total elapsed time in seconds. It should monotonically increase by the iteration time with each
                call.
            batches: Total batches seen per device. It should monotonically increase with each call.
            samples: Total samples seen per device. It should monotonically increase by the batch size with each call.
            lengths: Total length of the samples seen. It should monotonically increase by the lengths of a batch with
                each call.
            flops: Flops elapased per device since last ``update()`` call. You can easily compute this by using
                :func:`measure_flops` and multiplying it by the number of batches that have been processed.
                The value might be different in each device if the batch size is not the same.

        zExpected samples (z') to be greater or equal than batches ()NzExpected lengths (z') to be greater or equal than samples (zIf lengths are passed (z1), there needs to be the same number of samples ()
r   append
ValueErrorr   r   r   lenRuntimeErrorr   r   )r   r%   r&   r'   r#   r$   s         r    updatezThroughput.updateq   s   . 	

$W1':abiajjklmmW%W%  #5gY>efmenno!pqqMM  )4==!S%77"-c$--.@-A BT]]+,A/  KKut67 r"   c                 4   | j                   d   | j                  d   | j                  d   d}| j                  r| j                  d   |d<   | j                  dkD  }t        | j                         | j                   j                  k(  rF| j                   d   | j                   d   z
  }| j                  d   | j                  d   z
  }| j                  d   | j                  d   z
  }||z  }||z  }|j                  d| j                   d||z  d| j                   d|i       |r0|| j                  z  }|j                  ||| j                  z  d	       t        | j                        | j                  j                  k(  rM| j                  d   | j                  d   z
  }	|	|z  }
|
|d| j                   d
<   |r|
| j                  z  }||d
<   t        | j                        | j                  j                  k(  rt        | j                        | j                  d   z
  }| j                   d   | j                   d   z
  }||z  }|| j                  z  }|r||d<   ||d| j                   d<   | j                  r || j                  z  |d| j                   d<   |S )zCompute throughput metrics.)r%   r&   r'   r#   r   r   devicebatches_per_secsamples_per_sec)r2   r3   items_per_secflops_per_secmfu)r   r   r   r   r   r,   r   r.   r   r   sumr   )r   metricsadd_global_metricselapsed_timeelapsed_batcheselapsed_samplesdev_samples_per_secdev_batches_per_secr3   elapsed_lengthsdev_items_per_secr4   elapsed_flopsr5   dev_flops_per_secs                  r    computezThroughput.compute   s    JJrN}}R(}}R(

 ==!%r!2GI!__q0 tzz?djj///::b>DJJqM9L"mmB/$--2BBO"mmB/$--2BBO"1L"@"1L"@NN(8/L:X(8:M  ""5"G'6':T__'L  
 4==!T]]%9%99"&--"3dmmA6F"F$3l$B!BS& 0>?%$5$GM/<GO,t{{t{{111,t{{1~=M::b>DJJqM9L)L8M - ?!+8(>OGfT^^,M:;##8IDL`L`8`& 045r"   c                    | j                   j                          | j                  j                          | j                  j                          | j                  j                          | j
                  j                          y N)r   clearr   r   r   r   r   s    r    resetzThroughput.reset   sR    

r"   )Nr   d   /)r   N)__name__
__module____qualname____doc__r   floatintstrr!   r.   _THROUGHPUT_METRICSrC   rH    r"   r    r   r       s    ;| vy<'<CF<Y\<or<	<2 "&#'8 '8 	'8
 '8 #'8 }'8 
'8R2, 2hr"   r   c                   L     e Zd ZdZdddeddf fdZd
dee   dedefd	Z	 xZ
S )ThroughputMonitora  Computes throughput.

    This class will automatically keep a count of the number of log calls (``step``). But that can be modified as
    desired. For manual logging, using :class:`Throughput` directly might be desired.

    Example::

        logger = ...
        fabric = Fabric(logger=logger)
        throughput = ThroughputMonitor(fabric)
        t0 = time()
        for i in range(1, 100):
            do_work()
            if torch.cuda.is_available(): torch.cuda.synchronize()  # required or else time() won't be correct
            throughput.update(time=time() - t0, batches=i, samples=i)
            if i % 10 == 0:
                throughput.compute_and_log(step=i)

    Args:
        fabric: The Fabric object.
        \**kwargs: See available parameters in :class:`Throughput`

    fabricr   kwargsr   Nc                    |j                          t        |j                  j                        }t	        |j
                  |      }t        |   d||j                  d| || _	        d| _
        t        | j                        | _        t        | j                  i       | _        t        | j                  i       | _        t        | j                        | _        y )N)r   r   r0   )defaultrS   )_validate_launched_plugin_to_compute_dtypestrategy	precisionget_available_flopsr1   superr!   r   _fabricstepr   r.   rC   compute_and_logrH   )r   rV   rW   dtyper   	__class__s        r    r!   zThroughputMonitor.__init__   s    !!#()B)BC-fmmUCaVEVEVaZ`a	$T[[1%dllB?-d.B.BBO#DJJ/
r"   ra   c                     || j                   dz   n|| _          | j                  di |}| j                  j                  || j                          |S )zSee :meth:`Throughput.compute`

        Args:
            step: Can be used to override the logging step.
            \**kwargs: See available parameters in :meth:`Throughput.compute`

        r   )r8   ra   rS   )ra   rC   r`   log_dict)r   ra   rW   r8   s       r    rb   z!ThroughputMonitor.compute_and_log   sL     (,|TYY]	$,,((gDII>r"   rE   )rK   rL   rM   rN   r   r!   r   rP   rR   rb   __classcell__rd   s   @r    rU   rU      sA    00x 03 04 0HSM C L_ r"   rU   model
forward_fnloss_fnr   c                     ddl m}  |d      }|5  | |        n | |             j                          ddd       |j                         S # 1 sw Y   |j                         S xY w)a-  Utility to compute the total number of FLOPs used by a module during training or during inference.

    It's recommended to create a meta-device model for this:

    Example::

        with torch.device("meta"):
            model = MyModel()
            x = torch.randn(2, 32)

        model_fwd = lambda: model(x)
        fwd_flops = measure_flops(model, model_fwd)

        model_loss = lambda y: y.sum()
        fwd_and_bwd_flops = measure_flops(model, model_fwd, model_loss)

    Args:
        model: The model whose FLOPs should be measured.
        forward_fn: A function that runs ``forward`` on the model and returns the result.
        loss_fn: A function that computes the loss given the ``forward_fn`` output. If provided, the loss and `backward`
            FLOPs will be included in the result.

    r   )FlopCounterModeF)displayN)torch.utils.flop_counterrm   backwardget_total_flops)ri   rj   rk   rm   flop_counters        r    measure_flopsrs   
  s^    8 9"51L	?LJL!**,	 

 '')) 

 ''))s   &AA(	h200 sxm1g   =Bg  wBtfloat32g  2#Cg  4&kCg  4&k,C	h200 nvl1g  WHBg  WHBg  Cg  y`(Cg ?r'Ch100 nvlg  $^/lBg SCg SCg  />2,Ch100 sxmg  wBg  $^/lBg SBg  />2C	h100 pcieg   vHBg   vHBg  |Bg  |Cg @.CCzrtx 4090g  $Bg bCint4g bCzrtx 4080g  m%Bg )Bg )Czrtx 4080 superg  :Bg  :Bg  :Cl4g  ĎBg  $ Bg  $ Bg  $ Bl40g  ˓Bg  ˓Bg  ˓Bg  ˓Ca100g  ꤡBg  2Bg  2Bg  2Bg  2Ca6000g  \EBg  \EBg gBg @$ Ca40g  xBg  xBg bcBg Ca10gg  P`Bg  4&kBg  4&kBg  4&kBg  4&kBzrtx 3090 tig  @0Bg  @0Bg  @0Czrtx 3090g  Pb0Bg  q$Bg  q$ Czrtx 3080 tig  cBg  cBg  YBg  jZBg  .Bg  .Bg  IvvBg  gH|Bg  gH|Bg  }wBg  Bg  Bg  Bg   h_Bg  HBg  HBg  HBg  HBg  `cԩBg  H`Bg  3Bg  3Bg   xHBg   xHBg   xHBg FBg  Bg  <vBg  /|Bg  /|Bg  pkGBg  pkGBg  *Bg  *Bg  P`Bg  ᎬBg  BwBg  BwBg  BwBg  @YԝBg  @YԭB)zrtx 3080zrtx 3070t4quadro rtx 5000zrtx 2080 superzrtx 2080 tizrtx 2080zrtx 2070 super	titan rtxv100 sxm	v100 pcie
v100s pcie_CUDA_FLOPSg  聰vBg  ӊBg  CBg  `teB)v2v3v4	v5litepodr1   rc   c                 V   | j                   dk(  rmt        j                  j                  |       }|j	                         }d|v rd|v rd}nd|v rd}nd|v rd|v rd	}nd
|v rd}nd|v sd|v rd}nd|v r	d|v rdnd}nd|v r+|j                  d      d   }d}d|v rd}nd|v rd}d| | }nUd|v rd}nNd|v rd}nGd|v rd}n@d|v rd}n9d|v rd}n2d |v rd }n+d!|v rd!}n$d"|v rd#}nd$|v rd%}nd&|v rd'}nt        d(|       y)|t        vrt        d(|d*|       y)t        |   }|t        j                  u r&d+d,l	m
}  |       rt        j                         d-k7  rd.}||vrt        |d/|        y)t        ||         S | j                   d0k(  rd+d1lm} |rd+d2lm}	 nd+d2lm}	 |	j%                         }
|
j'                  d3      xs |
d4   j                  d5      d+   }|j	                         }t)        |t*              sJ |t,        vrt        d6|d7|        y)t        t,        |         S y))8zReturns the available theoretical FLOPs.

    This is an optimistic upper limit that could only be achievable if only thick matmuls were run in a benchmark
    environment.

    cudah200sxm1rt   nvl1rv   h100hbm3rx   nvlrw   pciehbm2ery   r{   teslar|   zgeforce rtx     r_   z supertiz tizrtx r~   r}   r   r   r   r   r   zv100-sxmr   z	v100-pcier   z
v100s-pcier   zFLOPs not found for Nz
, chip is r   )_is_ampere_or_laterhighestru   z does not support xla)_XLA_GREATER_EQUAL_2_1)tpuTYPEACCELERATOR_TYPE-zFLOPs not found for TPU z with )typetorchr   get_device_namelowersplitr   r   float32"lightning_fabric.accelerators.cudar   get_float32_matmul_precisionrP   !lightning_fabric.accelerators.xlar   torch_xla._internalr   torch_xla.experimentalget_tpu_envget
isinstancerQ   
_TPU_FLOPS)r1   rc   device_namechipnumberextradtype_to_flopsr   r   r   tpu_envs              r    r^   r^   "  s    {{fjj008  "T>~"4"t^~!$!47d?"T\#tO5Dd"ZZ_Q'FE$ &%)D_Dt^Dd]Dt^DT\D$&$DD D4DD DT!D 1+AB{"1+
4(ST$T*EMM!N"$)K)K)MQZ)Z"&k_,>ugFG>%()){{eL!/2//#kk&)VW5G-H-N-Ns-STU-V  "+s+++z!5k_F5'RS:d#$$! r"   pluginr   c                     ddl m}m}m}m}m}m}m}m}m	}	 t        | |      st        d|        t        | |      r| j                  S t        | ||f      r| j                  S t        | |      rt        j                  S t        | |	|f      r| j                   S t        | |      rt        j"                  S t        | |      r(| j$                  j&                  xs t        j(                  S t        | |      rt        j(                  S t+        |       )Nr   )	BitsandbytesPrecisionDeepSpeedPrecisionDoublePrecisionFSDPPrecisionHalfPrecisionMixedPrecisionr   TransformerEnginePrecisionXLAPrecisionz!Expected a precision plugin, got )lightning_fabric.pluginsr   r   r   r   r   r   r   r   r   r   r-   rc   _desired_input_dtyper   double_desired_dtypeint8mixed_precision_configreduce_dtyper   NotImplementedError)
r   r   r   r   r   r   r   r   r   r   s
             r    r[   r[   }  s    
 
 
 fi(>vhGHH&/0||&=.9:***&/*||&<);<=$$$&45zz&-(,,99JU]]J&)$}}
f
%%r"   T)boundc                        e Zd ZdZdeddf fdZedee   fd       Z	e
deddfd       Ze
d	ed
eddfd       Z xZS )r   zjCustom fixed size list that only supports right-append and ensures that all values increase monotonically.r   r   Nc                 0    t         |           || _        y rE   )r_   r!   r   )r   r   rd   s     r    r!   z_MonotonicWindow.__init__  s    r"   c                 *    t        |       dkD  r| d   S y )Nr   r0   )r,   rG   s    r    lastz_MonotonicWindow.last  s    t9q=8Or"   xc                     | j                   }|||k\  rt        d| d|       t        j                  | |       t	        |       | j
                  kD  r| d= y y )Nz&Expected the value to increase, last: z, current: r   )r   r+   listr*   r,   r   )r   r   r   s      r    r*   z_MonotonicWindow.append  s\    yy	EdV;WXVYZ[[D!t9t{{"Q #r"   keyvaluec                     t        d      )Nz__setitem__ is not supported)r   )r   r   r   s      r    __setitem__z_MonotonicWindow.__setitem__  s     ""@AAr"   )rK   rL   rM   rN   rP   r!   propertyr   r   r   r
   r*   r   r   rg   rh   s   @r    r   r     s    ts t  hqk  
  d   Bs B3 B4 B Br"   r   rE   )-collectionsr   typingr   r   r   r   r   r	   r   typing_extensionsr
   $lightning_fabric.utilities.rank_zeror   r   lightning_fabricr   r   r   dictrQ   rP   rO   rR   r   rU   nnModuleTensorrs   float64r   bfloat16float16r   r   rc   __annotations__r   r1   r^   r[   r   r   r   rS   r"   r    <module>r      s    I I  & O'23c5j 112 
s sl1
 1n AE$*88??$*U\\)*$* h~u||;<=$* 		$*N`@ vvFv

F`@ vvFv

F`@( uxH	y

I)`@8 wwHx

I9`@H wwFv

GI`@\ wGw

H	]`@l wGw

Hm`@| wGw

H}`@L 	wEv

FM`@\ 
wGv

F]`@r vwFv

Fs`@B wGw

HC`@R 
wGw

HS`@d wGv

Fe`@t uEu

Fu`@D wGw

FE`@T wGw

HU`@f 	wGw

F 	wGw

H 	vu

F	 	ww
 	ww

H	 	ww

H	 	ww

H	 	vw

F	 	ww

F	 	vwv 	tuv 	vwvw`@T#tE#u{{"23U:;;< `N 


X% X%U5;;;K5L X%QYZ]Q^ X%v&[ &U[[ &B CuBtAw Br"   