# %% import os procid = int(os.environ.get("SLURM_PROCID", 0)) localid = int(os.environ.get("SLURM_LOCALID", 0)) world_size = int(os.environ.get("SLURM_JOB_NUM_NODES", 1)) * int( os.environ.get("SLURM_NTASKS_PER_NODE", 1) ) verbose = False assert world_size > 0, "WORLD_SIZE is 0" print(f"PROCID: {procid}, LOCALID: {localid}, WORLD_SIZE: {world_size}") # %% from suno_utils.diffusion import generation as diffusion_gen from suno_utils.tasks.upsample_engine import UpsampleEngine, Request from suno_utils.tasks.dac_vae_fixed_25hz import decode_stream_to_full_audio, encode, decode import torch import numpy as np from tqdm import tqdm from suno_utils.audio import Audio torch.cuda.set_device(localid) # %% import json id_path = "/home/christian/code/christian/metadata/v45_splits/ids_keep_sets_v11.json" all_ids = [] with open(id_path, "r") as f: ids_keep_sets = json.load(f) for k, v in ids_keep_sets.items(): print(k, len(v)) all_ids.extend(v) # print an example print(ids_keep_sets["discogs_subset"][0]) print(len(all_ids)) all_ids = all_ids[procid::world_size] def get_audio_path(id): path = f"/app2/suno/data/raw_audio_opus_v0/{id}.opus" if os.path.exists(path): return path else: return None # !ls /app/suno/data/auk_v0 # !head /app/suno/data/auk_v0/metas_v2_val.jsonl # !ls /app2/suno/data/raw_audio_opus_v0/ZrGWefub8RU.opus # %% import numpy as np diffusion_gen.preload_models( dit_model_filepath="/app/suno/checkpoints/2025-05-26_22-49-34_s8646/last_ckpt_infer.pt", # 12 stems codec_filepath="s3://suno-data/minz/models/dac_vae_tuned_25hz.pth", compile=True, ) engine = UpsampleEngine(min_chunk_size=25 * 30) # %% def gen_stem( audio: Audio, stem_type_cfg_scale=1.0, tags="extract [split_karaoke]", steps=4, seed=3, codec_scale_factor=0.4, scale_ctx_vector=True, noise_ctx_level=0.0, infill_prefix_latents=None, infill_suffix_latents=None, ): vae = encode(audio) gen_cfg = diffusion_gen.DiffusionGenerationConfig( lyrics=tags, steps=steps, seed=seed, codec_scale_factor=codec_scale_factor, scale_ctx_vector=scale_ctx_vector, noise_ctx_level=noise_ctx_level, text_cfg_coef=stem_type_cfg_scale, infill_prefix_latents=infill_prefix_latents, infill_suffix_latents=infill_suffix_latents, drop_semantic_tokens=True, ) if verbose: print(gen_cfg) request = Request( id="dummy", generation_config=gen_cfg, tokens=np.zeros((vae.shape[0], 1)), input_tokens_finished=True, stem_ctx_latents=vae, ) result = engine.run_request(request, tqdm_enabled=verbose) vae_latents = torch.concat(result.vae_latents) if verbose: print(f"vae_latents: {vae_latents.shape}") audios = [] for i in tqdm(range(vae_latents.shape[1]), desc="Decoding stems", disable=not verbose): # audios.append(decode(vae_latents[:, i])) if vae_latents.shape[0] > 25 * 60: audios.append(decode_stream_to_full_audio(vae_latents[:, i], n_stride_tokens=25 * 60)) else: audios.append(decode(vae_latents[:, i])) return audios categories = [ "Vocals", "Backing_Vocals", "Drums", "Bass", "Guitar", "Keyboard", "Percussion", "Strings", "Synth", "FX", "Brass", "Woodwinds", ] # %% OUT_DIR = "/app2/suno/data/sft_stems_12_output_v11" os.makedirs(OUT_DIR, exist_ok=True) def write_stem(args): id, category, stem = args if stem.loudness < -45: return None out_path = f"{OUT_DIR}/{id}_{category}.opus" stem.write_opus(out_path) return category import threading def process_id(id): audio = get_audio_path(id) if audio is None: return None audio = Audio.from_file(audio, n_channels=2) stems = gen_stem(audio, steps=8) # Prepare arguments for parallel processing write_args = [(id, category, stem) for category, stem in zip(categories, stems)] # Use threading to write stems in parallel results = [None] * len(write_args) threads = [] def write_stem_thread(i, args): results[i] = write_stem(args) for i, args in enumerate(write_args): thread = threading.Thread(target=write_stem_thread, args=(i, args)) threads.append(thread) thread.start() for thread in threads: thread.join() # Filter out None results found_categories = [result for result in results if result is not None] return found_categories # %% for id in tqdm(all_ids, mininterval=600): try: process_id(id) except Exception as e: print(f"Error processing {id}: {e}") continue