Ë
    (täið%  ã                   ó\  — d dl Z d dlmZ d dlmZ d dlZddlmZmZ ddl	m
Z
 ddlmZ dd	ej                  d
ej                  dz  dej                  fd„Z	 	 ddej                  dej                  ded
ej                  dz  dej                  f
d„Ze G d„ de
«      «       Z G d„ dee«      Zy)é    N)Ú	dataclass)ÚLiteralé   )ÚConfigMixinÚregister_to_config)Ú
BaseOutputé   )ÚSchedulerMixinÚtÚ	generatorÚreturnc                 óJ  — |�|j                   n| j                   }t        j                  | |¬«      j                  dd|¬«      j	                  | j                   «      }t        j
                  t        j
                  |j                  d«      «       j                  d«      «       S )aš  
    Generate Gumbel noise for sampling.

    Args:
        t (`torch.Tensor`):
            Input tensor to match the shape and dtype of the output noise.
        generator (`torch.Generator`, *optional*):
            A random number generator for reproducible sampling.

    Returns:
        `torch.Tensor`:
            Gumbel-distributed noise with the same shape, dtype, and device as the input tensor.
    ©Údevicer   r	   ©r   ç#B’¡œÇ;)r   ÚtorchÚ
zeros_likeÚuniform_ÚtoÚlogÚclamp)r   r   r   Únoises       úu/Volumes/fast/ai/experiments/MLX_z-image/.venv/lib/python3.12/site-packages/diffusers/schedulers/scheduling_amused.pyÚgumbel_noiser      s‚   € ð "+Ð!6ˆY×Ò¸A¿H¹H€FÜ×Ñ˜Q vÔ.×7Ñ7¸¸1È	Ð7ÓR×UÑUÐVW×V^ÑV^Ó_€EÜ�I‰IœŸ	™	 %§+¡+¨eÓ"4Ó5Ð5×<Ñ<¸UÓCÓDÐDÐDó    Úmask_lenÚprobsÚtemperaturec                 ó  — t        j                  |j                  d«      «      |t        ||¬«      z  z   }t        j                  |d¬«      j
                  }t        j                  |d| j                  «       «      }||k  }|S )a‘  
    Mask tokens by selecting the top-k lowest confidence scores with temperature-based randomness.

    Args:
        mask_len (`torch.Tensor`):
            Number of tokens to mask per sample in the batch.
        probs (`torch.Tensor`):
            Probability scores for each token.
        temperature (`float`, *optional*, defaults to 1.0):
            Temperature parameter for controlling randomness in the masking process.
        generator (`torch.Generator`, *optional*):
            A random number generator for reproducible sampling.

    Returns:
        `torch.Tensor`:
            Boolean mask indicating which tokens should be masked.
    r   r   éÿÿÿÿ©Údimr	   )r   r   r   r   ÚsortÚvaluesÚgatherÚlong)r   r   r   r   Ú
confidenceÚsorted_confidenceÚcut_offÚmaskings           r   Úmask_by_random_topkr,      sl   € ô. —‘˜5Ÿ;™; uÓ-Ó.°¼|ÈEÐ]fÔ?gÑ1gÑg€JÜŸ
™
 :°2Ô6×=Ñ=ÐÜ�l‰lÐ,¨a°·±³ÓA€GØ˜7Ñ"€GØ€Nr   c                   óX   — e Zd ZU dZej
                  ed<   dZej                  dz  ed<   y)ÚAmusedSchedulerOutputa½  
    Output class for the scheduler's `step` function output.

    Args:
        prev_sample (`torch.LongTensor` of shape `(batch_size, height, width)` or `(batch_size, sequence_length)`):
            Computed sample `(x_{t-1})` of previous timestep with token IDs. `prev_sample` should be used as next model
            input in the denoising loop.
        pred_original_sample (`torch.LongTensor` of shape `(batch_size, height, width)` or `(batch_size, sequence_length)`, *optional*):
            The predicted fully denoised sample `(x_{0})` with token IDs based on the model output from the current
            timestep. `pred_original_sample` can be used to preview progress or for guidance.
    Úprev_sampleNÚpred_original_sample)	Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚTensorÚ__annotations__r0   Ú	Generator© r   r   r.   r.   =   s'   … ñ
ð —‘ÓØ37Ð˜%Ÿ/™/¨DÑ0Ô7r   r.   c                   ó¤  — e Zd ZU dZdZej                  dz  ed<   ej                  dz  ed<   e	 dde	de
d   fd	„«       Z	 	 dd
e	de	ee	e	f   z  ee	   z  deej                  z  fd„Z	 	 	 ddej"                  de	dej$                  dedej                  dz  dedeez  fd„Z	 ddej$                  de	dej                  dz  dej$                  fd„Zy)ÚAmusedScheduleraË  
    A scheduler for masked token generation as used in [`AmusedPipeline`].

    This scheduler iteratively unmasks tokens based on their confidence scores, following either a cosine or linear
    schedule. Unlike traditional diffusion schedulers that work with continuous pixel values, this scheduler operates
    on discrete token IDs, making it suitable for autoregressive and non-autoregressive masked token generation models.

    This scheduler inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the
    generic methods the library implements for all schedulers such as loading and saving.

    Args:
        mask_token_id (`int`):
            The token ID used to represent masked tokens in the sequence.
        masking_schedule (`Literal["cosine", "linear"]`, *optional*, defaults to `"cosine"`):
            The schedule type for determining the mask ratio at each timestep. Can be either `"cosine"` or `"linear"`.
    r	   NÚtemperaturesÚ	timestepsÚmask_token_idÚmasking_schedule)ÚcosineÚlinearc                 ó    — d | _         d | _        y ©N)r;   r<   )Úselfr=   r>   s      r   Ú__init__zAmusedScheduler.__init__f   s   € ð !ˆÔØˆ�r   Únum_inference_stepsr   r   c                 ó  — t        j                  ||¬«      j                  d«      | _        t	        |t
        t        f«      r%t        j                  |d   |d   ||¬«      | _        y t        j                  |d||¬«      | _        y )Nr   r   r	   g{®Gáz„?)	r   ÚarangeÚflipr<   Ú
isinstanceÚtupleÚlistÚlinspacer;   )rC   rE   r   r   s       r   Úset_timestepszAmusedScheduler.set_timestepso   sj   € ô Ÿ™Ð&9À&ÔI×NÑNÈqÓQˆŒä�k¤E¬4 =Ô1Ü %§¡¨{¸1©~¸{È1¹~ÐObÐkqÔ rˆDÕä %§¡¨{¸DÐBUÐ^dÔ eˆDÕr   Úmodel_outputÚtimestepÚsampleÚstarting_mask_ratior   Úreturn_dictr   c                 óˆ  — |j                   dk(  xr |j                   dk(  }|rM|j                  \  }}	}
}|j                  ||
|z  «      }|j                  ||	|
|z  «      j                  ddd«      }|| j                  j
                  k(  }|j                  d¬«      }|j                  }|�|j                  |j                  «      n|}|j                  j                  dk(  r-|j                  t        j                  k7  r|j                  «       }|j                  d|j                  d«      «      }t        j                  |d|¬	«      j                  |¬
«      } |d d …df   j                   |j                  d d Ž }t        j"                  |||«      }|dk(  r|}�nò|j                  d   }| j$                  |k(  j'                  «       }|dz   t)        | j$                  «      z  }| j                  j*                  dk(  r*t        j,                  |t.        j0                  z  dz  «      }nA| j                  j*                  dk(  rd|z
  }n"t3        d| j                  j*                  › �«      ‚||z  }||z  j5                  «       }t        j6                  |j9                  dd¬«      dz
  |«      }t        j:                  t        j<                  dg|j                  ¬
«      |«      }t        j>                  |d|d d …d d …d f   «      d d …d d …df   }t        j"                  ||t        j@                  |j                  «      j:                  «      }tC        ||| jD                  |   |«      }t        j"                  || j                  j
                  |«      }|r&|j                  
«      }|j                  ||
|«      }|s||fS tG        ||«      S )Né   é   r   r   r	   r!   r"   Úcpur   r   r?   r@   úunknown masking schedule T)r#   Úkeepdim)$ÚndimÚshapeÚreshapeÚpermuteÚconfigr=   Úsoftmaxr   r   ÚtypeÚdtyper   Úfloat32ÚfloatÚsizeÚmultinomialÚviewÚwherer<   ÚnonzeroÚlenr>   ÚcosÚmathÚpiÚ
ValueErrorÚfloorÚminÚsumÚmaxÚtensorr&   Úfinfor,   r;   r.   )rC   rN   rO   rP   rQ   r   rR   Útwo_dim_inputÚ
batch_sizeÚcodebook_sizeÚheightÚwidthÚunknown_mapr   r   Úprobs_r0   r/   Úseq_lenÚstep_idxÚratioÚ
mask_ratior   Úselected_probsr+   s                            r   ÚstepzAmusedScheduler.step|   sS  € ð Ÿ™ qÑ(ÒC¨\×->Ñ->À!Ñ-CˆáØ7C×7IÑ7IÑ4ˆJ˜ v¨uØ—^‘^ J°¸±Ó?ˆFØ'×/Ñ/°
¸MÈ6ÐTYÉ>ÓZ×bÑbÐcdÐfgÐijÓkˆLà §¡× 9Ñ 9Ñ9ˆà×$Ñ$¨Ð$Ó,ˆà—‘ˆØ/8Ð/D�—‘˜)×*Ñ*Ô+È%ˆØ�=‰=×Ñ Ò&¨6¯<©<¼5¿=¹=Ò+HØ—\‘\“^ˆFØ—‘  E§J¡J¨r£NÓ3ˆÜ$×0Ñ0°¸ÀiÔP×SÑSÐ[aÐSÓbÐØ>Ð3²A°q°DÑ9×>Ñ>ÀÇÁÈCÈRÐ@PÐQÐÜ$Ÿ{™{¨;Ð8LÈfÓUÐà�qŠ=Ø.ŠKà—l‘l 1‘oˆGØŸ™¨(Ñ2×;Ñ;Ó=ˆHØ ‘\¤S¨¯©Ó%8Ñ8ˆEà�{‰{×+Ñ+¨xÒ7Ü"ŸY™Y u¬t¯w©w¡¸Ñ':Ó;‘
Ø—‘×-Ñ-°Ò9Ø ™Y‘
ä Ð#<¸T¿[¹[×=YÑ=YÐ<ZÐ![Ó\Ð\à,¨zÑ9ˆJà *Ñ,×3Ñ3Ó5ˆHä—y‘y §¡°RÀ Ó!FÈÑ!JÈHÓUˆHä—y‘y¤§¡¨q¨c¸,×:MÑ:MÔ!NÐPXÓYˆHä"Ÿ\™\¨%°Ð5IÊ!ÊQÐPTÈ*Ñ5UÓVÒWXÒZ[Ð]^ÐW^Ñ_ˆNä"Ÿ[™[¨°nÄeÇkÁkÐR`×RfÑRfÓFg×FkÑFkÓlˆNä)¨(°NÀD×DUÑDUÐV^ÑD_ÐajÓkˆGô  Ÿ+™+ g¨t¯{©{×/HÑ/HÐJ^Ó_ˆKáØ%×-Ñ-¨j¸&À%ÓHˆKØ#7×#?Ñ#?À
ÈFÐTYÓ#ZÐ áØÐ!5Ð6Ð6ä$ [Ð2FÓGÐGr   c                 ó|  — | j                   |k(  j                  «       }|dz   t        | j                   «      z  }| j                  j                  dk(  r*t        j                  |t        j                  z  dz  «      }nA| j                  j                  dk(  rd|z
  }n"t        d| j                  j                  › �«      ‚t        j                  |j                  |�|j                  n|j                  |¬«      j                  |j                  «      |k  }|j                  «       }| j                  j                  ||<   |S )a“  
        Add noise to a sample by randomly masking tokens according to the masking schedule.

        Args:
            sample (`torch.LongTensor`):
                The input sample containing token IDs to be partially masked.
            timesteps (`int`):
                The timestep that determines how much masking to apply. Higher timesteps result in more masking.
            generator (`torch.Generator`, *optional*):
                A random number generator for reproducible masking.

        Returns:
            `torch.LongTensor`:
                The sample with some tokens replaced by `mask_token_id` according to the masking schedule.
        r	   r?   r   r@   rW   )r   r   )r<   rg   rh   r]   r>   r   ri   rj   rk   rl   ÚrandrZ   r   r   Úcloner=   )	rC   rP   r<   r   r{   r|   r}   Úmask_indicesÚmasked_samples	            r   Ú	add_noisezAmusedScheduler.add_noiseÁ   s  € ð* —N‘N iÑ/×8Ñ8Ó:ˆØ˜A‘¤ T§^¡^Ó!4Ñ4ˆà�;‰;×'Ñ'¨8Ò3ÜŸ™ 5¬4¯7©7¡?°QÑ#6Ó7‰JØ�[‰[×)Ñ)¨XÒ5Ø˜U™‰JäÐ8¸¿¹×9UÑ9UÐ8VÐWÓXÐXô �J‰JØ—‘¸Ð9N Y×%5Ò%5ÐTZ×TaÑTaÐmvôç‰b�—‘ÓØñð 	ð Ÿ™›ˆà&*§k¡k×&?Ñ&?ˆ�lÑ#àÐr   )r?   ))r   r   N)ç      ð?NTrB   )r1   r2   r3   r4   Úorderr   r7   r6   r   Úintr   rD   rJ   rK   Ústrr   rM   r5   Ú
LongTensorrb   Úboolr.   r   r…   r8   r   r   r:   r:   O   sk  … ñð" €Eà—/‘/ DÑ(Ó(Ø�‰ Ñ%Ó%àð 9Añàðð "Ð"4Ñ5òó ðð :@Ø%)ñ	fà ðfð ˜5  c ™?Ñ*¨T°#©YÑ6ðfð �e—l‘lÑ"ó	fð$ &)Ø,0Ø ñCHà—l‘lðCHð ðCHð × Ñ ð	CHð
 #ðCHð —?‘? TÑ)ðCHð ðCHð 
 Ñ	&óCHðR -1ñ	*à× Ñ ð*ð ð*ð —?‘? TÑ)ð	*ð
 
×	Ñ	ô*r   r:   rB   )r†   N)rj   Údataclassesr   Útypingr   r   Úconfiguration_utilsr   r   Úutilsr   Úscheduling_utilsr
   r5   r7   r   rb   r,   r.   r:   r8   r   r   Ú<module>r‘      sÄ   ðÛ Ý !Ý ã ç AÝ Ý ,ñE�E—L‘Lð E¨U¯_©_¸tÑ-Cð EÈuÏ|É|ó Eð, Ø(,ñ	Ø�l‰lðà�<‰<ðð ðð �‰ Ñ%ð	ð
 ‡\�\óð< ô8˜Jó 8ó ð8ô"\�n kõ \r   