
     ~j                     r   S r SSKr\R                  R                  SS5        SSKrSSKJr  SSKJr	  \	R                  " 5         \R                  " \R                  S   \R                  S9r\R                  5         \" \R                  5      S:  a  \" \R                  S   5      OS	r\" \R                  5      S
:  a  \" \R                  S
   5      OSr\R*                  " S5        \R,                  " SS\5      r\R,                  " S\S5      r\R2                  " 5          \" \\R4                  " \/5      \SS9S   rSSS5        \R9                  S5      R$                  R;                  5       R=                  5       R?                  S5      RA                  S5        \R9                  S5      R;                  5       R=                  5       R?                  S5      RA                  S5        \R9                  S5      R$                  R;                  5       R=                  5       R?                  S5      RA                  S5        \!" S\"" \RF                  5      S\S\"" \RF                  5      S\RI                  S5      RK                  5       RM                  5       RO                  5       -  5        g! , (       d  f       GNo= f)aK  Torch reference of the Music3 flow matching DiT.

Imports MiniMaxMusic3Transformer1DModel from the reference project ../../diffusers.
Dumps seeded random inputs and the predicted velocity for the GGML parity.
Run from tests/ directory. All paths relative to CWD.

Usage:
    ./mm3-dit-ref.py <checkpoint_dir> [t_latent] [timestep]
    Nz../../diffusers/src)MiniMaxMusic3Transformer1DModel)logging   )torch_dtype         gffffff?i	     i   F)return_dictfloat32zparity/dit_xt.binzparity/dit_cond.binzparity/dit_ref.binxttz-> velzrms %.6f)(__doc__syspathinserttorch8diffusers.models.transformers.transformer_minimax_music3r   diffusers.utilsr   diffusers_loggingdisable_progress_barfrom_pretrainedargvr   modelevallenintTfloatt_valmanual_seedrandnr   condno_gradtensorvelsqueeze
contiguousnumpyastypetofileprinttupleshapepowmeansqrtitem     mm3-dit-ref.py<module>r6      s    ( )  9  & & ('77QVQ^Q^_ 

CHH)Cr!#((ma/chhqkT   # 
[[C{{1a
]]_
ELL%)4U
CA
FC  

1    " " $ + +I 6 = =>Q R Q    " " $ + +I 6 = =>S T A      # # % , ,Y 7 > >?S T dE"((OS%53CZRURYRYZ[R\RaRaRcRhRhRjRoRoRqEq r _s   !J''
J6