
    xiWQ                        S r SSKrSSKrSSKrSSKJrJrJrJrJ	r	  SSK
Jr  SSKJr  SSKrSSKrSSKJrJr   SSKJr  Sr " S S\5      rS\\   S\\\R4                  4   4S jr " S S\(       a  \O\5      r " S S\5      rS\\   S\\\4   4S jr " S S\(       a  \O\5      r S\S\	\\\\4      \\\4   4   4S jr!g! \ a!    S	r\R.                  " S
5         " S S5      r Nf = f)z
PyTorch Lightning DataModule for LoRA Training

Handles data loading and preprocessing for training ACE-Step LoRA adapters.
Supports both raw audio loading and preprocessed tensor loading.
    N)OptionalListDictAnyTuple)logger)	safe_path)Dataset
DataLoader)LightningDataModuleTFz?Lightning not installed. Training module will not be available.c                       \ rS rSrSrg)r       N)__name__
__module____qualname____firstlineno____static_attributes__r       S/mnt/workspace/acestep.cpp/tests/../../ACE-Step-1.5/acestep/training/data_module.pyr   r      s    r   r   c                   |    \ rS rSrSrS\4S jrS\S\\   4S jrS\	4S jr
S	\	S\\\R                  4   4S
 jrSrg)PreprocessedTensorDataset#   a  Dataset that loads preprocessed tensor files.

This is the recommended dataset for training as all tensors are pre-computed:
- target_latents: VAE-encoded audio [T, 64]
- encoder_hidden_states: Condition encoder output [L, D]
- encoder_attention_mask: Condition mask [L]
- context_latents: Source context [T, 65]
- attention_mask: Audio latent mask [T]

No VAE/text encoder needed during training - just load tensors directly!

tensor_dirc                    [        U5      n[        R                  R                  U5      (       d  [	        SU 35      eX l        / U l        [        SU R
                  S9n[        R                  R                  U5      (       ax  [        US5       n[        R                  " U5      nSSS5        WR                  S/ 5      nU H4  nU R                  U5      nUc  M  U R                  R                  U5        M6     Os[        R                  " U R
                  5       HO  nUR                  S5      (       d  M  US:w  d  M#  U R                  R                  [        X@R
                  S95        MQ     U R                   V	s/ s H+  n	[        R                  R                  U	5      (       d  M)  U	PM-     sn	U l        [#        U R                   5      [#        U R                  5      :w  aC  [$        R&                  " S[#        U R                  5      [#        U R                   5      -
   S	35        [$        R(                  " S
[#        U R                   5       SU R
                   35        g! , (       d  f       GN= fs  sn	f )zInitialize from a directory of preprocessed .pt files.

Args:
    tensor_dir: Directory containing preprocessed .pt files and manifest.json
    
Raises:
    ValueError: If tensor_dir is not an existing directory or escapes safe root.
zNot an existing directory: zmanifest.jsonbaserNsamplesz.ptzSome tensor files not found: z missingzPreprocessedTensorDataset: z samples from )r	   ospathisdir
ValueErrorr   sample_pathsexistsopenjsonloadget_resolve_manifest_pathappendlistdirendswithvalid_pathslenr   warninginfo)
selfr   validated_dirmanifest_pathfmanifest	raw_pathsrawresolvedps
             r   __init__"PreprocessedTensorDataset.__init__0   s    "*-ww}}]++::,GHH'') "/H77>>-((mS)Q99Q< * Y3I 66s;'%%,,X6 ! ZZ0::e$$o)=%%,,!!//: 1 (,'8'8N'8!BGGNN1<MA'8Nt C(9(9$::NN/t(()C0@0@,AAB(L
 	)#d.>.>*?)@ AOO$&	
1 *)  Os   I<(I1(I1
I.r8   returnc                     [        XR                  S9n[        R                  R	                  U5      (       a  U$   [        U5      n[        R                  R	                  U5      (       a  [        R                  " SU 35        U$  [        R                  " SU 35        g! [
         a     Nrf = f! [
         a     N6f = f)u  Resolve a single manifest sample path to a validated absolute path.

Tries ``base=tensor_dir`` first (correct for new manifests that store
paths relative to the tensor directory).  If the resulting path does
not exist on disk, falls back to resolving against the global safe
root (backward compat for legacy manifests that stored CWD-relative
paths like ``./datasets/…/foo.pt``).

Returns:
    Validated absolute path, or ``None`` if the path cannot be
    resolved safely.
r   z-Resolved legacy manifest path via safe root: z%Skipping unresolvable manifest path: N)	r	   r   r    r!   r%   r#   r   debugr0   )r2   r8   childs      r   r*   0PreprocessedTensorDataset._resolve_manifest_path_   s    	c8Eww~~e$$ %	cNEww~~e$$CC5I 	 % 	>seDE  		  		s#   8B" A	B2 "
B/.B/2
B?>B?c                 ,    [        U R                  5      $ N)r/   r.   r2   s    r   __len__!PreprocessedTensorDataset.__len__   s    4##$$r   idxc           	          U R                   U   n[        R                  " USSS9nUS   US   US   US   US   UR                  S	0 5      S
.$ )zkLoad a preprocessed tensor file.

Returns:
    Dictionary containing all pre-computed tensors for training
cpuT)map_locationweights_onlytarget_latentsattention_maskencoder_hidden_statesencoder_attention_maskcontext_latentsmetadatarL   rM   rN   rO   rP   rQ   )r.   torchr(   r)   )r2   rG   tensor_pathdatas       r   __getitem__%PreprocessedTensorDataset.__getitem__   sm     &&s+zz+EM ##34"#34%)*A%B&*+C&D#$56R0
 	
r   )r$   r   r.   N)r   r   r   r   __doc__strr;   r   r*   intrE   r   rS   TensorrV   r   r   r   r   r   r   #   sY    
-
3 -
^!# !(3- !F% %
s 
tC,='> 
r   r   batchr=   c           
         [        S U  5       5      n[        S U  5       5      n/ n/ n/ n/ n/ nU  GH	  nUS   n	U	R                  S   U:  aD  U	R                  XR                  S   -
  U	R                  S   5      n
[        R                  " X/SS9n	UR                  U	5        US   nUR                  S   U:  a6  UR                  XR                  S   -
  5      n
[        R                  " X/SS9nUR                  U5        US   nUR                  S   U:  aD  UR                  XR                  S   -
  UR                  S   5      n
[        R                  " X/SS9nUR                  U5        US	   nUR                  S   U:  aD  UR                  X-R                  S   -
  UR                  S   5      n
[        R                  " X/SS9nUR                  U5        US
   nUR                  S   U:  a6  UR                  X.R                  S   -
  5      n
[        R                  " X/SS9nUR                  U5        GM     [        R                  " U5      [        R                  " U5      [        R                  " U5      [        R                  " U5      [        R                  " U5      U  Vs/ s H  oS   PM	     snS.$ s  snf )a  Collate function for preprocessed tensor batches.

Handles variable-length tensors by padding to the longest in the batch.

Args:
    batch: List of sample dictionaries with pre-computed tensors
    
Returns:
    Batched dictionary with all tensors stacked
c              3   D   #    U  H  oS    R                   S   v   M     g7f)rL   r   Nshape.0ss     r   	<genexpr>-collate_preprocessed_batch.<locals>.<genexpr>   s     Eu!+,2215u    c              3   D   #    U  H  oS    R                   S   v   M     g7f)rN   r   Nr_   ra   s     r   rd   re      s     Mu!34::1=urf   rL   r      )dimrM   rP   rN   rO   rQ   rR   )maxr`   	new_zerosrS   catr+   stack)r\   max_latent_lenmax_encoder_lenrL   attention_masksrN   encoder_attention_masksrP   sampletlpadamclehseamrc   s                   r   collate_preprocessed_batchry      sz    EuEENMuMMO NO O$%88A;',,~;RXXa[ICB9!,Bb! $%88A;',,~;<CB9!,Br" %&88A;',,~;RXXa[ICB9!,Br" ,-99Q</)--))A, >		!MC))SJA.C$$S) -.99Q</)--))A, >?C))SJA.C&&s+E J  ++n5++o6!&-B!C"'++.E"F ;;7,12Eqz]E2  3s   ,J?c                      ^  \ rS rSrSr       SS\S\S\S\S\S\S	\S
\4U 4S jjjr	SS\
\   4S jjrS\4S jrS\
\   4S jrSrU =r$ )PreprocessedDataModule   zDataModule for preprocessed tensor files.

This is the recommended DataModule for training. It loads pre-computed tensors
directly without needing VAE, text encoder, or condition encoder at training time.
r   
batch_sizenum_workers
pin_memoryprefetch_factorpersistent_workerspin_memory_device	val_splitc	                    > [         (       a  [        T	U ]	  5         Xl        X l        X0l        X@l        XPl        X`l        Xpl	        Xl
        SU l        SU l        g)a4  Initialize the data module.

Args:
    tensor_dir: Directory containing preprocessed .pt files
    batch_size: Training batch size
    num_workers: Number of data loading workers
    pin_memory: Whether to pin memory for faster GPU transfer
    val_split: Fraction of data for validation (0 = no validation)
N)LIGHTNING_AVAILABLEsuperr;   r   r}   r~   r   r   r   r   r   train_datasetval_dataset)
r2   r   r}   r~   r   r   r   r   r   	__class__s
            r   r;   PreprocessedDataModule.__init__   sR    ( G$$&$."4!2"!r   stagec                 ~   US:X  d  Uc  [        U R                  5      nU R                  S:  a  [        U5      S:  ar  [	        S[        [        U5      U R                  -  5      5      n[        U5      U-
  n[        R                  R                  R                  X$U/5      u  U l
        U l        gX l
        SU l        gg)zSetup datasets.fitNr   rh   )r   r   r   r/   rj   rZ   rS   utilsrU   random_splitr   r   )r2   r   full_datasetn_valn_trains        r   setupPreprocessedDataModule.setup  s    E>U]4T__EL ~~!c,&7!&;As3|#4t~~#EFGl+e37<{{7G7G7T7T E"284"D$4 &2"#'  +r   r=   c                 J   U R                   S:X  a  SOU R                  nU R                   S:X  a  SOU R                  n[        U R                  U R
                  SU R                   U R                  [        SUUS9	nU R                  (       a  U R                  US'   [        S0 UD6$ )zCreate training dataloader.r   NFT)	datasetr}   shuffler~   r   
collate_fn	drop_lastr   r   r   r   )
r~   r   r   dictr   r}   r   ry   r   r   r2   r   r   kwargss       r   train_dataloader'PreprocessedDataModule.train_dataloader  s    "&"2"2a"7$T=Q=Q&*&6&6!&;UAXAX&&((1+1

 !!*.*@*@F&'#F##r   c                 d   U R                   c  gU R                  S:X  a  SOU R                  nU R                  S:X  a  SOU R                  n[	        U R                   U R
                  SU R                  U R                  [        UUS9nU R                  (       a  U R                  US'   [        S0 UD6$ )zCreate validation dataloader.Nr   F)r   r}   r   r~   r   r   r   r   r   r   )
r   r~   r   r   r   r}   r   ry   r   r   r   s       r   val_dataloader%PreprocessedDataModule.val_dataloader+  s    #"&"2"2a"7$T=Q=Q&*&6&6!&;UAXAX$$((1+1	
 !!*.*@*@F&'#F##r   )
r}   r~   r   r   r   r   r   r   r   r   )rh      T   T         rC   )r   r   r   r   rX   rY   rZ   boolfloatr;   r   r   r   r   r   r   __classcell__r   s   @r   r{   r{      s      #'!#! !  !  	! 
 !  !  !!  !  !  ! F(8C= ($$* $&$ 4 $ $r   r{   c                       \ rS rSrSr  SS\\\\4      S\	S\
4S jjrS\\\\4      4S jrS\
4S	 jrS
\
S\\\R                  4   4S jrSrg)AceStepTrainingDatasetiD  aj  Dataset for ACE-Step LoRA training from raw audio.

DEPRECATED: Use PreprocessedTensorDataset instead for better performance.

Audio Format Requirements (handled automatically):
- Sample rate: 48kHz (resampled if different)
- Channels: Stereo (2 channels, mono is duplicated)
- Max duration: 240 seconds (4 minutes)
- Min duration: 5 seconds (padded if shorter)
r   max_durationtarget_sample_ratec                     Xl         X l        X0l        X@l        U R	                  5       U l        [        R                  " S[        U R
                  5       S35        g)zInitialize the dataset.zDataset initialized with z valid samplesN)	r   dit_handlerr   r   _validate_samplesvalid_samplesr   r1   r/   )r2   r   r   r   r   s        r   r;   AceStepTrainingDataset.__init__P  sN     &("4!335/D4F4F0G/HWXr   r=   c                 >   / n[        U R                  5       H  u  p#UR                  SS5      nU(       d  [        R                  " SU S35        M:   [        U5      n[        R                  R                  U5      (       d  [        R                  " SU SU 35        M  UR                  S5      (       d  [        R                  " SU S35        M  0 UESU0EnUR                  U5        M     U$ ! [         a!    [        R                  " SU SU 35         GM  f = f)	zAValidate and filter samples, resolving audio paths to safe paths.
audio_pathr   zSample z: Missing audio_pathz: Rejected unsafe path: z: Audio file not found: captionz: Missing caption)	enumerater   r)   r   r0   r	   r#   r    r!   isfiler+   )r2   validirr   r   	validateds         r   r   (AceStepTrainingDataset._validate_samples`  s   "4<<0IAL"5J+?@A%j1	
 77>>),,+CJ<PQ::i((+<=> 98i8FLL - 10 !  +CJ<PQs   C11&DDc                 ,    [        U R                  5      $ rC   )r/   r   rD   s    r   rE   AceStepTrainingDataset.__len__}  s    4%%&&r   rG   c                 F   U R                   U   nUS   n[        R                  " U5      u  pEXPR                  :w  a1  [        R                  R                  XPR                  5      nU" U5      nUR                  S   S:X  a  UR                  SS5      nOUR                  S   S:  a  USS2SS24   n[        U R                  U R                  -  5      nUR                  S   U:  a  USS2SU24   n[        SU R                  -  5      nUR                  S   U:  a=  XR                  S   -
  n	[        R                  R                  R                  USU	45      nUUR                  SS5      UR                  S	S
5      UR                  SS5      UR                  S	S
5      UR                  S5      UR                  SS5      UR                  SS5      UR                  SUR                  S   U R                  -  5      UR                  SS5      UR                  SS5      S.US.$ )zGet a single training sample.r   r   rh   r   Ng      @r   r   lyricsz[Instrumental]bpmkeyscaletimesignaturedurationlanguageunknownis_instrumentalT)r   r   r   r   r   r   r   r   )audior   r   rQ   r   )r   
torchaudior(   r   
transformsResampler`   repeatrZ   r   rS   nn
functionalrt   r)   )
r2   rG   rr   r   r   sr	resamplermax_samplesmin_samplespaddings
             r   rV   "AceStepTrainingDataset.__getitem__  s   ##C(L)
OOJ/	 ((("--66r;R;RSIe$E ;;q>QLLA&E[[^a"1"a%LE $++d.E.EEF;;q>K'!\k\/*E# 7 778;;q>K'!KKN2GHH''++EAw<@E zz)R0jj+;<!::i4 **X/?@zz%("JJz26!'OR!@"JJz5;;q>DD[D[3[\"JJz9=#)::.?#F	 %
 	
r   )r   r   r   r   r   N)      n@i  )r   r   r   r   rX   r   r   rY   r   r   rZ   r;   r   rE   rS   r[   rV   r   r   r   r   r   r   D  s    	 $"'Yd38n%Y 	Y
  Y 4S#X#7 :' '+
s +
tC,='> +
r   r   c           
      v   [        S U  5       5      n/ n/ nU  H  nUS   nUR                  S   nXa:  a0  X-
  n[        R                  R                  R                  USU45      nUR                  U5        [        R                  " U5      nXa:  a  SXS& UR                  U5        M     [        R                  " U5      [        R                  " U5      U  V	s/ s H  oS   PM	     sn	U  V	s/ s H  oS   PM	     sn	U  V	s/ s H  oS   PM	     sn	U  V	s/ s H  oS	   PM	     sn	S
.$ s  sn	f s  sn	f s  sn	f s  sn	f )z0Collate function for raw audio batches (legacy).c              3   D   #    U  H  oS    R                   S   v   M     g7f)r   rh   Nr_   )rb   rr   s     r   rd   )collate_training_batch.<locals>.<genexpr>  s     ?v/''*rf   r   rh   r   Nr   r   rQ   r   )r   rM   captionsr   rQ   audio_paths)	rj   r`   rS   r   r   rt   r+   onesrm   )
r\   max_lenpadded_audiorp   rr   r   	audio_lenr   maskrc   s
             r   collate_training_batchr     s&   ???GLOwKKN	)GHH''++EAw<@EE"zz'" Dt$   \*++o6+015ay\51(-.1X;.,12Eqz]E2167A,7  2.27s   D',D,
 D1D6c                      ^  \ rS rSrSr     SS\\\\4      S\	S\	S\
S\S\4U 4S	 jjjrSS
\\   4S jjrS\4S jrS\\   4S jrSrU =r$ )AceStepDataModulei  ztDataModule for raw audio loading (legacy).

DEPRECATED: Use PreprocessedDataModule for better training performance.
r   r}   r~   r   r   r   c                    > [         (       a  [        TU ]	  5         Xl        X l        X0l        X@l        XPl        X`l        Xpl	        S U l
        S U l        g rC   )r   r   r;   r   r   r}   r~   r   r   r   r   r   )	r2   r   r   r}   r~   r   r   r   r   s	           r   r;   AceStepDataModule.__init__  sL     G&$&$("!r   r   c                    US:X  d  UGce  U R                   S:  Ga  [        U R                  5      S:  Ga  [        S[	        [        U R                  5      U R                   -  5      5      n[        [        [        U R                  5      5      5      n[        R                  " U5        US U nX2S  nU Vs/ s H  o`R                  U   PM     nnU Vs/ s H  o`R                  U   PM     nn[        XpR                  U R                  5      U l        [        XR                  U R                  5      U l        g [        U R                  U R                  U R                  5      U l        S U l        g g s  snf s  snf )Nr   r   rh   )r   r/   r   rj   rZ   listrangerandomr   r   r   r   r   r   )	r2   r   r   indicesval_indicestrain_indicesr   train_samplesval_sampless	            r   r   AceStepDataModule.setup  s4   E>U]~~!c$,,&7!&;As3t||#4t~~#EFGuS%678w'%fuo ':G H-Qa- H8CD1||AD%;!#3#3T5F5F&" $:!1!143D3D$  &<LL$"2"2D4E4E&" $( / + !IDs   5E1E6r=   c           
      x    [        U R                  U R                  SU R                  U R                  [
        SS9$ )NT)r}   r   r~   r   r   r   )r   r   r}   r~   r   r   rD   s    r   r   "AceStepDataModule.train_dataloader  s8    ((-
 	
r   c           	          U R                   c  g [        U R                   U R                  SU R                  U R                  [
        S9$ )NF)r}   r   r~   r   r   )r   r   r}   r~   r   r   rD   s    r   r    AceStepDataModule.val_dataloader  sD    #((-
 	
r   )	r}   r   r   r~   r   r   r   r   r   )rh   r   Tr   r   rC   )r   r   r   r   rX   r   r   rY   r   rZ   r   r   r;   r   r   r   r   r   r   r   r   s   @r   r   r     s     # d38n%  	 
          0(8C= (4	
* 	

 4 
 
r   r   	json_pathc                 @   [        U 5      n[        R                  R                  U5      (       d  [	        SU  35      e[        USSS9 n[        R                  " U5      nSSS5        WR                  S0 5      nUR                  S/ 5      nXT4$ ! , (       d  f       N5= f)zLoad a dataset from JSON file.

Args:
    json_path: Path to the JSON dataset file.

Returns:
    Tuple of (samples list, metadata dict).

Raises:
    ValueError: If json_path does not point to an existing file or escapes safe root.
zDataset JSON file not found: r   zutf-8)encodingNrQ   r   )	r	   r    r!   r   r#   r&   r'   r(   r)   )r   r   r5   rU   rQ   r   s         r   load_dataset_from_jsonr     s     )$I77>>)$$8DEE	iw	/1yy| 
0 xx
B'Hhhy"%G 
0	/s   	B
B)"rX   r    r'   r   typingr   r   r   r   r   logurur   acestep.training.path_safetyr	   rS   r   torch.utils.datar
   r   lightning.pytorchr   r   ImportErrorr0   r   rY   r[   ry   objectr{   r   r   r   r   r   r   r   <module>r      s+   
   3 3  2   05r
 r
jAd4j AT#u||:K5L AHa$4G0V a$Pg
W g
T$t* c3h @N
/B+ N
bc eDc3h4H$sTWx.4X.Y Q  
NNTU 	s   C $C/.C/