
    ^j                     `    d dl Z d dlZd dlmZ d dlmZ d dlmZmZm	Z	 g dZ
d Zd Zd Zd	 Zy)
    N)deepcopy)nn)
Conv2dSameBatchNormAct2dLinear)extract_layer	set_layeradapt_model_from_stringadapt_model_from_filec                 "   |j                  d      }| }t        | d      r|d   dk7  r| j                  }t        | d      s|d   dk(  r|dd }|D ]=  }t        ||      r,|j                         st	        ||      },|t        |         };|c S  |S )zExtract a layer from a model using dot-separated path.

    Args:
        model: PyTorch model.
        layer: Dot-separated layer path (e.g., 'layer1.0.conv1').

    Returns:
        Extracted module.
    .moduler      N)splithasattrr   isdigitgetattrint)modellayerr   ls       ]/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/timm/models/_prune.pyr   r      s     KKEFuhE!H$85(#aH(<ab	 6199; +AM M    c                    |j                  d      }| }t        | d      r|d   dk7  r| j                  }d}|}|D ]?  }t        ||      s|j                         st	        ||      }n|t        |         }|dz  }A |dz  }|d| D ]-  }|j                         st	        ||      } |t        |         }/ ||   }t        |||       y)zSet a layer in a model using dot-separated path.

    Args:
        model: PyTorch model.
        layer: Dot-separated layer path.
        val: New value for the layer.
    r   r   r   r   N)r   r   r   r   r   r   setattr)r   r   valr   	lst_indexmodule2r   s          r   r	   r	   '   s     KKEFuhE!H$8IG 7A99;!'1-!#a&/NI NI:I $yy{VQ'FCF^F	$
 	iAFAsr   c                    d}i }|j                  |      }|D ]T  }|j                  d      }|d   }|d   dd j                  d      }|d   dk7  s9|D cg c]  }t        |       c}||<   V t        | j                               j                  }	t        | j                               j
                  }
|	|
d}t        |       }| j                         D ]P  \  }}t        | |      }t        |t        j                        st        |t              rt        |t              rt        }nt        j                  }||d	z      }|d   }|d   }d}|j                  dkD  r|}|} |d|||j                  |j                  d
u|j                   |j"                  ||j$                  d|}t'        |||       t        |t(              rit)        ||d	z      d   f|j*                  |j,                  |j.                  dd|}|j0                  |_        |j2                  |_        t'        |||       Wt        |t        j4                        rQt        j4                  d||d	z      d   |j*                  |j,                  |j.                  dd|}t'        |||       t        |t        j6                        s||d	z      d   }t7        d||j8                  |j                  d
ud|}t'        |||       t;        |d      s)t=        |dd      |j>                  k(  r||_         ||_        S |jC                          | jC                          |S c c}w )a  Adapt a model to pruned structure from string specification.

    Args:
        parent_module: Original model to adapt.
        model_string: String containing layer shapes for pruned model.

    Returns:
        Adapted model with pruned layer dimensions.
    z***:r   r   , )devicedtypez.weightN)in_channelsout_channelskernel_sizebiaspaddingdilationgroupsstrideT)epsmomentumaffinetrack_running_stats)num_featuresr.   r/   r0   r1   )in_featuresout_featuresr)   r2   head_hidden_size )"r   r   next
parametersr$   r%   r   named_modulesr   
isinstancer   Conv2dr   r,   r(   r)   r*   r+   r-   r	   r   r.   r/   r0   dropactBatchNorm2dr   r4   r   r   r2   r5   eval)parent_modulemodel_string	separator
state_dict	lst_shapekkeyshapeir$   r%   dd
new_modulenm
old_moduleconvsr&   r'   gnew_convnew_bnr2   new_fcs                            r   r
   r
   F   sN    IJ""9-I 6GGCLd!Qr
  %8r>/45!s1v5JsO6 -**,-44F))+,22EU	+B-(J++- =71"=!4
j")),
:z0R*j1!yy1y=)AA$KQ4LA  1$* 
')&22__D0"**#,,!((
 
H j!X.
N3#1y=)!,NN#,,!(($( F %//FK#FJj!V,
BNN3^^ 'I6q9NN#,,!(($( F j!V,
BII.%a)m4Q7L ('44__D0 	F j!V,z>2:'91=AXAXX2>J/*6
'{=7~ OOU 6s   Mc                     t        j                  t        t        j                  j                  d|dz               }t        | |j                  d      j                               S )zAdapt a model to pruned structure from file specification.

    Args:
        parent_module: Original model to adapt.
        model_variant: Name of pruned model variant file.

    Returns:
        Adapted model with pruned layer dimensions.
    _prunedz.txtzutf-8)	pkgutilget_data__name__ospathjoinr
   decodestrip)r@   model_variant
adapt_datas      r   r   r      sL     !!(BGGLLMTZDZ,[\J"=*2C2CG2L2R2R2TUUr   )rY   rV   copyr   torchr   timm.layersr   r   r   __all__r   r	   r
   r   r6   r   r   <module>rd      s3    	    : :
\6>\~Vr   