Untitled

 avatar
unknown
python
10 months ago
5.8 kB
15
Indexable
def validate_audio_quality(
    original_files: List[Path], metadata: Dict, output_dir: Path
) -> Tuple[Dict, List[Path]]:
    """
    Validates audio quality by reconstructing with XY model and scoring with AudioBox Aesthetics.
    Removes samples that are below the quality threshold. This is done in batches to conserve memory.
    """
    print("Starting audio quality validation stage...")
    
    # --- 1. Load XY Models and Encode Audio ---
    print("Loading XY Tokenizer models for reconstruction...")
    gpu_ids = list(range(torch.cuda.device_count()))
    xy_models = prepare_xy_models(
        gpu_ids,
        config_path="./XY_Tokenizer/config/xy_tokenizer_32k_config.yaml",
        ckpt_path="./XY_Tokenizer/weights/xy_tokenizer.ckpt",
    )

    if not xy_models:
        print("ERROR: Failed to load XY Tokenizer models. Skipping quality validation.")
        return metadata, original_files

    recon_dir = output_dir / "recon"
    safe_mkdir(recon_dir)
    encode_and_reconstruct_audio_files(original_files, xy_models, recon_dir)

    # Offload XY models
    del xy_models
    torch.cuda.empty_cache()
    print("XY Tokenizer models offloaded.")

    # --- 3. Score and Filter in Batches ---
    print("Scoring and filtering reconstructed audio with AudioBox Aesthetics model...")
    bad_samples = set()
    good_sample_scores = {}
    scores_log_path = output_dir / "scores.jsonl"
    
    try:
        predictor = initialize_predictor()
        recon_files = sorted(list(recon_dir.glob("*.wav")))
        
        if not recon_files:
            print("No reconstructed files found to score.")
            shutil.rmtree(recon_dir, ignore_errors=True)
            return metadata, original_files
        
        batch_size = 1
        with tqdm(total=len(recon_files), desc="Scoring audio quality") as pbar:
            for i in range(0, len(recon_files), batch_size):
                batch_files = recon_files[i:i + batch_size]
                paths_for_predictor = [{"path": str(p)} for p in batch_files]

                try:
                    scores = predictor.forward(paths_for_predictor)
                except Exception as e:
                    tqdm.write(f"\n[ERROR] Failed to score batch: {e}. Skipping batch.")
                    pbar.update(len(batch_files))
                    continue

                for j, recon_file in enumerate(batch_files):
                    sample_id = recon_file.stem
                    score_data = scores[j]

                    with open(scores_log_path, 'a') as f:
                        f.write(json.dumps({"sample_id": sample_id, "scores": score_data}) + '\n')

                    is_bad = False
                    reasons = []
                    
                    pc_score = score_data.get("PC", 10.0) # Default high to fail if not present
                    ce_score = score_data.get("CE", 0.0)
                    cu_score = score_data.get("CU", 0.0)

                    # 1. Hard check for PC
                    if pc_score > PC_THRESHOLD:
                        is_bad = True
                        reasons.append(f"PC_above_threshold:{pc_score:.2f}")
                    # 2. If PC is fine, check if at least one of CE or CU passes
                    else:
                        if ce_score >= CE_THRESHOLD or cu_score >= CU_THRESHOLD:
                            is_bad = False  # This is a good sample
                        else:
                            # Both CE and CU failed
                            is_bad = True
                            if ce_score < CE_THRESHOLD:
                                reasons.append(f"CE_below_threshold:{ce_score:.2f}")
                            if cu_score < CU_THRESHOLD:
                                reasons.append(f"CU_below_threshold:{cu_score:.2f}")

                    if is_bad:
                        bad_samples.add(sample_id)
                    else:
                        good_sample_scores[sample_id] = score_data
                
                pbar.update(len(batch_files))

                # Clear cache after each batch to prevent memory accumulation
                torch.cuda.empty_cache()

    except Exception as e:
        print(f"ERROR: Failed to run AudioBox Aesthetics model: {e}. Skipping quality validation.")
        if recon_dir.exists():
            shutil.rmtree(recon_dir, ignore_errors=True)
        return metadata, original_files
    
    # --- 4. Finalize based on scoring results ---
    if bad_samples:
        print(f"Found {len(bad_samples)} bad quality samples to remove.")
    else:
        print("No bad quality samples found.")
    
    # Clean up recon dir now that scoring is complete
    shutil.rmtree(recon_dir, ignore_errors=True)
    # Explicitly clear GPU memory after AudioBox model
    del predictor
    torch.cuda.empty_cache()
    print("Cleaned up reconstruction files and cleared GPU memory after quality validation.")
        
    # Find original files to remove
    files_to_remove = [f for f in original_files if f.stem in bad_samples]
    
    # Remove original audio files
    for f in files_to_remove:
        if f.exists():
            f.unlink()

    # Filter metadata and create list of good audio files
    filtered_metadata = {k: v for k, v in metadata.items() if k not in bad_samples}
    good_audio_files = [f for f in original_files if f.stem not in bad_samples]
    
    # Add scores to the metadata of good samples
    for sample_id, scores_data in good_sample_scores.items():
        if sample_id in filtered_metadata:
            filtered_metadata[sample_id]['quality_scores'] = scores_data
    
    if bad_samples:
        print(f"Removed {len(bad_samples)} samples. Metadata count changed from {len(metadata)} to {len(filtered_metadata)}.")
    
    return filtered_metadata, good_audio_files
Editor is loading...
Leave a Comment