import json import boto3 import statistics from typing import Dict, List, Any import matplotlib.pyplot as plt import numpy as np import math import pandas as pd from tqdm import tqdm from suno_utils.audio import Audio def load_json_from_s3(bucket_name: str, key: str) -> Dict[str, Any]: """Load JSON data from S3 bucket""" s3 = boto3.client("s3") try: response = s3.get_object(Bucket=bucket_name, Key=key) content = response["Body"].read().decode("utf-8") return json.loads(content) except Exception as e: # print(f"Error loading {key}: {e}") return None def process_wer_data( ids_data: Dict[ str, List[str], ], bucket_name: str = "suno-data-uploads", s3_prefix: str = "tasks/feature_eval/cover_persona/2025_07_11-16_20_50/", ) -> Dict[str, Any]: """Process WER data for all files""" results = {} flat_data = [] all_wers = [] # Collect all WER values for overall stats all_data = [] # Collect all data for overall stats for group_id, file_ids in tqdm(ids_data.items()): group_wers = [] group_data = [] for file_id in file_ids: # Construct S3 key s3_key = f"{s3_prefix}{file_id}_infill_wer.json" # Load JSON from S3 data = load_json_from_s3(bucket_name, s3_key) if data and "wer" in data: data["s3_id"] = file_id flat_data.append(data) group_wers.append(data["wer"]) group_data.append(data) # Add to overall collections all_wers.append(data["wer"]) all_data.append(data) else: pass # print(f"Missing or invalid data for {file_id}") # Calculate statistics for this group if group_wers: results[group_id] = { "wers": group_wers, "mean_wer": statistics.mean(group_wers), "median_wer": statistics.median(group_wers), "min_wer": min(group_wers), "max_wer": max(group_wers), "std_wer": statistics.stdev(group_wers) if len(group_wers) > 1 else 0, "count": len(group_wers), "data": group_data, # Include full data if needed } # Add overall statistics as "all" entry if all_wers: results["all"] = { "wers": all_wers, "mean_wer": statistics.mean(all_wers), "median_wer": statistics.median(all_wers), "min_wer": min(all_wers), "max_wer": max(all_wers), "std_wer": statistics.stdev(all_wers) if len(all_wers) > 1 else 0, "count": len(all_wers), "data": all_data, } return results, flat_data def get_infill_data(timestamp): infill_path = f"/home/sara/glockenspiel/suno_utils/task_eval/modal_runs/infill_mappings_{timestamp}.json" with open(infill_path, "r") as f: ids_data = json.load(f) wer_results, flat_data = process_wer_data( ids_data, s3_prefix=f"tasks/feature_eval/cover_persona/{timestamp}/" ) return wer_results, flat_data def get_dur_success(flat_data): df = pd.DataFrame(flat_data) counts = df["duration_matches"].value_counts() print(f"{len(df)} items") percentages = counts / len(df) return percentages def get_closest_bucket(infill_dur): buckets = [8, 15, 25] bucket_distance = { bucket_dur: abs(infill_dur - bucket_dur) for bucket_dur in buckets } min_key = min(bucket_distance, key=bucket_distance.get) return min_key def get_wer_by_duration(flat_data): df = pd.DataFrame(flat_data) df["duration_bucket"] = df["infill_dur_s"].apply(get_closest_bucket) return df.groupby("duration_bucket")["wer"].mean() timestamp = "2025_07_14-20_38_56" results_30b, flat_30b = get_infill_data(timestamp) print(get_dur_success(flat_30b))