#!/usr/bin/env python3 """Analyze trends across dataset versions and create visualizations.""" import json import sys from pathlib import Path from typing import Dict, List, Any import matplotlib.pyplot as plt import matplotlib matplotlib.use("Agg") # Non-interactive backend import numpy as np # Add parent directory to path sys.path.append(str(Path(__file__).parent.parent.parent)) class TrendAnalyzer: """Analyze trends across dataset versions.""" def __init__(self, output_dir: Path): """Initialize analyzer. Args: output_dir: Base output directory containing version subdirectories """ self.output_dir = Path(output_dir) # Put viz at same level as outputs self.viz_dir = self.output_dir.parent / "viz" self.viz_dir.mkdir(exist_ok=True) def load_version_data(self, version: str) -> Dict[str, Any]: """Load analysis data for a specific version. Args: version: Version name (e.g., 'v9') Returns: Dictionary with train, val, and other analysis results """ # Find the most recent full analysis directory for this version version_dirs = sorted(self.output_dir.glob(f"{version}_full_*")) if not version_dirs: print(f"No analysis found for {version}") return None latest_dir = version_dirs[-1] print(f"Loading {version} from {latest_dir}") data = {"version": version} # Load train data train_dir = latest_dir / "train" if train_dir.exists(): data["train"] = self._load_analysis_files(train_dir) # Load val data val_dir = latest_dir / "val" if val_dir.exists(): data["val"] = self._load_analysis_files(val_dir) # Only include versions that have both train and val completed if "train" not in data or "val" not in data: print(f"Skipping {version} - incomplete analysis (missing train or val)") return None return data def _load_analysis_files(self, analysis_dir: Path) -> Dict[str, Any]: """Load all analysis JSON files from a directory.""" results = {} for json_file in analysis_dir.glob("*_analysis_*.json"): analyzer_name = json_file.stem.split("_analysis_")[0] with open(json_file) as f: results[analyzer_name] = json.load(f) return results def collect_all_versions(self) -> List[Dict[str, Any]]: """Collect data for all available versions.""" versions = ["v0", "v1", "v2", "v3", "v4", "v5", "v6", "v7", "v8", "v9"] all_data = [] for version in versions: data = self.load_version_data(version) if data: all_data.append(data) return all_data def plot_lists_trends(self, all_data: List[Dict]): """Plot trends in lists (artists, covers, playlists) over versions.""" fig, axes = plt.subplots(3, 2, figsize=(16, 18)) versions = [d["version"] for d in all_data] categories = ["artist_ids", "cover_ids", "playlist_ids", "stems"] colors = ["blue", "orange", "green", "red"] # 1-4. Individual plots for each category (absolute counts) for idx, (cat, color) in enumerate(zip(categories, colors)): ax = axes[idx // 2, idx % 2] counts = [] for d in all_data: if "train" in d and "lists" in d["train"]: cat_data = d["train"]["lists"].get(cat, {}) cat_summary = cat_data.get("summary", {}) counts.append(cat_summary.get("records_with_field", 0)) else: counts.append(0) ax.plot(versions, counts, marker="o", color=color, linewidth=2, markersize=8) ax.set_xlabel("Version") ax.set_ylabel("Number of Records") ax.set_title(f'{cat.replace("_", " ").title()} Over Versions (Absolute)') ax.grid(True, alpha=0.3) # Add percentage on each point total_records = [] for d in all_data: if "train" in d and "lists" in d["train"]: total_records.append(d["train"]["lists"]["summary"].get("total_records", 1)) else: total_records.append(1) for i, (v, c, t) in enumerate(zip(versions, counts, total_records)): pct = 100 * c / t if t > 0 else 0 ax.text(i, c, f"{pct:.1f}%", ha="center", va="bottom", fontsize=8) # 5. Compound combinations count (log scale) ax = axes[2, 0] compound_counts = [] for d in all_data: if "train" in d and "lists" in d["train"]: compound_data = d["train"]["lists"].get("compound_combinations", {}) compound_counts.append(len(compound_data.get("all_combinations", []))) else: compound_counts.append(0) ax.bar(versions, compound_counts, color="skyblue", edgecolor="black") ax.set_xlabel("Version") ax.set_ylabel("Number of Unique Combinations (log scale)") ax.set_yscale("log") ax.set_title("Unique Compound List Combinations Over Versions") ax.grid(True, alpha=0.3, axis="y") # 6. Top combinations (latest version) ax = axes[2, 1] if all_data and "train" in all_data[-1] and "lists" in all_data[-1]["train"]: compound_data = all_data[-1]["train"]["lists"].get("compound_combinations", {}) combos = compound_data.get("all_combinations", [])[:10] names = [c["combination"] for c in combos] counts = [c["count"] for c in combos] ax.barh(range(len(names)), counts, color="lightcoral", edgecolor="black") ax.set_yticks(range(len(names))) ax.set_yticklabels(names) ax.set_xlabel("Count") ax.set_title(f'Top 10 Compound Combinations ({all_data[-1]["version"]})') ax.grid(True, alpha=0.3, axis="x") plt.tight_layout() plt.savefig(self.viz_dir / "lists_trends.png", dpi=150, bbox_inches="tight") print(f"Saved: {self.viz_dir / 'lists_trends.png'}") plt.close() def plot_tags_trends(self, all_data: List[Dict]): """Plot trends in tags over versions.""" fig, axes = plt.subplots(2, 2, figsize=(16, 12)) versions = [d["version"] for d in all_data] # 1. Average tags per song ax = axes[0, 0] avg_tags = [] for d in all_data: if "train" in d and "tags" in d["train"]: dist = d["train"]["tags"].get("tags_per_song_distribution", {}) avg_tags.append(dist.get("mean", 0)) else: avg_tags.append(0) ax.plot(versions, avg_tags, marker="o", color="blue", linewidth=2) ax.set_xlabel("Version") ax.set_ylabel("Average Tags per Song") ax.set_title("Average Tags per Song Over Versions") ax.grid(True, alpha=0.3) # 2. Percentage with tags ax = axes[0, 1] pct_with_tags = [] for d in all_data: if "train" in d and "tags" in d["train"]: total = d["train"]["tags"]["summary"].get("total_records", 1) with_tags = d["train"]["tags"]["summary"].get("records_with_tags", 0) pct_with_tags.append(100 * with_tags / total if total > 0 else 0) else: pct_with_tags.append(0) ax.plot(versions, pct_with_tags, marker="o", color="green", linewidth=2) ax.set_xlabel("Version") ax.set_ylabel("Percentage (%)") ax.set_title("Percentage of Songs with Tags Over Versions") ax.grid(True, alpha=0.3) # 3. Unique tags count ax = axes[1, 0] unique_tags = [] for d in all_data: if "train" in d and "tags" in d["train"]: unique_tags.append(d["train"]["tags"]["summary"].get("unique_tags", 0)) else: unique_tags.append(0) ax.bar(versions, unique_tags, color="coral", edgecolor="black") ax.set_xlabel("Version") ax.set_ylabel("Count") ax.set_title("Unique Tags Over Versions") ax.grid(True, alpha=0.3, axis="y") # 4. Total tag occurrences ax = axes[1, 1] total_occurrences = [] for d in all_data: if "train" in d and "tags" in d["train"]: total_occurrences.append(d["train"]["tags"]["summary"].get("total_tag_occurrences", 0)) else: total_occurrences.append(0) ax.bar(versions, total_occurrences, color="lightblue", edgecolor="black") ax.set_xlabel("Version") ax.set_ylabel("Total Occurrences") ax.set_title("Total Tag Occurrences Over Versions") ax.grid(True, alpha=0.3, axis="y") plt.tight_layout() plt.savefig(self.viz_dir / "tags_trends.png", dpi=150, bbox_inches="tight") print(f"Saved: {self.viz_dir / 'tags_trends.png'}") plt.close() def plot_stems_trends(self, all_data: List[Dict]): """Plot trends in stems over versions.""" fig, axes = plt.subplots(2, 2, figsize=(16, 12)) versions = [d["version"] for d in all_data] # 1. Absolute count with stems ax = axes[0, 0] with_stems = [] with_captions = [] for d in all_data: if "train" in d and "stems" in d["train"]: with_stems.append(d["train"]["stems"]["summary"].get("records_with_stems", 0)) with_captions.append( d["train"]["stems"]["summary"].get("records_with_stems_captions", 0) ) else: with_stems.append(0) with_captions.append(0) ax.plot(versions, with_stems, marker="o", label="With Stems", linewidth=2) ax.plot(versions, with_captions, marker="s", label="With Stems Captions", linewidth=2) ax.set_xlabel("Version") ax.set_ylabel("Number of Records") ax.set_title("Records with Stems Over Versions (Absolute)") ax.legend() ax.grid(True, alpha=0.3) # 2. Percentage with stems ax = axes[0, 1] pct_stems = [] pct_captions = [] for d in all_data: if "train" in d and "stems" in d["train"]: total = d["train"]["stems"]["summary"].get("total_records", 1) stems = d["train"]["stems"]["summary"].get("records_with_stems", 0) captions = d["train"]["stems"]["summary"].get("records_with_stems_captions", 0) pct_stems.append(100 * stems / total if total > 0 else 0) pct_captions.append(100 * captions / total if total > 0 else 0) else: pct_stems.append(0) pct_captions.append(0) ax.plot(versions, pct_stems, marker="o", label="With Stems", linewidth=2) ax.plot(versions, pct_captions, marker="s", label="With Stems Captions", linewidth=2) ax.set_xlabel("Version") ax.set_ylabel("Percentage (%)") ax.set_title("Percentage with Stems Over Versions") ax.legend() ax.grid(True, alpha=0.3) # 3. Average stem count per song (for songs with stems) ax = axes[1, 0] avg_stem_count = [] for d in all_data: if "train" in d and "stems" in d["train"]: dist = d["train"]["stems"].get("stem_count_distribution", {}) avg_stem_count.append(dist.get("mean", 0)) else: avg_stem_count.append(0) ax.plot(versions, avg_stem_count, marker="o", color="purple", linewidth=2) ax.set_xlabel("Version") ax.set_ylabel("Average Stem Count") ax.set_title("Average Stems per Song (for songs with stems)") ax.grid(True, alpha=0.3) # 4. Unique stem names ax = axes[1, 1] unique_stems = [] for d in all_data: if "train" in d and "stems" in d["train"]: unique_stems.append(d["train"]["stems"]["summary"].get("unique_stem_names", 0)) else: unique_stems.append(0) ax.bar(versions, unique_stems, color="mediumpurple", edgecolor="black") ax.set_xlabel("Version") ax.set_ylabel("Count") ax.set_title("Unique Stem Names Over Versions") ax.grid(True, alpha=0.3, axis="y") plt.tight_layout() plt.savefig(self.viz_dir / "stems_trends.png", dpi=150, bbox_inches="tight") print(f"Saved: {self.viz_dir / 'stems_trends.png'}") plt.close() def plot_text_trends(self, all_data: List[Dict]): """Plot trends in text over versions.""" fig, axes = plt.subplots(2, 2, figsize=(16, 12)) versions = [d["version"] for d in all_data] # 1. Average text length ax = axes[0, 0] avg_text_len = [] for d in all_data: if "train" in d and "text" in d["train"]: dist = d["train"]["text"].get("text_length_distribution", {}) avg_text_len.append(dist.get("mean", 0)) else: avg_text_len.append(0) ax.plot(versions, avg_text_len, marker="o", color="darkgreen", linewidth=2) ax.set_xlabel("Version") ax.set_ylabel("Average Length (characters)") ax.set_title("Average Text Length Over Versions") ax.grid(True, alpha=0.3) # 2. Percentage with text ax = axes[0, 1] pct_with_text = [] for d in all_data: if "train" in d and "text" in d["train"]: total = d["train"]["text"]["summary"].get("total_records", 1) with_text = d["train"]["text"]["summary"].get("records_with_text", 0) pct_with_text.append(100 * with_text / total if total > 0 else 0) else: pct_with_text.append(0) ax.plot(versions, pct_with_text, marker="o", color="darkorange", linewidth=2) ax.set_xlabel("Version") ax.set_ylabel("Percentage (%)") ax.set_title("Percentage with Text Over Versions") ax.grid(True, alpha=0.3) # 3. Records with text (absolute) ax = axes[1, 0] with_text = [] for d in all_data: if "train" in d and "text" in d["train"]: with_text.append(d["train"]["text"]["summary"].get("records_with_text", 0)) else: with_text.append(0) ax.bar(versions, with_text, color="lightgreen", edgecolor="black") ax.set_xlabel("Version") ax.set_ylabel("Number of Records") ax.set_title("Records with Text Over Versions") ax.grid(True, alpha=0.3, axis="y") # 4. Bracket patterns (latest version) ax = axes[1, 1] if all_data and "train" in all_data[-1] and "text" in all_data[-1]["train"]: brackets = all_data[-1]["train"]["text"].get("bracket_patterns_count", {}) if brackets: names = list(brackets.keys())[:8] counts = [brackets[k] for k in names] ax.barh(range(len(names)), counts, color="lightcoral", edgecolor="black") ax.set_yticks(range(len(names))) ax.set_yticklabels(names) ax.set_xlabel("Count") ax.set_title(f'Bracket Patterns ({all_data[-1]["version"]})') ax.grid(True, alpha=0.3, axis="x") plt.tight_layout() plt.savefig(self.viz_dir / "text_trends.png", dpi=150, bbox_inches="tight") print(f"Saved: {self.viz_dir / 'text_trends.png'}") plt.close() def plot_validation_trends(self, all_data: List[Dict]): """Plot trends in validation set over versions.""" fig, axes = plt.subplots(1, 2, figsize=(14, 5)) versions = [d["version"] for d in all_data] # 1. Validation set size ax = axes[0] val_sizes = [] for d in all_data: if "val" in d and "id" in d["val"]: val_sizes.append(d["val"]["id"]["summary"].get("total_records", 0)) else: val_sizes.append(0) ax.bar(versions, val_sizes, color="steelblue", edgecolor="black") ax.set_xlabel("Version") ax.set_ylabel("Number of Records") ax.set_title("Validation Set Size Over Versions") ax.grid(True, alpha=0.3, axis="y") # 2. Validation percentage of total ax = axes[1] train_sizes = [] for d in all_data: if "train" in d and "id" in d["train"]: train_sizes.append(d["train"]["id"]["summary"].get("total_records", 0)) else: train_sizes.append(0) # Calculate validation percentages val_percentages = [] for train_size, val_size in zip(train_sizes, val_sizes): total = train_size + val_size if total > 0: val_percentages.append(100 * val_size / total) else: val_percentages.append(0) ax.bar(versions, val_percentages, color="steelblue", edgecolor="black") ax.set_xlabel("Version") ax.set_ylabel("Validation Percentage (%)") ax.set_title("Validation Set as % of Total Data Over Versions") # Add percentage labels on bars for i, (v, pct) in enumerate(zip(versions, val_percentages)): ax.text(i, pct, f"{pct:.3f}%", ha="center", va="bottom", fontsize=9) ax.grid(True, alpha=0.3, axis="y") plt.tight_layout() plt.savefig(self.viz_dir / "validation_trends.png", dpi=150, bbox_inches="tight") print(f"Saved: {self.viz_dir / 'validation_trends.png'}") plt.close() def plot_duplicate_ids_trends(self, all_data: List[Dict]): """Plot trends in duplicate IDs over versions.""" fig, axes = plt.subplots(1, 2, figsize=(14, 5)) versions = [d["version"] for d in all_data] # 1. Duplicate IDs count ax = axes[0] dup_ids = [] for d in all_data: if "train" in d and "id" in d["train"]: dup_ids.append(d["train"]["id"]["summary"].get("duplicate_ids", 0)) else: dup_ids.append(0) ax.bar(versions, dup_ids, color="salmon", edgecolor="black") ax.set_xlabel("Version") ax.set_ylabel("Number of Duplicate IDs") ax.set_title("Duplicate IDs Over Versions") ax.grid(True, alpha=0.3, axis="y") # 2. Percentage of duplicates ax = axes[1] dup_pct = [] for d in all_data: if "train" in d and "id" in d["train"]: total = d["train"]["id"]["summary"].get("unique_ids", 1) dups = d["train"]["id"]["summary"].get("duplicate_ids", 0) dup_pct.append(100 * dups / total if total > 0 else 0) else: dup_pct.append(0) ax.bar(versions, dup_pct, color="lightcoral", edgecolor="black") ax.set_xlabel("Version") ax.set_ylabel("Percentage (%)") ax.set_title("Duplicate IDs Percentage Over Versions") ax.grid(True, alpha=0.3, axis="y") plt.tight_layout() plt.savefig(self.viz_dir / "duplicate_ids_trends.png", dpi=150, bbox_inches="tight") print(f"Saved: {self.viz_dir / 'duplicate_ids_trends.png'}") plt.close() def plot_s3_categories_trends(self, all_data: List[Dict]): """Plot s3_categories trends: pie charts per version + multi-line chart across versions.""" if not all_data: return versions = [d["version"] for d in all_data] # Extract s3_categories data for all versions all_categories_data = [] for d in all_data: if "train" in d and "paths" in d["train"]: s3_cats = d["train"]["paths"].get("s3_categories", {}) all_categories_data.append(s3_cats) else: all_categories_data.append({}) # Get unique category names across all versions all_category_names = set() for cats in all_categories_data: all_category_names.update(cats.keys()) all_category_names = sorted(all_category_names) # Create a color map for consistent coloring across all plots colors_map = { "bundles/v0/discogs": "#1f77b4", "bundles/v0/pond5": "#ff7f0e", "bundles/v0/genius": "#2ca02c", "bundles/v0/covers": "#d62728", "bundles/v0/discogs_subset_50k": "#9467bd", "bundles/v0/deezer": "#8c564b", "bundles/v0/imslp": "#e377c2", } # Part 1: Pie charts for each version (3 rows x 3 cols for 9 versions) num_versions = len(all_data) fig, axes = plt.subplots(3, 3, figsize=(18, 18)) axes = axes.flatten() for idx, (version, cats_data) in enumerate(zip(versions, all_categories_data)): ax = axes[idx] if not cats_data: ax.text(0.5, 0.5, "No data", ha="center", va="center", transform=ax.transAxes) ax.set_title(f"{version}") ax.axis("off") continue # Sort by count descending sorted_cats = sorted(cats_data.items(), key=lambda x: x[1], reverse=True) labels = [cat.split("/")[-1] for cat, _ in sorted_cats] # Use short names sizes = [count for _, count in sorted_cats] colors = [colors_map.get(cat, "#gray") for cat, _ in sorted_cats] # Create pie chart wedges, texts, autotexts = ax.pie( sizes, labels=labels, colors=colors, autopct="%1.1f%%", startangle=90 ) # Make percentage text smaller for autotext in autotexts: autotext.set_color("white") autotext.set_fontsize(8) autotext.set_weight("bold") ax.set_title(f"{version} S3 Categories", fontsize=12, fontweight="bold") # Hide unused subplots for idx in range(num_versions, len(axes)): axes[idx].axis("off") plt.tight_layout() plt.savefig(self.viz_dir / "s3_categories_pie_charts.png", dpi=150, bbox_inches="tight") print(f"Saved: {self.viz_dir / 's3_categories_pie_charts.png'}") plt.close() # Part 2: Multi-line chart showing percentage trends across versions fig, ax = plt.subplots(figsize=(14, 8)) # Calculate percentages for each category across versions for category in all_category_names: percentages = [] for cats_data in all_categories_data: total = sum(cats_data.values()) if cats_data else 1 count = cats_data.get(category, 0) pct = 100 * count / total if total > 0 else 0 percentages.append(pct) # Plot line for this category short_name = category.split("/")[-1] if "/" in category else category color = colors_map.get(category, None) ax.plot( versions, percentages, marker="o", linewidth=2, markersize=8, label=short_name, color=color, ) # Add percentage labels on each point for i, (v, pct) in enumerate(zip(versions, percentages)): if pct > 0.5: # Only show labels for visible percentages ax.text(i, pct, f"{pct:.1f}%", ha="center", va="bottom", fontsize=7) ax.set_xlabel("Version", fontsize=12) ax.set_ylabel("Percentage (%)", fontsize=12) ax.set_title("S3 Category Distribution Trends Over Versions", fontsize=14, fontweight="bold") ax.legend(loc="best", fontsize=10) ax.grid(True, alpha=0.3) ax.set_ylim(0, 100) plt.tight_layout() plt.savefig(self.viz_dir / "s3_categories_trends.png", dpi=150, bbox_inches="tight") print(f"Saved: {self.viz_dir / 's3_categories_trends.png'}") plt.close() def plot_distributions(self, all_data: List[Dict]): """Plot distributions for all versions.""" if not all_data: return # Generate distribution plots for each version for version_data in all_data: self._plot_single_version_distribution(version_data) def _plot_single_version_distribution(self, version_data: Dict): """Plot distributions for a single version.""" version = version_data["version"] if "train" not in version_data: return fig, axes = plt.subplots(2, 3, figsize=(18, 12)) # 1. Text length distribution ax = axes[0, 0] if "text" in version_data["train"]: dist = version_data["train"]["text"].get("text_length_distribution", {}) if dist: # Draw horizontal bars showing data density between percentiles min_val, p25, median, p75, p95, max_val = ( dist.get("min", 0), dist.get("p25", 0), dist.get("median", 0), dist.get("p75", 0), dist.get("p95", 0), dist.get("max", 0), ) bar_height = 0.6 y_pos = 0.5 # 0-25% (min to p25): 25% of data ax.barh( y_pos, p25 - min_val, left=min_val, height=bar_height, color="lightblue", alpha=0.3, edgecolor="blue", ) # 25-50% (p25 to median): 25% of data ax.barh( y_pos, median - p25, left=p25, height=bar_height, color="lightgreen", alpha=0.4, edgecolor="green", ) # 50-75% (median to p75): 25% of data ax.barh( y_pos, p75 - median, left=median, height=bar_height, color="lightyellow", alpha=0.4, edgecolor="orange", ) # 75-95% (p75 to p95): 20% of data ax.barh( y_pos, p95 - p75, left=p75, height=bar_height, color="lightcoral", alpha=0.3, edgecolor="red", ) # 95-100% (p95 to max): 5% of data ax.barh( y_pos, max_val - p95, left=p95, height=bar_height, color="pink", alpha=0.2, edgecolor="darkred", ) ax.axvline( dist.get("mean", 0), color="r", linestyle="--", linewidth=2, label=f'Mean: {dist.get("mean", 0):.0f}', ) ax.axvline( median, color="g", linestyle="--", linewidth=2, label=f'Median: {dist.get("median", 0):.0f}', ) ax.axvline( p25, color="orange", linestyle=":", linewidth=1.5, alpha=0.7, label=f'P25: {dist.get("p25", 0):.0f}', ) ax.axvline( p75, color="purple", linestyle=":", linewidth=1.5, alpha=0.7, label=f'P75: {dist.get("p75", 0):.0f}', ) ax.axvline( p95, color="brown", linestyle=":", linewidth=1.5, alpha=0.7, label=f'P95: {dist.get("p95", 0):.0f}', ) stats_text = f'Min: {min_val:.0f}, Max: {max_val:.0f}, P99: {dist.get("p99", 0):.0f}' ax.text( 0.98, 0.98, stats_text, transform=ax.transAxes, fontsize=8, verticalalignment="top", horizontalalignment="right", bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), ) ax.set_xlabel("Text Length (characters, log scale)") ax.set_ylabel("Frequency") ax.set_title(f"Text Length Distribution ({version})") ax.set_xscale("log") ax.set_ylim(0, 1) ax.set_yticks([]) ax.legend(fontsize=8, loc="upper left") ax.grid(True, alpha=0.3) # 2. Tags per song distribution ax = axes[0, 1] if "tags" in version_data["train"]: histogram = version_data["train"]["tags"].get("tags_per_song_histogram", {}) if histogram: counts = sorted([(int(k), v) for k, v in histogram.items() if int(k) <= 30]) if counts: tags, freqs = zip(*counts) ax.bar(tags, freqs, color="skyblue", edgecolor="black") ax.set_xlabel("Number of Tags") ax.set_ylabel("Frequency (log scale)") ax.set_yscale("log") ax.set_title(f"Tags per Song Distribution ({version})") ax.grid(True, alpha=0.3, axis="y") # 3. Stem count distribution ax = axes[0, 2] if "stems" in version_data["train"]: dist = version_data["train"]["stems"].get("stem_count_distribution", {}) if dist: # Draw horizontal bars showing data density between percentiles min_val, p25, median, p75, p95, max_val = ( dist.get("min", 0), dist.get("p25", 0), dist.get("median", 0), dist.get("p75", 0), dist.get("p95", 0), dist.get("max", 0), ) bar_height = 0.6 y_pos = 0.5 ax.barh( y_pos, p25 - min_val, left=min_val, height=bar_height, color="lightblue", alpha=0.3, edgecolor="blue", ) ax.barh( y_pos, median - p25, left=p25, height=bar_height, color="lightgreen", alpha=0.4, edgecolor="green", ) ax.barh( y_pos, p75 - median, left=median, height=bar_height, color="lightyellow", alpha=0.4, edgecolor="orange", ) ax.barh( y_pos, p95 - p75, left=p75, height=bar_height, color="lightcoral", alpha=0.3, edgecolor="red", ) ax.barh( y_pos, max_val - p95, left=p95, height=bar_height, color="pink", alpha=0.2, edgecolor="darkred", ) ax.axvline( dist.get("mean", 0), color="r", linestyle="--", linewidth=2, label=f'Mean: {dist.get("mean", 0):.1f}', ) ax.axvline( median, color="g", linestyle="--", linewidth=2, label=f'Median: {dist.get("median", 0):.0f}', ) ax.axvline( p25, color="orange", linestyle=":", linewidth=1.5, alpha=0.7, label=f'P25: {dist.get("p25", 0):.0f}', ) ax.axvline( p75, color="purple", linestyle=":", linewidth=1.5, alpha=0.7, label=f'P75: {dist.get("p75", 0):.0f}', ) ax.axvline( p95, color="brown", linestyle=":", linewidth=1.5, alpha=0.7, label=f'P95: {dist.get("p95", 0):.0f}', ) stats_text = f'Min: {min_val:.0f}, Max: {max_val:.0f}, P99: {dist.get("p99", 0):.0f}' ax.text( 0.98, 0.98, stats_text, transform=ax.transAxes, fontsize=8, verticalalignment="top", horizontalalignment="right", bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), ) ax.set_xlabel("Number of Stems") ax.set_ylabel("Frequency") ax.set_title(f"Stem Count Distribution ({version})") ax.set_ylim(0, 1) ax.set_yticks([]) ax.legend(fontsize=8, loc="upper left") ax.grid(True, alpha=0.3) # 4. Language distribution ax = axes[1, 0] if "language" in version_data["train"]: top_langs = version_data["train"]["language"].get("top_languages", [])[:10] if top_langs: langs, counts = zip(*top_langs) ax.barh(range(len(langs)), counts, color="lightgreen", edgecolor="black") ax.set_yticks(range(len(langs))) ax.set_yticklabels(langs) ax.set_xlabel("Count") ax.set_title(f"Top 10 Languages ({version})") ax.grid(True, alpha=0.3, axis="x") # 5. ID length distribution ax = axes[1, 1] if "id" in version_data["train"]: dist = version_data["train"]["id"].get("id_length_distribution", {}) if dist: # Draw horizontal bars showing data density between percentiles min_val, p25, median, p75, p95, max_val = ( dist.get("min", 0), dist.get("p25", 0), dist.get("median", 0), dist.get("p75", 0), dist.get("p95", 0), dist.get("max", 0), ) bar_height = 0.6 y_pos = 0.5 ax.barh( y_pos, p25 - min_val, left=min_val, height=bar_height, color="lightblue", alpha=0.3, edgecolor="blue", ) ax.barh( y_pos, median - p25, left=p25, height=bar_height, color="lightgreen", alpha=0.4, edgecolor="green", ) ax.barh( y_pos, p75 - median, left=median, height=bar_height, color="lightyellow", alpha=0.4, edgecolor="orange", ) ax.barh( y_pos, p95 - p75, left=p75, height=bar_height, color="lightcoral", alpha=0.3, edgecolor="red", ) ax.barh( y_pos, max_val - p95, left=p95, height=bar_height, color="pink", alpha=0.2, edgecolor="darkred", ) ax.axvline( dist.get("mean", 0), color="r", linestyle="--", linewidth=2, label=f'Mean: {dist.get("mean", 0):.1f}', ) ax.axvline( median, color="g", linestyle="--", linewidth=2, label=f'Median: {dist.get("median", 0):.0f}', ) ax.axvline( p25, color="orange", linestyle=":", linewidth=1.5, alpha=0.7, label=f'P25: {dist.get("p25", 0):.0f}', ) ax.axvline( p75, color="purple", linestyle=":", linewidth=1.5, alpha=0.7, label=f'P75: {dist.get("p75", 0):.0f}', ) ax.axvline( p95, color="brown", linestyle=":", linewidth=1.5, alpha=0.7, label=f'P95: {dist.get("p95", 0):.0f}', ) stats_text = f'Min: {min_val:.0f}, Max: {max_val:.0f}, P99: {dist.get("p99", 0):.0f}' ax.text( 0.98, 0.98, stats_text, transform=ax.transAxes, fontsize=8, verticalalignment="top", horizontalalignment="right", bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), ) ax.set_xlabel("ID Length (characters)") ax.set_ylabel("Frequency") ax.set_title(f"ID Length Distribution ({version})") ax.set_ylim(0, 1) ax.set_yticks([]) ax.legend(fontsize=8, loc="upper left") ax.grid(True, alpha=0.3) # 6. Weight distribution ax = axes[1, 2] if "weight" in version_data["train"]: dist = version_data["train"]["weight"].get("weight_distribution", {}) if dist: # Draw horizontal bars showing data density between percentiles min_val, p25, median, p75, p95, max_val = ( dist.get("min", 0), dist.get("p25", 0), dist.get("median", 0), dist.get("p75", 0), dist.get("p95", 0), dist.get("max", 0), ) bar_height = 0.6 y_pos = 0.5 ax.barh( y_pos, p25 - min_val, left=min_val, height=bar_height, color="lightblue", alpha=0.3, edgecolor="blue", ) ax.barh( y_pos, median - p25, left=p25, height=bar_height, color="lightgreen", alpha=0.4, edgecolor="green", ) ax.barh( y_pos, p75 - median, left=median, height=bar_height, color="lightyellow", alpha=0.4, edgecolor="orange", ) ax.barh( y_pos, p95 - p75, left=p75, height=bar_height, color="lightcoral", alpha=0.3, edgecolor="red", ) ax.barh( y_pos, max_val - p95, left=p95, height=bar_height, color="pink", alpha=0.2, edgecolor="darkred", ) ax.axvline( dist.get("mean", 0), color="r", linestyle="--", linewidth=2, label=f'Mean: {dist.get("mean", 0):.3f}', ) ax.axvline( median, color="g", linestyle="--", linewidth=2, label=f'Median: {dist.get("median", 0):.3f}', ) ax.axvline( p25, color="orange", linestyle=":", linewidth=1.5, alpha=0.7, label=f'P25: {dist.get("p25", 0):.3f}', ) ax.axvline( p75, color="purple", linestyle=":", linewidth=1.5, alpha=0.7, label=f'P75: {dist.get("p75", 0):.3f}', ) ax.axvline( p95, color="brown", linestyle=":", linewidth=1.5, alpha=0.7, label=f'P95: {dist.get("p95", 0):.3f}', ) stats_text = f'Min: {min_val:.3f}, Max: {max_val:.3f}, P99: {dist.get("p99", 0):.3f}' ax.text( 0.98, 0.98, stats_text, transform=ax.transAxes, fontsize=8, verticalalignment="top", horizontalalignment="right", bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), ) ax.set_xlabel("Weight") ax.set_ylabel("Frequency") ax.set_title(f"Weight Distribution ({version})") ax.set_ylim(0, 1) ax.set_yticks([]) ax.legend(fontsize=8, loc="upper left") ax.grid(True, alpha=0.3) plt.tight_layout() plt.savefig(self.viz_dir / f"distributions_{version}.png", dpi=150, bbox_inches="tight") print(f"Saved: {self.viz_dir / f'distributions_{version}.png'}") plt.close() def generate_all_visualizations(self): """Generate all trend visualizations.""" print("Collecting data from all versions...") all_data = self.collect_all_versions() if not all_data: print("No data found!") return print(f"Found data for {len(all_data)} versions: {[d['version'] for d in all_data]}") print("\nGenerating visualizations...") self.plot_lists_trends(all_data) self.plot_tags_trends(all_data) self.plot_stems_trends(all_data) self.plot_text_trends(all_data) self.plot_validation_trends(all_data) self.plot_duplicate_ids_trends(all_data) self.plot_s3_categories_trends(all_data) self.plot_distributions(all_data) print(f"\n✓ All visualizations saved to: {self.viz_dir}") def main(): """Main entry point.""" import argparse parser = argparse.ArgumentParser(description="Analyze trends across dataset versions") parser.add_argument( "--output-dir", default="/home/vibert/data/suno_data_monitor/outputs", help="Base output directory", ) args = parser.parse_args() analyzer = TrendAnalyzer(args.output_dir) analyzer.generate_all_visualizations() if __name__ == "__main__": main()