a
    Ú
þd,  ã                   @   sŽ   d dl mZ d dl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lZd dlZG dd	„ d	ƒZddd„Zddd„ZdS )é    )Úabsolute_importN)Únn)ÚOrderedDict)ÚVariable)Úzoom)Útqdmc                   @   s¸   e Zd Zdd„ Zddddddddddd	d
ddgfdd„Zd-dd„Zdd„ Zdd„ Zdd„ Zdd„ Z	dd„ Z
dd„ Zdd„ Zdd „ Zd!d"„ Zd#d$„ Zd%d&„ Zd'd(„ Zd)d*„ Zd.d+d,„ZdS )/ÚTrainerc                 C   s   | j S ©N)Ú
model_name©Úself© r   úV/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/lpips/trainer.pyÚname   s    zTrainer.nameÚlpipsZalexZLabFNTg-Cëâ6?ç      à?z0.1r   c                 C   sª  || _ || _|| _|| _|
| _|	| _d||f | _| jdkr`tj|
 ||d|	||d|dd
| _np| jdkr~tj||dd| _nR| jdv r tj	||d	| _d
| _n0| jdv rÂtj
||d	| _d| _ntd| j ƒ‚t| j ¡ ƒ| _| jr4t ¡ | _|  jt| jj ¡ ƒ7  _|| _|| _tjj| j||dfd| _n
| j ¡  |r„| j |d ¡ tjj| j|d| _| jr„| jj|d d| _|r¦tdƒ t | j¡ tdƒ dS )aÚ  
        INPUTS
            model - ['lpips'] for linearly calibrated network
                    ['baseline'] for off-the-shelf network
                    ['L2'] for L2 distance in Lab colorspace
                    ['SSIM'] for ssim in RGB colorspace
            net - ['squeeze','alex','vgg']
            model_path - if None, will look in weights/[NET_NAME].pth
            colorspace - ['Lab','RGB'] colorspace to use for L2 and SSIM
            use_gpu - bool - whether or not to use a GPU
            printNet - bool - whether or not to print network architecture out
            spatial - bool - whether to output an array containing varying distances across spatial dimensions
            is_train - bool - [True] for training mode
            lr - float - initial learning rate
            beta1 - float - initial momentum term for adam
            version - 0.1 for latest, 0.0 was original (with a bug)
            gpu_ids - int array - [0] by default, gpus to use
        z%s [%s]r   TF)
Z
pretrainedÚnetÚversionr   ÚspatialÚ	pnet_randÚ	pnet_tuneZuse_dropoutÚ
model_pathZ	eval_modeZbaseline)r   r   r   )ÚL2Úl2)Úuse_gpuÚ
colorspacer   )ÚDSSIMZdssimÚSSIMZssimr   zModel [%s] not recognized.g+‡ÙÎ÷ï?)ÚlrZbetasr   )Z
device_ids©Zdevicez----------- Networks initialized -------------z/-----------------------------------------------N)r   Úgpu_idsÚmodelr   Úis_trainr   r
   r   ZLPIPSr   r   Ú
ValueErrorÚlistÚ
parametersZBCERankingLossÚrankLossr   Úold_lrÚtorchZoptimZAdamÚoptimizer_netÚevalÚtor   ZDataParallelÚprintZnetworksZprint_network)r   r!   r   r   r   r   r   r   ZprintNetr   r"   r   Zbeta1r   r    r   r   r   Ú
initialize   sL    
þ





zTrainer.initializec                 C   s   | j j|||dS )zö Function computes the distance between image patches in0 and in1
        INPUTS
            in0, in1 - torch.Tensor object of shape Nx3xXxY - image patch scaled to [-1,1]
        OUTPUT
            computed distances between in0 and in1
        )ÚretPerLayer)r   Úforward)r   Zin0Zin1r.   r   r   r   r/   V   s    zTrainer.forwardc                 C   s0   |   ¡  | j ¡  |  ¡  | j ¡  |  ¡  d S r	   )Úforward_trainr)   Z	zero_gradÚbackward_trainÚstepÚclamp_weightsr   r   r   r   Úoptimize_parametersa   s
    

zTrainer.optimize_parametersc                 C   s>   | j  ¡ D ].}t|dƒr
|jdkr
tj|jjdd|j_q
d S )NÚweight)é   r6   r   )Úmin)r   ÚmodulesÚhasattrZkernel_sizer(   Úclampr5   Údata)r   Úmoduler   r   r   r3   h   s    zTrainer.clamp_weightsc                 C   sº   |d | _ |d | _|d | _|d | _| jr†| j j| jd d| _ | jj| jd d| _| jj| jd d| _| jj| jd d| _t| j dd| _t| jdd| _	t| jdd| _
d S )	NÚrefÚp0Úp1Újudger   r   T)Zrequires_grad)Z	input_refZinput_p0Zinput_p1Úinput_judger   r+   r    r   Úvar_refÚvar_p0Úvar_p1)r   r;   r   r   r   Ú	set_inputm   s    



zTrainer.set_inputc                 C   s|   |   | j| j¡| _|   | j| j¡| _|  | j| j| j¡| _t	d| j ƒ 
| j ¡ ¡| _| j  | j| j| jd d ¡| _| jS )Nç      ð?g       @)r/   rB   rC   Úd0rD   Úd1Úcompute_accuracyrA   Úacc_rr   ÚviewÚsizeZ	var_judger&   Ú
loss_totalr   r   r   r   r0   }   s     zTrainer.forward_trainc                 C   s   t  | j¡ ¡  d S r	   )r(   ÚmeanrM   Zbackwardr   r   r   r   r1   ˆ   s    zTrainer.backward_trainc                 C   s>   ||k   ¡ j ¡  ¡ }|  ¡  ¡  ¡ }|| d| d|   S )z) d0, d1 are Variables, judge is a Tensor r6   )Úcpur;   ÚnumpyÚflatten)r   rG   rH   r@   Zd1_lt_d0Z	judge_perr   r   r   rI   ‹   s    zTrainer.compute_accuracyc                 C   sF   t d| jj ¡  ¡ fd| jfgƒ}| ¡ D ]}t || ¡||< q*|S )NrM   rJ   )	r   rM   r;   rO   rP   rJ   ÚkeysÚnprN   )r   ZretDictÚkeyr   r   r   Úget_current_errors‘   s    ÿzTrainer.get_current_errorsc                 C   s”   d| j j ¡ d  }t | j j¡}t | jj¡}t | jj¡}t|||dgdd}t|||dgdd}t|||dgdd}td|fd|fd|fgƒS )	Né   é   r6   r   )Úorderr=   r>   r?   )	rB   r;   rL   r   Z	tensor2imrC   rD   r   r   )r   Zzoom_factorZref_imgZp0_imgZp1_imgZref_img_visZ
p0_img_visZ
p1_img_visr   r   r   Úget_current_visualsš   s    þzTrainer.get_current_visualsc                 C   sF   | j r|  | jj|d|¡ n|  | j|d|¡ |  | jj|d|¡ d S )NÚ Zrank)r   Úsave_networkr   r<   r&   )r   ÚpathÚlabelr   r   r   Úsave©   s    zTrainer.savec                 C   s.   d||f }t j ||¡}t | ¡ |¡ d S )Nú%s_net_%s.pth)Úosr\   Újoinr(   r^   Z
state_dict)r   Únetworkr\   Únetwork_labelÚepoch_labelÚsave_filenameÚ	save_pathr   r   r   r[   ±   s    zTrainer.save_networkc                 C   s<   d||f }t j | j|¡}td| ƒ | t |¡¡ d S )Nr_   zLoading network from %s)r`   r\   ra   Úsave_dirr,   Zload_state_dictr(   Úload)r   rb   rc   rd   re   rf   r   r   r   Úload_network·   s    zTrainer.load_networkc                 C   sH   | j | }| j| }| jjD ]}||d< qtdt| j|f ƒ || _d S )Nr   zupdate lr [%s] decay: %f -> %f)r   r'   r)   Zparam_groupsr,   Útype)r   Znepoch_decayZlrdr   Zparam_groupr   r   r   Úupdate_learning_rate½   s    


zTrainer.update_learning_ratec                 C   s   | j S r	   )Zimage_pathsr   r   r   r   Úget_image_pathsÈ   s    zTrainer.get_image_pathsc                 C   s:   t  tj | jd¡|¡ t jtj | jd¡|gdd d S )NZ	done_flagz%i)Úfmt)rS   r^   r`   r\   ra   rg   Zsavetxt)r   Úflagr   r   r   Ú	save_doneË   s    zTrainer.save_done)F)F)Ú__name__Ú
__module__Ú__qualname__r   r-   r/   r4   r3   rE   r0   r1   rI   rU   rY   r^   r[   ri   rk   rl   ro   r   r   r   r   r      s(   þ
C
	r   rZ   c                 C   sî   g }g }g }t |  ¡ |dD ]p}|||d |d ƒj ¡  ¡  ¡  ¡ 7 }|||d |d ƒj ¡  ¡  ¡  ¡ 7 }||d  ¡  ¡  ¡  ¡ 7 }qt |¡}t |¡}t |¡}||k d|  ||k |  ||kd  }t 	|¡t
||||dfS )	a   Function computes Two Alternative Forced Choice (2AFC) score using
        distance function 'func' in dataset 'data_loader'
    INPUTS
        data_loader - CustomDatasetDataLoader object - contains a TwoAFCDataset inside
        func - callable distance function - calling d=func(in0,in1) should take 2
            pytorch tensors with shape Nx3xXxY, and return numpy array of length N
    OUTPUTS
        [0] - 2AFC score in [0,1], fraction of time func agrees with human evaluators
        [1] - dictionary with following elements
            d0s,d1s - N arrays containing distances between reference patch to perturbed patches 
            gts - N array in [0,1], preferred patch selected by human evaluators
                (closer to "0" for left patch p0, "1" for right patch p1,
                "0.6" means 60pct people preferred right patch, 40pct preferred left)
            scores - N array in [0,1], corresponding to what percentage function agreed with humans
    CONSTS
        N - number of test triplets in data_loader
    ©Údescr=   r>   r?   r@   rF   r   )Úd0sÚd1sÚgtsÚscores)r   Ú	load_datar;   rO   rP   rQ   ÚtolistrS   ÚarrayrN   Údict)Údata_loaderÚfuncr   ru   rv   rw   r;   rx   r   r   r   Úscore_2afc_datasetÐ   s    ((


(r   c                 C   sæ   g }g }t |  ¡ |dD ]D}|||d |d ƒj ¡  ¡  ¡ 7 }||d  ¡  ¡  ¡  ¡ 7 }qt |¡}t |¡}t 	|¡}|| }|| }	t 
|	¡}
t 
d|	 ¡}t |	¡|
 }|
|
|  }|
|
|  }t ||¡}|t||dfS )aê   Function computes JND score using distance function 'func' in dataset 'data_loader'
    INPUTS
        data_loader - CustomDatasetDataLoader object - contains a JNDDataset inside
        func - callable distance function - calling d=func(in0,in1) should take 2
            pytorch tensors with shape Nx3xXxY, and return pytorch array of length N
    OUTPUTS
        [0] - JND score in [0,1], mAP score (area under precision-recall curve)
        [1] - dictionary with following elements
            ds - N array containing distances between two patches shown to human evaluator
            sames - N array containing fraction of people who thought the two patches were identical
    CONSTS
        N - number of test triplets in data_loader
    rs   r>   r?   Zsamer6   )ÚdsÚsames)r   ry   r;   rO   rP   rz   rQ   rS   r{   ZargsortZcumsumÚsumr   Zvoc_apr|   )r}   r~   r   r€   rw   r;   r   Zsorted_indsZ	ds_sortedZsames_sortedZTPsZFPsZFNsZprecsZrecsZscorer   r   r   Úscore_jnd_datasetó   s"    $



rƒ   )rZ   )rZ   )Ú
__future__r   rP   rS   r(   r   Úcollectionsr   Ztorch.autogradr   Zscipy.ndimager   r   r   r`   r   r   rƒ   r   r   r   r   Ú<module>   s    B
#