import numpy as np import os from tqdm import tqdm import torch from suno_utils.audio import Audio import json import random from suno_utils.tasks.ditto import Ditto if __name__ == "__main__": avialbe_device = f"cuda:{os.environ['CUDA_VISIBLE_DEVICES']}" print(avialbe_device) with open("/home/tony/Data/Preference/7b_v2/7b_before_recode_20240412_s3.json", "r") as fp: all_s3_ids = json.load(fp) print(f"total jobs are, {len(all_s3_ids)}") my_ditto = Ditto( music_encoder_name="musicfm_concat", latent_dim=128, model_path="/home/tony/Data/Ditto/ditto.pt", music_encoder_path="/home/tony/Data/Ditto/music_encoder.pt", is_flash=False ) my_ditto = my_ditto.eval().cuda() print("finish loading ditto") def get_ditto_music_emb(test_id): s3_url = f"s3://suno-data-uploads/studio/uploads/{test_id}.mp3" start = 0 dur = 120 audio = Audio.from_s3(s3_url, n_channels=1, sample_rate=24000) audio = audio.get_segment(from_s=start, to_s=start + dur) wav = torch.tensor(audio.array_float).unsqueeze(0).cuda() emb = my_ditto.music_to_latent(wav)[0].detach().cpu().numpy() return emb def get_ditto_text_emb(test_text): emb = my_ditto.text_to_latent("[CLS]" + test_text)[0].detach().cpu().numpy() return emb def get_ditto_and_save(test_id): output_path = f"/app/suno/data/dpo/ditto_npz/{test_id}.npy" if os.path.exists(output_path): return with open(output_path, "w") as fp: fp.write("") try: emb = get_ditto_music_emb(test_id) np.save(output_path, emb) except Exception as e: print(f"failed {test_id} with {str(e)}") os.remove(output_path) print("start working") random.shuffle(all_s3_ids) for filename in tqdm(all_s3_ids): get_ditto_and_save(filename) print("finish working")