a
    d6                     @   sx   d dl Z d dlZd dlZddlmZmZ ddlmZmZm	Z	 e	ddrXd dl
m  mZ dd ZG d	d
 d
ejjZdS )    N   )AcceleratorStateGradientState)DistributedType
honor_typeis_tpu_availableF)Zcheck_devicec                    sh   t | ttfr&t|  fdd| D S t | trNt|  fdd|  D S t | tjrd| 	 S | S )Nc                 3   s   | ]}t | V  qd S Nmove_to_device).0tdevice ]/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/accelerate/optimizer.py	<genexpr>       z!move_to_device.<locals>.<genexpr>c                    s   i | ]\}}|t | qS r   r	   )r   kvr   r   r   
<dictcomp>    r   z"move_to_device.<locals>.<dictcomp>)

isinstancelisttupler   dicttypeitemstorchZTensorto)stater   r   r   r   r
      s    

r
   c                   @   s   e Zd ZdZd#ddZedd Zejdd Zed	d
 Zejdd
 Zedd Z	e	jdd Z	dd Z
dd Zdd Zd$ddZd%ddZdd Zedd Zedd Zdd  Zd!d" ZdS )&AcceleratedOptimizera  
    Internal wrapper around a torch optimizer.

    Conditionally will perform `step` and `zero_grad` if gradients should be synchronized when performing gradient
    accumulation.

    Args:
        optimizer (`torch.optim.optimizer.Optimizer`):
            The optimizer to wrap.
        device_placement (`bool`, *optional*, defaults to `True`):
            Whether or not the optimizer should handle device placement. If so, it will place the state dictionary of
            `optimizer` on the right device.
        scaler (`torch.cuda.amp.grad_scaler.GradScaler`, *optional*):
            The scaler to use in the step function if training with mixed precision.
    TNc                 C   sz   || _ || _t | _t | _|| _d| _d | _|rv| j 	 }| jj
tjkr\t|| jj nt|| jj}| j | d S )NF)	optimizerscalerr   accelerator_stater   gradient_statedevice_placement_is_overflow_last_scale
state_dictdistributed_typer   TPUxmsend_cpu_data_to_devicer   r
   load_state_dict)selfr    r$   r!   r'   r   r   r   __init__7   s    
zAcceleratedOptimizer.__init__c                 C   s   | j jS r   r    r   r-   r   r   r   r   I   s    zAcceleratedOptimizer.statec                 C   s   || j _d S r   r/   r-   r   r   r   r   r   M   s    c                 C   s   | j jS r   r    param_groupsr0   r   r   r   r3   Q   s    z!AcceleratedOptimizer.param_groupsc                 C   s   || j _d S r   r2   )r-   r3   r   r   r   r3   U   s    c                 C   s   | j jS r   r    defaultsr0   r   r   r   r5   Y   s    zAcceleratedOptimizer.defaultsc                 C   s   || j _d S r   r4   )r-   r5   r   r   r   r5   ]   s    c                 C   s   | j | d S r   )r    add_param_group)r-   param_groupr   r   r   r6   a   s    z$AcceleratedOptimizer.add_param_groupc                 C   s4   | j jtjkr$| jr$t|| j j | j	| d S r   )
r"   r(   r   r)   r$   r*   r+   r   r    r,   )r-   r'   r   r   r   r,   d   s    z$AcceleratedOptimizer.load_state_dictc                 C   s
   | j  S r   )r    r'   r0   r   r   r   r'   i   s    zAcceleratedOptimizer.state_dictc                 C   sZ   | j jrVdt| jjjv }|r<|d u r,d}| jj|d n|d urLtd| j  d S )Nset_to_noneF)r8   zJ`set_to_none` for Optimizer.zero_grad` is not supported by this optimizer.)r#   sync_gradientsinspect	signaturer    	zero_grad
parameters
ValueError)r-   r8   Z
accept_argr   r   r   r<   l   s    zAcceleratedOptimizer.zero_gradc                 C   s   | j jr| jjtjkr<|d ur&d|ini }tj| j|d np| j	d urd}| j
d u rd| j	 | _
d}| j	| j| | j	  | j	 }|s|| j
k | _|| _
n| j| d S )Nclosure)optimizer_argsFT)r#   r9   r"   r(   r   r)   r*   Zoptimizer_stepr    r!   r&   Z	get_scalestepupdater%   )r-   r?   r@   Z	new_scaleZscale_afterr   r   r   rA   x   s     



zAcceleratedOptimizer.stepc                    s,   | j jD ]} fdd|d D |d< qd S )Nc                    s   g | ]}  ||qS r   )get)r   pparameters_mapr   r   
<listcomp>   r   z;AcceleratedOptimizer._switch_parameters.<locals>.<listcomp>paramsr2   )r-   rF   r7   r   rE   r   _switch_parameters   s    z'AcceleratedOptimizer._switch_parametersc                 C   s   t dt | jS )zTWhether or not the optimizer step was done, or skipped because of gradient overflow.zThe `is_overflow` property is deprecated and will be removed in version 1.0 of Accelerate use `optimizer.step_was_skipped` instead.)warningswarnFutureWarningr%   r0   r   r   r   is_overflow   s
    z AcceleratedOptimizer.is_overflowc                 C   s   | j S )z.Whether or not the optimizer step was skipped.)r%   r0   r   r   r   step_was_skipped   s    z%AcceleratedOptimizer.step_was_skippedc                 C   s
   | j  S r   )__dict__copyr0   r   r   r   __getstate__   s    z!AcceleratedOptimizer.__getstate__c                 C   s   | j | d S r   )rO   rB   r1   r   r   r   __setstate__   s    z!AcceleratedOptimizer.__setstate__)TN)N)N)__name__
__module____qualname____doc__r.   propertyr   setterr3   r5   r6   r,   r'   r<   rA   rI   rM   rN   rQ   rR   r   r   r   r   r   &   s4   









	
r   )r:   rJ   r   r   r   r   utilsr   r   r   Ztorch_xla.core.xla_modelcoreZ	xla_modelr*   r
   ZoptimZ	Optimizerr   r   r   r   r   <module>   s   

