
num_covers = 97
print(f"Found {len(metas)} songs...")
# do only english songs
metas = [meta for meta in metas if meta["lang"] == "en"]
print(f"Found {len(metas)} english songs...")

covers_dir = "covers-20240816"
os.makedirs(f"/home/christian/code/christian/notebooks/outputs/{covers_dir}", exist_ok=True)

for n in range(100):
    uid = uuid.uuid4()
    rand_meta_idx = np.random.randint(len(metas))
    meta = metas[rand_meta_idx]
    text = meta["lyrics"]

    # tags of the source song
    text_tags = meta["tags_text"]
    text_tags = [tag.replace("Genius", "") for tag in text_tags]
    text_tags = [tag.strip() for tag in text_tags]
    print("source tags:", text_tags)

    # download audio from s3
    filename = os.path.basename(meta["audio_filepath"])
    os.system(f"aws s3 cp {meta['audio_filepath']} /home/christian/code/christian/tmp2/{filename}")
    audio, sr = torchaudio.load(f"/home/christian/code/christian/tmp2/{filename}")

    # crop first 30 sec
    source_dur_sec = 60.0
    start_idx = np.random.randint(0, int(sr*10))
    end_idx = start_idx + int(sr*source_dur_sec)
    start_sec = start_idx / sr
    audio = audio[:, start_idx:end_idx]
    IPython.display.display(IPython.display.Audio(data=audio, rate=sr))
    torchaudio.save(f"/home/christian/code/christian/notebooks/outputs/{covers_dir}/{uid}-source.mp3", audio, sr)

    # get random tags from a different song
    rand_meta_idx = np.random.randint(len(metas))
    target_meta = metas[rand_meta_idx]
    target_tags = target_meta["tags_text"]
    target_tags = [tag.replace("Genius", "") for tag in target_tags]
    target_tags = [tag.strip() for tag in target_tags]
    print("target tags:", target_tags)

    # encode the audio
    cover_arr = process_audio(Audio.from_s3(meta["audio_filepath"], n_channels=2))
    print(cover_arr.shape)

    num_tokens = int(RATE_HZ * source_dur_sec)
    start_token = int(RATE_HZ * start_sec)
    end_token = start_token + num_tokens
    cover_arr = cover_arr[start_token:end_token,:]
    print(cover_arr.shape)

    # randomly sample from 1.0 to 6.0
    cfg_coef = 1.3
    cfg_coef_tags = np.random.choice([1.3, 6.0])

    gconf = GenerationConfig(
        text=text,
        text_tags=target_tags,
        cover_arr=cover_arr,
        cfg_coef=cfg_coef,
        cfg_coef_tags=cfg_coef_tags,
        n_repeat_tags=3,
        text_start_control_tags="{start}",
        text_end_control_tags="{end}",
        n_batch=1,
        min_eos_p=0.1,
        min_text_offset=0,
        eos_pad_duration_s=0,
        max_gen_duration_s=60,
    )

    example_meta = {
        "uid": str(uid),
        "source_start_sec": start_sec,
        "source_end_sec": start_sec + 60.0,
        "source_tags" : text_tags,
        "source_s3_filepath" : meta["audio_filepath"],
        "target_tags" : target_tags,
        "lyrics" : text,
        "cfg_coef" : cfg_coef,
        "cfg_coef_tags" : cfg_coef_tags,
    }

    with open(f"/home/christian/code/christian/notebooks/outputs/{covers_dir}/{uid}-meta.json", "w") as f:
        json.dump(example_meta, f)

    requests = [
        make_request(f"{i}", gconf, engine.model.config, engine.tokenizer)
        for i in range(N_BATCH)
    ]
    # print(requests[0])
    jobs = engine.run_request(requests, tqdm_enabled=True)

    for job_idx, job in enumerate(jobs):
        stream = engine.token_generator(job)

        # stack stream
        #stream_arr = torch.stack(stream_arr)
        #print(stream_arr.shape)

        # convert to numpy and save as npz
        #stream_arr = stream_arr.cpu().numpy()
        #np.savez(f"/home/christian/code/christian/notebooks/outputs/{covers_dir}/{uid}-cover-{job_idx}.npz", stream_arr)
        
        audio = codec_decode_stream_to_full_audio(align_codes(stream, cfg))
        #audio.play()
        audio.write_mp3(f"/home/christian/code/christian/notebooks/outputs/{covers_dir}/{uid}-cover-{job_idx}.mp3")

