CoolFace
Apppublic

BlueWaveSemi45/DramaboxCPU

sourceHugging Faceotherupdated 4mo agoView on Hugging Face
0likes
preprocess.py352 linesDownload Raw Back to src
1#!/usr/bin/env python32"""3Preprocess TTS datasets for LTX-2.3 audio-only LoRA fine-tuning.4 5Takes paired (audio, transcript) data and produces the format expected by6the LTX trainer:7    .precomputed/8    ├── latents/sample_N.pt         # Dummy video latents (minimal)9    ├── conditions/sample_N.pt      # Text embeddings from Gemma10    └── audio_latents/sample_N.pt   # Audio VAE-encoded latents11 12Supports multiple dataset formats:13  - gemini_synthetic: index.txt with ~-separated fields (id~speaker~lang~sr~samples~dur~phonemes~text)14  - libriheavy: index_ft.txt with ~-separated fields (id~speaker~lang~samples~dur~phonemes~text)15  - manifest: JSON/JSONL with {"audio_filepath": ..., "text": ...}16  - tsv: TSV file with audio_path<TAB>text columns17 18Usage:19    python preprocess_tts_data.py \20        --dataset-type gemini_synthetic \21        --index /mnt/large-datasets/gemini_synthetic_dataset/conversational_dataset_pp/index.txt \22        --audio-dir /mnt/large-datasets/gemini_synthetic_dataset/conversational_dataset_pp/wavs \23        --output-dir /mnt/persistent0/manmay/tts_training_data \24        --max-samples 10000 \25        --max-duration 20.0 \26        --min-duration 3.027"""28 29import argparse30import json31import logging32import os33import sys34from pathlib import Path35 36import torch37import torchaudio38 39REPO_DIR = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))40sys.path.insert(0, os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "ltx2"))41# ltx-pipelines on path via ltx2/42 43MODEL_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))44GEMMA_DIR = os.environ.get("GEMMA_DIR", "gemma-3-12b-it-qat-q4_0-unquantized")45 46 47def parse_args():48    p = argparse.ArgumentParser(description="Preprocess TTS data for LTX-2.3 fine-tuning")49    p.add_argument("--dataset-type", required=True,50                   choices=["gemini_synthetic", "libriheavy", "manifest", "tsv"],51                   help="Dataset format type")52    p.add_argument("--index", required=True, help="Path to index/manifest file")53    p.add_argument("--audio-dir", default=None,54                   help="Base directory for audio files (if paths in index are relative)")55    p.add_argument("--output-dir", required=True, help="Output directory for preprocessed data")56    p.add_argument("--checkpoint", default=os.path.join(MODEL_DIR, "ltx-2.3-22b-distilled.safetensors"))57    p.add_argument("--gemma-root", default=GEMMA_DIR)58    p.add_argument("--max-samples", type=int, default=0, help="Max samples to process (0=all)")59    p.add_argument("--max-duration", type=float, default=20.0, help="Max audio duration in seconds")60    p.add_argument("--min-duration", type=float, default=2.0, help="Min audio duration in seconds")61    p.add_argument("--batch-size", type=int, default=8, help="Batch size for text encoding")62    p.add_argument("--skip-existing", action="store_true", help="Skip already processed samples")63    p.add_argument("--audio-only-ckpt", default=None,64                   help="Audio-only checkpoint for VAE encoding (optional, uses full ckpt if not set)")65    p.add_argument("--shard", type=int, default=0, help="Shard index (for parallel processing)")66    p.add_argument("--num-shards", type=int, default=1, help="Total number of shards")67    p.add_argument("--gpu", type=int, default=None, help="GPU device index to use")68    return p.parse_args()69 70 71def parse_gemini_synthetic(index_path: str, audio_dir: str | None) -> list[dict]:72    """Parse gemini_synthetic format: id~speaker~lang~sr~samples~dur~phonemes~text"""73    samples = []74    with open(index_path) as f:75        for line in f:76            parts = line.strip().split("~")77            if len(parts) < 7:78                continue79            file_id = parts[0]80            text = parts[-1]  # Last field is always the text81            sr = int(parts[3])82            n_samples = int(parts[4])83            duration = n_samples / sr84 85            # Find audio file86            if audio_dir:87                # Try common extensions88                for ext in [".flac", ".wav", ".mp3"]:89                    audio_path = os.path.join(audio_dir, file_id + ext)90                    if os.path.exists(audio_path):91                        break92                else:93                    continue94            else:95                audio_path = file_id96 97            samples.append({98                "id": file_id,99                "audio_path": audio_path,100                "text": text,101                "duration": duration,102            })103    return samples104 105 106def parse_libriheavy(index_path: str, audio_dir: str | None) -> list[dict]:107    """Parse libriheavy format: id~speaker~lang~samples~dur~phonemes~text"""108    samples = []109    with open(index_path) as f:110        for line in f:111            parts = line.strip().split("~")112            if len(parts) < 7:113                continue114            file_id = parts[0]115            text = parts[-1]116            n_samples = int(parts[3])117            duration = int(parts[4]) / 1000.0  # milliseconds to seconds118 119            if audio_dir:120                for ext in [".flac", ".wav", ".mp3"]:121                    audio_path = os.path.join(audio_dir, file_id + ext)122                    if os.path.exists(audio_path):123                        break124                else:125                    continue126            else:127                audio_path = file_id128 129            samples.append({130                "id": file_id,131                "audio_path": audio_path,132                "text": text,133                "duration": duration,134            })135    return samples136 137 138def parse_manifest(index_path: str, audio_dir: str | None) -> list[dict]:139    """Parse JSON/JSONL manifest with audio_filepath and text fields."""140    samples = []141    with open(index_path) as f:142        for line in f:143            entry = json.loads(line.strip())144            audio_path = entry.get("audio_filepath", entry.get("audio_path", ""))145            text = entry.get("text", entry.get("transcript", ""))146            duration = entry.get("duration", 0.0)147 148            if audio_dir and not os.path.isabs(audio_path):149                audio_path = os.path.join(audio_dir, audio_path)150 151            if os.path.exists(audio_path) and text:152                samples.append({153                    "id": Path(audio_path).stem,154                    "audio_path": audio_path,155                    "text": text,156                    "duration": duration,157                })158    return samples159 160 161def parse_tsv(index_path: str, audio_dir: str | None) -> list[dict]:162    """Parse TSV file with audio_path<TAB>text."""163    samples = []164    with open(index_path) as f:165        for line in f:166            parts = line.strip().split("\t")167            if len(parts) < 2:168                continue169            audio_path, text = parts[0], parts[1]170            if audio_dir and not os.path.isabs(audio_path):171                audio_path = os.path.join(audio_dir, audio_path)172            if os.path.exists(audio_path):173                samples.append({174                    "id": Path(audio_path).stem,175                    "audio_path": audio_path,176                    "text": text,177                    "duration": 0.0,178                })179    return samples180 181 182PARSERS = {183    "gemini_synthetic": parse_gemini_synthetic,184    "libriheavy": parse_libriheavy,185    "manifest": parse_manifest,186    "tsv": parse_tsv,187}188 189 190@torch.inference_mode()191def main():192    logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")193    args = parse_args()194 195    from ltx_core.model.audio_vae import encode_audio as vae_encode_audio196    from ltx_core.types import Audio197    from ltx_pipelines.utils.blocks import AudioConditioner198    from ltx_pipelines.utils.media_io import decode_audio_from_file199    from ltx_trainer.model_loader import load_text_encoder, load_embeddings_processor200 201    if args.gpu is not None:202        os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu)203    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")204    dtype = torch.bfloat16205 206    # Create output directories207    out = Path(args.output_dir)208    (out / "latents").mkdir(parents=True, exist_ok=True)209    (out / "conditions").mkdir(parents=True, exist_ok=True)210    (out / "audio_latents").mkdir(parents=True, exist_ok=True)211 212    # Parse dataset213    logging.info(f"Parsing {args.dataset_type} dataset from {args.index}...")214    samples = PARSERS[args.dataset_type](args.index, args.audio_dir)215    logging.info(f"Found {len(samples)} samples")216 217    # Filter by duration218    before = len(samples)219    samples = [s for s in samples if args.min_duration <= s["duration"] <= args.max_duration]220    logging.info(f"After duration filter [{args.min_duration}s, {args.max_duration}s]: {len(samples)} (dropped {before - len(samples)})")221 222    if args.max_samples > 0:223        samples = samples[:args.max_samples]224        logging.info(f"Limiting to {len(samples)} samples")225 226    # Assign global indices before sharding227    for i, s in enumerate(samples):228        s["global_idx"] = i229 230    # Shard the data for parallel processing231    if args.num_shards > 1:232        total = len(samples)233        samples = samples[args.shard::args.num_shards]234        logging.info(f"Shard {args.shard}/{args.num_shards}: {len(samples)} samples (of {total} total)")235 236    # ── Step 1: Encode text with Gemma (Blocks 1+2 only) ──237    # The trainer runs Block 3 (embeddings processor/connectors) during training,238    # so we only precompute Blocks 1+2 here (Gemma LLM + feature extractor).239    logging.info("Loading text encoder (Gemma + feature extractor)...")240    text_encoder = load_text_encoder(args.gemma_root, device=device, dtype=dtype)241 242    # Load feature extractor on CPU first to save GPU memory, then move to device243    logging.info("Loading feature extractor (on CPU first to save GPU memory)...")244    emb_proc = load_embeddings_processor(args.checkpoint, device="cpu", dtype=dtype)245    text_encoder.feature_extractor = emb_proc.feature_extractor.to(device)246    del emb_proc247    torch.cuda.empty_cache()248 249    logging.info("Encoding text prompts (Blocks 1+2: Gemma + feature extractor)...")250    for i, sample in enumerate(samples):251        gidx = sample["global_idx"]252        cond_path = out / "conditions" / f"sample_{gidx:06d}.pt"253        if args.skip_existing and cond_path.exists():254            continue255 256        text = sample["text"]257        # Run Blocks 1+2: Gemma LLM → feature extractor258        hidden_states, attention_mask = text_encoder.encode(text)259        video_feats, audio_feats = text_encoder.feature_extractor(260            hidden_states, attention_mask, "left"261        )262 263        torch.save({264            "video_prompt_embeds": video_feats.squeeze(0).cpu(),265            "audio_prompt_embeds": audio_feats.squeeze(0).cpu() if audio_feats is not None else video_feats.squeeze(0).cpu(),266            "prompt_attention_mask": attention_mask.squeeze(0).bool().cpu(),267        }, cond_path)268 269        if i % 100 == 0:270            logging.info(f"  Text encoding: {i}/{len(samples)}")271 272    del text_encoder273    torch.cuda.empty_cache()274 275    # ── Step 2: Encode audio with Audio VAE ──276    ckpt_for_vae = args.audio_only_ckpt or args.checkpoint277    logging.info(f"Loading audio VAE from {ckpt_for_vae}...")278 279    ac = AudioConditioner(checkpoint_path=ckpt_for_vae, dtype=dtype, device=device)280 281    logging.info("Encoding audio samples...")282    for idx, sample in enumerate(samples):283        gidx = sample["global_idx"]284        audio_path = out / "audio_latents" / f"sample_{gidx:06d}.pt"285        if args.skip_existing and audio_path.exists():286            continue287 288        try:289            # Load audio290            voice = decode_audio_from_file(sample["audio_path"], device, 0.0, args.max_duration)291            if voice is None:292                logging.warning(f"  Skipping {sample['id']}: no audio")293                continue294 295            w = voice.waveform296            if w.dim() == 2:297                if w.shape[0] == 1:298                    w = w.repeat(2, 1)299                w = w.unsqueeze(0)300            elif w.dim() == 3 and w.shape[1] == 1:301                w = w.repeat(1, 2, 1)302            voice = Audio(waveform=w, sampling_rate=voice.sampling_rate)303 304            # Encode through Audio VAE305            audio_latent = ac(lambda enc: vae_encode_audio(voice, enc, None))306 307            # Save audio latent308            torch.save({309                "latents": audio_latent.squeeze(0).cpu(),  # [C=8, T, F=16]310                "sample_rate": 16000,311            }, audio_path)312 313        except Exception as e:314            logging.warning(f"  Skipping {sample['id']}: {e}")315            continue316 317        if idx % 100 == 0:318            logging.info(f"  Audio encoding: {idx}/{len(samples)}")319 320    del ac321    torch.cuda.empty_cache()322 323    # ── Step 3: Create dummy video latents ──324    logging.info("Creating dummy video latents...")325    # Minimal video: 1 frame, 64x64 = 2x2 in latent space326    dummy_video = {327        "latents": torch.zeros(128, 1, 2, 2),328        "num_frames": 1,329        "height": 2,330        "width": 2,331        "fps": 24.0,332    }333    for idx, sample in enumerate(samples):334        gidx = sample["global_idx"]335        latent_path = out / "latents" / f"sample_{gidx:06d}.pt"336        if args.skip_existing and latent_path.exists():337            continue338        torch.save(dummy_video, latent_path)339 340    # ── Summary ──341    n_audio = len(list((out / "audio_latents").glob("*.pt")))342    n_cond = len(list((out / "conditions").glob("*.pt")))343    n_lat = len(list((out / "latents").glob("*.pt")))344    logging.info(f"\nDone! Output: {args.output_dir}")345    logging.info(f"  audio_latents: {n_audio} files")346    logging.info(f"  conditions:    {n_cond} files")347    logging.info(f"  latents:       {n_lat} files")348 349 350if __name__ == "__main__":351    main()352