import os import uuid import torch import torchaudio import numpy as np import matplotlib.pyplot as plt from suno_utils.audio import Audio from suno_utils.utils.text import ( write_jsonl, read_jsonl, write_json, read_json, normalize_whitespace, ) if __name__ == "__main__": metas_filepath = "/app/suno/data/chirp_v4_ft_sm/multi/metas_tr_quality.jsonl" # metas_filepath = "metadata/deezer_metas_audio_quality.jsonl" metas = read_jsonl(metas_filepath) subset = os.path.basename(metas_filepath).replace(".jsonl", "") os.makedirs("outputs/audio", exist_ok=True) os.makedirs(f"outputs/audio/{subset}", exist_ok=True) save_audio = False results = {} scores_list = [] for meta in metas: if "audio_quality" in meta: quant = float(meta["audio_quality"]["quantification"]) pref = float(meta["audio_quality"]["preference"]) score = float(meta["audio_quality"]["score"]) # score = ((pref * 2) - 1) * (quant + 1) # sign = 1 if pref > 0.5 else -1 uid = uuid.uuid4() results[uid] = { "s3_filepath": meta["s3_filepath"], "score": score, "id": meta["id"], # "score": float(meta["audio_quality"]["quantification"]), # "start_s": meta["start_s"], # "end_s": meta["end_s"], } scores_list.append(score) print( f"min: {np.min(scores_list)} max: {np.max(scores_list)} mean: {np.mean(scores_list)}" ) print() fig, ax = plt.subplots() counts, bins, patches = ax.hist(scores_list, bins=5, edgecolor="black") print(bins) # Add text annotations for count, patch in zip(counts, patches): height = patch.get_height() ax.text( patch.get_x() + patch.get_width() / 2.0, height, int(height), ha="center", va="bottom", ) plt.yscale("log") plt.savefig(f"outputs/{subset}_scores.png", dpi=300) # save the audio from the top N worse and best audio for mode in ["worst"]: if mode == "worst": sorted_results = dict( sorted(results.items(), key=lambda item: item[1]["score"]) ) else: sorted_results = dict( sorted(results.items(), key=lambda item: item[1]["score"], reverse=True) ) for idx, (uid, result) in enumerate(sorted_results.items()): if idx < 25: print(result["score"], result["s3_filepath"]) # start_frame = int(result["start_s"] * sample_rate) # end_frame = int(result["end_s"] * sample_rate) # audio = audio[:, start_frame:end_frame] if save_audio: s3_id = os.path.basename(result["s3_filepath"]).split(".")[0] audio = Audio.from_s3(result["s3_filepath"], n_channels=2) sample_rate = audio.sample_rate audio = torch.from_numpy(audio.array_float) torchaudio.save( f"""outputs/audio/{subset}/{result["score"]:0.2f}-{s3_id}.mp3""", audio, sample_rate, compression=torchaudio.io.CodecConfig(bit_rate=320_000), ) print()