+
    &jH                     d   ^ RI t ^ RIt^ RIt^ RIt^ RIHt ^ RIHt ^ RIH	t	 ^ RI
HtHt ^ RIt^ RIHt ^ RIHt ^ RIHu Ht ^ RIHu Ht ^ RIHtHt ^ RIHt ^ RIHtH t  . R#Ot!R R lt"R R lt#R R lt$R$R R llt%R R lt&R R lt']PP                  ]PR                  ]PT                  ]PV                  ]PX                  ]PZ                  ]P\                  ]P^                  ]P`                  ]Pb                  ]P^                  ]Pd                  ]Pf                  .t4]Pj                  ]Pl                  .t7]PP                  ]Pp                  ]PR                  ]Pr                  ]PT                  R /t:R R lt;R R lt< ! R R	4      t=R%R lt>R R lt? ! R  R
4      t@R]P                  3R! R" lltBR# )&    N)defaultdict)Iterable)Enum)Anycast)ArgumentTarget)	ShapeProp)fuse_conv_bn_evalfuse_linear_bn_evalMklSubgraph	UnionFindc                R    V ^8  d   QhR\         R\        \         \         3,          /# )   targetreturn)strtuple)formats   "z/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/torch/fx/experimental/optimization.py__annotate__r   %   s"     - - -sCx -    c                X    V P                  R^4      Ev rV'       d   V^ ,          V3# RV3# )zd
Splits a qualname into parent path and last atom.
For example, `foo.bar.baz` -> (`foo.bar`, `baz`)
. )rsplit)r   parentnames   &  r   _parent_namer   %   s1    
 MM#q)MV6!9,,B,,r   c                    V ^8  d   QhR\         \        ,          R\        P                  R\        \
        \        3,          /# )r   patternnodemodules)r   typefxNodedictr   r   )r   s   "r   r   r   /   s4      d^#%7759#s(^r   c                    \        VP                  4      ^ 8X  d   R# VP                  ^ ,          V3p\        W4       F  w  rE\        V\        P
                  4      '       g    R# VP                  R8w  d    R# \        VP                  \        4      '       g    R# VP                  V9  d    R# \        W%P                  ,          4      VJg   K   R# 	  R# )r   Fcall_moduleT)
lenargszip
isinstancer%   r&   opr   r   r$   )r!   r"   r#   nodesexpected_typecurrent_nodes   &&&   r   matches_module_patternr2   /   s     499~"&))A,!5E'*7':#,00??m+,--s33g-++,-]B (; r   c                    V ^8  d   QhR\         P                  R\        \        \        3,          R\
        P                  P                  /# )r   r"   r#   
new_module)r%   r&   r'   r   r   torchnnModule)r   s   "r   r   r   C   s8     4 4
''4 cN48=4r   c                     \        V P                  \        4      '       g"   \        R \	        V P                  4       24      h\        V P                  4      w  r4W!V P                  &   \        W,          WB4       R# )Expected str target, got N)r-   r   r   AssertionErrorr$   r   setattr)r"   r#   r4   parent_namer   s   &&&  r   replace_node_moduler=   C   s\     dkk3''8dkk9J8KLMM$T[[1K%DKKG $3r   c                    V ^8  d   QhR\         P                  P                  R\         P                  P                  /# r   modelr   )r5   r6   r7   )r   s   "r   r   r   M   s*     %/ %/ %/588?? %/r   c                   \         P                  \         P                  3\         P                  \         P                  3\         P
                  \         P                  3\         P                  \         P                  3.pV'       g   \        P                  ! V 4      p V'       d+   \        V \        P                  P                  4      '       g   \        P                  ! V 4      pMT p\        VP!                  4       4      p\        P                  ! VP"                  4      pV EFE  pVP$                   EF0  p\'        WxV4      '       g   K  \)        VP*                  ^ ,          P,                  4      ^8  d   KE  WXP*                  ^ ,          P.                  ,          p	WXP.                  ,          p
V
P0                  '       g   K  V^ ,          \         P                  \         P                  \         P
                  39   d   \3        W4      pM\5        W4      p\7        VP*                  ^ ,          W[4       VP9                  VP*                  ^ ,          4       VP;                  V4       EK3  	  EKH  	  \        P                  ! WF4      # )z
Fuses convolution/BN and linear/BN layers for inference purposes.
Will deepcopy your model by default, but can modify the model inplace as well.
)r6   Conv1dBatchNorm1dConv2dBatchNorm2dConv3dBatchNorm3dLinearcopydeepcopyr-   r5   r%   GraphModulesymbolic_tracer'   named_modulesgraphr/   r2   r*   r+   usersr   track_running_statsr   r   r=   replace_all_uses_with
erase_node)r@   inplaceno_tracepatternsfx_modelr#   	new_graphr!   r"   first_layerbnfused_layers   &&&         r   fuser[   M   s    
BNN#	BNN#	BNN#	BNN#	H e$:eUXX-A-ABB$$U+8))+,Ghnn-IOOD%gW==tyy|))*Q.%iil&9&9:[[)---1:"))RYY		!BB"3K"DK"5k"FK#DIIaL'G**499Q<8$$T* $ " >>(..r   c                X    V ^8  d   QhR\         P                  R\         P                  /# r?   )r6   r7   )r   s   "r   r   r   u   s"     0 0")) 0		 0r   c                    \         P                  ! V 4      p ! R R\        P                   P                  4      pV! V4      P	                  4       # )z-
Removes all dropout layers from the module.
c                   >   a a ] tR t^{t oV3R lV 3R lltRtVtV ;t# )&remove_dropout.<locals>.DropoutRemoverc                \   < V ^8  d   QhRS[ RS[S[R3,          RS[S[S[3,          RS[/# )r   r   r+   .kwargsr   )r	   r   r   r'   r   r   )r   __classdict__s   "r   r   3remove_dropout.<locals>.DropoutRemover.__annotate__|   sE     	A 	A 	A(-hm(<	AFJ3PS8n	A	Ar   c                   < \        V P                  V,          \        P                  4      '       d1   \	        V4      ^8w  d   \        R\	        V4       24      hV^ ,          # \        SV `  WV4      # )   z Expected 1 arg for Dropout, got )r-   
submodulesr6   Dropoutr*   r:   superr)   )selfr   r+   ra   	__class__s   &&&&r   r)   2remove_dropout.<locals>.DropoutRemover.call_module|   s]     $//&12::>>t9>(+KCPTI;)WXXAww*6@@r    )__name__
__module____qualname____firstlineno__r)   __static_attributes____classdictcell____classcell__)rj   rb   s   @@r   DropoutRemoverr_   {   s     	A 	A 	Ar   rt   )r%   rL   r5   Transformer	transform)r@   rV   rt   s   &  r   remove_dropoutrw   u   sB       'H	A-- 	A (#--//r   c          	          V ^8  d   QhR\         P                  R\        \        P                  ,          R\        \        P                  ,          R\        \        P                  ,          /# )r   orig_moduler/   inputsoutputs)r6   r7   listr%   r&   )r   s   "r   r   r      sL     2 22=2 M2 "'']	2r   c                r  a	 \         P                  ! 4       p/ o	V F#  pVP                  VP                  4      pVS	V&   K%  	  V F  pVP	                  VV	3R l4      pVS	V&   K   	  TP                  V Uu. uF  pS	V,          NK  	  up4       VP                  4        \         P                  ! W4      # u upi )zy
Given lists of nodes from an existing graph that represent a subgraph, returns a submodule that executes that subgraph.
c                    < SV ,          # Nrl   )xenvs   &r   <lambda>"extract_subgraph.<locals>.<lambda>   s	    s1vr   )r%   Graphplaceholderr   	node_copyoutputlintrK   )
ry   r/   rz   r{   rW   inputnew_noder"   r   r   s
   &&&&     @r   extract_subgraphr      s     
I"$C((4E
  &&t-=>D	  8fc&kk89NN>>+11 9s   5B4c                 .    \         P                  ! V 4      # r   )	th_mkldnnMkldnnBatchNorm)a_s   &&r   r   r      s    !:!:1!=r   c                    V ^8  d   QhR\         \        P                  ,          R\        \        \
        P                  3,          /# )r   r/   r#   r|   r%   r&   r'   r   r6   r7   )r   s   "r   r   r      s/      T"''] T#ryy.5I r   c                $   / pV  EF  pVP                   R8X  g   K  \        VP                  \        4      '       g"   \	        R\        VP                  4       24      hWP                  ,          p\        V4      \        9   g   K  \        \        V4      ,          ! V\        P                  4      p\        V\        P                  4      '       g   \	        R\        V4       24      h\        P                  ! V4      W%&   \        W1V4       EK	  	  V# )z
For each node, if it's a module that can be preconverted into MKLDNN,
then we do so and create a mapping to allow us to convert from the MKLDNN
version of the module to the original.
r)   r9   zExpected nn.Module, got )r.   r-   r   r   r:   r$   
mkldnn_mapr5   floatr6   r7   rI   rJ   r=   )r/   r#   old_modulesr"   
cur_moduler4   s   &&    r   modules_to_mkldnnr      s     /1K77m#dkk3//$'@dkkAR@S%TUU -JJ:-'Z(89*ekkR
!*bii88(+CDDTCU)VWW*.--
*C'#D:>  r   c                    V ^8  d   QhR\         \        P                  ,          R\        \        \
        P                  3,          R\        \
        P                  \
        P                  3,          /# )r   r/   r#   r   r   )r   s   "r   r   r      sR     L L=L#ryy.!L bii*+Lr   c                   V  F  pVP                   R8X  g   K  \        VP                  \        4      '       g"   \	        R\        VP                  4       24      hWP                  ,          pWB9   g   Kq  \        W1W$,          4       K  	  R# )zU
Maps each module that's been changed with `modules_to_mkldnn` back to its
original.
r)   r9   N)r.   r-   r   r   r:   r$   r=   )r/   r#   r   r"   r   s   &&&  r   reset_modulesr      sg     77m#dkk3//$'@dkkAR@S%TUU -J(#D;3JK r   c                   2   a  ] tR t^t o V 3R lR ltRtV tR# )r   c                4   < V ^8  d   QhRS[ P                  /# )r   fx_graph)r%   r   )r   rb   s   "r   r   MklSubgraph.__annotate__   s     + + +r   c                <    Wn         . V n        . V n        . V n        R # r   )r   r/   start_nodes	end_nodes)ri   r   s   &&r   __init__MklSubgraph.__init__   s     $&
*,(*r   )r   r   r/   r   N)rm   rn   ro   rp   r   rq   rr   rb   s   @r   r   r      s     + +r   c                2   a aaaa RoRoR V VVVV3R llpV# )a?  
This generates a heuristic that can be passed into `optimize_for_inference` that
determines whether a subgraph should be run in MKL by running it with the example_inputs.

Example usage:
    heuristic = gen_mkl_autotuner(example_inputs, iters=10)
    fast_model = optimization.optimize_for_inference(model, heuristic)
Nc                0    V ^8  d   QhR\         R\        /# r   rN   r   r   bool)r   s   "r   r   'gen_mkl_autotuner.<locals>.__annotate__   s      &  &  &  &r   c                   <aa V P                   pS
fG   V P                  P                  o
V P                  P                  o\	        S
4      P                  S	4       V Uu. uF#  p\        P                  ! VP                  4      NK%  	  upo\        \        \        P                  ,          V P                   Uu. uF  q"P                  ^ ,          NK  	  up4      p\        S
V P                   W4      oVV3R lpV! VV3R l4      p\#        SP$                  P                   \'        SP)                  4       4      S4       V! VV3R l4      pWV8  # u upi u upi )Nc                    < \        S4       F
  pV ! 4        K  	  \        P                  ! 4       p\        S4       F
  pV ! 4        K  	  \        P                  ! 4       V,
          # r   )rangetime)fr   beginiterswarmups   &  r   	benchmark?gen_mkl_autotuner.<locals>.use_mkl_heuristic.<locals>.benchmark  sE    6] #IIKE5\ "99;&&r   c                     < S! S U u. uF  q P                  4       NK  	  up !   U u. uF  q P                  4       NK  	  up # u up i u up i r   )	to_mkldnnto_dense)isample_inputs	submodules    r   r   >gen_mkl_autotuner.<locals>.use_mkl_heuristic.<locals>.<lambda>
  s?    &/1WA++-1W&X&X

&X1Ws
   AAc                     < S! S !  # r   rl   )r   r   s   r   r   r     s
    	=(Ar   )r   r   owning_moduler   r
   	propagater5   randnshaper   r|   r%   r&   r   r+   r   r/   r   rN   r'   rM   )rN   input_nodesr"   output_argsr   mkl_timeno_mkl_timer   r   example_inputsrV   r   r   r   s   &      @@r   use_mkl_heuristic,gen_mkl_autotuner.<locals>.use_mkl_heuristic   s   ''~~33H..44Kh)).9=HI[TTZZ0[I4=EOO*TOD99Q<<O*TU$Xu{{KU		' 
 	OO!!((*+		
   AB%%3 J*Ts   )E7E
rl   )r   r   r   r   rV   r   s   fff @@r   gen_mkl_autotunerr      s"     HK &  &D r   c                0    V ^8  d   QhR\         R\        /# r   r   )r   s   "r   r   r     s        +  $  r   c                2    \        V P                  4      ^8  # )z
This is a heuristic that can be passed into `optimize_for_inference` that
determines whether a subgraph should be run in MKL by checking if there
are more than 2 nodes in it
)r*   r/   )rN   s   &r   use_mkl_lengthr     s     u{{ar   c                   \   a  ] tR tRt o R tV 3R lR ltV 3R lR ltV 3R lR ltR	tV t	R
# )r   i$  c                B    R .V,          V n         ^ .V,          V n        R # r   r   size)ri   ns   &&r   r   UnionFind.__init__%  s    )-
 !sQw	r   c                    < V ^8  d   QhRS[ /# )r   vint)r   rb   s   "r   r   UnionFind.__annotate__)  s      # r   c                @    WP                   V&   ^V P                  V&   R# )re   Nr   )ri   r   s   &&r   make_setUnionFind.make_set)  s    A		!r   c                &   < V ^8  d   QhRS[ RS[ /# )r   r   r   r   )r   rb   s   "r   r   r   -  s     ) )c )c )r   c                    V P                   V,          pW8X  d   V# Vf   \        R4      hV P                  V4      V P                   V&   \        \        V P                   V,          4      # )NzParent is None)r   r:   findr   r   )ri   r   pars   && r   r   UnionFind.find-  sT    kk!n8H; !1223ACQ((r   c                &   < V ^8  d   QhRS[ RS[ /# )r   r   br   )r   rb   s   "r   r   r   6  s     % %c %c %r   c                *   V P                  V4      V P                  V4      r!W8X  d   V# V P                  V,          V P                  V,          8  d   Y!r!WP                  V&   V P                  V;;,          V P                  V,          ,          uu&   R # r   )r   r   r   )ri   r   r   s   &&&r   joinUnionFind.join6  se    yy|TYYq\16H99Q<$))A,&qA		!		!$r   r   N)
rm   rn   ro   rp   r   r   r   r   rq   rr   r   s   @r   r   r   $  s(     ' ) )% %r   c                    V ^8  d   QhR\         P                  P                  R\        \        \
        3,          R,          R\        \        P                  ,          R\         P                  P                  /# )r   r@   pass_configNtracerr   )	r5   r6   r7   r'   r   r   r$   r%   Tracer)r   s   "r   r   r   @  s[     r r88??rc3h$&r Or XX__	rr   c                  aa RRRRRR\         //pVf   / pVP                  V4       VR,          '       d   \        V 4      p VR,          '       d   \        V 4      p VR,          RJ d   V # \	        VR,          \
        4      '       g   \        R4      hRVR,          9  d   \        R	4      hVR,          R,          pV! 4       pVP                  \        P                  ! V 4      4      o\        P                  ! VP                  S4       \        V P                  4       4      p ! R
 R\        4      p\        SP                   4       EF^  pVP"                  p	VP$                  R8X  d   WhP&                  ,          p
\)        V
4      \*        9   d   VP,                  p	\/        V
P1                  4       R4      pVe[   VP2                  \4        P6                  8w  d   \9        R4      hVP:                  \4        P:                  ! R4      8w  d   \9        R4      hMTVP$                  R8X  dD   VP&                  \*        9   d   VP,                  p	M!VP&                  \<        9   d   VP>                  p	WP"                  8w  g   EK3  WP>                  8X  dR   \@        ;QJ d&    R VPB                   4       F  '       g   K   RM	  RM! R VPB                   4       4      '       g   EK  SPE                  V4      ;_uu_ 4        \        PF                  ! VPB                  V3R l4      pRRR4       \I        \J        \        PL                  PN                  ,          X4      Vn!        SPQ                  V4      ;_uu_ 4        SPS                  RRV34      pVPU                  V4       V3Vn!        RRR4       EKa  	  \W        \        SP                   4      V4      pVSn,        SP                    F  pVP$                  R8X  g   K  VP&                  R8X  g   K)  VPB                  ^ ,          p\        VPZ                  4      pV FK  pVP$                  R8X  g   K  VP&                  R8X  g   K)  VPU                  V4       SP]                  V4       KM  	  \_        VPZ                  4      ^ 8X  g   K  SP]                  V4       K  	  \_        SP                   4      p\a        V4      oV3R lp\c        SP                   4       EF  w  ppVP$                  R8X  d,   VP&                  R8X  d   VVn2        SPg                  V4       KC  VP$                  R8X  dX   VP&                  R8X  dG   V! VPB                  ^ ,          4      f   \9        R4      hV! VPB                  ^ ,          4      Vn4        K  VPj                   Uu. uF9  p\	        V\        Pl                  4      '       g   K%  V! V4      f   K1  V! V4      NK;  	  pp\_        V4      ^ 8X  d   EK  \@        ;QJ d    R V 4       F  '       g   K   RM	  RM! R V 4       4      '       d   \9        R4      h\o        V4      pV^ ,          Vn8        VR,           F  pSPs                  V^ ,          V4       K  	  EK  	  \u        V3R l4      pSP                    F  p\w        VR4      '       d<   VSPy                  VPp                  4      ,          P                   P{                  V4       \w        VR4      '       d<   VSPy                  VPd                  4      ,          P|                  P{                  V4       \w        VR4      '       g   K  VSPy                  VPh                  4      ,          P~                  P{                  V4       K  	  VP                  4        F  pV! V4      '       d   K  VP|                  VP~                  ,            F8  pVPB                  ^ ,          pVPU                  V4       SP]                  V4       K:  	  \        VP                   Wn4       K  	  ^ pSP                    F0  pVP&                  R8X  g   VP&                  R8X  g   K'  V^,          pK2  	  \        P                  ! \        4      P                  RV4       SP                  4        \        P                  ! V S4      pV#   + '       g   i     EL; i  + '       g   i     EK  ; iu upi ) a  
Performs a set of optimization passes to optimize a model for the
purposes of inference. Specifically, the passes that are run are:
1. Conv/BN fusion
2. Dropout removal
3. MKL layout optimizations

The third optimization takes a function `use_mkl_heuristic` that's used
to determine whether a subgraph should be explicitly run in MKL layout.

Note: As FX does not currently handle aliasing, this pass currently
assumes nothing aliases. If that isn't true, use at your own risk.
conv_bn_fuseTrw   mkldnn_layout_optimize	heuristicNFz+mkldnn_layout_optimize config is not a dictz4Heuristic not found in mkldnn_layout_optimize configc                   "    ] tR tRt^t^t^tRtR# )*optimize_for_inference.<locals>.MklSupportil  rl   N)rm   rn   ro   rp   NOYESUNKNOWNrq   rl   r   r   
MklSupportr   l  s    r   r   r)   z)this pass is only for torch.float modulescpuz!this pass is only for CPU modulescall_functionc              3   >   "   T F  qP                   R 8H  x  K  	  R# 5i)r   N)r   ).0args   & r   	<genexpr>)optimize_for_inference.<locals>.<genexpr>  s     Iy::3ys   c                 *   < SP                  R V 34      # )r   )call_method)r   r   s   &r   r   (optimize_for_inference.<locals>.<lambda>  s    )=)=kA4)Pr   r   r   r   c                    < \        V R 4      '       d   SP                  V P                  4      # \        V R4      '       d   SP                  V P                  4      # R# )colorstart_colorN)hasattrr   r   r   )r   ufs   &r   	get_color)optimize_for_inference.<locals>.get_color  sF    1g77177##1m$$771==))r   z!Expected color for to_dense inputc              3   (   "   T F  qR J x  K
  	  R # 5ir   rl   )r   r   s   & r   r   r     s     1j9js   zFound None in cur_colors:re   NNc                     < \        S 4      # r   )r   )r   s   r   r   r     s
    H@Ur   r   r   	end_colorzmkldnn conversions: %s)Gr   updater[   rw   r-   r'   RuntimeErrortracerI   rJ   r%   rK   rootrM   r   r|   r/   r   r.   r   r$   mkldnn_supportedr   next
parametersdtyper5   r   r:   devicemkldnn_supported_unknownr   anyr+   inserting_beforemap_argr   r   r"   r   inserting_aftercreate_noderQ   r   r   rO   rR   r*   r   	enumerater   r   r  all_input_nodesr&   sortedr   r   r   r   r   appendr   r   valuesr   logging	getLoggerrm   infor   )r@   r   r   default_pass_configr   
cur_tracerr#   r   r"   supports_mkldnnr   sample_parametermkldnn_argsdense_xr   prv_noderO   user	num_nodesr  cur_idxr   
cur_colorsother_colormkldnn_graphsrN   prvmkldnn_conversionsresultr   r  s   &&&                          @@r   optimize_for_inferencer.  @  s   & 	$ ;"?
 {+>**U+,,u%34=)*BCTJJHII-.FGGQRR+,DEkRJe 45HNN:??H-$()<)<)>$?GT  X^^$$--77m# -JJ#33",..#'
(=(=(?#F #/'--<,G  (..%,,u2EE,-PQQWW'{{..",.. 88","4"4mm+"4"44sItyyIsssItyyIII**400 jjIIP 1
 U277#3#34kBDI))$//"..}j4'R**73 $w 0/? %J $D$8'BK&H 77m#z(Ayy|H$E77m+{0J..x8''-  4::!###D)  HNN#I	9	B$ #8>>277m#{(B&DKK WW%$++*C1&.$%HII&tyy|4DN ---Aa)  Q< 	!-   :!#s1j1sss1j111$%?@@
+J#ADJ)"~~
1{3  .- 32 -88U,VM4!!"''$**-.44;;DA4''"''$"2"234@@GGM4%%"''$..12<<CCDI  %%' ''))EOO;;iil**3/##D) < %++w< ( ;;+%
)B!#  h$$%=?QRMMO^^E8,FMK 100 0//fs*   3&b8..c""c!	c!c!8c	c)r2   r=   r[   rw   r   r   r   r   r   r   r   r.  )FF)
   re   )CrI   r  operatorr   collectionsr   collections.abcr   enumr   typingr   r   r5   torch.fxr%   torch.nnr6   torch.nn.functional
functionalFtorch.utils.mkldnnutilsmkldnnr   torch.fx.noder   r	   torch.fx.passes.shape_propr
   torch.nn.utils.fusionr   r   __all__r   r2   r=   r[   rw   r   rD   rH   rE   ReLU	MaxPool2d	AvgPool2dAdaptiveAvgPool2drelu	transposesigmoid
avg_pool2dadaptive_avg_pool2dr  addmulr  MkldnnConv2dMkldnnLinearr   r   r   r   r   r   r   r   r.  rl   r   r   <module>rN     sP       # $        & & * 0 H -(4%/P0(2. IIIINNGGLLLL	JJ	OO	MMFFLL & %LL(,,7 IIy%%IIy%%NN=
,L$+ +.b % %< *. iir rr   