
    xiWQ                        d Z ddlZddlZddlZddlmZmZmZmZm	Z	 ddl
mZ ddlmZ ddlZddlZddlmZmZ 	 ddlmZ dZ G d de      Zdee   deeej4                  f   fdZ G d derene      Z G d de      Zdee   deeef   fdZ G d derene      Z dede	eeeef      eeef   f   fdZ!y# e$ r! d	Z ej.                  d
        G d d      ZY w xY w)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                       e Zd Zy)r   N)__name__
__module____qualname__     ;/mnt/workspace/ACE-Step-1.5/acestep/training/data_module.pyr   r      s    r   r   c                   l    e Zd ZdZdefdZdedee   fdZdefdZ	dede
eej                  f   fd	Zy
)PreprocessedTensorDataseta  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                    t        |      }t        j                  j                  |      st	        d|       || _        g | _        t        d| j
                        }t        j                  j                  |      rst        |d      5 }t        j                  |      }ddd       j                  dg       }|D ]1  }| j                  |      }|| j                  j                  |       3 nlt        j                  | j
                        D ]J  }|j                  d      s|dk7  s| j                  j                  t        || j
                               L | j                  D 	cg c]$  }	t        j                  j                  |	      s#|	& c}	| _        t#        | j                         t#        | j                        k7  rBt%        j&                  dt#        | j                        t#        | j                         z
   d	       t%        j(                  d
t#        | j                          d| j
                          y# 1 sw Y   xY wc c}	w )a"  Initialize 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__z"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*$IIIr4   returnc                 p   	 t        || j                        }t        j                  j	                  |      r|S 	 	 t        |      }t        j                  j	                  |      rt        j                  d|        |S 	 t        j                  d|        y# t
        $ r Y jw xY w# t
        $ r Y 3w xY w)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   debugr,   )r.   r4   childs      r   r&   z0PreprocessedTensorDataset._resolve_manifest_path_   s    	c8Eww~~e$ %	cNEww~~e$CC5I 	 % 	>seDE  		  		s#   7B AB) 	B&%B&)	B54B5c                 ,    t        | j                        S N)r+   r*   r.   s    r   __len__z!PreprocessedTensorDataset.__len__   s    4##$$r   idxc           	          | j                   |   }t        j                  |dd      }|d   |d   |d   |d   |d   |j                  d	i       d
S )zLoad 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rE   rF   rG   rH   rI   rJ   )r*   torchr$   r%   )r.   r@   tensor_pathdatas       r   __getitem__z%PreprocessedTensorDataset.__getitem__   sm     &&s+zz+EM ##34"#34%)*A%B&*+C&D#$56R0
 	
r   N)r   r   r   __doc__strr7   r   r&   intr?   r   rL   TensorrO   r   r   r   r   r   #   sY    
-
3 -
^!# !(3- !F% %
s 
tC,='> 
r   r   batchr8   c           
         t        d | D              }t        d | D              }g }g }g }g }g }| D ]  }|d   }	|	j                  d   |k  rH|	j                  ||	j                  d   z
  |	j                  d         }
t        j                  |	|
gd      }	|j                  |	       |d   }|j                  d   |k  r:|j                  ||j                  d   z
        }
t        j                  ||
gd      }|j                  |       |d   }|j                  d   |k  rH|j                  ||j                  d   z
  |j                  d         }
t        j                  ||
gd      }|j                  |       |d	   }|j                  d   |k  rH|j                  ||j                  d   z
  |j                  d         }
t        j                  ||
gd      }|j                  |       |d
   }|j                  d   |k  r:|j                  ||j                  d   z
        }
t        j                  ||
gd      }|j                  |        t        j                  |      t        j                  |      t        j                  |      t        j                  |      t        j                  |      | D cg c]  }|d   	 c}dS c c}w )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   @   K   | ]  }|d    j                   d     yw)rE   r   Nshape.0ss     r   	<genexpr>z-collate_preprocessed_batch.<locals>.<genexpr>   s"     Eu!+,2215u   c              3   @   K   | ]  }|d    j                   d     yw)rG   r   NrW   rY   s     r   r\   z-collate_preprocessed_batch.<locals>.<genexpr>   s"     Mu!!34::1=ur]   rE   r      )dimrF   rI   rG   rH   rJ   rK   )maxrX   	new_zerosrL   catr'   stack)rT   max_latent_lenmax_encoder_lenrE   attention_masksrG   encoder_attention_masksrI   sampletlpadamclehseamr[   s                   r   collate_preprocessed_batchrp      s    EuEENMuMMO NO O$%88A;',,~;RXXa[ICB9!,Bb! $%88A;',,~;<CB9!,Br" %&88A;',,~;RXXa[ICB9!,Br" ,-99Q</)--#))A, >		!MC))S#JA.C$$S) -.99Q</)--#))A, >?C))S#JA.C&&s+E J  ++n5++o6!&-B!C"'++.E"F ;;7,12EqQz]E2  3s   4Kc                        e Zd ZdZ	 	 	 	 	 	 	 ddededededededed	ef fd
Zdde	e   fdZ
defdZde	e   fdZ xZS )PreprocessedDataModulezDataModule 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	                     t         rt        	| 	          || _        || _        || _        || _        || _        || _        || _	        || _
        d| _        d| _        y)at  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superr7   r   rs   rt   ru   rv   rw   rx   ry   train_datasetval_dataset)
r.   r   rs   rt   ru   rv   rw   rx   ry   	__class__s
            r   r7   zPreprocessedDataModule.__init__   s_    ( G$$&$."4!2"!r   stagec                 z   |dk(  s|t        | j                        }| j                  dkD  rt        |      dkD  rst	        dt        t        |      | j                  z              }t        |      |z
  }t        j                  j                  j                  |||g      \  | _
        | _        y|| _
        d| _        yy)zSetup datasets.fitNr   r_   )r   r   ry   r+   ra   rR   rL   utilsrN   random_splitr}   r~   )r.   r   full_datasetn_valn_trains        r   setupzPreprocessedDataModule.setup  s    E>U]4T__EL ~~!c,&7!&;As3|#4t~~#EFGl+e37<{{7G7G7T7T 7E"284"D$4 &2"#'  +r   r8   c                 B   | j                   dk(  rdn| j                  }| j                   dk(  rdn| j                  }t        | j                  | j
                  d| j                   | j                  t        d||	      }| j                  r| j                  |d<   t        di |S )zCreate training dataloader.r   NFT)	datasetrs   shufflert   ru   
collate_fn	drop_lastrv   rw   rx   r   )
rt   rv   rw   dictr}   rs   ru   rp   rx   r   r.   rv   rw   kwargss       r   train_dataloaderz'PreprocessedDataModule.train_dataloader  s    "&"2"2a"7$T=Q=Q&*&6&6!&;UAXAX&&((1+1

 !!*.*@*@F&'#F##r   c           
      Z   | j                   y| j                  dk(  rdn| j                  }| j                  dk(  rdn| j                  }t	        | j                   | j
                  d| j                  | j                  t        ||      }| j                  r| j                  |d<   t        di |S )zCreate validation dataloader.Nr   F)r   rs   r   rt   ru   r   rv   rw   rx   r   )
r~   rt   rv   rw   r   rs   ru   rp   rx   r   r   s       r   val_dataloaderz%PreprocessedDataModule.val_dataloader+  s    #"&"2"2a"7$T=Q=Q&*&6&6!&;UAXAX$$((1+1	
 !!*.*@*@F&'#F##r   )r_      T   T         r=   )r   r   r   rP   rQ   rR   boolfloatr7   r   r   r   r   r   __classcell__r   s   @r   rr   rr      s      #'!#! !  !  	! 
 !  !  !!  !  ! F(8C= ($$* $&$ 4 $r   rr   c                       e Zd ZdZ	 	 ddeeeef      dede	fdZ
deeeef      fdZde	fdZd	e	deeej                  f   fd
Zy)AceStepTrainingDataseta  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                     || _         || _        || _        || _        | j	                         | _        t        j                  dt        | j
                         d       y)zInitialize the dataset.zDataset initialized with z valid samplesN)	r   dit_handlerr   r   _validate_samplesvalid_samplesr   r-   r+   )r.   r   r   r   r   s        r   r7   zAceStepTrainingDataset.__init__P  sV     &("4!335/D4F4F0G/HWXr   r8   c                 
   g }t        | j                        D ]  \  }}|j                  dd      }|st        j                  d| d       4	 t        |      }t        j                  j                  |      st        j                  d| d|        {|j                  d      st        j                  d| d       i |d|i}|j                  |        |S # t        $ r t        j                  d| d|        Y w xY w)	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   r,   r	   r   r   r   isfiler'   )r.   validiri   r   	validateds         r   r   z(AceStepTrainingDataset._validate_samples`  s   "4<<0IAvL"5J+?@A%j1	
 77>>),+CJ<PQ::i(+<=> 98i8FLL - 10 !  +CJ<PQs   C$DDc                 ,    t        | j                        S r=   )r+   r   r>   s    r   r?   zAceStepTrainingDataset.__len__}  s    4%%&&r   r@   c                 B   | j                   |   }|d   }t        j                  |      \  }}|| j                  k7  r2t        j                  j                  || j                        } ||      }|j                  d   dk(  r|j                  dd      }n|j                  d   dkD  r|ddddf   }t        | j                  | j                  z        }|j                  d   |kD  r|ddd|f   }t        d| j                  z        }|j                  d   |k  r>||j                  d   z
  }	t        j                  j                  j                  |d|	f      }||j                  dd      |j                  d	d
      |j                  dd      |j                  d	d
      |j                  d      |j                  dd      |j                  dd      |j                  d|j                  d   | j                  z        |j                  dd      |j                  dd      d|dS )zGet a single training sample.r   r   r_   r   Ng      @r   r   lyricsz[Instrumental]bpmkeyscaletimesignaturedurationlanguageunknownis_instrumentalT)r   r   r   r   r   r   r   r   )audior   r   rJ   r   )r   
torchaudior$   r   
transformsResamplerX   repeatrR   r   rL   nn
functionalrk   r%   )
r.   r@   ri   r   r   sr	resamplermax_samplesmin_samplespaddings
             r   rO   z"AceStepTrainingDataset.__getitem__  s   ##C(L)
OOJ/	r ((("--66r4;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'!EKKN2GHH''++EAw<@E zz)R0jj+;<!::i4 **X/?@zz%("JJz26!'OR!@"JJz5;;q>DD[D[3[\"JJz9=#)::.?#F	 %
 	
r   N)      n@i  )r   r   r   rP   r   r   rQ   r   r   rR   r7   r   r?   rL   rS   rO   r   r   r   r   r   D  s    	 $"'Yd38n%Y 	Y
  Y 4S#X#7 :' '+
s +
tC,='> +
r   r   c           
      `   t        d | D              }g }g }| D ]  }|d   }|j                  d   }||k  r1||z
  }t        j                  j                  j                  |d|f      }|j                  |       t        j                  |      }||k  rd||d |j                  |        t        j                  |      t        j                  |      | D 	cg c]  }	|	d   	 c}	| D 	cg c]  }	|	d   	 c}	| D 	cg c]  }	|	d   	 c}	| D 	cg c]  }	|	d	   	 c}	d
S c c}	w c c}	w c c}	w c c}	w )z0Collate function for raw audio batches (legacy).c              3   @   K   | ]  }|d    j                   d     yw)r   r_   NrW   )rZ   ri   s     r   r\   z)collate_training_batch.<locals>.<genexpr>  s      ?v&/''*r]   r   r_   r   Nr   r   rJ   r   )r   rF   captionsr   rJ   audio_paths)	ra   rX   rL   r   r   rk   r'   onesrd   )
rT   max_lenpadded_audiorg   ri   r   	audio_lenr   maskr[   s
             r   collate_training_batchr     s6   ???GLOwKKN	w	)GHH''++EAw<@EE"zz'"w Dt$   \*++o6+015aQy\51(-.11X;.,12EqQz]E2167A,7  2.27s   D'D!
9D&D+c                        e Zd ZdZ	 	 	 	 	 ddeeeef      dedede	de
de
f fdZdd	ee   fd
ZdefdZdee   fdZ xZS )AceStepDataModulezDataModule for raw audio loading (legacy).
    
    DEPRECATED: Use PreprocessedDataModule for better training performance.
    r   rs   rt   ru   r   ry   c                     t         rt        | 	          || _        || _        || _        || _        || _        || _        || _	        d | _
        d | _        y r=   )r{   r|   r7   r   r   rs   rt   ru   r   ry   r}   r~   )	r.   r   r   rs   rt   ru   r   ry   r   s	           r   r7   zAceStepDataModule.__init__  sW     G&$&$("!r   r   c                    |dk(  s|a| j                   dkD  rt        | j                        dkD  r t        dt	        t        | j                        | j                   z              }t        t        t        | j                                    }t        j                  |       |d | }||d  }|D cg c]  }| j                  |    }}|D cg c]  }| j                  |    }}t        || j                  | j                        | _        t        || j                  | j                        | _        y t        | j                  | j                  | j                        | _        d | _        y y c c}w c c}w )Nr   r   r_   )ry   r+   r   ra   rR   listrangerandomr   r   r   r   r}   r~   )	r.   r   r   indicesval_indicestrain_indicesr   train_samplesval_sampless	            r   r   zAceStepDataModule.setup  s>   E>U]~~!c$,,&7!&;As3t||#4t~~#EFGuS%678w'%fuo ':G H-Qa- H8CD1t||AD%;!4#3#3T5F5F&" $:!1!143D3D$  &<LL$"2"2D4E4E&" $( / + !IDs   1E+E0r8   c           	      ~    t        | j                  | j                  d| j                  | j                  t
        d      S )NT)rs   r   rt   ru   r   r   )r   r}   rs   rt   ru   r   r>   s    r   r   z"AceStepDataModule.train_dataloader  s8    ((-
 	
r   c                     | j                   y t        | j                   | j                  d| j                  | j                  t
              S )NF)rs   r   rt   ru   r   )r~   r   rs   rt   ru   r   r>   s    r   r   z AceStepDataModule.val_dataloader  sD    #((-
 	
r   )r_   r   Tr   r   r=   )r   r   r   rP   r   r   rQ   r   rR   r   r   r7   r   r   r   r   r   r   r   s   @r   r   r     s     # d38n%  	 
        0(8C= (4	
* 	

 4 
r   r   	json_pathc                 2   t        |       }t        j                  j                  |      st	        d|        t        |dd      5 }t        j                  |      }ddd       j                  di       }|j                  dg       }||fS # 1 sw Y   1xY w)a  Load 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)encodingNrJ   r   )	r	   r   r   r   r   r"   r#   r$   r%   )r   r   r1   rN   rJ   r   s         r   load_dataset_from_jsonr     s     )$I77>>)$8DEE	iw	/1yy| 
0 xx
B'Hhhy"%GH 
0	/s   BB)"rP   r   r#   r   typingr   r   r   r   r   logurur   acestep.training.path_safetyr	   rL   r   torch.utils.datar
   r   lightning.pytorchr   r{   ImportErrorr,   r   rQ   rS   rp   objectrr   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  FNNTU 	s   B; ;#C! C!