
    iύ                       d Z ddlmZ ddlZddl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ZddlmZ ddlmZ dd	lmZ  G d
 d      Z	 d(	 	 	 	 	 	 	 d)dZ	 	 	 d*	 	 	 	 	 	 	 	 	 d+dZ	 d,	 	 	 d-dZ	 	 	 d.	 	 	 	 	 	 	 	 	 d/dZ	 	 d0	 	 	 	 	 d1dZ	 	 	 d2	 	 	 	 	 	 	 	 	 d3dZ	 	 d4	 	 	 	 	 	 	 d5dZedk(  rddlZ ej@                  d      Z!e!jE                  dg dd       e!jE                  ddd       e!jE                  ddd !       e!jE                  d"dd#!       e!jG                         Z$ ee$jJ                  e$jL                  e$jN                  e$jP                   $       d6d%Z)d7d&Z*d8d'Z+y)9z
MLX Model Weight Conversion System

Converts pretrained PyTorch HTDemucs models to MLX format with proper
weight layout transformations for Conv1d/Conv2d layers.
    )annotationsN)datetime)Path)version   )MIN_MLX_VERSION)MLX_MODEL_REGISTRYc                  2    e Zd ZdZdddZd	dZd
dZddZy)BagOfModelsMLXz
    MLX wrapper for ensemble of models with weighted averaging.

    This mirrors the PyTorch BagOfModels but operates on MLX arrays.
    Weights are per-source: weights[model_idx][source_idx]
    Nc                   || _         |d   j                  | _        |d   j                  | _        |d   j                  | _        |&|D cg c]  }dgt	        | j                        z   }}|| _        dgt	        | j                        z  | _        |D ],  }t        |      D ]  \  }}| j                  |xx   |z  cc<    . y c c}w )Nr         ?g        )modelssources
samplerateaudio_channelslenweightstotals	enumerate)selfr   r   _model_weightssrc_idxws          k/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/demucs_mlx/mlx_convert.py__init__zBagOfModelsMLX.__init__   s    ay(( )..$Qi66?:@A&Qus4<<00&GA ec$,,//$M'6
G$)$ 7 % Bs   
 Cc                z   d}t        | j                  | j                        D ]L  \  }} ||      }t        j                  |      j                  dt        |      dd      }||z  }||}H||z   }N t        j                  | j                        j                  dt        | j                        dd      }||z  }|S )z=Apply all models and average outputs with per-source weights.Nr   )zipr   r   mxarrayreshaper   r   )r   x	estimatesmodelr   outweight_arraytotals_arrays           r   __call__zBagOfModelsMLX.__call__0   s    	$'T\\$B E=(C 88M2::1c->PRSUVWL$C 	%O	 %C xx,44QDKK8H!QO,	    c                    t        | j                        D ci c]  \  }}d| |j                          c}}S c c}}w )z!Return state dict for all models.model_)r   r   
state_dict)r   ir$   s      r   r,   zBagOfModelsMLX.state_dictH   sJ     &dkk2
25 QCL%**,,2
 	
 
s   ;c                l    t        | j                        D ]  \  }}|j                  |d|            y)zLoad state dict for all models.r+   N)r   r   load_state_dict)r   stater-   r$   s       r   r/   zBagOfModelsMLX.load_state_dictO   s2    !$++.HAu!!%&"56 /r)   N)r   ztp.Listr   z$tp.Optional[tp.List[tp.List[float]]])r"   mx.arrayreturnr2   )r3   tp.Dict)r0   r4   )__name__
__module____qualname____doc__r   r(   r,   r/    r)   r   r   r      s    *"0
7r)   r   c                    |s| S |dk(  rt        j                  | d      S |dk(  rt        j                  | d      S |dk(  rt        j                  | d      S |dk(  rt        j                  | d      S t        d	|       )
z6Convert convolution weight from PyTorch to MLX layout.conv1d)r      r   conv_transpose1d)r   r<   r   conv2d)r   r<      r   conv_transpose2d)r   r<   r?   r   zUnknown conv_type: )np	transpose
ValueError)weight	conv_typerB   s      r   convert_conv_weightrF   U   s     H||FI..	(	(||FI..	h	||FL11	(	(||FL11.yk:;;r)   c           	         ddl }i }i }||j                         D ]  \  }}t        ||j                  j                        r	d|| d<   /t        ||j                  j
                        r	d|| d<   Xt        ||j                  j                        r	d|| d<   t        ||j                  j                        sd|| d<    | j                         D ]@  \  }	}
|
j                         j                         j                         }d}d}d	|	v xrN d
|	j                         v xs: d|	j                         v xs& d|	j                         v xs d|	j                         v }|	|v r||	   }d}nX|rVt        |j                        }d|	j                         v xs d|	j                         v }|dk(  r	|rdnd}d}n|dk(  r|rdnd}d}|r8|r6t        ||      }|r(t!        d|	 d|
j                   d|j                          t#        j$                  |      ||	<   C i }t'        |j)                               D ]n  }	d|	v r|	j+                  d      }t        |      dk  r(|d   }|dvr2|d   }|j-                  d      sIdj/                  |dd d|gz         }||vsg||	   ||<   p |j1                  |       i }t'        |j)                               D ]  }	d|	vr|	j+                  dd      \  }}|j3                  d       }|r|dt        d         }d!|vrD|j5                  d!d      \  }}|j7                         sjt9        |      }|d"vrz|rd#nd$}|j-                  d%      r$|d&k(  rd'nd(}| d| d| d| }||vs||	   ||<   |||f}|j;                  |i       } ||	   | |<    |j                         D ]R  \  \  }}}} | j=                  d)      }!| j=                  d*      }"|!|"1|"|!n|!|"n|!|"z   }#| d| d| d+}||vsN|#||<   T i }$t'        |j)                               D ]  }	d,|	v sd-|	v r|	j?                  d,d.      }d/|	v rtA        j$                  ||	         }%|%j                  d   dz  }&|%d|&ddf   }'|%|&d|&z  ddf   }(|%d|&z  dddf   })|j?                  d/d0      }t#        j$                  |'      |$| d1<   t#        j$                  |(      |$| d2<   t#        j$                  |)      |$| d3<   d4|	v rtA        j$                  ||	         }#|#j                  d   dz  }&|#d|& }*|#|&d|&z   }+|#d|&z  d },|j?                  d4d0      }t#        j$                  |*      |$| d5<   t#        j$                  |+      |$| d6<   t#        j$                  |,      |$| d7<   d8|	v s||	   |$|<   d9|	v sd:|	v s|	j?                  d;d<      }||	   |$|<    |j1                  |$       |r|S tC        d=      )>zV
    Convert PyTorch state dict to MLX format with proper layout transformations.
    r   Nr;   z.weightr=   r>   r@   FrD   convrewrite	upsamplerdownsamplerTconv_trrB   r?      z  Transposed : u    → z.gn..r<   )rD   biasnormgnz.lstm.r   _reverse_l)	weight_ih	weight_hhbias_ihbias_hhbackward_lstmsforward_lstmsweight_rW   WxWhrY   rZ   z.bias	self_attn
cross_attnattnz.in_proj_weight z.query_proj.weightz.key_proj.weightz.value_proj.weightz.in_proj_biasz.query_proj.biasz.key_proj.biasz.value_proj.biasz
.out_proj.z.norm_out.weightz.norm_out.biasz
.norm_out.z.norm_out.gn.z5Nested conversion is not supported; use flatten=True.)"torchnamed_modules
isinstancennConv1dConvTranspose1dConv2dConvTranspose2ditemsdetachcpunumpylowerr   shaperF   printr   r    listkeyssplit
startswithjoinupdateendswithrsplitisdigitint
setdefaultgetreplacerA   NotImplementedError)-torch_stateverboseflattentorch_modelrd   flat_mlx_statemodule_param_typesmodule_namemodulenameparamnp_paramneeds_transposerE   is_conv_like_weightndimis_transposenorm_wrapper_fixespartslastprevnew_name	lstm_biasprefixrest
is_reversebase	layer_strlayerdir_namemlx_namekeyentryrY   rZ   rQ   transformer_fixesrD   	embed_dimquery_weight
key_weightvalue_weight
query_biaskey_bias
value_biass-                                                r   convert_state_dictr   j   s    N,.#.#<#<#>K&%((//2>F"k]'#:;FEHH$<$<=>P"k]'#:;FEHHOO4>F"k]'#:;FEHH$<$<=>P"k]'#:; $? #((*e<<>%%'--/  	  ,tzz|# +$**,&+DJJL(+ djjl* 	 %%*40I"O x~~&D$

4Stzz|8SLqy2>.H	"&2>.H	"& y*8Y?HdV2ekk]%?OPQ  "xx1tO +T ^((*+T>

3u:>Ry))Ry??6"xxcr
dD\ 9:H~-/=d/C"8, , ,- I^((*+4zz(A.]]:.
)#j/)*Dt++dA.i  "IGG'1#??9%#{2tH 8*AeWAhZ@H~-+9$+?x(8U+C((b1E(.E$K3 ,6 -6OO,=(!55))I&))I&?w!/w7?QX[bQbXQxj%6>)'+N8$ -> ^((*+$,$"6||K8H D(."67"LLOq0	%jyj!m4#Ia	k$91$<=
%a	klAo6''(92>AC,AW!TF*<"=>?Axx
?S!TF*:";<AC,AW!TF*<"=> D(xxt 45 JJqMQ.	!*9-
	!I+6!!I+,/
''<?Axx
?S!TF*:";<=?XXh=O!TF."9:?Axx
?S!TF*:";<%.<T.B!(+4'+;t+C||L/BH*8*>h'C ,F +,
U
VVr)   c                p   ddl m} ddlm} ddlm} | j                  j                  }rt        d| d       t        | d      r| j                  \  }}nt        d| d	      fd
}|dk(  rH|j                  d      rt        d      |j                  d      rt        d       ||i  |||      }	nX|dk(  r ||i  |||      }	nC|dk(  r0d|v rd|vr|d   |d<   d|v rd|vr|d   |d<    ||i  |||      }	nt        d|       t        | d      rt        |	d      r| j                  |	_        r&t        dt        | j!                                d       | j!                         }
t#        |
d|       }rt        d       t%        |	|       rt        d       rt        d|        |	S )z0
    Convert a single PyTorch model to MLX.
    r   	DemucsMLX
HDemucsMLXHTDemucsMLXConverting ..._init_args_kwargszModel z doesn't have _init_args_kwargsc                X   t        j                  |       }t        |j                  j	                               |j                         D ci c]  \  }}|v s|| }}}r?t        fd|j	                         D              }|rt        d| j                   d|        |S c c}}w )Nc              3  ,   K   | ]  }|vs|  y wr1   r9   .0kalloweds     r   	<genexpr>z?convert_single_model.<locals>._filter_kwargs.<locals>.<genexpr>-       M(81AW<LQ(8   	"  Dropping unsupported kwargs for rN   	inspect	signatureset
parametersrt   rl   sortedrr   r5   )	
target_cls	in_kwargssigr   vfiltereddroppedr   r   s	          @r   _filter_kwargsz,convert_single_model.<locals>._filter_kwargs(  s    
+cnn))+,%.__%6G%6TQ!w,AqD%6GM	(8MMG::;N;N:OrRYQZ[\ Hs   B&B&HTDemucst_sparse_self_attnz2Sparse self-attention not supported in MLX backendt_sparse_cross_attnz3Sparse cross-attention not supported in MLX backendHDemucsDemucsgelugelu_actgluglu_actz#MLX conversion not implemented for segmentz parameters...T)r   r   r   z   Using manual weight loading...z  Loaded parameters manuallyu   ✓ Converted )
mlx_demucsr   mlx_hdemucsr   mlx_htdemucsr   	__class__r5   rr   hasattrr   rC   r~   r   r   r   r,   r   _load_weights_into_model)r   r   r   r   r   model_classargskwargsr   	mlx_modelr   r   s    `          r   convert_single_modelr     s    &')''00KK=,- {/0"44f6+.MNOO j ::*+QRR::+,RSSMV)LM				!Kz6(JK			 V
& 8!'F:F?y6 &uF9tI~i'HI	!1+?
 	
 {I&79i+H'//	 C 6 6 89:.IJ((*K'WdN 01Y7,-{m,-r)   
model_namec                
    	 ddl m} ddlm} | t
        vr,t        d|  dt        t        j                                      |d}t        j                  |d	
       t
        |    }|r6t        d       t        d|  d       t        d|d           t        d       |rt        d        ||       }t        ||      r~|r"t        dt        |j                         d       |j                  }t        |d      r|j                   |j                   }	nUt        |j"                        }
|D cg c]  }dg|
z  
 }	}n,|rt        d       |g}t        |j"                        }
dg|
z  g}	|rt        dt        |       d       g }t%        |      D ]Q  \  }}|r,t        |      dkD  rt        d|dz    dt        |       d       t'        ||      }|j)                  |       S t        ||      rD|rt        dt        |       d|	 d       t+        ||	      }d }t-        |d         j.                  }n|d   }t-        |      j.                  }d}|d   j0                  \  }}t3        |      }t        |d   d!      r|d   j4                  |d!<   g g  g |D ]  }|j0                  \  }}t3        |      }t        |d!      r|j4                  |d!<   j)                  t        |              j)                  |       j)                  t-        |      j.                          |rt        d"       | |||||j7                         t8        t        |      t        ||      r|	ndt;        j<                         j?                         |d#   d$}t        ||      rXtA        fd%dd D              }tA         fd& dd D              }|s|r
|d'<    |d(<   tA        fd)dd D              r|d*<   t        jB                  jE                  ||  d+      }tG        |d,      5 }tI        jJ                  ||       ddd       |r@t        jB                  jM                  |      d-z  }t        d.|        t        d/|d0d1       |r_|rt        d2       	 tO        |d   |d   |       d	|d3<   tG        |d,      5 }tI        jJ                  ||       ddd       |rt        d4       |r$t        d7       t        d8|        t        d       |S # t        $ r t	        d      dw xY wc c}w # 1 sw Y   xY w# 1 sw Y   hxY w# tP        $ r}|rt        d5|        d6|d3<   Y d}~d}~ww xY w)9zH
    Convert Demucs/HDemucs/HTDemucs PyTorch weights to MLX format.
    r   )BagOfModels)	get_modelz_Model conversion requires the [convert] extras. Install with: pip install 'demucs-mlx[convert]'NzUnknown model: z. Available: ./mlx_checkpointsT)exist_okzF======================================================================r   z to MLX formatzDescription: descriptionz
1. Loading PyTorch model(s)...z   Found bag with  modelsr   r   z   Loaded single modelz
2. Converting z model(s)...r   z

   Model /:)r   z
3. Creating ensemble with z model(s) and weights r   r   r   z
4. Saving MLX checkpoint...
signatures)r   r   sub_model_classr   r   r0   mlx_version
num_modelsr   conversion_datetorch_signaturesc              3  .   K   | ]  }d    |k7    ywr   Nr9   )r   aper_model_argss     r   r   z+convert_htdemucs_weights.<locals>.<genexpr>  s     M:LQ.+q0:L   c              3  .   K   | ]  }d    |k7    ywr   r9   )r   r   per_model_kwargss     r   r   z+convert_htdemucs_weights.<locals>.<genexpr>  s     S>R,Q/14>Rr   r   r   c              3  .   K   | ]  }d    |k7    ywr   r9   )r   cper_model_classs     r   r   z+convert_htdemucs_weights.<locals>.<genexpr>  s     D0C1q!Q&0Cr   r   _mlx.pklwbi   z   Saved to: z   File size: z.1fz MBz!
5. Running verification tests...verification_passedu      ✓ Verification passedu      ✗ Verification failed: FzG
======================================================================u   ✓ Conversion complete: ))demucs.applyr   demucs.pretrainedr   ImportErrorr	   rC   rs   rt   osmakedirsrr   rf   r   r   r   r   r   r   r   appendr   typer5   r   dictr   r,   r   r   now	isoformatanypathrw   openpickledumpgetsizeverify_conversion	Exception)!r   
output_dirverifyr   r   r   configr   torch_modelsr   num_sourcesr   
mlx_modelsr-   tmr   final_modelr   r   r   r   tm_args	tm_kwargs
checkpointargs_differkwargs_differoutput_pathffile_size_mber   r   r   s!                                 @@@r   convert_htdemucs_weightsr  `  su   ,/ ++j\ *16689:<
 	

 (
KK
T*
+FhJ<~67f]3456h01J'K+{+&s;+=+='>&?wGH"));	*{/B/B/N!))Gk112K4@ALqu{*LGA*+#}+--.5;&' \!2 3<@AJ<(2s<(1,K!uAc,&7%8:;(W=	)$	 ) +{+.s:.? @&is, %Z9&z!}-66 m;'00?44LD&&\F|A	*(O33yNO11O	2y!#%::Ii d7m,	*tBx001  -. !"*'')&*o(kB7#<<>335"<0J +{+M.:LMMS>Nqr>RSS-+9J'(-=J)*D0CDD,;J()'',,zj\+BCK	k4	 AJ" 
! ww{3{Ck]+,|C04567
	6l1oz!}gN04J,-k4(AJ* )23 o)+78h[  >
 	L BV 
!	  )(  	64QC8905J,-	6sM   T 0TT$%T< 2T0	T< T$T-0T95T< <	U$UU$c                   ddl }ddlm}m} |rt	        d       |j                  ddd      }t        j                  j                   ||j                                     }|j                         5  | j                           | |      }	ddd       t        |d      r|j                           ||      }
t        j                  |
        ||
      }	|z
  j                         j                         j                         }|	|z
  j                         j                         j                         }|	j                         j                         j                         }||d	z   z  }|rNt	        d
|d       t	        d|d       t	        d|d       t	        dt!        |j"                                ||kD  rt%        d|dd|d      y# 1 sw Y   FxY w)z+Verify MLX conversion by comparing outputs.r   N)from_dlpack	to_dlpackz   Testing with random input...r   r<   i evalg:0yE>z   Max absolute difference: z.2ez   Mean absolute difference: z   Relative error: z   Output shape: z$Verification failed: relative error z > T)rd   torch.utils.dlpackr!  r"  rr   randnr   core
contiguousno_gradr#  r   absmaxitemmeantuplerq   rC   )r   r   	tolerancer   rd   r!  r"  torch_input	mlx_inputtorch_output
mlx_outputmlx_output_torchmax_diff	mean_diff	torch_max	rel_errors                   r   r  r    s    9/0++aI.K##Ik.D.D.F$GHI	";/ 
 y&!9%J GGJ":.//446::<AACH 00557<<>CCEI  "&&(--/II,-I,XcN;<-i_=>#Ic?34!%(8(>(>"?!@AB929S/YsOT
 	
 ; 
s   0GGr   c                   ddl m} ddlm} ddlm} d,d}t        j                  j                  ||  d      }t        j                  j                  ||  d      }	t        j                  j                  |      r>t        j                  j                  |	      r|rt        d|        	 t        | ||	      S t        j                  j                  ||  d      }t        j                  j                  |      rO|rt        d|        t        |d      5 }t        j                  |      }ddd       t!        j#                  d      t$              r't        d       |rt'        | |d|      S t)        d      |d   }|d   }|d   }|dk(  r$|d   }|d   } ||j#                  d            }|j#                  d      }|j#                  d      }|j#                  d      }g }t+        |      D ]  }|}|}|r|t-        |      k  rt/        ||         }|r|t-        |      k  r||   }|}|r|t-        |      k  r |||         }|dk(  r	 ||i |}n*|d k(  r	 ||i |}n|d!k(  r	 ||i |}nt)        d"|       |j1                  |        t3        ||      }|j5                  |d#          nL|dk(  r	 ||i |}n*|d k(  r	 ||i |}n|d!k(  r	 ||i |}nt)        d$|       |j5                  |d#          |rt        d%|        |dk(  r#|j6                  D ]  }|j9                           |S |j9                          |S |r/|rt        d&|  d'       t;        | |d|(       t'        | |d|      S t=        d)| d*|  d+      # t        $ r}
|rt        d
|
 d       Y d}
~
d}
~
ww xY w# 1 sw Y   xY w)-z/Load MLX model from cache or convert if needed.r   r   r   r   c                4    | sy| dv r| S d| v ryd| v ryd| v ryyNr   >   r   r   r   r   r   r   r   r   r9   r   s    r   _normalize_model_classz.load_mlx_model.<locals>._normalize_model_class9  :     ==K tr)   .safetensors_config.jsonz-Loading from safetensors (preferred format): )	cache_dirr   z%Warning: Safetensors loading failed (z), falling back to pickleNr   zLoading cached MLX model: rbr   z/Warning: MLX version mismatch. Re-converting...F)auto_convertr   zMLX version mismatchr   r   r   r   r   r   r   r   r   r   r   r   r   Unknown sub-model class: r0   Unknown model class:    ✓ Loaded z"No cached model found. Converting r   r  r  r   zNo cached MLX model found at z . Run convert_htdemucs_weights('z	') first.r   tp.Optional[str]r3   str)r   r   r   r   r   r   r   r  rw   existsrr   load_mlx_model_from_safetensorsr  r  r	  load_version_ltr~   r   load_mlx_modelrC   ranger   r-  r  r   r/   r   r#  r  FileNotFoundError)r   r@  rB  r   r   r   r   r<  safetensors_pathsafetensors_config_pathr  
cache_pathr  r  model_class_namer   r   r   r   r   r   r   r   r   r-   
model_argsmodel_kwargsr   r$   r  s                                 r   rN  rN  .  s    &') ww||I*\/JK ggll9L6QR	ww~~&'BGGNN;R,SABRASTU	\2i  iJ<x)@AJ	ww~~j!.zl;<*d#qQJ $ z~~m4oFCD%j)%Y`aa344%m4&!H%//#L1J +G4Z^^DU5VWO'^^,<=N)~~.@A(nn->?OF:&!
%!a#n*=&=!&~a'8!9J#C0@,A(A#3A#6L-"q3+?'?"89K"LK-/'D|DE L0&
ClCE K/%zB\BE$'@%NOOe$' '* )9K''
7(;<  =0)4:6:!\1($9&9![0'88 #89I8J!KLL''
7(;<K 0123//$++

 ,
  	6zl#FG 	! 		
 	w
 	

  +J< 8--7L	C
 	
A  	\=aS@YZ[	\ $#s$   ,N N7	N4N//N47Oc           	        ddl }ddl}ddlm} ddlm} ddlm} ddlm	} fd}	d	 }
d2d
}|j                  j                  ||  d      }|j                  j                  ||  d      }|j                  j                  |      st        d| d|        |j                  j                  |      st        d| d|        rt        d|        t        |d      5 }|j!                  |      }ddd       t#        j%                  d      t&              r&t        d|j%                  d       dt&         d       n=|j%                  d      t&        k7  r%t        d|j%                  d       dt&         d       rt        d        ||      }rt        dt)        |       d       |d   }|d   rt+        |d         nd}|d   }|j%                  d      }|j%                  d      }|j%                  d      }|d    }|j%                  d!      } ||j%                  d"            }|d#k(  rNrt        d$| d%       g }t-        |      D ]  }|}|}||t)        |      k  rt+        ||         }||t)        |      k  r||   }|}||t)        |      k  r |||         }|d&k(  r ||i  |	||      }n@|d'k(  r ||i  |	||      }n+|d(k(  r |
|      }  ||i  |	||       }nt/        d)|       d*| d+}!i }"|j1                         D ]*  \  }#}$|#j3                  |!      s|#t)        |!      d }%|$|"|%<   , t5        ||"       |j7                  |        t9        ||      }&rt        d,| d-       nrt        d.| d/       |d&k(  r ||i  |	||      }&n@|d'k(  r ||i  |	||      }&n+|d(k(  r |
|      }  ||i  |	||       }&nt/        d0|       t5        |&|       rt        d1|        |d#k(  r#|&j:                  D ]  }|j=                           |&S |&j=                          |&S # 1 sw Y   JxY w)3a  
    Load MLX model from safetensors format (faster, safer than pickle).

    This function loads models converted with convert_to_safetensors.py.
    Benefits:
    - 10-16x faster loading via lazy loading
    - 40%+ less memory usage
    - No PyTorch dependency for inference
    - Safer format (no arbitrary code execution)

    Args:
        model_name: Model name (e.g., 'htdemucs')
        cache_dir: Directory containing .safetensors files
        verbose: Print loading information

    Returns:
        MLX model instance with loaded weights

    Raises:
        FileNotFoundError: If safetensors or config files not found
        ValueError: If model configuration is invalid
    r   N)	load_filer   r   r   r   c                Z   dd l } |j                  |       }t        |j                  j	                               |j                         D ci c]  \  }}|v s|| }}}	r?t        fd|j	                         D              }|rt        d| j                   d|        |S c c}}w )Nr   c              3  ,   K   | ]  }|vs|  y wr1   r9   r   s     r   r   zJload_mlx_model_from_safetensors.<locals>._filter_kwargs.<locals>.<genexpr>  r   r   r   rN   r   )
r   r   r   r   r   r   r   r   r   r   s
           @r   r   z7load_mlx_model_from_safetensors.<locals>._filter_kwargs  s    g
+cnn))+,%.__%6G%6TQ!w,AqD%6GM	(8MMG::;N;N:OrRYQZ[\ Hs   B'B'c                \    t        |       }d|v rd|vr|d   |d<   d|v rd|vr|d   |d<   |S )Nr   r   r   r   )r  )r   r%   s     r   _normalize_demucs_kwargszAload_mlx_model_from_safetensors.<locals>._normalize_demucs_kwargs  sF    9oS=Zs2!&kC
OC<IS0 ZC	N
r)   c                4    | sy| dv r| S d| v ryd| v ryd| v ryyr:  r9   r;  s    r   r<  z?load_mlx_model_from_safetensors.<locals>._normalize_model_class  r=  r)   r>  r?  zSafetensors file not found: z6
Please run: python scripts/convert_to_safetensors.py zConfig file not found: z$Loading MLX model from safetensors: rr   z%Warning: Checkpoint created with MLX z, but using z. Consider reconverting.zLoading weights...zLoaded z weight tensorsr   r   r9   r   r   r   r   r   r   r   r   zReconstructing bag with z
 models...r   r   r   rC  r+   rO   u   ✓ Loaded BagOfModelsMLX with r   zReconstructing r   rD  rE  rG  )jsonr   safetensors.mlxrX  r   r   r   r   r   r   r  rw   rJ  rP  rr   r  rL  rM  r~   r   r   r-  rO  rC   rl   rv   r   r  r   r   r#  )'r   r@  r   r_  r   rX  r   r   r   r   r\  r<  rQ  config_pathr  r  weights_dictrT  r   r   r   r   r   r   bag_weightsr   r   r-   rU  rV  r   r$   demucs_kwargsmodel_prefixflat_model_stater   valueoriginal_keyr  s'     `                                    r   rK  rK    s   6 )%')	 ww||I*\/JK'',,yZL*EFK 77>>*+*+;*< =DDN<Q
 	

 77>>+&%k] 3DDN<Q
 	

 45E4FGH 
k3	11 
  6::m,o>3FJJ}4M3N O() *%&	

 
M	"o	53FJJ}4M3N O() *%&	
 "#-.LL)*/:; m,$*6N5 DHFZZ 01Nzz"45jj!23O%J**Y'K,VZZ8I-JKO++,ZL
CD z"AJ!L)a#n2E.E">!#45
+C8H4I0I/2)K*q33G/G4_Q5GHm+#Z]>+|3\],"J[.\2Z[+ 8 F!:Z	=1YZ #<[M!JKK $A3a=L!*002
U>>,/#&s<'8'9#:L5:$\2	 3 %U,<=MM% G #L %V[93J<wGH O$4#5S9:},%tS~k6/RSK-$dQnZ.PQK,4V<M#TV^I}-UVK45E4FGHH l;K 0123 ++ ''EJJL (
  	e 
 	s    P::Q__main__z-Convert HTDemucs PyTorch models to MLX format)r   )htdemucshtdemucs_fthtdemucs_6szModel to convert)choiceshelpz--output-dirz$Output directory for MLX checkpoints)defaultrn  z--verify
store_truez'Run verification tests after conversion)actionrn  z--quietzSuppress progress outputrF  c                    | sy	 t        j                  |       t        j                  |      k  S # t        $ r | |k7  cY S w xY w)NT)r   parser  )r   bs     r   rM  rM    sB    }}Q'--"222 Avs   *0 A Ac                     t        t              j                         } t        | j                        }d|v r,t        |d |j                  d        }|j                         r|S | j                  d   S )NvariantsrM   )r   __file__resolvers   r   indexrJ  parents)herer   roots      r   _model_root_dirr}    sa    >!!#DEUU3EKK
345;;=K<<?r)   c                f    | j                         }dfd	 ||       | j                  |       y)zCLoad flat weights into MLX model state (handles MLX conv wrappers).c                   t        | t              r!| j                         D ]  \  }}|xr |dv xr t        |t              }|r|}n|r| d| n|}t        |t              rJd|v xr% t        |d   t              xr d|d   v xs d|d   v }|r |d   |||       ~ ||||       t        |t              rdt	        |      D ]U  \  }	}
| d|	 }t        |
t              r-t        |
j                               dgk(  r |
d   ||d       J |
|||       W ||v s||   | |<    y t        | t              r't	        |       D ]  \  }	}
| d|	 } |
|||        y y )	N)rH   rL   rI   rO   rH   rD   rQ   )inside_sequentiallayersT)rf   r  rl   rs   r   rt   )
model_dict	flat_dictr   r  r   rg  is_sequential_convpath_for_contenthas_conv_wrapperr-   r+  idx_pathcopy_weights_from_flats               r   r  z8_load_weights_into_model.<locals>.copy_weights_from_flat  s   j$'(..0
U&7 '=%(,J%J'=%/t%< # &'-$<B&3%'8$eT*(.% )]&0v&E)]'/5='@'[FeTZmD[ % (.!&M96F.?A /!9.>.?A  t,#,U#34&6%7q#<%dD1d499;6GH:6U2 $X	8268 3 $i2CE $4 (94*34D*E
3I 1L 
D)$Z04$XQqc?&)X&79 1 *r)   N)rc   F)r,   rx   )r$   flat_weightsmodel_stater  s      @r   r   r     s/    ""$K-9^ ;5	LLr)   )T)rD   
np.ndarrayrE   rI  rB   boolr3   r  )FFN)
r   ztp.Dict[str, tp.Any]r   r  r   r  r   ztp.Optional[tp.Any]r3   r4   )F)r   r  r3   tp.Any)NFT)
r   rI  r  rH  r  r  r   r  r3   rI  )g-C6?T)r.  floatr   r  r3   r  )r   TF)
r   rI  r@  rI  rB  r  r   r  r3   r  )r   F)r   rI  r@  rI  r   r  r3   r  )r   rH  rt  rI  r3   r  )r3   r   )r  ztp.Dict[str, mx.array]),r8   
__future__r   r   r   r	  typingtpr   pathlibr   mlx.corer&  r   ro   rA   	packagingr   mlx_backendr   mlx_registryr	   r   rF   r   r   r  r  rN  rK  r5   argparseArgumentParserparseradd_argument
parse_argsr   r   r  r  quietrM  r}  r   r9   r)   r   <module>r     s   #  	        ( ,:7 :7@ <<< < 	<. '+	eW%eWeW eW %	eW
 eWT KK K` $(	ZZ Z Z 	Z
 	Z@ 	. . 	.
 
.f )	E
E
E
 E
 	E

 E
T )JJJ J 	JZ z$X$$CF :  
 #3  
 6  
 '   D??{{JJ	4r)   