a
    d[                     @   sH   d dl Z d dlmZ d dlmZ ddlmZ e G dd deZ	dS )    N)
functional)MODEL_REGISTRY   )SRModelc                   @   s   e Zd Zdd ZdS )SwinIRModelc           	      C   sZ  | j d d }| j dd}d\}}| j \}}}}|| dkrN|||  }|| dkrf|||  }t| jd|d|fd}t| dr| j  t	
  | || _W d    n1 s0    Y  nH| j  t	
  | || _W d    n1 s0    Y  | j  | j \}}}}| jd d d d d|||  d|||  f | _d S )	NZ	network_gwindow_sizescaler   )r   r   r   Zreflect	net_g_ema)optgetZlqsizeFpadhasattrr	   evaltorchZno_gradoutputZnet_gZtrain)	selfr   r   Z	mod_pad_hZ	mod_pad_w_hwimg r   d/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/basicsr/models/swinir_model.pytest   s&    


,

*
zSwinIRModel.testN)__name__
__module____qualname__r   r   r   r   r   r      s   r   )
r   Ztorch.nnr   r   Zbasicsr.utils.registryr   Zsr_modelr   registerr   r   r   r   r   <module>   s
   