
    (tis                         d dl Z d dlmZ d dlmZ d dlZd dlZddlm	Z	m
Z
 ddlmZ ddlmZmZ e G d d	e             Z G d
 dee	      Zy)    N)	dataclass)Literal   )ConfigMixinregister_to_config)SchedulerMixin)
BaseOutput	deprecatec                       e Zd ZU ej                  ed<   dZej                  dz  ed<   dZej                  dz  ed<   dZe	dz  ed<   y)HeliosSchedulerOutputprev_sampleNmodel_outputslast_sample
this_order)
__name__
__module____qualname__torchFloatTensor__annotations__r   r   r   int     u/Volumes/fast/ai/experiments/MLX_z-image/.venv/lib/python3.12/site-packages/diffusers/schedulers/scheduling_helios.pyr   r      sF    """.2M5$$t+2,0K""T)0!Jd
!r   r   c            %       ~   e Zd Zg ZdZedddg ddddd	d
dd
g dd
dddfdededededede	de
dede	de
de	dee   dede	de
de	ded    f"d!       Zd" Zd# Zed$        Zed%        ZdUd&efd'Zd( Z	 	 	 	 	 dVd)ed*edz  d+e
ej,                  z  d,e	dz  d-e	dz  d.e	fd/Zd-ed0ed1ej0                  fd2Zd3 Zd4 ZdWd5Zd6 Z	 	 	 	 	 	 dXd7ej<                  d8eej<                  z  d9ej<                  d:ej>                  dz  d0ej<                  dz  d;ej<                  dz  d<e	d=e e!z  fd>Z"d? Z#ddd@d7ej0                  d9ej0                  d0ej0                  d=ej0                  fdAZ$dddddBd7ej0                  d9ej0                  dCed0ej0                  d;ej0                  d=ej0                  fdDZ%ddddddEdFej0                  dGej0                  dHej0                  dCedIej0                  d0ej0                  d=ej0                  fdJZ&	 	 	 	 	 	 	 	 	 	 	 	 dYd7ej0                  d8eej0                  z  d9ej0                  d<e	dKedLedIej0                  d0ej0                  d;ej0                  dMedNedOedPej0                  d=e e!z  fdQZ'	 	 	 	 dZd7ej<                  d8eej<                  z  d9ej<                  d:ej>                  dz  d<e	d=e e!z  fdRZ(dS Z)dT Z*y)[HeliosScheduler   i        ?   )r   UUUUUU?gUUUUUU?r   r    Fflow_predictionr   Tbh2Nunipcexponentialnum_train_timestepsshiftstagesstage_rangegammathresholdingprediction_typesolver_order
predict_x0solver_typelower_order_finaldisable_correctorsolver_puse_flow_sigmasscheduler_typeuse_dynamic_shiftingtime_shift_type)r$   linearc                    i | _         i | _        i | _        i | _        i | _        i | _        | j                          | j                  d   j                         | _	        | j                  d   j                         | _
        || _        |
dvr1|
dv r| j                  d       nt        |
 d| j                         |	| _        d g|z  | _        d g|z  | _        d| _        || _        || _        d | _        d | _        d | _        y )Nr   )bh1r"   )midpointheunlogrhor"   )r.   z is not implemented for )timestep_ratiostimesteps_per_stagesigmas_per_stagestart_sigmas
end_sigmasori_start_sigmasinit_sigmas_for_each_stagesigmasitem	sigma_min	sigma_maxr)   r   NotImplementedError	__class__r-   r   timestep_listlower_order_numsr0   r1   r   _step_index_begin_index)selfr%   r&   r'   r(   r)   r*   r+   r,   r-   r.   r/   r0   r1   r2   r3   r4   r5   s                     r   __init__zHeliosScheduler.__init__'   s   ,  "#%  " " 	'')R--/Q,,.
n,<<''E':)[M9QRVR`R`Qa*bcc$"Vl2"Vl2 !!2  r   c                    | j                   j                  }| j                   j                  }t        j                  dd|z  |dz         }d|z
  }t        j
                  ||z  d|dz
  |z  z   z        dd j                         }t        j                  |      }||z  j                         }d| _
        d| _        || _        |j                  d      | _        y)z<
        initialize the global timesteps and sigmas
        r   r   Nr8   cpu)configr%   r&   nplinspaceflipcopyr   
from_numpyclonerL   rM   	timestepstorD   )rN   r%   r&   alphasrD   rY   s         r   init_sigmaszHeliosScheduler.init_sigmasZ   s     #kk==!!Q$7 79Lq9PQv1	V/C+CDEcrJOOQ!!&)1188:	 "ii&r   c                    | j                          g }| j                  j                  }| j                  j                  }| j                  j                  }t        |      D ]  }t        ||   |z        }t        |d      }t        ||dz      |z        }t        ||      }| j                  |   j                         }||k  r| j                  |   j                         nd}	|| j                  |<   |dk7  rJd|z
  }
| j                  j                  }dt        j                  dd|z  z         d|
z
  z  |
z   z  |
z  }d|z
  }|j                  ||	z
         || j                   |<   |	| j"                  |<    t%        |      }t        |      D ]K  }|dk(  rd}nt%        |d|       |z  }||dz
  k(  rd}nt%        |d|dz          |z  }||f| j&                  |<   M t        |      D ]  }| j&                  |   }t        | j(                  t        |d   |z           d      }| j(                  t        t        |d   |z        |dz
           }t+        j,                  |||dz         }t/        |t0        j2                        r|dd nt1        j4                  |dd       | j6                  |<   t+        j,                  dd|dz         }t1        j4                  |dd       | j8                  |<    y)	z3
        Init the timesteps for each stage
        r   r   g        Ng?i  r8   g+?)r\   rR   r'   r%   r(   ranger   maxminrD   rE   rB   r)   mathsqrtappendr@   rA   sumr=   rY   rS   rT   
isinstancer   TensorrW   r>   r?   )rN   stage_distancer'   training_stepsr(   i_sstart_indice
end_indicestart_sigma	end_sigma	ori_sigmar)   corrected_sigmatot_distancestart_ratio	end_ratiotimestep_ratiotimestep_maxtimestep_minrY   stage_sigmass                        r   rC   z*HeliosScheduler.init_sigmas_for_each_stagel   s    	##88kk-- =C{3/.@AL|Q/L[q1NBCJZ8J++l388:K:D~:UJ/446[^I)4D!!#&axO	))#$		!q5y/(Ba)m(TW`(`#aen"n/1!!+	"9:%0Dc"#,DOOC ' !, >*=Cax!!.#"67,Ffqj .	yq 9:\I	)4i(@D  % ! =C!11#6Nt~~c.2Cn2T.UVX[\L>>#c.2Cn2T.UWehiWi*jkLL,QR@RSI",Y"E	#25K[K[\efigi\jKk $$S) ;;ua!1CDL).)9)9,s:K)LD!!#& !r   c                     | j                   S )zg
        The index counter for current timestep. It will increase 1 after each scheduler step.
        )rL   rN   s    r   
step_indexzHeliosScheduler.step_index   s    
 r   c                     | j                   S )zq
        The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
        rM   rx   s    r   begin_indexzHeliosScheduler.begin_index   s    
    r   r|   c                     || _         y)z
        Sets the begin index for the scheduler. This function should be run from pipeline before the inference.

        Args:
            begin_index (`int`):
                The begin index for the scheduler.
        Nr{   )rN   r|   s     r   set_begin_indexzHeliosScheduler.set_begin_index   s     (r   c                 4    || j                   j                  z  S NrR   r%   )rN   sigmas     r   _sigma_to_tzHeliosScheduler._sigma_to_t   s    t{{6666r   num_inference_stepsstage_indexdevicerD   muis_amplify_first_chunkc                    | j                   j                  dk(  r|r	|dz  dz   }n|dz   }|| _        | j                          | j                   j                  dk(  r|t        j                  dd| j                   j                  z  |dz         dd j                  t
        j                        }| j                   j                  dk7  r?| j                   j                  rJ | j                  | j                   j                  d|      }|| j                   j                  z  j                         }t        j                  |      }n| j                   |   }t        j                  |d   j#                         |d   j#                         |      }| j$                  |   }	t        j                  |	d   j#                         |	d   j#                         |      }
t        j                  |
      }t        j                  |      j'                  |      | _        t        j*                  |t        j,                  d      g      j'                  |      | _        d| _        | j3                          | j                   j                  dk(  rJ| j(                  dd | _        t        j*                  | j.                  dd	 | j.                  dd g      | _        | j                   j                  r| j                   j                  dk(  sJ | j                  |d| j.                        | _        | j                   j                  dk(  r,| j.                  dd | j                   j                  z  | _        y| j                   |   j5                         | j.                  dd | j                   |   j7                         | j                   |   j5                         z
  z  z   | _        yy)
zA
        Setting the timesteps and sigmas for each stage
        dmdr   r   Nr8   r   r   r   )rR   r3   r   r\   r'   rS   rT   r%   astypefloat32r&   r4   
time_shiftrV   r   rW   r>   rE   r?   rZ   rY   catzerosrD   rL   reset_scheduler_historyr`   r_   )rN   r   r   r   rD   r   r   rY   stage_timestepsrv   ratioss              r   set_timestepszHeliosScheduler.set_timesteps   s.    ;;%%.%&9A&=&A#&9A&=##6 ;;"~QDKK,K,K(KM`cdMdefigijqqJJ ;;$$+#{{????!__T[[->->VLF$++"A"AAGGII%%f-F"66{CO"'')#((*#I  00=L[[a!5!5!7b9I9N9N9PRefF%%f-F)))477v7FiiQ 89<<F<K$$&;;%%.!^^CR0DN))T[["%5t{{237G$HIDK;;++;;$$+++//"c4;;?DK{{!!Q&!%Sb!1DKK4S4S!S!%!9!9+!F!J!J!Lt{{[^\^O_,,[9==?$BZBZ[fBgBkBkBmmP " ,r   r   tc                     | j                   j                  dk(  r| j                  |||      S | j                   j                  dk(  r| j                  |||      S y)a  
        Apply time shifting to the sigmas.

        Args:
            mu (`float`):
                The mu parameter for the time shift.
            sigma (`float`):
                The sigma parameter for the time shift.
            t (`torch.Tensor`):
                The input timesteps.

        Returns:
            `torch.Tensor`:
                The time-shifted timesteps.
        r$   r6   N)rR   r5   _time_shift_exponential_time_shift_linearrN   r   r   r   s       r   r   zHeliosScheduler.time_shift  sW      ;;&&-7//E1==[[((H4**2ua88 5r   c                 p    t        j                  |      t        j                  |      d|z  dz
  |z  z   z  S Nr   )ra   expr   s       r   r   z'HeliosScheduler._time_shift_exponential  s/    xx|txx|q1uqyU.BBCCr   c                 $    ||d|z  dz
  |z  z   z  S r   r   r   s       r   r   z"HeliosScheduler._time_shift_linear  s    R1q519..//r   c                     || j                   }||k(  j                         }t        |      dkD  rdnd}||   j                         S )Nr   r   )rY   nonzerolenrE   )rN   timestepschedule_timestepsindicesposs        r   index_for_timestepz"HeliosScheduler.index_for_timestep!  sL    %!%%1::< w<!#as|  ""r   c                     | j                   Vt        |t        j                        r%|j	                  | j
                  j                        }| j                  |      | _        y | j                  | _        y r   )
r|   re   r   rf   rZ   rY   r   r   rL   rM   )rN   r   s     r   _init_step_indexz HeliosScheduler._init_step_index/  sU    #(ELL1#;;t~~'<'<=#66x@D#00Dr   model_outputr   sample	generator
sigma_nextreturn_dictreturnc                 6   |d u |d u k(  sJ d       |Q|Ot        |t              s4t        |t        j                        st        |t        j                        rt        d      | j                  d| _        |j                  t        j                        }|7|5| j                  | j                     }| j                  | j                  dz      }|||z
  |z  z   }|j                  |j                        }| xj                  dz  c_        |s|fS t        |      S )Nz:sigma and sigma_next must both be None or both be not NonezPassing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to `EulerDiscreteScheduler.step()` is not supported. Make sure to pass one of the `scheduler.timesteps` as a timestep.r   r   )r   )re   r   r   	IntTensor
LongTensor
ValueErrorry   rL   rZ   r   rD   dtyper   )	rN   r   r   r   r   r   r   r   r   s	            r   
step_eulerzHeliosScheduler.step_euler7  s    :#56t8tt6=Z/8S)h8h(8(89 K  ??" D 5==)=Z/KK0ET__q%89J
U 2lBB "nn\%7%78 	A>!$==r   c                     | j                   j                  r d|z
  }t        j                  |d      }||fS d|dz  dz   dz  z  }||z  }||fS )Nr   g:0yE>)r`   r         ?)rR   r2   r   clamp)rN   r   alpha_tsigma_ts       r   _sigma_to_alpha_sigma_tz'HeliosScheduler._sigma_to_alpha_sigma_ti  sb    ;;&&%iGkk%T2G
  E1HqLS01GgoGr   r   r   c                "   t        |      dkD  r|d   n|j                  dd      }|t        |      dkD  r|d   }nt        d      |t        ddd       d	}|d
}| j                  | j
                     }| j                  |      \  }}	| j                  r| j                  j                  dk(  r||	|z  z
  |z  }
n| j                  j                  dk(  r|}
n| j                  j                  dk(  r||z  |	|z  z
  }
nc| j                  j                  dk(  r'|r| j                  | j
                     }	n|}	||	|z  z
  }
n#t        d| j                  j                   d      | j                  j                  r| j                  |
      }
|
S | j                  j                  dk(  r|S | j                  j                  dk(  r|||z  z
  |	z  }|S | j                  j                  dk(  r||z  |	|z  z   }|S t        d| j                  j                   d      )a  
        Convert the model output to the corresponding type the UniPC algorithm needs.

        Args:
            model_output (`torch.Tensor`):
                The direct output from the learned diffusion model.
            timestep (`int`):
                The current discrete timestep in the diffusion chain.
            sample (`torch.Tensor`):
                A current instance of a sample created by the diffusion process.

        Returns:
            `torch.Tensor`:
                The converted model output.
        r   r   Nr   /missing `sample` as a required keyword argumentrY   1.0.0zPassing `timesteps` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`FTepsilonr   v_predictionr!   zprediction_type given as zj must be one of `epsilon`, `sample`, `v_prediction`, or `flow_prediction` for the UniPCMultistepScheduler.zW must be one of `epsilon`, `sample`, or `v_prediction` for the UniPCMultistepScheduler.)r   popr   r
   rD   ry   r   r-   rR   r+   r*   _threshold_sample)rN   r   r   r   argskwargsr   flagr   r   x0_predr   s               r   convert_model_outputz$HeliosScheduler.convert_model_outputs  s+   . "$i!m47J1M>4y1}a !RSS Z =DKK0E77>??{{**i7!Gl$::gE,,8&,,>!F*W|-CC,,0AA"kk$//:G#G 7\#99 /0K0K/L M\ \ 
 {{''009N{{**i7##,,8!Gl$::gE,,>!L07V3CC /0K0K/L MG G r   )r   orderr   r   r   c                	   t        |      dkD  r|d   n|j                  dd      }|t        |      dkD  r|d   }nt        d      |t        |      dkD  r|d   }nt        d      |t        ddd	       | j                  }	| j
                  d
   }
|	d
   }|}| j                  r)| j                  j                  ||
|      j                  }|S |8|6| j                  | j                  dz      | j                  | j                     }}n||}}| j                  |      \  }}| j                  |      \  }}t        j                  |      t        j                  |      z
  }t        j                  |      t        j                  |      z
  }||z
  }|j                  }g }g }t        d|      D ]  }| j                  |z
  }|	|dz       }| j                  | j                  |         \  }}t        j                  |      t        j                  |      z
  }||z
  |z  }|j!                  |       |j!                  ||z
  |z          |j!                  d       t        j"                  ||      }g }g } | j$                  r| n|}!t        j&                  |!      }"|"|!z  dz
  }#d}$| j(                  j*                  dk(  r|!}%n9| j(                  j*                  dk(  rt        j&                  |!      }%n
t-               t        d|dz         D ]T  }|j!                  t        j.                  ||dz
               | j!                  |#|$z  |%z         |$|dz   z  }$|#|!z  d|$z  z
  }#V t        j0                  |      }t        j"                  | |      } t        |      dkD  rt        j0                  |d      }|dk(  r$t        j"                  dg|j2                  |      }&nWt        j4                  j7                  |dd
dd
f   | dd
       j9                  |      j9                  |j2                        }&nd}| j$                  r9||z  |z  ||"z  |z  z
  }'|t        j:                  d&|      }(nd}(|'||%z  |(z  z
  }n8||z  |z  ||"z  |z  z
  }'|t        j:                  d&|      }(nd}(|'||%z  |(z  z
  }|j9                  |j2                        }|S )a  
        One step for the UniP (B(h) version). Alternatively, `self.solver_p` is used if is specified.

        Args:
            model_output (`torch.Tensor`):
                The direct output from the learned diffusion model at the current timestep.
            prev_timestep (`int`):
                The previous discrete timestep in the diffusion chain.
            sample (`torch.Tensor`):
                A current instance of a sample created by the diffusion process.
            order (`int`):
                The order of UniP at this timestep (corresponds to the *p* in UniPC-p).

        Returns:
            `torch.Tensor`:
                The sample tensor at the previous timestep.
        r   prev_timestepNr   r   r   .missing `order` as a required keyword argumentr   zPassing `prev_timestep` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`r8   r   r   r9   r"   dimr   r   r   k,bkc...->bc...)r   r   r   r
   r   rJ   r1   stepr   rD   ry   r   r   logr   r^   rc   tensorr-   expm1rR   r.   rH   powstackr   linalgsolverZ   einsum))rN   r   r   r   r   r   r   r   r   model_output_lists0m0xx_tr   sigma_s0r   alpha_s0lambda_t	lambda_s0hr   rksD1sisimialpha_sisigma_si	lambda_sirkRbhhh_phi_1h_phi_kfactorial_iB_hrhos_px_t_pred_ress)                                            r   multistep_uni_p_bh_updatez)HeliosScheduler.multistep_uni_p_bh_update  s   6 $'t9q=QfjjRV6W>4y1}a !RSS=4y1}Q !QRR$ ^
 !..#r"==--$$\2q9EECJ%- $DOOa,? @$++dooB^XG *EXG77@!99(C(99W%		'(::IIh'%))H*==	y q%A1$B"QU8,B!%!=!=dkk"o!NHh		(+eii.AAIi'1,BJJrNJJR2~& ! 	

3ll3v.??aR++b/B,";;""e+C[[$$-++b/C%''q%!)$AHHUYYsAE*+HHW{*S011q5 KlQ_4G	 % KKNLL6*s8a<++cq)Czse1776J++Acrc3B3hK3B@CCFKNNqwwWC??X%)Gg,=,BBD <<(963G311CX%)Gg,=,BBD <<(963G311CffQWWo
r   )r   this_sampler   sigma_beforer   this_model_outputr   r   r   c                	   t        |      dkD  r|d   n|j                  dd      }	|t        |      dkD  r|d   }nt        d      |t        |      dkD  r|d   }nt        d      |t        |      dkD  r|d   }nt        d	      |	t        dd
d       | j                  }
|
d   }|}|}|}|8|6| j
                  | j                     | j
                  | j                  dz
     }}n||}}| j                  |      \  }}| j                  |      \  }}t        j                  |      t        j                  |      z
  }t        j                  |      t        j                  |      z
  }||z
  }|j                  }g }g }t        d|      D ]  }| j                  |dz   z
  }|
|dz       }| j                  | j
                  |         \  }}t        j                  |      t        j                  |      z
  }||z
  |z  }|j                  |       |j                  ||z
  |z          |j                  d       t        j                  ||      }g } g }!| j                  r| n|}"t        j                  |"      }#|#|"z  dz
  }$d}%| j                   j"                  dk(  r|"}&n9| j                   j"                  dk(  rt        j                  |"      }&n
t%               t        d|dz         D ]T  }| j                  t        j&                  ||dz
               |!j                  |$|%z  |&z         |%|dz   z  }%|$|"z  d|%z  z
  }$V t        j(                  |       } t        j                  |!|      }!t        |      dkD  rt        j(                  |d      }nd}|dk(  r$t        j                  dg|j*                  |      }'nHt        j,                  j/                  | |!      j1                  |      j1                  |j*                        }'| j                  rJ||z  |z  ||#z  |z  z
  }(|t        j2                  d|'dd |      })nd})||z
  }*|(||&z  |)|'d   |*z  z   z  z
  }nI||z  |z  ||#z  |z  z
  }(|t        j2                  d|'dd |      })nd})||z
  }*|(||&z  |)|'d   |*z  z   z  z
  }|j1                  |j*                        }|S )a  
        One step for the UniC (B(h) version).

        Args:
            this_model_output (`torch.Tensor`):
                The model outputs at `x_t`.
            this_timestep (`int`):
                The current timestep `t`.
            last_sample (`torch.Tensor`):
                The generated sample before the last predictor `x_{t-1}`.
            this_sample (`torch.Tensor`):
                The generated sample after the last predictor `x_{t}`.
            order (`int`):
                The `p` of UniC-p at this step. The effective order of accuracy should be `order + 1`.

        Returns:
            `torch.Tensor`:
                The corrected sample tensor at the current timestep.
        r   this_timestepNr   z4missing `last_sample` as a required keyword argumentr   z4missing `this_sample` as a required keyword argumentr   r   r   zPassing `this_timestep` is deprecated and has no effect as model output conversion is now handled via an internal counter `self.step_index`r8   r   r   r9   r"   r   r   r   r   )r   r   r   r
   r   rD   ry   r   r   r   r   r^   rc   r   r-   r   rR   r.   rH   r   r   r   r   r   rZ   r   )+rN   r   r   r   r   r   r   r   r   r   r   r   r   r   model_tr   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   rhos_cr   corr_resD1_ts+                                              r   multistep_uni_c_bh_updatez)HeliosScheduler.multistep_uni_c_bh_updateI  s   < $'t9q=QfjjRV6W4y1}"1g !WXX4y1}"1g !WXX=4y1}Q !QRR$ ^ !..r"#EM $DOO <dkk$//\]J]>^XG %|XG77@!99(C(99W%		'(::IIh'%))H*==	y ##q%AAE*B"QU8,B!%!=!=dkk"o!NHh		(+eii.AAIi'1,BJJrNJJR2~& ! 	

3ll3v.??aR++b/B,";;""e+C[[$$-++b/C%''q%!)$AHHUYYsAE*+HHW{*S011q5 KlQ_4G	 % KKNLL6*s8a<++cq)CC A:\\3%qwwvFF\\''1-008;;AGGDF??X%)Gg,=,BBD <<(96#2;LR<D3(VBZ$5F*FGGCX%)Gg,=,BBD <<(96#2;LR<D3(VBZ$5F*FGGCffQWWo
r   r   rJ   cus_step_indexcus_lower_order_numcus_this_ordercus_last_samplec                    | j                   t        d      |
| j                  d| _        n|
| _        ||| _        ||| _        ||| _        | j                  dkD  xr+ | j                  dz
  | j                  vxr | j                  d u}| j                  |||      }|||d d | _	        |d d | _
        |r+| j                  || j                  || j
                  ||      }||||d<   |dd  | _	        |dd  | _
        nt        | j                  j                  dz
        D ]@  }| j                  |dz      | j                  |<   | j                  |dz      | j                  |<   B || j                  d<   || j                  d<   | j                  j                  rAt!        | j                  j                  t#        | j$                        | j                  z
        }n| j                  j                  }t!        || j                  dz         | _        | j
                  dkD  sJ || _        | j'                  ||| j
                  ||	      }|8| j                  | j                  j                  k  r| xj                  dz  c_        |
| xj                  dz  c_        |s||| j                  | j
                  fS t)        ||| j                  | j
                        S )	NzaNumber of inference steps is 'None', you need to run 'set_timesteps' after creating the schedulerr   r   r   r8   )r   r   r   r   r   r   )r   r   r   r   r   )r   r   r   r   )r   r   ry   rL   rK   r   r   r0   r   r   rJ   r   r^   rR   r,   r/   r`   r   rY   r   r   )rN   r   r   r   r   r   rJ   r   r   r   r  r  r  r  use_correctormodel_output_convertr   r   r   s                      r   
step_unipczHeliosScheduler.step_unipc  s     ##+s  !&#$ -D*$7D!%,DO&.D OOavDOOa$7t?U?U$UvZ^ZjZjrvZv 	
  $88f\a8b$)B!.s!3D!.s!3D33"6 ,,"oo) 4 F $)B 4M"!.qr!2D!.qr!2D4;;33a78(,(:(:1q5(A""1%(,(:(:1q5(A""1% 9 &:Dr"%-Dr";;((T[[55s4>>7JT__7\]J11Jj$*?*?!*CD"""!44%//! 5 
 &$$t{{'?'??%%*% !!0@0@$//RR$#'((	
 	
r   c                     | j                   j                  dk(  r| j                  |||||      S | j                   j                  dk(  r| j                  ||||      S t        )Neuler)r   r   r   r   r   r#   )r   r   r   r   )rR   r3   r   r  rH   )rN   r   r   r   r   r   s         r   r   zHeliosScheduler.step>  sy     ;;%%0??)!#' #   [[''72??)!'	 #   &%r   c                 $   d g| j                   j                  z  | _        d g| j                   j                  z  | _        d| _        | j                   j
                  | _        | j                   j                  | _        d | _        d | _        d | _	        y )Nr   )
rR   r,   r   rJ   rK   r0   r1   r   rL   rM   rx   s    r   r   z'HeliosScheduler.reset_scheduler_historyX  sw    "Vdkk&>&>>"Vdkk&>&>> !!%!>!>,, r   c                 .    | j                   j                  S r   r   rx   s    r   __len__zHeliosScheduler.__len__b  s    {{...r   )r   )NNNNFr   )NNNNNT)NNTNNNNNNNNN)NNNT)+r   r   r   _compatiblesr   r   r   floatlistboolstrr   r   rO   r\   rC   propertyry   r|   r~   r   r   r   r   rf   r   r   r   r   r   r   	Generatorr   tupler   r   r   r   r   r  r   r   r  r   r   r   r   r   #   s    LE $(0"0 "&')#' $%%*<I'0! 0! 0! 	0!
 0! 0! 0! 0! 0! 0! 0!  0!  90! !0!  !0!" #0!$ #%0!& !!89'0! 0!d'$:Mx     ! !(3 (7 #'%)"',= = 4Z= ell"	=
 t= 4K= !%=@9U 95 9U\\ 9,D0#1 /3$(,0*./3 />''/> %+++/> !!	/>
 ??T)/>   4'/> %%,/> /> 
	&/>d   $"NllN 	N
 ||N 
Nh  $"#'DllD 	D
 D ||D LLD 
DT %)$(%)"L <<L \\	L
 \\L L llL ||L 
Lb (,# ""%)"#'"#'"(,d
lld
 $d
 	d

 d
 d
 d
 lld
 ||d
 LLd
 d
 !d
 d
 d
 
	&d
T /3$(,0 &''& %+++& !!	&
 ??T)& & 
	&&4!/r   r   )ra   dataclassesr   typingr   numpyrS   r   configuration_utilsr   r   schedulers.scheduling_utilsr   utilsr	   r
   r   r   r   r   r   <module>r     sJ     !    A 8 ) "J " "@/nk @/r   